@@ -7,7 +7,7 @@ use axum::http::HeaderMap;
77use axum:: response:: Response ;
88use either:: Either ;
99use futures:: stream:: { SplitSink , SplitStream } ;
10- use futures:: { SinkExt , Stream , StreamExt } ;
10+ use futures:: { Sink , SinkExt , Stream , StreamExt } ;
1111use serde_json:: Value ;
1212use tokio_util:: sync:: CancellationToken ;
1313use 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) ]
363410mod 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 ( ) {
0 commit comments