rara/packages/go/rag/handler.go
toki fdc86c7ff4
Some checks are pending
ci / validate (push) Waiting to run
initial commit
2026-07-18 18:41:17 +09:00

240 lines
7 KiB
Go

package rag
import (
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"github.com/google/uuid"
)
type Handler struct {
resolver ReleaseResolver
engine Engine
}
type wireRequest struct {
ProjectID string `json:"project_id"`
KnowledgeBaseID string `json:"knowledge_base_id"`
ReleaseAlias string `json:"release_alias"`
Query string `json:"query"`
TopK int `json:"top_k"`
Filters json.RawMessage `json:"filters"`
Trace *bool `json:"trace"`
Stream bool `json:"stream"`
Generation json.RawMessage `json:"generation"`
}
type releaseIdentity struct {
ID uuid.UUID `json:"release_id"`
Alias string `json:"release_alias"`
Version int32 `json:"release_version"`
}
type retrieveResponse struct {
TraceID uuid.UUID `json:"trace_id"`
Release releaseIdentity `json:"release"`
Evidence []Evidence `json:"evidence"`
Diagnostics map[string]any `json:"diagnostics,omitempty"`
}
type answerResponse struct {
TraceID uuid.UUID `json:"trace_id"`
Release releaseIdentity `json:"release"`
Answer string `json:"answer"`
Evidence []Evidence `json:"evidence"`
Usage map[string]any `json:"usage,omitempty"`
}
type errorResponse struct {
Code string `json:"code"`
Message string `json:"message"`
TraceID *uuid.UUID `json:"trace_id,omitempty"`
Details map[string]any `json:"details,omitempty"`
}
func NewHandler(resolver ReleaseResolver, engine Engine) *Handler {
return &Handler{resolver: resolver, engine: engine}
}
func Register(mux *http.ServeMux, handler *Handler) {
mux.HandleFunc("POST /v1/retrieve", handler.retrieve)
mux.HandleFunc("POST /v1/answer", handler.answer)
}
func (handler *Handler) retrieve(writer http.ResponseWriter, request *http.Request) {
parsed, err := decodeRequest(writer, request)
if err != nil {
writeError(writer, http.StatusBadRequest, "invalid_request", err.Error(), nil)
return
}
release, err := handler.resolver.Resolve(
request.Context(),
parsed.Request.ProjectID,
parsed.Request.KnowledgeBaseID,
parsed.Request.ReleaseAlias,
)
if err != nil {
handleServiceError(writer, err, nil)
return
}
result, err := handler.engine.Retrieve(request.Context(), release, parsed.Request)
if err != nil {
handleServiceError(writer, err, nil)
return
}
traceID := uuid.New()
writeResponse(writer, http.StatusOK, retrieveResponse{
TraceID: traceID,
Release: identity(release),
Evidence: result.Evidence,
Diagnostics: result.Diagnostics,
})
}
func (handler *Handler) answer(writer http.ResponseWriter, request *http.Request) {
parsed, err := decodeRequest(writer, request)
if err != nil {
writeError(writer, http.StatusBadRequest, "invalid_request", err.Error(), nil)
return
}
release, err := handler.resolver.Resolve(
request.Context(),
parsed.Request.ProjectID,
parsed.Request.KnowledgeBaseID,
parsed.Request.ReleaseAlias,
)
if err != nil {
handleServiceError(writer, err, nil)
return
}
result, err := handler.engine.Answer(request.Context(), release, parsed)
if err != nil {
if parsed.Stream {
writeSSEError(writer, err)
return
}
handleServiceError(writer, err, nil)
return
}
traceID := uuid.New()
if parsed.Stream {
writeSSECompleted(writer, traceID, release, result)
return
}
writeResponse(writer, http.StatusOK, answerResponse{
TraceID: traceID,
Release: identity(release),
Answer: result.Answer,
Evidence: result.Evidence,
Usage: result.Usage,
})
}
func decodeRequest(writer http.ResponseWriter, request *http.Request) (AnswerRequest, error) {
request.Body = http.MaxBytesReader(writer, request.Body, 1<<20)
decoder := json.NewDecoder(request.Body)
decoder.DisallowUnknownFields()
var wire wireRequest
if err := decoder.Decode(&wire); err != nil {
return AnswerRequest{}, fmt.Errorf("decode JSON: %w", err)
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
if err == nil {
return AnswerRequest{}, errors.New("request body must contain one JSON object")
}
return AnswerRequest{}, fmt.Errorf("decode trailing JSON: %w", err)
}
projectID, err := uuid.Parse(wire.ProjectID)
if err != nil {
return AnswerRequest{}, errors.New("project_id must be a UUID")
}
knowledgeBaseID, err := uuid.Parse(wire.KnowledgeBaseID)
if err != nil {
return AnswerRequest{}, errors.New("knowledge_base_id must be a UUID")
}
if strings.TrimSpace(wire.Query) == "" {
return AnswerRequest{}, errors.New("query is required")
}
if wire.ReleaseAlias == "" {
wire.ReleaseAlias = "production"
}
if wire.TopK == 0 {
wire.TopK = 10
}
if wire.TopK < 1 || wire.TopK > 100 {
return AnswerRequest{}, errors.New("top_k must be between 1 and 100")
}
trace := true
if wire.Trace != nil {
trace = *wire.Trace
}
return AnswerRequest{
Request: Request{
ProjectID: projectID,
KnowledgeBaseID: knowledgeBaseID,
ReleaseAlias: wire.ReleaseAlias,
Query: wire.Query,
TopK: wire.TopK,
Filters: wire.Filters,
Trace: trace,
},
Stream: wire.Stream,
Generation: wire.Generation,
}, nil
}
func identity(release Release) releaseIdentity {
return releaseIdentity{ID: release.ID, Alias: release.Alias, Version: release.Version}
}
func handleServiceError(writer http.ResponseWriter, err error, traceID *uuid.UUID) {
switch {
case errors.Is(err, ErrNotFound):
writeError(writer, http.StatusNotFound, "release_not_found", err.Error(), traceID)
case errors.Is(err, ErrUnavailable):
writeError(writer, http.StatusServiceUnavailable, "integration_unavailable", err.Error(), traceID)
default:
writeError(writer, http.StatusInternalServerError, "internal_error", "request failed", traceID)
}
}
func writeResponse(writer http.ResponseWriter, status int, value any) {
writer.Header().Set("Content-Type", "application/json")
writer.WriteHeader(status)
_ = json.NewEncoder(writer).Encode(value)
}
func writeError(
writer http.ResponseWriter,
status int,
code string,
message string,
traceID *uuid.UUID,
) {
writeResponse(writer, status, errorResponse{Code: code, Message: message, TraceID: traceID})
}
func writeSSEError(writer http.ResponseWriter, err error) {
writer.Header().Set("Content-Type", "text/event-stream")
writer.Header().Set("Cache-Control", "no-cache")
payload, _ := json.Marshal(errorResponse{Code: "integration_unavailable", Message: err.Error()})
_, _ = fmt.Fprintf(writer, "event: error\ndata: %s\n\n", payload)
}
func writeSSECompleted(writer http.ResponseWriter, traceID uuid.UUID, release Release, result AnswerResult) {
writer.Header().Set("Content-Type", "text/event-stream")
writer.Header().Set("Cache-Control", "no-cache")
payload, _ := json.Marshal(answerResponse{
TraceID: traceID,
Release: identity(release),
Answer: result.Answer,
Evidence: result.Evidence,
Usage: result.Usage,
})
_, _ = fmt.Fprintf(writer, "event: completed\ndata: %s\n\n", payload)
}