1use std::future::Future;
20use std::pin::Pin;
21use std::sync::Arc;
22
23use bevy_ecs::entity::Entity;
24use tokio::runtime::Handle;
25use tokio::sync::Mutex;
26use tokio::sync::Notify;
27use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender};
28use tokio::task::JoinHandle;
29
30pub type ToolExecFuture = Pin<Box<dyn Future<Output = Vec<(String, String)>> + Send>>;
34
35pub type BoxedToolExec = Box<dyn FnOnce() -> ToolExecFuture + Send>;
39
40pub struct ToolJob {
42 pub entity: Entity,
44 pub exec: BoxedToolExec,
46 pub cancel: crate::cancel::CancelToken,
49}
50
51pub struct ToolOutcome {
54 pub entity: Entity,
56 pub results: Vec<(String, String)>,
58 pub elapsed: std::time::Duration,
62}
63
64pub type SharedJobRx = Arc<Mutex<UnboundedReceiver<ToolJob>>>;
67
68pub async fn tool_worker(
74 jobs: SharedJobRx,
75 results: UnboundedSender<ToolOutcome>,
76 wake: Arc<Notify>,
77) {
78 loop {
79 let next = {
81 let mut rx = jobs.lock().await;
82 rx.recv().await
83 };
84 let Some(ToolJob {
85 entity,
86 exec,
87 cancel,
88 }) = next
89 else {
90 return; };
92 let started = std::time::Instant::now();
99 let out = tokio::select! {
100 biased;
101 _ = cancel.cancelled() => continue,
102 out = exec() => out,
103 };
104 let _ = results.send(ToolOutcome {
106 entity,
107 results: out,
108 elapsed: started.elapsed(),
109 });
110 wake.notify_one();
111 }
112}
113
114pub fn spawn_tool_pool(
118 runtime: &Handle,
119 jobs: UnboundedReceiver<ToolJob>,
120 results: UnboundedSender<ToolOutcome>,
121 wake: Arc<Notify>,
122 workers: usize,
123) -> Vec<JoinHandle<()>> {
124 let shared: SharedJobRx = Arc::new(Mutex::new(jobs));
125 (0..workers.max(1))
126 .map(|_| runtime.spawn(tool_worker(shared.clone(), results.clone(), wake.clone())))
127 .collect()
128}
129
130#[cfg(test)]
131mod tests {
132 use super::*;
133 use tokio::sync::mpsc;
134
135 fn job(entity: u32, pairs: Vec<(&'static str, &'static str)>) -> ToolJob {
136 job_with(entity, pairs, crate::cancel::CancelToken::new())
137 }
138
139 fn job_with(
140 entity: u32,
141 pairs: Vec<(&'static str, &'static str)>,
142 cancel: crate::cancel::CancelToken,
143 ) -> ToolJob {
144 ToolJob {
145 entity: Entity::from_raw_u32(entity).expect("index came from a live entity id"),
146 exec: Box::new(move || {
147 Box::pin(async move {
148 pairs
149 .into_iter()
150 .map(|(a, b)| (a.to_string(), b.to_string()))
151 .collect()
152 })
153 }),
154 cancel,
155 }
156 }
157
158 fn held_job(
162 entity: u32,
163 started: Arc<Notify>,
164 release: Arc<Notify>,
165 cancel: crate::cancel::CancelToken,
166 ) -> ToolJob {
167 ToolJob {
168 entity: Entity::from_raw_u32(entity).expect("index came from a live entity id"),
169 exec: Box::new(move || {
170 Box::pin(async move {
171 started.notify_one();
176 release.notified().await;
177 vec![("held".to_string(), "done".to_string())]
178 })
179 }),
180 cancel,
181 }
182 }
183
184 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
187 async fn a_released_batch_completes_normally() {
188 let (jtx, jrx) = mpsc::unbounded_channel();
189 let (rtx, mut rrx) = mpsc::unbounded_channel();
190 let wake = Arc::new(Notify::new());
191 let started = Arc::new(Notify::new());
192 let release = Arc::new(Notify::new());
193
194 jtx.send(held_job(
195 7,
196 started.clone(),
197 release.clone(),
198 crate::cancel::CancelToken::new(),
199 ))
200 .unwrap();
201 drop(jtx);
202
203 let handles = spawn_tool_pool(&Handle::current(), jrx, rtx, wake, 1);
204 tokio::time::timeout(std::time::Duration::from_secs(5), started.notified())
205 .await
206 .expect("the batch started");
207 release.notify_one();
208
209 let out = tokio::time::timeout(std::time::Duration::from_secs(5), rrx.recv())
210 .await
211 .expect("the batch finished")
212 .expect("an outcome arrived");
213 assert_eq!(out.results, vec![("held".to_string(), "done".to_string())]);
214 for h in handles {
215 let _ = h.await;
216 }
217 }
218
219 fn shared(rx: UnboundedReceiver<ToolJob>) -> SharedJobRx {
220 Arc::new(Mutex::new(rx))
221 }
222
223 #[tokio::test]
224 async fn worker_processes_jobs_in_order_then_exits_on_close() {
225 let (jtx, jrx) = mpsc::unbounded_channel();
226 let (rtx, mut rrx) = mpsc::unbounded_channel();
227 let wake = Arc::new(Notify::new());
228
229 jtx.send(job(1, vec![("c1", "r1")])).unwrap();
230 jtx.send(job(2, vec![("c2", "r2")])).unwrap();
231 drop(jtx); tool_worker(shared(jrx), rtx, wake).await;
234
235 let first = rrx.try_recv().unwrap();
236 assert_eq!(
237 first.entity,
238 Entity::from_raw_u32(1).expect("a small literal index is always a valid entity id")
239 );
240 assert_eq!(first.results, vec![("c1".to_string(), "r1".to_string())]);
241 let second = rrx.try_recv().unwrap();
242 assert_eq!(
243 second.entity,
244 Entity::from_raw_u32(2).expect("a small literal index is always a valid entity id")
245 );
246 assert!(rrx.try_recv().is_err()); }
248
249 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
257 async fn a_cancelled_batch_is_abandoned_and_frees_the_worker() {
258 let (jtx, jrx) = mpsc::unbounded_channel();
259 let (rtx, mut rrx) = mpsc::unbounded_channel();
260 let wake = Arc::new(Notify::new());
261 let cancel = crate::cancel::CancelToken::new();
262
263 let started = Arc::new(Notify::new());
265 let release = Arc::new(Notify::new());
266 jtx.send(held_job(
267 1,
268 started.clone(),
269 release.clone(),
270 cancel.clone(),
271 ))
272 .unwrap();
273 jtx.send(job(2, vec![("c2", "r2")])).unwrap();
275 drop(jtx);
276
277 let handles = spawn_tool_pool(&Handle::current(), jrx, rtx, wake, 1);
278 tokio::time::timeout(std::time::Duration::from_secs(5), started.notified())
280 .await
281 .expect("the batch started");
282 cancel.cancel();
283
284 let next = tokio::time::timeout(std::time::Duration::from_secs(5), rrx.recv())
286 .await
287 .expect("the worker was freed by the cancel")
288 .expect("an outcome arrived");
289 assert_eq!(
290 next.entity,
291 Entity::from_raw_u32(2).expect("a small literal index is always a valid entity id"),
292 "the queued batch ran"
293 );
294 assert!(
295 rrx.try_recv().is_err(),
296 "and the cancelled batch reported no results"
297 );
298 for h in handles {
299 let _ = h.await;
300 }
301 }
302
303 #[tokio::test]
304 async fn worker_survives_dropped_results_receiver() {
305 let (jtx, jrx) = mpsc::unbounded_channel();
306 let (rtx, rrx) = mpsc::unbounded_channel();
307 drop(rrx); let wake = Arc::new(Notify::new());
309
310 jtx.send(job(9, vec![("c", "r")])).unwrap();
311 drop(jtx);
312 tool_worker(shared(jrx), rtx, wake).await;
314 }
315
316 #[tokio::test(flavor = "multi_thread", worker_threads = 3)]
317 async fn pool_runs_batches_concurrently_and_all_exit_on_close() {
318 use std::sync::atomic::{AtomicUsize, Ordering};
319 use std::time::Duration;
320
321 let (jtx, jrx) = mpsc::unbounded_channel();
322 let (rtx, mut rrx) = mpsc::unbounded_channel();
323 let wake = Arc::new(Notify::new());
324
325 let arrived = Arc::new(AtomicUsize::new(0));
329 let go = Arc::new(Notify::new());
330 for i in 1..=3u32 {
331 let arrived = arrived.clone();
332 let go = go.clone();
333 jtx.send(ToolJob {
334 entity: Entity::from_raw_u32(i).expect("index came from a live entity id"),
335 exec: Box::new(move || {
336 Box::pin(async move {
337 if arrived.fetch_add(1, Ordering::SeqCst) + 1 == 3 {
338 go.notify_waiters();
339 }
340 while arrived.load(Ordering::SeqCst) < 3 {
342 go.notified().await;
343 }
344 vec![("c".to_string(), "r".to_string())]
345 })
346 }),
347 cancel: crate::cancel::CancelToken::new(),
348 })
349 .unwrap();
350 }
351 drop(jtx);
352
353 let handles = spawn_tool_pool(&Handle::current(), jrx, rtx, wake, 3);
354 for _ in 0..3 {
356 tokio::time::timeout(Duration::from_secs(5), rrx.recv())
357 .await
358 .expect("all batches complete concurrently")
359 .expect("outcome present");
360 }
361 for h in handles {
362 h.await.unwrap(); }
364 }
365
366 #[tokio::test(flavor = "multi_thread")]
367 async fn spawn_tool_pool_clamps_zero_to_one_worker() {
368 let (jtx, jrx) = mpsc::unbounded_channel();
369 let (rtx, mut rrx) = mpsc::unbounded_channel();
370 let wake = Arc::new(Notify::new());
371 jtx.send(job(7, vec![("c", "r")])).unwrap();
372 drop(jtx);
373 let handles = spawn_tool_pool(&Handle::current(), jrx, rtx, wake, 0);
374 assert_eq!(handles.len(), 1); let out = rrx.recv().await.unwrap();
376 assert_eq!(
377 out.entity,
378 Entity::from_raw_u32(7).expect("a small literal index is always a valid entity id")
379 );
380 for h in handles {
381 h.await.unwrap();
382 }
383 }
384}