From 2a6b87aca083a5bf498ac1f68a1b636c500d7aaa Mon Sep 17 00:00:00 2001 From: Luis Pater Date: Thu, 3 Sep 2026 21:16:55 +0800 Subject: [PATCH] feat(openai): send periodic ping control frames during responses websocket streaming - Add `writePing` to responses websocket writer to emit Ping control frames. - Send periodic keep-alive Ping frames based on streaming configuration during response forwarding. - Reset keep-alive interval upon receiving data chunks and abort session if ping write fails. Closes: #5413 --- internal/config/sdk_config.go | 3 +- sdk/api/handlers/handlers.go | 2 +- .../openai/openai_responses_websocket.go | 12 + .../openai_responses_websocket_forward.go | 28 +- .../openai/openai_responses_websocket_test.go | 327 ++++++++++++++++++ 5 files changed, 368 insertions(+), 4 deletions(-) diff --git a/internal/config/sdk_config.go b/internal/config/sdk_config.go index 5c971d69..1c5a1bb6 100644 --- a/internal/config/sdk_config.go +++ b/internal/config/sdk_config.go @@ -74,7 +74,8 @@ type ClaudeCodeConfig struct { // StreamingConfig holds server streaming behavior configuration. type StreamingConfig struct { - // KeepAliveSeconds controls how often the server emits SSE heartbeats (": keep-alive\n\n"). + // KeepAliveSeconds controls how often the server emits SSE heartbeats (": keep-alive\n\n") + // or WebSocket Ping control frames. // <= 0 disables keep-alives. Default is 0. KeepAliveSeconds int `yaml:"keepalive-seconds,omitempty" json:"keepalive-seconds,omitempty"` diff --git a/sdk/api/handlers/handlers.go b/sdk/api/handlers/handlers.go index cf31c5e0..78170df3 100644 --- a/sdk/api/handlers/handlers.go +++ b/sdk/api/handlers/handlers.go @@ -108,7 +108,7 @@ func BuildErrorResponseBody(status int, errText string) []byte { return payload } -// StreamingKeepAliveInterval returns the SSE keep-alive interval for this server. +// StreamingKeepAliveInterval returns the streaming keep-alive interval for this server (SSE heartbeats and WebSocket Ping frames). // Returning 0 disables keep-alives (default when unset). func StreamingKeepAliveInterval(cfg *config.SDKConfig) time.Duration { seconds := defaultStreamingKeepAliveSeconds diff --git a/sdk/api/handlers/openai/openai_responses_websocket.go b/sdk/api/handlers/openai/openai_responses_websocket.go index 21cfa077..a8c2e368 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket.go +++ b/sdk/api/handlers/openai/openai_responses_websocket.go @@ -149,6 +149,18 @@ func (w *responsesWebsocketWriter) closeWithoutError() (bool, error) { return true, w.conn.Close() } +func (w *responsesWebsocketWriter) writePing() error { + if w == nil || w.conn == nil { + return errors.New("responses websocket: writer is nil") + } + w.writeMu.Lock() + defer w.writeMu.Unlock() + if w.closing.Load() { + return websocket.ErrCloseSent + } + return w.conn.WriteControl(websocket.PingMessage, nil, time.Time{}) +} + func (w *responsesWebsocketWriter) closeWithPayload(payload []byte) (bool, error) { if w == nil || w.conn == nil { return false, nil diff --git a/sdk/api/handlers/openai/openai_responses_websocket_forward.go b/sdk/api/handlers/openai/openai_responses_websocket_forward.go index 49603edb..cba76837 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket_forward.go +++ b/sdk/api/handlers/openai/openai_responses_websocket_forward.go @@ -21,8 +21,9 @@ import ( ) type responsesWebsocketForwardOptions struct { - toolCacheTurn *responsesWebsocketToolCacheTurn - suppressError func(*interfaces.ErrorMessage) bool + toolCacheTurn *responsesWebsocketToolCacheTurn + suppressError func(*interfaces.ErrorMessage) bool + keepAliveInterval *time.Duration } func (h *OpenAIResponsesAPIHandler) forwardResponsesWebsocket( @@ -51,11 +52,31 @@ func (h *OpenAIResponsesAPIHandler) forwardResponsesWebsocket( downstreamSessionKey = websocketDownstreamSessionKey(c.Request) } + var keepAliveTicker *time.Ticker + var keepAliveC <-chan time.Time + keepAliveInterval := time.Duration(0) + if h != nil { + keepAliveInterval = handlers.StreamingKeepAliveInterval(h.Cfg) + } + if opts.keepAliveInterval != nil { + keepAliveInterval = *opts.keepAliveInterval + } + if keepAliveInterval > 0 { + keepAliveTicker = time.NewTicker(keepAliveInterval) + defer keepAliveTicker.Stop() + keepAliveC = keepAliveTicker.C + } + for { select { case <-c.Request.Context().Done(): cancel(c.Request.Context().Err()) return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), nil, c.Request.Context().Err() + case <-keepAliveC: + if errPing := writer.writePing(); errPing != nil { + cancel(errPing) + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), nil, errPing + } case errMsg, ok := <-errs: if !ok { errs = nil @@ -111,6 +132,9 @@ func (h *OpenAIResponsesAPIHandler) forwardResponsesWebsocket( cancel(nil) return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), nil, nil } + if keepAliveTicker != nil && keepAliveInterval > 0 { + keepAliveTicker.Reset(keepAliveInterval) + } payloads := websocketJSONPayloadsFromChunk(chunk) for i := range payloads { diff --git a/sdk/api/handlers/openai/openai_responses_websocket_test.go b/sdk/api/handlers/openai/openai_responses_websocket_test.go index 2079dcf7..9499b69e 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket_test.go +++ b/sdk/api/handlers/openai/openai_responses_websocket_test.go @@ -5801,3 +5801,330 @@ func TestNormalizeSubsequentRequestAssistantInputTriggersTranscriptReplacement(t t.Fatalf("input[0].id = %q, want %q", input[0].Get("id").String(), "msg-3") } } + +func TestForwardResponsesWebsocketEmitsPeriodicPingControlFrames(t *testing.T) { + gin.SetMode(gin.TestMode) + + serverErrCh := make(chan error, 1) + data := make(chan []byte) + errCh := make(chan *interfaces.ErrorMessage) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := responsesWebsocketUpgrader.Upgrade(w, r, nil) + if err != nil { + serverErrCh <- err + return + } + defer func() { _ = conn.Close() }() + + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Request = r + + cfg := &sdkconfig.SDKConfig{ + Streaming: sdkconfig.StreamingConfig{ + KeepAliveSeconds: 1, + }, + } + h := NewOpenAIResponsesAPIHandler(handlers.NewBaseAPIHandlers(cfg, nil)) + + _, _, _, errMsg, errForward := h.forwardResponsesWebsocket( + ctx, + newResponsesWebsocketWriter(conn), + func(...interface{}) {}, + data, + errCh, + newInMemoryWebsocketTimelineLog(), + "session-keepalive-test", + ) + if errMsg != nil { + serverErrCh <- fmt.Errorf("unexpected error message: %v", errMsg.Error) + return + } + if errForward != nil { + serverErrCh <- fmt.Errorf("unexpected forward error: %v", errForward) + return + } + serverErrCh <- nil + })) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + clientConn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = clientConn.Close() }() + + pingReceived := make(chan struct{}, 1) + clientConn.SetPingHandler(func(appData string) error { + select { + case pingReceived <- struct{}{}: + default: + } + return clientConn.WriteControl(websocket.PongMessage, []byte(appData), time.Now().Add(time.Second)) + }) + + clientDone := make(chan struct{}) + go func() { + defer close(clientDone) + for { + _, _, errRead := clientConn.ReadMessage() + if errRead != nil { + return + } + } + }() + + select { + case <-pingReceived: + // Received expected Ping control frame while upstream is waiting. + case <-time.After(1500 * time.Millisecond): + t.Fatal("expected websocket Ping control frame during upstream wait, got none") + } + + // Unblock forwardResponsesWebsocket with terminal completion. + data <- []byte(`{"type":"response.done","response":{"id":"resp-ping-1","output":[]}}`) + close(data) + close(errCh) + + select { + case serverErr := <-serverErrCh: + if serverErr != nil { + t.Fatalf("server error: %v", serverErr) + } + case <-time.After(2 * time.Second): + t.Fatal("server timed out completing forwardResponsesWebsocket") + } + + <-clientDone +} + +func TestForwardResponsesWebsocketPingKeepAliveOptionsOverride(t *testing.T) { + gin.SetMode(gin.TestMode) + + serverErrCh := make(chan error, 1) + data := make(chan []byte) + errCh := make(chan *interfaces.ErrorMessage) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := responsesWebsocketUpgrader.Upgrade(w, r, nil) + if err != nil { + serverErrCh <- err + return + } + defer func() { _ = conn.Close() }() + + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Request = r + + interval := 20 * time.Millisecond + h := NewOpenAIResponsesAPIHandler(handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, nil)) + + _, _, _, errMsg, errForward := h.forwardResponsesWebsocket( + ctx, + newResponsesWebsocketWriter(conn), + func(...interface{}) {}, + data, + errCh, + newInMemoryWebsocketTimelineLog(), + "session-keepalive-override", + responsesWebsocketForwardOptions{ + keepAliveInterval: &interval, + }, + ) + if errMsg != nil { + serverErrCh <- fmt.Errorf("unexpected error message: %v", errMsg.Error) + return + } + if errForward != nil { + serverErrCh <- fmt.Errorf("unexpected forward error: %v", errForward) + return + } + serverErrCh <- nil + })) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + clientConn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = clientConn.Close() }() + + pingReceived := make(chan struct{}, 1) + clientConn.SetPingHandler(func(appData string) error { + select { + case pingReceived <- struct{}{}: + default: + } + return clientConn.WriteControl(websocket.PongMessage, []byte(appData), time.Now().Add(time.Second)) + }) + + clientDone := make(chan struct{}) + go func() { + defer close(clientDone) + for { + _, _, errRead := clientConn.ReadMessage() + if errRead != nil { + return + } + } + }() + + select { + case <-pingReceived: + case <-time.After(500 * time.Millisecond): + t.Fatal("expected websocket Ping control frame via option override, got none") + } + + data <- []byte(`{"type":"response.done","response":{"id":"resp-override","output":[]}}`) + close(data) + close(errCh) + + select { + case serverErr := <-serverErrCh: + if serverErr != nil { + t.Fatalf("server error: %v", serverErr) + } + case <-time.After(2 * time.Second): + t.Fatal("server timed out completing forwardResponsesWebsocket") + } + + <-clientDone +} + +func TestResponsesWebsocketWriterWritePing(t *testing.T) { + // Nil writer check + var nilWriter *responsesWebsocketWriter + if err := nilWriter.writePing(); err == nil { + t.Fatal("expected error on nil writer.writePing(), got nil") + } + + // Active connection and closed writer checks + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := responsesWebsocketUpgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer func() { _ = conn.Close() }() + + writer := newResponsesWebsocketWriter(conn) + if errPing := writer.writePing(); errPing != nil { + t.Errorf("writePing() error = %v, want nil", errPing) + } + + writer.closing.Store(true) + if errPing := writer.writePing(); !errors.Is(errPing, websocket.ErrCloseSent) { + t.Errorf("writePing() after closing = %v, want ErrCloseSent", errPing) + } + })) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + clientConn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = clientConn.Close() }() + + pingReceived := make(chan struct{}, 1) + clientConn.SetPingHandler(func(string) error { + select { + case pingReceived <- struct{}{}: + default: + } + return nil + }) + + go func() { + for { + if _, _, errRead := clientConn.ReadMessage(); errRead != nil { + return + } + } + }() + + select { + case <-pingReceived: + case <-time.After(500 * time.Millisecond): + t.Fatal("expected ping from writer.writePing(), got none") + } +} + +func TestForwardResponsesWebsocketPingWriteFailureAbortsSession(t *testing.T) { + gin.SetMode(gin.TestMode) + + serverErrCh := make(chan error, 1) + cancelledCh := make(chan error, 1) + data := make(chan []byte) + errCh := make(chan *interfaces.ErrorMessage) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := responsesWebsocketUpgrader.Upgrade(w, r, nil) + if err != nil { + serverErrCh <- err + return + } + defer func() { _ = conn.Close() }() + + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Request = r + + interval := 10 * time.Millisecond + h := NewOpenAIResponsesAPIHandler(handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, nil)) + + _, _, _, errMsg, errForward := h.forwardResponsesWebsocket( + ctx, + newResponsesWebsocketWriter(conn), + func(errs ...interface{}) { + if len(errs) > 0 { + if errVal, ok := errs[0].(error); ok { + cancelledCh <- errVal + return + } + } + cancelledCh <- nil + }, + data, + errCh, + newInMemoryWebsocketTimelineLog(), + "session-ping-fail", + responsesWebsocketForwardOptions{ + keepAliveInterval: &interval, + }, + ) + if errMsg != nil { + serverErrCh <- fmt.Errorf("unexpected error message: %v", errMsg.Error) + return + } + serverErrCh <- errForward + })) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + clientConn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + + // Close client connection immediately so the next server Ping write fails. + _ = clientConn.Close() + + select { + case serverErr := <-serverErrCh: + if serverErr == nil { + t.Fatal("expected error on server ping write failure, got nil") + } + case <-time.After(2 * time.Second): + t.Fatal("server timed out awaiting ping write abort") + } + + select { + case cancelErr := <-cancelledCh: + if cancelErr == nil { + t.Fatal("expected cancel callback to be invoked with ping error, got nil") + } + case <-time.After(time.Second): + t.Fatal("timed out awaiting cancel callback on ping write failure") + } +} -- 2.51.2