Merge pull request #4163 from bestony/fix/openai-ws-ingress-lifecycle

fix(openai-ws): bound ingress session lifecycle
This commit is contained in:
Wesley Liddick
2026-07-13 16:07:53 +08:00
committed by GitHub
19 changed files with 933 additions and 12 deletions
+21
View File
@@ -596,6 +596,27 @@ If you disable URL validation or response header filtering, harden your network
- Enforce TLS-only outbound traffic
- Strip sensitive upstream response headers at the proxy
#### OpenAI Responses WebSocket ingress limits
`gateway.openai_ws` bounds the lifetime and aggregate count of client-facing
Responses WebSocket sessions. These safeguards apply independently from
per-turn user and account concurrency slots, which are released between turns.
```yaml
gateway:
openai_ws:
# Close a client socket idle between completed turns; 0 disables this safeguard.
ingress_inter_turn_idle_timeout_seconds: 300
# Distributed API-key limit for live client ingress sessions; 0 disables it.
max_ingress_connections_per_api_key: 64
```
The connection cap is coordinated through Redis using a 60-second lease that
is refreshed every 20 seconds. A process that cannot confirm a lease for a
full lease lifetime closes its local WebSocket rather than continuing outside
the global cap. Use `http_bridge` for client-WebSocket/upstream-HTTP operation
when rolling out or mitigating upstream WebSocket issues.
#### ⚠️ Important: Creating the Admin Account
The initial admin account is **only created via the setup wizard** (served at `http://<host>:8080` on first run). The `default.admin_email` / `default.admin_password` fields in `config.yaml` are **not used** to create it — they exist in the template for historical reasons.
+14
View File
@@ -923,6 +923,12 @@ type GatewayOpenAIWSConfig struct {
ModeRouterV2Enabled bool `mapstructure:"mode_router_v2_enabled"`
// IngressModeDefault: ingress 默认模式(off/ctx_pool/passthrough/http_bridge)
IngressModeDefault string `mapstructure:"ingress_mode_default"`
// IngressInterTurnIdleTimeoutSeconds bounds the time a client may remain idle
// between completed ingress turns. Zero disables this protection.
IngressInterTurnIdleTimeoutSeconds int `mapstructure:"ingress_inter_turn_idle_timeout_seconds"`
// MaxIngressConnectionsPerAPIKey bounds live client WebSocket ingress sessions
// per API key across all instances. Zero disables this protection.
MaxIngressConnectionsPerAPIKey int `mapstructure:"max_ingress_connections_per_api_key"`
// Enabled: 全局总开关(默认 true)
Enabled bool `mapstructure:"enabled"`
// OAuthEnabled: 是否允许 OpenAI OAuth 账号使用 WS
@@ -1945,6 +1951,8 @@ func setDefaults() {
viper.SetDefault("gateway.openai_ws.enabled", true)
viper.SetDefault("gateway.openai_ws.mode_router_v2_enabled", false)
viper.SetDefault("gateway.openai_ws.ingress_mode_default", "ctx_pool")
viper.SetDefault("gateway.openai_ws.ingress_inter_turn_idle_timeout_seconds", 300)
viper.SetDefault("gateway.openai_ws.max_ingress_connections_per_api_key", 64)
viper.SetDefault("gateway.openai_ws.oauth_enabled", true)
viper.SetDefault("gateway.openai_ws.apikey_enabled", true)
viper.SetDefault("gateway.openai_ws.force_http", false)
@@ -2717,6 +2725,12 @@ func (c *Config) Validate() error {
if c.Gateway.OpenAIWS.MaxConnsPerAccount <= 0 {
return fmt.Errorf("gateway.openai_ws.max_conns_per_account must be positive")
}
if c.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds < 0 {
return fmt.Errorf("gateway.openai_ws.ingress_inter_turn_idle_timeout_seconds must be non-negative")
}
if c.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey < 0 {
return fmt.Errorf("gateway.openai_ws.max_ingress_connections_per_api_key must be non-negative")
}
if c.Gateway.OpenAIWS.MinIdlePerAccount < 0 {
return fmt.Errorf("gateway.openai_ws.min_idle_per_account must be non-negative")
}
+16
View File
@@ -182,6 +182,12 @@ func TestLoadDefaultOpenAIWSConfig(t *testing.T) {
if cfg.Gateway.OpenAIWS.IngressModeDefault != "ctx_pool" {
t.Fatalf("Gateway.OpenAIWS.IngressModeDefault = %q, want %q", cfg.Gateway.OpenAIWS.IngressModeDefault, "ctx_pool")
}
if cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds != 300 {
t.Fatalf("Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = %d, want 300", cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds)
}
if cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey != 64 {
t.Fatalf("Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = %d, want 64", cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey)
}
}
func TestLoadDefaultOpenAICompactModel(t *testing.T) {
@@ -1640,6 +1646,16 @@ func TestValidateConfig_OpenAIWSRules(t *testing.T) {
mutate: func(c *Config) { c.Gateway.OpenAIWS.MaxConnsPerAccount = 0 },
wantErr: "gateway.openai_ws.max_conns_per_account",
},
{
name: "ingress_inter_turn_idle_timeout_seconds 不能为负数",
mutate: func(c *Config) { c.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = -1 },
wantErr: "gateway.openai_ws.ingress_inter_turn_idle_timeout_seconds",
},
{
name: "max_ingress_connections_per_api_key 不能为负数",
mutate: func(c *Config) { c.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = -1 },
wantErr: "gateway.openai_ws.max_ingress_connections_per_api_key",
},
{
name: "min_idle_per_account 不能为负数",
mutate: func(c *Config) { c.Gateway.OpenAIWS.MinIdlePerAccount = -1 },
@@ -220,6 +220,15 @@ func (h *ConcurrencyHelper) TryAcquireUserSlotForAPIKey(ctx context.Context, use
return h.withAPIKeySlot(ctx, apiKeyID, releaseFunc), true, nil
}
// AcquireOpenAIWSIngressLease bounds the whole client WebSocket lifecycle,
// independently from per-turn user and account slots.
func (h *ConcurrencyHelper) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int) (*service.OpenAIWSIngressLease, bool, error) {
if h == nil || h.concurrencyService == nil {
return nil, false, fmt.Errorf("concurrency service is unavailable")
}
return h.concurrencyService.AcquireOpenAIWSIngressLease(ctx, apiKeyID, maxConnections)
}
// TryAcquireAccountSlot 尝试立即获取账号并发槽位。
// 返回值: (releaseFunc, acquired, error)
func (h *ConcurrencyHelper) TryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (func(), bool, error) {
@@ -11,10 +11,13 @@ import (
)
type concurrencyCacheMock struct {
acquireUserSlotFn func(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error)
acquireAccountSlotFn func(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error)
releaseUserCalled int32
releaseAccountCalled int32
acquireUserSlotFn func(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error)
acquireAccountSlotFn func(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error)
acquireIngressLeaseFn func(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error)
releaseIngressLeaseFn func(ctx context.Context, apiKeyID int64, leaseID string) error
releaseUserCalled int32
releaseAccountCalled int32
releaseIngressCalled int32
}
func (m *concurrencyCacheMock) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) {
@@ -97,6 +100,25 @@ func (m *concurrencyCacheMock) CleanupStaleProcessSlots(ctx context.Context, act
return nil
}
func (m *concurrencyCacheMock) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error) {
if m.acquireIngressLeaseFn != nil {
return m.acquireIngressLeaseFn(ctx, apiKeyID, maxConnections, leaseID)
}
return false, nil
}
func (m *concurrencyCacheMock) RefreshOpenAIWSIngressLease(context.Context, int64, string) (bool, error) {
return true, nil
}
func (m *concurrencyCacheMock) ReleaseOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) error {
atomic.AddInt32(&m.releaseIngressCalled, 1)
if m.releaseIngressLeaseFn != nil {
return m.releaseIngressLeaseFn(ctx, apiKeyID, leaseID)
}
return nil
}
func TestConcurrencyHelper_TryAcquireUserSlot(t *testing.T) {
cache := &concurrencyCacheMock{
acquireUserSlotFn: func(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) {
@@ -1299,10 +1299,36 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
wsConn.SetReadLimit(service.ResolveOpenAIWSClientReadLimitBytes(h.cfg))
ctx := c.Request.Context()
maxIngressConnections := 0
if h.cfg != nil {
maxIngressConnections = h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey
}
ingressLease, ingressLeaseAcquired, ingressLeaseErr := h.concurrencyHelper.AcquireOpenAIWSIngressLease(ctx, apiKey.ID, maxIngressConnections)
if ingressLeaseErr != nil {
reqLog.Error("openai.websocket_ingress_lease_acquire_failed", zap.Error(ingressLeaseErr))
closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to reserve websocket ingress capacity")
return
}
if !ingressLeaseAcquired {
reqLog.Info("openai.websocket_ingress_capacity_rejected", zap.Int("max_ingress_connections_per_api_key", maxIngressConnections))
closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "too many open websocket connections, please retry later")
return
}
if ingressLease != nil {
defer ingressLease.Release()
ctx = ingressLease.Context()
c.Request = c.Request.WithContext(ctx)
}
readCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
msgType, firstMessage, err := wsConn.Read(readCtx)
cancel()
if err != nil {
if errors.Is(context.Cause(ctx), service.ErrOpenAIWSIngressLeaseLost) {
reqLog.Warn("openai.websocket_ingress_lease_lost_before_first_message", zap.Error(err))
closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "websocket ingress capacity lease lost; please reconnect")
return
}
closeStatus, closeReason := summarizeWSCloseErrorForLog(err)
reqLog.Warn("openai.websocket_read_first_message_failed",
zap.Error(err),
@@ -1692,6 +1718,25 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
continue
}
if errors.Is(context.Cause(ctx), service.ErrOpenAIWSIngressLeaseLost) {
reqLog.Warn("openai.websocket_ingress_lease_lost",
zap.Int64("account_id", account.ID),
zap.Error(err),
)
closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "websocket ingress capacity lease lost; please reconnect")
return
}
var closeErr *service.OpenAIWSClientCloseError
if errors.As(err, &closeErr) && closeErr.StatusCode() == coderws.StatusNormalClosure {
reqLog.Info("openai.websocket_ingress_closed_normally",
zap.Int64("account_id", account.ID),
zap.String("reason", closeErr.Reason()),
)
closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason())
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
closeStatus, closeReason := summarizeWSCloseErrorForLog(err)
reqLog.Warn("openai.websocket_proxy_failed",
@@ -1700,7 +1745,6 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
zap.String("close_status", closeStatus),
zap.String("close_reason", closeReason),
)
var closeErr *service.OpenAIWSClientCloseError
if errors.As(err, &closeErr) {
closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason())
return
@@ -8,6 +8,7 @@ import (
"net/http/httptest"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
@@ -711,6 +712,68 @@ func TestOpenAIResponsesWebSocket_InvalidUpgradeDoesNotSetTransport(t *testing.T
require.Equal(t, service.OpenAIClientTransportUnknown, service.GetOpenAIClientTransport(c))
}
func TestOpenAIResponsesWebSocket_IngressCapacityRejected(t *testing.T) {
gin.SetMode(gin.TestMode)
cache := &concurrencyCacheMock{
acquireIngressLeaseFn: func(context.Context, int64, int, string) (bool, error) {
return false, nil
},
}
h := newOpenAIHandlerForPreviousResponseIDValidation(t, cache)
h.cfg = &config.Config{}
h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = 1
wsServer := newOpenAIWSHandlerTestServer(t, h, middleware.AuthSubject{UserID: 1, Concurrency: 1})
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http")+"/openai/v1/responses", nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, _, err = clientConn.Read(readCtx)
cancelRead()
var closeErr coderws.CloseError
require.ErrorAs(t, err, &closeErr)
require.Equal(t, coderws.StatusTryAgainLater, closeErr.Code)
}
func TestOpenAIResponsesWebSocket_IngressLeaseReleasedOnEarlyReturn(t *testing.T) {
gin.SetMode(gin.TestMode)
cache := &concurrencyCacheMock{
acquireIngressLeaseFn: func(context.Context, int64, int, string) (bool, error) {
return true, nil
},
}
h := newOpenAIHandlerForPreviousResponseIDValidation(t, cache)
h.cfg = &config.Config{}
h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = 1
wsServer := newOpenAIWSHandlerTestServer(t, h, middleware.AuthSubject{UserID: 1, Concurrency: 1})
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http")+"/openai/v1/responses", nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageBinary, []byte("not a response.create frame"))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, _, err = clientConn.Read(readCtx)
cancelRead()
var closeErr coderws.CloseError
require.ErrorAs(t, err, &closeErr)
require.Equal(t, coderws.StatusPolicyViolation, closeErr.Code)
require.Eventually(t, func() bool {
return atomic.LoadInt32(&cache.releaseIngressCalled) == 1
}, time.Second, 10*time.Millisecond)
}
func TestOpenAIResponsesWebSocket_RejectsMessageIDAsPreviousResponseID(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -30,6 +30,10 @@ const (
userSlotKeyPrefix = "concurrency:user:"
// 格式: concurrency:api_key:{apiKeyID}
apiKeySlotKeyPrefix = "concurrency:api_key:"
// API-key-scoped client WebSocket ingress leases use a shorter TTL than
// ordinary request slots, because idle ingress sessions do not hold a turn slot.
openAIWSIngressLeaseKeyPrefix = "concurrency:openai_ws_ingress:api_key:"
openAIWSIngressLeaseTTLSeconds = 60
// 等待队列计数器格式: concurrency:wait:{userID}
waitQueueKeyPrefix = "concurrency:wait:"
// 账号级等待队列计数器格式: wait:account:{accountID}
@@ -138,6 +142,49 @@ var (
return 1
`)
// acquireOpenAIWSIngressLeaseScript atomically reaps crashed members and
// acquires or refreshes one API-key-scoped ingress lease using Redis TIME.
acquireOpenAIWSIngressLeaseScript = redis.NewScript(`
redis.replicate_commands()
local key = KEYS[1]
local maxConnections = tonumber(ARGV[1])
local ttl = tonumber(ARGV[2])
local leaseID = ARGV[3]
local now = tonumber(redis.call('TIME')[1])
local expireBefore = now - ttl
redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore)
if redis.call('ZSCORE', key, leaseID) ~= false then
redis.call('ZADD', key, now, leaseID)
redis.call('EXPIRE', key, ttl)
return 1
end
if redis.call('ZCARD', key) < maxConnections then
redis.call('ZADD', key, now, leaseID)
redis.call('EXPIRE', key, ttl)
return 1
end
return 0
`)
// refreshOpenAIWSIngressLeaseScript does not recreate a missing member: a
// process that lost its lease must terminate its local WebSocket instead of
// silently continuing beyond the distributed cap.
refreshOpenAIWSIngressLeaseScript = redis.NewScript(`
redis.replicate_commands()
local key = KEYS[1]
local ttl = tonumber(ARGV[1])
local leaseID = ARGV[2]
local now = tonumber(redis.call('TIME')[1])
local expireBefore = now - ttl
redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore)
if redis.call('ZSCORE', key, leaseID) == false then
return 0
end
redis.call('ZADD', key, now, leaseID)
redis.call('EXPIRE', key, ttl)
return 1
`)
// incrementWaitScript - refreshes TTL on each increment to keep queue depth accurate
// KEYS[1] = wait queue key
// ARGV[1] = maxWait
@@ -283,6 +330,10 @@ func apiKeySlotKey(apiKeyID int64) string {
return fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID)
}
func openAIWSIngressLeaseKey(apiKeyID int64) string {
return fmt.Sprintf("%s%d", openAIWSIngressLeaseKeyPrefix, apiKeyID)
}
func waitQueueKey(userID int64) string {
return fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID)
}
@@ -623,6 +674,48 @@ func (c *concurrencyCache) ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64
return c.rdb.ZRem(ctx, key, requestID).Err()
}
func (c *concurrencyCache) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error) {
if c == nil || c.rdb == nil || apiKeyID <= 0 || maxConnections <= 0 || leaseID == "" {
return false, nil
}
result, err := acquireOpenAIWSIngressLeaseScript.Run(
ctx,
c.rdb,
[]string{openAIWSIngressLeaseKey(apiKeyID)},
maxConnections,
openAIWSIngressLeaseTTLSeconds,
leaseID,
).Int()
if err != nil {
return false, err
}
return result == 1, nil
}
func (c *concurrencyCache) RefreshOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) (bool, error) {
if c == nil || c.rdb == nil || apiKeyID <= 0 || leaseID == "" {
return false, nil
}
result, err := refreshOpenAIWSIngressLeaseScript.Run(
ctx,
c.rdb,
[]string{openAIWSIngressLeaseKey(apiKeyID)},
openAIWSIngressLeaseTTLSeconds,
leaseID,
).Int()
if err != nil {
return false, err
}
return result == 1, nil
}
func (c *concurrencyCache) ReleaseOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) error {
if c == nil || c.rdb == nil || apiKeyID <= 0 || leaseID == "" {
return nil
}
return c.rdb.ZRem(ctx, openAIWSIngressLeaseKey(apiKeyID), leaseID).Err()
}
func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) {
if len(apiKeyIDs) == 0 {
return map[int64]int{}, nil
@@ -50,6 +50,53 @@ func (s *ConcurrencyCacheSuite) apiKeyConcurrencyCache() apiKeyConcurrencyCacheF
return cache
}
func (s *ConcurrencyCacheSuite) TestOpenAIWSIngressAPIKeySlot_HardLimitRefreshAndRelease() {
apiKeyID := int64(9011)
firstLeaseID := "ingress-first"
secondLeaseID := "ingress-second"
ok, err := s.rawCache.AcquireOpenAIWSIngressLease(s.ctx, apiKeyID, 1, firstLeaseID)
require.NoError(s.T(), err)
require.True(s.T(), ok)
ok, err = s.rawCache.AcquireOpenAIWSIngressLease(s.ctx, apiKeyID, 1, secondLeaseID)
require.NoError(s.T(), err)
require.False(s.T(), ok, "a second live session must not exceed the API key limit")
ok, err = s.rawCache.RefreshOpenAIWSIngressLease(s.ctx, apiKeyID, firstLeaseID)
require.NoError(s.T(), err)
require.True(s.T(), ok, "the current owner must be able to refresh its lease")
require.NoError(s.T(), s.rawCache.ReleaseOpenAIWSIngressLease(s.ctx, apiKeyID, firstLeaseID))
ok, err = s.rawCache.AcquireOpenAIWSIngressLease(s.ctx, apiKeyID, 1, secondLeaseID)
require.NoError(s.T(), err)
require.True(s.T(), ok, "released capacity must become available immediately")
}
func (s *ConcurrencyCacheSuite) TestOpenAIWSIngressAPIKeySlot_ReapsCrashedLeaseWithoutDeletingLiveOtherInstance() {
apiKeyID := int64(9012)
key := openAIWSIngressLeaseKey(apiKeyID)
now, err := s.rawCache.redisUnixSeconds(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, key,
redis.Z{Score: float64(now - openAIWSIngressLeaseTTLSeconds - 1), Member: "crashed-instance"},
redis.Z{Score: float64(now), Member: "live-other-instance"},
).Err())
require.NoError(s.T(), s.rdb.Expire(s.ctx, key, time.Duration(openAIWSIngressLeaseTTLSeconds)*time.Second).Err())
ok, err := s.rawCache.AcquireOpenAIWSIngressLease(s.ctx, apiKeyID, 2, "new-instance")
require.NoError(s.T(), err)
require.True(s.T(), ok, "the crashed member should be reaped before enforcing the limit")
_, err = s.rdb.ZScore(s.ctx, key, "crashed-instance").Result()
require.ErrorIs(s.T(), err, redis.Nil)
_, err = s.rdb.ZScore(s.ctx, key, "live-other-instance").Result()
require.NoError(s.T(), err, "a live lease owned by another instance must be preserved")
count, err := s.rdb.ZCard(s.ctx, key).Result()
require.NoError(s.T(), err)
require.Equal(s.T(), int64(2), count)
}
func (s *ConcurrencyCacheSuite) TestAccountSlot_AcquireAndRelease() {
accountID := int64(10)
reqID1, reqID2, reqID3 := "req1", "req2", "req3"
@@ -6,6 +6,7 @@ import (
"crypto/sha256"
"encoding/binary"
"encoding/hex"
"errors"
"os"
"strconv"
"sync"
@@ -13,6 +14,7 @@ import (
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"go.uber.org/zap"
"golang.org/x/sync/singleflight"
)
@@ -59,6 +61,131 @@ type APIKeyConcurrencyCache interface {
GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error)
}
// OpenAIWSIngressLeaseCache owns the short-lived distributed lease used to
// bound live client WebSocket sessions. It is deliberately independent of the
// request-slot namespace: idle ingress connections do not occupy turn slots.
type OpenAIWSIngressLeaseCache interface {
AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error)
RefreshOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) (bool, error)
ReleaseOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) error
}
const (
openAIWSIngressLeaseTTL = 60 * time.Second
openAIWSIngressLeaseRefreshInterval = 20 * time.Second
openAIWSIngressLeaseOperationTO = 2 * time.Second
)
var ErrOpenAIWSIngressLeaseLost = errors.New("openai websocket ingress lease lost")
// OpenAIWSIngressLease keeps a Redis-backed ingress lease alive and cancels
// its context if Redis cannot confirm ownership for a full lease lifetime.
// Call Release on every handler exit to reclaim capacity immediately.
type OpenAIWSIngressLease struct {
ctx context.Context
cancel context.CancelCauseFunc
cache OpenAIWSIngressLeaseCache
apiKeyID int64
leaseID string
stopOnce sync.Once
stopCh chan struct{}
refreshDone chan struct{}
}
func (l *OpenAIWSIngressLease) Context() context.Context {
if l == nil || l.ctx == nil {
return context.Background()
}
return l.ctx
}
func (l *OpenAIWSIngressLease) Release() {
if l == nil {
return
}
l.stopOnce.Do(func() {
if l.stopCh != nil {
close(l.stopCh)
}
if l.cancel != nil {
l.cancel(nil)
}
if l.refreshDone != nil {
<-l.refreshDone
}
if l.cache == nil || l.apiKeyID <= 0 || l.leaseID == "" {
return
}
releaseCtx, releaseCancel := context.WithTimeout(context.Background(), openAIWSIngressLeaseOperationTO)
defer releaseCancel()
if err := l.cache.ReleaseOpenAIWSIngressLease(releaseCtx, l.apiKeyID, l.leaseID); err != nil {
logger.L().Warn("openai_ws_ingress_lease_release_failed",
zap.Int64("api_key_id", l.apiKeyID),
zap.Error(err),
)
}
})
}
func (l *OpenAIWSIngressLease) refreshLoop() {
defer func() {
if l != nil && l.refreshDone != nil {
close(l.refreshDone)
}
}()
if l == nil || l.cache == nil {
return
}
ticker := time.NewTicker(openAIWSIngressLeaseRefreshInterval)
defer ticker.Stop()
lastConfirmedAt := time.Now()
for {
select {
case <-l.ctx.Done():
return
case <-l.stopCh:
return
case <-ticker.C:
var lost bool
lastConfirmedAt, lost = l.refresh(lastConfirmedAt)
if lost {
l.cancel(ErrOpenAIWSIngressLeaseLost)
return
}
}
}
}
// refresh confirms the lease is still owned. A missing member is an immediate
// lease loss; transient Redis errors are tolerated only for one full lease TTL.
func (l *OpenAIWSIngressLease) refresh(lastConfirmedAt time.Time) (time.Time, bool) {
refreshCtx, refreshCancel := context.WithTimeout(context.Background(), openAIWSIngressLeaseOperationTO)
owned, err := l.cache.RefreshOpenAIWSIngressLease(refreshCtx, l.apiKeyID, l.leaseID)
refreshCancel()
if err == nil && owned {
return time.Now(), false
}
if err == nil {
err = ErrOpenAIWSIngressLeaseLost
}
elapsed := time.Since(lastConfirmedAt)
logger.L().Warn("openai_ws_ingress_lease_refresh_failed",
zap.Int64("api_key_id", l.apiKeyID),
zap.Duration("unconfirmed_for", elapsed),
zap.Error(err),
)
if errors.Is(err, ErrOpenAIWSIngressLeaseLost) || elapsed >= openAIWSIngressLeaseTTL {
logger.L().Error("openai_ws_ingress_lease_lost",
zap.Int64("api_key_id", l.apiKeyID),
zap.Duration("unconfirmed_for", elapsed),
zap.Error(err),
)
return lastConfirmedAt, true
}
return lastConfirmedAt, false
}
var (
requestIDPrefix = initRequestIDPrefix()
requestIDCounter atomic.Uint64
@@ -125,6 +252,47 @@ func NewConcurrencyService(cache ConcurrencyCache) *ConcurrencyService {
return svc
}
// AcquireOpenAIWSIngressLease atomically reserves one live ingress connection
// for an API key. A non-positive limit explicitly disables this protection.
func (s *ConcurrencyService) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int) (*OpenAIWSIngressLease, bool, error) {
if maxConnections <= 0 {
return nil, true, nil
}
if s == nil || s.cache == nil || apiKeyID <= 0 {
return nil, false, errors.New("openai websocket ingress lease cache is unavailable")
}
cache, ok := s.cache.(OpenAIWSIngressLeaseCache)
if !ok {
return nil, false, errors.New("openai websocket ingress lease cache is unsupported")
}
leaseID := generateRequestID()
baseCtx := context.Background()
if ctx != nil {
baseCtx = context.WithoutCancel(ctx)
}
acquireCtx, acquireCancel := context.WithTimeout(baseCtx, openAIWSIngressLeaseOperationTO)
acquired, err := cache.AcquireOpenAIWSIngressLease(acquireCtx, apiKeyID, maxConnections, leaseID)
acquireCancel()
if err != nil || !acquired {
return nil, acquired, err
}
if ctx == nil {
ctx = context.Background()
}
leaseCtx, leaseCancel := context.WithCancelCause(ctx)
lease := &OpenAIWSIngressLease{
ctx: leaseCtx,
cancel: leaseCancel,
cache: cache,
apiKeyID: apiKeyID,
leaseID: leaseID,
stopCh: make(chan struct{}),
refreshDone: make(chan struct{}),
}
go lease.refreshLoop()
return lease, true, nil
}
// SetAccountLoadBatchCacheTTL 设置账号负载批量读取的极短 TTL 缓存;非正数表示禁用缓存。
func (s *ConcurrencyService) SetAccountLoadBatchCacheTTL(ttl time.Duration) {
if s == nil {
@@ -45,7 +45,47 @@ type stubConcurrencyCacheForTest struct {
releasedAPIKeyRequestIDs []string
}
type ingressLeaseCacheForTest struct {
stubConcurrencyCacheForTest
acquireIngressResult bool
acquireIngressErr error
acquireIngressFn func(context.Context, int64, int, string) (bool, error)
refreshIngressResult bool
refreshIngressErr error
refreshIngressFn func(context.Context, int64, string) (bool, error)
releaseIngressErr error
releaseIngressFn func(context.Context, int64, string) error
acquireIngressCalls int
refreshIngressCalls int
releaseIngressCalls int
}
func (c *ingressLeaseCacheForTest) AcquireOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, maxConnections int, leaseID string) (bool, error) {
c.acquireIngressCalls++
if c.acquireIngressFn != nil {
return c.acquireIngressFn(ctx, apiKeyID, maxConnections, leaseID)
}
return c.acquireIngressResult, c.acquireIngressErr
}
func (c *ingressLeaseCacheForTest) RefreshOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) (bool, error) {
c.refreshIngressCalls++
if c.refreshIngressFn != nil {
return c.refreshIngressFn(ctx, apiKeyID, leaseID)
}
return c.refreshIngressResult, c.refreshIngressErr
}
func (c *ingressLeaseCacheForTest) ReleaseOpenAIWSIngressLease(ctx context.Context, apiKeyID int64, leaseID string) error {
c.releaseIngressCalls++
if c.releaseIngressFn != nil {
return c.releaseIngressFn(ctx, apiKeyID, leaseID)
}
return c.releaseIngressErr
}
var _ ConcurrencyCache = (*stubConcurrencyCacheForTest)(nil)
var _ OpenAIWSIngressLeaseCache = (*ingressLeaseCacheForTest)(nil)
func (c *stubConcurrencyCacheForTest) AcquireAccountSlot(_ context.Context, _ int64, _ int, _ string) (bool, error) {
return c.acquireResult, c.acquireErr
@@ -285,6 +325,114 @@ func TestGetAPIKeyConcurrencyBatch_Fallbacks(t *testing.T) {
})
}
func TestAcquireOpenAIWSIngressLease(t *testing.T) {
t.Run("zero value release is safe", func(t *testing.T) {
var lease OpenAIWSIngressLease
require.NotPanics(t, lease.Release)
})
t.Run("disabled", func(t *testing.T) {
cache := &ingressLeaseCacheForTest{}
lease, acquired, err := NewConcurrencyService(cache).AcquireOpenAIWSIngressLease(nil, 1, 0)
require.NoError(t, err)
require.True(t, acquired)
require.Nil(t, lease)
require.Zero(t, cache.acquireIngressCalls)
})
t.Run("unsupported cache fails closed", func(t *testing.T) {
lease, acquired, err := NewConcurrencyService(&stubConcurrencyCacheForTest{}).AcquireOpenAIWSIngressLease(context.Background(), 1, 1)
require.Error(t, err)
require.False(t, acquired)
require.Nil(t, lease)
})
t.Run("capacity rejected", func(t *testing.T) {
cache := &ingressLeaseCacheForTest{acquireIngressResult: false}
lease, acquired, err := NewConcurrencyService(cache).AcquireOpenAIWSIngressLease(context.Background(), 1, 1)
require.NoError(t, err)
require.False(t, acquired)
require.Nil(t, lease)
})
t.Run("release returns capacity", func(t *testing.T) {
cache := &ingressLeaseCacheForTest{acquireIngressResult: true, refreshIngressResult: true}
lease, acquired, err := NewConcurrencyService(cache).AcquireOpenAIWSIngressLease(nil, 1, 1)
require.NoError(t, err)
require.True(t, acquired)
require.NotNil(t, lease)
lease.Release()
lease.Release()
require.Equal(t, 1, cache.releaseIngressCalls)
})
}
func TestOpenAIWSIngressLeaseRefreshLoss(t *testing.T) {
t.Run("missing lease is lost immediately", func(t *testing.T) {
cache := &ingressLeaseCacheForTest{refreshIngressResult: false}
lease := &OpenAIWSIngressLease{cache: cache, apiKeyID: 1, leaseID: "missing"}
_, lost := lease.refresh(time.Now())
require.True(t, lost)
require.Equal(t, 1, cache.refreshIngressCalls)
})
t.Run("persistent redis errors lose lease after ttl", func(t *testing.T) {
cache := &ingressLeaseCacheForTest{refreshIngressErr: errors.New("redis unavailable")}
lease := &OpenAIWSIngressLease{cache: cache, apiKeyID: 1, leaseID: "unconfirmed"}
_, lost := lease.refresh(time.Now().Add(-openAIWSIngressLeaseTTL))
require.True(t, lost)
require.Equal(t, 1, cache.refreshIngressCalls)
})
}
func TestOpenAIWSIngressLeaseReleaseWaitsForInFlightRefresh(t *testing.T) {
refreshStarted := make(chan struct{})
allowRefresh := make(chan struct{})
cache := &ingressLeaseCacheForTest{
refreshIngressFn: func(context.Context, int64, string) (bool, error) {
close(refreshStarted)
<-allowRefresh
return true, nil
},
}
ctx, cancel := context.WithCancelCause(context.Background())
lease := &OpenAIWSIngressLease{
ctx: ctx,
cancel: cancel,
cache: cache,
apiKeyID: 1,
leaseID: "in-flight-refresh",
stopCh: make(chan struct{}),
refreshDone: make(chan struct{}),
}
go func() {
defer close(lease.refreshDone)
_, _ = lease.refresh(time.Now())
}()
<-refreshStarted
released := make(chan struct{})
go func() {
lease.Release()
close(released)
}()
select {
case <-released:
t.Fatal("release returned before the in-flight refresh completed")
case <-time.After(20 * time.Millisecond):
}
require.Zero(t, cache.releaseIngressCalls)
close(allowRefresh)
select {
case <-released:
case <-time.After(time.Second):
t.Fatal("release did not complete after the refresh returned")
}
require.Equal(t, 1, cache.releaseIngressCalls)
}
func TestGenerateRequestID_UsesStablePrefixAndMonotonicCounter(t *testing.T) {
id1 := generateRequestID()
id2 := generateRequestID()
@@ -39,6 +39,13 @@ type openAIWSClientConn interface {
Close() error
}
// openAIWSIdlePingCapable is intentionally separate from openAIWSClientConn.
// A pool probe happens while no goroutine is reading an idle connection, which
// is not safe for every WebSocket implementation.
type openAIWSIdlePingCapable interface {
SupportsIdlePingWithoutReader() bool
}
// openAIWSClientDialer 抽象 WS 建连器。
type openAIWSClientDialer interface {
Dial(ctx context.Context, wsURL string, headers http.Header, proxyURL string) (openAIWSClientConn, int, http.Header, error)
@@ -301,6 +308,14 @@ func (c *coderOpenAIWSClientConn) Ping(ctx context.Context) error {
return c.conn.Ping(ctx)
}
// SupportsIdlePingWithoutReader reports the actual coder/websocket contract.
// Conn.Ping waits for a pong, while control frames are only consumed by Read.
// The pool deliberately has no reader on an idle connection, so using Ping as
// a health probe would deterministically time out a healthy socket.
func (*coderOpenAIWSClientConn) SupportsIdlePingWithoutReader() bool {
return false
}
func (c *coderOpenAIWSClientConn) Close() error {
if c == nil || c.conn == nil {
return nil
@@ -110,3 +110,7 @@ func TestCoderOpenAIWSClientDialer_ProxyTransportTLSHandshakeTimeout(t *testing.
require.NotNil(t, transport)
require.Equal(t, 10*time.Second, transport.TLSHandshakeTimeout)
}
func TestCoderOpenAIWSClientConn_DoesNotSupportIdlePingWithoutReader(t *testing.T) {
require.False(t, (&coderOpenAIWSClientConn{}).SupportsIdlePingWithoutReader())
}
@@ -18,6 +18,13 @@ import (
"github.com/tidwall/sjson"
)
func (s *OpenAIGatewayService) openAIWSIngressInterTurnIdleTimeout() time.Duration {
if s == nil || s.cfg == nil || s.cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds <= 0 {
return 0
}
return time.Duration(s.cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds) * time.Second
}
func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
ctx context.Context,
c *gin.Context,
@@ -360,8 +367,23 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
readClientMessage := func() ([]byte, error) {
msgType, payload, readErr := clientConn.Read(ctx)
readCtx := ctx
idleTimeout := s.openAIWSIngressInterTurnIdleTimeout()
cancelRead := func() {}
if idleTimeout > 0 {
readCtx, cancelRead = context.WithTimeout(ctx, idleTimeout)
}
msgType, payload, readErr := clientConn.Read(readCtx)
cancelRead()
if readErr != nil {
if idleTimeout > 0 && errors.Is(readErr, context.DeadlineExceeded) && ctx.Err() == nil {
logOpenAIWSModeInfo("ingress_ws_inter_turn_idle_timeout account_id=%d timeout_seconds=%d", account.ID, int(idleTimeout.Seconds()))
return nil, NewOpenAIWSClientCloseError(
coderws.StatusNormalClosure,
"websocket idle timeout",
readErr,
)
}
return nil, readErr
}
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
@@ -1318,7 +1340,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
unpinSessionConn(sessionConnID)
}
}
shouldPreflightPing := turn > 1 && sessionLease != nil && turnRetry == 0
shouldPreflightPing := turn > 1 && sessionLease != nil && sessionLease.SupportsIdlePingWithoutReader() && turnRetry == 0
if shouldPreflightPing && openAIWSIngressPreflightPingIdle > 0 && !lastTurnFinishedAt.IsZero() {
if time.Since(lastTurnFinishedAt) < openAIWSIngressPreflightPingIdle {
shouldPreflightPing = false
@@ -164,6 +164,111 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossT
require.Len(t, captureConn.writes, 2, "应向同一上游连接发送两轮 response.create")
}
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_IdleTimeoutReleasesStoreDisabledSession(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 1
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
captureConn := &openAIWSCaptureConn{events: [][]byte{
[]byte(`{"type":"response.completed","response":{"id":"resp_idle_timeout","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
}}
captureDialer := &openAIWSCaptureDialer{conn: captureConn}
pool := newOpenAIWSConnPool(cfg)
pool.setClientDialerForTest(captureDialer)
defer pool.Close()
svc := &OpenAIGatewayService{
cfg: cfg,
httpUpstream: &httpUpstreamRecorder{},
cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
toolCorrector: NewCodexToolCorrector(),
openaiWSPool: pool,
}
account := &Account{
ID: 116,
Name: "openai-ingress-idle-timeout",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"responses_websockets_v2_enabled": true},
}
serverErrCh := make(chan error, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
serverErrCh <- err
return
}
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
_, firstMessage, err := conn.Read(readCtx)
cancelRead()
if err != nil {
serverErrCh <- err
return
}
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
ginCtx.Request = r.Clone(r.Context())
serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
}))
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false}`))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, event, err := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, err)
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
select {
case proxyErr := <-serverErrCh:
var closeErr *OpenAIWSClientCloseError
require.ErrorAs(t, proxyErr, &closeErr)
require.Equal(t, coderws.StatusNormalClosure, closeErr.StatusCode())
require.Equal(t, "websocket idle timeout", closeErr.Reason())
case <-time.After(4 * time.Second):
t.Fatal("timed out waiting for idle ingress session to close")
}
ap, ok := pool.getAccountPool(account.ID)
require.True(t, ok)
ap.mu.Lock()
require.Empty(t, ap.pinnedConns, "idle close must unpin a store=false session")
for _, conn := range ap.conns {
require.False(t, conn.isLeased(), "idle close must release the upstream lease")
}
ap.mu.Unlock()
}
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_FollowupCreateCanOmitModel(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -632,3 +632,100 @@ func TestOpenAIWSHTTPBridgeKeepsContinuationFramesOnHTTPWithoutPreviousResponseI
require.Equal(t, 0, captureDialer.DialCount())
require.Empty(t, captureConn.writes)
}
func TestOpenAIWSHTTPBridge_IdleTimeoutClosesClientSession(t *testing.T) {
gin.SetMode(gin.TestMode)
sseBody := strings.Join([]string{
`data: {"type":"response.completed","response":{"id":"resp_bridge_idle","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`,
"",
}, "\n")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(sseBody)),
}}
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true
cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1
cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 1
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
svc := &OpenAIGatewayService{
cfg: cfg,
httpUpstream: upstream,
cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
toolCorrector: NewCodexToolCorrector(),
}
account := &Account{
ID: 20,
Name: "api-key-bridge-idle-timeout",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Credentials: map[string]any{"api_key": "sk-upstream"},
Extra: map[string]any{"responses_websockets_v2_enabled": true},
Concurrency: 1,
Status: StatusActive,
Schedulable: true,
}
errCh := make(chan error, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
errCh <- err
return
}
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
_, firstMessage, err := conn.Read(readCtx)
cancelRead()
if err != nil {
errCh <- err
return
}
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
ginCtx.Request = r.Clone(r.Context())
errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
}))
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false,"input":"hello"}`))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, event, err := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, err)
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
select {
case proxyErr := <-errCh:
var closeErr *OpenAIWSClientCloseError
require.ErrorAs(t, proxyErr, &closeErr)
require.Equal(t, coderws.StatusNormalClosure, closeErr.StatusCode())
require.Equal(t, "websocket idle timeout", closeErr.Reason())
case <-time.After(4 * time.Second):
t.Fatal("timed out waiting for idle HTTP bridge session to close")
}
require.Len(t, upstream.bodies, 1, "an idle client must not leave a continuation request running")
}
+21 -3
View File
@@ -203,6 +203,14 @@ func (l *openAIWSConnLease) PingWithTimeout(timeout time.Duration) error {
return conn.pingWithTimeout(timeout)
}
func (l *openAIWSConnLease) SupportsIdlePingWithoutReader() bool {
conn, err := l.activeConn()
if err != nil {
return false
}
return conn.supportsIdlePingWithoutReader()
}
func (l *openAIWSConnLease) MarkBroken() {
if l == nil || l.pool == nil || l.conn == nil || l.released.Load() {
return
@@ -437,6 +445,16 @@ func (c *openAIWSConn) pingWithTimeout(timeout time.Duration) error {
return nil
}
func (c *openAIWSConn) supportsIdlePingWithoutReader() bool {
if c == nil || c.ws == nil {
return false
}
capable, ok := c.ws.(openAIWSIdlePingCapable)
// Test and alternate implementations keep the historical probe behavior
// unless they explicitly declare it unsafe.
return !ok || capable.SupportsIdlePingWithoutReader()
}
func (c *openAIWSConn) touch() {
if c == nil {
return
@@ -707,7 +725,7 @@ func (p *openAIWSConnPool) runBackgroundPingSweep() {
g.SetLimit(10)
for _, item := range candidates {
item := item
if item.conn == nil || item.conn.isLeased() || item.conn.waiters.Load() > 0 {
if item.conn == nil || item.conn.isLeased() || item.conn.waiters.Load() > 0 || !item.conn.supportsIdlePingWithoutReader() {
continue
}
g.Go(func() error {
@@ -1613,7 +1631,7 @@ func (p *openAIWSConnPool) nextConnID(accountID int64) string {
}
func (p *openAIWSConnPool) shouldHealthCheckConn(conn *openAIWSConn) bool {
if conn == nil {
if conn == nil || !conn.supportsIdlePingWithoutReader() {
return false
}
return conn.idleDuration(time.Now()) >= openAIWSConnHealthCheckIdle
@@ -1669,7 +1687,7 @@ func (p *openAIWSConnPool) effectiveMaxConnsByAccount(account *Account) int {
if account.Concurrency <= 0 {
return 0
}
return account.Concurrency
return min(account.Concurrency, hardCap)
}
if account == nil || !p.dynamicMaxConnsEnabled() {
return hardCap
@@ -739,7 +739,7 @@ func TestOpenAIWSConnPool_EffectiveMaxConnsDisabledFallbackHardCap(t *testing.T)
require.Equal(t, 8, pool.effectiveMaxConnsByAccount(account), "关闭动态模式后应保持旧行为")
}
func TestOpenAIWSConnPool_EffectiveMaxConnsByAccount_ModeRouterV2UsesAccountConcurrency(t *testing.T) {
func TestOpenAIWSConnPool_EffectiveMaxConnsByAccount_ModeRouterV2RespectsHardCap(t *testing.T) {
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 8
@@ -750,7 +750,7 @@ func TestOpenAIWSConnPool_EffectiveMaxConnsByAccount_ModeRouterV2UsesAccountConc
pool := newOpenAIWSConnPool(cfg)
high := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 20}
require.Equal(t, 20, pool.effectiveMaxConnsByAccount(high), "v2 路径应直接使用账号并发数作为池上限")
require.Equal(t, 8, pool.effectiveMaxConnsByAccount(high), "v2 路径也必须受连接池硬上限约束")
nonPositive := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 0}
require.Equal(t, 0, pool.effectiveMaxConnsByAccount(nonPositive), "并发数<=0 时应不可调度")
@@ -1319,6 +1319,9 @@ func TestOpenAIWSConnPool_UtilityBranches(t *testing.T) {
conn := newOpenAIWSConn("health", 1, &openAIWSFakeConn{}, nil)
conn.lastUsedNano.Store(time.Now().Add(-openAIWSConnHealthCheckIdle - time.Second).UnixNano())
require.True(t, pool.shouldHealthCheckConn(conn))
unsafeConn := newOpenAIWSConn("unsafe_health", 1, &openAIWSIdlePingUnsupportedConn{}, nil)
unsafeConn.lastUsedNano.Store(time.Now().Add(-openAIWSConnHealthCheckIdle - time.Second).UnixNano())
require.False(t, pool.shouldHealthCheckConn(unsafeConn))
}
func TestOpenAIWSConn_LeaseAndTimeHelpers_NilAndClosedBranches(t *testing.T) {
@@ -1610,6 +1613,14 @@ type openAIWSPingBlockingConn struct {
release <-chan struct{}
}
type openAIWSIdlePingUnsupportedConn struct {
openAIWSFakeConn
}
func (c *openAIWSIdlePingUnsupportedConn) SupportsIdlePingWithoutReader() bool {
return false
}
func (c *openAIWSPingBlockingConn) WriteJSON(context.Context, any) error {
return nil
}
+4
View File
@@ -254,6 +254,10 @@ gateway:
# ingress 默认模式:off|ctx_pool|passthrough|http_bridge(仅 mode_router_v2_enabled=true 生效)
# 兼容旧值:shared/dedicated 会按 ctx_pool 处理。
ingress_mode_default: ctx_pool
# Close a client WebSocket that stays idle between completed turns (seconds). Set 0 to disable.
ingress_inter_turn_idle_timeout_seconds: 300
# Limit live client WebSocket ingress sessions per API key across all instances. Set 0 to disable.
max_ingress_connections_per_api_key: 64
# 全局总开关,默认 true;关闭时所有请求保持原有 HTTP/SSE 路由
enabled: true
# 按账号类型细分开关