diff --git a/backend/internal/handler/openai_gateway_compact_body_signal_test.go b/backend/internal/handler/openai_gateway_compact_body_signal_test.go index a4bfb90466..a44d47c856 100644 --- a/backend/internal/handler/openai_gateway_compact_body_signal_test.go +++ b/backend/internal/handler/openai_gateway_compact_body_signal_test.go @@ -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) } -// 回归 #3875:body-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 compact(Codex v1 unary 协议)即使 body 带 stream:true 也不标记, // 保持 JSON 写回行为不变。 func TestNormalizeOpenAIResponsesCompactRequest_PathBasedStreamTrueNotMarked(t *testing.T) { diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 83e644d857..a4d7ee7b18 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -580,21 +580,33 @@ func isBareOpenAIResponsesPath(c *gin.Context) bool { return strings.HasSuffix(normalizedPath, "/responses") } -// normalizeOpenAIResponsesCompactRequest 统一处理两种入站 compact 形态: -// path-based(POST /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) diff --git a/backend/internal/service/openai_compact_body_signal.go b/backend/internal/service/openai_compact_body_signal.go index fce62046c1..ce561b0c5a 100644 --- a/backend/internal/service/openai_compact_body_signal.go +++ b/backend/internal/service/openai_compact_body_signal.go @@ -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 diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 0d03d83dcf..c3b1b996f1 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -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, } diff --git a/backend/internal/service/openai_gpt56_max_test.go b/backend/internal/service/openai_gpt56_max_test.go index 272eb16ff0..cbca2ff3ee 100644 --- a/backend/internal/service/openai_gpt56_max_test.go +++ b/backend/internal/service/openai_gpt56_max_test.go @@ -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) } diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go index de8ecf030b..60790b9e4c 100644 --- a/backend/internal/service/openai_oauth_passthrough_test.go +++ b/backend/internal/service/openai_oauth_passthrough_test.go @@ -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")) } diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index a4d47218e7..5830444815 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -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 { diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index adae109e09..bb4ac2242c 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -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 到 context,apiKeyID=0。 require.Equal(t, isolateOpenAISessionID(0, "sess-oauth-1"), captureDialer.lastHeaders.Get("session_id")) diff --git a/backend/internal/service/openai_ws_pool.go b/backend/internal/service/openai_ws_pool.go index 5950e02841..329908e762 100644 --- a/backend/internal/service/openai_ws_pool.go +++ b/backend/internal/service/openai_ws_pool.go @@ -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 diff --git a/backend/internal/service/openai_ws_pool_test.go b/backend/internal/service/openai_ws_pool_test.go index b2683ee041..ae9b94ce4a 100644 --- a/backend/internal/service/openai_ws_pool_test.go +++ b/backend/internal/service/openai_ws_pool_test.go @@ -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