Skip to main content

agent_client_protocol/jsonrpc/
run.rs

1//! Run trait for background tasks that run alongside a connection.
2//!
3//! Run implementations are composable background tasks that run while a connection is active.
4//! They're used for things like MCP tool handlers that need to receive calls through
5//! channels and invoke user-provided closures.
6
7use std::future::Future;
8use std::marker::PhantomData;
9use std::sync::{Arc, Mutex};
10
11use futures::FutureExt;
12use futures::future::{Either, select};
13
14use crate::{
15    ConnectionTo,
16    jsonrpc::{ConnectionContext, RawConnectionContext, connection_context},
17    role::Role,
18};
19
20/// Cleanup policy installed by the boundary that owns a runner composition.
21/// The first error is visible even while a chain retains its sibling for cleanup.
22#[derive(Clone)]
23pub(crate) struct RunnerErrorScope {
24    error: Arc<Mutex<Option<crate::Error>>>,
25    close: Arc<dyn Fn() + Send + Sync>,
26    cleanup: futures::future::Shared<futures::future::BoxFuture<'static, ()>>,
27}
28
29impl std::fmt::Debug for RunnerErrorScope {
30    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31        formatter
32            .debug_struct("RunnerErrorScope")
33            .finish_non_exhaustive()
34    }
35}
36
37impl RunnerErrorScope {
38    pub(crate) fn new(
39        close: impl Fn() + Send + Sync + 'static,
40        cleanup: impl Future<Output = ()> + Send + 'static,
41    ) -> Self {
42        Self {
43            error: Arc::default(),
44            close: Arc::new(close),
45            cleanup: cleanup.boxed().shared(),
46        }
47    }
48
49    pub(crate) fn error(&self) -> Option<crate::Error> {
50        self.error
51            .lock()
52            .expect("runner error mutex poisoned")
53            .clone()
54    }
55
56    pub(crate) fn finish(
57        &self,
58        error: crate::Error,
59    ) -> futures::future::Shared<futures::future::BoxFuture<'static, ()>> {
60        let first_error = {
61            let mut stored = self.error.lock().expect("runner error mutex poisoned");
62            if stored.is_none() {
63                *stored = Some(error);
64                true
65            } else {
66                false
67            }
68        };
69        if first_error {
70            (self.close)();
71        }
72        self.cleanup.clone()
73    }
74}
75
76/// A background task that runs alongside a connection.
77///
78/// `RunIn<R>` means "run in the context of being role R". The task receives
79/// a `ConnectionTo<R::Counterpart>` for communicating with the other side.
80///
81/// Implementations are composed using [`ChainRun`] and run in parallel
82/// when the connection is active.
83pub trait RunWithConnectionTo<Counterpart: Role>: Send {
84    /// Run this task to completion.
85    fn run_with_connection_to(
86        self,
87        cx: ConnectionTo<Counterpart>,
88    ) -> impl Future<Output = Result<(), crate::Error>> + Send;
89}
90
91/// A no-op RunIn that completes immediately.
92#[derive(Debug, Default)]
93pub struct NullRun;
94
95impl<Counterpart: Role> RunWithConnectionTo<Counterpart> for NullRun {
96    fn run_with_connection_to(
97        self,
98        _cx: ConnectionTo<Counterpart>,
99    ) -> impl Future<Output = Result<(), crate::Error>> + Send {
100        std::future::ready(Ok(()))
101    }
102}
103
104/// Chains two RunIn implementations to run in parallel.
105#[derive(Debug)]
106pub struct ChainRun<A, B> {
107    a: A,
108    b: B,
109}
110
111impl<A, B> ChainRun<A, B> {
112    /// Create a new chained RunIn from two RunIn implementations.
113    pub fn new(a: A, b: B) -> Self {
114        Self { a, b }
115    }
116}
117
118impl<Counterpart: Role, A, B> RunWithConnectionTo<Counterpart> for ChainRun<A, B>
119where
120    A: RunWithConnectionTo<Counterpart>,
121    B: RunWithConnectionTo<Counterpart>,
122{
123    async fn run_with_connection_to(
124        self,
125        cx: ConnectionTo<Counterpart>,
126    ) -> Result<(), crate::Error> {
127        // Box the futures to avoid stack overflow with deeply nested RunIn chains
128        let a_fut = Box::pin(self.a.run_with_connection_to(cx.clone()));
129        let b_fut = Box::pin(self.b.run_with_connection_to(cx.clone()));
130        match select(a_fut, b_fut).await {
131            Either::Left((Ok(()), b)) => b.await,
132            Either::Right((Ok(()), a)) => a.await,
133            Either::Left((Err(error), b)) => {
134                finish_runner_after_error(b, &cx, error.clone()).await;
135                Err(error)
136            }
137            Either::Right((Err(error), a)) => {
138                finish_runner_after_error(a, &cx, error.clone()).await;
139                Err(error)
140            }
141        }
142    }
143}
144
145async fn finish_runner_after_error<R: Role>(
146    runner: impl Future<Output = Result<(), crate::Error>>,
147    cx: &ConnectionTo<R>,
148    error: crate::Error,
149) {
150    // A sibling runner can own the actual scoped operation. Keep it polled
151    // through the owning boundary's cleanup, never arbitrary infinite user work.
152    match select(Box::pin(runner), Box::pin(cx.finish_runner_error(error))).await {
153        Either::Left((_, cleanup)) => cleanup.await,
154        Either::Right(((), _)) => {}
155    }
156}
157
158/// A RunIn created from a closure via [`with_spawned`](crate::Builder::with_spawned).
159pub struct SpawnedRun<F, Context = RawConnectionContext> {
160    task_fn: F,
161    location: &'static std::panic::Location<'static>,
162    context: PhantomData<fn() -> Context>,
163}
164
165impl<F, Context> SpawnedRun<F, Context> {
166    /// Create a new spawned RunIn from a closure.
167    pub fn new(location: &'static std::panic::Location<'static>, task_fn: F) -> Self {
168        Self {
169            task_fn,
170            location,
171            context: PhantomData,
172        }
173    }
174}
175
176impl<Counterpart, F, Fut, Context> RunWithConnectionTo<Counterpart> for SpawnedRun<F, Context>
177where
178    Counterpart: Role,
179    Context: ConnectionContext,
180    F: FnOnce(Context::Connection<Counterpart>) -> Fut + Send,
181    Fut: Future<Output = Result<(), crate::Error>> + Send,
182{
183    async fn run_with_connection_to(
184        self,
185        connection: ConnectionTo<Counterpart>,
186    ) -> Result<(), crate::Error> {
187        let location = self.location;
188        (self.task_fn)(connection_context::from_raw::<Context, _>(connection))
189            .await
190            .map_err(|err| {
191                let data = err.data.clone();
192                err.data(serde_json::json!({
193                    "spawned_at": format!("{}:{}:{}", location.file(), location.line(), location.column()),
194                    "data": data,
195                }))
196            })
197    }
198}
199
200#[cfg(test)]
201mod tests {
202    use super::RunnerErrorScope;
203    use std::sync::{
204        Arc,
205        atomic::{AtomicUsize, Ordering},
206    };
207
208    #[test]
209    fn repeated_runner_errors_close_and_join_scope_once() {
210        let closes = Arc::new(AtomicUsize::new(0));
211        let joins = Arc::new(AtomicUsize::new(0));
212        let close_count = closes.clone();
213        let join_count = joins.clone();
214        let scope = RunnerErrorScope::new(
215            move || {
216                close_count.fetch_add(1, Ordering::SeqCst);
217            },
218            async move {
219                join_count.fetch_add(1, Ordering::SeqCst);
220            },
221        );
222        let first_error = crate::Error::invalid_params().data("first");
223        let first = scope.finish(first_error.clone());
224        let second = scope.clone().finish(crate::Error::internal_error());
225
226        assert_eq!(scope.error(), Some(first_error));
227        assert_eq!(closes.load(Ordering::SeqCst), 1);
228        futures::executor::block_on(futures::future::join(first, second));
229        assert_eq!(joins.load(Ordering::SeqCst), 1);
230    }
231}