1use futures::{
4 SinkExt, StreamExt,
5 channel::{mpsc, oneshot},
6 future::{self, BoxFuture, Either},
7};
8use schemars::JsonSchema;
9use serde::{Serialize, de::DeserializeOwned};
10use std::sync::{Arc, Mutex};
11
12use crate::{ConnectionTo, Error, Role, RunWithConnectionTo};
13
14use super::{McpConnectionTo, McpTool};
15
16struct ToolCall<P, R, MyRole: Role> {
17 params: P,
18 mcp_connection: McpConnectionTo<MyRole>,
19 result_tx: futures::channel::oneshot::Sender<Result<R, Error>>,
20 done_tx: oneshot::Sender<()>,
21}
22
23struct QueuedCall<P, R, MyRole: Role>(Arc<Mutex<Option<ToolCall<P, R, MyRole>>>>);
28
29impl<P, R, MyRole: Role> QueuedCall<P, R, MyRole> {
30 fn share(&self) -> Self {
31 Self(self.0.clone())
32 }
33
34 fn take(&self) -> Option<ToolCall<P, R, MyRole>> {
35 self.0.lock().unwrap().take()
36 }
37}
38
39impl<P, R, MyRole: Role> Drop for QueuedCall<P, R, MyRole> {
40 fn drop(&mut self) {
41 if let Some(ToolCall {
42 params,
43 mcp_connection,
44 result_tx,
45 done_tx,
46 }) = self.take()
47 {
48 drop(params);
51 drop(mcp_connection);
52 drop(result_tx);
53 let _finished = done_tx.send(());
54 }
55 }
56}
57
58struct CallResult<P, R, MyRole: Role> {
59 result_rx: oneshot::Receiver<Result<R, Error>>,
62 queued_call: QueuedCall<P, R, MyRole>,
63}
64
65async fn run_call<R>(
68 future: impl Future<Output = Result<R, Error>>,
69 mut result_tx: oneshot::Sender<Result<R, Error>>,
70 done_tx: oneshot::Sender<()>,
71) {
72 let result = {
73 let cancelled = result_tx.cancellation();
74 futures::pin_mut!(future, cancelled);
75 match future::select(cancelled, future).await {
76 Either::Left(_) => None,
77 Either::Right((result, _)) => Some(result),
78 }
79 };
80 if let Some(result) = result {
81 drop(result_tx.send(result));
83 }
84 let _finished = done_tx.send(());
85}
86
87struct ToolFnMutRunner<F, P, R, Counterpart: Role> {
88 func: F,
89 call_rx: mpsc::Receiver<QueuedCall<P, R, Counterpart>>,
90 tool_future_fn: Box<
91 dyn for<'a> Fn(
92 &'a mut F,
93 P,
94 McpConnectionTo<Counterpart>,
95 ) -> BoxFuture<'a, Result<R, Error>>
96 + Send,
97 >,
98}
99
100impl<F, P, R, Counterpart, Counterpart1> RunWithConnectionTo<Counterpart1>
101 for ToolFnMutRunner<F, P, R, Counterpart>
102where
103 Counterpart: Role,
104 Counterpart1: Role,
105 P: Send,
106 R: Send,
107 F: Send,
108{
109 async fn run_with_connection_to(
110 self,
111 _connection: ConnectionTo<Counterpart1>,
112 ) -> Result<(), Error> {
113 let ToolFnMutRunner {
114 mut func,
115 mut call_rx,
116 tool_future_fn,
117 } = self;
118 while let Some(queued_call) = call_rx.next().await {
119 let Some(ToolCall {
120 params,
121 mcp_connection,
122 result_tx,
123 done_tx,
124 }) = queued_call.take()
125 else {
126 continue;
127 };
128 if result_tx.is_canceled() {
129 drop(params);
130 drop(mcp_connection);
131 let _finished = done_tx.send(());
132 continue;
133 }
134 run_call(
135 tool_future_fn(&mut func, params, mcp_connection),
136 result_tx,
137 done_tx,
138 )
139 .await;
140 }
141 Ok(())
142 }
143}
144
145struct ToolFnRunner<F, P, R, Counterpart: Role> {
146 func: F,
147 call_rx: mpsc::Receiver<QueuedCall<P, R, Counterpart>>,
148 tool_future_fn: Box<
149 dyn for<'a> Fn(&'a F, P, McpConnectionTo<Counterpart>) -> BoxFuture<'a, Result<R, Error>>
150 + Send
151 + Sync,
152 >,
153}
154
155impl<F, P, R, Counterpart, Counterpart1> RunWithConnectionTo<Counterpart1>
156 for ToolFnRunner<F, P, R, Counterpart>
157where
158 Counterpart: Role,
159 Counterpart1: Role,
160 P: Send,
161 R: Send,
162 F: Send + Sync,
163{
164 async fn run_with_connection_to(
165 self,
166 _connection: ConnectionTo<Counterpart1>,
167 ) -> Result<(), Error> {
168 let ToolFnRunner {
169 func,
170 call_rx,
171 tool_future_fn,
172 } = self;
173 crate::util::process_stream_concurrently(
174 call_rx,
175 async |tool_call| {
176 fn hack<'a, F, P, R, MyRole>(
177 func: &'a F,
178 params: P,
179 mcp_connection: McpConnectionTo<MyRole>,
180 tool_future_fn: &'a (
181 dyn Fn(
182 &'a F,
183 P,
184 McpConnectionTo<MyRole>,
185 ) -> BoxFuture<'a, Result<R, Error>>
186 + Send
187 + Sync
188 ),
189 result_tx: oneshot::Sender<Result<R, Error>>,
190 done_tx: oneshot::Sender<()>,
191 ) -> BoxFuture<'a, ()>
192 where
193 MyRole: Role,
194 P: Send,
195 R: Send,
196 F: Send + Sync,
197 {
198 Box::pin(async move {
199 if result_tx.is_canceled() {
200 drop(params);
201 drop(mcp_connection);
202 let _finished = done_tx.send(());
203 return;
204 }
205 run_call(
206 tool_future_fn(func, params, mcp_connection),
207 result_tx,
208 done_tx,
209 )
210 .await;
211 })
212 }
213
214 let Some(ToolCall {
215 params,
216 mcp_connection,
217 result_tx,
218 done_tx,
219 }) = tool_call.take()
220 else {
221 return Ok(());
222 };
223
224 hack(
225 &func,
226 params,
227 mcp_connection,
228 &*tool_future_fn,
229 result_tx,
230 done_tx,
231 )
232 .await;
233 Ok(())
234 },
235 |a, b| Box::pin(a(b)),
236 )
237 .await
238 }
239}
240
241struct ToolFnTool<P, Ret, R: Role> {
242 name: String,
243 description: String,
244 call_tx: mpsc::Sender<QueuedCall<P, Ret, R>>,
245}
246
247impl<P, Ret, R> McpTool<R> for ToolFnTool<P, Ret, R>
248where
249 R: Role,
250 P: JsonSchema + DeserializeOwned + 'static + Send,
251 Ret: JsonSchema + Serialize + 'static + Send,
252{
253 type Input = P;
254 type Output = Ret;
255
256 fn name(&self) -> String {
257 self.name.clone()
258 }
259
260 fn description(&self) -> String {
261 self.description.clone()
262 }
263
264 async fn call_tool(&self, params: P, mcp_connection: McpConnectionTo<R>) -> Result<Ret, Error> {
265 let (result_tx, result_rx) = oneshot::channel();
266 let (done_tx, done_rx) = oneshot::channel();
267 #[cfg(feature = "unstable_mcp_over_acp")]
268 mcp_connection.register_cleanup(done_rx);
269 #[cfg(not(feature = "unstable_mcp_over_acp"))]
270 let _done_rx = done_rx;
271
272 let mut call = CallResult {
273 result_rx,
274 queued_call: QueuedCall(Arc::new(Mutex::new(Some(ToolCall {
275 params,
276 mcp_connection,
277 result_tx,
278 done_tx,
279 })))),
280 };
281 self.call_tx
282 .clone()
283 .send(call.queued_call.share())
284 .await
285 .map_err(crate::util::internal_error)?;
286
287 (&mut call.result_rx)
288 .await
289 .map_err(crate::util::internal_error)?
290 }
291}
292
293pub fn tool_fn_mut<P, Ret, F, Counterpart>(
297 name: impl ToString,
298 description: impl ToString,
299 func: F,
300 tool_future_fn: impl for<'a> Fn(
301 &'a mut F,
302 P,
303 McpConnectionTo<Counterpart>,
304 ) -> BoxFuture<'a, Result<Ret, Error>>
305 + Send
306 + 'static,
307) -> (
308 impl McpTool<Counterpart> + 'static,
309 impl RunWithConnectionTo<Counterpart>,
310)
311where
312 Counterpart: Role,
313 P: JsonSchema + DeserializeOwned + 'static + Send,
314 Ret: JsonSchema + Serialize + 'static + Send,
315 F: AsyncFnMut(P, McpConnectionTo<Counterpart>) -> Result<Ret, Error> + Send,
316{
317 let (call_tx, call_rx) = mpsc::channel(128);
318 (
319 ToolFnTool {
320 name: name.to_string(),
321 description: description.to_string(),
322 call_tx,
323 },
324 ToolFnMutRunner {
325 func,
326 call_rx,
327 tool_future_fn: Box::new(tool_future_fn),
328 },
329 )
330}
331
332pub fn tool_fn<P, Ret, F, Counterpart>(
334 name: impl ToString,
335 description: impl ToString,
336 func: F,
337 tool_future_fn: impl for<'a> Fn(
338 &'a F,
339 P,
340 McpConnectionTo<Counterpart>,
341 ) -> BoxFuture<'a, Result<Ret, Error>>
342 + Send
343 + Sync
344 + 'static,
345) -> (
346 impl McpTool<Counterpart> + 'static,
347 impl RunWithConnectionTo<Counterpart>,
348)
349where
350 Counterpart: Role,
351 P: JsonSchema + DeserializeOwned + 'static + Send,
352 Ret: JsonSchema + Serialize + 'static + Send,
353 F: AsyncFn(P, McpConnectionTo<Counterpart>) -> Result<Ret, Error> + Send + Sync + 'static,
354{
355 let (call_tx, call_rx) = mpsc::channel(128);
356 (
357 ToolFnTool {
358 name: name.to_string(),
359 description: description.to_string(),
360 call_tx,
361 },
362 ToolFnRunner {
363 func,
364 call_rx,
365 tool_future_fn: Box::new(tool_future_fn),
366 },
367 )
368}
369
370#[cfg(test)]
371mod tests {
372 use std::{
373 pin::Pin,
374 sync::Mutex,
375 task::{Context, Poll},
376 };
377
378 use futures::FutureExt as _;
379
380 use super::*;
381 use crate::{Channel, mcp_server::McpConnectionContext, role::mcp};
382
383 type ResultReceiver = CallResult<u32, u32, mcp::Client>;
384
385 #[derive(Default)]
386 struct State {
387 entered: Mutex<Vec<u32>>,
388 dropped: Mutex<Vec<u32>>,
389 discard_result: Mutex<Option<ResultReceiver>>,
390 }
391
392 struct UserFuture<'a> {
394 state: &'a State,
395 id: u32,
396 }
397
398 impl Future for UserFuture<'_> {
399 type Output = Result<u32, Error>;
400
401 fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
402 if self.id == 0 {
403 Poll::Pending
404 } else {
405 drop(self.state.discard_result.lock().unwrap().take());
409 Poll::Ready(Ok(self.id))
410 }
411 }
412 }
413
414 impl Drop for UserFuture<'_> {
415 fn drop(&mut self) {
416 self.state.dropped.lock().unwrap().push(self.id);
417 }
418 }
419
420 #[derive(Clone, Copy)]
421 enum Mode {
422 Mutable,
423 Concurrent,
424 }
425
426 fn runner(
427 mode: Mode,
428 state: &State,
429 call_rx: mpsc::Receiver<QueuedCall<u32, u32, mcp::Client>>,
430 connection: ConnectionTo<mcp::Client>,
431 ) -> BoxFuture<'_, Result<(), Error>> {
432 match mode {
433 Mode::Mutable => Box::pin(
434 ToolFnMutRunner {
435 func: (state, Vec::<u32>::new()),
438 call_rx,
439 tool_future_fn: Box::new(|func, id, _connection| {
440 func.0.entered.lock().unwrap().push(id);
441 Box::pin(async move {
442 let result = UserFuture { state: func.0, id }.await;
443 func.1.push(id);
444 result
445 })
446 }),
447 }
448 .run_with_connection_to(connection),
449 ),
450 Mode::Concurrent => Box::pin(
451 ToolFnRunner {
452 func: state,
453 call_rx,
454 tool_future_fn: Box::new(|state, id, _connection| {
455 state.entered.lock().unwrap().push(id);
457 Box::pin(UserFuture { state, id })
458 }),
459 }
460 .run_with_connection_to(connection),
461 ),
462 }
463 }
464
465 async fn enqueue(
466 tool: &ToolFnTool<u32, u32, mcp::Client>,
467 id: u32,
468 connection: &McpConnectionTo<mcp::Client>,
469 ) -> ResultReceiver {
470 let (result_tx, result_rx) = oneshot::channel();
471 let (done_tx, done_rx) = oneshot::channel();
472 drop(done_rx);
473 let call = CallResult {
474 result_rx,
475 queued_call: QueuedCall(Arc::new(Mutex::new(Some(ToolCall {
476 params: id,
477 mcp_connection: connection.clone(),
478 result_tx,
479 done_tx,
480 })))),
481 };
482 tool.call_tx
485 .clone()
486 .send(call.queued_call.share())
487 .await
488 .unwrap();
489 call
490 }
491
492 fn assert_pending(future: impl Future) {
493 assert!(future.now_or_never().is_none());
494 }
495
496 #[derive(Clone, Copy)]
497 enum Case {
498 Running,
499 Queued,
500 DeliveryRace,
501 ConcurrentProgress,
502 }
503
504 fn check(mode: Mode, case: Case) {
505 let (channel, _peer) = Channel::duplex();
506 futures::executor::block_on(mcp::Server.builder().connect_with(
507 channel,
508 async |connection| {
509 let context = McpConnectionTo {
510 context: McpConnectionContext::Standalone,
511 connection: connection.clone(),
512 #[cfg(feature = "unstable_mcp_over_acp")]
513 cleanup: None,
514 };
515 let state = State::default();
516 let (call_tx, call_rx) = mpsc::channel(128);
517 let tool = ToolFnTool {
518 name: "test".into(),
519 description: "test".into(),
520 call_tx,
521 };
522 let mut runner = runner(mode, &state, call_rx, connection);
523
524 match case {
525 Case::Running | Case::ConcurrentProgress => {
526 let mut first = Box::pin(tool.call_tool(0, context.clone()));
527 assert_pending(first.as_mut());
528 assert_pending(runner.as_mut());
529 assert_eq!(*state.entered.lock().unwrap(), [0]);
530 assert!(state.dropped.lock().unwrap().is_empty());
531
532 if matches!(case, Case::ConcurrentProgress) {
533 let mut next = enqueue(&tool, 1, &context).await;
534 assert_pending(runner.as_mut());
535 assert_eq!(next.result_rx.try_recv().unwrap().unwrap().unwrap(), 1);
536 assert_eq!(*state.dropped.lock().unwrap(), [1]);
539 }
540
541 drop(first);
542 assert_pending(runner.as_mut());
543 assert!(state.dropped.lock().unwrap().contains(&0));
544 }
545 Case::Queued => {
546 drop(enqueue(&tool, 2, &context).await);
547 assert!(state.entered.lock().unwrap().is_empty());
548
549 let first = enqueue(&tool, 0, &context).await;
550 assert_pending(runner.as_mut());
551 assert_eq!(*state.entered.lock().unwrap(), [0]);
552 assert!(state.dropped.lock().unwrap().is_empty());
553
554 let queued_context = context.clone();
558 #[cfg(feature = "unstable_mcp_over_acp")]
559 let queued_context = McpConnectionTo {
560 cleanup: Some(Arc::new(Mutex::new(Vec::new()))),
561 ..queued_context
562 };
563 let mut queued = Box::pin(tool.call_tool(3, queued_context.clone()));
564 assert_pending(queued.as_mut());
565 drop(queued);
566 #[cfg(feature = "unstable_mcp_over_acp")]
567 queued_context.wait_cleanup().now_or_never().unwrap();
568 assert_eq!(*state.entered.lock().unwrap(), [0]);
569 assert!(state.dropped.lock().unwrap().is_empty());
570 drop(first);
571 assert_pending(runner.as_mut());
572 assert_eq!(*state.entered.lock().unwrap(), [0]);
573 assert_eq!(*state.dropped.lock().unwrap(), [0]);
574 }
575 Case::DeliveryRace => {
576 let receiver = enqueue(&tool, 4, &context).await;
577 *state.discard_result.lock().unwrap() = Some(receiver);
578 assert_pending(runner.as_mut());
579 assert!(state.discard_result.lock().unwrap().is_none());
580 assert_eq!(*state.entered.lock().unwrap(), [4]);
581 assert_eq!(*state.dropped.lock().unwrap(), [4]);
582 }
583 }
584
585 let mut next = Box::pin(tool.call_tool(5, context));
587 assert_pending(next.as_mut());
588 assert_pending(runner.as_mut());
589 assert_eq!(next.now_or_never().unwrap().unwrap(), 5);
590 assert_eq!(state.entered.lock().unwrap().last(), Some(&5));
591 assert_eq!(state.dropped.lock().unwrap().last(), Some(&5));
592 drop(tool);
593 runner.now_or_never().unwrap().unwrap();
594 Ok(())
595 },
596 ))
597 .unwrap();
598 }
599
600 #[test]
601 fn mutable_running_cancellation_drops_user_future_and_allows_next_call() {
602 check(Mode::Mutable, Case::Running);
603 }
604
605 #[test]
606 fn concurrent_running_cancellation_drops_user_future_and_allows_next_call() {
607 check(Mode::Concurrent, Case::Running);
608 }
609
610 #[test]
611 fn mutable_cancelled_queued_calls_never_enter_closure() {
612 check(Mode::Mutable, Case::Queued);
613 }
614
615 #[test]
616 fn concurrent_cancelled_queued_calls_never_enter_closure() {
617 check(Mode::Concurrent, Case::Queued);
618 }
619
620 #[test]
621 fn mutable_failed_result_delivery_does_not_stop_runner() {
622 check(Mode::Mutable, Case::DeliveryRace);
623 }
624
625 #[test]
626 fn concurrent_failed_result_delivery_does_not_stop_runner() {
627 check(Mode::Concurrent, Case::DeliveryRace);
628 }
629
630 #[test]
631 fn concurrent_borrowed_futures_make_independent_progress() {
632 check(Mode::Concurrent, Case::ConcurrentProgress);
633 }
634
635 #[cfg(feature = "unstable_mcp_over_acp")]
636 #[derive(serde::Deserialize, JsonSchema)]
637 struct DropParams {
638 #[serde(skip)]
639 #[schemars(skip)]
640 on_drop: Option<Box<dyn FnOnce() + Send>>,
641 }
642
643 #[cfg(feature = "unstable_mcp_over_acp")]
644 impl Drop for DropParams {
645 fn drop(&mut self) {
646 if let Some(on_drop) = self.on_drop.take() {
647 on_drop();
648 }
649 }
650 }
651
652 #[cfg(feature = "unstable_mcp_over_acp")]
653 #[test]
654 fn cancellation_during_enqueue_destroys_payload_before_cleanup_ack() {
655 let (channel, _peer) = Channel::duplex();
656 futures::executor::block_on(mcp::Server.builder().connect_with(
657 channel,
658 async |connection| {
659 type QueueState = Mutex<Option<ToolCall<DropParams, u32, mcp::Client>>>;
660
661 let cleanup = Arc::new(Mutex::new(Vec::<oneshot::Receiver<()>>::new()));
662 let context = McpConnectionTo {
663 context: McpConnectionContext::Standalone,
664 connection,
665 cleanup: Some(cleanup.clone()),
666 };
667 let (call_tx, mut call_rx) = mpsc::channel(0);
670 let tool = ToolFnTool::<DropParams, u32, _> {
671 name: "test".into(),
672 description: "test".into(),
673 call_tx,
674 };
675 let dropped = Arc::new(Mutex::new(false));
676 let queue_state = Arc::new(Mutex::new(None::<std::sync::Weak<QueueState>>));
677 let params = DropParams {
678 on_drop: Some(Box::new({
679 let dropped = dropped.clone();
680 let cleanup = cleanup.clone();
681 let queue_state = queue_state.clone();
682 move || {
683 assert!(
684 cleanup.lock().unwrap()[0].try_recv().unwrap().is_none(),
685 "cleanup ack preceded queued params destructor"
686 );
687 let queue = queue_state
688 .lock()
689 .unwrap()
690 .as_ref()
691 .unwrap()
692 .upgrade()
693 .unwrap();
694 assert!(
695 queue.try_lock().is_ok(),
696 "params destructor ran under queue lock"
697 );
698 *dropped.lock().unwrap() = true;
699 }
700 })),
701 };
702 let mut call = Box::pin(tool.call_tool(params, context.clone()));
703 assert_pending(call.as_mut());
704 assert_eq!(cleanup.lock().unwrap().len(), 1);
705 let queued = call_rx.next().now_or_never().unwrap().unwrap();
706 *queue_state.lock().unwrap() = Some(Arc::downgrade(&queued.0));
707 drop(call);
710 assert!(*dropped.lock().unwrap());
711 assert!(
712 queued.take().is_none(),
713 "cancelled payload retained by queue"
714 );
715 assert_eq!(
716 Arc::strong_count(&cleanup),
717 2,
718 "queued host context survived cleanup acknowledgment"
719 );
720 context.wait_cleanup().now_or_never().unwrap();
721 Ok(())
722 },
723 ))
724 .unwrap();
725 }
726
727 #[test]
728 fn running_cleanup_ack_follows_actual_future_destructor() {
729 struct PendingUser {
730 done: Arc<Mutex<oneshot::Receiver<()>>>,
731 dropped: Arc<Mutex<bool>>,
732 }
733 impl Future for PendingUser {
734 type Output = Result<(), Error>;
735 fn poll(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Self::Output> {
736 Poll::Pending
737 }
738 }
739 impl Drop for PendingUser {
740 fn drop(&mut self) {
741 assert!(self.done.lock().unwrap().try_recv().unwrap().is_none());
742 *self.dropped.lock().unwrap() = true;
743 }
744 }
745 let (result_tx, result_rx) = oneshot::channel();
746 let (done_tx, done_rx) = oneshot::channel();
747 let done = Arc::new(Mutex::new(done_rx));
748 let dropped = Arc::new(Mutex::new(false));
749 let mut call = Box::pin(run_call(
750 PendingUser {
751 done: done.clone(),
752 dropped: dropped.clone(),
753 },
754 result_tx,
755 done_tx,
756 ));
757 assert_pending(call.as_mut());
758 drop(result_rx);
759 call.now_or_never().unwrap();
760 assert!(*dropped.lock().unwrap());
761 assert_eq!(done.lock().unwrap().try_recv().unwrap(), Some(()));
762 }
763}