code-mode: implement and verify stdio session host

This commit is contained in:
Channing Conger
2026-06-16 05:10:04 +00:00
parent 208395eaa1
commit 2720658f1d
9 changed files with 1427 additions and 39 deletions

9
codex-rs/Cargo.lock generated
View File

@@ -2510,6 +2510,15 @@ dependencies = [
[[package]]
name = "codex-code-mode-host"
version = "0.0.0"
dependencies = [
"codex-code-mode",
"codex-code-mode-protocol",
"codex-utils-cargo-bin",
"pretty_assertions",
"serde_json",
"tokio",
"tokio-util",
]
[[package]]
name = "codex-code-mode-protocol"

View File

@@ -10,3 +10,14 @@ path = "src/main.rs"
[lints]
workspace = true
[dependencies]
codex-code-mode = { workspace = true }
codex-code-mode-protocol = { workspace = true }
serde_json = { workspace = true }
tokio = { workspace = true, features = ["io-std", "io-util", "macros", "process", "rt-multi-thread", "sync", "time"] }
tokio-util = { workspace = true, features = ["rt"] }
[dev-dependencies]
codex-utils-cargo-bin = { workspace = true }
pretty_assertions = { workspace = true }

View File

@@ -0,0 +1,126 @@
use std::time::Duration;
use codex_code_mode::session_runtime as runtime;
use codex_code_mode_protocol::wire;
pub(super) fn create_cell_request(request: wire::CreateCellRequest) -> runtime::CreateCellRequest {
runtime::CreateCellRequest {
tool_call_id: request.tool_call_id,
enabled_tools: request
.enabled_tools
.into_iter()
.map(|definition| runtime::ToolDefinition {
name: definition.name,
tool_name: runtime::ToolName {
name: definition.tool_name.name,
namespace: definition.tool_name.namespace,
},
description: definition.description,
kind: tool_kind(definition.kind),
})
.collect(),
source: request.source,
}
}
pub(super) fn observe_mode(mode: wire::ObserveMode) -> runtime::ObserveMode {
match mode {
wire::ObserveMode::YieldAfter { duration_ms } => {
runtime::ObserveMode::YieldAfter(Duration::from_millis(duration_ms))
}
wire::ObserveMode::PendingFrontier => runtime::ObserveMode::PendingFrontier,
}
}
pub(super) fn runtime_cell_id(cell_id: &wire::CellId) -> runtime::CellId {
runtime::CellId::new(cell_id.as_str())
}
pub(super) fn wire_cell_id(cell_id: &runtime::CellId) -> wire::CellId {
wire::CellId::new(cell_id.as_str())
}
pub(super) fn cell_event(event: runtime::CellEvent) -> wire::CellEvent {
match event {
runtime::CellEvent::Yielded { content_items } => wire::CellEvent::Yielded {
content_items: content_items.into_iter().map(output_item).collect(),
},
runtime::CellEvent::Pending {
content_items,
pending_tool_call_ids,
} => wire::CellEvent::Pending {
content_items: content_items.into_iter().map(output_item).collect(),
pending_tool_call_ids,
},
runtime::CellEvent::Completed {
content_items,
error_text,
} => wire::CellEvent::Completed {
content_items: content_items.into_iter().map(output_item).collect(),
error_text,
},
runtime::CellEvent::Terminated { content_items } => wire::CellEvent::Terminated {
content_items: content_items.into_iter().map(output_item).collect(),
},
}
}
pub(super) fn runtime_error(error: runtime::Error) -> wire::Error {
match error {
runtime::Error::ShuttingDown => wire::Error::ShuttingDown,
runtime::Error::DuplicateCell(cell_id) => wire::Error::DuplicateCell {
cell_id: wire_cell_id(&cell_id),
},
runtime::Error::MissingCell(cell_id) => wire::Error::MissingCell {
cell_id: wire_cell_id(&cell_id),
},
runtime::Error::ClosedCell(cell_id) => wire::Error::ClosedCell {
cell_id: wire_cell_id(&cell_id),
},
runtime::Error::BusyObserver(cell_id) => wire::Error::BusyObserver {
cell_id: wire_cell_id(&cell_id),
},
runtime::Error::AlreadyTerminating(cell_id) => wire::Error::AlreadyTerminating {
cell_id: wire_cell_id(&cell_id),
},
runtime::Error::Runtime(message) => wire::Error::Runtime { message },
}
}
pub(super) fn nested_tool_call(invocation: runtime::NestedToolCall) -> wire::NestedToolCall {
wire::NestedToolCall {
cell_id: wire_cell_id(&invocation.cell_id),
runtime_tool_call_id: invocation.runtime_tool_call_id,
tool_name: wire::ToolName {
name: invocation.tool_name.name,
namespace: invocation.tool_name.namespace,
},
tool_kind: match invocation.tool_kind {
runtime::ToolKind::Function => wire::ToolKind::Function,
runtime::ToolKind::Freeform => wire::ToolKind::Freeform,
},
input: invocation.input,
}
}
fn tool_kind(kind: wire::ToolKind) -> runtime::ToolKind {
match kind {
wire::ToolKind::Function => runtime::ToolKind::Function,
wire::ToolKind::Freeform => runtime::ToolKind::Freeform,
}
}
fn output_item(item: runtime::OutputItem) -> wire::OutputItem {
match item {
runtime::OutputItem::Text { text } => wire::OutputItem::Text { text },
runtime::OutputItem::Image { image_url, detail } => wire::OutputItem::Image {
image_url,
detail: detail.map(|detail| match detail {
runtime::ImageDetail::Auto => wire::ImageDetail::Auto,
runtime::ImageDetail::Low => wire::ImageDetail::Low,
runtime::ImageDetail::High => wire::ImageDetail::High,
runtime::ImageDetail::Original => wire::ImageDetail::Original,
}),
},
}
}

View File

@@ -1 +1,322 @@
fn main() {}
use std::collections::HashMap;
use std::future::Future;
use std::sync::Arc;
use std::sync::Mutex as StdMutex;
use std::sync::PoisonError;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::AtomicU64;
use std::sync::atomic::Ordering;
use codex_code_mode_protocol::wire::ClientMessage;
use codex_code_mode_protocol::wire::Error;
use codex_code_mode_protocol::wire::HostMessage;
use codex_code_mode_protocol::wire::HostRequest;
use codex_code_mode_protocol::wire::HostResponse;
use codex_code_mode_protocol::wire::RequestId;
use codex_code_mode_protocol::wire::SessionId;
use codex_code_mode_protocol::wire::WireResult;
use codex_code_mode_protocol::wire::read_frame;
use codex_code_mode_protocol::wire::write_frame;
use tokio::sync::Mutex;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
mod convert;
mod peer;
mod session;
use peer::HostPeer;
use session::HostSession;
const IPC_CHANNEL_CAPACITY: usize = 128;
#[tokio::main(flavor = "current_thread")]
async fn main() {
if let Err(err) = run().await {
eprintln!("codex-code-mode-host failed: {err}");
std::process::exit(1);
}
}
async fn run() -> Result<(), String> {
let (outgoing_tx, mut outgoing_rx) = mpsc::channel(IPC_CHANNEL_CAPACITY);
let peer = Arc::new(HostPeer::new(outgoing_tx));
let state = Arc::new(HostState {
sessions: StdMutex::new(HashMap::new()),
pending_requests: Mutex::new(HashMap::new()),
next_session_id: AtomicU64::new(1),
closing: AtomicBool::new(false),
peer: Arc::clone(&peer),
});
let writer = tokio::spawn(async move {
let mut stdout = tokio::io::stdout();
while let Some(message) = outgoing_rx.recv().await {
write_frame(&mut stdout, &message)
.await
.map_err(|err| err.to_string())?;
}
Ok::<(), String>(())
});
let input_result = async {
let mut stdin = tokio::io::stdin();
while let Some(message) = read_frame::<_, ClientMessage>(&mut stdin)
.await
.map_err(|err| err.to_string())?
{
match message {
ClientMessage::Request { id, request } => {
state.spawn_request(id, request).await;
}
ClientMessage::CancelRequest { id } => state.cancel_request(id).await,
ClientMessage::CallbackResponse { id, response } => {
peer.complete(id, response).await;
}
}
}
Ok::<(), String>(())
}
.await;
peer.disconnect().await;
state.disconnect().await;
drop(state);
drop(peer);
let writer_result = writer.await.map_err(|err| err.to_string())?;
input_result?;
writer_result
}
struct HostState {
sessions: StdMutex<HashMap<SessionId, HostSession>>,
pending_requests: Mutex<HashMap<RequestId, CancellationToken>>,
next_session_id: AtomicU64,
closing: AtomicBool,
peer: Arc<HostPeer>,
}
impl HostState {
async fn spawn_request(self: &Arc<Self>, request_id: RequestId, request: HostRequest) {
if self.closing.load(Ordering::Acquire) {
self.send_response(request_id, Err(Error::ShuttingDown), None)
.await;
return;
}
let cancellation_token = CancellationToken::new();
let duplicate = {
let mut pending_requests = self.pending_requests.lock().await;
if let std::collections::hash_map::Entry::Vacant(e) = pending_requests.entry(request_id)
{
e.insert(cancellation_token.clone());
false
} else {
true
}
};
if duplicate {
self.send_response(
request_id,
Err(Error::InvalidRequest {
message: format!("duplicate request id {request_id}"),
}),
None,
)
.await;
return;
}
let state = Arc::clone(self);
spawn_critical_request_task(async move {
state
.handle_request(request_id, request, cancellation_token)
.await;
state.pending_requests.lock().await.remove(&request_id);
});
}
async fn cancel_request(&self, request_id: RequestId) {
if let Some(cancellation_token) = self.pending_requests.lock().await.get(&request_id) {
cancellation_token.cancel();
}
}
async fn handle_request(
&self,
request_id: RequestId,
request: HostRequest,
cancellation_token: CancellationToken,
) {
match request {
HostRequest::CreateSession => {
if cancellation_token.is_cancelled() || self.closing.load(Ordering::Acquire) {
return;
}
let session_id = self.next_session_id.fetch_add(1, Ordering::Relaxed);
let session = HostSession::new(session_id, Arc::clone(&self.peer));
let inserted = {
let mut sessions = self.sessions.lock().unwrap_or_else(PoisonError::into_inner);
if self.closing.load(Ordering::Acquire) {
false
} else {
sessions.insert(session_id, session);
true
}
};
if !inserted {
return;
}
self.send_response(
request_id,
Ok(HostResponse::SessionCreated { session_id }),
Some(&cancellation_token),
)
.await;
}
HostRequest::ShutdownSession { session_id } => {
let session = self
.sessions
.lock()
.unwrap_or_else(PoisonError::into_inner)
.remove(&session_id);
let result = match session {
Some(session) => session
.shutdown()
.await
.map(|()| HostResponse::SessionShutdown)
.map_err(convert::runtime_error),
None => Err(Error::MissingSession { session_id }),
};
self.send_response(request_id, result, Some(&cancellation_token))
.await;
}
HostRequest::CreateCell {
session_id,
request,
} => {
let Some(session) = self.session(session_id) else {
self.send_response(
request_id,
Err(Error::MissingSession { session_id }),
Some(&cancellation_token),
)
.await;
return;
};
let result = session
.create_cell(convert::create_cell_request(request))
.await
.map(|cell_id| HostResponse::CellCreated {
cell_id: convert::wire_cell_id(&cell_id),
})
.map_err(convert::runtime_error);
self.send_response(request_id, result, Some(&cancellation_token))
.await;
}
HostRequest::Observe {
session_id,
cell_id,
mode,
} => {
let Some(session) = self.session(session_id) else {
self.send_response(
request_id,
Err(Error::MissingSession { session_id }),
Some(&cancellation_token),
)
.await;
return;
};
let runtime_cell_id = convert::runtime_cell_id(&cell_id);
let result = session
.observe(&runtime_cell_id, convert::observe_mode(mode))
.await
.map(convert::cell_event)
.map(|event| HostResponse::Observed { event })
.map_err(convert::runtime_error);
self.send_response(request_id, result, Some(&cancellation_token))
.await;
}
HostRequest::Terminate {
session_id,
cell_id,
} => {
let Some(session) = self.session(session_id) else {
self.send_response(
request_id,
Err(Error::MissingSession { session_id }),
Some(&cancellation_token),
)
.await;
return;
};
let runtime_cell_id = convert::runtime_cell_id(&cell_id);
let result = session
.terminate(&runtime_cell_id)
.await
.map(convert::cell_event)
.map(|event| HostResponse::Observed { event })
.map_err(convert::runtime_error);
self.send_response(request_id, result, Some(&cancellation_token))
.await;
}
}
}
fn session(&self, session_id: SessionId) -> Option<HostSession> {
self.sessions
.lock()
.unwrap_or_else(PoisonError::into_inner)
.get(&session_id)
.cloned()
}
async fn send_response(
&self,
request_id: RequestId,
result: Result<HostResponse, Error>,
cancellation_token: Option<&CancellationToken>,
) {
if cancellation_token.is_some_and(CancellationToken::is_cancelled) {
return;
}
self.peer
.send(HostMessage::Response {
id: request_id,
result: WireResult::from_result(result),
})
.await;
}
async fn disconnect(&self) {
self.closing.store(true, Ordering::Release);
for cancellation_token in self.pending_requests.lock().await.values() {
cancellation_token.cancel();
}
let sessions = self
.sessions
.lock()
.unwrap_or_else(PoisonError::into_inner)
.drain()
.map(|(_, session)| session)
.collect::<Vec<_>>();
for session in sessions {
let _ = session.shutdown().await;
}
while !self.pending_requests.lock().await.is_empty() {
tokio::task::yield_now().await;
}
}
}
fn spawn_critical_request_task(future: impl Future<Output = ()> + Send + 'static) {
let task = tokio::spawn(future);
tokio::spawn(async move {
if let Err(err) = task.await
&& err.is_panic()
{
eprintln!("code-mode host request task panicked: {err}");
std::process::exit(1);
}
});
}

View File

@@ -0,0 +1,146 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use std::sync::atomic::Ordering;
use codex_code_mode_protocol::wire::CallbackId;
use codex_code_mode_protocol::wire::CallbackRequest;
use codex_code_mode_protocol::wire::CallbackResponse;
use codex_code_mode_protocol::wire::HostMessage;
use codex_code_mode_protocol::wire::SessionId;
use tokio::sync::Mutex;
use tokio::sync::mpsc;
use tokio::sync::oneshot;
use tokio_util::sync::CancellationToken;
pub(super) struct HostPeer {
outgoing_tx: mpsc::Sender<HostMessage>,
pending_callbacks: Mutex<HashMap<CallbackId, oneshot::Sender<CallbackResponse>>>,
next_callback_id: AtomicU64,
connection_closed: CancellationToken,
}
impl HostPeer {
pub(super) fn new(outgoing_tx: mpsc::Sender<HostMessage>) -> Self {
Self {
outgoing_tx,
pending_callbacks: Mutex::new(HashMap::new()),
next_callback_id: AtomicU64::new(1),
connection_closed: CancellationToken::new(),
}
}
pub(super) async fn send(&self, message: HostMessage) -> bool {
tokio::select! {
sent = self.outgoing_tx.send(message) => sent.is_ok(),
_ = self.connection_closed.cancelled() => false,
}
}
pub(super) fn send_nowait(self: &Arc<Self>, message: HostMessage) {
match self.outgoing_tx.try_send(message) {
Ok(()) | Err(mpsc::error::TrySendError::Closed(_)) => {}
Err(mpsc::error::TrySendError::Full(message)) => {
let peer = Arc::clone(self);
tokio::spawn(async move {
peer.send(message).await;
});
}
}
}
pub(super) async fn call(
self: &Arc<Self>,
session_id: SessionId,
request: CallbackRequest,
cancellation_token: CancellationToken,
) -> Result<CallbackResponse, String> {
if self.connection_closed.is_cancelled() {
return Err("code-mode client connection closed".to_string());
}
let id = self.next_callback_id.fetch_add(1, Ordering::Relaxed);
let (response_tx, response_rx) = oneshot::channel();
self.pending_callbacks.lock().await.insert(id, response_tx);
let mut pending_callback = PendingCallback::new(Arc::clone(self), id);
let request_sent = tokio::select! {
sent = self.outgoing_tx.send(HostMessage::CallbackRequest {
id,
session_id,
request,
}) => sent.is_ok(),
_ = cancellation_token.cancelled() => false,
_ = self.connection_closed.cancelled() => false,
};
if !request_sent {
self.pending_callbacks.lock().await.remove(&id);
pending_callback.disarm();
return Err(if cancellation_token.is_cancelled() {
"code mode callback cancelled".to_string()
} else {
"code-mode client connection closed".to_string()
});
}
tokio::select! {
response = response_rx => {
pending_callback.disarm();
response.map_err(|_| {
"code-mode client closed before returning callback output".to_string()
})
},
_ = cancellation_token.cancelled() => {
if self.pending_callbacks.lock().await.remove(&id).is_some() {
self.send(HostMessage::CancelCallback { id }).await;
}
pending_callback.disarm();
Err("code mode callback cancelled".to_string())
}
_ = self.connection_closed.cancelled() => {
self.pending_callbacks.lock().await.remove(&id);
pending_callback.disarm();
Err("code-mode client connection closed".to_string())
}
}
}
pub(super) async fn complete(&self, id: CallbackId, response: CallbackResponse) {
if let Some(response_tx) = self.pending_callbacks.lock().await.remove(&id) {
let _ = response_tx.send(response);
}
}
pub(super) async fn disconnect(&self) {
self.connection_closed.cancel();
self.pending_callbacks.lock().await.clear();
}
}
struct PendingCallback {
peer: Arc<HostPeer>,
id: Option<CallbackId>,
}
impl PendingCallback {
fn new(peer: Arc<HostPeer>, id: CallbackId) -> Self {
Self { peer, id: Some(id) }
}
fn disarm(&mut self) {
self.id = None;
}
}
impl Drop for PendingCallback {
fn drop(&mut self) {
let Some(id) = self.id.take() else {
return;
};
let peer = Arc::clone(&self.peer);
tokio::spawn(async move {
if peer.pending_callbacks.lock().await.remove(&id).is_some() {
peer.send(HostMessage::CancelCallback { id }).await;
}
});
}
}

View File

@@ -0,0 +1,120 @@
use std::sync::Arc;
use codex_code_mode::SessionRuntime;
use codex_code_mode::session_runtime as runtime;
use codex_code_mode_protocol::wire::CallbackRequest;
use codex_code_mode_protocol::wire::CallbackResponse;
use codex_code_mode_protocol::wire::HostMessage;
use codex_code_mode_protocol::wire::SessionId;
use serde_json::Value as JsonValue;
use tokio_util::sync::CancellationToken;
use crate::convert;
use crate::peer::HostPeer;
#[derive(Clone)]
pub(super) struct HostSession {
runtime: Arc<SessionRuntime<RemoteDelegate>>,
}
impl HostSession {
pub(super) fn new(session_id: SessionId, peer: Arc<HostPeer>) -> Self {
let delegate = Arc::new(RemoteDelegate { session_id, peer });
Self {
runtime: Arc::new(SessionRuntime::new(delegate)),
}
}
pub(super) async fn create_cell(
&self,
request: runtime::CreateCellRequest,
) -> Result<runtime::CellId, runtime::Error> {
self.runtime.create_cell(request).await
}
pub(super) async fn observe(
&self,
cell_id: &runtime::CellId,
mode: runtime::ObserveMode,
) -> Result<runtime::CellEvent, runtime::Error> {
self.runtime.observe(cell_id, mode).await
}
pub(super) async fn terminate(
&self,
cell_id: &runtime::CellId,
) -> Result<runtime::CellEvent, runtime::Error> {
self.runtime.terminate(cell_id).await
}
pub(super) async fn shutdown(&self) -> Result<(), runtime::Error> {
self.runtime.shutdown().await
}
}
struct RemoteDelegate {
session_id: SessionId,
peer: Arc<HostPeer>,
}
impl runtime::SessionRuntimeDelegate for RemoteDelegate {
async fn invoke_tool(
&self,
invocation: runtime::NestedToolCall,
cancellation_token: CancellationToken,
) -> Result<JsonValue, String> {
match self
.peer
.call(
self.session_id,
CallbackRequest::InvokeTool {
invocation: convert::nested_tool_call(invocation),
},
cancellation_token,
)
.await?
{
CallbackResponse::ToolResult { result } => Ok(result),
CallbackResponse::ToolError { error_text } => Err(error_text),
CallbackResponse::NotificationDelivered
| CallbackResponse::NotificationError { .. } => {
Err("code-mode client returned an invalid tool response".to_string())
}
}
}
async fn notify(
&self,
call_id: String,
cell_id: runtime::CellId,
text: String,
cancellation_token: CancellationToken,
) -> Result<(), String> {
match self
.peer
.call(
self.session_id,
CallbackRequest::Notify {
call_id,
cell_id: convert::wire_cell_id(&cell_id),
text,
},
cancellation_token,
)
.await?
{
CallbackResponse::NotificationDelivered => Ok(()),
CallbackResponse::NotificationError { error_text } => Err(error_text),
CallbackResponse::ToolResult { .. } | CallbackResponse::ToolError { .. } => {
Err("code-mode client returned an invalid notification response".to_string())
}
}
}
fn cell_closed(&self, cell_id: &runtime::CellId) {
self.peer.send_nowait(HostMessage::CellClosed {
session_id: self.session_id,
cell_id: convert::wire_cell_id(cell_id),
});
}
}

View File

@@ -0,0 +1,684 @@
#![allow(clippy::expect_used)]
use std::process::Stdio;
use std::time::Duration;
use codex_code_mode_protocol::wire::*;
use pretty_assertions::assert_eq;
use serde_json::json;
use tokio::process::Child;
use tokio::process::ChildStdin;
use tokio::process::Command;
use tokio::sync::mpsc;
struct HostHarness {
child: Child,
stdin: Option<ChildStdin>,
messages_rx: mpsc::UnboundedReceiver<Result<HostMessage, String>>,
}
impl HostHarness {
fn spawn() -> Self {
let mut child = Command::new(
codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("resolve host binary"),
)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::inherit())
.kill_on_drop(true)
.spawn()
.expect("spawn host");
let stdin = child.stdin.take().expect("host stdin");
let mut stdout = child.stdout.take().expect("host stdout");
let (messages_tx, messages_rx) = mpsc::unbounded_channel();
tokio::spawn(async move {
loop {
match read_frame(&mut stdout).await {
Ok(Some(message)) => {
if messages_tx.send(Ok(message)).is_err() {
break;
}
}
Ok(None) => break,
Err(err) => {
let _ = messages_tx.send(Err(err.to_string()));
break;
}
}
}
});
Self {
child,
stdin: Some(stdin),
messages_rx,
}
}
async fn send(&mut self, message: ClientMessage) {
write_frame(self.stdin.as_mut().expect("host stdin open"), &message)
.await
.expect("write host message");
}
async fn recv(&mut self) -> HostMessage {
tokio::time::timeout(Duration::from_secs(5), self.messages_rx.recv())
.await
.expect("host message timeout")
.expect("host output closed")
.expect("read host message")
}
async fn assert_no_message(&mut self, duration: Duration) {
assert!(
tokio::time::timeout(duration, self.messages_rx.recv())
.await
.is_err(),
"received an unexpected host message"
);
}
async fn create_session(&mut self) -> SessionId {
self.send(ClientMessage::Request {
id: 1,
request: HostRequest::CreateSession,
})
.await;
match self.recv().await {
HostMessage::Response {
id: 1,
result:
WireResult::Ok {
value: HostResponse::SessionCreated { session_id },
},
} => session_id,
message => panic!("unexpected create-session response: {message:?}"),
}
}
async fn shutdown_session(&mut self, session_id: SessionId, request_id: RequestId) {
self.send(ClientMessage::Request {
id: request_id,
request: HostRequest::ShutdownSession { session_id },
})
.await;
assert_eq!(
self.recv().await,
HostMessage::Response {
id: request_id,
result: WireResult::Ok {
value: HostResponse::SessionShutdown,
},
}
);
}
async fn finish(mut self) {
self.stdin.take();
let status = tokio::time::timeout(Duration::from_secs(5), self.child.wait())
.await
.expect("host exit timeout")
.expect("wait for host");
assert!(status.success(), "host exited with {status}");
}
}
fn create_cell_request(source: &str) -> CreateCellRequest {
CreateCellRequest {
tool_call_id: "call-1".to_string(),
enabled_tools: Vec::new(),
source: source.to_string(),
}
}
async fn create_cell(
host: &mut HostHarness,
session_id: SessionId,
request_id: RequestId,
request: CreateCellRequest,
) -> CellId {
host.send(ClientMessage::Request {
id: request_id,
request: HostRequest::CreateCell {
session_id,
request,
},
})
.await;
match host.recv().await {
HostMessage::Response {
id,
result:
WireResult::Ok {
value: HostResponse::CellCreated { cell_id },
},
} if id == request_id => cell_id,
message => panic!("unexpected create-cell response: {message:?}"),
}
}
async fn observe_cell(
host: &mut HostHarness,
session_id: SessionId,
request_id: RequestId,
cell_id: CellId,
mode: ObserveMode,
) {
host.send(ClientMessage::Request {
id: request_id,
request: HostRequest::Observe {
session_id,
cell_id,
mode,
},
})
.await;
}
async fn recv_observation(host: &mut HostHarness, request_id: RequestId) -> CellEvent {
match host.recv().await {
HostMessage::Response {
id,
result:
WireResult::Ok {
value: HostResponse::Observed { event },
},
} if id == request_id => event,
message => panic!("unexpected observation response: {message:?}"),
}
}
async fn recv_terminal_observation(
host: &mut HostHarness,
session_id: SessionId,
request_id: RequestId,
cell_id: &CellId,
) -> CellEvent {
let mut event = None;
let mut cell_closed = false;
for _ in 0..2 {
match host.recv().await {
HostMessage::Response {
id,
result:
WireResult::Ok {
value: HostResponse::Observed { event: response },
},
} if id == request_id => event = Some(response),
HostMessage::CellClosed {
session_id: closed_session_id,
cell_id: closed_cell_id,
} => {
assert_eq!(closed_session_id, session_id);
assert_eq!(&closed_cell_id, cell_id);
cell_closed = true;
}
message => panic!("unexpected terminal observation message: {message:?}"),
}
}
assert!(cell_closed);
event.expect("terminal observation response")
}
#[tokio::test]
async fn yields_resumes_and_closes_over_stdio() {
let mut host = HostHarness::spawn();
let session_id = host.create_session().await;
let cell_id = create_cell(
&mut host,
session_id,
2,
create_cell_request(concat!(
"await new Promise(resolve => setTimeout(resolve, 100));",
r#"text("before"); yield_control(); text("after");"#,
)),
)
.await;
observe_cell(
&mut host,
session_id,
3,
cell_id.clone(),
ObserveMode::YieldAfter {
duration_ms: 60_000,
},
)
.await;
let mut yielded = None;
let mut saw_cell_closed = false;
for _ in 0..2 {
match host.recv().await {
HostMessage::Response {
id: 3,
result:
WireResult::Ok {
value: HostResponse::Observed { event },
},
} => yielded = Some(event),
HostMessage::CellClosed {
session_id: closed_session_id,
cell_id: closed_cell_id,
} => {
assert_eq!(closed_session_id, session_id);
assert_eq!(closed_cell_id, cell_id);
saw_cell_closed = true;
}
message => panic!("unexpected initial observation message: {message:?}"),
}
}
assert_eq!(
yielded,
Some(CellEvent::Yielded {
content_items: vec![OutputItem::Text {
text: "before".to_string(),
}],
})
);
assert!(saw_cell_closed);
observe_cell(
&mut host,
session_id,
4,
cell_id.clone(),
ObserveMode::YieldAfter { duration_ms: 1 },
)
.await;
assert_eq!(
recv_observation(&mut host, 4).await,
CellEvent::Completed {
content_items: vec![OutputItem::Text {
text: "after".to_string(),
}],
error_text: None,
}
);
host.shutdown_session(session_id, 5).await;
host.finish().await;
}
#[tokio::test]
async fn termination_resolves_an_active_observer() {
let mut host = HostHarness::spawn();
let session_id = host.create_session().await;
let cell_id = create_cell(
&mut host,
session_id,
2,
create_cell_request("await new Promise(() => {});"),
)
.await;
observe_cell(
&mut host,
session_id,
3,
cell_id.clone(),
ObserveMode::YieldAfter {
duration_ms: 60_000,
},
)
.await;
host.assert_no_message(Duration::from_millis(100)).await;
host.send(ClientMessage::Request {
id: 4,
request: HostRequest::Terminate {
session_id,
cell_id: cell_id.clone(),
},
})
.await;
let mut saw_observation_response = false;
let mut saw_termination_response = false;
let mut saw_cell_closed = false;
for _ in 0..3 {
match host.recv().await {
HostMessage::Response {
id: 3,
result:
WireResult::Ok {
value:
HostResponse::Observed {
event: CellEvent::Terminated { content_items },
},
},
} => {
assert!(content_items.is_empty());
saw_observation_response = true;
}
HostMessage::Response {
id: 4,
result:
WireResult::Ok {
value:
HostResponse::Observed {
event: CellEvent::Terminated { content_items },
},
},
} => {
assert!(content_items.is_empty());
saw_termination_response = true;
}
HostMessage::CellClosed {
session_id: closed_session_id,
cell_id: closed_cell_id,
} => {
assert_eq!(closed_session_id, session_id);
assert_eq!(closed_cell_id, cell_id);
saw_cell_closed = true;
}
message => panic!("unexpected termination message: {message:?}"),
}
}
assert!(saw_observation_response);
assert!(saw_termination_response);
assert!(saw_cell_closed);
host.shutdown_session(session_id, 5).await;
host.finish().await;
}
#[tokio::test]
async fn forwards_tool_and_notification_callbacks() {
let mut host = HostHarness::spawn();
let session_id = host.create_session().await;
let mut request = create_cell_request(
r#"
notify("note");
const value = await tools.echo({ value: "input" });
text(value.value);
"#,
);
request.enabled_tools.push(ToolDefinition {
name: "echo".to_string(),
tool_name: ToolName {
name: "echo".to_string(),
namespace: None,
},
description: String::new(),
kind: ToolKind::Function,
});
let cell_id = create_cell(&mut host, session_id, 2, request).await;
observe_cell(
&mut host,
session_id,
3,
cell_id.clone(),
ObserveMode::YieldAfter {
duration_ms: 60_000,
},
)
.await;
for _ in 0..2 {
match host.recv().await {
HostMessage::CallbackRequest {
id,
session_id: callback_session_id,
request: CallbackRequest::Notify { text, .. },
} => {
assert_eq!(callback_session_id, session_id);
assert_eq!(text, "note");
host.send(ClientMessage::CallbackResponse {
id,
response: CallbackResponse::NotificationDelivered,
})
.await;
}
HostMessage::CallbackRequest {
id,
session_id: callback_session_id,
request: CallbackRequest::InvokeTool { invocation },
} => {
assert_eq!(callback_session_id, session_id);
assert_eq!(invocation.input, Some(json!({"value": "input"})));
host.send(ClientMessage::CallbackResponse {
id,
response: CallbackResponse::ToolResult {
result: json!({"value": "output"}),
},
})
.await;
}
message => panic!("unexpected callback message: {message:?}"),
}
}
assert_eq!(
recv_terminal_observation(&mut host, session_id, 3, &cell_id).await,
CellEvent::Completed {
content_items: vec![OutputItem::Text {
text: "output".to_string(),
}],
error_text: None,
}
);
host.shutdown_session(session_id, 4).await;
host.finish().await;
}
#[tokio::test]
async fn pending_frontier_rejects_a_second_observer_and_termination_preempts_the_first() {
let mut host = HostHarness::spawn();
let session_id = host.create_session().await;
let cell_id = create_cell(
&mut host,
session_id,
2,
create_cell_request("await new Promise(() => {});"),
)
.await;
observe_cell(
&mut host,
session_id,
3,
cell_id.clone(),
ObserveMode::PendingFrontier,
)
.await;
assert_eq!(
recv_observation(&mut host, 3).await,
CellEvent::Pending {
content_items: Vec::new(),
pending_tool_call_ids: Vec::new(),
}
);
host.send(ClientMessage::Request {
id: 4,
request: HostRequest::Observe {
session_id,
cell_id: cell_id.clone(),
mode: ObserveMode::YieldAfter {
duration_ms: 60_000,
},
},
})
.await;
host.assert_no_message(Duration::from_millis(100)).await;
host.send(ClientMessage::Request {
id: 5,
request: HostRequest::Observe {
session_id,
cell_id: cell_id.clone(),
mode: ObserveMode::YieldAfter {
duration_ms: 60_000,
},
},
})
.await;
assert_eq!(
host.recv().await,
HostMessage::Response {
id: 5,
result: WireResult::Err {
error: Error::BusyObserver {
cell_id: cell_id.clone(),
},
},
}
);
host.send(ClientMessage::Request {
id: 6,
request: HostRequest::Terminate {
session_id,
cell_id: cell_id.clone(),
},
})
.await;
let terminated = HostResponse::Observed {
event: CellEvent::Terminated {
content_items: Vec::new(),
},
};
let mut response_ids = Vec::new();
let mut saw_cell_closed = false;
for _ in 0..3 {
match host.recv().await {
HostMessage::Response {
id,
result: WireResult::Ok { value },
} => {
assert_eq!(value, terminated);
response_ids.push(id);
}
HostMessage::CellClosed {
session_id: closed_session_id,
cell_id: closed_cell_id,
} => {
assert_eq!(closed_session_id, session_id);
assert_eq!(closed_cell_id, cell_id);
saw_cell_closed = true;
}
message => panic!("unexpected termination response: {message:?}"),
}
}
response_ids.sort_unstable();
assert_eq!(response_ids, vec![4, 6]);
assert!(saw_cell_closed);
host.shutdown_session(session_id, 7).await;
host.finish().await;
}
#[tokio::test]
async fn shutdown_cancels_callbacks_before_acknowledging_the_session() {
let mut host = HostHarness::spawn();
let session_id = host.create_session().await;
let mut request = create_cell_request("await tools.echo({});");
request.enabled_tools.push(ToolDefinition {
name: "echo".to_string(),
tool_name: ToolName {
name: "echo".to_string(),
namespace: None,
},
description: String::new(),
kind: ToolKind::Function,
});
let cell_id = create_cell(&mut host, session_id, 2, request).await;
observe_cell(
&mut host,
session_id,
3,
cell_id.clone(),
ObserveMode::YieldAfter {
duration_ms: 60_000,
},
)
.await;
let callback_id = match host.recv().await {
HostMessage::CallbackRequest {
id,
request: CallbackRequest::InvokeTool { .. },
..
} => id,
message => panic!("unexpected callback message: {message:?}"),
};
host.send(ClientMessage::Request {
id: 4,
request: HostRequest::ShutdownSession { session_id },
})
.await;
let mut saw_callback_cancellation = false;
let mut saw_terminated_response = false;
let mut saw_shutdown_response = false;
let mut saw_cell_closed = false;
for _ in 0..4 {
match host.recv().await {
HostMessage::CancelCallback { id } => {
assert_eq!(id, callback_id);
saw_callback_cancellation = true;
}
HostMessage::Response {
id: 3,
result:
WireResult::Ok {
value:
HostResponse::Observed {
event: CellEvent::Terminated { content_items },
},
},
} => {
assert!(content_items.is_empty());
saw_terminated_response = true;
}
HostMessage::Response {
id: 4,
result:
WireResult::Ok {
value: HostResponse::SessionShutdown,
},
} => saw_shutdown_response = true,
HostMessage::CellClosed {
session_id: closed_session_id,
cell_id: closed_cell_id,
} => {
assert_eq!(closed_session_id, session_id);
assert_eq!(closed_cell_id, cell_id);
saw_cell_closed = true;
}
message => panic!("unexpected shutdown message: {message:?}"),
}
}
assert!(saw_callback_cancellation);
assert!(saw_terminated_response);
assert!(saw_shutdown_response);
assert!(saw_cell_closed);
host.finish().await;
}
#[tokio::test]
async fn disconnect_exits_with_a_pending_callback() {
let mut host = HostHarness::spawn();
let session_id = host.create_session().await;
let mut request = create_cell_request("await tools.echo({});");
request.enabled_tools.push(ToolDefinition {
name: "echo".to_string(),
tool_name: ToolName {
name: "echo".to_string(),
namespace: None,
},
description: String::new(),
kind: ToolKind::Function,
});
let cell_id = create_cell(&mut host, session_id, 2, request).await;
observe_cell(
&mut host,
session_id,
3,
cell_id,
ObserveMode::YieldAfter {
duration_ms: 60_000,
},
)
.await;
assert!(matches!(
host.recv().await,
HostMessage::CallbackRequest {
request: CallbackRequest::InvokeTool { .. },
..
}
));
host.finish().await;
}

