mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:48:43 +08:00
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:
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user