fix(claude): strip cache control from deferred tools

This commit is contained in:
wucm667
2026-08-11 16:57:38 +08:00
parent 1e618dbc29
commit 9c36b75a7d
7 changed files with 141 additions and 4 deletions
@@ -769,6 +769,30 @@ 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":"ordinary","custom":{"defer_loading":false},"cache_control":{"type":"ephemeral"}},{"name":"malformed","custom":{"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.True(t, gjson.GetBytes(wireBody, "tools.1.cache_control").Exists())
require.True(t, gjson.GetBytes(wireBody, "tools.2.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.True(t, gjson.GetBytes(countBody, "tools.1.cache_control").Exists())
require.True(t, gjson.GetBytes(countBody, "tools.2.cache_control").Exists())
}
func TestGatewayService_AnthropicOAuth_NotAffectedByAPIKeyPassthroughToggle(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
@@ -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 != "" {
@@ -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)
@@ -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 {
@@ -248,14 +248,17 @@ 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 不允许 custom.defer_loading=true 的工具携带 cache_control,
// 因此会先清理所有延迟加载工具上的客户端断点。其余行为对齐 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 +267,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 +298,28 @@ func applyToolsLastCacheBreakpoint(body []byte) []byte {
return body
}
func isDeferredLoadingTool(tool gjson.Result) bool {
return 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) 等价)。
@@ -2,6 +2,7 @@ package service
import (
"context"
"fmt"
"strings"
"testing"
@@ -149,6 +150,33 @@ 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","custom":{"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":"true","custom":{"defer_loading":true},"cache_control":{"type":"ephemeral"}},{"name":"false","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"}}]}`)
out := stripDeferredToolCacheControl(body)
require.False(t, gjson.GetBytes(out, "tools.0.cache_control").Exists())
for idx := 1; idx < 5; 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)
@@ -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