diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go index 52362dd1..4b6a47eb 100644 --- a/backend/internal/handler/gateway_helper.go +++ b/backend/internal/handler/gateway_helper.go @@ -18,18 +18,12 @@ import ( // claudeCodeValidator is a singleton validator for Claude Code client detection var claudeCodeValidator = service.NewClaudeCodeValidator() -const claudeCodeParsedRequestContextKey = "claude_code_parsed_request" - // SetClaudeCodeClientContext 检查请求是否来自 Claude Code 客户端,并设置到 context 中 // 返回更新后的 context func SetClaudeCodeClientContext(c *gin.Context, body []byte, parsedReq *service.ParsedRequest) { if c == nil || c.Request == nil { return } - if parsedReq != nil { - c.Set(claudeCodeParsedRequestContextKey, parsedReq) - } - ua := c.GetHeader("User-Agent") // Fast path:非 Claude CLI UA 直接判定 false,避免热路径二次 JSON 反序列化。 if !claudeCodeValidator.ValidateUserAgent(ua) { @@ -45,9 +39,6 @@ func SetClaudeCodeClientContext(c *gin.Context, body []byte, parsedReq *service. } else { // 仅在确认为 Claude CLI 且 messages 路径时再做 body 解析。 bodyMap := claudeCodeBodyMapFromParsedRequest(parsedReq) - if bodyMap == nil { - bodyMap = claudeCodeBodyMapFromContextCache(c) - } if bodyMap == nil && len(body) > 0 { _ = json.Unmarshal(body, &bodyMap) } @@ -87,24 +78,6 @@ func claudeCodeBodyMapFromParsedRequest(parsedReq *service.ParsedRequest) map[st return bodyMap } -func claudeCodeBodyMapFromContextCache(c *gin.Context) map[string]any { - if c == nil { - return nil - } - if bodyMap := service.CachedOpenAIParsedRequestBody(c); bodyMap != nil { - return bodyMap - } - if cached, ok := c.Get(claudeCodeParsedRequestContextKey); ok { - switch v := cached.(type) { - case *service.ParsedRequest: - return claudeCodeBodyMapFromParsedRequest(v) - case service.ParsedRequest: - return claudeCodeBodyMapFromParsedRequest(&v) - } - } - return nil -} - // 并发槽位等待相关常量 // // 性能优化说明: diff --git a/backend/internal/handler/gateway_helper_hotpath_test.go b/backend/internal/handler/gateway_helper_hotpath_test.go index d973bedb..a6b6a429 100644 --- a/backend/internal/handler/gateway_helper_hotpath_test.go +++ b/backend/internal/handler/gateway_helper_hotpath_test.go @@ -177,7 +177,7 @@ func TestSetClaudeCodeClientContext_FastPathAndStrictPath(t *testing.T) { }) } -func TestSetClaudeCodeClientContext_ReuseParsedRequestAndContextCache(t *testing.T) { +func TestSetClaudeCodeClientContext_ReuseParsedRequest(t *testing.T) { t.Run("reuse parsed request without body unmarshal", func(t *testing.T) { c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") c.Request.Header.Set("User-Agent", "claude-cli/1.0.1") @@ -192,24 +192,6 @@ func TestSetClaudeCodeClientContext_ReuseParsedRequestAndContextCache(t *testing SetClaudeCodeClientContext(c, []byte(`{invalid`), parsedReq) require.True(t, service.IsClaudeCodeClient(c.Request.Context())) }) - - t.Run("reuse context cache without body unmarshal", func(t *testing.T) { - c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") - c.Request.Header.Set("User-Agent", "claude-cli/1.0.1") - c.Request.Header.Set("X-App", "claude-code") - c.Request.Header.Set("anthropic-beta", "message-batches-2024-09-24") - c.Request.Header.Set("anthropic-version", "2023-06-01") - service.CacheOpenAIParsedRequestBody(c, []byte(`{invalid`), map[string]any{ - "model": "claude-3-5-sonnet-20241022", - "system": []any{ - map[string]any{"text": "You are Claude Code, Anthropic's official CLI for Claude."}, - }, - "metadata": map[string]any{"user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"}, - }) - - SetClaudeCodeClientContext(c, []byte(`{invalid`), nil) - require.True(t, service.IsClaudeCodeClient(c.Request.Context())) - }) } func TestWaitForSlotWithPingTimeout_AccountAndUserAcquire(t *testing.T) { diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index ac312dbb..0aa477b0 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -948,19 +948,12 @@ func (h *OpenAIGatewayHandler) validateFunctionCallOutputRequest(c *gin.Context, return true } - var reqBody map[string]any - if err := json.Unmarshal(body, &reqBody); err != nil { - // 保持原有容错语义:解析失败时跳过预校验,沿用后续上游校验结果。 - return true - } - - service.CacheOpenAIParsedRequestBody(c, body, reqBody) - validation := service.ValidateFunctionCallOutputContext(reqBody) + validation := service.ValidateFunctionCallOutputContextBytes(body) if !validation.HasFunctionCallOutput { return true } - previousResponseID, _ := reqBody["previous_response_id"].(string) + previousResponseID := gjson.GetBytes(body, "previous_response_id").String() if strings.TrimSpace(previousResponseID) != "" || validation.HasToolCallContext { return true } diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index f8c34031..62182655 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -36,13 +36,6 @@ import ( "go.uber.org/zap" ) -// openAIParsedRequestBodyCache 绑定 body 指纹,避免 handler 预解析的旧 body 污染后续 forwardBody。 -type openAIParsedRequestBodyCache struct { - bodyHash uint64 - bodyLen int - reqBody map[string]any -} - const ( // ChatGPT internal API for OAuth accounts chatgptCodexURL = "https://chatgpt.com/backend-api/codex/responses" @@ -53,8 +46,6 @@ const ( // codex_cli_only 拒绝时单个请求头日志长度上限(字符) codexCLIOnlyHeaderValueMaxBytes = 256 - // OpenAIParsedRequestBodyKey 缓存 handler 侧已解析的请求体,避免重复解析。 - OpenAIParsedRequestBodyKey = "openai_parsed_request_body" // OpenAI WS Mode 失败后的重连次数上限(不含首次尝试)。 // 与 Codex 客户端保持一致:失败后最多重连 5 次。 openAIWSReconnectRetryLimit = 5 @@ -3099,8 +3090,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco }) s.handleFailoverSideEffects(ctx, resp, account, upstreamModel) - // reqBody 会被本次账号尝试原地修改,failover 前必须释放,避免下一账号复用脏 map。 - releaseOpenAIParsedRequestBody(c) return nil, &UpstreamFailoverError{ StatusCode: resp.StatusCode, ResponseBody: respBody, @@ -3113,7 +3102,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco reasoningEffort := extractOpenAIReasoningEffort(reqBody, originalModel) serviceTier := extractOpenAIServiceTier(reqBody) - releaseOpenAIParsedRequestBody(c) + // 上游接受后只保留计费需要的标量,避免响应处理期间继续保活完整 input/tools map。 + reqBody = nil // Handle normal response var usage *OpenAIUsage @@ -6843,60 +6833,14 @@ func isEmptyBase64DataURI(raw string) bool { return strings.TrimSpace(strings.TrimPrefix(rest, "base64,")) == "" } -func getOpenAIRequestBodyMap(c *gin.Context, body []byte) (map[string]any, error) { - // 同一个 gin.Context 内 failover/渠道映射可能传入新 body,缓存必须先校验 body 指纹。 - bodyHash := xxhash.Sum64(body) - bodyLen := len(body) - if c != nil { - if cached, ok := c.Get(OpenAIParsedRequestBodyKey); ok { - if cache, ok := cached.(openAIParsedRequestBodyCache); ok && cache.reqBody != nil && cache.bodyLen == bodyLen && cache.bodyHash == bodyHash { - return cache.reqBody, nil - } - } - } - +func getOpenAIRequestBodyMap(_ *gin.Context, body []byte) (map[string]any, error) { var reqBody map[string]any if err := json.Unmarshal(body, &reqBody); err != nil { return nil, fmt.Errorf("parse request: %w", err) } - if c != nil { - c.Set(OpenAIParsedRequestBodyKey, openAIParsedRequestBodyCache{bodyHash: bodyHash, bodyLen: bodyLen, reqBody: reqBody}) - } return reqBody, nil } -// CacheOpenAIParsedRequestBody 仅缓存与当前 body 绑定的解析结果。 -func CacheOpenAIParsedRequestBody(c *gin.Context, body []byte, reqBody map[string]any) { - if c == nil || reqBody == nil { - return - } - c.Set(OpenAIParsedRequestBodyKey, openAIParsedRequestBodyCache{ - bodyHash: xxhash.Sum64(body), - bodyLen: len(body), - reqBody: reqBody, - }) -} - -// CachedOpenAIParsedRequestBody 只给同请求内不关心 body 参数的轻量识别逻辑使用。 -func CachedOpenAIParsedRequestBody(c *gin.Context) map[string]any { - if c == nil { - return nil - } - if cached, ok := c.Get(OpenAIParsedRequestBodyKey); ok { - if cache, ok := cached.(openAIParsedRequestBodyCache); ok { - return cache.reqBody - } - } - return nil -} - -func releaseOpenAIParsedRequestBody(c *gin.Context) { - if c == nil { - return - } - delete(c.Keys, OpenAIParsedRequestBodyKey) -} - func extractOpenAIReasoningEffort(reqBody map[string]any, requestedModel string) *string { if value, present := getOpenAIReasoningEffortFromReqBody(reqBody); present { if value == "" { diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go index 2ff72e2f..af0d21c4 100644 --- a/backend/internal/service/openai_gateway_service_hotpath_test.go +++ b/backend/internal/service/openai_gateway_service_hotpath_test.go @@ -106,26 +106,13 @@ func TestExtractOpenAIReasoningEffortFromBody(t *testing.T) { } } -func TestGetOpenAIRequestBodyMap_UsesContextCache(t *testing.T) { - gin.SetMode(gin.TestMode) - rec := httptest.NewRecorder() - c, _ := gin.CreateTestContext(rec) - - cached := map[string]any{"model": "cached-model", "stream": true} - CacheOpenAIParsedRequestBody(c, []byte(`{invalid-json`), cached) - - got, err := getOpenAIRequestBodyMap(c, []byte(`{invalid-json`)) - require.NoError(t, err) - require.Equal(t, cached, got) -} - -func TestGetOpenAIRequestBodyMap_ParseErrorWithoutCache(t *testing.T) { +func TestGetOpenAIRequestBodyMap_ParseError(t *testing.T) { _, err := getOpenAIRequestBodyMap(nil, []byte(`{invalid-json`)) require.Error(t, err) require.Contains(t, err.Error(), "parse request") } -func TestGetOpenAIRequestBodyMap_WriteBackContextCache(t *testing.T) { +func TestGetOpenAIRequestBodyMap_DoesNotWriteContextCache(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() c, _ := gin.CreateTestContext(rec) @@ -133,20 +120,7 @@ func TestGetOpenAIRequestBodyMap_WriteBackContextCache(t *testing.T) { got, err := getOpenAIRequestBodyMap(c, []byte(`{"model":"gpt-5","stream":true}`)) require.NoError(t, err) require.Equal(t, "gpt-5", got["model"]) - - require.Equal(t, got, CachedOpenAIParsedRequestBody(c)) -} - -func TestGetOpenAIRequestBodyMap_IgnoresCacheForDifferentBody(t *testing.T) { - gin.SetMode(gin.TestMode) - rec := httptest.NewRecorder() - c, _ := gin.CreateTestContext(rec) - - CacheOpenAIParsedRequestBody(c, []byte(`{"model":"cached-model"}`), map[string]any{"model": "cached-model"}) - - got, err := getOpenAIRequestBodyMap(c, []byte(`{"model":"forward-model"}`)) - require.NoError(t, err) - require.Equal(t, "forward-model", got["model"]) + require.Empty(t, c.Keys) } func TestSanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(t *testing.T) { diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go index 7d503f5a..36064549 100644 --- a/backend/internal/service/openai_tool_continuation.go +++ b/backend/internal/service/openai_tool_continuation.go @@ -1,6 +1,10 @@ package service -import "strings" +import ( + "strings" + + "github.com/tidwall/gjson" +) // ToolContinuationSignals 聚合工具续链相关信号,避免重复遍历 input。 type ToolContinuationSignals struct { @@ -150,6 +154,63 @@ func AnalyzeToolContinuationSignals(reqBody map[string]any) ToolContinuationSign return signals } +// ValidateFunctionCallOutputContextBytes 基于 raw JSON 校验工具输出续链,避免 handler 预校验阶段全量解码大 input。 +func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValidation { + result := FunctionCallOutputValidation{} + input := gjson.GetBytes(body, "input") + if !input.IsArray() { + return result + } + + var callIDs map[string]struct{} + var referenceIDs map[string]struct{} + input.ForEach(func(_, item gjson.Result) bool { + if !item.IsObject() { + return true + } + itemType := item.Get("type").String() + switch { + case isCodexToolCallOutputItemType(itemType): + result.HasFunctionCallOutput = true + callID := strings.TrimSpace(item.Get("call_id").String()) + if callID == "" { + result.HasFunctionCallOutputMissingCallID = true + return true + } + if callIDs == nil { + callIDs = make(map[string]struct{}) + } + callIDs[callID] = struct{}{} + case isCodexToolCallContextItemType(itemType): + if strings.TrimSpace(item.Get("call_id").String()) != "" { + result.HasToolCallContext = true + } + case itemType == "item_reference": + idValue := strings.TrimSpace(item.Get("id").String()) + if idValue == "" { + return true + } + if referenceIDs == nil { + referenceIDs = make(map[string]struct{}) + } + referenceIDs[idValue] = struct{}{} + } + return !(result.HasFunctionCallOutput && result.HasToolCallContext) + }) + if !result.HasFunctionCallOutput || result.HasToolCallContext || len(callIDs) == 0 || len(referenceIDs) == 0 { + return result + } + allReferenced := true + for callID := range callIDs { + if _, ok := referenceIDs[callID]; !ok { + allReferenced = false + break + } + } + result.HasItemReferenceForAllCallIDs = allReferenced + return result +} + // ValidateFunctionCallOutputContext 为 handler 提供低开销校验结果: // 1) 无工具输出直接返回 // 2) 若已存在工具调用上下文则提前返回 diff --git a/backend/internal/service/openai_tool_continuation_test.go b/backend/internal/service/openai_tool_continuation_test.go index 0e0552f6..4610652b 100644 --- a/backend/internal/service/openai_tool_continuation_test.go +++ b/backend/internal/service/openai_tool_continuation_test.go @@ -1,6 +1,7 @@ package service import ( + "encoding/json" "testing" "github.com/stretchr/testify/require" @@ -118,3 +119,68 @@ func TestHasItemReferenceForCallIDs(t *testing.T) { require.True(t, HasItemReferenceForCallIDs(req, []string{"call_1", "call_2"})) require.False(t, HasItemReferenceForCallIDs(req, []string{"call_1", "call_3"})) } + +func TestValidateFunctionCallOutputContextBytesMatchesMapValidation(t *testing.T) { + // handler 预校验走 raw JSON 扫描,语义必须与 service 内部 map 校验保持一致。 + cases := []struct { + name string + body map[string]any + }{ + { + name: "no_input", + body: map[string]any{"model": "gpt-5.4"}, + }, + { + name: "missing_call_id", + body: map[string]any{"input": []any{map[string]any{"type": "function_call_output"}}}, + }, + { + name: "call_id_without_reference", + body: map[string]any{"input": []any{map[string]any{"type": "function_call_output", "call_id": "call_1"}}}, + }, + { + name: "matching_reference", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call_output", "call_id": "call_1"}, + map[string]any{"type": "item_reference", "id": "call_1"}, + }}, + }, + { + name: "partial_reference", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call_output", "call_id": "call_1"}, + map[string]any{"type": "tool_search_output", "call_id": "call_2"}, + map[string]any{"type": "item_reference", "id": "call_1"}, + }}, + }, + { + name: "tool_context", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call_output", "call_id": "call_1"}, + map[string]any{"type": "function_call", "call_id": "call_1"}, + }}, + }, + { + name: "all_codex_tool_outputs", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call_output", "call_id": "call_function"}, + map[string]any{"type": "tool_search_output", "call_id": "call_search"}, + map[string]any{"type": "custom_tool_call_output", "call_id": "call_custom"}, + map[string]any{"type": "mcp_tool_call_output", "call_id": "call_mcp"}, + map[string]any{"type": "item_reference", "id": "call_function"}, + map[string]any{"type": "item_reference", "id": "call_search"}, + map[string]any{"type": "item_reference", "id": "call_custom"}, + map[string]any{"type": "item_reference", "id": "call_mcp"}, + }}, + }, + } + + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + bodyBytes, err := json.Marshal(tt.body) + require.NoError(t, err) + + require.Equal(t, ValidateFunctionCallOutputContext(tt.body), ValidateFunctionCallOutputContextBytes(bodyBytes)) + }) + } +}