mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-22 06:40:21 +08:00
Merge pull request #3971 from yardbirds0/codex/fix-remote-compaction-v2
fix: 保留 remote_compaction_v2 原生 Responses 请求链路
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 到 context,apiKeyID=0。
|
||||
require.Equal(t, isolateOpenAISessionID(0, "sess-oauth-1"), captureDialer.lastHeaders.Get("session_id"))
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user