Skip to content

Commit a6764e8

Browse files
committed
修复 OAuth/SetupToken 转发请求体重排并增加调试开关
1 parent 9f6ab6b commit a6764e8

8 files changed

Lines changed: 722 additions & 368 deletions

backend/internal/pkg/antigravity/request_transformer.go

Lines changed: 4 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -275,21 +275,6 @@ func filterOpenCodePrompt(text string) string {
275275
return ""
276276
}
277277

278-
// systemBlockFilterPrefixes 需要从 system 中过滤的文本前缀列表
279-
var systemBlockFilterPrefixes = []string{
280-
"x-anthropic-billing-header",
281-
}
282-
283-
// filterSystemBlockByPrefix 如果文本匹配过滤前缀,返回空字符串
284-
func filterSystemBlockByPrefix(text string) string {
285-
for _, prefix := range systemBlockFilterPrefixes {
286-
if strings.HasPrefix(text, prefix) {
287-
return ""
288-
}
289-
}
290-
return text
291-
}
292-
293278
// buildSystemInstruction 构建 systemInstruction(与 Antigravity-Manager 保持一致)
294279
func buildSystemInstruction(system json.RawMessage, modelName string, opts TransformOptions, tools []ClaudeTool) *GeminiContent {
295280
var parts []GeminiPart
@@ -306,8 +291,8 @@ func buildSystemInstruction(system json.RawMessage, modelName string, opts Trans
306291
if strings.Contains(sysStr, "You are Antigravity") {
307292
userHasAntigravityIdentity = true
308293
}
309-
// 过滤 OpenCode 默认提示词和黑名单前缀
310-
filtered := filterSystemBlockByPrefix(filterOpenCodePrompt(sysStr))
294+
// 过滤 OpenCode 默认提示词
295+
filtered := filterOpenCodePrompt(sysStr)
311296
if filtered != "" {
312297
userSystemParts = append(userSystemParts, GeminiPart{Text: filtered})
313298
}
@@ -321,8 +306,8 @@ func buildSystemInstruction(system json.RawMessage, modelName string, opts Trans
321306
if strings.Contains(block.Text, "You are Antigravity") {
322307
userHasAntigravityIdentity = true
323308
}
324-
// 过滤 OpenCode 默认提示词和黑名单前缀
325-
filtered := filterSystemBlockByPrefix(filterOpenCodePrompt(block.Text))
309+
// 过滤 OpenCode 默认提示词
310+
filtered := filterOpenCodePrompt(block.Text)
326311
if filtered != "" {
327312
userSystemParts = append(userSystemParts, GeminiPart{Text: filtered})
328313
}

backend/internal/pkg/antigravity/request_transformer_test.go

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,10 @@ package antigravity
22

33
import (
44
"encoding/json"
5+
"strings"
56
"testing"
7+
8+
"github.com/stretchr/testify/require"
69
)
710

811
// TestBuildParts_ThinkingBlockWithoutSignature 测试thinking block无signature时的处理
@@ -349,3 +352,51 @@ func TestBuildGenerationConfig_ThinkingDynamicBudget(t *testing.T) {
349352
})
350353
}
351354
}
355+
356+
func TestTransformClaudeToGeminiWithOptions_PreservesBillingHeaderSystemBlock(t *testing.T) {
357+
tests := []struct {
358+
name string
359+
system json.RawMessage
360+
}{
361+
{
362+
name: "system array",
363+
system: json.RawMessage(`[{"type":"text","text":"x-anthropic-billing-header keep"}]`),
364+
},
365+
{
366+
name: "system string",
367+
system: json.RawMessage(`"x-anthropic-billing-header keep"`),
368+
},
369+
}
370+
371+
for _, tt := range tests {
372+
t.Run(tt.name, func(t *testing.T) {
373+
claudeReq := &ClaudeRequest{
374+
Model: "claude-3-5-sonnet-latest",
375+
System: tt.system,
376+
Messages: []ClaudeMessage{
377+
{
378+
Role: "user",
379+
Content: json.RawMessage(`[{"type":"text","text":"hello"}]`),
380+
},
381+
},
382+
}
383+
384+
body, err := TransformClaudeToGeminiWithOptions(claudeReq, "project-1", "gemini-2.5-flash", DefaultTransformOptions())
385+
require.NoError(t, err)
386+
387+
var req V1InternalRequest
388+
require.NoError(t, json.Unmarshal(body, &req))
389+
require.NotNil(t, req.Request.SystemInstruction)
390+
391+
found := false
392+
for _, part := range req.Request.SystemInstruction.Parts {
393+
if strings.Contains(part.Text, "x-anthropic-billing-header keep") {
394+
found = true
395+
break
396+
}
397+
}
398+
399+
require.True(t, found, "转换后的 systemInstruction 应保留 x-anthropic-billing-header 内容")
400+
})
401+
}
402+
}

