From bde723ae7dedc87754228a2476a7a4cfbf05b4b8 Mon Sep 17 00:00:00 2001 From: Channing Conger Date: Wed, 12 Aug 2026 21:36:05 +0000 Subject: [PATCH] Reconnect gRPC code-mode sessions after host restarts (#38257) ## What changed - Reopen a cached code-mode session when its gRPC host stops, while serializing concurrent reconnection attempts and coordinating shutdown. - Scope cell IDs to the new host generation so callbacks remain consistent and stale `wait` or `terminate` requests are rejected. - Accept both `unix://` and `unix:` endpoints for gRPC hosts on Unix systems. ## Testing - Cover host restart recovery, concurrent execution after reconnection, generation-aware callbacks and cell operations, stale cell rejection, and Unix socket execution. GitOrigin-RevId: 548e168fdcef7f7d54bd32262e614886bf7bdd32 --- codex-rs/code-mode-host/tests/grpc.rs | 227 +++++++++++++++++ .../code-mode/src/grpc_session/generation.rs | 124 +++++++++ .../src/grpc_session/generation_tests.rs | 230 +++++++++++++++++ codex-rs/code-mode/src/grpc_session/mod.rs | 18 +- .../code-mode/src/grpc_session/reconnect.rs | 237 ++++++++++++++++++ .../code-mode/src/grpc_session/transport.rs | 10 + 6 files changed, 841 insertions(+), 5 deletions(-) create mode 100644 codex-rs/code-mode/src/grpc_session/generation.rs create mode 100644 codex-rs/code-mode/src/grpc_session/generation_tests.rs create mode 100644 codex-rs/code-mode/src/grpc_session/reconnect.rs diff --git a/codex-rs/code-mode-host/tests/grpc.rs b/codex-rs/code-mode-host/tests/grpc.rs index 5cdc69939a..2d30be4a2c 100644 --- a/codex-rs/code-mode-host/tests/grpc.rs +++ b/codex-rs/code-mode-host/tests/grpc.rs @@ -21,16 +21,26 @@ use codex_code_mode::ToolDefinition; use codex_code_mode::ToolInvocationFuture; use codex_code_mode::WaitOutcome; use codex_code_mode::WaitRequest; +#[cfg(unix)] +use codex_code_mode_host::GrpcCodeModeHost; use codex_code_mode_protocol::grpc; use codex_code_mode_protocol::grpc::code_mode_host_client::CodeModeHostClient; +#[cfg(unix)] +use codex_code_mode_protocol::grpc::code_mode_host_server::CodeModeHostServer; use codex_protocol::ToolName; use futures::FutureExt; use pretty_assertions::assert_eq; use serde_json::json; +#[cfg(unix)] +use tokio::net::UnixListener; use tokio::sync::Semaphore; use tokio::time::timeout; +#[cfg(unix)] +use tokio_stream::wrappers::UnixListenerStream; use tokio_util::sync::CancellationToken; use tonic::Code; +#[cfg(unix)] +use tonic::transport::Server; #[path = "support/host.rs"] mod host; @@ -899,3 +909,220 @@ async fn dropping_a_grpc_lease_retires_its_server_session() -> Result<()> { .context("dropping the gRPC lease did not retire its server session")??; Ok(()) } + +#[tokio::test] +async fn cached_session_recovers_after_a_remote_host_restarts() -> Result<()> { + let mut original = HostHarness::start("grpc://127.0.0.1:0").await?; + let listen_url = original + .endpoint + .replacen("http://", "grpc://", /*count*/ 1); + let provider = GrpcCodeModeSessionProvider::new(original.endpoint.clone()); + let delegate = Arc::new(RecordingDelegate::default()); + let session = provider + .create_session(delegate.clone()) + .await + .map_err(anyhow::Error::msg)?; + + let mut pending = request("await tools.echo({generation: 1}); await new Promise(() => {});"); + pending.enabled_tools = vec![tool("echo")]; + pending.yield_time_ms = Some(/*value*/ 1); + let started = session.execute(pending).await.map_err(anyhow::Error::msg)?; + let old_cell_id = started.cell_id.clone(); + assert_eq!(old_cell_id, cell_id("1")); + assert!(matches!( + started.initial_response().await, + Ok(RuntimeResponse::Yielded { .. }) + )); + + let interrupted_wait = start_active_wait( + Arc::clone(&session), + WaitRequest { + cell_id: old_cell_id.clone(), + yield_time_ms: 60_000, + }, + ) + .await?; + timeout(TEST_TIMEOUT, async { + while delegate + .invocations + .lock() + .unwrap_or_else(PoisonError::into_inner) + .is_empty() + { + tokio::task::yield_now().await; + } + }) + .await + .context("original host did not dispatch its tool callback")?; + original + ._child + .kill() + .await + .context("stop the original gRPC host")?; + assert!( + timeout(TEST_TIMEOUT, interrupted_wait) + .await + .context("host loss did not interrupt the pending wait")? + .context("interrupted wait task panicked")? + .is_err() + ); + timeout(TEST_TIMEOUT, async { + loop { + if delegate + .closed_cells + .lock() + .unwrap_or_else(PoisonError::into_inner) + .contains(&old_cell_id) + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .context("host loss did not retire the original generation's cell")?; + + let _replacement = HostHarness::start(&listen_url).await?; + let mut callback = request( + r#"const result = await tools.echo({generation: 2}); notify("reconnected"); text(result.value);"#, + ); + callback.tool_call_id = "reconnected-call".to_string(); + callback.enabled_tools = vec![tool("echo")]; + let (callback_response, concurrent_response) = tokio::join!( + execute(&session, callback), + execute(&session, request(r#"text("concurrent")"#)), + ); + let callback_response = callback_response?; + let RuntimeResponse::Result { + cell_id: callback_cell_id, + .. + } = &callback_response + else { + anyhow::bail!("reconnected tool call did not complete"); + }; + let callback_cell_id = callback_cell_id.clone(); + assert_eq!( + callback_response, + text_response(callback_cell_id.as_str(), "output") + ); + let concurrent_response = concurrent_response?; + let RuntimeResponse::Result { + cell_id: concurrent_cell_id, + .. + } = &concurrent_response + else { + anyhow::bail!("concurrent reconnected cell did not complete"); + }; + let concurrent_cell_id = concurrent_cell_id.clone(); + assert_eq!( + concurrent_response, + text_response(concurrent_cell_id.as_str(), "concurrent") + ); + let mut replacement_cell_ids = [callback_cell_id.as_str(), concurrent_cell_id.as_str()]; + replacement_cell_ids.sort_unstable(); + assert_eq!(replacement_cell_ids, ["g2:1", "g2:2"]); + assert_eq!( + delegate + .invocations + .lock() + .unwrap_or_else(PoisonError::into_inner) + .iter() + .map(|invocation| (invocation.cell_id.clone(), invocation.input.clone())) + .collect::>(), + vec![ + (old_cell_id.clone(), Some(json!({ "generation": 1 }))), + (callback_cell_id.clone(), Some(json!({ "generation": 2 }))), + ] + ); + assert_eq!( + *delegate + .notifications + .lock() + .unwrap_or_else(PoisonError::into_inner), + vec![( + "reconnected-call".to_string(), + callback_cell_id, + "reconnected".to_string(), + )] + ); + + let mut pending = request("await new Promise(() => {});"); + pending.yield_time_ms = Some(/*value*/ 1); + let started = session.execute(pending).await.map_err(anyhow::Error::msg)?; + let replacement_cell_id = started.cell_id.clone(); + assert_eq!(replacement_cell_id, cell_id("g2:3")); + assert_eq!( + started.initial_response().await, + Ok(RuntimeResponse::Yielded { + cell_id: replacement_cell_id.clone(), + content_items: Vec::new(), + }) + ); + assert_eq!( + session + .wait(WaitRequest { + cell_id: replacement_cell_id.clone(), + yield_time_ms: 1, + }) + .await + .map_err(anyhow::Error::msg)?, + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: replacement_cell_id.clone(), + content_items: Vec::new(), + }) + ); + assert_eq!( + session + .terminate(replacement_cell_id.clone()) + .await + .map_err(anyhow::Error::msg)?, + WaitOutcome::LiveCell(RuntimeResponse::Terminated { + cell_id: replacement_cell_id, + content_items: Vec::new(), + }) + ); + + let stale_wait = session + .wait(WaitRequest { + cell_id: old_cell_id.clone(), + yield_time_ms: 1, + }) + .await + .unwrap_err(); + assert!(stale_wait.contains("stale code-mode host generation")); + let stale_termination = session.terminate(old_cell_id).await.unwrap_err(); + assert!(stale_termination.contains("stale code-mode host generation")); + session.shutdown().await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[cfg(unix)] +#[tokio::test] +async fn unix_socket_endpoints_execute_code_mode_cells() -> Result<()> { + let directory = tempfile::tempdir().context("create Unix socket directory")?; + let socket_path = directory.path().join("grpc.sock"); + let listener = UnixListener::bind(&socket_path).context("bind code-mode Unix socket")?; + let server = tokio::spawn( + Server::builder() + .add_service(CodeModeHostServer::new(GrpcCodeModeHost::new())) + .serve_with_incoming(UnixListenerStream::new(listener)), + ); + + for endpoint in [ + format!("unix://{}", socket_path.display()), + format!("unix:{}", socket_path.display()), + ] { + let session = GrpcCodeModeSessionProvider::new(endpoint) + .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .await + .map_err(anyhow::Error::msg)?; + assert_eq!( + execute(&session, request(r#"text("unix socket")"#)).await?, + text_response("1", "unix socket") + ); + session.shutdown().await.map_err(anyhow::Error::msg)?; + } + + server.abort(); + Ok(()) +} diff --git a/codex-rs/code-mode/src/grpc_session/generation.rs b/codex-rs/code-mode/src/grpc_session/generation.rs new file mode 100644 index 0000000000..a81f1b4bea --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/generation.rs @@ -0,0 +1,124 @@ +use std::sync::Arc; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeNestedToolCall; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::NotificationFuture; +use codex_code_mode_protocol::RuntimeResponse; +use codex_code_mode_protocol::StartedCell; +use codex_code_mode_protocol::ToolInvocationFuture; +use codex_code_mode_protocol::WaitOutcome; +use tokio_util::sync::CancellationToken; + +pub(super) struct GenerationDelegate { + pub(super) delegate: Arc, + pub(super) generation: u64, +} + +impl CodeModeSessionDelegate for GenerationDelegate { + fn invoke_tool<'a>( + &'a self, + mut invocation: CodeModeNestedToolCall, + cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + invocation.cell_id = public_cell_id(self.generation, &invocation.cell_id); + self.delegate.invoke_tool(invocation, cancellation_token) + } + + fn notify<'a>( + &'a self, + call_id: String, + cell_id: CellId, + text: String, + cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + self.delegate.notify( + call_id, + public_cell_id(self.generation, &cell_id), + text, + cancellation_token, + ) + } + + fn cell_closed(&self, cell_id: &CellId) { + self.delegate + .cell_closed(&public_cell_id(self.generation, cell_id)); + } +} + +fn public_cell_id(generation: u64, cell_id: &CellId) -> CellId { + if generation == 1 { + cell_id.clone() + } else { + CellId::new(format!("g{generation}:{cell_id}")) + } +} + +pub(super) fn remote_cell_id(generation: u64, cell_id: &CellId) -> Result { + if generation == 1 { + return Ok(cell_id.clone()); + } + + let prefix = format!("g{generation}:"); + cell_id + .as_str() + .strip_prefix(&prefix) + .map(|cell_id| CellId::new(cell_id.to_string())) + .ok_or_else(|| "cell belongs to a stale code-mode host generation".to_string()) +} + +pub(super) fn public_started_cell(generation: u64, started: StartedCell) -> StartedCell { + if generation == 1 { + return started; + } + let cell_id = public_cell_id(generation, &started.cell_id); + StartedCell::from_future(cell_id, async move { + started + .initial_response() + .await + .map(|response| public_runtime_response(generation, response)) + }) +} + +fn public_runtime_response(generation: u64, response: RuntimeResponse) -> RuntimeResponse { + match response { + RuntimeResponse::Yielded { + cell_id, + content_items, + } => RuntimeResponse::Yielded { + cell_id: public_cell_id(generation, &cell_id), + content_items, + }, + RuntimeResponse::Terminated { + cell_id, + content_items, + } => RuntimeResponse::Terminated { + cell_id: public_cell_id(generation, &cell_id), + content_items, + }, + RuntimeResponse::Result { + cell_id, + content_items, + error_text, + } => RuntimeResponse::Result { + cell_id: public_cell_id(generation, &cell_id), + content_items, + error_text, + }, + } +} + +pub(super) fn public_wait_outcome(generation: u64, outcome: WaitOutcome) -> WaitOutcome { + match outcome { + WaitOutcome::LiveCell(response) => { + WaitOutcome::LiveCell(public_runtime_response(generation, response)) + } + WaitOutcome::MissingCell(response) => { + WaitOutcome::MissingCell(public_runtime_response(generation, response)) + } + } +} + +#[cfg(test)] +#[path = "generation_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode/src/grpc_session/generation_tests.rs b/codex-rs/code-mode/src/grpc_session/generation_tests.rs new file mode 100644 index 0000000000..3bf1a064dd --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/generation_tests.rs @@ -0,0 +1,230 @@ +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeNestedToolCall; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::CodeModeToolKind; +use codex_code_mode_protocol::FunctionCallOutputContentItem; +use codex_code_mode_protocol::NotificationFuture; +use codex_code_mode_protocol::RuntimeResponse; +use codex_code_mode_protocol::StartedCell; +use codex_code_mode_protocol::ToolInvocationFuture; +use codex_code_mode_protocol::WaitOutcome; +use codex_protocol::ToolName; +use pretty_assertions::assert_eq; +use serde_json::json; +use tokio::sync::oneshot; +use tokio_util::sync::CancellationToken; + +use super::GenerationDelegate; +use super::public_cell_id; +use super::public_started_cell; +use super::public_wait_outcome; +use super::remote_cell_id; + +#[derive(Default)] +struct RecordingDelegate { + calls: Mutex>, + notifications: Mutex>, + closed: Mutex>, +} + +impl CodeModeSessionDelegate for RecordingDelegate { + fn invoke_tool<'a>( + &'a self, + invocation: CodeModeNestedToolCall, + _cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + self.calls.lock().expect("calls lock").push(invocation); + Box::pin(async { Ok(json!({ "ok": true })) }) + } + + fn notify<'a>( + &'a self, + call_id: String, + cell_id: CellId, + text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + self.notifications + .lock() + .expect("notifications lock") + .push((call_id, cell_id, text)); + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, cell_id: &CellId) { + self.closed + .lock() + .expect("closed cells lock") + .push(cell_id.clone()); + } +} + +#[test] +fn first_generation_preserves_existing_cell_ids() { + let cell_id = CellId::new("42".to_string()); + + assert_eq!(public_cell_id(/*generation*/ 1, &cell_id), cell_id); + assert_eq!(remote_cell_id(/*generation*/ 1, &cell_id), Ok(cell_id)); +} + +#[test] +fn first_generation_preserves_opaque_cell_ids_that_resemble_generation_prefixes() { + for value in ["graphics:1", "g2:42"] { + let cell_id = CellId::new(value.to_string()); + + assert_eq!(public_cell_id(/*generation*/ 1, &cell_id), cell_id); + assert_eq!(remote_cell_id(/*generation*/ 1, &cell_id), Ok(cell_id)); + } +} + +#[test] +fn later_generations_prefix_public_ids_and_strip_wire_ids() { + let wire_id = CellId::new("42".to_string()); + let public_id = CellId::new("g2:42".to_string()); + + assert_eq!(public_cell_id(/*generation*/ 2, &wire_id), public_id); + assert_eq!(remote_cell_id(/*generation*/ 2, &public_id), Ok(wire_id)); +} + +#[test] +fn stale_generation_ids_are_rejected_after_reconnection() { + for cell_id in [ + "42".to_string(), + "g1:42".to_string(), + "g3:42".to_string(), + "x".repeat(10_000), + ] { + assert_eq!( + remote_cell_id(/*generation*/ 2, &CellId::new(cell_id)), + Err("cell belongs to a stale code-mode host generation".to_string()) + ); + } +} + +#[tokio::test] +async fn reconnect_maps_every_delegate_callback_to_its_generation() { + let recording = Arc::new(RecordingDelegate::default()); + let delegate = GenerationDelegate { + delegate: recording.clone(), + generation: 2, + }; + let wire_id = CellId::new("42".to_string()); + let public_id = CellId::new("g2:42".to_string()); + let invocation = CodeModeNestedToolCall { + cell_id: wire_id.clone(), + runtime_tool_call_id: "runtime-call".to_string(), + tool_name: ToolName::plain("echo"), + tool_kind: CodeModeToolKind::Function, + input: Some(json!({ "value": true })), + }; + + assert_eq!( + delegate + .invoke_tool(invocation.clone(), CancellationToken::new()) + .await, + Ok(json!({ "ok": true })) + ); + assert_eq!( + delegate + .notify( + "outer-call".to_string(), + wire_id.clone(), + "notice".to_string(), + CancellationToken::new(), + ) + .await, + Ok(()) + ); + delegate.cell_closed(&wire_id); + + assert_eq!( + *recording.calls.lock().expect("calls lock"), + vec![CodeModeNestedToolCall { + cell_id: public_id.clone(), + ..invocation + }] + ); + assert_eq!( + *recording.notifications.lock().expect("notifications lock"), + vec![( + "outer-call".to_string(), + public_id.clone(), + "notice".to_string() + )] + ); + assert_eq!( + *recording.closed.lock().expect("closed cells lock"), + vec![public_id] + ); +} + +#[tokio::test] +async fn reconnected_execution_maps_started_and_initial_response_ids() { + let (response_tx, response_rx) = oneshot::channel(); + let wire_id = CellId::new("42".to_string()); + let public_id = CellId::new("g2:42".to_string()); + let claimed = Arc::new(AtomicBool::new(false)); + let initial_response_claimed = Arc::clone(&claimed); + let response = RuntimeResponse::Result { + cell_id: wire_id.clone(), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "result".to_string(), + }], + error_text: None, + }; + let started = StartedCell::from_future(wire_id, async move { + initial_response_claimed.store(true, Ordering::Release); + response_rx.await.expect("receive initial response") + }); + let started = public_started_cell(/*generation*/ 2, started); + + assert_eq!(started.cell_id, public_id); + tokio::task::yield_now().await; + assert!(!claimed.load(Ordering::Acquire)); + response_tx + .send(Ok(response)) + .expect("send initial response"); + assert_eq!( + started.initial_response().await, + Ok(RuntimeResponse::Result { + cell_id: public_id, + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "result".to_string(), + }], + error_text: None, + }) + ); +} + +#[test] +fn reconnected_wait_maps_live_and_missing_outcomes() { + let public_id = CellId::new("g2:42".to_string()); + let yielded = RuntimeResponse::Yielded { + cell_id: CellId::new("42".to_string()), + content_items: Vec::new(), + }; + let terminated = RuntimeResponse::Terminated { + cell_id: CellId::new("42".to_string()), + content_items: Vec::new(), + }; + + assert_eq!( + public_wait_outcome(/*generation*/ 2, WaitOutcome::LiveCell(yielded)), + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: public_id.clone(), + content_items: Vec::new(), + }) + ); + assert_eq!( + public_wait_outcome(/*generation*/ 2, WaitOutcome::MissingCell(terminated)), + WaitOutcome::MissingCell(RuntimeResponse::Terminated { + cell_id: public_id, + content_items: Vec::new(), + }) + ); +} diff --git a/codex-rs/code-mode/src/grpc_session/mod.rs b/codex-rs/code-mode/src/grpc_session/mod.rs index 43a325f5f0..a7932bbd2f 100644 --- a/codex-rs/code-mode/src/grpc_session/mod.rs +++ b/codex-rs/code-mode/src/grpc_session/mod.rs @@ -39,12 +39,16 @@ mod callbacks; mod completion; mod conversion; mod deadline; +mod generation; mod operations; +mod reconnect; mod state; mod transport; type GrpcClient = CodeModeHostClient; +const SHUTDOWN_ERROR: &str = "code mode session is shutting down"; + /// Creates code-mode sessions over an HTTP/2 gRPC connection. #[derive(Clone)] pub struct GrpcCodeModeSessionProvider { @@ -52,7 +56,7 @@ pub struct GrpcCodeModeSessionProvider { } impl GrpcCodeModeSessionProvider { - /// Connects lazily to an `http://` or `https://` gRPC endpoint. + /// Connects lazily to an `http://`, `https://`, or `unix://` gRPC endpoint. pub fn new(endpoint: impl Into) -> Self { Self::with_http_client_factory( endpoint, @@ -166,9 +170,13 @@ impl CodeModeSessionProvider for GrpcCodeModeSessionProvider { limits: CodeModeSessionCellExecutionLimits, ) -> CodeModeSessionProviderFuture<'a> { Box::pin(async move { - self.open_binding(delegate, limits) - .await - .map(|session| session as _) + let session = Arc::new(reconnect::ReconnectableSession::new( + self.clone(), + delegate, + limits, + )); + session.initialize().await?; + Ok(session as _) }) } } @@ -246,7 +254,7 @@ impl SessionInner { .unwrap_or_else(PoisonError::into_inner) .require_open()?; if self.shutdown_requested.load(Ordering::Acquire) { - return Err("code mode session is shutting down".to_string()); + return Err(SHUTDOWN_ERROR.to_string()); } Ok(()) } diff --git a/codex-rs/code-mode/src/grpc_session/reconnect.rs b/codex-rs/code-mode/src/grpc_session/reconnect.rs new file mode 100644 index 0000000000..e901fde6f4 --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/reconnect.rs @@ -0,0 +1,237 @@ +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSession; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::CodeModeSessionResultFuture; +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 tokio::sync::Semaphore; +use tokio::sync::watch; +use tokio_util::sync::CancellationToken; + +use super::GrpcCodeModeSession; +use super::GrpcCodeModeSessionProvider; +use super::SHUTDOWN_ERROR; +use super::generation; +use super::generation::GenerationDelegate; +use crate::remote_session::ShutdownResultReceiver; +use crate::remote_session::wait_for_watch; + +pub(super) struct ReconnectableSession { + inner: Arc, +} + +struct ReconnectInner { + provider: GrpcCodeModeSessionProvider, + delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, + binding: Mutex>, + opening_permit: Semaphore, + next_generation: AtomicU64, + shutdown_requested: CancellationToken, + shutdown_result: Mutex>, +} + +#[derive(Clone)] +struct SessionBinding { + session: Arc, + generation: u64, +} + +impl ReconnectableSession { + pub(super) fn new( + provider: GrpcCodeModeSessionProvider, + delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, + ) -> Self { + Self { + inner: Arc::new(ReconnectInner { + provider, + delegate, + limits, + binding: Mutex::new(None), + opening_permit: Semaphore::new(/*permits*/ 1), + next_generation: AtomicU64::new(1), + shutdown_requested: CancellationToken::new(), + shutdown_result: Mutex::new(None), + }), + } + } + + pub(super) async fn initialize(&self) -> Result<(), String> { + self.inner.get_or_open_binding().await.map(|_| ()) + } +} + +impl CodeModeSession for ReconnectableSession { + fn execute<'a>( + &'a self, + request: ExecuteRequest, + ) -> CodeModeSessionResultFuture<'a, StartedCell> { + Box::pin(async move { + let binding = self.inner.get_or_open_binding().await?; + let started = binding.session.execute(request).await?; + Ok(generation::public_started_cell(binding.generation, started)) + }) + } + + fn wait<'a>(&'a self, request: WaitRequest) -> CodeModeSessionResultFuture<'a, WaitOutcome> { + Box::pin(async move { + let binding = self.inner.get_or_open_binding().await?; + let request = WaitRequest { + cell_id: generation::remote_cell_id(binding.generation, &request.cell_id)?, + yield_time_ms: request.yield_time_ms, + }; + let outcome = binding.session.wait(request).await?; + Ok(generation::public_wait_outcome(binding.generation, outcome)) + }) + } + + fn terminate<'a>(&'a self, cell_id: CellId) -> CodeModeSessionResultFuture<'a, WaitOutcome> { + Box::pin(async move { + let binding = self.inner.get_or_open_binding().await?; + let cell_id = generation::remote_cell_id(binding.generation, &cell_id)?; + let outcome = binding.session.terminate(cell_id).await?; + Ok(generation::public_wait_outcome(binding.generation, outcome)) + }) + } + + fn shutdown<'a>(&'a self) -> CodeModeSessionResultFuture<'a, ()> { + Box::pin(wait_for_watch(self.inner.request_shutdown())) + } +} + +impl Drop for ReconnectableSession { + fn drop(&mut self) { + if tokio::runtime::Handle::try_current().is_ok() { + self.inner.request_shutdown(); + } + } +} + +impl ReconnectInner { + async fn get_or_open_binding(&self) -> Result { + if self.shutdown_requested.is_cancelled() { + return Err(SHUTDOWN_ERROR.to_string()); + } + if let Some(binding) = self.live_binding() { + return Ok(binding); + } + + let _opening_permit = tokio::select! { + biased; + _ = self.shutdown_requested.cancelled() => { + return Err(SHUTDOWN_ERROR.to_string()); + } + permit = self.opening_permit.acquire() => permit + .map_err(|_| "gRPC code-mode session opening coordinator closed".to_string())?, + }; + if self.shutdown_requested.is_cancelled() { + return Err(SHUTDOWN_ERROR.to_string()); + } + if let Some(binding) = self.live_binding() { + return Ok(binding); + } + + let previous_binding = self + .binding + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone(); + if let Some(binding) = previous_binding { + wait_for_watch(binding.session.inner.request_shutdown()).await?; + } + + let generation = self.next_generation.fetch_add(1, Ordering::Relaxed); + let delegate = Arc::new(GenerationDelegate { + delegate: Arc::clone(&self.delegate), + generation, + }); + let session = tokio::select! { + biased; + _ = self.shutdown_requested.cancelled() => { + return Err(SHUTDOWN_ERROR.to_string()); + } + session = self.provider.open_binding(delegate, self.limits.clone()) => session?, + }; + let binding = SessionBinding { + session, + generation, + }; + let published = { + let mut current = self.binding.lock().unwrap_or_else(PoisonError::into_inner); + if self.shutdown_requested.is_cancelled() { + false + } else { + *current = Some(binding.clone()); + true + } + }; + if !published { + let _ = wait_for_watch(binding.session.inner.request_shutdown()).await; + return Err(SHUTDOWN_ERROR.to_string()); + } + Ok(binding) + } + + fn live_binding(&self) -> Option { + self.binding + .lock() + .unwrap_or_else(PoisonError::into_inner) + .as_ref() + .filter(|binding| !binding.session.inner.stopped.is_cancelled()) + .cloned() + } + + fn request_shutdown(self: &Arc) -> ShutdownResultReceiver { + { + let binding = self.binding.lock().unwrap_or_else(PoisonError::into_inner); + self.shutdown_requested.cancel(); + if let Some(binding) = binding.as_ref() { + binding.session.inner.request_shutdown(); + } + } + let mut result = self + .shutdown_result + .lock() + .unwrap_or_else(PoisonError::into_inner); + if let Some(receiver) = result.as_ref() { + return receiver.clone(); + } + + let (sender, receiver) = watch::channel(None); + *result = Some(receiver.clone()); + let inner = Arc::clone(self); + tokio::spawn(async move { + let opening_permit = match inner.opening_permit.acquire().await { + Ok(permit) => permit, + Err(_) => { + sender.send_replace(Some(Err( + "gRPC code-mode session opening coordinator closed".to_string(), + ))); + return; + } + }; + let binding = inner + .binding + .lock() + .unwrap_or_else(PoisonError::into_inner) + .take(); + drop(opening_permit); + let result = match binding { + Some(binding) => wait_for_watch(binding.session.inner.request_shutdown()).await, + None => Ok(()), + }; + sender.send_replace(Some(result)); + }); + receiver + } +} diff --git a/codex-rs/code-mode/src/grpc_session/transport.rs b/codex-rs/code-mode/src/grpc_session/transport.rs index ce98f8e188..0ca665125b 100644 --- a/codex-rs/code-mode/src/grpc_session/transport.rs +++ b/codex-rs/code-mode/src/grpc_session/transport.rs @@ -10,6 +10,7 @@ use tonic::codegen::http::Request; use tonic::codegen::http::Response; use tonic::codegen::http::Uri; use tonic::transport::Channel; +use tonic::transport::Endpoint; use tower::ServiceExt; use tower::service_fn; use tower::util::BoxCloneSyncService; @@ -53,6 +54,15 @@ impl SharedTransport { self.client .get_or_try_init(|| async { let client = match &self.endpoint { + TransportEndpoint::Url { endpoint, .. } if endpoint.starts_with("unix:") => { + let channel = Endpoint::from_shared(endpoint.clone()) + .map_err(|error| { + format!("invalid gRPC code-mode Unix socket endpoint: {error}") + })? + .connect_lazy(); + let transport = channel.map_err(io::Error::other); + CodeModeHostClient::new(BoxCloneSyncService::new(transport)) + } TransportEndpoint::Url { endpoint, http_client_factory,