Merge pull request #5177 from r266-tech/fix-nonpassthrough-write-context

fix(openai-ws): preserve terminal event on lease loss
This commit is contained in:
Wesley Liddick
2026-08-04 16:28:22 +08:00
committed by GitHub
4 changed files with 229 additions and 3 deletions
@@ -1612,7 +1612,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
reqLog.Info("openai.websocket_ingress_started")
clientIP := ip.GetClientIP(c)
userAgent := strings.TrimSpace(c.GetHeader("User-Agent"))
ctx := c.Request.Context()
clientLifecycleCtx := c.Request.Context()
ctx := clientLifecycleCtx
maxIngressConnections := 0
if h.cfg != nil {
maxIngressConnections = h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey
@@ -2017,6 +2018,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
// openAIWSTurnPricing 的注释——绝不能用建连时刻初始化。
var turnPricing openAIWSTurnPricing
hooks := &service.OpenAIWSIngressHooks{
ClientLifecycleContext: clientLifecycleCtx,
InitialRequestModel: reqModel,
MaxReasoningEffort: maxReasoningEffort,
ReasoningEffortMappings: reasoningEffortMappings,
@@ -207,6 +207,10 @@ func (e *OpenAIWSClientCloseError) Reason() string {
// OpenAIWSIngressHooks 定义入站 WS 每个 turn 的生命周期回调。
type OpenAIWSIngressHooks struct {
// ClientLifecycleContext is the request context before an ingress lease
// adds its independent cancellation signal. Downstream writes bind to it
// so shutdown and disconnect cancellation remain direct during lease loss.
ClientLifecycleContext context.Context
// InitialRequestModel is the client-facing model from the first frame,
// before channel or account mapping. Ingress modes preserve it for usage
// attribution while MapRequestModel determines the upstream model.
@@ -25,6 +25,21 @@ func (s *OpenAIGatewayService) openAIWSIngressInterTurnIdleTimeout() time.Durati
return time.Duration(s.cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds) * time.Second
}
// newOpenAIWSDownstreamWriteContext binds writes directly to the client
// lifecycle while excluding the separate ingress-lease cancellation signal.
// This lets a lease-loss path finish its current client write before
// ReadOpenAIWSClientMessage sends the retryable close frame.
func newOpenAIWSDownstreamWriteContext(controlCtx context.Context, hooks *OpenAIWSIngressHooks, timeout time.Duration) (context.Context, context.CancelFunc) {
writeParent := controlCtx
if hooks != nil && hooks.ClientLifecycleContext != nil {
writeParent = hooks.ClientLifecycleContext
}
if writeParent == nil {
writeParent = context.Background()
}
return context.WithTimeout(writeParent, timeout)
}
func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
ctx context.Context,
c *gin.Context,
@@ -369,7 +384,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
// the kernel send buffer before any close frame is queued.
eventBytes := buildOpenAIFastPolicyBlockedWSEvent(blocked)
if eventBytes != nil {
writeCtx, cancel := context.WithTimeout(ctx, s.openAIWSWriteTimeout())
writeCtx, cancel := newOpenAIWSDownstreamWriteContext(ctx, hooks, s.openAIWSWriteTimeout())
_ = clientConn.Write(writeCtx, coderws.MessageText, eventBytes)
cancel()
}
@@ -396,7 +411,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
writeClientMessage := func(message []byte) error {
writeCtx, cancel := context.WithTimeout(ctx, s.openAIWSWriteTimeout())
writeCtx, cancel := newOpenAIWSDownstreamWriteContext(ctx, hooks, s.openAIWSWriteTimeout())
defer cancel()
return clientConn.Write(writeCtx, coderws.MessageText, message)
}
@@ -18,6 +18,80 @@ import (
"github.com/tidwall/gjson"
)
type openAIWSLeaseLossAfterReadConn struct {
*openAIWSCaptureConn
cancel context.CancelCauseFunc
once sync.Once
}
func (c *openAIWSLeaseLossAfterReadConn) ReadMessage(ctx context.Context) ([]byte, error) {
message, err := c.openAIWSCaptureConn.ReadMessage(ctx)
if err == nil {
c.once.Do(func() {
c.cancel(ErrOpenAIWSIngressLeaseLost)
})
}
return message, err
}
type openAIWSSingleConnDialer struct {
conn openAIWSClientConn
}
func (d *openAIWSSingleConnDialer) Dial(
ctx context.Context,
wsURL string,
headers http.Header,
proxyURL string,
) (openAIWSClientConn, int, http.Header, error) {
return d.conn, 0, nil, nil
}
func TestOpenAIWSDownstreamWriteContext_CancellationOwnership(t *testing.T) {
t.Run("pre-canceled ordinary context is canceled before return", func(t *testing.T) {
controlCtx, cancelControl := context.WithCancelCause(context.Background())
cancelControl(context.Canceled)
writeCtx, cancelWrite := newOpenAIWSDownstreamWriteContext(controlCtx, nil, time.Second)
defer cancelWrite()
require.ErrorIs(t, writeCtx.Err(), context.Canceled)
})
t.Run("lease loss keeps current write alive", func(t *testing.T) {
lifecycleCtx, cancelLifecycle := context.WithCancelCause(context.Background())
controlCtx, cancelControl := context.WithCancelCause(lifecycleCtx)
hooks := &OpenAIWSIngressHooks{ClientLifecycleContext: lifecycleCtx}
writeCtx, cancelWrite := newOpenAIWSDownstreamWriteContext(controlCtx, hooks, time.Second)
defer cancelWrite()
cancelControl(ErrOpenAIWSIngressLeaseLost)
select {
case <-writeCtx.Done():
t.Fatalf("lease loss unexpectedly canceled downstream write: %v", writeCtx.Err())
case <-time.After(20 * time.Millisecond):
}
clientDisconnected := errors.New("client disconnected")
cancelLifecycle(clientDisconnected)
<-writeCtx.Done()
require.ErrorIs(t, context.Cause(writeCtx), clientDisconnected)
})
t.Run("ordinary cancellation is direct and preserves cause", func(t *testing.T) {
lifecycleCtx, cancelLifecycle := context.WithCancelCause(context.Background())
controlCtx, cancelControl := context.WithCancelCause(lifecycleCtx)
defer cancelControl(context.Canceled)
hooks := &OpenAIWSIngressHooks{ClientLifecycleContext: lifecycleCtx}
writeCtx, cancelWrite := newOpenAIWSDownstreamWriteContext(controlCtx, hooks, time.Second)
defer cancelWrite()
serverShutdown := errors.New("server shutdown")
cancelLifecycle(serverShutdown)
require.ErrorIs(t, writeCtx.Err(), context.Canceled)
require.ErrorIs(t, context.Cause(writeCtx), serverShutdown)
})
}
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossTurns(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -169,6 +243,137 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossT
require.Len(t, captureConn.writes, 2, "应向同一上游连接发送两轮 response.create")
}
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_LeaseLossSendsRetryClose(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
lifecycleCtx, cancelLifecycle := context.WithCancelCause(context.Background())
defer cancelLifecycle(context.Canceled)
controlCtx, cancelControl := context.WithCancelCause(lifecycleCtx)
upstreamConn := &openAIWSLeaseLossAfterReadConn{
openAIWSCaptureConn: &openAIWSCaptureConn{events: [][]byte{
[]byte(`{"type":"response.completed","response":{"id":"resp_lease_loss","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
}},
cancel: cancelControl,
}
pool := newOpenAIWSConnPool(cfg)
pool.setClientDialerForTest(&openAIWSSingleConnDialer{conn: upstreamConn})
defer pool.Close()
svc := &OpenAIGatewayService{
cfg: cfg,
httpUpstream: &httpUpstreamRecorder{},
cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
toolCorrector: NewCodexToolCorrector(),
openaiWSPool: pool,
}
account := &Account{
ID: 118,
Name: "openai-ingress-lease-loss",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"responses_websockets_v2_enabled": true},
}
serverErrCh := make(chan error, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
serverErrCh <- err
return
}
defer func() {
_ = conn.CloseNow()
}()
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
req := r.Clone(controlCtx)
req.Header = req.Header.Clone()
req.Header.Set("User-Agent", "unit-test-agent/1.0")
ginCtx.Request = req
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
msgType, firstMessage, readErr := conn.Read(readCtx)
cancelRead()
if readErr != nil {
serverErrCh <- readErr
return
}
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
serverErrCh <- errors.New("unsupported websocket client message type")
return
}
serverErrCh <- svc.ProxyResponsesWebSocketFromClient(
controlCtx,
ginCtx,
conn,
account,
"sk-test",
firstMessage,
&OpenAIWSIngressHooks{ClientLifecycleContext: lifecycleCtx},
)
}))
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() {
_ = clientConn.CloseNow()
}()
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false}`))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
msgType, event, err := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, err)
require.Equal(t, coderws.MessageText, msgType)
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
closeReadCtx, cancelCloseRead := context.WithTimeout(context.Background(), 3*time.Second)
_, _, err = clientConn.Read(closeReadCtx)
cancelCloseRead()
var closeErr coderws.CloseError
require.ErrorAs(t, err, &closeErr)
require.Equal(t, coderws.StatusTryAgainLater, closeErr.Code)
require.Equal(t, "websocket ingress capacity lease lost; please reconnect", closeErr.Reason)
select {
case serverErr := <-serverErrCh:
var clientCloseErr *OpenAIWSClientCloseError
require.ErrorAs(t, serverErr, &clientCloseErr)
require.Equal(t, coderws.StatusTryAgainLater, clientCloseErr.StatusCode())
require.ErrorIs(t, serverErr, ErrOpenAIWSIngressLeaseLost)
case <-time.After(3 * time.Second):
t.Fatal("ingress lease-loss reader did not exit")
}
}
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_IdleTimeoutReleasesStoreDisabledSession(t *testing.T) {
gin.SetMode(gin.TestMode)