From 1e557a554e8fa130bbf5c50de91a7fbb3487f962 Mon Sep 17 00:00:00 2001 From: Channing Conger Date: Tue, 11 Aug 2026 17:17:58 +0000 Subject: [PATCH] Add gRPC-backed code-mode sessions (#38041) ## What changed - Add `GrpcCodeModeSessionProvider` for opening code-mode sessions over HTTP/2 or an existing `tonic` channel. - Support execution, waiting, termination, per-session limits, cell-closure callbacks, and graceful shutdown over the gRPC protocol. - Bound transport waits and error messages, validate host identifiers and responses, and clean up abandoned executions and observers. ## Testing - Add end-to-end TCP tests covering session persistence, cancellation, concurrent waits, shutdown, cell cleanup, and independent yield limits. - Add unit coverage for protocol conversion, deadlines, and session lifecycle state. GitOrigin-RevId: d4729ce608ad4b42a99744b07e1f230e46cb24ec --- codex-rs/Cargo.lock | 3 + .../src/grpc/robustness_tests.rs | 26 + .../code-mode-host/src/grpc/validation.rs | 2 +- codex-rs/code-mode-host/tests/grpc.rs | 645 ++++++++++++++++++ .../src/grpc/codex.code_mode.v1.proto | 3 +- codex-rs/code-mode-protocol/src/grpc/mod.rs | 2 + codex-rs/code-mode-protocol/src/session.rs | 30 +- codex-rs/code-mode/Cargo.toml | 3 + .../code-mode/src/grpc_session/callbacks.rs | 70 ++ .../code-mode/src/grpc_session/conversion.rs | 134 ++++ .../src/grpc_session/conversion_tests.rs | 168 +++++ .../code-mode/src/grpc_session/deadline.rs | 77 +++ .../src/grpc_session/deadline_tests.rs | 131 ++++ codex-rs/code-mode/src/grpc_session/mod.rs | 317 +++++++++ .../code-mode/src/grpc_session/operations.rs | 391 +++++++++++ codex-rs/code-mode/src/grpc_session/state.rs | 149 ++++ .../code-mode/src/grpc_session/state_tests.rs | 89 +++ .../code-mode/src/grpc_session/transport.rs | 46 ++ codex-rs/code-mode/src/lib.rs | 2 + codex-rs/code-mode/src/remote_session.rs | 4 +- 20 files changed, 2275 insertions(+), 17 deletions(-) create mode 100644 codex-rs/code-mode-host/tests/grpc.rs create mode 100644 codex-rs/code-mode/src/grpc_session/callbacks.rs create mode 100644 codex-rs/code-mode/src/grpc_session/conversion.rs create mode 100644 codex-rs/code-mode/src/grpc_session/conversion_tests.rs create mode 100644 codex-rs/code-mode/src/grpc_session/deadline.rs create mode 100644 codex-rs/code-mode/src/grpc_session/deadline_tests.rs create mode 100644 codex-rs/code-mode/src/grpc_session/mod.rs create mode 100644 codex-rs/code-mode/src/grpc_session/operations.rs create mode 100644 codex-rs/code-mode/src/grpc_session/state.rs create mode 100644 codex-rs/code-mode/src/grpc_session/state_tests.rs create mode 100644 codex-rs/code-mode/src/grpc_session/transport.rs diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 2e04b753e5..312cde8937 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -2564,10 +2564,13 @@ dependencies = [ "codex-websocket-client", "futures", "pretty_assertions", + "serde_json", "tokio", "tokio-tungstenite", "tokio-util", + "tonic", "tracing", + "uuid", ] [[package]] diff --git a/codex-rs/code-mode-host/src/grpc/robustness_tests.rs b/codex-rs/code-mode-host/src/grpc/robustness_tests.rs index ced3290edd..eb9b21b086 100644 --- a/codex-rs/code-mode-host/src/grpc/robustness_tests.rs +++ b/codex-rs/code-mode-host/src/grpc/robustness_tests.rs @@ -18,6 +18,7 @@ use tonic::Request; use tonic::Status; use uuid::Uuid; +use super::ExecutionAdmission; use super::GrpcCodeModeHost; use super::tests::execute_events; use super::tests::execute_request; @@ -125,6 +126,31 @@ async fn rejects_oversized_identifiers_tool_metadata_and_subscription_filters() assert!(host.state.session(&session_id).is_ok()); } +#[tokio::test] +async fn dropping_execution_before_admission_releases_its_reservation() { + let host = GrpcCodeModeHost::new(); + let (session_id, _events) = open_session(&host).await; + let session = host.state.session(&session_id).expect("open session"); + let execution_id = "execution-abandoned-before-admission".to_string(); + session + .reserve_execution(&execution_id) + .expect("reserve execution"); + + drop(ExecutionAdmission { + session: Arc::clone(&session), + execution_id: Some(execution_id.clone()), + }); + + let error = session + .admit_execution( + execution_id, + "cell".to_string(), + host.state.cell_permit().expect("reserve cell permit"), + ) + .expect_err("abandoned execution must not admit a runtime cell"); + assert_eq!(error.code(), Code::Cancelled); +} + #[tokio::test] async fn dropping_an_unread_buffered_execution_outcome_retires_its_cell() { let host = GrpcCodeModeHost::new(); diff --git a/codex-rs/code-mode-host/src/grpc/validation.rs b/codex-rs/code-mode-host/src/grpc/validation.rs index 7328a96f33..0a45ad5995 100644 --- a/codex-rs/code-mode-host/src/grpc/validation.rs +++ b/codex-rs/code-mode-host/src/grpc/validation.rs @@ -2,7 +2,7 @@ use codex_code_mode_protocol::grpc as proto; use tonic::Status; use uuid::Uuid; -pub(super) const MAX_IDENTIFIER_BYTES: usize = 256; +pub(super) use codex_code_mode_protocol::grpc::MAX_IDENTIFIER_BYTES; pub(super) const MAX_TOOL_FILTERS: usize = 64; pub(super) const MAX_TOOL_DEFINITIONS: usize = 1_024; pub(super) const MAX_TOOL_DESCRIPTION_BYTES: usize = 16 * 1_024; diff --git a/codex-rs/code-mode-host/tests/grpc.rs b/codex-rs/code-mode-host/tests/grpc.rs new file mode 100644 index 0000000000..168f3b7b2c --- /dev/null +++ b/codex-rs/code-mode-host/tests/grpc.rs @@ -0,0 +1,645 @@ +use std::sync::Arc; +use std::sync::PoisonError; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use codex_code_mode::CodeModeSession; +use codex_code_mode::CodeModeSessionCellExecutionLimits; +use codex_code_mode::CodeModeSessionProvider; +use codex_code_mode::ExecuteRequest; +use codex_code_mode::FunctionCallOutputContentItem; +use codex_code_mode::GrpcCodeModeSessionProvider; +use codex_code_mode::NoopCodeModeSessionDelegate; +use codex_code_mode::RuntimeResponse; +use codex_code_mode::WaitOutcome; +use codex_code_mode::WaitRequest; +use codex_code_mode_protocol::grpc; +use codex_code_mode_protocol::grpc::code_mode_host_client::CodeModeHostClient; +use futures::FutureExt; +use pretty_assertions::assert_eq; +use tokio::time::timeout; +use tonic::Code; + +#[path = "support/host.rs"] +mod host; +#[path = "support/recording_delegate.rs"] +mod recording_delegate; + +use host::HostHarness; +use recording_delegate::RecordingDelegate; +use recording_delegate::cell_id; + +const TEST_TIMEOUT: Duration = Duration::from_secs(20); + +fn request(source: &str) -> ExecuteRequest { + ExecuteRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: Vec::new(), + source: source.to_string(), + yield_time_ms: Some(/*value*/ 5_000), + max_output_tokens: Some(/*value*/ 1_000), + } +} + +fn text_response(cell: &str, value: &str) -> RuntimeResponse { + RuntimeResponse::Result { + cell_id: cell_id(cell), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: value.to_string(), + }], + error_text: None, + } +} + +async fn execute( + session: &Arc, + request: ExecuteRequest, +) -> Result { + timeout(TEST_TIMEOUT, async { + session + .execute(request) + .await + .map_err(anyhow::Error::msg)? + .initial_response() + .await + .map_err(anyhow::Error::msg) + }) + .await + .context("timed out executing gRPC code-mode cell")? +} + +async fn start_active_wait( + session: Arc, + request: WaitRequest, +) -> Result>> { + let (admitted_tx, admitted_rx) = tokio::sync::oneshot::channel(); + let wait = tokio::spawn(async move { + let mut wait = session.wait(request); + match wait.as_mut().now_or_never() { + Some(result) => result, + None => { + let _ = admitted_tx.send(()); + wait.await + } + } + }); + timeout(TEST_TIMEOUT, admitted_rx) + .await + .context("timed out waiting for observer admission")? + .context("wait completed before its observer became active")?; + Ok(wait) +} + +#[tokio::test] +async fn tcp_session_persists_values_and_reports_cell_closure() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + assert!(host.endpoint.starts_with("http://127.0.0.1:")); + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let delegate = Arc::new(RecordingDelegate::default()); + let session = provider + .create_session(delegate.clone()) + .await + .map_err(anyhow::Error::msg)?; + + assert_eq!( + execute(&session, request(r#"store("key", "persisted");"#)).await?, + RuntimeResponse::Result { + cell_id: cell_id("1"), + content_items: Vec::new(), + error_text: None, + } + ); + + assert_eq!( + execute(&session, request(r#"text(String(load("key")));"#)).await?, + text_response("2", "persisted") + ); + + session.shutdown().await.map_err(anyhow::Error::msg)?; + assert_eq!( + *delegate + .closed_cells + .lock() + .unwrap_or_else(PoisonError::into_inner), + vec![cell_id("1"), cell_id("2")] + ); + Ok(()) +} + +#[tokio::test] +async fn shutdown_immediately_rejects_new_operations() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let session = provider + .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .await + .map_err(anyhow::Error::msg)?; + + let shutdown = session.shutdown(); + let expected = "code mode session is shutting down".to_string(); + assert_eq!( + session.execute(request("text('too late');")).await.err(), + Some(expected.clone()) + ); + assert_eq!( + session + .wait(WaitRequest { + cell_id: cell_id("missing"), + yield_time_ms: 1, + }) + .await, + Err(expected.clone()) + ); + assert_eq!(session.terminate(cell_id("missing")).await, Err(expected)); + + shutdown.await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[tokio::test] +async fn cancelling_execution_before_admission_keeps_the_session_usable() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let delegate = Arc::new(RecordingDelegate::default()); + let session = provider + .create_session(delegate.clone()) + .await + .map_err(anyhow::Error::msg)?; + let mut pending = request("await new Promise(() => {});"); + pending.yield_time_ms = Some(/*value*/ 1); + + assert!(session.execute(pending).now_or_never().is_none()); + + let abandoned_cell = cell_id("1"); + timeout(TEST_TIMEOUT, async { + loop { + if delegate + .closed_cells + .lock() + .unwrap_or_else(PoisonError::into_inner) + .contains(&abandoned_cell) + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .context("cancelled execution was never admitted and cleaned up")?; + + timeout(TEST_TIMEOUT, async { + loop { + match session + .wait(WaitRequest { + cell_id: abandoned_cell.clone(), + yield_time_ms: 1, + }) + .await + { + Ok(WaitOutcome::MissingCell(_)) + | Ok(WaitOutcome::LiveCell(RuntimeResponse::Terminated { .. })) => break Ok(()), + Ok(WaitOutcome::LiveCell(RuntimeResponse::Yielded { .. })) => { + tokio::task::yield_now().await; + } + Ok(outcome) => anyhow::bail!("unexpected abandoned-cell outcome: {outcome:?}"), + Err(error) if error.contains("already has an active observer") => { + tokio::task::yield_now().await; + } + Err(error) => break Err(anyhow::Error::msg(error)), + } + } + }) + .await + .context("cancelled execution leaked its remote cell")??; + + assert_eq!( + execute(&session, request(r#"text("still alive");"#)).await?, + text_response("2", "still alive") + ); + session.shutdown().await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[tokio::test] +async fn dropping_a_started_cell_off_runtime_terminates_its_buffered_remote_execution() -> Result<()> +{ + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let delegate = Arc::new(RecordingDelegate::default()); + let session = provider + .create_session(delegate.clone()) + .await + .map_err(anyhow::Error::msg)?; + 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 abandoned_cell = started.cell_id.clone(); + + timeout(TEST_TIMEOUT, async { + loop { + match session + .wait(WaitRequest { + cell_id: abandoned_cell.clone(), + yield_time_ms: 1, + }) + .await + { + Ok(WaitOutcome::LiveCell(RuntimeResponse::Yielded { .. })) => break Ok(()), + Err(error) if error.contains("already has an active observer") => { + tokio::task::yield_now().await; + } + Ok(outcome) => anyhow::bail!("unexpected execution outcome: {outcome:?}"), + Err(error) => break Err(anyhow::Error::msg(error)), + } + } + }) + .await + .context("execution never produced a buffered initial response")??; + + std::thread::spawn(move || drop(started)) + .join() + .map_err(|_| anyhow::anyhow!("dropping a started cell outside Tokio panicked"))?; + + timeout(TEST_TIMEOUT, async { + loop { + match session + .wait(WaitRequest { + cell_id: abandoned_cell.clone(), + yield_time_ms: 1, + }) + .await + { + Ok(WaitOutcome::MissingCell(_)) + | Ok(WaitOutcome::LiveCell(RuntimeResponse::Terminated { .. })) => break Ok(()), + Ok(WaitOutcome::LiveCell(RuntimeResponse::Yielded { .. })) => { + tokio::task::yield_now().await; + } + Ok(outcome) => anyhow::bail!("unexpected abandoned-cell outcome: {outcome:?}"), + Err(error) if error.contains("already has an active observer") => { + tokio::task::yield_now().await; + } + Err(error) => break Err(anyhow::Error::msg(error)), + } + } + }) + .await + .context("dropping a started cell did not terminate its buffered remote execution")??; + + assert!( + delegate + .closed_cells + .lock() + .unwrap_or_else(PoisonError::into_inner) + .contains(&abandoned_cell) + ); + + assert_eq!( + execute(&session, request(r#"text("still alive");"#)).await?, + text_response("2", "still alive") + ); + session.shutdown().await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[tokio::test] +async fn dropping_an_initial_response_terminates_its_pending_remote_execution() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let delegate = Arc::new(RecordingDelegate::default()); + let session = provider + .create_session(delegate.clone()) + .await + .map_err(anyhow::Error::msg)?; + let mut pending = request("await new Promise(() => {});"); + pending.yield_time_ms = Some(/*value*/ 60_000); + let started = session.execute(pending).await.map_err(anyhow::Error::msg)?; + let abandoned_cell = started.cell_id.clone(); + let initial_response = tokio::spawn(started.initial_response()); + tokio::task::yield_now().await; + initial_response.abort(); + let _ = initial_response.await; + + timeout(TEST_TIMEOUT, async { + loop { + if delegate + .closed_cells + .lock() + .unwrap_or_else(PoisonError::into_inner) + .contains(&abandoned_cell) + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .context("dropping an initial response did not terminate its pending remote execution")?; + + assert_eq!( + execute(&session, request(r#"text("still alive");"#)).await?, + text_response("2", "still alive") + ); + session.shutdown().await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[tokio::test] +async fn concurrent_wait_rejects_without_displacing_the_active_observer() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let session = provider + .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .await + .map_err(anyhow::Error::msg)?; + 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 running_cell = started.cell_id.clone(); + assert_eq!( + started + .initial_response() + .await + .map_err(anyhow::Error::msg)?, + RuntimeResponse::Yielded { + cell_id: running_cell.clone(), + content_items: Vec::new(), + } + ); + + let first_wait = start_active_wait( + Arc::clone(&session), + WaitRequest { + cell_id: running_cell.clone(), + yield_time_ms: 100, + }, + ) + .await?; + + assert_eq!( + timeout( + Duration::from_secs(/*secs*/ 1), + session.wait(WaitRequest { + cell_id: running_cell.clone(), + yield_time_ms: 60_000, + }), + ) + .await + .context("concurrent wait did not reject immediately")? + .unwrap_err(), + format!("exec cell {running_cell} already has an active observer") + ); + + assert_eq!( + timeout(TEST_TIMEOUT, first_wait) + .await + .context("active wait was displaced by the rejected observer")?? + .map_err(anyhow::Error::msg)?, + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: running_cell.clone(), + content_items: Vec::new(), + }) + ); + assert_eq!( + session + .terminate(running_cell.clone()) + .await + .map_err(anyhow::Error::msg)?, + WaitOutcome::LiveCell(RuntimeResponse::Terminated { + cell_id: running_cell, + content_items: Vec::new(), + }) + ); + + session.shutdown().await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[tokio::test] +async fn dropping_a_wait_retires_its_observer_before_the_next_wait() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let session = provider + .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .await + .map_err(anyhow::Error::msg)?; + 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 running_cell = started.cell_id.clone(); + assert_eq!( + started + .initial_response() + .await + .map_err(anyhow::Error::msg)?, + RuntimeResponse::Yielded { + cell_id: running_cell.clone(), + content_items: Vec::new(), + } + ); + + let first_wait = start_active_wait( + Arc::clone(&session), + WaitRequest { + cell_id: running_cell.clone(), + yield_time_ms: 60_000, + }, + ) + .await?; + first_wait.abort(); + let _ = first_wait.await; + + assert_eq!( + timeout( + TEST_TIMEOUT, + session.wait(WaitRequest { + cell_id: running_cell.clone(), + yield_time_ms: 1, + }), + ) + .await + .context("replacement wait did not observe cancellation retirement")? + .map_err(anyhow::Error::msg)?, + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: running_cell.clone(), + content_items: Vec::new(), + }) + ); + assert_eq!( + session + .terminate(running_cell.clone()) + .await + .map_err(anyhow::Error::msg)?, + WaitOutcome::LiveCell(RuntimeResponse::Terminated { + cell_id: running_cell, + content_items: Vec::new(), + }) + ); + session.shutdown().await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[tokio::test] +async fn dropping_a_session_off_runtime_retires_its_active_cells() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let delegate = Arc::new(RecordingDelegate::default()); + let session = provider + .create_session(delegate.clone()) + .await + .map_err(anyhow::Error::msg)?; + let mut pending = request("await new Promise(() => {});"); + pending.yield_time_ms = Some(/*value*/ 1); + assert_eq!( + execute(&session, pending).await?, + RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + } + ); + + std::thread::spawn(move || drop(session)) + .join() + .map_err(|_| anyhow::anyhow!("dropping a session outside Tokio panicked"))?; + + timeout(TEST_TIMEOUT, async { + loop { + if delegate + .closed_cells + .lock() + .unwrap_or_else(PoisonError::into_inner) + .contains(&cell_id("1")) + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .context("dropping a session outside Tokio did not retire its active cell")?; + + Ok(()) +} + +#[tokio::test] +async fn sessions_enforce_independent_yield_limits() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let provider = GrpcCodeModeSessionProvider::new(host.endpoint); + let limited = provider + .create_session_with_limits( + Arc::new(NoopCodeModeSessionDelegate), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(/*value*/ 1), + max_heap_size_bytes: Some(/*value*/ 16 * 1024 * 1024), + }, + ) + .await + .map_err(anyhow::Error::msg)?; + let other = provider + .create_session_with_limits( + Arc::new(NoopCodeModeSessionDelegate), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(/*value*/ 1_000), + max_heap_size_bytes: None, + }, + ) + .await + .map_err(anyhow::Error::msg)?; + + let mut pending = request("await new Promise(() => {});"); + pending.yield_time_ms = Some(/*value*/ 60_000); + assert_eq!( + execute(&limited, pending).await?, + RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + } + ); + assert_eq!( + timeout( + TEST_TIMEOUT, + limited.wait(WaitRequest { + cell_id: cell_id("1"), + yield_time_ms: 60_000, + }), + ) + .await + .context("session yield limit did not bound an explicit wait")? + .map_err(anyhow::Error::msg)?, + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + }) + ); + assert_eq!( + execute( + &other, + request(r#"await new Promise(resolve => setTimeout(resolve, 25)); text("isolated");"#), + ) + .await?, + text_response("1", "isolated") + ); + assert_eq!( + limited + .terminate(cell_id("1")) + .await + .map_err(anyhow::Error::msg)?, + WaitOutcome::LiveCell(RuntimeResponse::Terminated { + cell_id: cell_id("1"), + content_items: Vec::new(), + }) + ); + assert_eq!( + execute(&limited, request(r#"text("recovered");"#)).await?, + text_response("2", "recovered") + ); + limited.shutdown().await.map_err(anyhow::Error::msg)?; + other.shutdown().await.map_err(anyhow::Error::msg)?; + Ok(()) +} + +#[tokio::test] +async fn dropping_a_grpc_lease_retires_its_server_session() -> Result<()> { + let host = HostHarness::start("grpc://127.0.0.1:0").await?; + let mut client = CodeModeHostClient::connect(host.endpoint) + .await + .context("connect raw gRPC client")?; + let mut lease = client + .open_session(grpc::OpenSessionRequest { + cell_execution_limits: None, + }) + .await + .context("open raw gRPC session")? + .into_inner(); + let first = lease + .message() + .await + .context("read raw gRPC session opening")? + .context("raw gRPC session ended before its opening event")?; + let Some(grpc::session_event::Event::Opened(opened)) = first.event else { + anyhow::bail!("raw gRPC session did not start with an opening event"); + }; + drop(lease); + + timeout(TEST_TIMEOUT, async { + loop { + match client + .subscribe_to_tool_calls(grpc::SubscribeToToolCallsRequest { + session_id: opened.session_id.clone(), + tool_names: Vec::new(), + }) + .await + { + Ok(response) => drop(response), + Err(status) if status.code() == Code::NotFound => return Ok::<_, anyhow::Error>(()), + Err(status) => { + anyhow::bail!("unexpected session status after lease drop: {status}") + } + } + tokio::task::yield_now().await; + } + }) + .await + .context("dropping the gRPC lease did not retire its server session")??; + Ok(()) +} diff --git a/codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto b/codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto index 25db26de3b..ece24730cc 100644 --- a/codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto +++ b/codex-rs/code-mode-protocol/src/grpc/codex.code_mode.v1.proto @@ -134,7 +134,8 @@ message CellClosed { string execution_id = 1; string cell_id = 2; - // All lower sequences must be observed before the client retires the cell. + // Last tool-call sequence issued before closure. Clients may retire the cell + // immediately and reject tool calls delivered after its closure. uint64 final_tool_call_sequence = 3; } diff --git a/codex-rs/code-mode-protocol/src/grpc/mod.rs b/codex-rs/code-mode-protocol/src/grpc/mod.rs index 566c5cf969..08c88badea 100644 --- a/codex-rs/code-mode-protocol/src/grpc/mod.rs +++ b/codex-rs/code-mode-protocol/src/grpc/mod.rs @@ -3,3 +3,5 @@ pub use code_mode_proto::codex::code_mode::v1::*; #[cfg(not(codex_bazel))] tonic::include_proto!("codex.code_mode.v1"); + +pub const MAX_IDENTIFIER_BYTES: usize = 256; diff --git a/codex-rs/code-mode-protocol/src/session.rs b/codex-rs/code-mode-protocol/src/session.rs index 617aed5969..ffba25005c 100644 --- a/codex-rs/code-mode-protocol/src/session.rs +++ b/codex-rs/code-mode-protocol/src/session.rs @@ -62,27 +62,31 @@ pub struct StartedCell { impl StartedCell { pub fn new(cell_id: CellId, initial_response_rx: oneshot::Receiver) -> Self { - Self { - cell_id, - initial_response: Box::pin(async move { - initial_response_rx - .await - .map_err(|_| "exec runtime ended unexpectedly".to_string()) - }), - } + Self::from_future(cell_id, async move { + initial_response_rx + .await + .map_err(|_| "exec runtime ended unexpectedly".to_string()) + }) } pub fn from_result_receiver( cell_id: CellId, initial_response_rx: oneshot::Receiver>, + ) -> Self { + Self::from_future(cell_id, async move { + initial_response_rx + .await + .map_err(|_| "exec runtime ended unexpectedly".to_string())? + }) + } + + pub fn from_future( + cell_id: CellId, + initial_response: impl Future> + Send + 'static, ) -> Self { Self { cell_id, - initial_response: Box::pin(async move { - initial_response_rx - .await - .map_err(|_| "exec runtime ended unexpectedly".to_string())? - }), + initial_response: Box::pin(initial_response), } } diff --git a/codex-rs/code-mode/Cargo.toml b/codex-rs/code-mode/Cargo.toml index 80121599fa..e669d7e9ba 100644 --- a/codex-rs/code-mode/Cargo.toml +++ b/codex-rs/code-mode/Cargo.toml @@ -19,10 +19,13 @@ codex-install-context = { workspace = true } codex-protocol = { workspace = true } codex-websocket-client = { workspace = true } futures = { workspace = true } +serde_json = { workspace = true } tokio = { workspace = true, features = ["io-util", "macros", "net", "process", "rt", "sync", "time"] } tokio-tungstenite = { workspace = true } tokio-util = { workspace = true, features = ["rt"] } +tonic = { workspace = true } tracing = { workspace = true } +uuid = { workspace = true, features = ["v4"] } [dev-dependencies] pretty_assertions = { workspace = true } diff --git a/codex-rs/code-mode/src/grpc_session/callbacks.rs b/codex-rs/code-mode/src/grpc_session/callbacks.rs new file mode 100644 index 0000000000..0a87ffd3e5 --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/callbacks.rs @@ -0,0 +1,70 @@ +use std::sync::Arc; +use std::sync::PoisonError; +use std::sync::atomic::Ordering; + +use codex_code_mode_protocol::grpc; + +use super::SessionInner; + +impl SessionInner { + pub(super) fn spawn_session_events( + self: &Arc, + mut events: tonic::Streaming, + ) { + let inner = Arc::clone(self); + self.stream_tasks.spawn(async move { + loop { + let event = tokio::select! { + biased; + _ = inner.stopped.cancelled() => return, + event = events.message() => event, + }; + match event { + Ok(Some(event)) => { + if let Err(error) = inner.handle_session_event(event) { + inner.fail(error); + return; + } + } + Ok(None) => { + if !inner.shutdown_requested.load(Ordering::Acquire) { + inner.fail( + "gRPC code-mode session lease closed unexpectedly".to_string(), + ); + } + return; + } + Err(error) => { + if !inner.shutdown_requested.load(Ordering::Acquire) { + inner.fail(super::deadline::failure("session lease", error)); + } + return; + } + } + } + }); + } + + fn handle_session_event(&self, event: grpc::SessionEvent) -> Result<(), String> { + match event + .event + .ok_or_else(|| "gRPC code-mode host sent an empty session event".to_string())? + { + grpc::session_event::Event::Opened(_) => { + Err("gRPC code-mode host repeated the session opening event".to_string()) + } + grpc::session_event::Event::ToolCallCancelled(_) + | grpc::session_event::Event::Notification(_) + | grpc::session_event::Event::NotificationCancelled(_) => Ok(()), + grpc::session_event::Event::CellClosed(closed) => { + let cell = self + .state + .lock() + .unwrap_or_else(PoisonError::into_inner) + .close_cell(closed)?; + self.report_closed_cell(cell); + Ok(()) + } + } + } +} diff --git a/codex-rs/code-mode/src/grpc_session/conversion.rs b/codex-rs/code-mode/src/grpc_session/conversion.rs new file mode 100644 index 0000000000..9c93346e04 --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/conversion.rs @@ -0,0 +1,134 @@ +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeToolKind; +use codex_code_mode_protocol::ExecuteRequest; +use codex_code_mode_protocol::FunctionCallOutputContentItem; +use codex_code_mode_protocol::ImageDetail; +use codex_code_mode_protocol::RuntimeResponse; +use codex_code_mode_protocol::ToolDefinition; +use codex_code_mode_protocol::WaitOutcome; +use codex_code_mode_protocol::grpc; + +pub(super) fn execute_request( + session_id: &str, + execution_id: String, + request: ExecuteRequest, +) -> Result { + Ok(grpc::ExecuteRequest { + session_id: session_id.to_string(), + execution_id, + tool_call_id: request.tool_call_id, + source: request.source, + enabled_tools: request + .enabled_tools + .into_iter() + .map(tool_definition) + .collect::, _>>()?, + yield_time_ms: request.yield_time_ms, + max_output_tokens: request + .max_output_tokens + .map(u64::try_from) + .transpose() + .map_err(|error| format!("invalid code-mode output token limit: {error}"))?, + }) +} + +fn tool_definition(definition: ToolDefinition) -> Result { + Ok(grpc::ToolDefinition { + name: definition.name, + tool_name: Some(grpc::ToolName { + name: definition.tool_name.name, + namespace: definition.tool_name.namespace, + }), + description: definition.description, + kind: match definition.kind { + CodeModeToolKind::Function => grpc::ToolKind::Function as i32, + CodeModeToolKind::Freeform => grpc::ToolKind::Freeform as i32, + }, + input_schema_json: definition + .input_schema + .map(|schema| serde_json::to_vec(&schema)) + .transpose() + .map_err(|error| format!("failed to encode code-mode tool input schema: {error}"))?, + output_schema_json: definition + .output_schema + .map(|schema| serde_json::to_vec(&schema)) + .transpose() + .map_err(|error| format!("failed to encode code-mode tool output schema: {error}"))?, + }) +} + +pub(super) fn runtime_response(outcome: grpc::ExecutionOutcome) -> Result { + super::validate_identifier(&outcome.cell_id, "cell ID")?; + let cell_id = CellId::new(outcome.cell_id); + let content_items = outcome + .content_items + .into_iter() + .map(content_item) + .collect::, _>>()?; + match outcome + .outcome + .ok_or_else(|| "code-mode execution omitted its outcome".to_string())? + { + grpc::execution_outcome::Outcome::Yielded(_) => Ok(RuntimeResponse::Yielded { + cell_id, + content_items, + }), + grpc::execution_outcome::Outcome::Terminated(_) => Ok(RuntimeResponse::Terminated { + cell_id, + content_items, + }), + grpc::execution_outcome::Outcome::Completed(completed) => Ok(RuntimeResponse::Result { + cell_id, + content_items, + error_text: completed.error_text, + }), + } +} + +pub(super) fn wait_outcome(response: grpc::WaitResponse) -> Result { + match response + .state + .ok_or_else(|| "code-mode wait omitted its outcome".to_string())? + { + grpc::wait_response::State::LiveCell(response) => { + runtime_response(response).map(WaitOutcome::LiveCell) + } + grpc::wait_response::State::MissingCell(response) => { + runtime_response(response).map(WaitOutcome::MissingCell) + } + } +} + +fn content_item(item: grpc::ContentItem) -> Result { + match item + .item + .ok_or_else(|| "code-mode execution returned an empty content item".to_string())? + { + grpc::content_item::Item::Text(text) => { + Ok(FunctionCallOutputContentItem::InputText { text: text.text }) + } + grpc::content_item::Item::Image(image) => Ok(FunctionCallOutputContentItem::InputImage { + image_url: image.image_url, + detail: image.detail.map(image_detail).transpose()?, + }), + grpc::content_item::Item::Audio(audio) => Ok(FunctionCallOutputContentItem::InputAudio { + audio_url: audio.audio_url, + }), + } +} + +fn image_detail(value: i32) -> Result { + match grpc::ImageDetail::try_from(value) { + Ok(grpc::ImageDetail::Auto) => Ok(ImageDetail::Auto), + Ok(grpc::ImageDetail::Low) => Ok(ImageDetail::Low), + Ok(grpc::ImageDetail::High) => Ok(ImageDetail::High), + Ok(grpc::ImageDetail::Original) => Ok(ImageDetail::Original), + Ok(grpc::ImageDetail::Unspecified) | Err(_) => { + Err(format!("code-mode image has invalid detail value {value}")) + } + } +} + +#[cfg(test)] +#[path = "conversion_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode/src/grpc_session/conversion_tests.rs b/codex-rs/code-mode/src/grpc_session/conversion_tests.rs new file mode 100644 index 0000000000..33fd34e1a3 --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/conversion_tests.rs @@ -0,0 +1,168 @@ +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeToolKind; +use codex_code_mode_protocol::ExecuteRequest; +use codex_code_mode_protocol::FunctionCallOutputContentItem; +use codex_code_mode_protocol::ImageDetail; +use codex_code_mode_protocol::RuntimeResponse; +use codex_code_mode_protocol::ToolDefinition; +use codex_code_mode_protocol::WaitOutcome; +use codex_code_mode_protocol::grpc; +use codex_protocol::ToolName; +use pretty_assertions::assert_eq; +use serde_json::json; + +use super::execute_request; +use super::runtime_response; +use super::wait_outcome; + +#[test] +fn execute_request_preserves_tool_schemas_namespaces_and_limits() { + let request = ExecuteRequest { + tool_call_id: "outer".to_string(), + enabled_tools: vec![ToolDefinition { + name: "search".to_string(), + tool_name: ToolName::namespaced("work", "search"), + description: "search the workspace".to_string(), + kind: CodeModeToolKind::Freeform, + input_schema: Some(json!({"type": "object"})), + output_schema: Some(json!({"type": "string"})), + }], + source: "text('hello')".to_string(), + yield_time_ms: Some(25), + max_output_tokens: Some(128), + }; + + assert_eq!( + execute_request("session", "execution".to_string(), request), + Ok(grpc::ExecuteRequest { + session_id: "session".to_string(), + execution_id: "execution".to_string(), + tool_call_id: "outer".to_string(), + source: "text('hello')".to_string(), + enabled_tools: vec![grpc::ToolDefinition { + name: "search".to_string(), + tool_name: Some(grpc::ToolName { + name: "search".to_string(), + namespace: Some("work".to_string()), + }), + description: "search the workspace".to_string(), + kind: grpc::ToolKind::Freeform as i32, + input_schema_json: Some(br#"{"type":"object"}"#.to_vec()), + output_schema_json: Some(br#"{"type":"string"}"#.to_vec()), + }], + yield_time_ms: Some(25), + max_output_tokens: Some(128), + }) + ); +} + +#[test] +fn runtime_response_decodes_mixed_content_items() { + let outcome = grpc::ExecutionOutcome { + cell_id: "cell".to_string(), + content_items: vec![ + grpc::ContentItem { + item: Some(grpc::content_item::Item::Text(grpc::TextContent { + text: "hello".to_string(), + })), + }, + grpc::ContentItem { + item: Some(grpc::content_item::Item::Image(grpc::ImageContent { + image_url: "data:image/png;base64,AA==".to_string(), + detail: Some(grpc::ImageDetail::Original as i32), + })), + }, + grpc::ContentItem { + item: Some(grpc::content_item::Item::Audio(grpc::AudioContent { + audio_url: "data:audio/wav;base64,AA==".to_string(), + })), + }, + ], + outcome: Some(grpc::execution_outcome::Outcome::Completed( + grpc::ExecutionCompleted { + error_text: Some("warning".to_string()), + }, + )), + }; + + assert_eq!( + runtime_response(outcome), + Ok(RuntimeResponse::Result { + cell_id: CellId::new("cell".to_string()), + content_items: vec![ + FunctionCallOutputContentItem::InputText { + text: "hello".to_string(), + }, + FunctionCallOutputContentItem::InputImage { + image_url: "data:image/png;base64,AA==".to_string(), + detail: Some(ImageDetail::Original), + }, + FunctionCallOutputContentItem::InputAudio { + audio_url: "data:audio/wav;base64,AA==".to_string(), + }, + ], + error_text: Some("warning".to_string()), + }) + ); +} + +#[test] +fn wait_outcome_preserves_missing_cell_state() { + let response = grpc::WaitResponse { + state: Some(grpc::wait_response::State::MissingCell( + grpc::ExecutionOutcome { + cell_id: "missing".to_string(), + content_items: Vec::new(), + outcome: Some(grpc::execution_outcome::Outcome::Terminated( + grpc::ExecutionTerminated {}, + )), + }, + )), + }; + + assert_eq!( + wait_outcome(response), + Ok(WaitOutcome::MissingCell(RuntimeResponse::Terminated { + cell_id: CellId::new("missing".to_string()), + content_items: Vec::new(), + })) + ); +} + +#[test] +fn oversized_response_cell_ids_are_rejected() { + let response = grpc::ExecutionOutcome { + cell_id: "x".repeat(grpc::MAX_IDENTIFIER_BYTES + 1), + content_items: Vec::new(), + outcome: Some(grpc::execution_outcome::Outcome::Yielded( + grpc::ExecutionYielded {}, + )), + }; + + assert_eq!( + runtime_response(response), + Err(format!( + "gRPC code-mode host returned cell ID exceeding {} bytes", + grpc::MAX_IDENTIFIER_BYTES + )) + ); +} + +#[test] +fn invalid_output_enums_and_missing_oneofs_are_rejected() { + let invalid_image = grpc::ExecutionOutcome { + cell_id: "cell".to_string(), + content_items: vec![grpc::ContentItem { + item: Some(grpc::content_item::Item::Image(grpc::ImageContent { + image_url: "image".to_string(), + detail: Some(grpc::ImageDetail::Unspecified as i32), + })), + }], + outcome: Some(grpc::execution_outcome::Outcome::Yielded( + grpc::ExecutionYielded {}, + )), + }; + + assert!(runtime_response(invalid_image).is_err()); + assert!(wait_outcome(grpc::WaitResponse { state: None }).is_err()); +} diff --git a/codex-rs/code-mode/src/grpc_session/deadline.rs b/codex-rs/code-mode/src/grpc_session/deadline.rs new file mode 100644 index 0000000000..266b1ed5e5 --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/deadline.rs @@ -0,0 +1,77 @@ +use std::fmt::Display; +use std::future::Future; +use std::time::Duration; + +use super::SessionInner; + +const TRANSPORT_TIMEOUT: Duration = Duration::from_secs(60); +const MAX_ERROR_BYTES: usize = 512; + +pub(super) async fn startup( + operation: &str, + request: impl Future>, +) -> Result { + match enforce(operation, Duration::ZERO, request).await { + Ok(result) => Ok(result), + Err(RequestError::Failed(error)) => Err(failure(operation, error)), + Err(RequestError::TimedOut(reason)) => Err(reason), + } +} + +pub(super) async fn request( + session: &SessionInner, + operation: &str, + runtime_timeout: Duration, + request: impl Future>, +) -> Result { + let result = tokio::select! { + biased; + _ = session.stopped.cancelled() => { + return Err("gRPC code-mode session closed".to_string()); + } + result = enforce(operation, runtime_timeout, request) => result, + }; + + match result { + Ok(value) => Ok(value), + Err(RequestError::Failed(error)) => Err(failure(operation, error)), + Err(RequestError::TimedOut(reason)) => { + session.fail(reason.clone()); + Err(reason) + } + } +} + +pub(super) fn failure(operation: &str, error: impl Display) -> String { + let mut message = format!("gRPC code-mode {operation} failed: {error}"); + if message.len() > MAX_ERROR_BYTES { + let boundary = message.floor_char_boundary(MAX_ERROR_BYTES - "...".len()); + message.truncate(boundary); + message.push_str("..."); + } + message +} + +async fn enforce( + operation: &str, + runtime_timeout: Duration, + request: impl Future>, +) -> Result> { + let timeout = runtime_timeout.saturating_add(TRANSPORT_TIMEOUT); + match tokio::time::timeout(timeout, request).await { + Ok(Ok(value)) => Ok(value), + Ok(Err(error)) => Err(RequestError::Failed(error)), + Err(_) => Err(RequestError::TimedOut(format!( + "gRPC code-mode host timed out waiting for {operation} response" + ))), + } +} + +enum RequestError { + Failed(E), + TimedOut(String), +} + +#[cfg(test)] +#[path = "deadline_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode/src/grpc_session/deadline_tests.rs b/codex-rs/code-mode/src/grpc_session/deadline_tests.rs new file mode 100644 index 0000000000..1d034af7b1 --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/deadline_tests.rs @@ -0,0 +1,131 @@ +use std::future::pending; +use std::sync::Arc; +use std::time::Duration; + +use codex_code_mode_protocol::DEFAULT_EXEC_YIELD_TIME_MS; +use pretty_assertions::assert_eq; + +use super::MAX_ERROR_BYTES; +use super::RequestError; +use super::enforce; +use super::startup; + +#[tokio::test(start_paused = true)] +async fn stalled_transport_fails_after_its_deadline() { + let task = tokio::spawn(enforce( + "termination", + Duration::ZERO, + pending::>(), + )); + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(61)).await; + + assert!(matches!( + task.await.expect("deadline task"), + Err(RequestError::TimedOut(message)) + if message == "gRPC code-mode host timed out waiting for termination response" + )); +} + +#[tokio::test(start_paused = true)] +async fn requested_runtime_duration_is_added_to_the_transport_deadline() { + let task = tokio::spawn(enforce( + "wait", + Duration::from_secs(120), + pending::>(), + )); + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(61)).await; + tokio::task::yield_now().await; + assert!(!task.is_finished()); + + tokio::time::advance(Duration::from_secs(120)).await; + assert!(matches!( + task.await.expect("deadline task"), + Err(RequestError::TimedOut(_)) + )); +} + +#[tokio::test(start_paused = true)] +async fn default_execution_yield_and_grace_extend_the_outcome_deadline() { + let runtime_timeout = + Duration::from_millis(DEFAULT_EXEC_YIELD_TIME_MS).saturating_add(Duration::from_secs(1)); + let task = tokio::spawn(enforce( + "execution outcome", + runtime_timeout, + pending::>(), + )); + tokio::task::yield_now().await; + + tokio::time::advance(Duration::from_secs(70)).await; + tokio::task::yield_now().await; + assert!(!task.is_finished()); + + tokio::time::advance(Duration::from_secs(2)).await; + assert!(matches!( + task.await.expect("execution outcome deadline task"), + Err(RequestError::TimedOut(message)) + if message == "gRPC code-mode host timed out waiting for execution outcome response" + )); +} + +#[tokio::test] +async fn transport_status_is_preserved() { + let result = enforce("wait", Duration::ZERO, async { + Err::<(), _>(tonic::Status::not_found("missing")) + }) + .await; + + match result { + Err(RequestError::Failed(error)) => { + assert_eq!(error.code(), tonic::Code::NotFound); + assert_eq!(error.message(), "missing"); + } + _ => panic!("expected the original gRPC status"), + } +} + +#[tokio::test] +async fn transport_status_messages_are_bounded_at_utf8_boundaries() { + let error = startup("session opening", async { + Err::<(), _>(tonic::Status::internal("🦀".repeat(MAX_ERROR_BYTES))) + }) + .await + .expect_err("oversized gRPC status must fail"); + + assert!(error.len() <= MAX_ERROR_BYTES); + assert!(error.starts_with("gRPC code-mode session opening failed:")); + assert!(error.ends_with("...")); +} + +#[tokio::test(start_paused = true)] +async fn stalled_channel_acquisition_times_out_and_remains_retryable() { + let channel = Arc::new(tokio::sync::OnceCell::new()); + let stalled_channel = Arc::clone(&channel); + let stalled = tokio::spawn(async move { + startup("transport connection", async { + stalled_channel + .get_or_try_init(pending::>) + .await + .copied() + }) + .await + }); + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(61)).await; + + assert_eq!( + stalled.await.expect("channel connection task"), + Err("gRPC code-mode host timed out waiting for transport connection response".to_string()) + ); + assert_eq!( + startup("transport connection", async { + channel + .get_or_try_init(|| async { Ok::<_, String>(42usize) }) + .await + .copied() + }) + .await, + Ok(42) + ); +} diff --git a/codex-rs/code-mode/src/grpc_session/mod.rs b/codex-rs/code-mode/src/grpc_session/mod.rs new file mode 100644 index 0000000000..cfa9a1c8ac --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/mod.rs @@ -0,0 +1,317 @@ +use std::collections::HashMap; +use std::panic::AssertUnwindSafe; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::PoisonError; +use std::sync::Weak; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::time::Duration; + +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::CodeModeSessionProvider; +use codex_code_mode_protocol::CodeModeSessionProviderFuture; +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 codex_code_mode_protocol::grpc; +use codex_code_mode_protocol::grpc::code_mode_host_client::CodeModeHostClient; +use codex_code_mode_protocol::host::MAX_FRAME_BYTES; +use tokio::sync::watch; +use tokio_util::sync::CancellationToken; +use tokio_util::task::TaskTracker; +use tonic::transport::Channel; + +use self::operations::WaitSlot; +use self::state::SessionState; +use self::transport::SharedTransport; +use crate::remote_session::ShutdownResultReceiver; +use crate::remote_session::wait_for_watch; + +mod callbacks; +mod conversion; +mod deadline; +mod operations; +mod state; +mod transport; + +type GrpcClient = CodeModeHostClient; + +/// Creates code-mode sessions over an HTTP/2 gRPC connection. +#[derive(Clone)] +pub struct GrpcCodeModeSessionProvider { + transport: Arc, +} + +impl GrpcCodeModeSessionProvider { + /// Connects lazily to an `http://` gRPC endpoint. + pub fn new(endpoint: impl Into) -> Self { + Self::from_transport(SharedTransport::new(endpoint.into())) + } + + /// Uses an existing channel, including channels backed by custom transports. + pub fn with_channel(channel: Channel) -> Self { + Self::from_transport(SharedTransport::with_channel(channel)) + } + + fn from_transport(transport: SharedTransport) -> Self { + Self { + transport: Arc::new(transport), + } + } + + async fn open_binding( + &self, + delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, + ) -> Result, String> { + let channel = deadline::startup("transport connection", self.transport.channel()).await?; + let mut client = grpc_client(channel); + let limits = grpc::SessionCellExecutionLimits { + max_yield_time_ms: limits.max_yield_time_ms, + max_heap_size_bytes: limits + .max_heap_size_bytes + .map(u64::try_from) + .transpose() + .map_err(|error| format!("invalid code-mode heap size limit: {error}"))?, + }; + let cell_execution_limits = (limits.max_yield_time_ms.is_some() + || limits.max_heap_size_bytes.is_some()) + .then_some(limits); + let mut lease = deadline::startup( + "session opening", + client.open_session(grpc::OpenSessionRequest { + cell_execution_limits, + }), + ) + .await? + .into_inner(); + let first = deadline::startup("session lease opening", lease.message()) + .await? + .ok_or_else(|| "gRPC code-mode session lease ended before opening".to_string())?; + let Some(grpc::session_event::Event::Opened(opened)) = first.event else { + return Err("gRPC code-mode session lease omitted its opening event".to_string()); + }; + validate_identifier(&opened.session_id, "session ID")?; + + let inner = Arc::new(SessionInner { + id: opened.session_id, + client, + delegate, + runtime: tokio::runtime::Handle::current(), + state: Mutex::new(SessionState::default()), + wait_slots: Mutex::new(HashMap::new()), + shutdown_requested: AtomicBool::new(false), + shutdown_result: Mutex::new(None), + stopped: CancellationToken::new(), + stream_tasks: TaskTracker::new(), + _transport: Arc::clone(&self.transport), + }); + let mut opening = OpeningSession { + inner: Some(Arc::clone(&inner)), + }; + inner.spawn_session_events(lease); + inner.require_open()?; + opening.inner = None; + Ok(Arc::new(GrpcCodeModeSession { inner })) + } +} + +impl CodeModeSessionProvider for GrpcCodeModeSessionProvider { + fn create_session<'a>( + &'a self, + delegate: Arc, + ) -> CodeModeSessionProviderFuture<'a> { + self.create_session_with_limits(delegate, CodeModeSessionCellExecutionLimits::default()) + } + + fn create_session_with_limits<'a>( + &'a self, + delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, + ) -> CodeModeSessionProviderFuture<'a> { + Box::pin(async move { + self.open_binding(delegate, limits) + .await + .map(|session| session as _) + }) + } +} + +struct GrpcCodeModeSession { + inner: Arc, +} + +struct OpeningSession { + inner: Option>, +} + +impl Drop for OpeningSession { + fn drop(&mut self) { + let Some(inner) = self.inner.take() else { + return; + }; + if tokio::runtime::Handle::try_current().is_ok() { + inner.request_shutdown(); + } else { + inner.close_state(/*failure*/ None); + } + } +} + +impl CodeModeSession for GrpcCodeModeSession { + fn execute<'a>( + &'a self, + request: ExecuteRequest, + ) -> CodeModeSessionResultFuture<'a, StartedCell> { + Box::pin(self.inner.execute(request)) + } + + fn wait<'a>(&'a self, request: WaitRequest) -> CodeModeSessionResultFuture<'a, WaitOutcome> { + Box::pin(self.inner.wait(request)) + } + + fn terminate<'a>(&'a self, cell_id: CellId) -> CodeModeSessionResultFuture<'a, WaitOutcome> { + Box::pin(self.inner.terminate(cell_id)) + } + + fn shutdown<'a>(&'a self) -> CodeModeSessionResultFuture<'a, ()> { + Box::pin(wait_for_watch(self.inner.request_shutdown())) + } +} + +impl Drop for GrpcCodeModeSession { + fn drop(&mut self) { + self.inner.request_shutdown(); + } +} + +pub(super) struct SessionInner { + pub(super) id: String, + pub(super) client: GrpcClient, + pub(super) delegate: Arc, + runtime: tokio::runtime::Handle, + state: Mutex, + wait_slots: Mutex>>, + shutdown_requested: AtomicBool, + shutdown_result: Mutex>, + pub(super) stopped: CancellationToken, + stream_tasks: TaskTracker, + _transport: Arc, +} + +impl SessionInner { + pub(super) fn client(&self) -> GrpcClient { + self.client.clone() + } + + pub(super) fn require_open(&self) -> Result<(), String> { + self.state + .lock() + .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()); + } + Ok(()) + } + + pub(super) fn report_closed_cell(&self, cell_id: Option) { + if let Some(cell_id) = cell_id { + self.wait_slots + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(&cell_id); + let _ = std::panic::catch_unwind(AssertUnwindSafe(|| { + self.delegate.cell_closed(&cell_id); + })); + } + } + + pub(super) fn fail(&self, reason: String) { + self.close_state(Some(reason)); + } + + fn close_state(&self, failure: Option) { + let cells = self + .state + .lock() + .unwrap_or_else(PoisonError::into_inner) + .close(failure); + self.stopped.cancel(); + self.stream_tasks.close(); + for cell_id in cells { + self.report_closed_cell(Some(cell_id)); + } + } + + fn request_shutdown(self: &Arc) -> ShutdownResultReceiver { + self.shutdown_requested.store(true, Ordering::Release); + 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); + self.runtime.spawn(async move { + let result = inner.drive_shutdown().await; + sender.send_replace(Some(result)); + }); + receiver + } + + async fn drive_shutdown(&self) -> Result<(), String> { + let is_open = self + .state + .lock() + .unwrap_or_else(PoisonError::into_inner) + .require_open() + .is_ok(); + let result = if is_open { + let mut client = self.client(); + deadline::request( + self, + "session shutdown", + Duration::ZERO, + client.close_session(grpc::CloseSessionRequest { + session_id: self.id.clone(), + }), + ) + .await + .map(|_| ()) + } else { + Ok(()) + }; + self.close_state(/*failure*/ None); + self.stream_tasks.wait().await; + result + } +} + +fn validate_identifier(value: &str, field: &str) -> Result<(), String> { + if value.is_empty() { + return Err(format!("gRPC code-mode host returned an empty {field}")); + } + if value.len() > grpc::MAX_IDENTIFIER_BYTES { + return Err(format!( + "gRPC code-mode host returned {field} exceeding {} bytes", + grpc::MAX_IDENTIFIER_BYTES + )); + } + Ok(()) +} + +fn grpc_client(channel: Channel) -> GrpcClient { + CodeModeHostClient::new(channel) + .max_decoding_message_size(MAX_FRAME_BYTES) + .max_encoding_message_size(MAX_FRAME_BYTES) +} diff --git a/codex-rs/code-mode/src/grpc_session/operations.rs b/codex-rs/code-mode/src/grpc_session/operations.rs new file mode 100644 index 0000000000..7513e06c26 --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/operations.rs @@ -0,0 +1,391 @@ +use std::sync::Arc; +use std::sync::PoisonError; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::DEFAULT_EXEC_YIELD_TIME_MS; +use codex_code_mode_protocol::ExecuteRequest; +use codex_code_mode_protocol::RuntimeResponse; +use codex_code_mode_protocol::StartedCell; +use codex_code_mode_protocol::WaitOutcome; +use codex_code_mode_protocol::WaitRequest; +use codex_code_mode_protocol::grpc; +use tokio::sync::OwnedMutexGuard; +use tokio::sync::oneshot; +use tracing::debug; +use uuid::Uuid; + +use super::SessionInner; +use super::conversion; +use super::deadline; + +pub(super) struct WaitSlot { + lock: Arc>, + active: AtomicBool, +} + +struct ExecutionOwnership { + session: Arc, + execution_id: String, + armed: bool, +} + +impl Drop for ExecutionOwnership { + fn drop(&mut self) { + if !self.armed { + return; + } + let cell = self + .session + .state + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove_execution(&self.execution_id); + let Some(cell_id) = cell else { + return; + }; + self.session.report_closed_cell(Some(cell_id.clone())); + if self.session.stopped.is_cancelled() { + return; + } + let session = Arc::clone(&self.session); + self.session.runtime.spawn(async move { + if let Err(error) = session.terminate(cell_id).await + && !session.stopped.is_cancelled() + { + debug!("abandoned code-mode execution termination raced closure: {error}"); + } + }); + } +} + +impl SessionInner { + pub(super) async fn execute( + self: &Arc, + request: ExecuteRequest, + ) -> Result { + self.require_open()?; + let execution_id = Uuid::new_v4().to_string(); + let request = conversion::execute_request(&self.id, execution_id.clone(), request)?; + self.state + .lock() + .unwrap_or_else(PoisonError::into_inner) + .begin_execution(execution_id.clone())?; + let ownership = ExecutionOwnership { + session: Arc::clone(self), + execution_id, + armed: true, + }; + let (started_tx, started_rx) = oneshot::channel(); + let inner = Arc::clone(self); + self.stream_tasks.spawn(async move { + inner.drive_execution(request, ownership, started_tx).await; + }); + started_rx + .await + .map_err(|_| "gRPC code-mode execution driver ended unexpectedly".to_string())? + } + + async fn drive_execution( + self: Arc, + request: grpc::ExecuteRequest, + ownership: ExecutionOwnership, + started_tx: oneshot::Sender>, + ) { + let runtime_timeout = + Duration::from_millis(request.yield_time_ms.unwrap_or(DEFAULT_EXEC_YIELD_TIME_MS)) + .saturating_add(Duration::from_secs(1)); + let opening = async { + let mut client = self.client(); + let mut stream = + deadline::request(&self, "execution", Duration::ZERO, client.execute(request)) + .await? + .into_inner(); + let first = deadline::request( + &self, + "execution starting event", + Duration::ZERO, + stream.message(), + ) + .await? + .ok_or_else(|| { + "gRPC code-mode execution ended before its starting event".to_string() + })?; + let Some(grpc::execute_event::Event::Started(started)) = first.event else { + return Err("gRPC code-mode execution omitted its starting event".to_string()); + }; + super::validate_identifier(&started.execution_id, "execution ID")?; + if started.execution_id != ownership.execution_id { + let error = format!( + "gRPC code-mode execution returned ID {} instead of {}", + started.execution_id, ownership.execution_id + ); + self.fail(error.clone()); + return Err(error); + } + + let admission = self + .state + .lock() + .unwrap_or_else(PoisonError::into_inner) + .admit_execution(&ownership.execution_id, &started.cell_id); + if let Err(error) = admission { + self.fail(error.clone()); + return Err(error); + } + + Ok((CellId::new(started.cell_id), stream)) + } + .await; + let (cell_id, stream) = match opening { + Ok(opening) => opening, + Err(error) => { + let _ = started_tx.send(Err(error)); + return; + } + }; + let (response_tx, response_rx) = oneshot::channel(); + let mut claim = ownership; + let started = StartedCell::from_future(cell_id.clone(), async move { + let closure = claim + .session + .state + .lock() + .unwrap_or_else(PoisonError::into_inner) + .mark_execution_ready(&claim.execution_id); + match closure { + Ok(cell) => claim.session.report_closed_cell(cell), + Err(error) => { + claim.session.fail(error.clone()); + return Err(error); + } + } + let response = response_rx + .await + .map_err(|_| "exec runtime ended unexpectedly".to_string())??; + claim.armed = false; + drop(claim); + Ok(response) + }); + if started_tx.send(Ok(started)).is_err() { + return; + } + self.drive_execution_outcome(cell_id, stream, response_tx, runtime_timeout) + .await; + } + + async fn drive_execution_outcome( + self: Arc, + cell_id: CellId, + mut stream: tonic::Streaming, + mut response_tx: oneshot::Sender>, + runtime_timeout: Duration, + ) { + let outcome = tokio::select! { + biased; + _ = response_tx.closed() => return, + outcome = deadline::request( + &self, + "execution outcome", + runtime_timeout, + stream.message(), + ) => match outcome { + Ok(Some(grpc::ExecuteEvent { + event: Some(grpc::execute_event::Event::Outcome(outcome)), + })) => conversion::runtime_response(outcome), + Ok(Some(_)) => { + Err("gRPC code-mode execution returned an unexpected event".to_string()) + } + Ok(None) => Err("gRPC code-mode execution omitted its initial outcome".to_string()), + Err(error) => Err(error), + }, + }; + let outcome = outcome.and_then(|response| { + if runtime_response_cell_id(&response) != &cell_id { + let error = format!( + "gRPC code-mode execution returned cell {} instead of {cell_id}", + runtime_response_cell_id(&response) + ); + self.fail(error.clone()); + Err(error) + } else { + Ok(response) + } + }); + let _ = response_tx.send(outcome); + } + + pub(super) async fn wait( + self: &Arc, + request: WaitRequest, + ) -> Result { + self.require_open()?; + let slot = { + let mut slots = self + .wait_slots + .lock() + .unwrap_or_else(PoisonError::into_inner); + slots.retain(|_, slot| slot.strong_count() != 0); + let slot = match slots + .get(&request.cell_id) + .and_then(std::sync::Weak::upgrade) + { + Some(slot) => slot, + None => { + let slot = Arc::new(WaitSlot { + lock: Arc::new(tokio::sync::Mutex::new(())), + active: AtomicBool::new(false), + }); + slots.insert(request.cell_id.clone(), Arc::downgrade(&slot)); + slot + } + }; + if slot.active.swap(true, Ordering::AcqRel) { + return Err(format!( + "exec cell {} already has an active observer", + request.cell_id + )); + } + slot + }; + let lock = Arc::clone(&slot.lock); + let mut cancellation = WaitCancellation { + session: Arc::clone(self), + slot: Some(slot), + wait_id: None, + permit: None, + }; + let permit = lock.lock_owned().await; + self.require_open()?; + let wait_id = Uuid::new_v4().to_string(); + cancellation.wait_id = Some(wait_id.clone()); + cancellation.permit = Some(permit); + let expected_cell_id = request.cell_id; + let runtime_timeout = + Duration::from_millis(request.yield_time_ms).saturating_add(Duration::from_secs(1)); + let request = grpc::WaitRequest { + session_id: self.id.clone(), + cell_id: expected_cell_id.as_str().to_string(), + wait_id, + yield_time_ms: request.yield_time_ms, + }; + let mut client = self.client(); + let response = deadline::request(self, "wait", runtime_timeout, client.wait(request)).await; + cancellation.disarm(); + self.prune_wait_slots(); + let outcome = conversion::wait_outcome(response?.into_inner())?; + self.validate_wait_cell(&expected_cell_id, outcome) + } + + pub(super) async fn terminate(&self, cell_id: CellId) -> Result { + self.require_open()?; + let mut client = self.client(); + let response = deadline::request( + self, + "termination", + Duration::ZERO, + client.terminate(grpc::TerminateRequest { + session_id: self.id.clone(), + cell_id: cell_id.as_str().to_string(), + }), + ) + .await? + .into_inner(); + let outcome = conversion::wait_outcome(response)?; + self.validate_wait_cell(&cell_id, outcome) + } + + fn validate_wait_cell( + &self, + expected_cell_id: &CellId, + outcome: WaitOutcome, + ) -> Result { + let actual_cell_id = match &outcome { + WaitOutcome::LiveCell(response) | WaitOutcome::MissingCell(response) => { + runtime_response_cell_id(response) + } + }; + if actual_cell_id != expected_cell_id { + let error = format!( + "gRPC code-mode host returned cell {actual_cell_id} instead of {expected_cell_id}" + ); + self.fail(error.clone()); + return Err(error); + } + Ok(outcome) + } + + fn prune_wait_slots(&self) { + self.wait_slots + .lock() + .unwrap_or_else(PoisonError::into_inner) + .retain(|_, slot| slot.strong_count() != 0); + } +} + +fn runtime_response_cell_id(response: &RuntimeResponse) -> &CellId { + match response { + RuntimeResponse::Yielded { cell_id, .. } + | RuntimeResponse::Terminated { cell_id, .. } + | RuntimeResponse::Result { cell_id, .. } => cell_id, + } +} + +struct WaitCancellation { + session: Arc, + slot: Option>, + wait_id: Option, + permit: Option>, +} + +impl WaitCancellation { + fn disarm(&mut self) { + self.wait_id = None; + if let Some(slot) = self.slot.take() { + slot.active.store(false, Ordering::Release); + } + self.permit = None; + } +} + +impl Drop for WaitCancellation { + fn drop(&mut self) { + let slot = self.slot.take(); + if let Some(slot) = slot.as_ref() { + slot.active.store(false, Ordering::Release); + } + let Some(wait_id) = self.wait_id.take() else { + return; + }; + let permit = self.permit.take(); + if self.session.stopped.is_cancelled() { + return; + } + let session = Arc::clone(&self.session); + self.session.runtime.spawn(async move { + let mut client = session.client(); + let result = deadline::request( + &session, + "wait cancellation", + Duration::ZERO, + client.cancel_wait(grpc::CancelWaitRequest { + session_id: session.id.clone(), + wait_id, + }), + ) + .await; + if let Err(error) = result + && !session.stopped.is_cancelled() + { + session.fail(format!( + "failed to retire canceled gRPC code-mode wait: {error}" + )); + } + drop(permit); + drop(slot); + session.prune_wait_slots(); + }); + } +} diff --git a/codex-rs/code-mode/src/grpc_session/state.rs b/codex-rs/code-mode/src/grpc_session/state.rs new file mode 100644 index 0000000000..55ecb708ae --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/state.rs @@ -0,0 +1,149 @@ +use std::collections::HashMap; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::grpc; + +#[derive(Default)] +struct ExecutionRecord { + cell_id: Option, + started: bool, + ready: bool, + closed: bool, +} + +impl ExecutionRecord { + fn accept_cell(&mut self, cell_id: &str) -> Result<(), String> { + super::validate_identifier(cell_id, "cell ID")?; + if let Some(current) = self.cell_id.as_ref() { + if current.as_str() != cell_id { + return Err(format!( + "code-mode execution changed cell ID from {current} to {cell_id}" + )); + } + } else { + self.cell_id = Some(CellId::new(cell_id.to_string())); + } + Ok(()) + } +} + +#[derive(Default)] +pub(super) struct SessionState { + executions: HashMap, + failure: Option, + closed: bool, +} + +impl SessionState { + pub(super) fn require_open(&self) -> Result<(), String> { + if self.closed { + return Err(self + .failure + .clone() + .unwrap_or_else(|| "code-mode gRPC session is closed".to_string())); + } + Ok(()) + } + + pub(super) fn begin_execution(&mut self, execution_id: String) -> Result<(), String> { + self.require_open()?; + if execution_id.is_empty() || self.executions.contains_key(&execution_id) { + return Err("code-mode execution ID was empty or reused".to_string()); + } + self.executions + .insert(execution_id, ExecutionRecord::default()); + Ok(()) + } + + pub(super) fn admit_execution( + &mut self, + execution_id: &str, + cell_id: &str, + ) -> Result<(), String> { + self.require_open()?; + if self.executions.iter().any(|(id, execution)| { + id != execution_id + && execution + .cell_id + .as_ref() + .is_some_and(|current| current.as_str() == cell_id) + }) { + return Err(format!("code-mode host reused active cell ID {cell_id}")); + } + let execution = self + .executions + .get_mut(execution_id) + .ok_or_else(|| format!("unknown code-mode execution {execution_id}"))?; + if execution.started { + return Err(format!("code-mode execution {execution_id} started twice")); + } + execution.accept_cell(cell_id)?; + execution.started = true; + Ok(()) + } + + pub(super) fn mark_execution_ready( + &mut self, + execution_id: &str, + ) -> Result, String> { + self.require_open()?; + let execution = self + .executions + .get_mut(execution_id) + .ok_or_else(|| format!("unknown code-mode execution {execution_id}"))?; + if !execution.started || execution.ready { + return Err(format!( + "code-mode execution {execution_id} was not ready to be claimed" + )); + } + execution.ready = true; + Ok(self.close_execution_if_ready(execution_id)) + } + + pub(super) fn close_cell( + &mut self, + closed: grpc::CellClosed, + ) -> Result, String> { + self.require_open()?; + let Some(execution) = self.executions.get_mut(&closed.execution_id) else { + return Ok(None); + }; + execution.accept_cell(&closed.cell_id)?; + if execution.closed { + return Err(format!( + "code-mode host returned an invalid closure for cell {}", + closed.cell_id + )); + } + execution.closed = true; + Ok(self.close_execution_if_ready(&closed.execution_id)) + } + + pub(super) fn close(&mut self, failure: Option) -> Vec { + if self.closed { + return Vec::new(); + } + self.closed = true; + self.failure = failure; + self.executions + .drain() + .filter_map(|(_, execution)| execution.cell_id) + .collect() + } + + fn close_execution_if_ready(&mut self, execution_id: &str) -> Option { + self.executions + .get(execution_id) + .is_some_and(|execution| execution.started && execution.ready && execution.closed) + .then(|| self.remove_execution(execution_id)) + .flatten() + } + + pub(super) fn remove_execution(&mut self, execution_id: &str) -> Option { + self.executions.remove(execution_id)?.cell_id + } +} + +#[cfg(test)] +#[path = "state_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode/src/grpc_session/state_tests.rs b/codex-rs/code-mode/src/grpc_session/state_tests.rs new file mode 100644 index 0000000000..d1d4e05cae --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/state_tests.rs @@ -0,0 +1,89 @@ +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::grpc; +use pretty_assertions::assert_eq; + +use super::SessionState; + +#[test] +fn cell_closure_waits_until_the_started_cell_is_claimed() { + let mut state = SessionState::default(); + state + .begin_execution("execution".to_string()) + .expect("register execution"); + + assert_eq!( + state + .close_cell(grpc::CellClosed { + execution_id: "execution".to_string(), + cell_id: "cell".to_string(), + final_tool_call_sequence: 0, + }) + .expect("record early cell closure"), + None + ); + state + .admit_execution("execution", "cell") + .expect("admit started cell"); + assert_eq!( + state + .mark_execution_ready("execution") + .expect("claim started cell"), + Some(CellId::new("cell".to_string())) + ); +} + +#[test] +fn oversized_cell_ids_are_rejected_before_admission() { + let mut state = SessionState::default(); + state + .begin_execution("execution".to_string()) + .expect("register execution"); + + assert_eq!( + state.admit_execution("execution", &"x".repeat(grpc::MAX_IDENTIFIER_BYTES + 1)), + Err(format!( + "gRPC code-mode host returned cell ID exceeding {} bytes", + grpc::MAX_IDENTIFIER_BYTES + )) + ); + assert_eq!(state.remove_execution("execution"), None); +} + +#[test] +fn abandonment_before_start_ignores_later_cell_closure() { + let mut state = SessionState::default(); + state + .begin_execution("execution".to_string()) + .expect("register execution"); + + assert_eq!(state.remove_execution("execution"), None); + assert_eq!( + state + .close_cell(grpc::CellClosed { + execution_id: "execution".to_string(), + cell_id: "cell".to_string(), + final_tool_call_sequence: 0, + }) + .expect("ignore closure for abandoned execution"), + None + ); + assert!(state.close(/*failure*/ None).is_empty()); +} + +#[test] +fn disconnect_returns_each_live_cell_once() { + let mut state = SessionState::default(); + state + .begin_execution("execution".to_string()) + .expect("register execution"); + state + .admit_execution("execution", "cell") + .expect("admit execution"); + + assert_eq!( + state.close(Some("lease closed".to_string())), + vec![CellId::new("cell".to_string())] + ); + assert!(state.close(/*failure*/ None).is_empty()); + assert_eq!(state.require_open(), Err("lease closed".to_string())); +} diff --git a/codex-rs/code-mode/src/grpc_session/transport.rs b/codex-rs/code-mode/src/grpc_session/transport.rs new file mode 100644 index 0000000000..23ef21408b --- /dev/null +++ b/codex-rs/code-mode/src/grpc_session/transport.rs @@ -0,0 +1,46 @@ +use tonic::transport::Channel; +use tonic::transport::Endpoint; + +pub(super) struct SharedTransport { + endpoint: TransportEndpoint, + channel: tokio::sync::OnceCell, +} + +enum TransportEndpoint { + Url(String), + Connected(Channel), +} + +impl SharedTransport { + pub(super) fn new(endpoint: String) -> Self { + Self { + endpoint: TransportEndpoint::Url(endpoint), + channel: tokio::sync::OnceCell::new(), + } + } + + pub(super) fn with_channel(channel: Channel) -> Self { + Self { + endpoint: TransportEndpoint::Connected(channel), + channel: tokio::sync::OnceCell::new(), + } + } + + pub(super) async fn channel(&self) -> Result { + self.channel + .get_or_try_init(|| async { + match &self.endpoint { + TransportEndpoint::Url(endpoint) => Endpoint::from_shared(endpoint.clone()) + .map_err(|error| format!("invalid gRPC code-mode host endpoint: {error}"))? + .connect() + .await + .map_err(|error| { + format!("failed to connect to gRPC code-mode host: {error}") + }), + TransportEndpoint::Connected(channel) => Ok(channel.clone()), + } + }) + .await + .cloned() + } +} diff --git a/codex-rs/code-mode/src/lib.rs b/codex-rs/code-mode/src/lib.rs index 9e7d7d09e0..efbe08e1fa 100644 --- a/codex-rs/code-mode/src/lib.rs +++ b/codex-rs/code-mode/src/lib.rs @@ -1,6 +1,8 @@ +mod grpc_session; mod remote_session; pub use codex_code_mode_protocol::*; +pub use grpc_session::GrpcCodeModeSessionProvider; pub use remote_session::DisabledCodeModeSessionProvider; pub use remote_session::ProcessOwnedCodeModeSession; pub use remote_session::ProcessOwnedCodeModeSessionProvider; diff --git a/codex-rs/code-mode/src/remote_session.rs b/codex-rs/code-mode/src/remote_session.rs index 8685cfb739..502642129f 100644 --- a/codex-rs/code-mode/src/remote_session.rs +++ b/codex-rs/code-mode/src/remote_session.rs @@ -32,7 +32,7 @@ use crate::NoopCodeModeSessionDelegate; mod connection; -type ShutdownResultReceiver = watch::Receiver>>; +pub(crate) type ShutdownResultReceiver = watch::Receiver>>; /// Creates code-mode sessions backed by one lazily spawned process host. pub struct ProcessOwnedCodeModeSessionProvider { @@ -547,7 +547,7 @@ enum ShutdownAction { Close(SessionBinding), } -async fn wait_for_watch( +pub(crate) async fn wait_for_watch( mut result_rx: watch::Receiver>>, ) -> Result where