diff --git a/codex-rs/exec-server/src/relay.rs b/codex-rs/exec-server/src/relay.rs index ae07d7652..b0f9ed7a6 100644 --- a/codex-rs/exec-server/src/relay.rs +++ b/codex-rs/exec-server/src/relay.rs @@ -21,6 +21,7 @@ use crate::connection::CHANNEL_CAPACITY; use crate::connection::JsonRpcConnection; use crate::connection::JsonRpcConnectionEvent; use crate::connection::JsonRpcTransport; +use crate::connection::WEBSOCKET_KEEPALIVE_INTERVAL; use crate::relay_proto::RelayData; use crate::relay_proto::RelayMessageFrame; use crate::relay_proto::RelayResume; @@ -141,6 +142,37 @@ fn jsonrpc_payload(message: &JSONRPCMessage) -> Result, ExecServerError> serde_json::to_vec(message).map_err(ExecServerError::Json) } +enum RelayEventSendError { + IncomingClosed, + WebSocketClosed, +} + +async fn send_event_with_keepalive( + websocket: &mut T, + keepalive: &mut tokio::time::Interval, + incoming_tx: &mpsc::Sender, + event: JsonRpcConnectionEvent, +) -> Result<(), RelayEventSendError> +where + T: Sink + Unpin, +{ + let send = incoming_tx.send(event); + tokio::pin!(send); + loop { + tokio::select! { + result = &mut send => { + return result.map_err(|_| RelayEventSendError::IncomingClosed); + } + _ = keepalive.tick() => { + websocket + .send(Message::Ping(Vec::new().into())) + .await + .map_err(|_| RelayEventSendError::WebSocketClosed)?; + } + } + } +} + pub(crate) fn harness_connection_from_websocket( stream: T, connection_label: String, @@ -168,6 +200,11 @@ where return; } + let mut keepalive = tokio::time::interval_at( + tokio::time::Instant::now() + WEBSOCKET_KEEPALIVE_INTERVAL, + WEBSOCKET_KEEPALIVE_INTERVAL, + ); + keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); let mut next_seq = 0u32; loop { tokio::select! { @@ -193,6 +230,12 @@ where break; } } + _ = keepalive.tick() => { + if websocket.send(Message::Ping(Vec::new().into())).await.is_err() { + let _ = disconnected_tx.send(true); + break; + } + } incoming_message = websocket.next() => { match incoming_message { Some(Ok(Message::Binary(payload))) => { @@ -226,12 +269,20 @@ where match kind { RelayFrameBodyKind::Data => match frame.into_jsonrpc_message() { Ok(message) => { - if incoming_tx - .send(JsonRpcConnectionEvent::Message(message)) - .await - .is_err() + match send_event_with_keepalive( + &mut websocket, + &mut keepalive, + &incoming_tx, + JsonRpcConnectionEvent::Message(message), + ) + .await { - break; + Ok(()) => {} + Err(RelayEventSendError::IncomingClosed) => break, + Err(RelayEventSendError::WebSocketClosed) => { + let _ = disconnected_tx.send(true); + break; + } } } Err(err) => { @@ -308,8 +359,13 @@ pub(crate) async fn run_multiplexed_environment( let (physical_outgoing_tx, mut physical_outgoing_rx) = mpsc::channel::>(CHANNEL_CAPACITY); + let mut keepalive = tokio::time::interval_at( + tokio::time::Instant::now() + WEBSOCKET_KEEPALIVE_INTERVAL, + WEBSOCKET_KEEPALIVE_INTERVAL, + ); + keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); let mut streams: HashMap = HashMap::new(); - loop { + 'websocket: loop { let frame = tokio::select! { maybe_encoded = physical_outgoing_rx.recv() => { let Some(encoded) = maybe_encoded else { @@ -320,6 +376,12 @@ pub(crate) async fn run_multiplexed_environment( } continue; } + _ = keepalive.tick() => { + if websocket.send(Message::Ping(Vec::new().into())).await.is_err() { + break; + } + continue; + } incoming_message = websocket.next() => match incoming_message { Some(Ok(Message::Binary(payload))) => { match decode_relay_message_frame(payload.as_ref()) { @@ -368,13 +430,20 @@ pub(crate) async fn run_multiplexed_environment( physical_outgoing_tx.clone(), ) }); - if stream - .incoming_tx - .send(JsonRpcConnectionEvent::Message(message)) - .await - .is_err() + let incoming_tx = stream.incoming_tx.clone(); + match send_event_with_keepalive( + &mut websocket, + &mut keepalive, + &incoming_tx, + JsonRpcConnectionEvent::Message(message), + ) + .await { - streams.remove(&stream_id); + Ok(()) => {} + Err(RelayEventSendError::IncomingClosed) => { + streams.remove(&stream_id); + } + Err(RelayEventSendError::WebSocketClosed) => break 'websocket, } } RelayFrameBodyKind::Reset => { @@ -468,12 +537,14 @@ mod tests { use std::task::Poll; use std::time::Duration; + use codex_app_server_protocol::JSONRPCError; use codex_app_server_protocol::JSONRPCRequest; use codex_app_server_protocol::RequestId; use futures::Sink; use futures::Stream; use futures::channel::mpsc as futures_mpsc; use futures::task::AtomicWaker; + use pretty_assertions::assert_eq; use tokio::net::TcpListener; use tokio::time::timeout; use tokio_tungstenite::accept_async; @@ -483,11 +554,15 @@ mod tests { use super::*; #[tokio::test] - async fn harness_connection_receives_relay_data() -> anyhow::Result<()> { + async fn harness_connection_sends_keepalive_and_receives_relay_data() -> anyhow::Result<()> { let (client_websocket, mut server_websocket) = websocket_pair().await?; let mut connection = harness_connection_from_websocket(client_websocket, "test".to_string()); let stream_id = read_resume_stream_id(&mut server_websocket).await?; + read_keepalive_ping(&mut server_websocket).await?; + server_websocket + .send(Message::Pong(b"keepalive".to_vec().into())) + .await?; let message = test_jsonrpc_message(); server_websocket @@ -509,6 +584,126 @@ mod tests { Ok(()) } + #[tokio::test] + async fn multiplexed_environment_sends_keepalive_and_processes_relay_data() -> anyhow::Result<()> + { + let (client_websocket, mut server_websocket) = websocket_pair().await?; + let runtime_paths = crate::ExecServerRuntimePaths::new( + std::env::current_exe()?, + /*codex_linux_sandbox_exe*/ None, + ) + .map_err(anyhow::Error::from)?; + let environment_task = tokio::spawn(run_multiplexed_environment( + client_websocket, + ConnectionProcessor::new(runtime_paths), + )); + + read_keepalive_ping(&mut server_websocket).await?; + server_websocket + .send(Message::Pong(b"keepalive".to_vec().into())) + .await?; + let stream_id = "test-stream".to_string(); + let request = test_jsonrpc_message(); + server_websocket + .send(Message::Binary( + encode_relay_message_frame(&RelayMessageFrame::data( + stream_id.clone(), + /*seq*/ 0, + jsonrpc_payload(&request)?, + )) + .into(), + )) + .await?; + + let response_frame = timeout(Duration::from_secs(1), async { + loop { + let message = server_websocket + .next() + .await + .expect("websocket should stay open")?; + match message { + Message::Binary(payload) => { + return decode_relay_message_frame(payload.as_ref()) + .map_err(anyhow::Error::from); + } + Message::Ping(_) | Message::Pong(_) | Message::Frame(_) => {} + Message::Text(text) => { + anyhow::bail!("expected binary relay frame, got text: {text}"); + } + Message::Close(_) => { + anyhow::bail!("websocket closed before relay response"); + } + } + } + }) + .await??; + let expected_message = JSONRPCMessage::Error(JSONRPCError { + id: RequestId::Integer(1), + error: crate::rpc::method_not_found( + "exec-server stub does not implement `test` yet".to_string(), + ), + }); + assert_eq!( + response_frame, + RelayMessageFrame::data( + stream_id, + /*seq*/ 0, + jsonrpc_payload(&expected_message)?, + ) + ); + + environment_task.abort(); + let _ = environment_task.await; + Ok(()) + } + + #[tokio::test] + async fn send_event_with_keepalive_pings_while_incoming_queue_is_full() -> anyhow::Result<()> { + let (mut websocket, _control, mut outbound_rx) = + ControlledWebSocket::new(/*write_ready*/ true); + let (incoming_tx, mut incoming_rx) = mpsc::channel(/*buffer*/ 1); + let message = test_jsonrpc_message(); + let expected_message = message.clone(); + incoming_tx + .send(JsonRpcConnectionEvent::MalformedMessage { + reason: "first".to_string(), + }) + .await?; + let mut keepalive = tokio::time::interval_at( + tokio::time::Instant::now() + WEBSOCKET_KEEPALIVE_INTERVAL, + WEBSOCKET_KEEPALIVE_INTERVAL, + ); + keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + + let send_task = tokio::spawn(async move { + send_event_with_keepalive( + &mut websocket, + &mut keepalive, + &incoming_tx, + JsonRpcConnectionEvent::Message(message), + ) + .await + }); + + assert!(matches!( + timeout(Duration::from_secs(1), outbound_rx.next()).await?, + Some(Message::Ping(_)) + )); + assert!(matches!( + incoming_rx.recv().await, + Some(JsonRpcConnectionEvent::MalformedMessage { reason }) if reason == "first" + )); + assert!(matches!( + timeout(Duration::from_secs(1), send_task).await??, + Ok(()) + )); + assert!(matches!( + incoming_rx.recv().await, + Some(JsonRpcConnectionEvent::Message(actual)) if actual == expected_message + )); + Ok(()) + } + #[tokio::test] async fn harness_connection_reports_text_frames_as_malformed() -> anyhow::Result<()> { let (client_websocket, mut server_websocket) = websocket_pair().await?; @@ -612,6 +807,21 @@ mod tests { Ok(frame.stream_id) } + async fn read_keepalive_ping( + websocket: &mut WebSocketStream, + ) -> anyhow::Result<()> { + loop { + let Some(message) = timeout(Duration::from_secs(1), websocket.next()).await? else { + anyhow::bail!("websocket closed before keepalive ping"); + }; + match message? { + Message::Ping(_) => return Ok(()), + Message::Binary(_) | Message::Text(_) | Message::Pong(_) | Message::Frame(_) => {} + Message::Close(_) => anyhow::bail!("websocket closed before keepalive ping"), + } + } + } + fn test_jsonrpc_message() -> JSONRPCMessage { JSONRPCMessage::Request(JSONRPCRequest { id: RequestId::Integer(1),