diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go index eba08a113b..2fb7c4baa3 100644 --- a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go +++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go @@ -769,6 +769,32 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_BuildRequestRejectsInvalidBas require.Error(t, err) } +func TestGatewayService_AnthropicAPIKeyPassthrough_StripsDeferredToolCacheControl(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + svc := &GatewayService{cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}}} + account := &Account{Platform: PlatformAnthropic, Type: AccountTypeAPIKey} + body := []byte(`{"tools":[{"name":"deferred","custom":{"defer_loading":true},"cache_control":{"type":"ephemeral"}},{"name":"top-level-deferred","defer_loading":true,"cache_control":{"type":"ephemeral"}},{"name":"ordinary","defer_loading":false,"cache_control":{"type":"ephemeral"}},{"name":"malformed","defer_loading":"true","cache_control":{"type":"ephemeral"}}]}`) + + _, wireBody, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(context.Background(), c, account, body, "k") + require.NoError(t, err) + require.False(t, gjson.GetBytes(wireBody, "tools.0.cache_control").Exists()) + require.False(t, gjson.GetBytes(wireBody, "tools.1.cache_control").Exists()) + require.True(t, gjson.GetBytes(wireBody, "tools.2.cache_control").Exists()) + require.True(t, gjson.GetBytes(wireBody, "tools.3.cache_control").Exists()) + + countReq, err := svc.buildCountTokensRequestAnthropicAPIKeyPassthrough(context.Background(), c, account, body, "k") + require.NoError(t, err) + countBody, err := io.ReadAll(countReq.Body) + require.NoError(t, err) + require.False(t, gjson.GetBytes(countBody, "tools.0.cache_control").Exists()) + require.False(t, gjson.GetBytes(countBody, "tools.1.cache_control").Exists()) + require.True(t, gjson.GetBytes(countBody, "tools.2.cache_control").Exists()) + require.True(t, gjson.GetBytes(countBody, "tools.3.cache_control").Exists()) +} + func TestGatewayService_AnthropicOAuth_NotAffectedByAPIKeyPassthroughToggle(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() diff --git a/backend/internal/service/gateway_anthropic_passthrough.go b/backend/internal/service/gateway_anthropic_passthrough.go index def4d34648..533e780017 100644 --- a/backend/internal/service/gateway_anthropic_passthrough.go +++ b/backend/internal/service/gateway_anthropic_passthrough.go @@ -310,6 +310,7 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough( body []byte, token string, ) (*http.Request, []byte, error) { + body = stripDeferredToolCacheControl(body) targetURL := claudeAPIURL baseURL := account.GetBaseURL() if baseURL != "" { diff --git a/backend/internal/service/gateway_context_management_test.go b/backend/internal/service/gateway_context_management_test.go index bdeb2c600b..b7ab3e9429 100644 --- a/backend/internal/service/gateway_context_management_test.go +++ b/backend/internal/service/gateway_context_management_test.go @@ -4,6 +4,7 @@ package service import ( "context" + "fmt" "io" "net/http" "net/http/httptest" @@ -650,6 +651,51 @@ func TestBuildCountTokensRequest_APIKeyHaiku_StripsContextManagementEndToEnd(t * "count_tokens API-key + 客户端未带 beta token → body strip") } +func TestBuildCountTokensRequest_StripsCacheControlOnlyFromLiteralDeferredTools(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"claude-haiku-4-5","messages":[],"tools":[{"name":"deferred","custom":{"defer_loading":true},"cache_control":{"type":"ephemeral"}},{"name":"ordinary","custom":{"defer_loading":false},"cache_control":{"type":"ephemeral"}},{"name":"string","custom":{"defer_loading":"true"},"cache_control":{"type":"ephemeral"}},{"name":"number","custom":{"defer_loading":1},"cache_control":{"type":"ephemeral"}},{"name":"object","custom":{"defer_loading":{}},"cache_control":{"type":"ephemeral"}}]}`) + + tests := []struct { + name string + account *Account + token string + tokenType string + }{ + { + name: "generic API key", + account: &Account{Platform: PlatformAnthropic, Type: AccountTypeAPIKey}, + token: "sk-ant-test", + tokenType: "apikey", + }, + { + name: "recognized Claude Code OAuth without mimicry", + account: &Account{Platform: PlatformAnthropic, Type: AccountTypeOAuth}, + token: "oauth-token", + tokenType: "oauth", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil) + svc := &GatewayService{cfg: &config.Config{}} + + req, wireBody, err := svc.buildCountTokensRequest( + context.Background(), c, tt.account, body, + tt.token, tt.tokenType, "claude-haiku-4-5", false, + ) + require.NoError(t, err) + require.False(t, gjson.GetBytes(wireBody, "tools.0.cache_control").Exists()) + for idx := 1; idx < 5; idx++ { + require.Equal(t, "ephemeral", gjson.GetBytes(wireBody, fmt.Sprintf("tools.%d.cache_control.type", idx)).String()) + } + require.JSONEq(t, string(wireBody), string(readUpstreamBodyForTest(t, req))) + }) + } +} + // count_tokens passthrough preserve 测试 func TestBuildCountTokensRequestAnthropicAPIKeyPassthrough_PreservesContextManagementWhenClientHeaderHasBeta(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/gateway_count_tokens.go b/backend/internal/service/gateway_count_tokens.go index 7b66ac9596..003cbb7c1e 100644 --- a/backend/internal/service/gateway_count_tokens.go +++ b/backend/internal/service/gateway_count_tokens.go @@ -366,6 +366,7 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough( body []byte, token string, ) (*http.Request, error) { + body = stripDeferredToolCacheControl(body) targetURL := claudeAPICountTokensURL baseURL := account.GetBaseURL() if baseURL != "" { @@ -429,6 +430,7 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough( // buildCountTokensRequest 构建 count_tokens 上游请求 func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, mimicClaudeCode bool) (*http.Request, []byte, error) { + body = stripDeferredToolCacheControl(body) // 确定目标 URL targetURL := claudeAPICountTokensURL if account.Type == AccountTypeAPIKey { diff --git a/backend/internal/service/gateway_tool_rewrite.go b/backend/internal/service/gateway_tool_rewrite.go index da62daf961..d1b6993723 100644 --- a/backend/internal/service/gateway_tool_rewrite.go +++ b/backend/internal/service/gateway_tool_rewrite.go @@ -248,14 +248,18 @@ func applyToolNameRewriteToBody(body []byte, rw *ToolNameRewrite) []byte { return body } -// applyToolsLastCacheBreakpoint 在 tools 数组最后一个工具上注入 cache_control -// 断点,对齐 Parrot `tools[-1]["cache_control"] = {"type":"ephemeral","ttl":"1h"}` -// 行为,但 ttl 按本仓规则: +// applyToolsLastCacheBreakpoint 在最后一个非延迟加载工具上注入 cache_control +// 断点。Anthropic 不允许 defer_loading=true 的工具携带 cache_control, +// 因此会先清理所有延迟加载工具上的客户端断点。兼容官方顶层字段和 +// Claude Code 使用的 custom.defer_loading 字段。其余行为对齐 Parrot +// `tools[-1]["cache_control"] = {"type":"ephemeral","ttl":"1h"}`, +// 但 ttl 按本仓规则: // - 客户端已为该 tool 显式设置 cache_control.ttl → 完全透传不覆盖 // - 否则注入 {"type":"ephemeral","ttl": claude.DefaultCacheControlTTL} // // 纯副作用函数,tools 不存在或为空数组时 no-op。 func applyToolsLastCacheBreakpoint(body []byte) []byte { + body = stripDeferredToolCacheControl(body) tools := gjson.GetBytes(body, "tools") if !tools.IsArray() { return body @@ -264,7 +268,17 @@ func applyToolsLastCacheBreakpoint(body []byte) []byte { if len(arr) == 0 { return body } - lastIdx := len(arr) - 1 + lastIdx := -1 + for idx, tool := range arr { + if isDeferredLoadingTool(tool) { + continue + } + lastIdx = idx + } + if lastIdx == -1 { + return body + } + existingCC := arr[lastIdx].Get("cache_control") if existingCC.Exists() && existingCC.Get("ttl").String() != "" { @@ -285,6 +299,29 @@ func applyToolsLastCacheBreakpoint(body []byte) []byte { return body } +func isDeferredLoadingTool(tool gjson.Result) bool { + return tool.Get("defer_loading").Type == gjson.True || + tool.Get("custom.defer_loading").Type == gjson.True +} + +// stripDeferredToolCacheControl removes the cache marker Anthropic rejects on +// deferred tools. Only the literal JSON boolean true enables deferred loading. +func stripDeferredToolCacheControl(body []byte) []byte { + tools := gjson.GetBytes(body, "tools") + if !tools.IsArray() { + return body + } + for idx, tool := range tools.Array() { + if !isDeferredLoadingTool(tool) || !tool.Get("cache_control").Exists() { + continue + } + if next, err := sjson.DeleteBytes(body, fmt.Sprintf("tools.%d.cache_control", idx)); err == nil { + body = next + } + } + return body +} + // restoreToolNamesInBytes 对 bytes chunk 做逆向还原:假名 → 真名。 // 按 ReverseOrdered 的假名长度倒序逐个 bytes.Replace,防止子串冲突 // (与 Parrot _restore_tool_names_in_chunk 的 sorted(..., reverse=True) 等价)。 diff --git a/backend/internal/service/gateway_tool_rewrite_test.go b/backend/internal/service/gateway_tool_rewrite_test.go index 9e6f6806da..c894ef708c 100644 --- a/backend/internal/service/gateway_tool_rewrite_test.go +++ b/backend/internal/service/gateway_tool_rewrite_test.go @@ -2,6 +2,7 @@ package service import ( "context" + "fmt" "strings" "testing" @@ -149,6 +150,34 @@ func TestApplyToolsLastCacheBreakpoint_PassesThroughClientTTL(t *testing.T) { require.Equal(t, "1h", gjson.GetBytes(out, "tools.0.cache_control.ttl").String()) } +func TestApplyToolsLastCacheBreakpoint_StripsDeferredToolCacheControl(t *testing.T) { + body := []byte(`{"tools":[{"name":"a","custom":{"defer_loading":true},"cache_control":{"type":"ephemeral","ttl":"1h"}},{"name":"b","custom":{"defer_loading":true}}]}`) + out := applyToolsLastCacheBreakpoint(body) + + require.False(t, gjson.GetBytes(out, "tools.0.cache_control").Exists()) + require.False(t, gjson.GetBytes(out, "tools.1.cache_control").Exists()) +} + +func TestApplyToolsLastCacheBreakpoint_SkipsDeferredFinalTool(t *testing.T) { + body := []byte(`{"tools":[{"name":"a","input_schema":{}},{"name":"b","defer_loading":true}]}`) + out := applyToolsLastCacheBreakpoint(body) + + require.Equal(t, "ephemeral", gjson.GetBytes(out, "tools.0.cache_control.type").String()) + require.Equal(t, "5m", gjson.GetBytes(out, "tools.0.cache_control.ttl").String()) + require.False(t, gjson.GetBytes(out, "tools.1.cache_control").Exists()) +} + +func TestApplyToolsLastCacheBreakpoint_OnlyLiteralTrueIsDeferred(t *testing.T) { + body := []byte(`{"tools":[{"name":"custom-true","custom":{"defer_loading":true},"cache_control":{"type":"ephemeral"}},{"name":"top-level-true","defer_loading":true,"cache_control":{"type":"ephemeral"}},{"name":"false","defer_loading":false,"cache_control":{"type":"ephemeral"}},{"name":"string","defer_loading":"true","cache_control":{"type":"ephemeral"}},{"name":"number","defer_loading":1,"cache_control":{"type":"ephemeral"}},{"name":"object","defer_loading":{},"cache_control":{"type":"ephemeral"}}]}`) + out := stripDeferredToolCacheControl(body) + + require.False(t, gjson.GetBytes(out, "tools.0.cache_control").Exists()) + require.False(t, gjson.GetBytes(out, "tools.1.cache_control").Exists()) + for idx := 2; idx < 6; idx++ { + require.Equal(t, "ephemeral", gjson.GetBytes(out, fmt.Sprintf("tools.%d.cache_control.type", idx)).String()) + } +} + func TestStripMessageCacheControl(t *testing.T) { body := []byte(`{"messages":[{"role":"user","content":[{"type":"text","text":"hi","cache_control":{"type":"ephemeral"}}]}]}`) out := stripMessageCacheControl(body) diff --git a/backend/internal/service/gateway_upstream_request.go b/backend/internal/service/gateway_upstream_request.go index c5d962f20b..0a8efbdd5e 100644 --- a/backend/internal/service/gateway_upstream_request.go +++ b/backend/internal/service/gateway_upstream_request.go @@ -19,6 +19,7 @@ import ( ) func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, reqStream bool, mimicClaudeCode bool) (*http.Request, []byte, error) { + body = stripDeferredToolCacheControl(body) if account.Platform == PlatformAnthropic && account.Type == AccountTypeServiceAccount { req, err := s.buildUpstreamRequestAnthropicVertex(ctx, c, account, body, token, modelID, reqStream) return req, body, err