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) }