mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge pull request #4163 from bestony/fix/openai-ws-ingress-lifecycle
fix(openai-ws): bound ingress session lifecycle
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
# 按账号类型细分开关
|
||||
|
||||
Reference in New Issue
Block a user