fix(openai-ws): close ingress reads cleanly

This commit is contained in:
王鹏
2026-07-15 22:15:44 +08:00
parent eb2b8632de
commit 4f641208a0
7 changed files with 322 additions and 47 deletions
@@ -1360,6 +1360,28 @@ 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()
maxIngressConnections := 0
if h.cfg != nil {
maxIngressConnections = h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey
}
ingressLease, ingressLeaseAcquired, ingressLeaseErr := h.concurrencyHelper.AcquireOpenAIWSIngressLease(ctx, apiKey.ID, maxIngressConnections)
if ingressLeaseErr != nil {
reqLog.Error("openai.websocket_ingress_lease_acquire_failed", zap.Error(ingressLeaseErr))
h.errorResponse(c, http.StatusServiceUnavailable, "service_unavailable", "WebSocket ingress capacity is temporarily unavailable")
return
}
if !ingressLeaseAcquired {
reqLog.Info("openai.websocket_ingress_capacity_rejected", zap.Int("max_ingress_connections_per_api_key", maxIngressConnections))
c.Header("Retry-After", "5")
h.errorResponse(c, http.StatusTooManyRequests, "rate_limit_error", "Too many open WebSocket connections, please retry later")
return
}
if ingressLease != nil {
defer ingressLease.Release()
ctx = ingressLease.Context()
c.Request = c.Request.WithContext(ctx)
}
wsConn, err := coderws.Accept(c.Writer, c.Request, &coderws.AcceptOptions{
CompressionMode: coderws.CompressionContextTakeover,
@@ -1381,32 +1403,14 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
}()
wsConn.SetReadLimit(service.ResolveOpenAIWSClientReadLimitBytes(h.cfg))
ctx := c.Request.Context()
maxIngressConnections := 0
if h.cfg != nil {
maxIngressConnections = h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey
}
ingressLease, ingressLeaseAcquired, ingressLeaseErr := h.concurrencyHelper.AcquireOpenAIWSIngressLease(ctx, apiKey.ID, maxIngressConnections)
if ingressLeaseErr != nil {
reqLog.Error("openai.websocket_ingress_lease_acquire_failed", zap.Error(ingressLeaseErr))
closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to reserve websocket ingress capacity")
return
}
if !ingressLeaseAcquired {
reqLog.Info("openai.websocket_ingress_capacity_rejected", zap.Int("max_ingress_connections_per_api_key", maxIngressConnections))
closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "too many open websocket connections, please retry later")
return
}
if ingressLease != nil {
defer ingressLease.Release()
ctx = ingressLease.Context()
c.Request = c.Request.WithContext(ctx)
}
firstMessageTimeout := service.ResolveOpenAIWSClientFirstMessageTimeout(h.cfg)
readCtx, cancel := context.WithTimeout(ctx, firstMessageTimeout)
msgType, firstMessage, err := wsConn.Read(readCtx)
cancel()
msgType, firstMessage, err := service.ReadOpenAIWSClientMessage(
ctx,
wsConn,
firstMessageTimeout,
coderws.StatusPolicyViolation,
"missing first response.create message",
)
if err != nil {
if errors.Is(context.Cause(ctx), service.ErrOpenAIWSIngressLeaseLost) {
reqLog.Warn("openai.websocket_ingress_lease_lost_before_first_message", zap.Error(err))
@@ -820,17 +820,36 @@ func TestOpenAIResponsesWebSocket_IngressCapacityRejected(t *testing.T) {
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http")+"/openai/v1/responses", nil)
clientConn, response, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http")+"/openai/v1/responses", nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
require.Error(t, err)
require.Nil(t, clientConn)
require.NotNil(t, response)
require.Equal(t, http.StatusTooManyRequests, response.StatusCode)
_ = response.Body.Close()
}
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
_, _, err = clientConn.Read(readCtx)
cancelRead()
var closeErr coderws.CloseError
require.ErrorAs(t, err, &closeErr)
require.Equal(t, coderws.StatusTryAgainLater, closeErr.Code)
func TestOpenAIResponsesWebSocket_IngressLeaseBackendUnavailableBeforeUpgrade(t *testing.T) {
gin.SetMode(gin.TestMode)
cache := &concurrencyCacheMock{
acquireIngressLeaseFn: func(context.Context, int64, int, string) (bool, error) {
return false, errors.New("redis unavailable")
},
}
h := newOpenAIHandlerForPreviousResponseIDValidation(t, cache)
h.cfg = &config.Config{}
h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = 1
wsServer := newOpenAIWSHandlerTestServer(t, h, middleware.AuthSubject{UserID: 1, Concurrency: 1})
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, response, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http")+"/openai/v1/responses", nil)
cancelDial()
require.Error(t, err)
require.Nil(t, clientConn)
require.NotNil(t, response)
require.Equal(t, http.StatusServiceUnavailable, response.StatusCode)
_ = response.Body.Close()
}
func TestOpenAIResponsesWebSocket_FirstMessageTimeoutUsesConfig(t *testing.T) {
@@ -902,6 +921,33 @@ func TestOpenAIResponsesWebSocket_IngressLeaseReleasedOnEarlyReturn(t *testing.T
}, time.Second, 10*time.Millisecond)
}
func TestOpenAIResponsesWebSocket_IngressLeaseReleasedWhenUpgradeFails(t *testing.T) {
gin.SetMode(gin.TestMode)
cache := &concurrencyCacheMock{
acquireIngressLeaseFn: func(context.Context, int64, int, string) (bool, error) {
return true, nil
},
}
h := newOpenAIHandlerForPreviousResponseIDValidation(t, cache)
h.cfg = &config.Config{}
h.cfg.Gateway.OpenAIWS.MaxIngressConnectionsPerAPIKey = 1
wsServer := newOpenAIWSHandlerTestServer(t, h, middleware.AuthSubject{UserID: 1, Concurrency: 1})
defer wsServer.Close()
req, err := http.NewRequest(http.MethodGet, wsServer.URL+"/openai/v1/responses", nil)
require.NoError(t, err)
req.Header.Set("Upgrade", "websocket")
req.Header.Set("Connection", "Upgrade")
req.Header.Set("Sec-WebSocket-Version", "13")
resp, err := http.DefaultClient.Do(req)
require.NoError(t, err)
_ = resp.Body.Close()
require.NotEqual(t, http.StatusSwitchingProtocols, resp.StatusCode)
require.Eventually(t, func() bool {
return atomic.LoadInt32(&cache.releaseIngressCalled) == 1
}, time.Second, 10*time.Millisecond)
}
func TestOpenAIResponsesWebSocket_RejectsMessageIDAsPreviousResponseID(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -0,0 +1,70 @@
package service
import (
"context"
"errors"
"time"
coderws "github.com/coder/websocket"
)
type openAIWSClientReadResult struct {
messageType coderws.MessageType
payload []byte
err error
}
// ReadOpenAIWSClientMessage keeps one reader alive while control events send
// their close frame, then closes the transport and joins that reader.
func ReadOpenAIWSClientMessage(
controlCtx context.Context,
conn *coderws.Conn,
timeout time.Duration,
timeoutStatus coderws.StatusCode,
timeoutReason string,
) (coderws.MessageType, []byte, error) {
if conn == nil {
return 0, nil, errors.New("openai websocket client connection is nil")
}
if controlCtx == nil {
controlCtx = context.Background()
}
readDone := make(chan openAIWSClientReadResult, 1)
go func() {
messageType, payload, err := conn.Read(context.Background())
readDone <- openAIWSClientReadResult{messageType: messageType, payload: payload, err: err}
}()
var timeoutCh <-chan time.Time
var timer *time.Timer
if timeout > 0 {
timer = time.NewTimer(timeout)
timeoutCh = timer.C
defer timer.Stop()
}
closeAndJoin := func(status coderws.StatusCode, reason string, cause error) (coderws.MessageType, []byte, error) {
_ = conn.Close(status, reason)
_ = conn.CloseNow()
<-readDone
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,
)
}
return closeAndJoin(coderws.StatusGoingAway, "websocket request canceled", cause)
}
}
@@ -0,0 +1,143 @@
package service
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
coderws "github.com/coder/websocket"
"github.com/stretchr/testify/require"
)
func TestReadOpenAIWSClientMessage_ControlCloseFrames(t *testing.T) {
tests := []struct {
name string
timeout time.Duration
timeoutStatus coderws.StatusCode
timeoutReason string
cancelCause error
wantStatus coderws.StatusCode
wantReason string
}{
{
name: "inter-turn idle sends normal close",
timeout: 25 * time.Millisecond,
timeoutStatus: coderws.StatusNormalClosure,
timeoutReason: "websocket idle timeout",
wantStatus: coderws.StatusNormalClosure,
wantReason: "websocket idle timeout",
},
{
name: "first message timeout sends policy close",
timeout: 25 * time.Millisecond,
timeoutStatus: coderws.StatusPolicyViolation,
timeoutReason: "missing first response.create message",
wantStatus: coderws.StatusPolicyViolation,
wantReason: "missing first response.create message",
},
{
name: "lease loss sends retry close",
cancelCause: ErrOpenAIWSIngressLeaseLost,
wantStatus: coderws.StatusTryAgainLater,
wantReason: "websocket ingress capacity lease lost; please reconnect",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
controlCtx, cancelControl := context.WithCancelCause(context.Background())
defer cancelControl(context.Canceled)
serverResult := make(chan error, 1)
readStarted := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, nil)
if err != nil {
serverResult <- err
return
}
defer func() { _ = conn.CloseNow() }()
close(readStarted)
_, _, err = ReadOpenAIWSClientMessage(
controlCtx,
conn,
tt.timeout,
tt.timeoutStatus,
tt.timeoutReason,
)
serverResult <- err
}))
defer server.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(server.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
<-readStarted
if tt.cancelCause != nil {
cancelControl(tt.cancelCause)
}
readCtx, cancelRead := context.WithTimeout(context.Background(), time.Second)
_, _, err = clientConn.Read(readCtx)
cancelRead()
var clientClose coderws.CloseError
require.ErrorAs(t, err, &clientClose)
require.Equal(t, tt.wantStatus, clientClose.Code)
require.Equal(t, tt.wantReason, clientClose.Reason)
select {
case serverErr := <-serverResult:
var closeErr *OpenAIWSClientCloseError
require.ErrorAs(t, serverErr, &closeErr)
require.Equal(t, tt.wantStatus, closeErr.StatusCode())
require.Equal(t, tt.wantReason, closeErr.Reason())
case <-time.After(time.Second):
t.Fatal("server read goroutine did not exit after close handshake")
}
})
}
}
func TestReadOpenAIWSClientMessage_ParentCancellationStillJoinsRead(t *testing.T) {
controlCtx, cancelControl := context.WithCancelCause(context.Background())
serverResult := make(chan error, 1)
readStarted := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, nil)
if err != nil {
serverResult <- err
return
}
defer func() { _ = conn.CloseNow() }()
close(readStarted)
_, _, err = ReadOpenAIWSClientMessage(controlCtx, conn, 0, 0, "")
serverResult <- err
}))
defer server.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(server.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
<-readStarted
cancelControl(errors.New("server shutting down"))
readCtx, cancelRead := context.WithTimeout(context.Background(), time.Second)
_, _, err = clientConn.Read(readCtx)
cancelRead()
var clientClose coderws.CloseError
require.ErrorAs(t, err, &clientClose)
require.Equal(t, coderws.StatusGoingAway, clientClose.Code)
require.Equal(t, "websocket request canceled", clientClose.Reason)
select {
case <-serverResult:
case <-time.After(time.Second):
t.Fatal("server read goroutine leaked after parent cancellation")
}
}
@@ -382,22 +382,18 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
readClientMessage := func() ([]byte, error) {
readCtx := ctx
idleTimeout := s.openAIWSIngressInterTurnIdleTimeout()
cancelRead := func() {}
if idleTimeout > 0 {
readCtx, cancelRead = context.WithTimeout(ctx, idleTimeout)
}
msgType, payload, readErr := clientConn.Read(readCtx)
cancelRead()
msgType, payload, readErr := ReadOpenAIWSClientMessage(
ctx,
clientConn,
idleTimeout,
coderws.StatusNormalClosure,
"websocket idle timeout",
)
if readErr != nil {
if idleTimeout > 0 && errors.Is(readErr, context.DeadlineExceeded) && ctx.Err() == nil {
var closeErr *OpenAIWSClientCloseError
if errors.As(readErr, &closeErr) && closeErr.StatusCode() == coderws.StatusNormalClosure {
logOpenAIWSModeInfo("ingress_ws_inter_turn_idle_timeout account_id=%d timeout_seconds=%d", account.ID, int(idleTimeout.Seconds()))
return nil, NewOpenAIWSClientCloseError(
coderws.StatusNormalClosure,
"websocket idle timeout",
readErr,
)
}
return nil, readErr
}
@@ -249,6 +249,14 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_IdleTimeoutRelea
require.NoError(t, err)
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 clientClose coderws.CloseError
require.ErrorAs(t, err, &clientClose)
require.Equal(t, coderws.StatusNormalClosure, clientClose.Code)
require.Equal(t, "websocket idle timeout", clientClose.Reason)
select {
case proxyErr := <-serverErrCh:
var closeErr *OpenAIWSClientCloseError
@@ -775,6 +775,14 @@ func TestOpenAIWSHTTPBridge_IdleTimeoutClosesClientSession(t *testing.T) {
require.NoError(t, err)
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 clientClose coderws.CloseError
require.ErrorAs(t, err, &clientClose)
require.Equal(t, coderws.StatusNormalClosure, clientClose.Code)
require.Equal(t, "websocket idle timeout", clientClose.Reason)
select {
case proxyErr := <-errCh:
var closeErr *OpenAIWSClientCloseError