1use std::future::Future;
36use std::pin::Pin;
37use std::sync::Arc;
38use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
39
40use bevy_ecs::entity::Entity;
41use tokio::runtime::Handle;
42use tokio::sync::Notify;
43use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender};
44use tokio::sync::{OwnedSemaphorePermit, Semaphore};
45use tokio::task::{JoinHandle, JoinSet};
46
47use crate::inference_pool::expect_permit;
48
49pub type ToolExecFuture = Pin<Box<dyn Future<Output = Vec<(String, String)>> + Send>>;
53
54pub type BoxedToolExec = Box<dyn FnOnce() -> ToolExecFuture + Send>;
58
59pub struct ToolJob {
61 pub entity: Entity,
63 pub exec: BoxedToolExec,
65 pub cancel: crate::cancel::CancelToken,
68}
69
70pub struct ToolOutcome {
73 pub entity: Entity,
75 pub results: Vec<(String, String)>,
77 pub elapsed: std::time::Duration,
81}
82
83#[derive(Debug)]
90pub struct ToolLaneStats {
91 queued: AtomicUsize,
92 busy: AtomicUsize,
93 parked: AtomicUsize,
94 workers: AtomicUsize,
97}
98
99impl ToolLaneStats {
100 pub fn new(workers: usize) -> Self {
102 Self {
103 queued: AtomicUsize::new(0),
104 busy: AtomicUsize::new(0),
105 parked: AtomicUsize::new(0),
106 workers: AtomicUsize::new(workers.max(1)),
107 }
108 }
109
110 pub fn enqueued(&self) {
112 self.queued.fetch_add(1, Ordering::Relaxed);
113 }
114
115 fn abandoned(&self) {
118 self.queued.fetch_sub(1, Ordering::Relaxed);
119 }
120
121 fn started(&self) {
123 self.queued.fetch_sub(1, Ordering::Relaxed);
124 self.busy.fetch_add(1, Ordering::Relaxed);
125 }
126
127 fn finished(&self) {
129 self.busy.fetch_sub(1, Ordering::Relaxed);
130 }
131
132 fn began_park(&self) {
135 self.busy.fetch_sub(1, Ordering::Relaxed);
136 self.parked.fetch_add(1, Ordering::Relaxed);
137 }
138
139 fn resumed(&self) {
141 self.parked.fetch_sub(1, Ordering::Relaxed);
142 self.busy.fetch_add(1, Ordering::Relaxed);
143 }
144
145 fn ended_park(&self) {
148 self.parked.fetch_sub(1, Ordering::Relaxed);
149 }
150
151 pub fn queued(&self) -> usize {
153 self.queued.load(Ordering::Relaxed)
154 }
155
156 pub fn busy(&self) -> usize {
158 self.busy.load(Ordering::Relaxed)
159 }
160
161 pub fn parked(&self) -> usize {
163 self.parked.load(Ordering::Relaxed)
164 }
165
166 pub fn workers(&self) -> usize {
168 self.workers.load(Ordering::Relaxed)
169 }
170
171 fn widen(&self, extra: usize) {
173 self.workers.fetch_add(extra, Ordering::Relaxed);
174 }
175
176 fn narrowed(&self, taken: usize) {
178 self.workers.fetch_sub(taken, Ordering::Relaxed);
179 }
180
181 #[must_use]
184 pub fn is_saturated(&self) -> bool {
185 self.busy() >= self.workers() && self.queued() > 0
186 }
187}
188
189pub struct ToolLane {
192 permits: Arc<Semaphore>,
194 stats: Arc<ToolLaneStats>,
196 results: UnboundedSender<ToolOutcome>,
198 wake: Arc<Notify>,
201 runtime: Handle,
203}
204
205impl ToolLane {
206 pub fn new(
209 runtime: Handle,
210 results: UnboundedSender<ToolOutcome>,
211 wake: Arc<Notify>,
212 concurrency: usize,
213 stats: Arc<ToolLaneStats>,
214 ) -> Arc<Self> {
215 Arc::new(Self {
216 permits: Arc::new(Semaphore::new(concurrency.max(1))),
217 stats,
218 results,
219 wake,
220 runtime,
221 })
222 }
223
224 pub fn serve(self: &Arc<Self>, jobs: UnboundedReceiver<ToolJob>) -> JoinHandle<()> {
228 let lane = self.clone();
229 self.runtime.clone().spawn(serve_lane(lane, jobs))
230 }
231
232 pub fn relieve(&self, extra: usize) -> usize {
241 if extra == 0 {
242 return 0;
243 }
244 self.permits.add_permits(extra);
245 self.stats.widen(extra);
246 extra
247 }
248
249 pub fn narrow(&self, upto: usize) -> usize {
257 if upto == 0 {
258 return 0;
259 }
260 let taken = self.permits.forget_permits(upto);
261 self.stats.narrowed(taken);
262 taken
263 }
264}
265
266async fn serve_lane(lane: Arc<ToolLane>, mut jobs: UnboundedReceiver<ToolJob>) {
269 let mut batches = JoinSet::new();
270 loop {
271 tokio::select! {
272 job = jobs.recv() => match job {
273 Some(job) => {
274 batches.spawn_on(run_batch(lane.clone(), job), &lane.runtime);
275 }
276 None => break, },
278 Some(_) = batches.join_next(), if !batches.is_empty() => {}
282 }
283 }
284 while batches.join_next().await.is_some() {}
285}
286
287async fn run_batch(lane: Arc<ToolLane>, job: ToolJob) {
290 let ToolJob {
291 entity,
292 exec,
293 cancel,
294 } = job;
295 let permit = tokio::select! {
298 biased;
299 _ = cancel.cancelled() => {
300 lane.stats.abandoned();
301 return;
302 }
303 permit = lane.permits.clone().acquire_owned() => expect_permit(permit),
304 };
305 lane.stats.started();
306 let ticket = Arc::new(LaneTicket::new(lane.clone(), permit));
307 let started = std::time::Instant::now();
308 let out = LANE_TICKET
313 .scope(ticket, async move {
314 tokio::select! {
315 biased;
316 _ = cancel.cancelled() => None,
317 out = exec() => Some(out),
318 }
319 })
320 .await;
321 let Some(out) = out else { return };
324 let _ = lane.results.send(ToolOutcome {
326 entity,
327 results: out,
328 elapsed: started.elapsed(),
329 });
330 lane.wake.notify_one();
331}
332
333tokio::task_local! {
334 static LANE_TICKET: Arc<LaneTicket>;
342}
343
344struct LaneTicket {
350 lane: Arc<ToolLane>,
351 permit: std::sync::Mutex<Option<OwnedSemaphorePermit>>,
353 parked: AtomicBool,
355}
356
357impl LaneTicket {
358 fn new(lane: Arc<ToolLane>, permit: OwnedSemaphorePermit) -> Self {
359 Self {
360 lane,
361 permit: std::sync::Mutex::new(Some(permit)),
362 parked: AtomicBool::new(false),
363 }
364 }
365
366 fn release(&self) {
368 let held = self.take_permit();
369 drop(held);
373 self.parked.store(true, Ordering::Relaxed);
374 self.lane.stats.began_park();
375 self.lane.wake.notify_one();
376 }
377
378 async fn reacquire(&self) {
382 let permit = expect_permit(self.lane.permits.clone().acquire_owned().await);
383 self.lane.stats.resumed();
386 self.parked.store(false, Ordering::Relaxed);
387 *self
388 .permit
389 .lock()
390 .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(permit);
391 }
392
393 fn take_permit(&self) -> Option<OwnedSemaphorePermit> {
394 self.permit
395 .lock()
396 .unwrap_or_else(std::sync::PoisonError::into_inner)
397 .take()
398 }
399}
400
401impl Drop for LaneTicket {
402 fn drop(&mut self) {
403 drop(self.take_permit());
405 match self.parked.load(Ordering::Relaxed) {
406 true => self.lane.stats.ended_park(),
407 false => self.lane.stats.finished(),
408 }
409 self.lane.wake.notify_one();
410 }
411}
412
413pub async fn off_lane<T>(fut: impl Future<Output = T>) -> T {
424 let Ok(ticket) = LANE_TICKET.try_with(Arc::clone) else {
425 return fut.await;
426 };
427 ticket.release();
428 let out = fut.await;
429 ticket.reacquire().await;
430 out
431}
432
433#[cfg(test)]
434mod tests {
435 use super::*;
436 use std::time::Duration;
437 use tokio::sync::mpsc;
438
439 struct Harness {
442 lane: Arc<ToolLane>,
443 jobs: Option<UnboundedSender<ToolJob>>,
445 outcomes: mpsc::UnboundedReceiver<ToolOutcome>,
446 serving: Option<JoinHandle<()>>,
447 stats: Arc<ToolLaneStats>,
448 }
449
450 impl Harness {
451 fn new(concurrency: usize) -> Self {
452 let (jobs, job_rx) = mpsc::unbounded_channel();
453 let (result_tx, outcomes) = mpsc::unbounded_channel();
454 let stats = Arc::new(ToolLaneStats::new(concurrency));
455 let lane = ToolLane::new(
456 Handle::current(),
457 result_tx,
458 Arc::new(Notify::new()),
459 concurrency,
460 stats.clone(),
461 );
462 let serving = lane.serve(job_rx);
463 Self {
464 lane,
465 jobs: Some(jobs),
466 outcomes,
467 serving: Some(serving),
468 stats,
469 }
470 }
471
472 fn submit(&self, job: ToolJob) {
474 self.stats.enqueued();
475 self.sender().send(job).expect("the lane is serving");
476 }
477
478 fn sender(&self) -> &UnboundedSender<ToolJob> {
479 self.jobs.as_ref().expect("the lane is still open")
480 }
481
482 async fn drain(&mut self) {
484 drop(self.jobs.take());
485 let serving = self.serving.take().expect("the lane was serving");
486 timeout(serving).await.expect("the lane task ended");
487 }
488
489 async fn next_outcome(&mut self) -> ToolOutcome {
490 timeout(self.outcomes.recv())
491 .await
492 .expect("an outcome arrived")
493 }
494
495 async fn next_indices(&mut self, n: usize) -> Vec<u64> {
501 let mut seen = Vec::new();
502 for _ in 0..n {
503 seen.push(self.next_outcome().await.entity.to_bits());
504 }
505 seen.sort_unstable();
506 seen
507 }
508 }
509
510 async fn timeout<T>(fut: impl Future<Output = T>) -> T {
513 tokio::time::timeout(Duration::from_secs(30), fut)
514 .await
515 .expect("the lane made progress")
516 }
517
518 fn sorted_bits(entities: &[Entity]) -> Vec<u64> {
521 let mut bits: Vec<u64> = entities.iter().map(|e| e.to_bits()).collect();
522 bits.sort_unstable();
523 bits
524 }
525
526 fn entity(index: u32) -> Entity {
527 Entity::from_raw_u32(index).expect("a small literal index is a valid entity id")
528 }
529
530 fn job(index: u32, pairs: Vec<(&'static str, &'static str)>) -> ToolJob {
531 job_with(index, pairs, crate::cancel::CancelToken::new())
532 }
533
534 fn job_with(
535 index: u32,
536 pairs: Vec<(&'static str, &'static str)>,
537 cancel: crate::cancel::CancelToken,
538 ) -> ToolJob {
539 ToolJob {
540 entity: entity(index),
541 exec: Box::new(move || {
542 Box::pin(async move {
543 pairs
544 .into_iter()
545 .map(|(a, b)| (a.to_string(), b.to_string()))
546 .collect()
547 })
548 }),
549 cancel,
550 }
551 }
552
553 fn held_job(
560 index: u32,
561 started: Arc<Notify>,
562 release: Arc<Notify>,
563 cancel: crate::cancel::CancelToken,
564 ) -> ToolJob {
565 ToolJob {
566 entity: entity(index),
567 exec: Box::new(move || {
568 Box::pin(async move {
569 started.notify_one();
570 release.notified().await;
571 vec![("held".to_string(), "done".to_string())]
572 })
573 }),
574 cancel,
575 }
576 }
577
578 fn parking_job(
585 index: u32,
586 started: Arc<Notify>,
587 release: Arc<Notify>,
588 cancel: crate::cancel::CancelToken,
589 ) -> ToolJob {
590 ToolJob {
591 entity: entity(index),
592 exec: Box::new(move || {
593 Box::pin(async move {
594 off_lane(async move {
595 started.notify_one();
596 release.notified().await;
597 })
598 .await;
599 vec![("parked".to_string(), "done".to_string())]
600 })
601 }),
602 cancel,
603 }
604 }
605
606 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
607 async fn the_lane_runs_batches_and_reports_them() {
608 let mut h = Harness::new(1);
609 h.submit(job(1, vec![("c", "r")]));
610 h.submit(job(2, vec![("c", "r")]));
611
612 let first = h.next_outcome().await;
613 assert_eq!(
614 first.results,
615 vec![("c".to_string(), "r".to_string())],
616 "the batch reported its call"
617 );
618 let mut seen = vec![first.entity.to_bits()];
619 seen.extend(h.next_indices(1).await);
620 seen.sort_unstable();
621 assert_eq!(
622 seen,
623 sorted_bits(&[entity(1), entity(2)]),
624 "both batches were reported"
625 );
626
627 h.drain().await;
628 assert!(h.outcomes.try_recv().is_err(), "no more outcomes");
629 }
630
631 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
638 async fn a_parked_batch_lets_the_batch_it_waits_on_run() {
639 let mut h = Harness::new(1);
640 let started = Arc::new(Notify::new());
641 let release = Arc::new(Notify::new());
642
643 h.submit(parking_job(
644 1,
645 started.clone(),
646 release.clone(),
647 crate::cancel::CancelToken::new(),
648 ));
649 timeout(started.notified()).await;
650 assert_eq!(
651 (h.stats.busy(), h.stats.parked()),
652 (0, 1),
653 "the waiter gave the lane back"
654 );
655
656 let releaser = release.clone();
659 h.submit(ToolJob {
660 entity: entity(2),
661 exec: Box::new(move || {
662 Box::pin(async move {
663 releaser.notify_one();
664 vec![("c2".to_string(), "r2".to_string())]
665 })
666 }),
667 cancel: crate::cancel::CancelToken::new(),
668 });
669
670 assert_eq!(
671 h.next_indices(2).await,
672 sorted_bits(&[entity(1), entity(2)]),
673 "both batches finished"
674 );
675
676 h.drain().await;
677 }
678
679 #[tokio::test(flavor = "multi_thread", worker_threads = 3)]
682 async fn a_resumed_batch_takes_a_permit_again() {
683 let mut h = Harness::new(1);
684 let started = Arc::new(Notify::new());
685 let release = Arc::new(Notify::new());
686 h.submit(parking_job(
687 1,
688 started.clone(),
689 release.clone(),
690 crate::cancel::CancelToken::new(),
691 ));
692 timeout(started.notified()).await;
693
694 let held_started = Arc::new(Notify::new());
696 let held_release = Arc::new(Notify::new());
697 h.submit(held_job(
698 2,
699 held_started.clone(),
700 held_release.clone(),
701 crate::cancel::CancelToken::new(),
702 ));
703 timeout(held_started.notified()).await;
704 assert_eq!(h.stats.busy(), 1, "the lane is full again");
705
706 release.notify_one();
710 assert!(
711 tokio::time::timeout(Duration::from_millis(250), h.outcomes.recv())
712 .await
713 .is_err(),
714 "the resumed batch waited for a permit instead of running"
715 );
716
717 held_release.notify_one();
718 let first = h.next_outcome().await;
719 assert_eq!(first.entity, entity(2), "the holder finished first");
720 let second = h.next_outcome().await;
721 assert_eq!(second.entity, entity(1), "then the resumed batch");
722
723 h.drain().await;
724 assert_eq!((h.stats.busy(), h.stats.parked()), (0, 0));
725 }
726
727 #[tokio::test]
729 async fn off_lane_outside_the_lane_just_awaits() {
730 assert_eq!(off_lane(async { 7 }).await, 7);
731 }
732
733 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
736 async fn a_cancelled_batch_is_abandoned_and_frees_the_lane() {
737 let mut h = Harness::new(1);
738 let cancel = crate::cancel::CancelToken::new();
739 let started = Arc::new(Notify::new());
740 let release = Arc::new(Notify::new());
741 h.submit(held_job(1, started.clone(), release, cancel.clone()));
742 timeout(started.notified()).await;
743 assert_eq!((h.stats.queued(), h.stats.busy()), (0, 1));
744
745 h.submit(job(2, vec![("c2", "r2")]));
747 cancel.cancel();
748
749 let next = h.next_outcome().await;
750 assert_eq!(next.entity, entity(2), "the queued batch ran");
751 h.drain().await;
752 assert!(
753 h.outcomes.try_recv().is_err(),
754 "the cancelled batch reported nothing"
755 );
756 assert_eq!(h.stats.busy(), 0, "and gave its permit back");
757 }
758
759 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
762 async fn a_cancelled_parked_batch_leaves_the_counters_straight() {
763 let mut h = Harness::new(1);
764 let cancel = crate::cancel::CancelToken::new();
765 let started = Arc::new(Notify::new());
766 let release = Arc::new(Notify::new());
767 h.submit(parking_job(1, started.clone(), release, cancel.clone()));
768 timeout(started.notified()).await;
769 assert_eq!((h.stats.busy(), h.stats.parked()), (0, 1));
770
771 cancel.cancel();
772 h.submit(job(2, vec![("c2", "r2")]));
773 let next = h.next_outcome().await;
774 assert_eq!(next.entity, entity(2));
775
776 h.drain().await;
777 assert_eq!((h.stats.busy(), h.stats.parked()), (0, 0));
778 }
779
780 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
783 async fn a_batch_cancelled_while_queued_never_runs() {
784 let mut h = Harness::new(1);
785 let blocker = crate::cancel::CancelToken::new();
786 let started = Arc::new(Notify::new());
787 let release = Arc::new(Notify::new());
788 h.submit(held_job(1, started.clone(), release.clone(), blocker));
789 timeout(started.notified()).await;
790
791 let cancel = crate::cancel::CancelToken::new();
792 h.submit(job_with(2, vec![("c2", "r2")], cancel.clone()));
793 cancel.cancel();
794 release.notify_one();
795
796 let first = h.next_outcome().await;
797 assert_eq!(first.entity, entity(1));
798 h.drain().await;
799 assert!(
800 h.outcomes.try_recv().is_err(),
801 "the cancelled batch never produced results"
802 );
803 assert_eq!(h.stats.queued(), 0, "and left the queue count clean");
804 }
805
806 #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
807 async fn the_lane_runs_batches_concurrently_up_to_its_cap() {
808 let h = Harness::new(3);
809 let barrier = Arc::new(tokio::sync::Barrier::new(3));
815 for i in 1..=3u32 {
816 let barrier = barrier.clone();
817 h.submit(ToolJob {
818 entity: entity(i),
819 exec: Box::new(move || {
820 Box::pin(async move {
821 barrier.wait().await;
822 vec![("c".to_string(), "r".to_string())]
823 })
824 }),
825 cancel: crate::cancel::CancelToken::new(),
826 });
827 }
828 let mut h = h;
829 h.drain().await;
830 for _ in 0..3 {
831 timeout(h.outcomes.recv()).await.expect("outcome present");
832 }
833 }
834
835 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
836 async fn the_lane_survives_a_dropped_outcome_receiver() {
837 let mut h = Harness::new(1);
838 h.submit(job(9, vec![("c", "r")]));
839 h.outcomes.close();
842 h.drain().await;
843 }
844
845 #[tokio::test(flavor = "multi_thread", worker_threads = 3)]
847 async fn relief_widens_the_lane() {
848 let mut h = Harness::new(1);
849 let started = Arc::new(Notify::new());
850 let release = Arc::new(Notify::new());
851 h.submit(held_job(
852 1,
853 started.clone(),
854 release.clone(),
855 crate::cancel::CancelToken::new(),
856 ));
857 timeout(started.notified()).await;
858 h.submit(job(2, vec![("c2", "r2")]));
859 assert!(h.stats.is_saturated(), "full, with a batch behind it");
860
861 assert_eq!(h.lane.relieve(0), 0, "relieving nothing changes nothing");
862 assert_eq!(h.lane.relieve(1), 1);
863 assert_eq!(h.stats.workers(), 2, "the cap moved with the permits");
864
865 let freed = h.next_outcome().await;
866 assert_eq!(freed.entity, entity(2), "the queued batch got in");
867
868 release.notify_one();
869 let held = h.next_outcome().await;
870 assert_eq!(held.entity, entity(1));
871 h.drain().await;
872 }
873
874 #[tokio::test]
878 async fn narrow_reclaims_idle_permits_and_never_busy_ones() {
879 let mut h = Harness::new(1);
880 assert_eq!(h.lane.relieve(2), 2);
881 assert_eq!(h.stats.workers(), 3);
882
883 assert_eq!(h.lane.narrow(0), 0, "narrowing nothing changes nothing");
886 assert_eq!(h.lane.narrow(1), 1);
887 assert_eq!(h.stats.workers(), 2);
888
889 let started_a = Arc::new(Notify::new());
892 let release_a = Arc::new(Notify::new());
893 h.submit(held_job(
894 1,
895 started_a.clone(),
896 release_a.clone(),
897 crate::cancel::CancelToken::new(),
898 ));
899 timeout(started_a.notified()).await;
900 let started_b = Arc::new(Notify::new());
901 let release_b = Arc::new(Notify::new());
902 h.submit(held_job(
903 2,
904 started_b.clone(),
905 release_b.clone(),
906 crate::cancel::CancelToken::new(),
907 ));
908 timeout(started_b.notified()).await;
909 assert_eq!(h.lane.narrow(1), 0, "a busy lane keeps its permits");
910 assert_eq!(h.stats.workers(), 2);
911
912 release_a.notify_one();
913 release_b.notify_one();
914 h.next_outcome().await;
915 h.next_outcome().await;
916 h.drain().await;
917 }
918
919 #[tokio::test]
923 async fn a_zero_width_lane_is_clamped_to_one() {
924 assert_eq!(ToolLaneStats::new(0).workers(), 1);
925 let mut h = Harness::new(0);
926 h.submit(job(7, vec![("c", "r")]));
927 assert_eq!(h.next_outcome().await.entity, entity(7));
928 h.drain().await;
929 }
930
931 #[test]
932 fn lane_stats_track_queue_depth_and_saturation() {
933 let stats = ToolLaneStats::new(2);
934 assert_eq!((stats.queued(), stats.busy(), stats.parked()), (0, 0, 0));
935 assert!(!stats.is_saturated(), "an idle lane is not saturated");
936
937 stats.enqueued();
938 stats.enqueued();
939 stats.enqueued();
940 assert_eq!(stats.queued(), 3);
941 stats.started();
943 stats.started();
944 assert_eq!((stats.queued(), stats.busy()), (1, 2));
945 assert!(
946 stats.is_saturated(),
947 "the lane is full with a batch still queued"
948 );
949
950 stats.began_park();
952 assert_eq!((stats.busy(), stats.parked()), (1, 1));
953 assert!(!stats.is_saturated(), "parked capacity is capacity");
954 stats.resumed();
955 assert_eq!((stats.busy(), stats.parked()), (2, 0));
956
957 stats.began_park();
958 stats.ended_park();
959 assert_eq!((stats.busy(), stats.parked()), (1, 0));
960
961 stats.finished();
962 stats.abandoned();
963 assert_eq!((stats.queued(), stats.busy()), (0, 0));
964 }
965}