diff options
| author | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
| commit | 3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch) | |
| tree | 517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/provider | |
| parent | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (diff) | |
feat: complete commercial developer workflowspublish-commercial-control-plane
Add tenant-safe usage observability, prepaid billing controls, API key lifecycle management, Embeddings metering, configurable billing alerts, and resilient provider health propagation. Harden Stripe failure handling, migrations, readiness, and the authenticated control-plane UI with end-to-end verification evidence.
Diffstat (limited to 'internal/provider')
| -rw-r--r-- | internal/provider/forwarder.go | 24 | ||||
| -rw-r--r-- | internal/provider/forwarder_test.go | 18 |
2 files changed, 38 insertions, 4 deletions
diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go index d7d850d..f816f8a 100644 --- a/internal/provider/forwarder.go +++ b/internal/provider/forwarder.go @@ -19,9 +19,10 @@ import ( ) type Result struct { - Response *http.Response - Route domain.Route - Attempts int + Response *http.Response + Route domain.Route + Attempts int + AttemptStartedAt time.Time } type Forwarder struct { @@ -83,7 +84,7 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque lastErr = fmt.Errorf("upstream %s returned %d", route.Provider.ID, response.StatusCode) continue } - return Result{Response: response, Route: route, Attempts: attempts}, nil + return Result{Response: response, Route: route, Attempts: attempts, AttemptStartedAt: attemptStarted}, nil } if lastErr == nil { lastErr = errors.New("all upstream routes failed") @@ -99,6 +100,19 @@ func (f *Forwarder) observe(modelID string, route domain.Route, statusCode int, providerhealth.Observation{StatusCode: statusCode, Latency: latency, Failed: failed}) } +// ObserveTTFT feeds the first user-visible output latency into adaptive route +// selection. It is deliberately separate from the header/circuit observation +// because streaming TTFT is only known after response forwarding begins. +func (f *Forwarder) ObserveTTFT(modelID string, route domain.Route, latency time.Duration) { + if f.health == nil || latency <= 0 { + return + } + f.health.ObserveTTFT(providerhealth.RouteKey{ModelID: modelID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}, latency) + if f.metrics != nil { + f.metrics.UpstreamTTFT(latency) + } +} + func rewriteModel(body []byte, upstreamModel string) ([]byte, error) { return rewriteRequest(body, upstreamModel, domain.ProtocolOpenAI) } @@ -141,6 +155,8 @@ func rewriteRequestWithWireAPI(body []byte, upstreamModel string, protocol domai func endpointURL(provider domain.Provider, _ domain.Protocol) string { baseURL := strings.TrimRight(provider.BaseURL, "/") switch provider.EffectiveWireAPI() { + case "embeddings": + return baseURL + "/embeddings" case "responses": return baseURL + "/responses" case "messages": diff --git a/internal/provider/forwarder_test.go b/internal/provider/forwarder_test.go index e9751ce..2ae81af 100644 --- a/internal/provider/forwarder_test.go +++ b/internal/provider/forwarder_test.go @@ -58,3 +58,21 @@ func TestResponsesWireAPIUsesResponsesEndpointWithoutChatStreamOptions(t *testin t.Fatalf("endpoint URL = %q", got) } } + +func TestEmbeddingsWireAPIUsesEmbeddingsEndpoint(t *testing.T) { + result, err := rewriteRequestWithWireAPI([]byte(`{"model":"public/model","input":["one","two"]}`), "embedding-upstream", domain.ProtocolOpenAIEmbeddings, "embeddings") + if err != nil { + t.Fatal(err) + } + var body map[string]json.RawMessage + if err := json.Unmarshal(result, &body); err != nil { + t.Fatal(err) + } + if string(body["model"]) != `"embedding-upstream"` || string(body["input"]) != `["one","two"]` { + t.Fatalf("unexpected rewritten body: %s", result) + } + provider := domain.Provider{BaseURL: "https://example.test/v1", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings"} + if got := endpointURL(provider, domain.ProtocolOpenAIEmbeddings); got != "https://example.test/v1/embeddings" { + t.Fatalf("endpoint URL = %q", got) + } +} |
