Add gRPC-backed code-mode sessions (#38041)

## What changed

- Add `GrpcCodeModeSessionProvider` for opening code-mode sessions over HTTP/2 or an existing `tonic` channel.
- Support execution, waiting, termination, per-session limits, cell-closure callbacks, and graceful shutdown over the gRPC protocol.
- Bound transport waits and error messages, validate host identifiers and responses, and clean up abandoned executions and observers.

## Testing

- Add end-to-end TCP tests covering session persistence, cancellation, concurrent waits, shutdown, cell cleanup, and independent yield limits.
- Add unit coverage for protocol conversion, deadlines, and session lifecycle state.

GitOrigin-RevId: d4729ce608ad4b42a99744b07e1f230e46cb24ec
This commit is contained in:
Channing Conger
2026-08-11 17:17:58 +00:00
committed by copyberry
parent b2543af02b
commit 1e557a554e
20 changed files with 2275 additions and 17 deletions

3
codex-rs/Cargo.lock generated
View File

@@ -2564,10 +2564,13 @@ dependencies = [
"codex-websocket-client",
"futures",
"pretty_assertions",
"serde_json",
"tokio",
"tokio-tungstenite",
"tokio-util",
"tonic",
"tracing",
"uuid",
]
[[package]]

View File

@@ -18,6 +18,7 @@ use tonic::Request;
use tonic::Status;
use uuid::Uuid;
use super::ExecutionAdmission;
use super::GrpcCodeModeHost;
use super::tests::execute_events;
use super::tests::execute_request;
@@ -125,6 +126,31 @@ async fn rejects_oversized_identifiers_tool_metadata_and_subscription_filters()
assert!(host.state.session(&session_id).is_ok());
}
#[tokio::test]
async fn dropping_execution_before_admission_releases_its_reservation() {
let host = GrpcCodeModeHost::new();
let (session_id, _events) = open_session(&host).await;
let session = host.state.session(&session_id).expect("open session");
let execution_id = "execution-abandoned-before-admission".to_string();
session
.reserve_execution(&execution_id)
.expect("reserve execution");
drop(ExecutionAdmission {
session: Arc::clone(&session),
execution_id: Some(execution_id.clone()),
});
let error = session
.admit_execution(
execution_id,
"cell".to_string(),
host.state.cell_permit().expect("reserve cell permit"),
)
.expect_err("abandoned execution must not admit a runtime cell");
assert_eq!(error.code(), Code::Cancelled);
}
#[tokio::test]
async fn dropping_an_unread_buffered_execution_outcome_retires_its_cell() {
let host = GrpcCodeModeHost::new();

View File

@@ -2,7 +2,7 @@ use codex_code_mode_protocol::grpc as proto;
use tonic::Status;
use uuid::Uuid;
pub(super) const MAX_IDENTIFIER_BYTES: usize = 256;
pub(super) use codex_code_mode_protocol::grpc::MAX_IDENTIFIER_BYTES;
pub(super) const MAX_TOOL_FILTERS: usize = 64;
pub(super) const MAX_TOOL_DEFINITIONS: usize = 1_024;
pub(super) const MAX_TOOL_DESCRIPTION_BYTES: usize = 16 * 1_024;

View File

@@ -0,0 +1,645 @@
use std::sync::Arc;
use std::sync::PoisonError;
use std::time::Duration;
use anyhow::Context;
use anyhow::Result;
use codex_code_mode::CodeModeSession;
use codex_code_mode::CodeModeSessionCellExecutionLimits;
use codex_code_mode::CodeModeSessionProvider;
use codex_code_mode::ExecuteRequest;
use codex_code_mode::FunctionCallOutputContentItem;
use codex_code_mode::GrpcCodeModeSessionProvider;
use codex_code_mode::NoopCodeModeSessionDelegate;
use codex_code_mode::RuntimeResponse;
use codex_code_mode::WaitOutcome;
use codex_code_mode::WaitRequest;
use codex_code_mode_protocol::grpc;
use codex_code_mode_protocol::grpc::code_mode_host_client::CodeModeHostClient;
use futures::FutureExt;
use pretty_assertions::assert_eq;
use tokio::time::timeout;
use tonic::Code;
#[path = "support/host.rs"]
mod host;
#[path = "support/recording_delegate.rs"]
mod recording_delegate;
use host::HostHarness;
use recording_delegate::RecordingDelegate;
use recording_delegate::cell_id;
const TEST_TIMEOUT: Duration = Duration::from_secs(20);
fn request(source: &str) -> ExecuteRequest {
ExecuteRequest {
tool_call_id: "call-1".to_string(),
enabled_tools: Vec::new(),
source: source.to_string(),
yield_time_ms: Some(/*value*/ 5_000),
max_output_tokens: Some(/*value*/ 1_000),
}
}
fn text_response(cell: &str, value: &str) -> RuntimeResponse {
RuntimeResponse::Result {
cell_id: cell_id(cell),
content_items: vec![FunctionCallOutputContentItem::InputText {
text: value.to_string(),
}],
error_text: None,
}
}
async fn execute(
session: &Arc<dyn CodeModeSession>,
request: ExecuteRequest,
) -> Result<RuntimeResponse> {
timeout(TEST_TIMEOUT, async {
session
.execute(request)
.await
.map_err(anyhow::Error::msg)?
.initial_response()
.await
.map_err(anyhow::Error::msg)
})
.await
.context("timed out executing gRPC code-mode cell")?
}
async fn start_active_wait(
session: Arc<dyn CodeModeSession>,
request: WaitRequest,
) -> Result<tokio::task::JoinHandle<std::result::Result<WaitOutcome, String>>> {
let (admitted_tx, admitted_rx) = tokio::sync::oneshot::channel();
let wait = tokio::spawn(async move {
let mut wait = session.wait(request);
match wait.as_mut().now_or_never() {
Some(result) => result,
None => {
let _ = admitted_tx.send(());
wait.await
}
}
});
timeout(TEST_TIMEOUT, admitted_rx)
.await
.context("timed out waiting for observer admission")?
.context("wait completed before its observer became active")?;
Ok(wait)
}
#[tokio::test]
async fn tcp_session_persists_values_and_reports_cell_closure() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
assert!(host.endpoint.starts_with("http://127.0.0.1:"));
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let delegate = Arc::new(RecordingDelegate::default());
let session = provider
.create_session(delegate.clone())
.await
.map_err(anyhow::Error::msg)?;
assert_eq!(
execute(&session, request(r#"store("key", "persisted");"#)).await?,
RuntimeResponse::Result {
cell_id: cell_id("1"),
content_items: Vec::new(),
error_text: None,
}
);
assert_eq!(
execute(&session, request(r#"text(String(load("key")));"#)).await?,
text_response("2", "persisted")
);
session.shutdown().await.map_err(anyhow::Error::msg)?;
assert_eq!(
*delegate
.closed_cells
.lock()
.unwrap_or_else(PoisonError::into_inner),
vec![cell_id("1"), cell_id("2")]
);
Ok(())
}
#[tokio::test]
async fn shutdown_immediately_rejects_new_operations() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let session = provider
.create_session(Arc::new(NoopCodeModeSessionDelegate))
.await
.map_err(anyhow::Error::msg)?;
let shutdown = session.shutdown();
let expected = "code mode session is shutting down".to_string();
assert_eq!(
session.execute(request("text('too late');")).await.err(),
Some(expected.clone())
);
assert_eq!(
session
.wait(WaitRequest {
cell_id: cell_id("missing"),
yield_time_ms: 1,
})
.await,
Err(expected.clone())
);
assert_eq!(session.terminate(cell_id("missing")).await, Err(expected));
shutdown.await.map_err(anyhow::Error::msg)?;
Ok(())
}
#[tokio::test]
async fn cancelling_execution_before_admission_keeps_the_session_usable() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let delegate = Arc::new(RecordingDelegate::default());
let session = provider
.create_session(delegate.clone())
.await
.map_err(anyhow::Error::msg)?;
let mut pending = request("await new Promise(() => {});");
pending.yield_time_ms = Some(/*value*/ 1);
assert!(session.execute(pending).now_or_never().is_none());
let abandoned_cell = cell_id("1");
timeout(TEST_TIMEOUT, async {
loop {
if delegate
.closed_cells
.lock()
.unwrap_or_else(PoisonError::into_inner)
.contains(&abandoned_cell)
{
break;
}
tokio::task::yield_now().await;
}
})
.await
.context("cancelled execution was never admitted and cleaned up")?;
timeout(TEST_TIMEOUT, async {
loop {
match session
.wait(WaitRequest {
cell_id: abandoned_cell.clone(),
yield_time_ms: 1,
})
.await
{
Ok(WaitOutcome::MissingCell(_))
| Ok(WaitOutcome::LiveCell(RuntimeResponse::Terminated { .. })) => break Ok(()),
Ok(WaitOutcome::LiveCell(RuntimeResponse::Yielded { .. })) => {
tokio::task::yield_now().await;
}
Ok(outcome) => anyhow::bail!("unexpected abandoned-cell outcome: {outcome:?}"),
Err(error) if error.contains("already has an active observer") => {
tokio::task::yield_now().await;
}
Err(error) => break Err(anyhow::Error::msg(error)),
}
}
})
.await
.context("cancelled execution leaked its remote cell")??;
assert_eq!(
execute(&session, request(r#"text("still alive");"#)).await?,
text_response("2", "still alive")
);
session.shutdown().await.map_err(anyhow::Error::msg)?;
Ok(())
}
#[tokio::test]
async fn dropping_a_started_cell_off_runtime_terminates_its_buffered_remote_execution() -> Result<()>
{
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let delegate = Arc::new(RecordingDelegate::default());
let session = provider
.create_session(delegate.clone())
.await
.map_err(anyhow::Error::msg)?;
let mut pending = request("await new Promise(() => {});");
pending.yield_time_ms = Some(/*value*/ 1);
let started = session.execute(pending).await.map_err(anyhow::Error::msg)?;
let abandoned_cell = started.cell_id.clone();
timeout(TEST_TIMEOUT, async {
loop {
match session
.wait(WaitRequest {
cell_id: abandoned_cell.clone(),
yield_time_ms: 1,
})
.await
{
Ok(WaitOutcome::LiveCell(RuntimeResponse::Yielded { .. })) => break Ok(()),
Err(error) if error.contains("already has an active observer") => {
tokio::task::yield_now().await;
}
Ok(outcome) => anyhow::bail!("unexpected execution outcome: {outcome:?}"),
Err(error) => break Err(anyhow::Error::msg(error)),
}
}
})
.await
.context("execution never produced a buffered initial response")??;
std::thread::spawn(move || drop(started))
.join()
.map_err(|_| anyhow::anyhow!("dropping a started cell outside Tokio panicked"))?;
timeout(TEST_TIMEOUT, async {
loop {
match session
.wait(WaitRequest {
cell_id: abandoned_cell.clone(),
yield_time_ms: 1,
})
.await
{
Ok(WaitOutcome::MissingCell(_))
| Ok(WaitOutcome::LiveCell(RuntimeResponse::Terminated { .. })) => break Ok(()),
Ok(WaitOutcome::LiveCell(RuntimeResponse::Yielded { .. })) => {
tokio::task::yield_now().await;
}
Ok(outcome) => anyhow::bail!("unexpected abandoned-cell outcome: {outcome:?}"),
Err(error) if error.contains("already has an active observer") => {
tokio::task::yield_now().await;
}
Err(error) => break Err(anyhow::Error::msg(error)),
}
}
})
.await
.context("dropping a started cell did not terminate its buffered remote execution")??;
assert!(
delegate
.closed_cells
.lock()
.unwrap_or_else(PoisonError::into_inner)
.contains(&abandoned_cell)
);
assert_eq!(
execute(&session, request(r#"text("still alive");"#)).await?,
text_response("2", "still alive")
);
session.shutdown().await.map_err(anyhow::Error::msg)?;
Ok(())
}
#[tokio::test]
async fn dropping_an_initial_response_terminates_its_pending_remote_execution() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let delegate = Arc::new(RecordingDelegate::default());
let session = provider
.create_session(delegate.clone())
.await
.map_err(anyhow::Error::msg)?;
let mut pending = request("await new Promise(() => {});");
pending.yield_time_ms = Some(/*value*/ 60_000);
let started = session.execute(pending).await.map_err(anyhow::Error::msg)?;
let abandoned_cell = started.cell_id.clone();
let initial_response = tokio::spawn(started.initial_response());
tokio::task::yield_now().await;
initial_response.abort();
let _ = initial_response.await;
timeout(TEST_TIMEOUT, async {
loop {
if delegate
.closed_cells
.lock()
.unwrap_or_else(PoisonError::into_inner)
.contains(&abandoned_cell)
{
break;
}
tokio::task::yield_now().await;
}
})
.await
.context("dropping an initial response did not terminate its pending remote execution")?;
assert_eq!(
execute(&session, request(r#"text("still alive");"#)).await?,
text_response("2", "still alive")
);
session.shutdown().await.map_err(anyhow::Error::msg)?;
Ok(())
}
#[tokio::test]
async fn concurrent_wait_rejects_without_displacing_the_active_observer() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let session = provider
.create_session(Arc::new(NoopCodeModeSessionDelegate))
.await
.map_err(anyhow::Error::msg)?;
let mut pending = request("await new Promise(() => {});");
pending.yield_time_ms = Some(/*value*/ 1);
let started = session.execute(pending).await.map_err(anyhow::Error::msg)?;
let running_cell = started.cell_id.clone();
assert_eq!(
started
.initial_response()
.await
.map_err(anyhow::Error::msg)?,
RuntimeResponse::Yielded {
cell_id: running_cell.clone(),
content_items: Vec::new(),
}
);
let first_wait = start_active_wait(
Arc::clone(&session),
WaitRequest {
cell_id: running_cell.clone(),
yield_time_ms: 100,
},
)
.await?;
assert_eq!(
timeout(
Duration::from_secs(/*secs*/ 1),
session.wait(WaitRequest {
cell_id: running_cell.clone(),
yield_time_ms: 60_000,
}),
)
.await
.context("concurrent wait did not reject immediately")?
.unwrap_err(),
format!("exec cell {running_cell} already has an active observer")
);
assert_eq!(
timeout(TEST_TIMEOUT, first_wait)
.await
.context("active wait was displaced by the rejected observer")??
.map_err(anyhow::Error::msg)?,
WaitOutcome::LiveCell(RuntimeResponse::Yielded {
cell_id: running_cell.clone(),
content_items: Vec::new(),
})
);
assert_eq!(
session
.terminate(running_cell.clone())
.await
.map_err(anyhow::Error::msg)?,
WaitOutcome::LiveCell(RuntimeResponse::Terminated {
cell_id: running_cell,
content_items: Vec::new(),
})
);
session.shutdown().await.map_err(anyhow::Error::msg)?;
Ok(())
}
#[tokio::test]
async fn dropping_a_wait_retires_its_observer_before_the_next_wait() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let session = provider
.create_session(Arc::new(NoopCodeModeSessionDelegate))
.await
.map_err(anyhow::Error::msg)?;
let mut pending = request("await new Promise(() => {});");
pending.yield_time_ms = Some(/*value*/ 1);
let started = session.execute(pending).await.map_err(anyhow::Error::msg)?;
let running_cell = started.cell_id.clone();
assert_eq!(
started
.initial_response()
.await
.map_err(anyhow::Error::msg)?,
RuntimeResponse::Yielded {
cell_id: running_cell.clone(),
content_items: Vec::new(),
}
);
let first_wait = start_active_wait(
Arc::clone(&session),
WaitRequest {
cell_id: running_cell.clone(),
yield_time_ms: 60_000,
},
)
.await?;
first_wait.abort();
let _ = first_wait.await;
assert_eq!(
timeout(
TEST_TIMEOUT,
session.wait(WaitRequest {
cell_id: running_cell.clone(),
yield_time_ms: 1,
}),
)
.await
.context("replacement wait did not observe cancellation retirement")?
.map_err(anyhow::Error::msg)?,
WaitOutcome::LiveCell(RuntimeResponse::Yielded {
cell_id: running_cell.clone(),
content_items: Vec::new(),
})
);
assert_eq!(
session
.terminate(running_cell.clone())
.await
.map_err(anyhow::Error::msg)?,
WaitOutcome::LiveCell(RuntimeResponse::Terminated {
cell_id: running_cell,
content_items: Vec::new(),
})
);
session.shutdown().await.map_err(anyhow::Error::msg)?;
Ok(())
}
#[tokio::test]
async fn dropping_a_session_off_runtime_retires_its_active_cells() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let delegate = Arc::new(RecordingDelegate::default());
let session = provider
.create_session(delegate.clone())
.await
.map_err(anyhow::Error::msg)?;
let mut pending = request("await new Promise(() => {});");
pending.yield_time_ms = Some(/*value*/ 1);
assert_eq!(
execute(&session, pending).await?,
RuntimeResponse::Yielded {
cell_id: cell_id("1"),
content_items: Vec::new(),
}
);
std::thread::spawn(move || drop(session))
.join()
.map_err(|_| anyhow::anyhow!("dropping a session outside Tokio panicked"))?;
timeout(TEST_TIMEOUT, async {
loop {
if delegate
.closed_cells
.lock()
.unwrap_or_else(PoisonError::into_inner)
.contains(&cell_id("1"))
{
break;
}
tokio::task::yield_now().await;
}
})
.await
.context("dropping a session outside Tokio did not retire its active cell")?;
Ok(())
}
#[tokio::test]
async fn sessions_enforce_independent_yield_limits() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let provider = GrpcCodeModeSessionProvider::new(host.endpoint);
let limited = provider
.create_session_with_limits(
Arc::new(NoopCodeModeSessionDelegate),
CodeModeSessionCellExecutionLimits {
max_yield_time_ms: Some(/*value*/ 1),
max_heap_size_bytes: Some(/*value*/ 16 * 1024 * 1024),
},
)
.await
.map_err(anyhow::Error::msg)?;
let other = provider
.create_session_with_limits(
Arc::new(NoopCodeModeSessionDelegate),
CodeModeSessionCellExecutionLimits {
max_yield_time_ms: Some(/*value*/ 1_000),
max_heap_size_bytes: None,
},
)
.await
.map_err(anyhow::Error::msg)?;
let mut pending = request("await new Promise(() => {});");
pending.yield_time_ms = Some(/*value*/ 60_000);
assert_eq!(
execute(&limited, pending).await?,
RuntimeResponse::Yielded {
cell_id: cell_id("1"),
content_items: Vec::new(),
}
);
assert_eq!(
timeout(
TEST_TIMEOUT,
limited.wait(WaitRequest {
cell_id: cell_id("1"),
yield_time_ms: 60_000,
}),
)
.await
.context("session yield limit did not bound an explicit wait")?
.map_err(anyhow::Error::msg)?,
WaitOutcome::LiveCell(RuntimeResponse::Yielded {
cell_id: cell_id("1"),
content_items: Vec::new(),
})
);
assert_eq!(
execute(
&other,
request(r#"await new Promise(resolve => setTimeout(resolve, 25)); text("isolated");"#),
)
.await?,
text_response("1", "isolated")
);
assert_eq!(
limited
.terminate(cell_id("1"))
.await
.map_err(anyhow::Error::msg)?,
WaitOutcome::LiveCell(RuntimeResponse::Terminated {
cell_id: cell_id("1"),
content_items: Vec::new(),
})
);
assert_eq!(
execute(&limited, request(r#"text("recovered");"#)).await?,
text_response("2", "recovered")
);
limited.shutdown().await.map_err(anyhow::Error::msg)?;
other.shutdown().await.map_err(anyhow::Error::msg)?;
Ok(())
}
#[tokio::test]
async fn dropping_a_grpc_lease_retires_its_server_session() -> Result<()> {
let host = HostHarness::start("grpc://127.0.0.1:0").await?;
let mut client = CodeModeHostClient::connect(host.endpoint)
.await
.context("connect raw gRPC client")?;
let mut lease = client
.open_session(grpc::OpenSessionRequest {
cell_execution_limits: None,
})
.await
.context("open raw gRPC session")?
.into_inner();
let first = lease
.message()
.await
.context("read raw gRPC session opening")?
.context("raw gRPC session ended before its opening event")?;
let Some(grpc::session_event::Event::Opened(opened)) = first.event else {
anyhow::bail!("raw gRPC session did not start with an opening event");
};
drop(lease);
timeout(TEST_TIMEOUT, async {
loop {
match client
.subscribe_to_tool_calls(grpc::SubscribeToToolCallsRequest {
session_id: opened.session_id.clone(),
tool_names: Vec::new(),
})
.await
{
Ok(response) => drop(response),
Err(status) if status.code() == Code::NotFound => return Ok::<_, anyhow::Error>(()),
Err(status) => {
anyhow::bail!("unexpected session status after lease drop: {status}")
}
}
tokio::task::yield_now().await;
}
})
.await
.context("dropping the gRPC lease did not retire its server session")??;
Ok(())
}

View File

@@ -134,7 +134,8 @@ message CellClosed {
string execution_id = 1;
string cell_id = 2;
// All lower sequences must be observed before the client retires the cell.
// Last tool-call sequence issued before closure. Clients may retire the cell
// immediately and reject tool calls delivered after its closure.
uint64 final_tool_call_sequence = 3;
}

View File

@@ -3,3 +3,5 @@ pub use code_mode_proto::codex::code_mode::v1::*;
#[cfg(not(codex_bazel))]
tonic::include_proto!("codex.code_mode.v1");
pub const MAX_IDENTIFIER_BYTES: usize = 256;

View File

@@ -62,27 +62,31 @@ pub struct StartedCell {
impl StartedCell {
pub fn new(cell_id: CellId, initial_response_rx: oneshot::Receiver<RuntimeResponse>) -> Self {
Self {
cell_id,
initial_response: Box::pin(async move {
initial_response_rx
.await
.map_err(|_| "exec runtime ended unexpectedly".to_string())
}),
}
Self::from_future(cell_id, async move {
initial_response_rx
.await
.map_err(|_| "exec runtime ended unexpectedly".to_string())
})
}
pub fn from_result_receiver(
cell_id: CellId,
initial_response_rx: oneshot::Receiver<Result<RuntimeResponse, String>>,
) -> Self {
Self::from_future(cell_id, async move {
initial_response_rx
.await
.map_err(|_| "exec runtime ended unexpectedly".to_string())?
})
}
pub fn from_future(
cell_id: CellId,
initial_response: impl Future<Output = Result<RuntimeResponse, String>> + Send + 'static,
) -> Self {
Self {
cell_id,
initial_response: Box::pin(async move {
initial_response_rx
.await
.map_err(|_| "exec runtime ended unexpectedly".to_string())?
}),
initial_response: Box::pin(initial_response),
}
}

View File

@@ -19,10 +19,13 @@ codex-install-context = { workspace = true }
codex-protocol = { workspace = true }
codex-websocket-client = { workspace = true }
futures = { workspace = true }
serde_json = { workspace = true }
tokio = { workspace = true, features = ["io-util", "macros", "net", "process", "rt", "sync", "time"] }
tokio-tungstenite = { workspace = true }
tokio-util = { workspace = true, features = ["rt"] }
tonic = { workspace = true }
tracing = { workspace = true }
uuid = { workspace = true, features = ["v4"] }
[dev-dependencies]
pretty_assertions = { workspace = true }

View File

@@ -0,0 +1,70 @@
use std::sync::Arc;
use std::sync::PoisonError;
use std::sync::atomic::Ordering;
use codex_code_mode_protocol::grpc;
use super::SessionInner;
impl SessionInner {
pub(super) fn spawn_session_events(
self: &Arc<Self>,
mut events: tonic::Streaming<grpc::SessionEvent>,
) {
let inner = Arc::clone(self);
self.stream_tasks.spawn(async move {
loop {
let event = tokio::select! {
biased;
_ = inner.stopped.cancelled() => return,
event = events.message() => event,
};
match event {
Ok(Some(event)) => {
if let Err(error) = inner.handle_session_event(event) {
inner.fail(error);
return;
}
}
Ok(None) => {
if !inner.shutdown_requested.load(Ordering::Acquire) {
inner.fail(
"gRPC code-mode session lease closed unexpectedly".to_string(),
);
}
return;
}
Err(error) => {
if !inner.shutdown_requested.load(Ordering::Acquire) {
inner.fail(super::deadline::failure("session lease", error));
}
return;
}
}
}
});
}
fn handle_session_event(&self, event: grpc::SessionEvent) -> Result<(), String> {
match event
.event
.ok_or_else(|| "gRPC code-mode host sent an empty session event".to_string())?
{
grpc::session_event::Event::Opened(_) => {
Err("gRPC code-mode host repeated the session opening event".to_string())
}
grpc::session_event::Event::ToolCallCancelled(_)
| grpc::session_event::Event::Notification(_)
| grpc::session_event::Event::NotificationCancelled(_) => Ok(()),
grpc::session_event::Event::CellClosed(closed) => {
let cell = self
.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.close_cell(closed)?;
self.report_closed_cell(cell);
Ok(())
}
}
}
}

View File

@@ -0,0 +1,134 @@
use codex_code_mode_protocol::CellId;
use codex_code_mode_protocol::CodeModeToolKind;
use codex_code_mode_protocol::ExecuteRequest;
use codex_code_mode_protocol::FunctionCallOutputContentItem;
use codex_code_mode_protocol::ImageDetail;
use codex_code_mode_protocol::RuntimeResponse;
use codex_code_mode_protocol::ToolDefinition;
use codex_code_mode_protocol::WaitOutcome;
use codex_code_mode_protocol::grpc;
pub(super) fn execute_request(
session_id: &str,
execution_id: String,
request: ExecuteRequest,
) -> Result<grpc::ExecuteRequest, String> {
Ok(grpc::ExecuteRequest {
session_id: session_id.to_string(),
execution_id,
tool_call_id: request.tool_call_id,
source: request.source,
enabled_tools: request
.enabled_tools
.into_iter()
.map(tool_definition)
.collect::<Result<Vec<_>, _>>()?,
yield_time_ms: request.yield_time_ms,
max_output_tokens: request
.max_output_tokens
.map(u64::try_from)
.transpose()
.map_err(|error| format!("invalid code-mode output token limit: {error}"))?,
})
}
fn tool_definition(definition: ToolDefinition) -> Result<grpc::ToolDefinition, String> {
Ok(grpc::ToolDefinition {
name: definition.name,
tool_name: Some(grpc::ToolName {
name: definition.tool_name.name,
namespace: definition.tool_name.namespace,
}),
description: definition.description,
kind: match definition.kind {
CodeModeToolKind::Function => grpc::ToolKind::Function as i32,
CodeModeToolKind::Freeform => grpc::ToolKind::Freeform as i32,
},
input_schema_json: definition
.input_schema
.map(|schema| serde_json::to_vec(&schema))
.transpose()
.map_err(|error| format!("failed to encode code-mode tool input schema: {error}"))?,
output_schema_json: definition
.output_schema
.map(|schema| serde_json::to_vec(&schema))
.transpose()
.map_err(|error| format!("failed to encode code-mode tool output schema: {error}"))?,
})
}
pub(super) fn runtime_response(outcome: grpc::ExecutionOutcome) -> Result<RuntimeResponse, String> {
super::validate_identifier(&outcome.cell_id, "cell ID")?;
let cell_id = CellId::new(outcome.cell_id);
let content_items = outcome
.content_items
.into_iter()
.map(content_item)
.collect::<Result<Vec<_>, _>>()?;
match outcome
.outcome
.ok_or_else(|| "code-mode execution omitted its outcome".to_string())?
{
grpc::execution_outcome::Outcome::Yielded(_) => Ok(RuntimeResponse::Yielded {
cell_id,
content_items,
}),
grpc::execution_outcome::Outcome::Terminated(_) => Ok(RuntimeResponse::Terminated {
cell_id,
content_items,
}),
grpc::execution_outcome::Outcome::Completed(completed) => Ok(RuntimeResponse::Result {
cell_id,
content_items,
error_text: completed.error_text,
}),
}
}
pub(super) fn wait_outcome(response: grpc::WaitResponse) -> Result<WaitOutcome, String> {
match response
.state
.ok_or_else(|| "code-mode wait omitted its outcome".to_string())?
{
grpc::wait_response::State::LiveCell(response) => {
runtime_response(response).map(WaitOutcome::LiveCell)
}
grpc::wait_response::State::MissingCell(response) => {
runtime_response(response).map(WaitOutcome::MissingCell)
}
}
}
fn content_item(item: grpc::ContentItem) -> Result<FunctionCallOutputContentItem, String> {
match item
.item
.ok_or_else(|| "code-mode execution returned an empty content item".to_string())?
{
grpc::content_item::Item::Text(text) => {
Ok(FunctionCallOutputContentItem::InputText { text: text.text })
}
grpc::content_item::Item::Image(image) => Ok(FunctionCallOutputContentItem::InputImage {
image_url: image.image_url,
detail: image.detail.map(image_detail).transpose()?,
}),
grpc::content_item::Item::Audio(audio) => Ok(FunctionCallOutputContentItem::InputAudio {
audio_url: audio.audio_url,
}),
}
}
fn image_detail(value: i32) -> Result<ImageDetail, String> {
match grpc::ImageDetail::try_from(value) {
Ok(grpc::ImageDetail::Auto) => Ok(ImageDetail::Auto),
Ok(grpc::ImageDetail::Low) => Ok(ImageDetail::Low),
Ok(grpc::ImageDetail::High) => Ok(ImageDetail::High),
Ok(grpc::ImageDetail::Original) => Ok(ImageDetail::Original),
Ok(grpc::ImageDetail::Unspecified) | Err(_) => {
Err(format!("code-mode image has invalid detail value {value}"))
}
}
}
#[cfg(test)]
#[path = "conversion_tests.rs"]
mod tests;

View File

@@ -0,0 +1,168 @@
use codex_code_mode_protocol::CellId;
use codex_code_mode_protocol::CodeModeToolKind;
use codex_code_mode_protocol::ExecuteRequest;
use codex_code_mode_protocol::FunctionCallOutputContentItem;
use codex_code_mode_protocol::ImageDetail;
use codex_code_mode_protocol::RuntimeResponse;
use codex_code_mode_protocol::ToolDefinition;
use codex_code_mode_protocol::WaitOutcome;
use codex_code_mode_protocol::grpc;
use codex_protocol::ToolName;
use pretty_assertions::assert_eq;
use serde_json::json;
use super::execute_request;
use super::runtime_response;
use super::wait_outcome;
#[test]
fn execute_request_preserves_tool_schemas_namespaces_and_limits() {
let request = ExecuteRequest {
tool_call_id: "outer".to_string(),
enabled_tools: vec![ToolDefinition {
name: "search".to_string(),
tool_name: ToolName::namespaced("work", "search"),
description: "search the workspace".to_string(),
kind: CodeModeToolKind::Freeform,
input_schema: Some(json!({"type": "object"})),
output_schema: Some(json!({"type": "string"})),
}],
source: "text('hello')".to_string(),
yield_time_ms: Some(25),
max_output_tokens: Some(128),
};
assert_eq!(
execute_request("session", "execution".to_string(), request),
Ok(grpc::ExecuteRequest {
session_id: "session".to_string(),
execution_id: "execution".to_string(),
tool_call_id: "outer".to_string(),
source: "text('hello')".to_string(),
enabled_tools: vec![grpc::ToolDefinition {
name: "search".to_string(),
tool_name: Some(grpc::ToolName {
name: "search".to_string(),
namespace: Some("work".to_string()),
}),
description: "search the workspace".to_string(),
kind: grpc::ToolKind::Freeform as i32,
input_schema_json: Some(br#"{"type":"object"}"#.to_vec()),
output_schema_json: Some(br#"{"type":"string"}"#.to_vec()),
}],
yield_time_ms: Some(25),
max_output_tokens: Some(128),
})
);
}
#[test]
fn runtime_response_decodes_mixed_content_items() {
let outcome = grpc::ExecutionOutcome {
cell_id: "cell".to_string(),
content_items: vec![
grpc::ContentItem {
item: Some(grpc::content_item::Item::Text(grpc::TextContent {
text: "hello".to_string(),
})),
},
grpc::ContentItem {
item: Some(grpc::content_item::Item::Image(grpc::ImageContent {
image_url: "data:image/png;base64,AA==".to_string(),
detail: Some(grpc::ImageDetail::Original as i32),
})),
},
grpc::ContentItem {
item: Some(grpc::content_item::Item::Audio(grpc::AudioContent {
audio_url: "data:audio/wav;base64,AA==".to_string(),
})),
},
],
outcome: Some(grpc::execution_outcome::Outcome::Completed(
grpc::ExecutionCompleted {
error_text: Some("warning".to_string()),
},
)),
};
assert_eq!(
runtime_response(outcome),
Ok(RuntimeResponse::Result {
cell_id: CellId::new("cell".to_string()),
content_items: vec![
FunctionCallOutputContentItem::InputText {
text: "hello".to_string(),
},
FunctionCallOutputContentItem::InputImage {
image_url: "data:image/png;base64,AA==".to_string(),
detail: Some(ImageDetail::Original),
},
FunctionCallOutputContentItem::InputAudio {
audio_url: "data:audio/wav;base64,AA==".to_string(),
},
],
error_text: Some("warning".to_string()),
})
);
}
#[test]
fn wait_outcome_preserves_missing_cell_state() {
let response = grpc::WaitResponse {
state: Some(grpc::wait_response::State::MissingCell(
grpc::ExecutionOutcome {
cell_id: "missing".to_string(),
content_items: Vec::new(),
outcome: Some(grpc::execution_outcome::Outcome::Terminated(
grpc::ExecutionTerminated {},
)),
},
)),
};
assert_eq!(
wait_outcome(response),
Ok(WaitOutcome::MissingCell(RuntimeResponse::Terminated {
cell_id: CellId::new("missing".to_string()),
content_items: Vec::new(),
}))
);
}
#[test]
fn oversized_response_cell_ids_are_rejected() {
let response = grpc::ExecutionOutcome {
cell_id: "x".repeat(grpc::MAX_IDENTIFIER_BYTES + 1),
content_items: Vec::new(),
outcome: Some(grpc::execution_outcome::Outcome::Yielded(
grpc::ExecutionYielded {},
)),
};
assert_eq!(
runtime_response(response),
Err(format!(
"gRPC code-mode host returned cell ID exceeding {} bytes",
grpc::MAX_IDENTIFIER_BYTES
))
);
}
#[test]
fn invalid_output_enums_and_missing_oneofs_are_rejected() {
let invalid_image = grpc::ExecutionOutcome {
cell_id: "cell".to_string(),
content_items: vec![grpc::ContentItem {
item: Some(grpc::content_item::Item::Image(grpc::ImageContent {
image_url: "image".to_string(),
detail: Some(grpc::ImageDetail::Unspecified as i32),
})),
}],
outcome: Some(grpc::execution_outcome::Outcome::Yielded(
grpc::ExecutionYielded {},
)),
};
assert!(runtime_response(invalid_image).is_err());
assert!(wait_outcome(grpc::WaitResponse { state: None }).is_err());
}

View File

@@ -0,0 +1,77 @@
use std::fmt::Display;
use std::future::Future;
use std::time::Duration;
use super::SessionInner;
const TRANSPORT_TIMEOUT: Duration = Duration::from_secs(60);
const MAX_ERROR_BYTES: usize = 512;
pub(super) async fn startup<T, E: Display>(
operation: &str,
request: impl Future<Output = Result<T, E>>,
) -> Result<T, String> {
match enforce(operation, Duration::ZERO, request).await {
Ok(result) => Ok(result),
Err(RequestError::Failed(error)) => Err(failure(operation, error)),
Err(RequestError::TimedOut(reason)) => Err(reason),
}
}
pub(super) async fn request<T>(
session: &SessionInner,
operation: &str,
runtime_timeout: Duration,
request: impl Future<Output = Result<T, tonic::Status>>,
) -> Result<T, String> {
let result = tokio::select! {
biased;
_ = session.stopped.cancelled() => {
return Err("gRPC code-mode session closed".to_string());
}
result = enforce(operation, runtime_timeout, request) => result,
};
match result {
Ok(value) => Ok(value),
Err(RequestError::Failed(error)) => Err(failure(operation, error)),
Err(RequestError::TimedOut(reason)) => {
session.fail(reason.clone());
Err(reason)
}
}
}
pub(super) fn failure(operation: &str, error: impl Display) -> String {
let mut message = format!("gRPC code-mode {operation} failed: {error}");
if message.len() > MAX_ERROR_BYTES {
let boundary = message.floor_char_boundary(MAX_ERROR_BYTES - "...".len());
message.truncate(boundary);
message.push_str("...");
}
message
}
async fn enforce<T, E>(
operation: &str,
runtime_timeout: Duration,
request: impl Future<Output = Result<T, E>>,
) -> Result<T, RequestError<E>> {
let timeout = runtime_timeout.saturating_add(TRANSPORT_TIMEOUT);
match tokio::time::timeout(timeout, request).await {
Ok(Ok(value)) => Ok(value),
Ok(Err(error)) => Err(RequestError::Failed(error)),
Err(_) => Err(RequestError::TimedOut(format!(
"gRPC code-mode host timed out waiting for {operation} response"
))),
}
}
enum RequestError<E> {
Failed(E),
TimedOut(String),
}
#[cfg(test)]
#[path = "deadline_tests.rs"]
mod tests;

View File

@@ -0,0 +1,131 @@
use std::future::pending;
use std::sync::Arc;
use std::time::Duration;
use codex_code_mode_protocol::DEFAULT_EXEC_YIELD_TIME_MS;
use pretty_assertions::assert_eq;
use super::MAX_ERROR_BYTES;
use super::RequestError;
use super::enforce;
use super::startup;
#[tokio::test(start_paused = true)]
async fn stalled_transport_fails_after_its_deadline() {
let task = tokio::spawn(enforce(
"termination",
Duration::ZERO,
pending::<Result<(), tonic::Status>>(),
));
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(61)).await;
assert!(matches!(
task.await.expect("deadline task"),
Err(RequestError::TimedOut(message))
if message == "gRPC code-mode host timed out waiting for termination response"
));
}
#[tokio::test(start_paused = true)]
async fn requested_runtime_duration_is_added_to_the_transport_deadline() {
let task = tokio::spawn(enforce(
"wait",
Duration::from_secs(120),
pending::<Result<(), tonic::Status>>(),
));
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(61)).await;
tokio::task::yield_now().await;
assert!(!task.is_finished());
tokio::time::advance(Duration::from_secs(120)).await;
assert!(matches!(
task.await.expect("deadline task"),
Err(RequestError::TimedOut(_))
));
}
#[tokio::test(start_paused = true)]
async fn default_execution_yield_and_grace_extend_the_outcome_deadline() {
let runtime_timeout =
Duration::from_millis(DEFAULT_EXEC_YIELD_TIME_MS).saturating_add(Duration::from_secs(1));
let task = tokio::spawn(enforce(
"execution outcome",
runtime_timeout,
pending::<Result<(), tonic::Status>>(),
));
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(70)).await;
tokio::task::yield_now().await;
assert!(!task.is_finished());
tokio::time::advance(Duration::from_secs(2)).await;
assert!(matches!(
task.await.expect("execution outcome deadline task"),
Err(RequestError::TimedOut(message))
if message == "gRPC code-mode host timed out waiting for execution outcome response"
));
}
#[tokio::test]
async fn transport_status_is_preserved() {
let result = enforce("wait", Duration::ZERO, async {
Err::<(), _>(tonic::Status::not_found("missing"))
})
.await;
match result {
Err(RequestError::Failed(error)) => {
assert_eq!(error.code(), tonic::Code::NotFound);
assert_eq!(error.message(), "missing");
}
_ => panic!("expected the original gRPC status"),
}
}
#[tokio::test]
async fn transport_status_messages_are_bounded_at_utf8_boundaries() {
let error = startup("session opening", async {
Err::<(), _>(tonic::Status::internal("🦀".repeat(MAX_ERROR_BYTES)))
})
.await
.expect_err("oversized gRPC status must fail");
assert!(error.len() <= MAX_ERROR_BYTES);
assert!(error.starts_with("gRPC code-mode session opening failed:"));
assert!(error.ends_with("..."));
}
#[tokio::test(start_paused = true)]
async fn stalled_channel_acquisition_times_out_and_remains_retryable() {
let channel = Arc::new(tokio::sync::OnceCell::new());
let stalled_channel = Arc::clone(&channel);
let stalled = tokio::spawn(async move {
startup("transport connection", async {
stalled_channel
.get_or_try_init(pending::<Result<usize, String>>)
.await
.copied()
})
.await
});
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(61)).await;
assert_eq!(
stalled.await.expect("channel connection task"),
Err("gRPC code-mode host timed out waiting for transport connection response".to_string())
);
assert_eq!(
startup("transport connection", async {
channel
.get_or_try_init(|| async { Ok::<_, String>(42usize) })
.await
.copied()
})
.await,
Ok(42)
);
}

View File

@@ -0,0 +1,317 @@
use std::collections::HashMap;
use std::panic::AssertUnwindSafe;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::PoisonError;
use std::sync::Weak;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering;
use std::time::Duration;
use codex_code_mode_protocol::CellId;
use codex_code_mode_protocol::CodeModeSession;
use codex_code_mode_protocol::CodeModeSessionCellExecutionLimits;
use codex_code_mode_protocol::CodeModeSessionDelegate;
use codex_code_mode_protocol::CodeModeSessionProvider;
use codex_code_mode_protocol::CodeModeSessionProviderFuture;
use codex_code_mode_protocol::CodeModeSessionResultFuture;
use codex_code_mode_protocol::ExecuteRequest;
use codex_code_mode_protocol::StartedCell;
use codex_code_mode_protocol::WaitOutcome;
use codex_code_mode_protocol::WaitRequest;
use codex_code_mode_protocol::grpc;
use codex_code_mode_protocol::grpc::code_mode_host_client::CodeModeHostClient;
use codex_code_mode_protocol::host::MAX_FRAME_BYTES;
use tokio::sync::watch;
use tokio_util::sync::CancellationToken;
use tokio_util::task::TaskTracker;
use tonic::transport::Channel;
use self::operations::WaitSlot;
use self::state::SessionState;
use self::transport::SharedTransport;
use crate::remote_session::ShutdownResultReceiver;
use crate::remote_session::wait_for_watch;
mod callbacks;
mod conversion;
mod deadline;
mod operations;
mod state;
mod transport;
type GrpcClient = CodeModeHostClient<Channel>;
/// Creates code-mode sessions over an HTTP/2 gRPC connection.
#[derive(Clone)]
pub struct GrpcCodeModeSessionProvider {
transport: Arc<SharedTransport>,
}
impl GrpcCodeModeSessionProvider {
/// Connects lazily to an `http://` gRPC endpoint.
pub fn new(endpoint: impl Into<String>) -> Self {
Self::from_transport(SharedTransport::new(endpoint.into()))
}
/// Uses an existing channel, including channels backed by custom transports.
pub fn with_channel(channel: Channel) -> Self {
Self::from_transport(SharedTransport::with_channel(channel))
}
fn from_transport(transport: SharedTransport) -> Self {
Self {
transport: Arc::new(transport),
}
}
async fn open_binding(
&self,
delegate: Arc<dyn CodeModeSessionDelegate>,
limits: CodeModeSessionCellExecutionLimits,
) -> Result<Arc<GrpcCodeModeSession>, String> {
let channel = deadline::startup("transport connection", self.transport.channel()).await?;
let mut client = grpc_client(channel);
let limits = grpc::SessionCellExecutionLimits {
max_yield_time_ms: limits.max_yield_time_ms,
max_heap_size_bytes: limits
.max_heap_size_bytes
.map(u64::try_from)
.transpose()
.map_err(|error| format!("invalid code-mode heap size limit: {error}"))?,
};
let cell_execution_limits = (limits.max_yield_time_ms.is_some()
|| limits.max_heap_size_bytes.is_some())
.then_some(limits);
let mut lease = deadline::startup(
"session opening",
client.open_session(grpc::OpenSessionRequest {
cell_execution_limits,
}),
)
.await?
.into_inner();
let first = deadline::startup("session lease opening", lease.message())
.await?
.ok_or_else(|| "gRPC code-mode session lease ended before opening".to_string())?;
let Some(grpc::session_event::Event::Opened(opened)) = first.event else {
return Err("gRPC code-mode session lease omitted its opening event".to_string());
};
validate_identifier(&opened.session_id, "session ID")?;
let inner = Arc::new(SessionInner {
id: opened.session_id,
client,
delegate,
runtime: tokio::runtime::Handle::current(),
state: Mutex::new(SessionState::default()),
wait_slots: Mutex::new(HashMap::new()),
shutdown_requested: AtomicBool::new(false),
shutdown_result: Mutex::new(None),
stopped: CancellationToken::new(),
stream_tasks: TaskTracker::new(),
_transport: Arc::clone(&self.transport),
});
let mut opening = OpeningSession {
inner: Some(Arc::clone(&inner)),
};
inner.spawn_session_events(lease);
inner.require_open()?;
opening.inner = None;
Ok(Arc::new(GrpcCodeModeSession { inner }))
}
}
impl CodeModeSessionProvider for GrpcCodeModeSessionProvider {
fn create_session<'a>(
&'a self,
delegate: Arc<dyn CodeModeSessionDelegate>,
) -> CodeModeSessionProviderFuture<'a> {
self.create_session_with_limits(delegate, CodeModeSessionCellExecutionLimits::default())
}
fn create_session_with_limits<'a>(
&'a self,
delegate: Arc<dyn CodeModeSessionDelegate>,
limits: CodeModeSessionCellExecutionLimits,
) -> CodeModeSessionProviderFuture<'a> {
Box::pin(async move {
self.open_binding(delegate, limits)
.await
.map(|session| session as _)
})
}
}
struct GrpcCodeModeSession {
inner: Arc<SessionInner>,
}
struct OpeningSession {
inner: Option<Arc<SessionInner>>,
}
impl Drop for OpeningSession {
fn drop(&mut self) {
let Some(inner) = self.inner.take() else {
return;
};
if tokio::runtime::Handle::try_current().is_ok() {
inner.request_shutdown();
} else {
inner.close_state(/*failure*/ None);
}
}
}
impl CodeModeSession for GrpcCodeModeSession {
fn execute<'a>(
&'a self,
request: ExecuteRequest,
) -> CodeModeSessionResultFuture<'a, StartedCell> {
Box::pin(self.inner.execute(request))
}
fn wait<'a>(&'a self, request: WaitRequest) -> CodeModeSessionResultFuture<'a, WaitOutcome> {
Box::pin(self.inner.wait(request))
}
fn terminate<'a>(&'a self, cell_id: CellId) -> CodeModeSessionResultFuture<'a, WaitOutcome> {
Box::pin(self.inner.terminate(cell_id))
}
fn shutdown<'a>(&'a self) -> CodeModeSessionResultFuture<'a, ()> {
Box::pin(wait_for_watch(self.inner.request_shutdown()))
}
}
impl Drop for GrpcCodeModeSession {
fn drop(&mut self) {
self.inner.request_shutdown();
}
}
pub(super) struct SessionInner {
pub(super) id: String,
pub(super) client: GrpcClient,
pub(super) delegate: Arc<dyn CodeModeSessionDelegate>,
runtime: tokio::runtime::Handle,
state: Mutex<SessionState>,
wait_slots: Mutex<HashMap<CellId, Weak<WaitSlot>>>,
shutdown_requested: AtomicBool,
shutdown_result: Mutex<Option<ShutdownResultReceiver>>,
pub(super) stopped: CancellationToken,
stream_tasks: TaskTracker,
_transport: Arc<SharedTransport>,
}
impl SessionInner {
pub(super) fn client(&self) -> GrpcClient {
self.client.clone()
}
pub(super) fn require_open(&self) -> Result<(), String> {
self.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.require_open()?;
if self.shutdown_requested.load(Ordering::Acquire) {
return Err("code mode session is shutting down".to_string());
}
Ok(())
}
pub(super) fn report_closed_cell(&self, cell_id: Option<CellId>) {
if let Some(cell_id) = cell_id {
self.wait_slots
.lock()
.unwrap_or_else(PoisonError::into_inner)
.remove(&cell_id);
let _ = std::panic::catch_unwind(AssertUnwindSafe(|| {
self.delegate.cell_closed(&cell_id);
}));
}
}
pub(super) fn fail(&self, reason: String) {
self.close_state(Some(reason));
}
fn close_state(&self, failure: Option<String>) {
let cells = self
.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.close(failure);
self.stopped.cancel();
self.stream_tasks.close();
for cell_id in cells {
self.report_closed_cell(Some(cell_id));
}
}
fn request_shutdown(self: &Arc<Self>) -> ShutdownResultReceiver {
self.shutdown_requested.store(true, Ordering::Release);
let mut result = self
.shutdown_result
.lock()
.unwrap_or_else(PoisonError::into_inner);
if let Some(receiver) = result.as_ref() {
return receiver.clone();
}
let (sender, receiver) = watch::channel(None);
*result = Some(receiver.clone());
let inner = Arc::clone(self);
self.runtime.spawn(async move {
let result = inner.drive_shutdown().await;
sender.send_replace(Some(result));
});
receiver
}
async fn drive_shutdown(&self) -> Result<(), String> {
let is_open = self
.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.require_open()
.is_ok();
let result = if is_open {
let mut client = self.client();
deadline::request(
self,
"session shutdown",
Duration::ZERO,
client.close_session(grpc::CloseSessionRequest {
session_id: self.id.clone(),
}),
)
.await
.map(|_| ())
} else {
Ok(())
};
self.close_state(/*failure*/ None);
self.stream_tasks.wait().await;
result
}
}
fn validate_identifier(value: &str, field: &str) -> Result<(), String> {
if value.is_empty() {
return Err(format!("gRPC code-mode host returned an empty {field}"));
}
if value.len() > grpc::MAX_IDENTIFIER_BYTES {
return Err(format!(
"gRPC code-mode host returned {field} exceeding {} bytes",
grpc::MAX_IDENTIFIER_BYTES
));
}
Ok(())
}
fn grpc_client(channel: Channel) -> GrpcClient {
CodeModeHostClient::new(channel)
.max_decoding_message_size(MAX_FRAME_BYTES)
.max_encoding_message_size(MAX_FRAME_BYTES)
}

