Skip to content

Commit bbec97d

Browse files
fix: complete WebSocket shutdown handshake (vllm-project#141)
## Summary - Complete the WebSocket close handshake after draining an active response during shutdown. - Reject post-shutdown WebSocket requests before they can start inference. - Make the shutdown regression deterministic with a ping/pong receipt barrier and strict close acknowledgment. This fixes the Linux CI race observed at caf6bf1 where unread client frames could make socket teardown surface ECONNRESET after the active response completed. ## Test Plan - `cargo fmt -- --check` - `cargo clippy --all-targets -- -D warnings` - `cargo test` - `uvx pre-commit run --all-files` - 100 repeated runs of `test_websocket_shutdown_drains_active_response_before_closing` - gstack `/review` with specialist, adversarial, and structured re-review gates - Claude read-only worktree review --------- Signed-off-by: Francisco Javier Arceo <farceo@redhat.com>
1 parent 113b870 commit bbec97d

5 files changed

Lines changed: 481 additions & 29 deletions

File tree

.gitignore

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,9 @@ build/
1414
# Rust
1515
target/
1616

17+
# Local worktrees
18+
.worktrees/
19+
1720
# IDE
1821
.idea/
1922
.vscode/

crates/agentic-server/src/handler/websocket/responses.rs

Lines changed: 163 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@ use axum::http::HeaderMap;
77
use axum::response::Response;
88
use either::Either;
99
use futures::stream::{SplitSink, SplitStream};
10-
use futures::{SinkExt, Stream, StreamExt};
10+
use futures::{Sink, SinkExt, Stream, StreamExt};
1111
use serde_json::Value;
1212
use tokio_util::sync::CancellationToken;
1313
use tracing::{debug, warn};
@@ -49,10 +49,7 @@ async fn responses_ws_loop(socket: WebSocket, state: AppState, headers: HeaderMa
4949
let text = if let Some(buffered) = queue.pop_front() {
5050
buffered
5151
} else {
52-
let message = tokio::select! {
53-
() = shutdown_token.cancelled() => break,
54-
message = receiver.next() => message,
55-
};
52+
let message = next_ws_message(&shutdown_token, &mut receiver).await;
5653

5754
let Some(message) = message else {
5855
break;
@@ -100,10 +97,56 @@ async fn responses_ws_loop(socket: WebSocket, state: AppState, headers: HeaderMa
10097
}
10198
}
10299
}
100+
close_ws(&mut sender, &mut receiver).await;
101+
debug!("responses websocket session closed");
102+
}
103+
104+
async fn next_ws_message<Receiver>(
105+
shutdown_token: &CancellationToken,
106+
receiver: &mut Receiver,
107+
) -> Option<Receiver::Item>
108+
where
109+
Receiver: Stream + Unpin,
110+
{
111+
tokio::select! {
112+
biased;
113+
() = shutdown_token.cancelled() => None,
114+
message = receiver.next() => {
115+
if shutdown_token.is_cancelled() {
116+
None
117+
} else {
118+
message
119+
}
120+
},
121+
}
122+
}
123+
124+
fn keep_if_running<T>(shutdown_token: &CancellationToken, value: T) -> Option<T> {
125+
(!shutdown_token.is_cancelled()).then_some(value)
126+
}
127+
128+
async fn close_ws<Sender, Receiver, SendError, ReceiveError>(sender: &mut Sender, receiver: &mut Receiver)
129+
where
130+
Sender: Sink<Message, Error = SendError> + Unpin,
131+
Receiver: Stream<Item = Result<Message, ReceiveError>> + Unpin,
132+
SendError: std::fmt::Display,
133+
ReceiveError: std::fmt::Display,
134+
{
103135
if let Err(error) = sender.close().await {
104-
debug!(%error, "failed to close responses websocket cleanly");
136+
debug!(%error, "failed to send responses websocket close frame");
137+
return;
138+
}
139+
140+
while let Some(message) = receiver.next().await {
141+
match message {
142+
Ok(Message::Close(_)) => break,
143+
Ok(Message::Text(_) | Message::Binary(_) | Message::Ping(_) | Message::Pong(_)) => {}
144+
Err(error) => {
145+
debug!(%error, "responses websocket close handshake receive failed");
146+
break;
147+
}
148+
}
105149
}
106-
debug!("responses websocket session closed");
107150
}
108151

109152
/// Process one `response.create` message.
@@ -153,6 +196,10 @@ async fn handle_ws_text(
153196
.with_auth(auth)
154197
.run()
155198
.await?;
199+
let Some(result) = keep_if_running(shutdown_token, result) else {
200+
debug!("discarded websocket response initialized during shutdown");
201+
return Ok(());
202+
};
156203
let Either::Right(stream) = result else {
157204
return Err(WsError::Executor(ExecutorError::InvalidRequest(
158205
"websocket response.create must produce a stream".to_owned(),
@@ -361,9 +408,116 @@ async fn send_ws_json(sender: &mut WsSender, value: Value) -> Result<(), WsError
361408

362409
#[cfg(test)]
363410
mod tests {
364-
use futures::stream;
411+
use std::pin::Pin;
412+
use std::task::{Context, Poll};
413+
414+
use axum::extract::ws::Message;
415+
use futures::{Sink, Stream, StreamExt, sink, stream};
416+
use tokio_util::sync::CancellationToken;
417+
418+
use super::{ShutdownInput, close_ws, keep_if_running, next_shutdown_input, next_ws_message};
419+
420+
struct CloseErrorSink;
421+
422+
struct CancellingStream {
423+
shutdown_token: CancellationToken,
424+
item: Option<&'static str>,
425+
}
426+
427+
impl Stream for CancellingStream {
428+
type Item = &'static str;
429+
430+
fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
431+
self.shutdown_token.cancel();
432+
Poll::Ready(self.item.take())
433+
}
434+
}
365435

366-
use super::{ShutdownInput, next_shutdown_input};
436+
impl Sink<Message> for CloseErrorSink {
437+
type Error = &'static str;
438+
439+
fn poll_ready(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
440+
Poll::Ready(Ok(()))
441+
}
442+
443+
fn start_send(self: Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
444+
Ok(())
445+
}
446+
447+
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
448+
Poll::Ready(Ok(()))
449+
}
450+
451+
fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
452+
Poll::Ready(Err("close failed"))
453+
}
454+
}
455+
456+
#[tokio::test]
457+
async fn cancelled_shutdown_wins_over_ready_websocket_message() {
458+
let shutdown_token = CancellationToken::new();
459+
shutdown_token.cancel();
460+
let mut receiver = stream::iter(["must remain unread"]);
461+
462+
assert!(next_ws_message(&shutdown_token, &mut receiver).await.is_none());
463+
assert_eq!(receiver.next().await, Some("must remain unread"));
464+
}
465+
466+
#[tokio::test]
467+
async fn cancellation_during_receive_discards_websocket_message() {
468+
let shutdown_token = CancellationToken::new();
469+
let mut receiver = CancellingStream {
470+
shutdown_token: shutdown_token.clone(),
471+
item: Some("must be discarded"),
472+
};
473+
474+
assert!(next_ws_message(&shutdown_token, &mut receiver).await.is_none());
475+
assert!(shutdown_token.is_cancelled());
476+
assert_eq!(receiver.next().await, None);
477+
}
478+
479+
#[test]
480+
fn cancellation_after_request_setup_discards_unpolled_stream() {
481+
let shutdown_token = CancellationToken::new();
482+
shutdown_token.cancel();
483+
484+
assert_eq!(keep_if_running(&shutdown_token, "unpolled stream"), None);
485+
}
486+
487+
#[tokio::test]
488+
async fn close_ws_ignores_late_frames_until_peer_close() {
489+
let mut sender = sink::drain();
490+
let mut receiver = stream::iter([
491+
Ok::<_, &'static str>(Message::Text("late request".into())),
492+
Ok(Message::Binary(vec![1].into())),
493+
Ok(Message::Close(None)),
494+
Err("must remain unread"),
495+
]);
496+
497+
close_ws(&mut sender, &mut receiver).await;
498+
499+
assert!(matches!(receiver.next().await, Some(Err("must remain unread"))));
500+
}
501+
502+
#[tokio::test]
503+
async fn close_ws_returns_without_reading_when_close_send_fails() {
504+
let mut sender = CloseErrorSink;
505+
let mut receiver = stream::iter([Ok::<_, &'static str>(Message::Close(None))]);
506+
507+
close_ws(&mut sender, &mut receiver).await;
508+
509+
assert!(matches!(receiver.next().await, Some(Ok(Message::Close(None)))));
510+
}
511+
512+
#[tokio::test]
513+
async fn close_ws_stops_reading_after_receive_error() {
514+
let mut sender = sink::drain();
515+
let mut receiver = stream::iter([Err::<Message, _>("receive failed"), Ok(Message::Close(None))]);
516+
517+
close_ws(&mut sender, &mut receiver).await;
518+
519+
assert!(matches!(receiver.next().await, Some(Ok(Message::Close(None)))));
520+
}
367521

368522
#[tokio::test]
369523
async fn shutdown_input_priority_alternates_when_both_streams_are_ready() {

crates/agentic-server/tests/responses_websocket_test.rs

Lines changed: 41 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -344,6 +344,39 @@ async fn recv_close_or_end(ws: &mut WebSocketStream<MaybeTlsStream<TcpStream>>)
344344
}
345345
}
346346

347+
async fn recv_clean_close(ws: &mut WebSocketStream<MaybeTlsStream<TcpStream>>) {
348+
let message = tokio::time::timeout(std::time::Duration::from_secs(2), ws.next())
349+
.await
350+
.expect("timed out waiting for clean websocket close");
351+
match message {
352+
Some(Ok(Message::Close(_))) => ws.flush().await.expect("failed to acknowledge websocket close"),
353+
None => panic!("websocket ended without a close frame"),
354+
Some(Err(error)) => panic!("websocket close failed: {error}"),
355+
Some(Ok(message)) => panic!("expected websocket close, got {message:?}"),
356+
}
357+
}
358+
359+
async fn send_ping_and_wait_for_pong(ws: &mut WebSocketStream<MaybeTlsStream<TcpStream>>, payload: Bytes) {
360+
ws.send(Message::Ping(payload.clone())).await.unwrap();
361+
loop {
362+
let message = tokio::time::timeout(std::time::Duration::from_secs(2), ws.next())
363+
.await
364+
.expect("timed out waiting for websocket pong")
365+
.expect("websocket should yield a message")
366+
.expect("websocket message should be ok");
367+
match message {
368+
Message::Pong(actual) => {
369+
assert_eq!(actual, payload);
370+
break;
371+
}
372+
Message::Ping(_) | Message::Frame(_) => {}
373+
Message::Text(text) => panic!("unexpected text before pong: {text}"),
374+
Message::Close(frame) => panic!("websocket closed before pong: {frame:?}"),
375+
Message::Binary(_) => panic!("unexpected binary websocket message"),
376+
}
377+
}
378+
}
379+
347380
async fn send_json(ws: &mut WebSocketStream<MaybeTlsStream<TcpStream>>, value: Value) {
348381
ws.send(Message::Text(value.to_string().into())).await.unwrap();
349382
}
@@ -1231,25 +1264,7 @@ async fn test_websocket_ping_returns_pong_without_upstream_request() {
12311264
let (gateway_url, _gateway) = spawn_gateway(fixture.state.clone()).await;
12321265
let mut ws = connect_responses_ws(&gateway_url).await;
12331266

1234-
ws.send(Message::Ping(Bytes::from_static(b"ping"))).await.unwrap();
1235-
1236-
loop {
1237-
let message = tokio::time::timeout(std::time::Duration::from_secs(2), ws.next())
1238-
.await
1239-
.expect("timed out waiting for websocket pong")
1240-
.expect("websocket should yield a message")
1241-
.expect("websocket message should be ok");
1242-
match message {
1243-
Message::Pong(payload) => {
1244-
assert_eq!(payload, Bytes::from_static(b"ping"));
1245-
break;
1246-
}
1247-
Message::Ping(_) | Message::Frame(_) => {}
1248-
Message::Text(text) => panic!("unexpected text websocket message: {text}"),
1249-
Message::Close(frame) => panic!("websocket closed before pong: {frame:?}"),
1250-
Message::Binary(_) => panic!("unexpected binary websocket message"),
1251-
}
1252-
}
1267+
send_ping_and_wait_for_pong(&mut ws, Bytes::from_static(b"ping")).await;
12531268

12541269
assert!(mock.request_bodies().await.is_empty());
12551270
}
@@ -1274,6 +1289,7 @@ async fn test_websocket_shutdown_drains_active_response_before_closing() {
12741289
MockResponsesServer::start_gated(sse_response("resp_upstream_shutdown", "msg_upstream_shutdown", "DONE")).await;
12751290
let fixture = storage_backed_state(&mock.url).await;
12761291
let shutdown_token = fixture.state.shutdown_token.clone();
1292+
let websocket_tracker = fixture.state.websocket_tracker.clone();
12771293
let (gateway_url, _gateway) = spawn_gateway(fixture.state.clone()).await;
12781294
let mut ws = connect_responses_ws(&gateway_url).await;
12791295

@@ -1302,11 +1318,16 @@ async fn test_websocket_shutdown_drains_active_response_before_closing() {
13021318
}),
13031319
)
13041320
.await;
1321+
let barrier = Bytes::from_static(b"shutdown-request-received");
1322+
send_ping_and_wait_for_pong(&mut ws, barrier).await;
13051323
release.send(()).unwrap();
13061324

13071325
let events = recv_until_completed(&mut ws).await;
13081326
assert_eq!(events.last().unwrap()["type"], "response.completed");
1309-
recv_close_or_end(&mut ws).await;
1327+
recv_clean_close(&mut ws).await;
1328+
tokio::time::timeout(std::time::Duration::from_secs(2), websocket_tracker.wait_until_idle())
1329+
.await
1330+
.expect("server did not receive the websocket close acknowledgement");
13101331
assert_eq!(mock.request_bodies().await.len(), 1);
13111332
}
13121333

0 commit comments

Comments
 (0)