diff --git a/codex-rs/code-mode-host/src/host_tests.rs b/codex-rs/code-mode-host/src/host_tests.rs index 3e78ae6911..07a2d4d717 100644 --- a/codex-rs/code-mode-host/src/host_tests.rs +++ b/codex-rs/code-mode-host/src/host_tests.rs @@ -8,6 +8,7 @@ use std::task::Context; use std::task::Poll; use std::time::Duration; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::host::Capability; use codex_code_mode_protocol::host::CapabilitySet; use codex_code_mode_protocol::host::ClientHello; @@ -23,6 +24,7 @@ use codex_code_mode_protocol::host::HostResponse; 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::SESSION_RESOURCE_LIMITS_CAPABILITY; use codex_code_mode_protocol::host::SessionId; use codex_code_mode_protocol::host::SupportedProtocolVersions; use codex_code_mode_protocol::host::WireExecuteRequest; @@ -129,6 +131,7 @@ async fn handshake_and_multiple_session_lifecycles_are_ordered() { id: request_id, request: HostRequest::OpenSession { session_id: session_id(id), + cell_execution_limits: None, }, }) .await @@ -199,6 +202,7 @@ async fn disconnect_cancels_a_backpressured_host_writer() { id: request_id(/*value*/ 1), request: HostRequest::OpenSession { session_id: session_id("backpressured-session"), + cell_execution_limits: None, }, }) .await @@ -357,6 +361,56 @@ async fn optional_dual_websocket_capability_falls_back_to_a_single_connection() host.await.expect("host task").expect("host connection"); } +#[tokio::test] +async fn session_resource_limits_are_negotiated_when_optional_or_required() { + let dual_capability = + Capability::new(DUAL_WEBSOCKET_CAPABILITY).expect("dual websocket capability"); + let resource_limits_capability = + Capability::new(SESSION_RESOURCE_LIMITS_CAPABILITY).expect("resource limits capability"); + let resource_limits = + CapabilitySet::try_new([resource_limits_capability.clone()]).expect("host capabilities"); + + for (required, optional) in [ + ( + CapabilitySet::empty(), + CapabilitySet::try_new([dual_capability, resource_limits_capability]) + .expect("optional capabilities"), + ), + (resource_limits.clone(), CapabilitySet::empty()), + ] { + let (host_stream, client_stream) = tokio::io::duplex(/*max_buf_size*/ 1024); + let (host_reader, host_writer) = tokio::io::split(host_stream); + let (client_reader, client_writer) = tokio::io::split(client_stream); + let host = tokio::spawn(run(host_reader, host_writer)); + let mut reader = FramedReader::new(client_reader); + let mut writer = FramedWriter::new(client_writer); + + writer + .write(&ClientToHost::ClientHello( + ClientHello::new( + SupportedProtocolVersions::try_new([ProtocolVersion::V1]) + .expect("supported versions"), + required, + optional, + ) + .expect("client hello"), + )) + .await + .expect("write hello"); + assert_eq!( + reader.read::().await.expect("host hello"), + Some(HostToClient::HostHello(HostHello::new( + ProtocolVersion::V1, + resource_limits.clone(), + ))) + ); + + drop(writer); + drop(reader); + host.await.expect("host task").expect("host connection"); + } +} + #[tokio::test] async fn incompatible_or_invalid_handshake_is_rejected() { let (host_stream, client_stream) = tokio::io::duplex(/*max_buf_size*/ 1024); @@ -393,6 +447,7 @@ async fn incompatible_or_invalid_handshake_is_rejected() { id: request_id(/*value*/ 1), request: HostRequest::OpenSession { session_id: session_id("session-1"), + cell_execution_limits: None, }, }) .await @@ -458,6 +513,7 @@ async fn session_id_cannot_be_reused_after_shutdown() { request_id(/*value*/ 1), HostRequest::OpenSession { session_id: id.clone(), + cell_execution_limits: None, }, ), ( @@ -483,7 +539,10 @@ async fn session_id_cannot_be_reused_after_shutdown() { writer .write(&ClientToHost::Request { id: request_id(/*value*/ 3), - request: HostRequest::OpenSession { session_id: id }, + request: HostRequest::OpenSession { + session_id: id, + cell_execution_limits: None, + }, }) .await .expect("reuse session ID"); @@ -568,7 +627,10 @@ async fn execute_request_id_remains_active_until_initial_response() { }); let session_id = session_id("session-1"); state - .open_session(session_id.clone()) + .open_session( + session_id.clone(), + CodeModeSessionCellExecutionLimits::default(), + ) .expect("open session"); let request_id = request_id(/*value*/ 1); @@ -627,7 +689,10 @@ async fn active_cell_limit_rejects_execute_without_disconnecting() { }; let session_id = session_id("session-1"); state - .open_session(session_id.clone()) + .open_session( + session_id.clone(), + CodeModeSessionCellExecutionLimits::default(), + ) .expect("open session"); let request_id = request_id(/*value*/ 1); diff --git a/codex-rs/code-mode-host/src/lib.rs b/codex-rs/code-mode-host/src/lib.rs index 88d2e161f5..dd4cf24bb6 100644 --- a/codex-rs/code-mode-host/src/lib.rs +++ b/codex-rs/code-mode-host/src/lib.rs @@ -10,6 +10,7 @@ use std::time::Duration; use anyhow::Context; use anyhow::Result; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::host::Capability; use codex_code_mode_protocol::host::CapabilitySet; use codex_code_mode_protocol::host::ClientToHost; @@ -23,6 +24,7 @@ use codex_code_mode_protocol::host::HostToClient; use codex_code_mode_protocol::host::MAX_PENDING_DELEGATE_CALLS; use codex_code_mode_protocol::host::ProtocolVersion; use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::SESSION_RESOURCE_LIMITS_CAPABILITY; use codex_code_mode_protocol::host::SessionId; use codex_code_mode_protocol::host::SupportedProtocolVersions; use codex_code_mode_protocol::host::TransportLane; @@ -314,11 +316,21 @@ async fn negotiate( } else { None }; - let host_capabilities = if registration.is_some() { - CapabilitySet::try_new([dual_capability])? - } else { - CapabilitySet::empty() - }; + let resource_limits_capability = Capability::new(SESSION_RESOURCE_LIMITS_CAPABILITY)?; + let resource_limits_requested = client_hello + .required_capabilities() + .contains(&resource_limits_capability) + || client_hello + .optional_capabilities() + .contains(&resource_limits_capability); + let host_capabilities = CapabilitySet::try_new( + [ + registration.is_some().then_some(dual_capability), + resource_limits_requested.then_some(resource_limits_capability), + ] + .into_iter() + .flatten(), + )?; if let Some(capability) = client_hello .required_capabilities() .iter() @@ -419,10 +431,16 @@ impl HostState { return; } match request { - HostRequest::OpenSession { session_id } => { - let result = self - .open_session(session_id.clone()) - .map(|()| HostResponse::SessionReady { session_id }); + HostRequest::OpenSession { + session_id, + cell_execution_limits, + } => { + let result = CodeModeSessionCellExecutionLimits::try_from( + cell_execution_limits.unwrap_or_default(), + ) + .map_err(|error| format!("invalid code-mode session execution limits: {error}")) + .and_then(|limits| self.open_session(session_id.clone(), limits)) + .map(|()| HostResponse::SessionReady { session_id }); self.respond(request_id, result); } HostRequest::Execute { @@ -537,7 +555,11 @@ impl HostState { } } - fn open_session(&self, session_id: SessionId) -> Result<(), String> { + fn open_session( + &self, + session_id: SessionId, + cell_execution_limits: CodeModeSessionCellExecutionLimits, + ) -> Result<(), String> { let mut sessions = self.sessions.lock().unwrap_or_else(PoisonError::into_inner); if sessions.contains_key(&session_id) { return Err(format!( @@ -571,6 +593,7 @@ impl HostState { InProcessCodeModeSession::with_delegate_and_task_failure_handler( delegate, task_failure_handler, + cell_execution_limits, ), ), ); diff --git a/codex-rs/code-mode-host/tests/stdio.rs b/codex-rs/code-mode-host/tests/stdio.rs index 8edcbcaa8d..4ae5bfe1e2 100644 --- a/codex-rs/code-mode-host/tests/stdio.rs +++ b/codex-rs/code-mode-host/tests/stdio.rs @@ -12,6 +12,7 @@ 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::CodeModeSessionCellExecutionLimits; use codex_code_mode::CodeModeSessionDelegate; use codex_code_mode::CodeModeSessionProvider; use codex_code_mode::CodeModeToolKind; @@ -260,6 +261,74 @@ async fn next_callback_event( .expect("callback event stream closed") } +#[tokio::test] +async fn session_execution_limits_are_isolated_on_a_shared_process_host() { + let provider: Arc = + Arc::new(ProcessOwnedCodeModeSessionProvider::with_host_program( + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("host binary"), + )); + let limited = provider + .create_session_with_limits( + Arc::new(RecordingDelegate::default()), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(1), + max_heap_size_bytes: None, + }, + ) + .await + .expect("create limited session"); + let other = provider + .create_session_with_limits( + Arc::new(RecordingDelegate::default()), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(1_000), + max_heap_size_bytes: None, + }, + ) + .await + .expect("create independently limited session"); + + let response = tokio::time::timeout( + Duration::from_secs(5), + execute(&limited, execute_request("await new Promise(() => {});")), + ) + .await + .expect("session limit should bound the default execution wait"); + assert_eq!( + response, + RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + } + ); + assert_eq!( + tokio::time::timeout( + Duration::from_secs(5), + limited.wait(WaitRequest { + cell_id: cell_id("1"), + yield_time_ms: 60_000, + }), + ) + .await + .expect("session limit should bound explicit waits") + .expect("wait for yielded cell"), + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + }) + ); + + limited + .terminate(cell_id("1")) + .await + .expect("terminate cell"); + limited.shutdown().await.expect("shutdown limited session"); + other + .shutdown() + .await + .expect("shutdown independent session"); +} + #[tokio::test] async fn remote_session_persists_values_forwards_delegates_and_controls_cells() { let provider = ProcessOwnedCodeModeSessionProvider::with_host_program( diff --git a/codex-rs/code-mode-host/tests/websocket.rs b/codex-rs/code-mode-host/tests/websocket.rs index 030f83432d..9ff50e1b84 100644 --- a/codex-rs/code-mode-host/tests/websocket.rs +++ b/codex-rs/code-mode-host/tests/websocket.rs @@ -6,6 +6,7 @@ use anyhow::Context; use anyhow::Result; use codex_code_mode::CellId; use codex_code_mode::CodeModeNestedToolCall; +use codex_code_mode::CodeModeSessionCellExecutionLimits; use codex_code_mode::CodeModeSessionDelegate; use codex_code_mode::CodeModeSessionProvider; use codex_code_mode::CodeModeToolKind; @@ -32,6 +33,7 @@ use codex_code_mode_protocol::host::HostToClient; use codex_code_mode_protocol::host::MAX_FRAME_BYTES; use codex_code_mode_protocol::host::ProtocolVersion; use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::SESSION_RESOURCE_LIMITS_CAPABILITY; use codex_code_mode_protocol::host::SessionId; use codex_code_mode_protocol::host::SupportedProtocolVersions; use codex_code_mode_protocol::host::WireContentItem; @@ -225,16 +227,18 @@ impl HostClient { async fn negotiate_dual(&mut self, websocket_url: &str) -> Result { let capability = Capability::new(DUAL_WEBSOCKET_CAPABILITY)?; + let resource_limits = Capability::new(SESSION_RESOURCE_LIMITS_CAPABILITY)?; let hello = ClientHello::new( SupportedProtocolVersions::try_new([ProtocolVersion::V1])?, CapabilitySet::empty(), - CapabilitySet::try_new([capability.clone()])?, + CapabilitySet::try_new([capability.clone(), resource_limits.clone()])?, )?; self.send(&ClientToHost::ClientHello(hello)).await?; let HostToClient::HostHello(hello) = self.read().await? else { anyhow::bail!("expected code-mode host hello"); }; assert!(hello.capabilities().contains(&capability)); + assert!(hello.capabilities().contains(&resource_limits)); let token = hello .bulk_connection_token() .context("dual websocket handshake omitted its pairing token")?; @@ -257,6 +261,7 @@ impl HostClient { id, request: HostRequest::OpenSession { session_id: session_id.clone(), + cell_execution_limits: None, }, }) .await?; @@ -416,7 +421,13 @@ async fn production_websocket_client_runs_nested_tools_while_other_sessions_prog .await .map_err(anyhow::Error::msg)?; let fast_session = provider - .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .create_session_with_limits( + Arc::new(NoopCodeModeSessionDelegate), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(5_000), + max_heap_size_bytes: None, + }, + ) .await .map_err(anyhow::Error::msg)?; diff --git a/codex-rs/code-mode-protocol/src/host/host_tests.rs b/codex-rs/code-mode-protocol/src/host/host_tests.rs index 240028a570..7404cb6d4d 100644 --- a/codex-rs/code-mode-protocol/src/host/host_tests.rs +++ b/codex-rs/code-mode-protocol/src/host/host_tests.rs @@ -30,11 +30,13 @@ use super::WireImageDetail; use super::WireNestedToolCall; use super::WireResult; use super::WireRuntimeResponse; +use super::WireSessionCellExecutionLimits; use super::WireToolDefinition; use super::WireToolKind; use super::WireToolName; use super::WireWaitOutcome; use super::WireWaitRequest; +use crate::CodeModeSessionCellExecutionLimits; use crate::ExecuteRequest; fn session_id() -> SessionId { @@ -363,6 +365,61 @@ fn handshake_v1_variants_are_pinned() { } } +#[test] +fn open_session_serializes_optional_cell_execution_limits() { + assert_wire_round_trip( + HostRequest::OpenSession { + session_id: session_id(), + cell_execution_limits: Some(WireSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: Some(16 * 1024 * 1024), + }), + }, + json!({ + "method": "session/open", + "sessionId": "session-1", + "cellExecutionLimits": { + "maxYieldTimeMs": 250, + "maxHeapSizeBytes": 16 * 1024 * 1024, + }, + }), + ); +} + +#[test] +fn session_cell_execution_limits_convert_between_domain_and_wire() { + let domain_limits = CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: Some(16_usize * 1024 * 1024), + }; + let wire_limits = WireSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: Some(16_u64 * 1024 * 1024), + }; + + assert_eq!( + WireSessionCellExecutionLimits::try_from(domain_limits.clone()) + .expect("domain limits convert to wire limits"), + wire_limits + ); + assert_eq!( + CodeModeSessionCellExecutionLimits::try_from(wire_limits) + .expect("wire limits convert to domain limits"), + domain_limits + ); +} + +#[cfg(target_pointer_width = "32")] +#[test] +fn session_cell_execution_limits_reject_heap_sizes_that_exceed_usize() { + let wire_limits = WireSessionCellExecutionLimits { + max_yield_time_ms: None, + max_heap_size_bytes: Some(u64::from(u32::MAX) + 1), + }; + + assert!(CodeModeSessionCellExecutionLimits::try_from(wire_limits).is_err()); +} + #[test] fn client_to_host_v1_variants_are_pinned() { let execute_request = execute_request(); @@ -371,6 +428,7 @@ fn client_to_host_v1_variants_are_pinned() { request_id(/*value*/ 1), HostRequest::OpenSession { session_id: session_id(), + cell_execution_limits: None, }, json!({ "method": "session/open", "sessionId": "session-1" }), ), @@ -810,6 +868,17 @@ fn every_nested_v1_object_rejects_unknown_fields() { })) .is_err() ); + assert!( + serde_json::from_value::(json!({ + "method": "session/open", + "sessionId": "session-1", + "cellExecutionLimits": { + "maxYieldTimeMs": 250, + "unexpected": true, + }, + })) + .is_err() + ); assert!( serde_json::from_value::(json!({ "tool_call_id": "call-1", diff --git a/codex-rs/code-mode-protocol/src/host/message.rs b/codex-rs/code-mode-protocol/src/host/message.rs index e6918ea4fd..4b922f5ce2 100644 --- a/codex-rs/code-mode-protocol/src/host/message.rs +++ b/codex-rs/code-mode-protocol/src/host/message.rs @@ -17,6 +17,7 @@ use super::WireCellId; use super::WireExecuteRequest; use super::WireNestedToolCall; use super::WireRuntimeResponse; +use super::WireSessionCellExecutionLimits; use super::WireWaitOutcome; use super::WireWaitRequest; @@ -242,7 +243,11 @@ impl HostToClient { #[serde(deny_unknown_fields, tag = "method", rename_all_fields = "camelCase")] pub enum HostRequest { #[serde(rename = "session/open")] - OpenSession { session_id: SessionId }, + OpenSession { + session_id: SessionId, + #[serde(default, skip_serializing_if = "Option::is_none")] + cell_execution_limits: Option, + }, #[serde(rename = "session/execute")] Execute { session_id: SessionId, diff --git a/codex-rs/code-mode-protocol/src/host/mod.rs b/codex-rs/code-mode-protocol/src/host/mod.rs index 6169c61a17..4059230213 100644 --- a/codex-rs/code-mode-protocol/src/host/mod.rs +++ b/codex-rs/code-mode-protocol/src/host/mod.rs @@ -16,6 +16,9 @@ pub const MAX_PENDING_DELEGATE_CALLS: usize = 1_024; /// Optional second WebSocket carrying delegate callbacks and their responses. pub const DUAL_WEBSOCKET_CAPABILITY: &str = "dual-websocket-v1"; +/// Negotiated support for cell execution resource limits on `session/open`. +pub const SESSION_RESOURCE_LIMITS_CAPABILITY: &str = "session-cell-execution-resource-limits"; + /// Selects one socket of a negotiated dual-WebSocket connection. #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub enum TransportLane { @@ -44,6 +47,7 @@ pub use payload::WireExecuteRequest; pub use payload::WireImageDetail; pub use payload::WireNestedToolCall; pub use payload::WireRuntimeResponse; +pub use payload::WireSessionCellExecutionLimits; pub use payload::WireToolDefinition; pub use payload::WireToolKind; pub use payload::WireToolName; diff --git a/codex-rs/code-mode-protocol/src/host/payload.rs b/codex-rs/code-mode-protocol/src/host/payload.rs index 7cb15db379..cee3501e3a 100644 --- a/codex-rs/code-mode-protocol/src/host/payload.rs +++ b/codex-rs/code-mode-protocol/src/host/payload.rs @@ -7,6 +7,7 @@ use serde_json::Value as JsonValue; use crate::CellId; use crate::CodeModeNestedToolCall; +use crate::CodeModeSessionCellExecutionLimits; use crate::CodeModeToolKind; use crate::ExecuteRequest; use crate::FunctionCallOutputContentItem; @@ -16,6 +17,38 @@ use crate::ToolDefinition; use crate::WaitOutcome; use crate::WaitRequest; +/// The per-cell execution limits carried by a V1 session-open request. +#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)] +#[serde(deny_unknown_fields, rename_all = "camelCase")] +pub struct WireSessionCellExecutionLimits { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_yield_time_ms: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_heap_size_bytes: Option, +} + +impl TryFrom for WireSessionCellExecutionLimits { + type Error = TryFromIntError; + + fn try_from(value: CodeModeSessionCellExecutionLimits) -> Result { + Ok(Self { + max_yield_time_ms: value.max_yield_time_ms, + max_heap_size_bytes: value.max_heap_size_bytes.map(u64::try_from).transpose()?, + }) + } +} + +impl TryFrom for CodeModeSessionCellExecutionLimits { + type Error = TryFromIntError; + + fn try_from(value: WireSessionCellExecutionLimits) -> Result { + Ok(Self { + max_yield_time_ms: value.max_yield_time_ms, + max_heap_size_bytes: value.max_heap_size_bytes.map(usize::try_from).transpose()?, + }) + } +} + /// A cell identifier with a wire representation owned by protocol V1. #[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] #[serde(transparent)] diff --git a/codex-rs/code-mode-protocol/src/lib.rs b/codex-rs/code-mode-protocol/src/lib.rs index 8c84ab6b12..c7133f5228 100644 --- a/codex-rs/code-mode-protocol/src/lib.rs +++ b/codex-rs/code-mode-protocol/src/lib.rs @@ -34,6 +34,7 @@ pub use runtime::WaitToPendingOutcome; pub use runtime::WaitToPendingRequest; pub use session::CellId; pub use session::CodeModeSession; +pub use session::CodeModeSessionCellExecutionLimits; pub use session::CodeModeSessionDelegate; pub use session::CodeModeSessionProvider; pub use session::CodeModeSessionProviderFuture; diff --git a/codex-rs/code-mode-protocol/src/session.rs b/codex-rs/code-mode-protocol/src/session.rs index 01953f0abb..617aed5969 100644 --- a/codex-rs/code-mode-protocol/src/session.rs +++ b/codex-rs/code-mode-protocol/src/session.rs @@ -23,6 +23,13 @@ pub type ToolInvocationFuture<'a> = Pin> + Send + 'a>>; pub type NotificationFuture<'a> = Pin> + Send + 'a>>; +/// Optional resource limits shared by every cell in one code-mode session. +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct CodeModeSessionCellExecutionLimits { + pub max_yield_time_ms: Option, + pub max_heap_size_bytes: Option, +} + #[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] pub struct CellId(String); @@ -164,6 +171,24 @@ pub trait CodeModeSessionProvider: Send + Sync { &'a self, delegate: Arc, ) -> CodeModeSessionProviderFuture<'a>; + + /// Creates a session whose cells share the supplied execution limits. + /// + /// Existing providers remain compatible with unlimited sessions, but must + /// explicitly implement this method before accepting non-default limits. + fn create_session_with_limits<'a>( + &'a self, + delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, + ) -> CodeModeSessionProviderFuture<'a> { + if limits == CodeModeSessionCellExecutionLimits::default() { + self.create_session(delegate) + } else { + Box::pin(async { + Err("code-mode session provider does not support resource limits".to_string()) + }) + } + } } #[cfg(test)] diff --git a/codex-rs/code-mode-runtime/src/service.rs b/codex-rs/code-mode-runtime/src/service.rs index 7ca05eedaf..366f494263 100644 --- a/codex-rs/code-mode-runtime/src/service.rs +++ b/codex-rs/code-mode-runtime/src/service.rs @@ -4,6 +4,7 @@ use std::time::Duration; use codex_code_mode_protocol::CellId; use codex_code_mode_protocol::CodeModeNestedToolCall; use codex_code_mode_protocol::CodeModeSession; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::CodeModeSessionDelegate; use codex_code_mode_protocol::CodeModeSessionResultFuture; use codex_code_mode_protocol::CodeModeToolKind; @@ -29,17 +30,9 @@ use crate::session_runtime::SessionRuntime; const YIELD_GRACE_PERIOD: Duration = Duration::from_secs(1); const MIN_YIELD_TIME_FOR_GRACE: Duration = Duration::from_secs(10); -fn yield_timeout(yield_time_ms: u64) -> Duration { - let yield_time = Duration::from_millis(yield_time_ms); - if yield_time >= MIN_YIELD_TIME_FOR_GRACE { - yield_time.saturating_add(YIELD_GRACE_PERIOD) - } else { - yield_time - } -} - pub struct InProcessCodeModeSession { runtime: SessionRuntime, + cell_execution_limits: CodeModeSessionCellExecutionLimits, } impl InProcessCodeModeSession { @@ -48,20 +41,36 @@ impl InProcessCodeModeSession { } pub fn with_delegate(delegate: Arc) -> Self { + Self::with_delegate_and_limits(delegate, CodeModeSessionCellExecutionLimits::default()) + } + + pub fn with_delegate_and_limits( + delegate: Arc, + cell_execution_limits: CodeModeSessionCellExecutionLimits, + ) -> Self { Self { runtime: SessionRuntime::new(Arc::new(ProtocolDelegate { delegate })), + cell_execution_limits: CodeModeSessionCellExecutionLimits { + max_heap_size_bytes: None, + ..cell_execution_limits + }, } } pub fn with_delegate_and_task_failure_handler( delegate: Arc, task_failure_handler: Arc, + cell_execution_limits: CodeModeSessionCellExecutionLimits, ) -> Self { Self { runtime: SessionRuntime::new_with_task_failure_handler( Arc::new(ProtocolDelegate { delegate }), Some(task_failure_handler), ), + cell_execution_limits: CodeModeSessionCellExecutionLimits { + max_heap_size_bytes: None, + ..cell_execution_limits + }, } } @@ -71,7 +80,7 @@ impl InProcessCodeModeSession { .runtime .execute( runtime_request(request), - runtime::ObserveMode::YieldAfter(yield_timeout(yield_time_ms)), + runtime::ObserveMode::YieldAfter(self.resolve_yield_timeout(yield_time_ms)), ) .await .map_err(|error| error.to_string())?; @@ -126,7 +135,7 @@ impl InProcessCodeModeSession { .runtime .begin_observe( &runtime_cell_id, - runtime::ObserveMode::YieldAfter(yield_timeout(yield_time_ms)), + runtime::ObserveMode::YieldAfter(self.resolve_yield_timeout(yield_time_ms)), ) .await { @@ -185,6 +194,20 @@ impl InProcessCodeModeSession { .await .map_err(|error| error.to_string()) } + + fn resolve_yield_timeout(&self, yield_time_ms: u64) -> Duration { + let yield_time = Duration::from_millis(yield_time_ms); + let timeout = if yield_time >= MIN_YIELD_TIME_FOR_GRACE { + yield_time.saturating_add(YIELD_GRACE_PERIOD) + } else { + yield_time + }; + + self.cell_execution_limits + .max_yield_time_ms + .map(Duration::from_millis) + .map_or(timeout, |limit| timeout.min(limit)) + } } impl Default for InProcessCodeModeSession { diff --git a/codex-rs/code-mode-runtime/src/service_tests.rs b/codex-rs/code-mode-runtime/src/service_tests.rs index 0e5c3677f8..8ec8cbe15f 100644 --- a/codex-rs/code-mode-runtime/src/service_tests.rs +++ b/codex-rs/code-mode-runtime/src/service_tests.rs @@ -12,12 +12,12 @@ use super::WaitOutcome; use super::WaitRequest; use super::WaitToPendingOutcome; use super::WaitToPendingRequest; -use super::yield_timeout; use crate::CodeModeToolKind; use crate::ExecuteRequest; use crate::ExecuteToPendingOutcome; use crate::FunctionCallOutputContentItem; use crate::ToolDefinition; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::NotificationFuture; use codex_code_mode_protocol::ToolInvocationFuture; use codex_protocol::ToolName; @@ -27,15 +27,36 @@ use tokio::sync::Notify; use tokio_util::sync::CancellationToken; #[test] -fn yield_timeout_adds_grace_only_at_ten_seconds() { - assert_eq!( - yield_timeout(/*yield_time_ms*/ 9_999), - Duration::from_millis(9_999) - ); - assert_eq!( - yield_timeout(/*yield_time_ms*/ 10_000), - Duration::from_secs(11) - ); +fn resolve_yield_timeout_applies_grace_before_session_limits() { + for (max_yield_time_ms, requested_yield_time_ms, expected_timeout) in [ + (None, 0, Duration::ZERO), + (None, 9_999, Duration::from_millis(9_999)), + (None, 10_000, Duration::from_secs(11)), + (None, 10_001, Duration::from_millis(11_001)), + (Some(0), 0, Duration::ZERO), + (Some(0), 10_000, Duration::ZERO), + (Some(5_000), 9_999, Duration::from_secs(5)), + (Some(10_000), 10_000, Duration::from_secs(10)), + (Some(10_500), 10_000, Duration::from_millis(10_500)), + (Some(11_000), 10_000, Duration::from_secs(11)), + (Some(12_000), 10_000, Duration::from_secs(11)), + (Some(10_500), 5_000, Duration::from_secs(5)), + (Some(u64::MAX), u64::MAX, Duration::from_millis(u64::MAX)), + ] { + let session = InProcessCodeModeSession::with_delegate_and_limits( + Arc::new(ReleasableToolDelegate::default()), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms, + max_heap_size_bytes: None, + }, + ); + + assert_eq!( + session.resolve_yield_timeout(requested_yield_time_ms), + expected_timeout, + "requested {requested_yield_time_ms} ms with limit {max_yield_time_ms:?}" + ); + } } #[tokio::test(start_paused = true)] @@ -68,6 +89,79 @@ async fn execute_waits_for_nested_tool_during_yield_grace() { ); } +#[tokio::test(start_paused = true)] +async fn execute_and_wait_clamp_yield_grace_without_stopping_the_cell() { + let delegate = Arc::new(ReleasableToolDelegate::default()); + let service = InProcessCodeModeSession::with_delegate_and_limits( + delegate.clone(), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(/*value*/ 10_000), + max_heap_size_bytes: None, + }, + ); + let started = service + .execute(ExecuteRequest { + enabled_tools: vec![echo_tool()], + source: r#"await tools.echo({}); text("done");"#.to_string(), + yield_time_ms: None, + ..execute_request("") + }) + .await + .unwrap(); + let initial_response = tokio::spawn(started.initial_response()); + wait_until_tool_started(&delegate).await; + + tokio::time::advance(Duration::from_millis(9_999)).await; + assert!(!initial_response.is_finished()); + tokio::time::advance(Duration::from_millis(/*millis*/ 1)).await; + wait_until_finished(&initial_response).await; + assert_eq!( + initial_response.await.unwrap().unwrap(), + RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + } + ); + + let wait_response = service + .begin_wait(WaitRequest { + cell_id: cell_id("1"), + yield_time_ms: 10_000, + }) + .await; + let wait_response = tokio::spawn(wait_response); + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(/*secs*/ 10)).await; + wait_until_finished(&wait_response).await; + assert_eq!( + wait_response.await.unwrap().unwrap(), + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + }) + ); + + delegate.release_tool(); + let completion = service + .begin_wait(WaitRequest { + cell_id: cell_id("1"), + yield_time_ms: 10_000, + }) + .await; + let completion = tokio::spawn(completion); + wait_until_finished(&completion).await; + assert_eq!( + completion.await.unwrap().unwrap(), + WaitOutcome::LiveCell(RuntimeResponse::Result { + cell_id: cell_id("1"), + content_items: vec![FunctionCallOutputContentItem::InputText { + text: "done".to_string(), + }], + error_text: None, + }) + ); +} + #[tokio::test(start_paused = true)] async fn wait_waits_for_nested_tool_during_yield_grace() { let delegate = Arc::new(ReleasableToolDelegate::default()); @@ -113,6 +207,76 @@ async fn wait_waits_for_nested_tool_during_yield_grace() { ); } +#[tokio::test(start_paused = true)] +async fn zero_yield_limit_is_immediate_and_scoped_to_its_session() { + let zero_delegate = Arc::new(ReleasableToolDelegate::default()); + let zero_session = InProcessCodeModeSession::with_delegate_and_limits( + zero_delegate.clone(), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(/*value*/ 0), + max_heap_size_bytes: None, + }, + ); + let limited_delegate = Arc::new(ReleasableToolDelegate::default()); + let limited_session = InProcessCodeModeSession::with_delegate_and_limits( + limited_delegate.clone(), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(/*value*/ 10), + max_heap_size_bytes: None, + }, + ); + let request = ExecuteRequest { + enabled_tools: vec![echo_tool()], + source: "await tools.echo({});".to_string(), + yield_time_ms: Some(/*value*/ 60_000), + ..execute_request("") + }; + let zero_started = zero_session.execute(request.clone()).await.unwrap(); + let limited_started = limited_session.execute(request).await.unwrap(); + let zero_response = tokio::spawn(zero_started.initial_response()); + let limited_response = tokio::spawn(limited_started.initial_response()); + wait_until_tool_started(&zero_delegate).await; + wait_until_tool_started(&limited_delegate).await; + wait_until_finished(&zero_response).await; + assert!(!limited_response.is_finished()); + assert_eq!( + zero_response.await.unwrap().unwrap(), + RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + } + ); + + let zero_wait = zero_session + .begin_wait(WaitRequest { + cell_id: cell_id("1"), + yield_time_ms: 60_000, + }) + .await; + let zero_wait = tokio::spawn(zero_wait); + wait_until_finished(&zero_wait).await; + assert_eq!( + zero_wait.await.unwrap().unwrap(), + WaitOutcome::LiveCell(RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + }) + ); + + tokio::time::advance(Duration::from_millis(/*millis*/ 10)).await; + wait_until_finished(&limited_response).await; + assert_eq!( + limited_response.await.unwrap().unwrap(), + RuntimeResponse::Yielded { + cell_id: cell_id("1"), + content_items: Vec::new(), + } + ); + + zero_session.shutdown().await.unwrap(); + limited_session.shutdown().await.unwrap(); +} + async fn wait_until_finished(task: &tokio::task::JoinHandle) { for _ in 0..10_000 { if task.is_finished() { diff --git a/codex-rs/code-mode/src/remote_session.rs b/codex-rs/code-mode/src/remote_session.rs index 40d4b73934..8685cfb739 100644 --- a/codex-rs/code-mode/src/remote_session.rs +++ b/codex-rs/code-mode/src/remote_session.rs @@ -8,6 +8,7 @@ use std::sync::atomic::Ordering; use codex_code_mode_protocol::CellId; use codex_code_mode_protocol::CodeModeSession; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::CodeModeSessionDelegate; use codex_code_mode_protocol::CodeModeSessionProvider; use codex_code_mode_protocol::CodeModeSessionProviderFuture; @@ -85,7 +86,15 @@ impl CodeModeSessionProvider for ProcessOwnedCodeModeSessionProvider { &'a self, delegate: Arc, ) -> CodeModeSessionProviderFuture<'a> { - Box::pin(create_host_session(delegate, self.process_host())) + self.create_session_with_limits(delegate, CodeModeSessionCellExecutionLimits::default()) + } + + fn create_session_with_limits<'a>( + &'a self, + delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, + ) -> CodeModeSessionProviderFuture<'a> { + Box::pin(create_host_session(delegate, self.process_host(), limits)) } } @@ -100,6 +109,14 @@ impl CodeModeSessionProvider for DisabledCodeModeSessionProvider { ) -> CodeModeSessionProviderFuture<'a> { Box::pin(async { Err("code-mode host is disabled".to_string()) }) } + + fn create_session_with_limits<'a>( + &'a self, + delegate: Arc, + _limits: CodeModeSessionCellExecutionLimits, + ) -> CodeModeSessionProviderFuture<'a> { + self.create_session(delegate) + } } impl WebSocketCodeModeSessionProvider { @@ -129,15 +146,28 @@ impl CodeModeSessionProvider for WebSocketCodeModeSessionProvider { &'a self, delegate: Arc, ) -> CodeModeSessionProviderFuture<'a> { - Box::pin(create_host_session(delegate, Arc::clone(&self.host))) + self.create_session_with_limits(delegate, CodeModeSessionCellExecutionLimits::default()) + } + + fn create_session_with_limits<'a>( + &'a self, + delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, + ) -> CodeModeSessionProviderFuture<'a> { + Box::pin(create_host_session( + delegate, + Arc::clone(&self.host), + limits, + )) } } async fn create_host_session( delegate: Arc, host: Arc, + limits: CodeModeSessionCellExecutionLimits, ) -> Result, String> { - let session = ProcessOwnedCodeModeSession::with_host(delegate, host); + let session = ProcessOwnedCodeModeSession::with_host(delegate, host, limits); session.connection().await?; Ok(Arc::new(session)) } @@ -244,6 +274,7 @@ struct SessionBinding { struct SessionInner { host: Arc, delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, state: StdMutex, next_generation: AtomicU64, shutdown_requested: AtomicBool, @@ -263,14 +294,20 @@ impl ProcessOwnedCodeModeSession { Arc::new(OwnedCodeModeHost::new( InstallContext::current().code_mode_host_program(), )), + CodeModeSessionCellExecutionLimits::default(), ) } - fn with_host(delegate: Arc, host: Arc) -> Self { + fn with_host( + delegate: Arc, + host: Arc, + limits: CodeModeSessionCellExecutionLimits, + ) -> Self { Self { inner: Arc::new(SessionInner { host, delegate, + limits, state: StdMutex::new(SessionState::New), next_generation: AtomicU64::new(1), shutdown_requested: AtomicBool::new(false), @@ -361,7 +398,11 @@ impl SessionInner { let result = match self.host.connection().await { Ok(connection) => { let cleanup = connection - .open_session(remote.clone(), Arc::clone(&self.delegate)) + .open_session( + remote.clone(), + Arc::clone(&self.delegate), + self.limits.clone(), + ) .await; cleanup.map(|cleanup| SessionBinding { connection, diff --git a/codex-rs/code-mode/src/remote_session/connection.rs b/codex-rs/code-mode/src/remote_session/connection.rs index d08c8fbbc4..ff281e5647 100644 --- a/codex-rs/code-mode/src/remote_session/connection.rs +++ b/codex-rs/code-mode/src/remote_session/connection.rs @@ -10,6 +10,7 @@ use std::sync::atomic::Ordering; use std::time::Duration; use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::CodeModeSessionDelegate; use codex_code_mode_protocol::ExecuteRequest; use codex_code_mode_protocol::StartedCell; @@ -28,6 +29,7 @@ use codex_code_mode_protocol::host::MAX_FRAME_BYTES; use codex_code_mode_protocol::host::MAX_PENDING_DELEGATE_CALLS; use codex_code_mode_protocol::host::ProtocolVersion; use codex_code_mode_protocol::host::RequestId; +use codex_code_mode_protocol::host::SESSION_RESOURCE_LIMITS_CAPABILITY; use codex_code_mode_protocol::host::SupportedProtocolVersions; use codex_code_mode_protocol::host::TransportLane; use codex_http_client::HttpClientFactory; @@ -120,6 +122,7 @@ pub(super) struct Connection { alive: Arc, failure: Arc>>, cancellation: CancellationToken, + capabilities: CapabilitySet, } struct CallerCancellation { @@ -285,11 +288,14 @@ impl Connection { let handshake = async { let dual_capability = Capability::new(DUAL_WEBSOCKET_CAPABILITY).map_err(|error| error.to_string())?; + let session_limits_capability = Capability::new(SESSION_RESOURCE_LIMITS_CAPABILITY) + .map_err(|error| error.to_string())?; let optional_capabilities = if bulk_connection_options.is_some() { - CapabilitySet::try_new([dual_capability.clone()]) + CapabilitySet::try_new([dual_capability.clone(), session_limits_capability]) .map_err(|error| error.to_string())? } else { - CapabilitySet::empty() + CapabilitySet::try_new([session_limits_capability]) + .map_err(|error| error.to_string())? }; let hello = ClientHello::new( SupportedProtocolVersions::try_new([ProtocolVersion::V1]) @@ -310,7 +316,8 @@ impl Connection { Some(HostToClient::HostHello(hello)) if hello.selected_version() == ProtocolVersion::V1 => { - if hello.capabilities().contains(&dual_capability) { + let capabilities = hello.capabilities().clone(); + let bulk_token = if capabilities.contains(&dual_capability) { hello .bulk_connection_token() .map(str::to_string) @@ -318,12 +325,15 @@ impl Connection { "code-mode host advertised dual websockets without a pairing token" .to_string() }) - .map(Some) + .map(Some)? } else if hello.bulk_connection_token().is_some() { - Err("code-mode host returned an unexpected bulk pairing token".to_string()) + return Err( + "code-mode host returned an unexpected bulk pairing token".to_string() + ); } else { - Ok(None) - } + None + }; + Ok((capabilities, bulk_token)) } Some(HostToClient::HandshakeRejected { reason }) => { Err(format!("code-mode host rejected the handshake: {reason:?}")) @@ -344,8 +354,8 @@ impl Connection { )); } }; - let bulk_token = match handshake_result { - Ok(token) => token, + let (capabilities, bulk_token) = match handshake_result { + Ok(negotiated) => negotiated, Err(err) => { let _ = writer.close().await; owner.close().await; @@ -460,6 +470,7 @@ impl Connection { alive, failure, cancellation, + capabilities, }) } @@ -478,13 +489,25 @@ impl Connection { &self, session: RemoteSession, delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, ) -> Result { + if limits != CodeModeSessionCellExecutionLimits::default() + && !self + .capabilities + .iter() + .any(|capability| capability.as_str() == SESSION_RESOURCE_LIMITS_CAPABILITY) + { + return Err(format!( + "code-mode host does not support session resource limits: missing `{SESSION_RESOURCE_LIMITS_CAPABILITY}` capability" + )); + } let cleanup = SessionCleanup::new(); let cancellation = CallerCancellation::new(); let (response_tx, response_rx) = oneshot::channel(); self.send(DriverCommand::OpenSession { session, delegate, + limits, cleanup: cleanup.clone(), caller_cancellation: cancellation.token(), response_tx, diff --git a/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs b/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs index 4bfb713625..446c9c6abc 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver/commands.rs @@ -1,6 +1,7 @@ use std::sync::Arc; use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::CodeModeSessionDelegate; use codex_code_mode_protocol::ExecuteRequest; use codex_code_mode_protocol::WaitOutcome; @@ -8,6 +9,7 @@ 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::WireSessionCellExecutionLimits; use codex_code_mode_protocol::host::WireWaitRequest; use tokio::sync::oneshot; use tokio_util::sync::CancellationToken; @@ -28,10 +30,18 @@ impl ConnectionDriver { DriverCommand::OpenSession { session, delegate, + limits, cleanup, caller_cancellation, response_tx, - } => self.open_session(session, delegate, cleanup, caller_cancellation, response_tx), + } => self.open_session( + session, + delegate, + limits, + cleanup, + caller_cancellation, + response_tx, + ), DriverCommand::Execute { session, request, @@ -60,6 +70,7 @@ impl ConnectionDriver { &mut self, session: RemoteSession, delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, cleanup: super::cleanup::SessionCleanup, caller_cancellation: CancellationToken, response_tx: oneshot::Sender>, @@ -71,6 +82,15 @@ impl ConnectionDriver { ))); return true; } + let limits = match WireSessionCellExecutionLimits::try_from(limits) { + Ok(limits) => limits, + Err(error) => { + let _ = response_tx.send(Err(format!( + "failed to encode code-mode session execution limits: {error}" + ))); + return true; + } + }; let request_id = match self.requests.allocate_id() { Ok(id) => id, Err(err) => { @@ -82,6 +102,8 @@ impl ConnectionDriver { id: request_id, request: HostRequest::OpenSession { session_id: session.id.clone(), + cell_execution_limits: (limits != WireSessionCellExecutionLimits::default()) + .then_some(limits), }, }; let frame = match EncodedFrame::encode(&message) { 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 index d4f6a102f3..7b561b330a 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver/types.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver/types.rs @@ -1,6 +1,7 @@ use std::sync::Arc; use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::CodeModeSessionDelegate; use codex_code_mode_protocol::ExecuteRequest; use codex_code_mode_protocol::RuntimeResponse; @@ -30,6 +31,7 @@ pub(in crate::remote_session::connection) enum DriverCommand { OpenSession { session: RemoteSession, delegate: Arc, + limits: CodeModeSessionCellExecutionLimits, cleanup: SessionCleanup, caller_cancellation: CancellationToken, response_tx: oneshot::Sender>, diff --git a/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs b/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs index 2dc84eeeec..3f51607b44 100644 --- a/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs +++ b/codex-rs/code-mode/src/remote_session/connection/driver_tests.rs @@ -7,16 +7,19 @@ use std::time::Duration; use codex_code_mode_protocol::CellId; use codex_code_mode_protocol::CodeModeNestedToolCall; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; 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::CapabilitySet; 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::HostRequest; use codex_code_mode_protocol::host::HostResponse; use codex_code_mode_protocol::host::HostToClient; use codex_code_mode_protocol::host::MAX_PENDING_DELEGATE_CALLS; @@ -25,6 +28,7 @@ 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::WireSessionCellExecutionLimits; use codex_code_mode_protocol::host::WireWaitOutcome; use codex_protocol::ToolName; use pretty_assertions::assert_eq; @@ -95,6 +99,7 @@ impl DriverHarness { .send(DriverCommand::OpenSession { session: session.clone(), delegate, + limits: Default::default(), cleanup: cleanup.clone(), caller_cancellation: CancellationToken::new(), response_tx, @@ -190,6 +195,50 @@ impl Drop for DriverHarness { } } +#[tokio::test] +async fn open_session_includes_nondefault_cell_execution_limits() { + let mut harness = DriverHarness::start(); + let session = remote_session(); + let limits = CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: Some(16 * 1024 * 1024), + }; + let (response_tx, _response_rx) = oneshot::channel(); + + harness + .command_tx + .send(DriverCommand::OpenSession { + session: session.clone(), + delegate: Arc::new(RecordingDelegate::default()), + limits, + cleanup: SessionCleanup::new(), + caller_cancellation: CancellationToken::new(), + response_tx, + }) + .await + .expect("limited session open command"); + let frame = harness + .outgoing_rx + .recv() + .await + .expect("limited session open frame"); + + assert_eq!( + EncodedFrame::decode_framed::(&frame.into_framed_bytes()) + .expect("decode limited session open request"), + ClientToHost::Request { + id: RequestId::new(/*value*/ 1), + request: HostRequest::OpenSession { + session_id: session.id, + cell_execution_limits: Some(WireSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: Some(16 * 1024 * 1024), + }), + }, + } + ); +} + #[derive(Default)] struct RecordingDelegate { closed_cells: StdMutex>, @@ -518,6 +567,7 @@ async fn dropped_open_waiter_shuts_down_committed_session() { .send(DriverCommand::OpenSession { session: session.clone(), delegate: Arc::new(RecordingDelegate::default()), + limits: Default::default(), cleanup, caller_cancellation: CancellationToken::new(), response_tx: open_tx, @@ -1352,6 +1402,7 @@ async fn queued_remote_wait_times_out_and_invalidates_the_connection() { alive: Arc::clone(&harness.alive), failure: Arc::clone(&harness.failure), cancellation: harness.cancellation.clone(), + capabilities: CapabilitySet::empty(), }; let response = tokio::spawn(async move { let result = connection @@ -1397,6 +1448,7 @@ async fn queued_remote_termination_times_out_and_invalidates_the_connection() { alive: Arc::clone(&harness.alive), failure: Arc::clone(&harness.failure), cancellation: harness.cancellation.clone(), + capabilities: CapabilitySet::empty(), }; let response = tokio::spawn(async move { let result = connection diff --git a/codex-rs/code-mode/src/remote_session_tests.rs b/codex-rs/code-mode/src/remote_session_tests.rs index 32cd548003..d5fd559648 100644 --- a/codex-rs/code-mode/src/remote_session_tests.rs +++ b/codex-rs/code-mode/src/remote_session_tests.rs @@ -3,6 +3,7 @@ use std::path::PathBuf; use std::sync::Arc; use std::time::Duration; +use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits; use codex_code_mode_protocol::CodeModeSessionProvider; use codex_code_mode_protocol::ExecuteRequest; use codex_code_mode_protocol::FunctionCallOutputContentItem; @@ -17,6 +18,7 @@ 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::ProtocolVersion; +use codex_code_mode_protocol::host::SESSION_RESOURCE_LIMITS_CAPABILITY; use codex_code_mode_protocol::host::WireCellId; use codex_code_mode_protocol::host::WireContentItem; use codex_code_mode_protocol::host::WireResult; @@ -117,13 +119,19 @@ async fn websocket_provider_executes_over_shared_connector() { let request = EncodedFrame::decode_framed::(&frame) .expect("websocket test host should decode a framed protocol message"); let responses = match request { - ClientToHost::ClientHello(_) => vec![HostToClient::HostHello(HostHello::new( - ProtocolVersion::V1, - CapabilitySet::empty(), - ))], + ClientToHost::ClientHello(hello) => { + let capability = Capability::new(SESSION_RESOURCE_LIMITS_CAPABILITY) + .expect("session-limit capability"); + assert!(hello.optional_capabilities().contains(&capability)); + assert_eq!(hello.required_capabilities(), &CapabilitySet::empty()); + vec![HostToClient::HostHello(HostHello::new( + ProtocolVersion::V1, + CapabilitySet::empty(), + ))] + } ClientToHost::Request { id, - request: HostRequest::OpenSession { session_id }, + request: HostRequest::OpenSession { session_id, .. }, } => vec![HostToClient::Response { id, result: WireResult::Ok { @@ -192,6 +200,32 @@ async fn websocket_provider_executes_over_shared_connector() { .create_session(Arc::new(NoopCodeModeSessionDelegate)) .await .expect("shared websocket connector should open a code-mode session"); + let error = provider + .create_session_with_limits( + Arc::new(NoopCodeModeSessionDelegate), + CodeModeSessionCellExecutionLimits { + max_yield_time_ms: Some(250), + max_heap_size_bytes: None, + }, + ) + .await + .err() + .expect("legacy host should reject a limited session"); + assert_eq!( + error, + format!( + "code-mode host does not support session resource limits: missing `{SESSION_RESOURCE_LIMITS_CAPABILITY}` capability" + ) + ); + let second_session = provider + .create_session(Arc::new(NoopCodeModeSessionDelegate)) + .await + .expect("rejecting limited sessions should preserve the shared legacy-host connection"); + second_session + .shutdown() + .await + .expect("second unlimited session should shut down"); + drop(second_session); let response = session .execute(ExecuteRequest { tool_call_id: "shared-websocket".to_string(), @@ -261,6 +295,13 @@ async fn websocket_provider_fails_when_a_negotiated_bulk_connection_is_unavailab let capability = Capability::new(DUAL_WEBSOCKET_CAPABILITY).expect("dual websocket capability"); assert!(hello.optional_capabilities().contains(&capability)); + let session_limits_capability = + Capability::new(SESSION_RESOURCE_LIMITS_CAPABILITY).expect("session-limit capability"); + assert!( + hello + .optional_capabilities() + .contains(&session_limits_capability) + ); let hello = HostToClient::HostHello( HostHello::new( ProtocolVersion::V1,