refactor(gateway): introduce request body refs

This commit is contained in:
name 2026-05-29 21:05:47 +08:00
parent f18451e56f
commit d8cbf9ab5c
11 changed files with 95 additions and 61 deletions

View File

@ -154,7 +154,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
setOpsRequestContext(c, "", false) setOpsRequestContext(c, "", false)
parsedReq, err := service.ParseGatewayRequest(body, domain.PlatformAnthropic) bodyRef := service.NewRequestBodyRef(body)
parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic)
if err != nil { if err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return return
@ -746,11 +747,11 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 应用渠道模型映射到请求 // 应用渠道模型映射到请求
if channelMapping.Mapped { if channelMapping.Mapped {
parsedReq.Model = channelMapping.MappedModel 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 必需字段 // Bedrock CC 兼容:渠道模型映射后,清理 Anthropic API 专有字段、注入 Bedrock 必需字段
parsedReq.Body = h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body, parsedReq.Model, account, apiKey.GroupID) parsedReq.Body.Replace(h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body.Bytes(), parsedReq.Model, account, apiKey.GroupID))
body = parsedReq.Body body = parsedReq.Body.Bytes()
// 转发请求 - 根据账号平台分流 // 转发请求 - 根据账号平台分流
c.Set("parsed_request", parsedReq) c.Set("parsed_request", parsedReq)
@ -1683,7 +1684,8 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) {
setOpsRequestContext(c, "", false) setOpsRequestContext(c, "", false)
parsedReq, err := service.ParseGatewayRequest(body, domain.PlatformAnthropic) bodyRef := service.NewRequestBodyRef(body)
parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic)
if err != nil { if err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return return

View File

@ -151,9 +151,10 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
} }
// Parse request for session hash // Parse request for session hash
parsedReq, _ := service.ParseGatewayRequest(body, "chat_completions") bodyRef := service.NewRequestBodyRef(body)
parsedReq, _ := service.ParseGatewayRequest(bodyRef, "chat_completions")
if parsedReq == nil { 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{ parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c), ClientIP: ip.GetClientIP(c),

View File

@ -156,9 +156,10 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
} }
// Parse request for session hash // Parse request for session hash
parsedReq, _ := service.ParseGatewayRequest(body, "responses") bodyRef := service.NewRequestBodyRef(body)
parsedReq, _ := service.ParseGatewayRequest(bodyRef, "responses")
if parsedReq == nil { 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{ parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c), ClientIP: ip.GetClientIP(c),

View File

@ -262,7 +262,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
sessionHash := extractGeminiCLISessionHash(c, body) sessionHash := extractGeminiCLISessionHash(c, body)
if sessionHash == "" { if sessionHash == "" {
// Fallback: 使用通用的会话哈希生成逻辑(适用于其他客户端) // Fallback: 使用通用的会话哈希生成逻辑(适用于其他客户端)
parsedReq, _ := service.ParseGatewayRequest(body, domain.PlatformGemini) parsedReq, _ := service.ParseGatewayRequest(service.NewRequestBodyRef(body), domain.PlatformGemini)
if parsedReq != nil { if parsedReq != nil {
parsedReq.SessionContext = &service.SessionContext{ parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c), ClientIP: ip.GetClientIP(c),

View File

@ -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"}]}]}`) 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{ parsed := &ParsedRequest{
Body: body, Body: NewRequestBodyRef(body),
Model: "claude-3-7-sonnet-20250219", Model: "claude-3-7-sonnet-20250219",
Stream: true, 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"}}`) body := []byte(`{"model":"claude-3-5-sonnet-latest","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}],"thinking":{"type":"enabled"}}`)
parsed := &ParsedRequest{ parsed := &ParsedRequest{
Body: body, Body: NewRequestBodyRef(body),
Model: "claude-3-5-sonnet-latest", 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"}]}]}`) body := []byte(`{"model":"` + tt.model + `","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
parsed := &ParsedRequest{ parsed := &ParsedRequest{
Body: body, Body: NewRequestBodyRef(body),
Model: tt.model, Model: tt.model,
} }
@ -429,7 +429,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingPreservesOtherFie
// 包含复杂字段的请求体:system、thinking、messages // 包含复杂字段的请求体: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}`) 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{ parsed := &ParsedRequest{
Body: body, Body: NewRequestBodyRef(body),
Model: "claude-sonnet-4-20250514", 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}}`) 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{ parsed := &ParsedRequest{
Body: body, Body: NewRequestBodyRef(body),
Model: "claude-sonnet-4-20250514", Model: "claude-sonnet-4-20250514",
} }
@ -547,7 +547,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_EmptyModelSkipsMapping(t *tes
body := []byte(`{"messages":[{"role":"user","content":"hello"}]}`) body := []byte(`{"messages":[{"role":"user","content":"hello"}]}`)
parsed := &ParsedRequest{ parsed := &ParsedRequest{
Body: body, Body: NewRequestBodyRef(body),
Model: "", // 空模型 Model: "", // 空模型
} }
@ -636,7 +636,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_CountTokens404PassthroughNotE
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil) c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil)
body := []byte(`{"model":"claude-sonnet-4-5-20250929","messages":[{"role":"user","content":"hi"}]}`) 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{ upstream := &anthropicHTTPUpstreamRecorder{
resp: &http.Response{ resp: &http.Response{
@ -767,7 +767,7 @@ func TestGatewayService_AnthropicOAuth_ForwardPreservesBillingHeaderSystemBlock(
c, _ := gin.CreateTestContext(rec) c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) 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) require.NoError(t, err)
upstream := &anthropicHTTPUpstreamRecorder{ upstream := &anthropicHTTPUpstreamRecorder{

View File

@ -51,6 +51,35 @@ type SessionContext struct {
APIKeyID int64 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 保存网关请求的预解析结果 // ParsedRequest 保存网关请求的预解析结果
// //
// 性能优化说明: // 性能优化说明:
@ -64,7 +93,7 @@ type SessionContext struct {
// 2. 将解析结果 ParsedRequest 传递给 Service 层 // 2. 将解析结果 ParsedRequest 传递给 Service 层
// 3. 避免重复 json.Unmarshal,减少 CPU 和内存开销 // 3. 避免重复 json.Unmarshal,减少 CPU 和内存开销
type ParsedRequest struct { type ParsedRequest struct {
Body []byte // 原始请求体(保留用于转发) Body *RequestBodyRef // 原始请求体引用(保留用于转发)
Model string // 请求的模型名称 Model string // 请求的模型名称
Stream bool // 是否为流式请求 Stream bool // 是否为流式请求
MetadataUserID string // metadata.user_id(用于会话亲和) MetadataUserID string // metadata.user_id(用于会话亲和)
@ -130,17 +159,18 @@ func normalizeSessionUserAgentFallback(raw string) string {
// ParseGatewayRequest 解析网关请求体并返回结构化结果。 // ParseGatewayRequest 解析网关请求体并返回结构化结果。
// protocol 指定请求协议格式(domain.PlatformAnthropic / domain.PlatformGemini), // protocol 指定请求协议格式(domain.PlatformAnthropic / domain.PlatformGemini),
// 不同协议使用不同的 system/messages 字段名。 // 不同协议使用不同的 system/messages 字段名。
func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) { func ParseGatewayRequest(body *RequestBodyRef, protocol string) (*ParsedRequest, error) {
bodyBytes := body.Bytes()
// 保持与旧实现一致:请求体必须是合法 JSON。 // 保持与旧实现一致:请求体必须是合法 JSON。
// 注意:gjson.GetBytes 对非法 JSON 不会报错,因此需要显式校验。 // 注意:gjson.GetBytes 对非法 JSON 不会报错,因此需要显式校验。
if !gjson.ValidBytes(body) { if !gjson.ValidBytes(bodyBytes) {
return nil, fmt.Errorf("invalid json") return nil, fmt.Errorf("invalid json")
} }
// 性能: // 性能:
// - gjson.GetBytes 会把匹配的 Raw/Str 安全复制成 string(对于巨大 messages 会产生额外拷贝)。 // - gjson.GetBytes 会把匹配的 Raw/Str 安全复制成 string(对于巨大 messages 会产生额外拷贝)。
// - 这里将 body 通过 unsafe 零拷贝视为 string,仅在本函数内使用,且 body 不会被修改。 // - 这里将 body 通过 unsafe 零拷贝视为 string,仅在本函数内使用,且 body 不会被修改。
jsonStr := *(*string)(unsafe.Pointer(&body)) jsonStr := *(*string)(unsafe.Pointer(&bodyBytes))
parsed := &ParsedRequest{ parsed := &ParsedRequest{
Body: body, Body: body,
@ -197,7 +227,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) {
// Gemini 原生格式: systemInstruction.parts / contents // Gemini 原生格式: systemInstruction.parts / contents
if sysParts := gjson.Get(jsonStr, "systemInstruction.parts"); sysParts.Exists() && sysParts.IsArray() { if sysParts := gjson.Get(jsonStr, "systemInstruction.parts"); sysParts.Exists() && sysParts.IsArray() {
var parts []any 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 return nil, err
} }
parsed.System = parts 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() { if contents := gjson.Get(jsonStr, "contents"); contents.Exists() && contents.IsArray() {
var msgs []any 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 return nil, err
} }
parsed.Messages = msgs parsed.Messages = msgs
@ -224,7 +254,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) {
parsed.System = sys.String() parsed.System = sys.String()
default: default:
var system any 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 return nil, err
} }
parsed.System = system 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() { if msgs := gjson.Get(jsonStr, "messages"); msgs.Exists() && msgs.IsArray() {
var messages []any 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 return nil, err
} }
parsed.Messages = messages parsed.Messages = messages

View File

@ -14,7 +14,7 @@ import (
func TestParseGatewayRequest(t *testing.T) { 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"}]}`) 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.NoError(t, err)
require.Equal(t, "claude-3-7-sonnet", parsed.Model) require.Equal(t, "claude-3-7-sonnet", parsed.Model)
require.True(t, parsed.Stream) require.True(t, parsed.Stream)
@ -27,7 +27,7 @@ func TestParseGatewayRequest(t *testing.T) {
func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) { func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) {
body := []byte(`{"model":"claude-sonnet-4-5","thinking":{"type":"enabled"},"messages":[{"content":"hi"}]}`) 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.NoError(t, err)
require.Equal(t, "claude-sonnet-4-5", parsed.Model) require.Equal(t, "claude-sonnet-4-5", parsed.Model)
require.True(t, parsed.ThinkingEnabled) require.True(t, parsed.ThinkingEnabled)
@ -35,7 +35,7 @@ func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) {
func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) { func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) {
body := []byte(`{"model":"claude-sonnet-4-5","thinking":{"type":"adaptive"},"messages":[{"content":"hi"}]}`) 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.NoError(t, err)
require.Equal(t, "claude-sonnet-4-5", parsed.Model) require.Equal(t, "claude-sonnet-4-5", parsed.Model)
require.True(t, parsed.ThinkingEnabled) require.True(t, parsed.ThinkingEnabled)
@ -43,21 +43,21 @@ func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) {
func TestParseGatewayRequest_MaxTokens(t *testing.T) { func TestParseGatewayRequest_MaxTokens(t *testing.T) {
body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1}`) body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1}`)
parsed, err := ParseGatewayRequest(body, "") parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, 1, parsed.MaxTokens) require.Equal(t, 1, parsed.MaxTokens)
} }
func TestParseGatewayRequest_MaxTokensNonIntegralIgnored(t *testing.T) { func TestParseGatewayRequest_MaxTokensNonIntegralIgnored(t *testing.T) {
body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1.5}`) 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.NoError(t, err)
require.Equal(t, 0, parsed.MaxTokens) require.Equal(t, 0, parsed.MaxTokens)
} }
func TestParseGatewayRequest_SystemNull(t *testing.T) { func TestParseGatewayRequest_SystemNull(t *testing.T) {
body := []byte(`{"model":"claude-3","system":null}`) body := []byte(`{"model":"claude-3","system":null}`)
parsed, err := ParseGatewayRequest(body, "") parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err) require.NoError(t, err)
// 显式传入 system:null 也应视为“字段已存在”,避免默认 system 被注入。 // 显式传入 system:null 也应视为“字段已存在”,避免默认 system 被注入。
require.True(t, parsed.HasSystem) require.True(t, parsed.HasSystem)
@ -66,13 +66,13 @@ func TestParseGatewayRequest_SystemNull(t *testing.T) {
func TestParseGatewayRequest_InvalidModelType(t *testing.T) { func TestParseGatewayRequest_InvalidModelType(t *testing.T) {
body := []byte(`{"model":123}`) body := []byte(`{"model":123}`)
_, err := ParseGatewayRequest(body, "") _, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.Error(t, err) require.Error(t, err)
} }
func TestParseGatewayRequest_InvalidStreamType(t *testing.T) { func TestParseGatewayRequest_InvalidStreamType(t *testing.T) {
body := []byte(`{"stream":"true"}`) body := []byte(`{"stream":"true"}`)
_, err := ParseGatewayRequest(body, "") _, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.Error(t, err) require.Error(t, err)
} }
@ -86,7 +86,7 @@ func TestParseGatewayRequest_GeminiContents(t *testing.T) {
{"role": "user", "parts": [{"text": "How are you?"}]} {"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.NoError(t, err)
require.Len(t, parsed.Messages, 3, "should parse contents as Messages") require.Len(t, parsed.Messages, 3, "should parse contents as Messages")
require.False(t, parsed.HasSystem, "Gemini format should not set HasSystem") 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"}]} {"role": "user", "parts": [{"text": "Hello"}]}
] ]
}`) }`)
parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err) require.NoError(t, err)
require.NotNil(t, parsed.System, "should parse systemInstruction.parts as System") require.NotNil(t, parsed.System, "should parse systemInstruction.parts as System")
parts, ok := parsed.System.([]any) parts, ok := parsed.System.([]any)
@ -119,7 +119,7 @@ func TestParseGatewayRequest_GeminiWithModel(t *testing.T) {
"model": "gemini-2.5-pro", "model": "gemini-2.5-pro",
"contents": [{"role": "user", "parts": [{"text": "test"}]}] "contents": [{"role": "user", "parts": [{"text": "test"}]}]
}`) }`)
parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, "gemini-2.5-pro", parsed.Model) require.Equal(t, "gemini-2.5-pro", parsed.Model)
require.Len(t, parsed.Messages, 1) require.Len(t, parsed.Messages, 1)
@ -132,7 +132,7 @@ func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) {
"messages": [{"role": "user", "content": "ignored"}], "messages": [{"role": "user", "content": "ignored"}],
"contents": [{"role": "user", "parts": [{"text": "real content"}]}] "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.NoError(t, err)
require.False(t, parsed.HasSystem, "Gemini protocol should not parse Anthropic system field") require.False(t, parsed.HasSystem, "Gemini protocol should not parse Anthropic system field")
require.Nil(t, parsed.System, "no systemInstruction = nil System") 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) { func TestParseGatewayRequest_GeminiEmptyContents(t *testing.T) {
body := []byte(`{"contents": []}`) body := []byte(`{"contents": []}`)
parsed, err := ParseGatewayRequest(body, domain.PlatformGemini) parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err) require.NoError(t, err)
require.Empty(t, parsed.Messages) require.Empty(t, parsed.Messages)
} }
func TestParseGatewayRequest_GeminiNoContents(t *testing.T) { func TestParseGatewayRequest_GeminiNoContents(t *testing.T) {
body := []byte(`{"model": "gemini-2.5-flash"}`) 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.NoError(t, err)
require.Nil(t, parsed.Messages) require.Nil(t, parsed.Messages)
require.Equal(t, "gemini-2.5-flash", parsed.Model) 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"}]}], "contents": [{"role": "user", "parts": [{"text": "ignored"}]}],
"systemInstruction": {"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.NoError(t, err)
require.True(t, parsed.HasSystem) require.True(t, parsed.HasSystem)
require.Equal(t, "real system", parsed.System) require.Equal(t, "real system", parsed.System)
@ -897,7 +897,7 @@ func TestParseGatewayRequest_TypeValidation(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
_, err := ParseGatewayRequest([]byte(tt.body), "") _, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
if tt.wantErr { if tt.wantErr {
require.Error(t, err) require.Error(t, err)
if tt.errSubstr != "" { if tt.errSubstr != "" {
@ -959,7 +959,7 @@ func TestParseGatewayRequest_OptionalFieldsMissing(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { 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.NoError(t, err)
require.Equal(t, tt.wantModel, parsed.Model) require.Equal(t, tt.wantModel, parsed.Model)
@ -1023,7 +1023,7 @@ func TestParseGatewayRequest_MaxTokensBoundary(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
parsed, err := ParseGatewayRequest([]byte(tt.body), "") parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
if tt.wantErr { if tt.wantErr {
require.Error(t, err) require.Error(t, err)
return return
@ -1040,7 +1040,7 @@ func TestParseGatewayRequest_MaxTokensBoundary(t *testing.T) {
// 核心路径:先 Unmarshal 到 map[string]any,再逐字段提取。 // 核心路径:先 Unmarshal 到 map[string]any,再逐字段提取。
func parseGatewayRequestOld(body []byte, protocol string) (*ParsedRequest, error) { func parseGatewayRequestOld(body []byte, protocol string) (*ParsedRequest, error) {
parsed := &ParsedRequest{ parsed := &ParsedRequest{
Body: body, Body: NewRequestBodyRef(body),
} }
var req map[string]any var req map[string]any
@ -1151,7 +1151,7 @@ func BenchmarkParseGatewayRequest_New_Small(b *testing.B) {
b.SetBytes(int64(len(data))) b.SetBytes(int64(len(data)))
b.ResetTimer() b.ResetTimer()
for i := 0; i < b.N; i++ { 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 { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { 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.NoError(t, err)
require.Equal(t, tt.wantEffort, parsed.OutputEffort) require.Equal(t, tt.wantEffort, parsed.OutputEffort)
}) })
@ -1245,6 +1245,6 @@ func BenchmarkParseGatewayRequest_New_Large(b *testing.B) {
b.SetBytes(int64(len(data))) b.SetBytes(int64(len(data)))
b.ResetTimer() b.ResetTimer()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
_, _ = ParseGatewayRequest(data, "") _, _ = ParseGatewayRequest(NewRequestBodyRef(data), "")
} }
} }

View File

@ -4400,12 +4400,12 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
} }
// Web Search 模拟:纯 web_search 请求时,直接调用搜索 API 构造响应 // 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) return s.handleWebSearchEmulation(ctx, c, account, parsed)
} }
if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() { if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() {
passthroughBody := parsed.Body passthroughBody := parsed.Body.Bytes()
passthroughModel := parsed.Model passthroughModel := parsed.Model
if passthroughModel != "" { if passthroughModel != "" {
if mappedModel := account.GetMappedModel(passthroughModel); mappedModel != 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) c.Set(betaPolicyFilterSetKey, filterSet)
} }
body := parsed.Body body := parsed.Body.Bytes()
reqModel := parsed.Model reqModel := parsed.Model
reqStream := parsed.Stream reqStream := parsed.Stream
originalModel := reqModel originalModel := reqModel
@ -5735,7 +5735,7 @@ func (s *GatewayService) forwardBedrock(
) (*ForwardResult, error) { ) (*ForwardResult, error) {
reqModel := parsed.Model reqModel := parsed.Model
reqStream := parsed.Stream reqStream := parsed.Stream
body := parsed.Body body := parsed.Body.Bytes()
region := bedrockRuntimeRegion(account) region := bedrockRuntimeRegion(account)
mappedModel, ok := ResolveBedrockModelID(account, reqModel) mappedModel, ok := ResolveBedrockModelID(account, reqModel)
@ -9172,7 +9172,7 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
} }
if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() { if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() {
passthroughBody := parsed.Body passthroughBody := parsed.Body.Bytes()
if reqModel := parsed.Model; reqModel != "" { if reqModel := parsed.Model; reqModel != "" {
if mappedModel := account.GetMappedModel(reqModel); mappedModel != reqModel { if mappedModel := account.GetMappedModel(reqModel); mappedModel != reqModel {
passthroughBody = s.replaceModelInBody(passthroughBody, mappedModel) passthroughBody = s.replaceModelInBody(passthroughBody, mappedModel)
@ -9188,7 +9188,7 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
return nil return nil
} }
body := parsed.Body body := parsed.Body.Bytes()
reqModel := parsed.Model reqModel := parsed.Model
// Pre-filter: strip empty text blocks to prevent upstream 400. // Pre-filter: strip empty text blocks to prevent upstream 400.

View File

@ -14,7 +14,7 @@ func BenchmarkGenerateSessionHash_Metadata(b *testing.B) {
b.ReportAllocs() b.ReportAllocs()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
parsed, err := ParseGatewayRequest(body, "") parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
if err != nil { if err != nil {
b.Fatalf("解析请求失败: %v", err) b.Fatalf("解析请求失败: %v", err)
} }

View File

@ -150,7 +150,7 @@ func (s *GatewayService) handleWebSearchEmulation(
parsed.OnUpstreamAccepted() parsed.OnUpstreamAccepted()
} }
query := extractSearchQueryFromBody(parsed.Body) query := extractSearchQueryFromBody(parsed.Body.Bytes())
if query == "" { if query == "" {
return nil, fmt.Errorf("web search emulation: no query found in messages") return nil, fmt.Errorf("web search emulation: no query found in messages")
} }

View File

@ -1198,7 +1198,7 @@ func TestGenerateSessionHash_GeminiMultiTurnHashNotSticky(t *testing.T) {
hashes := make([]string, 3) hashes := make([]string, 3)
for i, body := range [][]byte{round1Body, round2Body, round3Body} { for i, body := range [][]byte{round1Body, round2Body, round3Body} {
parsed, err := ParseGatewayRequest(body, "gemini") parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
require.NoError(t, err) require.NoError(t, err)
parsed.SessionContext = ctx parsed.SessionContext = ctx
hashes[i] = svc.GenerateSessionHash(parsed) 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") require.NotEqual(t, hashes[0], hashes[2], "round 1 vs 3 hash should differ")
// 同一轮重试应产生相同 hash // 同一轮重试应产生相同 hash
parsed1Again, err := ParseGatewayRequest(round2Body, "gemini") parsed1Again, err := ParseGatewayRequest(NewRequestBodyRef(round2Body), "gemini")
require.NoError(t, err) require.NoError(t, err)
parsed1Again.SessionContext = ctx parsed1Again.SessionContext = ctx
h2Again := svc.GenerateSessionHash(parsed1Again) 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) require.NoError(t, err)
parsed.SessionContext = &SessionContext{ parsed.SessionContext = &SessionContext{
ClientIP: "10.0.0.1", 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") require.NotEmpty(t, h, "end-to-end Gemini flow should produce a hash")
// 同一请求再次解析应产生相同 hash // 同一请求再次解析应产生相同 hash
parsed2, err := ParseGatewayRequest(body, "gemini") parsed2, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
require.NoError(t, err) require.NoError(t, err)
parsed2.SessionContext = &SessionContext{ parsed2.SessionContext = &SessionContext{
ClientIP: "10.0.0.1", 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") require.Equal(t, h, h2, "same request should produce same hash")
// 不同用户发送相同请求应产生不同 hash // 不同用户发送相同请求应产生不同 hash
parsed3, err := ParseGatewayRequest(body, "gemini") parsed3, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
require.NoError(t, err) require.NoError(t, err)
parsed3.SessionContext = &SessionContext{ parsed3.SessionContext = &SessionContext{
ClientIP: "10.0.0.2", ClientIP: "10.0.0.2",