1use std::{future::Future, path::Path};
2
3use futures::{
4 channel::oneshot,
5 future::{self, Either},
6};
7
8use crate::{
9 Agent, Client, ConnectionTo, DynamicHandlerGuard, JsonRpcRequest, Responder, SentRequest,
10 V2ConnectionTo,
11 jsonrpc::run::{NullRun, RunWithConnectionTo},
12 role::{HasPeer, acp::ProxySessionMessages},
13 schema::v2,
14};
15
16#[cfg(feature = "unstable_mcp_over_acp")]
17use crate::{jsonrpc::run::ChainRun, mcp_server::McpServer};
18
19use super::{
20 SessionCleanup, close_session_registrations, drive_session_cleanup,
21 run_attached_session_runner, session_cleanup, session_runner_connection, wait_session_cleanup,
22};
23
24async fn run_pending_session_setup<Counterpart, Run>(
25 connection: ConnectionTo<Counterpart>,
26 run: Run,
27 started_tx: oneshot::Sender<Result<(), crate::Error>>,
28 promotion_rx: oneshot::Receiver<()>,
29 cleanup: SessionCleanup,
30) -> Result<(), crate::Error>
31where
32 Counterpart: HasPeer<Agent>,
33 Run: RunWithConnectionTo<Counterpart>,
34{
35 let (connection, error_scope) = session_runner_connection(connection, &cleanup);
36 let mut run = Box::pin(run.run_with_connection_to(connection));
37 let first_poll =
38 future::poll_fn(|cx| std::task::Poll::Ready(std::future::Future::poll(run.as_mut(), cx)))
39 .await;
40 if let Some(error) = error_scope.error() {
43 close_session_registrations(&cleanup);
44 if first_poll.is_pending() {
45 drop(drive_session_cleanup(run, cleanup).await);
46 } else {
47 wait_session_cleanup(cleanup).await;
48 }
49 drop(started_tx.send(Err(error)));
50 return Ok(());
51 }
52 let readiness = match &first_poll {
53 std::task::Poll::Ready(result) => result.clone(),
54 std::task::Poll::Pending => Ok(()),
55 };
56 drop(started_tx.send(readiness));
57
58 match first_poll {
59 std::task::Poll::Ready(Ok(())) => {
60 if promotion_rx.await.is_err() {
61 close_session_registrations(&cleanup);
62 wait_session_cleanup(cleanup).await;
63 }
64 Ok(())
65 }
66 std::task::Poll::Ready(Err(_)) => {
70 close_session_registrations(&cleanup);
71 wait_session_cleanup(cleanup).await;
72 Ok(())
73 }
74 std::task::Poll::Pending => match future::select(run, promotion_rx).await {
75 Either::Left((result, promotion_rx)) => {
76 if result.is_err() || promotion_rx.await.is_err() {
81 close_session_registrations(&cleanup);
82 wait_session_cleanup(cleanup).await;
83 }
84 result
85 }
86 Either::Right((Ok(()), run)) => {
87 run_attached_session_runner(run, cleanup, error_scope).await
88 }
89 Either::Right((Err(_), run)) => {
90 close_session_registrations(&cleanup);
91 run_attached_session_runner(run, cleanup, error_scope).await
92 }
93 },
94 }
95}
96
97fn send_session_setup<Counterpart, Request, Run>(
98 connection: V2ConnectionTo<Counterpart>,
99 request: Request,
100 dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
101 run: Run,
102 ordered: bool,
103) -> SentRequest<Request::Response>
104where
105 Counterpart: HasPeer<Agent>,
106 Request: JsonRpcRequest + 'static,
107 Request::Response: 'static,
108 Run: RunWithConnectionTo<Counterpart> + 'static,
109{
110 let raw_connection = connection.raw_connection().clone();
111 if dynamic_handler_registrations.is_empty() {
112 drop(run);
113 if ordered {
114 raw_connection.send_ordered_request_to(Agent, request)
115 } else {
116 raw_connection.send_request_to(Agent, request)
117 }
118 } else {
119 let handlers_ready = raw_connection.dynamic_handler_barrier();
120 let (runner_started_tx, runner_started_rx) = oneshot::channel();
121 let (promotion_tx, promotion_rx) = oneshot::channel();
122 let cleanup = session_cleanup(&dynamic_handler_registrations);
123 let runner_started = match raw_connection.spawn(run_pending_session_setup(
124 raw_connection.clone(),
125 run,
126 runner_started_tx,
127 promotion_rx,
128 cleanup,
129 )) {
130 Ok(()) => Either::Left(async move {
131 runner_started_rx.await.map_err(|error| {
132 crate::util::internal_error(format!(
133 "session setup runner stopped before its initial poll: {error}"
134 ))
135 })?
136 }),
137 Err(error) => Either::Right(future::ready(Err(error))),
138 };
139 let readiness = async move {
140 future::try_join(handlers_ready, runner_started).await?;
141 Ok(())
142 };
143 let response_hook = move |_response: &Request::Response| {
144 promotion_tx.send(()).map_err(|()| {
145 crate::util::internal_error("session setup runner stopped before setup completed")
146 })?;
147 dynamic_handler_registrations
148 .into_iter()
149 .for_each(DynamicHandlerGuard::detach);
150 Ok(())
151 };
152
153 if ordered {
154 raw_connection.send_ordered_request_to_with_response_hook_after(
155 Agent,
156 request,
157 readiness,
158 response_hook,
159 )
160 } else {
161 raw_connection.send_request_to_with_response_hook_after(
162 Agent,
163 request,
164 readiness,
165 response_hook,
166 )
167 }
168 }
169}
170
171impl<Counterpart> V2ConnectionTo<Counterpart>
172where
173 Counterpart: HasPeer<Agent>,
174{
175 pub fn build_session(&self, cwd: impl AsRef<Path>) -> V2SessionBuilder<Counterpart> {
177 V2SessionBuilder::new(self, v2::NewSessionRequest::new(cwd.as_ref()))
178 }
179
180 pub fn build_session_cwd(&self) -> Result<V2SessionBuilder<Counterpart>, crate::Error> {
184 let cwd = std::env::current_dir().map_err(|error| {
185 crate::Error::internal_error().data(format!("cannot get current directory: {error}"))
186 })?;
187 Ok(self.build_session(cwd))
188 }
189
190 pub fn build_session_from(
192 &self,
193 request: v2::NewSessionRequest,
194 ) -> V2SessionBuilder<Counterpart> {
195 V2SessionBuilder::new(self, request)
196 }
197
198 #[cfg(feature = "unstable_session_fork")]
204 pub fn fork_session(
205 &self,
206 session_id: impl Into<v2::SessionId>,
207 cwd: impl AsRef<Path>,
208 ) -> V2ForkSessionBuilder<Counterpart> {
209 self.fork_session_from(v2::ForkSessionRequest::new(session_id, cwd.as_ref()))
210 }
211
212 #[cfg(feature = "unstable_session_fork")]
218 pub fn fork_session_from(
219 &self,
220 request: v2::ForkSessionRequest,
221 ) -> V2ForkSessionBuilder<Counterpart> {
222 V2ForkSessionBuilder::new(self, request)
223 }
224
225 pub fn resume_session(
231 &self,
232 session_id: impl Into<v2::SessionId>,
233 cwd: impl AsRef<Path>,
234 ) -> V2ResumeSessionBuilder<Counterpart> {
235 self.resume_session_from(v2::ResumeSessionRequest::new(session_id, cwd.as_ref()))
236 }
237
238 pub fn resume_session_from(
245 &self,
246 request: v2::ResumeSessionRequest,
247 ) -> V2ResumeSessionBuilder<Counterpart> {
248 V2ResumeSessionBuilder::new(self, request)
249 }
250}
251
252#[must_use = "call `start_session` or `on_proxy_session_start` to send the `session/new` request"]
263#[derive(Debug)]
264pub struct V2SessionBuilder<Counterpart, Run = NullRun>
265where
266 Counterpart: HasPeer<Agent>,
267 Run: RunWithConnectionTo<Counterpart>,
268{
269 connection: V2ConnectionTo<Counterpart>,
270 request: v2::NewSessionRequest,
271 dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
272 run: Run,
273}
274
275impl<Counterpart> V2SessionBuilder<Counterpart, NullRun>
276where
277 Counterpart: HasPeer<Agent>,
278{
279 fn new(connection: &V2ConnectionTo<Counterpart>, request: v2::NewSessionRequest) -> Self {
280 Self {
281 connection: connection.clone(),
282 request,
283 dynamic_handler_registrations: Vec::new(),
284 run: NullRun,
285 }
286 }
287}
288
289impl<Counterpart, Run> V2SessionBuilder<Counterpart, Run>
290where
291 Counterpart: HasPeer<Agent>,
292 Run: RunWithConnectionTo<Counterpart>,
293{
294 #[cfg(feature = "unstable_mcp_over_acp")]
302 pub fn with_mcp_server<McpRun>(
303 mut self,
304 mcp_server: McpServer<Counterpart, McpRun>,
305 ) -> Result<V2SessionBuilder<Counterpart, ChainRun<Run, McpRun>>, crate::Error>
306 where
307 McpRun: RunWithConnectionTo<Counterpart>,
308 {
309 let (handler, mcp_run) = mcp_server.into_v2_handler_and_runner();
310 self.dynamic_handler_registrations
311 .push(handler.into_dynamic_handler(&mut self.request.mcp_servers, &self.connection)?);
312 Ok(V2SessionBuilder {
313 connection: self.connection,
314 request: self.request,
315 dynamic_handler_registrations: self.dynamic_handler_registrations,
316 run: ChainRun::new(self.run, mcp_run),
317 })
318 }
319
320 fn send_new_session(self, ordered: bool) -> SentRequest<v2::NewSessionResponse>
321 where
322 Run: 'static,
323 {
324 let Self {
325 connection,
326 request,
327 dynamic_handler_registrations,
328 run,
329 } = self;
330 send_session_setup(
331 connection,
332 request,
333 dynamic_handler_registrations,
334 run,
335 ordered,
336 )
337 }
338
339 pub fn start_session(self) -> SentRequest<OpenedV2Session<Counterpart, v2::NewSessionResponse>>
352 where
353 Run: 'static,
354 {
355 let session_connection = self.connection.clone();
356 self.send_new_session(false).map(move |response| {
357 let session = V2Session {
358 session_id: response.session_id.clone(),
359 connection: session_connection,
360 };
361 Ok(OpenedV2Session { session, response })
362 })
363 }
364
365 pub fn on_proxy_session_start<F, Fut>(
377 self,
378 responder: Responder<v2::NewSessionResponse>,
379 op: F,
380 ) -> Result<(), crate::Error>
381 where
382 Counterpart: HasPeer<Client>,
383 Run: 'static,
384 F: FnOnce(OpenedV2Session<Counterpart, v2::NewSessionResponse>) -> Fut + Send + 'static,
385 Fut: Future<Output = Result<(), crate::Error>> + Send,
386 {
387 let session_connection = self.connection.clone();
388 self.send_new_session(true)
389 .forward_cancellation_from(responder.cancellation())
390 .on_receiving_ok_result(responder, async move |response, responder| {
391 let session_id = response.session_id.clone();
392 let raw_connection = session_connection.raw_connection();
393 let route = match raw_connection.add_dynamic_handler(ProxySessionMessages::new(
394 crate::schema::v1::SessionId::from(session_id.0.clone()),
395 )) {
396 Ok(route) => route,
397 Err(error) => return responder.respond_with_error(error),
398 };
399
400 let opened = OpenedV2Session {
401 session: V2Session {
402 session_id,
403 connection: session_connection.clone(),
404 },
405 response: response.clone(),
406 };
407 responder.respond(response)?;
408 route.detach();
409 raw_connection.spawn(async move { op(opened).await })
410 })
411 }
412}
413
414#[cfg(feature = "unstable_session_fork")]
427#[must_use = "call `start_session` or `on_proxy_session_start` to send the `session/fork` request"]
428#[derive(Debug)]
429pub struct V2ForkSessionBuilder<Counterpart, Run = NullRun>
430where
431 Counterpart: HasPeer<Agent>,
432 Run: RunWithConnectionTo<Counterpart>,
433{
434 connection: V2ConnectionTo<Counterpart>,
435 request: v2::ForkSessionRequest,
436 dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
437 run: Run,
438}
439
440#[cfg(feature = "unstable_session_fork")]
441impl<Counterpart> V2ForkSessionBuilder<Counterpart, NullRun>
442where
443 Counterpart: HasPeer<Agent>,
444{
445 fn new(connection: &V2ConnectionTo<Counterpart>, request: v2::ForkSessionRequest) -> Self {
446 Self {
447 connection: connection.clone(),
448 request,
449 dynamic_handler_registrations: Vec::new(),
450 run: NullRun,
451 }
452 }
453}
454
455#[cfg(feature = "unstable_session_fork")]
456impl<Counterpart, Run> V2ForkSessionBuilder<Counterpart, Run>
457where
458 Counterpart: HasPeer<Agent>,
459 Run: RunWithConnectionTo<Counterpart>,
460{
461 #[cfg(feature = "unstable_mcp_over_acp")]
470 pub fn with_mcp_server<McpRun>(
471 mut self,
472 mcp_server: McpServer<Counterpart, McpRun>,
473 ) -> Result<V2ForkSessionBuilder<Counterpart, ChainRun<Run, McpRun>>, crate::Error>
474 where
475 McpRun: RunWithConnectionTo<Counterpart>,
476 {
477 let (handler, mcp_run) = mcp_server.into_v2_handler_and_runner();
478 self.dynamic_handler_registrations
479 .push(handler.into_dynamic_handler(&mut self.request.mcp_servers, &self.connection)?);
480 Ok(V2ForkSessionBuilder {
481 connection: self.connection,
482 request: self.request,
483 dynamic_handler_registrations: self.dynamic_handler_registrations,
484 run: ChainRun::new(self.run, mcp_run),
485 })
486 }
487
488 fn send_fork_session(self, ordered: bool) -> SentRequest<v2::ForkSessionResponse>
489 where
490 Run: 'static,
491 {
492 let Self {
493 connection,
494 request,
495 dynamic_handler_registrations,
496 run,
497 } = self;
498 send_session_setup(
499 connection,
500 request,
501 dynamic_handler_registrations,
502 run,
503 ordered,
504 )
505 }
506
507 pub fn start_session(self) -> SentRequest<OpenedV2Session<Counterpart, v2::ForkSessionResponse>>
521 where
522 Run: 'static,
523 {
524 let session_connection = self.connection.clone();
525 self.send_fork_session(false).map(move |response| {
526 let session = V2Session {
527 session_id: response.session_id.clone(),
528 connection: session_connection,
529 };
530 Ok(OpenedV2Session { session, response })
531 })
532 }
533
534 pub fn on_proxy_session_start<F, Fut>(
547 self,
548 responder: Responder<v2::ForkSessionResponse>,
549 op: F,
550 ) -> Result<(), crate::Error>
551 where
552 Counterpart: HasPeer<Client>,
553 Run: 'static,
554 F: FnOnce(OpenedV2Session<Counterpart, v2::ForkSessionResponse>) -> Fut + Send + 'static,
555 Fut: Future<Output = Result<(), crate::Error>> + Send,
556 {
557 let session_connection = self.connection.clone();
558 self.send_fork_session(true)
559 .forward_cancellation_from(responder.cancellation())
560 .on_receiving_ok_result(responder, async move |response, responder| {
561 let session_id = response.session_id.clone();
562 let raw_connection = session_connection.raw_connection();
563 let route = match raw_connection.add_dynamic_handler(ProxySessionMessages::new(
564 crate::schema::v1::SessionId::from(session_id.0.clone()),
565 )) {
566 Ok(route) => route,
567 Err(error) => return responder.respond_with_error(error),
568 };
569
570 let opened = OpenedV2Session {
571 session: V2Session {
572 session_id,
573 connection: session_connection.clone(),
574 },
575 response: response.clone(),
576 };
577 responder.respond(response)?;
578 route.detach();
579 raw_connection.spawn(async move { op(opened).await })
580 })
581 }
582}
583
584#[must_use = "call `start_session` or `on_proxy_session_start` to send the `session/resume` request"]
596#[derive(Debug)]
597pub struct V2ResumeSessionBuilder<Counterpart, Run = NullRun>
598where
599 Counterpart: HasPeer<Agent>,
600 Run: RunWithConnectionTo<Counterpart>,
601{
602 connection: V2ConnectionTo<Counterpart>,
603 request: v2::ResumeSessionRequest,
604 dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
605 run: Run,
606}
607
608impl<Counterpart> V2ResumeSessionBuilder<Counterpart, NullRun>
609where
610 Counterpart: HasPeer<Agent>,
611{
612 fn new(connection: &V2ConnectionTo<Counterpart>, request: v2::ResumeSessionRequest) -> Self {
613 Self {
614 connection: connection.clone(),
615 request,
616 dynamic_handler_registrations: Vec::new(),
617 run: NullRun,
618 }
619 }
620}
621
622impl<Counterpart, Run> V2ResumeSessionBuilder<Counterpart, Run>
623where
624 Counterpart: HasPeer<Agent>,
625 Run: RunWithConnectionTo<Counterpart>,
626{
627 #[cfg(feature = "unstable_mcp_over_acp")]
636 pub fn with_mcp_server<McpRun>(
637 mut self,
638 mcp_server: McpServer<Counterpart, McpRun>,
639 ) -> Result<V2ResumeSessionBuilder<Counterpart, ChainRun<Run, McpRun>>, crate::Error>
640 where
641 McpRun: RunWithConnectionTo<Counterpart>,
642 {
643 let (handler, mcp_run) = mcp_server.into_v2_handler_and_runner();
644 self.dynamic_handler_registrations
645 .push(handler.into_dynamic_handler(&mut self.request.mcp_servers, &self.connection)?);
646 Ok(V2ResumeSessionBuilder {
647 connection: self.connection,
648 request: self.request,
649 dynamic_handler_registrations: self.dynamic_handler_registrations,
650 run: ChainRun::new(self.run, mcp_run),
651 })
652 }
653
654 fn send_resume_session(self, ordered: bool) -> SentRequest<v2::ResumeSessionResponse>
655 where
656 Run: 'static,
657 {
658 let Self {
659 connection,
660 request,
661 dynamic_handler_registrations,
662 run,
663 } = self;
664 send_session_setup(
665 connection,
666 request,
667 dynamic_handler_registrations,
668 run,
669 ordered,
670 )
671 }
672
673 pub fn start_session(
686 self,
687 ) -> SentRequest<OpenedV2Session<Counterpart, v2::ResumeSessionResponse>>
688 where
689 Run: 'static,
690 {
691 let session_id = self.request.session_id.clone();
692 let session_connection = self.connection.clone();
693 self.send_resume_session(false).map(move |response| {
694 let session = V2Session {
695 session_id,
696 connection: session_connection,
697 };
698 Ok(OpenedV2Session { session, response })
699 })
700 }
701
702 pub fn on_proxy_session_start<F, Fut>(
716 mut self,
717 responder: Responder<v2::ResumeSessionResponse>,
718 op: F,
719 ) -> Result<(), crate::Error>
720 where
721 Counterpart: HasPeer<Client>,
722 Run: 'static,
723 F: FnOnce(OpenedV2Session<Counterpart, v2::ResumeSessionResponse>) -> Fut + Send + 'static,
724 Fut: Future<Output = Result<(), crate::Error>> + Send,
725 {
726 let session_id = self.request.session_id.clone();
727 let session_connection = self.connection.clone();
728 self.dynamic_handler_registrations.push(
729 session_connection
730 .raw_connection()
731 .add_dynamic_handler(ProxySessionMessages::new(
732 crate::schema::v1::SessionId::from(session_id.0.clone()),
733 ))?,
734 );
735
736 self.send_resume_session(true)
737 .forward_cancellation_from(responder.cancellation())
738 .on_receiving_ok_result(responder, async move |response, responder| {
739 let opened = OpenedV2Session {
740 session: V2Session {
741 session_id,
742 connection: session_connection.clone(),
743 },
744 response: response.clone(),
745 };
746 responder.respond(response)?;
747 session_connection
748 .raw_connection()
749 .spawn(async move { op(opened).await })
750 })
751 }
752}
753
754#[derive(Debug)]
761pub struct OpenedV2Session<Link, Response>
762where
763 Link: HasPeer<Agent>,
764{
765 session: V2Session<Link>,
766 response: Response,
767}
768
769impl<Link, Response> OpenedV2Session<Link, Response>
770where
771 Link: HasPeer<Agent>,
772{
773 pub fn session(&self) -> &V2Session<Link> {
775 &self.session
776 }
777
778 pub fn response(&self) -> &Response {
780 &self.response
781 }
782
783 pub fn into_parts(self) -> (V2Session<Link>, Response) {
785 (self.session, self.response)
786 }
787
788 pub fn into_session(self) -> V2Session<Link> {
790 self.session
791 }
792}
793
794#[derive(Debug, Clone)]
801pub struct V2Session<Link>
802where
803 Link: HasPeer<Agent>,
804{
805 session_id: v2::SessionId,
806 connection: V2ConnectionTo<Link>,
807}
808
809impl<Link> V2Session<Link>
810where
811 Link: HasPeer<Agent>,
812{
813 pub fn session_id(&self) -> &v2::SessionId {
815 &self.session_id
816 }
817
818 pub fn connection(&self) -> &V2ConnectionTo<Link> {
820 &self.connection
821 }
822
823 pub fn send_prompt(&self, prompt: impl ToString) -> SentRequest<v2::PromptResponse> {
829 self.send_prompt_blocks(vec![prompt.to_string().into()])
830 }
831
832 pub fn send_prompt_blocks(
838 &self,
839 prompt: Vec<v2::ContentBlock>,
840 ) -> SentRequest<v2::PromptResponse> {
841 self.connection.send_request_to(
842 Agent,
843 v2::PromptRequest::new(self.session_id.clone(), prompt),
844 )
845 }
846
847 pub fn cancel_active_work(&self) -> Result<(), crate::Error> {
856 self.connection.send_notification_to(
857 Agent,
858 v2::CancelSessionNotification::new(self.session_id.clone()),
859 )
860 }
861
862 pub fn set_config_option(
867 &self,
868 config_id: impl Into<v2::SessionConfigId>,
869 value: impl Into<v2::SessionConfigOptionValue>,
870 ) -> SentRequest<v2::SetSessionConfigOptionResponse> {
871 self.connection.send_request_to(
872 Agent,
873 v2::SetSessionConfigOptionRequest::new(self.session_id.clone(), config_id, value),
874 )
875 }
876
877 pub fn close(&self) -> SentRequest<v2::CloseSessionResponse> {
882 self.connection
883 .send_request_to(Agent, v2::CloseSessionRequest::new(self.session_id.clone()))
884 }
885}