diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index eb5c4a42..79bed8b9 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -154,7 +154,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) { setOpsRequestContext(c, "", false) - parsedReq, err := service.ParseGatewayRequest(body, domain.PlatformAnthropic) + bodyRef := service.NewRequestBodyRef(body) + parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic) if err != nil { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return @@ -746,11 +747,11 @@ func (h *GatewayHandler) Messages(c *gin.Context) { // 应用渠道模型映射到请求 if channelMapping.Mapped { parsedReq.Model = channelMapping.MappedModel - parsedReq.Body = h.gatewayService.ReplaceModelInBody(parsedReq.Body, channelMapping.MappedModel) + parsedReq.Body.Replace(h.gatewayService.ReplaceModelInBody(parsedReq.Body.Bytes(), channelMapping.MappedModel)) } // Bedrock CC 兼容:渠道模型映射后,清理 Anthropic API 专有字段、注入 Bedrock 必需字段 - parsedReq.Body = h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body, parsedReq.Model, account, apiKey.GroupID) - body = parsedReq.Body + parsedReq.Body.Replace(h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body.Bytes(), parsedReq.Model, account, apiKey.GroupID)) + body = parsedReq.Body.Bytes() // 转发请求 - 根据账号平台分流 c.Set("parsed_request", parsedReq) @@ -1683,7 +1684,8 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) { setOpsRequestContext(c, "", false) - parsedReq, err := service.ParseGatewayRequest(body, domain.PlatformAnthropic) + bodyRef := service.NewRequestBodyRef(body) + parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic) if err != nil { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") return diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index daf6e6ea..719700aa 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -151,9 +151,10 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { } // Parse request for session hash - parsedReq, _ := service.ParseGatewayRequest(body, "chat_completions") + bodyRef := service.NewRequestBodyRef(body) + parsedReq, _ := service.ParseGatewayRequest(bodyRef, "chat_completions") if parsedReq == nil { - parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: body} + parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: bodyRef} } parsedReq.SessionContext = &service.SessionContext{ ClientIP: ip.GetClientIP(c), diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index f57b9989..49f80d19 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -156,9 +156,10 @@ func (h *GatewayHandler) Responses(c *gin.Context) { } // Parse request for session hash - parsedReq, _ := service.ParseGatewayRequest(body, "responses") + bodyRef := service.NewRequestBodyRef(body) + parsedReq, _ := service.ParseGatewayRequest(bodyRef, "responses") if parsedReq == nil { - parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: body} + parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: bodyRef} } parsedReq.SessionContext = &service.SessionContext{ ClientIP: ip.GetClientIP(c), diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index 0b33ca3e..5d8e6fa8 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -262,7 +262,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { sessionHash := extractGeminiCLISessionHash(c, body) if sessionHash == "" { // Fallback: 使用通用的会话哈希生成逻辑(适用于其他客户端) - parsedReq, _ := service.ParseGatewayRequest(body, domain.PlatformGemini) + parsedReq, _ := service.ParseGatewayRequest(service.NewRequestBodyRef(body), domain.PlatformGemini) if parsedReq != nil { parsedReq.SessionContext = &service.SessionContext{ ClientIP: ip.GetClientIP(c), diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go index 9062c517..a67a3dc2 100644 --- a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go +++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go @@ -112,7 +112,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardStreamPreservesBodyAnd body := []byte(`{"model":"claude-3-7-sonnet-20250219","stream":true,"system":[{"type":"text","text":"x-anthropic-billing-header keep"}],"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`) parsed := &ParsedRequest{ - Body: body, + Body: NewRequestBodyRef(body), Model: "claude-3-7-sonnet-20250219", Stream: true, } @@ -202,7 +202,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardCountTokensPreservesBo body := []byte(`{"model":"claude-3-5-sonnet-latest","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}],"thinking":{"type":"enabled"}}`) parsed := &ParsedRequest{ - Body: body, + Body: NewRequestBodyRef(body), Model: "claude-3-5-sonnet-latest", } @@ -344,7 +344,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingEdgeCases(t *test body := []byte(`{"model":"` + tt.model + `","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`) parsed := &ParsedRequest{ - Body: body, + Body: NewRequestBodyRef(body), Model: tt.model, } @@ -429,7 +429,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingPreservesOtherFie // 包含复杂字段的请求体:system、thinking、messages body := []byte(`{"model":"claude-sonnet-4-20250514","system":[{"type":"text","text":"You are a helpful assistant."}],"messages":[{"role":"user","content":[{"type":"text","text":"hello world"}]}],"thinking":{"type":"enabled","budget_tokens":5000},"max_tokens":1024}`) parsed := &ParsedRequest{ - Body: body, + Body: NewRequestBodyRef(body), Model: "claude-sonnet-4-20250514", } @@ -485,7 +485,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_CountTokensFiltersGenerationF body := []byte(`{"model":"claude-sonnet-4-20250514","system":[{"type":"text","text":"sys"}],"messages":[{"role":"user","content":"hello"}],"tools":[{"name":"tool","input_schema":{"type":"object"}}],"temperature":0.7,"top_p":0.9,"top_k":40,"stream":true,"stop_sequences":["END"],"max_tokens":1024,"thinking":{"type":"enabled","budget_tokens":5000}}`) parsed := &ParsedRequest{ - Body: body, + Body: NewRequestBodyRef(body), Model: "claude-sonnet-4-20250514", } @@ -547,7 +547,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_EmptyModelSkipsMapping(t *tes body := []byte(`{"messages":[{"role":"user","content":"hello"}]}`) parsed := &ParsedRequest{ - Body: body, + Body: NewRequestBodyRef(body), Model: "", // 空模型 } @@ -636,7 +636,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_CountTokens404PassthroughNotE c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil) body := []byte(`{"model":"claude-sonnet-4-5-20250929","messages":[{"role":"user","content":"hi"}]}`) - parsed := &ParsedRequest{Body: body, Model: "claude-sonnet-4-5-20250929"} + parsed := &ParsedRequest{Body: NewRequestBodyRef(body), Model: "claude-sonnet-4-5-20250929"} upstream := &anthropicHTTPUpstreamRecorder{ resp: &http.Response{ @@ -767,7 +767,7 @@ func TestGatewayService_AnthropicOAuth_ForwardPreservesBillingHeaderSystemBlock( c, _ := gin.CreateTestContext(rec) c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) - parsed, err := ParseGatewayRequest([]byte(tt.body), PlatformAnthropic) + parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), PlatformAnthropic) require.NoError(t, err) upstream := &anthropicHTTPUpstreamRecorder{ diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go index 91f7601c..819bb0a8 100644 --- a/backend/internal/service/gateway_request.go +++ b/backend/internal/service/gateway_request.go @@ -51,6 +51,35 @@ type SessionContext struct { APIKeyID int64 } +type RequestBodyRef struct { + data []byte +} + +func NewRequestBodyRef(data []byte) *RequestBodyRef { + return &RequestBodyRef{data: data} +} + +func (b *RequestBodyRef) Bytes() []byte { + if b == nil { + return nil + } + return b.data +} + +func (b *RequestBodyRef) Len() int { + if b == nil { + return 0 + } + return len(b.data) +} + +func (b *RequestBodyRef) Replace(data []byte) { + if b == nil { + return + } + b.data = data +} + // ParsedRequest 保存网关请求的预解析结果 // // 性能优化说明: @@ -64,7 +93,7 @@ type SessionContext struct { // 2. 将解析结果 ParsedRequest 传递给 Service 层 // 3. 避免重复 json.Unmarshal,减少 CPU 和内存开销 type ParsedRequest struct { - Body []byte // 原始请求体(保留用于转发) + Body *RequestBodyRef // 原始请求体引用(保留用于转发) Model string // 请求的模型名称 Stream bool // 是否为流式请求 MetadataUserID string // metadata.user_id(用于会话亲和) @@ -130,17 +159,18 @@ func normalizeSessionUserAgentFallback(raw string) string { // ParseGatewayRequest 解析网关请求体并返回结构化结果。 // protocol 指定请求协议格式(domain.PlatformAnthropic / domain.PlatformGemini), // 不同协议使用不同的 system/messages 字段名。 -func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) { +func ParseGatewayRequest(body *RequestBodyRef, protocol string) (*ParsedRequest, error) { + bodyBytes := body.Bytes() // 保持与旧实现一致:请求体必须是合法 JSON。 // 注意:gjson.GetBytes 对非法 JSON 不会报错,因此需要显式校验。 - if !gjson.ValidBytes(body) { + if !gjson.ValidBytes(bodyBytes) { return nil, fmt.Errorf("invalid json") } // 性能: // - gjson.GetBytes 会把匹配的 Raw/Str 安全复制成 string(对于巨大 messages 会产生额外拷贝)。 // - 这里将 body 通过 unsafe 零拷贝视为 string,仅在本函数内使用,且 body 不会被修改。 - jsonStr := *(*string)(unsafe.Pointer(&body)) + jsonStr := *(*string)(unsafe.Pointer(&bodyBytes)) parsed := &ParsedRequest{ Body: body, @@ -197,7 +227,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) { // Gemini 原生格式: systemInstruction.parts / contents if sysParts := gjson.Get(jsonStr, "systemInstruction.parts"); sysParts.Exists() && sysParts.IsArray() { var parts []any - if err := json.Unmarshal(sliceRawFromBody(body, sysParts), &parts); err != nil { + if err := json.Unmarshal(sliceRawFromBody(bodyBytes, sysParts), &parts); err != nil { return nil, err } parsed.System = parts @@ -205,7 +235,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) { if contents := gjson.Get(jsonStr, "contents"); contents.Exists() && contents.IsArray() { var msgs []any - if err := json.Unmarshal(sliceRawFromBody(body, contents), &msgs); err != nil { + if err := json.Unmarshal(sliceRawFromBody(bodyBytes, contents), &msgs); err != nil { return nil, err } parsed.Messages = msgs @@ -224,7 +254,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) { parsed.System = sys.String() default: var system any - if err := json.Unmarshal(sliceRawFromBody(body, sys), &system); err != nil { + if err := json.Unmarshal(sliceRawFromBody(bodyBytes, sys), &system); err != nil { return nil, err } parsed.System = system @@ -233,7 +263,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) { if msgs := gjson.Get(jsonStr, "messages"); msgs.Exists() && msgs.IsArray() { var messages []any - if err := json.Unmarshal(sliceRawFromBody(body, msgs), &messages); err != nil { + if err := json.Unmarshal(sliceRawFromBody(bodyBytes, msgs), &messages); err != nil { return nil, err } parsed.Messages = messages diff --git a/backend/internal/service/gateway_request_test.go b/backend/internal/service/gateway_request_test.go index 045dc66c..d415b871 100644 --- a/backend/internal/service/gateway_request_test.go +++ b/backend/internal/service/gateway_request_test.go @@ -14,7 +14,7 @@ import ( func TestParseGatewayRequest(t *testing.T) { body := []byte(`{"model":"claude-3-7-sonnet","stream":true,"metadata":{"user_id":"session_123e4567-e89b-12d3-a456-426614174000"},"system":[{"type":"text","text":"hello","cache_control":{"type":"ephemeral"}}],"messages":[{"content":"hi"}]}`) - parsed, err := ParseGatewayRequest(body, "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "") require.NoError(t, err) require.Equal(t, "claude-3-7-sonnet", parsed.Model) require.True(t, parsed.Stream) @@ -27,7 +27,7 @@ func TestParseGatewayRequest(t *testing.T) { func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) { body := []byte(`{"model":"claude-sonnet-4-5","thinking":{"type":"enabled"},"messages":[{"content":"hi"}]}`) - parsed, err := ParseGatewayRequest(body, "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "") require.NoError(t, err) require.Equal(t, "claude-sonnet-4-5", parsed.Model) require.True(t, parsed.ThinkingEnabled) @@ -35,7 +35,7 @@ func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) { func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) { body := []byte(`{"model":"claude-sonnet-4-5","thinking":{"type":"adaptive"},"messages":[{"content":"hi"}]}`) - parsed, err := ParseGatewayRequest(body, "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "") require.NoError(t, err) require.Equal(t, "claude-sonnet-4-5", parsed.Model) require.True(t, parsed.ThinkingEnabled) @@ -43,21 +43,21 @@ func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) { func TestParseGatewayRequest_MaxTokens(t *testing.T) { body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1}`) - parsed, err := ParseGatewayRequest(body, "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "") require.NoError(t, err) require.Equal(t, 1, parsed.MaxTokens) } func TestParseGatewayRequest_MaxTokensNonIntegralIgnored(t *testing.T) { body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1.5}`) - parsed, err := ParseGatewayRequest(body, "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "") require.NoError(t, err) require.Equal(t, 0, parsed.MaxTokens) } func TestParseGatewayRequest_SystemNull(t *testing.T) { body := []byte(`{"model":"claude-3","system":null}`) - parsed, err := ParseGatewayRequest(body, "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "") require.NoError(t, err) // 显式传入 system:null 也应视为“字段已存在”,避免默认 system 被注入。 require.True(t, parsed.HasSystem) @@ -66,13 +66,13 @@ func TestParseGatewayRequest_SystemNull(t *testing.T) { func TestParseGatewayRequest_InvalidModelType(t *testing.T) { body := []byte(`{"model":123}`) - _, err := ParseGatewayRequest(body, "") + _, err := ParseGatewayRequest(NewRequestBodyRef(body), "") require.Error(t, err) } func TestParseGatewayRequest_InvalidStreamType(t *testing.T) { body := []byte(`{"stream":"true"}`) - _, err := ParseGatewayRequest(body, "") + _, err := ParseGatewayRequest(NewRequestBodyRef(body), "") require.Error(t, err) } @@ -86,7 +86,7 @@ func TestParseGatewayRequest_GeminiContents(t *testing.T) { {"role": "user", "parts": [{"text": "How are you?"}]} ] }`) - parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini) require.NoError(t, err) require.Len(t, parsed.Messages, 3, "should parse contents as Messages") require.False(t, parsed.HasSystem, "Gemini format should not set HasSystem") @@ -102,7 +102,7 @@ func TestParseGatewayRequest_GeminiSystemInstruction(t *testing.T) { {"role": "user", "parts": [{"text": "Hello"}]} ] }`) - parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini) require.NoError(t, err) require.NotNil(t, parsed.System, "should parse systemInstruction.parts as System") parts, ok := parsed.System.([]any) @@ -119,7 +119,7 @@ func TestParseGatewayRequest_GeminiWithModel(t *testing.T) { "model": "gemini-2.5-pro", "contents": [{"role": "user", "parts": [{"text": "test"}]}] }`) - parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini) require.NoError(t, err) require.Equal(t, "gemini-2.5-pro", parsed.Model) require.Len(t, parsed.Messages, 1) @@ -132,7 +132,7 @@ func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) { "messages": [{"role": "user", "content": "ignored"}], "contents": [{"role": "user", "parts": [{"text": "real content"}]}] }`) - parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini) require.NoError(t, err) require.False(t, parsed.HasSystem, "Gemini protocol should not parse Anthropic system field") require.Nil(t, parsed.System, "no systemInstruction = nil System") @@ -141,14 +141,14 @@ func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) { func TestParseGatewayRequest_GeminiEmptyContents(t *testing.T) { body := []byte(`{"contents": []}`) - parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini) require.NoError(t, err) require.Empty(t, parsed.Messages) } func TestParseGatewayRequest_GeminiNoContents(t *testing.T) { body := []byte(`{"model": "gemini-2.5-flash"}`) - parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini) require.NoError(t, err) require.Nil(t, parsed.Messages) require.Equal(t, "gemini-2.5-flash", parsed.Model) @@ -162,7 +162,7 @@ func TestParseGatewayRequest_AnthropicIgnoresGeminiFields(t *testing.T) { "contents": [{"role": "user", "parts": [{"text": "ignored"}]}], "systemInstruction": {"parts": [{"text": "ignored"}]} }`) - parsed, err := ParseGatewayRequest(body, domain.PlatformAnthropic) + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic) require.NoError(t, err) require.True(t, parsed.HasSystem) require.Equal(t, "real system", parsed.System) @@ -897,7 +897,7 @@ func TestParseGatewayRequest_TypeValidation(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - _, err := ParseGatewayRequest([]byte(tt.body), "") + _, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "") if tt.wantErr { require.Error(t, err) if tt.errSubstr != "" { @@ -959,7 +959,7 @@ func TestParseGatewayRequest_OptionalFieldsMissing(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - parsed, err := ParseGatewayRequest([]byte(tt.body), "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "") require.NoError(t, err) require.Equal(t, tt.wantModel, parsed.Model) @@ -1023,7 +1023,7 @@ func TestParseGatewayRequest_MaxTokensBoundary(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - parsed, err := ParseGatewayRequest([]byte(tt.body), "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "") if tt.wantErr { require.Error(t, err) return @@ -1040,7 +1040,7 @@ func TestParseGatewayRequest_MaxTokensBoundary(t *testing.T) { // 核心路径:先 Unmarshal 到 map[string]any,再逐字段提取。 func parseGatewayRequestOld(body []byte, protocol string) (*ParsedRequest, error) { parsed := &ParsedRequest{ - Body: body, + Body: NewRequestBodyRef(body), } var req map[string]any @@ -1151,7 +1151,7 @@ func BenchmarkParseGatewayRequest_New_Small(b *testing.B) { b.SetBytes(int64(len(data))) b.ResetTimer() for i := 0; i < b.N; i++ { - _, _ = ParseGatewayRequest(data, "") + _, _ = ParseGatewayRequest(NewRequestBodyRef(data), "") } } @@ -1203,7 +1203,7 @@ func TestParseGatewayRequest_OutputEffort(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - parsed, err := ParseGatewayRequest([]byte(tt.body), "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "") require.NoError(t, err) require.Equal(t, tt.wantEffort, parsed.OutputEffort) }) @@ -1245,6 +1245,6 @@ func BenchmarkParseGatewayRequest_New_Large(b *testing.B) { b.SetBytes(int64(len(data))) b.ResetTimer() for i := 0; i < b.N; i++ { - _, _ = ParseGatewayRequest(data, "") + _, _ = ParseGatewayRequest(NewRequestBodyRef(data), "") } } diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index f807f3ec..8f55bf13 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -4400,12 +4400,12 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A } // Web Search 模拟:纯 web_search 请求时,直接调用搜索 API 构造响应 - if account != nil && s.shouldEmulateWebSearch(ctx, account, parsed.GroupID, parsed.Body) { + if account != nil && s.shouldEmulateWebSearch(ctx, account, parsed.GroupID, parsed.Body.Bytes()) { return s.handleWebSearchEmulation(ctx, c, account, parsed) } if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() { - passthroughBody := parsed.Body + passthroughBody := parsed.Body.Bytes() passthroughModel := parsed.Model if passthroughModel != "" { if mappedModel := account.GetMappedModel(passthroughModel); mappedModel != passthroughModel { @@ -4441,7 +4441,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A c.Set(betaPolicyFilterSetKey, filterSet) } - body := parsed.Body + body := parsed.Body.Bytes() reqModel := parsed.Model reqStream := parsed.Stream originalModel := reqModel @@ -5735,7 +5735,7 @@ func (s *GatewayService) forwardBedrock( ) (*ForwardResult, error) { reqModel := parsed.Model reqStream := parsed.Stream - body := parsed.Body + body := parsed.Body.Bytes() region := bedrockRuntimeRegion(account) mappedModel, ok := ResolveBedrockModelID(account, reqModel) @@ -9172,7 +9172,7 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context, } if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() { - passthroughBody := parsed.Body + passthroughBody := parsed.Body.Bytes() if reqModel := parsed.Model; reqModel != "" { if mappedModel := account.GetMappedModel(reqModel); mappedModel != reqModel { passthroughBody = s.replaceModelInBody(passthroughBody, mappedModel) @@ -9188,7 +9188,7 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context, return nil } - body := parsed.Body + body := parsed.Body.Bytes() reqModel := parsed.Model // Pre-filter: strip empty text blocks to prevent upstream 400. diff --git a/backend/internal/service/gateway_service_benchmark_test.go b/backend/internal/service/gateway_service_benchmark_test.go index c9c4d3dd..5637680b 100644 --- a/backend/internal/service/gateway_service_benchmark_test.go +++ b/backend/internal/service/gateway_service_benchmark_test.go @@ -14,7 +14,7 @@ func BenchmarkGenerateSessionHash_Metadata(b *testing.B) { b.ReportAllocs() for i := 0; i < b.N; i++ { - parsed, err := ParseGatewayRequest(body, "") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "") if err != nil { b.Fatalf("解析请求失败: %v", err) } diff --git a/backend/internal/service/gateway_websearch_emulation.go b/backend/internal/service/gateway_websearch_emulation.go index a42b5585..2f9c8e0c 100644 --- a/backend/internal/service/gateway_websearch_emulation.go +++ b/backend/internal/service/gateway_websearch_emulation.go @@ -150,7 +150,7 @@ func (s *GatewayService) handleWebSearchEmulation( parsed.OnUpstreamAccepted() } - query := extractSearchQueryFromBody(parsed.Body) + query := extractSearchQueryFromBody(parsed.Body.Bytes()) if query == "" { return nil, fmt.Errorf("web search emulation: no query found in messages") } diff --git a/backend/internal/service/generate_session_hash_test.go b/backend/internal/service/generate_session_hash_test.go index 39679c3d..8f3258b7 100644 --- a/backend/internal/service/generate_session_hash_test.go +++ b/backend/internal/service/generate_session_hash_test.go @@ -1198,7 +1198,7 @@ func TestGenerateSessionHash_GeminiMultiTurnHashNotSticky(t *testing.T) { hashes := make([]string, 3) for i, body := range [][]byte{round1Body, round2Body, round3Body} { - parsed, err := ParseGatewayRequest(body, "gemini") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini") require.NoError(t, err) parsed.SessionContext = ctx hashes[i] = svc.GenerateSessionHash(parsed) @@ -1211,7 +1211,7 @@ func TestGenerateSessionHash_GeminiMultiTurnHashNotSticky(t *testing.T) { require.NotEqual(t, hashes[0], hashes[2], "round 1 vs 3 hash should differ") // 同一轮重试应产生相同 hash - parsed1Again, err := ParseGatewayRequest(round2Body, "gemini") + parsed1Again, err := ParseGatewayRequest(NewRequestBodyRef(round2Body), "gemini") require.NoError(t, err) parsed1Again.SessionContext = ctx h2Again := svc.GenerateSessionHash(parsed1Again) @@ -1234,7 +1234,7 @@ func TestGenerateSessionHash_GeminiEndToEnd(t *testing.T) { ] }`) - parsed, err := ParseGatewayRequest(body, "gemini") + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini") require.NoError(t, err) parsed.SessionContext = &SessionContext{ ClientIP: "10.0.0.1", @@ -1246,7 +1246,7 @@ func TestGenerateSessionHash_GeminiEndToEnd(t *testing.T) { require.NotEmpty(t, h, "end-to-end Gemini flow should produce a hash") // 同一请求再次解析应产生相同 hash - parsed2, err := ParseGatewayRequest(body, "gemini") + parsed2, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini") require.NoError(t, err) parsed2.SessionContext = &SessionContext{ ClientIP: "10.0.0.1", @@ -1258,7 +1258,7 @@ func TestGenerateSessionHash_GeminiEndToEnd(t *testing.T) { require.Equal(t, h, h2, "same request should produce same hash") // 不同用户发送相同请求应产生不同 hash - parsed3, err := ParseGatewayRequest(body, "gemini") + parsed3, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini") require.NoError(t, err) parsed3.SessionContext = &SessionContext{ ClientIP: "10.0.0.2",