mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:18:25 +08:00
fix(openai): send OAuth routing hints
This commit is contained in:
@@ -1110,6 +1110,7 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin.
|
||||
|
||||
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)
|
||||
account.ApplyHeaderOverrides(req.Header)
|
||||
setOpenAICodexRoutingHintFromBody(req.Header, account, body)
|
||||
|
||||
return req, nil
|
||||
}
|
||||
|
||||
@@ -453,6 +453,7 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough(
|
||||
|
||||
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)
|
||||
account.ApplyHeaderOverrides(req.Header)
|
||||
setOpenAICodexRoutingHintFromBody(req.Header, account, body)
|
||||
|
||||
return req, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/tidwall/gjson"
|
||||
"golang.org/x/net/http/httpguts"
|
||||
)
|
||||
|
||||
const openAICodexRoutingHintHeader = "x-codex-routing-hint"
|
||||
|
||||
// setOpenAICodexRoutingHint mirrors the Codex backend routing-hint contract for
|
||||
// OpenAI OAuth requests. The request model must already be the final upstream
|
||||
// slug and serviceTier must already reflect any local policy rewrite/filter.
|
||||
func setOpenAICodexRoutingHint(headers http.Header, account *Account, model string, serviceTier string) {
|
||||
if headers == nil || account == nil || !account.IsOpenAIOAuth() {
|
||||
return
|
||||
}
|
||||
|
||||
// OAuth headers are built afresh for every request/dial, but delete first so
|
||||
// callers updating a cached header set cannot retain a stale routing hint.
|
||||
headers.Del(openAICodexRoutingHintHeader)
|
||||
model = strings.TrimSpace(model)
|
||||
if model == "" {
|
||||
return
|
||||
}
|
||||
|
||||
// Codex treats "default" as an explicit standard-routing sentinel, not as a
|
||||
// service tier sent to the backend. Fast follows the gateway's existing
|
||||
// canonicalization and therefore becomes "priority"; flex stays "flex".
|
||||
canonicalTier := normalizedOpenAIServiceTierValue(serviceTier)
|
||||
// This backport has no Codex model-catalog snapshot with which to validate
|
||||
// arbitrary tier ids. Keep the hint to the two effective tiers Codex itself
|
||||
// selects; default, missing, and other gateway-compatible API values remain
|
||||
// model-only rather than expanding the ChatGPT routing protocol here.
|
||||
switch canonicalTier {
|
||||
case OpenAIFastTierPriority, OpenAIFastTierFlex:
|
||||
default:
|
||||
canonicalTier = ""
|
||||
}
|
||||
|
||||
hint := "model=" + model
|
||||
if canonicalTier != "" {
|
||||
hint += ";tier=" + canonicalTier
|
||||
}
|
||||
if !httpguts.ValidHeaderFieldValue(hint) {
|
||||
return
|
||||
}
|
||||
headers.Set(openAICodexRoutingHintHeader, hint)
|
||||
}
|
||||
|
||||
func setOpenAICodexRoutingHintFromBody(headers http.Header, account *Account, body []byte) {
|
||||
fields := gjson.GetManyBytes(body, "model", "service_tier")
|
||||
setOpenAICodexRoutingHint(headers, account, fields[0].String(), fields[1].String())
|
||||
}
|
||||
@@ -0,0 +1,207 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSetOpenAICodexRoutingHintCanonicalizesOfficialServiceTiers(t *testing.T) {
|
||||
oauthAccount := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}
|
||||
tests := []struct {
|
||||
name string
|
||||
model string
|
||||
serviceTier string
|
||||
want string
|
||||
}{
|
||||
{name: "fast alias", model: "gpt-5.6", serviceTier: " fast ", want: "model=gpt-5.6;tier=priority"},
|
||||
{name: "priority", model: "gpt-5.6", serviceTier: "priority", want: "model=gpt-5.6;tier=priority"},
|
||||
{name: "flex", model: "gpt-5.6", serviceTier: "flex", want: "model=gpt-5.6;tier=flex"},
|
||||
{name: "explicit default sentinel", model: "gpt-5.6", serviceTier: "default", want: "model=gpt-5.6"},
|
||||
{name: "omitted tier", model: "gpt-5.6", want: "model=gpt-5.6"},
|
||||
{name: "auto is not expanded without catalog support", model: "gpt-5.6", serviceTier: "auto", want: "model=gpt-5.6"},
|
||||
{name: "scale is not expanded without catalog support", model: "gpt-5.6", serviceTier: "scale", want: "model=gpt-5.6"},
|
||||
{name: "unknown tier does not expand protocol", model: "gpt-5.6", serviceTier: "turbo", want: "model=gpt-5.6"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
headers := make(http.Header)
|
||||
setOpenAICodexRoutingHint(headers, oauthAccount, tt.model, tt.serviceTier)
|
||||
require.Equal(t, tt.want, headers.Get(openAICodexRoutingHintHeader))
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("invalid header value is omitted", func(t *testing.T) {
|
||||
headers := make(http.Header)
|
||||
setOpenAICodexRoutingHint(headers, oauthAccount, "gpt-5.6\ninvalid", "priority")
|
||||
require.Empty(t, headers.Get(openAICodexRoutingHintHeader))
|
||||
})
|
||||
|
||||
t.Run("api key is untouched", func(t *testing.T) {
|
||||
headers := make(http.Header)
|
||||
headers.Set(openAICodexRoutingHintHeader, "caller-owned")
|
||||
setOpenAICodexRoutingHint(headers, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, "gpt-5.6", "priority")
|
||||
require.Equal(t, "caller-owned", headers.Get(openAICodexRoutingHintHeader))
|
||||
})
|
||||
}
|
||||
|
||||
func TestOpenAIOAuthHTTPBuildersSendRoutingHintFromFinalBody(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
oauthAccount := &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"chatgpt_account_id": "test-account",
|
||||
},
|
||||
}
|
||||
svc := &OpenAIGatewayService{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
body []byte
|
||||
want string
|
||||
}{
|
||||
{name: "fast", body: []byte(`{"model":"gpt-5.6-codex","service_tier":"fast"}`), want: "model=gpt-5.6-codex;tier=priority"},
|
||||
{name: "flex", body: []byte(`{"model":"gpt-5.6-codex","service_tier":"flex"}`), want: "model=gpt-5.6-codex;tier=flex"},
|
||||
{name: "default", body: []byte(`{"model":"gpt-5.6-codex","service_tier":"default"}`), want: "model=gpt-5.6-codex"},
|
||||
{name: "omitted", body: []byte(`{"model":"gpt-5.6-codex"}`), want: "model=gpt-5.6-codex"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
for _, passthrough := range []bool{false, true} {
|
||||
mode := "ordinary"
|
||||
if passthrough {
|
||||
mode = "passthrough"
|
||||
}
|
||||
t.Run(mode, func(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(tt.body))
|
||||
|
||||
var req *http.Request
|
||||
var err error
|
||||
if passthrough {
|
||||
req, err = svc.buildUpstreamRequestOpenAIPassthrough(context.Background(), c, oauthAccount, tt.body, "test-token")
|
||||
} else {
|
||||
req, err = svc.buildUpstreamRequest(context.Background(), c, oauthAccount, tt.body, "test-token", false, "", true)
|
||||
}
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.want, req.Header.Get(openAICodexRoutingHintHeader))
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildOpenAIWSHeadersSendsOAuthRoutingHintOnly(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
|
||||
svc := &OpenAIGatewayService{}
|
||||
decision := OpenAIWSProtocolDecision{Transport: OpenAIUpstreamTransportResponsesWebsocketV2}
|
||||
|
||||
build := func(t *testing.T, account *Account, tier string) http.Header {
|
||||
headers, _, err := svc.buildOpenAIWSHeaders(
|
||||
context.Background(),
|
||||
c,
|
||||
account,
|
||||
"test-token",
|
||||
decision,
|
||||
true,
|
||||
"",
|
||||
"",
|
||||
"",
|
||||
"gpt-5.6-codex",
|
||||
tier,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
return headers
|
||||
}
|
||||
|
||||
oauthAccount := &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"chatgpt_account_id": "test-account",
|
||||
},
|
||||
}
|
||||
require.Equal(t, "model=gpt-5.6-codex;tier=priority", build(t, oauthAccount, "fast").Get(openAICodexRoutingHintHeader))
|
||||
require.Equal(t, "model=gpt-5.6-codex", build(t, oauthAccount, "default").Get(openAICodexRoutingHintHeader))
|
||||
require.Empty(t, build(t, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, "priority").Get(openAICodexRoutingHintHeader))
|
||||
}
|
||||
|
||||
func TestOpenAIWSConnPoolDoesNotReuseDifferentRoutingHints(t *testing.T) {
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 4
|
||||
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
|
||||
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 4
|
||||
|
||||
pool := newOpenAIWSConnPool(cfg)
|
||||
dialer := &openAIWSCountingDialer{}
|
||||
pool.setClientDialerForTest(dialer)
|
||||
account := &Account{ID: 913, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
|
||||
|
||||
acquire := func(t *testing.T, hint string) *openAIWSConnLease {
|
||||
headers := make(http.Header)
|
||||
headers.Set(openAICodexRoutingHintHeader, hint)
|
||||
lease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{
|
||||
Account: account,
|
||||
WSURL: "wss://example.com/v1/responses",
|
||||
Headers: headers,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, lease)
|
||||
return lease
|
||||
}
|
||||
|
||||
priority := acquire(t, "model=gpt-5.6-codex;tier=priority")
|
||||
priorityConnID := priority.ConnID()
|
||||
priority.Release()
|
||||
flexHeaders := make(http.Header)
|
||||
flexHeaders.Set(openAICodexRoutingHintHeader, "model=gpt-5.6-codex;tier=flex")
|
||||
_, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{
|
||||
Account: account,
|
||||
WSURL: "wss://example.com/v1/responses",
|
||||
Headers: flexHeaders,
|
||||
PreferredConnID: priorityConnID,
|
||||
ForcePreferredConn: true,
|
||||
})
|
||||
require.ErrorIs(t, err, errOpenAIWSPreferredConnUnavailable)
|
||||
|
||||
priorityAgain := acquire(t, "model=gpt-5.6-codex;tier=priority")
|
||||
require.True(t, priorityAgain.Reused())
|
||||
require.Equal(t, priorityConnID, priorityAgain.ConnID())
|
||||
priorityAgain.Release()
|
||||
|
||||
flex := acquire(t, "model=gpt-5.6-codex;tier=flex")
|
||||
require.False(t, flex.Reused())
|
||||
require.NotEqual(t, priorityConnID, flex.ConnID())
|
||||
flex.Release()
|
||||
|
||||
otherModel := acquire(t, "model=gpt-5.5-codex;tier=priority")
|
||||
require.False(t, otherModel.Reused())
|
||||
require.NotEqual(t, priorityConnID, otherModel.ConnID())
|
||||
otherModel.Release()
|
||||
|
||||
defaultTier := acquire(t, "model=gpt-5.6-codex")
|
||||
require.False(t, defaultTier.Reused())
|
||||
require.NotEqual(t, priorityConnID, defaultTier.ConnID())
|
||||
defaultConnID := defaultTier.ConnID()
|
||||
defaultTier.Release()
|
||||
|
||||
defaultAgain := acquire(t, "model=gpt-5.6-codex")
|
||||
require.True(t, defaultAgain.Reused())
|
||||
require.Equal(t, defaultConnID, defaultAgain.ConnID())
|
||||
defaultAgain.Release()
|
||||
|
||||
require.Equal(t, 4, dialer.DialCount())
|
||||
}
|
||||
@@ -617,7 +617,20 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
}
|
||||
}
|
||||
|
||||
wsHeaders, _, buildHdrErr := s.buildOpenAIWSHeaders(ctx, c, account, token, wsDecision, isCodexCLI, turnState, strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)), firstPayload.promptCacheKey)
|
||||
firstRoutingFields := gjson.GetManyBytes(firstPayload.payloadRaw, "model", "service_tier")
|
||||
wsHeaders, _, buildHdrErr := s.buildOpenAIWSHeaders(
|
||||
ctx,
|
||||
c,
|
||||
account,
|
||||
token,
|
||||
wsDecision,
|
||||
isCodexCLI,
|
||||
turnState,
|
||||
strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)),
|
||||
firstPayload.promptCacheKey,
|
||||
firstRoutingFields[0].String(),
|
||||
firstRoutingFields[1].String(),
|
||||
)
|
||||
if buildHdrErr != nil {
|
||||
return fmt.Errorf("build ws headers: %w", buildHdrErr)
|
||||
}
|
||||
@@ -1632,16 +1645,30 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
if parseErr != nil {
|
||||
return parseErr
|
||||
}
|
||||
nextRoutingFields := gjson.GetManyBytes(nextPayload.payloadRaw, "model", "service_tier")
|
||||
if nextPayload.promptCacheKey != "" {
|
||||
// ingress 会话在整个客户端 WS 生命周期内复用同一上游连接;
|
||||
// prompt_cache_key 对握手头的更新仅在未来需要重新建连时生效。
|
||||
updatedHeaders, _, updHdrErr := s.buildOpenAIWSHeaders(ctx, c, account, token, wsDecision, isCodexCLI, turnState, strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)), nextPayload.promptCacheKey)
|
||||
updatedHeaders, _, updHdrErr := s.buildOpenAIWSHeaders(
|
||||
ctx,
|
||||
c,
|
||||
account,
|
||||
token,
|
||||
wsDecision,
|
||||
isCodexCLI,
|
||||
turnState,
|
||||
strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)),
|
||||
nextPayload.promptCacheKey,
|
||||
nextRoutingFields[0].String(),
|
||||
nextRoutingFields[1].String(),
|
||||
)
|
||||
if updHdrErr != nil {
|
||||
logOpenAIWSModeInfo("ingress_ws_update_headers_failed account_id=%d err=%v", account.ID, updHdrErr)
|
||||
} else {
|
||||
baseAcquireReq.Headers = updatedHeaders
|
||||
}
|
||||
}
|
||||
setOpenAICodexRoutingHint(baseAcquireReq.Headers, account, nextRoutingFields[0].String(), nextRoutingFields[1].String())
|
||||
if nextPayload.previousResponseID != "" {
|
||||
expectedPrev := strings.TrimSpace(lastTurnResponseID)
|
||||
chainedFromLast := expectedPrev != "" && nextPayload.previousResponseID == expectedPrev
|
||||
|
||||
@@ -75,6 +75,8 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders(
|
||||
turnState string,
|
||||
turnMetadata string,
|
||||
promptCacheKey string,
|
||||
routingModel string,
|
||||
routingServiceTier string,
|
||||
) (http.Header, openAIWSSessionHeaderResolution, error) {
|
||||
headers := make(http.Header)
|
||||
if account == nil || !account.IsOpenAIAgentIdentity() {
|
||||
@@ -157,6 +159,7 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders(
|
||||
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)。
|
||||
// 覆盖所有 WS 模式(ctx_pool/dedicated/passthrough)的握手头。
|
||||
account.ApplyHeaderOverrides(headers)
|
||||
setOpenAICodexRoutingHint(headers, account, routingModel, routingServiceTier)
|
||||
|
||||
return headers, sessionResolution, nil
|
||||
}
|
||||
|
||||
@@ -414,6 +414,8 @@ func TestOpenAIGatewayService_BuildOpenAIWSHeadersPreservesCodexIdentity(t *test
|
||||
"",
|
||||
"",
|
||||
"",
|
||||
"",
|
||||
"",
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -136,7 +136,19 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2(
|
||||
storeDisabledConnMode := s.openAIWSStoreDisabledConnMode()
|
||||
forceNewConnByPolicy := shouldForceNewConnOnStoreDisabled(storeDisabledConnMode, lastFailureReason)
|
||||
forceNewConn := forceNewConnByPolicy && storeDisabled && previousResponseID == "" && sessionHash != "" && preferredConnID == ""
|
||||
wsHeaders, sessionResolution, buildHdrErr := s.buildOpenAIWSHeaders(ctx, c, account, token, decision, isCodexCLI, turnState, turnMetadata, promptCacheKey)
|
||||
wsHeaders, sessionResolution, buildHdrErr := s.buildOpenAIWSHeaders(
|
||||
ctx,
|
||||
c,
|
||||
account,
|
||||
token,
|
||||
decision,
|
||||
isCodexCLI,
|
||||
turnState,
|
||||
turnMetadata,
|
||||
promptCacheKey,
|
||||
openAIWSPayloadString(payload, "model"),
|
||||
openAIWSPayloadString(payload, "service_tier"),
|
||||
)
|
||||
if buildHdrErr != nil {
|
||||
return nil, fmt.Errorf("build ws headers: %w", buildHdrErr)
|
||||
}
|
||||
|
||||
@@ -76,6 +76,11 @@ type openAIWSAcquireRequest struct {
|
||||
ForcePreferredConn bool
|
||||
}
|
||||
|
||||
type openAIWSHandshakeCompatibilityKey struct {
|
||||
betaFeatures string
|
||||
routingHint string
|
||||
}
|
||||
|
||||
type openAIWSConnLease struct {
|
||||
pool *openAIWSConnPool
|
||||
accountID int64
|
||||
@@ -240,8 +245,8 @@ type openAIWSConn struct {
|
||||
id string
|
||||
ws openAIWSClientConn
|
||||
|
||||
handshakeHeaders http.Header
|
||||
betaFeatures string
|
||||
handshakeHeaders http.Header
|
||||
handshakeCompatibility openAIWSHandshakeCompatibilityKey
|
||||
|
||||
leaseCh chan struct{}
|
||||
closedCh chan struct{}
|
||||
@@ -525,8 +530,8 @@ 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) matchesHandshakeCompatibility(compatibility openAIWSHandshakeCompatibilityKey) bool {
|
||||
return c != nil && c.handshakeCompatibility == compatibility
|
||||
}
|
||||
|
||||
func (c *openAIWSConn) isPrewarmed() bool {
|
||||
@@ -838,7 +843,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque
|
||||
|
||||
retryAcquire:
|
||||
accountID := req.Account.ID
|
||||
betaFeatures := normalizeOpenAIWSBetaFeatures(req.Headers)
|
||||
compatibility := normalizeOpenAIWSHandshakeCompatibility(req.Headers)
|
||||
effectiveMaxConns := p.effectiveMaxConnsByAccount(req.Account)
|
||||
if effectiveMaxConns <= 0 {
|
||||
return nil, errOpenAIWSConnQueueFull
|
||||
@@ -866,7 +871,7 @@ retryAcquire:
|
||||
return nil, errOpenAIWSPreferredConnUnavailable
|
||||
}
|
||||
preferredConn, ok := ap.conns[preferredConnID]
|
||||
if !ok || !preferredConn.matchesBetaFeatures(betaFeatures) {
|
||||
if !ok || !preferredConn.matchesHandshakeCompatibility(compatibility) {
|
||||
p.recordConnPickDuration(time.Since(pickStartedAt))
|
||||
ap.mu.Unlock()
|
||||
closeOpenAIWSConns(evicted)
|
||||
@@ -947,7 +952,7 @@ retryAcquire:
|
||||
}
|
||||
|
||||
if preferredConnID != "" {
|
||||
if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) && conn.tryAcquire() {
|
||||
if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesHandshakeCompatibility(compatibility) && conn.tryAcquire() {
|
||||
connPick := time.Since(pickStartedAt)
|
||||
p.recordConnPickDuration(connPick)
|
||||
ap.mu.Unlock()
|
||||
@@ -969,7 +974,7 @@ retryAcquire:
|
||||
}
|
||||
}
|
||||
|
||||
best := p.pickLeastBusyConnLocked(ap, "", betaFeatures)
|
||||
best := p.pickLeastBusyConnLocked(ap, "", compatibility)
|
||||
if best != nil && best.tryAcquire() {
|
||||
connPick := time.Since(pickStartedAt)
|
||||
p.recordConnPickDuration(connPick)
|
||||
@@ -991,7 +996,7 @@ retryAcquire:
|
||||
return lease, nil
|
||||
}
|
||||
for _, conn := range ap.conns {
|
||||
if conn == nil || conn == best || !conn.matchesBetaFeatures(betaFeatures) {
|
||||
if conn == nil || conn == best || !conn.matchesHandshakeCompatibility(compatibility) {
|
||||
continue
|
||||
}
|
||||
if conn.tryAcquire() {
|
||||
@@ -1018,8 +1023,8 @@ retryAcquire:
|
||||
}
|
||||
|
||||
if !req.ForceNewConn && len(ap.conns)+ap.creating >= effectiveMaxConns {
|
||||
compatible := p.pickLeastBusyConnLocked(ap, "", betaFeatures)
|
||||
if idle := p.pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap, betaFeatures); idle != nil {
|
||||
compatible := p.pickLeastBusyConnLocked(ap, "", compatibility)
|
||||
if idle := p.pickOldestIdleConnWithDifferentHandshakeCompatibilityLocked(ap, compatibility); idle != nil {
|
||||
delete(ap.conns, idle.id)
|
||||
evicted = append(evicted, idle)
|
||||
p.metrics.scaleDownTotal.Add(1)
|
||||
@@ -1100,7 +1105,7 @@ retryAcquire:
|
||||
return nil, errOpenAIWSConnQueueFull
|
||||
}
|
||||
|
||||
target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID, betaFeatures)
|
||||
target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID, compatibility)
|
||||
connPick := time.Since(pickStartedAt)
|
||||
p.recordConnPickDuration(connPick)
|
||||
if target == nil {
|
||||
@@ -1173,13 +1178,16 @@ func (p *openAIWSConnPool) pickOldestIdleConnLocked(ap *openAIWSAccountPool) *op
|
||||
return oldest
|
||||
}
|
||||
|
||||
func (p *openAIWSConnPool) pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap *openAIWSAccountPool, betaFeatures string) *openAIWSConn {
|
||||
func (p *openAIWSConnPool) pickOldestIdleConnWithDifferentHandshakeCompatibilityLocked(
|
||||
ap *openAIWSAccountPool,
|
||||
compatibility openAIWSHandshakeCompatibilityKey,
|
||||
) *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) {
|
||||
if conn == nil || conn.matchesHandshakeCompatibility(compatibility) || conn.isLeased() || conn.waiters.Load() > 0 || p.isConnPinnedLocked(ap, conn.id) {
|
||||
continue
|
||||
}
|
||||
if oldest == nil || conn.lastUsedAt().Before(oldest.lastUsedAt()) {
|
||||
@@ -1330,13 +1338,17 @@ func (p *openAIWSConnPool) cleanupAccountLocked(ap *openAIWSAccountPool, now tim
|
||||
return evicted
|
||||
}
|
||||
|
||||
func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, preferredConnID, betaFeatures string) *openAIWSConn {
|
||||
func (p *openAIWSConnPool) pickLeastBusyConnLocked(
|
||||
ap *openAIWSAccountPool,
|
||||
preferredConnID string,
|
||||
compatibility openAIWSHandshakeCompatibilityKey,
|
||||
) *openAIWSConn {
|
||||
if ap == nil || len(ap.conns) == 0 {
|
||||
return nil
|
||||
}
|
||||
preferredConnID = stringsTrim(preferredConnID)
|
||||
if preferredConnID != "" {
|
||||
if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) {
|
||||
if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesHandshakeCompatibility(compatibility) {
|
||||
return conn
|
||||
}
|
||||
}
|
||||
@@ -1344,7 +1356,7 @@ func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, pref
|
||||
var bestWaiters int32
|
||||
var bestLastUsed time.Time
|
||||
for _, conn := range ap.conns {
|
||||
if conn == nil || !conn.matchesBetaFeatures(betaFeatures) {
|
||||
if conn == nil || !conn.matchesHandshakeCompatibility(compatibility) {
|
||||
continue
|
||||
}
|
||||
waiters := conn.waiters.Load()
|
||||
@@ -1676,7 +1688,7 @@ func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequ
|
||||
}
|
||||
id := p.nextConnID(req.Account.ID)
|
||||
pooledConn := newOpenAIWSConn(id, req.Account.ID, conn, handshakeHeaders)
|
||||
pooledConn.betaFeatures = normalizeOpenAIWSBetaFeatures(req.Headers)
|
||||
pooledConn.handshakeCompatibility = normalizeOpenAIWSHandshakeCompatibility(req.Headers)
|
||||
return pooledConn, nil
|
||||
}
|
||||
|
||||
@@ -1880,6 +1892,13 @@ func normalizeOpenAIWSBetaFeatures(headers http.Header) string {
|
||||
return strings.Join(normalized, ",")
|
||||
}
|
||||
|
||||
func normalizeOpenAIWSHandshakeCompatibility(headers http.Header) openAIWSHandshakeCompatibilityKey {
|
||||
return openAIWSHandshakeCompatibilityKey{
|
||||
betaFeatures: normalizeOpenAIWSBetaFeatures(headers),
|
||||
routingHint: strings.TrimSpace(headers.Get(openAICodexRoutingHintHeader)),
|
||||
}
|
||||
}
|
||||
|
||||
func cloneHeader(src http.Header) http.Header {
|
||||
if src == nil {
|
||||
return nil
|
||||
|
||||
@@ -800,7 +800,19 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
turnState = strings.TrimSpace(c.GetHeader(openAIWSTurnStateHeader))
|
||||
turnMetadata = strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader))
|
||||
}
|
||||
headers, _, buildHdrErr := s.buildOpenAIWSHeaders(ctx, c, account, token, wsDecision, isCodexCLI, turnState, turnMetadata, promptCacheKey)
|
||||
headers, _, buildHdrErr := s.buildOpenAIWSHeaders(
|
||||
ctx,
|
||||
c,
|
||||
account,
|
||||
token,
|
||||
wsDecision,
|
||||
isCodexCLI,
|
||||
turnState,
|
||||
turnMetadata,
|
||||
promptCacheKey,
|
||||
gjson.GetBytes(firstClientMessage, "model").String(),
|
||||
gjson.GetBytes(firstClientMessage, "service_tier").String(),
|
||||
)
|
||||
if buildHdrErr != nil {
|
||||
return fmt.Errorf("build ws headers: %w", buildHdrErr)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user