backend/internal/service/gateway_anthropic_apikey_passthrough_test.go

Lines changed: 77 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -688,6 +688,83 @@ func TestGatewayService_AnthropicOAuth_NotAffectedByAPIKeyPassthroughToggle(t *t
688688
require.Contains(t, req.Header.Get("anthropic-beta"), claude.BetaOAuth, "OAuth 链路仍应按原逻辑补齐 oauth beta")
689689
}
690690

691+
func TestGatewayService_AnthropicOAuth_ForwardPreservesBillingHeaderSystemBlock(t *testing.T) {
692+
gin.SetMode(gin.TestMode)
693+
694+
tests := []struct {
695+
name string
696+
body string
697+
}{
698+
{
699+
name: "system array",
700+
body: `{"model":"claude-3-5-sonnet-latest","system":[{"type":"text","text":"x-anthropic-billing-header keep"}],"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`,
701+
},
702+
{
703+
name: "system string",
704+
body: `{"model":"claude-3-5-sonnet-latest","system":"x-anthropic-billing-header keep","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`,
705+
},
706+
}
707+
708+
for _, tt := range tests {
709+
t.Run(tt.name, func(t *testing.T) {
710+
rec := httptest.NewRecorder()
711+
c, _ := gin.CreateTestContext(rec)
712+
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
713+
714+
parsed, err := ParseGatewayRequest([]byte(tt.body), PlatformAnthropic)
715+
require.NoError(t, err)
716+
717+
upstream := &anthropicHTTPUpstreamRecorder{
718+
resp: &http.Response{
719+
StatusCode: http.StatusOK,
720+
Header: http.Header{
721+
"Content-Type": []string{"application/json"},
722+
"x-request-id": []string{"rid-oauth-preserve"},
723+
},
724+
Body: io.NopCloser(strings.NewReader(`{"id":"msg_1","type":"message","role":"assistant","model":"claude-3-5-sonnet-20241022","content":[{"type":"text","text":"ok"}],"usage":{"input_tokens":12,"output_tokens":7}}`)),
725+
},
726+
}
727+
728+
cfg := &config.Config{
729+
Gateway: config.GatewayConfig{
730+
MaxLineSize: defaultMaxLineSize,
731+
},
732+
}
733+
svc := &GatewayService{
734+
cfg: cfg,
735+
responseHeaderFilter: compileResponseHeaderFilter(cfg),
736+
httpUpstream: upstream,
737+
rateLimitService: &RateLimitService{},
738+
deferredService: &DeferredService{},
739+
}
740+
741+
account := &Account{
742+
ID: 301,
743+
Name: "anthropic-oauth-preserve",
744+
Platform: PlatformAnthropic,
745+
Type: AccountTypeOAuth,
746+
Concurrency: 1,
747+
Credentials: map[string]any{
748+
"access_token": "oauth-token",
749+
},
750+
Status: StatusActive,
751+
Schedulable: true,
752+
}
753+
754+
result, err := svc.Forward(context.Background(), c, account, parsed)
755+
require.NoError(t, err)
756+
require.NotNil(t, result)
757+
require.NotNil(t, upstream.lastReq)
758+
require.Equal(t, "Bearer oauth-token", upstream.lastReq.Header.Get("authorization"))
759+
require.Contains(t, upstream.lastReq.Header.Get("anthropic-beta"), claude.BetaOAuth)
760+
761+
system := gjson.GetBytes(upstream.lastBody, "system")
762+
require.True(t, system.Exists())
763+
require.Contains(t, system.Raw, "x-anthropic-billing-header keep")
764+
})
765+
}
766+
}
767+
691768
func TestGatewayService_AnthropicAPIKeyPassthrough_StreamingStillCollectsUsageAfterClientDisconnect(t *testing.T) {
692769
gin.SetMode(gin.TestMode)
693770

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,72 @@
1+
package service
2+
3+
import (
4+
"strings"
5+
"testing"
6+
7+
"github.com/Wei-Shaw/sub2api/internal/pkg/claude"
8+
"github.com/stretchr/testify/require"
9+
)
10+
11+
func assertJSONTokenOrder(t *testing.T, body string, tokens ...string) {
12+
t.Helper()
13+
14+
last := -1
15+
for _, token := range tokens {
16+
pos := strings.Index(body, token)
17+
require.NotEqualf(t, -1, pos, "missing token %s in body %s", token, body)
18+
require.Greaterf(t, pos, last, "token %s should appear after previous tokens in body %s", token, body)
19+
last = pos
20+
}
21+
}
22+
23+
func TestReplaceModelInBody_PreservesTopLevelFieldOrder(t *testing.T) {
24+
svc := &GatewayService{}
25+
body := []byte(`{"alpha":1,"model":"claude-3-5-sonnet-latest","messages":[],"omega":2}`)
26+
27+
result := svc.replaceModelInBody(body, "claude-3-5-sonnet-20241022")
28+
resultStr := string(result)
29+
30+
assertJSONTokenOrder(t, resultStr, `"alpha"`, `"model"`, `"messages"`, `"omega"`)
31+
require.Contains(t, resultStr, `"model":"claude-3-5-sonnet-20241022"`)
32+
}
33+
34+
func TestNormalizeClaudeOAuthRequestBody_PreservesTopLevelFieldOrder(t *testing.T) {
35+
body := []byte(`{"alpha":1,"model":"claude-3-5-sonnet-latest","temperature":0.2,"system":"You are OpenCode, the best coding agent on the planet.","messages":[],"tool_choice":{"type":"auto"},"omega":2}`)
36+
37+
result, modelID := normalizeClaudeOAuthRequestBody(body, "claude-3-5-sonnet-latest", claudeOAuthNormalizeOptions{
38+
injectMetadata: true,
39+
metadataUserID: "user-1",
40+
})
41+
resultStr := string(result)
42+
43+
require.Equal(t, claude.NormalizeModelID("claude-3-5-sonnet-latest"), modelID)
44+
assertJSONTokenOrder(t, resultStr, `"alpha"`, `"model"`, `"system"`, `"messages"`, `"omega"`, `"tools"`, `"metadata"`)
45+
require.NotContains(t, resultStr, `"temperature"`)
46+
require.NotContains(t, resultStr, `"tool_choice"`)
47+
require.Contains(t, resultStr, `"system":"`+claudeCodeSystemPrompt+`"`)
48+
require.Contains(t, resultStr, `"tools":[]`)
49+
require.Contains(t, resultStr, `"metadata":{"user_id":"user-1"}`)
50+
}
51+
52+
func TestInjectClaudeCodePrompt_PreservesFieldOrder(t *testing.T) {
53+
body := []byte(`{"alpha":1,"system":[{"id":"block-1","type":"text","text":"Custom"}],"messages":[],"omega":2}`)
54+
55+
result := injectClaudeCodePrompt(body, []any{
56+
map[string]any{"id": "block-1", "type": "text", "text": "Custom"},
57+
})
58+
resultStr := string(result)
59+
60+
assertJSONTokenOrder(t, resultStr, `"alpha"`, `"system"`, `"messages"`, `"omega"`)
61+
require.Contains(t, resultStr, `{"id":"block-1","type":"text","text":"`+claudeCodeSystemPrompt+`\n\nCustom"}`)
62+
}
63+
64+
func TestEnforceCacheControlLimit_PreservesTopLevelFieldOrder(t *testing.T) {
65+
body := []byte(`{"alpha":1,"system":[{"type":"text","text":"s1","cache_control":{"type":"ephemeral"}},{"type":"text","text":"s2","cache_control":{"type":"ephemeral"}}],"messages":[{"role":"user","content":[{"type":"text","text":"m1","cache_control":{"type":"ephemeral"}},{"type":"text","text":"m2","cache_control":{"type":"ephemeral"}},{"type":"text","text":"m3","cache_control":{"type":"ephemeral"}}]}],"omega":2}`)
66+
67+
result := enforceCacheControlLimit(body)
68+
resultStr := string(result)
69+
70+
assertJSONTokenOrder(t, resultStr, `"alpha"`, `"system"`, `"messages"`, `"omega"`)
71+
require.Equal(t, 4, strings.Count(resultStr, `"cache_control"`))
72+
}
Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,34 @@
1+
package service
2+
3+
import "testing"
4+
5+
func TestDebugGatewayBodyLoggingEnabled(t *testing.T) {
6+
t.Run("default disabled", func(t *testing.T) {
7+
t.Setenv(debugGatewayBodyEnv, "")
8+
if debugGatewayBodyLoggingEnabled() {
9+
t.Fatalf("expected debug gateway body logging to be disabled by default")
10+
}
11+
})
12+
13+
t.Run("enabled with true-like values", func(t *testing.T) {
14+
for _, value := range []string{"1", "true", "TRUE", "yes", "on"} {
15+
t.Run(value, func(t *testing.T) {
16+
t.Setenv(debugGatewayBodyEnv, value)
17+
if !debugGatewayBodyLoggingEnabled() {
18+
t.Fatalf("expected debug gateway body logging to be enabled for %q", value)
19+
}
20+
})
21+
}
22+
})
23+
24+
t.Run("disabled with other values", func(t *testing.T) {
25+
for _, value := range []string{"0", "false", "off", "debug"} {
26+
t.Run(value, func(t *testing.T) {
27+
t.Setenv(debugGatewayBodyEnv, value)
28+
if debugGatewayBodyLoggingEnabled() {
29+
t.Fatalf("expected debug gateway body logging to be disabled for %q", value)
30+
}
31+
})
32+
}
33+
})
34+
}

0 commit comments

Comments
 (0)