diff --git a/internal/runtime/executor/claude_executor_test.go b/internal/runtime/executor/claude_executor_test.go index 1148429f2..d8c4448d1 100644 --- a/internal/runtime/executor/claude_executor_test.go +++ b/internal/runtime/executor/claude_executor_test.go @@ -2004,120 +2004,49 @@ func TestClaudeCountTokensBetasForCredentialMatchesNativeOAuth220(t *testing.T) } } -func TestShouldFinalizeClaudeCountTokensCCHSkipsDirectAnthropic(t *testing.T) { - if shouldFinalizeClaudeCountTokensCCH(true, true) { - t.Fatal("direct Anthropic count_tokens must not receive CPA CCH") +func TestShouldUseClaudeUpstreamTokenCount(t *testing.T) { + tests := []struct { + name string + apiKey string + baseURL string + want bool + }{ + {name: "official OAuth", apiKey: "sk-ant-oat-official", baseURL: "https://api.anthropic.com", want: true}, + {name: "official API key", apiKey: "key-official", baseURL: "https://api.anthropic.com:443", want: true}, + {name: "custom OAuth", apiKey: "sk-ant-oat-custom", baseURL: "https://gateway.example"}, + {name: "custom API key", apiKey: "key-custom", baseURL: "https://gateway.example"}, + {name: "lookalike host", apiKey: "sk-ant-oat-lookalike", baseURL: "https://api.anthropic.com.example"}, + {name: "insecure official host", apiKey: "sk-ant-oat-http", baseURL: "http://api.anthropic.com"}, + {name: "missing credential", baseURL: "https://api.anthropic.com"}, } - if !shouldFinalizeClaudeCountTokensCCH(true, false) { - t.Fatal("custom-gateway count_tokens should retain existing CCH behavior") + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + if got := shouldUseClaudeUpstreamTokenCount(test.apiKey, test.baseURL); got != test.want { + t.Fatalf("shouldUseClaudeUpstreamTokenCount() = %v, want %v", got, test.want) + } + }) } } -func TestClaudeExecutor_CountTokensOAuthUsesUpstreamCLIShape(t *testing.T) { - var upstreamAlias string - var upstreamBody []byte - var upstreamHeaders http.Header - var upstreamPath string - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - body, _ := io.ReadAll(r.Body) - upstreamBody = bytes.Clone(body) - upstreamHeaders = r.Header.Clone() - upstreamPath = r.URL.RequestURI() - upstreamAlias = gjson.GetBytes(body, "tools.0.name").String() - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"input_tokens":7}`)) - })) - defer server.Close() - - deviceIDs := []string{ - "0000000000000000000000000000000000000000000000000000000000000000", - } - executor := NewClaudeExecutor(&config.Config{}) - auth := &cliproxyauth.Auth{ - ID: "oauth-mcp-count-tokens", - Attributes: map[string]string{ - "api_key": "sk-ant-oat-mcp-count-tokens", - "base_url": server.URL, - }, - Metadata: map[string]any{ - "account_uuid": "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa", - claudeauth.ClaudeDeviceIDsMetadataKey: deviceIDs, - }, - } - payload := []byte(`{"model":"claude-opus-5","system":"count-system-prompt","messages":[{"role":"user","content":"search"}],"tools":[{"name":"search_web","input_schema":{"type":"object"}}]}`) - resp, errCount := executor.CountTokens(context.Background(), auth, cliproxyexecutor.Request{ - Model: "claude-opus-5", - Payload: payload, - }, cliproxyexecutor.Options{ - SourceFormat: sdktranslator.FormatClaude, - Metadata: map[string]any{ - cliproxyexecutor.ExecutionSessionMetadataKey: "count-agent-conversation", - }, - }) - if errCount != nil { - t.Fatalf("CountTokens() error = %v", errCount) - } - if upstreamPath != "/v1/messages/count_tokens?beta=true" { - t.Fatalf("upstream count_tokens path = %q, want beta endpoint", upstreamPath) - } - if !helps.IsClaudeMCPToolName(upstreamAlias) { - t.Fatalf("upstream count_tokens tool name = %q, want mcp__ alias", upstreamAlias) - } - if got := upstreamHeaders.Get("User-Agent"); got != "claude-cli/2.1.220 (external, cli)" { - t.Fatalf("count_tokens User-Agent = %q, want CLI identity", got) - } - // count_tokens carries its own much smaller profile, not the inference baseline. - wantBetas := claudeCountTokensBetasForCredential(true) - if got := upstreamHeaders.Get("Anthropic-Beta"); got != wantBetas { - t.Fatalf("count_tokens Anthropic-Beta = %q, want %q", got, wantBetas) - } - // Claude Code omits X-Stainless-Timeout on count_tokens. - if got := upstreamHeaders.Get("X-Stainless-Timeout"); got != "" { - t.Fatalf("count_tokens X-Stainless-Timeout = %q, want it absent", got) - } - if got := gjson.GetBytes(upstreamBody, "system.1.text").String(); got != claudeCodeCLIIdentity { - t.Fatalf("count_tokens system.1.text = %q, want official CLI identity", got) - } - if got := gjson.GetBytes(upstreamBody, "system.0.text").String(); !strings.Contains(got, "cc_entrypoint=cli;") { - t.Fatalf("count_tokens billing attribution = %q, want cli", got) - } - if got := gjson.GetBytes(upstreamBody, "system.#").Int(); got != 2 { - t.Fatalf("count_tokens system block count = %d, want 2", got) - } - content := gjson.GetBytes(upstreamBody, "messages.0.content").Array() - if len(content) != 2 { - t.Fatalf("count_tokens first user content has %d blocks, want currentDate and user text", len(content)) - } - assertClaudeCodeCurrentDateBlock(t, content[0]) - assertEphemeralUserTextBlock(t, content[1], "search") - assertClaudeMidConversationSystemMessage(t, upstreamBody, 1, "count-system-prompt") - if _, ok := claudeBillingCCHDigitsOffset(upstreamBody); !ok { - t.Fatalf("count_tokens Claude OAuth custom BaseURL body is missing CCH: %s", upstreamBody) - } - assertClaudeCountTokensIdentity(t, upstreamBody, upstreamHeaders) - if got := gjson.GetBytes(resp.Payload, "input_tokens").Int(); got != 7 { - t.Fatalf("input_tokens = %d, want 7", got) - } -} - -func TestClaudeExecutor_LegacySystemReminderAcrossMessagesStreamAndCountTokens(t *testing.T) { +func TestClaudeExecutor_LegacySystemReminderAcrossMessagesAndStream(t *testing.T) { var mu sync.Mutex captured := make(map[string][]byte) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { body, _ := io.ReadAll(r.Body) - kind := "messages" if strings.Contains(r.URL.Path, "count_tokens") { - kind = "count_tokens" - } else if gjson.GetBytes(body, "stream").Bool() { + t.Errorf("custom OAuth count_tokens unexpectedly reached upstream: %s", r.URL.Path) + w.WriteHeader(http.StatusInternalServerError) + return + } + kind := "messages" + if gjson.GetBytes(body, "stream").Bool() { kind = "stream" } mu.Lock() captured[kind] = bytes.Clone(body) mu.Unlock() switch kind { - case "count_tokens": - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"input_tokens":7}`)) case "stream": w.Header().Set("Content-Type", "text/event-stream") _, _ = w.Write([]byte("event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n")) @@ -2155,11 +2084,15 @@ func TestClaudeExecutor_LegacySystemReminderAcrossMessagesStreamAndCountTokens(t }, cliproxyexecutor.Options{SourceFormat: sdktranslator.FormatClaude}); errExecute != nil { t.Fatalf("Execute() error = %v", errExecute) } - if _, errCount := executor.CountTokens(context.Background(), auth, cliproxyexecutor.Request{ + countResp, errCount := executor.CountTokens(context.Background(), auth, cliproxyexecutor.Request{ Model: "claude-opus-4-6", Payload: makePayload("count-user", false), - }, cliproxyexecutor.Options{SourceFormat: sdktranslator.FormatClaude}); errCount != nil { + }, cliproxyexecutor.Options{SourceFormat: sdktranslator.FormatClaude}) + if errCount != nil { t.Fatalf("CountTokens() error = %v", errCount) } + if got := gjson.GetBytes(countResp.Payload, "input_tokens").Int(); got <= 0 { + t.Fatalf("local count_tokens input_tokens = %d, want positive estimate", got) + } streamResult, errStream := executor.ExecuteStream(context.Background(), auth, cliproxyexecutor.Request{ Model: "claude-opus-4-6", Payload: makePayload("stream-user", true), }, cliproxyexecutor.Options{SourceFormat: sdktranslator.FormatClaude}) @@ -2174,12 +2107,11 @@ func TestClaudeExecutor_LegacySystemReminderAcrossMessagesStreamAndCountTokens(t mu.Lock() bodies := map[string][]byte{ - "messages": bytes.Clone(captured["messages"]), - "count_tokens": bytes.Clone(captured["count_tokens"]), - "stream": bytes.Clone(captured["stream"]), + "messages": bytes.Clone(captured["messages"]), + "stream": bytes.Clone(captured["stream"]), } mu.Unlock() - for kind, wantUser := range map[string]string{"messages": "messages-user", "count_tokens": "count-user", "stream": "stream-user"} { + for kind, wantUser := range map[string]string{"messages": "messages-user", "stream": "stream-user"} { body := bodies[kind] if len(body) == 0 { t.Fatalf("missing %s upstream capture", kind) @@ -2495,6 +2427,7 @@ func TestClaudeExecutor_CountTokensCountsLocallyWithoutUpstreamRequest(t *testin apiKey string }{ {name: "custom API key", apiKey: "key-123"}, + {name: "custom OAuth", apiKey: "sk-ant-oat-custom"}, } for _, testCase := range testCases { diff --git a/internal/runtime/executor/claude_executor_tokens.go b/internal/runtime/executor/claude_executor_tokens.go index bf0a1b189..1b962fa57 100644 --- a/internal/runtime/executor/claude_executor_tokens.go +++ b/internal/runtime/executor/claude_executor_tokens.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "net/http" + "strings" "github.com/router-for-me/CLIProxyAPI/v7/internal/runtime/executor/helps" "github.com/router-for-me/CLIProxyAPI/v7/internal/thinking" @@ -22,11 +23,10 @@ func (e *ClaudeExecutor) CountTokens(ctx context.Context, auth *cliproxyauth.Aut if baseURL == "" { baseURL = "https://api.anthropic.com" } - endpoint := fmt.Sprintf("%s/v1/messages/count_tokens?beta=true", baseURL) - // Claude Code uses the native endpoint for OAuth and first-party Anthropic - // credentials. Keep local counting for custom API-key gateways that may not - // implement count_tokens. - if apiKey != "" && claudeCCHSigningEnabled(apiKey, claudeCCHUpstreamAnthropic, endpoint) { + // Only Anthropic's first-party origin has the measured native count_tokens + // contract. Every custom/third-party base URL keeps local estimation, + // regardless of whether the credential is OAuth or an API key. + if shouldUseClaudeUpstreamTokenCount(apiKey, baseURL) { return e.countTokensUpstream(ctx, auth, req, opts) } @@ -111,12 +111,11 @@ func validateClaudeTokenCountRequest(body []byte) error { return nil } -// countTokensUpstream preserves native token counting for Claude-compatible -// providers that expose their own count_tokens endpoint. -func shouldFinalizeClaudeCountTokensCCH(cchSigning, directAnthropic bool) bool { - return cchSigning && !directAnthropic +func shouldUseClaudeUpstreamTokenCount(apiKey, baseURL string) bool { + return strings.TrimSpace(apiKey) != "" && isAnthropicUpstreamBase(baseURL) } +// countTokensUpstream preserves Anthropic's native token-counting contract. func (e *ClaudeExecutor) countTokensUpstream(ctx context.Context, auth *cliproxyauth.Auth, req cliproxyexecutor.Request, opts cliproxyexecutor.Options) (cliproxyexecutor.Response, error) { baseModel := thinking.ParseSuffix(req.Model).ModelName upstreamModel := e.upstreamModel(baseModel) @@ -127,7 +126,6 @@ func (e *ClaudeExecutor) countTokensUpstream(ctx context.Context, auth *cliproxy } url := fmt.Sprintf("%s/v1/messages/count_tokens?beta=true", baseURL) oauthToken := isClaudeOAuthToken(apiKey) - cchSigning := claudeCCHSigningEnabled(apiKey, claudeCCHUpstreamAnthropic, url) from := opts.SourceFormat responseFormat := cliproxyexecutor.ResponseFormatOrSource(opts) @@ -179,7 +177,7 @@ func (e *ClaudeExecutor) countTokensUpstream(ctx context.Context, auth *cliproxy body, apiKey, confirmedClaudeCode, - cchSigning, + false, ) if errCloaking != nil { return cliproxyexecutor.Response{}, errCloaking @@ -208,15 +206,6 @@ func (e *ClaudeExecutor) countTokensUpstream(ctx context.Context, auth *cliproxy body, _ = sjson.DeleteBytes(body, "context_management") body, _ = sjson.DeleteBytes(body, "diagnostics") } - if shouldFinalizeClaudeCountTokensCCH(cchSigning, directAnthropic) { - fallbackBilling := claudeCCHFallbackBillingHeader(ctx, e.cfg, body, claudeCodeDetection.Entrypoint) - var errCCH error - body, errCCH = finalizeAnthropicMessagesBodyCCH(body, fallbackBilling) - if errCCH != nil { - return cliproxyexecutor.Response{}, fmt.Errorf("finalize Claude CCH: %w", errCCH) - } - } - httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) if err != nil { return cliproxyexecutor.Response{}, err