diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index ff0265d4d..9f3e69e3c 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -2519,7 +2519,10 @@ dependencies = [ "anyhow", "codex-code-mode", "codex-code-mode-protocol", + "codex-protocol", + "codex-utils-cargo-bin", "pretty_assertions", + "serde_json", "tokio", "tokio-util", ] diff --git a/codex-rs/code-mode-host/Cargo.toml b/codex-rs/code-mode-host/Cargo.toml index d37bd4101..90b043b09 100644 --- a/codex-rs/code-mode-host/Cargo.toml +++ b/codex-rs/code-mode-host/Cargo.toml @@ -24,4 +24,7 @@ tokio = { workspace = true, features = ["io-std", "io-util", "macros", "rt", "sy tokio-util = { workspace = true, features = ["rt"] } [dev-dependencies] +codex-protocol = { workspace = true } +codex-utils-cargo-bin = { workspace = true } pretty_assertions = { workspace = true } +serde_json = { workspace = true } diff --git a/codex-rs/code-mode-host/tests/stdio.rs b/codex-rs/code-mode-host/tests/stdio.rs new file mode 100644 index 000000000..e7b750616 --- /dev/null +++ b/codex-rs/code-mode-host/tests/stdio.rs @@ -0,0 +1,874 @@ +#![allow(clippy::expect_used)] + +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; +use std::time::Duration; + +#[cfg(unix)] +use std::os::unix::fs::PermissionsExt; + +use codex_code_mode::CellId; +use codex_code_mode::CodeModeNestedToolCall; +use codex_code_mode::CodeModeSession; +use codex_code_mode::CodeModeSessionDelegate; +use codex_code_mode::CodeModeSessionProvider; +use codex_code_mode::CodeModeToolKind; +use codex_code_mode::ExecuteRequest; +use codex_code_mode::FunctionCallOutputContentItem; +use codex_code_mode::NotificationFuture; +use codex_code_mode::ProcessOwnedCodeModeSessionProvider; +use codex_code_mode::RuntimeResponse; +use codex_code_mode::ToolDefinition; +use codex_code_mode::ToolInvocationFuture; +use codex_code_mode::WaitOutcome; +use codex_code_mode::WaitRequest; +use codex_code_mode::host::MAX_FRAME_BYTES; +use codex_protocol::ToolName; +use pretty_assertions::assert_eq; +use serde_json::json; +use tokio::sync::Semaphore; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +#[derive(Default)] +struct RecordingDelegate { + invocations: Mutex>, + notifications: Mutex>, + closed_cells: Mutex>, +} + +#[derive(Debug, Eq, PartialEq)] +enum CallbackEvent { + Started(String), + Cancelled(String), + CellClosed(CellId), +} + +struct CancellationDelegate { + events_tx: mpsc::UnboundedSender, + fast_tool_release: Semaphore, + slow_tool_started: Semaphore, + hold_slow_cleanup: AtomicBool, + slow_cleanup_release: Semaphore, +} + +struct OversizedResultDelegate; + +impl CodeModeSessionDelegate for OversizedResultDelegate { + fn invoke_tool<'a>( + &'a self, + _invocation: CodeModeNestedToolCall, + _cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + Box::pin(async { Ok(json!("x".repeat(MAX_FRAME_BYTES))) }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, _cell_id: &CellId) {} +} + +impl CancellationDelegate { + fn new() -> (Arc, mpsc::UnboundedReceiver) { + let (events_tx, events_rx) = mpsc::unbounded_channel(); + ( + Arc::new(Self { + events_tx, + fast_tool_release: Semaphore::new(/*permits*/ 0), + slow_tool_started: Semaphore::new(/*permits*/ 0), + hold_slow_cleanup: AtomicBool::new(false), + slow_cleanup_release: Semaphore::new(/*permits*/ 0), + }), + events_rx, + ) + } + + #[cfg(unix)] + fn hold_slow_cleanup(&self) { + self.hold_slow_cleanup.store(true, Ordering::Release); + } + + #[cfg(unix)] + fn release_slow_cleanup(&self) { + self.slow_cleanup_release.add_permits(1); + } +} + +impl CodeModeSessionDelegate for CancellationDelegate { + fn invoke_tool<'a>( + &'a self, + invocation: CodeModeNestedToolCall, + cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + Box::pin(async move { + let tool_name = invocation.tool_name.name.clone(); + if tool_name == "tool_call_barrier" { + let permit = self + .slow_tool_started + .acquire() + .await + .map_err(|_| "slow tool barrier closed".to_string())?; + permit.forget(); + return Ok(json!({ "tool": tool_name })); + } + let _ = self + .events_tx + .send(CallbackEvent::Started(tool_name.clone())); + if tool_name == "tool_call_slow" { + self.slow_tool_started.add_permits(1); + cancellation_token.cancelled().await; + let _ = self.events_tx.send(CallbackEvent::Cancelled(tool_name)); + if self.hold_slow_cleanup.load(Ordering::Acquire) { + let permit = self + .slow_cleanup_release + .acquire() + .await + .map_err(|_| "slow tool cleanup release closed".to_string())?; + permit.forget(); + } + return Err("slow tool cancelled".to_string()); + } + let permit = self + .fast_tool_release + .acquire() + .await + .map_err(|_| "fast tool release closed".to_string())?; + permit.forget(); + Ok(json!({ "tool": tool_name })) + }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, cell_id: &CellId) { + let _ = self + .events_tx + .send(CallbackEvent::CellClosed(cell_id.clone())); + } +} + +impl CodeModeSessionDelegate for RecordingDelegate { + fn invoke_tool<'a>( + &'a self, + invocation: CodeModeNestedToolCall, + _cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + self.invocations + .lock() + .expect("invocations lock") + .push(invocation); + Box::pin(async { Ok(json!({ "value": "output" })) }) + } + + 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_cells + .lock() + .expect("closed cells lock") + .push(cell_id.clone()); + } +} + +fn cell_id(value: &str) -> CellId { + CellId::new(value.to_string()) +} + +fn execute_request(source: &str) -> ExecuteRequest { + ExecuteRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: Vec::new(), + source: source.to_string(), + yield_time_ms: None, + max_output_tokens: None, + } +} + +async fn execute(session: &Arc, request: ExecuteRequest) -> RuntimeResponse { + session + .execute(request) + .await + .expect("start execution") + .initial_response() + .await + .expect("initial response") +} + +async fn execute_to_terminal( + session: &Arc, + request: ExecuteRequest, +) -> RuntimeResponse { + let started = session.execute(request).await.expect("start execution"); + let mut response = started.initial_response().await.expect("initial response"); + loop { + match response { + RuntimeResponse::Yielded { cell_id, .. } => { + response = match session + .wait(WaitRequest { + cell_id, + yield_time_ms: 60_000, + }) + .await + .expect("wait for terminal response") + { + WaitOutcome::LiveCell(response) | WaitOutcome::MissingCell(response) => { + response + } + }; + } + response => return response, + } + } +} + +async fn next_callback_event( + events_rx: &mut mpsc::UnboundedReceiver, +) -> CallbackEvent { + tokio::time::timeout(Duration::from_secs(5), events_rx.recv()) + .await + .expect("callback event timeout") + .expect("callback event stream closed") +} + +#[tokio::test] +async fn remote_session_persists_values_forwards_delegates_and_controls_cells() { + let provider = ProcessOwnedCodeModeSessionProvider::with_host_program( + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("host binary"), + ); + let delegate = Arc::new(RecordingDelegate::default()); + let session = provider + .create_session(delegate.clone()) + .await + .expect("create remote session"); + + assert_eq!( + execute(&session, execute_request(r#"store("key", "persisted");"#),).await, + RuntimeResponse::Result { + cell_id: cell_id("1"), + content_items: Vec::new(), + error_text: None, + } + ); + + let mut callback_request = execute_request( + r#" +const result = await tools.echo({ value: String(load("key")) }); +notify("notice"); +text(result.value); +"#, + ); + callback_request.tool_call_id = "call-2".to_string(); + callback_request.enabled_tools = vec![ToolDefinition { + name: "echo".to_string(), + tool_name: ToolName::plain("echo"), + description: String::new(), + kind: CodeModeToolKind::Function, + input_schema: None, + output_schema: None, + }]; + assert_eq!( + execute(&session, callback_request).await, + RuntimeResponse::Result { + cell_id: cell_id("2"), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "output".to_string(), + }], + error_text: None, + } + ); + assert_eq!( + *delegate.invocations.lock().expect("invocations lock"), + vec![CodeModeNestedToolCall { + cell_id: cell_id("2"), + runtime_tool_call_id: "tool-1".to_string(), + tool_name: ToolName::plain("echo"), + tool_kind: CodeModeToolKind::Function, + input: Some(json!({ "value": "persisted" })), + }] + ); + assert_eq!( + *delegate.notifications.lock().expect("notifications lock"), + vec![("call-2".to_string(), cell_id("2"), "notice".to_string())] + ); + + let mut pending_request = execute_request("await new Promise(() => {});"); + pending_request.tool_call_id = "call-3".to_string(); + pending_request.yield_time_ms = Some(1); + assert_eq!( + execute(&session, pending_request).await, + RuntimeResponse::Yielded { + cell_id: cell_id("3"), + content_items: Vec::new(), + } + ); + assert_eq!( + session + .wait(WaitRequest { + cell_id: cell_id("3"), + yield_time_ms: 1, + }) + .await + .expect("wait for cell"), + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: cell_id("3"), + content_items: Vec::new(), + }) + ); + assert_eq!( + session + .terminate(cell_id("3")) + .await + .expect("terminate cell"), + WaitOutcome::LiveCell(RuntimeResponse::Terminated { + cell_id: cell_id("3"), + content_items: Vec::new(), + }) + ); + + session.shutdown().await.expect("shutdown remote session"); + assert_eq!( + *delegate.closed_cells.lock().expect("closed cells lock"), + vec![cell_id("1"), cell_id("2"), cell_id("3")] + ); +} + +#[tokio::test] +async fn dropping_long_wait_releases_observer_before_next_wait() { + let provider = ProcessOwnedCodeModeSessionProvider::with_host_program( + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("host binary"), + ); + let session = provider + .create_session(Arc::new(RecordingDelegate::default())) + .await + .expect("create remote session"); + let mut request = execute_request("await new Promise(() => {});"); + request.yield_time_ms = Some(1); + let started = session.execute(request).await.expect("start execution"); + let running_cell_id = started.cell_id.clone(); + assert_eq!( + started.initial_response().await.expect("initial response"), + RuntimeResponse::Yielded { + cell_id: running_cell_id.clone(), + content_items: Vec::new(), + } + ); + + let wait_session = Arc::clone(&session); + let wait_cell_id = running_cell_id.clone(); + let first_wait = tokio::spawn(async move { + wait_session + .wait(WaitRequest { + cell_id: wait_cell_id, + yield_time_ms: 60_000, + }) + .await + }); + tokio::time::sleep(Duration::from_millis(100)).await; + first_wait.abort(); + let _ = first_wait.await; + + assert_eq!( + tokio::time::timeout( + Duration::from_secs(2), + session.wait(WaitRequest { + cell_id: running_cell_id.clone(), + yield_time_ms: 1, + }) + ) + .await + .expect("second wait timeout") + .expect("second wait"), + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: running_cell_id.clone(), + content_items: Vec::new(), + }) + ); + session + .terminate(running_cell_id) + .await + .expect("terminate cell"); + session.shutdown().await.expect("shutdown remote session"); +} + +#[tokio::test] +async fn unawaited_slow_tool_is_cancelled_after_parallel_tools_complete() { + let provider = ProcessOwnedCodeModeSessionProvider::with_host_program( + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("host binary"), + ); + let (delegate, mut events_rx) = CancellationDelegate::new(); + let session = provider + .create_session(delegate.clone()) + .await + .expect("create remote session"); + let mut request = execute_request( + r#" +await (async () => { +text("hello world"); +yield_control(); +await Promise.all([ + tools.tool_call_a({}), + tools.tool_call_b({}), +]); +text("hello"); +tools.tool_call_slow({}); +await tools.tool_call_barrier({}); +return; +})(); +"#, + ); + request.enabled_tools = [ + "tool_call_a", + "tool_call_b", + "tool_call_slow", + "tool_call_barrier", + ] + .into_iter() + .map(|name| ToolDefinition { + name: name.to_string(), + tool_name: ToolName::plain(name), + description: String::new(), + kind: CodeModeToolKind::Function, + input_schema: None, + output_schema: None, + }) + .collect(); + + let started = session.execute(request).await.expect("start execution"); + let running_cell_id = started.cell_id.clone(); + assert_eq!( + started.initial_response().await.expect("initial response"), + RuntimeResponse::Yielded { + cell_id: running_cell_id.clone(), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "hello world".to_string(), + }], + } + ); + + let wait_session = Arc::clone(&session); + let wait_cell_id = running_cell_id.clone(); + let wait_task = tokio::spawn(async move { + wait_session + .wait(WaitRequest { + cell_id: wait_cell_id, + yield_time_ms: 60_000, + }) + .await + }); + + let mut parallel_tools = vec![ + next_callback_event(&mut events_rx).await, + next_callback_event(&mut events_rx).await, + ]; + parallel_tools.sort_by(|left, right| format!("{left:?}").cmp(&format!("{right:?}"))); + assert_eq!( + parallel_tools, + vec![ + CallbackEvent::Started("tool_call_a".to_string()), + CallbackEvent::Started("tool_call_b".to_string()), + ] + ); + delegate.fast_tool_release.add_permits(2); + + assert_eq!( + next_callback_event(&mut events_rx).await, + CallbackEvent::Started("tool_call_slow".to_string()) + ); + assert_eq!( + next_callback_event(&mut events_rx).await, + CallbackEvent::Cancelled("tool_call_slow".to_string()) + ); + assert_eq!( + next_callback_event(&mut events_rx).await, + CallbackEvent::CellClosed(running_cell_id.clone()) + ); + assert_eq!( + wait_task + .await + .expect("wait task") + .expect("wait for terminal response"), + WaitOutcome::LiveCell(RuntimeResponse::Result { + cell_id: running_cell_id, + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "hello".to_string(), + }], + error_text: None, + }) + ); + session.shutdown().await.expect("shutdown remote session"); +} + +#[tokio::test] +async fn oversized_execute_request_does_not_close_the_shared_host() { + let provider = ProcessOwnedCodeModeSessionProvider::with_host_program( + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("host binary"), + ); + let session = provider + .create_session(Arc::new(RecordingDelegate::default())) + .await + .expect("create remote session"); + let error = session + .execute(execute_request(&"x".repeat(MAX_FRAME_BYTES))) + .await + .err() + .expect("oversized execute should fail"); + assert!( + error.contains("IPC frame limit"), + "unexpected error: {error}" + ); + + assert_eq!( + execute(&session, execute_request(r#"text("still alive");"#)).await, + RuntimeResponse::Result { + cell_id: cell_id("1"), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "still alive".to_string(), + }], + error_text: None, + } + ); + session.shutdown().await.expect("shutdown remote session"); +} + +#[tokio::test] +async fn oversized_delegate_payloads_fail_only_the_tool_call() { + let provider = ProcessOwnedCodeModeSessionProvider::with_host_program( + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("host binary"), + ); + let session = provider + .create_session(Arc::new(OversizedResultDelegate)) + .await + .expect("create remote session"); + let tool = |name: &str| ToolDefinition { + name: name.to_string(), + tool_name: ToolName::plain(name), + description: String::new(), + kind: CodeModeToolKind::Function, + input_schema: None, + output_schema: None, + }; + + let mut oversized_argument = execute_request(&format!( + r#" +try {{ + await tools.big_argument({{ value: "x".repeat({MAX_FRAME_BYTES}) }}); +}} catch (_) {{ + text("argument rejected"); +}} +"# + )); + oversized_argument.enabled_tools = vec![tool("big_argument")]; + oversized_argument.yield_time_ms = Some(60_000); + assert_eq!( + execute_to_terminal(&session, oversized_argument).await, + RuntimeResponse::Result { + cell_id: cell_id("1"), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "argument rejected".to_string(), + }], + error_text: None, + } + ); + + let mut oversized_result = execute_request( + r#" +try { + await tools.big_result({}); +} catch (_) { + text("result rejected"); +} +"#, + ); + oversized_result.enabled_tools = vec![tool("big_result")]; + oversized_result.yield_time_ms = Some(60_000); + assert_eq!( + execute_to_terminal(&session, oversized_result).await, + RuntimeResponse::Result { + cell_id: cell_id("2"), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "result rejected".to_string(), + }], + error_text: None, + } + ); + + assert_eq!( + execute(&session, execute_request(r#"text("still alive");"#)).await, + RuntimeResponse::Result { + cell_id: cell_id("3"), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "still alive".to_string(), + }], + error_text: None, + } + ); + session.shutdown().await.expect("shutdown remote session"); +} + +#[tokio::test] +async fn oversized_initial_response_does_not_close_the_shared_host() { + let provider = ProcessOwnedCodeModeSessionProvider::with_host_program( + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("host binary"), + ); + let session = provider + .create_session(Arc::new(RecordingDelegate::default())) + .await + .expect("create remote session"); + let started = session + .execute(execute_request(&format!( + r#"text("x".repeat({MAX_FRAME_BYTES}));"# + ))) + .await + .expect("start oversized response"); + let error = started + .initial_response() + .await + .expect_err("oversized initial response should fail"); + assert!( + error.contains("IPC frame limit"), + "unexpected error: {error}" + ); + + assert_eq!( + execute(&session, execute_request(r#"text("still alive");"#)).await, + RuntimeResponse::Result { + cell_id: cell_id("2"), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "still alive".to_string(), + }], + error_text: None, + } + ); + session.shutdown().await.expect("shutdown remote session"); +} + +#[cfg(unix)] +#[tokio::test] +async fn child_process_loss_cleans_up_and_rebuilds_the_shared_host() { + let host_program = + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("host binary"); + let proxy_dir = + std::env::temp_dir().join(format!("codex-code-mode-host-loss-{}", std::process::id())); + let proxy_program = proxy_dir.join("host-proxy.sh"); + let pid_path = proxy_dir.join("host.pid"); + let _ = std::fs::remove_dir_all(&proxy_dir); + std::fs::create_dir_all(&proxy_dir).expect("create host proxy directory"); + std::fs::write( + &proxy_program, + format!( + "#!/bin/sh\nprintf '%s\\n' \"$$\" > '{}'\nexec '{}'\n", + pid_path.display(), + host_program.display() + ), + ) + .expect("write host proxy"); + let mut permissions = std::fs::metadata(&proxy_program) + .expect("host proxy metadata") + .permissions(); + permissions.set_mode(/*mode*/ 0o700); + std::fs::set_permissions(&proxy_program, permissions).expect("make host proxy executable"); + + let provider = ProcessOwnedCodeModeSessionProvider::with_host_program(proxy_program); + let (delegate_a, mut events_a) = CancellationDelegate::new(); + delegate_a.hold_slow_cleanup(); + let delegate_b = Arc::new(RecordingDelegate::default()); + let session_a = provider + .create_session(delegate_a.clone()) + .await + .expect("create first remote session"); + let session_b = provider + .create_session(delegate_b.clone()) + .await + .expect("create second remote session"); + + let mut request_a = execute_request("await tools.tool_call_slow({});"); + request_a.yield_time_ms = Some(1); + request_a.enabled_tools = vec![ToolDefinition { + name: "tool_call_slow".to_string(), + tool_name: ToolName::plain("tool_call_slow"), + description: String::new(), + kind: CodeModeToolKind::Function, + input_schema: None, + output_schema: None, + }]; + let started_a = session_a + .execute(request_a) + .await + .expect("start first cell"); + let cell_a = started_a.cell_id.clone(); + assert_eq!( + started_a + .initial_response() + .await + .expect("first initial response"), + RuntimeResponse::Yielded { + cell_id: cell_a.clone(), + content_items: Vec::new(), + } + ); + assert_eq!( + next_callback_event(&mut events_a).await, + CallbackEvent::Started("tool_call_slow".to_string()) + ); + + let mut request_b = execute_request("await new Promise(() => {});"); + request_b.yield_time_ms = Some(1); + let started_b = session_b + .execute(request_b) + .await + .expect("start second cell"); + let cell_b = started_b.cell_id.clone(); + assert_eq!( + started_b + .initial_response() + .await + .expect("second initial response"), + RuntimeResponse::Yielded { + cell_id: cell_b.clone(), + content_items: Vec::new(), + } + ); + + let wait_a_session = Arc::clone(&session_a); + let wait_a_cell = cell_a.clone(); + let wait_a = tokio::spawn(async move { + wait_a_session + .wait(WaitRequest { + cell_id: wait_a_cell, + yield_time_ms: 60_000, + }) + .await + }); + let wait_b_session = Arc::clone(&session_b); + let wait_b_cell = cell_b.clone(); + let wait_b = tokio::spawn(async move { + wait_b_session + .wait(WaitRequest { + cell_id: wait_b_cell, + yield_time_ms: 60_000, + }) + .await + }); + tokio::time::sleep(Duration::from_millis(100)).await; + + let pid = tokio::time::timeout(Duration::from_secs(5), async { + loop { + if let Ok(pid) = std::fs::read_to_string(&pid_path) + && let Ok(pid) = pid.trim().parse::() + { + break pid; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("host pid timeout"); + let kill_status = std::process::Command::new("kill") + .args(["-KILL", &pid.to_string()]) + .status() + .expect("kill host process"); + assert!(kill_status.success()); + + assert!( + tokio::time::timeout(Duration::from_secs(5), wait_a) + .await + .expect("first wait failure timeout") + .expect("first wait task") + .is_err() + ); + assert!( + tokio::time::timeout(Duration::from_secs(5), wait_b) + .await + .expect("second wait failure timeout") + .expect("second wait task") + .is_err() + ); + let closure_events = [ + next_callback_event(&mut events_a).await, + next_callback_event(&mut events_a).await, + ]; + assert!(closure_events.contains(&CallbackEvent::Cancelled("tool_call_slow".to_string()))); + assert!(closure_events.contains(&CallbackEvent::CellClosed(cell_a.clone()))); + tokio::time::timeout(Duration::from_secs(5), async { + loop { + if delegate_b + .closed_cells + .lock() + .expect("closed cells lock") + .contains(&cell_b) + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("unrelated session cleanup timeout"); + + assert_eq!( + execute(&session_b, execute_request(r#"text("replacement");"#)).await, + RuntimeResponse::Result { + cell_id: cell_id("g2:1"), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "replacement".to_string(), + }], + error_text: None, + } + ); + let stale_error = session_b + .wait(WaitRequest { + cell_id: cell_b.clone(), + yield_time_ms: 1, + }) + .await + .expect_err("stale cell should be rejected"); + assert!(stale_error.contains("stale code-mode host generation")); + + tokio::time::timeout(Duration::from_secs(5), session_a.shutdown()) + .await + .expect("failed session shutdown timeout") + .expect("shutdown failed session"); + tokio::time::timeout(Duration::from_secs(5), session_b.shutdown()) + .await + .expect("unrelated session shutdown timeout") + .expect("shutdown replacement session"); + + delegate_a.release_slow_cleanup(); + tokio::task::yield_now().await; + assert!(matches!( + events_a.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + + std::fs::remove_dir_all(proxy_dir).expect("remove host proxy directory"); +} diff --git a/codex-rs/code-mode/Cargo.toml b/codex-rs/code-mode/Cargo.toml index 6b91e64fe..011fee6eb 100644 --- a/codex-rs/code-mode/Cargo.toml +++ b/codex-rs/code-mode/Cargo.toml @@ -21,7 +21,7 @@ codex-protocol = { workspace = true } deno_core_icudata = { workspace = true } futures = { workspace = true } serde_json = { workspace = true } -tokio = { workspace = true, features = ["macros", "rt", "sync", "time"] } +tokio = { workspace = true, features = ["io-util", "macros", "process", "rt", "sync", "time"] } tokio-util = { workspace = true, features = ["rt"] } tracing = { workspace = true } v8 = { workspace = true } diff --git a/codex-rs/code-mode/src/lib.rs b/codex-rs/code-mode/src/lib.rs index 60c02e4fa..653d7b160 100644 --- a/codex-rs/code-mode/src/lib.rs +++ b/codex-rs/code-mode/src/lib.rs @@ -1,4 +1,5 @@ mod cell_actor; +mod remote_session; mod runtime; mod service; mod session_runtime; @@ -6,6 +7,8 @@ mod session_runtime; pub(crate) type TaskFailureHandler = std::sync::Arc; pub use codex_code_mode_protocol::*; +pub use remote_session::ProcessOwnedCodeModeSession; +pub use remote_session::ProcessOwnedCodeModeSessionProvider; pub use service::InProcessCodeModeSession; pub use service::InProcessCodeModeSessionProvider; pub use service::NoopCodeModeSessionDelegate; diff --git a/codex-rs/code-mode/src/remote_session.rs b/codex-rs/code-mode/src/remote_session.rs new file mode 100644 index 000000000..d942b377c --- /dev/null +++ b/codex-rs/code-mode/src/remote_session.rs @@ -0,0 +1,496 @@ +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::atomic::AtomicBool; +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::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::host::SessionId; +use tokio::sync::Semaphore; +use tokio::sync::watch; + +use self::connection::Connection; +use self::connection::RemoteSession; +use self::connection::SessionCleanup; +use crate::NoopCodeModeSessionDelegate; + +mod connection; + +const CODE_MODE_HOST_PATH_ENV: &str = "CODEX_CODE_MODE_HOST_PATH"; + +type ShutdownResultReceiver = watch::Receiver>>; + +/// Creates code-mode sessions backed by one lazily spawned process host. +pub struct ProcessOwnedCodeModeSessionProvider { + host_program: PathBuf, + process_host: StdMutex>>, +} + +impl ProcessOwnedCodeModeSessionProvider { + pub fn with_host_program(host_program: PathBuf) -> Self { + Self { + host_program, + process_host: StdMutex::new(None), + } + } + + fn process_host(&self) -> Arc { + let mut process_host = self + .process_host + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(process_host) = process_host.as_ref() { + return Arc::clone(process_host); + } + + let new_process_host = Arc::new(OwnedProcessHost::new(self.host_program.clone())); + *process_host = Some(Arc::clone(&new_process_host)); + new_process_host + } +} + +impl Default for ProcessOwnedCodeModeSessionProvider { + fn default() -> Self { + Self::with_host_program(default_host_program()) + } +} + +impl CodeModeSessionProvider for ProcessOwnedCodeModeSessionProvider { + fn create_session<'a>( + &'a self, + delegate: Arc, + ) -> CodeModeSessionProviderFuture<'a> { + let session = ProcessOwnedCodeModeSession::with_process_host(delegate, self.process_host()); + Box::pin(async move { + session.connection().await?; + let session: Arc = Arc::new(session); + Ok(session) + }) + } +} + +struct OwnedProcessHost { + host_program: PathBuf, + connection: StdMutex>>, + spawn_permit: Semaphore, + next_session_id: AtomicU64, +} + +impl OwnedProcessHost { + fn new(host_program: PathBuf) -> Self { + Self { + host_program, + connection: StdMutex::new(None), + spawn_permit: Semaphore::new(/*permits*/ 1), + next_session_id: AtomicU64::new(1), + } + } + + async fn connection(&self) -> Result, String> { + if let Some(connection) = self.live_connection() { + return Ok(connection); + } + + let _spawn_permit = self + .spawn_permit + .acquire() + .await + .map_err(|_| "code-mode host spawn coordinator closed".to_string())?; + if let Some(connection) = self.live_connection() { + return Ok(connection); + } + let new_connection = Arc::new(Connection::spawn(&self.host_program).await?); + *self + .connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(Arc::clone(&new_connection)); + Ok(new_connection) + } + + fn live_connection(&self) -> Option> { + self.connection + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .as_ref() + .filter(|connection| connection.is_alive()) + .cloned() + } + + fn allocate_session_id(&self) -> SessionId { + let value = self.next_session_id.fetch_add(1, Ordering::Relaxed); + match SessionId::new(format!("session-{value}")) { + Ok(session_id) => session_id, + Err(_) => unreachable!("a generated code-mode session ID is nonempty"), + } + } +} + +enum SessionState { + New, + Opening { + remote: RemoteSession, + result_rx: watch::Receiver>>, + }, + Open(SessionBinding), + Closing, + Closed, +} + +#[derive(Clone)] +struct SessionBinding { + connection: Arc, + remote: RemoteSession, + cleanup: SessionCleanup, +} + +struct SessionInner { + process_host: Arc, + delegate: Arc, + state: StdMutex, + next_generation: AtomicU64, + shutdown_requested: AtomicBool, + shutdown_result: StdMutex>, + retired_cleanups: StdMutex>, +} + +/// A logical code-mode session assigned to a process-owned host. +pub struct ProcessOwnedCodeModeSession { + inner: Arc, +} + +impl ProcessOwnedCodeModeSession { + pub fn new() -> Self { + Self::with_process_host( + Arc::new(NoopCodeModeSessionDelegate), + Arc::new(OwnedProcessHost::new(default_host_program())), + ) + } + + fn with_process_host( + delegate: Arc, + process_host: Arc, + ) -> Self { + Self { + inner: Arc::new(SessionInner { + process_host, + delegate, + state: StdMutex::new(SessionState::New), + next_generation: AtomicU64::new(1), + shutdown_requested: AtomicBool::new(false), + shutdown_result: StdMutex::new(None), + retired_cleanups: StdMutex::new(Vec::new()), + }), + } + } + + async fn connection(&self) -> Result { + self.inner.connection().await + } + + pub async fn execute(&self, request: ExecuteRequest) -> Result { + let binding = self.connection().await?; + binding.connection.execute(binding.remote, request).await + } + + pub async fn wait(&self, request: WaitRequest) -> Result { + let binding = self.connection().await?; + binding.connection.wait(binding.remote, request).await + } + + pub async fn terminate(&self, cell_id: CellId) -> Result { + let binding = self.connection().await?; + binding.connection.terminate(binding.remote, cell_id).await + } + + pub async fn shutdown(&self) -> Result<(), String> { + wait_for_watch(self.inner.request_shutdown()).await + } +} + +impl SessionInner { + async fn connection(self: &Arc) -> Result { + loop { + if self.shutdown_requested.load(Ordering::Acquire) { + return Err("code mode session is shutting down".to_string()); + } + let (result_rx, start) = { + let mut state = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + match &*state { + SessionState::New => { + let generation = self.next_generation.fetch_add(1, Ordering::Relaxed); + let remote = RemoteSession { + id: self.process_host.allocate_session_id(), + generation, + }; + let (result_tx, result_rx) = watch::channel(None); + *state = SessionState::Opening { + remote: remote.clone(), + result_rx: result_rx.clone(), + }; + (result_rx, Some((remote, result_tx))) + } + SessionState::Opening { result_rx, .. } => (result_rx.clone(), None), + SessionState::Open(binding) if binding.connection.is_alive() => { + return Ok(binding.clone()); + } + SessionState::Open(binding) => { + self.retain_cleanup(binding.cleanup.clone()); + *state = SessionState::New; + continue; + } + SessionState::Closing | SessionState::Closed => { + return Err("code mode session is shutting down".to_string()); + } + } + }; + if let Some((remote, result_tx)) = start { + let inner = Arc::clone(self); + tokio::spawn(async move { + inner.open(remote, result_tx).await; + }); + } + return wait_for_watch(result_rx).await; + } + } + + async fn open( + self: Arc, + remote: RemoteSession, + result_tx: watch::Sender>>, + ) { + let result = match self.process_host.connection().await { + Ok(connection) => { + let cleanup = connection + .open_session(remote.clone(), Arc::clone(&self.delegate)) + .await; + cleanup.map(|cleanup| SessionBinding { + connection, + remote: remote.clone(), + cleanup, + }) + } + Err(err) => Err(err), + }; + { + let mut state = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if matches!( + &*state, + SessionState::Opening { + remote: opening_remote, + .. + } if opening_remote == &remote + ) { + *state = match &result { + Ok(binding) => SessionState::Open(binding.clone()), + Err(_) => SessionState::New, + }; + } + } + result_tx.send_replace(Some(result)); + } + + fn request_shutdown(self: &Arc) -> ShutdownResultReceiver { + self.shutdown_requested.store(true, Ordering::Release); + let mut shutdown_result = self + .shutdown_result + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(result_rx) = shutdown_result.as_ref() { + return result_rx.clone(); + } + let (result_tx, result_rx) = watch::channel(None); + *shutdown_result = Some(result_rx.clone()); + let inner = Arc::clone(self); + tokio::spawn(async move { + let result = inner.drive_shutdown().await; + result_tx.send_replace(Some(result)); + }); + result_rx + } + + async fn drive_shutdown(self: &Arc) -> Result<(), String> { + loop { + let action = { + let mut state = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + match &*state { + SessionState::New => { + *state = SessionState::Closed; + ShutdownAction::Finish + } + SessionState::Opening { result_rx, .. } => { + ShutdownAction::WaitForOpen(result_rx.clone()) + } + SessionState::Open(binding) if !binding.connection.is_alive() => { + let cleanup = binding.cleanup.clone(); + *state = SessionState::Closing; + ShutdownAction::WaitForSessionCleanup(cleanup) + } + SessionState::Open(binding) => { + let binding = binding.clone(); + *state = SessionState::Closing; + ShutdownAction::Close(binding) + } + SessionState::Closing => { + return Err("code-mode session shutdown driver entered twice".to_string()); + } + SessionState::Closed => return Ok(()), + } + }; + match action { + ShutdownAction::WaitForOpen(result_rx) => { + let _ = wait_for_watch(result_rx).await; + } + ShutdownAction::Finish => { + self.wait_for_retired_cleanups().await; + return Ok(()); + } + ShutdownAction::WaitForSessionCleanup(cleanup) => { + cleanup.wait().await; + self.wait_for_retired_cleanups().await; + *self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = SessionState::Closed; + return Ok(()); + } + ShutdownAction::Close(binding) => { + let result = binding.connection.shutdown_session(binding.remote).await; + if result.is_err() && !binding.connection.is_alive() { + binding.cleanup.wait().await; + } + self.wait_for_retired_cleanups().await; + *self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) = SessionState::Closed; + return result; + } + } + } + } + + fn retain_cleanup(&self, cleanup: SessionCleanup) { + let mut retired = self + .retired_cleanups + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + retired.retain(|cleanup| !cleanup.is_complete()); + if !cleanup.is_complete() { + retired.push(cleanup); + } + } + + async fn wait_for_retired_cleanups(&self) { + let retired = std::mem::take( + &mut *self + .retired_cleanups + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner), + ); + for cleanup in retired { + cleanup.wait().await; + } + } +} + +enum ShutdownAction { + WaitForOpen(watch::Receiver>>), + Finish, + WaitForSessionCleanup(SessionCleanup), + Close(SessionBinding), +} + +async fn wait_for_watch( + mut result_rx: watch::Receiver>>, +) -> Result +where + T: Clone, +{ + loop { + if let Some(result) = result_rx.borrow().clone() { + return result; + } + result_rx + .changed() + .await + .map_err(|_| "code-mode session transition stopped".to_string())?; + } +} + +impl Drop for ProcessOwnedCodeModeSession { + fn drop(&mut self) { + if tokio::runtime::Handle::try_current().is_ok() { + self.inner.request_shutdown(); + } + } +} + +impl Default for ProcessOwnedCodeModeSession { + fn default() -> Self { + Self::new() + } +} + +impl CodeModeSession for ProcessOwnedCodeModeSession { + fn execute<'a>( + &'a self, + request: ExecuteRequest, + ) -> CodeModeSessionResultFuture<'a, StartedCell> { + Box::pin(ProcessOwnedCodeModeSession::execute(self, request)) + } + + fn wait<'a>(&'a self, request: WaitRequest) -> CodeModeSessionResultFuture<'a, WaitOutcome> { + Box::pin(ProcessOwnedCodeModeSession::wait(self, request)) + } + + fn terminate<'a>(&'a self, cell_id: CellId) -> CodeModeSessionResultFuture<'a, WaitOutcome> { + Box::pin(ProcessOwnedCodeModeSession::terminate(self, cell_id)) + } + + fn shutdown<'a>(&'a self) -> CodeModeSessionResultFuture<'a, ()> { + Box::pin(ProcessOwnedCodeModeSession::shutdown(self)) + } +} + +fn default_host_program() -> PathBuf { + if let Some(path) = std::env::var_os(CODE_MODE_HOST_PATH_ENV) { + return PathBuf::from(path); + } + let executable_name = if cfg!(windows) { + "codex-code-mode-host.exe" + } else { + "codex-code-mode-host" + }; + if let Ok(current_exe) = std::env::current_exe() + && let Some(parent) = current_exe.parent() + { + let sibling = parent.join(executable_name); + if sibling.is_file() { + return sibling; + } + } + PathBuf::from(executable_name) +} + +#[cfg(test)] +#[path = "remote_session_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode/src/remote_session/connection.rs b/codex-rs/code-mode/src/remote_session/connection.rs new file mode 100644 index 000000000..7154784de --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection.rs @@ -0,0 +1,458 @@ +use std::path::Path; +use std::process::Stdio; +use std::sync::Arc; +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::CodeModeSessionDelegate; +use codex_code_mode_protocol::ExecuteRequest; +use codex_code_mode_protocol::StartedCell; +use codex_code_mode_protocol::WaitOutcome; +use codex_code_mode_protocol::WaitRequest; +use codex_code_mode_protocol::host::CapabilitySet; +use codex_code_mode_protocol::host::ClientHello; +use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::EncodedFrame; +use codex_code_mode_protocol::host::FramedReader; +use codex_code_mode_protocol::host::FramedWriter; +use codex_code_mode_protocol::host::HostToClient; +use codex_code_mode_protocol::host::ProtocolVersion; +use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::SupportedProtocolVersions; +use tokio::io::AsyncBufReadExt; +use tokio::io::BufReader; +use tokio::process::Child; +use tokio::process::Command; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::task::JoinHandle; +use tokio_util::sync::CancellationToken; +use tracing::debug; +use tracing::warn; + +use self::driver::ConnectionDriver; +use self::driver::DriverCommand; +use self::driver::DriverEvent; +use self::driver::DriverLifecycle; +pub(super) use self::driver::RemoteSession; +pub(super) use self::driver::SessionCleanup; +use self::reader::drive_reader; + +mod driver; +mod reader; + +const IPC_CHANNEL_CAPACITY: usize = 128; +const HOST_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10); + +pub(super) struct Connection { + command_tx: mpsc::Sender, + execute_claim_tx: mpsc::UnboundedSender, + alive: Arc, + failure: Arc>>, + cancellation: CancellationToken, +} + +struct CallerCancellation { + token: CancellationToken, + armed: bool, +} + +struct ConnectionSupervisor { + child: Child, + event_tx: mpsc::Sender, + cancellation: CancellationToken, + alive: Arc, + failure: Arc>>, + driver_task: JoinHandle<()>, + reader_task: JoinHandle>, + writer_task: JoinHandle>, +} + +impl CallerCancellation { + fn new() -> Self { + Self { + token: CancellationToken::new(), + armed: true, + } + } + + fn token(&self) -> CancellationToken { + self.token.clone() + } + + fn disarm(mut self) { + self.armed = false; + } +} + +impl Drop for CallerCancellation { + fn drop(&mut self) { + if self.armed { + self.token.cancel(); + } + } +} + +impl Connection { + pub(super) async fn spawn(host_program: &Path) -> Result { + let mut command = Command::new(host_program); + #[cfg(unix)] + command.process_group(0); + let mut child = command + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .kill_on_drop(true) + .spawn() + .map_err(|err| { + format!( + "failed to spawn code-mode host {}: {err}", + host_program.display() + ) + })?; + + if let Some(stderr) = child.stderr.take() { + tokio::spawn(async move { + let mut lines = BufReader::new(stderr).lines(); + loop { + match lines.next_line().await { + Ok(Some(line)) => debug!("code-mode host stderr: {line}"), + Ok(None) => break, + Err(err) => { + warn!("failed to read code-mode host stderr: {err}"); + break; + } + } + } + }); + } + + let stdin = child + .stdin + .take() + .ok_or_else(|| "spawned code-mode host has no stdin".to_string())?; + let stdout = child + .stdout + .take() + .ok_or_else(|| "spawned code-mode host has no stdout".to_string())?; + let mut reader = FramedReader::new(stdout); + let mut writer = FramedWriter::new(stdin); + let handshake = async { + let hello = ClientHello::new( + SupportedProtocolVersions::try_new([ProtocolVersion::V1]) + .map_err(|err| err.to_string())?, + CapabilitySet::empty(), + CapabilitySet::empty(), + ) + .map_err(|err| err.to_string())?; + writer + .write(&ClientToHost::ClientHello(hello)) + .await + .map_err(|err| format!("failed to write code-mode host hello: {err}"))?; + match reader + .read::() + .await + .map_err(|err| format!("failed to read code-mode host hello: {err}"))? + { + Some(HostToClient::HostHello(hello)) + if hello.selected_version() == ProtocolVersion::V1 => + { + Ok(()) + } + Some(HostToClient::HandshakeRejected { reason }) => { + Err(format!("code-mode host rejected the handshake: {reason:?}")) + } + Some(message) => Err(format!( + "code-mode host returned an invalid handshake response: {message:?}" + )), + None => Err("code-mode host exited during handshake".to_string()), + } + }; + let handshake_result = match tokio::time::timeout(HOST_HANDSHAKE_TIMEOUT, handshake).await { + Ok(result) => result, + Err(_) => { + kill_and_reap(&mut child).await; + return Err("timed out negotiating with the code-mode host".to_string()); + } + }; + if let Err(err) = handshake_result { + kill_and_reap(&mut child).await; + return Err(err); + } + + let (command_tx, command_rx) = mpsc::channel(IPC_CHANNEL_CAPACITY); + let (event_tx, event_rx) = mpsc::channel(IPC_CHANNEL_CAPACITY); + let (outgoing_tx, mut outgoing_rx) = mpsc::channel::(IPC_CHANNEL_CAPACITY); + let cancellation = CancellationToken::new(); + let alive = Arc::new(AtomicBool::new(true)); + let failure = Arc::new(std::sync::Mutex::new(None)); + + let writer_cancellation = cancellation.clone(); + let writer_task = tokio::spawn(async move { + loop { + tokio::select! { + _ = writer_cancellation.cancelled() => return Ok(()), + frame = outgoing_rx.recv() => { + let Some(frame) = frame else { + return Err("code-mode host outgoing stream closed".to_string()); + }; + if let Err(err) = writer.write_frame(&frame).await { + return Err(format!("failed to write code-mode host message: {err}")); + } + } + } + } + }); + + let reader_events = event_tx.clone(); + let reader_cancellation = cancellation.clone(); + let reader_task = + tokio::spawn( + async move { drive_reader(reader, reader_events, reader_cancellation).await }, + ); + + let (driver, execute_claim_tx) = ConnectionDriver::new( + command_rx, + event_rx, + event_tx.clone(), + outgoing_tx, + DriverLifecycle { + alive: Arc::clone(&alive), + failure: Arc::clone(&failure), + cancellation: cancellation.clone(), + }, + ); + let driver_task = tokio::spawn(driver.run()); + tokio::spawn( + ConnectionSupervisor { + child, + event_tx, + cancellation: cancellation.clone(), + alive: Arc::clone(&alive), + failure: Arc::clone(&failure), + driver_task, + reader_task, + writer_task, + } + .run(), + ); + + Ok(Self { + command_tx, + execute_claim_tx, + alive, + failure, + cancellation, + }) + } + + pub(super) fn is_alive(&self) -> bool { + if self.command_tx.is_closed() { + mark_connection_dead( + &self.alive, + &self.failure, + "code-mode connection driver closed".to_string(), + ); + } + self.alive.load(Ordering::Acquire) + } + + pub(super) async fn open_session( + &self, + session: RemoteSession, + delegate: Arc, + ) -> Result { + let cleanup = SessionCleanup::new(); + let cancellation = CallerCancellation::new(); + let (response_tx, response_rx) = oneshot::channel(); + self.send(DriverCommand::OpenSession { + session, + delegate, + cleanup: cleanup.clone(), + caller_cancellation: cancellation.token(), + response_tx, + }) + .await?; + let result = self.receive(response_rx).await; + cancellation.disarm(); + result?; + Ok(cleanup) + } + + pub(super) async fn execute( + &self, + session: RemoteSession, + request: ExecuteRequest, + ) -> Result { + let cancellation = CallerCancellation::new(); + let (response_tx, response_rx) = oneshot::channel(); + self.send(DriverCommand::Execute { + session, + request, + caller_cancellation: cancellation.token(), + response_tx, + }) + .await?; + let delivered = match self.receive(response_rx).await { + Ok(delivered) => delivered, + Err(err) => { + cancellation.disarm(); + return Err(err); + } + }; + self.execute_claim_tx + .send(delivered.request_id) + .map_err(|_| self.failure_message())?; + cancellation.disarm(); + Ok(delivered.started) + } + + pub(super) async fn wait( + &self, + session: RemoteSession, + request: WaitRequest, + ) -> Result { + let cancellation = CallerCancellation::new(); + let (response_tx, response_rx) = oneshot::channel(); + self.send(DriverCommand::Wait { + session, + request, + caller_cancellation: cancellation.token(), + response_tx, + }) + .await?; + let result = self.receive(response_rx).await; + cancellation.disarm(); + result + } + + pub(super) async fn terminate( + &self, + session: RemoteSession, + cell_id: CellId, + ) -> Result { + let (response_tx, response_rx) = oneshot::channel(); + self.send(DriverCommand::Terminate { + session, + cell_id, + response_tx, + }) + .await?; + self.receive(response_rx).await + } + + pub(super) async fn shutdown_session(&self, session: RemoteSession) -> Result<(), String> { + let (response_tx, response_rx) = oneshot::channel(); + self.send(DriverCommand::ShutdownSession { + session, + response_tx, + }) + .await?; + self.receive(response_rx).await + } + + async fn send(&self, command: DriverCommand) -> Result<(), String> { + if !self.is_alive() { + return Err(self.failure_message()); + } + self.command_tx + .send(command) + .await + .map_err(|_| self.failure_message()) + } + + async fn receive( + &self, + response_rx: oneshot::Receiver>, + ) -> Result { + response_rx.await.map_err(|_| self.failure_message())? + } + + fn failure_message(&self) -> String { + self.failure + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + .unwrap_or_else(|| "code-mode host connection closed".to_string()) + } +} + +impl Drop for Connection { + fn drop(&mut self) { + mark_connection_dead( + &self.alive, + &self.failure, + "code-mode host connection closed".to_string(), + ); + self.cancellation.cancel(); + } +} + +impl ConnectionSupervisor { + async fn run(mut self) { + let mut child_exited = false; + let reason = tokio::select! { + biased; + _ = self.cancellation.cancelled() => failure_message(&self.failure), + result = &mut self.driver_task => match result { + Ok(()) => "code-mode connection driver exited unexpectedly".to_string(), + Err(err) => format!("code-mode connection driver task failed: {err}"), + }, + result = &mut self.reader_task => task_failure("reader", result), + result = &mut self.writer_task => task_failure("writer", result), + result = self.child.wait() => { + child_exited = true; + match result { + Ok(status) => format!("code-mode host exited with status {status}"), + Err(err) => format!("failed waiting for code-mode host: {err}"), + } + } + }; + mark_connection_dead(&self.alive, &self.failure, reason.clone()); + let _ = self.event_tx.try_send(DriverEvent::Failed(reason)); + self.cancellation.cancel(); + if !child_exited { + kill_and_reap(&mut self.child).await; + } + } +} + +fn task_failure( + task_name: &str, + result: Result, tokio::task::JoinError>, +) -> String { + match result { + Ok(Ok(())) => format!("code-mode connection {task_name} exited unexpectedly"), + Ok(Err(err)) => err, + Err(err) => format!("code-mode connection {task_name} task failed: {err}"), + } +} + +fn mark_connection_dead( + alive: &AtomicBool, + failure: &std::sync::Mutex>, + reason: String, +) { + alive.store(false, Ordering::Release); + let mut failure = failure + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if failure.is_none() { + *failure = Some(reason); + } +} + +fn failure_message(failure: &std::sync::Mutex>) -> String { + failure + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() + .unwrap_or_else(|| "code-mode host connection closed".to_string()) +} + +async fn kill_and_reap(child: &mut Child) { + let _ = child.start_kill(); + let _ = child.wait().await; +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver.rs b/codex-rs/code-mode/src/remote_session/connection/driver.rs new file mode 100644 index 000000000..102ad4758 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver.rs @@ -0,0 +1,179 @@ +use std::panic::AssertUnwindSafe; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::host::EncodedFrame; +use codex_code_mode_protocol::host::RequestId; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +pub(in crate::remote_session) use self::cleanup::SessionCleanup; +use self::delegate_runtime::DelegateRuntime; +use self::request_tracker::RequestTracker; +use self::session_registry::SessionRegistry; +pub(super) use self::types::DriverCommand; +pub(super) use self::types::DriverEvent; +pub(in crate::remote_session) use self::types::RemoteSession; + +mod cell_ids; +mod cleanup; +mod commands; +mod delegate_runtime; +mod request_tracker; +mod responses; +mod session_registry; +mod types; + +pub(super) struct DriverLifecycle { + pub(super) alive: Arc, + pub(super) failure: Arc>>, + pub(super) cancellation: CancellationToken, +} + +pub(super) struct ConnectionDriver { + command_rx: mpsc::Receiver, + event_rx: mpsc::Receiver, + event_tx: mpsc::Sender, + execute_claim_rx: mpsc::UnboundedReceiver, + outgoing_tx: mpsc::Sender, + requests: RequestTracker, + sessions: SessionRegistry, + delegates: DelegateRuntime, + alive: Arc, + failure: Arc>>, + cancellation: CancellationToken, + failed: bool, +} + +impl ConnectionDriver { + pub(super) fn new( + command_rx: mpsc::Receiver, + event_rx: mpsc::Receiver, + event_tx: mpsc::Sender, + outgoing_tx: mpsc::Sender, + lifecycle: DriverLifecycle, + ) -> (Self, mpsc::UnboundedSender) { + let (execute_claim_tx, execute_claim_rx) = mpsc::unbounded_channel(); + ( + Self { + command_rx, + event_rx, + event_tx: event_tx.clone(), + execute_claim_rx, + outgoing_tx, + requests: RequestTracker::new(), + sessions: SessionRegistry::new(), + delegates: DelegateRuntime::new(event_tx), + alive: lifecycle.alive, + failure: lifecycle.failure, + cancellation: lifecycle.cancellation, + failed: false, + }, + execute_claim_tx, + ) + } + + pub(super) async fn run(mut self) { + loop { + tokio::select! { + biased; + _ = self.cancellation.cancelled() => { + self.fail("code-mode host connection closed".to_string()); + return; + } + event = self.event_rx.recv() => { + let Some(event) = event else { + self.fail("code-mode host event stream closed".to_string()); + return; + }; + if !self.cancel_dropped_callers() || !self.handle_event(event) { + return; + } + } + claim = self.execute_claim_rx.recv() => { + let Some(request_id) = claim else { + self.fail("code-mode execute claim stream closed".to_string()); + return; + }; + self.requests.claim_execute(request_id); + } + command = self.command_rx.recv() => { + let Some(command) = command else { + self.fail("code-mode host command stream closed".to_string()); + return; + }; + if !self.cancel_dropped_callers() || !self.handle_command(command) { + return; + } + } + } + } + } + + fn handle_event(&mut self, event: DriverEvent) -> bool { + let keep_running = match event { + DriverEvent::HostMessage(message) => self.handle_host_message(message), + DriverEvent::DelegateCompleted { id, result } => self.complete_delegate(id, result), + DriverEvent::RequestCancelled(id) => self.cancel_request(id), + DriverEvent::Failed(reason) => { + self.fail(reason); + false + } + }; + if keep_running { + self.flush_deferred_waits() + } else { + false + } + } + + fn queue_frame(&mut self, frame: EncodedFrame) -> bool { + match self.outgoing_tx.try_send(frame) { + Ok(()) => true, + Err(mpsc::error::TrySendError::Full(_)) => { + self.fail("code-mode host outgoing queue is full".to_string()); + false + } + Err(mpsc::error::TrySendError::Closed(_)) => { + self.fail("code-mode host writer closed".to_string()); + false + } + } + } + + fn fail(&mut self, reason: String) { + if self.failed { + return; + } + self.failed = true; + self.alive.store(false, Ordering::Release); + let reason = { + let mut failure = self + .failure + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + failure.get_or_insert(reason).clone() + }; + self.requests.fail_all(&reason); + let failed_sessions = self.sessions.drain(); + self.delegates.fail_all(failed_sessions); + self.cancellation.cancel(); + } +} + +impl Drop for ConnectionDriver { + fn drop(&mut self) { + self.fail("code-mode connection driver stopped unexpectedly".to_string()); + } +} + +fn notify_cell_closed(delegate: &Arc, cell_id: &CellId) { + let _ = std::panic::catch_unwind(AssertUnwindSafe(|| delegate.cell_closed(cell_id))); +} + +#[cfg(test)] +#[path = "driver_tests.rs"] +mod tests; diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/cell_ids.rs b/codex-rs/code-mode/src/remote_session/connection/driver/cell_ids.rs new file mode 100644 index 000000000..666f6edda --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver/cell_ids.rs @@ -0,0 +1,111 @@ +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::RuntimeResponse; +use codex_code_mode_protocol::WaitOutcome; +use codex_code_mode_protocol::WaitRequest; +use codex_code_mode_protocol::host::WireCellId; +use codex_code_mode_protocol::host::WireRuntimeResponse; +use codex_code_mode_protocol::host::WireWaitOutcome; +use codex_code_mode_protocol::host::WireWaitRequest; + +use super::RemoteSession; + +pub(super) fn public_cell_id(generation: u64, cell_id: &WireCellId) -> CellId { + if generation == 1 { + CellId::new(cell_id.as_str().to_string()) + } else { + CellId::new(format!("g{generation}:{}", cell_id.as_str())) + } +} + +pub(super) fn public_cell_id_from_protocol(generation: u64, cell_id: &CellId) -> CellId { + public_cell_id(generation, &WireCellId::new(cell_id.as_str())) +} + +pub(super) fn remote_cell_id( + session: &RemoteSession, + cell_id: &CellId, +) -> Result { + if session.generation == 1 { + if cell_id.as_str().starts_with('g') && cell_id.as_str().contains(':') { + return Err(format!( + "cell {cell_id} belongs to a stale code-mode host generation" + )); + } + return Ok(WireCellId::new(cell_id.as_str())); + } + let prefix = format!("g{}:", session.generation); + let Some(remote_id) = cell_id.as_str().strip_prefix(&prefix) else { + return Err(format!( + "cell {cell_id} belongs to a stale code-mode host generation" + )); + }; + Ok(WireCellId::new(remote_id)) +} + +pub(super) fn remote_wait_request( + session: &RemoteSession, + request: WaitRequest, +) -> Result { + Ok(WireWaitRequest { + cell_id: remote_cell_id(session, &request.cell_id)?, + yield_time_ms: request.yield_time_ms, + }) +} + +pub(super) fn public_runtime_response( + generation: u64, + response: RuntimeResponse, +) -> RuntimeResponse { + match response { + RuntimeResponse::Yielded { + cell_id, + content_items, + } => RuntimeResponse::Yielded { + cell_id: public_cell_id_from_protocol(generation, &cell_id), + content_items, + }, + RuntimeResponse::Terminated { + cell_id, + content_items, + } => RuntimeResponse::Terminated { + cell_id: public_cell_id_from_protocol(generation, &cell_id), + content_items, + }, + RuntimeResponse::Result { + cell_id, + content_items, + error_text, + } => RuntimeResponse::Result { + cell_id: public_cell_id_from_protocol(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)) + } + } +} + +pub(super) fn runtime_response_cell_id(response: &WireRuntimeResponse) -> &WireCellId { + match response { + WireRuntimeResponse::Yielded { cell_id, .. } + | WireRuntimeResponse::Terminated { cell_id, .. } + | WireRuntimeResponse::Result { cell_id, .. } => cell_id, + } +} + +pub(super) fn wait_outcome_cell_id(outcome: &WireWaitOutcome) -> &WireCellId { + match outcome { + WireWaitOutcome::LiveCell(response) | WireWaitOutcome::MissingCell(response) => { + runtime_response_cell_id(response) + } + } +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/cleanup.rs b/codex-rs/code-mode/src/remote_session/connection/driver/cleanup.rs new file mode 100644 index 000000000..b994b6488 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver/cleanup.rs @@ -0,0 +1,40 @@ +use std::sync::Arc; + +use tokio_util::sync::CancellationToken; + +use super::notify_cell_closed; +use super::session_registry::CellOwner; + +struct CleanupInner { + complete: CancellationToken, +} + +#[derive(Clone)] +pub(in crate::remote_session) struct SessionCleanup { + inner: Arc, +} + +impl SessionCleanup { + pub(in crate::remote_session) fn new() -> Self { + Self { + inner: Arc::new(CleanupInner { + complete: CancellationToken::new(), + }), + } + } + + pub(super) fn fail(&self, cells: Vec) { + for owner in cells { + notify_cell_closed(&owner.delegate, &owner.cell_id); + } + self.inner.complete.cancel(); + } + + pub(in crate::remote_session) async fn wait(&self) { + self.inner.complete.cancelled().await; + } + + pub(in crate::remote_session) fn is_complete(&self) -> bool { + self.inner.complete.is_cancelled() + } +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs b/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs new file mode 100644 index 000000000..c9cba9c32 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs @@ -0,0 +1,298 @@ +use std::sync::Arc; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::ExecuteRequest; +use codex_code_mode_protocol::WaitOutcome; +use codex_code_mode_protocol::WaitRequest; +use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::EncodedFrame; +use codex_code_mode_protocol::host::HostRequest; +use codex_code_mode_protocol::host::WireWaitRequest; +use tokio::sync::oneshot; +use tokio_util::sync::CancellationToken; + +use super::ConnectionDriver; +use super::cell_ids::remote_cell_id; +use super::cell_ids::remote_wait_request; +use super::types::CancellableRequest; +use super::types::DeferredWait; +use super::types::DeliveredExecute; +use super::types::DriverCommand; +use super::types::PendingRequest; +use super::types::RemoteSession; + +impl ConnectionDriver { + pub(super) fn handle_command(&mut self, command: DriverCommand) -> bool { + match command { + DriverCommand::OpenSession { + session, + delegate, + cleanup, + caller_cancellation, + response_tx, + } => self.open_session(session, delegate, cleanup, caller_cancellation, response_tx), + DriverCommand::Execute { + session, + request, + caller_cancellation, + response_tx, + } => self.execute(session, request, caller_cancellation, response_tx), + DriverCommand::Wait { + session, + request, + caller_cancellation, + response_tx, + } => self.wait(session, request, caller_cancellation, response_tx), + DriverCommand::Terminate { + session, + cell_id, + response_tx, + } => self.terminate(session, cell_id, response_tx), + DriverCommand::ShutdownSession { + session, + response_tx, + } => self.shutdown_session(session, response_tx), + } + } + + fn open_session( + &mut self, + session: RemoteSession, + delegate: Arc, + cleanup: super::cleanup::SessionCleanup, + caller_cancellation: CancellationToken, + response_tx: oneshot::Sender>, + ) -> bool { + if self.sessions.contains(&session.id) || self.requests.contains_pending_open(&session) { + let _ = response_tx.send(Err(format!( + "code-mode session {} is already open", + session.id + ))); + return true; + } + let request_id = match self.requests.allocate_id() { + Ok(id) => id, + Err(err) => { + let _ = response_tx.send(Err(err)); + return false; + } + }; + let message = ClientToHost::Request { + id: request_id, + request: HostRequest::OpenSession { + session_id: session.id.clone(), + }, + }; + let frame = match EncodedFrame::encode(&message) { + Ok(frame) => frame, + Err(err) => { + let _ = response_tx.send(Err(format!( + "failed to encode code-mode open-session request: {err}" + ))); + return true; + } + }; + let cancellation = CancellableRequest::new(caller_cancellation); + self.requests.insert_pending( + request_id, + PendingRequest::OpenSession { + session, + delegate, + cleanup, + cancellation, + response_tx, + }, + &self.event_tx, + ); + self.queue_frame(frame) + } + + fn execute( + &mut self, + session: RemoteSession, + request: ExecuteRequest, + caller_cancellation: CancellationToken, + response_tx: oneshot::Sender>, + ) -> bool { + if let Err(err) = self.sessions.require_ready(&session) { + let _ = response_tx.send(Err(err)); + return true; + } + let request = match request.try_into() { + Ok(request) => request, + Err(err) => { + let _ = response_tx.send(Err(format!( + "failed to encode code-mode execute request: {err}" + ))); + return true; + } + }; + let request_id = match self.requests.allocate_id() { + Ok(id) => id, + Err(err) => { + let _ = response_tx.send(Err(err)); + return false; + } + }; + let message = ClientToHost::Request { + id: request_id, + request: HostRequest::Execute { + session_id: session.id.clone(), + request, + }, + }; + let frame = match EncodedFrame::encode(&message) { + Ok(frame) => frame, + Err(err) => { + let _ = response_tx.send(Err(format!( + "code-mode execute request exceeds the IPC frame limit: {err}" + ))); + return true; + } + }; + let (initial_response_tx, initial_response_rx) = oneshot::channel(); + let cancellation = CancellableRequest::new(caller_cancellation); + self.requests.insert_pending( + request_id, + PendingRequest::Execute { + session, + response_tx, + initial_response_tx, + initial_response_rx, + cancellation, + }, + &self.event_tx, + ); + self.queue_frame(frame) + } + + fn wait( + &mut self, + session: RemoteSession, + request: WaitRequest, + caller_cancellation: CancellationToken, + response_tx: oneshot::Sender>, + ) -> bool { + if let Err(err) = self.sessions.require_ready(&session) { + let _ = response_tx.send(Err(err)); + return true; + } + let request = match remote_wait_request(&session, request) { + Ok(request) => request, + Err(err) => { + let _ = response_tx.send(Err(err)); + return true; + } + }; + if self.requests.has_cancelled_wait(&session, &request.cell_id) { + self.requests.push_deferred_wait(DeferredWait { + session, + request, + caller_cancellation, + response_tx, + }); + return true; + } + self.start_wait(session, request, caller_cancellation, response_tx) + } + + pub(super) fn start_wait( + &mut self, + session: RemoteSession, + request: WireWaitRequest, + caller_cancellation: CancellationToken, + response_tx: oneshot::Sender>, + ) -> bool { + let cell_id = request.cell_id.clone(); + self.send_request( + HostRequest::Wait { + session_id: session.id.clone(), + request, + }, + PendingRequest::Wait { + session, + cell_id, + cancellation: CancellableRequest::new(caller_cancellation), + response_tx, + }, + ) + } + + fn terminate( + &mut self, + session: RemoteSession, + cell_id: CellId, + response_tx: oneshot::Sender>, + ) -> bool { + if let Err(err) = self.sessions.require_ready(&session) { + let _ = response_tx.send(Err(err)); + return true; + } + let cell_id = match remote_cell_id(&session, &cell_id) { + Ok(cell_id) => cell_id, + Err(err) => { + let _ = response_tx.send(Err(err)); + return true; + } + }; + let pending_cell_id = cell_id.clone(); + self.send_request( + HostRequest::Terminate { + session_id: session.id.clone(), + cell_id, + }, + PendingRequest::Terminate { + session, + cell_id: pending_cell_id, + response_tx, + }, + ) + } + + fn shutdown_session( + &mut self, + session: RemoteSession, + response_tx: oneshot::Sender>, + ) -> bool { + if let Err(err) = self.sessions.begin_shutdown(&session) { + let _ = response_tx.send(Err(err)); + return true; + } + self.send_request( + HostRequest::ShutdownSession { + session_id: session.id.clone(), + }, + PendingRequest::ShutdownSession { + session, + response_tx, + }, + ) + } + + pub(super) fn send_request(&mut self, request: HostRequest, pending: PendingRequest) -> bool { + let request_id = match self.requests.allocate_id() { + Ok(id) => id, + Err(err) => { + pending.fail(err); + return false; + } + }; + let message = ClientToHost::Request { + id: request_id, + request, + }; + let frame = match EncodedFrame::encode(&message) { + Ok(frame) => frame, + Err(err) => { + pending.fail(format!( + "code-mode request exceeds the IPC frame limit: {err}" + )); + return true; + } + }; + self.requests + .insert_pending(request_id, pending, &self.event_tx); + self.queue_frame(frame) + } +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/delegate_runtime.rs b/codex-rs/code-mode/src/remote_session/connection/driver/delegate_runtime.rs new file mode 100644 index 000000000..c79270661 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver/delegate_runtime.rs @@ -0,0 +1,339 @@ +//! Client-side delegate task and closure lifecycle. +//! +//! Cancellation revokes the task's completion path before removing its active-call state. The +//! delegate future may finish later, but it can no longer send a response or affect cell closure. + +use std::collections::HashMap; +use std::collections::HashSet; +use std::collections::VecDeque; + +use codex_code_mode_protocol::CodeModeNestedToolCall; +use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::DelegateRequest; +use codex_code_mode_protocol::host::DelegateRequestId; +use codex_code_mode_protocol::host::DelegateResponse; +use codex_code_mode_protocol::host::EncodedFrame; +use codex_code_mode_protocol::host::SessionId; +use codex_code_mode_protocol::host::WireCellId; +use codex_code_mode_protocol::host::WireResult; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +use super::ConnectionDriver; +use super::notify_cell_closed; +use super::session_registry::CellOwner; +use super::session_registry::DelegateTarget; +use super::session_registry::FailedSession; +use super::types::DriverEvent; + +const MAX_RECENT_DELEGATE_REQUEST_IDS: usize = 4096; + +#[derive(Clone, Eq, Hash, PartialEq)] +struct CellKey { + session_id: codex_code_mode_protocol::host::SessionId, + cell_id: codex_code_mode_protocol::CellId, +} + +impl CellKey { + fn for_owner(owner: &CellOwner) -> Self { + Self { + session_id: owner.session_id.clone(), + cell_id: owner.cell_id.clone(), + } + } +} + +struct DelegateCall { + cell: CellKey, + cancellation: CancellationToken, + completion_stop: CancellationToken, +} + +impl DelegateCall { + fn revoke(&self) { + self.cancellation.cancel(); + self.completion_stop.cancel(); + } +} + +enum DelegateTask { + InvokeTool(CodeModeNestedToolCall), + Notify { + call_id: String, + cell_id: codex_code_mode_protocol::CellId, + text: String, + }, +} + +pub(super) struct DelegateEffects { + pub(super) response: Option<(DelegateRequestId, Result)>, + pub(super) closed_cells: Vec, +} + +impl DelegateEffects { + fn empty() -> Self { + Self { + response: None, + closed_cells: Vec::new(), + } + } + + fn append(&mut self, mut other: Self) { + debug_assert!(self.response.is_none()); + self.response = other.response.take(); + self.closed_cells.append(&mut other.closed_cells); + } +} + +pub(super) struct DelegateRuntime { + calls: HashMap, + seen_requests: HashSet, + request_order: VecDeque, + event_tx: mpsc::Sender, +} + +impl DelegateRuntime { + pub(super) fn new(event_tx: mpsc::Sender) -> Self { + Self { + calls: HashMap::new(), + seen_requests: HashSet::new(), + request_order: VecDeque::new(), + event_tx, + } + } + + pub(super) fn start( + &mut self, + id: DelegateRequestId, + target: DelegateTarget, + request: DelegateRequest, + ) -> Result<(), String> { + if self.calls.contains_key(&id) || self.seen_requests.contains(&id) { + return Err(format!("duplicate code-mode delegate request ID {id:?}")); + } + self.remember_request(id); + let cancellation = CancellationToken::new(); + let task_request = match request { + DelegateRequest::InvokeTool { invocation } => { + let mut invocation: CodeModeNestedToolCall = invocation.into(); + invocation.cell_id = target.cell_id.clone(); + DelegateTask::InvokeTool(invocation) + } + DelegateRequest::Notify { + call_id, + cell_id: _, + text, + } => DelegateTask::Notify { + call_id, + cell_id: target.cell_id.clone(), + text, + }, + }; + let delegate = target.delegate; + let task_cancellation = cancellation.clone(); + let delegate_task = tokio::spawn(async move { + match task_request { + DelegateTask::InvokeTool(invocation) => delegate + .invoke_tool(invocation, task_cancellation) + .await + .map(|result| DelegateResponse::ToolResult { result }), + DelegateTask::Notify { + call_id, + cell_id, + text, + } => delegate + .notify(call_id, cell_id, text, task_cancellation) + .await + .map(|()| DelegateResponse::NotificationDelivered), + } + }); + let completion_stop = CancellationToken::new(); + self.calls.insert( + id, + DelegateCall { + cell: CellKey { + session_id: target.session_id, + cell_id: target.cell_id, + }, + cancellation, + completion_stop: completion_stop.clone(), + }, + ); + let event_tx = self.event_tx.clone(); + tokio::spawn(async move { + let result = tokio::select! { + biased; + _ = completion_stop.cancelled() => return, + result = delegate_task => match result { + Ok(result) => result, + Err(err) => Err(format!("code-mode delegate task failed: {err}")), + }, + }; + tokio::select! { + biased; + _ = completion_stop.cancelled() => {} + _ = event_tx.send(DriverEvent::DelegateCompleted { id, result }) => {} + } + }); + Ok(()) + } + + pub(super) fn cancel(&mut self, id: DelegateRequestId) { + if let Some(call) = self.calls.remove(&id) { + call.revoke(); + } + } + + pub(super) fn complete( + &mut self, + id: DelegateRequestId, + result: Result, + ) -> DelegateEffects { + if self.calls.remove(&id).is_none() { + return DelegateEffects::empty(); + } + let mut effects = DelegateEffects::empty(); + effects.response = Some((id, result)); + effects + } + + pub(super) fn close_cell(&mut self, owner: CellOwner) -> DelegateEffects { + let key = CellKey::for_owner(&owner); + self.calls.retain(|_, call| { + if call.cell != key { + return true; + } + call.revoke(); + false + }); + let mut effects = DelegateEffects::empty(); + effects.closed_cells.push(owner); + effects + } + + pub(super) fn close_cells(&mut self, owners: Vec) -> DelegateEffects { + let mut effects = DelegateEffects::empty(); + for owner in owners { + effects.append(self.close_cell(owner)); + } + effects + } + + pub(super) fn fail_all(&mut self, failed_sessions: Vec) { + for (_, call) in self.calls.drain() { + call.revoke(); + } + for session in failed_sessions { + session.cleanup.fail(session.cells); + } + } + + fn remember_request(&mut self, id: DelegateRequestId) { + self.seen_requests.insert(id); + self.request_order.push_back(id); + while self.request_order.len() > MAX_RECENT_DELEGATE_REQUEST_IDS { + if let Some(expired) = self.request_order.pop_front() { + self.seen_requests.remove(&expired); + } + } + } +} + +impl ConnectionDriver { + pub(super) fn start_delegate( + &mut self, + id: DelegateRequestId, + session_id: SessionId, + request: DelegateRequest, + ) -> bool { + let wire_cell_id = match &request { + DelegateRequest::InvokeTool { invocation } => &invocation.cell_id, + DelegateRequest::Notify { cell_id, .. } => cell_id, + }; + let target = match self.sessions.delegate_target(&session_id, wire_cell_id) { + Ok(target) => target, + Err(err) => { + self.fail(err); + return false; + } + }; + if let Err(err) = self.delegates.start(id, target, request) { + self.fail(err); + return false; + } + true + } + + pub(super) fn complete_delegate( + &mut self, + id: DelegateRequestId, + result: Result, + ) -> bool { + let effects = self.delegates.complete(id, result); + self.apply_delegate_effects(effects) + } + + fn send_delegate_response( + &mut self, + id: DelegateRequestId, + result: Result, + ) -> bool { + let message = ClientToHost::DelegateResponse { + id, + result: WireResult::from_result(result), + }; + let frame = match EncodedFrame::encode(&message) { + Ok(frame) => frame, + Err(err) => { + let fallback = ClientToHost::DelegateResponse { + id, + result: WireResult::Err { + message: format!( + "code-mode delegate response exceeds the IPC frame limit: {err}" + ), + }, + }; + match EncodedFrame::encode(&fallback) { + Ok(frame) => frame, + Err(fallback_err) => { + self.fail(format!( + "failed to encode code-mode delegate error response: {fallback_err}" + )); + return false; + } + } + } + }; + self.queue_frame(frame) + } + + pub(super) fn close_cell(&mut self, session_id: SessionId, cell_id: WireCellId) -> bool { + let owner = match self.sessions.remove_cell(&session_id, &cell_id) { + Ok(owner) => owner, + Err(err) => { + self.fail(err); + return false; + } + }; + let effects = self.delegates.close_cell(owner); + self.apply_delegate_effects(effects) + } + + pub(super) fn close_session_locally(&mut self, session_id: &SessionId) -> DelegateEffects { + self.requests.remove_unclaimed_for_session(session_id); + let owners = self.sessions.remove_session(session_id); + self.delegates.close_cells(owners) + } + + pub(super) fn apply_delegate_effects(&mut self, effects: DelegateEffects) -> bool { + if let Some((id, result)) = effects.response + && !self.send_delegate_response(id, result) + { + return false; + } + for closed in effects.closed_cells { + notify_cell_closed(&closed.delegate, &closed.cell_id); + } + true + } +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/request_tracker.rs b/codex-rs/code-mode/src/remote_session/connection/driver/request_tracker.rs new file mode 100644 index 000000000..0e7fc0881 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver/request_tracker.rs @@ -0,0 +1,186 @@ +use std::collections::HashMap; +use std::collections::VecDeque; + +use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::SessionId; +use codex_code_mode_protocol::host::WireCellId; +use tokio::sync::mpsc; + +use super::types::DeferredWait; +use super::types::DriverEvent; +use super::types::InitialResponse; +use super::types::PendingRequest; +use super::types::RemoteSession; +use super::types::UnclaimedExecute; + +pub(super) enum CancellationAction { + Send(RequestId), + Terminate { + request_id: RequestId, + execute: UnclaimedExecute, + }, +} + +pub(super) struct RequestTracker { + pending: HashMap, + unclaimed_executes: HashMap, + initial_responses: HashMap, + deferred_waits: VecDeque, + next_request_id: i64, +} + +impl RequestTracker { + pub(super) fn new() -> Self { + Self { + pending: HashMap::new(), + unclaimed_executes: HashMap::new(), + initial_responses: HashMap::new(), + deferred_waits: VecDeque::new(), + next_request_id: 1, + } + } + + pub(super) fn contains_pending_open(&self, session: &RemoteSession) -> bool { + self.pending.values().any(|pending| { + matches!( + pending, + PendingRequest::OpenSession { + session: pending_session, + .. + } if pending_session.id == session.id + ) + }) + } + + pub(super) fn allocate_id(&mut self) -> Result { + let id = self.next_request_id; + self.next_request_id = self + .next_request_id + .checked_add(1) + .ok_or_else(|| "code-mode host request ID space exhausted".to_string())?; + Ok(RequestId::new(id)) + } + + pub(super) fn insert_pending( + &mut self, + id: RequestId, + pending: PendingRequest, + event_tx: &mpsc::Sender, + ) { + self.pending.insert(id, pending); + if let Some(cancellation) = self + .pending + .get_mut(&id) + .and_then(PendingRequest::cancellation_mut) + { + cancellation.spawn_watcher(id, event_tx.clone()); + } + } + + pub(super) fn remove_pending(&mut self, id: RequestId) -> Option { + self.pending.remove(&id) + } + + pub(super) fn insert_initial_response(&mut self, id: RequestId, response: InitialResponse) { + self.initial_responses.insert(id, response); + } + + pub(super) fn remove_initial_response(&mut self, id: RequestId) -> Option { + self.initial_responses.remove(&id) + } + + pub(super) fn insert_unclaimed_execute(&mut self, id: RequestId, execute: UnclaimedExecute) { + self.unclaimed_executes.insert(id, execute); + } + + pub(super) fn claim_execute(&mut self, id: RequestId) { + self.unclaimed_executes.remove(&id); + } + + pub(super) fn collect_cancellations(&mut self) -> Vec { + let mut actions = self + .pending + .iter_mut() + .filter_map(|(id, pending)| { + let cancellation = pending.cancellation_mut()?; + (cancellation.is_cancelled() && cancellation.mark_reported()) + .then_some(CancellationAction::Send(*id)) + }) + .collect::>(); + actions.extend( + self.unclaimed_executes + .extract_if(|_, execute| { + execute.cancellation.is_cancelled() && execute.cancellation.mark_reported() + }) + .map(|(request_id, execute)| CancellationAction::Terminate { + request_id, + execute, + }), + ); + actions + } + + pub(super) fn mark_cancelled(&mut self, id: RequestId) -> Option { + if let Some(cancellation) = self + .pending + .get_mut(&id) + .and_then(PendingRequest::cancellation_mut) + { + return cancellation + .mark_reported() + .then_some(CancellationAction::Send(id)); + } + let execute = self.unclaimed_executes.get_mut(&id)?; + if !execute.cancellation.mark_reported() { + return None; + } + self.unclaimed_executes + .remove(&id) + .map(|execute| CancellationAction::Terminate { + request_id: id, + execute, + }) + } + + pub(super) fn has_cancelled_wait(&self, session: &RemoteSession, cell_id: &WireCellId) -> bool { + self.pending.values().any(|pending| { + matches!( + pending, + PendingRequest::Wait { + session: pending_session, + cell_id: pending_cell_id, + cancellation, + .. + } if pending_session == session + && pending_cell_id == cell_id + && cancellation.is_cancelled() + ) + }) + } + + pub(super) fn push_deferred_wait(&mut self, wait: DeferredWait) { + self.deferred_waits.push_back(wait); + } + + pub(super) fn take_deferred_waits(&mut self) -> VecDeque { + std::mem::take(&mut self.deferred_waits) + } + + pub(super) fn remove_unclaimed_for_session(&mut self, session_id: &SessionId) { + self.unclaimed_executes + .retain(|_, execute| &execute.session.id != session_id); + } + + pub(super) fn fail_all(&mut self, reason: &str) { + for (_, pending) in self.pending.drain() { + pending.fail(reason.to_string()); + } + self.unclaimed_executes.clear(); + for (_, initial) in self.initial_responses.drain() { + let _ = initial.response_tx.send(Err(reason.to_string())); + } + for wait in self.deferred_waits.drain(..) { + let _ = wait.response_tx.send(Err(reason.to_string())); + } + } +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/responses.rs b/codex-rs/code-mode/src/remote_session/connection/driver/responses.rs new file mode 100644 index 000000000..19c78ec00 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver/responses.rs @@ -0,0 +1,397 @@ +use codex_code_mode_protocol::StartedCell; +use codex_code_mode_protocol::host::ClientToHost; +use codex_code_mode_protocol::host::EncodedFrame; +use codex_code_mode_protocol::host::HostRequest; +use codex_code_mode_protocol::host::HostResponse; +use codex_code_mode_protocol::host::HostToClient; +use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::WireCellId; +use tokio::sync::oneshot; + +use super::ConnectionDriver; +use super::cell_ids::public_runtime_response; +use super::cell_ids::public_wait_outcome; +use super::cell_ids::runtime_response_cell_id; +use super::cell_ids::wait_outcome_cell_id; +use super::request_tracker::CancellationAction; +use super::session_registry::CellAdmissionError; +use super::types::DeliveredExecute; +use super::types::InitialResponse; +use super::types::PendingRequest; +use super::types::RemoteSession; +use super::types::UnclaimedExecute; + +impl ConnectionDriver { + pub(super) fn flush_deferred_waits(&mut self) -> bool { + let mut deferred = self.requests.take_deferred_waits(); + while let Some(wait) = deferred.pop_front() { + if wait.caller_cancellation.is_cancelled() { + let _ = wait + .response_tx + .send(Err("code-mode request cancelled".to_string())); + continue; + } + if self + .requests + .has_cancelled_wait(&wait.session, &wait.request.cell_id) + { + self.requests.push_deferred_wait(wait); + continue; + } + if !self.start_wait( + wait.session, + wait.request, + wait.caller_cancellation, + wait.response_tx, + ) { + for wait in deferred { + let _ = wait + .response_tx + .send(Err("code-mode host connection closed".to_string())); + } + return false; + } + } + true + } + + pub(super) fn handle_host_message(&mut self, message: HostToClient) -> bool { + match message { + HostToClient::Response { id, result } => { + self.complete_request(id, result.into_result()) + } + HostToClient::InitialResponse { id, result } => { + self.complete_initial_response(id, result.into_result()) + } + HostToClient::DelegateRequest { + id, + session_id, + request, + } => self.start_delegate(id, session_id, request), + HostToClient::CancelDelegateRequest { id } => { + self.delegates.cancel(id); + true + } + HostToClient::CellClosed { + session_id, + cell_id, + } => self.close_cell(session_id, cell_id), + HostToClient::HostHello(_) | HostToClient::HandshakeRejected { .. } => { + self.fail("code-mode host sent a second handshake response".to_string()); + false + } + } + } + + fn complete_request(&mut self, id: RequestId, result: Result) -> bool { + let Some(pending) = self.requests.remove_pending(id) else { + self.fail(format!("code-mode host returned unknown request ID {id:?}")); + return false; + }; + match pending { + PendingRequest::OpenSession { + session, + delegate, + cleanup, + cancellation, + response_tx, + } => match result { + Ok(HostResponse::SessionReady { session_id }) if session_id == session.id => { + let abandoned = cancellation.is_cancelled() || response_tx.is_closed(); + self.sessions + .insert_ready(session.clone(), delegate, cleanup); + if abandoned || response_tx.send(Ok(())).is_err() { + return self.shutdown_abandoned_session(session); + } + } + Ok(_) => { + let reason = + "code-mode host returned an invalid open-session response".to_string(); + let _ = response_tx.send(Err(reason.clone())); + self.fail(reason); + return false; + } + Err(err) => { + let _ = response_tx.send(Err(err)); + } + }, + PendingRequest::Execute { + session, + response_tx, + initial_response_tx, + initial_response_rx, + cancellation, + } => match result { + Ok(HostResponse::ExecutionStarted { cell_id }) => { + // The host owns a checked, never-reused ID sequence. Retain only live + // IDs so client memory scales with concurrency, not session lifetime. + let remote_cell_id = cell_id.clone(); + let public_id = match self.sessions.admit_cell(&session, cell_id) { + Ok(public_id) => public_id, + Err(CellAdmissionError::MissingSession) => { + let _ = response_tx + .send(Err("code-mode session closed during execute".to_string())); + return true; + } + Err(CellAdmissionError::DuplicateCell) => { + let reason = format!( + "code-mode host reused live cell {} in session {}", + remote_cell_id.as_str(), + session.id + ); + let _ = response_tx.send(Err(reason.clone())); + self.fail(reason); + return false; + } + }; + self.requests.insert_initial_response( + id, + InitialResponse { + generation: session.generation, + cell_id: remote_cell_id.clone(), + response_tx: initial_response_tx, + }, + ); + let started = StartedCell::from_result_receiver(public_id, initial_response_rx); + if cancellation.is_cancelled() || response_tx.is_closed() { + return self.terminate_abandoned_cell(session, remote_cell_id); + } + let delivered = DeliveredExecute { + request_id: id, + started, + }; + if response_tx.send(Ok(delivered)).is_err() { + return self.terminate_abandoned_cell(session, remote_cell_id); + } + self.requests.insert_unclaimed_execute( + id, + UnclaimedExecute { + session, + cell_id: remote_cell_id, + cancellation, + }, + ); + } + Ok(_) => { + let reason = "code-mode host returned an invalid execute response".to_string(); + let _ = response_tx.send(Err(reason.clone())); + self.fail(reason); + return false; + } + Err(err) => { + let _ = response_tx.send(Err(err)); + } + }, + PendingRequest::Wait { + session, + cell_id, + cancellation: _, + response_tx, + } => { + let result = match result { + Ok(HostResponse::WaitCompleted { outcome }) => { + if wait_outcome_cell_id(&outcome) != &cell_id { + let reason = format!( + "code-mode host returned cell {} for request targeting {}", + wait_outcome_cell_id(&outcome).as_str(), + cell_id.as_str() + ); + let _ = response_tx.send(Err(reason.clone())); + self.fail(reason); + return false; + } + Ok(public_wait_outcome(session.generation, outcome.into())) + } + Ok(_) => { + let reason = "code-mode host returned an invalid cell response".to_string(); + let _ = response_tx.send(Err(reason.clone())); + self.fail(reason); + return false; + } + Err(err) => Err(err), + }; + let _ = response_tx.send(result); + } + PendingRequest::Terminate { + session, + cell_id, + response_tx, + } => { + let result = match result { + Ok(HostResponse::WaitCompleted { outcome }) => { + if wait_outcome_cell_id(&outcome) != &cell_id { + let reason = format!( + "code-mode host returned cell {} for request targeting {}", + wait_outcome_cell_id(&outcome).as_str(), + cell_id.as_str() + ); + let _ = response_tx.send(Err(reason.clone())); + self.fail(reason); + return false; + } + public_wait_outcome(session.generation, outcome.into()) + } + Ok(_) => { + let reason = "code-mode host returned an invalid cell response".to_string(); + let _ = response_tx.send(Err(reason.clone())); + self.fail(reason); + return false; + } + Err(err) => { + let _ = response_tx.send(Err(err)); + return true; + } + }; + let _ = response_tx.send(Ok(result)); + } + PendingRequest::ShutdownSession { + session, + response_tx, + } => match result { + Ok(HostResponse::SessionClosed { session_id }) if session_id == session.id => { + let effects = self.close_session_locally(&session.id); + if !self.apply_delegate_effects(effects) { + return false; + } + let _ = response_tx.send(Ok(())); + } + Ok(_) => { + let err = "code-mode host returned an invalid shutdown response".to_string(); + let _ = response_tx.send(Err(err.clone())); + self.fail(err); + return false; + } + Err(err) => { + let _ = response_tx.send(Err(err.clone())); + self.fail(err); + return false; + } + }, + } + true + } + + pub(super) fn cancel_dropped_callers(&mut self) -> bool { + for action in self.requests.collect_cancellations() { + if !self.apply_cancellation(action) { + return false; + } + } + true + } + + pub(super) fn cancel_request(&mut self, id: RequestId) -> bool { + self.requests + .mark_cancelled(id) + .is_none_or(|action| self.apply_cancellation(action)) + } + + fn apply_cancellation(&mut self, action: CancellationAction) -> bool { + match action { + CancellationAction::Send(id) => self.send_cancel_request(id), + CancellationAction::Terminate { + request_id, + execute, + } => { + if !self.send_cancel_request(request_id) { + return false; + } + self.terminate_abandoned_cell(execute.session, execute.cell_id) + } + } + } + + fn send_cancel_request(&mut self, id: RequestId) -> bool { + let frame = match EncodedFrame::encode(&ClientToHost::CancelRequest { id }) { + Ok(frame) => frame, + Err(err) => { + self.fail(format!( + "failed to encode code-mode cancellation request: {err}" + )); + return false; + } + }; + self.queue_frame(frame) + } + + fn shutdown_abandoned_session(&mut self, session: RemoteSession) -> bool { + let Some(should_shutdown) = self.sessions.begin_abandoned_shutdown(&session.id) else { + self.fail(format!( + "code-mode host committed abandoned session {} without local state", + session.id + )); + return false; + }; + if !should_shutdown { + return true; + } + let (response_tx, response_rx) = oneshot::channel(); + drop(response_rx); + self.send_request( + HostRequest::ShutdownSession { + session_id: session.id.clone(), + }, + PendingRequest::ShutdownSession { + session, + response_tx, + }, + ) + } + + fn terminate_abandoned_cell(&mut self, session: RemoteSession, cell_id: WireCellId) -> bool { + let Some(is_closing) = self.sessions.is_closing(&session.id) else { + self.fail(format!( + "code-mode host admitted an abandoned cell in unknown session {}", + session.id + )); + return false; + }; + if is_closing { + return true; + } + let (response_tx, response_rx) = oneshot::channel(); + drop(response_rx); + self.send_request( + HostRequest::Terminate { + session_id: session.id.clone(), + cell_id: cell_id.clone(), + }, + PendingRequest::Terminate { + session, + cell_id, + response_tx, + }, + ) + } + + fn complete_initial_response( + &mut self, + id: RequestId, + result: Result, + ) -> bool { + let Some(initial) = self.requests.remove_initial_response(id) else { + self.fail(format!( + "code-mode host returned initial response for unknown request ID {id:?}" + )); + return false; + }; + let response = match result { + Ok(response) if runtime_response_cell_id(&response) == &initial.cell_id => { + Ok(public_runtime_response(initial.generation, response.into())) + } + Ok(response) => { + let reason = format!( + "code-mode host returned initial response for cell {} instead of {}", + runtime_response_cell_id(&response).as_str(), + initial.cell_id.as_str() + ); + let _ = initial.response_tx.send(Err(reason.clone())); + self.fail(reason); + return false; + } + Err(err) => Err(err), + }; + let _ = initial.response_tx.send(response); + true + } +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/session_registry.rs b/codex-rs/code-mode/src/remote_session/connection/driver/session_registry.rs new file mode 100644 index 000000000..a3d1c82ae --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver/session_registry.rs @@ -0,0 +1,222 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::host::SessionId; +use codex_code_mode_protocol::host::WireCellId; + +use super::cell_ids::public_cell_id; +use super::cleanup::SessionCleanup; +use super::types::RemoteSession; + +pub(super) struct CellOwner { + pub(super) session_id: SessionId, + pub(super) cell_id: CellId, + pub(super) delegate: Arc, +} + +pub(super) struct DelegateTarget { + pub(super) session_id: SessionId, + pub(super) cell_id: CellId, + pub(super) delegate: Arc, +} + +pub(super) struct FailedSession { + pub(super) cleanup: SessionCleanup, + pub(super) cells: Vec, +} + +pub(super) enum CellAdmissionError { + MissingSession, + DuplicateCell, +} + +struct SessionRecord { + remote: RemoteSession, + delegate: Arc, + cleanup: SessionCleanup, + phase: SessionPhase, + cells: HashMap, +} + +#[derive(Clone, Copy, Eq, PartialEq)] +enum SessionPhase { + Ready, + Closing, +} + +pub(super) struct SessionRegistry { + records: HashMap, +} + +impl SessionRegistry { + pub(super) fn new() -> Self { + Self { + records: HashMap::new(), + } + } + + pub(super) fn contains(&self, session_id: &SessionId) -> bool { + self.records.contains_key(session_id) + } + + pub(super) fn insert_ready( + &mut self, + session: RemoteSession, + delegate: Arc, + cleanup: SessionCleanup, + ) { + self.records.insert( + session.id.clone(), + SessionRecord { + remote: session, + delegate, + cleanup, + phase: SessionPhase::Ready, + cells: HashMap::new(), + }, + ); + } + + pub(super) fn require_ready(&self, session: &RemoteSession) -> Result<(), String> { + let record = self + .records + .get(&session.id) + .ok_or_else(|| format!("unknown code-mode session {}", session.id))?; + if record.remote != *session { + return Err("stale code-mode session generation".to_string()); + } + if record.phase != SessionPhase::Ready { + return Err("code-mode session is shutting down".to_string()); + } + Ok(()) + } + + pub(super) fn begin_shutdown(&mut self, session: &RemoteSession) -> Result<(), String> { + let record = self + .records + .get_mut(&session.id) + .ok_or_else(|| format!("unknown code-mode session {}", session.id))?; + if record.remote != *session { + return Err("stale code-mode session generation".to_string()); + } + if record.phase == SessionPhase::Closing { + return Err("code-mode session is already closing".to_string()); + } + record.phase = SessionPhase::Closing; + Ok(()) + } + + pub(super) fn begin_abandoned_shutdown(&mut self, session_id: &SessionId) -> Option { + let record = self.records.get_mut(session_id)?; + if record.phase == SessionPhase::Closing { + return Some(false); + } + record.phase = SessionPhase::Closing; + Some(true) + } + + pub(super) fn is_closing(&self, session_id: &SessionId) -> Option { + self.records + .get(session_id) + .map(|record| record.phase == SessionPhase::Closing) + } + + pub(super) fn admit_cell( + &mut self, + session: &RemoteSession, + cell_id: WireCellId, + ) -> Result { + let Some(record) = self.records.get_mut(&session.id) else { + return Err(CellAdmissionError::MissingSession); + }; + if record.cells.contains_key(&cell_id) { + return Err(CellAdmissionError::DuplicateCell); + } + let public_id = public_cell_id(session.generation, &cell_id); + record.cells.insert(cell_id, public_id.clone()); + Ok(public_id) + } + + pub(super) fn delegate_target( + &self, + session_id: &SessionId, + cell_id: &WireCellId, + ) -> Result { + let session = self + .records + .get(session_id) + .ok_or_else(|| format!("code-mode host delegated for unknown session {session_id}"))?; + let public_id = session.cells.get(cell_id).cloned().ok_or_else(|| { + format!( + "code-mode host delegated for unknown cell {} in session {session_id}", + cell_id.as_str() + ) + })?; + Ok(DelegateTarget { + session_id: session_id.clone(), + cell_id: public_id, + delegate: Arc::clone(&session.delegate), + }) + } + + pub(super) fn remove_cell( + &mut self, + session_id: &SessionId, + cell_id: &WireCellId, + ) -> Result { + let session = self.records.get_mut(session_id).ok_or_else(|| { + format!( + "code-mode host closed cell {} in unknown session {session_id}", + cell_id.as_str() + ) + })?; + let public_id = session + .cells + .remove(cell_id) + .ok_or_else(|| format!("code-mode host closed unknown cell in session {session_id}"))?; + Ok(CellOwner { + session_id: session_id.clone(), + cell_id: public_id, + delegate: Arc::clone(&session.delegate), + }) + } + + pub(super) fn remove_session(&mut self, session_id: &SessionId) -> Vec { + let Some(session) = self.records.remove(session_id) else { + return Vec::new(); + }; + session + .cells + .into_values() + .map(|cell_id| CellOwner { + session_id: session_id.clone(), + cell_id, + delegate: Arc::clone(&session.delegate), + }) + .collect() + } + + pub(super) fn drain(&mut self) -> Vec { + let sessions = std::mem::take(&mut self.records); + sessions + .into_iter() + .map(|(session_id, session)| { + let cells = session + .cells + .into_values() + .map(|cell_id| CellOwner { + session_id: session_id.clone(), + cell_id, + delegate: Arc::clone(&session.delegate), + }) + .collect(); + FailedSession { + cleanup: session.cleanup, + cells, + } + }) + .collect() + } +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/types.rs b/codex-rs/code-mode/src/remote_session/connection/driver/types.rs new file mode 100644 index 000000000..d4f6a102f --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver/types.rs @@ -0,0 +1,196 @@ +use std::sync::Arc; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSessionDelegate; +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::host::DelegateRequestId; +use codex_code_mode_protocol::host::DelegateResponse; +use codex_code_mode_protocol::host::HostToClient; +use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::SessionId; +use codex_code_mode_protocol::host::WireCellId; +use codex_code_mode_protocol::host::WireWaitRequest; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio_util::sync::CancellationToken; + +use super::cleanup::SessionCleanup; + +#[derive(Clone, Debug, Eq, PartialEq)] +pub(in crate::remote_session) struct RemoteSession { + pub(in crate::remote_session) id: SessionId, + pub(in crate::remote_session) generation: u64, +} + +pub(in crate::remote_session::connection) enum DriverCommand { + OpenSession { + session: RemoteSession, + delegate: Arc, + cleanup: SessionCleanup, + caller_cancellation: CancellationToken, + response_tx: oneshot::Sender>, + }, + Execute { + session: RemoteSession, + request: ExecuteRequest, + caller_cancellation: CancellationToken, + response_tx: oneshot::Sender>, + }, + Wait { + session: RemoteSession, + request: WaitRequest, + caller_cancellation: CancellationToken, + response_tx: oneshot::Sender>, + }, + Terminate { + session: RemoteSession, + cell_id: CellId, + response_tx: oneshot::Sender>, + }, + ShutdownSession { + session: RemoteSession, + response_tx: oneshot::Sender>, + }, +} + +pub(in crate::remote_session::connection) enum DriverEvent { + HostMessage(HostToClient), + DelegateCompleted { + id: DelegateRequestId, + result: Result, + }, + RequestCancelled(RequestId), + Failed(String), +} + +pub(super) struct CancellableRequest { + caller_cancellation: CancellationToken, + watcher_stop: CancellationToken, + reported: bool, +} + +impl CancellableRequest { + pub(super) fn new(caller_cancellation: CancellationToken) -> Self { + Self { + caller_cancellation, + watcher_stop: CancellationToken::new(), + reported: false, + } + } + + pub(super) fn is_cancelled(&self) -> bool { + self.caller_cancellation.is_cancelled() + } + + pub(super) fn mark_reported(&mut self) -> bool { + if self.reported { + return false; + } + self.reported = true; + true + } + + pub(super) fn spawn_watcher(&self, id: RequestId, event_tx: mpsc::Sender) { + let caller_cancellation = self.caller_cancellation.clone(); + let watcher_stop = self.watcher_stop.clone(); + tokio::spawn(async move { + tokio::select! { + _ = caller_cancellation.cancelled() => { + let _ = event_tx.send(DriverEvent::RequestCancelled(id)).await; + } + _ = watcher_stop.cancelled() => {} + } + }); + } +} + +impl Drop for CancellableRequest { + fn drop(&mut self) { + self.watcher_stop.cancel(); + } +} + +pub(super) struct InitialResponse { + pub(super) generation: u64, + pub(super) cell_id: WireCellId, + pub(super) response_tx: oneshot::Sender>, +} + +pub(in crate::remote_session::connection) struct DeliveredExecute { + pub(in crate::remote_session::connection) request_id: RequestId, + pub(in crate::remote_session::connection) started: StartedCell, +} + +pub(super) struct UnclaimedExecute { + pub(super) session: RemoteSession, + pub(super) cell_id: WireCellId, + pub(super) cancellation: CancellableRequest, +} + +pub(super) enum PendingRequest { + OpenSession { + session: RemoteSession, + delegate: Arc, + cleanup: SessionCleanup, + cancellation: CancellableRequest, + response_tx: oneshot::Sender>, + }, + Execute { + session: RemoteSession, + response_tx: oneshot::Sender>, + initial_response_tx: oneshot::Sender>, + initial_response_rx: oneshot::Receiver>, + cancellation: CancellableRequest, + }, + Wait { + session: RemoteSession, + cell_id: WireCellId, + cancellation: CancellableRequest, + response_tx: oneshot::Sender>, + }, + Terminate { + session: RemoteSession, + cell_id: WireCellId, + response_tx: oneshot::Sender>, + }, + ShutdownSession { + session: RemoteSession, + response_tx: oneshot::Sender>, + }, +} + +pub(super) struct DeferredWait { + pub(super) session: RemoteSession, + pub(super) request: WireWaitRequest, + pub(super) caller_cancellation: CancellationToken, + pub(super) response_tx: oneshot::Sender>, +} + +impl PendingRequest { + pub(super) fn cancellation_mut(&mut self) -> Option<&mut CancellableRequest> { + match self { + Self::OpenSession { cancellation, .. } + | Self::Execute { cancellation, .. } + | Self::Wait { cancellation, .. } => Some(cancellation), + Self::Terminate { .. } | Self::ShutdownSession { .. } => None, + } + } + + pub(super) fn fail(self, reason: String) { + match self { + Self::OpenSession { response_tx, .. } | Self::ShutdownSession { response_tx, .. } => { + let _ = response_tx.send(Err(reason)); + } + Self::Execute { response_tx, .. } => { + let _ = response_tx.send(Err(reason)); + } + Self::Wait { response_tx, .. } | Self::Terminate { response_tx, .. } => { + let _ = response_tx.send(Err(reason)); + } + } + } +} diff --git a/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs b/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs new file mode 100644 index 000000000..97ed9b553 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs @@ -0,0 +1,1514 @@ +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeNestedToolCall; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::ExecuteRequest; +use codex_code_mode_protocol::NotificationFuture; +use codex_code_mode_protocol::ToolInvocationFuture; +use codex_code_mode_protocol::WaitRequest; +use codex_code_mode_protocol::host::DelegateRequest; +use codex_code_mode_protocol::host::DelegateRequestId; +use codex_code_mode_protocol::host::HostResponse; +use codex_code_mode_protocol::host::HostToClient; +use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::SessionId; +use codex_code_mode_protocol::host::WireNestedToolCall; +use codex_code_mode_protocol::host::WireResult; +use codex_code_mode_protocol::host::WireRuntimeResponse; +use codex_code_mode_protocol::host::WireWaitOutcome; +use codex_protocol::ToolName; +use pretty_assertions::assert_eq; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio_util::sync::CancellationToken; + +use super::ConnectionDriver; +use super::DriverCommand; +use super::DriverEvent; +use super::DriverLifecycle; +use super::RemoteSession; +use super::SessionCleanup; + +struct DriverHarness { + command_tx: mpsc::Sender, + event_tx: mpsc::Sender, + execute_claim_tx: mpsc::UnboundedSender, + outgoing_rx: mpsc::Receiver, + cancellation: CancellationToken, + alive: Arc, + driver_task: tokio::task::JoinHandle<()>, +} + +impl DriverHarness { + fn start() -> Self { + let (command_tx, command_rx) = mpsc::channel(/*max_capacity*/ 16); + let (event_tx, event_rx) = mpsc::channel(/*max_capacity*/ 16); + let (outgoing_tx, outgoing_rx) = mpsc::channel(/*max_capacity*/ 16); + let cancellation = CancellationToken::new(); + let alive = Arc::new(AtomicBool::new(true)); + let (driver, execute_claim_tx) = ConnectionDriver::new( + command_rx, + event_rx, + event_tx.clone(), + outgoing_tx, + DriverLifecycle { + alive: Arc::clone(&alive), + failure: Arc::new(StdMutex::new(None)), + cancellation: cancellation.clone(), + }, + ); + let driver_task = tokio::spawn(driver.run()); + Self { + command_tx, + event_tx, + execute_claim_tx, + outgoing_rx, + cancellation, + alive, + driver_task, + } + } + + async fn open( + &mut self, + session: RemoteSession, + delegate: Arc, + ) -> SessionCleanup { + let cleanup = SessionCleanup::new(); + let (response_tx, response_rx) = oneshot::channel(); + self.command_tx + .send(DriverCommand::OpenSession { + session: session.clone(), + delegate, + cleanup: cleanup.clone(), + caller_cancellation: CancellationToken::new(), + response_tx, + }) + .await + .expect("open command"); + self.outgoing_rx.recv().await.expect("open frame"); + self.event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 1), + result: WireResult::Ok { + value: HostResponse::SessionReady { + session_id: session.id, + }, + }, + })) + .await + .expect("open response"); + response_rx + .await + .expect("open reply") + .expect("open session"); + cleanup + } + + async fn start_cell( + &mut self, + session: RemoteSession, + request_id: i64, + cell_id: &str, + ) -> codex_code_mode_protocol::StartedCell { + let (response_tx, response_rx) = oneshot::channel(); + self.command_tx + .send(DriverCommand::Execute { + session, + request: ExecuteRequest { + tool_call_id: format!("call-{request_id}"), + enabled_tools: Vec::new(), + source: "await new Promise(() => {})".to_string(), + yield_time_ms: Some(1), + max_output_tokens: None, + }, + caller_cancellation: CancellationToken::new(), + response_tx, + }) + .await + .expect("execute command"); + self.outgoing_rx.recv().await.expect("execute frame"); + self.event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(request_id), + result: WireResult::Ok { + value: HostResponse::ExecutionStarted { + cell_id: CellId::new(cell_id.to_string()).into(), + }, + }, + })) + .await + .expect("execute response"); + let delivered = response_rx + .await + .expect("execute reply") + .expect("execute session"); + self.execute_claim_tx + .send(delivered.request_id) + .expect("claim execute"); + delivered.started + } + + async fn start_tool_delegate(&self, session: &RemoteSession, id: DelegateRequestId) { + self.event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id, + session_id: session.id.clone(), + request: DelegateRequest::InvokeTool { + invocation: WireNestedToolCall { + cell_id: CellId::new("1".to_string()).into(), + runtime_tool_call_id: "tool-1".to_string(), + tool_name: ToolName::plain("slow").into(), + tool_kind: codex_code_mode_protocol::CodeModeToolKind::Function.into(), + input: None, + }, + }, + })) + .await + .expect("delegate request"); + } +} + +impl Drop for DriverHarness { + fn drop(&mut self) { + self.cancellation.cancel(); + } +} + +#[derive(Default)] +struct RecordingDelegate { + closed_cells: StdMutex>, + invocations: AtomicUsize, + notifications: AtomicUsize, +} + +struct PanickingDelegate; + +#[derive(Debug, Eq, PartialEq)] +enum HeldDelegateEvent { + Started, + Cancelled, + Finished, + CellClosed(CellId), +} + +struct HeldDelegate { + events_tx: mpsc::UnboundedSender, + release: CancellationToken, +} + +impl HeldDelegate { + fn new() -> ( + Arc, + mpsc::UnboundedReceiver, + CancellationToken, + ) { + let (events_tx, events_rx) = mpsc::unbounded_channel(); + let release = CancellationToken::new(); + ( + Arc::new(Self { + events_tx, + release: release.clone(), + }), + events_rx, + release, + ) + } +} + +impl CodeModeSessionDelegate for HeldDelegate { + fn invoke_tool<'a>( + &'a self, + _invocation: CodeModeNestedToolCall, + cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + let events_tx = self.events_tx.clone(); + let release = self.release.clone(); + Box::pin(async move { + let _ = events_tx.send(HeldDelegateEvent::Started); + cancellation_token.cancelled().await; + let _ = events_tx.send(HeldDelegateEvent::Cancelled); + release.cancelled().await; + let _ = events_tx.send(HeldDelegateEvent::Finished); + Err("cancelled".to_string()) + }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, cell_id: &CellId) { + let _ = self + .events_tx + .send(HeldDelegateEvent::CellClosed(cell_id.clone())); + } +} + +impl CodeModeSessionDelegate for PanickingDelegate { + fn invoke_tool<'a>( + &'a self, + _invocation: CodeModeNestedToolCall, + _cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + Box::pin(async { panic!("delegate panic probe") }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, _cell_id: &CellId) {} +} + +impl CodeModeSessionDelegate for RecordingDelegate { + fn invoke_tool<'a>( + &'a self, + _invocation: CodeModeNestedToolCall, + cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + self.invocations.fetch_add(1, Ordering::Relaxed); + Box::pin(async move { + cancellation_token.cancelled().await; + Err("cancelled".to_string()) + }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + self.notifications.fetch_add(1, Ordering::Relaxed); + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, cell_id: &CellId) { + self.closed_cells + .lock() + .expect("closed cells lock") + .push(cell_id.clone()); + } +} + +fn remote_session() -> RemoteSession { + RemoteSession { + id: SessionId::new("session-1").expect("session ID"), + generation: 1, + } +} + +async fn next_held_delegate_event( + events_rx: &mut mpsc::UnboundedReceiver, +) -> HeldDelegateEvent { + tokio::time::timeout(Duration::from_secs(1), events_rx.recv()) + .await + .expect("delegate event timeout") + .expect("delegate event stream") +} + +#[tokio::test] +async fn dropped_open_waiter_shuts_down_committed_session() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let (open_tx, open_rx) = oneshot::channel(); + let cleanup = SessionCleanup::new(); + harness + .command_tx + .send(DriverCommand::OpenSession { + session: session.clone(), + delegate: Arc::new(RecordingDelegate::default()), + cleanup, + caller_cancellation: CancellationToken::new(), + response_tx: open_tx, + }) + .await + .expect("open command"); + drop(open_rx); + harness.outgoing_rx.recv().await.expect("open frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 1), + result: WireResult::Ok { + value: HostResponse::SessionReady { + session_id: session.id.clone(), + }, + }, + })) + .await + .expect("open response"); + harness + .outgoing_rx + .recv() + .await + .expect("abandoned session shutdown frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: HostResponse::SessionClosed { + session_id: session.id.clone(), + }, + }, + })) + .await + .expect("shutdown response"); + + let (execute_tx, execute_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Execute { + session: session.clone(), + request: ExecuteRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: Vec::new(), + source: "text('ok')".to_string(), + yield_time_ms: None, + max_output_tokens: None, + }, + caller_cancellation: CancellationToken::new(), + response_tx: execute_tx, + }) + .await + .expect("execute command"); + assert_eq!( + execute_rx + .await + .expect("execute reply") + .err() + .expect("closed session should reject execute"), + "unknown code-mode session session-1" + ); +} + +#[tokio::test] +async fn delegate_cancel_is_best_effort_and_sends_no_late_response() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let request_id = DelegateRequestId::new(/*value*/ 7); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id: request_id, + session_id: session.id.clone(), + request: DelegateRequest::InvokeTool { + invocation: WireNestedToolCall { + cell_id: CellId::new("1".to_string()).into(), + runtime_tool_call_id: "tool-1".to_string(), + tool_name: ToolName::plain("slow").into(), + tool_kind: codex_code_mode_protocol::CodeModeToolKind::Function.into(), + input: None, + }, + }, + })) + .await + .expect("delegate request"); + harness + .event_tx + .send(DriverEvent::HostMessage( + HostToClient::CancelDelegateRequest { id: request_id }, + )) + .await + .expect("delegate cancel"); + tokio::task::yield_now().await; + assert!(matches!( + harness.outgoing_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id: request_id, + session_id: session.id, + request: DelegateRequest::Notify { + call_id: "notify-reused".to_string(), + cell_id: CellId::new("1".to_string()).into(), + text: "duplicate".to_string(), + }, + })) + .await + .expect("reused delegate request"); + tokio::task::yield_now().await; + + assert!(!harness.alive.load(Ordering::Acquire)); + assert_eq!(delegate.invocations.load(Ordering::Relaxed), 1); + assert_eq!(delegate.notifications.load(Ordering::Relaxed), 0); +} + +#[tokio::test] +async fn terminate_closes_cell_without_waiting_for_delegate_cleanup() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let (delegate, mut events_rx, release) = HeldDelegate::new(); + harness.open(session.clone(), delegate).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let delegate_id = DelegateRequestId::new(/*value*/ 7); + harness.start_tool_delegate(&session, delegate_id).await; + assert_eq!( + next_held_delegate_event(&mut events_rx).await, + HeldDelegateEvent::Started + ); + + let (response_tx, response_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Terminate { + session: session.clone(), + cell_id: CellId::new("1".to_string()), + response_tx, + }) + .await + .expect("terminate command"); + harness.outgoing_rx.recv().await.expect("terminate frame"); + harness + .event_tx + .send(DriverEvent::HostMessage( + HostToClient::CancelDelegateRequest { id: delegate_id }, + )) + .await + .expect("delegate cancel"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::CellClosed { + session_id: session.id, + cell_id: CellId::new("1".to_string()).into(), + })) + .await + .expect("cell close"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Ok { + value: HostResponse::WaitCompleted { + outcome: WireWaitOutcome::LiveCell(WireRuntimeResponse::Terminated { + cell_id: CellId::new("1".to_string()).into(), + content_items: Vec::new(), + }), + }, + }, + })) + .await + .expect("terminate response"); + + let closure_events = [ + next_held_delegate_event(&mut events_rx).await, + next_held_delegate_event(&mut events_rx).await, + ]; + assert!(closure_events.contains(&HeldDelegateEvent::Cancelled)); + assert!(closure_events.contains(&HeldDelegateEvent::CellClosed(CellId::new("1".to_string())))); + assert_eq!( + response_rx.await.expect("terminate reply"), + Ok(codex_code_mode_protocol::WaitOutcome::LiveCell( + codex_code_mode_protocol::RuntimeResponse::Terminated { + cell_id: CellId::new("1".to_string()), + content_items: Vec::new(), + } + )) + ); + assert!(matches!( + events_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) + )); + + release.cancel(); + assert_eq!( + next_held_delegate_event(&mut events_rx).await, + HeldDelegateEvent::Finished + ); + assert!(matches!( + events_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) + )); + assert!(matches!( + harness.outgoing_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + assert!(harness.alive.load(Ordering::Acquire)); +} + +#[tokio::test] +async fn shutdown_closes_cell_without_waiting_for_delegate_cleanup() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let (delegate, mut events_rx, release) = HeldDelegate::new(); + harness.open(session.clone(), delegate).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let delegate_id = DelegateRequestId::new(/*value*/ 7); + harness.start_tool_delegate(&session, delegate_id).await; + assert_eq!( + next_held_delegate_event(&mut events_rx).await, + HeldDelegateEvent::Started + ); + + let (response_tx, response_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::ShutdownSession { + session: session.clone(), + response_tx, + }) + .await + .expect("shutdown command"); + harness.outgoing_rx.recv().await.expect("shutdown frame"); + harness + .event_tx + .send(DriverEvent::HostMessage( + HostToClient::CancelDelegateRequest { id: delegate_id }, + )) + .await + .expect("delegate cancel"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::CellClosed { + session_id: session.id.clone(), + cell_id: CellId::new("1".to_string()).into(), + })) + .await + .expect("cell close"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Ok { + value: HostResponse::SessionClosed { + session_id: session.id, + }, + }, + })) + .await + .expect("shutdown response"); + + let closure_events = [ + next_held_delegate_event(&mut events_rx).await, + next_held_delegate_event(&mut events_rx).await, + ]; + assert!(closure_events.contains(&HeldDelegateEvent::Cancelled)); + assert!(closure_events.contains(&HeldDelegateEvent::CellClosed(CellId::new("1".to_string())))); + assert_eq!(response_rx.await.expect("shutdown reply"), Ok(())); + assert!(matches!( + events_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + + release.cancel(); + assert_eq!( + next_held_delegate_event(&mut events_rx).await, + HeldDelegateEvent::Finished + ); + assert!(matches!( + events_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) + )); + assert!(matches!( + harness.outgoing_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + assert!(harness.alive.load(Ordering::Acquire)); +} + +#[tokio::test] +async fn completed_delegate_request_id_cannot_be_reused() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let request_id = DelegateRequestId::new(/*value*/ 7); + let request = || DelegateRequest::Notify { + call_id: "notify-1".to_string(), + cell_id: CellId::new("1".to_string()).into(), + text: "once".to_string(), + }; + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id: request_id, + session_id: session.id.clone(), + request: request(), + })) + .await + .expect("delegate request"); + harness + .outgoing_rx + .recv() + .await + .expect("delegate response frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id: request_id, + session_id: session.id, + request: request(), + })) + .await + .expect("reused delegate request"); + tokio::task::yield_now().await; + + assert!(!harness.alive.load(Ordering::Acquire)); + assert_eq!(delegate.notifications.load(Ordering::Relaxed), 1); +} + +#[tokio::test] +async fn delegate_task_panic_becomes_tool_error_without_killing_connection() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + harness + .open(session.clone(), Arc::new(PanickingDelegate)) + .await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id: DelegateRequestId::new(/*value*/ 7), + session_id: session.id.clone(), + request: DelegateRequest::InvokeTool { + invocation: WireNestedToolCall { + cell_id: CellId::new("1".to_string()).into(), + runtime_tool_call_id: "tool-1".to_string(), + tool_name: ToolName::plain("panic").into(), + tool_kind: codex_code_mode_protocol::CodeModeToolKind::Function.into(), + input: None, + }, + }, + })) + .await + .expect("delegate request"); + tokio::time::timeout(Duration::from_secs(1), harness.outgoing_rx.recv()) + .await + .expect("delegate response timeout") + .expect("delegate response frame"); + + assert!(harness.alive.load(Ordering::Acquire)); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::CellClosed { + session_id: session.id, + cell_id: CellId::new("1".to_string()).into(), + })) + .await + .expect("cell close"); +} + +#[tokio::test] +async fn delegate_for_unknown_cell_fails_connection_without_invocation() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id: DelegateRequestId::new(/*value*/ 7), + session_id: session.id, + request: DelegateRequest::InvokeTool { + invocation: WireNestedToolCall { + cell_id: CellId::new("missing".to_string()).into(), + runtime_tool_call_id: "tool-1".to_string(), + tool_name: ToolName::plain("slow").into(), + tool_kind: codex_code_mode_protocol::CodeModeToolKind::Function.into(), + input: None, + }, + }, + })) + .await + .expect("delegate request"); + tokio::task::yield_now().await; + + assert!(!harness.alive.load(Ordering::Acquire)); + assert_eq!(delegate.invocations.load(Ordering::Relaxed), 0); +} + +#[tokio::test] +async fn delegate_after_cell_close_fails_connection_without_invocation() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::CellClosed { + session_id: session.id.clone(), + cell_id: CellId::new("1".to_string()).into(), + })) + .await + .expect("cell close"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::DelegateRequest { + id: DelegateRequestId::new(/*value*/ 7), + session_id: session.id, + request: DelegateRequest::Notify { + call_id: "notify-1".to_string(), + cell_id: CellId::new("1".to_string()).into(), + text: "late".to_string(), + }, + })) + .await + .expect("delegate request"); + tokio::task::yield_now().await; + + assert!(!harness.alive.load(Ordering::Acquire)); + assert_eq!(delegate.invocations.load(Ordering::Relaxed), 0); +} + +#[tokio::test] +async fn mismatched_initial_response_fails_connection_and_closes_cell_once() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let started = harness.start_cell(session, /*request_id*/ 2, "1").await; + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::InitialResponse { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: WireRuntimeResponse::Yielded { + cell_id: CellId::new("2".to_string()).into(), + content_items: Vec::new(), + }, + }, + })) + .await + .expect("initial response"); + + assert!(started.initial_response().await.is_err()); + assert!(!harness.alive.load(Ordering::Acquire)); + assert_eq!( + *delegate.closed_cells.lock().expect("closed cells lock"), + vec![CellId::new("1".to_string())] + ); +} + +#[tokio::test] +async fn mismatched_wait_response_fails_connection() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let (response_tx, response_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Wait { + session, + request: WaitRequest { + cell_id: CellId::new("1".to_string()), + yield_time_ms: 1, + }, + caller_cancellation: CancellationToken::new(), + response_tx, + }) + .await + .expect("wait command"); + harness.outgoing_rx.recv().await.expect("wait frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Ok { + value: HostResponse::WaitCompleted { + outcome: WireWaitOutcome::LiveCell(WireRuntimeResponse::Yielded { + cell_id: CellId::new("2".to_string()).into(), + content_items: Vec::new(), + }), + }, + }, + })) + .await + .expect("wait response"); + + assert!(response_rx.await.expect("wait reply").is_err()); + assert!(!harness.alive.load(Ordering::Acquire)); + assert_eq!( + *delegate.closed_cells.lock().expect("closed cells lock"), + vec![CellId::new("1".to_string())] + ); +} + +#[tokio::test] +async fn mismatched_terminate_response_fails_connection() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let (response_tx, response_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Terminate { + session, + cell_id: CellId::new("1".to_string()), + response_tx, + }) + .await + .expect("terminate command"); + harness.outgoing_rx.recv().await.expect("terminate frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Ok { + value: HostResponse::WaitCompleted { + outcome: WireWaitOutcome::MissingCell(WireRuntimeResponse::Terminated { + cell_id: CellId::new("2".to_string()).into(), + content_items: Vec::new(), + }), + }, + }, + })) + .await + .expect("terminate response"); + + assert!(response_rx.await.expect("terminate reply").is_err()); + assert!(!harness.alive.load(Ordering::Acquire)); + assert_eq!( + *delegate.closed_cells.lock().expect("closed cells lock"), + vec![CellId::new("1".to_string())] + ); +} + +#[tokio::test] +async fn remote_wait_accepts_durations_longer_than_five_minutes() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + harness + .open(session.clone(), Arc::new(RecordingDelegate::default())) + .await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let (response_tx, response_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Wait { + session, + request: WaitRequest { + cell_id: CellId::new("1".to_string()), + yield_time_ms: 300_001, + }, + caller_cancellation: CancellationToken::new(), + response_tx, + }) + .await + .expect("wait command"); + tokio::time::timeout(Duration::from_secs(1), harness.outgoing_rx.recv()) + .await + .expect("wait frame timeout") + .expect("wait frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Ok { + value: HostResponse::WaitCompleted { + outcome: WireWaitOutcome::LiveCell(WireRuntimeResponse::Yielded { + cell_id: CellId::new("1".to_string()).into(), + content_items: Vec::new(), + }), + }, + }, + })) + .await + .expect("wait response"); + + assert_eq!( + response_rx.await.expect("wait reply"), + Ok(codex_code_mode_protocol::WaitOutcome::LiveCell( + codex_code_mode_protocol::RuntimeResponse::Yielded { + cell_id: CellId::new("1".to_string()), + content_items: Vec::new(), + } + )) + ); +} + +#[tokio::test] +async fn cancelled_wait_is_retired_before_next_wait_is_sent() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + harness + .open(session.clone(), Arc::new(RecordingDelegate::default())) + .await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let first_cancellation = CancellationToken::new(); + let (first_tx, first_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Wait { + session: session.clone(), + request: WaitRequest { + cell_id: CellId::new("1".to_string()), + yield_time_ms: 60_000, + }, + caller_cancellation: first_cancellation.clone(), + response_tx: first_tx, + }) + .await + .expect("first wait command"); + harness.outgoing_rx.recv().await.expect("first wait frame"); + first_cancellation.cancel(); + drop(first_rx); + + let (second_tx, second_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Wait { + session, + request: WaitRequest { + cell_id: CellId::new("1".to_string()), + yield_time_ms: 1, + }, + caller_cancellation: CancellationToken::new(), + response_tx: second_tx, + }) + .await + .expect("second wait command"); + harness + .outgoing_rx + .recv() + .await + .expect("cancel request frame"); + assert!(matches!( + harness.outgoing_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Err { + message: "code-mode request cancelled".to_string(), + }, + })) + .await + .expect("cancelled wait response"); + harness.outgoing_rx.recv().await.expect("second wait frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 4), + result: WireResult::Ok { + value: HostResponse::WaitCompleted { + outcome: WireWaitOutcome::LiveCell(WireRuntimeResponse::Yielded { + cell_id: CellId::new("1".to_string()).into(), + content_items: Vec::new(), + }), + }, + }, + })) + .await + .expect("second wait response"); + + assert_eq!( + second_rx.await.expect("second wait reply"), + Ok(codex_code_mode_protocol::WaitOutcome::LiveCell( + codex_code_mode_protocol::RuntimeResponse::Yielded { + cell_id: CellId::new("1".to_string()), + content_items: Vec::new(), + } + )) + ); +} + +#[tokio::test] +async fn abandoned_execute_is_tracked_and_terminated_after_admission() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let cancellation = CancellationToken::new(); + let (execute_tx, execute_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Execute { + session: session.clone(), + request: ExecuteRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: Vec::new(), + source: "await new Promise(() => {})".to_string(), + yield_time_ms: Some(1), + max_output_tokens: None, + }, + caller_cancellation: cancellation.clone(), + response_tx: execute_tx, + }) + .await + .expect("execute command"); + harness.outgoing_rx.recv().await.expect("execute frame"); + cancellation.cancel(); + drop(execute_rx); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: HostResponse::ExecutionStarted { + cell_id: CellId::new("1".to_string()).into(), + }, + }, + })) + .await + .expect("execute response"); + + harness + .outgoing_rx + .recv() + .await + .expect("execute cancellation frame"); + harness + .outgoing_rx + .recv() + .await + .expect("abandoned cell termination frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::InitialResponse { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: WireRuntimeResponse::Terminated { + cell_id: CellId::new("1".to_string()).into(), + content_items: Vec::new(), + }, + }, + })) + .await + .expect("initial response"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Ok { + value: HostResponse::WaitCompleted { + outcome: WireWaitOutcome::LiveCell(WireRuntimeResponse::Terminated { + cell_id: CellId::new("1".to_string()).into(), + content_items: Vec::new(), + }), + }, + }, + })) + .await + .expect("terminate response"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::CellClosed { + session_id: session.id, + cell_id: CellId::new("1".to_string()).into(), + })) + .await + .expect("cell close"); + tokio::task::yield_now().await; + + assert!(harness.alive.load(Ordering::Acquire)); + assert_eq!( + *delegate.closed_cells.lock().expect("closed cells lock"), + vec![CellId::new("1".to_string())] + ); +} + +#[tokio::test] +async fn delivered_but_unclaimed_execute_is_terminated_when_the_caller_is_cancelled() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let cancellation = CancellationToken::new(); + let (execute_tx, execute_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Execute { + session: session.clone(), + request: ExecuteRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: Vec::new(), + source: "await new Promise(() => {})".to_string(), + yield_time_ms: Some(1), + max_output_tokens: None, + }, + caller_cancellation: cancellation.clone(), + response_tx: execute_tx, + }) + .await + .expect("execute command"); + harness.outgoing_rx.recv().await.expect("execute frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: HostResponse::ExecutionStarted { + cell_id: CellId::new("1".to_string()).into(), + }, + }, + })) + .await + .expect("execute response"); + let delivered = execute_rx + .await + .expect("execute reply") + .expect("delivered execute"); + assert_eq!(delivered.request_id, RequestId::new(/*value*/ 2)); + cancellation.cancel(); + + harness + .outgoing_rx + .recv() + .await + .expect("execute cancellation frame"); + harness + .outgoing_rx + .recv() + .await + .expect("unclaimed cell termination frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::InitialResponse { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: WireRuntimeResponse::Terminated { + cell_id: CellId::new("1".to_string()).into(), + content_items: Vec::new(), + }, + }, + })) + .await + .expect("initial response"); + assert!(delivered.started.initial_response().await.is_ok()); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 3), + result: WireResult::Ok { + value: HostResponse::WaitCompleted { + outcome: WireWaitOutcome::LiveCell(WireRuntimeResponse::Terminated { + cell_id: CellId::new("1".to_string()).into(), + content_items: Vec::new(), + }), + }, + }, + })) + .await + .expect("terminate response"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::CellClosed { + session_id: session.id, + cell_id: CellId::new("1".to_string()).into(), + })) + .await + .expect("cell close"); + tokio::task::yield_now().await; + + assert!(harness.alive.load(Ordering::Acquire)); + assert_eq!( + *delegate.closed_cells.lock().expect("closed cells lock"), + vec![CellId::new("1".to_string())] + ); +} + +#[tokio::test] +async fn session_accepts_more_than_4096_cells_without_growing_a_tombstone_set() { + const CELL_COUNT: usize = 4097; + + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + + for sequence in 1..=CELL_COUNT { + let request_id = i64::try_from(sequence).expect("cell sequence fits in i64") + 1; + let cell_id = sequence.to_string(); + let started = harness + .start_cell(session.clone(), request_id, &cell_id) + .await; + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::InitialResponse { + id: RequestId::new(request_id), + result: WireResult::Ok { + value: WireRuntimeResponse::Yielded { + cell_id: CellId::new(cell_id.clone()).into(), + content_items: Vec::new(), + }, + }, + })) + .await + .expect("initial response"); + assert!(started.initial_response().await.is_ok()); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::CellClosed { + session_id: session.id.clone(), + cell_id: CellId::new(cell_id).into(), + })) + .await + .expect("cell close"); + } + + tokio::time::timeout(Duration::from_secs(1), async { + while delegate + .closed_cells + .lock() + .expect("closed cells lock") + .len() + != CELL_COUNT + { + tokio::task::yield_now().await; + } + }) + .await + .expect("cell close callbacks timeout"); + assert!(harness.alive.load(Ordering::Acquire)); +} + +#[tokio::test] +async fn connection_failure_closes_every_live_cell_once() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + let cleanup = harness.open(session.clone(), delegate.clone()).await; + let (execute_tx, execute_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Execute { + session, + request: ExecuteRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: Vec::new(), + source: "await new Promise(() => {})".to_string(), + yield_time_ms: Some(1), + max_output_tokens: None, + }, + caller_cancellation: CancellationToken::new(), + response_tx: execute_tx, + }) + .await + .expect("execute command"); + harness.outgoing_rx.recv().await.expect("execute frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: HostResponse::ExecutionStarted { + cell_id: CellId::new("1".to_string()).into(), + }, + }, + })) + .await + .expect("execute response"); + let _started = execute_rx + .await + .expect("execute reply") + .expect("execute session"); + harness + .event_tx + .send(DriverEvent::Failed("host crashed".to_string())) + .await + .expect("failure event"); + tokio::time::timeout(Duration::from_secs(1), cleanup.wait()) + .await + .expect("session cleanup timeout"); + assert_eq!( + *delegate.closed_cells.lock().expect("closed cells lock"), + vec![CellId::new("1".to_string())] + ); +} + +#[tokio::test] +async fn session_cleanup_does_not_wait_for_delegate_completion() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let (delegate, mut events_rx, release) = HeldDelegate::new(); + let cleanup = harness.open(session.clone(), delegate).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + harness + .start_tool_delegate(&session, DelegateRequestId::new(/*value*/ 7)) + .await; + assert_eq!( + next_held_delegate_event(&mut events_rx).await, + HeldDelegateEvent::Started + ); + + harness + .event_tx + .send(DriverEvent::Failed("host crashed".to_string())) + .await + .expect("failure event"); + let closure_events = [ + next_held_delegate_event(&mut events_rx).await, + next_held_delegate_event(&mut events_rx).await, + ]; + assert!(closure_events.contains(&HeldDelegateEvent::Cancelled)); + assert!(closure_events.contains(&HeldDelegateEvent::CellClosed(CellId::new("1".to_string())))); + tokio::time::timeout(Duration::from_secs(1), cleanup.wait()) + .await + .expect("session cleanup timeout"); + assert!(matches!( + events_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + + release.cancel(); + assert_eq!( + next_held_delegate_event(&mut events_rx).await, + HeldDelegateEvent::Finished + ); + assert!(matches!( + events_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) + )); +} + +#[tokio::test] +async fn aborting_driver_marks_connection_dead_and_closes_cells() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let delegate = Arc::new(RecordingDelegate::default()); + harness.open(session.clone(), delegate.clone()).await; + let _started = harness + .start_cell(session.clone(), /*request_id*/ 2, "1") + .await; + let (wait_tx, wait_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Wait { + session, + request: WaitRequest { + cell_id: CellId::new("1".to_string()), + yield_time_ms: 60_000, + }, + caller_cancellation: CancellationToken::new(), + response_tx: wait_tx, + }) + .await + .expect("wait command"); + harness.outgoing_rx.recv().await.expect("wait frame"); + + harness.driver_task.abort(); + for _ in 0..10 { + if !harness.alive.load(Ordering::Acquire) { + break; + } + tokio::task::yield_now().await; + } + + assert!(!harness.alive.load(Ordering::Acquire)); + assert!(harness.cancellation.is_cancelled()); + assert!(wait_rx.await.expect("wait failure").is_err()); + assert_eq!( + *delegate.closed_cells.lock().expect("closed cells lock"), + vec![CellId::new("1".to_string())] + ); +} + +#[tokio::test] +async fn dropped_shutdown_waiter_does_not_abort_remote_cleanup() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + harness + .open(session.clone(), Arc::new(RecordingDelegate::default())) + .await; + let (shutdown_tx, shutdown_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::ShutdownSession { + session: session.clone(), + response_tx: shutdown_tx, + }) + .await + .expect("shutdown command"); + drop(shutdown_rx); + harness.outgoing_rx.recv().await.expect("shutdown frame"); + harness + .event_tx + .send(DriverEvent::HostMessage(HostToClient::Response { + id: RequestId::new(/*value*/ 2), + result: WireResult::Ok { + value: HostResponse::SessionClosed { + session_id: session.id.clone(), + }, + }, + })) + .await + .expect("shutdown response"); + + let (execute_tx, execute_rx) = oneshot::channel(); + harness + .command_tx + .send(DriverCommand::Execute { + session, + request: ExecuteRequest { + tool_call_id: "call-2".to_string(), + enabled_tools: Vec::new(), + source: "text('unreachable')".to_string(), + yield_time_ms: None, + max_output_tokens: None, + }, + caller_cancellation: CancellationToken::new(), + response_tx: execute_tx, + }) + .await + .expect("execute command"); + assert_eq!( + execute_rx + .await + .expect("execute reply") + .err() + .expect("closed session should reject execute"), + "unknown code-mode session session-1" + ); + assert!(matches!( + harness.outgoing_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); +} diff --git a/codex-rs/code-mode/src/remote_session/connection/reader.rs b/codex-rs/code-mode/src/remote_session/connection/reader.rs new file mode 100644 index 000000000..fbfc94de4 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session/connection/reader.rs @@ -0,0 +1,29 @@ +use codex_code_mode_protocol::host::FramedReader; +use codex_code_mode_protocol::host::HostToClient; +use tokio::process::ChildStdout; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +use super::driver::DriverEvent; + +pub(super) async fn drive_reader( + mut reader: FramedReader, + events: mpsc::Sender, + cancellation: CancellationToken, +) -> Result<(), String> { + loop { + let message = tokio::select! { + _ = cancellation.cancelled() => return Ok(()), + result = reader.read::() => result, + }; + let message = match message { + Ok(Some(message)) => message, + Ok(None) => return Err("code-mode host closed its stdout".to_string()), + Err(err) => return Err(format!("failed to read code-mode host message: {err}")), + }; + events + .send(DriverEvent::HostMessage(message)) + .await + .map_err(|_| "code-mode connection driver closed".to_string())?; + } +} diff --git a/codex-rs/code-mode/src/remote_session_tests.rs b/codex-rs/code-mode/src/remote_session_tests.rs new file mode 100644 index 000000000..d97120755 --- /dev/null +++ b/codex-rs/code-mode/src/remote_session_tests.rs @@ -0,0 +1,52 @@ +use std::sync::Arc; + +use codex_code_mode_protocol::CodeModeSessionProvider; + +use super::ProcessOwnedCodeModeSession; +use super::ProcessOwnedCodeModeSessionProvider; +use crate::NoopCodeModeSessionDelegate; + +#[test] +fn provider_reuses_its_live_process_host() { + let provider = ProcessOwnedCodeModeSessionProvider::default(); + + let first = provider.process_host(); + let second = provider.process_host(); + + assert!(Arc::ptr_eq(&first, &second)); +} + +#[tokio::test] +async fn provider_reports_host_spawn_failure() { + let provider = ProcessOwnedCodeModeSessionProvider::with_host_program( + "codex-code-mode-host-does-not-exist".into(), + ); + + let error = provider + .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .await + .err() + .expect("session creation should fail"); + + assert!(error.contains("failed to spawn code-mode host")); +} + +#[tokio::test] +async fn shutdown_before_open_does_not_spawn_the_host() { + let session = ProcessOwnedCodeModeSession::new(); + + session.shutdown().await.expect("shutdown session"); + let error = session + .execute(codex_code_mode_protocol::ExecuteRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: Vec::new(), + source: "text('unreachable')".to_string(), + yield_time_ms: None, + max_output_tokens: None, + }) + .await + .err() + .expect("shutdown session should reject execution"); + + assert_eq!(error, "code mode session is shutting down"); +}