Merge pull request #3971 from yardbirds0/codex/fix-remote-compaction-v2

fix: 保留 remote_compaction_v2 原生 Responses 请求链路
This commit is contained in:
Wesley Liddick
2026-07-13 08:54:53 +08:00
committed by GitHub
10 changed files with 505 additions and 83 deletions
@@ -23,46 +23,61 @@ func newCompactBodySignalTestContext(t *testing.T, path string, body []byte) *gi
return c
}
// body-signal 提升后必须与 path-based compact 走同一条链路:
// path 改写、requireCompact 判定、stream/store/prompt_cache_key 归一化删除。
// 回归防护:若 stream 字段存活,Forward 会用流式 handler 解析 compact 的
// JSON 响应,导致 "stream ended before a terminal event" 的换号 failover 风暴。
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalPromoted(t *testing.T) {
func TestNormalizeOpenAIResponsesCompactRequest_RemoteV2StaysOnResponses(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{
"model":"gpt-5.5",
"model":"gpt-5.6-sol",
"stream":true,
"store":true,
"prompt_cache_key":"pck-signal-1",
"reasoning":{"effort":"max","context":"all_turns"},
"input":[
{"type":"message","role":"user","content":"hello"},
{"type":"compaction_trigger"}
]
}`)
c := newCompactBodySignalTestContext(t, "/v1/responses", body)
c.Request.Header.Set("x-codex-beta-features", "responses_websockets_v2, remote_compaction_v2, another_feature")
normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok)
require.Equal(t, "/v1/responses/compact", c.Request.URL.Path)
require.True(t, isOpenAIRemoteCompactPath(c))
require.False(t, gjson.GetBytes(normalized, "stream").Exists())
require.False(t, gjson.GetBytes(normalized, "store").Exists())
require.False(t, gjson.GetBytes(normalized, "prompt_cache_key").Exists())
require.Equal(t, "gpt-5.5", gjson.GetBytes(normalized, "model").String())
require.True(t, gjson.GetBytes(normalized, "input").IsArray())
require.Equal(t, "/v1/responses", c.Request.URL.Path)
require.False(t, isOpenAIRemoteCompactPath(c))
require.Equal(t, body, normalized)
require.True(t, gjson.GetBytes(normalized, "stream").Bool())
require.True(t, gjson.GetBytes(normalized, "store").Bool())
require.Equal(t, "pck-signal-1", gjson.GetBytes(normalized, "prompt_cache_key").String())
require.Equal(t, "max", gjson.GetBytes(normalized, "reasoning.effort").String())
require.Equal(t, "all_turns", gjson.GetBytes(normalized, "reasoning.context").String())
reqStream, streamOK := parseOpenAICompatibleStream(normalized)
require.True(t, streamOK)
require.False(t, reqStream)
require.True(t, reqStream)
seed, exists := c.Get(service.OpenAICompactSessionSeedKeyForTest())
require.True(t, exists)
require.Equal(t, "pck-signal-1", seed)
_, seedExists := c.Get(service.OpenAICompactSessionSeedKeyForTest())
require.False(t, seedExists)
_, streamMarkerExists := c.Get(service.OpenAICompactClientStreamKeyForTest())
require.False(t, streamMarkerExists)
}
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalTrailingSlash(t *testing.T) {
func TestNormalizeOpenAIResponsesCompactRequest_RemoteV2PathAliasesStayOnResponses(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.6-sol","stream":true,"input":[{"type":"compaction_trigger"}]}`)
for _, path := range []string{"/v1/responses/", "/backend-api/codex/responses"} {
t.Run(path, func(t *testing.T) {
c := newCompactBodySignalTestContext(t, path, body)
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok)
require.Equal(t, path, c.Request.URL.Path)
require.Equal(t, body, normalized)
})
}
}
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalTrailingSlashPromoted(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`)
c := newCompactBodySignalTestContext(t, "/v1/responses/", body)
@@ -82,6 +97,64 @@ func TestNormalizeOpenAIResponsesCompactRequest_CodexDirectAliasPromoted(t *test
require.Equal(t, "/backend-api/codex/responses/compact", c.Request.URL.Path)
}
func TestNormalizeOpenAIResponsesCompactRequest_NonRemoteV2BodySignalPromoted(t *testing.T) {
h := &OpenAIGatewayHandler{}
tests := []struct {
name string
body []byte
betaHeader string
wantMarked bool
}{
{
name: "no_header",
body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`),
wantMarked: true,
},
{
name: "unrelated_header",
body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`),
betaHeader: "responses_websockets_v2",
wantMarked: true,
},
{
name: "wrong_case_header",
body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`),
betaHeader: "REMOTE_COMPACTION_V2",
wantMarked: true,
},
{
name: "stream_false",
body: []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"compaction_trigger"}]}`),
betaHeader: "remote_compaction_v2",
},
{
name: "stream_absent",
body: []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`),
betaHeader: "remote_compaction_v2",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
c := newCompactBodySignalTestContext(t, "/v1/responses", tt.body)
if tt.betaHeader != "" {
c.Request.Header.Set("x-codex-beta-features", tt.betaHeader)
}
normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), tt.body)
require.True(t, ok)
require.Equal(t, "/v1/responses/compact", c.Request.URL.Path)
require.False(t, gjson.GetBytes(normalized, "stream").Exists())
marked, exists := c.Get(service.OpenAICompactClientStreamKeyForTest())
require.Equal(t, tt.wantMarked, exists)
if tt.wantMarked {
require.Equal(t, true, marked)
}
})
}
}
func TestNormalizeOpenAIResponsesCompactRequest_NoTriggerUntouched(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"message","role":"user","content":"hello"}]}`)
@@ -99,6 +172,7 @@ func TestNormalizeOpenAIResponsesCompactRequest_PathBasedNoDoubleSuffix(t *testi
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.5","stream":true,"store":true,"input":[{"type":"message","role":"user","content":"hello"}]}`)
c := newCompactBodySignalTestContext(t, "/v1/responses/compact", body)
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok)
@@ -118,36 +192,6 @@ func TestNormalizeOpenAIResponsesCompactRequest_SubpathNotPromoted(t *testing.T)
require.Equal(t, body, normalized)
}
// 回归 #3875body-signal 原始请求 stream:true 时必须标记 client-stream
// 供响应写回阶段把上游 unary JSON 合成回 Codex remote compact v2 所需的 SSE。
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalStreamTrueMarksClientStream(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`)
c := newCompactBodySignalTestContext(t, "/v1/responses", body)
_, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok)
marked, exists := c.Get(service.OpenAICompactClientStreamKeyForTest())
require.True(t, exists)
require.Equal(t, true, marked)
}
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalStreamFalseNotMarked(t *testing.T) {
h := &OpenAIGatewayHandler{}
for name, body := range map[string][]byte{
"stream_false": []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"compaction_trigger"}]}`),
"stream_absent": []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`),
} {
c := newCompactBodySignalTestContext(t, "/v1/responses", body)
_, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok, name)
require.Equal(t, "/v1/responses/compact", c.Request.URL.Path, name)
_, exists := c.Get(service.OpenAICompactClientStreamKeyForTest())
require.False(t, exists, "case %s 不应标记 client-stream", name)
}
}
// path-based compactCodex v1 unary 协议)即使 body 带 stream:true 也不标记,
// 保持 JSON 写回行为不变。
func TestNormalizeOpenAIResponsesCompactRequest_PathBasedStreamTrueNotMarked(t *testing.T) {
@@ -580,21 +580,33 @@ func isBareOpenAIResponsesPath(c *gin.Context) bool {
return strings.HasSuffix(normalizedPath, "/responses")
}
// normalizeOpenAIResponsesCompactRequest 统一处理两种入站 compact 形态:
// path-basedPOST /v1/responses/compact)与 Codex remote compact v2 的
// body-signal(普通 POST /v1/responses 的 input 中携带 type=compaction_trigger
// 见 #3777)。body-signal 命中时在 stream 解析、compact body 归一化与
// requireCompact 调度判定之前改写 URL path,使后续全部链路(含 passthrough
// 分支与上游 URL 构建)与 path-based 完全一致。
func isOpenAIRemoteCompactionV2Request(c *gin.Context, body []byte) bool {
stream, valid := parseOpenAICompatibleStream(body)
if !valid || !stream || c == nil || c.Request == nil {
return false
}
for _, header := range c.Request.Header.Values("x-codex-beta-features") {
for _, feature := range strings.Split(header, ",") {
if strings.TrimSpace(feature) == "remote_compaction_v2" {
return true
}
}
}
return false
}
// normalizeOpenAIResponsesCompactRequest keeps Codex remote compaction v2 on
// its native streaming /responses wire and preserves the legacy body-signal
// promotion for clients that do not explicitly advertise that protocol.
// 返回归一化后的 body;ok=false 表示错误响应已写出,调用方应直接 return。
func (h *OpenAIGatewayHandler) normalizeOpenAIResponsesCompactRequest(c *gin.Context, reqLog *zap.Logger, body []byte) ([]byte, bool) {
isCompactRequest := service.IsOpenAIResponsesCompactPathForTest(c)
if !isCompactRequest && isBareOpenAIResponsesPath(c) && service.HasCompactionTriggerInInput(body) {
if isOpenAIRemoteCompactionV2Request(c, body) {
return body, true
}
c.Request.URL.Path = strings.TrimRight(c.Request.URL.Path, "/") + "/compact"
isCompactRequest = true
// Codex remote compact v2 的原始请求是流式 /responses:白名单归一化会删除
// stream 并让上游走 unary JSON,但客户端仍按 SSE 消费响应。记录原始
// stream 意图,响应写回阶段据此把 JSON 合成回 SSE(#3875)。
clientStream := gjson.GetBytes(body, "stream").Bool()
if clientStream {
service.MarkOpenAICompactClientStream(c)
@@ -2,18 +2,10 @@ package service
import "github.com/tidwall/gjson"
// HasCompactionTriggerInInput detects the Codex remote compact v2 body signal:
// an input item with type "compaction_trigger". When the client sends this
// inside a normal POST /v1/responses (instead of POST /v1/responses/compact),
// the request must still be treated as a compact request — otherwise the
// upstream path, model mapping, and body normalization are all wrong, causing
// Codex to receive a non-compact response and fail with:
//
// "remote compaction v2 expected exactly one compaction output item, got 0"
//
// The gateway handler promotes such requests by rewriting the URL path to the
// compact form before stream parsing, compact body normalization, and
// compact-capable account scheduling, so both inbound forms share one code path.
// HasCompactionTriggerInInput detects an input item with
// type="compaction_trigger". The handler combines this body signal with the
// request path, stream flag, and Codex beta feature header to distinguish the
// native remote compaction v2 wire from the legacy /responses/compact bridge.
func HasCompactionTriggerInInput(body []byte) bool {
if len(body) == 0 {
return false
@@ -66,6 +66,7 @@ var openaiAllowedHeaders = map[string]bool{
"user-agent": true,
"originator": true,
"session_id": true,
"x-codex-beta-features": true,
"x-codex-turn-state": true,
"x-codex-turn-metadata": true,
}
@@ -81,6 +82,7 @@ var openaiPassthroughAllowedHeaders = map[string]bool{
"user-agent": true,
"originator": true,
"session_id": true,
"x-codex-beta-features": true,
"x-codex-turn-state": true,
"x-codex-turn-metadata": true,
}
@@ -223,13 +223,17 @@ func TestOpenAIGatewayServiceForwardOAuthCompactDowngradesMaxEffort(t *testing.T
require.Equal(t, "xhigh", *result.ReasoningEffort)
}
func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing.T) {
func TestOpenAIGatewayServiceForwardOAuthRemoteCompactV2PreservesResponsesWire(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{
resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"compaction\",\"encrypted_content\":\"summary\"}}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":1,\"output_tokens\":2}}}\n\n" +
"data: [DONE]\n\n",
)),
},
}
cfg := &config.Config{}
@@ -244,6 +248,9 @@ func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing.
Credentials: map[string]any{
"access_token": "oauth-token",
"chatgpt_account_id": "chatgpt-acc",
"compact_model_mapping": map[string]any{
"gpt-5.6-sol": "gpt-5.6-sol-openai-compact",
},
},
Status: StatusActive,
Schedulable: true,
@@ -251,16 +258,82 @@ func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing.
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
body := []byte(`{"model":"gpt-5.6-sol","instructions":"response-test","input":"hello","reasoning":{"effort":"max"}}`)
body := []byte(`{"model":"gpt-5.6-sol","stream":true,"instructions":"response-test","input":[{"type":"message","role":"user","content":"hello"},{"type":"compaction_trigger"}],"reasoning":{"effort":"max","context":"all_turns"}}`)
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, upstream.lastReq)
require.Equal(t, chatgptCodexURL, upstream.lastReq.URL.String())
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String())
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
require.Equal(t, "compaction_trigger", gjson.GetBytes(upstream.lastBody, "input.#(type==\"compaction_trigger\").type").String())
require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String())
require.Equal(t, "all_turns", gjson.GetBytes(upstream.lastBody, "reasoning.context").String())
require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features"))
require.Contains(t, rec.Body.String(), `"type":"compaction"`)
require.Contains(t, rec.Body.String(), `"encrypted_content":"summary"`)
require.NotNil(t, result.ReasoningEffort)
require.Equal(t, "max", *result.ReasoningEffort)
}
func TestOpenAIGatewayServiceForwardAPIKeyRemoteCompactV2PreservesResponsesWire(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{
resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"compaction\",\"encrypted_content\":\"summary\"}}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":1,\"output_tokens\":2}}}\n\n" +
"data: [DONE]\n\n",
)),
},
}
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
account := &Account{
ID: 11,
Name: "openai-apikey-responses",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "https://example.com/v1",
"compact_model_mapping": map[string]any{
"gpt-5.6-sol": "gpt-5.6-sol-openai-compact",
},
},
Extra: map[string]any{"use_responses_api": true},
Status: StatusActive,
Schedulable: true,
}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
body := []byte(`{"model":"gpt-5.6-sol","stream":true,"instructions":"response-test","input":[{"type":"message","role":"user","content":"hello"},{"type":"compaction_trigger"}],"reasoning":{"effort":"max","context":"all_turns"}}`)
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, upstream.lastReq)
require.Equal(t, "https://example.com/v1/responses", upstream.lastReq.URL.String())
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String())
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
require.Equal(t, "compaction_trigger", gjson.GetBytes(upstream.lastBody, "input.#(type==\"compaction_trigger\").type").String())
require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String())
require.Equal(t, "all_turns", gjson.GetBytes(upstream.lastBody, "reasoning.context").String())
require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features"))
require.Contains(t, rec.Body.String(), `"type":"compaction"`)
require.Contains(t, rec.Body.String(), `"encrypted_content":"summary"`)
require.NotNil(t, result.ReasoningEffort)
require.Equal(t, "max", *result.ReasoningEffort)
}
@@ -347,6 +347,7 @@ func TestOpenAIGatewayService_OAuthPassthrough_StreamKeepsToolNameAndBodyNormali
c.Request.Header.Set("Accept-Encoding", "gzip")
c.Request.Header.Set("Proxy-Authorization", "Basic abc")
c.Request.Header.Set("X-Test", "keep")
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
originalBody := []byte(`{"model":"gpt-5.2","stream":true,"store":true,"instructions":"local-test-instructions","input":[{"type":"text","text":"hi"}]}`)
@@ -409,6 +410,7 @@ func TestOpenAIGatewayService_OAuthPassthrough_StreamKeepsToolNameAndBodyNormali
require.Empty(t, upstream.lastReq.Header.Get("Accept-Encoding"))
require.Empty(t, upstream.lastReq.Header.Get("Proxy-Authorization"))
require.Empty(t, upstream.lastReq.Header.Get("X-Test"))
require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features"))
// 3) required OAuth headers are present
require.Equal(t, "chatgpt.com", upstream.lastReq.Host)
@@ -1373,6 +1375,7 @@ func TestOpenAIGatewayService_APIKeyPassthrough_PreservesBodyAndUsesResponsesEnd
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil))
c.Request.Header.Set("User-Agent", "curl/8.0")
c.Request.Header.Set("X-Test", "keep")
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
originalBody := []byte(`{"model":"gpt-5.2","stream":false,"service_tier":"flex","max_output_tokens":128,"input":[{"type":"text","text":"hi"}]}`)
resp := &http.Response{
@@ -1410,6 +1413,7 @@ func TestOpenAIGatewayService_APIKeyPassthrough_PreservesBodyAndUsesResponsesEnd
require.Equal(t, "https://api.openai.com/v1/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer sk-api-key", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "curl/8.0", upstream.lastReq.Header.Get("User-Agent"))
require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features"))
require.Empty(t, upstream.lastReq.Header.Get("X-Test"))
}
@@ -74,6 +74,11 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders(
if v := strings.TrimSpace(c.Request.Header.Get("accept-language")); v != "" {
headers.Set("accept-language", v)
}
for _, value := range c.Request.Header.Values("x-codex-beta-features") {
if value = strings.TrimSpace(value); value != "" {
headers.Add("x-codex-beta-features", value)
}
}
}
// OAuth 账号:将 apiKeyID 混入 session 标识符,防止跨用户会话碰撞。
if account != nil && account.Type == AccountTypeOAuth {
@@ -602,6 +602,7 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T
c.Request.Header.Set("User-Agent", "codex_cli_rs/0.98.0")
c.Request.Header.Set("session_id", "sess-oauth-1")
c.Request.Header.Set("conversation_id", "conv-oauth-1")
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
@@ -661,6 +662,7 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T
require.True(t, gjson.Get(requestJSON, "stream").Exists(), "WSv2 payload 应保留 stream 字段")
require.True(t, gjson.Get(requestJSON, "stream").Bool(), "OAuth Codex 规范化后应强制 stream=true")
require.Equal(t, openAIWSBetaV2Value, captureDialer.lastHeaders.Get("OpenAI-Beta"))
require.Equal(t, "remote_compaction_v2", captureDialer.lastHeaders.Get("x-codex-beta-features"))
// OAuth 账号的 session_id/conversation_id 应被 isolateOpenAISessionID 隔离,
// 测试中未设置 api_key 到 contextapiKeyID=0。
require.Equal(t, isolateOpenAISessionID(0, "sess-oauth-1"), captureDialer.lastHeaders.Get("session_id"))
+132 -9
View File
@@ -218,6 +218,9 @@ func (l *openAIWSConnLease) Release() {
return
}
l.conn.release()
if l.pool != nil {
l.pool.notifyAccountPoolChanged(l.accountID)
}
}
type openAIWSConn struct {
@@ -225,6 +228,7 @@ type openAIWSConn struct {
ws openAIWSClientConn
handshakeHeaders http.Header
betaFeatures string
leaseCh chan struct{}
closedCh chan struct{}
@@ -498,6 +502,10 @@ func (c *openAIWSConn) handshakeHeader(name string) string {
return strings.TrimSpace(c.handshakeHeaders.Get(strings.TrimSpace(name)))
}
func (c *openAIWSConn) matchesBetaFeatures(betaFeatures string) bool {
return c != nil && c.betaFeatures == betaFeatures
}
func (c *openAIWSConn) isPrewarmed() bool {
if c == nil {
return false
@@ -516,6 +524,7 @@ type openAIWSAccountPool struct {
mu sync.Mutex
conns map[string]*openAIWSConn
pinnedConns map[string]int
changedCh chan struct{}
creating int
lastCleanupAt time.Time
lastAcquire *openAIWSAcquireRequest
@@ -525,6 +534,23 @@ type openAIWSAccountPool struct {
prewarmFailAt time.Time
}
func (ap *openAIWSAccountPool) changeChannelLocked() chan struct{} {
if ap.changedCh == nil {
ap.changedCh = make(chan struct{})
}
return ap.changedCh
}
func (ap *openAIWSAccountPool) signalChangedLocked() {
if ap == nil {
return
}
if ap.changedCh != nil {
close(ap.changedCh)
}
ap.changedCh = make(chan struct{})
}
type OpenAIWSPoolMetricsSnapshot struct {
AcquireTotal int64
AcquireReuseTotal int64
@@ -786,7 +812,9 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
return nil, errors.New("ws url is empty")
}
retryAcquire:
accountID := req.Account.ID
betaFeatures := normalizeOpenAIWSBetaFeatures(req.Headers)
effectiveMaxConns := p.effectiveMaxConnsByAccount(req.Account)
if effectiveMaxConns <= 0 {
return nil, errOpenAIWSConnQueueFull
@@ -814,7 +842,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
return nil, errOpenAIWSPreferredConnUnavailable
}
preferredConn, ok := ap.conns[preferredConnID]
if !ok || preferredConn == nil {
if !ok || !preferredConn.matchesBetaFeatures(betaFeatures) {
p.recordConnPickDuration(time.Since(pickStartedAt))
ap.mu.Unlock()
closeOpenAIWSConns(evicted)
@@ -895,7 +923,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
}
if preferredConnID != "" {
if conn, ok := ap.conns[preferredConnID]; ok && conn.tryAcquire() {
if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) && conn.tryAcquire() {
connPick := time.Since(pickStartedAt)
p.recordConnPickDuration(connPick)
ap.mu.Unlock()
@@ -917,7 +945,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
}
}
best := p.pickLeastBusyConnLocked(ap, "")
best := p.pickLeastBusyConnLocked(ap, "", betaFeatures)
if best != nil && best.tryAcquire() {
connPick := time.Since(pickStartedAt)
p.recordConnPickDuration(connPick)
@@ -939,7 +967,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
return lease, nil
}
for _, conn := range ap.conns {
if conn == nil || conn == best {
if conn == nil || conn == best || !conn.matchesBetaFeatures(betaFeatures) {
continue
}
if conn.tryAcquire() {
@@ -965,6 +993,37 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
}
}
if !req.ForceNewConn && len(ap.conns)+ap.creating >= effectiveMaxConns {
compatible := p.pickLeastBusyConnLocked(ap, "", betaFeatures)
if idle := p.pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap, betaFeatures); idle != nil {
delete(ap.conns, idle.id)
evicted = append(evicted, idle)
p.metrics.scaleDownTotal.Add(1)
} else if compatible == nil {
hasConnection := false
for _, conn := range ap.conns {
if conn != nil {
hasConnection = true
break
}
}
if !hasConnection && ap.creating == 0 {
ap.mu.Unlock()
closeOpenAIWSConns(evicted)
return nil, errOpenAIWSConnClosed
}
changedCh := ap.changeChannelLocked()
ap.mu.Unlock()
closeOpenAIWSConns(evicted)
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-changedCh:
goto retryAcquire
}
}
}
if req.ForceNewConn && len(ap.conns)+ap.creating >= effectiveMaxConns {
if idle := p.pickOldestIdleConnLocked(ap); idle != nil {
delete(ap.conns, idle.id)
@@ -988,6 +1047,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
if dialErr != nil {
ap.prewarmFails++
ap.prewarmFailAt = time.Now()
ap.signalChangedLocked()
ap.mu.Unlock()
return nil, dialErr
}
@@ -1016,7 +1076,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
return nil, errOpenAIWSConnQueueFull
}
target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID)
target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID, betaFeatures)
connPick := time.Since(pickStartedAt)
p.recordConnPickDuration(connPick)
if target == nil {
@@ -1089,6 +1149,22 @@ func (p *openAIWSConnPool) pickOldestIdleConnLocked(ap *openAIWSAccountPool) *op
return oldest
}
func (p *openAIWSConnPool) pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap *openAIWSAccountPool, betaFeatures string) *openAIWSConn {
if ap == nil || len(ap.conns) == 0 {
return nil
}
var oldest *openAIWSConn
for _, conn := range ap.conns {
if conn == nil || conn.matchesBetaFeatures(betaFeatures) || conn.isLeased() || conn.waiters.Load() > 0 || p.isConnPinnedLocked(ap, conn.id) {
continue
}
if oldest == nil || conn.lastUsedAt().Before(oldest.lastUsedAt()) {
oldest = conn
}
}
return oldest
}
func (p *openAIWSConnPool) getOrCreateAccountPool(accountID int64) *openAIWSAccountPool {
if p == nil || accountID <= 0 {
return nil
@@ -1101,6 +1177,7 @@ func (p *openAIWSConnPool) getOrCreateAccountPool(accountID int64) *openAIWSAcco
ap := &openAIWSAccountPool{
conns: make(map[string]*openAIWSConn),
pinnedConns: make(map[string]int),
changedCh: make(chan struct{}),
}
actual, _ := p.accounts.LoadOrStore(accountID, ap)
if typed, ok := actual.(*openAIWSAccountPool); ok && typed != nil {
@@ -1126,6 +1203,16 @@ func (p *openAIWSConnPool) getAccountPool(accountID int64) (*openAIWSAccountPool
return ap, typed && ap != nil
}
func (p *openAIWSConnPool) notifyAccountPoolChanged(accountID int64) {
ap, ok := p.getAccountPool(accountID)
if !ok || ap == nil {
return
}
ap.mu.Lock()
ap.signalChangedLocked()
ap.mu.Unlock()
}
func (p *openAIWSConnPool) isConnPinnedLocked(ap *openAIWSAccountPool, connID string) bool {
if ap == nil || connID == "" || len(ap.pinnedConns) == 0 {
return false
@@ -1212,17 +1299,20 @@ func (p *openAIWSConnPool) cleanupAccountLocked(ap *openAIWSAccountPool, now tim
p.metrics.scaleDownTotal.Add(int64(redundant))
}
}
if len(evicted) > 0 {
ap.signalChangedLocked()
}
return evicted
}
func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, preferredConnID string) *openAIWSConn {
func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, preferredConnID, betaFeatures string) *openAIWSConn {
if ap == nil || len(ap.conns) == 0 {
return nil
}
preferredConnID = stringsTrim(preferredConnID)
if preferredConnID != "" {
if conn, ok := ap.conns[preferredConnID]; ok {
if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) {
return conn
}
}
@@ -1230,7 +1320,7 @@ func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, pref
var bestWaiters int32
var bestLastUsed time.Time
for _, conn := range ap.conns {
if conn == nil {
if conn == nil || !conn.matchesBetaFeatures(betaFeatures) {
continue
}
waiters := conn.waiters.Load()
@@ -1395,10 +1485,12 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ
if err != nil {
ap.prewarmFails++
ap.prewarmFailAt = time.Now()
ap.signalChangedLocked()
ap.mu.Unlock()
continue
}
if len(ap.conns) >= p.effectiveMaxConnsByAccount(req.Account) {
ap.signalChangedLocked()
ap.mu.Unlock()
conn.close()
continue
@@ -1406,6 +1498,7 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ
ap.conns[conn.id] = conn
ap.prewarmFails = 0
ap.prewarmFailAt = time.Time{}
ap.signalChangedLocked()
ap.mu.Unlock()
}
}
@@ -1424,6 +1517,7 @@ func (p *openAIWSConnPool) evictConn(accountID int64, connID string) {
if len(ap.pinnedConns) > 0 {
delete(ap.pinnedConns, connID)
}
ap.signalChangedLocked()
}
ap.mu.Unlock()
}
@@ -1476,9 +1570,11 @@ func (p *openAIWSConnPool) UnpinConn(accountID int64, connID string) {
count := ap.pinnedConns[connID]
if count <= 1 {
delete(ap.pinnedConns, connID)
ap.signalChangedLocked()
return
}
ap.pinnedConns[connID] = count - 1
ap.signalChangedLocked()
}
func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequest) (*openAIWSConn, error) {
@@ -1501,7 +1597,9 @@ func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequ
}
}
id := p.nextConnID(req.Account.ID)
return newOpenAIWSConn(id, req.Account.ID, conn, handshakeHeaders), nil
pooledConn := newOpenAIWSConn(id, req.Account.ID, conn, handshakeHeaders)
pooledConn.betaFeatures = normalizeOpenAIWSBetaFeatures(req.Headers)
return pooledConn, nil
}
func (p *openAIWSConnPool) nextConnID(accountID int64) string {
@@ -1679,6 +1777,31 @@ func cloneOpenAIWSAcquireRequestPtr(req *openAIWSAcquireRequest) *openAIWSAcquir
return &copied
}
func normalizeOpenAIWSBetaFeatures(headers http.Header) string {
features := make(map[string]struct{})
for name, values := range headers {
if !strings.EqualFold(strings.TrimSpace(name), "x-codex-beta-features") {
continue
}
for _, value := range values {
for _, feature := range strings.Split(value, ",") {
if feature = strings.TrimSpace(feature); feature != "" {
features[feature] = struct{}{}
}
}
}
}
if len(features) == 0 {
return ""
}
normalized := make([]string, 0, len(features))
for feature := range features {
normalized = append(normalized, feature)
}
sort.Strings(normalized)
return strings.Join(normalized, ",")
}
func cloneHeader(src http.Header) http.Header {
if src == nil {
return nil
@@ -342,6 +342,171 @@ func TestOpenAIWSConnPool_ForceNewConnSkipsReuse(t *testing.T) {
require.Equal(t, 2, dialer.DialCount(), "ForceNewConn=true 时应跳过空闲连接复用并新建连接")
}
func TestOpenAIWSConnPool_AcquireReusesOnlyMatchingBetaFeatures(t *testing.T) {
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2
pool := newOpenAIWSConnPool(cfg)
dialer := &openAIWSCountingDialer{}
pool.setClientDialerForTest(dialer)
account := &Account{ID: 128, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
baseReq := openAIWSAcquireRequest{
Account: account,
WSURL: "wss://example.com/v1/responses",
}
plainLease, err := pool.Acquire(context.Background(), baseReq)
require.NoError(t, err)
plainConnID := plainLease.ConnID()
plainLease.Release()
betaReq := baseReq
betaReq.Headers = http.Header{"X-Codex-Beta-Features": {" remote_compaction_v2 ", " responses_websockets_v2 "}}
betaLease, err := pool.Acquire(context.Background(), betaReq)
require.NoError(t, err)
require.False(t, betaLease.Reused())
require.NotEqual(t, plainConnID, betaLease.ConnID())
betaConnID := betaLease.ConnID()
betaLease.Release()
reorderedReq := baseReq
reorderedReq.Headers = http.Header{"X-Codex-Beta-Features": {"responses_websockets_v2,remote_compaction_v2"}}
reorderedLease, err := pool.Acquire(context.Background(), reorderedReq)
require.NoError(t, err)
require.True(t, reorderedLease.Reused())
require.Equal(t, betaConnID, reorderedLease.ConnID())
reorderedLease.Release()
_, err = pool.Acquire(context.Background(), openAIWSAcquireRequest{
Account: account,
WSURL: baseReq.WSURL,
Headers: betaReq.Headers,
PreferredConnID: plainConnID,
ForcePreferredConn: true,
})
require.ErrorIs(t, err, errOpenAIWSPreferredConnUnavailable)
plainLease, err = pool.Acquire(context.Background(), baseReq)
require.NoError(t, err)
require.True(t, plainLease.Reused())
require.Equal(t, plainConnID, plainLease.ConnID())
plainLease.Release()
require.Equal(t, 2, dialer.DialCount())
}
func TestOpenAIWSConnPool_AcquireReplacesIdleConnWithDifferentBetaFeatures(t *testing.T) {
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
pool := newOpenAIWSConnPool(cfg)
dialer := &openAIWSCountingDialer{}
pool.setClientDialerForTest(dialer)
account := &Account{ID: 129, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
plainLease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{
Account: account,
WSURL: "wss://example.com/v1/responses",
})
require.NoError(t, err)
plainConnID := plainLease.ConnID()
plainLease.Release()
betaLease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{
Account: account,
WSURL: "wss://example.com/v1/responses",
Headers: http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}},
})
require.NoError(t, err)
require.False(t, betaLease.Reused())
require.NotEqual(t, plainConnID, betaLease.ConnID())
betaLease.Release()
require.Equal(t, 2, dialer.DialCount())
}
func TestOpenAIWSConnPool_AcquireWaitsForBusyIncompatibleConnection(t *testing.T) {
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
pool := newOpenAIWSConnPool(cfg)
dialer := &openAIWSCountingDialer{}
pool.setClientDialerForTest(dialer)
account := &Account{ID: 130, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
baseReq := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"}
plainLease, err := pool.Acquire(context.Background(), baseReq)
require.NoError(t, err)
plainConnID := plainLease.ConnID()
type acquireResult struct {
lease *openAIWSConnLease
err error
}
resultCh := make(chan acquireResult, 1)
var done atomic.Bool
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
go func() {
betaReq := baseReq
betaReq.Headers = http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}}
lease, acquireErr := pool.Acquire(ctx, betaReq)
resultCh <- acquireResult{lease: lease, err: acquireErr}
done.Store(true)
}()
require.Never(t, done.Load, 50*time.Millisecond, 5*time.Millisecond)
plainLease.Release()
result := <-resultCh
require.NoError(t, result.err)
require.NotNil(t, result.lease)
require.False(t, result.lease.Reused())
require.NotEqual(t, plainConnID, result.lease.ConnID())
result.lease.Release()
require.Equal(t, 2, dialer.DialCount())
}
func TestOpenAIWSConnPool_AcquireReplacesIncompatibleIdleWhenMatchingBusy(t *testing.T) {
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2
pool := newOpenAIWSConnPool(cfg)
dialer := &openAIWSCountingDialer{}
pool.setClientDialerForTest(dialer)
account := &Account{ID: 131, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
baseReq := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"}
plainLease, err := pool.Acquire(context.Background(), baseReq)
require.NoError(t, err)
plainConnID := plainLease.ConnID()
plainLease.Release()
betaReq := baseReq
betaReq.Headers = http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}}
busyBetaLease, err := pool.Acquire(context.Background(), betaReq)
require.NoError(t, err)
secondBetaLease, err := pool.Acquire(context.Background(), betaReq)
require.NoError(t, err)
require.False(t, secondBetaLease.Reused())
require.NotEqual(t, plainConnID, secondBetaLease.ConnID())
require.NotEqual(t, busyBetaLease.ConnID(), secondBetaLease.ConnID())
secondBetaLease.Release()
busyBetaLease.Release()
require.Equal(t, 3, dialer.DialCount())
}
func TestOpenAIWSConnPool_AcquireForcePreferredConnUnavailable(t *testing.T) {
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2