mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
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:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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