use std::{future::Future, path::Path};
use futures::{
channel::oneshot,
future::{self, Either},
};
use crate::{
Agent, Client, ConnectionTo, DynamicHandlerGuard, JsonRpcRequest, Responder, SentRequest,
V2ConnectionTo,
jsonrpc::run::{NullRun, RunWithConnectionTo},
role::{HasPeer, acp::ProxySessionMessages},
schema::v2,
};
#[cfg(feature = "unstable_mcp_over_acp")]
use crate::{jsonrpc::run::ChainRun, mcp_server::McpServer};
async fn run_pending_session_setup<Counterpart, Run>(
connection: ConnectionTo<Counterpart>,
run: Run,
started_tx: oneshot::Sender<Result<(), crate::Error>>,
promotion_rx: oneshot::Receiver<()>,
) -> Result<(), crate::Error>
where
Counterpart: HasPeer<Agent>,
Run: RunWithConnectionTo<Counterpart>,
{
let mut run = Box::pin(run.run_with_connection_to(connection));
let first_poll =
future::poll_fn(|cx| std::task::Poll::Ready(std::future::Future::poll(run.as_mut(), cx)))
.await;
let readiness = match &first_poll {
std::task::Poll::Ready(result) => result.clone(),
std::task::Poll::Pending => Ok(()),
};
drop(started_tx.send(readiness));
match first_poll {
std::task::Poll::Ready(Ok(())) => {
let _ = promotion_rx.await;
Ok(())
}
std::task::Poll::Ready(Err(_)) => Ok(()),
std::task::Poll::Pending => match future::select(run, promotion_rx).await {
Either::Left((result, promotion_rx)) => {
result?;
let _ = promotion_rx.await;
Ok(())
}
Either::Right((Ok(()), run)) => run.await,
Either::Right((Err(_), _run)) => Ok(()),
},
}
}
fn send_session_setup<Counterpart, Request, Run>(
connection: V2ConnectionTo<Counterpart>,
request: Request,
dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
run: Run,
ordered: bool,
) -> SentRequest<Request::Response>
where
Counterpart: HasPeer<Agent>,
Request: JsonRpcRequest + 'static,
Request::Response: 'static,
Run: RunWithConnectionTo<Counterpart> + 'static,
{
let raw_connection = connection.raw_connection().clone();
if dynamic_handler_registrations.is_empty() {
drop(run);
if ordered {
raw_connection.send_ordered_request_to(Agent, request)
} else {
raw_connection.send_request_to(Agent, request)
}
} else {
let handlers_ready = raw_connection.dynamic_handler_barrier();
let (runner_started_tx, runner_started_rx) = oneshot::channel();
let (promotion_tx, promotion_rx) = oneshot::channel();
let runner_started = match raw_connection.spawn(run_pending_session_setup(
raw_connection.clone(),
run,
runner_started_tx,
promotion_rx,
)) {
Ok(()) => Either::Left(async move {
runner_started_rx.await.map_err(|error| {
crate::util::internal_error(format!(
"session setup runner stopped before its initial poll: {error}"
))
})?
}),
Err(error) => Either::Right(future::ready(Err(error))),
};
let readiness = async move {
future::try_join(handlers_ready, runner_started).await?;
Ok(())
};
let response_hook = move |_response: &Request::Response| {
promotion_tx.send(()).map_err(|()| {
crate::util::internal_error("session setup runner stopped before setup completed")
})?;
dynamic_handler_registrations
.into_iter()
.for_each(DynamicHandlerGuard::detach);
Ok(())
};
if ordered {
raw_connection.send_ordered_request_to_with_response_hook_after(
Agent,
request,
readiness,
response_hook,
)
} else {
raw_connection.send_request_to_with_response_hook_after(
Agent,
request,
readiness,
response_hook,
)
}
}
}
impl<Counterpart> V2ConnectionTo<Counterpart>
where
Counterpart: HasPeer<Agent>,
{
pub fn build_session(&self, cwd: impl AsRef<Path>) -> V2SessionBuilder<Counterpart> {
V2SessionBuilder::new(self, v2::NewSessionRequest::new(cwd.as_ref()))
}
pub fn build_session_cwd(&self) -> Result<V2SessionBuilder<Counterpart>, crate::Error> {
let cwd = std::env::current_dir().map_err(|error| {
crate::Error::internal_error().data(format!("cannot get current directory: {error}"))
})?;
Ok(self.build_session(cwd))
}
pub fn build_session_from(
&self,
request: v2::NewSessionRequest,
) -> V2SessionBuilder<Counterpart> {
V2SessionBuilder::new(self, request)
}
#[cfg(feature = "unstable_session_fork")]
pub fn fork_session(
&self,
session_id: impl Into<v2::SessionId>,
cwd: impl AsRef<Path>,
) -> V2ForkSessionBuilder<Counterpart> {
self.fork_session_from(v2::ForkSessionRequest::new(session_id, cwd.as_ref()))
}
#[cfg(feature = "unstable_session_fork")]
pub fn fork_session_from(
&self,
request: v2::ForkSessionRequest,
) -> V2ForkSessionBuilder<Counterpart> {
V2ForkSessionBuilder::new(self, request)
}
pub fn resume_session(
&self,
session_id: impl Into<v2::SessionId>,
cwd: impl AsRef<Path>,
) -> V2ResumeSessionBuilder<Counterpart> {
self.resume_session_from(v2::ResumeSessionRequest::new(session_id, cwd.as_ref()))
}
pub fn resume_session_from(
&self,
request: v2::ResumeSessionRequest,
) -> V2ResumeSessionBuilder<Counterpart> {
V2ResumeSessionBuilder::new(self, request)
}
}
#[must_use = "call `start_session` or `on_proxy_session_start` to send the `session/new` request"]
#[derive(Debug)]
pub struct V2SessionBuilder<Counterpart, Run = NullRun>
where
Counterpart: HasPeer<Agent>,
Run: RunWithConnectionTo<Counterpart>,
{
connection: V2ConnectionTo<Counterpart>,
request: v2::NewSessionRequest,
dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
run: Run,
}
impl<Counterpart> V2SessionBuilder<Counterpart, NullRun>
where
Counterpart: HasPeer<Agent>,
{
fn new(connection: &V2ConnectionTo<Counterpart>, request: v2::NewSessionRequest) -> Self {
Self {
connection: connection.clone(),
request,
dynamic_handler_registrations: Vec::new(),
run: NullRun,
}
}
}
impl<Counterpart, Run> V2SessionBuilder<Counterpart, Run>
where
Counterpart: HasPeer<Agent>,
Run: RunWithConnectionTo<Counterpart>,
{
#[cfg(feature = "unstable_mcp_over_acp")]
pub fn with_mcp_server<McpRun>(
mut self,
mcp_server: McpServer<Counterpart, McpRun>,
) -> Result<V2SessionBuilder<Counterpart, ChainRun<Run, McpRun>>, crate::Error>
where
McpRun: RunWithConnectionTo<Counterpart>,
{
let (handler, mcp_run) = mcp_server.into_v2_handler_and_runner();
self.dynamic_handler_registrations
.push(handler.into_dynamic_handler(&mut self.request.mcp_servers, &self.connection)?);
Ok(V2SessionBuilder {
connection: self.connection,
request: self.request,
dynamic_handler_registrations: self.dynamic_handler_registrations,
run: ChainRun::new(self.run, mcp_run),
})
}
fn send_new_session(self, ordered: bool) -> SentRequest<v2::NewSessionResponse>
where
Run: 'static,
{
let Self {
connection,
request,
dynamic_handler_registrations,
run,
} = self;
send_session_setup(
connection,
request,
dynamic_handler_registrations,
run,
ordered,
)
}
pub fn start_session(self) -> SentRequest<OpenedV2Session<Counterpart, v2::NewSessionResponse>>
where
Run: 'static,
{
let session_connection = self.connection.clone();
self.send_new_session(false).map(move |response| {
let session = V2Session {
session_id: response.session_id.clone(),
connection: session_connection,
};
Ok(OpenedV2Session { session, response })
})
}
pub fn on_proxy_session_start<F, Fut>(
self,
responder: Responder<v2::NewSessionResponse>,
op: F,
) -> Result<(), crate::Error>
where
Counterpart: HasPeer<Client>,
Run: 'static,
F: FnOnce(OpenedV2Session<Counterpart, v2::NewSessionResponse>) -> Fut + Send + 'static,
Fut: Future<Output = Result<(), crate::Error>> + Send,
{
let session_connection = self.connection.clone();
self.send_new_session(true)
.forward_cancellation_from(responder.cancellation())
.on_receiving_ok_result(responder, async move |response, responder| {
let session_id = response.session_id.clone();
let raw_connection = session_connection.raw_connection();
let route = match raw_connection.add_dynamic_handler(ProxySessionMessages::new(
crate::schema::v1::SessionId::from(session_id.0.clone()),
)) {
Ok(route) => route,
Err(error) => return responder.respond_with_error(error),
};
let opened = OpenedV2Session {
session: V2Session {
session_id,
connection: session_connection.clone(),
},
response: response.clone(),
};
responder.respond(response)?;
route.detach();
raw_connection.spawn(async move { op(opened).await })
})
}
}
#[cfg(feature = "unstable_session_fork")]
#[must_use = "call `start_session` or `on_proxy_session_start` to send the `session/fork` request"]
#[derive(Debug)]
pub struct V2ForkSessionBuilder<Counterpart, Run = NullRun>
where
Counterpart: HasPeer<Agent>,
Run: RunWithConnectionTo<Counterpart>,
{
connection: V2ConnectionTo<Counterpart>,
request: v2::ForkSessionRequest,
dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
run: Run,
}
#[cfg(feature = "unstable_session_fork")]
impl<Counterpart> V2ForkSessionBuilder<Counterpart, NullRun>
where
Counterpart: HasPeer<Agent>,
{
fn new(connection: &V2ConnectionTo<Counterpart>, request: v2::ForkSessionRequest) -> Self {
Self {
connection: connection.clone(),
request,
dynamic_handler_registrations: Vec::new(),
run: NullRun,
}
}
}
#[cfg(feature = "unstable_session_fork")]
impl<Counterpart, Run> V2ForkSessionBuilder<Counterpart, Run>
where
Counterpart: HasPeer<Agent>,
Run: RunWithConnectionTo<Counterpart>,
{
#[cfg(feature = "unstable_mcp_over_acp")]
pub fn with_mcp_server<McpRun>(
mut self,
mcp_server: McpServer<Counterpart, McpRun>,
) -> Result<V2ForkSessionBuilder<Counterpart, ChainRun<Run, McpRun>>, crate::Error>
where
McpRun: RunWithConnectionTo<Counterpart>,
{
let (handler, mcp_run) = mcp_server.into_v2_handler_and_runner();
self.dynamic_handler_registrations
.push(handler.into_dynamic_handler(&mut self.request.mcp_servers, &self.connection)?);
Ok(V2ForkSessionBuilder {
connection: self.connection,
request: self.request,
dynamic_handler_registrations: self.dynamic_handler_registrations,
run: ChainRun::new(self.run, mcp_run),
})
}
fn send_fork_session(self, ordered: bool) -> SentRequest<v2::ForkSessionResponse>
where
Run: 'static,
{
let Self {
connection,
request,
dynamic_handler_registrations,
run,
} = self;
send_session_setup(
connection,
request,
dynamic_handler_registrations,
run,
ordered,
)
}
pub fn start_session(self) -> SentRequest<OpenedV2Session<Counterpart, v2::ForkSessionResponse>>
where
Run: 'static,
{
let session_connection = self.connection.clone();
self.send_fork_session(false).map(move |response| {
let session = V2Session {
session_id: response.session_id.clone(),
connection: session_connection,
};
Ok(OpenedV2Session { session, response })
})
}
pub fn on_proxy_session_start<F, Fut>(
self,
responder: Responder<v2::ForkSessionResponse>,
op: F,
) -> Result<(), crate::Error>
where
Counterpart: HasPeer<Client>,
Run: 'static,
F: FnOnce(OpenedV2Session<Counterpart, v2::ForkSessionResponse>) -> Fut + Send + 'static,
Fut: Future<Output = Result<(), crate::Error>> + Send,
{
let session_connection = self.connection.clone();
self.send_fork_session(true)
.forward_cancellation_from(responder.cancellation())
.on_receiving_ok_result(responder, async move |response, responder| {
let session_id = response.session_id.clone();
let raw_connection = session_connection.raw_connection();
let route = match raw_connection.add_dynamic_handler(ProxySessionMessages::new(
crate::schema::v1::SessionId::from(session_id.0.clone()),
)) {
Ok(route) => route,
Err(error) => return responder.respond_with_error(error),
};
let opened = OpenedV2Session {
session: V2Session {
session_id,
connection: session_connection.clone(),
},
response: response.clone(),
};
responder.respond(response)?;
route.detach();
raw_connection.spawn(async move { op(opened).await })
})
}
}
#[must_use = "call `start_session` or `on_proxy_session_start` to send the `session/resume` request"]
#[derive(Debug)]
pub struct V2ResumeSessionBuilder<Counterpart, Run = NullRun>
where
Counterpart: HasPeer<Agent>,
Run: RunWithConnectionTo<Counterpart>,
{
connection: V2ConnectionTo<Counterpart>,
request: v2::ResumeSessionRequest,
dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
run: Run,
}
impl<Counterpart> V2ResumeSessionBuilder<Counterpart, NullRun>
where
Counterpart: HasPeer<Agent>,
{
fn new(connection: &V2ConnectionTo<Counterpart>, request: v2::ResumeSessionRequest) -> Self {
Self {
connection: connection.clone(),
request,
dynamic_handler_registrations: Vec::new(),
run: NullRun,
}
}
}
impl<Counterpart, Run> V2ResumeSessionBuilder<Counterpart, Run>
where
Counterpart: HasPeer<Agent>,
Run: RunWithConnectionTo<Counterpart>,
{
#[cfg(feature = "unstable_mcp_over_acp")]
pub fn with_mcp_server<McpRun>(
mut self,
mcp_server: McpServer<Counterpart, McpRun>,
) -> Result<V2ResumeSessionBuilder<Counterpart, ChainRun<Run, McpRun>>, crate::Error>
where
McpRun: RunWithConnectionTo<Counterpart>,
{
let (handler, mcp_run) = mcp_server.into_v2_handler_and_runner();
self.dynamic_handler_registrations
.push(handler.into_dynamic_handler(&mut self.request.mcp_servers, &self.connection)?);
Ok(V2ResumeSessionBuilder {
connection: self.connection,
request: self.request,
dynamic_handler_registrations: self.dynamic_handler_registrations,
run: ChainRun::new(self.run, mcp_run),
})
}
fn send_resume_session(self, ordered: bool) -> SentRequest<v2::ResumeSessionResponse>
where
Run: 'static,
{
let Self {
connection,
request,
dynamic_handler_registrations,
run,
} = self;
send_session_setup(
connection,
request,
dynamic_handler_registrations,
run,
ordered,
)
}
pub fn start_session(
self,
) -> SentRequest<OpenedV2Session<Counterpart, v2::ResumeSessionResponse>>
where
Run: 'static,
{
let session_id = self.request.session_id.clone();
let session_connection = self.connection.clone();
self.send_resume_session(false).map(move |response| {
let session = V2Session {
session_id,
connection: session_connection,
};
Ok(OpenedV2Session { session, response })
})
}
pub fn on_proxy_session_start<F, Fut>(
mut self,
responder: Responder<v2::ResumeSessionResponse>,
op: F,
) -> Result<(), crate::Error>
where
Counterpart: HasPeer<Client>,
Run: 'static,
F: FnOnce(OpenedV2Session<Counterpart, v2::ResumeSessionResponse>) -> Fut + Send + 'static,
Fut: Future<Output = Result<(), crate::Error>> + Send,
{
let session_id = self.request.session_id.clone();
let session_connection = self.connection.clone();
self.dynamic_handler_registrations.push(
session_connection
.raw_connection()
.add_dynamic_handler(ProxySessionMessages::new(
crate::schema::v1::SessionId::from(session_id.0.clone()),
))?,
);
self.send_resume_session(true)
.forward_cancellation_from(responder.cancellation())
.on_receiving_ok_result(responder, async move |response, responder| {
let opened = OpenedV2Session {
session: V2Session {
session_id,
connection: session_connection.clone(),
},
response: response.clone(),
};
responder.respond(response)?;
session_connection
.raw_connection()
.spawn(async move { op(opened).await })
})
}
}
#[derive(Debug)]
pub struct OpenedV2Session<Link, Response>
where
Link: HasPeer<Agent>,
{
session: V2Session<Link>,
response: Response,
}
impl<Link, Response> OpenedV2Session<Link, Response>
where
Link: HasPeer<Agent>,
{
pub fn session(&self) -> &V2Session<Link> {
&self.session
}
pub fn response(&self) -> &Response {
&self.response
}
pub fn into_parts(self) -> (V2Session<Link>, Response) {
(self.session, self.response)
}
pub fn into_session(self) -> V2Session<Link> {
self.session
}
}
#[derive(Debug, Clone)]
pub struct V2Session<Link>
where
Link: HasPeer<Agent>,
{
session_id: v2::SessionId,
connection: V2ConnectionTo<Link>,
}
impl<Link> V2Session<Link>
where
Link: HasPeer<Agent>,
{
pub fn session_id(&self) -> &v2::SessionId {
&self.session_id
}
pub fn connection(&self) -> &V2ConnectionTo<Link> {
&self.connection
}
pub fn send_prompt(&self, prompt: impl ToString) -> SentRequest<v2::PromptResponse> {
self.send_prompt_blocks(vec![prompt.to_string().into()])
}
pub fn send_prompt_blocks(
&self,
prompt: Vec<v2::ContentBlock>,
) -> SentRequest<v2::PromptResponse> {
self.connection.send_request_to(
Agent,
v2::PromptRequest::new(self.session_id.clone(), prompt),
)
}
pub fn cancel_active_work(&self) -> Result<(), crate::Error> {
self.connection.send_notification_to(
Agent,
v2::CancelSessionNotification::new(self.session_id.clone()),
)
}
pub fn set_config_option(
&self,
config_id: impl Into<v2::SessionConfigId>,
value: impl Into<v2::SessionConfigOptionValue>,
) -> SentRequest<v2::SetSessionConfigOptionResponse> {
self.connection.send_request_to(
Agent,
v2::SetSessionConfigOptionRequest::new(self.session_id.clone(), config_id, value),
)
}
pub fn close(&self) -> SentRequest<v2::CloseSessionResponse> {
self.connection
.send_request_to(Agent, v2::CloseSessionRequest::new(self.session_id.clone()))
}
}