agent_client_protocol/jsonrpc/
run.rs1use 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#[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
76pub trait RunWithConnectionTo<Counterpart: Role>: Send {
84 fn run_with_connection_to(
86 self,
87 cx: ConnectionTo<Counterpart>,
88 ) -> impl Future<Output = Result<(), crate::Error>> + Send;
89}
90
91#[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#[derive(Debug)]
106pub struct ChainRun<A, B> {
107 a: A,
108 b: B,
109}
110
111impl<A, B> ChainRun<A, B> {
112 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 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 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
158pub 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 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}