View File

@@ -0,0 +1,391 @@
use std::sync::Arc;
use std::sync::PoisonError;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering;
use std::time::Duration;
use codex_code_mode_protocol::CellId;
use codex_code_mode_protocol::DEFAULT_EXEC_YIELD_TIME_MS;
use codex_code_mode_protocol::ExecuteRequest;
use codex_code_mode_protocol::RuntimeResponse;
use codex_code_mode_protocol::StartedCell;
use codex_code_mode_protocol::WaitOutcome;
use codex_code_mode_protocol::WaitRequest;
use codex_code_mode_protocol::grpc;
use tokio::sync::OwnedMutexGuard;
use tokio::sync::oneshot;
use tracing::debug;
use uuid::Uuid;
use super::SessionInner;
use super::conversion;
use super::deadline;
pub(super) struct WaitSlot {
lock: Arc<tokio::sync::Mutex<()>>,
active: AtomicBool,
}
struct ExecutionOwnership {
session: Arc<SessionInner>,
execution_id: String,
armed: bool,
}
impl Drop for ExecutionOwnership {
fn drop(&mut self) {
if !self.armed {
return;
}
let cell = self
.session
.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.remove_execution(&self.execution_id);
let Some(cell_id) = cell else {
return;
};
self.session.report_closed_cell(Some(cell_id.clone()));
if self.session.stopped.is_cancelled() {
return;
}
let session = Arc::clone(&self.session);
self.session.runtime.spawn(async move {
if let Err(error) = session.terminate(cell_id).await
&& !session.stopped.is_cancelled()
{
debug!("abandoned code-mode execution termination raced closure: {error}");
}
});
}
}
impl SessionInner {
pub(super) async fn execute(
self: &Arc<Self>,
request: ExecuteRequest,
) -> Result<StartedCell, String> {
self.require_open()?;
let execution_id = Uuid::new_v4().to_string();
let request = conversion::execute_request(&self.id, execution_id.clone(), request)?;
self.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.begin_execution(execution_id.clone())?;
let ownership = ExecutionOwnership {
session: Arc::clone(self),
execution_id,
armed: true,
};
let (started_tx, started_rx) = oneshot::channel();
let inner = Arc::clone(self);
self.stream_tasks.spawn(async move {
inner.drive_execution(request, ownership, started_tx).await;
});
started_rx
.await
.map_err(|_| "gRPC code-mode execution driver ended unexpectedly".to_string())?
}
async fn drive_execution(
self: Arc<Self>,
request: grpc::ExecuteRequest,
ownership: ExecutionOwnership,
started_tx: oneshot::Sender<Result<StartedCell, String>>,
) {
let runtime_timeout =
Duration::from_millis(request.yield_time_ms.unwrap_or(DEFAULT_EXEC_YIELD_TIME_MS))
.saturating_add(Duration::from_secs(1));
let opening = async {
let mut client = self.client();
let mut stream =
deadline::request(&self, "execution", Duration::ZERO, client.execute(request))
.await?
.into_inner();
let first = deadline::request(
&self,
"execution starting event",
Duration::ZERO,
stream.message(),
)
.await?
.ok_or_else(|| {
"gRPC code-mode execution ended before its starting event".to_string()
})?;
let Some(grpc::execute_event::Event::Started(started)) = first.event else {
return Err("gRPC code-mode execution omitted its starting event".to_string());
};
super::validate_identifier(&started.execution_id, "execution ID")?;
if started.execution_id != ownership.execution_id {
let error = format!(
"gRPC code-mode execution returned ID {} instead of {}",
started.execution_id, ownership.execution_id
);
self.fail(error.clone());
return Err(error);
}
let admission = self
.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.admit_execution(&ownership.execution_id, &started.cell_id);
if let Err(error) = admission {
self.fail(error.clone());
return Err(error);
}
Ok((CellId::new(started.cell_id), stream))
}
.await;
let (cell_id, stream) = match opening {
Ok(opening) => opening,
Err(error) => {
let _ = started_tx.send(Err(error));
return;
}
};
let (response_tx, response_rx) = oneshot::channel();
let mut claim = ownership;
let started = StartedCell::from_future(cell_id.clone(), async move {
let closure = claim
.session
.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.mark_execution_ready(&claim.execution_id);
match closure {
Ok(cell) => claim.session.report_closed_cell(cell),
Err(error) => {
claim.session.fail(error.clone());
return Err(error);
}
}
let response = response_rx
.await
.map_err(|_| "exec runtime ended unexpectedly".to_string())??;
claim.armed = false;
drop(claim);
Ok(response)
});
if started_tx.send(Ok(started)).is_err() {
return;
}
self.drive_execution_outcome(cell_id, stream, response_tx, runtime_timeout)
.await;
}
async fn drive_execution_outcome(
self: Arc<Self>,
cell_id: CellId,
mut stream: tonic::Streaming<grpc::ExecuteEvent>,
mut response_tx: oneshot::Sender<Result<RuntimeResponse, String>>,
runtime_timeout: Duration,
) {
let outcome = tokio::select! {
biased;
_ = response_tx.closed() => return,
outcome = deadline::request(
&self,
"execution outcome",
runtime_timeout,
stream.message(),
) => match outcome {
Ok(Some(grpc::ExecuteEvent {
event: Some(grpc::execute_event::Event::Outcome(outcome)),
})) => conversion::runtime_response(outcome),
Ok(Some(_)) => {
Err("gRPC code-mode execution returned an unexpected event".to_string())
}
Ok(None) => Err("gRPC code-mode execution omitted its initial outcome".to_string()),
Err(error) => Err(error),
},
};
let outcome = outcome.and_then(|response| {
if runtime_response_cell_id(&response) != &cell_id {
let error = format!(
"gRPC code-mode execution returned cell {} instead of {cell_id}",
runtime_response_cell_id(&response)
);
self.fail(error.clone());
Err(error)
} else {
Ok(response)
}
});
let _ = response_tx.send(outcome);
}
pub(super) async fn wait(
self: &Arc<Self>,
request: WaitRequest,
) -> Result<WaitOutcome, String> {
self.require_open()?;
let slot = {
let mut slots = self
.wait_slots
.lock()
.unwrap_or_else(PoisonError::into_inner);
slots.retain(|_, slot| slot.strong_count() != 0);
let slot = match slots
.get(&request.cell_id)
.and_then(std::sync::Weak::upgrade)
{
Some(slot) => slot,
None => {
let slot = Arc::new(WaitSlot {
lock: Arc::new(tokio::sync::Mutex::new(())),
active: AtomicBool::new(false),
});
slots.insert(request.cell_id.clone(), Arc::downgrade(&slot));
slot
}
};
if slot.active.swap(true, Ordering::AcqRel) {
return Err(format!(
"exec cell {} already has an active observer",
request.cell_id
));
}
slot
};
let lock = Arc::clone(&slot.lock);
let mut cancellation = WaitCancellation {
session: Arc::clone(self),
slot: Some(slot),
wait_id: None,
permit: None,
};
let permit = lock.lock_owned().await;
self.require_open()?;
let wait_id = Uuid::new_v4().to_string();
cancellation.wait_id = Some(wait_id.clone());
cancellation.permit = Some(permit);
let expected_cell_id = request.cell_id;
let runtime_timeout =
Duration::from_millis(request.yield_time_ms).saturating_add(Duration::from_secs(1));
let request = grpc::WaitRequest {
session_id: self.id.clone(),
cell_id: expected_cell_id.as_str().to_string(),
wait_id,
yield_time_ms: request.yield_time_ms,
};
let mut client = self.client();
let response = deadline::request(self, "wait", runtime_timeout, client.wait(request)).await;
cancellation.disarm();
self.prune_wait_slots();
let outcome = conversion::wait_outcome(response?.into_inner())?;
self.validate_wait_cell(&expected_cell_id, outcome)
}
pub(super) async fn terminate(&self, cell_id: CellId) -> Result<WaitOutcome, String> {
self.require_open()?;
let mut client = self.client();
let response = deadline::request(
self,
"termination",
Duration::ZERO,
client.terminate(grpc::TerminateRequest {
session_id: self.id.clone(),
cell_id: cell_id.as_str().to_string(),
}),
)
.await?
.into_inner();
let outcome = conversion::wait_outcome(response)?;
self.validate_wait_cell(&cell_id, outcome)
}
fn validate_wait_cell(
&self,
expected_cell_id: &CellId,
outcome: WaitOutcome,
) -> Result<WaitOutcome, String> {
let actual_cell_id = match &outcome {
WaitOutcome::LiveCell(response) | WaitOutcome::MissingCell(response) => {
runtime_response_cell_id(response)
}
};
if actual_cell_id != expected_cell_id {
let error = format!(
"gRPC code-mode host returned cell {actual_cell_id} instead of {expected_cell_id}"
);
self.fail(error.clone());
return Err(error);
}
Ok(outcome)
}
fn prune_wait_slots(&self) {
self.wait_slots
.lock()
.unwrap_or_else(PoisonError::into_inner)
.retain(|_, slot| slot.strong_count() != 0);
}
}
fn runtime_response_cell_id(response: &RuntimeResponse) -> &CellId {
match response {
RuntimeResponse::Yielded { cell_id, .. }
| RuntimeResponse::Terminated { cell_id, .. }
| RuntimeResponse::Result { cell_id, .. } => cell_id,
}
}
struct WaitCancellation {
session: Arc<SessionInner>,
slot: Option<Arc<WaitSlot>>,
wait_id: Option<String>,
permit: Option<OwnedMutexGuard<()>>,
}
impl WaitCancellation {
fn disarm(&mut self) {
self.wait_id = None;
if let Some(slot) = self.slot.take() {
slot.active.store(false, Ordering::Release);
}
self.permit = None;
}
}
impl Drop for WaitCancellation {
fn drop(&mut self) {
let slot = self.slot.take();
if let Some(slot) = slot.as_ref() {
slot.active.store(false, Ordering::Release);
}
let Some(wait_id) = self.wait_id.take() else {
return;
};
let permit = self.permit.take();
if self.session.stopped.is_cancelled() {
return;
}
let session = Arc::clone(&self.session);
self.session.runtime.spawn(async move {
let mut client = session.client();
let result = deadline::request(
&session,
"wait cancellation",
Duration::ZERO,
client.cancel_wait(grpc::CancelWaitRequest {
session_id: session.id.clone(),
wait_id,
}),
)
.await;
if let Err(error) = result
&& !session.stopped.is_cancelled()
{
session.fail(format!(
"failed to retire canceled gRPC code-mode wait: {error}"
));
}
drop(permit);
drop(slot);
session.prune_wait_slots();
});
}
}

