mirror of
https://github.com/openai/codex.git
synced 2026-09-06 15:29:32 +00:00
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:
committed by
copyberry
parent
b2543af02b
commit
1e557a554e
3
codex-rs/Cargo.lock
generated
3
codex-rs/Cargo.lock
generated
@@ -2564,10 +2564,13 @@ dependencies = [
|
||||
"codex-websocket-client",
|
||||
"futures",
|
||||
"pretty_assertions",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"tokio-util",
|
||||
"tonic",
|
||||
"tracing",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
|
||||
645
codex-rs/code-mode-host/tests/grpc.rs
Normal file
645
codex-rs/code-mode-host/tests/grpc.rs
Normal 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(())
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 }
|
||||
|
||||
70
codex-rs/code-mode/src/grpc_session/callbacks.rs
Normal file
70
codex-rs/code-mode/src/grpc_session/callbacks.rs
Normal 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(())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
134
codex-rs/code-mode/src/grpc_session/conversion.rs
Normal file
134
codex-rs/code-mode/src/grpc_session/conversion.rs
Normal 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;
|
||||
168
codex-rs/code-mode/src/grpc_session/conversion_tests.rs
Normal file
168
codex-rs/code-mode/src/grpc_session/conversion_tests.rs
Normal 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());
|
||||
}
|
||||
77
codex-rs/code-mode/src/grpc_session/deadline.rs
Normal file
77
codex-rs/code-mode/src/grpc_session/deadline.rs
Normal 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;
|
||||
131
codex-rs/code-mode/src/grpc_session/deadline_tests.rs
Normal file
131
codex-rs/code-mode/src/grpc_session/deadline_tests.rs
Normal 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)
|
||||
);
|
||||
}
|
||||
317
codex-rs/code-mode/src/grpc_session/mod.rs
Normal file
317
codex-rs/code-mode/src/grpc_session/mod.rs
Normal 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)
|
||||
}
|
||||
391
codex-rs/code-mode/src/grpc_session/operations.rs
Normal file
391
codex-rs/code-mode/src/grpc_session/operations.rs
Normal 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();
|
||||
});
|
||||
}
|
||||
}
|
||||
149
codex-rs/code-mode/src/grpc_session/state.rs
Normal file
149
codex-rs/code-mode/src/grpc_session/state.rs
Normal 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;
|
||||
89
codex-rs/code-mode/src/grpc_session/state_tests.rs
Normal file
89
codex-rs/code-mode/src/grpc_session/state_tests.rs
Normal 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()));
|
||||
}
|
||||
46
codex-rs/code-mode/src/grpc_session/transport.rs
Normal file
46
codex-rs/code-mode/src/grpc_session/transport.rs
Normal 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()
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user