Skip to main content

a3s_code_core/
task_scheduler.rs

1//! Agent-wide admission scheduler for top-level and background work.
2//!
3//! Session transcript admission remains single-flight. This scheduler adds a
4//! shared capacity boundary across every session created by one [`Agent`](crate::Agent),
5//! using `a3s-lane`'s stable priority queue for exact priority/FIFO ordering.
6
7use 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/// Relative importance of work admitted through an agent's shared scheduler.
22///
23/// Lower values run first. `Urgent` is reserved for explicit host control
24/// actions and never participates in aging. Older work from the other classes
25/// can age up to `Interactive`, but never ahead of `Urgent`.
26#[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/// Agent-wide task scheduler settings.
72#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
73#[serde(rename_all = "camelCase")]
74pub struct TaskSchedulerConfig {
75    /// Maximum number of independently admitted tasks across all sessions.
76    #[serde(default = "default_max_active", alias = "max_active")]
77    pub max_active: usize,
78    /// Time before queued work is promoted by one priority level.
79    #[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    /// Validate configuration before starting the scheduler actor.
94    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/// Scheduler admission failures.
118#[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/// Counts grouped by the stable public priority classes.
129#[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/// Point-in-time scheduler occupancy for hosts and diagnostics.
152#[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/// Shared actor handle. One instance belongs to each `Agent`.
164#[derive(Debug)]
165pub struct TaskScheduler {
166    tx: mpsc::UnboundedSender<SchedulerMessage>,
167    next_id: AtomicU64,
168    closed: Arc<AtomicBool>,
169}
170
171impl TaskScheduler {
172    /// Start a scheduler on the current Tokio runtime.
173    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    /// Wait until this task owns one global execution slot.
186    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    /// Return a consistent actor-owned occupancy snapshot.
229    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    /// Reject pending work and wait for already-admitted leases to finish.
238    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
249/// RAII ownership of one globally admitted execution slot.
250pub struct TaskLease {
251    id: u64,
252    tx: mpsc::UnboundedSender<SchedulerMessage>,
253    released: bool,
254}
255
256impl TaskLease {
257    /// Stable admission identifier, useful for tracing.
258    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        // Re-insertion gives Lane fresh sequence numbers. Insert in original
414        // sequence order so work that ages into the same class remains FIFO.
415        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}