diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go index dc59611cea..59442fcbca 100644 --- a/backend/internal/service/gateway_request.go +++ b/backend/internal/service/gateway_request.go @@ -119,6 +119,7 @@ func clearGatewayRequestDerivedState(parsed *ParsedRequest) { parsed.MaxTokens = 0 parsed.systemRange = missingJSONRange() parsed.messagesRange = missingJSONRange() + parsed.inputRange = missingJSONRange() } func clearGatewayRequestRanges(parsed *ParsedRequest) { @@ -128,6 +129,7 @@ func clearGatewayRequestRanges(parsed *ParsedRequest) { parsed.HasSystem = false parsed.systemRange = missingJSONRange() parsed.messagesRange = missingJSONRange() + parsed.inputRange = missingJSONRange() } func setGatewayRequestRanges(parsed *ParsedRequest, protocol string, jsonStr string) { @@ -150,6 +152,11 @@ func setGatewayRequestRanges(parsed *ParsedRequest, protocol string, jsonStr str if msgs := gjson.Get(jsonStr, "messages"); msgs.Exists() && msgs.IsArray() { parsed.messagesRange = rangeFromResult(msgs) } + if protocol == "responses" { + if input := gjson.Get(jsonStr, "input"); input.Exists() { + parsed.inputRange = rangeFromResult(input) + } + } } } @@ -235,6 +242,7 @@ type ParsedRequest struct { protocol string // 当前 Body 的协议格式,用于 Body 替换后刷新 raw range systemRange jsonRange // system/systemInstruction.parts 的 raw JSON 范围,绑定 Body 当前内容 messagesRange jsonRange // messages/contents 的 raw JSON 范围,绑定 Body 当前内容 + inputRange jsonRange // Responses API input 的 raw JSON 范围,绑定 Body 当前内容 // GroupID 请求所属分组 ID(来自 API Key) GroupID *int64 @@ -317,6 +325,10 @@ func (p *ParsedRequest) MessagesRaw() []byte { return p.raw(p.messagesRange) } +func (p *ParsedRequest) InputRaw() []byte { + return p.raw(p.inputRange) +} + func (p *ParsedRequest) DecodeSystem(dst any) error { raw := p.SystemRaw() if len(raw) == 0 { diff --git a/backend/internal/service/gateway_request_test.go b/backend/internal/service/gateway_request_test.go index 288c031c37..7a1b6ef948 100644 --- a/backend/internal/service/gateway_request_test.go +++ b/backend/internal/service/gateway_request_test.go @@ -77,6 +77,15 @@ func TestParseGatewayRequest_InvalidStreamType(t *testing.T) { require.Error(t, err) } +func TestParseGatewayRequest_ResponsesInput(t *testing.T) { + body := []byte(`{"model":"gpt-5.1","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"hello"}]}]}`) + parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "responses") + require.NoError(t, err) + require.NotEmpty(t, parsed.InputRaw()) + require.Nil(t, parsed.MessagesRaw()) + require.Equal(t, "hello", gjson.ParseBytes(parsed.InputRaw()).Get("0.content.0.text").String()) +} + // ============ Gemini 原生格式解析测试 ============ func TestParseGatewayRequest_GeminiContents(t *testing.T) { diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 531c93e4e7..a1c9790262 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -770,7 +770,11 @@ func (s *GatewayService) GenerateSessionHash(parsed *ParsedRequest) string { if systemText := extractTextFromSystemRaw(parsed.SystemRaw()); systemText != "" { _, _ = combined.WriteString(systemText) } + contentStart := combined.Len() appendMessageTextsFromRaw(&combined, parsed.MessagesRaw()) + if combined.Len() == contentStart { + appendResponsesSessionAnchorFromRaw(&combined, parsed.InputRaw()) + } if combined.Len() > 0 { hash := s.hashContent(combined.String()) slog.Info("sticky.hash_source", @@ -929,6 +933,65 @@ func appendMessageTextsFromRaw(builder *strings.Builder, raw []byte) { }) } +func appendResponsesSessionAnchorFromRaw(builder *strings.Builder, raw []byte) { + if builder == nil || len(raw) == 0 { + return + } + input := parseRawJSONView(raw) + if input.Type == gjson.String { + _, _ = builder.WriteString(input.String()) + return + } + if !input.IsArray() { + return + } + + input.ForEach(func(_, item gjson.Result) bool { + if item.Type == gjson.String { + _, _ = builder.WriteString(item.String()) + return false + } + + switch item.Get("role").String() { + case "system", "developer": + appendResponsesContentText(builder, item.Get("content")) + case "user": + appendResponsesContentText(builder, item.Get("content")) + return false + default: + if item.Get("type").String() == "input_text" { + if text := item.Get("text").String(); text != "" { + _, _ = builder.WriteString(text) + } + return false + } + } + return true + }) +} + +func appendResponsesContentText(builder *strings.Builder, content gjson.Result) { + if builder == nil || !content.Exists() { + return + } + if content.Type == gjson.String { + _, _ = builder.WriteString(content.String()) + return + } + if !content.IsArray() { + return + } + content.ForEach(func(_, part gjson.Result) bool { + switch part.Get("type").String() { + case "input_text", "text": + if text := part.Get("text").String(); text != "" { + _, _ = builder.WriteString(text) + } + } + return true + }) +} + func extractCacheableTextFromSystemRaw(raw []byte) string { system := parseRawJSONView(raw) if !system.IsArray() { diff --git a/backend/internal/service/generate_session_hash_test.go b/backend/internal/service/generate_session_hash_test.go index 5ed3f0ae99..135c71fa86 100644 --- a/backend/internal/service/generate_session_hash_test.go +++ b/backend/internal/service/generate_session_hash_test.go @@ -26,6 +26,14 @@ func mustParseGeminiSessionHashRequest(t *testing.T, body string, ctx *SessionCo return parsed } +func mustParseResponsesSessionHashRequest(t *testing.T, body string, ctx *SessionContext) *ParsedRequest { + t.Helper() + parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(body)), "responses") + require.NoError(t, err) + parsed.SessionContext = ctx + return parsed +} + func anthropicSessionBody(system any, messages []any, metadataUserID string) string { body := map[string]any{} if system != nil { @@ -217,6 +225,60 @@ func TestGenerateSessionHash_ContinuousConversation_SameRoundSameHash(t *testing require.Equal(t, h1, h2, "same conversation state should produce identical hash on retry") } +func TestGenerateSessionHash_ResponsesDifferentInputProducesDifferentHash(t *testing.T) { + svc := &GatewayService{} + ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "codex_cli_rs/0.1.0", APIKeyID: 1} + first := mustParseResponsesSessionHashRequest(t, `{"input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"help me with Go"}]}]}`, ctx) + second := mustParseResponsesSessionHashRequest(t, `{"input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"help me with Python"}]}]}`, ctx) + + h1 := svc.GenerateSessionHash(first) + h2 := svc.GenerateSessionHash(second) + require.NotEmpty(t, h1) + require.NotEmpty(t, h2) + require.NotEqual(t, h1, h2, "different Responses input should produce different hashes for the same client") +} + +func TestGenerateSessionHash_ResponsesGrowingInputKeepsStableHash(t *testing.T) { + svc := &GatewayService{} + ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "codex_cli_rs/0.1.0", APIKeyID: 1} + round1 := mustParseResponsesSessionHashRequest(t, `{"input":[{"type":"message","role":"developer","content":[{"type":"input_text","text":"Be concise."}]},{"type":"message","role":"user","content":[{"type":"input_text","text":"help me with Go"}]}]}`, ctx) + round2 := mustParseResponsesSessionHashRequest(t, `{"input":[{"type":"message","role":"developer","content":[{"type":"input_text","text":"Be concise."}]},{"type":"message","role":"user","content":[{"type":"input_text","text":"help me with Go"}]},{"type":"message","role":"assistant","content":[{"type":"output_text","text":"Sure."}]},{"type":"message","role":"user","content":[{"type":"input_text","text":"add tests"}]}]}`, ctx) + + h1 := svc.GenerateSessionHash(round1) + h2 := svc.GenerateSessionHash(round2) + require.NotEmpty(t, h1) + require.Equal(t, h1, h2, "Responses input growth should preserve the hash when the conversation prefix is stable") +} + +func TestGenerateSessionHash_MessagesPathIgnoresResponsesInput(t *testing.T) { + svc := &GatewayService{} + ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1} + first := mustParseResponsesSessionHashRequest(t, `{"messages":[{"role":"user","content":"hello"}],"input":"first"}`, ctx) + second := mustParseResponsesSessionHashRequest(t, `{"messages":[{"role":"user","content":"hello"}],"input":"second"}`, ctx) + + h1 := svc.GenerateSessionHash(first) + h2 := svc.GenerateSessionHash(second) + require.Equal(t, h1, h2, "existing messages fallback should remain authoritative when messages contain text") +} + +func TestGenerateSessionHash_ResponsesInputDoesNotOverrideHigherPrioritySources(t *testing.T) { + svc := &GatewayService{} + ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1} + + t.Run("metadata user id", func(t *testing.T) { + metadata := "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000" + parsed := mustParseResponsesSessionHashRequest(t, `{"metadata":{"user_id":"`+metadata+`"},"input":"hello"}`, ctx) + require.Equal(t, "123e4567-e89b-12d3-a456-426614174000", svc.GenerateSessionHash(parsed)) + }) + + t.Run("cache control", func(t *testing.T) { + body := `{"system":[{"type":"text","text":"stable cache anchor","cache_control":{"type":"ephemeral"}}],"input":"hello"}` + first := mustParseResponsesSessionHashRequest(t, body, ctx) + second := mustParseResponsesSessionHashRequest(t, body, &SessionContext{ClientIP: "9.8.7.6", UserAgent: "other", APIKeyID: 2}) + require.Equal(t, svc.GenerateSessionHash(first), svc.GenerateSessionHash(second)) + }) +} + func TestGenerateSessionHash_MessageRollback(t *testing.T) { svc := &GatewayService{} ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}