View File

@@ -30,7 +30,6 @@ pub(crate) use self::types::CellToolCall;
pub(crate) use self::types::CompletionCommit;
use self::types::CompletionDelivery;
use self::types::ObservationDelivery;
use crate::runtime::PendingRuntimeMode;
use crate::runtime::RuntimeCommand;
use crate::runtime::RuntimeControlCommand;
use crate::runtime::RuntimeEvent;
@@ -52,12 +51,8 @@ impl CellActor {
) -> Result<(CellHandle, impl Future<Output = ()> + Send + 'static), String> {
let (event_tx, event_rx) = mpsc::unbounded_channel();
let (command_tx, command_rx) = mpsc::unbounded_channel();
let (runtime_tx, runtime_control_tx, runtime_terminate_handle) = spawn_runtime(
stored_values,
runtime_request(request),
event_tx,
PendingRuntimeMode::PauseUntilResumed,
)?;
let (runtime_tx, runtime_control_tx, runtime_terminate_handle) =
spawn_runtime(stored_values, runtime_request(request), event_tx)?;
let handle = CellHandle::new(command_tx, Arc::clone(&cell_state));
let task = run_cell(
host,

View File

@@ -29,13 +29,6 @@ pub(crate) enum RuntimeCommand {
Terminate,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub(crate) enum PendingRuntimeMode {
#[cfg(test)]
Continue,
PauseUntilResumed,
}
#[derive(Debug)]
pub(crate) enum RuntimeControlCommand {
Continue,
@@ -69,7 +62,6 @@ pub(crate) fn spawn_runtime(
stored_values: HashMap<String, JsonValue>,
request: CreateCellRequest,
event_tx: mpsc::UnboundedSender<RuntimeEvent>,
pending_mode: PendingRuntimeMode,
) -> Result<
(
std_mpsc::Sender<RuntimeCommand>,
@@ -102,7 +94,6 @@ pub(crate) fn spawn_runtime(
event_tx,
command_rx,
control_rx,
pending_mode,
isolate_handle_tx,
runtime_command_tx,
);
@@ -165,7 +156,6 @@ fn run_runtime(
event_tx: mpsc::UnboundedSender<RuntimeEvent>,
command_rx: std_mpsc::Receiver<RuntimeCommand>,
control_rx: std_mpsc::Receiver<RuntimeControlCommand>,
pending_mode: PendingRuntimeMode,
isolate_handle_tx: std_mpsc::SyncSender<v8::IsolateHandle>,
runtime_command_tx: std_mpsc::Sender<RuntimeCommand>,
) {
@@ -221,9 +211,7 @@ fn run_runtime(
}
let mut pending_promise = pending_promise;
while let Some(command) =
next_runtime_command(&event_tx, &command_rx, &control_rx, pending_mode)
{
while let Some(command) = next_runtime_command(&event_tx, &command_rx, &control_rx) {
match command {
RuntimeCommand::Terminate => break,
RuntimeCommand::ToolResponse { id, result } => {
@@ -276,7 +264,6 @@ fn next_runtime_command(
event_tx: &mpsc::UnboundedSender<RuntimeEvent>,
command_rx: &std_mpsc::Receiver<RuntimeCommand>,
control_rx: &std_mpsc::Receiver<RuntimeControlCommand>,
pending_mode: PendingRuntimeMode,
) -> Option<RuntimeCommand> {
loop {
match command_rx.try_recv() {
@@ -286,14 +273,10 @@ fn next_runtime_command(
}
let _ = event_tx.send(RuntimeEvent::Pending);
match pending_mode {
#[cfg(test)]
PendingRuntimeMode::Continue => return command_rx.recv().ok(),
PendingRuntimeMode::PauseUntilResumed => match control_rx.recv().ok()? {
RuntimeControlCommand::Continue => return command_rx.recv().ok(),
RuntimeControlCommand::Resume => continue,
RuntimeControlCommand::Terminate => return Some(RuntimeCommand::Terminate),
},
match control_rx.recv().ok()? {
RuntimeControlCommand::Continue => return command_rx.recv().ok(),
RuntimeControlCommand::Resume => continue,
RuntimeControlCommand::Terminate => return Some(RuntimeCommand::Terminate),
}
}
}
@@ -331,7 +314,6 @@ mod tests {
use tokio::sync::mpsc;
use super::CreateCellRequest;
use super::PendingRuntimeMode;
use super::RuntimeCommand;
use super::RuntimeControlCommand;
use super::RuntimeEvent;
@@ -349,13 +331,8 @@ mod tests {
#[tokio::test]
async fn terminate_execution_stops_cpu_bound_module() {
let (event_tx, mut event_rx) = mpsc::unbounded_channel();
let (_runtime_tx, _runtime_control_tx, runtime_terminate_handle) = spawn_runtime(
HashMap::new(),
execute_request("while (true) {}"),
event_tx,
PendingRuntimeMode::Continue,
)
.unwrap();
let (_runtime_tx, _runtime_control_tx, runtime_terminate_handle) =
spawn_runtime(HashMap::new(), execute_request("while (true) {}"), event_tx).unwrap();
let started_event = tokio::time::timeout(Duration::from_secs(1), event_rx.recv())
.await
@@ -395,7 +372,6 @@ await new Promise(() => {});
"#,
),
event_tx,
PendingRuntimeMode::PauseUntilResumed,
)
.unwrap();