fix(gateway): anchor responses fallback to input

This commit is contained in:
wucm667
2026-06-09 13:53:19 +08:00
parent 434af38fd5
commit a67b10f468
4 changed files with 146 additions and 0 deletions
@@ -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 {
@@ -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) {
@@ -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() {
@@ -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}