diff --git a/README.md b/README.md index 7e3b9ed6de..6bb4068644 100644 --- a/README.md +++ b/README.md @@ -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://: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. diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index df3afb6c7e..8e081bac34 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -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") } diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index 32aff543af..489bc346b8 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -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 }, diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go index 48110da93f..8489b17f6d 100644 --- a/backend/internal/handler/gateway_helper.go +++ b/backend/internal/handler/gateway_helper.go @@ -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) { diff --git a/backend/internal/handler/gateway_helper_fastpath_test.go b/backend/internal/handler/gateway_helper_fastpath_test.go index fecb9b071d..7ae9cb513e 100644 --- a/backend/internal/handler/gateway_helper_fastpath_test.go +++ b/backend/internal/handler/gateway_helper_fastpath_test.go @@ -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) { diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index afa2a5073a..7d6dc2c17a 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -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 diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index b7f43079ef..781c16b392 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -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) diff --git a/backend/internal/repository/concurrency_cache.go b/backend/internal/repository/concurrency_cache.go index b657c1ce8f..5341d411e9 100644 --- a/backend/internal/repository/concurrency_cache.go +++ b/backend/internal/repository/concurrency_cache.go @@ -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 diff --git a/backend/internal/repository/concurrency_cache_integration_test.go b/backend/internal/repository/concurrency_cache_integration_test.go index f7e27d1118..02821159ee 100644 --- a/backend/internal/repository/concurrency_cache_integration_test.go +++ b/backend/internal/repository/concurrency_cache_integration_test.go @@ -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" diff --git a/backend/internal/service/concurrency_service.go b/backend/internal/service/concurrency_service.go index f2f2aade89..df18379704 100644 --- a/backend/internal/service/concurrency_service.go +++ b/backend/internal/service/concurrency_service.go @@ -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 { diff --git a/backend/internal/service/concurrency_service_test.go b/backend/internal/service/concurrency_service_test.go index 3f358bbe6a..d079c7e60a 100644 --- a/backend/internal/service/concurrency_service_test.go +++ b/backend/internal/service/concurrency_service_test.go @@ -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() diff --git a/backend/internal/service/openai_ws_client.go b/backend/internal/service/openai_ws_client.go index 80b7553083..e336abcbd6 100644 --- a/backend/internal/service/openai_ws_client.go +++ b/backend/internal/service/openai_ws_client.go @@ -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 diff --git a/backend/internal/service/openai_ws_client_test.go b/backend/internal/service/openai_ws_client_test.go index a88d626651..95614cdbfe 100644 --- a/backend/internal/service/openai_ws_client_test.go +++ b/backend/internal/service/openai_ws_client_test.go @@ -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()) +} diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index a2af4b760f..504a7302f2 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -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 diff --git a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go index 9ae1b855ea..9d150856fd 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go @@ -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) diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index da2ae77917..0105c7a331 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -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") +} diff --git a/backend/internal/service/openai_ws_pool.go b/backend/internal/service/openai_ws_pool.go index 329908e762..8affe2a930 100644 --- a/backend/internal/service/openai_ws_pool.go +++ b/backend/internal/service/openai_ws_pool.go @@ -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 diff --git a/backend/internal/service/openai_ws_pool_test.go b/backend/internal/service/openai_ws_pool_test.go index ae9b94ce4a..8d339359ee 100644 --- a/backend/internal/service/openai_ws_pool_test.go +++ b/backend/internal/service/openai_ws_pool_test.go @@ -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 } diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml index eff4bfb598..954263d9c4 100644 --- a/deploy/config.example.yaml +++ b/deploy/config.example.yaml @@ -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 # 按账号类型细分开关