diff --git a/codex-rs/code-mode-host/Cargo.toml b/codex-rs/code-mode-host/Cargo.toml index 71c1237d20..685397aadf 100644 --- a/codex-rs/code-mode-host/Cargo.toml +++ b/codex-rs/code-mode-host/Cargo.toml @@ -29,7 +29,7 @@ serde_json = { workspace = true } tokio = { workspace = true, features = ["io-std", "io-util", "macros", "net", "process", "rt", "sync", "time"] } tokio-stream = { workspace = true } tokio-util = { workspace = true, features = ["rt"] } -tonic = { workspace = true } +tonic = { workspace = true, features = ["router", "transport"] } tracing = { workspace = true } tracing-subscriber = { workspace = true } uuid = { workspace = true, features = ["v4"] } diff --git a/codex-rs/code-mode-host/src/grpc_transport.rs b/codex-rs/code-mode-host/src/grpc_transport.rs new file mode 100644 index 0000000000..665e8de55f --- /dev/null +++ b/codex-rs/code-mode-host/src/grpc_transport.rs @@ -0,0 +1,43 @@ +use std::io; +use std::io::Write as _; +use std::net::SocketAddr; + +use anyhow::Context; +use anyhow::Result; +use codex_code_mode_protocol::grpc::code_mode_host_server::CodeModeHostServer; +use codex_code_mode_protocol::host::MAX_FRAME_BYTES; +use tokio::net::TcpListener; +use tonic::transport::Server; +use tonic::transport::server::TcpIncoming; +use tracing::info; + +use crate::GrpcCodeModeHost; + +pub(super) async fn run_tcp_listener(bind_address: SocketAddr) -> Result<()> { + let listener = bind_tcp_listener(bind_address).await?; + let local_address = listener + .local_addr() + .context("failed to read code-mode gRPC listen address")?; + info!("codex-code-mode-host listening on http://{local_address}"); + println!("http://{local_address}"); + io::stdout() + .flush() + .context("failed to publish code-mode gRPC listen address")?; + + Server::builder() + .add_service( + CodeModeHostServer::new(GrpcCodeModeHost::new()) + .max_decoding_message_size(MAX_FRAME_BYTES) + .max_encoding_message_size(MAX_FRAME_BYTES), + ) + .serve_with_incoming(listener) + .await + .context("code-mode gRPC TCP listener failed") +} + +pub(super) async fn bind_tcp_listener(bind_address: SocketAddr) -> Result { + let listener = TcpListener::bind(bind_address) + .await + .with_context(|| format!("failed to bind code-mode gRPC host to {bind_address}"))?; + Ok(TcpIncoming::from(listener).with_nodelay(Some(true))) +} diff --git a/codex-rs/code-mode-host/src/lib.rs b/codex-rs/code-mode-host/src/lib.rs index 05b119d142..34e13af94c 100644 --- a/codex-rs/code-mode-host/src/lib.rs +++ b/codex-rs/code-mode-host/src/lib.rs @@ -50,6 +50,7 @@ pub use self::transport::DEFAULT_LISTEN_URL; mod delegate; mod grpc; +mod grpc_transport; mod peer; mod transport; diff --git a/codex-rs/code-mode-host/src/main.rs b/codex-rs/code-mode-host/src/main.rs index 50ec048288..c7da2080c3 100644 --- a/codex-rs/code-mode-host/src/main.rs +++ b/codex-rs/code-mode-host/src/main.rs @@ -2,7 +2,7 @@ use clap::Parser; #[derive(Debug, Parser)] struct Cli { - /// Transport endpoint: `stdio`, `stdio://`, or `ws://IP:PORT`. + /// Transport endpoint: `stdio`, `stdio://`, `ws://IP:PORT`, or `grpc://IP:PORT`. #[arg( long, value_name = "URL", diff --git a/codex-rs/code-mode-host/src/transport.rs b/codex-rs/code-mode-host/src/transport.rs index d4dddd081a..3817f3b605 100644 --- a/codex-rs/code-mode-host/src/transport.rs +++ b/codex-rs/code-mode-host/src/transport.rs @@ -49,6 +49,7 @@ use uuid::Uuid; use crate::HostLimits; use crate::MAX_IN_FLIGHT_REQUESTS; +use crate::grpc_transport; /// The default transport retains the standalone host's original stdio behavior. pub const DEFAULT_LISTEN_URL: &str = "stdio"; @@ -63,6 +64,7 @@ type BoxedWriter = Box; enum ListenTransport { Stdio, WebSocket(SocketAddr), + Grpc(SocketAddr), } pub(crate) enum ConnectionReader { @@ -204,6 +206,7 @@ pub(crate) async fn run_transport(listen_url: &str) -> Result<()> { match parse_listen_url(listen_url)? { ListenTransport::Stdio => crate::run_stdio().await, ListenTransport::WebSocket(bind_address) => run_websocket_listener(bind_address).await, + ListenTransport::Grpc(bind_address) => grpc_transport::run_tcp_listener(bind_address).await, } } @@ -221,8 +224,17 @@ fn parse_listen_url(listen_url: &str) -> Result { }); } + if let Some(socket_addr) = listen_url.strip_prefix("grpc://") { + return socket_addr + .parse::() + .map(ListenTransport::Grpc) + .with_context(|| { + format!("invalid gRPC --listen URL `{listen_url}`; expected `grpc://IP:PORT`") + }); + } + anyhow::bail!( - "unsupported --listen URL `{listen_url}`; expected `ws://IP:PORT`, `stdio`, or `stdio://`" + "unsupported --listen URL `{listen_url}`; expected `ws://IP:PORT`, `grpc://IP:PORT`, `stdio`, or `stdio://`" ); } diff --git a/codex-rs/code-mode-host/src/transport_tests.rs b/codex-rs/code-mode-host/src/transport_tests.rs index 39a3bc5084..e1982e1864 100644 --- a/codex-rs/code-mode-host/src/transport_tests.rs +++ b/codex-rs/code-mode-host/src/transport_tests.rs @@ -1,6 +1,7 @@ use std::net::SocketAddr; use axum::serve::Listener; +use futures::StreamExt; use pretty_assertions::assert_eq; use tokio::net::TcpStream; @@ -9,6 +10,7 @@ use super::ListenTransport; use super::MAX_PENDING_BULK_CONNECTIONS; use super::bind_websocket_listener; use super::parse_listen_url; +use crate::grpc_transport::bind_tcp_listener; #[tokio::test] async fn websocket_listener_disables_nagle() { @@ -33,6 +35,31 @@ async fn websocket_listener_disables_nagle() { assert!(stream); } +#[tokio::test] +async fn grpc_listener_disables_nagle() { + let bind_address = "127.0.0.1:0" + .parse() + .expect("gRPC test listener should have a valid bind address"); + let mut listener = bind_tcp_listener(bind_address) + .await + .expect("gRPC test listener should bind"); + let local_addr = listener + .local_addr() + .expect("gRPC test listener should have a local address"); + let _client = TcpStream::connect(local_addr) + .await + .expect("gRPC test client should connect"); + let nodelay = listener + .next() + .await + .expect("gRPC test listener should accept a connection") + .expect("gRPC test listener should return a valid socket") + .nodelay() + .expect("accepted gRPC socket should expose TCP_NODELAY"); + + assert!(nodelay); +} + #[test] fn bulk_connection_registration_cleans_up_when_dropped() { let registry = BulkConnectionRegistry::default(); diff --git a/codex-rs/code-mode-host/tests/grpc_tcp.rs b/codex-rs/code-mode-host/tests/grpc_tcp.rs new file mode 100644 index 0000000000..5532b2b2a9 --- /dev/null +++ b/codex-rs/code-mode-host/tests/grpc_tcp.rs @@ -0,0 +1,58 @@ +use std::process::Stdio; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use codex_code_mode_protocol::grpc; +use codex_code_mode_protocol::grpc::code_mode_host_client::CodeModeHostClient; +use tokio::io::AsyncBufReadExt; +use tokio::io::BufReader; +use tokio::process::Command; +use tokio::time::timeout; +use tonic::transport::Endpoint; + +const TEST_TIMEOUT: Duration = Duration::from_secs(/*secs*/ 10); + +#[tokio::test] +async fn tcp_listener_opens_a_grpc_session() -> Result<()> { + let mut host = Command::new(codex_utils_cargo_bin::cargo_bin("codex-code-mode-host")?) + .args(["--listen", "grpc://127.0.0.1:0"]) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .kill_on_drop(/*kill_on_drop*/ true) + .spawn() + .context("failed to start gRPC code-mode host")?; + let stdout = host + .stdout + .take() + .context("gRPC code-mode host stdout is unavailable")?; + let mut stdout = BufReader::new(stdout); + let mut endpoint = String::new(); + timeout(TEST_TIMEOUT, stdout.read_line(&mut endpoint)) + .await + .context("gRPC code-mode host did not publish its endpoint")??; + let endpoint = Endpoint::from_shared(endpoint.trim().to_string()) + .context("gRPC code-mode host published an invalid endpoint")? + .connect_timeout(TEST_TIMEOUT) + .timeout(TEST_TIMEOUT); + let mut client = CodeModeHostClient::connect(endpoint) + .await + .context("failed to connect to gRPC code-mode host")?; + let mut events = client + .open_session(grpc::OpenSessionRequest { + cell_execution_limits: None, + }) + .await + .context("failed to open gRPC code-mode session")? + .into_inner(); + let event = timeout(TEST_TIMEOUT, events.message()) + .await + .context("timed out waiting for gRPC code-mode session event")? + .context("failed to read gRPC code-mode session event")? + .context("gRPC code-mode session ended before opening")?; + assert!(matches!( + event.event, + Some(grpc::session_event::Event::Opened(opened)) if !opened.session_id.is_empty() + )); + Ok(()) +}