Skip to content

Commit ab4e8b2

Browse files
author
QTom
committed
fix(gateway): 防止 OpenAI Codex 跨用户串流
根因:多个用户共享同一 OAuth 账号时,conversation_id/session_id 头 未做用户隔离,导致上游 chatgpt.com 将不同用户的请求关联到同一会话。 HTTP SSE 修复: - 新增 isolateOpenAISessionID(apiKeyID, raw),将 API Key ID 混入 session 标识符(xxhash),确保不同 Key 的用户产生不同上游会话 - buildUpstreamRequest: OAuth 分支先 Del 客户端透传的 session 头, 再用隔离值覆盖 - buildUpstreamRequestOpenAIPassthrough: 透传路径同样隔离 - ForwardAsAnthropic: Anthropic Messages 兼容路径同步修复 - buildOpenAIWSHeaders: WS 路径的 OAuth session 头同步隔离
1 parent 474165d commit ab4e8b2

5 files changed

Lines changed: 119 additions & 25 deletions

File tree

backend/internal/service/openai_gateway_messages.go

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -107,10 +107,11 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
107107
return nil, fmt.Errorf("build upstream request: %w", err)
108108
}
109109

110-
// Override session_id with a deterministic UUID derived from the sticky
111-
// session key (buildUpstreamRequest may have set it to the raw value).
110+
// Override session_id with a deterministic UUID derived from the isolated
111+
// session key, ensuring different API keys produce different upstream sessions.
112112
if promptCacheKey != "" {
113-
upstreamReq.Header.Set("session_id", generateSessionUUID(promptCacheKey))
113+
apiKeyID := getAPIKeyIDFromContext(c)
114+
upstreamReq.Header.Set("session_id", generateSessionUUID(isolateOpenAISessionID(apiKeyID, promptCacheKey)))
114115
}
115116

116117
// 7. Send request

backend/internal/service/openai_gateway_service.go

