fix(openai-ws): enforce passthrough turn lifecycle

This commit is contained in:
王鹏
2026-07-18 19:17:53 +08:00
parent b1a6b80267
commit f0e0b7e6d8
7 changed files with 1493 additions and 39 deletions
@@ -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)
@@ -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)
}
}
@@ -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)
@@ -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,
@@ -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()
@@ -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,
@@ -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")
}
}