summaryrefslogtreecommitdiff
path: root/internal/httpapi
diff options
context:
space:
mode:
Diffstat (limited to 'internal/httpapi')
-rw-r--r--internal/httpapi/api.go62
-rw-r--r--internal/httpapi/api_test.go161
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,