1use std::{future::Future, marker::PhantomData, path::Path, sync::Arc};
2
3use futures::channel::{mpsc, oneshot};
4use futures::future::{self, Either};
5
6use crate::{
7 Agent, Client, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, JsonRpcRequest, Responder,
8 Role,
9 jsonrpc::{
10 DynamicHandlerCleanup, DynamicHandlerGuard,
11 run::{NullRun, RunWithConnectionTo, RunnerErrorScope},
12 },
13 role::{HasPeer, acp::ProxySessionMessages},
14 schema::v1::{
15 ContentBlock, ContentChunk, LoadSessionRequest, LoadSessionResponse, Meta,
16 NewSessionRequest, NewSessionResponse, PromptRequest, PromptResponse, ResumeSessionRequest,
17 ResumeSessionResponse, SessionConfigOption, SessionId, SessionModeState,
18 SessionNotification, SessionUpdate, StopReason,
19 },
20 util::{MatchDispatch, MatchDispatchFrom},
21};
22
23#[cfg(feature = "unstable_mcp_over_acp")]
24use crate::{jsonrpc::run::ChainRun, mcp_server::McpServer};
25
26#[cfg(feature = "unstable_protocol_v2")]
27mod v2;
28#[cfg(feature = "unstable_protocol_v2")]
29pub use v2::*;
30
31type SessionCleanup = Vec<Arc<dyn DynamicHandlerCleanup>>;
32
33fn session_cleanup<R: Role>(guards: &[DynamicHandlerGuard<R>]) -> SessionCleanup {
34 guards
35 .iter()
36 .filter_map(DynamicHandlerGuard::cleanup)
37 .collect()
38}
39
40fn close_session_registrations(cleanup: &SessionCleanup) {
41 for registration in cleanup {
42 registration.close();
43 }
44}
45
46async fn wait_session_cleanup(cleanup: SessionCleanup) {
47 future::join_all(cleanup.iter().map(|registration| registration.wait())).await;
48}
49
50fn session_runner_connection<R: Role>(
51 connection: ConnectionTo<R>,
52 cleanup: &SessionCleanup,
53) -> (ConnectionTo<R>, RunnerErrorScope) {
54 let closing = cleanup.clone();
55 let scope = RunnerErrorScope::new(
56 move || close_session_registrations(&closing),
57 wait_session_cleanup(cleanup.clone()),
58 );
59 (connection.with_runner_error_scope(scope.clone()), scope)
60}
61
62async fn drive_session_cleanup(
65 run: impl Future<Output = Result<(), crate::Error>>,
66 cleanup: SessionCleanup,
67) -> Result<(), crate::Error> {
68 match future::select(
69 Box::pin(wait_session_cleanup(cleanup.clone())),
70 Box::pin(run),
71 )
72 .await
73 {
74 Either::Left(((), _run)) => Ok(()),
75 Either::Right((result, waiting)) => {
76 if result.is_err() {
77 close_session_registrations(&cleanup);
78 }
79 waiting.await;
80 result
81 }
82 }
83}
84
85async fn run_attached_session_runner(
86 run: impl Future<Output = Result<(), crate::Error>>,
87 cleanup: SessionCleanup,
88 error_scope: RunnerErrorScope,
89) -> Result<(), crate::Error> {
90 let result = if cleanup.is_empty() {
91 run.await
92 } else {
93 drive_session_cleanup(run, cleanup).await
94 };
95 error_scope.error().map_or(result, Err)
98}
99
100async fn run_session_scope<T>(
101 run: impl Future<Output = Result<(), crate::Error>>,
102 op: impl Future<Output = Result<T, crate::Error>>,
103 cleanup: SessionCleanup,
104 error_scope: RunnerErrorScope,
105) -> Result<T, crate::Error> {
106 let result = match future::select(Box::pin(run), Box::pin(op)).await {
107 Either::Left((run_result, op)) => {
108 let result = match run_result {
109 Ok(()) => op.await,
110 Err(error) => {
111 drop(op);
112 Err(error)
113 }
114 };
115 close_session_registrations(&cleanup);
116 wait_session_cleanup(cleanup).await;
117 result
118 }
119 Either::Right((result, run)) => {
120 close_session_registrations(&cleanup);
121 let runner_result = drive_session_cleanup(run, cleanup).await;
122 result.and_then(|value| runner_result.map(|()| value))
124 }
125 };
126 result.and_then(|value| error_scope.error().map_or(Ok(value), Err))
130}
131
132#[derive(Debug)]
134pub struct Blocking;
135impl SessionBlockState for Blocking {}
136
137#[derive(Debug)]
139pub struct NonBlocking;
140impl SessionBlockState for NonBlocking {}
141
142pub trait SessionBlockState: Send + 'static + Sync + std::fmt::Debug {}
145
146impl<Counterpart: Role> ConnectionTo<Counterpart>
147where
148 Counterpart: HasPeer<Agent>,
149{
150 pub fn build_session(&self, cwd: impl AsRef<Path>) -> SessionBuilder<Counterpart, NullRun> {
155 SessionBuilder::new(self, NewSessionRequest::new(cwd.as_ref()))
156 }
157
158 pub fn build_session_cwd(&self) -> Result<SessionBuilder<Counterpart, NullRun>, crate::Error> {
165 let cwd = std::env::current_dir().map_err(|e| {
166 crate::Error::internal_error().data(format!("cannot get current directory: {e}"))
167 })?;
168 Ok(self.build_session(cwd))
169 }
170
171 pub fn build_session_from(
176 &self,
177 request: NewSessionRequest,
178 ) -> SessionBuilder<Counterpart, NullRun> {
179 SessionBuilder::new(self, request)
180 }
181
182 pub fn load_session(
191 &self,
192 session_id: impl Into<SessionId>,
193 cwd: impl AsRef<Path>,
194 ) -> RestoreSessionBuilder<Counterpart, LoadSessionRequest> {
195 self.load_session_from(LoadSessionRequest::new(session_id, cwd.as_ref()))
196 }
197
198 pub fn load_session_from(
204 &self,
205 request: LoadSessionRequest,
206 ) -> RestoreSessionBuilder<Counterpart, LoadSessionRequest> {
207 RestoreSessionBuilder::new(self, request)
208 }
209
210 pub fn resume_session(
217 &self,
218 session_id: impl Into<SessionId>,
219 cwd: impl AsRef<Path>,
220 ) -> RestoreSessionBuilder<Counterpart, ResumeSessionRequest> {
221 self.resume_session_from(ResumeSessionRequest::new(session_id, cwd.as_ref()))
222 }
223
224 pub fn resume_session_from(
230 &self,
231 request: ResumeSessionRequest,
232 ) -> RestoreSessionBuilder<Counterpart, ResumeSessionRequest> {
233 RestoreSessionBuilder::new(self, request)
234 }
235
236 pub(crate) fn attach_session<'runner>(
247 &self,
248 response: NewSessionResponse,
249 mcp_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
250 ) -> Result<ActiveSession<'runner, Counterpart>, crate::Error> {
251 let NewSessionResponse {
252 session_id,
253 modes,
254 config_options,
255 meta,
256 ..
257 } = response;
258
259 let prepared = self.prepare_session_routing(&session_id)?;
260 Ok(prepared.into_active_session(
261 self.clone(),
262 session_id,
263 modes,
264 config_options,
265 meta,
266 mcp_handler_registrations,
267 ))
268 }
269
270 fn prepare_session_routing(
275 &self,
276 session_id: &SessionId,
277 ) -> Result<PreparedSession<Counterpart>, crate::Error> {
278 let (update_tx, update_rx) = mpsc::unbounded();
279 let handler = ActiveSessionHandler::new(session_id.clone(), update_tx.clone());
280 let session_handler_registration = self.add_dynamic_handler(handler)?;
281
282 Ok(PreparedSession {
283 update_rx,
284 update_tx,
285 session_handler_registration,
286 })
287 }
288}
289
290struct PreparedSession<Counterpart: Role>
292where
293 Counterpart: HasPeer<Agent>,
294{
295 update_rx: mpsc::UnboundedReceiver<SessionMessage>,
296 update_tx: mpsc::UnboundedSender<SessionMessage>,
297 session_handler_registration: DynamicHandlerGuard<Counterpart>,
298}
299
300impl<Counterpart> PreparedSession<Counterpart>
301where
302 Counterpart: HasPeer<Agent>,
303{
304 fn into_active_session<'runner>(
305 self,
306 connection: ConnectionTo<Counterpart>,
307 session_id: SessionId,
308 modes: Option<SessionModeState>,
309 config_options: Option<Vec<SessionConfigOption>>,
310 meta: Option<Meta>,
311 mcp_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
312 ) -> ActiveSession<'runner, Counterpart> {
313 ActiveSession {
314 session_id,
315 modes,
316 config_options,
317 meta,
318 update_rx: self.update_rx,
319 update_tx: self.update_tx,
320 connection,
321 session_handler_registration: self.session_handler_registration,
322 mcp_handler_registrations,
323 _runner: PhantomData,
324 }
325 }
326}
327
328trait RestoreRequest: JsonRpcRequest {
330 fn session_id(&self) -> &SessionId;
331 fn response_modes(response: &Self::Response) -> Option<SessionModeState>;
332 fn response_config_options(response: &Self::Response) -> Option<Vec<SessionConfigOption>>;
333 fn response_meta(response: &Self::Response) -> Option<Meta>;
334}
335
336impl RestoreRequest for LoadSessionRequest {
337 fn session_id(&self) -> &SessionId {
338 &self.session_id
339 }
340
341 fn response_modes(response: &Self::Response) -> Option<SessionModeState> {
342 response.modes.clone()
343 }
344
345 fn response_config_options(response: &Self::Response) -> Option<Vec<SessionConfigOption>> {
346 response.config_options.clone()
347 }
348
349 fn response_meta(response: &Self::Response) -> Option<Meta> {
350 response.meta.clone()
351 }
352}
353
354impl RestoreRequest for ResumeSessionRequest {
355 fn session_id(&self) -> &SessionId {
356 &self.session_id
357 }
358
359 fn response_modes(response: &Self::Response) -> Option<SessionModeState> {
360 response.modes.clone()
361 }
362
363 fn response_config_options(response: &Self::Response) -> Option<Vec<SessionConfigOption>> {
364 response.config_options.clone()
365 }
366
367 fn response_meta(response: &Self::Response) -> Option<Meta> {
368 response.meta.clone()
369 }
370}
371
372#[must_use = "use `start_session` or `on_session_start` to restore the session"]
389#[derive(Debug)]
390pub struct RestoreSessionBuilder<Counterpart, Request, BlockState = NonBlocking>
391where
392 Counterpart: HasPeer<Agent>,
393 BlockState: SessionBlockState,
394{
395 connection: ConnectionTo<Counterpart>,
396 request: Request,
397 block_state: PhantomData<BlockState>,
398}
399
400impl<Counterpart, Request> RestoreSessionBuilder<Counterpart, Request, NonBlocking>
401where
402 Counterpart: HasPeer<Agent>,
403{
404 fn new(connection: &ConnectionTo<Counterpart>, request: Request) -> Self {
405 Self {
406 connection: connection.clone(),
407 request,
408 block_state: PhantomData,
409 }
410 }
411
412 pub fn block_task(self) -> RestoreSessionBuilder<Counterpart, Request, Blocking> {
416 RestoreSessionBuilder {
417 connection: self.connection,
418 request: self.request,
419 block_state: PhantomData,
420 }
421 }
422}
423
424fn restored_session<Counterpart, Request>(
425 connection: ConnectionTo<Counterpart>,
426 session_id: SessionId,
427 prepared: PreparedSession<Counterpart>,
428 response: Request::Response,
429) -> RestoredSession<'static, Counterpart, Request::Response>
430where
431 Counterpart: HasPeer<Agent>,
432 Request: RestoreRequest,
433{
434 let session = prepared.into_active_session(
435 connection,
436 session_id,
437 Request::response_modes(&response),
438 Request::response_config_options(&response),
439 Request::response_meta(&response),
440 Vec::new(),
441 );
442
443 RestoredSession { session, response }
444}
445
446fn on_restore_session_start<Counterpart, Request, F, Fut>(
447 builder: RestoreSessionBuilder<Counterpart, Request>,
448 op: F,
449) -> Result<(), crate::Error>
450where
451 Counterpart: HasPeer<Agent>,
452 Request: RestoreRequest,
453 F: FnOnce(RestoredSession<'static, Counterpart, Request::Response>) -> Fut + Send + 'static,
454 Fut: Future<Output = Result<(), crate::Error>> + Send,
455{
456 ensure_v1_session_protocol(&builder.connection)?;
457
458 let RestoreSessionBuilder {
459 connection,
460 request,
461 block_state: _,
462 } = builder;
463 let session_id = request.session_id().clone();
464 let prepared = connection.prepare_session_routing(&session_id)?;
465 let routing_ready = connection.dynamic_handler_barrier();
466
467 connection
468 .send_ordered_request_to_after(Agent, request, routing_ready)
469 .on_receiving_result({
470 let connection = connection.clone();
471 async move |result| {
472 let response = result?;
473 let restored = restored_session::<_, Request>(
474 connection.clone(),
475 session_id,
476 prepared,
477 response,
478 );
479 connection.spawn(async move { op(restored).await })
480 }
481 })
482}
483
484async fn start_restored_session<Counterpart, Request>(
485 builder: RestoreSessionBuilder<Counterpart, Request, Blocking>,
486) -> Result<RestoredSession<'static, Counterpart, Request::Response>, crate::Error>
487where
488 Counterpart: HasPeer<Agent>,
489 Request: RestoreRequest,
490{
491 ensure_v1_session_protocol(&builder.connection)?;
492
493 let RestoreSessionBuilder {
494 connection,
495 request,
496 block_state: _,
497 } = builder;
498 let session_id = request.session_id().clone();
499 let prepared = connection.prepare_session_routing(&session_id)?;
500 let routing_ready = connection.dynamic_handler_barrier();
501 let session_connection = connection.clone();
502
503 connection
504 .send_ordered_request_to_after(Agent, request, routing_ready)
505 .block_task_with_ordered_result(move |result| {
506 let response = result?;
507 Ok(restored_session::<_, Request>(
508 session_connection,
509 session_id,
510 prepared,
511 response,
512 ))
513 })
514 .await
515}
516
517impl<Counterpart> RestoreSessionBuilder<Counterpart, LoadSessionRequest>
518where
519 Counterpart: HasPeer<Agent>,
520{
521 pub fn on_session_start<F, Fut>(self, op: F) -> Result<(), crate::Error>
528 where
529 F: FnOnce(RestoredSession<'static, Counterpart, LoadSessionResponse>) -> Fut
530 + Send
531 + 'static,
532 Fut: Future<Output = Result<(), crate::Error>> + Send,
533 {
534 on_restore_session_start(self, op)
535 }
536}
537
538impl<Counterpart> RestoreSessionBuilder<Counterpart, ResumeSessionRequest>
539where
540 Counterpart: HasPeer<Agent>,
541{
542 pub fn on_session_start<F, Fut>(self, op: F) -> Result<(), crate::Error>
548 where
549 F: FnOnce(RestoredSession<'static, Counterpart, ResumeSessionResponse>) -> Fut
550 + Send
551 + 'static,
552 Fut: Future<Output = Result<(), crate::Error>> + Send,
553 {
554 on_restore_session_start(self, op)
555 }
556}
557
558impl<Counterpart> RestoreSessionBuilder<Counterpart, LoadSessionRequest, Blocking>
559where
560 Counterpart: HasPeer<Agent>,
561{
562 pub async fn start_session(
569 self,
570 ) -> Result<RestoredSession<'static, Counterpart, LoadSessionResponse>, crate::Error> {
571 start_restored_session(self).await
572 }
573}
574
575impl<Counterpart> RestoreSessionBuilder<Counterpart, ResumeSessionRequest, Blocking>
576where
577 Counterpart: HasPeer<Agent>,
578{
579 pub async fn start_session(
586 self,
587 ) -> Result<RestoredSession<'static, Counterpart, ResumeSessionResponse>, crate::Error> {
588 start_restored_session(self).await
589 }
590}
591
592pub struct RestoredSession<'runner, Link, Response>
600where
601 Link: HasPeer<Agent>,
602{
603 session: ActiveSession<'runner, Link>,
604 response: Response,
605}
606
607impl<'runner, Link, Response> RestoredSession<'runner, Link, Response>
608where
609 Link: HasPeer<Agent>,
610{
611 pub fn session(&self) -> &ActiveSession<'runner, Link> {
613 &self.session
614 }
615
616 pub fn session_mut(&mut self) -> &mut ActiveSession<'runner, Link> {
618 &mut self.session
619 }
620
621 pub fn response(&self) -> &Response {
623 &self.response
624 }
625
626 pub fn into_parts(self) -> (ActiveSession<'runner, Link>, Response) {
628 (self.session, self.response)
629 }
630
631 pub fn into_session(self) -> ActiveSession<'runner, Link> {
633 self.session
634 }
635}
636
637impl<Link, Response> std::fmt::Debug for RestoredSession<'_, Link, Response>
638where
639 Link: HasPeer<Agent>,
640 Response: std::fmt::Debug,
641{
642 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
643 formatter
644 .debug_struct("RestoredSession")
645 .field("session_id", self.session.session_id())
646 .field("response", &self.response)
647 .finish()
648 }
649}
650
651#[must_use = "use `start_session`, `run_until`, or `on_session_start` to start the session"]
659#[derive(Debug)]
660pub struct SessionBuilder<
661 Counterpart,
662 Run: RunWithConnectionTo<Counterpart> = NullRun,
663 BlockState: SessionBlockState = NonBlocking,
664> where
665 Counterpart: HasPeer<Agent>,
666{
667 connection: ConnectionTo<Counterpart>,
668 request: NewSessionRequest,
669 dynamic_handler_registrations: Vec<DynamicHandlerGuard<Counterpart>>,
670 run: Run,
671 block_state: PhantomData<BlockState>,
672}
673
674impl<Counterpart> SessionBuilder<Counterpart, NullRun, NonBlocking>
675where
676 Counterpart: HasPeer<Agent>,
677{
678 fn new(connection: &ConnectionTo<Counterpart>, request: NewSessionRequest) -> Self {
679 SessionBuilder {
680 connection: connection.clone(),
681 request,
682 dynamic_handler_registrations: Vec::default(),
683 run: NullRun,
684 block_state: PhantomData,
685 }
686 }
687}
688
689impl<Counterpart, R, BlockState> SessionBuilder<Counterpart, R, BlockState>
690where
691 Counterpart: HasPeer<Agent>,
692 R: RunWithConnectionTo<Counterpart>,
693 BlockState: SessionBlockState,
694{
695 #[cfg(feature = "unstable_mcp_over_acp")]
697 pub fn with_mcp_server<McpRun>(
698 mut self,
699 mcp_server: McpServer<Counterpart, McpRun>,
700 ) -> Result<SessionBuilder<Counterpart, ChainRun<R, McpRun>, BlockState>, crate::Error>
701 where
702 McpRun: RunWithConnectionTo<Counterpart>,
703 {
704 let (handler, mcp_run) = mcp_server.into_handler_and_runner();
705 self.dynamic_handler_registrations
706 .push(handler.into_dynamic_handler(&mut self.request, &self.connection)?);
707 Ok(SessionBuilder {
708 connection: self.connection,
709 request: self.request,
710 dynamic_handler_registrations: self.dynamic_handler_registrations,
711 run: ChainRun::new(self.run, mcp_run),
712 block_state: self.block_state,
713 })
714 }
715
716 pub fn on_session_start<F, Fut>(self, op: F) -> Result<(), crate::Error>
761 where
762 R: 'static,
763 F: FnOnce(ActiveSession<'static, Counterpart>) -> Fut + Send + 'static,
764 Fut: Future<Output = Result<(), crate::Error>> + Send,
765 {
766 ensure_v1_session_protocol(&self.connection)?;
767
768 let Self {
769 connection,
770 request,
771 dynamic_handler_registrations,
772 run,
773 block_state: _,
774 } = self;
775
776 let cleanup = session_cleanup(&dynamic_handler_registrations);
777 let (runner_connection, error_scope) =
778 session_runner_connection(connection.clone(), &cleanup);
779 connection.spawn(run_attached_session_runner(
780 run.run_with_connection_to(runner_connection),
781 cleanup,
782 error_scope,
783 ))?;
784
785 connection
786 .send_ordered_request_to(Agent, request)
787 .on_receiving_result({
788 let connection = connection.clone();
789 async move |result| {
790 let response = result?;
791
792 let active_session =
793 connection.attach_session(response, dynamic_handler_registrations)?;
794
795 connection.spawn(async move { op(active_session).await })
796 }
797 })
798 }
799
800 pub fn on_proxy_session_start<F, Fut>(
853 self,
854 responder: Responder<NewSessionResponse>,
855 op: F,
856 ) -> Result<(), crate::Error>
857 where
858 F: FnOnce(SessionId) -> Fut + Send + 'static,
859 Fut: Future<Output = Result<(), crate::Error>> + Send,
860 Counterpart: HasPeer<Client>,
861 R: 'static,
862 {
863 ensure_v1_session_protocol(&self.connection)?;
864
865 let Self {
866 connection,
867 request,
868 dynamic_handler_registrations,
869 run,
870 block_state: _,
871 } = self;
872
873 let cleanup = session_cleanup(&dynamic_handler_registrations);
874 let (runner_connection, error_scope) =
875 session_runner_connection(connection.clone(), &cleanup);
876 connection.spawn(run_attached_session_runner(
877 run.run_with_connection_to(runner_connection),
878 cleanup,
879 error_scope,
880 ))?;
881
882 let sent = connection.send_ordered_request_to(Agent, request);
884 let sent = sent.forward_cancellation_from(responder.cancellation());
885
886 sent.on_receiving_ok_result(responder, {
887 let connection = connection.clone();
888 async move |response, responder| {
889 let session_id = response.session_id.clone();
892 responder.respond(response)?;
893
894 connection
896 .add_dynamic_handler(ProxySessionMessages::new(session_id.clone()))?
897 .detach();
898
899 dynamic_handler_registrations
901 .into_iter()
902 .for_each(DynamicHandlerGuard::detach);
903
904 connection.spawn(async move { op(session_id).await })
905 }
906 })
907 }
908}
909
910impl<Counterpart, R> SessionBuilder<Counterpart, R, NonBlocking>
911where
912 Counterpart: HasPeer<Agent>,
913 R: RunWithConnectionTo<Counterpart>,
914{
915 pub fn block_task(self) -> SessionBuilder<Counterpart, R, Blocking> {
924 SessionBuilder {
925 connection: self.connection,
926 request: self.request,
927 dynamic_handler_registrations: self.dynamic_handler_registrations,
928 run: self.run,
929 block_state: PhantomData,
930 }
931 }
932}
933
934impl<Counterpart, R> SessionBuilder<Counterpart, R, Blocking>
935where
936 Counterpart: HasPeer<Agent>,
937 R: RunWithConnectionTo<Counterpart>,
938{
939 pub async fn run_until<T>(
950 self,
951 op: impl for<'runner> AsyncFnOnce(
952 ActiveSession<'runner, Counterpart>,
953 ) -> Result<T, crate::Error>,
954 ) -> Result<T, crate::Error> {
955 let Self {
956 connection,
957 request,
958 dynamic_handler_registrations,
959 run,
960 block_state: _,
961 } = self;
962
963 let cleanup = session_cleanup(&dynamic_handler_registrations);
964 let (runner_connection, error_scope) =
965 session_runner_connection(connection.clone(), &cleanup);
966 run_session_scope(
967 run.run_with_connection_to(runner_connection),
968 async move {
969 ensure_v1_session_protocol(&connection)?;
970 let response = connection
971 .send_request_to(Agent, request)
972 .block_task()
973 .await?;
974 let active_session =
975 connection.attach_session(response, dynamic_handler_registrations)?;
976 op(active_session).await
977 },
978 cleanup,
979 error_scope,
980 )
981 .await
982 }
983
984 pub async fn start_session(self) -> Result<ActiveSession<'static, Counterpart>, crate::Error>
994 where
995 R: 'static,
996 {
997 ensure_v1_session_protocol(&self.connection)?;
998
999 let Self {
1000 connection,
1001 request,
1002 dynamic_handler_registrations,
1003 run,
1004 block_state: _,
1005 } = self;
1006
1007 let (active_session_tx, active_session_rx) = oneshot::channel();
1008
1009 let cleanup = session_cleanup(&dynamic_handler_registrations);
1010 let (runner_connection, error_scope) =
1011 session_runner_connection(connection.clone(), &cleanup);
1012 connection.spawn(run_attached_session_runner(
1013 run.run_with_connection_to(runner_connection),
1014 cleanup,
1015 error_scope,
1016 ))?;
1017
1018 connection.clone().spawn(async move {
1019 let response = connection
1020 .send_request_to(Agent, request)
1021 .block_task()
1022 .await?;
1023
1024 let active_session =
1025 connection.attach_session(response, dynamic_handler_registrations)?;
1026
1027 active_session_tx
1028 .send(active_session)
1029 .map_err(|_| crate::Error::internal_error())?;
1030
1031 Ok(())
1032 })?;
1033
1034 active_session_rx
1035 .await
1036 .map_err(|_| crate::Error::internal_error())
1037 }
1038
1039 pub async fn start_session_proxy(
1057 self,
1058 responder: Responder<NewSessionResponse>,
1059 ) -> Result<SessionId, crate::Error>
1060 where
1061 Counterpart: HasPeer<Client>,
1062 R: 'static,
1063 {
1064 let active_session = self.start_session().await?;
1065 let session_id = active_session.session_id().clone();
1066 responder.respond(active_session.response())?;
1067 active_session.proxy_remaining_messages()?;
1068 Ok(session_id)
1069 }
1070}
1071
1072#[derive(Debug)]
1081pub struct ActiveSession<'runner, Link>
1082where
1083 Link: HasPeer<Agent>,
1084{
1085 session_id: SessionId,
1086 update_rx: mpsc::UnboundedReceiver<SessionMessage>,
1087 update_tx: mpsc::UnboundedSender<SessionMessage>,
1088 modes: Option<SessionModeState>,
1089 config_options: Option<Vec<SessionConfigOption>>,
1090 meta: Option<serde_json::Map<String, serde_json::Value>>,
1091 connection: ConnectionTo<Link>,
1092
1093 session_handler_registration: DynamicHandlerGuard<Link>,
1097
1098 mcp_handler_registrations: Vec<DynamicHandlerGuard<Link>>,
1102
1103 _runner: PhantomData<&'runner ()>,
1105}
1106
1107#[non_exhaustive]
1109#[derive(Debug)]
1110#[allow(
1111 clippy::large_enum_variant,
1112 reason = "Dispatch messages vastly outnumber StopReason; boxing would add a heap allocation"
1113)]
1114pub enum SessionMessage {
1115 SessionMessage(Dispatch),
1118
1119 StopReason(StopReason),
1121}
1122
1123impl<Link> ActiveSession<'_, Link>
1124where
1125 Link: HasPeer<Agent>,
1126{
1127 pub fn session_id(&self) -> &SessionId {
1129 &self.session_id
1130 }
1131
1132 pub fn modes(&self) -> Option<&SessionModeState> {
1134 self.modes.as_ref()
1135 }
1136
1137 pub fn config_options(&self) -> Option<&[SessionConfigOption]> {
1139 self.config_options.as_deref()
1140 }
1141
1142 pub fn meta(&self) -> Option<&serde_json::Map<String, serde_json::Value>> {
1144 self.meta.as_ref()
1145 }
1146
1147 pub fn response(&self) -> NewSessionResponse {
1152 NewSessionResponse::new(self.session_id.clone())
1153 .modes(self.modes.clone())
1154 .config_options(self.config_options.clone())
1155 .meta(self.meta.clone())
1156 }
1157
1158 pub fn connection(&self) -> &ConnectionTo<Link> {
1160 &self.connection
1161 }
1162
1163 pub fn send_prompt(&mut self, prompt: impl ToString) -> Result<(), crate::Error> {
1165 let update_tx = self.update_tx.clone();
1166 self.connection
1167 .send_ordered_request_to(
1168 Agent,
1169 PromptRequest::new(self.session_id.clone(), vec![prompt.to_string().into()]),
1170 )
1171 .on_receiving_result(async move |result| {
1172 let PromptResponse { stop_reason, .. } = result?;
1173
1174 update_tx
1175 .unbounded_send(SessionMessage::StopReason(stop_reason))
1176 .map_err(crate::util::internal_error)?;
1177
1178 Ok(())
1179 })
1180 }
1181
1182 pub async fn read_update(&mut self) -> Result<SessionMessage, crate::Error> {
1184 use futures::StreamExt;
1185 let message =
1186 self.update_rx.next().await.ok_or_else(|| {
1187 crate::util::internal_error("session channel closed unexpectedly")
1188 })?;
1189
1190 Ok(message)
1191 }
1192
1193 pub async fn read_to_string(&mut self) -> Result<String, crate::Error> {
1196 let mut output = String::new();
1197 loop {
1198 let update = self.read_update().await?;
1199 tracing::trace!(?update, "read_to_string update");
1200 match update {
1201 SessionMessage::SessionMessage(dispatch) => MatchDispatch::new(dispatch)
1202 .if_notification(async |notif: SessionNotification| match notif.update {
1203 SessionUpdate::AgentMessageChunk(ContentChunk {
1204 content: ContentBlock::Text(text),
1205 ..
1206 }) => {
1207 output.push_str(&text.text);
1208 Ok(())
1209 }
1210 _ => Ok(()),
1211 })
1212 .await
1213 .otherwise_ignore()?,
1214 SessionMessage::StopReason(_stop_reason) => break,
1215 }
1216 }
1217 Ok(output)
1218 }
1219}
1220
1221impl<Link> ActiveSession<'static, Link>
1222where
1223 Link: HasPeer<Agent>,
1224{
1225 pub fn proxy_remaining_messages(self) -> Result<(), crate::Error>
1254 where
1255 Link: HasPeer<Client>,
1256 {
1257 let ActiveSession {
1259 session_id,
1260 mut update_rx,
1261 update_tx,
1262 connection,
1263 session_handler_registration,
1264 mcp_handler_registrations,
1265 modes: _,
1267 config_options: _,
1268 meta: _,
1269 _runner,
1270 } = self;
1271
1272 drop(session_handler_registration);
1276
1277 drop(update_tx);
1281
1282 while let Ok(message) = update_rx.try_recv() {
1286 match message {
1287 SessionMessage::SessionMessage(dispatch) => {
1288 connection.send_proxied_message_to(Client, dispatch)?;
1290 }
1291 SessionMessage::StopReason(_) => {
1292 }
1294 }
1295 }
1296
1297 connection
1301 .add_dynamic_handler(ProxySessionMessages::new(session_id))?
1302 .detach();
1303
1304 for registration in mcp_handler_registrations {
1306 registration.detach();
1307 }
1308
1309 Ok(())
1310 }
1311}
1312
1313struct ActiveSessionHandler {
1314 session_id: SessionId,
1315 update_tx: mpsc::UnboundedSender<SessionMessage>,
1316}
1317
1318impl ActiveSessionHandler {
1319 pub fn new(session_id: SessionId, update_tx: mpsc::UnboundedSender<SessionMessage>) -> Self {
1320 Self {
1321 session_id,
1322 update_tx,
1323 }
1324 }
1325}
1326
1327impl<Counterpart: Role> HandleDispatchFrom<Counterpart> for ActiveSessionHandler
1328where
1329 Counterpart: HasPeer<Agent>,
1330{
1331 async fn handle_dispatch_from(
1332 &mut self,
1333 message: Dispatch,
1334 cx: ConnectionTo<Counterpart>,
1335 ) -> Result<Handled<Dispatch>, crate::Error> {
1336 tracing::trace!(
1338 ?message,
1339 handler_session_id = ?self.session_id,
1340 "ActiveSessionHandler::handle_dispatch"
1341 );
1342 MatchDispatchFrom::new(message, &cx)
1343 .if_dispatch_from(Agent, async |message| {
1344 if let Some(session_id) = message.get_session_id()? {
1345 tracing::trace!(
1346 message_session_id = ?session_id,
1347 handler_session_id = ?self.session_id,
1348 "ActiveSessionHandler::handle_dispatch"
1349 );
1350 if session_id == self.session_id {
1351 self.update_tx
1352 .unbounded_send(SessionMessage::SessionMessage(message))
1353 .map_err(crate::util::internal_error)?;
1354 return Ok(Handled::Yes);
1355 }
1356 }
1357
1358 Ok(Handled::No {
1360 message,
1361 retry: false,
1362 })
1363 })
1364 .await
1365 .done()
1366 }
1367
1368 fn describe_chain(&self) -> impl std::fmt::Debug {
1369 format!("ActiveSessionHandler({})", self.session_id)
1370 }
1371}
1372
1373#[cfg(not(feature = "unstable_protocol_v2"))]
1374#[allow(
1375 clippy::unnecessary_wraps,
1376 reason = "signature matches the feature-enabled protocol guard"
1377)]
1378fn ensure_v1_session_protocol<Counterpart: Role>(
1379 _connection: &ConnectionTo<Counterpart>,
1380) -> Result<(), crate::Error> {
1381 Ok(())
1382}
1383
1384#[cfg(feature = "unstable_protocol_v2")]
1385fn ensure_v1_session_protocol<Counterpart: Role>(
1386 connection: &ConnectionTo<Counterpart>,
1387) -> Result<(), crate::Error> {
1388 if connection.acp_protocol_version() != Some(crate::schema::ProtocolVersion::V2) {
1389 return Ok(());
1390 }
1391
1392 Err(crate::Error::invalid_request().data(
1393 "stable session builders use ACP protocol v1 types, but this is a protocol v2 connection; \
1394 use the `V2ConnectionTo` supplied to `Client.v2()` callbacks",
1395 ))
1396}