mirror of
https://github.com/openai/codex.git
synced 2026-09-20 12:47:38 +00:00
code-mode: implement and verify stdio session host
This commit is contained in:
9
codex-rs/Cargo.lock
generated
9
codex-rs/Cargo.lock
generated
@@ -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"
|
||||
|
||||
@@ -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 }
|
||||
|
||||
126
codex-rs/code-mode-host/src/convert.rs
Normal file
126
codex-rs/code-mode-host/src/convert.rs
Normal 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,
|
||||
}),
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
146
codex-rs/code-mode-host/src/peer.rs
Normal file
146
codex-rs/code-mode-host/src/peer.rs
Normal 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;
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
120
codex-rs/code-mode-host/src/session.rs
Normal file
120
codex-rs/code-mode-host/src/session.rs
Normal 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),
|
||||
});
|
||||
}
|
||||
}
|
||||
684
codex-rs/code-mode-host/tests/host.rs
Normal file
684
codex-rs/code-mode-host/tests/host.rs
Normal 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;
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
Reference in New Issue
Block a user