1use a3s_lane::{Priority, PriorityQueue};
8use serde::{Deserialize, Serialize};
9use std::collections::{HashMap, HashSet};
10use std::str::FromStr;
11use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
12use std::sync::Arc;
13use thiserror::Error;
14use tokio::sync::{mpsc, oneshot};
15use tokio::time::Instant;
16use tokio_util::sync::CancellationToken;
17
18const DEFAULT_MAX_ACTIVE: usize = 4;
19const DEFAULT_AGING_INTERVAL_MS: u64 = 30_000;
20
21#[derive(
27 Debug, Clone, Copy, Default, PartialEq, Eq, Hash, Serialize, Deserialize, PartialOrd, Ord,
28)]
29#[serde(rename_all = "camelCase")]
30#[repr(u8)]
31pub enum TaskPriority {
32 Urgent = 0,
33 #[default]
34 Interactive = 1,
35 Foreground = 2,
36 Background = 3,
37 Maintenance = 4,
38}
39
40impl TaskPriority {
41 const ALL: [Self; 5] = [
42 Self::Urgent,
43 Self::Interactive,
44 Self::Foreground,
45 Self::Background,
46 Self::Maintenance,
47 ];
48
49 fn lane_priority(self) -> Priority {
50 self as Priority
51 }
52}
53
54impl FromStr for TaskPriority {
55 type Err = TaskSchedulerError;
56
57 fn from_str(value: &str) -> Result<Self, Self::Err> {
58 match value.trim().to_ascii_lowercase().replace(['-', '_'], "").as_str() {
59 "urgent" => Ok(Self::Urgent),
60 "interactive" | "user" => Ok(Self::Interactive),
61 "foreground" => Ok(Self::Foreground),
62 "background" => Ok(Self::Background),
63 "maintenance" => Ok(Self::Maintenance),
64 _ => Err(TaskSchedulerError::InvalidConfig(format!(
65 "unknown task priority '{value}'; expected urgent, interactive, foreground, background, or maintenance"
66 ))),
67 }
68 }
69}
70
71#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
73#[serde(rename_all = "camelCase")]
74pub struct TaskSchedulerConfig {
75 #[serde(default = "default_max_active", alias = "max_active")]
77 pub max_active: usize,
78 #[serde(default = "default_aging_interval_ms", alias = "aging_interval_ms")]
80 pub aging_interval_ms: u64,
81}
82
83impl Default for TaskSchedulerConfig {
84 fn default() -> Self {
85 Self {
86 max_active: default_max_active(),
87 aging_interval_ms: default_aging_interval_ms(),
88 }
89 }
90}
91
92impl TaskSchedulerConfig {
93 pub fn validate(&self) -> Result<(), TaskSchedulerError> {
95 if self.max_active == 0 {
96 return Err(TaskSchedulerError::InvalidConfig(
97 "maxActive must be greater than zero".to_string(),
98 ));
99 }
100 if self.aging_interval_ms == 0 {
101 return Err(TaskSchedulerError::InvalidConfig(
102 "agingIntervalMs must be greater than zero".to_string(),
103 ));
104 }
105 Ok(())
106 }
107}
108
109const fn default_max_active() -> usize {
110 DEFAULT_MAX_ACTIVE
111}
112
113const fn default_aging_interval_ms() -> u64 {
114 DEFAULT_AGING_INTERVAL_MS
115}
116
117#[derive(Debug, Clone, Error, PartialEq, Eq)]
119pub enum TaskSchedulerError {
120 #[error("task scheduler configuration is invalid: {0}")]
121 InvalidConfig(String),
122 #[error("task admission was cancelled")]
123 Cancelled,
124 #[error("task scheduler is closed")]
125 Closed,
126}
127
128#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
130#[serde(rename_all = "camelCase")]
131pub struct TaskPriorityCounts {
132 pub urgent: usize,
133 pub interactive: usize,
134 pub foreground: usize,
135 pub background: usize,
136 pub maintenance: usize,
137}
138
139impl TaskPriorityCounts {
140 fn increment(&mut self, priority: TaskPriority) {
141 match priority {
142 TaskPriority::Urgent => self.urgent += 1,
143 TaskPriority::Interactive => self.interactive += 1,
144 TaskPriority::Foreground => self.foreground += 1,
145 TaskPriority::Background => self.background += 1,
146 TaskPriority::Maintenance => self.maintenance += 1,
147 }
148 }
149}
150
151#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
153#[serde(rename_all = "camelCase")]
154pub struct TaskSchedulerStats {
155 pub max_active: usize,
156 pub active: usize,
157 pub pending: usize,
158 pub active_by_priority: TaskPriorityCounts,
159 pub pending_by_priority: TaskPriorityCounts,
160 pub closed: bool,
161}
162
163#[derive(Debug)]
165pub struct TaskScheduler {
166 tx: mpsc::UnboundedSender<SchedulerMessage>,
167 next_id: AtomicU64,
168 closed: Arc<AtomicBool>,
169}
170
171impl TaskScheduler {
172 pub fn new(config: TaskSchedulerConfig) -> Result<Self, TaskSchedulerError> {
174 config.validate()?;
175 let (tx, rx) = mpsc::unbounded_channel();
176 let closed = Arc::new(AtomicBool::new(false));
177 tokio::spawn(run_scheduler(rx, config, Arc::clone(&closed)));
178 Ok(Self {
179 tx,
180 next_id: AtomicU64::new(1),
181 closed,
182 })
183 }
184
185 pub async fn acquire(
187 &self,
188 priority: TaskPriority,
189 label: impl Into<String>,
190 cancellation: &CancellationToken,
191 ) -> Result<TaskLease, TaskSchedulerError> {
192 if self.closed.load(Ordering::Acquire) {
193 return Err(TaskSchedulerError::Closed);
194 }
195 if cancellation.is_cancelled() {
196 return Err(TaskSchedulerError::Cancelled);
197 }
198
199 let id = self.next_id.fetch_add(1, Ordering::Relaxed);
200 let (ready_tx, ready_rx) = oneshot::channel();
201 self.tx
202 .send(SchedulerMessage::Enqueue(QueuedAdmission {
203 id,
204 priority,
205 label: label.into(),
206 enqueued_at: Instant::now(),
207 ready: ready_tx,
208 }))
209 .map_err(|_| TaskSchedulerError::Closed)?;
210
211 tokio::select! {
212 biased;
213 _ = cancellation.cancelled() => {
214 let _ = self.tx.send(SchedulerMessage::Cancel(id));
215 Err(TaskSchedulerError::Cancelled)
216 }
217 ready = ready_rx => {
218 ready.map_err(|_| TaskSchedulerError::Closed)??;
219 Ok(TaskLease {
220 id,
221 tx: self.tx.clone(),
222 released: false,
223 })
224 }
225 }
226 }
227
228 pub async fn stats(&self) -> Result<TaskSchedulerStats, TaskSchedulerError> {
230 let (tx, rx) = oneshot::channel();
231 self.tx
232 .send(SchedulerMessage::Stats(tx))
233 .map_err(|_| TaskSchedulerError::Closed)?;
234 rx.await.map_err(|_| TaskSchedulerError::Closed)
235 }
236
237 pub async fn shutdown(&self) {
239 if self.closed.swap(true, Ordering::AcqRel) {
240 return;
241 }
242 let (tx, rx) = oneshot::channel();
243 if self.tx.send(SchedulerMessage::Shutdown(tx)).is_ok() {
244 let _ = rx.await;
245 }
246 }
247}
248
249pub struct TaskLease {
251 id: u64,
252 tx: mpsc::UnboundedSender<SchedulerMessage>,
253 released: bool,
254}
255
256impl TaskLease {
257 pub fn id(&self) -> u64 {
259 self.id
260 }
261}
262
263impl Drop for TaskLease {
264 fn drop(&mut self) {
265 if !self.released {
266 self.released = true;
267 let _ = self.tx.send(SchedulerMessage::Release(self.id));
268 }
269 }
270}
271
272struct QueuedAdmission {
273 id: u64,
274 priority: TaskPriority,
275 label: String,
276 enqueued_at: Instant,
277 ready: oneshot::Sender<Result<(), TaskSchedulerError>>,
278}
279
280enum SchedulerMessage {
281 Enqueue(QueuedAdmission),
282 Cancel(u64),
283 Release(u64),
284 Stats(oneshot::Sender<TaskSchedulerStats>),
285 Shutdown(oneshot::Sender<()>),
286}
287
288struct SchedulerState {
289 config: TaskSchedulerConfig,
290 pending: PriorityQueue<QueuedAdmission>,
291 cancelled: HashSet<u64>,
292 active: HashMap<u64, TaskPriority>,
293 closing: bool,
294 shutdown_waiters: Vec<oneshot::Sender<()>>,
295}
296
297async fn run_scheduler(
298 mut rx: mpsc::UnboundedReceiver<SchedulerMessage>,
299 config: TaskSchedulerConfig,
300 closed: Arc<AtomicBool>,
301) {
302 let mut state = SchedulerState {
303 config,
304 pending: PriorityQueue::new(),
305 cancelled: HashSet::new(),
306 active: HashMap::new(),
307 closing: false,
308 shutdown_waiters: Vec::new(),
309 };
310
311 while let Some(message) = rx.recv().await {
312 match message {
313 SchedulerMessage::Enqueue(item) => {
314 if state.closing {
315 let _ = item.ready.send(Err(TaskSchedulerError::Closed));
316 } else {
317 state.pending.push(item.priority.lane_priority(), item);
318 state.dispatch();
319 }
320 }
321 SchedulerMessage::Cancel(id) => {
322 if state.active.remove(&id).is_none() {
323 state.cancelled.insert(id);
324 state.purge_cancelled();
325 }
326 state.dispatch();
327 state.finish_shutdown_if_idle();
328 }
329 SchedulerMessage::Release(id) => {
330 state.active.remove(&id);
331 state.dispatch();
332 state.finish_shutdown_if_idle();
333 }
334 SchedulerMessage::Stats(reply) => {
335 let _ = reply.send(state.snapshot());
336 }
337 SchedulerMessage::Shutdown(reply) => {
338 state.closing = true;
339 closed.store(true, Ordering::Release);
340 while let Some(item) = state.pending.pop() {
341 let item = item.into_value();
342 let _ = item.ready.send(Err(TaskSchedulerError::Closed));
343 }
344 state.cancelled.clear();
345 state.shutdown_waiters.push(reply);
346 state.finish_shutdown_if_idle();
347 }
348 }
349
350 if state.closing && state.active.is_empty() && state.shutdown_waiters.is_empty() {
351 break;
352 }
353 }
354
355 closed.store(true, Ordering::Release);
356}
357
358impl SchedulerState {
359 fn purge_cancelled(&mut self) {
360 if self.cancelled.is_empty() || self.pending.is_empty() {
361 return;
362 }
363 let mut retained = Vec::with_capacity(self.pending.len());
364 while let Some(item) = self.pending.pop() {
365 if self.cancelled.remove(&item.value().id) {
366 let item = item.into_value();
367 let _ = item.ready.send(Err(TaskSchedulerError::Cancelled));
368 } else {
369 retained.push(item);
370 }
371 }
372 for item in retained {
373 self.pending.restore(item);
374 }
375 }
376
377 fn dispatch(&mut self) {
378 if self.closing {
379 return;
380 }
381 self.apply_aging();
382 while self.active.len() < self.config.max_active {
383 let Some(item) = self.pending.pop() else {
384 break;
385 };
386 let item = item.into_value();
387 if self.cancelled.remove(&item.id) {
388 continue;
389 }
390
391 let id = item.id;
392 let priority = item.priority;
393 let label = item.label;
394 self.active.insert(id, priority);
395 if item.ready.send(Ok(())).is_err() {
396 self.active.remove(&id);
397 continue;
398 }
399 tracing::trace!(admission_id = id, ?priority, %label, "task admitted");
400 }
401 }
402
403 fn apply_aging(&mut self) {
404 if self.pending.is_empty() {
405 return;
406 }
407 let now = Instant::now();
408 let interval_ms = self.config.aging_interval_ms as u128;
409 let mut entries = Vec::with_capacity(self.pending.len());
410 while let Some(item) = self.pending.pop() {
411 entries.push((item.sequence(), item.into_value()));
412 }
413 entries.sort_by_key(|(sequence, _)| *sequence);
416 for (_, item) in entries {
417 let elapsed_ms = now.duration_since(item.enqueued_at).as_millis();
418 let levels = (elapsed_ms / interval_ms).min(u8::MAX as u128) as u8;
419 let effective = if item.priority == TaskPriority::Urgent {
420 TaskPriority::Urgent.lane_priority()
421 } else {
422 (item.priority as u8).saturating_sub(levels).max(1) as Priority
423 };
424 self.pending.push(effective, item);
425 }
426 }
427
428 fn snapshot(&self) -> TaskSchedulerStats {
429 let mut active_by_priority = TaskPriorityCounts::default();
430 for priority in self.active.values() {
431 active_by_priority.increment(*priority);
432 }
433 let mut pending_by_priority = TaskPriorityCounts::default();
434 for item in self.pending.ordered() {
435 pending_by_priority.increment(item.value().priority);
436 }
437 debug_assert_eq!(
438 TaskPriority::ALL
439 .iter()
440 .map(|priority| match priority {
441 TaskPriority::Urgent => active_by_priority.urgent,
442 TaskPriority::Interactive => active_by_priority.interactive,
443 TaskPriority::Foreground => active_by_priority.foreground,
444 TaskPriority::Background => active_by_priority.background,
445 TaskPriority::Maintenance => active_by_priority.maintenance,
446 })
447 .sum::<usize>(),
448 self.active.len()
449 );
450 TaskSchedulerStats {
451 max_active: self.config.max_active,
452 active: self.active.len(),
453 pending: self.pending.len(),
454 active_by_priority,
455 pending_by_priority,
456 closed: self.closing,
457 }
458 }
459
460 fn finish_shutdown_if_idle(&mut self) {
461 if self.closing && self.active.is_empty() {
462 for waiter in self.shutdown_waiters.drain(..) {
463 let _ = waiter.send(());
464 }
465 }
466 }
467}
468
469#[cfg(test)]
470mod tests {
471 use super::*;
472 use std::time::Duration;
473
474 fn scheduler(max_active: usize, aging_interval_ms: u64) -> TaskScheduler {
475 TaskScheduler::new(TaskSchedulerConfig {
476 max_active,
477 aging_interval_ms,
478 })
479 .unwrap()
480 }
481
482 #[test]
483 fn priority_names_are_stable_and_reject_unknown_values() {
484 assert_eq!("user".parse(), Ok(TaskPriority::Interactive));
485 assert_eq!("background".parse(), Ok(TaskPriority::Background));
486 assert!("eventually".parse::<TaskPriority>().is_err());
487 }
488
489 async fn wait_for_pending(scheduler: &TaskScheduler, expected: usize) {
490 for _ in 0..100 {
491 if scheduler.stats().await.unwrap().pending == expected {
492 return;
493 }
494 tokio::task::yield_now().await;
495 }
496 panic!("scheduler never reached {expected} pending tasks");
497 }
498
499 #[tokio::test]
500 async fn strict_priority_and_fifo_are_enforced_globally() {
501 let scheduler = Arc::new(scheduler(1, 60_000));
502 let blocker = scheduler
503 .acquire(
504 TaskPriority::Interactive,
505 "blocker",
506 &CancellationToken::new(),
507 )
508 .await
509 .unwrap();
510 let (order_tx, mut order_rx) = mpsc::unbounded_channel();
511
512 for (name, priority) in [
513 ("background", TaskPriority::Background),
514 ("interactive-1", TaskPriority::Interactive),
515 ("foreground", TaskPriority::Foreground),
516 ("interactive-2", TaskPriority::Interactive),
517 ("urgent", TaskPriority::Urgent),
518 ] {
519 let expected = scheduler.stats().await.unwrap().pending + 1;
520 let task_scheduler = Arc::clone(&scheduler);
521 let order_tx = order_tx.clone();
522 tokio::spawn(async move {
523 let lease = task_scheduler
524 .acquire(priority, name, &CancellationToken::new())
525 .await
526 .unwrap();
527 order_tx.send(name).unwrap();
528 drop(lease);
529 });
530 wait_for_pending(&scheduler, expected).await;
531 }
532
533 drop(blocker);
534 let mut actual = Vec::new();
535 for _ in 0..5 {
536 actual.push(order_rx.recv().await.unwrap());
537 }
538 assert_eq!(
539 actual,
540 [
541 "urgent",
542 "interactive-1",
543 "interactive-2",
544 "foreground",
545 "background"
546 ]
547 );
548 scheduler.shutdown().await;
549 }
550
551 #[tokio::test]
552 async fn cancellation_does_not_consume_capacity() {
553 let scheduler = Arc::new(scheduler(1, 60_000));
554 let blocker = scheduler
555 .acquire(
556 TaskPriority::Interactive,
557 "blocker",
558 &CancellationToken::new(),
559 )
560 .await
561 .unwrap();
562 let cancellation = CancellationToken::new();
563 let cancelled_task = {
564 let scheduler = Arc::clone(&scheduler);
565 let cancellation = cancellation.clone();
566 tokio::spawn(async move {
567 scheduler
568 .acquire(TaskPriority::Urgent, "cancelled", &cancellation)
569 .await
570 })
571 };
572 wait_for_pending(&scheduler, 1).await;
573 cancellation.cancel();
574 assert!(matches!(
575 cancelled_task.await.unwrap(),
576 Err(TaskSchedulerError::Cancelled)
577 ));
578
579 let next = {
580 let scheduler = Arc::clone(&scheduler);
581 tokio::spawn(async move {
582 scheduler
583 .acquire(TaskPriority::Background, "next", &CancellationToken::new())
584 .await
585 })
586 };
587 wait_for_pending(&scheduler, 1).await;
588 drop(blocker);
589 let lease = next.await.unwrap().unwrap();
590 assert_eq!(scheduler.stats().await.unwrap().active, 1);
591 drop(lease);
592 scheduler.shutdown().await;
593 }
594
595 #[tokio::test]
596 async fn aging_prevents_background_starvation() {
597 let scheduler = Arc::new(scheduler(1, 2));
598 let blocker = scheduler
599 .acquire(
600 TaskPriority::Interactive,
601 "blocker",
602 &CancellationToken::new(),
603 )
604 .await
605 .unwrap();
606 let (order_tx, mut order_rx) = mpsc::unbounded_channel();
607 let background = {
608 let scheduler = Arc::clone(&scheduler);
609 let order_tx = order_tx.clone();
610 tokio::spawn(async move {
611 let lease = scheduler
612 .acquire(
613 TaskPriority::Background,
614 "old-background",
615 &CancellationToken::new(),
616 )
617 .await
618 .unwrap();
619 order_tx.send("background").unwrap();
620 drop(lease);
621 })
622 };
623 wait_for_pending(&scheduler, 1).await;
624 tokio::time::sleep(Duration::from_millis(8)).await;
625 let interactive = {
626 let scheduler = Arc::clone(&scheduler);
627 let order_tx = order_tx.clone();
628 tokio::spawn(async move {
629 let lease = scheduler
630 .acquire(
631 TaskPriority::Interactive,
632 "new-interactive",
633 &CancellationToken::new(),
634 )
635 .await
636 .unwrap();
637 order_tx.send("interactive").unwrap();
638 drop(lease);
639 })
640 };
641 wait_for_pending(&scheduler, 2).await;
642 drop(blocker);
643
644 assert_eq!(order_rx.recv().await.unwrap(), "background");
645 assert_eq!(order_rx.recv().await.unwrap(), "interactive");
646 background.await.unwrap();
647 interactive.await.unwrap();
648 scheduler.shutdown().await;
649 }
650
651 #[tokio::test]
652 async fn shutdown_rejects_pending_and_waits_for_active_lease() {
653 let scheduler = Arc::new(scheduler(1, 60_000));
654 let blocker = scheduler
655 .acquire(
656 TaskPriority::Interactive,
657 "blocker",
658 &CancellationToken::new(),
659 )
660 .await
661 .unwrap();
662 let pending = {
663 let scheduler = Arc::clone(&scheduler);
664 tokio::spawn(async move {
665 scheduler
666 .acquire(
667 TaskPriority::Background,
668 "pending",
669 &CancellationToken::new(),
670 )
671 .await
672 })
673 };
674 wait_for_pending(&scheduler, 1).await;
675 let shutdown = {
676 let scheduler = Arc::clone(&scheduler);
677 tokio::spawn(async move { scheduler.shutdown().await })
678 };
679 assert!(matches!(
680 pending.await.unwrap(),
681 Err(TaskSchedulerError::Closed)
682 ));
683 assert!(!shutdown.is_finished());
684 drop(blocker);
685 shutdown.await.unwrap();
686 assert!(matches!(
687 scheduler
688 .acquire(TaskPriority::Urgent, "late", &CancellationToken::new())
689 .await,
690 Err(TaskSchedulerError::Closed)
691 ));
692 }
693
694 #[tokio::test]
695 async fn stats_report_base_priority_occupancy() {
696 let scheduler = Arc::new(scheduler(1, 60_000));
697 let blocker = scheduler
698 .acquire(
699 TaskPriority::Foreground,
700 "blocker",
701 &CancellationToken::new(),
702 )
703 .await
704 .unwrap();
705 let waiting = {
706 let scheduler = Arc::clone(&scheduler);
707 tokio::spawn(async move {
708 scheduler
709 .acquire(
710 TaskPriority::Maintenance,
711 "waiting",
712 &CancellationToken::new(),
713 )
714 .await
715 })
716 };
717 wait_for_pending(&scheduler, 1).await;
718 let stats = scheduler.stats().await.unwrap();
719 assert_eq!(stats.max_active, 1);
720 assert_eq!(stats.active_by_priority.foreground, 1);
721 assert_eq!(stats.pending_by_priority.maintenance, 1);
722 drop(blocker);
723 drop(waiting.await.unwrap().unwrap());
724 scheduler.shutdown().await;
725 }
726}