diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index a02f5efa4c..b787d0f976 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -1612,7 +1612,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { reqLog.Info("openai.websocket_ingress_started") clientIP := ip.GetClientIP(c) userAgent := strings.TrimSpace(c.GetHeader("User-Agent")) - ctx := c.Request.Context() + clientLifecycleCtx := c.Request.Context() + ctx := clientLifecycleCtx maxIngressConnections := 0 if h.cfg != nil { maxIngressConnections = h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey @@ -2017,6 +2018,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { // openAIWSTurnPricing 的注释——绝不能用建连时刻初始化。 var turnPricing openAIWSTurnPricing hooks := &service.OpenAIWSIngressHooks{ + ClientLifecycleContext: clientLifecycleCtx, InitialRequestModel: reqModel, MaxReasoningEffort: maxReasoningEffort, ReasoningEffortMappings: reasoningEffortMappings, diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 9c0869fbee..51898060ac 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -207,6 +207,10 @@ func (e *OpenAIWSClientCloseError) Reason() string { // OpenAIWSIngressHooks 定义入站 WS 每个 turn 的生命周期回调。 type OpenAIWSIngressHooks struct { + // ClientLifecycleContext is the request context before an ingress lease + // adds its independent cancellation signal. Downstream writes bind to it + // so shutdown and disconnect cancellation remain direct during lease loss. + ClientLifecycleContext context.Context // InitialRequestModel is the client-facing model from the first frame, // before channel or account mapping. Ingress modes preserve it for usage // attribution while MapRequestModel determines the upstream model. diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 23582c2cea..400ae714b8 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -25,6 +25,21 @@ func (s *OpenAIGatewayService) openAIWSIngressInterTurnIdleTimeout() time.Durati return time.Duration(s.cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds) * time.Second } +// newOpenAIWSDownstreamWriteContext binds writes directly to the client +// lifecycle while excluding the separate ingress-lease cancellation signal. +// This lets a lease-loss path finish its current client write before +// ReadOpenAIWSClientMessage sends the retryable close frame. +func newOpenAIWSDownstreamWriteContext(controlCtx context.Context, hooks *OpenAIWSIngressHooks, timeout time.Duration) (context.Context, context.CancelFunc) { + writeParent := controlCtx + if hooks != nil && hooks.ClientLifecycleContext != nil { + writeParent = hooks.ClientLifecycleContext + } + if writeParent == nil { + writeParent = context.Background() + } + return context.WithTimeout(writeParent, timeout) +} + func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( ctx context.Context, c *gin.Context, @@ -369,7 +384,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( // the kernel send buffer before any close frame is queued. eventBytes := buildOpenAIFastPolicyBlockedWSEvent(blocked) if eventBytes != nil { - writeCtx, cancel := context.WithTimeout(ctx, s.openAIWSWriteTimeout()) + writeCtx, cancel := newOpenAIWSDownstreamWriteContext(ctx, hooks, s.openAIWSWriteTimeout()) _ = clientConn.Write(writeCtx, coderws.MessageText, eventBytes) cancel() } @@ -396,7 +411,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } writeClientMessage := func(message []byte) error { - writeCtx, cancel := context.WithTimeout(ctx, s.openAIWSWriteTimeout()) + writeCtx, cancel := newOpenAIWSDownstreamWriteContext(ctx, hooks, s.openAIWSWriteTimeout()) defer cancel() return clientConn.Write(writeCtx, coderws.MessageText, message) } 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 c90a1ebc0c..6715a01ae4 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go @@ -18,6 +18,80 @@ import ( "github.com/tidwall/gjson" ) +type openAIWSLeaseLossAfterReadConn struct { + *openAIWSCaptureConn + cancel context.CancelCauseFunc + once sync.Once +} + +func (c *openAIWSLeaseLossAfterReadConn) ReadMessage(ctx context.Context) ([]byte, error) { + message, err := c.openAIWSCaptureConn.ReadMessage(ctx) + if err == nil { + c.once.Do(func() { + c.cancel(ErrOpenAIWSIngressLeaseLost) + }) + } + return message, err +} + +type openAIWSSingleConnDialer struct { + conn openAIWSClientConn +} + +func (d *openAIWSSingleConnDialer) Dial( + ctx context.Context, + wsURL string, + headers http.Header, + proxyURL string, +) (openAIWSClientConn, int, http.Header, error) { + return d.conn, 0, nil, nil +} + +func TestOpenAIWSDownstreamWriteContext_CancellationOwnership(t *testing.T) { + t.Run("pre-canceled ordinary context is canceled before return", func(t *testing.T) { + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + cancelControl(context.Canceled) + + writeCtx, cancelWrite := newOpenAIWSDownstreamWriteContext(controlCtx, nil, time.Second) + defer cancelWrite() + require.ErrorIs(t, writeCtx.Err(), context.Canceled) + }) + + t.Run("lease loss keeps current write alive", func(t *testing.T) { + lifecycleCtx, cancelLifecycle := context.WithCancelCause(context.Background()) + controlCtx, cancelControl := context.WithCancelCause(lifecycleCtx) + hooks := &OpenAIWSIngressHooks{ClientLifecycleContext: lifecycleCtx} + writeCtx, cancelWrite := newOpenAIWSDownstreamWriteContext(controlCtx, hooks, time.Second) + defer cancelWrite() + + cancelControl(ErrOpenAIWSIngressLeaseLost) + select { + case <-writeCtx.Done(): + t.Fatalf("lease loss unexpectedly canceled downstream write: %v", writeCtx.Err()) + case <-time.After(20 * time.Millisecond): + } + + clientDisconnected := errors.New("client disconnected") + cancelLifecycle(clientDisconnected) + <-writeCtx.Done() + require.ErrorIs(t, context.Cause(writeCtx), clientDisconnected) + }) + + t.Run("ordinary cancellation is direct and preserves cause", func(t *testing.T) { + lifecycleCtx, cancelLifecycle := context.WithCancelCause(context.Background()) + controlCtx, cancelControl := context.WithCancelCause(lifecycleCtx) + defer cancelControl(context.Canceled) + hooks := &OpenAIWSIngressHooks{ClientLifecycleContext: lifecycleCtx} + writeCtx, cancelWrite := newOpenAIWSDownstreamWriteContext(controlCtx, hooks, time.Second) + defer cancelWrite() + + serverShutdown := errors.New("server shutdown") + cancelLifecycle(serverShutdown) + require.ErrorIs(t, writeCtx.Err(), context.Canceled) + require.ErrorIs(t, context.Cause(writeCtx), serverShutdown) + }) +} + func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossTurns(t *testing.T) { gin.SetMode(gin.TestMode) @@ -169,6 +243,137 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossT require.Len(t, captureConn.writes, 2, "应向同一上游连接发送两轮 response.create") } +func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_LeaseLossSendsRetryClose(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.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 + + lifecycleCtx, cancelLifecycle := context.WithCancelCause(context.Background()) + defer cancelLifecycle(context.Canceled) + controlCtx, cancelControl := context.WithCancelCause(lifecycleCtx) + upstreamConn := &openAIWSLeaseLossAfterReadConn{ + openAIWSCaptureConn: &openAIWSCaptureConn{events: [][]byte{ + []byte(`{"type":"response.completed","response":{"id":"resp_lease_loss","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`), + }}, + cancel: cancelControl, + } + pool := newOpenAIWSConnPool(cfg) + pool.setClientDialerForTest(&openAIWSSingleConnDialer{conn: upstreamConn}) + defer pool.Close() + + svc := &OpenAIGatewayService{ + cfg: cfg, + httpUpstream: &httpUpstreamRecorder{}, + cache: &stubGatewayCache{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + toolCorrector: NewCodexToolCorrector(), + openaiWSPool: pool, + } + account := &Account{ + ID: 118, + Name: "openai-ingress-lease-loss", + 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() + }() + + rec := httptest.NewRecorder() + ginCtx, _ := gin.CreateTestContext(rec) + req := r.Clone(controlCtx) + req.Header = req.Header.Clone() + req.Header.Set("User-Agent", "unit-test-agent/1.0") + ginCtx.Request = req + + readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) + msgType, firstMessage, readErr := conn.Read(readCtx) + cancelRead() + if readErr != nil { + serverErrCh <- readErr + return + } + if msgType != coderws.MessageText && msgType != coderws.MessageBinary { + serverErrCh <- errors.New("unsupported websocket client message type") + return + } + + serverErrCh <- svc.ProxyResponsesWebSocketFromClient( + controlCtx, + ginCtx, + conn, + account, + "sk-test", + firstMessage, + &OpenAIWSIngressHooks{ClientLifecycleContext: lifecycleCtx}, + ) + })) + 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}`)) + cancelWrite() + require.NoError(t, err) + + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + msgType, event, err := clientConn.Read(readCtx) + cancelRead() + require.NoError(t, err) + require.Equal(t, coderws.MessageText, msgType) + require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String()) + + closeReadCtx, cancelCloseRead := context.WithTimeout(context.Background(), 3*time.Second) + _, _, err = clientConn.Read(closeReadCtx) + cancelCloseRead() + var closeErr coderws.CloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusTryAgainLater, closeErr.Code) + require.Equal(t, "websocket ingress capacity lease lost; please reconnect", closeErr.Reason) + + select { + case serverErr := <-serverErrCh: + var clientCloseErr *OpenAIWSClientCloseError + require.ErrorAs(t, serverErr, &clientCloseErr) + require.Equal(t, coderws.StatusTryAgainLater, clientCloseErr.StatusCode()) + require.ErrorIs(t, serverErr, ErrOpenAIWSIngressLeaseLost) + case <-time.After(3 * time.Second): + t.Fatal("ingress lease-loss reader did not exit") + } +} + func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_IdleTimeoutReleasesStoreDisabledSession(t *testing.T) { gin.SetMode(gin.TestMode)