240 lines
7 KiB
Go
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)
|
|
}
|