diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 4536b2188b..cfad1bd94a 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -2571,6 +2571,7 @@ dependencies = [ "tokio-util", "tracing", "tracing-subscriber", + "uuid", ] [[package]] diff --git a/codex-rs/code-mode-host/Cargo.toml b/codex-rs/code-mode-host/Cargo.toml index 5bad7553f3..2f1febc2fb 100644 --- a/codex-rs/code-mode-host/Cargo.toml +++ b/codex-rs/code-mode-host/Cargo.toml @@ -27,6 +27,7 @@ tokio = { workspace = true, features = ["io-std", "io-util", "macros", "net", "p tokio-util = { workspace = true, features = ["rt"] } tracing = { workspace = true } tracing-subscriber = { workspace = true } +uuid = { workspace = true, features = ["v4"] } [dev-dependencies] codex-code-mode = { workspace = true } diff --git a/codex-rs/code-mode-host/src/host_tests.rs b/codex-rs/code-mode-host/src/host_tests.rs index b3639e7cc0..3e78ae6911 100644 --- a/codex-rs/code-mode-host/src/host_tests.rs +++ b/codex-rs/code-mode-host/src/host_tests.rs @@ -12,6 +12,7 @@ use codex_code_mode_protocol::host::Capability; use codex_code_mode_protocol::host::CapabilitySet; use codex_code_mode_protocol::host::ClientHello; use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::DUAL_WEBSOCKET_CAPABILITY; use codex_code_mode_protocol::host::EncodedFrame; use codex_code_mode_protocol::host::FramedReader; use codex_code_mode_protocol::host::FramedWriter; @@ -33,6 +34,7 @@ use tokio::sync::mpsc; use tokio::sync::oneshot; use tokio_util::sync::CancellationToken; use tokio_util::task::TaskTracker; +use uuid::Uuid; use super::HostState; use super::MAX_ACTIVE_CELLS; @@ -41,8 +43,12 @@ use super::MAX_RECENT_REQUEST_IDS; use super::RequestKind; use super::RequestRegistry; use super::SeenSessionIds; +use super::negotiate; use super::peer::HostPeer; use super::run; +use super::transport::BulkConnectionRegistry; +use super::transport::ConnectionReader; +use super::transport::ConnectionWriter; fn client_hello( versions: impl IntoIterator, @@ -250,6 +256,107 @@ impl AsyncWrite for BlockingWriter { } } +#[tokio::test] +async fn failed_host_hello_removes_the_bulk_pairing_reservation() { + let (host_reader, client_writer) = tokio::io::duplex(/*max_buf_size*/ 4096); + let capability = Capability::new(DUAL_WEBSOCKET_CAPABILITY).expect("dual websocket capability"); + let hello = ClientToHost::ClientHello( + ClientHello::new( + SupportedProtocolVersions::try_new([ProtocolVersion::V1]).expect("supported versions"), + CapabilitySet::empty(), + CapabilitySet::try_new([capability]).expect("optional capabilities"), + ) + .expect("client hello"), + ); + FramedWriter::new(client_writer) + .write(&hello) + .await + .expect("write client hello"); + + let bytes = Arc::new(Mutex::new(Vec::new())); + let registry = BulkConnectionRegistry::default(); + let mut reader = ConnectionReader::from_reader(host_reader); + let mut writer = ConnectionWriter::from_writer(FailingHandshakeWriter { + bytes: Arc::clone(&bytes), + }); + let result = negotiate(&mut reader, &mut writer, Some(®istry)).await; + assert!(result.is_err()); + + let message = EncodedFrame::decode_framed::( + &bytes.lock().unwrap_or_else(PoisonError::into_inner), + ) + .expect("decode failed host hello"); + let HostToClient::HostHello(hello) = message else { + panic!("expected a host hello"); + }; + let token = hello + .bulk_connection_token() + .expect("dual websocket pairing token"); + let token = Uuid::parse_str(token).expect("UUID pairing token"); + assert!(registry.remove(token).is_none()); +} + +struct FailingHandshakeWriter { + bytes: Arc>>, +} + +impl AsyncWrite for FailingHandshakeWriter { + fn poll_write( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + bytes: &[u8], + ) -> Poll> { + self.bytes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .extend_from_slice(bytes); + Poll::Ready(Ok(bytes.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Err(std::io::Error::other("host hello write failed"))) + } + + fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } +} + +#[tokio::test] +async fn optional_dual_websocket_capability_falls_back_to_a_single_connection() { + let (host_stream, client_stream) = tokio::io::duplex(/*max_buf_size*/ 1024); + let (host_reader, host_writer) = tokio::io::split(host_stream); + let (client_reader, client_writer) = tokio::io::split(client_stream); + let host = tokio::spawn(run(host_reader, host_writer)); + let mut reader = FramedReader::new(client_reader); + let mut writer = FramedWriter::new(client_writer); + let capability = Capability::new(DUAL_WEBSOCKET_CAPABILITY).expect("dual websocket capability"); + + writer + .write(&ClientToHost::ClientHello( + ClientHello::new( + SupportedProtocolVersions::try_new([ProtocolVersion::V1]) + .expect("supported versions"), + CapabilitySet::empty(), + CapabilitySet::try_new([capability]).expect("optional capabilities"), + ) + .expect("client hello"), + )) + .await + .expect("write hello"); + assert_eq!( + reader.read::().await.expect("host hello"), + Some(HostToClient::HostHello(HostHello::new( + ProtocolVersion::V1, + CapabilitySet::empty(), + ))) + ); + + drop(writer); + drop(reader); + host.await.expect("host task").expect("host connection"); +} + #[tokio::test] async fn incompatible_or_invalid_handshake_is_rejected() { let (host_stream, client_stream) = tokio::io::duplex(/*max_buf_size*/ 1024); @@ -395,7 +502,7 @@ async fn session_id_cannot_be_reused_after_shutdown() { } #[test] -fn request_cancellation_tombstones_are_bounded() { +fn request_history_is_bounded() { let mut requests = RequestRegistry::default(); let duplicate = request_id(/*value*/ -1); requests @@ -411,10 +518,6 @@ fn request_cancellation_tombstones_are_bounded() { requests.cancel(id); requests.finish(id); } - for value in 10_000..20_000 { - requests.cancel(request_id(value)); - } - assert!(requests.active.is_empty()); assert_eq!(requests.recent.len(), MAX_RECENT_REQUEST_IDS); assert_eq!(requests.recent_order.len(), MAX_RECENT_REQUEST_IDS); diff --git a/codex-rs/code-mode-host/src/lib.rs b/codex-rs/code-mode-host/src/lib.rs index 83b545d7e7..88d2e161f5 100644 --- a/codex-rs/code-mode-host/src/lib.rs +++ b/codex-rs/code-mode-host/src/lib.rs @@ -10,18 +10,22 @@ use std::time::Duration; use anyhow::Context; use anyhow::Result; +use codex_code_mode_protocol::host::Capability; use codex_code_mode_protocol::host::CapabilitySet; use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::DUAL_WEBSOCKET_CAPABILITY; use codex_code_mode_protocol::host::EncodedFrame; use codex_code_mode_protocol::host::HandshakeRejectReason; use codex_code_mode_protocol::host::HostHello; use codex_code_mode_protocol::host::HostRequest; use codex_code_mode_protocol::host::HostResponse; use codex_code_mode_protocol::host::HostToClient; +use codex_code_mode_protocol::host::MAX_PENDING_DELEGATE_CALLS; use codex_code_mode_protocol::host::ProtocolVersion; use codex_code_mode_protocol::host::RequestId; use codex_code_mode_protocol::host::SessionId; use codex_code_mode_protocol::host::SupportedProtocolVersions; +use codex_code_mode_protocol::host::TransportLane; use codex_code_mode_runtime::InProcessCodeModeSession; use tokio::io::AsyncRead; use tokio::io::AsyncWrite; @@ -32,6 +36,8 @@ use tokio_util::task::TaskTracker; use self::delegate::RemoteDelegate; use self::peer::HostPeer; +use self::transport::BulkConnectionRegistration; +use self::transport::BulkConnectionRegistry; use self::transport::ConnectionReader; use self::transport::ConnectionWriter; @@ -46,6 +52,14 @@ const MAX_ACTIVE_CELLS: usize = 128; const MAX_RECENT_REQUEST_IDS: usize = 4096; const MAX_RECENT_SESSION_IDS: usize = 4096; const SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(5); +const OUTGOING_CHANNEL_CAPACITY: usize = 128; +const BULK_PAIRING_TIMEOUT: Duration = Duration::from_secs(10); + +enum NegotiatedConnection { + Rejected, + Single, + Dual(BulkConnectionRegistration), +} struct HostLimits { request_permits: Arc, @@ -81,6 +95,7 @@ where ConnectionReader::from_reader(reader), ConnectionWriter::from_writer(writer), Arc::new(HostLimits::new()), + /*bulk_connections*/ None, ) .await } @@ -89,13 +104,39 @@ async fn run_connection( mut reader: ConnectionReader, mut writer: ConnectionWriter, limits: Arc, + bulk_connections: Option, ) -> Result<()> { - if !negotiate(&mut reader, &mut writer).await? { - return Ok(()); - } - - let (outgoing_tx, mut outgoing_rx) = mpsc::channel::(/*max_capacity*/ 128); - let peer = Arc::new(HostPeer::new(outgoing_tx)); + let negotiated = negotiate(&mut reader, &mut writer, bulk_connections.as_ref()).await?; + let bulk_connection = match negotiated { + NegotiatedConnection::Rejected => return Ok(()), + NegotiatedConnection::Single => None, + NegotiatedConnection::Dual(mut registration) => { + match tokio::time::timeout(BULK_PAIRING_TIMEOUT, registration.receive()).await { + Ok(Ok(connection)) => Some(connection), + Ok(Err(_)) => { + anyhow::bail!("code-mode host bulk websocket pairing was abandoned"); + } + Err(_) => { + anyhow::bail!("timed out pairing code-mode host bulk websocket"); + } + } + } + }; + let (mut bulk_reader, bulk_writer) = match bulk_connection { + Some(connection) => (Some(connection.reader), Some(connection.writer)), + None => (None, None), + }; + let (outgoing_tx, outgoing_rx) = mpsc::channel::(OUTGOING_CHANNEL_CAPACITY); + let (bulk_tx, bulk_rx) = if bulk_writer.is_some() { + let (sender, receiver) = mpsc::channel::(MAX_PENDING_DELEGATE_CALLS); + (Some(sender), Some(receiver)) + } else { + (None, None) + }; + let peer = match bulk_tx { + Some(sender) => Arc::new(HostPeer::new(outgoing_tx).with_bulk_sender(sender)), + None => Arc::new(HostPeer::new(outgoing_tx)), + }; let state = Arc::new(HostState { sessions: Mutex::new(HashMap::new()), seen_session_ids: Mutex::new(SeenSessionIds::default()), @@ -108,25 +149,14 @@ async fn run_connection( }); let writer_disconnected = peer.disconnection_token(); let writer_task = tokio::spawn(async move { - loop { - tokio::select! { - _ = writer_disconnected.cancelled() => return Ok::<(), anyhow::Error>(()), - frame = outgoing_rx.recv() => { - let Some(frame) = frame else { - return Ok(()); - }; - let result = tokio::select! { - _ = writer_disconnected.cancelled() => return Ok(()), - result = writer.write_frame(frame) => result, - }; - if let Err(err) = result { - return Err( - anyhow::Error::new(err) - .context("failed to write code-mode host message") - ); - } - } - } + if let (Some(bulk_writer), Some(bulk_rx)) = (bulk_writer, bulk_rx) { + tokio::try_join!( + drive_writer(writer, outgoing_rx, writer_disconnected.clone()), + drive_writer(bulk_writer, bulk_rx, writer_disconnected) + )?; + Ok(()) + } else { + drive_writer(writer, outgoing_rx, writer_disconnected).await } }); let writer_peer = Arc::clone(&peer); @@ -147,14 +177,30 @@ async fn run_connection( let input_result = async { loop { - let message = tokio::select! { + let (message, lane) = tokio::select! { + // Session operations and shutdown must make progress even when bulk callbacks are ready. + biased; _ = peer.disconnected() => break, - message = reader.read() => message - .context("failed to read code-mode client message")?, + message = reader.read() => ( + message.context("failed to read code-mode client control message")?, + TransportLane::Control, + ), + message = async { + match &mut bulk_reader { + Some(reader) => reader.read().await, + None => Ok(None), + } + }, if bulk_reader.is_some() => ( + message.context("failed to read code-mode client bulk message")?, + TransportLane::Bulk, + ), }; let Some(message) = message else { break; }; + if bulk_reader.is_some() && !message.allows_transport_lane(lane) { + anyhow::bail!("code-mode client sent a message on the wrong websocket lane"); + } match message { ClientToHost::ClientHello(_) => { anyhow::bail!("received a second code-mode client hello"); @@ -195,13 +241,40 @@ async fn run_connection( Ok(()) } -async fn negotiate(reader: &mut ConnectionReader, writer: &mut ConnectionWriter) -> Result { +async fn drive_writer( + mut writer: ConnectionWriter, + mut outgoing: mpsc::Receiver, + disconnected: CancellationToken, +) -> Result<()> { + loop { + tokio::select! { + _ = disconnected.cancelled() => return Ok(()), + frame = outgoing.recv() => { + let Some(frame) = frame else { + return Ok(()); + }; + tokio::select! { + _ = disconnected.cancelled() => return Ok(()), + result = writer.write_frame(frame) => { + result.context("failed to write code-mode host message")?; + } + } + } + } + } +} + +async fn negotiate( + reader: &mut ConnectionReader, + writer: &mut ConnectionWriter, + bulk_connections: Option<&BulkConnectionRegistry>, +) -> Result { let Some(first_message) = reader .read() .await .context("failed to read code-mode client hello")? else { - return Ok(false); + return Ok(NegotiatedConnection::Rejected); }; let ClientToHost::ClientHello(client_hello) = first_message else { writer @@ -212,7 +285,7 @@ async fn negotiate(reader: &mut ConnectionReader, writer: &mut ConnectionWriter) }) .await .context("failed to reject invalid code-mode client hello")?; - return Ok(false); + return Ok(NegotiatedConnection::Rejected); }; let supported_versions = SupportedProtocolVersions::try_new([ProtocolVersion::V1])?; @@ -226,10 +299,26 @@ async fn negotiate(reader: &mut ConnectionReader, writer: &mut ConnectionWriter) }) .await .context("failed to reject incompatible code-mode client")?; - return Ok(false); + return Ok(NegotiatedConnection::Rejected); } - let host_capabilities = CapabilitySet::empty(); + let dual_capability = Capability::new(DUAL_WEBSOCKET_CAPABILITY)?; + let dual_requested = client_hello + .required_capabilities() + .contains(&dual_capability) + || client_hello + .optional_capabilities() + .contains(&dual_capability); + let registration = if dual_requested { + bulk_connections.and_then(BulkConnectionRegistry::reserve) + } else { + None + }; + let host_capabilities = if registration.is_some() { + CapabilitySet::try_new([dual_capability])? + } else { + CapabilitySet::empty() + }; if let Some(capability) = client_hello .required_capabilities() .iter() @@ -243,17 +332,26 @@ async fn negotiate(reader: &mut ConnectionReader, writer: &mut ConnectionWriter) }) .await .context("failed to reject unsupported code-mode capability")?; - return Ok(false); + return Ok(NegotiatedConnection::Rejected); } + let (hello, negotiated) = if let Some(registration) = registration { + ( + HostHello::new(ProtocolVersion::V1, host_capabilities) + .with_bulk_connection_token(registration.token().to_string()), + NegotiatedConnection::Dual(registration), + ) + } else { + ( + HostHello::new(ProtocolVersion::V1, host_capabilities), + NegotiatedConnection::Single, + ) + }; writer - .write(&HostToClient::HostHello(HostHello::new( - ProtocolVersion::V1, - host_capabilities, - ))) + .write(&HostToClient::HostHello(hello)) .await .context("failed to write code-mode host hello")?; - Ok(true) + Ok(negotiated) } struct HostState { @@ -584,7 +682,7 @@ impl RequestRegistry { Ok(cancellation) } - fn cancel(&self, request_id: RequestId) { + fn cancel(&mut self, request_id: RequestId) { if let Some(request) = self.active.get(&request_id) && request.kind.is_cancellable() { diff --git a/codex-rs/code-mode-host/src/peer.rs b/codex-rs/code-mode-host/src/peer.rs index 7912051802..02b0c263dc 100644 --- a/codex-rs/code-mode-host/src/peer.rs +++ b/codex-rs/code-mode-host/src/peer.rs @@ -16,6 +16,7 @@ use codex_code_mode_protocol::host::HostToClient; use codex_code_mode_protocol::host::MAX_PENDING_DELEGATE_CALLS; use codex_code_mode_protocol::host::RequestId; use codex_code_mode_protocol::host::SessionId; +use codex_code_mode_protocol::host::TransportLane; use codex_code_mode_protocol::host::WireResult; use tokio::sync::Mutex; use tokio::sync::Notify; @@ -29,6 +30,7 @@ const CELL_MESSAGE_CAPACITY: usize = 128; pub(super) struct HostPeer { outgoing_tx: mpsc::Sender, + bulk_tx: Option>, pending: Mutex>, delegate_permits: Arc, cell_routes: StdMutex>, @@ -62,6 +64,7 @@ impl HostPeer { pub(super) fn new(outgoing_tx: mpsc::Sender) -> Self { Self { outgoing_tx, + bulk_tx: None, pending: Mutex::new(HashMap::new()), delegate_permits: Arc::new(Semaphore::new(MAX_PENDING_DELEGATE_CALLS)), cell_routes: StdMutex::new(HashMap::new()), @@ -72,10 +75,16 @@ impl HostPeer { } } + pub(super) fn with_bulk_sender(mut self, bulk_tx: mpsc::Sender) -> Self { + self.bulk_tx = Some(bulk_tx); + self + } + pub(super) fn send(&self, message: HostToClient) -> Result<(), PeerSendError> { let frame = EncodedFrame::encode(&message) .map_err(|err| PeerSendError::Payload(err.to_string()))?; - self.send_frame(frame) + let lane = message.transport_lane(); + self.send_frame(frame, lane) } pub(super) fn respond( @@ -389,8 +398,12 @@ impl HostPeer { }); } - fn send_frame(&self, frame: EncodedFrame) -> Result<(), PeerSendError> { - match self.outgoing_tx.try_send(frame) { + fn send_frame(&self, frame: EncodedFrame, lane: TransportLane) -> Result<(), PeerSendError> { + let sender = match lane { + TransportLane::Control => &self.outgoing_tx, + TransportLane::Bulk => self.bulk_tx.as_ref().unwrap_or(&self.outgoing_tx), + }; + match sender.try_send(frame) { Ok(()) => Ok(()), Err(mpsc::error::TrySendError::Full(_)) => { self.disconnect(); diff --git a/codex-rs/code-mode-host/src/transport.rs b/codex-rs/code-mode-host/src/transport.rs index 3342fc217c..c06fca40e3 100644 --- a/codex-rs/code-mode-host/src/transport.rs +++ b/codex-rs/code-mode-host/src/transport.rs @@ -1,13 +1,17 @@ +use std::collections::HashMap; use std::io; use std::io::Write as _; use std::net::SocketAddr; use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; use anyhow::Context; use anyhow::Result; use axum::Router; use axum::body::Body; use axum::extract::ConnectInfo; +use axum::extract::Path; use axum::extract::State; use axum::extract::ws::Message; use axum::extract::ws::WebSocket; @@ -34,15 +38,19 @@ use futures::stream::SplitStream; use tokio::io::AsyncRead; use tokio::io::AsyncWrite; use tokio::net::TcpListener; +use tokio::sync::oneshot; use tracing::info; use tracing::warn; +use uuid::Uuid; use crate::HostLimits; +use crate::MAX_IN_FLIGHT_REQUESTS; /// The default transport retains the standalone host's original stdio behavior. pub const DEFAULT_LISTEN_URL: &str = "stdio"; const MAX_WEBSOCKET_FRAME_BYTES: usize = MAX_FRAME_BYTES + std::mem::size_of::(); +const MAX_PENDING_BULK_CONNECTIONS: usize = MAX_IN_FLIGHT_REQUESTS; type BoxedReader = Box; type BoxedWriter = Box; @@ -66,6 +74,64 @@ pub(crate) enum ConnectionWriter { #[derive(Clone)] struct WebSocketListenerState { limits: Arc, + bulk_connections: BulkConnectionRegistry, +} + +pub(crate) struct BulkConnection { + pub(crate) reader: ConnectionReader, + pub(crate) writer: ConnectionWriter, +} + +#[derive(Clone, Default)] +pub(crate) struct BulkConnectionRegistry { + pending: Arc>>>, +} + +pub(crate) struct BulkConnectionRegistration { + registry: BulkConnectionRegistry, + token: Uuid, + receiver: oneshot::Receiver, +} + +impl BulkConnectionRegistry { + pub(crate) fn reserve(&self) -> Option { + let mut pending = self.pending.lock().unwrap_or_else(PoisonError::into_inner); + if pending.len() >= MAX_PENDING_BULK_CONNECTIONS { + return None; + } + + let token = Uuid::new_v4(); + let (sender, receiver) = oneshot::channel(); + pending.insert(token, sender); + Some(BulkConnectionRegistration { + registry: self.clone(), + token, + receiver, + }) + } + + pub(crate) fn remove(&self, token: Uuid) -> Option> { + self.pending + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(&token) + } +} + +impl BulkConnectionRegistration { + pub(crate) fn token(&self) -> Uuid { + self.token + } + + pub(crate) async fn receive(&mut self) -> Result { + (&mut self.receiver).await + } +} + +impl Drop for BulkConnectionRegistration { + fn drop(&mut self) { + self.registry.remove(self.token); + } } impl ConnectionReader { @@ -165,6 +231,7 @@ async fn run_websocket_listener(bind_address: SocketAddr) -> Result<()> { .context("failed to read code-mode host websocket listen address")?; let state = WebSocketListenerState { limits: Arc::new(HostLimits::new()), + bulk_connections: BulkConnectionRegistry::default(), }; info!("codex-code-mode-host listening on ws://{local_addr}"); println!("ws://{local_addr}"); @@ -174,6 +241,7 @@ async fn run_websocket_listener(bind_address: SocketAddr) -> Result<()> { let router = Router::new() .route("/", any(websocket_upgrade_handler)) + .route("/bulk/{token}", any(bulk_websocket_upgrade_handler)) .route("/readyz", get(readiness_handler)) .layer(middleware::from_fn(reject_requests_with_origin_header)) .with_state(state); @@ -220,6 +288,7 @@ async fn websocket_upgrade_handler( ConnectionReader::WebSocket(reader), ConnectionWriter::WebSocket(writer), state.limits, + Some(state.bulk_connections), ) .await { @@ -228,6 +297,38 @@ async fn websocket_upgrade_handler( }) } +async fn bulk_websocket_upgrade_handler( + websocket: WebSocketUpgrade, + Path(token): Path, + ConnectInfo(peer_addr): ConnectInfo, + State(state): State, +) -> Response { + let Ok(token) = Uuid::parse_str(&token) else { + return StatusCode::NOT_FOUND.into_response(); + }; + let Some(pending) = state.bulk_connections.remove(token) else { + return StatusCode::NOT_FOUND.into_response(); + }; + + websocket + .max_frame_size(MAX_WEBSOCKET_FRAME_BYTES) + .max_message_size(MAX_WEBSOCKET_FRAME_BYTES) + .on_upgrade(move |stream| async move { + info!(%peer_addr, "code-mode host bulk websocket client connected"); + let (writer, reader) = stream.split(); + if pending + .send(BulkConnection { + reader: ConnectionReader::WebSocket(reader), + writer: ConnectionWriter::WebSocket(writer), + }) + .is_err() + { + warn!(%peer_addr, "code-mode host bulk websocket pairing expired"); + } + }) + .into_response() +} + #[cfg(test)] #[path = "transport_tests.rs"] mod tests; diff --git a/codex-rs/code-mode-host/src/transport_tests.rs b/codex-rs/code-mode-host/src/transport_tests.rs index b62b63b380..7b4955f488 100644 --- a/codex-rs/code-mode-host/src/transport_tests.rs +++ b/codex-rs/code-mode-host/src/transport_tests.rs @@ -2,9 +2,34 @@ use std::net::SocketAddr; use pretty_assertions::assert_eq; +use super::BulkConnectionRegistry; use super::ListenTransport; +use super::MAX_PENDING_BULK_CONNECTIONS; use super::parse_listen_url; +#[test] +fn bulk_connection_registration_cleans_up_when_dropped() { + let registry = BulkConnectionRegistry::default(); + let registration = registry.reserve().expect("bulk connection registration"); + let token = registration.token(); + + drop(registration); + + assert!(registry.remove(token).is_none()); +} + +#[test] +fn bulk_connection_registrations_are_bounded_and_released() { + let registry = BulkConnectionRegistry::default(); + let mut registrations = (0..MAX_PENDING_BULK_CONNECTIONS) + .map(|_| registry.reserve().expect("bulk connection registration")) + .collect::>(); + + assert!(registry.reserve().is_none()); + drop(registrations.pop()); + assert!(registry.reserve().is_some()); +} + #[test] fn parse_listen_url_accepts_stdio_transports() { assert_eq!( diff --git a/codex-rs/code-mode-host/tests/websocket.rs b/codex-rs/code-mode-host/tests/websocket.rs index a0f471e737..030f83432d 100644 --- a/codex-rs/code-mode-host/tests/websocket.rs +++ b/codex-rs/code-mode-host/tests/websocket.rs @@ -1,12 +1,27 @@ use std::process::Stdio; +use std::sync::Arc; use std::time::Duration; use anyhow::Context; use anyhow::Result; +use codex_code_mode::CellId; +use codex_code_mode::CodeModeNestedToolCall; +use codex_code_mode::CodeModeSessionDelegate; +use codex_code_mode::CodeModeSessionProvider; +use codex_code_mode::CodeModeToolKind; +use codex_code_mode::ExecuteRequest; +use codex_code_mode::FunctionCallOutputContentItem; +use codex_code_mode::NoopCodeModeSessionDelegate; +use codex_code_mode::NotificationFuture; +use codex_code_mode::RuntimeResponse; +use codex_code_mode::ToolDefinition; +use codex_code_mode::ToolInvocationFuture; +use codex_code_mode::WebSocketCodeModeSessionProvider; use codex_code_mode_protocol::host::Capability; use codex_code_mode_protocol::host::CapabilitySet; use codex_code_mode_protocol::host::ClientHello; use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::DUAL_WEBSOCKET_CAPABILITY; use codex_code_mode_protocol::host::DelegateRequest; use codex_code_mode_protocol::host::DelegateResponse; use codex_code_mode_protocol::host::EncodedFrame; @@ -26,6 +41,8 @@ use codex_code_mode_protocol::host::WireRuntimeResponse; use codex_code_mode_protocol::host::WireToolDefinition; use codex_code_mode_protocol::host::WireToolKind; use codex_code_mode_protocol::host::WireToolName; +use codex_code_mode_protocol::host::WireWaitRequest; +use codex_protocol::ToolName; use futures::SinkExt; use futures::StreamExt; use pretty_assertions::assert_eq; @@ -37,6 +54,7 @@ use tokio::io::BufReader; use tokio::net::TcpStream; use tokio::process::Child; use tokio::process::Command; +use tokio::sync::Semaphore; use tokio::time::timeout; use tokio_tungstenite::MaybeTlsStream; use tokio_tungstenite::WebSocketStream; @@ -49,6 +67,8 @@ use tokio_tungstenite::tungstenite::http::HeaderValue; use tokio_tungstenite::tungstenite::http::StatusCode; use tokio_tungstenite::tungstenite::http::header::ORIGIN; use tokio_tungstenite::tungstenite::protocol::WebSocketConfig; +use tokio_util::sync::CancellationToken; +use uuid::Uuid; const TEST_TIMEOUT: Duration = Duration::from_secs(10); const MAX_WEBSOCKET_FRAME_BYTES: usize = MAX_FRAME_BYTES + std::mem::size_of::(); @@ -62,6 +82,43 @@ struct HostClient { websocket: WebSocketStream>, } +struct LargeToolResultDelegate { + started: Semaphore, + release: Semaphore, +} + +impl CodeModeSessionDelegate for LargeToolResultDelegate { + fn invoke_tool<'a>( + &'a self, + invocation: CodeModeNestedToolCall, + _cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + Box::pin(async move { + assert_eq!(invocation.tool_name, ToolName::plain("large")); + self.started.add_permits(1); + let permit = self + .release + .acquire() + .await + .map_err(|_| "large tool release closed".to_string())?; + permit.forget(); + Ok(json!({ "value": "x".repeat(8 * 1024 * 1024) })) + }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, _cell_id: &CellId) {} +} + impl HostHarness { async fn start() -> Result { let host_program = codex_utils_cargo_bin::cargo_bin("codex-code-mode-host")?; @@ -166,6 +223,34 @@ impl HostClient { Ok(()) } + async fn negotiate_dual(&mut self, websocket_url: &str) -> Result { + let capability = Capability::new(DUAL_WEBSOCKET_CAPABILITY)?; + let hello = ClientHello::new( + SupportedProtocolVersions::try_new([ProtocolVersion::V1])?, + CapabilitySet::empty(), + CapabilitySet::try_new([capability.clone()])?, + )?; + self.send(&ClientToHost::ClientHello(hello)).await?; + let HostToClient::HostHello(hello) = self.read().await? else { + anyhow::bail!("expected code-mode host hello"); + }; + assert!(hello.capabilities().contains(&capability)); + let token = hello + .bulk_connection_token() + .context("dual websocket handshake omitted its pairing token")?; + let bulk_url = format!("{}/bulk/{token}", websocket_url.trim_end_matches('/')); + let config = WebSocketConfig::default() + .max_frame_size(Some(MAX_WEBSOCKET_FRAME_BYTES)) + .max_message_size(Some(MAX_WEBSOCKET_FRAME_BYTES)); + let (websocket, _) = timeout( + TEST_TIMEOUT, + connect_async_with_config(bulk_url, Some(config), /*disable_nagle*/ false), + ) + .await + .context("timed out connecting to code-mode host bulk websocket")??; + Ok(HostClient { websocket }) + } + async fn open_session(&mut self, session_id: SessionId) -> Result<()> { let id = RequestId::new(/*value*/ 1); self.send(&ClientToHost::Request { @@ -317,6 +402,418 @@ async fn websocket_listener_executes_cells_and_forwards_tool_callbacks() -> Resu Ok(()) } +#[tokio::test] +async fn production_websocket_client_runs_nested_tools_while_other_sessions_progress() -> Result<()> +{ + let host = HostHarness::start().await?; + let provider = WebSocketCodeModeSessionProvider::new(host.websocket_url.clone()); + let delegate = Arc::new(LargeToolResultDelegate { + started: Semaphore::new(/*permits*/ 0), + release: Semaphore::new(/*permits*/ 0), + }); + let slow_session = provider + .create_session(delegate.clone()) + .await + .map_err(anyhow::Error::msg)?; + let fast_session = provider + .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .await + .map_err(anyhow::Error::msg)?; + + let slow_cell = slow_session + .execute(ExecuteRequest { + tool_call_id: "large-tool".to_string(), + enabled_tools: vec![ToolDefinition { + name: "large".to_string(), + tool_name: ToolName::plain("large"), + description: String::new(), + kind: CodeModeToolKind::Function, + input_schema: None, + output_schema: None, + }], + source: r#"const result = await tools.large({ value: "request" }); text(String(result.value.length));"# + .to_string(), + yield_time_ms: Some(20_000), + max_output_tokens: Some(1_000), + }) + .await + .map_err(anyhow::Error::msg)?; + let started = timeout(TEST_TIMEOUT, delegate.started.acquire()) + .await + .context("large tool callback did not start")??; + started.forget(); + + let fast_response = timeout(TEST_TIMEOUT, async { + fast_session + .execute(ExecuteRequest { + tool_call_id: "fast-before-transfer".to_string(), + enabled_tools: Vec::new(), + source: r#"text("fast-before");"#.to_string(), + yield_time_ms: Some(5_000), + max_output_tokens: Some(1_000), + }) + .await + .map_err(anyhow::Error::msg)? + .initial_response() + .await + .map_err(anyhow::Error::msg) + }) + .await + .context("unrelated execution was blocked by the pending tool callback")??; + assert_eq!( + fast_response, + RuntimeResponse::Result { + cell_id: CellId::new("1".to_string()), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "fast-before".to_string(), + }], + error_text: None, + } + ); + + delegate.release.add_permits(1); + let slow_response = async { + timeout(TEST_TIMEOUT, slow_cell.initial_response()) + .await + .context("large tool result did not finish")? + .map_err(anyhow::Error::msg) + }; + let concurrent_response = async { + timeout(TEST_TIMEOUT, async { + fast_session + .execute(ExecuteRequest { + tool_call_id: "fast-during-transfer".to_string(), + enabled_tools: Vec::new(), + source: r#"text("fast-during");"#.to_string(), + yield_time_ms: Some(5_000), + max_output_tokens: Some(1_000), + }) + .await + .map_err(anyhow::Error::msg)? + .initial_response() + .await + .map_err(anyhow::Error::msg) + }) + .await + .context("unrelated execution was blocked by the large tool transfer")? + }; + tokio::pin!(slow_response); + tokio::pin!(concurrent_response); + let concurrent_response = tokio::select! { + response = &mut concurrent_response => response?, + response = &mut slow_response => { + response?; + anyhow::bail!("large tool result completed before the unrelated control response"); + } + }; + let slow_response = slow_response.await?; + assert_eq!( + slow_response, + RuntimeResponse::Result { + cell_id: CellId::new("1".to_string()), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "8388608".to_string(), + }], + error_text: None, + } + ); + assert_eq!( + concurrent_response, + RuntimeResponse::Result { + cell_id: CellId::new("2".to_string()), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "fast-during".to_string(), + }], + error_text: None, + } + ); + + slow_session.shutdown().await.map_err(anyhow::Error::msg)?; + fast_session.shutdown().await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[tokio::test] +async fn websocket_dual_connections_route_notifications_and_tool_callbacks_to_separate_lanes() +-> Result<()> { + let host = HostHarness::start().await?; + let mut control = host.connect().await?; + let mut bulk = control.negotiate_dual(&host.websocket_url).await?; + let session_id = SessionId::new("dual-websocket-session")?; + control.open_session(session_id.clone()).await?; + + let execute_id = RequestId::new(/*value*/ 2); + control + .send(&ClientToHost::Request { + id: execute_id, + request: HostRequest::Execute { + session_id: session_id.clone(), + request: WireExecuteRequest { + tool_call_id: "dual-websocket-call".to_string(), + enabled_tools: vec![WireToolDefinition { + name: "echo".to_string(), + tool_name: WireToolName { + name: "echo".to_string(), + namespace: None, + }, + description: String::new(), + kind: WireToolKind::Function, + input_schema: None, + output_schema: None, + }], + source: r#"notify("important"); const result = await tools.echo({ value: "ping" }); text(result.value);"# + .to_string(), + yield_time_ms: Some(5_000), + max_output_tokens: Some(1_000), + }, + }, + }) + .await?; + + let started = control.read().await?; + let HostToClient::Response { + id, + result: + WireResult::Ok { + value: HostResponse::ExecutionStarted { cell_id }, + }, + } = started + else { + anyhow::bail!("expected execution-started response on control lane, got {started:?}"); + }; + assert_eq!(id, execute_id); + + let notification = control.read().await?; + let HostToClient::DelegateRequest { + id: notification_id, + .. + } = ¬ification + else { + anyhow::bail!("expected notification on control lane, got {notification:?}"); + }; + let notification_id = *notification_id; + assert_eq!( + notification, + HostToClient::DelegateRequest { + id: notification_id, + session_id: session_id.clone(), + request: DelegateRequest::Notify { + call_id: "dual-websocket-call".to_string(), + cell_id: cell_id.clone(), + text: "important".to_string(), + }, + } + ); + control + .send(&ClientToHost::DelegateResponse { + id: notification_id, + result: WireResult::Ok { + value: DelegateResponse::NotificationDelivered, + }, + }) + .await?; + + let callback = bulk.read().await?; + let HostToClient::DelegateRequest { + id: delegate_id, + session_id: callback_session_id, + request: DelegateRequest::InvokeTool { invocation }, + } = callback + else { + anyhow::bail!("expected tool callback on bulk lane, got {callback:?}"); + }; + assert_eq!(callback_session_id, session_id); + assert_eq!(invocation.input, Some(json!({ "value": "ping" }))); + + bulk.send(&ClientToHost::DelegateResponse { + id: delegate_id, + result: WireResult::Ok { + value: DelegateResponse::ToolResult { + result: json!({ "value": "pong" }), + }, + }, + }) + .await?; + + assert_eq!( + control.read().await?, + HostToClient::InitialResponse { + id: execute_id, + result: WireResult::Ok { + value: WireRuntimeResponse::Result { + cell_id: cell_id.clone(), + content_items: vec![WireContentItem::InputText { + text: "pong".to_string(), + }], + error_text: None, + }, + }, + } + ); + + assert_eq!( + control.read().await?, + HostToClient::CellClosed { + session_id: session_id.clone(), + cell_id, + } + ); + + let shutdown_id = RequestId::new(/*value*/ 3); + control + .send(&ClientToHost::Request { + id: shutdown_id, + request: HostRequest::ShutdownSession { + session_id: session_id.clone(), + }, + }) + .await?; + assert_eq!( + control.read().await?, + HostToClient::Response { + id: shutdown_id, + result: WireResult::Ok { + value: HostResponse::SessionClosed { session_id }, + }, + } + ); + Ok(()) +} + +#[tokio::test] +async fn websocket_control_operations_bypass_an_incomplete_bulk_frame() -> Result<()> { + let host = HostHarness::start().await?; + let mut control = host.connect().await?; + let mut bulk = control.negotiate_dual(&host.websocket_url).await?; + let session_id = SessionId::new("bulk-priority-session")?; + control.open_session(session_id.clone()).await?; + + let bulk_message = ClientToHost::DelegateResponse { + id: codex_code_mode_protocol::host::DelegateRequestId::new(/*value*/ 999), + result: WireResult::Ok { + value: DelegateResponse::ToolResult { + result: json!({ "image": "x".repeat(1024 * 1024) }), + }, + }, + }; + let payload = EncodedFrame::encode(&bulk_message)?.into_framed_bytes(); + let mask = [0x13_u8, 0x37, 0xc0, 0xde]; + let mut partial_frame = vec![0x82_u8, 0xff]; + partial_frame.extend_from_slice(&(payload.len() as u64).to_be_bytes()); + partial_frame.extend_from_slice(&mask); + partial_frame.extend( + payload[..64 * 1024] + .iter() + .enumerate() + .map(|(index, byte)| byte ^ mask[index % mask.len()]), + ); + bulk.websocket + .get_mut() + .write_all(&partial_frame) + .await + .context("failed to start the intentionally incomplete bulk frame")?; + + let request_id = RequestId::new(/*value*/ 42); + control + .send(&ClientToHost::Request { + id: request_id, + request: HostRequest::Wait { + session_id: session_id.clone(), + request: WireWaitRequest { + cell_id: codex_code_mode_protocol::host::WireCellId::new("missing-cell"), + yield_time_ms: 10, + }, + }, + }) + .await?; + let response = timeout(TEST_TIMEOUT, control.read()) + .await + .context("control wait was blocked behind the incomplete bulk frame")??; + assert!(matches!( + response, + HostToClient::Response { + id, + result: WireResult::Ok { + value: HostResponse::WaitCompleted { .. }, + }, + } if id == request_id + )); + + let execute_id = RequestId::new(/*value*/ 43); + control + .send(&ClientToHost::Request { + id: execute_id, + request: HostRequest::Execute { + session_id, + request: WireExecuteRequest { + tool_call_id: "control-execute".to_string(), + enabled_tools: Vec::new(), + source: r#"text("fast");"#.to_string(), + yield_time_ms: Some(5_000), + max_output_tokens: Some(1_000), + }, + }, + }) + .await?; + for _ in 0..2 { + let response = timeout(TEST_TIMEOUT, control.read()) + .await + .context("control execute was blocked behind the incomplete bulk frame")??; + assert!(matches!( + response, + HostToClient::Response { id, .. } | HostToClient::InitialResponse { id, .. } + if id == execute_id + )); + } + Ok(()) +} + +#[tokio::test] +async fn websocket_bulk_lane_rejects_control_messages() -> Result<()> { + let host = HostHarness::start().await?; + let mut control = host.connect().await?; + let mut bulk = control.negotiate_dual(&host.websocket_url).await?; + let session_id = SessionId::new("wrong-lane-session")?; + control.open_session(session_id.clone()).await?; + + bulk.send(&ClientToHost::Request { + id: RequestId::new(/*value*/ 42), + request: HostRequest::Wait { + session_id, + request: WireWaitRequest { + cell_id: codex_code_mode_protocol::host::WireCellId::new("missing-cell"), + yield_time_ms: 10, + }, + }, + }) + .await?; + + let result = timeout(TEST_TIMEOUT, control.websocket.next()) + .await + .context("wrong-lane control message did not disconnect the paired sockets")?; + assert!( + matches!(result, None | Some(Ok(Message::Close(_))) | Some(Err(_))), + "wrong-lane message unexpectedly returned {result:?}" + ); + Ok(()) +} + +#[tokio::test] +async fn websocket_bulk_pairing_rejects_unknown_tokens() -> Result<()> { + let host = HostHarness::start().await?; + let unknown_token = Uuid::new_v4(); + let url = format!("{}/bulk/{unknown_token}", host.websocket_url); + let error = match connect_async(url).await { + Ok(_) => anyhow::bail!("unknown bulk pairing token should be rejected"), + Err(error) => error, + }; + let WebSocketError::Http(response) = error else { + anyhow::bail!("bulk pairing failed unexpectedly: {error}"); + }; + assert_eq!(response.status(), StatusCode::NOT_FOUND); + Ok(()) +} + #[tokio::test] async fn websocket_listener_accepts_frames_larger_than_default_websocket_limit() -> Result<()> { let host = HostHarness::start().await?; diff --git a/codex-rs/code-mode-protocol/src/host/host_tests.rs b/codex-rs/code-mode-protocol/src/host/host_tests.rs index dde38e87c9..240028a570 100644 --- a/codex-rs/code-mode-protocol/src/host/host_tests.rs +++ b/codex-rs/code-mode-protocol/src/host/host_tests.rs @@ -22,6 +22,7 @@ use super::ProtocolVersion; use super::RequestId; use super::SessionId; use super::SupportedProtocolVersions; +use super::TransportLane; use super::WireCellId; use super::WireContentItem; use super::WireExecuteRequest; @@ -72,6 +73,130 @@ where ); } +#[test] +fn dual_websocket_hello_preserves_the_pairing_token() { + assert_wire_round_trip( + HostToClient::HostHello( + HostHello::new( + ProtocolVersion::V1, + CapabilitySet::try_new([capability("dual-websocket-v1")]) + .expect("valid capabilities"), + ) + .with_bulk_connection_token("pairing-token".to_string()), + ), + json!({ + "type": "connection/ready", + "selectedVersion": 1, + "capabilities": ["dual-websocket-v1"], + "bulkConnectionToken": "pairing-token", + }), + ); +} + +#[test] +fn message_families_use_dedicated_transport_lanes() { + for (message, lane) in [ + ( + ClientToHost::CancelRequest { + id: request_id(/*value*/ 1), + }, + TransportLane::Control, + ), + ( + ClientToHost::DelegateResponse { + id: delegate_request_id(/*value*/ 1), + result: WireResult::Ok { + value: DelegateResponse::NotificationDelivered, + }, + }, + TransportLane::Control, + ), + ( + ClientToHost::DelegateResponse { + id: delegate_request_id(/*value*/ 2), + result: WireResult::Ok { + value: DelegateResponse::ToolResult { + result: json!({ "value": "tool result" }), + }, + }, + }, + TransportLane::Bulk, + ), + ( + ClientToHost::DelegateResponse { + id: delegate_request_id(/*value*/ 3), + result: WireResult::Err { + message: "delegate failed".to_string(), + }, + }, + TransportLane::Bulk, + ), + ] { + assert_eq!(message.transport_lane(), lane); + assert!(message.allows_transport_lane(lane)); + assert!(!message.allows_transport_lane(match lane { + TransportLane::Control => TransportLane::Bulk, + TransportLane::Bulk => TransportLane::Control, + })); + } + + for (message, lane) in [ + ( + HostToClient::Response { + id: request_id(/*value*/ 1), + result: WireResult::Err { + message: "x".repeat(128 * 1024), + }, + }, + TransportLane::Control, + ), + ( + HostToClient::DelegateRequest { + id: delegate_request_id(/*value*/ 1), + session_id: session_id(), + request: DelegateRequest::Notify { + call_id: "call-1".to_string(), + cell_id: cell_id("cell-1"), + text: "important".to_string(), + }, + }, + TransportLane::Control, + ), + ( + HostToClient::DelegateRequest { + id: delegate_request_id(/*value*/ 2), + session_id: session_id(), + request: DelegateRequest::InvokeTool { + invocation: WireNestedToolCall { + cell_id: cell_id("cell-1"), + runtime_tool_call_id: "runtime-call-1".to_string(), + tool_name: WireToolName { + name: "tool".to_string(), + namespace: None, + }, + tool_kind: WireToolKind::Function, + input: None, + }, + }, + }, + TransportLane::Bulk, + ), + ( + HostToClient::CancelDelegateRequest { + id: delegate_request_id(/*value*/ 1), + }, + TransportLane::Bulk, + ), + ] { + assert_eq!(message.transport_lane(), lane); + assert!(message.allows_transport_lane(lane)); + assert!(!message.allows_transport_lane(match lane { + TransportLane::Control => TransportLane::Bulk, + TransportLane::Bulk => TransportLane::Control, + })); + } +} + fn execute_request() -> WireExecuteRequest { WireExecuteRequest { tool_call_id: "call-1".to_string(), diff --git a/codex-rs/code-mode-protocol/src/host/message.rs b/codex-rs/code-mode-protocol/src/host/message.rs index 0e83c866e4..e6918ea4fd 100644 --- a/codex-rs/code-mode-protocol/src/host/message.rs +++ b/codex-rs/code-mode-protocol/src/host/message.rs @@ -12,6 +12,7 @@ use super::ProtocolVersion; use super::RequestId; use super::SessionId; use super::SupportedProtocolVersions; +use super::TransportLane; use super::WireCellId; use super::WireExecuteRequest; use super::WireNestedToolCall; @@ -105,6 +106,8 @@ impl std::error::Error for ClientHelloError {} pub struct HostHello { selected_version: ProtocolVersion, capabilities: CapabilitySet, + #[serde(default, skip_serializing_if = "Option::is_none")] + bulk_connection_token: Option, } impl HostHello { @@ -112,9 +115,15 @@ impl HostHello { Self { selected_version, capabilities, + bulk_connection_token: None, } } + pub fn with_bulk_connection_token(mut self, token: String) -> Self { + self.bulk_connection_token = Some(token); + self + } + pub fn selected_version(&self) -> ProtocolVersion { self.selected_version } @@ -122,6 +131,10 @@ impl HostHello { pub fn capabilities(&self) -> &CapabilitySet { &self.capabilities } + + pub fn bulk_connection_token(&self) -> Option<&str> { + self.bulk_connection_token.as_deref() + } } /// Messages sent from a client to the code-mode host. @@ -141,6 +154,30 @@ pub enum ClientToHost { }, } +impl ClientToHost { + /// Keeps notification acknowledgments with control traffic and tool results on the bulk lane. + pub fn transport_lane(&self) -> TransportLane { + match self { + Self::DelegateResponse { + result: + WireResult::Ok { + value: DelegateResponse::NotificationDelivered, + }, + .. + } + | Self::ClientHello(_) + | Self::Request { .. } + | Self::CancelRequest { .. } => TransportLane::Control, + Self::DelegateResponse { .. } => TransportLane::Bulk, + } + } + + /// Validates the message families accepted by each paired socket. + pub fn allows_transport_lane(&self, lane: TransportLane) -> bool { + self.transport_lane() == lane + } +} + /// Messages sent from the code-mode host to a client. #[derive(Debug, Deserialize, PartialEq, Serialize)] #[serde(deny_unknown_fields, tag = "type", rename_all_fields = "camelCase")] @@ -174,6 +211,33 @@ pub enum HostToClient { }, } +impl HostToClient { + /// Keeps notifications with control traffic and nested-tool callbacks on the bulk lane. + pub fn transport_lane(&self) -> TransportLane { + match self { + Self::DelegateRequest { + request: DelegateRequest::InvokeTool { .. }, + .. + } + | Self::CancelDelegateRequest { .. } => TransportLane::Bulk, + Self::DelegateRequest { + request: DelegateRequest::Notify { .. }, + .. + } + | Self::HostHello(_) + | Self::HandshakeRejected { .. } + | Self::Response { .. } + | Self::InitialResponse { .. } + | Self::CellClosed { .. } => TransportLane::Control, + } + } + + /// Rejects messages received on the wrong paired socket. + pub fn allows_transport_lane(&self, lane: TransportLane) -> bool { + self.transport_lane() == lane + } +} + #[derive(Clone, Debug, Deserialize, PartialEq, Serialize)] #[serde(deny_unknown_fields, tag = "method", rename_all_fields = "camelCase")] pub enum HostRequest { diff --git a/codex-rs/code-mode-protocol/src/host/mod.rs b/codex-rs/code-mode-protocol/src/host/mod.rs index 5c81b1d4a0..6169c61a17 100644 --- a/codex-rs/code-mode-protocol/src/host/mod.rs +++ b/codex-rs/code-mode-protocol/src/host/mod.rs @@ -1,9 +1,8 @@ //! Messages and framing for the code-mode host boundary. //! //! Protocol version 1 multiplexes session operations and delegate callbacks by -//! request ID over one ordered connection. It defines no optional capabilities -//! yet; capability names provide an extension point for later versions without -//! weakening the v1 decoder. +//! request ID over one ordered connection. WebSocket peers can negotiate a +//! separate bulk connection without changing the existing inner messages. mod codec; mod error; @@ -14,6 +13,16 @@ mod types; /// Maximum number of unresolved delegate callbacks allowed per host connection. pub const MAX_PENDING_DELEGATE_CALLS: usize = 1_024; +/// Optional second WebSocket carrying delegate callbacks and their responses. +pub const DUAL_WEBSOCKET_CAPABILITY: &str = "dual-websocket-v1"; + +/// Selects one socket of a negotiated dual-WebSocket connection. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum TransportLane { + Control, + Bulk, +} + pub use codec::EncodedFrame; pub use codec::FramedReader; pub use codec::FramedWriter; diff --git a/codex-rs/code-mode/src/remote_session/connection.rs b/codex-rs/code-mode/src/remote_session/connection.rs index 033d24ab57..30310f3e83 100644 --- a/codex-rs/code-mode/src/remote_session/connection.rs +++ b/codex-rs/code-mode/src/remote_session/connection.rs @@ -14,17 +14,21 @@ use codex_code_mode_protocol::ExecuteRequest; use codex_code_mode_protocol::StartedCell; use codex_code_mode_protocol::WaitOutcome; use codex_code_mode_protocol::WaitRequest; +use codex_code_mode_protocol::host::Capability; use codex_code_mode_protocol::host::CapabilitySet; use codex_code_mode_protocol::host::ClientHello; use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::DUAL_WEBSOCKET_CAPABILITY; use codex_code_mode_protocol::host::EncodedFrame; use codex_code_mode_protocol::host::FramedReader; use codex_code_mode_protocol::host::FramedWriter; use codex_code_mode_protocol::host::HostToClient; use codex_code_mode_protocol::host::MAX_FRAME_BYTES; +use codex_code_mode_protocol::host::MAX_PENDING_DELEGATE_CALLS; use codex_code_mode_protocol::host::ProtocolVersion; use codex_code_mode_protocol::host::RequestId; use codex_code_mode_protocol::host::SupportedProtocolVersions; +use codex_code_mode_protocol::host::TransportLane; use codex_http_client::HttpClientFactory; use codex_websocket_client::WebSocketConnector; use futures::StreamExt; @@ -36,6 +40,7 @@ use tokio::sync::mpsc; use tokio::sync::oneshot; use tokio::task::JoinHandle; use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::Uri; use tokio_tungstenite::tungstenite::protocol::WebSocketConfig; use tokio_util::sync::CancellationToken; use tracing::debug; @@ -68,6 +73,7 @@ pub(super) enum ConnectionError { host_program: PathBuf, error: io::Error, }, + BulkConnectionUnavailable(String), Other(String), } @@ -98,7 +104,9 @@ impl fmt::Display for ConnectionError { &host_program[suffix_start..] ) } - Self::Other(message) => formatter.write_str(message), + Self::BulkConnectionUnavailable(message) | Self::Other(message) => { + formatter.write_str(message) + } } } } @@ -132,6 +140,45 @@ enum ConnectionOwner { WebSocket, } +struct BulkConnectionOptions { + websocket_url: String, + http_client_factory: HttpClientFactory, +} + +impl BulkConnectionOptions { + async fn connect( + &self, + token: &str, + ) -> Result<(ConnectionReader, ConnectionWriter), ConnectionError> { + let control_uri = self.websocket_url.parse::().map_err(|error| { + ConnectionError::Other(format!( + "failed to build code-mode host bulk websocket URL: {error}" + )) + })?; + let bulk_path = format!("{}/bulk/{token}", control_uri.path().trim_end_matches('/')); + let bulk_path_and_query = match control_uri.query() { + Some(query) => format!("{bulk_path}?{query}"), + None => bulk_path, + }; + let mut bulk_uri_parts = control_uri.into_parts(); + bulk_uri_parts.path_and_query = Some(bulk_path_and_query.parse().map_err(|error| { + ConnectionError::Other(format!( + "failed to build code-mode host bulk websocket path: {error}" + )) + })?); + let bulk_url = Uri::from_parts(bulk_uri_parts) + .map_err(|error| { + ConnectionError::Other(format!( + "failed to build code-mode host bulk websocket URL: {error}" + )) + })? + .to_string(); + connect_websocket_transport(&bulk_url, &self.http_client_factory) + .await + .map_err(|error| ConnectionError::BulkConnectionUnavailable(error.to_string())) + } +} + impl CallerCancellation { fn new() -> Self { Self { @@ -202,6 +249,7 @@ impl Connection { ConnectionReader::Stdio(FramedReader::new(stdout)), ConnectionWriter::Stdio(FramedWriter::new(stdin)), ConnectionOwner::Process(Box::new(child)), + /*bulk_connection_options*/ None, ) .await } @@ -210,38 +258,17 @@ impl Connection { websocket_url: &str, http_client_factory: &HttpClientFactory, ) -> Result { - let request = websocket_url.into_client_request().map_err(|error| { - ConnectionError::Other(format!( - "failed to build code-mode host websocket request: {error}" - )) - })?; - let connector = WebSocketConnector::new(http_client_factory).map_err(|error| { - ConnectionError::Other(format!( - "failed to configure code-mode host websocket TLS: {error}" - )) - })?; - let websocket_config = WebSocketConfig::default() - .max_frame_size(Some(MAX_WEBSOCKET_FRAME_BYTES)) - .max_message_size(Some(MAX_WEBSOCKET_FRAME_BYTES)); - let (websocket, _) = tokio::time::timeout( - HOST_HANDSHAKE_TIMEOUT, - connector.connect(request, websocket_config), - ) - .await - .map_err(|_| { - ConnectionError::Other("timed out connecting to the code-mode host websocket".into()) - })? - .map_err(|error| { - ConnectionError::Other(format!( - "failed to connect to the code-mode host websocket: {error}" - )) - })?; - let (writer, reader) = websocket.split(); + let (reader, writer) = + connect_websocket_transport(websocket_url, http_client_factory).await?; Self::establish( - ConnectionReader::WebSocket(reader), - ConnectionWriter::WebSocket(writer), + reader, + writer, ConnectionOwner::WebSocket, + Some(BulkConnectionOptions { + websocket_url: websocket_url.to_string(), + http_client_factory: http_client_factory.clone(), + }), ) .await } @@ -250,13 +277,22 @@ impl Connection { mut reader: ConnectionReader, mut writer: ConnectionWriter, mut owner: ConnectionOwner, + bulk_connection_options: Option, ) -> Result { let handshake = async { + let dual_capability = + Capability::new(DUAL_WEBSOCKET_CAPABILITY).map_err(|error| error.to_string())?; + let optional_capabilities = if bulk_connection_options.is_some() { + CapabilitySet::try_new([dual_capability.clone()]) + .map_err(|error| error.to_string())? + } else { + CapabilitySet::empty() + }; let hello = ClientHello::new( SupportedProtocolVersions::try_new([ProtocolVersion::V1]) .map_err(|err| err.to_string())?, CapabilitySet::empty(), - CapabilitySet::empty(), + optional_capabilities, ) .map_err(|err| err.to_string())?; writer @@ -271,7 +307,20 @@ impl Connection { Some(HostToClient::HostHello(hello)) if hello.selected_version() == ProtocolVersion::V1 => { - Ok(()) + if hello.capabilities().contains(&dual_capability) { + hello + .bulk_connection_token() + .map(str::to_string) + .ok_or_else(|| { + "code-mode host advertised dual websockets without a pairing token" + .to_string() + }) + .map(Some) + } else if hello.bulk_connection_token().is_some() { + Err("code-mode host returned an unexpected bulk pairing token".to_string()) + } else { + Ok(None) + } } Some(HostToClient::HandshakeRejected { reason }) => { Err(format!("code-mode host rejected the handshake: {reason:?}")) @@ -292,51 +341,85 @@ impl Connection { )); } }; - if let Err(err) = handshake_result { - let _ = writer.close().await; - owner.close().await; - return Err(ConnectionError::Other(err)); - } + let bulk_token = match handshake_result { + Ok(token) => token, + Err(err) => { + let _ = writer.close().await; + owner.close().await; + return Err(ConnectionError::Other(err)); + } + }; + let (bulk_reader, bulk_writer) = if let Some(token) = bulk_token { + let Some(options) = bulk_connection_options else { + let _ = writer.close().await; + owner.close().await; + return Err(ConnectionError::Other( + "code-mode host negotiated an unsupported bulk websocket".to_string(), + )); + }; + match options.connect(&token).await { + Ok((reader, writer)) => (Some(reader), Some(writer)), + Err(error) => { + let _ = writer.close().await; + owner.close().await; + return Err(error); + } + } + } else { + (None, None) + }; let (command_tx, command_rx) = mpsc::channel(IPC_CHANNEL_CAPACITY); let (event_tx, event_rx) = mpsc::channel(IPC_CHANNEL_CAPACITY); - let (outgoing_tx, mut outgoing_rx) = mpsc::channel::(IPC_CHANNEL_CAPACITY); + let (outgoing_tx, outgoing_rx) = mpsc::channel::(IPC_CHANNEL_CAPACITY); + let (bulk_tx, bulk_rx) = if bulk_writer.is_some() { + let (sender, receiver) = mpsc::channel::(MAX_PENDING_DELEGATE_CALLS); + (Some(sender), Some(receiver)) + } else { + (None, None) + }; + let dual_websocket = bulk_writer.is_some(); let cancellation = CancellationToken::new(); let alive = Arc::new(AtomicBool::new(true)); let failure = Arc::new(std::sync::Mutex::new(None)); let writer_cancellation = cancellation.clone(); let writer_task = tokio::spawn(async move { - loop { - tokio::select! { - _ = writer_cancellation.cancelled() => { - return writer - .close() - .await - .map_err(|error| format!("failed to close code-mode host connection: {error}")); - } - frame = outgoing_rx.recv() => { - let Some(frame) = frame else { - return Err("code-mode host outgoing stream closed".to_string()); - }; - let result = tokio::select! { - _ = writer_cancellation.cancelled() => return Ok(()), - result = writer.write_frame(frame) => result, - }; - if let Err(err) = result { - return Err(format!("failed to write code-mode host message: {err}")); - } - } - } + if let (Some(bulk_writer), Some(bulk_rx)) = (bulk_writer, bulk_rx) { + tokio::try_join!( + drive_writer(writer, outgoing_rx, writer_cancellation.clone()), + drive_writer(bulk_writer, bulk_rx, writer_cancellation) + )?; + Ok(()) + } else { + drive_writer(writer, outgoing_rx, writer_cancellation).await } }); let reader_events = event_tx.clone(); let reader_cancellation = cancellation.clone(); - let reader_task = - tokio::spawn( - async move { drive_reader(reader, reader_events, reader_cancellation).await }, - ); + let reader_task = tokio::spawn(async move { + let lane = dual_websocket.then_some(TransportLane::Control); + if let Some(bulk_reader) = bulk_reader { + tokio::try_join!( + drive_reader( + reader, + reader_events.clone(), + reader_cancellation.clone(), + lane, + ), + drive_reader( + bulk_reader, + reader_events, + reader_cancellation, + Some(TransportLane::Bulk), + ) + )?; + Ok(()) + } else { + drive_reader(reader, reader_events, reader_cancellation, lane).await + } + }); let (driver, execute_claim_tx) = ConnectionDriver::new( command_rx, @@ -349,6 +432,10 @@ impl Connection { cancellation: cancellation.clone(), }, ); + let driver = match bulk_tx { + Some(sender) => driver.with_bulk_sender(sender), + None => driver, + }; let driver_task = tokio::spawn(driver.run()); tokio::spawn( ConnectionSupervisor { @@ -504,6 +591,73 @@ impl Connection { } } +async fn connect_websocket_transport( + websocket_url: &str, + http_client_factory: &HttpClientFactory, +) -> Result<(ConnectionReader, ConnectionWriter), ConnectionError> { + let request = websocket_url.into_client_request().map_err(|error| { + ConnectionError::Other(format!( + "failed to build code-mode host websocket request: {error}" + )) + })?; + let connector = WebSocketConnector::new(http_client_factory).map_err(|error| { + ConnectionError::Other(format!( + "failed to configure code-mode host websocket TLS: {error}" + )) + })?; + let websocket_config = WebSocketConfig::default() + .max_frame_size(Some(MAX_WEBSOCKET_FRAME_BYTES)) + .max_message_size(Some(MAX_WEBSOCKET_FRAME_BYTES)); + let (websocket, _) = tokio::time::timeout( + HOST_HANDSHAKE_TIMEOUT, + connector.connect(request, websocket_config), + ) + .await + .map_err(|_| { + ConnectionError::Other("timed out connecting to the code-mode host websocket".into()) + })? + .map_err(|error| { + ConnectionError::Other(format!( + "failed to connect to the code-mode host websocket: {error}" + )) + })?; + let (writer, reader) = websocket.split(); + Ok(( + ConnectionReader::WebSocket(reader), + ConnectionWriter::WebSocket(writer), + )) +} + +async fn drive_writer( + mut writer: ConnectionWriter, + mut outgoing: mpsc::Receiver, + cancellation: CancellationToken, +) -> Result<(), String> { + loop { + tokio::select! { + _ = cancellation.cancelled() => { + return writer + .close() + .await + .map_err(|error| format!("failed to close code-mode host connection: {error}")); + } + frame = outgoing.recv() => { + let Some(frame) = frame else { + return Err("code-mode host outgoing stream closed".to_string()); + }; + tokio::select! { + _ = cancellation.cancelled() => return Ok(()), + result = writer.write_frame(frame) => { + result.map_err(|error| { + format!("failed to write code-mode host message: {error}") + })?; + } + } + } + } + } +} + impl Drop for Connection { fn drop(&mut self) { mark_connection_dead( diff --git a/codex-rs/code-mode/src/remote_session/connection/driver.rs b/codex-rs/code-mode/src/remote_session/connection/driver.rs index 102ad47581..667a435805 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver.rs @@ -1,3 +1,4 @@ +use std::collections::VecDeque; use std::panic::AssertUnwindSafe; use std::sync::Arc; use std::sync::atomic::AtomicBool; @@ -6,7 +7,9 @@ use std::sync::atomic::Ordering; use codex_code_mode_protocol::CellId; use codex_code_mode_protocol::CodeModeSessionDelegate; use codex_code_mode_protocol::host::EncodedFrame; +use codex_code_mode_protocol::host::HostToClient; use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::TransportLane; use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; @@ -39,7 +42,9 @@ pub(super) struct ConnectionDriver { event_tx: mpsc::Sender, execute_claim_rx: mpsc::UnboundedReceiver, outgoing_tx: mpsc::Sender, + bulk_tx: Option>, requests: RequestTracker, + deferred_host_messages: VecDeque, sessions: SessionRegistry, delegates: DelegateRuntime, alive: Arc, @@ -64,7 +69,9 @@ impl ConnectionDriver { event_tx: event_tx.clone(), execute_claim_rx, outgoing_tx, + bulk_tx: None, requests: RequestTracker::new(), + deferred_host_messages: VecDeque::new(), sessions: SessionRegistry::new(), delegates: DelegateRuntime::new(event_tx), alive: lifecycle.alive, @@ -130,8 +137,17 @@ impl ConnectionDriver { } } - fn queue_frame(&mut self, frame: EncodedFrame) -> bool { - match self.outgoing_tx.try_send(frame) { + pub(super) fn with_bulk_sender(mut self, sender: mpsc::Sender) -> Self { + self.bulk_tx = Some(sender); + self + } + + fn queue_frame(&mut self, frame: EncodedFrame, lane: TransportLane) -> bool { + let sender = match lane { + TransportLane::Control => &self.outgoing_tx, + TransportLane::Bulk => self.bulk_tx.as_ref().unwrap_or(&self.outgoing_tx), + }; + match sender.try_send(frame) { Ok(()) => true, Err(mpsc::error::TrySendError::Full(_)) => { self.fail("code-mode host outgoing queue is full".to_string()); diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs b/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs index c9cba9c32d..4bfb713625 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs @@ -105,7 +105,8 @@ impl ConnectionDriver { }, &self.event_tx, ); - self.queue_frame(frame) + let lane = message.transport_lane(); + self.queue_frame(frame, lane) } fn execute( @@ -164,7 +165,8 @@ impl ConnectionDriver { }, &self.event_tx, ); - self.queue_frame(frame) + let lane = message.transport_lane(); + self.queue_frame(frame, lane) } fn wait( @@ -293,6 +295,7 @@ impl ConnectionDriver { }; self.requests .insert_pending(request_id, pending, &self.event_tx); - self.queue_frame(frame) + let lane = message.transport_lane(); + self.queue_frame(frame, lane) } } diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/delegate_runtime.rs b/codex-rs/code-mode/src/remote_session/connection/driver/delegate_runtime.rs index 94be7b1b64..fb01f8783e 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver/delegate_runtime.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver/delegate_runtime.rs @@ -261,10 +261,7 @@ impl ConnectionDriver { }; let target = match self.sessions.delegate_target(&session_id, wire_cell_id) { Ok(target) => target, - Err(err) => { - self.fail(err); - return false; - } + Err(err) => return self.send_delegate_response(id, Err(err)), }; match self.delegates.start(id, target, request) { Ok(()) => true, @@ -321,7 +318,8 @@ impl ConnectionDriver { } } }; - self.queue_frame(frame) + let lane = message.transport_lane(); + self.queue_frame(frame, lane) } pub(super) fn close_cell(&mut self, session_id: SessionId, cell_id: WireCellId) -> bool { diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/request_tracker.rs b/codex-rs/code-mode/src/remote_session/connection/driver/request_tracker.rs index 0e7fc0881d..d39981acd0 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver/request_tracker.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver/request_tracker.rs @@ -52,6 +52,15 @@ impl RequestTracker { }) } + pub(super) fn has_pending_execute_for_session(&self, session_id: &SessionId) -> bool { + self.pending.values().any(|request| { + matches!( + request, + PendingRequest::Execute { session, .. } if session.id == *session_id + ) + }) + } + pub(super) fn allocate_id(&mut self) -> Result { let id = self.next_request_id; self.next_request_id = self diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/responses.rs b/codex-rs/code-mode/src/remote_session/connection/driver/responses.rs index 19c78ec000..17f697362e 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver/responses.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver/responses.rs @@ -56,6 +56,62 @@ impl ConnectionDriver { } pub(super) fn handle_host_message(&mut self, message: HostToClient) -> bool { + if self.should_defer_host_message(&message) { + if self.deferred_host_messages.len() + >= codex_code_mode_protocol::host::MAX_PENDING_DELEGATE_CALLS + { + self.fail( + "code-mode host exceeded deferred cross-socket message limit".to_string(), + ); + return false; + } + self.deferred_host_messages.push_back(message); + return true; + } + if !self.dispatch_host_message(message) { + return false; + } + for _ in 0..self.deferred_host_messages.len() { + let Some(message) = self.deferred_host_messages.pop_front() else { + break; + }; + if self.should_defer_host_message(&message) { + self.deferred_host_messages.push_back(message); + } else if !self.dispatch_host_message(message) { + return false; + } + } + true + } + + fn should_defer_host_message(&self, message: &HostToClient) -> bool { + match message { + HostToClient::DelegateRequest { + session_id, + request, + .. + } => { + let cell_id = match request { + codex_code_mode_protocol::host::DelegateRequest::InvokeTool { invocation } => { + &invocation.cell_id + } + codex_code_mode_protocol::host::DelegateRequest::Notify { cell_id, .. } => { + cell_id + } + }; + !self.sessions.contains_cell(session_id, cell_id) + && self.requests.has_pending_execute_for_session(session_id) + } + HostToClient::Response { .. } + | HostToClient::InitialResponse { .. } + | HostToClient::CellClosed { .. } + | HostToClient::CancelDelegateRequest { .. } + | HostToClient::HostHello(_) + | HostToClient::HandshakeRejected { .. } => false, + } + } + + fn dispatch_host_message(&mut self, message: HostToClient) -> bool { match message { HostToClient::Response { id, result } => { self.complete_request(id, result.into_result()) @@ -69,6 +125,12 @@ impl ConnectionDriver { request, } => self.start_delegate(id, session_id, request), HostToClient::CancelDelegateRequest { id } => { + self.deferred_host_messages.retain(|message| { + !matches!( + message, + HostToClient::DelegateRequest { id: deferred_id, .. } if *deferred_id == id + ) + }); self.delegates.cancel(id); true } @@ -302,7 +364,8 @@ impl ConnectionDriver { } fn send_cancel_request(&mut self, id: RequestId) -> bool { - let frame = match EncodedFrame::encode(&ClientToHost::CancelRequest { id }) { + let message = ClientToHost::CancelRequest { id }; + let frame = match EncodedFrame::encode(&message) { Ok(frame) => frame, Err(err) => { self.fail(format!( @@ -311,7 +374,8 @@ impl ConnectionDriver { return false; } }; - self.queue_frame(frame) + let lane = message.transport_lane(); + self.queue_frame(frame, lane) } fn shutdown_abandoned_session(&mut self, session: RemoteSession) -> bool { diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/session_registry.rs b/codex-rs/code-mode/src/remote_session/connection/driver/session_registry.rs index a3d1c82ae6..8668baa27e 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver/session_registry.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver/session_registry.rs @@ -61,6 +61,12 @@ impl SessionRegistry { self.records.contains_key(session_id) } + pub(super) fn contains_cell(&self, session_id: &SessionId, cell_id: &WireCellId) -> bool { + self.records + .get(session_id) + .is_some_and(|session| session.cells.contains_key(cell_id)) + } + pub(super) fn insert_ready( &mut self, session: RemoteSession, diff --git a/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs b/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs index 9f0f3f7625..bb33002ec2 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs @@ -15,6 +15,7 @@ use codex_code_mode_protocol::WaitRequest; use codex_code_mode_protocol::host::ClientToHost; use codex_code_mode_protocol::host::DelegateRequest; use codex_code_mode_protocol::host::DelegateRequestId; +use codex_code_mode_protocol::host::DelegateResponse; use codex_code_mode_protocol::host::EncodedFrame; use codex_code_mode_protocol::host::HostResponse; use codex_code_mode_protocol::host::HostToClient; @@ -193,6 +194,11 @@ struct RecordingDelegate { struct PanickingDelegate; +struct LargeResultBurstDelegate { + started: AtomicUsize, + release: CancellationToken, +} + #[derive(Debug, Eq, PartialEq)] enum HeldDelegateEvent { Started, @@ -282,6 +288,35 @@ impl CodeModeSessionDelegate for PanickingDelegate { fn cell_closed(&self, _cell_id: &CellId) {} } +impl CodeModeSessionDelegate for LargeResultBurstDelegate { + fn invoke_tool<'a>( + &'a self, + _invocation: CodeModeNestedToolCall, + cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + self.started.fetch_add(1, Ordering::Release); + let release = self.release.clone(); + Box::pin(async move { + tokio::select! { + _ = cancellation_token.cancelled() => Err("cancelled".to_string()), + _ = release.cancelled() => Ok("x".repeat(256 * 1024).into()), + } + }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, _cell_id: &CellId) {} +} + impl CodeModeSessionDelegate for RecordingDelegate { fn invoke_tool<'a>( &'a self, @@ -330,6 +365,143 @@ async fn next_held_delegate_event( .expect("delegate event stream") } +#[tokio::test] +async fn deferred_delegates_follow_cell_readiness_and_cancellation() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + + let (first_response_tx, first_response_rx) = oneshot::channel(); + let (second_response_tx, second_response_rx) = oneshot::channel(); + for (tool_call_id, response_tx) in [ + ("first-cell", first_response_tx), + ("second-cell", second_response_tx), + ] { + harness + .command_tx + .send(DriverCommand::Execute { + session: session.clone(), + request: ExecuteRequest { + tool_call_id: tool_call_id.to_string(), + enabled_tools: Vec::new(), + source: "text('done')".to_string(), + yield_time_ms: Some(/*yield_time_ms*/ 1), + max_output_tokens: None, + }, + caller_cancellation: CancellationToken::new(), + response_tx, + }) + .await + .expect("execute command"); + harness.outgoing_rx.recv().await.expect("execute frame"); + } + + let first_cell_id = CellId::new("first-cell".to_string()); + let second_cell_id = CellId::new("second-cell".to_string()); + let first_delegate_id = DelegateRequestId::new(/*value*/ 7); + let second_delegate_id = DelegateRequestId::new(/*value*/ 8); + let cancelled_delegate_id = DelegateRequestId::new(/*value*/ 9); + for (delegate_id, cell_id) in [ + (first_delegate_id, first_cell_id.clone()), + (second_delegate_id, second_cell_id.clone()), + (cancelled_delegate_id, first_cell_id.clone()), + ] { + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id: delegate_id, + session_id: session.id.clone(), + request: DelegateRequest::Notify { + call_id: format!("notify-{}", cell_id.as_str()), + cell_id: (&cell_id).into(), + text: "hello".to_string(), + }, + })) + .await + .expect("early delegate request"); + } + + harness + .event_tx + .send(DriverEvent::HostMessage( + HostToClient::CancelDelegateRequest { + id: cancelled_delegate_id, + }, + )) + .await + .expect("deferred delegate cancellation"); + + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Ok { + value: HostResponse::ExecutionStarted { + cell_id: (&second_cell_id).into(), + }, + }, + })) + .await + .expect("second execution-started response"); + let _second_started = second_response_rx + .await + .expect("second execute response") + .expect("second started cell"); + let second_response = tokio::time::timeout(Duration::from_secs(1), harness.outgoing_rx.recv()) + .await + .expect("second delegate response timeout") + .expect("second delegate response frame"); + assert_eq!( + EncodedFrame::decode_framed::(&second_response.into_framed_bytes()) + .expect("decode second delegate response"), + ClientToHost::DelegateResponse { + id: second_delegate_id, + result: WireResult::Ok { + value: codex_code_mode_protocol::host::DelegateResponse::NotificationDelivered, + }, + } + ); + assert_eq!(delegate.notifications.load(Ordering::Relaxed), 1); + + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: HostResponse::ExecutionStarted { + cell_id: (&first_cell_id).into(), + }, + }, + })) + .await + .expect("first execution-started response"); + let _first_started = first_response_rx + .await + .expect("first execute response") + .expect("first started cell"); + let first_response = tokio::time::timeout(Duration::from_secs(1), harness.outgoing_rx.recv()) + .await + .expect("first delegate response timeout") + .expect("first delegate response frame"); + assert_eq!( + EncodedFrame::decode_framed::(&first_response.into_framed_bytes()) + .expect("decode first delegate response"), + ClientToHost::DelegateResponse { + id: first_delegate_id, + result: WireResult::Ok { + value: codex_code_mode_protocol::host::DelegateResponse::NotificationDelivered, + }, + } + ); + assert_eq!(delegate.notifications.load(Ordering::Relaxed), 2); + assert!(matches!( + harness.outgoing_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + assert!(harness.alive.load(Ordering::Relaxed)); +} + #[tokio::test] async fn dropped_open_waiter_shuts_down_committed_session() { let mut harness = DriverHarness::start(); @@ -465,6 +637,95 @@ async fn delegate_cancel_is_best_effort_and_sends_no_late_response() { assert_eq!(delegate.notifications.load(Ordering::Relaxed), 0); } +#[tokio::test] +async fn concurrent_large_delegate_results_do_not_disconnect_a_backpressured_bulk_lane() { + const CONCURRENT_RESULTS: usize = 129; + + let (command_tx, command_rx) = mpsc::channel(/*max_capacity*/ 16); + let (event_tx, event_rx) = mpsc::channel(/*max_capacity*/ 16); + let (outgoing_tx, outgoing_rx) = mpsc::channel(/*max_capacity*/ 16); + let (bulk_tx, mut bulk_rx) = mpsc::channel(MAX_PENDING_DELEGATE_CALLS); + let cancellation = CancellationToken::new(); + let alive = Arc::new(AtomicBool::new(true)); + let (driver, execute_claim_tx) = ConnectionDriver::new( + command_rx, + event_rx, + event_tx.clone(), + outgoing_tx, + DriverLifecycle { + alive: Arc::clone(&alive), + failure: Arc::new(StdMutex::new(None)), + cancellation: cancellation.clone(), + }, + ); + let driver_task = tokio::spawn(driver.with_bulk_sender(bulk_tx).run()); + let mut harness = DriverHarness { + command_tx, + event_tx, + execute_claim_tx, + outgoing_rx, + cancellation, + alive, + driver_task, + }; + let session = remote_session(); + let delegate = Arc::new(LargeResultBurstDelegate { + started: AtomicUsize::new(0), + release: CancellationToken::new(), + }); + harness.open(session.clone(), delegate.clone()).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + + for value in 1..=CONCURRENT_RESULTS { + harness + .start_tool_delegate(&session, DelegateRequestId::new(value as i64)) + .await; + } + tokio::time::timeout(Duration::from_secs(10), async { + while delegate.started.load(Ordering::Acquire) < CONCURRENT_RESULTS { + tokio::task::yield_now().await; + } + }) + .await + .expect("concurrent delegate calls should all start"); + + delegate.release.cancel(); + tokio::time::timeout(Duration::from_secs(10), async { + while bulk_rx.len() < CONCURRENT_RESULTS { + assert!( + harness.alive.load(Ordering::Acquire), + "bulk queue disconnected before accepting all concurrent tool results" + ); + tokio::task::yield_now().await; + } + }) + .await + .expect("concurrent large results should queue behind the blocked bulk writer"); + + let _unrelated = harness.start_cell(session, /*request_id*/ 3, "2").await; + assert!(harness.alive.load(Ordering::Acquire)); + + for _ in 0..CONCURRENT_RESULTS { + let frame = bulk_rx.recv().await.expect("queued bulk delegate result"); + let message = EncodedFrame::decode_framed::(&frame.into_framed_bytes()) + .expect("decode queued delegate result"); + let ClientToHost::DelegateResponse { + result: + WireResult::Ok { + value: DelegateResponse::ToolResult { result }, + }, + .. + } = message + else { + panic!("expected a successful large delegate result"); + }; + assert_eq!(result.as_str().map(str::len), Some(256 * 1024)); + } + assert!(harness.alive.load(Ordering::Acquire)); +} + #[tokio::test] async fn delegate_limit_returns_an_error_without_disconnecting() { let mut harness = DriverHarness::start(); @@ -796,16 +1057,17 @@ async fn delegate_task_panic_becomes_tool_error_without_killing_connection() { } #[tokio::test] -async fn delegate_for_unknown_cell_fails_connection_without_invocation() { +async fn delegate_for_unknown_cell_returns_error_without_invocation() { let mut harness = DriverHarness::start(); let session = remote_session(); let delegate = Arc::new(RecordingDelegate::default()); harness.open(session.clone(), delegate.clone()).await; + let id = DelegateRequestId::new(/*value*/ 7); harness .event_tx .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { - id: DelegateRequestId::new(/*value*/ 7), + id, session_id: session.id, request: DelegateRequest::InvokeTool { invocation: WireNestedToolCall { @@ -819,14 +1081,28 @@ async fn delegate_for_unknown_cell_fails_connection_without_invocation() { })) .await .expect("delegate request"); - tokio::task::yield_now().await; + let response = tokio::time::timeout(Duration::from_secs(1), harness.outgoing_rx.recv()) + .await + .expect("delegate response timeout") + .expect("delegate response frame"); - assert!(!harness.alive.load(Ordering::Acquire)); + assert_eq!( + EncodedFrame::decode_framed::(&response.into_framed_bytes()) + .expect("decode delegate response"), + ClientToHost::DelegateResponse { + id, + result: WireResult::Err { + message: "code-mode host delegated for unknown cell missing in session session-1" + .to_string(), + }, + } + ); + assert!(harness.alive.load(Ordering::Acquire)); assert_eq!(delegate.invocations.load(Ordering::Relaxed), 0); } #[tokio::test] -async fn delegate_after_cell_close_fails_connection_without_invocation() { +async fn delegate_after_cell_close_returns_error_without_invocation() { let mut harness = DriverHarness::start(); let session = remote_session(); let delegate = Arc::new(RecordingDelegate::default()); @@ -842,10 +1118,11 @@ async fn delegate_after_cell_close_fails_connection_without_invocation() { })) .await .expect("cell close"); + let id = DelegateRequestId::new(/*value*/ 7); harness .event_tx .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { - id: DelegateRequestId::new(/*value*/ 7), + id, session_id: session.id, request: DelegateRequest::Notify { call_id: "notify-1".to_string(), @@ -855,10 +1132,24 @@ async fn delegate_after_cell_close_fails_connection_without_invocation() { })) .await .expect("delegate request"); - tokio::task::yield_now().await; + let response = tokio::time::timeout(Duration::from_secs(1), harness.outgoing_rx.recv()) + .await + .expect("delegate response timeout") + .expect("delegate response frame"); - assert!(!harness.alive.load(Ordering::Acquire)); - assert_eq!(delegate.invocations.load(Ordering::Relaxed), 0); + assert_eq!( + EncodedFrame::decode_framed::(&response.into_framed_bytes()) + .expect("decode delegate response"), + ClientToHost::DelegateResponse { + id, + result: WireResult::Err { + message: "code-mode host delegated for unknown cell 1 in session session-1" + .to_string(), + }, + } + ); + assert!(harness.alive.load(Ordering::Acquire)); + assert_eq!(delegate.notifications.load(Ordering::Relaxed), 0); } #[tokio::test] diff --git a/codex-rs/code-mode/src/remote_session/connection/reader.rs b/codex-rs/code-mode/src/remote_session/connection/reader.rs index 8fb7cd8b29..8bb218e67c 100644 --- a/codex-rs/code-mode/src/remote_session/connection/reader.rs +++ b/codex-rs/code-mode/src/remote_session/connection/reader.rs @@ -1,3 +1,4 @@ +use codex_code_mode_protocol::host::TransportLane; use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; @@ -8,6 +9,7 @@ pub(super) async fn drive_reader( mut reader: ConnectionReader, events: mpsc::Sender, cancellation: CancellationToken, + lane: Option, ) -> Result<(), String> { loop { let message = tokio::select! { @@ -19,6 +21,11 @@ pub(super) async fn drive_reader( Ok(None) => return Err("code-mode host closed its stdout".to_string()), Err(err) => return Err(format!("failed to read code-mode host message: {err}")), }; + if let Some(lane) = lane + && !message.allows_transport_lane(lane) + { + return Err("code-mode host sent a message on the wrong websocket lane".to_string()); + } events .send(DriverEvent::HostMessage(message)) .await diff --git a/codex-rs/code-mode/src/remote_session_tests.rs b/codex-rs/code-mode/src/remote_session_tests.rs index 0429209157..32cd548003 100644 --- a/codex-rs/code-mode/src/remote_session_tests.rs +++ b/codex-rs/code-mode/src/remote_session_tests.rs @@ -7,8 +7,10 @@ use codex_code_mode_protocol::CodeModeSessionProvider; use codex_code_mode_protocol::ExecuteRequest; use codex_code_mode_protocol::FunctionCallOutputContentItem; use codex_code_mode_protocol::RuntimeResponse; +use codex_code_mode_protocol::host::Capability; use codex_code_mode_protocol::host::CapabilitySet; use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::DUAL_WEBSOCKET_CAPABILITY; use codex_code_mode_protocol::host::EncodedFrame; use codex_code_mode_protocol::host::HostHello; use codex_code_mode_protocol::host::HostRequest; @@ -24,6 +26,8 @@ use codex_http_client::OutboundProxyPolicy; use futures::SinkExt; use futures::StreamExt; use pretty_assertions::assert_eq; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; use tokio::net::TcpListener; use tokio::time::timeout; use tokio_tungstenite::accept_async; @@ -224,6 +228,91 @@ async fn websocket_provider_executes_over_shared_connector() { .expect("websocket test host task should succeed"); } +#[tokio::test] +async fn websocket_provider_fails_when_a_negotiated_bulk_connection_is_unavailable() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("websocket test listener should bind"); + let websocket_url = format!( + "ws://{}/?access_token=shared-token", + listener + .local_addr() + .expect("websocket test listener should have an address") + ); + let server = tokio::spawn(async move { + let (stream, _) = listener + .accept() + .await + .expect("websocket test host should accept the first connection"); + let mut control = accept_async(stream) + .await + .expect("first control websocket should connect"); + let frame = control + .next() + .await + .expect("first client hello") + .expect("first websocket frame") + .into_data(); + let ClientToHost::ClientHello(hello) = + EncodedFrame::decode_framed(&frame).expect("decode first client hello") + else { + panic!("expected first client hello"); + }; + let capability = + Capability::new(DUAL_WEBSOCKET_CAPABILITY).expect("dual websocket capability"); + assert!(hello.optional_capabilities().contains(&capability)); + let hello = HostToClient::HostHello( + HostHello::new( + ProtocolVersion::V1, + CapabilitySet::try_new([capability.clone()]).expect("host capabilities"), + ) + .with_bulk_connection_token("fallback-token".to_string()), + ); + let frame = EncodedFrame::encode(&hello).expect("encode dual host hello"); + control + .send(Message::Binary(frame.into_framed_bytes().into())) + .await + .expect("send dual host hello"); + + let (mut bulk, _) = listener + .accept() + .await + .expect("websocket test host should accept the bulk connection"); + let mut request = [0_u8; 1024]; + let request_len = bulk + .read(&mut request) + .await + .expect("read bulk websocket handshake"); + let request = std::str::from_utf8(&request[..request_len]).expect("bulk HTTP request"); + assert!( + request.starts_with("GET /bulk/fallback-token?access_token=shared-token HTTP/1.1\r\n") + ); + bulk.write_all(b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + .await + .expect("reject unavailable bulk websocket"); + drop(bulk); + drop(control); + }); + + let provider = WebSocketCodeModeSessionProvider::new(websocket_url); + let error = match provider + .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .await + { + Ok(_) => panic!("provider should reject an unavailable negotiated bulk websocket"), + Err(error) => error, + }; + assert!( + error.contains("404"), + "unexpected negotiated bulk websocket error: {error}" + ); + drop(provider); + timeout(Duration::from_secs(5), server) + .await + .expect("negotiated bulk websocket test host should disconnect promptly") + .expect("negotiated bulk websocket test host task should succeed"); +} + #[tokio::test] async fn provider_returns_missing_host_error() { let provider = ProcessOwnedCodeModeSessionProvider::with_host_program(