diff options
Diffstat (limited to 'internal/httpapi')
| -rw-r--r-- | internal/httpapi/api.go | 62 | ||||
| -rw-r--r-- | internal/httpapi/api_test.go | 161 |
2 files changed, 216 insertions, 7 deletions
diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go index e66b170..ea4c0c2 100644 --- a/internal/httpapi/api.go +++ b/internal/httpapi/api.go @@ -150,6 +150,8 @@ func (a *API) registerInference(mux *http.ServeMux) { mux.HandleFunc("POST /api/v1/chat/completions", a.openAIChat) mux.HandleFunc("POST /v1/responses", a.openAIResponses) mux.HandleFunc("POST /api/v1/responses", a.openAIResponses) + mux.HandleFunc("POST /v1/embeddings", a.openAIEmbeddings) + mux.HandleFunc("POST /api/v1/embeddings", a.openAIEmbeddings) mux.HandleFunc("GET /anthropic/v1/models", a.anthropicModels) mux.HandleFunc("GET /api/anthropic/v1/models", a.anthropicModels) @@ -171,6 +173,10 @@ func (a *API) openAIResponses(w http.ResponseWriter, r *http.Request) { a.serveInference(w, r, domain.ProtocolOpenAIResponses) } +func (a *API) openAIEmbeddings(w http.ResponseWriter, r *http.Request) { + a.serveInference(w, r, domain.ProtocolOpenAIEmbeddings) +} + func (a *API) anthropicMessages(w http.ResponseWriter, r *http.Request) { a.serveInference(w, r, domain.ProtocolAnthropic) } @@ -243,6 +249,16 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do return } publicModel := model.ID + if protocol == domain.ProtocolOpenAIEmbeddings { + if envelope.Stream { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Embeddings does not support streaming"}, requestID) + return + } + if !modelDeclaresCapability(model, "embeddings") { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support embeddings"}, requestID) + return + } + } requestedOutput := max(envelope.MaxTokens, envelope.MaxCompletionTokens, envelope.MaxOutputTokens) if model.MaxOutputTokens > 0 && requestedOutput > model.MaxOutputTokens { apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Requested output exceeds the model maximum"}, requestID) @@ -284,7 +300,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do policy, _ = a.limiter.Policy(principal.ProjectID) } if err := a.billingMeter.Authorize(r.Context(), billing.Authorization{ - RequestID: requestID, Principal: principal, Model: model, Body: body, Policy: policy, + RequestID: requestID, Principal: principal, Model: model, Protocol: protocol, Body: body, Policy: policy, }); err != nil { if errors.Is(err, billing.ErrInsufficientBalance) { apierror.Write(w, apierror.Error{Status: http.StatusPaymentRequired, Type: "insufficient_balance", Message: "Account balance is insufficient"}, requestID) @@ -293,6 +309,13 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do ErrorType: "insufficient_balance", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) return } + if errors.Is(err, billing.ErrDailyQuotaExceeded) { + apierror.Write(w, apierror.Error{Status: http.StatusTooManyRequests, Type: "daily_quota_exceeded", Message: "API key daily spend quota exceeded"}, requestID) + a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusTooManyRequests, Success: false, + ErrorType: "daily_quota_exceeded", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) + return + } if errors.Is(err, billing.ErrQuotaExceeded) { apierror.Write(w, apierror.Error{Status: http.StatusTooManyRequests, Type: "monthly_quota_exceeded", Message: "Project monthly spend quota exceeded"}, requestID) a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, @@ -341,6 +364,16 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do }) return } + if protocol == domain.ProtocolOpenAIEmbeddings && !strings.HasPrefix(strings.ToLower(result.Response.Header.Get("Content-Type")), "application/json") { + apierror.Write(w, apierror.Error{Status: http.StatusBadGateway, Type: "invalid_provider_response", Message: "Upstream provider returned a non-JSON Embeddings response"}, requestID) + a.finishUsage(r, domain.UsageEvent{ + RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, + PublicModel: publicModel, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, + Protocol: protocol, StatusCode: http.StatusBadGateway, Success: false, ErrorType: "invalid_provider_response", + Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), + }) + return + } stream := envelope.Stream || strings.HasPrefix(strings.ToLower(result.Response.Header.Get("Content-Type")), "text/event-stream") copyResponseHeaders(w.Header(), result.Response.Header, stream) @@ -354,6 +387,18 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do copyErr := copyResponse(w, result.Response.Body, observer, stream) usageResult := observer.Usage() usageReported := observer.Reported() + ttftMS := int64(0) + if firstOutputAt := observer.FirstOutputAt(); !firstOutputAt.IsZero() { + ttftMS = firstOutputAt.Sub(startedAt).Milliseconds() + if ttftMS < 1 { + ttftMS = 1 + } + providerStartedAt := result.AttemptStartedAt + if providerStartedAt.IsZero() { + providerStartedAt = startedAt + } + a.forwarder.ObserveTTFT(model.ID, result.Route, firstOutputAt.Sub(providerStartedAt)) + } success = copyErr == nil errorType := "" if copyErr != nil { @@ -363,7 +408,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, PublicModel: publicModel, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, Protocol: protocol, Stream: stream, StatusCode: result.Response.StatusCode, Success: success, - ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), Usage: usageResult, + ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), TTFTMS: ttftMS, Usage: usageResult, UsageReported: usageReported, }) a.logger.Info("inference_request", @@ -376,6 +421,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do "status", result.Response.StatusCode, "attempts", result.Attempts, "duration_ms", time.Since(startedAt).Milliseconds(), + "ttft_ms", ttftMS, ) } @@ -428,6 +474,7 @@ func modelProviderDescriptors(model domain.Model) []map[string]string { func (a *API) availableOpenAIModels(principal domain.Principal) []domain.Model { combined := append(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal), a.catalog.ModelsFor(domain.ProtocolOpenAIResponses, principal)...) + combined = append(combined, a.catalog.ModelsFor(domain.ProtocolOpenAIEmbeddings, principal)...) seen := make(map[string]struct{}, len(combined)) result := make([]domain.Model, 0, len(combined)) for _, model := range combined { @@ -527,6 +574,15 @@ func modelHasCapability(model domain.Model, wanted string) bool { } return false } + +func modelDeclaresCapability(model domain.Model, wanted string) bool { + for _, value := range model.Capabilities { + if value == wanted || value == "*" { + return true + } + } + return false +} func modelCreated(model domain.Model) int64 { if model.ReleasedAt != nil { return model.ReleasedAt.Unix() @@ -653,7 +709,7 @@ func copyResponse(w http.ResponseWriter, body io.Reader, observer io.Writer, str if stream { destination = &flushingWriter{writer: w, controller: http.NewResponseController(w)} } - _, err := io.CopyBuffer(io.MultiWriter(destination, observer), body, make([]byte, 32<<10)) + _, err := io.CopyBuffer(io.MultiWriter(observer, destination), body, make([]byte, 32<<10)) return err } diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go index cbb3967..906e136 100644 --- a/internal/httpapi/api_test.go +++ b/internal/httpapi/api_test.go @@ -74,10 +74,14 @@ func TestInferenceBrowserOriginCORS(t *testing.T) { type fakeBillingMeter struct { authorizeErr error + authorized chan billing.Authorization settled chan domain.UsageEvent } -func (m *fakeBillingMeter) Authorize(context.Context, billing.Authorization) error { +func (m *fakeBillingMeter) Authorize(_ context.Context, authorization billing.Authorization) error { + if m.authorized != nil { + m.authorized <- authorization + } return m.authorizeErr } @@ -132,7 +136,7 @@ func TestOpenAIProxyRewritesModelAndEmitsUsage(t *testing.T) { select { case event := <-sink.events: - if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || !event.Success { + if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || event.TTFTMS < 1 || !event.Success { t.Fatalf("unexpected usage event: %+v", event) } case <-time.After(time.Second): @@ -178,6 +182,53 @@ func TestProxyFailsOverBeforeWritingResponse(t *testing.T) { } } +func TestGatewayLearnsTTFTAndPrefersFasterProvider(t *testing.T) { + var fastCalls atomic.Int64 + fast := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + fastCalls.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + })) + defer fast.Close() + var slowCalls atomic.Int64 + slow := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + slowCalls.Add(1) + time.Sleep(25 * time.Millisecond) + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + })) + defer slow.Close() + + gateway, sink := newTestGateway(t, + []config.ProviderConfig{ + {ID: "fast", Protocol: domain.ProtocolOpenAI, BaseURL: fast.URL + "/v1", APIKey: "one"}, + {ID: "slow", Protocol: domain.ProtocolOpenAI, BaseURL: slow.URL + "/v1", APIKey: "two"}, + }, + []config.RouteConfig{ + {Provider: "fast", UpstreamModel: "model", Weight: 1}, + {Provider: "slow", UpstreamModel: "model", Weight: 1}, + }, + ) + defer gateway.Close() + + for range 20 { + response := postOpenAI(t, gateway.URL, false) + _, _ = io.Copy(io.Discard, response.Body) + _ = response.Body.Close() + if response.StatusCode != http.StatusOK { + t.Fatalf("unexpected gateway status: %d", response.StatusCode) + } + select { + case <-sink.events: + case <-time.After(time.Second): + t.Fatal("usage event was not emitted") + } + } + if fastCalls.Load() != 16 || slowCalls.Load() != 4 { + t.Fatalf("TTFT feedback was not applied with bounded exploration: fast=%d slow=%d", fastCalls.Load(), slowCalls.Load()) + } +} + func TestProviderSelectorPinsRouteAndKeepsCanonicalUsageModel(t *testing.T) { var primaryCalls atomic.Int64 primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -465,11 +516,106 @@ func TestOpenAIResponsesProxyRewritesModelAndEmitsUsage(t *testing.T) { t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) } event := <-sink.events - if event.Protocol != domain.ProtocolOpenAIResponses || event.UpstreamModel != "gpt-upstream" || event.Usage.TotalTokens != 13 || !event.UsageReported { + if event.Protocol != domain.ProtocolOpenAIResponses || event.UpstreamModel != "gpt-upstream" || event.Usage.TotalTokens != 13 || event.TTFTMS < 1 || !event.UsageReported { t.Fatalf("unexpected Responses usage event: %+v", event) } } +func TestOpenAIEmbeddingsProxyRewritesModelAndMetersInputTokens(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/embeddings" { + t.Errorf("unexpected path: %s", r.URL.Path) + } + if r.Header.Get("Authorization") != "Bearer upstream-secret" { + t.Errorf("unexpected upstream authorization: %q", r.Header.Get("Authorization")) + } + var request map[string]any + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Error(err) + } + if request["model"] != "embedding-upstream" || request["input"] != "hello vector" { + t.Errorf("unexpected Embeddings request: %+v", request) + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"object":"list","data":[{"object":"embedding","index":0,"embedding":[0.1,0.2]}],"model":"embedding-upstream","usage":{"prompt_tokens":8,"total_tokens":8}}`) + })) + defer upstream.Close() + + meter := &fakeBillingMeter{authorized: make(chan billing.Authorization, 1), settled: make(chan domain.UsageEvent, 1)} + gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ + ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: upstream.URL + "/v1", APIKey: "upstream-secret", + }}, []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-upstream", Weight: 1}}, meter) + defer gateway.Close() + + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello vector"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + payload, _ := io.ReadAll(response.Body) + t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) + } + authorization := <-meter.authorized + if authorization.Protocol != domain.ProtocolOpenAIEmbeddings { + t.Fatalf("authorization protocol = %q", authorization.Protocol) + } + event := <-meter.settled + if event.Protocol != domain.ProtocolOpenAIEmbeddings || event.UpstreamModel != "embedding-upstream" || + event.Usage.InputTokens != 8 || event.Usage.OutputTokens != 0 || event.Usage.TotalTokens != 8 || !event.UsageReported { + t.Fatalf("unexpected Embeddings usage event: %+v", event) + } +} + +func TestOpenAIEmbeddingsRequiresDeclaredCapability(t *testing.T) { + providerConfig := config.ProviderConfig{ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "secret"} + modelCatalog := catalog.New(config.Config{Providers: []config.ProviderConfig{providerConfig}, Models: []config.ModelConfig{{ + ID: "public/model", Routes: []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-upstream", Weight: 1}}, + }}}) + authenticator, err := auth.NewStatic(`[{"key":"client-secret","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","scopes":["inference"]}]`, false) + if err != nil { + t.Fatal(err) + } + api := New(Options{Authenticator: authenticator, Catalog: modelCatalog, Router: routing.New(modelCatalog), Metrics: &telemetry.Metrics{}, Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), MaxBodyBytes: 1 << 20}) + request := httptest.NewRequest(http.MethodPost, "/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response := httptest.NewRecorder() + api.Handler().ServeHTTP(response, request) + if response.Code != http.StatusBadRequest || !strings.Contains(response.Body.String(), "unsupported_capability") { + t.Fatalf("unexpected response %d: %s", response.Code, response.Body.String()) + } +} + +func TestOpenAIEmbeddingsRejectsNonJSONSuccessBeforeResponseStarts(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + _, _ = io.WriteString(w, "<!doctype html><title>provider console</title>") + })) + defer upstream.Close() + meter := &fakeBillingMeter{settled: make(chan domain.UsageEvent, 1)} + gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ + ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: upstream.URL, APIKey: "secret", + }}, []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-model", Weight: 1}}, meter) + defer gateway.Close() + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + body, _ := io.ReadAll(response.Body) + if response.StatusCode != http.StatusBadGateway || !strings.Contains(string(body), "invalid_provider_response") || strings.Contains(string(body), "provider console") { + t.Fatalf("unexpected response %d: %s", response.StatusCode, body) + } + event := <-meter.settled + if event.Success || event.StatusCode != http.StatusBadGateway || event.ErrorType != "invalid_provider_response" || event.UsageReported { + t.Fatalf("unexpected invalid provider usage event: %+v", event) + } +} + func TestInsufficientBalanceRejectsBeforeCallingUpstream(t *testing.T) { var calls atomic.Int64 upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -542,9 +688,16 @@ func newTestGatewayWithBilling(t *testing.T, providers []config.ProviderConfig, if err != nil { t.Fatal(err) } + modelConfig := config.ModelConfig{ID: "public/model", OwnedBy: "test", Routes: routes} + for _, providerConfig := range providers { + if providerConfig.WireAPI == "embeddings" { + modelConfig.Capabilities = []string{"embeddings"} + break + } + } cfg := config.Config{ Providers: providers, - Models: []config.ModelConfig{{ID: "public/model", OwnedBy: "test", Routes: routes}}, + Models: []config.ModelConfig{modelConfig}, UpstreamHTTP: config.UpstreamHTTPConfig{ MaxIdleConnections: 100, MaxIdleConnectionsPerHost: 20, IdleConnectionTimeoutSecs: 10, ResponseHeaderTimeoutSecs: 2, |
