From addd5ef1dcd7458ae7b97f312fb06418edb217f4 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Mon, 20 Jul 2026 22:43:11 +0800 Subject: [PATCH] [verified] fix: align sync cache billing after failover --- backend/internal/handler/gateway_handler.go | 3 ++ .../service/gateway_anthropic_passthrough.go | 19 ++++++++ .../gateway_non_streaming_response_test.go | 47 +++++++++++++++++++ 3 files changed, 69 insertions(+) diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index f5d99f3c3d..891878799a 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -797,6 +797,9 @@ func (h *GatewayHandler) Messages(c *gin.Context) { if fs.SwitchCount > 0 { requestCtx = service.WithAccountSwitchCount(requestCtx, fs.SwitchCount, h.metadataBridgeEnabled()) } + if fs.ForceCacheBilling { + requestCtx = service.WithForceCacheBilling(requestCtx) + } // 记录 Forward 前已写入字节数,Forward 后若增加则说明 SSE 内容已发,禁止 failover writerSizeBeforeForward := c.Writer.Size() if account.Platform == service.PlatformAntigravity && account.Type != service.AccountTypeAPIKey { diff --git a/backend/internal/service/gateway_anthropic_passthrough.go b/backend/internal/service/gateway_anthropic_passthrough.go index 7309e141aa..2eb7a97e3b 100644 --- a/backend/internal/service/gateway_anthropic_passthrough.go +++ b/backend/internal/service/gateway_anthropic_passthrough.go @@ -20,6 +20,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/tidwall/gjson" + "github.com/tidwall/sjson" "github.com/gin-gonic/gin" ) @@ -782,6 +783,12 @@ func (s *GatewayService) handleNonStreamingResponseAnthropicAPIKeyPassthrough( } usage := parseClaudeUsageFromResponseBody(body) + if IsForceCacheBilling(ctx) && usage.InputTokens > 0 { + body, err = classifyAnthropicResponseInputAsCacheRead(body, usage) + if err != nil { + return nil, err + } + } writeAnthropicPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) contentType := strings.TrimSpace(resp.Header.Get("Content-Type")) @@ -793,6 +800,18 @@ func (s *GatewayService) handleNonStreamingResponseAnthropicAPIKeyPassthrough( return usage, nil } +func classifyAnthropicResponseInputAsCacheRead(body []byte, usage *ClaudeUsage) ([]byte, error) { + classified, err := sjson.SetBytes(body, "usage.input_tokens", 0) + if err != nil { + return nil, fmt.Errorf("classify forced cache billing input tokens: %w", err) + } + classified, err = sjson.SetBytes(classified, "usage.cache_read_input_tokens", usage.CacheReadInputTokens+usage.InputTokens) + if err != nil { + return nil, fmt.Errorf("classify forced cache billing cache read tokens: %w", err) + } + return classified, nil +} + func writeAnthropicPassthroughResponseHeaders(dst http.Header, src http.Header, filter *responseheaders.CompiledHeaderFilter) { if dst == nil || src == nil { return diff --git a/backend/internal/service/gateway_non_streaming_response_test.go b/backend/internal/service/gateway_non_streaming_response_test.go index a812a62e46..8dc48cef1f 100644 --- a/backend/internal/service/gateway_non_streaming_response_test.go +++ b/backend/internal/service/gateway_non_streaming_response_test.go @@ -13,6 +13,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" ) type nonJSONTempUnschedAccountRepo struct { @@ -143,6 +144,52 @@ func TestHandleNonStreamingResponseAnthropicAPIKeyPassthrough_ValidJSONUnchanged require.JSONEq(t, string(body), rec.Body.String()) } +func TestHandleNonStreamingResponseAnthropicAPIKeyPassthrough_ForceCacheBillingResponse(t *testing.T) { + tests := []struct { + name string + body string + want string + }{ + { + name: "converts input tokens for downstream billing", + body: `{"id":"msg_1","type":"message","content":[{"type":"text","text":"unchanged"}],"usage":{"input_tokens":5,"output_tokens":3}}`, + want: `{"id":"msg_1","type":"message","content":[{"type":"text","text":"unchanged"}],"usage":{"input_tokens":0,"output_tokens":3,"cache_read_input_tokens":5}}`, + }, + { + name: "adds to genuine cache reads", + body: `{"id":"msg_2","type":"message","usage":{"input_tokens":5,"output_tokens":3,"cache_read_input_tokens":7,"cache_creation_input_tokens":11}}`, + want: `{"id":"msg_2","type":"message","usage":{"input_tokens":0,"output_tokens":3,"cache_read_input_tokens":12,"cache_creation_input_tokens":11}}`, + }, + { + name: "zero input leaves response unchanged", + body: `{"id":"msg_3","type":"message","usage":{"input_tokens":0,"output_tokens":3,"cache_read_input_tokens":7}}`, + want: `{"id":"msg_3","type":"message","usage":{"input_tokens":0,"output_tokens":3,"cache_read_input_tokens":7}}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(bytes.NewBufferString(tt.body)), + } + svc := &GatewayService{cfg: &config.Config{}} + + usage, err := svc.handleNonStreamingResponseAnthropicAPIKeyPassthrough(WithForceCacheBilling(context.Background()), resp, c, &Account{ID: 2}) + + require.NoError(t, err) + require.Equal(t, int(gjson.Get(tt.body, "usage.input_tokens").Int()), usage.InputTokens, "local accounting must retain the unclassified usage") + require.Equal(t, int(gjson.Get(tt.body, "usage.cache_read_input_tokens").Int()), usage.CacheReadInputTokens, "local accounting must convert exactly once in RecordUsage") + require.JSONEq(t, tt.want, rec.Body.String()) + }) + } +} + func TestHandleNonStreamingResponse_NonJSON2xxMatchesModelScopedTempUnschedulableRule(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder()