mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:38:22 +08:00
Merge pull request #5004 from wucm667/fix/issue-4990-deferred-tool-cache-control
fix(claude): strip cache control from deferred tools
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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,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) 等价)。
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user