View File

@@ -0,0 +1,149 @@
use std::collections::HashMap;
use codex_code_mode_protocol::CellId;
use codex_code_mode_protocol::grpc;
#[derive(Default)]
struct ExecutionRecord {
cell_id: Option<CellId>,
started: bool,
ready: bool,
closed: bool,
}
impl ExecutionRecord {
fn accept_cell(&mut self, cell_id: &str) -> Result<(), String> {
super::validate_identifier(cell_id, "cell ID")?;
if let Some(current) = self.cell_id.as_ref() {
if current.as_str() != cell_id {
return Err(format!(
"code-mode execution changed cell ID from {current} to {cell_id}"
));
}
} else {
self.cell_id = Some(CellId::new(cell_id.to_string()));
}
Ok(())
}
}
#[derive(Default)]
pub(super) struct SessionState {
executions: HashMap<String, ExecutionRecord>,
failure: Option<String>,
closed: bool,
}
impl SessionState {
pub(super) fn require_open(&self) -> Result<(), String> {
if self.closed {
return Err(self
.failure
.clone()
.unwrap_or_else(|| "code-mode gRPC session is closed".to_string()));
}
Ok(())
}
pub(super) fn begin_execution(&mut self, execution_id: String) -> Result<(), String> {
self.require_open()?;
if execution_id.is_empty() || self.executions.contains_key(&execution_id) {
return Err("code-mode execution ID was empty or reused".to_string());
}
self.executions
.insert(execution_id, ExecutionRecord::default());
Ok(())
}
pub(super) fn admit_execution(
&mut self,
execution_id: &str,
cell_id: &str,
) -> Result<(), String> {
self.require_open()?;
if self.executions.iter().any(|(id, execution)| {
id != execution_id
&& execution
.cell_id
.as_ref()
.is_some_and(|current| current.as_str() == cell_id)
}) {
return Err(format!("code-mode host reused active cell ID {cell_id}"));
}
let execution = self
.executions
.get_mut(execution_id)
.ok_or_else(|| format!("unknown code-mode execution {execution_id}"))?;
if execution.started {
return Err(format!("code-mode execution {execution_id} started twice"));
}
execution.accept_cell(cell_id)?;
execution.started = true;
Ok(())
}
pub(super) fn mark_execution_ready(
&mut self,
execution_id: &str,
) -> Result<Option<CellId>, String> {
self.require_open()?;
let execution = self
.executions
.get_mut(execution_id)
.ok_or_else(|| format!("unknown code-mode execution {execution_id}"))?;
if !execution.started || execution.ready {
return Err(format!(
"code-mode execution {execution_id} was not ready to be claimed"
));
}
execution.ready = true;
Ok(self.close_execution_if_ready(execution_id))
}
pub(super) fn close_cell(
&mut self,
closed: grpc::CellClosed,
) -> Result<Option<CellId>, String> {
self.require_open()?;
let Some(execution) = self.executions.get_mut(&closed.execution_id) else {
return Ok(None);
};
execution.accept_cell(&closed.cell_id)?;
if execution.closed {
return Err(format!(
"code-mode host returned an invalid closure for cell {}",
closed.cell_id
));
}
execution.closed = true;
Ok(self.close_execution_if_ready(&closed.execution_id))
}
pub(super) fn close(&mut self, failure: Option<String>) -> Vec<CellId> {
if self.closed {
return Vec::new();
}
self.closed = true;
self.failure = failure;
self.executions
.drain()
.filter_map(|(_, execution)| execution.cell_id)
.collect()
}
fn close_execution_if_ready(&mut self, execution_id: &str) -> Option<CellId> {
self.executions
.get(execution_id)
.is_some_and(|execution| execution.started && execution.ready && execution.closed)
.then(|| self.remove_execution(execution_id))
.flatten()
}
pub(super) fn remove_execution(&mut self, execution_id: &str) -> Option<CellId> {
self.executions.remove(execution_id)?.cell_id
}
}
#[cfg(test)]
#[path = "state_tests.rs"]
mod tests;

