From a5bfad04b3f3023bf289f60459cdb08c27b7c189 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Tue, 18 Aug 2026 16:40:40 +0200 Subject: [PATCH 1/5] feat(anthropic): support Claude subscription tokens and native /v1/messages forwarding --- .env.template | 2 + docs/advanced/anthropic-messages-api.mdx | 21 +++ docs/guides/claude-code.mdx | 36 ++++- docs/providers/anthropic.mdx | 24 ++++ internal/providers/anthropic/anthropic.go | 53 +++++++- .../providers/anthropic/anthropic_test.go | 114 ++++++++++++++++ internal/providers/keyring.go | 10 ++ internal/server/messages_handler.go | 3 + internal/server/messages_handler_test.go | 106 +++++++++++++++ internal/server/messages_native.go | 127 ++++++++++++++++++ internal/server/messages_native_test.go | 107 +++++++++++++++ internal/server/passthrough_support.go | 27 +++- internal/server/passthrough_support_test.go | 9 +- internal/usage/stream_observer.go | 72 +++++++++- internal/usage/stream_observer_test.go | 52 +++++++ 15 files changed, 742 insertions(+), 21 deletions(-) create mode 100644 internal/server/messages_native.go create mode 100644 internal/server/messages_native_test.go diff --git a/.env.template b/.env.template index 9f7c53e3f..8ea4db1d8 100644 --- a/.env.template +++ b/.env.template @@ -428,6 +428,8 @@ # OPENAI_BASE_URL=https://api.openai.com/v1 # Anthropic +# Accepts a Console API key (sk-ant-api...) or a Claude subscription OAuth +# token from `claude setup-token` (sk-ant-oat...; Claude Code traffic only). # ANTHROPIC_API_KEY=sk-ant-... # ANTHROPIC_BASE_URL=https://api.anthropic.com/v1 # Anthropic /v1/messages requires max_tokens. When the caller omits it, GoModel diff --git a/docs/advanced/anthropic-messages-api.mdx b/docs/advanced/anthropic-messages-api.mdx index e9c6028e3..0d01f7077 100644 --- a/docs/advanced/anthropic-messages-api.mdx +++ b/docs/advanced/anthropic-messages-api.mdx @@ -22,6 +22,22 @@ This differs from the [passthrough API](/features/passthrough-api): `/p/anthropi forwards bytes verbatim to the Anthropic upstream only, while the managed `/v1/messages` endpoint routes anywhere and is fully managed. +## Native forwarding to Anthropic + +When a `/v1/messages` request resolves to an **Anthropic** provider, GoModel +skips the translation round-trip and forwards the original request body +verbatim (rewriting only the `model` field when an alias resolved to a +different name), then relays the provider-native response or SSE stream +unchanged. This preserves everything the canonical translation cannot — +`cache_control` breakpoints, thinking-block signatures, `anthropic-beta` +headers — which coding agents like Claude Code depend on. Rate limits, +budgets, audit logging, and streaming usage tracking still apply. + +Native forwarding is automatic. Requests fall back to the translated pipeline +when a feature that operates on the canonical request is in play: guardrails +request patching, the response cache, or failover routing. Requests resolving +to any non-Anthropic provider always translate. + ## Supported endpoints | Endpoint | Behavior | @@ -154,6 +170,11 @@ is recorded under the `/v1/messages` endpoint so it can be filtered in the dashb ## Limitations +These limitations apply to the **translated** pipeline — requests routed to a +non-Anthropic provider, or to Anthropic with guardrails, response cache, or +failover engaged. Requests [natively forwarded to Anthropic](#native-forwarding-to-anthropic) +are preserved byte-exactly and none of the below applies. + `/v1/messages` translates through GoModel's canonical chat type. Anthropic-specific features that have no canonical equivalent are not preserved end to end: diff --git a/docs/guides/claude-code.mdx b/docs/guides/claude-code.mdx index d3e181a68..596971fb5 100644 --- a/docs/guides/claude-code.mdx +++ b/docs/guides/claude-code.mdx @@ -22,11 +22,14 @@ Flow: Claude Code can be routed through GoModel whether or not you personally use a Claude Code subscription. For gateway mode, Claude Code talks to GoModel with `ANTHROPIC_BASE_URL` and `ANTHROPIC_AUTH_TOKEN`. GoModel still needs its own - `ANTHROPIC_API_KEY` to reach Anthropic upstream. + `ANTHROPIC_API_KEY` to reach Anthropic upstream — either a Console API key + or a Claude subscription OAuth token. ## How to get `ANTHROPIC_API_KEY` +**Option A — Console API key (pay-as-you-go):** + 1. Open the Claude Console and sign in to your API account. 2. Go to account settings in Console, then create an API key. 3. Copy the key once and set it for GoModel as `ANTHROPIC_API_KEY`. @@ -39,6 +42,24 @@ settings. are separate. API key usage is billed as API usage. +**Option B — Claude subscription token (Pro/Max):** + +1. Run `claude setup-token` on a machine where Claude Code is logged in to + your subscription. +2. Copy the generated `sk-ant-oat01-...` token and set it for GoModel as + `ANTHROPIC_API_KEY`. + +GoModel detects the `sk-ant-oat` prefix and authenticates upstream with the +OAuth Bearer scheme automatically. Usage is covered by your subscription's +limits instead of API billing. + + + Anthropic authorizes subscription tokens **only for Claude Code traffic**. + This setup works because Claude Code remains the client and GoModel relays + its requests unchanged; pointing other clients or tools at the same gateway + credential is rejected upstream and is against Anthropic's terms of service. + + ## 1. Run GoModel Start GoModel with a master key and an Anthropic provider key: @@ -149,14 +170,23 @@ If the gateway is wired correctly, the response will contain `ok`. ## 3. Configure Claude Code to use GoModel -Point Claude Code at GoModel's Anthropic passthrough: +Point Claude Code at GoModel: ```bash -export ANTHROPIC_BASE_URL=http://localhost:8080/p/anthropic +export ANTHROPIC_BASE_URL=http://localhost:8080 export ANTHROPIC_AUTH_TOKEN=change-me export CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS=1 ``` +The managed base URL serves `/v1/messages` with +[native forwarding](/advanced/anthropic-messages-api#native-forwarding-to-anthropic): +requests that resolve to Anthropic are relayed byte-exactly (preserving +`cache_control`, thinking signatures, and beta headers) while rate limits, +budgets, audit, and usage tracking apply — and model aliases or virtual models +can route Claude Code to other providers. Alternatively, +`ANTHROPIC_BASE_URL=http://localhost:8080/p/anthropic` uses raw passthrough +pinned to the Anthropic upstream. + Short Claude Code doc summary: for gateway mode, set `ANTHROPIC_BASE_URL` to your gateway URL and `ANTHROPIC_AUTH_TOKEN` to your gateway token, then run Claude Code normally. See the official guide: diff --git a/docs/providers/anthropic.mdx b/docs/providers/anthropic.mdx index 6db639e80..0af4fc18f 100644 --- a/docs/providers/anthropic.mdx +++ b/docs/providers/anthropic.mdx @@ -32,6 +32,30 @@ providers: it, keeping the OpenAI-compatible surface lenient. +## Claude subscription (OAuth token) + +GoModel also accepts a Claude subscription OAuth token as the Anthropic +credential. Generate one with `claude setup-token` (requires a Claude +Pro/Max subscription and the Claude Code CLI) and set it as the provider key: + +```bash +ANTHROPIC_API_KEY=sk-ant-oat01-... +``` + +Tokens with the `sk-ant-oat` prefix are detected automatically: GoModel sends +them as `Authorization: Bearer` with the `oauth-2025-04-20` beta instead of +`x-api-key`. No extra configuration is needed. + + + Anthropic authorizes subscription OAuth tokens **only for Claude Code + traffic**. Use this to route your own Claude Code through GoModel (see the + [Claude Code guide](/guides/claude-code)); requests from other clients are + rejected upstream with "This credential is only authorized for use with + Claude Code". Endpoints outside the Claude Code surface (such as model + listing) may also be rejected — if provider model discovery fails, configure + the `models` list for the provider explicitly. + + ## Reasoning effort mapping GoModel accepts the OpenAI-shaped `"reasoning": {"effort": "..."}` object as diff --git a/internal/providers/anthropic/anthropic.go b/internal/providers/anthropic/anthropic.go index c25e6887d..a873b25ac 100644 --- a/internal/providers/anthropic/anthropic.go +++ b/internal/providers/anthropic/anthropic.go @@ -33,8 +33,21 @@ var Registration = providers.Registration{ const ( defaultBaseURL = "https://api.anthropic.com/v1" anthropicAPIVersion = "2023-06-01" + + // oauthTokenPrefix identifies Claude subscription OAuth tokens (created + // with `claude setup-token`). Anthropic only authorizes these credentials + // for Claude Code-shaped traffic; they authenticate with a Bearer header + // plus the oauth beta instead of x-api-key. + oauthTokenPrefix = "sk-ant-oat" + oauthBetaFlag = "oauth-2025-04-20" + + anthropicBetaHeader = "anthropic-beta" ) +func isOAuthToken(key string) bool { + return strings.HasPrefix(key, oauthTokenPrefix) +} + var allowedAnthropicImageMediaTypes = map[string]struct{}{ "image/jpeg": {}, "image/png": {}, @@ -158,7 +171,13 @@ func (p *Provider) getBatchResultEndpoints(batchID string) map[string]string { // setHeaders sets the required headers for Anthropic API requests. It runs once // per outbound request; identified sessions resolve to a stable key. func (p *Provider) setHeaders(req *http.Request) { - req.Header.Set("x-api-key", p.keys.NextForContext(req.Context())) + key := p.keys.NextForContext(req.Context()) + if isOAuthToken(key) { + req.Header.Set("Authorization", "Bearer "+key) + req.Header.Set(anthropicBetaHeader, oauthBetaFlag) + } else { + req.Header.Set("x-api-key", key) + } req.Header.Set("anthropic-version", anthropicAPIVersion) // Forward request ID if present in context @@ -167,12 +186,42 @@ func (p *Provider) setHeaders(req *http.Request) { } } +// ensureOAuthBeta returns headers with the oauth beta flag merged into a +// client-supplied anthropic-beta value. Forwarded headers override the ones set +// by setHeaders, so a client that sends its own beta list would otherwise drop +// the oauth flag subscription tokens require. Headers without an anthropic-beta +// entry are returned unchanged: setHeaders' value survives in that case. +func ensureOAuthBeta(headers http.Header) http.Header { + for name, values := range headers { + if !strings.EqualFold(strings.TrimSpace(name), anthropicBetaHeader) { + continue + } + for _, value := range values { + for flag := range strings.SplitSeq(value, ",") { + if strings.TrimSpace(flag) == oauthBetaFlag { + return headers + } + } + } + merged := make(http.Header, len(headers)) + maps.Copy(merged, headers) + merged[name] = append(append([]string{}, values...), oauthBetaFlag) + return merged + } + return headers +} + // Passthrough forwards an opaque Anthropic-native request without typed translation. func (p *Provider) Passthrough(ctx context.Context, req *core.PassthroughRequest) (*core.PassthroughResponse, error) { if req == nil { return nil, core.NewInvalidRequestError("passthrough request is required", nil) } + headers := req.Headers + if p.keys.Any(isOAuthToken) { + headers = ensureOAuthBeta(headers) + } + resp, err := p.client.DoPassthrough(ctx, llmclient.Request{ Method: req.Method, Endpoint: providers.PassthroughEndpoint(req.Endpoint), @@ -181,7 +230,7 @@ func (p *Provider) Passthrough(ctx context.Context, req *core.PassthroughRequest Stream: req.Stream, StreamUncertain: req.StreamUncertain, RawBodyReader: req.Body, - Headers: req.Headers, + Headers: headers, }) if err != nil { return nil, err diff --git a/internal/providers/anthropic/anthropic_test.go b/internal/providers/anthropic/anthropic_test.go index cd62511d5..74b85ee51 100644 --- a/internal/providers/anthropic/anthropic_test.go +++ b/internal/providers/anthropic/anthropic_test.go @@ -5199,6 +5199,120 @@ func TestPassthrough(t *testing.T) { } } +func TestSetHeadersOAuthToken(t *testing.T) { + tests := []struct { + name string + key string + wantAPIKey string + wantAuth string + wantBeta string + }{ + { + name: "api key uses x-api-key", + key: "sk-ant-api03-abc", + wantAPIKey: "sk-ant-api03-abc", + }, + { + name: "oauth token uses bearer and oauth beta", + key: "sk-ant-oat01-abc", + wantAuth: "Bearer sk-ant-oat01-abc", + wantBeta: oauthBetaFlag, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + p := &Provider{keys: providers.NewKeyring(tt.key)} + req := httptest.NewRequest(http.MethodPost, "/messages", nil) + p.setHeaders(req) + + if got := req.Header.Get("x-api-key"); got != tt.wantAPIKey { + t.Errorf("x-api-key = %q, want %q", got, tt.wantAPIKey) + } + if got := req.Header.Get("Authorization"); got != tt.wantAuth { + t.Errorf("Authorization = %q, want %q", got, tt.wantAuth) + } + if got := req.Header.Get(anthropicBetaHeader); got != tt.wantBeta { + t.Errorf("anthropic-beta = %q, want %q", got, tt.wantBeta) + } + if got := req.Header.Get("anthropic-version"); got != anthropicAPIVersion { + t.Errorf("anthropic-version = %q, want %q", got, anthropicAPIVersion) + } + }) + } +} + +func TestPassthroughOAuthToken(t *testing.T) { + tests := []struct { + name string + clientBeta string + wantBeta []string + }{ + { + name: "no client beta keeps provider oauth beta", + wantBeta: []string{oauthBetaFlag}, + }, + { + name: "client beta merged with oauth flag", + clientBeta: "claude-code-20250219,interleaved-thinking-2025-05-14", + wantBeta: []string{"claude-code-20250219,interleaved-thinking-2025-05-14", oauthBetaFlag}, + }, + { + name: "client beta already containing oauth flag unchanged", + clientBeta: "claude-code-20250219," + oauthBetaFlag, + wantBeta: []string{"claude-code-20250219," + oauthBetaFlag}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var gotAuth, gotAPIKey string + var gotBeta []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + gotAPIKey = r.Header.Get("x-api-key") + gotBeta = r.Header.Values(anthropicBetaHeader) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{}`)) + })) + defer server.Close() + + provider := NewWithHTTPClient("sk-ant-oat01-abc", server.Client(), llmclient.Hooks{}) + provider.SetBaseURL(server.URL) + + headers := http.Header{"Content-Type": {"application/json"}} + if tt.clientBeta != "" { + headers.Set(anthropicBetaHeader, tt.clientBeta) + } + resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ + Method: http.MethodPost, + Endpoint: "messages", + Body: io.NopCloser(strings.NewReader(`{"model":"claude-sonnet-5"}`)), + Headers: headers, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer func() { + _ = resp.Body.Close() + }() + + if gotAuth != "Bearer sk-ant-oat01-abc" { + t.Errorf("Authorization = %q, want Bearer token", gotAuth) + } + if gotAPIKey != "" { + t.Errorf("x-api-key = %q, want empty", gotAPIKey) + } + if len(gotBeta) != len(tt.wantBeta) { + t.Fatalf("anthropic-beta values = %v, want %v", gotBeta, tt.wantBeta) + } + for i := range gotBeta { + if gotBeta[i] != tt.wantBeta[i] { + t.Errorf("anthropic-beta[%d] = %q, want %q", i, gotBeta[i], tt.wantBeta[i]) + } + } + }) + } +} + func TestResolveDefaultMaxTokens(t *testing.T) { tests := []struct { name string diff --git a/internal/providers/keyring.go b/internal/providers/keyring.go index 9bcc510a5..65aa6db9d 100644 --- a/internal/providers/keyring.go +++ b/internal/providers/keyring.go @@ -5,6 +5,7 @@ import ( "context" "crypto/hmac" "crypto/sha256" + "slices" "sync/atomic" "github.com/enterpilot/gomodel/internal/core" @@ -144,6 +145,15 @@ func rendezvousKeyScore(sessionID, key string) [sha256.Size]byte { return score } +// Any reports whether at least one configured key satisfies pred. It never +// advances the rotation. An empty or nil ring reports false. +func (k *Keyring) Any(pred func(string) bool) bool { + if k == nil || pred == nil { + return false + } + return slices.ContainsFunc(k.keys, pred) +} + // Primary returns the first configured key without advancing the rotation. // It is the key to use where a stable identity matters more than spreading // load, and where an empty ring must stay empty. diff --git a/internal/server/messages_handler.go b/internal/server/messages_handler.go index 9d9e31188..8bf934dea 100644 --- a/internal/server/messages_handler.go +++ b/internal/server/messages_handler.go @@ -163,6 +163,9 @@ func (s *translatedInferenceService) Messages(c *echo.Context) error { } attachPreparedWorkflow(c, ctx, workflow) + if s.canForwardMessagesNatively(workflow) { + return s.dispatchMessagesNative(c, prepared, workflow) + } return handleWithCache(s, c, prepared, workflow, s.dispatchMessages) } diff --git a/internal/server/messages_handler_test.go b/internal/server/messages_handler_test.go index 4ad31c82d..7798e0e16 100644 --- a/internal/server/messages_handler_test.go +++ b/internal/server/messages_handler_test.go @@ -2,6 +2,7 @@ package server import ( "encoding/json" + "io" "net/http" "net/http/httptest" "strings" @@ -127,6 +128,111 @@ func TestMessages_Streaming(t *testing.T) { } } +func TestMessages_NativeAnthropicForwarding(t *testing.T) { + nativeResponse := `{"id":"msg_1","type":"message","role":"assistant","content":[{"type":"text","text":"native"}],"model":"claude-test","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}` + provider := &mockProvider{ + supportedModels: []string{"claude-test"}, + providerTypes: map[string]string{"claude-test": "anthropic"}, + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: map[string][]string{"Content-Type": {"application/json"}}, + Body: io.NopCloser(strings.NewReader(nativeResponse)), + }, + } + + e := echo.New() + handler := NewHandler(provider, nil, nil, nil) + + // cache_control does not survive the translated pipeline; the native path + // must forward it verbatim. + reqBody := `{"model":"claude-test","max_tokens":64,"messages":[{"role":"user","content":[{"type":"text","text":"Hi","cache_control":{"type":"ephemeral"}}]}]}` + req := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(reqBody)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("anthropic-beta", "claude-code-20250219") + rec := httptest.NewRecorder() + + if err := handler.Messages(e.NewContext(req, rec)); err != nil { + t.Fatalf("Messages: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, body = %s", rec.Code, rec.Body.String()) + } + if rec.Body.String() != nativeResponse { + t.Errorf("body = %s, want provider-native response verbatim", rec.Body.String()) + } + + if provider.lastPassthroughProvider != "anthropic" { + t.Fatalf("passthrough provider = %q, want anthropic", provider.lastPassthroughProvider) + } + forwarded := provider.lastPassthroughReq + if forwarded == nil { + t.Fatal("provider did not receive a passthrough request") + } + if forwarded.Endpoint != "messages" || forwarded.Method != http.MethodPost { + t.Errorf("endpoint = %q method = %q", forwarded.Endpoint, forwarded.Method) + } + forwardedBody, err := io.ReadAll(forwarded.Body) + if err != nil { + t.Fatalf("read forwarded body: %v", err) + } + if string(forwardedBody) != reqBody { + t.Errorf("forwarded body = %s, want original request verbatim", forwardedBody) + } + if got := forwarded.Headers.Get("anthropic-beta"); got != "claude-code-20250219" { + t.Errorf("forwarded anthropic-beta = %q", got) + } +} + +func TestRewriteMessagesModel(t *testing.T) { + tests := []struct { + name string + body string + model string + want string + }{ + { + name: "same model leaves body untouched", + body: `{"model":"claude-test","max_tokens":1,"messages":[]}`, + model: "claude-test", + want: `{"model":"claude-test","max_tokens":1,"messages":[]}`, + }, + { + name: "empty model leaves body untouched", + body: `{"model":"claude-test"}`, + model: "", + want: `{"model":"claude-test"}`, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := rewriteMessagesModel([]byte(tt.body), tt.model) + if err != nil { + t.Fatalf("rewriteMessagesModel: %v", err) + } + if string(got) != tt.want { + t.Errorf("body = %s, want %s", got, tt.want) + } + }) + } + + t.Run("alias rewrites model field", func(t *testing.T) { + got, err := rewriteMessagesModel([]byte(`{"model":"my-alias","max_tokens":1}`), "claude-test") + if err != nil { + t.Fatalf("rewriteMessagesModel: %v", err) + } + var fields map[string]any + if err := json.Unmarshal(got, &fields); err != nil { + t.Fatalf("unmarshal rewritten body: %v", err) + } + if fields["model"] != "claude-test" { + t.Errorf("model = %v, want claude-test", fields["model"]) + } + if fields["max_tokens"] != float64(1) { + t.Errorf("max_tokens = %v, want 1", fields["max_tokens"]) + } + }) +} + func TestMessages_InvalidRequestReturnsAnthropicError(t *testing.T) { provider := &mockProvider{supportedModels: []string{"claude-test"}} e := echo.New() diff --git a/internal/server/messages_native.go b/internal/server/messages_native.go new file mode 100644 index 000000000..9a222fb5c --- /dev/null +++ b/internal/server/messages_native.go @@ -0,0 +1,127 @@ +package server + +import ( + "bytes" + "io" + "net/http" + "strings" + + "github.com/goccy/go-json" + "github.com/labstack/echo/v5" + + "github.com/enterpilot/gomodel/internal/auditlog" + "github.com/enterpilot/gomodel/internal/core" +) + +const anthropicProviderType = "anthropic" + +// canForwardMessagesNatively reports whether a prepared /v1/messages request +// can skip the translated pipeline and be forwarded to the provider in its +// original Anthropic dialect. Native forwarding preserves fields the canonical +// translation cannot round-trip (cache_control breakpoints, thinking block +// signatures, anthropic-beta headers), which Claude Code clients depend on. +// Features that operate on the canonical translated request take precedence: +// requests using guardrails patching, response caching, or failover stay on +// the translated pipeline. +func (s *translatedInferenceService) canForwardMessagesNatively(workflow *core.Workflow) bool { + if workflow == nil || strings.TrimSpace(workflow.ProviderType) != anthropicProviderType { + return false + } + if s.translatedRequestPatcher != nil { + return false + } + if s.responseCache != nil && workflow.CacheEnabled() { + return false + } + if len(s.inference().FailoverSelectors(workflow)) > 0 { + return false + } + _, ok := s.provider.(core.RoutablePassthrough) + return ok +} + +// dispatchMessagesNative forwards the original Anthropic Messages body to the +// resolved Anthropic provider and relays the provider-native response (JSON or +// SSE) unchanged, with admission, audit, and streaming usage accounting. +func (s *translatedInferenceService) dispatchMessagesNative(c *echo.Context, req *core.ChatRequest, workflow *core.Workflow) error { + passthroughProvider, ok := s.provider.(core.RoutablePassthrough) + if !ok { + return handleError(c, core.NewInvalidRequestError("provider passthrough is not supported by the current provider router", nil)) + } + + body, err := requestBodyBytes(c) + if err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + body, err = rewriteMessagesModel(body, req.Model) + if err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + + s.observeLiveProviderAttempts(c, workflow) + + adm, err := enforceAdmission(c, s.rateLimiter, s.budgetChecker, rateLimitRouteFromWorkflow(workflow)) + if err != nil { + return handleError(c, err) + } + defer adm.release() + ctx := adm.dispatchContext(c.Request().Context()) + + providerName := "" + if workflow.Resolution != nil { + providerName = workflow.Resolution.ProviderName + } + + resp, err := passthroughProvider.Passthrough(ctx, anthropicProviderType, &core.PassthroughRequest{ + Method: http.MethodPost, + Endpoint: "messages", + Operation: "anthropic.messages", + Model: req.Model, + Stream: req.Stream, + Body: io.NopCloser(bytes.NewReader(body)), + Headers: buildPassthroughHeaders(ctx, c.Request().Header), + ProviderName: providerName, + }) + if err != nil { + return handleError(c, err) + } + + auditlog.EnrichEntryWithWorkflow(c, workflow) + info := &core.PassthroughRouteInfo{ + Provider: anthropicProviderType, + ProviderName: providerName, + NormalizedEndpoint: "messages", + SemanticOperation: "anthropic.messages", + GenAIOperation: "chat", + Stream: req.Stream, + AuditPath: "/v1/messages", + Model: req.Model, + } + return proxyPassthroughResponse(c, s.logger, s.usageLogger, s.pricingResolver, anthropicProviderType, providerName, "messages", info, resp) +} + +// rewriteMessagesModel returns body with its "model" field replaced by the +// resolved model, leaving the body untouched when it already matches so +// aliased/renamed models reach the provider under their real name. +func rewriteMessagesModel(body []byte, model string) ([]byte, error) { + if strings.TrimSpace(model) == "" { + return body, nil + } + var fields map[string]json.RawMessage + if err := json.Unmarshal(body, &fields); err != nil { + return nil, err + } + var current string + if raw, ok := fields["model"]; ok { + _ = json.Unmarshal(raw, ¤t) + } + if current == model { + return body, nil + } + encoded, err := json.Marshal(model) + if err != nil { + return nil, err + } + fields["model"] = encoded + return json.Marshal(fields) +} diff --git a/internal/server/messages_native_test.go b/internal/server/messages_native_test.go new file mode 100644 index 000000000..da4c7eecc --- /dev/null +++ b/internal/server/messages_native_test.go @@ -0,0 +1,107 @@ +package server + +import ( + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/labstack/echo/v5" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/usage" +) + +// A streaming /v1/messages request through the native forwarding path must +// record a usage entry combining message_start input tokens with the final +// message_delta output tokens. +func TestMessages_NativeStreamingLogsUsage(t *testing.T) { + anthropicSSE := strings.Join([]string{ + `event: message_start`, + `data: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-fable-5","usage":{"input_tokens":19560,"cache_creation_input_tokens":100,"cache_read_input_tokens":200,"output_tokens":3}}}`, + ``, + `event: content_block_start`, + `data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`, + ``, + `event: content_block_delta`, + `data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"ok"}}`, + ``, + `event: content_block_stop`, + `data: {"type":"content_block_stop","index":0}`, + ``, + `event: message_delta`, + `data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":31}}`, + ``, + `event: message_stop`, + `data: {"type":"message_stop"}`, + ``, + ``, + }, "\n") + + provider := &mockProvider{ + supportedModels: []string{"claude-fable-5"}, + providerTypes: map[string]string{"claude-fable-5": "anthropic"}, + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: map[string][]string{"Content-Type": {"text/event-stream; charset=utf-8"}}, + Body: io.NopCloser(strings.NewReader(anthropicSSE)), + }, + } + usageLogger := &collectingUsageLogger{config: usage.Config{Enabled: true}} + + e := echo.New() + handler := NewHandler(provider, nil, usageLogger, nil) + + reqBody := `{"model":"claude-fable-5","max_tokens":64,"stream":true,"messages":[{"role":"user","content":"Hi"}]}` + req := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(reqBody)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + + if err := handler.Messages(e.NewContext(req, rec)); err != nil { + t.Fatalf("Messages: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, body = %s", rec.Code, rec.Body.String()) + } + if provider.lastPassthroughReq == nil { + t.Fatal("native path was not taken") + } + if !strings.Contains(rec.Body.String(), "message_stop") { + t.Fatalf("stream not relayed: %s", rec.Body.String()) + } + + if len(usageLogger.entries) != 1 { + t.Fatalf("usage entries = %d, want 1", len(usageLogger.entries)) + } + entry := usageLogger.entries[0] + if entry.InputTokens != 19560 { + t.Errorf("InputTokens = %d, want 19560", entry.InputTokens) + } + if entry.OutputTokens != 31 { + t.Errorf("OutputTokens = %d, want 31", entry.OutputTokens) + } + if entry.RawData["cache_creation_input_tokens"] != 100 { + t.Errorf("cache_creation_input_tokens = %v, want 100", entry.RawData["cache_creation_input_tokens"]) + } + if entry.RawData["cache_read_input_tokens"] != 200 { + t.Errorf("cache_read_input_tokens = %v, want 200", entry.RawData["cache_read_input_tokens"]) + } +} + +// A forwarded Accept-Encoding would make the upstream body arrive compressed, +// blinding the SSE usage and audit observers; it must be stripped so the +// transport decompresses transparently. +func TestBuildPassthroughHeadersDropsAcceptEncoding(t *testing.T) { + src := http.Header{ + "Accept-Encoding": {"gzip, deflate, br, zstd"}, + "Anthropic-Beta": {"claude-code-20250219"}, + } + dst := buildPassthroughHeaders(t.Context(), src) + if got := dst.Get("Accept-Encoding"); got != "" { + t.Errorf("Accept-Encoding forwarded as %q, want stripped", got) + } + if got := dst.Get("Anthropic-Beta"); got != "claude-code-20250219" { + t.Errorf("Anthropic-Beta = %q, want preserved", got) + } +} diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index 2dd9cd80f..78b1f44a1 100644 --- a/internal/server/passthrough_support.go +++ b/internal/server/passthrough_support.go @@ -133,6 +133,14 @@ func skipPassthroughRequestHeader(key string, userPathHeader ...string) bool { if key == "" { return true } + // A forwarded Accept-Encoding makes Go's transport return the upstream + // body still compressed, which blinds the audit and usage SSE observers. + // Dropping it lets the transport negotiate gzip itself and hand back + // decoded bytes; the client then receives an uncompressed response with + // the Content-Encoding header removed by the transport. + if strings.EqualFold(key, "Accept-Encoding") { + return true + } if strings.EqualFold(key, core.UserPathHeader) { return true } @@ -243,6 +251,13 @@ func passthroughAuditPath(c *echo.Context, providerType, endpoint string, info * } func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerType, providerName, endpoint string, info *core.PassthroughRouteInfo, resp *core.PassthroughResponse) error { + return proxyPassthroughResponse(c, s.logger, s.usageLogger, s.pricingResolver, providerType, providerName, endpoint, info, resp) +} + +// proxyPassthroughResponse relays a provider-native response (JSON or SSE) to +// the client, attaching audit and usage stream observers. It is shared by the +// /p/ passthrough surface and the /v1/messages native forwarding path. +func proxyPassthroughResponse(c *echo.Context, logger auditlog.LoggerInterface, usageLogger usage.LoggerInterface, pricingResolver usage.PricingResolver, providerType, providerName, endpoint string, info *core.PassthroughRouteInfo, resp *core.PassthroughResponse) error { if resp == nil || resp.Body == nil { return handleError(c, core.NewProviderError(providerType, http.StatusBadGateway, "provider returned empty passthrough response", nil)) } @@ -269,17 +284,17 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT auditlog.MarkEntryAsStreaming(c, true) auditlog.EnrichEntryWithStream(c, true) workflow := core.GetWorkflow(c.Request().Context()) - auditEnabled := s.logger != nil && s.logger.Config().Enabled && (workflow == nil || workflow.AuditEnabled()) + auditEnabled := logger != nil && logger.Config().Enabled && (workflow == nil || workflow.AuditEnabled()) entry := auditlog.GetStreamEntryFromContext(c) if auditEnabled && entry != nil { - auditlog.PopulateRequestData(entry, c.Request(), s.logger.Config()) + auditlog.PopulateRequestData(entry, c.Request(), logger.Config()) } streamEntry := auditlog.CreateStreamEntry(c.Request().Context(), entry) if streamEntry != nil { streamEntry.StatusCode = resp.StatusCode } - if auditEnabled && streamEntry != nil && s.logger.Config().LogHeaders { + if auditEnabled && streamEntry != nil && logger.Config().LogHeaders { auditlog.PopulateResponseHeaders(streamEntry, c.Response().Header()) } @@ -297,12 +312,12 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT observers := make([]streaming.Observer, 0, 2) if auditEnabled && streamEntry != nil { - if observer := auditlog.NewStreamLogObserver(s.logger, streamEntry, auditPath); observer != nil { + if observer := auditlog.NewStreamLogObserver(logger, streamEntry, auditPath); observer != nil { observers = append(observers, observer) } } - if s.usageLogger != nil && s.usageLogger.Config().Enabled && (workflow == nil || workflow.UsageEnabled()) { - if observer := usage.NewStreamUsageObserver(s.usageLogger, model, providerType, requestID, usagePath, s.pricingResolver, core.UserPathFromContext(c.Request().Context())); observer != nil { + if usageLogger != nil && usageLogger.Config().Enabled && (workflow == nil || workflow.UsageEnabled()) { + if observer := usage.NewStreamUsageObserver(usageLogger, model, providerType, requestID, usagePath, pricingResolver, core.UserPathFromContext(c.Request().Context())); observer != nil { observer.SetProviderName(providerName) observer.SetSessionID(core.SessionIDFromContext(c.Request().Context())) observer.SetLabels(core.RequestLabelsFromContext(c.Request().Context())) diff --git a/internal/server/passthrough_support_test.go b/internal/server/passthrough_support_test.go index a8c2f00e8..cdc233a9e 100644 --- a/internal/server/passthrough_support_test.go +++ b/internal/server/passthrough_support_test.go @@ -3,6 +3,7 @@ package server import ( "context" "net/http" + "slices" "testing" "github.com/enterpilot/gomodel/internal/core" @@ -32,13 +33,7 @@ func TestBuildPassthroughHeadersSkipsConfiguredUserPathHeader(t *testing.T) { // and the default handler must not reject those requests before contacting the // upstream. Caught by greptile P1 on PR #701. func TestDefaultEnabledPassthroughProvidersIncludesHetzner(t *testing.T) { - found := false - for _, p := range defaultEnabledPassthroughProviders { - if p == "hetzner" { - found = true - break - } - } + found := slices.Contains(defaultEnabledPassthroughProviders, "hetzner") if !found { t.Fatalf("defaultEnabledPassthroughProviders = %v, want hetzner included", defaultEnabledPassthroughProviders) } diff --git a/internal/usage/stream_observer.go b/internal/usage/stream_observer.go index c8d8db110..46af62a78 100644 --- a/internal/usage/stream_observer.go +++ b/internal/usage/stream_observer.go @@ -100,9 +100,60 @@ func (o *StreamUsageObserver) WantsJSONEvent(raw []byte) bool { func (o *StreamUsageObserver) OnJSONEvent(chunk map[string]any) { entry := o.extractUsageFromEvent(chunk) - if entry != nil { - o.cachedEntry = entry + if entry == nil { + return + } + o.cachedEntry = o.mergeWithCachedEntry(entry) +} + +// mergeWithCachedEntry folds usage from an earlier event into the latest one. +// Providers that report usage across events (Anthropic: input tokens in +// message_start, output tokens in the final message_delta) need the pieces +// combined; providers with cumulative usage chunks (OpenAI) are unaffected +// because their later events carry every field. Costs are recomputed when the +// merge changes token counts. +func (o *StreamUsageObserver) mergeWithCachedEntry(entry *UsageEntry) *UsageEntry { + cached := o.cachedEntry + if cached == nil { + return entry + } + + merged := false + if entry.InputTokens == 0 && cached.InputTokens > 0 { + entry.InputTokens = cached.InputTokens + merged = true + } + if entry.OutputTokens == 0 && cached.OutputTokens > 0 { + entry.OutputTokens = cached.OutputTokens + merged = true + } + if entry.ProviderID == "" { + entry.ProviderID = cached.ProviderID + } + for key, value := range cached.RawData { + if entry.RawData == nil { + entry.RawData = make(map[string]any, len(cached.RawData)) + } + if _, exists := entry.RawData[key]; !exists { + entry.RawData[key] = value + merged = true + } + } + if !merged { + return entry } + + if entry.TotalTokens < entry.InputTokens+entry.OutputTokens { + entry.TotalTokens = entry.InputTokens + entry.OutputTokens + } + var pricingArgs []*core.ModelPricing + if o.pricingResolver != nil { + if p := o.pricingResolver.ResolvePricing(o.pricingModel(entry.Model), o.pricingProvider()); p != nil { + pricingArgs = append(pricingArgs, p) + } + } + applyUsageCosts(entry, o.provider, o.endpoint, pricingArgs...) + return entry } func (o *StreamUsageObserver) OnStreamClose() { @@ -125,7 +176,8 @@ func (o *StreamUsageObserver) extractUsageFromEvent(chunk map[string]any) *Usage usageRaw, ok := chunk["usage"] if !ok { - if eventType, _ := chunk["type"].(string); eventType == "response.completed" || eventType == "response.done" { + switch eventType, _ := chunk["type"].(string); eventType { + case "response.completed", "response.done": if response, respOK := chunk["response"].(map[string]any); respOK { usageRaw, ok = response["usage"] if id, idOK := response["id"].(string); idOK && id != "" { @@ -135,6 +187,20 @@ func (o *StreamUsageObserver) extractUsageFromEvent(chunk map[string]any) *Usage model = m } } + case "message_start": + // Anthropic-native streams report input_tokens (and cache token + // details) only here, nested under "message"; the final + // message_delta usually carries just output_tokens. Both events + // are extracted and merged in OnJSONEvent. + if message, msgOK := chunk["message"].(map[string]any); msgOK { + usageRaw, ok = message["usage"] + if id, idOK := message["id"].(string); idOK && id != "" { + providerID = id + } + if m, modelOK := message["model"].(string); modelOK && m != "" { + model = m + } + } } } if !ok { diff --git a/internal/usage/stream_observer_test.go b/internal/usage/stream_observer_test.go index d3f006f87..0f6d27079 100644 --- a/internal/usage/stream_observer_test.go +++ b/internal/usage/stream_observer_test.go @@ -717,3 +717,55 @@ data: [DONE] t.Errorf("RewriteCostSaved = %v, want nil without pricing", *entries[0].RewriteCostSaved) } } + +func TestStreamUsageObserverAnthropicNativeEvents(t *testing.T) { + logger := &trackingLogger{enabled: true} + observer := NewStreamUsageObserver(logger, "claude-fable-5", "anthropic", "req-native", "/v1/messages", nil) + + // Anthropic-native streams split usage across events: input tokens and + // cache details arrive in message_start, output tokens in the final + // message_delta. + observer.OnJSONEvent(map[string]any{ + "type": "message_start", + "message": map[string]any{ + "id": "msg_native", + "model": "claude-fable-5", + "usage": map[string]any{ + "input_tokens": float64(19560), + "output_tokens": float64(3), + "cache_creation_input_tokens": float64(100), + "cache_read_input_tokens": float64(200), + }, + }, + }) + observer.OnJSONEvent(map[string]any{ + "type": "message_delta", + "delta": map[string]any{"stop_reason": "end_turn"}, + "usage": map[string]any{"output_tokens": float64(31)}, + }) + observer.OnStreamClose() + + entries := logger.getEntries() + if len(entries) != 1 { + t.Fatalf("expected 1 entry, got %d", len(entries)) + } + entry := entries[0] + if entry.InputTokens != 19560 { + t.Errorf("InputTokens = %d, want 19560", entry.InputTokens) + } + if entry.OutputTokens != 31 { + t.Errorf("OutputTokens = %d, want 31", entry.OutputTokens) + } + if entry.TotalTokens != 19591 { + t.Errorf("TotalTokens = %d, want 19591", entry.TotalTokens) + } + if entry.ProviderID != "msg_native" { + t.Errorf("ProviderID = %q, want msg_native", entry.ProviderID) + } + if entry.RawData["cache_creation_input_tokens"] != 100 { + t.Errorf("cache_creation_input_tokens = %v, want 100", entry.RawData["cache_creation_input_tokens"]) + } + if entry.RawData["cache_read_input_tokens"] != 200 { + t.Errorf("cache_read_input_tokens = %v, want 200", entry.RawData["cache_read_input_tokens"]) + } +} From 9f3cfe28a0d2814f50e201e46d761c51b075391e Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Tue, 18 Aug 2026 17:10:06 +0200 Subject: [PATCH 2/5] fix(anthropic): address PR review findings for native messages forwarding - pin the passthrough credential so the oauth beta merge and auth header always describe the same key on mixed keyrings - splice only the model value when rewriting aliased native requests, preserving all other request bytes - docs: qualify byte-exact wording with the model rewrite; list Team/Enterprise as supported setup-token plans --- docs/advanced/anthropic-messages-api.mdx | 3 +- docs/guides/claude-code.mdx | 7 +- docs/providers/anthropic.mdx | 3 +- internal/providers/anthropic/anthropic.go | 21 +++++- .../providers/anthropic/anthropic_test.go | 69 +++++++++++++++++++ internal/providers/keyring.go | 10 --- internal/server/messages_handler_test.go | 23 ++++--- internal/server/messages_native.go | 65 ++++++++++++----- 8 files changed, 156 insertions(+), 45 deletions(-) diff --git a/docs/advanced/anthropic-messages-api.mdx b/docs/advanced/anthropic-messages-api.mdx index 0d01f7077..5695b0d12 100644 --- a/docs/advanced/anthropic-messages-api.mdx +++ b/docs/advanced/anthropic-messages-api.mdx @@ -173,7 +173,8 @@ is recorded under the `/v1/messages` endpoint so it can be filtered in the dashb These limitations apply to the **translated** pipeline — requests routed to a non-Anthropic provider, or to Anthropic with guardrails, response cache, or failover engaged. Requests [natively forwarded to Anthropic](#native-forwarding-to-anthropic) -are preserved byte-exactly and none of the below applies. +are preserved byte-for-byte apart from the `model` value when an alias +resolved to a different name, and none of the below applies. `/v1/messages` translates through GoModel's canonical chat type. Anthropic-specific features that have no canonical equivalent are not preserved end to end: diff --git a/docs/guides/claude-code.mdx b/docs/guides/claude-code.mdx index 596971fb5..0ecdd215d 100644 --- a/docs/guides/claude-code.mdx +++ b/docs/guides/claude-code.mdx @@ -42,7 +42,7 @@ settings. are separate. API key usage is billed as API usage. -**Option B — Claude subscription token (Pro/Max):** +**Option B — Claude subscription token (Pro, Max, Team, or Enterprise):** 1. Run `claude setup-token` on a machine where Claude Code is logged in to your subscription. @@ -180,8 +180,9 @@ export CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS=1 The managed base URL serves `/v1/messages` with [native forwarding](/advanced/anthropic-messages-api#native-forwarding-to-anthropic): -requests that resolve to Anthropic are relayed byte-exactly (preserving -`cache_control`, thinking signatures, and beta headers) while rate limits, +requests that resolve to Anthropic are relayed unchanged apart from model +alias resolution (preserving `cache_control`, thinking signatures, and beta +headers) while rate limits, budgets, audit, and usage tracking apply — and model aliases or virtual models can route Claude Code to other providers. Alternatively, `ANTHROPIC_BASE_URL=http://localhost:8080/p/anthropic` uses raw passthrough diff --git a/docs/providers/anthropic.mdx b/docs/providers/anthropic.mdx index 0af4fc18f..b49cdec68 100644 --- a/docs/providers/anthropic.mdx +++ b/docs/providers/anthropic.mdx @@ -36,7 +36,8 @@ providers: GoModel also accepts a Claude subscription OAuth token as the Anthropic credential. Generate one with `claude setup-token` (requires a Claude -Pro/Max subscription and the Claude Code CLI) and set it as the provider key: +subscription — Pro, Max, Team, or Enterprise — and the Claude Code CLI) and +set it as the provider key: ```bash ANTHROPIC_API_KEY=sk-ant-oat01-... diff --git a/internal/providers/anthropic/anthropic.go b/internal/providers/anthropic/anthropic.go index a873b25ac..29a75d287 100644 --- a/internal/providers/anthropic/anthropic.go +++ b/internal/providers/anthropic/anthropic.go @@ -168,10 +168,23 @@ func (p *Provider) getBatchResultEndpoints(batchID string) map[string]string { return cloned } +// pinnedKeyContextKey carries a credential selected before the header hook +// runs. Passthrough pins its key so header adaptation (the oauth beta merge) +// and the auth header always describe the same credential, even when the +// keyring mixes OAuth tokens and API keys. +type pinnedKeyContextKey struct{} + +func withPinnedKey(ctx context.Context, key string) context.Context { + return context.WithValue(ctx, pinnedKeyContextKey{}, key) +} + // setHeaders sets the required headers for Anthropic API requests. It runs once // per outbound request; identified sessions resolve to a stable key. func (p *Provider) setHeaders(req *http.Request) { - key := p.keys.NextForContext(req.Context()) + key, pinned := req.Context().Value(pinnedKeyContextKey{}).(string) + if !pinned { + key = p.keys.NextForContext(req.Context()) + } if isOAuthToken(key) { req.Header.Set("Authorization", "Bearer "+key) req.Header.Set(anthropicBetaHeader, oauthBetaFlag) @@ -217,8 +230,12 @@ func (p *Provider) Passthrough(ctx context.Context, req *core.PassthroughRequest return nil, core.NewInvalidRequestError("passthrough request is required", nil) } + // Select the credential once and pin it for setHeaders, so the beta + // merge below and the auth header are always based on the same key. + key := p.keys.NextForContext(ctx) + ctx = withPinnedKey(ctx, key) headers := req.Headers - if p.keys.Any(isOAuthToken) { + if isOAuthToken(key) { headers = ensureOAuthBeta(headers) } diff --git a/internal/providers/anthropic/anthropic_test.go b/internal/providers/anthropic/anthropic_test.go index 74b85ee51..b4f6274d8 100644 --- a/internal/providers/anthropic/anthropic_test.go +++ b/internal/providers/anthropic/anthropic_test.go @@ -5313,6 +5313,75 @@ func TestPassthroughOAuthToken(t *testing.T) { } } +// A keyring mixing OAuth tokens and API keys must keep each request +// self-consistent: the oauth beta merge and the auth header always describe +// the credential actually dispatched. +func TestPassthroughMixedKeyringConsistency(t *testing.T) { + type observed struct { + auth string + apiKey string + beta string + } + var got []observed + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + got = append(got, observed{ + auth: r.Header.Get("Authorization"), + apiKey: r.Header.Get("x-api-key"), + beta: strings.Join(r.Header.Values(anthropicBetaHeader), ","), + }) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{}`)) + })) + defer server.Close() + + p := &Provider{ + keys: providers.NewKeyring("sk-ant-oat01-a", "sk-ant-api03-b"), + batchResultEndpoints: make(map[string]map[string]string), + } + cfg := llmclient.DefaultConfig("anthropic", server.URL) + p.client = llmclient.NewWithHTTPClient(server.Client(), cfg, p.setHeaders) + + for range 4 { + headers := http.Header{"Content-Type": {"application/json"}} + headers.Set(anthropicBetaHeader, "claude-code-20250219") + resp, err := p.Passthrough(context.Background(), &core.PassthroughRequest{ + Method: http.MethodPost, + Endpoint: "messages", + Body: io.NopCloser(strings.NewReader(`{"model":"claude-sonnet-5"}`)), + Headers: headers, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + _ = resp.Body.Close() + } + + sawOAuth, sawAPIKey := false, false + for i, o := range got { + hasOAuthBeta := strings.Contains(o.beta, oauthBetaFlag) + switch { + case o.auth != "": + sawOAuth = true + if o.apiKey != "" { + t.Errorf("request %d: both Authorization and x-api-key set", i) + } + if !hasOAuthBeta { + t.Errorf("request %d: OAuth credential without oauth beta (beta = %q)", i, o.beta) + } + case o.apiKey != "": + sawAPIKey = true + if hasOAuthBeta { + t.Errorf("request %d: API key with oauth beta (beta = %q)", i, o.beta) + } + default: + t.Errorf("request %d: no credential sent", i) + } + } + if !sawOAuth || !sawAPIKey { + t.Fatalf("rotation did not cover both credentials (oauth=%v apiKey=%v)", sawOAuth, sawAPIKey) + } +} + func TestResolveDefaultMaxTokens(t *testing.T) { tests := []struct { name string diff --git a/internal/providers/keyring.go b/internal/providers/keyring.go index 65aa6db9d..9bcc510a5 100644 --- a/internal/providers/keyring.go +++ b/internal/providers/keyring.go @@ -5,7 +5,6 @@ import ( "context" "crypto/hmac" "crypto/sha256" - "slices" "sync/atomic" "github.com/enterpilot/gomodel/internal/core" @@ -145,15 +144,6 @@ func rendezvousKeyScore(sessionID, key string) [sha256.Size]byte { return score } -// Any reports whether at least one configured key satisfies pred. It never -// advances the rotation. An empty or nil ring reports false. -func (k *Keyring) Any(pred func(string) bool) bool { - if k == nil || pred == nil { - return false - } - return slices.ContainsFunc(k.keys, pred) -} - // Primary returns the first configured key without advancing the rotation. // It is the key to use where a stable identity matters more than spreading // load, and where an empty ring must stay empty. diff --git a/internal/server/messages_handler_test.go b/internal/server/messages_handler_test.go index 7798e0e16..700f680f6 100644 --- a/internal/server/messages_handler_test.go +++ b/internal/server/messages_handler_test.go @@ -215,20 +215,23 @@ func TestRewriteMessagesModel(t *testing.T) { }) } - t.Run("alias rewrites model field", func(t *testing.T) { - got, err := rewriteMessagesModel([]byte(`{"model":"my-alias","max_tokens":1}`), "claude-test") + t.Run("alias rewrites only the model value", func(t *testing.T) { + // Whitespace, member order, and numeric spelling elsewhere must + // survive the rewrite untouched. + body := `{"max_tokens": 1, "model" : "my-alias" , "temperature": 1.50}` + want := `{"max_tokens": 1, "model" : "claude-test" , "temperature": 1.50}` + got, err := rewriteMessagesModel([]byte(body), "claude-test") if err != nil { t.Fatalf("rewriteMessagesModel: %v", err) } - var fields map[string]any - if err := json.Unmarshal(got, &fields); err != nil { - t.Fatalf("unmarshal rewritten body: %v", err) + if string(got) != want { + t.Errorf("body = %s, want %s", got, want) } - if fields["model"] != "claude-test" { - t.Errorf("model = %v, want claude-test", fields["model"]) - } - if fields["max_tokens"] != float64(1) { - t.Errorf("max_tokens = %v, want 1", fields["max_tokens"]) + }) + + t.Run("non-object body errors", func(t *testing.T) { + if _, err := rewriteMessagesModel([]byte(`[1,2]`), "claude-test"); err == nil { + t.Fatal("expected error for non-object body") } }) } diff --git a/internal/server/messages_native.go b/internal/server/messages_native.go index 9a222fb5c..171878dea 100644 --- a/internal/server/messages_native.go +++ b/internal/server/messages_native.go @@ -2,11 +2,14 @@ package server import ( "bytes" + // encoding/json rather than goccy: rewriteMessagesModel needs the + // decoder's InputOffset to splice the model value in place. + "encoding/json" + "errors" "io" "net/http" "strings" - "github.com/goccy/go-json" "github.com/labstack/echo/v5" "github.com/enterpilot/gomodel/internal/auditlog" @@ -100,28 +103,54 @@ func (s *translatedInferenceService) dispatchMessagesNative(c *echo.Context, req return proxyPassthroughResponse(c, s.logger, s.usageLogger, s.pricingResolver, anthropicProviderType, providerName, "messages", info, resp) } -// rewriteMessagesModel returns body with its "model" field replaced by the -// resolved model, leaving the body untouched when it already matches so -// aliased/renamed models reach the provider under their real name. +// rewriteMessagesModel returns body with its top-level "model" value replaced +// by the resolved model so aliased/renamed models reach the provider under +// their real name. Only the model value's bytes are spliced; every other byte +// of the request is preserved. The body is returned unchanged when the model +// already matches or has no model field (upstream validation rejects that). func rewriteMessagesModel(body []byte, model string) ([]byte, error) { if strings.TrimSpace(model) == "" { return body, nil } - var fields map[string]json.RawMessage - if err := json.Unmarshal(body, &fields); err != nil { - return nil, err - } - var current string - if raw, ok := fields["model"]; ok { - _ = json.Unmarshal(raw, ¤t) - } - if current == model { - return body, nil - } - encoded, err := json.Marshal(model) + dec := json.NewDecoder(bytes.NewReader(body)) + tok, err := dec.Token() if err != nil { return nil, err } - fields["model"] = encoded - return json.Marshal(fields) + if delim, ok := tok.(json.Delim); !ok || delim != '{' { + return nil, errors.New("request body is not a JSON object") + } + for dec.More() { + keyTok, err := dec.Token() + if err != nil { + return nil, err + } + key, _ := keyTok.(string) + var raw json.RawMessage + if err := dec.Decode(&raw); err != nil { + return nil, err + } + if key != "model" { + continue + } + var current string + _ = json.Unmarshal(raw, ¤t) + if current == model { + return body, nil + } + encoded, err := json.Marshal(model) + if err != nil { + return nil, err + } + // The model value is a scalar, so raw holds its exact source bytes + // and InputOffset points just past them. + end := dec.InputOffset() + start := end - int64(len(raw)) + rewritten := make([]byte, 0, int64(len(body))-int64(len(raw))+int64(len(encoded))) + rewritten = append(rewritten, body[:start]...) + rewritten = append(rewritten, encoded...) + rewritten = append(rewritten, body[end:]...) + return rewritten, nil + } + return body, nil } From 0d06bfc49748f76b00e36eda4a2186132c4b1906 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Tue, 18 Aug 2026 17:19:04 +0200 Subject: [PATCH 3/5] fix(anthropic): rewrite the last duplicate model member in native messages requests --- internal/server/messages_handler_test.go | 14 +++++++ internal/server/messages_native.go | 48 ++++++++++++++---------- 2 files changed, 42 insertions(+), 20 deletions(-) diff --git a/internal/server/messages_handler_test.go b/internal/server/messages_handler_test.go index 700f680f6..dc84e0fa2 100644 --- a/internal/server/messages_handler_test.go +++ b/internal/server/messages_handler_test.go @@ -229,6 +229,20 @@ func TestRewriteMessagesModel(t *testing.T) { } }) + t.Run("duplicate model members rewrite the last one", func(t *testing.T) { + // Decoders keep the last duplicate member, so the rewrite must target + // it — rewriting the first would leave the effective model unchanged. + body := `{"model":"ignored","max_tokens":1,"model":"my-alias"}` + want := `{"model":"ignored","max_tokens":1,"model":"claude-test"}` + got, err := rewriteMessagesModel([]byte(body), "claude-test") + if err != nil { + t.Fatalf("rewriteMessagesModel: %v", err) + } + if string(got) != want { + t.Errorf("body = %s, want %s", got, want) + } + }) + t.Run("non-object body errors", func(t *testing.T) { if _, err := rewriteMessagesModel([]byte(`[1,2]`), "claude-test"); err == nil { t.Fatal("expected error for non-object body") diff --git a/internal/server/messages_native.go b/internal/server/messages_native.go index 171878dea..ecb00dbad 100644 --- a/internal/server/messages_native.go +++ b/internal/server/messages_native.go @@ -120,6 +120,11 @@ func rewriteMessagesModel(body []byte, model string) ([]byte, error) { if delim, ok := tok.(json.Delim); !ok || delim != '{' { return nil, errors.New("request body is not a JSON object") } + // Walk every top-level member and remember the span of the last "model" + // value: decoders keep the last duplicate member, so that is the one the + // resolved model came from and the one to rewrite. + var modelRaw json.RawMessage + var modelEnd int64 for dec.More() { keyTok, err := dec.Token() if err != nil { @@ -133,24 +138,27 @@ func rewriteMessagesModel(body []byte, model string) ([]byte, error) { if key != "model" { continue } - var current string - _ = json.Unmarshal(raw, ¤t) - if current == model { - return body, nil - } - encoded, err := json.Marshal(model) - if err != nil { - return nil, err - } - // The model value is a scalar, so raw holds its exact source bytes - // and InputOffset points just past them. - end := dec.InputOffset() - start := end - int64(len(raw)) - rewritten := make([]byte, 0, int64(len(body))-int64(len(raw))+int64(len(encoded))) - rewritten = append(rewritten, body[:start]...) - rewritten = append(rewritten, encoded...) - rewritten = append(rewritten, body[end:]...) - return rewritten, nil - } - return body, nil + modelRaw = raw + modelEnd = dec.InputOffset() + } + if modelRaw == nil { + return body, nil + } + var current string + _ = json.Unmarshal(modelRaw, ¤t) + if current == model { + return body, nil + } + encoded, err := json.Marshal(model) + if err != nil { + return nil, err + } + // The model value is a scalar, so modelRaw holds its exact source bytes + // and modelEnd points just past them. + start := modelEnd - int64(len(modelRaw)) + rewritten := make([]byte, 0, int64(len(body))-int64(len(modelRaw))+int64(len(encoded))) + rewritten = append(rewritten, body[:start]...) + rewritten = append(rewritten, encoded...) + rewritten = append(rewritten, body[modelEnd:]...) + return rewritten, nil } From 21f622c357484660962645ee784df22fcf672dce Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Tue, 18 Aug 2026 17:23:45 +0200 Subject: [PATCH 4/5] feat(server): deliver response feedback to ext rewriters on native messages forwarding Extensions that request response feedback (e.g. compression epochs) now observe the native /v1/messages SSE stream; the feedback observer merges Anthropic split usage events (message_start input/cache tokens with the final message_delta) instead of keeping only the last event. --- internal/server/messages_native.go | 20 ++++++- internal/server/messages_native_test.go | 78 +++++++++++++++++++++++++ internal/server/passthrough_support.go | 8 ++- internal/server/response_feedback.go | 7 +++ 4 files changed, 109 insertions(+), 4 deletions(-) diff --git a/internal/server/messages_native.go b/internal/server/messages_native.go index ecb00dbad..16eb4933c 100644 --- a/internal/server/messages_native.go +++ b/internal/server/messages_native.go @@ -12,8 +12,10 @@ import ( "github.com/labstack/echo/v5" + "github.com/enterpilot/gomodel/ext" "github.com/enterpilot/gomodel/internal/auditlog" "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/streaming" ) const anthropicProviderType = "anthropic" @@ -100,7 +102,23 @@ func (s *translatedInferenceService) dispatchMessagesNative(c *echo.Context, req AuditPath: "/v1/messages", Model: req.Model, } - return proxyPassthroughResponse(c, s.logger, s.usageLogger, s.pricingResolver, anthropicProviderType, providerName, "messages", info, resp) + // Request rewriters (ext extensions) that asked for response feedback + // observe the native SSE stream too; its Anthropic-native usage events + // are understood by the feedback observer. + var extraObservers []streaming.Observer + if hasResponseFeedbackObservers(c) { + extraObservers = append(extraObservers, &responseFeedbackStreamObserver{ + ctx: c.Request().Context(), + observers: responseFeedbackObservers(c), + requestID: requestIDFromContextOrHeader(c.Request()), + sessionID: core.SessionIDFromContext(c.Request().Context()), + endpoint: ext.Endpoint("/v1/messages"), + model: req.Model, + providerType: anthropicProviderType, + providerName: providerName, + }) + } + return proxyPassthroughResponse(c, s.logger, s.usageLogger, s.pricingResolver, anthropicProviderType, providerName, "messages", info, resp, extraObservers...) } // rewriteMessagesModel returns body with its top-level "model" value replaced diff --git a/internal/server/messages_native_test.go b/internal/server/messages_native_test.go index da4c7eecc..1838be0a8 100644 --- a/internal/server/messages_native_test.go +++ b/internal/server/messages_native_test.go @@ -1,6 +1,7 @@ package server import ( + "context" "io" "net/http" "net/http/httptest" @@ -9,6 +10,7 @@ import ( "github.com/labstack/echo/v5" + "github.com/enterpilot/gomodel/ext" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/usage" ) @@ -105,3 +107,79 @@ func TestBuildPassthroughHeadersDropsAcceptEncoding(t *testing.T) { t.Errorf("Anthropic-Beta = %q, want preserved", got) } } + +type recordingFeedbackObserver struct { + input, read, write int + observed bool + calls int +} + +func (r *recordingFeedbackObserver) ObserveResponse(_ context.Context, _ string, _ ext.Endpoint, _ string, _ string, _ string, _ string, inputTokens, cachedInputTokens, cacheWriteInputTokens int, usageObserved bool) { + r.calls++ + r.input = inputTokens + r.read = cachedInputTokens + r.write = cacheWriteInputTokens + r.observed = usageObserved +} + +// Extensions that requested response feedback must receive usage from the +// native SSE stream, with input and cache tokens (message_start) merged with +// the final message_delta. +func TestMessages_NativeStreamingNotifiesFeedbackObservers(t *testing.T) { + anthropicSSE := strings.Join([]string{ + `event: message_start`, + `data: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-fable-5","usage":{"input_tokens":19560,"cache_creation_input_tokens":100,"cache_read_input_tokens":200,"output_tokens":3}}}`, + ``, + `event: message_delta`, + `data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":31}}`, + ``, + `event: message_stop`, + `data: {"type":"message_stop"}`, + ``, + ``, + }, "\n") + + provider := &mockProvider{ + supportedModels: []string{"claude-fable-5"}, + providerTypes: map[string]string{"claude-fable-5": "anthropic"}, + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: map[string][]string{"Content-Type": {"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(anthropicSSE)), + }, + } + + e := echo.New() + handler := NewHandler(provider, nil, nil, nil) + + reqBody := `{"model":"claude-fable-5","max_tokens":64,"stream":true,"messages":[{"role":"user","content":"Hi"}]}` + req := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(reqBody)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + observer := &recordingFeedbackObserver{} + setResponseFeedbackObservers(c, []ext.ResponseFeedbackObserver{observer}) + + if err := handler.Messages(c); err != nil { + t.Fatalf("Messages: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, body = %s", rec.Code, rec.Body.String()) + } + if observer.calls != 1 { + t.Fatalf("ObserveResponse calls = %d, want 1", observer.calls) + } + if !observer.observed { + t.Error("usageObserved = false, want true") + } + if observer.input != 19560 { + t.Errorf("inputTokens = %d, want 19560", observer.input) + } + if observer.read != 200 { + t.Errorf("cachedInputTokens = %d, want 200", observer.read) + } + if observer.write != 100 { + t.Errorf("cacheWriteInputTokens = %d, want 100", observer.write) + } +} diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index 78b1f44a1..d7258b2f3 100644 --- a/internal/server/passthrough_support.go +++ b/internal/server/passthrough_support.go @@ -255,9 +255,10 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT } // proxyPassthroughResponse relays a provider-native response (JSON or SSE) to -// the client, attaching audit and usage stream observers. It is shared by the +// the client, attaching audit and usage stream observers plus any +// extraObservers the caller supplies for SSE responses. It is shared by the // /p/ passthrough surface and the /v1/messages native forwarding path. -func proxyPassthroughResponse(c *echo.Context, logger auditlog.LoggerInterface, usageLogger usage.LoggerInterface, pricingResolver usage.PricingResolver, providerType, providerName, endpoint string, info *core.PassthroughRouteInfo, resp *core.PassthroughResponse) error { +func proxyPassthroughResponse(c *echo.Context, logger auditlog.LoggerInterface, usageLogger usage.LoggerInterface, pricingResolver usage.PricingResolver, providerType, providerName, endpoint string, info *core.PassthroughRouteInfo, resp *core.PassthroughResponse, extraObservers ...streaming.Observer) error { if resp == nil || resp.Body == nil { return handleError(c, core.NewProviderError(providerType, http.StatusBadGateway, "provider returned empty passthrough response", nil)) } @@ -310,7 +311,7 @@ func proxyPassthroughResponse(c *echo.Context, logger auditlog.LoggerInterface, } model = resolvedModelFromWorkflow(workflow, model) - observers := make([]streaming.Observer, 0, 2) + observers := make([]streaming.Observer, 0, 2+len(extraObservers)) if auditEnabled && streamEntry != nil { if observer := auditlog.NewStreamLogObserver(logger, streamEntry, auditPath); observer != nil { observers = append(observers, observer) @@ -325,6 +326,7 @@ func proxyPassthroughResponse(c *echo.Context, logger auditlog.LoggerInterface, observers = append(observers, observer) } } + observers = append(observers, extraObservers...) wrappedStream := streaming.NewObservedSSEStream(resp.Body, observers...) if len(observers) > 0 { defer func() { diff --git a/internal/server/response_feedback.go b/internal/server/response_feedback.go index 5ba5865a7..9f7b6a7b4 100644 --- a/internal/server/response_feedback.go +++ b/internal/server/response_feedback.go @@ -139,6 +139,13 @@ func (o *responseFeedbackStreamObserver) OnJSONEvent(payload map[string]any) { if details, ok := nestedMap(usageMap["input_tokens_details"]); ok { usage.read = max(usage.read, firstNumericInt(details, "cached_tokens")) } + // Merge rather than replace: Anthropic-native streams split usage across + // events (input and cache tokens in message_start, output in the final + // message_delta), while cumulative dialects report full usage in their + // last event — max keeps both correct. + usage.input = max(usage.input, o.usage.input) + usage.read = max(usage.read, o.usage.read) + usage.write = max(usage.write, o.usage.write) o.usage = usage } From 5b1f458d8e571748d03eec690acf6e7e96c927e3 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Tue, 18 Aug 2026 18:00:22 +0200 Subject: [PATCH 5/5] refactor(server): unify native forwarding dispatch Extract the shared gate plumbing and Passthrough-then-proxy block used by the streaming chat fast path and the /v1/messages native path into nativeForwardingProvider and forwardNative. Pure consolidation; both paths keep their existing eligibility conditions and behavior. --- internal/server/messages_native.go | 20 ++----- internal/server/native_dispatch.go | 53 +++++++++++++++++++ .../server/translated_inference_service.go | 22 +++----- 3 files changed, 62 insertions(+), 33 deletions(-) create mode 100644 internal/server/native_dispatch.go diff --git a/internal/server/messages_native.go b/internal/server/messages_native.go index 16eb4933c..379b68ec4 100644 --- a/internal/server/messages_native.go +++ b/internal/server/messages_native.go @@ -32,16 +32,10 @@ func (s *translatedInferenceService) canForwardMessagesNatively(workflow *core.W if workflow == nil || strings.TrimSpace(workflow.ProviderType) != anthropicProviderType { return false } - if s.translatedRequestPatcher != nil { - return false - } if s.responseCache != nil && workflow.CacheEnabled() { return false } - if len(s.inference().FailoverSelectors(workflow)) > 0 { - return false - } - _, ok := s.provider.(core.RoutablePassthrough) + _, ok := s.nativeForwardingProvider(workflow) return ok } @@ -49,11 +43,6 @@ func (s *translatedInferenceService) canForwardMessagesNatively(workflow *core.W // resolved Anthropic provider and relays the provider-native response (JSON or // SSE) unchanged, with admission, audit, and streaming usage accounting. func (s *translatedInferenceService) dispatchMessagesNative(c *echo.Context, req *core.ChatRequest, workflow *core.Workflow) error { - passthroughProvider, ok := s.provider.(core.RoutablePassthrough) - if !ok { - return handleError(c, core.NewInvalidRequestError("provider passthrough is not supported by the current provider router", nil)) - } - body, err := requestBodyBytes(c) if err != nil { return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) @@ -77,7 +66,7 @@ func (s *translatedInferenceService) dispatchMessagesNative(c *echo.Context, req providerName = workflow.Resolution.ProviderName } - resp, err := passthroughProvider.Passthrough(ctx, anthropicProviderType, &core.PassthroughRequest{ + passthroughReq := &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "messages", Operation: "anthropic.messages", @@ -86,9 +75,6 @@ func (s *translatedInferenceService) dispatchMessagesNative(c *echo.Context, req Body: io.NopCloser(bytes.NewReader(body)), Headers: buildPassthroughHeaders(ctx, c.Request().Header), ProviderName: providerName, - }) - if err != nil { - return handleError(c, err) } auditlog.EnrichEntryWithWorkflow(c, workflow) @@ -118,7 +104,7 @@ func (s *translatedInferenceService) dispatchMessagesNative(c *echo.Context, req providerName: providerName, }) } - return proxyPassthroughResponse(c, s.logger, s.usageLogger, s.pricingResolver, anthropicProviderType, providerName, "messages", info, resp, extraObservers...) + return s.forwardNative(c, ctx, passthroughReq, info, extraObservers...) } // rewriteMessagesModel returns body with its top-level "model" value replaced diff --git a/internal/server/native_dispatch.go b/internal/server/native_dispatch.go new file mode 100644 index 000000000..a1707e157 --- /dev/null +++ b/internal/server/native_dispatch.go @@ -0,0 +1,53 @@ +package server + +import ( + "context" + + "github.com/labstack/echo/v5" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/streaming" +) + +// nativeForwardingProvider returns the provider router's passthrough surface +// when native dialect forwarding is possible for the workflow at all: no +// translated-request patcher is installed and no failover targets are +// configured (both operate on the canonical translated request). Dialect +// gates layer their endpoint-specific conditions on top of this. +func (s *translatedInferenceService) nativeForwardingProvider(workflow *core.Workflow) (core.RoutablePassthrough, bool) { + if workflow == nil { + return nil, false + } + if s.translatedRequestPatcher != nil { + return nil, false + } + if len(s.inference().FailoverSelectors(workflow)) > 0 { + return nil, false + } + passthroughProvider, ok := s.provider.(core.RoutablePassthrough) + return passthroughProvider, ok +} + +// forwardNative executes a prepared provider-native request through the +// passthrough surface and relays the provider response (JSON or SSE) +// unchanged, with audit and usage stream observers plus any extraObservers. +// req.ProviderName and info.Provider must already be set; info.Provider names +// the resolved provider type for routing and error attribution. +func (s *translatedInferenceService) forwardNative( + c *echo.Context, + ctx context.Context, + req *core.PassthroughRequest, + info *core.PassthroughRouteInfo, + extraObservers ...streaming.Observer, +) error { + passthroughProvider, ok := s.provider.(core.RoutablePassthrough) + if !ok { + return handleError(c, core.NewInvalidRequestError("provider passthrough is not supported by the current provider router", nil)) + } + + resp, err := passthroughProvider.Passthrough(ctx, info.Provider, req) + if err != nil { + return handleError(c, err) + } + return proxyPassthroughResponse(c, s.logger, s.usageLogger, s.pricingResolver, info.Provider, req.ProviderName, req.Endpoint, info, resp, extraObservers...) +} diff --git a/internal/server/translated_inference_service.go b/internal/server/translated_inference_service.go index 9331d24bb..2791c8e9b 100644 --- a/internal/server/translated_inference_service.go +++ b/internal/server/translated_inference_service.go @@ -443,8 +443,7 @@ func (s *translatedInferenceService) tryFastPathStreamingChatPassthrough(c *echo return false, nil } - passthroughProvider, ok := s.provider.(core.RoutablePassthrough) - if !ok { + if _, ok := s.nativeForwardingProvider(workflow); !ok { return false, nil } @@ -453,33 +452,24 @@ func (s *translatedInferenceService) tryFastPathStreamingChatPassthrough(c *echo const endpoint = "/chat/completions" providerType := strings.TrimSpace(workflow.ProviderType) - resp, err := passthroughProvider.Passthrough(ctx, providerType, &core.PassthroughRequest{ + model := resolvedModelFromWorkflow(workflow, req.Model) + passthroughReq := &core.PassthroughRequest{ Method: c.Request().Method, Endpoint: endpoint, Operation: llmclient.OperationChat, - Model: resolvedModelFromWorkflow(workflow, req.Model), + Model: model, Stream: req.Stream, Body: c.Request().Body, Headers: buildPassthroughHeaders(ctx, c.Request().Header), ProviderName: providerNameFromWorkflow(workflow), - }) - if err != nil { - return true, handleError(c, err) } - info := &core.PassthroughRouteInfo{ Provider: providerType, RawEndpoint: strings.TrimPrefix(endpoint, "/"), AuditPath: c.Request().URL.Path, - Model: resolvedModelFromWorkflow(workflow, req.Model), - } - passthrough := passthroughService{ - provider: s.provider, - logger: s.logger, - usageLogger: s.usageLogger, - pricingResolver: s.pricingResolver, + Model: model, } - return true, passthrough.proxyPassthroughResponse(c, providerType, providerNameFromWorkflow(workflow), endpoint, info, resp) + return true, s.forwardNative(c, ctx, passthroughReq, info) } func (s *translatedInferenceService) Embeddings(c *echo.Context) error {