mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
fix(gateway): anchor responses fallback to input
This commit is contained in:
@@ -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}
|
||||
|
||||
Reference in New Issue
Block a user