View File

@@ -0,0 +1,89 @@
use codex_code_mode_protocol::CellId;
use codex_code_mode_protocol::grpc;
use pretty_assertions::assert_eq;
use super::SessionState;
#[test]
fn cell_closure_waits_until_the_started_cell_is_claimed() {
let mut state = SessionState::default();
state
.begin_execution("execution".to_string())
.expect("register execution");
assert_eq!(
state
.close_cell(grpc::CellClosed {
execution_id: "execution".to_string(),
cell_id: "cell".to_string(),
final_tool_call_sequence: 0,
})
.expect("record early cell closure"),
None
);
state
.admit_execution("execution", "cell")
.expect("admit started cell");
assert_eq!(
state
.mark_execution_ready("execution")
.expect("claim started cell"),
Some(CellId::new("cell".to_string()))
);
}
#[test]
fn oversized_cell_ids_are_rejected_before_admission() {
let mut state = SessionState::default();
state
.begin_execution("execution".to_string())
.expect("register execution");
assert_eq!(
state.admit_execution("execution", &"x".repeat(grpc::MAX_IDENTIFIER_BYTES + 1)),
Err(format!(
"gRPC code-mode host returned cell ID exceeding {} bytes",
grpc::MAX_IDENTIFIER_BYTES
))
);
assert_eq!(state.remove_execution("execution"), None);
}
#[test]
fn abandonment_before_start_ignores_later_cell_closure() {
let mut state = SessionState::default();
state
.begin_execution("execution".to_string())
.expect("register execution");
assert_eq!(state.remove_execution("execution"), None);
assert_eq!(
state
.close_cell(grpc::CellClosed {
execution_id: "execution".to_string(),
cell_id: "cell".to_string(),
final_tool_call_sequence: 0,
})
.expect("ignore closure for abandoned execution"),
None
);
assert!(state.close(/*failure*/ None).is_empty());
}
#[test]
fn disconnect_returns_each_live_cell_once() {
let mut state = SessionState::default();
state
.begin_execution("execution".to_string())
.expect("register execution");
state
.admit_execution("execution", "cell")
.expect("admit execution");
assert_eq!(
state.close(Some("lease closed".to_string())),
vec![CellId::new("cell".to_string())]
);
assert!(state.close(/*failure*/ None).is_empty());
assert_eq!(state.require_open(), Err("lease closed".to_string()));
}

