From f0e0b7e6d84f05d0936fb281d0a50263db5eefad Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=8B=E9=B9=8F?= <2829624376@qq.com> Date: Wed, 15 Jul 2026 22:44:19 +0800 Subject: [PATCH] fix(openai-ws): enforce passthrough turn lifecycle --- .../handler/openai_gateway_handler_test.go | 206 ++++++ .../internal/service/openai_ws_client_read.go | 83 ++- .../service/openai_ws_v2/passthrough_relay.go | 42 +- .../passthrough_relay_internal_test.go | 6 + .../openai_ws_v2/passthrough_relay_test.go | 64 ++ .../openai_ws_v2_passthrough_adapter.go | 540 +++++++++++++++- ...openai_ws_v2_passthrough_lifecycle_test.go | 591 ++++++++++++++++++ 7 files changed, 1493 insertions(+), 39 deletions(-) create mode 100644 backend/internal/service/openai_ws_v2_passthrough_lifecycle_test.go diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index 931f3b3f6..9f66f4a3f 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -1875,6 +1875,212 @@ func TestOpenAIResponsesWebSocket_FailoverOnUpstreamUsageLimitEvent(t *testing.T require.Equal(t, []int64{int64(9902)}, accountRepo.rateLimitedIDs) } +func TestOpenAIResponsesWebSocket_FirstOutputTimeoutWithoutDownstreamReusesClientForOneFailover(t *testing.T) { + gin.SetMode(gin.TestMode) + + firstHitCh := make(chan []byte, 1) + secondHitCh := make(chan []byte, 1) + var firstConnections atomic.Int32 + var secondConnections atomic.Int32 + + firstUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + firstConnections.Add(1) + conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) + if err != nil { + return + } + defer func() { _ = conn.CloseNow() }() + + readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) + _, payload, readErr := conn.Read(readCtx) + cancelRead() + if readErr == nil { + firstHitCh <- payload + } + + select { + case <-r.Context().Done(): + case <-time.After(3 * time.Second): + } + })) + defer firstUpstream.Close() + + secondUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + secondConnections.Add(1) + conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) + if err != nil { + return + } + defer func() { _ = conn.CloseNow() }() + + readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) + _, payload, readErr := conn.Read(readCtx) + cancelRead() + if readErr == nil { + secondHitCh <- payload + } + + for _, event := range []string{ + `{"type":"response.created","response":{"id":"resp_ws_timeout_b","model":"gpt-5.1"}}`, + `{"type":"response.output_text.delta","response_id":"resp_ws_timeout_b","delta":"recovered"}`, + `{"type":"response.completed","response":{"id":"resp_ws_timeout_b","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`, + } { + writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second) + writeErr := conn.Write(writeCtx, coderws.MessageText, []byte(event)) + cancelWrite() + if writeErr != nil { + return + } + } + readCtx, cancelRead = context.WithTimeout(r.Context(), 3*time.Second) + _, _, _ = conn.Read(readCtx) + cancelRead() + })) + defer secondUpstream.Close() + + groupID := int64(4212) + accounts := []service.Account{ + { + ID: 9912, + Name: "openai-ws-first-semantic-timeout", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 1, + Credentials: map[string]any{"api_key": "sk-first", "base_url": firstUpstream.URL}, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + "openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough, + }, + }, + { + ID: 9913, + Name: "openai-ws-failover-healthy", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 2, + Credentials: map[string]any{"api_key": "sk-second", "base_url": secondUpstream.URL}, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + "openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough, + }, + }, + } + + cfg := &config.Config{} + cfg.RunMode = config.RunModeSimple + cfg.Default.RateMultiplier = 1 + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIFirstOutputTimeoutSeconds = 1 + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 3 + cfg.Gateway.MaxAccountSwitches = 3 + + accountRepo := &openAIWSFailoverHandlerAccountRepoStub{accounts: accounts} + rateLimitSvc := service.NewRateLimitService(accountRepo, nil, cfg, nil, nil) + billingCacheSvc := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil) + gatewaySvc := service.NewOpenAIGatewayService( + accountRepo, nil, nil, nil, nil, nil, nil, cfg, nil, nil, + service.NewBillingService(cfg, nil), rateLimitSvc, billingCacheSvc, + nil, &service.DeferredService{}, nil, nil, nil, nil, nil, nil, nil, + ) + cache := &concurrencyCacheMock{ + acquireUserSlotFn: func(context.Context, int64, int, string) (bool, error) { return true, nil }, + acquireAccountSlotFn: func(context.Context, int64, int, string) (bool, error) { + return true, nil + }, + } + h := &OpenAIGatewayHandler{ + gatewayService: gatewaySvc, + billingCacheService: billingCacheSvc, + apiKeyService: &service.APIKeyService{}, + concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(cache), SSEPingFormatNone, time.Second), + maxAccountSwitches: 3, + } + + apiKey := &service.APIKey{ + ID: 1812, + GroupID: &groupID, + User: &service.User{ID: 1712, Status: service.StatusActive}, + Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI, Status: service.StatusActive}, + } + handlerDone := make(chan struct{}) + router := gin.New() + router.Use(func(c *gin.Context) { + c.Set(string(middleware.ContextKeyAPIKey), apiKey) + c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.User.ID, Concurrency: 1}) + c.Next() + }) + router.GET("/openai/v1/responses", func(c *gin.Context) { + h.ResponsesWebSocket(c) + close(handlerDone) + }) + handlerServer := httptest.NewServer(router) + defer handlerServer.Close() + + dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) + clientConn, _, err := coderws.Dial( + dialCtx, + "ws"+strings.TrimPrefix(handlerServer.URL, "http")+"/openai/v1/responses", + &coderws.DialOptions{CompressionMode: coderws.CompressionContextTakeover}, + ) + 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) + + var eventTypes []string + readCtx, cancelRead := context.WithTimeout(context.Background(), 6*time.Second) + for { + _, event, readErr := clientConn.Read(readCtx) + require.NoError(t, readErr) + eventType := gjson.GetBytes(event, "type").String() + eventTypes = append(eventTypes, eventType) + if eventType == "response.completed" { + require.Equal(t, "resp_ws_timeout_b", gjson.GetBytes(event, "response.id").String()) + break + } + } + cancelRead() + require.Contains(t, eventTypes, "response.output_text.delta") + require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done")) + + select { + case <-handlerDone: + case <-time.After(3 * time.Second): + t.Fatal("websocket handler did not finish after healthy failover turn") + } + select { + case <-firstHitCh: + case <-time.After(3 * time.Second): + t.Fatal("first upstream did not receive replayable request") + } + select { + case <-secondHitCh: + case <-time.After(3 * time.Second): + t.Fatal("second upstream did not receive replayed request") + } + require.Equal(t, int32(1), firstConnections.Load()) + require.Equal(t, int32(1), secondConnections.Load()) + require.NotContains(t, accountRepo.rateLimitedIDs, int64(9913), "healthy failover account must not be penalized") +} + func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSUsageLogCase) openAIResponsesWSUsageLogResult { t.Helper() gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_ws_client_read.go b/backend/internal/service/openai_ws_client_read.go index da8d650be..4c9d07c89 100644 --- a/backend/internal/service/openai_ws_client_read.go +++ b/backend/internal/service/openai_ws_client_read.go @@ -22,6 +22,29 @@ func ReadOpenAIWSClientMessage( timeout time.Duration, timeoutStatus coderws.StatusCode, timeoutReason string, +) (coderws.MessageType, []byte, error) { + return readOpenAIWSClientMessageWithTimeoutStart( + controlCtx, + conn, + timeout, + timeoutStatus, + timeoutReason, + nil, + nil, + ) +} + +// readOpenAIWSClientMessageWithTimeoutStart supports readers whose timeout +// starts after a state transition, such as a completed passthrough turn. When +// timeoutActive is nil, a positive timeout starts immediately. +func readOpenAIWSClientMessageWithTimeoutStart( + controlCtx context.Context, + conn *coderws.Conn, + timeout time.Duration, + timeoutStatus coderws.StatusCode, + timeoutReason string, + timeoutStart <-chan struct{}, + timeoutActive func() bool, ) (coderws.MessageType, []byte, error) { if conn == nil { return 0, nil, errors.New("openai websocket client connection is nil") @@ -36,13 +59,33 @@ func ReadOpenAIWSClientMessage( readDone <- openAIWSClientReadResult{messageType: messageType, payload: payload, err: err} }() - var timeoutCh <-chan time.Time var timer *time.Timer - if timeout > 0 { - timer = time.NewTimer(timeout) + var timeoutCh <-chan time.Time + startTimeout := func() { + if timeout <= 0 || (timeoutActive != nil && !timeoutActive()) { + return + } + if timer == nil { + timer = time.NewTimer(timeout) + } else { + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + timer.Reset(timeout) + } timeoutCh = timer.C - defer timer.Stop() } + if timeoutActive == nil || timeoutActive() { + startTimeout() + } + defer func() { + if timer != nil { + timer.Stop() + } + }() closeAndJoin := func(status coderws.StatusCode, reason string, cause error) (coderws.MessageType, []byte, error) { _ = conn.Close(status, reason) @@ -51,20 +94,24 @@ func ReadOpenAIWSClientMessage( return 0, nil, NewOpenAIWSClientCloseError(status, reason, cause) } - select { - case result := <-readDone: - return result.messageType, result.payload, result.err - case <-timeoutCh: - return closeAndJoin(timeoutStatus, timeoutReason, context.DeadlineExceeded) - case <-controlCtx.Done(): - cause := context.Cause(controlCtx) - if errors.Is(cause, ErrOpenAIWSIngressLeaseLost) { - return closeAndJoin( - coderws.StatusTryAgainLater, - "websocket ingress capacity lease lost; please reconnect", - cause, - ) + for { + select { + case result := <-readDone: + return result.messageType, result.payload, result.err + case <-timeoutStart: + startTimeout() + case <-timeoutCh: + return closeAndJoin(timeoutStatus, timeoutReason, context.DeadlineExceeded) + case <-controlCtx.Done(): + cause := context.Cause(controlCtx) + if errors.Is(cause, ErrOpenAIWSIngressLeaseLost) { + return closeAndJoin( + coderws.StatusTryAgainLater, + "websocket ingress capacity lease lost; please reconnect", + cause, + ) + } + return closeAndJoin(coderws.StatusGoingAway, "websocket request canceled", cause) } - return closeAndJoin(coderws.StatusGoingAway, "websocket request canceled", cause) } } diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go index d41abaac3..4dae079b5 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go @@ -52,6 +52,7 @@ type RelayTurnResult struct { type RelayExit struct { Stage string Err error + Graceful bool WroteDownstream bool } @@ -65,6 +66,9 @@ type RelayOptions struct { OnUsageParseFailure func(eventType string, usageRaw string) OnTurnComplete func(turn RelayTurnResult) BeforeWriteClient func(msgType coderws.MessageType, payload []byte, wroteDownstream bool) error + BeforeClientWrite func(msgType coderws.MessageType, payload []byte) + AfterClientWrite func(msgType coderws.MessageType, payload []byte, writeErr error) + BeforeRelayCancel func(exit RelayExit) ReadClientFrame func(ctx context.Context, clientConn FrameConn) (coderws.MessageType, []byte, error) OnTrace func(event RelayTraceEvent) Now func() time.Time @@ -226,7 +230,9 @@ func Relay( options.OnUsageParseFailure, options.OnTurnComplete, options.BeforeWriteClient, - func() { + options.BeforeClientWrite, + options.AfterClientWrite, + func(msgType coderws.MessageType, payload []byte) { if options.StartClientAfterFirstDownstream { startClientReader() } @@ -241,6 +247,13 @@ func Relay( go runIdleWatchdog(relayCtx, nowFn, options.IdleTimeout, &lastActivity, onTrace, exitCh) firstExit := <-exitCh + // An outer ingress cancellation is a control-plane close, not a graceful + // upstream disconnect. Leave the client connection open here so the + // adapter can emit the precise lease/request close code. Internal + // relayCancel does not cancel ctx and therefore does not take this path. + if ctx.Err() != nil { + firstExit.graceful = false + } emitRelayTrace(onTrace, RelayTraceEvent{ Stage: "first_exit", Direction: relayDirectionFromStage(firstExit.stage), @@ -248,6 +261,14 @@ func Relay( WroteDownstream: firstExit.wroteDownstream, Error: relayErrorString(firstExit.err), }) + if options.BeforeRelayCancel != nil { + options.BeforeRelayCancel(RelayExit{ + Stage: firstExit.stage, + Err: firstExit.err, + Graceful: firstExit.graceful, + WroteDownstream: firstExit.wroteDownstream, + }) + } combinedWroteDownstream := firstExit.wroteDownstream secondExit := relayExitSignal{graceful: true} hasSecondExit := false @@ -422,7 +443,9 @@ func runUpstreamToClient( onUsageParseFailure func(eventType string, usageRaw string), onTurnComplete func(turn RelayTurnResult), beforeWriteClient func(msgType coderws.MessageType, payload []byte, wroteDownstream bool) error, - afterWriteClient func(), + beforeClientWrite func(msgType coderws.MessageType, payload []byte), + afterClientWrite func(msgType coderws.MessageType, payload []byte, writeErr error), + afterWriteClient func(msgType coderws.MessageType, payload []byte), dropDownstreamWrites *atomic.Bool, forwardedFrames *atomic.Int64, droppedFrames *atomic.Int64, @@ -498,21 +521,28 @@ func runUpstreamToClient( markActivity() continue } - if err := writeClient(msgType, payload); err != nil { + if beforeClientWrite != nil { + beforeClientWrite(msgType, payload) + } + writeErr := writeClient(msgType, payload) + if afterClientWrite != nil { + afterClientWrite(msgType, payload, writeErr) + } + if writeErr != nil { emitRelayTrace(onTrace, RelayTraceEvent{ Stage: "write_client_failed", Direction: "upstream_to_client", MessageType: relayMessageTypeString(msgType), PayloadBytes: len(payload), WroteDownstream: wroteDownstream, - Error: err.Error(), + Error: writeErr.Error(), }) - exitCh <- relayExitSignal{stage: "write_client", err: err, wroteDownstream: wroteDownstream} + exitCh <- relayExitSignal{stage: "write_client", err: writeErr, wroteDownstream: wroteDownstream} return } wroteDownstream = true if afterWriteClient != nil { - afterWriteClient() + afterWriteClient(msgType, payload) } if forwardedFrames != nil { forwardedFrames.Add(1) diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go index f92117a2b..6036be745 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go @@ -125,6 +125,8 @@ func TestRunUpstreamToClient_ErrorAndDropPaths(t *testing.T) { nil, nil, nil, + nil, + nil, drop, nil, nil, @@ -156,6 +158,8 @@ func TestRunUpstreamToClient_ErrorAndDropPaths(t *testing.T) { nil, nil, nil, + nil, + nil, drop, nil, nil, @@ -190,6 +194,8 @@ func TestRunUpstreamToClient_ErrorAndDropPaths(t *testing.T) { nil, nil, nil, + nil, + nil, drop, nil, dropped, diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_test.go index cdd41a058..a5d77fe0f 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay_test.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay_test.go @@ -31,6 +31,12 @@ type delayedReadFrameConn struct { once sync.Once } +type readStartSpyFrameConn struct { + base FrameConn + started chan struct{} + startOnce sync.Once +} + type closeSpyFrameConn struct { closeCalls atomic.Int32 } @@ -126,6 +132,19 @@ func (c *delayedReadFrameConn) Close() error { return c.base.Close() } +func (c *readStartSpyFrameConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) { + c.startOnce.Do(func() { close(c.started) }) + return c.base.ReadFrame(ctx) +} + +func (c *readStartSpyFrameConn) WriteFrame(ctx context.Context, msgType coderws.MessageType, payload []byte) error { + return c.base.WriteFrame(ctx, msgType, payload) +} + +func (c *readStartSpyFrameConn) Close() error { + return c.base.Close() +} + func (c *closeSpyFrameConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) { if ctx == nil { ctx = context.Background() @@ -661,6 +680,51 @@ func TestRelay_ContextCanceled(t *testing.T) { require.NotNil(t, relayExit) } +func TestRelay_DownstreamPreambleStartsClientReader(t *testing.T) { + clientBase := newPassthroughTestFrameConn(nil, false) + clientConn := &readStartSpyFrameConn{base: clientBase, started: make(chan struct{})} + upstreamConn := newPassthroughTestFrameConn(nil, false) + resultCh := make(chan *RelayExit, 1) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + go func() { + _, relayExit := Relay( + ctx, + clientConn, + upstreamConn, + []byte(`{"type":"response.create","model":"gpt-5.1"}`), + RelayOptions{ + StartClientAfterFirstDownstream: true, + }, + ) + resultCh <- relayExit + }() + + upstreamConn.readCh <- passthroughTestFrame{ + msgType: coderws.MessageText, + payload: []byte(`{"type":"response.created","response":{"id":"resp_semantic_gate"}}`), + } + require.Eventually(t, func() bool { return len(clientBase.Writes()) == 1 }, time.Second, 10*time.Millisecond) + select { + case <-clientConn.started: + case <-time.After(time.Second): + t.Fatal("response.created did not start the client reader") + } + + upstreamConn.readCh <- passthroughTestFrame{ + msgType: coderws.MessageText, + payload: []byte(`{"type":"response.completed","response":{"id":"resp_semantic_gate","usage":{"input_tokens":1,"output_tokens":1}}}`), + } + _ = upstreamConn.Close() + select { + case relayExit := <-resultCh: + require.Nil(t, relayExit) + case <-time.After(time.Second): + t.Fatal("relay did not finish after terminal event and upstream close") + } +} + func TestRelay_TraceEvents_ContainsLifecycleStages(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go index af1bf7456..f61cb572b 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go @@ -7,7 +7,9 @@ import ( "net/http" "net/url" "strings" + "sync" "sync/atomic" + "time" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" @@ -18,7 +20,11 @@ import ( ) type openAIWSClientFrameConn struct { - conn *coderws.Conn + conn *coderws.Conn + controlCtx context.Context + interTurnIdleTimeout time.Duration + interTurnStarted chan struct{} + waitingForNextTurn atomic.Bool } // openAIWSPolicyEnforcingFrameConn wraps a client-side FrameConn and runs @@ -193,16 +199,379 @@ func openAIWSPassthroughRequestModelFromSessionFrame(payload []byte) string { const openaiWSV2PassthroughModeFields = "ws_mode=passthrough ws_router=v2" -var _ openaiwsv2.FrameConn = (*openAIWSClientFrameConn)(nil) +var errOpenAIWSPassthroughFirstOutputTimeout = errors.New("openai websocket passthrough first output timeout") +var errOpenAIWSPassthroughActiveTurnTimeout = errors.New("openai websocket passthrough active turn read timeout") -func (c *openAIWSClientFrameConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) { - if c == nil || c.conn == nil { +type openAIWSPassthroughDeadlinePhase uint8 + +const ( + openAIWSPassthroughDeadlinePhaseFirstSemantic openAIWSPassthroughDeadlinePhase = iota + 1 + openAIWSPassthroughDeadlinePhaseActiveRead +) + +type openAIWSPassthroughFirstOutputDeadline struct { + timeout time.Duration + startedAt time.Time + requestModel string + reasoningEffort string + phase openAIWSPassthroughDeadlinePhase +} + +type openAIWSPassthroughFirstOutputTimeoutError struct { + deadline openAIWSPassthroughFirstOutputDeadline +} + +func (e *openAIWSPassthroughFirstOutputTimeoutError) Error() string { + return errOpenAIWSPassthroughFirstOutputTimeout.Error() +} + +func (e *openAIWSPassthroughFirstOutputTimeoutError) Unwrap() error { + return errOpenAIWSPassthroughFirstOutputTimeout +} + +type openAIWSPassthroughActiveTurnTimeoutError struct{} + +func (e *openAIWSPassthroughActiveTurnTimeoutError) Error() string { + return errOpenAIWSPassthroughActiveTurnTimeout.Error() +} + +func (e *openAIWSPassthroughActiveTurnTimeoutError) Unwrap() error { + return errOpenAIWSPassthroughActiveTurnTimeout +} + +type openAIWSPassthroughFirstOutputDeadlineState struct { + armed bool + generation uint64 + deadline openAIWSPassthroughFirstOutputDeadline +} + +type openAIWSPassthroughTurnLifecycle struct { + mu sync.Mutex + inFlight bool +} + +func newOpenAIWSPassthroughTurnLifecycle(inFlight bool) *openAIWSPassthroughTurnLifecycle { + return &openAIWSPassthroughTurnLifecycle{inFlight: inFlight} +} + +func (l *openAIWSPassthroughTurnLifecycle) beginResponseCreate(onAccepted func()) bool { + if l == nil { + return false + } + l.mu.Lock() + defer l.mu.Unlock() + if l.inFlight { + return false + } + l.inFlight = true + if onAccepted != nil { + onAccepted() + } + return true +} + +func (l *openAIWSPassthroughTurnLifecycle) cancelResponseCreate() { + if l == nil { + return + } + l.mu.Lock() + l.inFlight = false + l.mu.Unlock() +} + +func (l *openAIWSPassthroughTurnLifecycle) beginTerminalWrite() { + if l != nil { + l.mu.Lock() + } +} + +func (l *openAIWSPassthroughTurnLifecycle) finishTerminalWrite(succeeded bool, onSucceeded func()) { + if l == nil { + return + } + if succeeded { + if onSucceeded != nil { + onSucceeded() + } + l.inFlight = false + } + l.mu.Unlock() +} + +type openAIWSPassthroughFirstOutputFrameConn struct { + inner openaiwsv2.FrameConn + resolveDeadline func(payload []byte) openAIWSPassthroughFirstOutputDeadline + activeReadTimeout time.Duration + + mu sync.Mutex + state openAIWSPassthroughFirstOutputDeadlineState + deadlineChanged chan struct{} +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) { + if c == nil || c.inner == nil { return coderws.MessageText, nil, errOpenAIWSConnClosed } if ctx == nil { ctx = context.Background() } - return c.conn.Read(ctx) + + type readResult struct { + msgType coderws.MessageType + payload []byte + err error + } + readCtx, cancelRead := context.WithCancel(ctx) + readResultCh := make(chan readResult, 1) + go func() { + msgType, payload, err := c.inner.ReadFrame(readCtx) + readResultCh <- readResult{msgType: msgType, payload: payload, err: err} + }() + + var timer *time.Timer + var timerCh <-chan time.Time + resetTimer := func() { + state := c.deadlineState() + if timer != nil { + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + } + if !state.armed || state.deadline.timeout <= 0 { + timerCh = nil + return + } + remaining := time.Until(state.deadline.startedAt.Add(state.deadline.timeout)) + if remaining < 0 { + remaining = 0 + } + if timer == nil { + timer = time.NewTimer(remaining) + } else { + timer.Reset(remaining) + } + timerCh = timer.C + } + resetTimer() + + defer func() { + cancelRead() + if timer != nil { + timer.Stop() + } + }() + for { + select { + case result := <-readResultCh: + if result.err == nil { + c.observeUpstreamActivity(result.msgType, result.payload) + } + return result.msgType, result.payload, result.err + case <-c.deadlineChanged: + resetTimer() + case <-timerCh: + state := c.deadlineState() + if !state.armed || state.deadline.timeout <= 0 || time.Now().Before(state.deadline.startedAt.Add(state.deadline.timeout)) { + resetTimer() + continue + } + if ctx.Err() != nil { + cancelRead() + <-readResultCh + return coderws.MessageText, nil, ctx.Err() + } + cancelRead() + <-readResultCh + if state.deadline.phase == openAIWSPassthroughDeadlinePhaseActiveRead { + return coderws.MessageText, nil, &openAIWSPassthroughActiveTurnTimeoutError{} + } + return coderws.MessageText, nil, &openAIWSPassthroughFirstOutputTimeoutError{deadline: state.deadline} + case <-ctx.Done(): + cancelRead() + <-readResultCh + return coderws.MessageText, nil, ctx.Err() + } + } +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) WriteFrame(ctx context.Context, msgType coderws.MessageType, payload []byte) error { + if c == nil || c.inner == nil { + return errOpenAIWSConnClosed + } + generation := uint64(0) + if msgType == coderws.MessageText && strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" { + generation = c.armDeadline(payload) + } + if err := c.inner.WriteFrame(ctx, msgType, payload); err != nil { + c.disarmDeadline(generation) + return err + } + return nil +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) Close() error { + if c == nil || c.inner == nil { + return nil + } + return c.inner.Close() +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) armDeadline(payload []byte) uint64 { + if c == nil || c.resolveDeadline == nil { + return 0 + } + deadline := c.resolveDeadline(payload) + if deadline.timeout <= 0 { + return 0 + } + if deadline.startedAt.IsZero() { + deadline.startedAt = time.Now() + } + deadline.phase = openAIWSPassthroughDeadlinePhaseFirstSemantic + c.mu.Lock() + c.state.generation++ + generation := c.state.generation + c.state.armed = true + c.state.deadline = deadline + c.mu.Unlock() + c.notifyDeadlineChanged() + return generation +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) observeUpstreamActivity(msgType coderws.MessageType, payload []byte) { + if c == nil { + return + } + if msgType == coderws.MessageText && openAIWSPassthroughIsTerminalOutput(payload) { + c.disarmDeadline(0) + return + } + state := c.deadlineState() + if state.armed && state.deadline.phase == openAIWSPassthroughDeadlinePhaseActiveRead { + c.armActiveReadDeadline() + return + } + if msgType == coderws.MessageText && openAIWSPassthroughStartsSemanticOutput(payload) { + c.armActiveReadDeadline() + } +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) armActiveReadDeadline() { + if c == nil { + return + } + if c.activeReadTimeout <= 0 { + c.disarmDeadline(0) + return + } + c.mu.Lock() + c.state.generation++ + c.state.armed = true + c.state.deadline = openAIWSPassthroughFirstOutputDeadline{ + timeout: c.activeReadTimeout, + startedAt: time.Now(), + phase: openAIWSPassthroughDeadlinePhaseActiveRead, + } + c.mu.Unlock() + c.notifyDeadlineChanged() +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) disarmDeadline(generation uint64) { + if c == nil { + return + } + c.mu.Lock() + if !c.state.armed || (generation != 0 && generation != c.state.generation) { + c.mu.Unlock() + return + } + c.state.armed = false + c.mu.Unlock() + c.notifyDeadlineChanged() +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) deadlineState() openAIWSPassthroughFirstOutputDeadlineState { + if c == nil { + return openAIWSPassthroughFirstOutputDeadlineState{} + } + c.mu.Lock() + defer c.mu.Unlock() + return c.state +} + +func (c *openAIWSPassthroughFirstOutputFrameConn) notifyDeadlineChanged() { + if c == nil || c.deadlineChanged == nil { + return + } + select { + case c.deadlineChanged <- struct{}{}: + default: + } +} + +func openAIWSPassthroughStartsSemanticOutput(payload []byte) bool { + eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String()) + switch eventType { + case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled": + return true + case "", "response.created", "response.in_progress", "response.output_item.added", "response.output_item.done": + return false + } + return strings.Contains(eventType, ".delta") || + strings.HasPrefix(eventType, "response.output_text") || + strings.HasPrefix(eventType, "response.output") +} + +func openAIWSPassthroughIsTerminalOutput(payload []byte) bool { + switch strings.TrimSpace(gjson.GetBytes(payload, "type").String()) { + case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled": + return true + default: + return false + } +} + +var _ openaiwsv2.FrameConn = (*openAIWSClientFrameConn)(nil) +var _ openaiwsv2.FrameConn = (*openAIWSPassthroughFirstOutputFrameConn)(nil) + +func (c *openAIWSClientFrameConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) { + if c == nil || c.conn == nil { + return coderws.MessageText, nil, errOpenAIWSConnClosed + } + controlCtx := ctx + if c.controlCtx != nil { + controlCtx = c.controlCtx + } + msgType, payload, err := readOpenAIWSClientMessageWithTimeoutStart( + controlCtx, + c.conn, + c.interTurnIdleTimeout, + coderws.StatusNormalClosure, + "websocket idle timeout", + c.interTurnStarted, + func() bool { return c.waitingForNextTurn.Load() }, + ) + return msgType, payload, err +} + +func (c *openAIWSClientFrameConn) markTurnStarted() { + if c != nil { + c.waitingForNextTurn.Store(false) + } +} + +func (c *openAIWSClientFrameConn) markTurnCompleted() { + if c == nil { + return + } + c.waitingForNextTurn.Store(true) + select { + case c.interTurnStarted <- struct{}{}: + default: + } } func (c *openAIWSClientFrameConn) WriteFrame(ctx context.Context, msgType coderws.MessageType, payload []byte) error { @@ -429,19 +798,68 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( if !ok { return errors.New("openai ws passthrough upstream connection does not support frame relay") } + relayUpstreamFrameConn := &openAIWSPassthroughFirstOutputFrameConn{ + inner: upstreamFrameConn, + activeReadTimeout: s.openAIWSPassthroughIdleTimeout(), + deadlineChanged: make(chan struct{}, 1), + resolveDeadline: func(payload []byte) openAIWSPassthroughFirstOutputDeadline { + reasoningEffort := "" + if current := usageMeta.reasoningEffort.Load(); current != nil { + reasoningEffort = *current + } + timeout := s.openAIFirstOutputTimeout(reasoningEffort) + if timeout <= 0 { + timeout = s.openAIWSPassthroughIdleTimeout() + } + model := openAIWSPassthroughRequestModelForFrame(payload) + if model == "" { + model = usageMeta.requestModelForFrame(payload) + } + if model == "" { + model = requestModel + } + return openAIWSPassthroughFirstOutputDeadline{ + timeout: timeout, + startedAt: time.Now(), + requestModel: model, + reasoningEffort: reasoningEffort, + } + }, + } completedTurns := atomic.Int32{} + turnLifecycle := newOpenAIWSPassthroughTurnLifecycle(true) + clientFrameConn := &openAIWSClientFrameConn{ + conn: clientConn, + controlCtx: ctx, + interTurnIdleTimeout: s.openAIWSIngressInterTurnIdleTimeout(), + interTurnStarted: make(chan struct{}, 1), + } policyClientConn := &openAIWSPolicyEnforcingFrameConn{ - inner: &openAIWSClientFrameConn{conn: clientConn}, + inner: clientFrameConn, // 注意线程安全:filter 仅在 runClientToUpstream 这一条 // goroutine 中被调用(passthrough_relay.go: ReadFrame loop), // capturedSessionModel 的读写都发生在该 goroutine 内,因此无需 // 加锁/原子化。 - filter: func(msgType coderws.MessageType, payload []byte) ([]byte, *OpenAIFastBlockedError, error) { + filter: func(msgType coderws.MessageType, payload []byte) (out []byte, blocked *OpenAIFastBlockedError, filterErr error) { if msgType != coderws.MessageText { return payload, nil, nil } - if strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" { + eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String()) + isResponseCreate := eventType == "response.create" + acceptedTurn := false + if isResponseCreate { + if !turnLifecycle.beginResponseCreate(clientFrameConn.markTurnStarted) { + err := errors.New("overlapping response.create is not supported") + return payload, nil, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, err.Error(), err) + } + defer func() { + if !acceptedTurn { + turnLifecycle.cancelResponseCreate() + } + }() + } + if isResponseCreate { if account.IsOpenAIOAuth() && isOpenAIResponsesLiteWebSocketPayload(payload) { litePayload, _, liteErr := normalizeOpenAIResponsesLiteToolsPayload(payload) if liteErr != nil { @@ -450,7 +868,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( payload = litePayload } } - if strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" && hooks != nil && hooks.BeforeRequest != nil { + if isResponseCreate && hooks != nil && hooks.BeforeRequest != nil { turnNo := int(completedTurns.Load()) + 1 if turnNo < 2 { turnNo = 2 @@ -499,9 +917,9 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( // extractOpenAIServiceTierFromBody 返回 nil;这里有意 // 覆盖(Store(nil)),因为 OpenAI 上游对该帧实际不传 // service_tier 时按 default 处理,billing 应如实反映。 - if policyErr == nil && blocked == nil && - strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" { + if policyErr == nil && blocked == nil && isResponseCreate { usageMeta.updateFromResponseCreate(out, model, requestModelForThisFrame) + acceptedTurn = true } return out, blocked, policyErr }, @@ -521,7 +939,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( } upstreamFirstMessageSent := false firstWriteCtx, cancelFirstWrite := context.WithTimeout(ctx, s.openAIWSWriteTimeout()) - firstWriteErr := upstreamFrameConn.WriteFrame(firstWriteCtx, coderws.MessageText, firstClientMessage) + firstWriteErr := relayUpstreamFrameConn.WriteFrame(firstWriteCtx, coderws.MessageText, firstClientMessage) cancelFirstWrite() if firstWriteErr != nil { return wrapOpenAIWSIngressTurnError( @@ -550,11 +968,14 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( relayResult, relayExit := openaiwsv2.RunEntry(openaiwsv2.EntryInput{ Ctx: ctx, ClientConn: policyClientConn, - UpstreamConn: upstreamFrameConn, + UpstreamConn: relayUpstreamFrameConn, FirstClientMessage: firstClientMessage, Options: openaiwsv2.RelayOptions{ - WriteTimeout: s.openAIWSWriteTimeout(), - IdleTimeout: s.openAIWSPassthroughIdleTimeout(), + WriteTimeout: s.openAIWSWriteTimeout(), + // Passthrough idle is enforced only after a completed turn by + // clientFrameConn. The relay-wide activity watchdog would also + // terminate a healthy active upstream turn. + IdleTimeout: 0, FirstMessageType: coderws.MessageText, FirstMessageSent: upstreamFirstMessageSent, StartClientAfterFirstDownstream: true, @@ -603,6 +1024,27 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( hooks.AfterTurn(turnNo, turnResult, nil) } }, + BeforeClientWrite: func(msgType coderws.MessageType, payload []byte) { + if msgType == coderws.MessageText && openAIWSPassthroughIsTerminalOutput(payload) { + turnLifecycle.beginTerminalWrite() + } + }, + AfterClientWrite: func(msgType coderws.MessageType, payload []byte, writeErr error) { + if msgType == coderws.MessageText && openAIWSPassthroughIsTerminalOutput(payload) { + turnLifecycle.finishTerminalWrite(writeErr == nil, clientFrameConn.markTurnCompleted) + } + }, + BeforeRelayCancel: func(exit openaiwsv2.RelayExit) { + if context.Cause(ctx) != nil { + return + } + status, reason, ok := openAIWSPassthroughRelayClientClose(exit, int(completedTurns.Load())) + if !ok { + return + } + _ = clientConn.Close(status, reason) + _ = clientConn.CloseNow() + }, BeforeWriteClient: func(msgType coderws.MessageType, payload []byte, wroteDownstream bool) error { if msgType != coderws.MessageText { return nil @@ -650,6 +1092,17 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( }, }, }) + if cause := context.Cause(ctx); cause != nil { + status := coderws.StatusGoingAway + reason := "websocket request canceled" + if errors.Is(cause, ErrOpenAIWSIngressLeaseLost) { + status = coderws.StatusTryAgainLater + reason = "websocket ingress capacity lease lost; please reconnect" + } + _ = clientConn.Close(status, reason) + _ = clientConn.CloseNow() + return NewOpenAIWSClientCloseError(status, reason, cause) + } result := &OpenAIForwardResult{ RequestID: relayResult.RequestID, @@ -704,6 +1157,41 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( ) relayErr := relayExit.Err + var firstOutputTimeoutErr *openAIWSPassthroughFirstOutputTimeoutError + if errors.As(relayErr, &firstOutputTimeoutErr) { + deadline := firstOutputTimeoutErr.deadline + failoverErr := s.newOpenAIFirstOutputTimeoutError( + ctx, + c, + account, + deadline.startedAt, + deadline.requestModel, + deadline.reasoningEffort, + deadline.timeout, + "websocket_first_semantic_output", + handshakeHeaders, + ) + if turnCount == 0 && !relayExit.WroteDownstream { + relayErr = failoverErr + } else { + // The handler only retains the initial response.create across + // account attempts. Replaying it after a later-turn timeout would + // duplicate the first turn, so later turns end the client session. + relayErr = NewOpenAIWSClientCloseError( + coderws.StatusGoingAway, + "upstream produced no semantic output; please reconnect", + firstOutputTimeoutErr, + ) + } + } + var activeTurnTimeoutErr *openAIWSPassthroughActiveTurnTimeoutError + if errors.As(relayErr, &activeTurnTimeoutErr) { + relayErr = NewOpenAIWSClientCloseError( + coderws.StatusGoingAway, + "upstream websocket read timeout; please reconnect", + activeTurnTimeoutErr, + ) + } if relayExit.Stage == "idle_timeout" { relayErr = NewOpenAIWSClientCloseError( coderws.StatusPolicyViolation, @@ -722,6 +1210,28 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( return turnErr } +func openAIWSPassthroughRelayClientClose(exit openaiwsv2.RelayExit, completedTurns int) (coderws.StatusCode, string, bool) { + var closeErr *OpenAIWSClientCloseError + if errors.As(exit.Err, &closeErr) { + return closeErr.StatusCode(), closeErr.Reason(), true + } + var activeTurnTimeoutErr *openAIWSPassthroughActiveTurnTimeoutError + if errors.As(exit.Err, &activeTurnTimeoutErr) { + return coderws.StatusGoingAway, "upstream websocket read timeout; please reconnect", true + } + var firstOutputTimeoutErr *openAIWSPassthroughFirstOutputTimeoutError + if errors.As(exit.Err, &firstOutputTimeoutErr) { + if completedTurns > 0 || exit.WroteDownstream { + return coderws.StatusGoingAway, "upstream produced no semantic output; please reconnect", true + } + return 0, "", false + } + if !exit.Graceful && exit.Stage == "read_upstream" { + return coderws.StatusInternalError, "upstream websocket proxy failed", true + } + return 0, "", false +} + func (s *OpenAIGatewayService) mapOpenAIWSPassthroughDialError( err error, statusCode int, diff --git a/backend/internal/service/openai_ws_v2_passthrough_lifecycle_test.go b/backend/internal/service/openai_ws_v2_passthrough_lifecycle_test.go new file mode 100644 index 000000000..de3d5fdf2 --- /dev/null +++ b/backend/internal/service/openai_ws_v2_passthrough_lifecycle_test.go @@ -0,0 +1,591 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + coderws "github.com/coder/websocket" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +type stagedPassthroughFrame struct { + messageType coderws.MessageType + payload []byte +} + +type stagedPassthroughConn struct { + frames chan stagedPassthroughFrame + writes chan []byte + closed chan struct{} + closeOnce sync.Once +} + +func newStagedPassthroughConn() *stagedPassthroughConn { + return &stagedPassthroughConn{ + frames: make(chan stagedPassthroughFrame, 4), + writes: make(chan []byte, 4), + closed: make(chan struct{}), + } +} + +func (c *stagedPassthroughConn) Send(payload string) { + c.frames <- stagedPassthroughFrame{messageType: coderws.MessageText, payload: []byte(payload)} +} + +func (c *stagedPassthroughConn) WriteJSON(context.Context, any) error { return nil } + +func (c *stagedPassthroughConn) ReadMessage(ctx context.Context) ([]byte, error) { + _, payload, err := c.ReadFrame(ctx) + return payload, err +} + +func (c *stagedPassthroughConn) Ping(context.Context) error { return nil } + +func (c *stagedPassthroughConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) { + if ctx == nil { + ctx = context.Background() + } + select { + case <-ctx.Done(): + return coderws.MessageText, nil, ctx.Err() + case <-c.closed: + return coderws.MessageText, nil, errOpenAIWSConnClosed + case frame := <-c.frames: + return frame.messageType, append([]byte(nil), frame.payload...), nil + } +} + +func (c *stagedPassthroughConn) WriteFrame(ctx context.Context, _ coderws.MessageType, payload []byte) error { + if ctx == nil { + ctx = context.Background() + } + select { + case <-ctx.Done(): + return ctx.Err() + case <-c.closed: + return errOpenAIWSConnClosed + default: + } + var parsed any + if err := json.Unmarshal(payload, &parsed); err != nil { + return err + } + select { + case c.writes <- append([]byte(nil), payload...): + case <-ctx.Done(): + return ctx.Err() + case <-c.closed: + return errOpenAIWSConnClosed + } + return nil +} + +func (c *stagedPassthroughConn) Close() error { + c.closeOnce.Do(func() { close(c.closed) }) + return nil +} + +type stagedPassthroughDialer struct { + conn openAIWSClientConn +} + +func (d *stagedPassthroughDialer) Dial(context.Context, string, http.Header, string) (openAIWSClientConn, int, http.Header, error) { + return d.conn, http.StatusSwitchingProtocols, http.Header{}, nil +} + +func newPassthroughLifecycleService(cfg *config.Config, upstream *stagedPassthroughConn) *OpenAIGatewayService { + return &OpenAIGatewayService{ + cfg: cfg, + httpUpstream: &httpUpstreamRecorder{}, + cache: &stubGatewayCache{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + toolCorrector: NewCodexToolCorrector(), + openaiWSPassthroughDialer: &stagedPassthroughDialer{conn: upstream}, + } +} + +func passthroughLifecycleConfig() *config.Config { + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIFirstOutputTimeoutSeconds = 1 + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true + cfg.Gateway.OpenAIWS.IngressModeDefault = OpenAIWSIngressModeCtxPool + cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 1 + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 1 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 + return cfg +} + +func passthroughLifecycleAccount() *Account { + return &Account{ + ID: 901, + Name: "passthrough-lifecycle", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{"api_key": "sk-test"}, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_mode": OpenAIWSIngressModePassthrough, + }, + } +} + +func startPassthroughLifecycleServer( + t *testing.T, + controlCtx context.Context, + svc *OpenAIGatewayService, + account *Account, +) (*httptest.Server, <-chan error) { + t.Helper() + serverErr := make(chan error, 1) + server := 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 { + serverErr <- err + return + } + defer func() { _ = conn.CloseNow() }() + + msgType, firstMessage, err := ReadOpenAIWSClientMessage( + controlCtx, + conn, + 3*time.Second, + coderws.StatusPolicyViolation, + "missing first response.create message", + ) + if err != nil { + serverErr <- err + return + } + if msgType != coderws.MessageText { + serverErr <- errors.New("first message was not text") + return + } + + recorder := httptest.NewRecorder() + ginCtx, _ := gin.CreateTestContext(recorder) + req := r.Clone(controlCtx) + req.Header = req.Header.Clone() + ginCtx.Request = req + serverErr <- svc.ProxyResponsesWebSocketFromClient(controlCtx, ginCtx, conn, account, "sk-test", firstMessage, nil) + })) + return server, serverErr +} + +func dialPassthroughLifecycleClient(t *testing.T, server *httptest.Server) *coderws.Conn { + t.Helper() + dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) + clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(server.URL, "http"), nil) + cancelDial() + require.NoError(t, err) + 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) + return clientConn +} + +func readPassthroughLifecycleFrame(t *testing.T, clientConn *coderws.Conn, timeout time.Duration) ([]byte, error) { + t.Helper() + readCtx, cancelRead := context.WithTimeout(context.Background(), timeout) + _, payload, err := clientConn.Read(readCtx) + cancelRead() + return payload, err +} + +func requirePassthroughUpstreamWrite(t *testing.T, upstream *stagedPassthroughConn, timeout time.Duration) []byte { + t.Helper() + select { + case payload := <-upstream.writes: + return payload + case <-time.After(timeout): + t.Fatal("passthrough request was not forwarded upstream") + return nil + } +} + +func TestOpenAIWSPassthroughTurnLifecycle_SerializesTerminalCommitAndNextTurn(t *testing.T) { + clientFrameConn := &openAIWSClientFrameConn{interTurnStarted: make(chan struct{}, 1)} + clientFrameConn.markTurnCompleted() + lifecycle := newOpenAIWSPassthroughTurnLifecycle(true) + lifecycle.beginTerminalWrite() + + admitted := make(chan bool, 1) + go func() { + admitted <- lifecycle.beginResponseCreate(clientFrameConn.markTurnStarted) + }() + select { + case <-admitted: + t.Fatal("next response.create was admitted before terminal commit completed") + case <-time.After(50 * time.Millisecond): + } + + lifecycle.finishTerminalWrite(true, clientFrameConn.markTurnCompleted) + select { + case ok := <-admitted: + require.True(t, ok) + case <-time.After(time.Second): + t.Fatal("next response.create remained blocked after terminal commit") + } + require.False(t, clientFrameConn.waitingForNextTurn.Load(), "accepted next turn must win over terminal idle state") + + lifecycle = newOpenAIWSPassthroughTurnLifecycle(true) + lifecycle.beginTerminalWrite() + admitted = make(chan bool, 1) + go func() { + admitted <- lifecycle.beginResponseCreate(nil) + }() + lifecycle.finishTerminalWrite(false, func() { + t.Error("failed terminal write must not commit idle state") + }) + require.False(t, <-admitted, "failed terminal write must keep the current turn in flight") +} + +func TestPassthroughLifecycle_LeaseLossSendsRetryClose(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.created","response":{"id":"resp_lease","model":"gpt-5.1"}}`) + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + + event, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "response.created", gjson.GetBytes(event, "type").String()) + cancelControl(ErrOpenAIWSIngressLeaseLost) + + _, err = readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + 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: + case <-time.After(3 * time.Second): + t.Fatal("passthrough lease-loss reader did not exit") + } +} + +func TestPassthroughLifecycle_CompletedTurnStartsInterTurnIdle(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.completed","response":{"id":"resp_idle","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`) + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + + event, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String()) + _, err = readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + var closeErr coderws.CloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusNormalClosure, closeErr.Code) + require.Equal(t, "websocket idle timeout", closeErr.Reason) + select { + case <-serverErr: + case <-time.After(3 * time.Second): + t.Fatal("passthrough idle reader did not exit") + } +} + +func TestPassthroughLifecycle_ActiveTurnInactivityUsesReadTimeout(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.output_text.delta","response_id":"resp_active","delta":"hello"}`) + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + + delta, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String()) + _, err = readPassthroughLifecycleFrame(t, clientConn, 2500*time.Millisecond) + var websocketCloseErr coderws.CloseError + require.ErrorAs(t, err, &websocketCloseErr) + require.Equal(t, coderws.StatusGoingAway, websocketCloseErr.Code) + require.Equal(t, "upstream websocket read timeout; please reconnect", websocketCloseErr.Reason) + select { + case err := <-serverErr: + var closeErr *OpenAIWSClientCloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusGoingAway, closeErr.StatusCode()) + require.Equal(t, "upstream websocket read timeout; please reconnect", closeErr.Reason()) + case <-time.After(2500 * time.Millisecond): + t.Fatal("passthrough active turn remained unbounded after upstream activity stopped") + } +} + +func TestPassthroughLifecycle_PreambleAllowsPromptClientCancel(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + cfg := passthroughLifecycleConfig() + cfg.Gateway.OpenAIFirstOutputTimeoutSeconds = 3 + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.created","response":{"id":"resp_cancel","model":"gpt-5.1"}}`) + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(cfg, upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + require.Equal(t, "response.create", gjson.GetBytes(requirePassthroughUpstreamWrite(t, upstream, time.Second), "type").String()) + + created, err := readPassthroughLifecycleFrame(t, clientConn, time.Second) + require.NoError(t, err) + require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) + writeCtx, cancelWrite := context.WithTimeout(context.Background(), time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.cancel","response_id":"resp_cancel"}`)) + cancelWrite() + require.NoError(t, err) + cancelFrame := requirePassthroughUpstreamWrite(t, upstream, 500*time.Millisecond) + require.Equal(t, "response.cancel", gjson.GetBytes(cancelFrame, "type").String()) + + require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done")) + select { + case <-serverErr: + case <-time.After(3 * time.Second): + t.Fatal("passthrough cancel test did not exit") + } +} + +func TestPassthroughLifecycle_RejectsOverlappingResponseCreate(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + cfg := passthroughLifecycleConfig() + cfg.Gateway.OpenAIFirstOutputTimeoutSeconds = 3 + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.created","response":{"id":"resp_overlap_first","model":"gpt-5.1"}}`) + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(cfg, upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + require.Equal(t, "response.create", gjson.GetBytes(requirePassthroughUpstreamWrite(t, upstream, time.Second), "type").String()) + + created, err := readPassthroughLifecycleFrame(t, clientConn, time.Second) + require.NoError(t, err) + require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) + writeCtx, cancelWrite := context.WithTimeout(context.Background(), time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1"}`)) + cancelWrite() + require.NoError(t, err) + + _, err = readPassthroughLifecycleFrame(t, clientConn, time.Second) + var websocketCloseErr coderws.CloseError + require.ErrorAs(t, err, &websocketCloseErr) + require.Equal(t, coderws.StatusPolicyViolation, websocketCloseErr.Code) + require.Equal(t, "overlapping response.create is not supported", websocketCloseErr.Reason) + select { + case err := <-serverErr: + var closeErr *OpenAIWSClientCloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusPolicyViolation, closeErr.StatusCode()) + require.Equal(t, "overlapping response.create is not supported", closeErr.Reason()) + case <-time.After(3 * time.Second): + t.Fatal("overlapping response.create did not terminate passthrough") + } +} + +func TestPassthroughLifecycle_ActiveTurnActivityRefreshesReadTimeout(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.output_text.delta","response_id":"resp_active_refresh","delta":"one"}`) + go func() { + for _, event := range []string{ + `{"type":"response.output_text.delta","response_id":"resp_active_refresh","delta":"two"}`, + `{"type":"response.output_text.delta","response_id":"resp_active_refresh","delta":"three"}`, + `{"type":"response.completed","response":{"id":"resp_active_refresh","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":3}}}`, + } { + timer := time.NewTimer(600 * time.Millisecond) + <-timer.C + timer.Stop() + upstream.Send(event) + } + }() + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + + for _, wantType := range []string{ + "response.output_text.delta", + "response.output_text.delta", + "response.output_text.delta", + "response.completed", + } { + frame, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, wantType, gjson.GetBytes(frame, "type").String()) + } + require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done")) + select { + case <-serverErr: + case <-time.After(3 * time.Second): + t.Fatal("passthrough active-turn refresh test did not exit") + } +} + +func TestPassthroughLifecycle_TerminalSwitchesToInterTurnIdleTimeout(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + cfg := passthroughLifecycleConfig() + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 1 + cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 2 + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.completed","response":{"id":"resp_idle_first","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`) + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(cfg, upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + require.Equal(t, "response.create", gjson.GetBytes(requirePassthroughUpstreamWrite(t, upstream, 3*time.Second), "type").String()) + + completed, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "resp_idle_first", gjson.GetBytes(completed, "response.id").String()) + time.Sleep(1300 * time.Millisecond) + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","previous_response_id":"resp_idle_first"}`)) + cancelWrite() + require.NoError(t, err) + require.Equal(t, "response.create", gjson.GetBytes(requirePassthroughUpstreamWrite(t, upstream, 3*time.Second), "type").String()) + upstream.Send(`{"type":"response.completed","response":{"id":"resp_idle_second","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`) + completed, err = readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "resp_idle_second", gjson.GetBytes(completed, "response.id").String()) + _, err = readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + var websocketCloseErr coderws.CloseError + require.ErrorAs(t, err, &websocketCloseErr) + require.Equal(t, coderws.StatusNormalClosure, websocketCloseErr.Code) + require.Equal(t, "websocket idle timeout", websocketCloseErr.Reason) + + select { + case err := <-serverErr: + var closeErr *OpenAIWSClientCloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusNormalClosure, closeErr.StatusCode()) + require.Equal(t, "websocket idle timeout", closeErr.Reason()) + case <-time.After(3 * time.Second): + t.Fatal("passthrough terminal turn did not use inter-turn idle timeout") + } +} + +func TestPassthroughLifecycle_FirstOutputTimeoutRemainsBounded(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + upstream := newStagedPassthroughConn() + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + + select { + case err := <-serverErr: + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusGatewayTimeout, failoverErr.StatusCode) + require.Contains(t, string(failoverErr.ResponseBody), "first_output_timeout") + case <-time.After(2500 * time.Millisecond): + t.Fatal("passthrough first output was left unbounded") + } +} + +func TestPassthroughLifecycle_ResponseCreatedTimeoutClosesWithoutFailover(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.created","response":{"id":"resp_preamble","model":"gpt-5.1"}}`) + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + + created, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) + _, err = readPassthroughLifecycleFrame(t, clientConn, 2500*time.Millisecond) + var websocketCloseErr coderws.CloseError + require.ErrorAs(t, err, &websocketCloseErr) + require.Equal(t, coderws.StatusGoingAway, websocketCloseErr.Code) + require.Equal(t, "upstream produced no semantic output; please reconnect", websocketCloseErr.Reason) + select { + case err := <-serverErr: + var failoverErr *UpstreamFailoverError + require.NotErrorAs(t, err, &failoverErr) + var closeErr *OpenAIWSClientCloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusGoingAway, closeErr.StatusCode()) + require.Equal(t, "upstream produced no semantic output; please reconnect", closeErr.Reason()) + case <-time.After(2500 * time.Millisecond): + t.Fatal("response.created timeout did not close the passthrough connection") + } +} + +func TestPassthroughLifecycle_SecondTurnTimeoutIsNotFailoverSafe(t *testing.T) { + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.completed","response":{"id":"resp_first","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`) + server, serverErr := startPassthroughLifecycleServer(t, controlCtx, newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream), passthroughLifecycleAccount()) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + + completed, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String()) + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","previous_response_id":"resp_first"}`)) + cancelWrite() + require.NoError(t, err) + upstream.Send(`{"type":"response.created","response":{"id":"resp_second","model":"gpt-5.1"}}`) + + created, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "response.created", gjson.GetBytes(created, "type").String()) + _, err = readPassthroughLifecycleFrame(t, clientConn, 2500*time.Millisecond) + var websocketCloseErr coderws.CloseError + require.ErrorAs(t, err, &websocketCloseErr) + require.Equal(t, coderws.StatusGoingAway, websocketCloseErr.Code) + require.Equal(t, "upstream produced no semantic output; please reconnect", websocketCloseErr.Reason) + select { + case err := <-serverErr: + var failoverErr *UpstreamFailoverError + require.NotErrorAs(t, err, &failoverErr, "handler must not replay the initial request on another account for a later-turn timeout") + var closeErr *OpenAIWSClientCloseError + require.ErrorAs(t, err, &closeErr) + require.Equal(t, coderws.StatusGoingAway, closeErr.StatusCode()) + case <-time.After(2500 * time.Millisecond): + t.Fatal("second turn first semantic output was left unbounded") + } +}