Lines changed: 43 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ import (
2424
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
2525
"github.com/Wei-Shaw/sub2api/internal/util/responseheaders"
2626
"github.com/Wei-Shaw/sub2api/internal/util/urlvalidator"
27+
"github.com/cespare/xxhash/v2"
2728
"github.com/gin-gonic/gin"
2829
"github.com/google/uuid"
2930
"github.com/tidwall/gjson"
@@ -787,6 +788,20 @@ func getAPIKeyIDFromContext(c *gin.Context) int64 {
787788
return apiKey.ID
788789
}
789790

791+
// isolateOpenAISessionID 将 apiKeyID 混入 session 标识符,
792+
// 确保不同 API Key 的用户即使使用相同的原始 session_id/conversation_id,
793+
// 到达上游的标识符也不同,防止跨用户会话碰撞。
794+
func isolateOpenAISessionID(apiKeyID int64, raw string) string {
795+
raw = strings.TrimSpace(raw)
796+
if raw == "" {
797+
return ""
798+
}
799+
h := xxhash.New()
800+
_, _ = fmt.Fprintf(h, "k%d:", apiKeyID)
801+
_, _ = h.WriteString(raw)
802+
return fmt.Sprintf("%016x", h.Sum64())
803+
}
804+
790805
func logCodexCLIOnlyDetection(ctx context.Context, c *gin.Context, account *Account, apiKeyID int64, result CodexClientRestrictionDetectionResult, body []byte) {
791806
if !result.Enabled {
792807
return
@@ -2501,13 +2516,17 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough(
25012516
if chatgptAccountID := account.GetChatGPTAccountID(); chatgptAccountID != "" {
25022517
req.Header.Set("chatgpt-account-id", chatgptAccountID)
25032518
}
2519+
apiKeyID := getAPIKeyIDFromContext(c)
2520+
// 先保存客户端原始值,再做 compact 补充,避免后续统一隔离时读到已处理的值。
2521+
clientSessionID := strings.TrimSpace(req.Header.Get("session_id"))
2522+
clientConversationID := strings.TrimSpace(req.Header.Get("conversation_id"))
25042523
if isOpenAIResponsesCompactPath(c) {
25052524
req.Header.Set("accept", "application/json")
25062525
if req.Header.Get("version") == "" {
25072526
req.Header.Set("version", codexCLIVersion)
25082527
}
2509-
if req.Header.Get("session_id") == "" {
2510-
req.Header.Set("session_id", resolveOpenAICompactSessionID(c))
2528+
if clientSessionID == "" {
2529+
clientSessionID = resolveOpenAICompactSessionID(c)
25112530
}
25122531
} else if req.Header.Get("accept") == "" {
25132532
req.Header.Set("accept", "text/event-stream")
@@ -2518,13 +2537,18 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough(
25182537
if req.Header.Get("originator") == "" {
25192538
req.Header.Set("originator", "codex_cli_rs")
25202539
}
2521-
if promptCacheKey != "" {
2522-
if req.Header.Get("conversation_id") == "" {
2523-
req.Header.Set("conversation_id", promptCacheKey)
2524-
}
2525-
if req.Header.Get("session_id") == "" {
2526-
req.Header.Set("session_id", promptCacheKey)
2527-
}
2540+
// 用隔离后的 session 标识符覆盖客户端透传值,防止跨用户会话碰撞。
2541+
if clientSessionID == "" {
2542+
clientSessionID = promptCacheKey
2543+
}
2544+
if clientConversationID == "" {
2545+
clientConversationID = promptCacheKey
2546+
}
2547+
if clientSessionID != "" {
2548+
req.Header.Set("session_id", isolateOpenAISessionID(apiKeyID, clientSessionID))
2549+
}
2550+
if clientConversationID != "" {
2551+
req.Header.Set("conversation_id", isolateOpenAISessionID(apiKeyID, clientConversationID))
25282552
}
25292553
}
25302554

@@ -2887,22 +2911,27 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin.
28872911
}
28882912
}
28892913
if account.Type == AccountTypeOAuth {
2914+
// 清除客户端透传的 session 头,后续用隔离后的值重新设置,防止跨用户会话碰撞。
2915+
req.Header.Del("conversation_id")
2916+
req.Header.Del("session_id")
2917+
28902918
req.Header.Set("OpenAI-Beta", "responses=experimental")
28912919
req.Header.Set("originator", resolveOpenAIUpstreamOriginator(c, isCodexCLI))
2920+
apiKeyID := getAPIKeyIDFromContext(c)
28922921
if isOpenAIResponsesCompactPath(c) {
28932922
req.Header.Set("accept", "application/json")
28942923
if req.Header.Get("version") == "" {
28952924
req.Header.Set("version", codexCLIVersion)
28962925
}
2897-
if req.Header.Get("session_id") == "" {
2898-
req.Header.Set("session_id", resolveOpenAICompactSessionID(c))
2899-
}
2926+
compactSession := resolveOpenAICompactSessionID(c)
2927+
req.Header.Set("session_id", isolateOpenAISessionID(apiKeyID, compactSession))
29002928
} else {
29012929
req.Header.Set("accept", "text/event-stream")
29022930
}
29032931
if promptCacheKey != "" {
2904-
req.Header.Set("conversation_id", promptCacheKey)
2905-
req.Header.Set("session_id", promptCacheKey)
2932+
isolated := isolateOpenAISessionID(apiKeyID, promptCacheKey)
2933+
req.Header.Set("conversation_id", isolated)
2934+
req.Header.Set("session_id", isolated)
29062935
}
29072936
}
29082937

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,50 @@
1+
package service
2+
3+
import (
4+
"testing"
5+
6+
"github.com/stretchr/testify/assert"
7+
"github.com/stretchr/testify/require"
8+
)
9+
10+
func TestIsolateOpenAISessionID(t *testing.T) {
11+
t.Run("empty_raw_returns_empty", func(t *testing.T) {
12+
assert.Equal(t, "", isolateOpenAISessionID(1, ""))
13+
assert.Equal(t, "", isolateOpenAISessionID(1, " "))
14+
})
15+
16+
t.Run("deterministic", func(t *testing.T) {
17+
a := isolateOpenAISessionID(42, "sess_abc123")
18+
b := isolateOpenAISessionID(42, "sess_abc123")
19+
assert.Equal(t, a, b)
20+
})
21+
22+
t.Run("different_apiKeyID_different_result", func(t *testing.T) {
23+
a := isolateOpenAISessionID(1, "same_session")
24+
b := isolateOpenAISessionID(2, "same_session")
25+
require.NotEqual(t, a, b, "不同 API Key 使用相同 session_id 应产生不同隔离值")
26+
})
27+
28+
t.Run("different_raw_different_result", func(t *testing.T) {
29+
a := isolateOpenAISessionID(1, "session_a")
30+
b := isolateOpenAISessionID(1, "session_b")
31+
require.NotEqual(t, a, b)
32+
})
33+
34+
t.Run("format_is_16_hex_chars", func(t *testing.T) {
35+
result := isolateOpenAISessionID(99, "test_session")
36+
assert.Len(t, result, 16, "应为 16 字符的 hex 字符串")
37+
for _, ch := range result {
38+
assert.True(t, (ch >= '0' && ch <= '9') || (ch >= 'a' && ch <= 'f'),
39+
"应仅包含 hex 字符: %c", ch)
40+
}
41+
})
42+
43+
t.Run("zero_apiKeyID_still_works", func(t *testing.T) {
44+
result := isolateOpenAISessionID(0, "session")
45+
assert.NotEmpty(t, result)
46+
// apiKeyID=0 与 apiKeyID=1 应产生不同结果
47+
other := isolateOpenAISessionID(1, "session")
48+
assert.NotEqual(t, result, other)
49+
})
50+
}

backend/internal/service/openai_ws_forwarder.go

Lines changed: 16 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1124,11 +1124,22 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders(
11241124
headers.Set("accept-language", v)
11251125
}
11261126
}
1127-
if sessionResolution.SessionID != "" {
1128-
headers.Set("session_id", sessionResolution.SessionID)
1129-
}
1130-
if sessionResolution.ConversationID != "" {
1131-
headers.Set("conversation_id", sessionResolution.ConversationID)
1127+
// OAuth 账号:将 apiKeyID 混入 session 标识符,防止跨用户会话碰撞。
1128+
if account != nil && account.Type == AccountTypeOAuth {
1129+
apiKeyID := getAPIKeyIDFromContext(c)
1130+
if sessionResolution.SessionID != "" {
1131+
headers.Set("session_id", isolateOpenAISessionID(apiKeyID, sessionResolution.SessionID))
1132+
}
1133+
if sessionResolution.ConversationID != "" {
1134+
headers.Set("conversation_id", isolateOpenAISessionID(apiKeyID, sessionResolution.ConversationID))
1135+
}
1136+
} else {
1137+
if sessionResolution.SessionID != "" {
1138+
headers.Set("session_id", sessionResolution.SessionID)
1139+
}
1140+
if sessionResolution.ConversationID != "" {
1141+
headers.Set("conversation_id", sessionResolution.ConversationID)
1142+
}
11321143
}
11331144
if state := strings.TrimSpace(turnState); state != "" {
11341145
headers.Set(openAIWSTurnStateHeader, state)

backend/internal/service/openai_ws_forwarder_success_test.go

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -454,8 +454,10 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T
454454
require.True(t, gjson.Get(requestJSON, "stream").Exists(), "WSv2 payload 应保留 stream 字段")
455455
require.True(t, gjson.Get(requestJSON, "stream").Bool(), "OAuth Codex 规范化后应强制 stream=true")
456456
require.Equal(t, openAIWSBetaV2Value, captureDialer.lastHeaders.Get("OpenAI-Beta"))
457-
require.Equal(t, "sess-oauth-1", captureDialer.lastHeaders.Get("session_id"))
458-
require.Equal(t, "conv-oauth-1", captureDialer.lastHeaders.Get("conversation_id"))
457+
// OAuth 账号的 session_id/conversation_id 应被 isolateOpenAISessionID 隔离,
458+
// 测试中未设置 api_key 到 context,apiKeyID=0。
459+
require.Equal(t, isolateOpenAISessionID(0, "sess-oauth-1"), captureDialer.lastHeaders.Get("session_id"))
460+
require.Equal(t, isolateOpenAISessionID(0, "conv-oauth-1"), captureDialer.lastHeaders.Get("conversation_id"))
459461
}
460462

461463
func TestOpenAIGatewayService_Forward_WSv2_OAuthOriginatorCompatibility(t *testing.T) {
@@ -596,7 +598,8 @@ func TestOpenAIGatewayService_Forward_WSv2_HeaderSessionFallbackFromPromptCacheK
596598
require.NotNil(t, result)
597599
require.Equal(t, "resp_prompt_cache_key", result.RequestID)
598600

599-
require.Equal(t, "pcache_123", captureDialer.lastHeaders.Get("session_id"))
601+
// OAuth 账号的 session_id 应被 isolateOpenAISessionID 隔离(apiKeyID=0,未在 context 设置)。
602+
require.Equal(t, isolateOpenAISessionID(0, "pcache_123"), captureDialer.lastHeaders.Get("session_id"))
600603
require.Empty(t, captureDialer.lastHeaders.Get("conversation_id"))
601604
require.NotNil(t, captureConn.lastWrite)
602605
require.True(t, gjson.Get(requestToJSONString(captureConn.lastWrite), "stream").Exists())

0 commit comments

Comments
 (0)