View File

@@ -0,0 +1,46 @@
use tonic::transport::Channel;
use tonic::transport::Endpoint;
pub(super) struct SharedTransport {
endpoint: TransportEndpoint,
channel: tokio::sync::OnceCell<Channel>,
}
enum TransportEndpoint {
Url(String),
Connected(Channel),
}
impl SharedTransport {
pub(super) fn new(endpoint: String) -> Self {
Self {
endpoint: TransportEndpoint::Url(endpoint),
channel: tokio::sync::OnceCell::new(),
}
}
pub(super) fn with_channel(channel: Channel) -> Self {
Self {
endpoint: TransportEndpoint::Connected(channel),
channel: tokio::sync::OnceCell::new(),
}
}
pub(super) async fn channel(&self) -> Result<Channel, String> {
self.channel
.get_or_try_init(|| async {
match &self.endpoint {
TransportEndpoint::Url(endpoint) => Endpoint::from_shared(endpoint.clone())
.map_err(|error| format!("invalid gRPC code-mode host endpoint: {error}"))?
.connect()
.await
.map_err(|error| {
format!("failed to connect to gRPC code-mode host: {error}")
}),
TransportEndpoint::Connected(channel) => Ok(channel.clone()),
}
})
.await
.cloned()
}
}

View File

@@ -1,6 +1,8 @@
mod grpc_session;
mod remote_session;
pub use codex_code_mode_protocol::*;
pub use grpc_session::GrpcCodeModeSessionProvider;
pub use remote_session::DisabledCodeModeSessionProvider;
pub use remote_session::ProcessOwnedCodeModeSession;
pub use remote_session::ProcessOwnedCodeModeSessionProvider;

View File

@@ -32,7 +32,7 @@ use crate::NoopCodeModeSessionDelegate;
mod connection;
type ShutdownResultReceiver = watch::Receiver<Option<Result<(), String>>>;
pub(crate) type ShutdownResultReceiver = watch::Receiver<Option<Result<(), String>>>;
/// Creates code-mode sessions backed by one lazily spawned process host.
pub struct ProcessOwnedCodeModeSessionProvider {
@@ -547,7 +547,7 @@ enum ShutdownAction {
Close(SessionBinding),
}
async fn wait_for_watch<T>(
pub(crate) async fn wait_for_watch<T>(
mut result_rx: watch::Receiver<Option<Result<T, String>>>,
) -> Result<T, String>
where