Merge remote-tracking branch 'origin/main' into feat/grok-prompt-cache-identity

# Conflicts:
#	backend/internal/handler/endpoint.go
#	backend/internal/service/openai_gateway_grok_test.go
This commit is contained in:
shaw
2026-07-13 09:12:44 +08:00
37 changed files with 2063 additions and 193 deletions
+9 -5
View File
@@ -18,6 +18,7 @@ const (
EndpointMessages = "/v1/messages"
EndpointChatCompletions = "/v1/chat/completions"
EndpointEmbeddings = "/v1/embeddings"
EndpointAlphaSearch = "/v1/alpha/search"
EndpointResponses = "/v1/responses"
EndpointResponsesCompact = "/v1/responses/compact"
EndpointImagesGenerations = "/v1/images/generations"
@@ -75,6 +76,8 @@ func NormalizeInboundEndpoint(path string) string {
switch {
case strings.Contains(path, EndpointEmbeddings):
return EndpointEmbeddings
case strings.Contains(path, EndpointAlphaSearch) || isBareOrSubpathOf(strings.TrimRight(path, "/"), "/alpha/search") || isBareOrSubpathOf(strings.TrimRight(path, "/"), "/backend-api/codex/alpha/search"):
return EndpointAlphaSearch
case strings.Contains(path, EndpointChatCompletions):
return EndpointChatCompletions
case strings.Contains(path, EndpointMessages):
@@ -155,10 +158,11 @@ func isBareOrSubpathOf(path, root string) bool {
// account platform and the normalized inbound endpoint.
//
// Platform-specific rules:
// - OpenAI and Grok default to /v1/responses (with optional subpath
// such as /v1/responses/compact preserved from the raw URL). Grok raw Chat
// requests override this through the forwarding result consumed by
// resolveOpenAIUpstreamEndpoint.
// - OpenAI and Grok text compatibility routes forward to /v1/responses
// (with optional subpath such as /v1/responses/compact preserved from
// the raw URL); native endpoints such as embeddings and alpha search
// retain their paths. Grok raw Chat requests override this through the
// forwarding result consumed by resolveOpenAIUpstreamEndpoint.
// - Anthropic → /v1/messages
// - Gemini → /v1beta/models
// - Antigravity → /v1/messages (Claude) or gemini (Gemini)
@@ -169,7 +173,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string {
switch platform {
case service.PlatformOpenAI, service.PlatformGrok:
if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos {
if inbound == EndpointEmbeddings || inbound == EndpointAlphaSearch || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos {
return inbound
}
// OpenAI forwards everything to the Responses API.
@@ -25,6 +25,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
{"/v1/messages", EndpointMessages},
{"/v1/chat/completions", EndpointChatCompletions},
{"/v1/embeddings", EndpointEmbeddings},
{"/v1/alpha/search", EndpointAlphaSearch},
{"/v1/responses", EndpointResponses},
{"/v1/responses/compact", EndpointResponsesCompact},
{"/v1/responses/compact/detail", EndpointResponsesCompact},
@@ -50,11 +51,13 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
{"/responses", EndpointResponses},
{"/responses/compact", EndpointResponsesCompact},
{"/responses/compact/detail", EndpointResponsesCompact},
{"/alpha/search", EndpointAlphaSearch},
// Bare Codex direct alias route — root vs. compact.
{"/backend-api/codex/responses", EndpointResponses},
{"/backend-api/codex/responses/compact", EndpointResponsesCompact},
{"/backend-api/codex/responses/compact/detail", EndpointResponsesCompact},
{"/backend-api/codex/alpha/search", EndpointAlphaSearch},
// Must NOT generalize to arbitrary paths merely ending in
// "/responses" (or "/responses/compact") that are unrelated to
@@ -119,6 +122,7 @@ func TestDeriveUpstreamEndpoint(t *testing.T) {
{"openai from messages", EndpointMessages, "/v1/messages", service.PlatformOpenAI, EndpointResponses},
{"openai from completions", EndpointChatCompletions, "/v1/chat/completions", service.PlatformOpenAI, EndpointResponses},
{"openai embeddings", EndpointEmbeddings, "/v1/embeddings", service.PlatformOpenAI, EndpointEmbeddings},
{"openai alpha search", EndpointAlphaSearch, "/backend-api/codex/alpha/search", service.PlatformOpenAI, EndpointAlphaSearch},
{"openai image generations", EndpointImagesGenerations, "/v1/images/generations", service.PlatformOpenAI, EndpointImagesGenerations},
{"openai image edits", EndpointImagesEdits, "/openai/v1/images/edits", service.PlatformOpenAI, EndpointImagesEdits},
{"grok chat defaults to responses without runtime result", EndpointChatCompletions, "/v1/chat/completions", service.PlatformGrok, EndpointResponses},
@@ -0,0 +1,194 @@
package handler
import (
"errors"
"net/http"
"strconv"
"strings"
"time"
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"go.uber.org/zap"
)
// AlphaSearch proxies the standalone search endpoint used by Codex Responses Lite.
func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
streamStarted := false
defer h.recoverResponsesPanic(c, &streamStarted)
setOpenAIClientTransportHTTP(c)
requestStart := time.Now()
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok || apiKey.Group == nil {
h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
return
}
if apiKey.Group.Platform != service.PlatformOpenAI {
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex alpha search is only available for OpenAI groups")
return
}
subject, ok := middleware2.GetAuthSubjectFromContext(c)
if !ok {
h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found")
return
}
reqLog := requestLogger(
c,
"handler.openai_gateway.alpha_search",
zap.Int64("user_id", subject.UserID),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
)
if !h.ensureResponsesDependencies(c, reqLog) {
return
}
body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request)
if err != nil {
if maxErr, ok := extractMaxBytesError(err); ok {
h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
return
}
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
return
}
if len(body) == 0 {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty")
return
}
if !gjson.ValidBytes(body) {
logRequestBodyParseFailure(reqLog, body, nil)
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
}
modelResult := gjson.GetBytes(body, "model")
if !modelResult.Exists() || modelResult.Type != gjson.String || strings.TrimSpace(modelResult.String()) == "" {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
return
}
requestedModel := strings.TrimSpace(modelResult.String())
reqLog = reqLog.With(zap.String("model", requestedModel))
setOpsRequestContext(c, requestedModel, false)
setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, requestedModel)
forwardBody := openAIModelMappedBody(body, channelMapping.Mapped, channelMapping.MappedModel, h.gatewayService.ReplaceModelInBody)
subscription, _ := middleware2.GetSubscriptionFromContext(c)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
userRelease, acquired := h.acquireResponsesUserSlot(c, subject.UserID, subject.Concurrency, false, &streamStarted, reqLog)
if !acquired {
return
}
if userRelease != nil {
defer userRelease()
}
if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil {
status, code, message, retryAfter := billingErrorDetails(err)
if retryAfter > 0 {
c.Header("Retry-After", strconv.Itoa(retryAfter))
}
h.errorResponse(c, status, code, message)
return
}
searchID := strings.TrimSpace(gjson.GetBytes(body, "id").String())
sessionHash := h.gatewayService.GenerateSessionHashWithFallback(c, nil, searchID)
failedAccountIDs := make(map[int64]struct{})
var lastFailoverErr *service.UpstreamFailoverError
switchCount := 0
routingStart := time.Now()
for {
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(),
apiKey.GroupID,
"",
sessionHash,
requestedModel,
failedAccountIDs,
service.OpenAIUpstreamTransportHTTPSSE,
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
service.PlatformOpenAI,
)
if err != nil || selection == nil || selection.Account == nil {
if len(failedAccountIDs) == 0 {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestedModel, requestedModel, service.PlatformOpenAI)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
return
}
if lastFailoverErr != nil {
h.handleFailoverExhausted(c, lastFailoverErr, false)
} else {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
}
return
}
account := selection.Account
setOpsSelectedAccount(c, account.ID, account.Platform)
accountRelease, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog)
if !acquired {
return
}
service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds())
writerSizeBeforeForward := c.Writer.Size()
forwardStart := time.Now()
err = func() error {
if accountRelease != nil {
defer accountRelease()
}
return h.gatewayService.ForwardAlphaSearch(c.Request.Context(), c, account, forwardBody)
}()
service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, time.Since(forwardStart).Milliseconds())
if err == nil {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil)
return
}
var failoverErr *service.UpstreamFailoverError
if !errors.As(err, &failoverErr) {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
if c.Writer.Size() == writerSizeBeforeForward {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
}
reqLog.Warn("openai_alpha_search.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err))
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
if c.Writer.Size() != writerSizeBeforeForward {
h.handleFailoverExhausted(c, failoverErr, true)
return
}
h.gatewayService.RecordOpenAIAccountSwitch()
failedAccountIDs[account.ID] = struct{}{}
lastFailoverErr = failoverErr
if switchCount >= h.maxAccountSwitches {
h.handleFailoverExhausted(c, failoverErr, false)
return
}
switchCount++
if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) {
h.handleFailoverExhausted(c, failoverErr, false)
return
}
reqLog.Warn("openai_alpha_search.upstream_failover_switching",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("switch_count", switchCount),
)
}
}
@@ -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)
@@ -1,9 +1,15 @@
package handler
import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) {
@@ -57,3 +63,28 @@ func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) {
assert.Nil(t, p)
})
}
func TestOpsCaptureWriter_CompactKeepaliveRestoresOriginalWriter(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
outerStatus := -1
router.Use(func(c *gin.Context) {
c.Next()
outerStatus = c.Writer.Status()
})
router.Use(OpsErrorLoggerMiddleware(nil))
router.GET("/compact", func(c *gin.Context) {
service.MarkOpenAICompactClientStream(c)
stop := service.StartOpenAICompactSSEKeepalive(c, time.Hour)
defer stop()
c.Status(http.StatusOK)
})
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, "/compact", nil)
require.NotPanics(t, func() {
router.ServeHTTP(recorder, request)
})
require.Equal(t, http.StatusOK, outerStatus)
require.Equal(t, http.StatusOK, recorder.Code)
}
@@ -131,3 +131,56 @@ func TestWire_UnknownEventFallsBackToDefault(t *testing.T) {
})
require.Contains(t, m, "response")
}
func TestResponsesOutputUnmarshal_ToolSearchObjectArguments(t *testing.T) {
var item ResponsesOutput
require.NoError(t, json.Unmarshal([]byte(`{
"type":"tool_search_call",
"id":"item_1",
"call_id":"call_1",
"execution":"client",
"arguments":{"query":"gmail","limit":2}
}`), &item))
require.Equal(t, "tool_search_call", item.Type)
require.Equal(t, `{"query":"gmail","limit":2}`, item.Arguments)
wire, err := json.Marshal(item)
require.NoError(t, err)
var decoded map[string]any
require.NoError(t, json.Unmarshal(wire, &decoded))
args, ok := decoded["arguments"].(map[string]any)
require.True(t, ok, "tool_search_call arguments must remain an object")
require.Equal(t, "gmail", args["query"])
}
func TestResponsesResponseUnmarshal_ToolSearchObjectArguments(t *testing.T) {
var response ResponsesResponse
require.NoError(t, json.Unmarshal([]byte(`{
"id":"response_1",
"object":"response",
"status":"completed",
"output":[{
"type":"tool_search_call",
"id":"item_1",
"call_id":"call_1",
"arguments":{"query":"gmail"}
}]
}`), &response))
require.Len(t, response.Output, 1)
require.Equal(t, `{"query":"gmail"}`, response.Output[0].Arguments)
}
func TestResponsesStreamEventUnmarshal_ToolSearchObjectArguments(t *testing.T) {
var event ResponsesStreamEvent
require.NoError(t, json.Unmarshal([]byte(`{
"type":"response.output_item.done",
"item":{
"type":"tool_search_call",
"id":"item_1",
"call_id":"call_1",
"arguments":{"query":"gmail"}
}
}`), &event))
require.NotNil(t, event.Item)
require.Equal(t, `{"query":"gmail"}`, event.Item.Arguments)
}
+50
View File
@@ -353,6 +353,56 @@ func (o ResponsesOutput) MarshalJSON() ([]byte, error) {
return json.Marshal(m)
}
// UnmarshalJSON accepts both the Responses function-call string form and the
// tool_search_call object form for arguments. The bridge stores arguments as a
// string internally, so object arguments are retained as their raw JSON.
func (o *ResponsesOutput) UnmarshalJSON(data []byte) error {
type responsesOutputAlias ResponsesOutput
var kind struct {
Type string `json:"type"`
}
if err := json.Unmarshal(data, &kind); err != nil {
return err
}
if kind.Type != "tool_search_call" {
var decoded responsesOutputAlias
if err := json.Unmarshal(data, &decoded); err != nil {
return err
}
*o = ResponsesOutput(decoded)
return nil
}
var fields map[string]json.RawMessage
if err := json.Unmarshal(data, &fields); err != nil {
return err
}
arguments, hasArguments := fields["arguments"]
delete(fields, "arguments")
normalized, err := json.Marshal(fields)
if err != nil {
return err
}
var decoded responsesOutputAlias
if err := json.Unmarshal(normalized, &decoded); err != nil {
return err
}
*o = ResponsesOutput(decoded)
if !hasArguments || string(arguments) == "null" {
return nil
}
var argumentString string
if err := json.Unmarshal(arguments, &argumentString); err == nil {
o.Arguments = argumentString
} else {
o.Arguments = string(arguments)
}
return nil
}
// WebSearchAction describes the search action in a web_search_call output item.
type WebSearchAction struct {
Type string `json:"type,omitempty"` // "search"
@@ -147,6 +147,7 @@ func RegisterGatewayRoutes(
}
h.Gateway.Responses(c)
})
gateway.POST("/alpha/search", h.OpenAIGateway.AlphaSearch)
gateway.GET("/responses", func(c *gin.Context) {
h.OpenAIGateway.ResponsesWebSocket(c)
})
@@ -212,6 +213,7 @@ func RegisterGatewayRoutes(
}
r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler)
r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler)
r.POST("/alpha/search", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.AlphaSearch)
r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
h.OpenAIGateway.ResponsesWebSocket(c)
})
@@ -220,6 +222,7 @@ func RegisterGatewayRoutes(
{
codexDirect.POST("/responses", responsesHandler)
codexDirect.POST("/responses/*subpath", responsesHandler)
codexDirect.POST("/alpha/search", h.OpenAIGateway.AlphaSearch)
codexDirect.GET("/responses", func(c *gin.Context) {
h.OpenAIGateway.ResponsesWebSocket(c)
})
@@ -65,6 +65,36 @@ func TestGatewayRoutesOpenAIResponsesCompactPathIsRegistered(t *testing.T) {
}
}
func TestGatewayRoutesOpenAIAlphaSearchPathsAreRegistered(t *testing.T) {
router := newGatewayRoutesTestRouter()
registered := make(map[string]bool)
for _, route := range router.Routes() {
if route.Method == http.MethodPost {
registered[route.Path] = true
}
}
for _, path := range []string{
"/v1/alpha/search",
"/alpha/search",
"/backend-api/codex/alpha/search",
} {
require.True(t, registered[path], "POST %s should be registered", path)
}
}
func TestGatewayRoutesAlphaSearchRejectsNonOpenAIGroup(t *testing.T) {
router := newGatewayRoutesTestRouter(service.PlatformGrok)
req := httptest.NewRequest(http.MethodPost, "/v1/alpha/search", strings.NewReader(`{"model":"gpt-5.6-sol"}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNotFound, w.Code)
require.Contains(t, w.Body.String(), "only available for OpenAI groups")
}
func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) {
router := newGatewayRoutesTestRouter()
@@ -0,0 +1,151 @@
package service
import (
"bytes"
"context"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
)
const (
chatgptCodexAlphaSearchURL = "https://chatgpt.com/backend-api/codex/alpha/search"
openAIPlatformAlphaSearchURL = "https://api.openai.com/v1/alpha/search"
)
// ForwardAlphaSearch proxies Codex standalone web search without binding the
// evolving alpha request or response schema.
func (s *OpenAIGatewayService) ForwardAlphaSearch(ctx context.Context, c *gin.Context, account *Account, body []byte) error {
if s == nil || c == nil || account == nil {
return fmt.Errorf("service, context, and account are required")
}
modelResult := gjson.GetBytes(body, "model")
requestedModel := strings.TrimSpace(modelResult.String())
if modelResult.Type != gjson.String || requestedModel == "" {
return fmt.Errorf("model is required")
}
upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(requestedModel))
if upstreamModel != "" && upstreamModel != requestedModel {
body = ReplaceModelInBody(body, upstreamModel)
}
token, _, err := s.GetAccessToken(ctx, account)
if err != nil {
return err
}
req, err := s.buildOpenAIAlphaSearchRequest(ctx, c, account, body, token)
if err != nil {
return err
}
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
}
upstreamStart := time.Now()
resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency)
SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds())
if err != nil {
return s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true)
}
defer func() { _ = resp.Body.Close() }()
respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError)
if err != nil {
return fmt.Errorf("read alpha search response: %w", err)
}
if resp.StatusCode >= http.StatusBadRequest {
upstreamMessage := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMessage, respBody) {
resp.Body = io.NopCloser(bytes.NewReader(respBody))
s.handleFailoverSideEffects(ctx, resp, account, respBody, upstreamModel)
return &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
}
if !account.IsShadow() {
s.UpdateCodexUsageSnapshotFromHeaders(ctx, account.ID, resp.Header)
}
writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
contentType := resp.Header.Get("Content-Type")
if contentType == "" {
contentType = "application/json"
}
c.Data(resp.StatusCode, contentType, respBody)
return nil
}
func (s *OpenAIGatewayService) buildOpenAIAlphaSearchRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) {
clientBeta := ""
if c != nil {
clientBeta = c.GetHeader("OpenAI-Beta")
}
req, err := s.buildUpstreamRequestOpenAIPassthrough(ctx, c, account, body, token)
if err != nil {
return nil, err
}
targetURL, err := s.openAIAlphaSearchURL(account)
if err != nil {
return nil, err
}
parsedURL, err := url.Parse(targetURL)
if err != nil {
return nil, fmt.Errorf("parse alpha search URL: %w", err)
}
if c != nil && c.Request != nil && c.Request.URL != nil {
query := parsedURL.Query()
for key, values := range c.Request.URL.Query() {
for _, value := range values {
query.Add(key, value)
}
}
parsedURL.RawQuery = query.Encode()
}
req.URL = parsedURL
req.Header.Set("Accept", "application/json")
if clientBeta == "" {
req.Header.Del("OpenAI-Beta")
}
if version := strings.TrimSpace(c.GetHeader("Version")); version != "" {
req.Header.Set("Version", version)
} else if account.Type == AccountTypeOAuth {
req.Header.Set("Version", codexCLIVersion)
}
return req, nil
}
func (s *OpenAIGatewayService) openAIAlphaSearchURL(account *Account) (string, error) {
if account == nil {
return "", fmt.Errorf("account is required")
}
switch account.Type {
case AccountTypeOAuth:
return chatgptCodexAlphaSearchURL, nil
case AccountTypeAPIKey:
baseURL := account.GetOpenAIBaseURL()
if baseURL == "" {
return openAIPlatformAlphaSearchURL, nil
}
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
return "", err
}
return buildOpenAIEndpointURL(validatedURL, "/v1/alpha/search"), nil
default:
return "", fmt.Errorf("unsupported OpenAI account type: %s", account.Type)
}
}
@@ -0,0 +1,139 @@
package service
import (
"bytes"
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestForwardAlphaSearchOAuthPreservesWire(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{
"id":"search-session",
"model":"gpt-5.6-sol",
"reasoning":{"effort":"max","context":"all_turns"},
"input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"latest news"}]}],
"commands":{"search_query":[{"q":"OpenAI news","recency":1}]},
"settings":{"allowed_callers":["direct"],"external_web_access":true},
"max_output_tokens":2000,
"future_field":{"keep":true}
}`)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search?feature=standalone", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Header.Set("User-Agent", codexCLIUserAgent)
c.Request.Header.Set("Originator", "codex_cli_rs")
c.Request.Header.Set("Version", "0.144.1")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"encrypted_output":"ciphertext","output":"search result"}`)),
}}
service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
account := &Account{
ID: 42,
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "oauth-token",
"chatgpt_account_id": "chatgpt-account",
},
}
err := service.ForwardAlphaSearch(context.Background(), c, account, body)
require.NoError(t, err)
require.Equal(t, http.StatusOK, recorder.Code)
require.JSONEq(t, `{"encrypted_output":"ciphertext","output":"search result"}`, recorder.Body.String())
require.Equal(t, chatgptCodexAlphaSearchURL+"?feature=standalone", upstream.lastReq.URL.String())
require.Equal(t, "chatgpt.com", upstream.lastReq.Host)
require.Equal(t, "Bearer oauth-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "chatgpt-account", upstream.lastReq.Header.Get("chatgpt-account-id"))
require.Equal(t, "application/json", upstream.lastReq.Header.Get("Accept"))
require.Equal(t, "0.144.1", upstream.lastReq.Header.Get("Version"))
require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta"))
require.JSONEq(t, string(body), string(upstream.lastBody))
}
func TestForwardAlphaSearchAPIKeyMapsModelAndPassesThroughError(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"id":"search-session","model":"gpt-5.6-sol","commands":{"search_query":[{"q":"news"}]}}`)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/alpha/search", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
upstreamBody := `{"error":{"type":"invalid_request_error","message":"bad search"}}`
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusBadRequest,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(upstreamBody)),
}}
service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
account := &Account{
ID: 7,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "https://compat.example/v4",
"model_mapping": map[string]any{
"gpt-5.6-sol": "upstream-5.6",
},
},
}
err := service.ForwardAlphaSearch(context.Background(), c, account, body)
require.NoError(t, err)
require.Equal(t, http.StatusBadRequest, recorder.Code)
require.JSONEq(t, upstreamBody, recorder.Body.String())
require.Equal(t, "https://compat.example/v4/alpha/search", upstream.lastReq.URL.String())
require.Equal(t, "Bearer sk-test", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "upstream-5.6", gjson.GetBytes(upstream.lastBody, "model").String())
require.True(t, gjson.GetBytes(upstream.lastBody, "commands.search_query").IsArray())
}
func TestForwardAlphaSearchReturnsFailoverBeforeWriting(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"id":"search-session","model":"gpt-5.6-sol","commands":{}}`)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search", bytes.NewReader(body))
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusTooManyRequests,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)),
}}
service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
account := &Account{
ID: 8,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Credentials: map[string]any{
"api_key": "sk-test",
},
}
err := service.ForwardAlphaSearch(context.Background(), c, account, body)
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode)
require.Equal(t, openAIPlatformAlphaSearchURL, upstream.lastReq.URL.String())
require.False(t, c.Writer.Written())
require.Empty(t, recorder.Body.String())
}
@@ -114,14 +114,15 @@ func TestFilterCodexInput_OutputTypeKeepsItemID(t *testing.T) {
require.Equal(t, "o1", out["id"], "output item id should be preserved")
}
// TestFilterCodexInput_NonToolCallItemKeepsID ensures non-tool-call items
// (e.g. message) still keep their id when PreserveReferences is true.
// TestFilterCodexInput_NonToolCallItemKeepsID ensures items subject to neither
// the fc* (call-input) nor the msg* (message) prefix rule still keep their id
// when PreserveReferences is true.
// message is covered separately in openai_codex_message_item_id_test.go (#3981).
func TestFilterCodexInput_NonToolCallItemKeepsID(t *testing.T) {
input := []any{
map[string]any{
"type": "message",
"id": "item_msg_001",
"role": "user",
"type": "web_search_call",
"id": "ws_001",
},
}
@@ -130,7 +131,7 @@ func TestFilterCodexInput_NonToolCallItemKeepsID(t *testing.T) {
})
require.Len(t, filtered, 1)
msg, ok := filtered[0].(map[string]any)
item, ok := filtered[0].(map[string]any)
require.True(t, ok)
require.Equal(t, "item_msg_001", msg["id"], "non-tool-call items keep their id in preserve mode")
require.Equal(t, "ws_001", item["id"], "unconstrained items keep their id in preserve mode")
}
@@ -11,12 +11,32 @@ import (
// 若请求携带 version 且低于该值,上游直接 404(issue #3901,2026-07 实测)。
const codexUpstreamMinVersion = "0.144.0"
// ensureCodexIdentityHeaders 补齐 OAuth(ChatGPT 内部接口)出站请求所需的 Codex 身份头。
// 已有 User-Agent 与 version 保持不变,交给紧随其后的 enforceCodexIdentityHeaders
// 做官方身份配对与最低版本校正。
func ensureCodexIdentityHeaders(h http.Header) {
if h == nil {
return
}
if strings.TrimSpace(h.Get("user-agent")) == "" {
h.Set("user-agent", codexCLIUserAgent)
}
if strings.TrimSpace(h.Get("originator")) == "" {
h.Set("originator", "codex_cli_rs")
}
if strings.TrimSpace(h.Get("version")) == "" {
h.Set("version", codexCLIVersion)
}
h.Set("OpenAI-Beta", "responses=experimental")
}
// enforceCodexIdentityHeaders 收口 OAuth(ChatGPT 内部接口)出站请求的客户端身份头。
// 上游要求 originator 与 User-Agent 首段配套且为官方客户端标识,version 头(若携带)
// 不低于 0.144.0,任一不满足即 404(issue #3901)。以最终 User-Agent 为准推导配套
// originator;推导不出官方身份(第三方 UA / UA 缺失)时整体回退为默认 Codex CLI 身份。
//
// 仅对携带 originator 的请求生效——compat messages bridge 故意不带 originator,保持原样。
// 仅对携带 originator 的请求生效;需要从缺失身份头恢复的调用方应先调用
// ensureCodexIdentityHeaders。
// 必须在所有 User-Agent 改写(自定义 UA / ForceCodexCLI / 浏览器 UA 兜底)之后调用。
func enforceCodexIdentityHeaders(h http.Header) {
if h == nil || h.Get("originator") == "" {
@@ -7,6 +7,36 @@ import (
"github.com/stretchr/testify/require"
)
func TestEnsureCodexIdentityHeaders(t *testing.T) {
t.Run("补齐缺失身份头", func(t *testing.T) {
h := make(http.Header)
ensureCodexIdentityHeaders(h)
enforceCodexIdentityHeaders(h)
require.Equal(t, "codex_cli_rs", h.Get("originator"))
require.Equal(t, codexCLIUserAgent, h.Get("user-agent"))
require.Equal(t, codexCLIVersion, h.Get("version"))
require.Equal(t, "responses=experimental", h.Get("OpenAI-Beta"))
})
t.Run("保留已有官方UA和合法version并重新配对", func(t *testing.T) {
const tuiUA = "codex-tui/9.9.9 (Mac OS X 14.0; arm64) iTerm (codex-tui; 9.9.9)"
h := make(http.Header)
h.Set("user-agent", tuiUA)
h.Set("version", "9.9.9")
h.Set("OpenAI-Beta", "assistants=v2")
ensureCodexIdentityHeaders(h)
enforceCodexIdentityHeaders(h)
require.Equal(t, "codex-tui", h.Get("originator"))
require.Equal(t, tuiUA, h.Get("user-agent"))
require.Equal(t, "9.9.9", h.Get("version"))
require.Equal(t, "responses=experimental", h.Get("OpenAI-Beta"))
})
}
func TestEnforceCodexIdentityHeaders(t *testing.T) {
const tuiUA = "codex-tui/0.140.2 (Mac OS X 14.0; arm64) iTerm (codex-tui; 0.140.2)"
@@ -102,13 +132,14 @@ func TestEnforceCodexIdentityHeaders(t *testing.T) {
}
}
// compat messages bridge 故意不带 originator:收口必须保持 no-op,不得注入身份头。
// enforce 本身仍只负责收口:缺少 originator 时必须保持 no-op,由需要恢复身份的
// 调用方先显式调用 ensureCodexIdentityHeaders。
func TestEnforceCodexIdentityHeaders_NoOriginatorIsNoop(t *testing.T) {
h := make(http.Header)
h.Set("user-agent", "luna/1.0.0")
h.Set("user-agent", "third-party-client/1.0.0")
enforceCodexIdentityHeaders(h)
require.Empty(t, h.Get("originator"))
require.Equal(t, "luna/1.0.0", h.Get("user-agent"))
require.Equal(t, "third-party-client/1.0.0", h.Get("user-agent"))
}
@@ -0,0 +1,160 @@
//go:build unit
package service
import (
"testing"
"github.com/stretchr/testify/require"
)
// TestFilterCodexInput_StripsMessageItemID_WhenPreservingReferences
// verifies that message items with a non-msg id (e.g. item_*) have their id
// stripped even when PreserveReferences is true. OpenAI upstream requires
// message ids to begin with "msg" and rejects item_* with 400:
// "Expected an ID that begins with 'msg'." (#3981)
func TestFilterCodexInput_StripsMessageItemID_WhenPreservingReferences(t *testing.T) {
input := []any{
map[string]any{
"type": "message",
"id": "item_3bc5a3fa8ccde25f1c0000d4",
"role": "user",
"content": []any{
map[string]any{"type": "input_text", "text": "hello"},
},
},
}
filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{
PreserveReferences: true,
})
require.Len(t, filtered, 1)
msg, ok := filtered[0].(map[string]any)
require.True(t, ok)
require.Equal(t, "message", msg["type"])
_, hasID := msg["id"]
require.False(t, hasID, "item_* id should be stripped from message")
require.Equal(t, "user", msg["role"], "role must be preserved")
require.NotNil(t, msg["content"], "content must be preserved")
}
// TestFilterCodexInput_KeepsMsgID_WhenPreservingReferences
// verifies that message items with a valid msg* id are kept when
// PreserveReferences is true, so context references are not lost.
func TestFilterCodexInput_KeepsMsgID_WhenPreservingReferences(t *testing.T) {
input := []any{
map[string]any{
"type": "message",
"id": "msg_validID123",
"role": "assistant",
},
}
filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{
PreserveReferences: true,
})
require.Len(t, filtered, 1)
msg, ok := filtered[0].(map[string]any)
require.True(t, ok)
require.Equal(t, "msg_validID123", msg["id"], "valid msg* id must be preserved")
}
// TestFilterCodexInput_StripsMessageIDWhenNotPreservingReferences ensures the
// non-continuation path still drops every message id regardless of prefix.
func TestFilterCodexInput_StripsMessageIDWhenNotPreservingReferences(t *testing.T) {
for _, id := range []string{"item_abc", "msg_validID123"} {
input := []any{
map[string]any{
"type": "message",
"id": id,
"role": "user",
},
}
filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{
PreserveReferences: false,
})
require.Len(t, filtered, 1)
msg, ok := filtered[0].(map[string]any)
require.True(t, ok)
_, hasID := msg["id"]
require.False(t, hasID, "id %q should be stripped when not preserving references", id)
}
}
// TestFilterCodexInput_MessageIDStripDoesNotMutateInput ensures the original
// input map is not modified in place when the id is stripped.
func TestFilterCodexInput_MessageIDStripDoesNotMutateInput(t *testing.T) {
original := map[string]any{
"type": "message",
"id": "item_abc",
"role": "user",
}
filtered := filterCodexInputWithOptions([]any{original}, codexInputFilterOptions{
PreserveReferences: true,
})
require.Len(t, filtered, 1)
require.Equal(t, "item_abc", original["id"], "original input must not be mutated")
}
// TestFilterCodexInput_MessageStripKeepsFunctionCallBehavior guards against a
// regression of #3785: message and function_call id rules are independent.
func TestFilterCodexInput_MessageStripKeepsFunctionCallBehavior(t *testing.T) {
input := []any{
map[string]any{
"type": "message",
"id": "item_msg_001",
"role": "user",
},
map[string]any{
"type": "function_call",
"id": "fc_validID123",
"call_id": "fc_validID123",
"name": "bash",
},
map[string]any{
"type": "function_call",
"id": "item_A9v0SNfS3VaLrfX0j3y4xhyK",
"call_id": "fc_abc123",
"name": "bash",
},
map[string]any{
"type": "function_call_output",
"id": "o1",
"call_id": "fc_abc123",
"output": "done",
},
}
filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{
PreserveReferences: true,
})
require.Len(t, filtered, 4)
msg, ok := filtered[0].(map[string]any)
require.True(t, ok)
_, hasID := msg["id"]
require.False(t, hasID, "message item_* id should be stripped")
fcValid, ok := filtered[1].(map[string]any)
require.True(t, ok)
require.Equal(t, "fc_validID123", fcValid["id"], "valid fc* id must be preserved")
fcBad, ok := filtered[2].(map[string]any)
require.True(t, ok)
_, hasID = fcBad["id"]
require.False(t, hasID, "function_call item_* id should still be stripped")
require.Equal(t, "fc_abc123", fcBad["call_id"], "call_id pairing must survive")
out, ok := filtered[3].(map[string]any)
require.True(t, ok)
require.Equal(t, "o1", out["id"], "output item id should be preserved")
require.Equal(t, "fc_abc123", out["call_id"], "call_id pairing must survive")
}
@@ -1405,6 +1405,15 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an
ensureCopy()
delete(newItem, "id")
}
} else if typ == "message" {
// 同理,message 类 item 的 id 必须以 "msg" 开头(上游校验
// "Expected an ID that begins with 'msg'")。item_* 形式的 id
// 来自客户端回放,需要删除。
// 注意:不改写成 msg_*,改写出的 id 未必对应真实的上游对象。
if id, ok := m["id"].(string); ok && id != "" && !strings.HasPrefix(id, "msg") {
ensureCopy()
delete(newItem, "id")
}
}
filtered = append(filtered, newItem)
@@ -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
@@ -1,6 +1,9 @@
package service
import (
"bufio"
"errors"
"net"
"net/http"
"sync"
"time"
@@ -43,12 +46,14 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func
if c == nil || c.Writer == nil || interval <= 0 || !openAICompactClientWantsStream(c) {
return func() {}
}
originalWriter := c.Writer
k := &openAICompactSSEKeepalive{
writer: c.Writer,
writer: originalWriter,
stop: make(chan struct{}),
}
c.Set(openAICompactSSEKeepaliveKey, k)
c.Writer = &openAICompactKeepaliveWriter{ResponseWriter: c.Writer, k: k}
wrappedWriter := &openAICompactKeepaliveWriter{ResponseWriter: originalWriter, k: k}
c.Writer = wrappedWriter
var reqDone <-chan struct{}
if c.Request != nil {
@@ -71,7 +76,14 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func
timer.Reset(interval)
}
}()
return k.Stop
return func() {
k.Stop()
// Do not leave a pooled middleware writer reachable through the compact
// wrapper after the request finishes.
if current, ok := c.Writer.(*openAICompactKeepaliveWriter); ok && current == wrappedWriter {
c.Writer = originalWriter
}
}
}
// beat 在锁内提交(首次)响应头并写出一条 SSE 注释行;返回 false 表示心跳已
@@ -181,52 +193,105 @@ type openAICompactKeepaliveWriter struct {
// suspend 停拍心跳;幂等。任何响应构造(含 Header 访问——写响应必先操作
// 响应头)都视为请求侧接管 ResponseWriter。
func (w *openAICompactKeepaliveWriter) suspend() {
if w.k == nil {
return
}
w.k.Stop()
}
func (w *openAICompactKeepaliveWriter) Header() http.Header {
w.suspend()
if w.ResponseWriter == nil {
return http.Header{}
}
return w.ResponseWriter.Header()
}
func (w *openAICompactKeepaliveWriter) Write(data []byte) (int, error) {
w.suspend()
if w.ResponseWriter == nil {
return 0, nil
}
return w.ResponseWriter.Write(data)
}
func (w *openAICompactKeepaliveWriter) WriteString(s string) (int, error) {
w.suspend()
if w.ResponseWriter == nil {
return 0, nil
}
return w.ResponseWriter.WriteString(s)
}
func (w *openAICompactKeepaliveWriter) WriteHeader(code int) {
w.suspend()
if w.ResponseWriter == nil {
return
}
w.ResponseWriter.WriteHeader(code)
}
func (w *openAICompactKeepaliveWriter) WriteHeaderNow() {
w.suspend()
if w.ResponseWriter == nil {
return
}
w.ResponseWriter.WriteHeaderNow()
}
func (w *openAICompactKeepaliveWriter) Flush() {
w.suspend()
if w.ResponseWriter == nil {
return
}
w.ResponseWriter.Flush()
}
func (w *openAICompactKeepaliveWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
if w.ResponseWriter == nil {
return nil, nil, errors.New("response writer released")
}
return w.ResponseWriter.Hijack()
}
func (w *openAICompactKeepaliveWriter) CloseNotify() <-chan bool {
if w.ResponseWriter == nil {
ch := make(chan bool)
close(ch)
return ch
}
return w.ResponseWriter.CloseNotify()
}
func (w *openAICompactKeepaliveWriter) Pusher() http.Pusher {
if w.ResponseWriter == nil {
return nil
}
return w.ResponseWriter.Pusher()
}
func (w *openAICompactKeepaliveWriter) Status() int {
if w.k == nil || w.ResponseWriter == nil {
return 0
}
w.k.mu.Lock()
defer w.k.mu.Unlock()
return w.ResponseWriter.Status()
}
func (w *openAICompactKeepaliveWriter) Size() int {
if w.k == nil || w.ResponseWriter == nil {
return 0
}
w.k.mu.Lock()
defer w.k.mu.Unlock()
return w.ResponseWriter.Size()
}
func (w *openAICompactKeepaliveWriter) Written() bool {
if w.k == nil || w.ResponseWriter == nil {
return false
}
w.k.mu.Lock()
defer w.k.mu.Unlock()
return w.ResponseWriter.Written()
@@ -6,6 +6,7 @@ import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
@@ -141,6 +142,110 @@ func TestOpenAICompactKeepaliveWriter_RequestSideWriteSuspendsBeats(t *testing.T
require.Contains(t, rec.Body.String(), `{"error":"local reject"}`)
}
func TestOpenAICompactKeepaliveWriter_NilInnerWriter_NoPanic(t *testing.T) {
w := &openAICompactKeepaliveWriter{
k: &openAICompactSSEKeepalive{stop: make(chan struct{})},
}
w.ResponseWriter = nil
assert.NotPanics(t, func() {
assert.Equal(t, 0, w.Status())
})
assert.NotPanics(t, func() {
assert.Equal(t, 0, w.Size())
})
assert.NotPanics(t, func() {
assert.False(t, w.Written())
})
assert.NotPanics(t, func() {
assert.NotNil(t, w.Header())
})
assert.NotPanics(t, func() {
n, err := w.Write([]byte("test"))
assert.Equal(t, 0, n)
assert.NoError(t, err)
})
assert.NotPanics(t, func() {
n, err := w.WriteString("test")
assert.Equal(t, 0, n)
assert.NoError(t, err)
})
assert.NotPanics(t, func() {
w.WriteHeader(http.StatusOK)
})
assert.NotPanics(t, func() {
w.WriteHeaderNow()
})
assert.NotPanics(t, func() {
w.Flush()
})
assert.NotPanics(t, func() {
conn, rw, err := w.Hijack()
assert.Nil(t, conn)
assert.Nil(t, rw)
assert.Error(t, err)
})
assert.NotPanics(t, func() {
ch := w.CloseNotify()
assert.NotNil(t, ch)
})
assert.NotPanics(t, func() {
assert.Nil(t, w.Pusher())
})
}
func TestOpenAICompactKeepaliveWriter_NilKeepalive_NoPanic(t *testing.T) {
c, rec := newCompactBridgeTestContext(t, true)
w := &openAICompactKeepaliveWriter{ResponseWriter: c.Writer}
assert.NotPanics(t, func() {
assert.Equal(t, 0, w.Status())
})
assert.NotPanics(t, func() {
assert.Equal(t, 0, w.Size())
})
assert.NotPanics(t, func() {
assert.False(t, w.Written())
})
assert.NotPanics(t, func() {
w.Header().Set("X-Test", "ok")
})
assert.NotPanics(t, func() {
w.WriteHeader(http.StatusAccepted)
})
assert.NotPanics(t, func() {
n, err := w.WriteString("ok")
assert.Equal(t, 2, n)
assert.NoError(t, err)
})
assert.NotPanics(t, func() {
w.Flush()
})
require.Equal(t, "ok", rec.Header().Get("X-Test"))
require.Equal(t, "ok", rec.Body.String())
}
func TestOpenAICompactKeepaliveWriter_DelegatesWhenReady(t *testing.T) {
c, rec := newCompactBridgeTestContext(t, true)
stop := StartOpenAICompactSSEKeepalive(c, time.Hour)
defer stop()
w, ok := c.Writer.(*openAICompactKeepaliveWriter)
require.True(t, ok)
w.Header().Set("X-Test", "ok")
w.WriteHeader(http.StatusAccepted)
n, err := w.WriteString("ready")
require.NoError(t, err)
require.Equal(t, len("ready"), n)
require.Equal(t, http.StatusAccepted, w.Status())
require.Equal(t, len("ready"), w.Size())
require.True(t, w.Written())
require.Equal(t, "ok", rec.Header().Get("X-Test"))
require.Equal(t, "ready", rec.Body.String())
}
// fast policy block 在心跳提交后必须降级为 response.failed 终止事件。
func TestWriteOpenAIFastPolicyBlockedResponse_AfterKeepaliveCommit(t *testing.T) {
c, rec := newCompactBridgeTestContext(t, true)
@@ -837,8 +837,7 @@ func TestForwardAsAnthropic_ReusesOAuthCodexTurnState(t *testing.T) {
require.NoError(t, err)
require.NotNil(t, firstResult)
require.Empty(t, upstream.requests[0].Header.Get("x-codex-turn-state"))
require.Empty(t, upstream.requests[0].Header.Get("OpenAI-Beta"))
require.Empty(t, upstream.requests[0].Header.Get("originator"))
requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs")
secondBody := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"first"},{"role":"assistant","content":"ok"},{"role":"user","content":"second"}],"stream":false}`)
secondRec := httptest.NewRecorder()
@@ -852,12 +851,73 @@ func TestForwardAsAnthropic_ReusesOAuthCodexTurnState(t *testing.T) {
require.Equal(t, "turn_state_first", upstream.requests[1].Header.Get("x-codex-turn-state"))
require.Equal(t, generateSessionUUID(isolateOpenAISessionID(0, "stable-cache-key")), upstream.requests[1].Header.Get("session_id"))
require.Empty(t, upstream.requests[1].Header.Get("conversation_id"))
require.Empty(t, upstream.requests[1].Header.Get("OpenAI-Beta"))
require.Empty(t, upstream.requests[1].Header.Get("originator"))
requireOpenAIMessagesCodexIdentity(t, upstream.requests[1], codexCLIUserAgent, "codex_cli_rs")
require.False(t, gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").Exists())
require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists())
}
func TestForwardAsAnthropic_OAuthRestoresCodexIdentityHeaders(t *testing.T) {
gin.SetMode(gin.TestMode)
const tuiUA = "codex-tui/9.9.9 (Mac OS X 14.0; arm64) iTerm (codex-tui; 9.9.9)"
tests := []struct {
name string
userAgent string
originator string
wantUserAgent string
wantOriginator string
}{
{
name: "官方UA逐字保留并重新配对",
userAgent: tuiUA,
originator: "opencode",
wantUserAgent: tuiUA,
wantOriginator: "codex-tui",
},
{
name: "第三方UA回退为默认Codex身份",
userAgent: "third-party-client/1.0.0",
originator: "opencode",
wantUserAgent: codexCLIUserAgent,
wantOriginator: "codex_cli_rs",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
body := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"hello"}],"stream":false}`)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Header.Set("User-Agent", tt.userAgent)
c.Request.Header.Set("originator", tt.originator)
upstream := &httpUpstreamRecorder{resp: openAICompatSSECompletedResponse("resp_identity", "gpt-5.4")}
svc := &OpenAIGatewayService{
httpUpstream: upstream,
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
}
account := &Account{
ID: 1,
Name: "openai-oauth",
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "oauth-token",
"chatgpt_account_id": "chatgpt-acc",
},
}
result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "gpt-5.4")
require.NoError(t, err)
require.NotNil(t, result)
requireOpenAIMessagesCodexIdentity(t, upstream.lastReq, tt.wantUserAgent, tt.wantOriginator)
})
}
}
func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey(t *testing.T) {
t.Parallel()
gin.SetMode(gin.TestMode)
@@ -896,6 +956,7 @@ func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey
firstSessionID := upstream.requests[0].Header.Get("session_id")
require.NotEmpty(t, firstSessionID)
require.Empty(t, upstream.requests[0].Header.Get("x-codex-turn-state"))
requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs")
require.False(t, gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").Exists())
secondBody := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"first"},{"role":"assistant","content":"ok"},{"role":"user","content":"second"}],"stream":false}`)
@@ -910,6 +971,7 @@ func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey
require.Equal(t, firstSessionID, upstream.requests[1].Header.Get("session_id"))
require.Equal(t, "turn_state_digest_first", upstream.requests[1].Header.Get("x-codex-turn-state"))
require.Empty(t, upstream.requests[1].Header.Get("conversation_id"))
requireOpenAIMessagesCodexIdentity(t, upstream.requests[1], codexCLIUserAgent, "codex_cli_rs")
require.False(t, gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").Exists())
require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists())
}
@@ -1064,8 +1126,7 @@ func TestForwardAsAnthropic_OAuthKeepsSystemAsDeveloperInput(t *testing.T) {
instructions := gjson.GetBytes(upstream.lastBody, "instructions")
require.True(t, instructions.Exists())
require.Empty(t, instructions.String())
require.Empty(t, upstream.requests[0].Header.Get("OpenAI-Beta"))
require.Empty(t, upstream.requests[0].Header.Get("originator"))
requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs")
}
func TestForwardAsAnthropic_OAuthAddsClaudeCodeTodoGuardForCompatModel(t *testing.T) {
@@ -1202,6 +1263,15 @@ func openAICompatSSECompletedResponse(responseID, model string) *http.Response {
}
}
func requireOpenAIMessagesCodexIdentity(t *testing.T, req *http.Request, wantUserAgent, wantOriginator string) {
t.Helper()
require.NotNil(t, req)
require.Equal(t, wantUserAgent, req.Header.Get("User-Agent"))
require.Equal(t, wantOriginator, req.Header.Get("originator"))
require.Equal(t, codexCLIVersion, req.Header.Get("version"))
require.Equal(t, "responses=experimental", req.Header.Get("OpenAI-Beta"))
}
func openAICompatSSEResponseWithoutUsage(responseID, model string) *http.Response {
body := strings.Join([]string{
`data: {"type":"response.completed","response":{"id":"` + responseID + `","object":"response","model":"` + model + `","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}]}}`,
@@ -1034,6 +1034,8 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
body := []byte(`{"model":"grok","max_tokens":32,"stream":false,"messages":[{"role":"user","content":"hi"}]}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
c.Set("api_key", &APIKey{ID: 5401})
c.Request.Header.Set("OpenAI-Beta", "grok-experimental")
c.Request.Header.Set("originator", "opencode")
account := &Account{
ID: 54,
@@ -1065,6 +1067,9 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent"))
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.Equal(t, "grok-experimental", upstream.lastReq.Header.Get("OpenAI-Beta"))
require.Empty(t, upstream.lastReq.Header.Get("originator"))
require.Empty(t, upstream.lastReq.Header.Get("version"))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String())
require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader))
@@ -266,9 +266,8 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
// 6. Build upstream request
if account.Type == AccountTypeOAuth && account.Platform != PlatformGrok {
// Messages 兼容桥即使 body 未带 todo-guard/prompt_cache_key 标记(如映射到非
// gpt-5/codex 模型),也必须让 buildUpstreamRequest 走 bridge 分支:不带
// originator、User-Agent 逐字透传,避免身份收口(issue #3901)误改本路径
// 刻意最小化的请求形态(下方的 Del(OpenAI-Beta/originator) 兜底保持不变)。
// gpt-5/codex 模型),也必须让 buildUpstreamRequest 走 bridge 分支,以保留
// 既有 body/session/conversation 行为。身份头在 post-build 阶段统一恢复。
setOpenAICompatMessagesBridgeContext(c, true)
}
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
@@ -293,12 +292,16 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
}
}
if account.Type == AccountTypeOAuth && account.Platform != PlatformGrok {
// Anthropic Messages compatibility uses the ChatGPT Codex SSE endpoint.
// Match airgate-openai's request shape: the SSE endpoint does not need
// the Responses experimental beta header, and forcing originator can make
// ChatGPT select a different internal continuation path.
upstreamReq.Header.Del("OpenAI-Beta")
upstreamReq.Header.Del("originator")
// buildUpstreamRequest 保留 Messages bridge 的 body/session 兼容行为,并会先
// 清除身份头。真正发送前恢复完整 Codex 身份,避免 ChatGPT Codex 上游因缺失
// originator/OpenAI-Beta 返回 404(issue #3901)。
ensureCodexIdentityHeaders(upstreamReq.Header)
enforceCodexIdentityHeaders(upstreamReq.Header)
logger.L().Debug("openai messages: upstream identity restored",
zap.Int64("account_id", account.ID),
zap.String("upstream_model", upstreamModel),
zap.Bool("compat_identity_restored", true),
)
}
if account.Type == AccountTypeOAuth && promptCacheKey != "" && strings.TrimSpace(c.GetHeader("conversation_id")) == "" {
upstreamReq.Header.Del("conversation_id")
@@ -356,6 +356,8 @@ func TestForwardAsAnthropic_ResponsesSupportedAccountStillUsesResponsesEndpoint(
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Header.Set("User-Agent", "third-party-client/1.0.0")
c.Request.Header.Set("originator", "opencode")
upstreamBody := strings.Join([]string{
`data: {"type":"response.completed","response":{"id":"resp_native","object":"response","model":"gpt-5.4","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}`,
@@ -385,5 +387,9 @@ func TestForwardAsAnthropic_ResponsesSupportedAccountStillUsesResponsesEndpoint(
"responses-capable account must stay on /v1/responses, got %s", upstream.lastReq.URL.String())
require.True(t, gjson.GetBytes(upstream.lastBody, "input").Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "messages").Exists())
require.Equal(t, "third-party-client/1.0.0", upstream.lastReq.Header.Get("User-Agent"))
require.Equal(t, "opencode", upstream.lastReq.Header.Get("originator"))
require.Empty(t, upstream.lastReq.Header.Get("version"))
require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta"))
require.Equal(t, "ok", gjson.Get(rec.Body.String(), "content.0.text").String())
}
@@ -67,6 +67,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,
}
@@ -82,6 +83,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"))
+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