use std::sync::{Arc, Mutex};
use std::time::Duration;
use crate::config::schema::LimitsConfig;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Timer {
Backoff,
Sweep,
}
pub trait Clock: Send + Sync + 'static {
fn now(&self) -> i64;
fn set_timeout(&self, action: Timer, ms: i64) -> u64;
fn clear_timeout(&self, handle: u64);
}
#[derive(Clone)]
pub struct SystemClock {
on_fire: Arc<dyn Fn(Timer) + Send + Sync>,
}
impl SystemClock {
pub fn new(on_fire: impl Fn(Timer) + Send + Sync + 'static) -> Self {
Self {
on_fire: Arc::new(on_fire),
}
}
}
impl Clock for SystemClock {
fn now(&self) -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_or(0, |since| {
i64::try_from(since.as_millis()).unwrap_or(i64::MAX)
})
}
fn set_timeout(&self, action: Timer, ms: i64) -> u64 {
let on_fire = Arc::clone(&self.on_fire);
tokio::spawn(async move {
let ms = u64::try_from(ms).unwrap_or(0);
tokio::time::sleep(Duration::from_millis(ms)).await;
on_fire(action);
});
0
}
fn clear_timeout(&self, _handle: u64) {
}
}
#[derive(Debug)]
pub struct Ticket {
released_at: Mutex<Option<i64>>,
}
impl Ticket {
fn mark_released(&self, at: i64) -> bool {
let mut released = self.released_at.lock().expect("a ticket lock");
if released.is_some() {
return false;
}
*released = Some(at);
true
}
}
#[derive(Debug)]
pub enum SubmitOutcome {
Admitted { ticket: Ticket },
Queued { position: usize },
Rejected { reason: String },
}
pub struct QueueEntry {
pub session_id: String,
pub on_admitted: Box<dyn FnOnce(Ticket) + Send>,
pub on_expired: Box<dyn FnOnce() + Send>,
pub on_position_changed: Option<Box<dyn Fn(usize) + Send>>,
}
struct PendingEntry {
entry: QueueEntry,
enqueued_at: i64,
last_reported_position: usize,
}
struct SchedulerState {
in_flight: usize,
live_sessions: usize,
queue: Vec<PendingEntry>,
backoff_until: i64,
backoff_ms: i64,
backoff_timer: Option<u64>,
last_start_at: i64,
sweep_timer: Option<u64>,
}
pub const PAUSED_REASON: &str = "provider backoff";
pub struct Scheduler {
limits: LimitsConfig,
clock: Arc<dyn Clock>,
base_backoff_ms: i64,
max_backoff_ms: i64,
start_interval_ms: i64,
state: Mutex<SchedulerState>,
}
impl Scheduler {
pub fn start(
limits: LimitsConfig,
clock: Arc<dyn Clock>,
base_backoff_ms: i64,
max_backoff_ms: i64,
start_interval_ms: i64,
) -> Arc<Self> {
Arc::new(Self {
limits,
clock,
base_backoff_ms,
max_backoff_ms,
start_interval_ms,
state: Mutex::new(SchedulerState {
in_flight: 0,
live_sessions: 0,
queue: Vec::new(),
backoff_until: 0,
backoff_ms: base_backoff_ms,
backoff_timer: None,
last_start_at: 0,
sweep_timer: None,
}),
})
}
pub fn turns_in_flight(&self) -> usize {
self.state.lock().expect("the scheduler lock").in_flight
}
pub fn queue_length(&self) -> usize {
self.state.lock().expect("the scheduler lock").queue.len()
}
#[allow(
dead_code,
reason = "read by this module's tests, which assert on state the daemon never asks for"
)]
pub fn sessions(&self) -> usize {
self.state.lock().expect("the scheduler lock").live_sessions
}
pub fn paused_because(&self) -> Option<&'static str> {
let state = self.state.lock().expect("the scheduler lock");
(self.clock.now() < state.backoff_until).then_some(PAUSED_REASON)
}
#[allow(
dead_code,
reason = "read by this module's tests, which assert on state the daemon never asks for"
)]
pub fn backoff_remaining_ms(&self) -> i64 {
let state = self.state.lock().expect("the scheduler lock");
(state.backoff_until - self.clock.now()).max(0)
}
pub fn reserve_session(&self) -> Option<i64> {
let mut state = self.state.lock().expect("the scheduler lock");
if state.live_sessions >= self.limits.max_live_sessions as usize {
return None;
}
state.live_sessions += 1;
let now = self.clock.now();
let earliest = state.last_start_at + self.start_interval_ms;
let delay = (earliest - now).max(0);
state.last_start_at = now.max(earliest);
Some(delay)
}
pub fn release_session(&self) {
let mut state = self.state.lock().expect("the scheduler lock");
if state.live_sessions > 0 {
state.live_sessions -= 1;
}
}
pub fn session_refused_reason(&self) -> String {
format!(
"the session limit of {} is reached, so this message did not start one",
self.limits.max_live_sessions
)
}
pub fn try_admit(self: &Arc<Self>) -> Option<Ticket> {
if !self.can_admit_now() {
return None;
}
self.state.lock().expect("the scheduler lock").in_flight += 1;
Some(Ticket {
released_at: Mutex::new(None),
})
}
pub fn submit(self: &Arc<Self>, entry: QueueEntry) -> SubmitOutcome {
if let Some(ticket) = self.try_admit() {
return SubmitOutcome::Admitted { ticket };
}
let mut state = self.state.lock().expect("the scheduler lock");
if state.queue.len() >= self.limits.max_queue_length as usize {
return SubmitOutcome::Rejected {
reason: format!(
"the queue is full at {} waiting prompts, so this message was not accepted",
self.limits.max_queue_length
),
};
}
let position = state.queue.len() + 1;
state.queue.push(PendingEntry {
entry,
enqueued_at: self.clock.now(),
last_reported_position: position,
});
drop(state);
self.schedule_sweep();
SubmitOutcome::Queued { position }
}
pub fn release(self: &Arc<Self>, ticket: &Ticket) -> bool {
if !ticket.mark_released(self.clock.now()) {
return false;
}
{
let mut state = self.state.lock().expect("the scheduler lock");
if state.in_flight > 0 {
state.in_flight -= 1;
}
}
self.pump();
true
}
pub fn cancel_session(&self, session_id: &str) -> usize {
let mut state = self.state.lock().expect("the scheduler lock");
let before = state.queue.len();
state
.queue
.retain(|pending| pending.entry.session_id != session_id);
let removed = before - state.queue.len();
if removed > 0 {
drop(state);
self.report_positions();
}
removed
}
pub fn note_rate_limit(self: &Arc<Self>) {
let backoff_ms = {
let mut state = self.state.lock().expect("the scheduler lock");
let now = self.clock.now();
if now < state.backoff_until {
state.backoff_ms = (state.backoff_ms * 2).min(self.max_backoff_ms);
}
state.backoff_until = now + state.backoff_ms;
state.backoff_ms
};
let previous = {
let mut state = self.state.lock().expect("the scheduler lock");
state
.backoff_timer
.replace(self.clock.set_timeout(Timer::Backoff, backoff_ms))
};
if let Some(handle) = previous {
self.clock.clear_timeout(handle);
}
}
pub fn note_success(&self) {
let mut state = self.state.lock().expect("the scheduler lock");
if self.clock.now() < state.backoff_until {
return;
}
state.backoff_ms = self.base_backoff_ms.max(state.backoff_ms / 2);
}
pub fn expire_stale(&self) -> usize {
let expired: Vec<PendingEntry>;
{
let mut state = self.state.lock().expect("the scheduler lock");
#[expect(clippy::cast_possible_wrap)]
let wait = self.limits.max_queue_wait_ms as i64;
let cutoff = self.clock.now() - wait;
let due: Vec<usize> = state
.queue
.iter()
.enumerate()
.filter(|(_, pending)| pending.enqueued_at <= cutoff)
.map(|(index, _)| index)
.collect();
expired = due
.into_iter()
.rev()
.map(|index| state.queue.remove(index))
.collect();
if expired.is_empty() {
return 0;
}
}
let count = expired.len();
for pending in expired {
(pending.entry.on_expired)();
}
self.report_positions();
count
}
pub fn shutdown(&self) {
let mut state = self.state.lock().expect("the scheduler lock");
if let Some(handle) = state.backoff_timer.take() {
self.clock.clear_timeout(handle);
}
if let Some(handle) = state.sweep_timer.take() {
self.clock.clear_timeout(handle);
}
state.queue.clear();
}
pub fn timer_fired(self: &Arc<Self>, action: Timer) {
match action {
Timer::Backoff => {
{
let mut state = self.state.lock().expect("the scheduler lock");
state.backoff_timer = None;
}
self.pump();
}
Timer::Sweep => {
{
let mut state = self.state.lock().expect("the scheduler lock");
state.sweep_timer = None;
}
self.expire_stale();
let waiting = self.state.lock().expect("the scheduler lock").queue.len();
if waiting > 0 {
self.schedule_sweep();
}
}
}
}
fn can_admit_now(&self) -> bool {
let state = self.state.lock().expect("the scheduler lock");
state.in_flight < self.limits.max_concurrent_turns as usize
&& self.clock.now() >= state.backoff_until
}
fn pump(self: &Arc<Self>) {
loop {
let next = {
let mut state = self.state.lock().expect("the scheduler lock");
let backoff_over = self.clock.now() >= state.backoff_until;
if state.queue.is_empty()
|| state.in_flight >= self.limits.max_concurrent_turns as usize
|| !backoff_over
{
break;
}
let next = state.queue.remove(0);
state.in_flight += 1;
next
};
(next.entry.on_admitted)(Ticket {
released_at: Mutex::new(None),
});
}
self.report_positions();
}
fn report_positions(&self) {
let mut state = self.state.lock().expect("the scheduler lock");
for (index, pending) in state.queue.iter_mut().enumerate() {
let position = index + 1;
if pending.last_reported_position == position {
continue;
}
pending.last_reported_position = position;
if let Some(report) = &mut pending.entry.on_position_changed {
report(position);
}
}
}
fn schedule_sweep(self: &Arc<Self>) {
let mut state = self.state.lock().expect("the scheduler lock");
if state.sweep_timer.is_some() {
return;
}
state.sweep_timer = Some(self.clock.set_timeout(Timer::Sweep, self.sweep_interval()));
}
fn sweep_interval(&self) -> i64 {
{
#[expect(clippy::cast_possible_wrap)]
let interval = (self.limits.max_queue_wait_ms / 4) as i64;
interval.max(1)
}
}
}
#[cfg(test)]
mod tests;