fix(openai-ws): bound ingress session lifecycle

This commit is contained in:
Bestony@Homelab
2026-07-13 15:32:42 +08:00
parent 551e2570dd
commit c8cfc93632
19 changed files with 933 additions and 12 deletions
@@ -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)