summaryrefslogtreecommitdiff
path: root/internal/provider
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-06 15:58:57 +1200
committerChia <Chia@93.nz>2026-08-06 15:58:57 +1200
commit3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch)
tree517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/provider
parent41e322c53d7b4b796eb377d0df9c29ecd10ba431 (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.go24
-rw-r--r--internal/provider/forwarder_test.go18
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)
+ }
+}