use std::cmp::Ordering;
use std::collections::{BTreeSet, BinaryHeap};
use std::time::Duration;
use ordered_float::OrderedFloat;
use rustc_hash::{FxHashMap, FxHashSet};
use serde::{Deserialize, Serialize};
use tokio::time::Instant;
use super::config::RouterQueuePolicy;
use super::policy_config::{PolicyClassConfig, PolicyProfile};
use super::queue_admission::{
AdmissionDecision, AdmissionEvent, AdmissionId, AdmissionRequest, AdmissionTicket,
ClassAdmissionAction, PolicyClassAdmissionPolicies, PolicyClassAdmissionPolicy,
RequestProgress, RequestProgressUpdater, WorkerEligibility, WorkerPlacement,
};
use super::types::KvSchedulerError;
use crate::protocols::WorkerWithDpRank;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct QueueSnapshot {
pub raw_isl_tokens: usize,
pub cached_tokens: usize,
pub uncached_tokens: usize,
pub scheduling_cost_tokens: usize,
}
impl QueueSnapshot {
pub fn new(raw_isl_tokens: usize, cached_tokens: usize) -> Self {
let cached_tokens = cached_tokens.min(raw_isl_tokens);
Self {
raw_isl_tokens,
cached_tokens,
uncached_tokens: raw_isl_tokens.saturating_sub(cached_tokens),
scheduling_cost_tokens: raw_isl_tokens.saturating_sub(cached_tokens).max(1),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum QueueLimitKind {
Requests,
RawIslTokens,
CachedTokens,
}
impl std::fmt::Display for QueueLimitKind {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Requests => formatter.write_str("requests"),
Self::RawIslTokens => formatter.write_str("raw_isl_tokens"),
Self::CachedTokens => formatter.write_str("cached_tokens"),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, thiserror::Error)]
#[error(
"router policy class {policy_class:?} queue {limit_kind} limit reached \
(current={current}, limit={limit})"
)]
pub struct QueueRejection {
pub policy_class: String,
pub limit_kind: QueueLimitKind,
pub current: usize,
pub limit: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct PolicyQueueStats {
pub requests: usize,
pub raw_isl_tokens: usize,
pub cached_tokens: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct QueuePriority {
strict_priority: u32,
policy_score: OrderedFloat<f64>,
}
#[inline]
fn cmp_queue_order(
lhs_priority: QueuePriority,
lhs_enqueue_seq: u64,
rhs_priority: QueuePriority,
rhs_enqueue_seq: u64,
) -> Ordering {
lhs_priority
.strict_priority
.cmp(&rhs_priority.strict_priority)
.then_with(|| lhs_priority.policy_score.cmp(&rhs_priority.policy_score))
.then_with(|| rhs_enqueue_seq.cmp(&lhs_enqueue_seq))
}
pub struct PolicyQueueEntry<T> {
class_index: usize,
priority: QueuePriority,
enqueue_seq: u64,
snapshot: QueueSnapshot,
payload: T,
}
impl<T> PolicyQueueEntry<T> {
pub fn class_index(&self) -> usize {
self.class_index
}
pub fn snapshot(&self) -> QueueSnapshot {
self.snapshot
}
pub fn payload(&self) -> &T {
&self.payload
}
pub fn payload_mut(&mut self) -> &mut T {
&mut self.payload
}
pub fn into_payload(self) -> T {
self.payload
}
}
impl<T> Eq for PolicyQueueEntry<T> {}
impl<T> PartialEq for PolicyQueueEntry<T> {
fn eq(&self, other: &Self) -> bool {
self.priority == other.priority && self.enqueue_seq == other.enqueue_seq
}
}
impl<T> Ord for PolicyQueueEntry<T> {
fn cmp(&self, other: &Self) -> Ordering {
self.priority
.strict_priority
.cmp(&other.priority.strict_priority)
.then_with(|| self.priority.policy_score.cmp(&other.priority.policy_score))
.then_with(|| other.enqueue_seq.cmp(&self.enqueue_seq))
}
}
impl<T> PartialOrd for PolicyQueueEntry<T> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
#[derive(Debug, Clone, Copy)]
struct DispatchCandidate {
placement: WorkerPlacement,
cost: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct WorkerLaneHead {
worker: WorkerWithDpRank,
priority: QueuePriority,
enqueue_seq: u64,
}
impl WorkerLaneHead {
fn new<T>(worker: WorkerWithDpRank, entry: &PolicyQueueEntry<T>) -> Self {
Self {
worker,
priority: entry.priority,
enqueue_seq: entry.enqueue_seq,
}
}
}
impl Ord for WorkerLaneHead {
fn cmp(&self, other: &Self) -> Ordering {
cmp_queue_order(
self.priority,
self.enqueue_seq,
other.priority,
other.enqueue_seq,
)
.then_with(|| self.worker.cmp(&other.worker))
}
}
impl PartialOrd for WorkerLaneHead {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
struct PolicyClassQueue<T> {
config: PolicyClassConfig,
admission_policy: Option<ScheduledAdmissionPolicy>,
pending: BinaryHeap<PolicyQueueEntry<T>>,
deferred: FxHashMap<AdmissionId, PolicyQueueEntry<T>>,
stats: PolicyQueueStats,
deficit: usize,
ready_by_worker: FxHashMap<WorkerWithDpRank, BinaryHeap<PolicyQueueEntry<T>>>,
blocked_workers: FxHashSet<WorkerWithDpRank>,
candidate_worker_heads: BTreeSet<WorkerLaneHead>,
needs_blocked_worker_recheck: bool,
}
struct ScheduledAdmissionPolicy {
policy: Box<dyn PolicyClassAdmissionPolicy>,
reconcile_interval: Duration,
next_reconcile: Instant,
}
impl<T> PolicyClassQueue<T> {
fn ready_is_empty(&self) -> bool {
self.pending.is_empty() && self.ready_by_worker.is_empty()
}
fn entries(&self) -> impl Iterator<Item = &PolicyQueueEntry<T>> {
self.pending
.iter()
.chain(self.ready_by_worker.values().flat_map(|ready| ready.iter()))
.chain(self.deferred.values())
}
fn push_ready(&mut self, placement: WorkerPlacement, entry: PolicyQueueEntry<T>) {
match placement {
WorkerPlacement::Any => self.pending.push(entry),
WorkerPlacement::Exact(worker) => {
let ready = self.ready_by_worker.entry(worker).or_default();
let old_head = (!self.blocked_workers.contains(&worker))
.then(|| ready.peek().map(|entry| WorkerLaneHead::new(worker, entry)))
.flatten();
ready.push(entry);
if !self.blocked_workers.contains(&worker) {
let new_head = WorkerLaneHead::new(
worker,
ready.peek().expect("worker lane was just populated"),
);
if old_head != Some(new_head) {
if let Some(old_head) = old_head {
let removed = self.candidate_worker_heads.remove(&old_head);
debug_assert!(removed);
}
self.candidate_worker_heads.insert(new_head);
}
}
}
}
}
#[inline]
fn next_dispatchable(
&mut self,
class_index: usize,
is_dispatchable: &mut impl FnMut(usize, &PolicyClassConfig, &T) -> bool,
) -> Option<DispatchCandidate> {
if self.ready_by_worker.is_empty() {
debug_assert!(
self.blocked_workers.is_empty(),
"blocked workers must have a ready lane"
);
return self
.pending
.peek()
.filter(|entry| is_dispatchable(class_index, &self.config, entry.payload()))
.map(|entry| DispatchCandidate {
placement: WorkerPlacement::Any,
cost: entry.snapshot.scheduling_cost_tokens,
});
}
self.next_dispatchable_worker(class_index, is_dispatchable)
}
#[inline(never)]
fn next_dispatchable_worker(
&mut self,
class_index: usize,
is_dispatchable: &mut impl FnMut(usize, &PolicyClassConfig, &T) -> bool,
) -> Option<DispatchCandidate> {
if self.needs_blocked_worker_recheck {
self.needs_blocked_worker_recheck = false;
let Self {
config,
ready_by_worker,
blocked_workers,
candidate_worker_heads,
..
} = self;
blocked_workers.retain(|worker| {
let ready = ready_by_worker
.get(worker)
.expect("blocked worker lane vanished");
let head = ready.peek().expect("blocked worker lane is empty");
if is_dispatchable(class_index, config, head.payload()) {
candidate_worker_heads.insert(WorkerLaneHead::new(*worker, head));
false
} else {
true
}
});
}
let shared = self
.pending
.peek()
.filter(|entry| is_dispatchable(class_index, &self.config, entry.payload()))
.map(|entry| {
(
entry.priority,
entry.enqueue_seq,
entry.snapshot.scheduling_cost_tokens,
)
});
loop {
let Some(head) = self.candidate_worker_heads.last().copied() else {
return shared.map(|(_, _, cost)| DispatchCandidate {
placement: WorkerPlacement::Any,
cost,
});
};
let dispatchable = {
let ready = self
.ready_by_worker
.get(&head.worker)
.expect("indexed worker lane vanished");
is_dispatchable(
class_index,
&self.config,
ready
.peek()
.expect("indexed worker lane is empty")
.payload(),
)
};
if !dispatchable {
let removed = self.candidate_worker_heads.pop_last();
debug_assert_eq!(removed, Some(head));
self.blocked_workers.insert(head.worker);
continue;
}
let exact_cost = self.ready_by_worker[&head.worker]
.peek()
.expect("indexed worker lane is empty")
.snapshot
.scheduling_cost_tokens;
return match shared {
Some((priority, enqueue_seq, cost))
if cmp_queue_order(priority, enqueue_seq, head.priority, head.enqueue_seq)
== Ordering::Greater =>
{
Some(DispatchCandidate {
placement: WorkerPlacement::Any,
cost,
})
}
_ => Some(DispatchCandidate {
placement: WorkerPlacement::Exact(head.worker),
cost: exact_cost,
}),
};
}
}
fn pop_lane(&mut self, placement: WorkerPlacement) -> PolicyQueueEntry<T> {
match placement {
WorkerPlacement::Any => self.pending.pop().expect("policy class front vanished"),
WorkerPlacement::Exact(worker) => {
let ready = self
.ready_by_worker
.get_mut(&worker)
.expect("queue admission worker lane vanished");
debug_assert!(!self.blocked_workers.contains(&worker));
let head = WorkerLaneHead::new(
worker,
ready.peek().expect("queue admission worker head vanished"),
);
let removed = self.candidate_worker_heads.remove(&head);
debug_assert!(removed);
let entry = ready.pop().expect("queue admission worker head vanished");
let next_head = ready.peek().map(|entry| WorkerLaneHead::new(worker, entry));
if let Some(next_head) = next_head {
self.candidate_worker_heads.insert(next_head);
} else {
self.ready_by_worker.remove(&worker);
}
entry
}
}
}
fn recheck_worker(&mut self, worker: WorkerWithDpRank) {
if self.blocked_workers.remove(&worker) {
let ready = self
.ready_by_worker
.get(&worker)
.expect("blocked worker lane vanished");
self.candidate_worker_heads.insert(WorkerLaneHead::new(
worker,
ready.peek().expect("blocked worker lane is empty"),
));
}
}
fn recheck_all_workers(&mut self) {
self.needs_blocked_worker_recheck = true;
}
fn rebuild_worker_heads(&mut self) {
self.candidate_worker_heads.clear();
self.blocked_workers
.retain(|worker| self.ready_by_worker.contains_key(worker));
for (&worker, ready) in &self.ready_by_worker {
if !self.blocked_workers.contains(&worker) {
self.candidate_worker_heads.insert(WorkerLaneHead::new(
worker,
ready.peek().expect("worker lane is empty"),
));
}
}
}
}
pub struct PolicyQueue<T> {
classes: Vec<PolicyClassQueue<T>>,
round_cursor: usize,
carry_class: Option<usize>,
next_enqueue_seq: u64,
next_admission_id: u64,
pending_count: usize,
candidates: Vec<Option<DispatchCandidate>>,
}
enum QueueEntryState {
Ready(WorkerPlacement),
Deferred(AdmissionId),
}
impl<T> PolicyQueue<T> {
pub fn new(profile: PolicyProfile) -> Self {
let class_count = profile.classes().len();
Self {
classes: profile
.classes()
.iter()
.cloned()
.map(|config| PolicyClassQueue {
config,
admission_policy: None,
pending: BinaryHeap::new(),
deferred: FxHashMap::default(),
ready_by_worker: FxHashMap::default(),
blocked_workers: FxHashSet::default(),
candidate_worker_heads: BTreeSet::new(),
needs_blocked_worker_recheck: false,
stats: PolicyQueueStats::default(),
deficit: 0,
})
.collect(),
round_cursor: 0,
carry_class: None,
next_enqueue_seq: 0,
next_admission_id: 0,
pending_count: 0,
candidates: vec![None; class_count],
}
}
pub(crate) fn new_with_admission_policies(
profile: PolicyProfile,
queue_recheck_interval: Duration,
mut policies: PolicyClassAdmissionPolicies,
) -> Result<Self, KvSchedulerError> {
let mut queue = Self::new(profile);
let now = Instant::now();
for class in &mut queue.classes {
let Some(policy) = policies.remove(&class.config.name) else {
if class.config.admission.is_some() {
return Err(KvSchedulerError::InitFailed(format!(
"no admission policy registered for configured policy class {:?}",
class.config.name
)));
}
continue;
};
let reconcile_interval = policy
.reconcile_interval()
.map_or(queue_recheck_interval, |requested| {
requested.min(queue_recheck_interval)
});
if reconcile_interval.is_zero() {
return Err(KvSchedulerError::InitFailed(format!(
"admission policy for policy class {:?} has a zero reconcile interval",
class.config.name
)));
}
class.admission_policy = Some(ScheduledAdmissionPolicy {
policy,
reconcile_interval,
next_reconcile: now + reconcile_interval,
});
}
if let Some(class_name) = policies.keys().next() {
return Err(KvSchedulerError::InitFailed(format!(
"admission policy registered for unknown policy class {class_name:?}"
)));
}
Ok(queue)
}
pub(crate) fn has_admission_policy(&self, class_index: usize) -> bool {
self.classes[class_index].admission_policy.is_some()
}
pub(crate) fn admit(
&mut self,
class_index: usize,
session_id: Option<&str>,
context_tokens: usize,
worker_eligibility: WorkerEligibility,
) -> Option<(AdmissionTicket, RequestProgressUpdater, AdmissionDecision)> {
let policy = &mut self.classes[class_index].admission_policy.as_mut()?.policy;
let id = AdmissionId::new(self.next_admission_id);
self.next_admission_id = self.next_admission_id.wrapping_add(1);
let ticket = AdmissionTicket { class_index, id };
let (progress, updater) = RequestProgress::new(context_tokens);
let decision = policy.admit(AdmissionRequest::with_progress(
id,
session_id,
progress,
worker_eligibility,
));
Some((ticket, updater, decision))
}
pub(crate) fn dispatched(
&mut self,
ticket: AdmissionTicket,
worker: WorkerWithDpRank,
) -> Vec<ClassAdmissionAction> {
self.admission_event(
ticket,
AdmissionEvent::Dispatched {
id: ticket.id,
worker,
},
)
}
pub(crate) fn completed(
&mut self,
ticket: AdmissionTicket,
context_tokens: usize,
) -> Vec<ClassAdmissionAction> {
self.admission_event(
ticket,
AdmissionEvent::Completed {
id: ticket.id,
context_tokens,
},
)
}
pub(crate) fn aborted(&mut self, ticket: AdmissionTicket) -> Vec<ClassAdmissionAction> {
self.admission_event(ticket, AdmissionEvent::Aborted { id: ticket.id })
}
pub(crate) fn reconcile_admission(
&mut self,
now: Instant,
force: bool,
) -> Vec<ClassAdmissionAction> {
let mut actions = Vec::new();
for (class_index, class) in self.classes.iter_mut().enumerate() {
let Some(scheduled) = &mut class.admission_policy else {
continue;
};
if !force && now < scheduled.next_reconcile {
continue;
}
if now >= scheduled.next_reconcile {
scheduled.next_reconcile = now + scheduled.reconcile_interval;
}
actions.extend(
scheduled
.policy
.on_event(AdmissionEvent::Reconcile)
.into_iter()
.map(|action| ClassAdmissionAction {
class_index,
action,
}),
);
}
actions
}
fn admission_event(
&mut self,
ticket: AdmissionTicket,
event: AdmissionEvent,
) -> Vec<ClassAdmissionAction> {
let Some(scheduled) = &mut self.classes[ticket.class_index].admission_policy else {
return Vec::new();
};
scheduled
.policy
.on_event(event)
.into_iter()
.map(|action| ClassAdmissionAction {
class_index: ticket.class_index,
action,
})
.collect()
}
pub fn pending_count(&self) -> usize {
self.pending_count
}
pub(crate) fn has_ready(&self) -> bool {
self.classes.iter().any(|class| !class.ready_is_empty())
}
pub(crate) fn any_ready_head(
&self,
mut predicate: impl FnMut(usize, &PolicyClassConfig, &T) -> bool,
) -> bool {
self.classes.iter().enumerate().any(|(class_index, class)| {
class
.pending
.peek()
.is_some_and(|entry| predicate(class_index, &class.config, entry.payload()))
|| class.ready_by_worker.values().any(|ready| {
ready
.peek()
.is_some_and(|entry| predicate(class_index, &class.config, entry.payload()))
})
})
}
pub fn class_count(&self) -> usize {
self.classes.len()
}
pub fn class_config(&self, class_index: usize) -> &PolicyClassConfig {
&self.classes[class_index].config
}
pub fn class_stats(&self, class_index: usize) -> PolicyQueueStats {
self.classes[class_index].stats
}
pub(crate) fn recheck_worker(&mut self, worker: WorkerWithDpRank) {
for class in &mut self.classes {
class.recheck_worker(worker);
}
}
pub(crate) fn recheck_all_workers(&mut self) {
for class in &mut self.classes {
class.recheck_all_workers();
}
}
pub fn has_backlog(&self, class_index: usize) -> bool {
!self.classes[class_index].ready_is_empty()
}
pub fn entries(&self) -> impl Iterator<Item = &PolicyQueueEntry<T>> {
self.classes.iter().flat_map(PolicyClassQueue::entries)
}
pub fn retain(&mut self, mut keep: impl FnMut(&T) -> bool) {
for class_index in 0..self.classes.len() {
drop(self.take_if_in_class(class_index, |payload| !keep(payload)));
}
}
#[allow(clippy::too_many_arguments)]
pub fn enqueue(
&mut self,
class_index: usize,
worker_count: usize,
snapshot: QueueSnapshot,
arrival_offset_secs: f64,
priority_jump: f64,
strict_priority: u32,
placement: WorkerPlacement,
payload: T,
) -> Result<(), (QueueRejection, T)> {
self.enqueue_with_state(
class_index,
worker_count,
snapshot,
arrival_offset_secs,
priority_jump,
strict_priority,
QueueEntryState::Ready(placement),
payload,
)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn enqueue_deferred(
&mut self,
class_index: usize,
worker_count: usize,
snapshot: QueueSnapshot,
arrival_offset_secs: f64,
priority_jump: f64,
strict_priority: u32,
admission_id: AdmissionId,
payload: T,
) -> Result<(), (QueueRejection, T)> {
self.enqueue_with_state(
class_index,
worker_count,
snapshot,
arrival_offset_secs,
priority_jump,
strict_priority,
QueueEntryState::Deferred(admission_id),
payload,
)
}
#[allow(clippy::too_many_arguments)]
fn enqueue_with_state(
&mut self,
class_index: usize,
worker_count: usize,
snapshot: QueueSnapshot,
arrival_offset_secs: f64,
priority_jump: f64,
strict_priority: u32,
state: QueueEntryState,
payload: T,
) -> Result<(), (QueueRejection, T)> {
let class = &mut self.classes[class_index];
if let Some(rejection) = queue_rejection(class, worker_count) {
return Err((rejection, payload));
}
let entry = make_entry(
class_index,
snapshot,
arrival_offset_secs,
priority_jump,
strict_priority,
class.config.queue_policy,
self.next_enqueue_seq,
payload,
);
self.next_enqueue_seq = self.next_enqueue_seq.wrapping_add(1);
add_stats(&mut class.stats, snapshot);
match state {
QueueEntryState::Ready(placement) => class.push_ready(placement, entry),
QueueEntryState::Deferred(admission_id) => {
let replaced = class.deferred.insert(admission_id, entry);
debug_assert!(replaced.is_none(), "duplicate deferred admission ID");
}
}
self.pending_count += 1;
Ok(())
}
pub(crate) fn make_ready(
&mut self,
source_class_index: usize,
target_class_index: usize,
admission_id: AdmissionId,
placement: WorkerPlacement,
replacement_snapshot: Option<(QueueSnapshot, f64, f64)>,
) -> Option<QueueSnapshot> {
if source_class_index == target_class_index {
let class = &mut self.classes[source_class_index];
let mut entry = class.deferred.remove(&admission_id)?;
let old_snapshot = entry.snapshot;
if let Some((snapshot, arrival_offset_secs, priority_jump)) = replacement_snapshot {
subtract_stats(&mut class.stats, old_snapshot);
add_stats(&mut class.stats, snapshot);
entry.priority.policy_score = OrderedFloat(queue_policy_score(
class.config.queue_policy,
snapshot,
arrival_offset_secs,
priority_jump,
));
entry.snapshot = snapshot;
}
class.push_ready(placement, entry);
return Some(old_snapshot);
}
let source = &mut self.classes[source_class_index];
let mut entry = source.deferred.remove(&admission_id)?;
let old_snapshot = entry.snapshot;
subtract_stats(&mut source.stats, old_snapshot);
if source.ready_is_empty() {
source.deficit = 0;
}
let (snapshot, arrival_offset_secs, priority_jump) = replacement_snapshot
.expect("cross-class exact placement must replace the queue snapshot");
entry.class_index = target_class_index;
entry.snapshot = snapshot;
let target = &mut self.classes[target_class_index];
entry.priority.policy_score = OrderedFloat(queue_policy_score(
target.config.queue_policy,
snapshot,
arrival_offset_secs,
priority_jump,
));
add_stats(&mut target.stats, snapshot);
target.push_ready(placement, entry);
Some(old_snapshot)
}
pub(crate) fn deferred_payload_mut(
&mut self,
class_index: usize,
admission_id: AdmissionId,
) -> Option<&mut T> {
self.classes[class_index]
.deferred
.get_mut(&admission_id)
.map(PolicyQueueEntry::payload_mut)
}
pub(crate) fn remove_deferred(
&mut self,
class_index: usize,
admission_id: AdmissionId,
) -> Option<PolicyQueueEntry<T>> {
let class = &mut self.classes[class_index];
let entry = class.deferred.remove(&admission_id)?;
subtract_stats(&mut class.stats, entry.snapshot);
self.pending_count -= 1;
if class.ready_is_empty() {
class.deficit = 0;
}
Some(entry)
}
pub(crate) fn take_if_in_class(
&mut self,
class_index: usize,
mut predicate: impl FnMut(&T) -> bool,
) -> (Vec<PolicyQueueEntry<T>>, bool) {
let class = &mut self.classes[class_index];
let remove_sequences: FxHashSet<u64> = class
.entries()
.filter(|entry| predicate(entry.payload()))
.map(|entry| entry.enqueue_seq)
.collect();
if remove_sequences.is_empty() {
return (Vec::new(), false);
}
let removed_ready_head = class
.pending
.peek()
.is_some_and(|entry| remove_sequences.contains(&entry.enqueue_seq))
|| class.ready_by_worker.values().any(|ready| {
ready
.peek()
.is_some_and(|entry| remove_sequences.contains(&entry.enqueue_seq))
});
let mut removed = Vec::new();
let mut retained = Vec::with_capacity(class.pending.len());
for entry in class.pending.drain() {
if remove_sequences.contains(&entry.enqueue_seq) {
removed.push(entry);
} else {
retained.push(entry);
}
}
class.pending = BinaryHeap::from(retained);
removed.extend(
class
.deferred
.extract_if(|_, entry| remove_sequences.contains(&entry.enqueue_seq))
.map(|(_, entry)| entry),
);
class.ready_by_worker.retain(|_, ready| {
let mut retained = Vec::with_capacity(ready.len());
for entry in ready.drain() {
if remove_sequences.contains(&entry.enqueue_seq) {
removed.push(entry);
} else {
retained.push(entry);
}
}
*ready = BinaryHeap::from(retained);
!ready.is_empty()
});
class.rebuild_worker_heads();
for entry in &removed {
subtract_stats(&mut class.stats, entry.snapshot);
self.pending_count -= 1;
}
if class.ready_is_empty() {
class.deficit = 0;
}
(removed, removed_ready_head)
}
pub fn pop_next(
&mut self,
mut is_dispatchable: impl FnMut(usize, &PolicyClassConfig, &T) -> bool,
) -> Option<PolicyQueueEntry<T>> {
if self.pending_count == 0 {
self.carry_class = None;
return None;
}
let class_count = self.classes.len();
self.candidates.fill(None);
let carried_class = self.carry_class.take();
if let Some(class_index) = carried_class {
let class = &mut self.classes[class_index];
let candidate = class.next_dispatchable(class_index, &mut is_dispatchable);
if let Some(candidate) = candidate
&& candidate.cost <= class.deficit
{
return Some(self.pop_candidate(class_index, candidate));
}
self.candidates[class_index] = candidate;
if class.ready_is_empty() {
class.deficit = 0;
}
}
for offset in 0..class_count {
let class_index = (self.round_cursor + offset) % class_count;
let class = &mut self.classes[class_index];
let candidate = if carried_class == Some(class_index) {
self.candidates[class_index]
} else {
class.next_dispatchable(class_index, &mut is_dispatchable)
};
let Some(candidate) = candidate else {
if class.ready_is_empty() {
class.deficit = 0;
}
continue;
};
self.candidates[class_index] = Some(candidate);
if candidate.cost <= class.deficit {
return Some(self.pop_candidate(class_index, candidate));
}
class.deficit = class.deficit.saturating_add(class.config.quantum);
if candidate.cost <= class.deficit {
return Some(self.pop_candidate(class_index, candidate));
}
}
let rounds = self
.candidates
.iter()
.enumerate()
.filter_map(|(class_index, candidate)| {
let candidate = candidate.as_ref()?;
let class = &self.classes[class_index];
let missing = candidate.cost.saturating_sub(class.deficit);
Some(missing.div_ceil(class.config.quantum))
})
.min()?;
for (class_index, candidate) in self.candidates.iter().enumerate() {
if candidate.is_none() {
continue;
}
let class = &mut self.classes[class_index];
class.deficit = class
.deficit
.saturating_add(class.config.quantum.saturating_mul(rounds));
}
for offset in 0..class_count {
let class_index = (self.round_cursor + offset) % class_count;
let class = &self.classes[class_index];
if let Some(candidate) = self.candidates[class_index]
&& candidate.cost <= class.deficit
{
return Some(self.pop_candidate(class_index, candidate));
}
}
None
}
pub fn drain(self) -> impl Iterator<Item = PolicyQueueEntry<T>> {
self.classes.into_iter().flat_map(|class| {
class
.pending
.into_iter()
.chain(class.ready_by_worker.into_values().flatten())
.chain(class.deferred.into_values())
})
}
fn pop_candidate(
&mut self,
class_index: usize,
candidate: DispatchCandidate,
) -> PolicyQueueEntry<T> {
self.round_cursor = (class_index + 1) % self.classes.len();
let class = &mut self.classes[class_index];
let entry = class.pop_lane(candidate.placement);
class.deficit = class
.deficit
.saturating_sub(entry.snapshot.scheduling_cost_tokens);
subtract_stats(&mut class.stats, entry.snapshot);
self.pending_count -= 1;
if class.ready_is_empty() {
class.deficit = 0;
} else {
self.carry_class = (class.deficit > 0).then_some(class_index);
}
entry
}
}
#[allow(clippy::too_many_arguments)]
fn make_entry<T>(
class_index: usize,
snapshot: QueueSnapshot,
arrival_offset_secs: f64,
priority_jump: f64,
strict_priority: u32,
queue_policy: RouterQueuePolicy,
enqueue_seq: u64,
payload: T,
) -> PolicyQueueEntry<T> {
let policy_score =
queue_policy_score(queue_policy, snapshot, arrival_offset_secs, priority_jump);
PolicyQueueEntry {
class_index,
priority: QueuePriority {
strict_priority,
policy_score: OrderedFloat(policy_score),
},
enqueue_seq,
snapshot,
payload,
}
}
fn queue_policy_score(
queue_policy: RouterQueuePolicy,
snapshot: QueueSnapshot,
arrival_offset_secs: f64,
priority_jump: f64,
) -> f64 {
match queue_policy {
RouterQueuePolicy::Fcfs => priority_jump.max(0.0) - arrival_offset_secs.max(0.0),
RouterQueuePolicy::Wspt => {
(1.0 + priority_jump.max(0.0)) / snapshot.scheduling_cost_tokens as f64
}
RouterQueuePolicy::Lcfs => priority_jump.max(0.0) + arrival_offset_secs.max(0.0),
}
}
fn queue_rejection<T>(class: &PolicyClassQueue<T>, worker_count: usize) -> Option<QueueRejection> {
for (limit_kind, current, limit_per_worker) in [
(
QueueLimitKind::Requests,
class.stats.requests,
class.config.request_queue_limit_per_worker,
),
(
QueueLimitKind::RawIslTokens,
class.stats.raw_isl_tokens,
class.config.raw_isl_token_queue_limit_per_worker,
),
(
QueueLimitKind::CachedTokens,
class.stats.cached_tokens,
class.config.cached_token_queue_limit_per_worker,
),
] {
let limit = limit_per_worker.map(|limit| limit.saturating_mul(worker_count));
if limit.is_some_and(|limit| current >= limit) {
return Some(QueueRejection {
policy_class: class.config.name.clone(),
limit_kind,
current,
limit: limit.expect("checked as some"),
});
}
}
None
}
fn add_stats(stats: &mut PolicyQueueStats, snapshot: QueueSnapshot) {
stats.requests += 1;
stats.raw_isl_tokens = stats.raw_isl_tokens.saturating_add(snapshot.raw_isl_tokens);
stats.cached_tokens = stats.cached_tokens.saturating_add(snapshot.cached_tokens);
}
fn subtract_stats(stats: &mut PolicyQueueStats, snapshot: QueueSnapshot) {
stats.requests = stats.requests.saturating_sub(1);
stats.raw_isl_tokens = stats.raw_isl_tokens.saturating_sub(snapshot.raw_isl_tokens);
stats.cached_tokens = stats.cached_tokens.saturating_sub(snapshot.cached_tokens);
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
use super::*;
use crate::scheduling::RouterPolicyConfig;
use crate::scheduling::queue_admission::{AdmissionAction, WorkerEligibilitySnapshot};
struct ReadyPolicy;
impl PolicyClassAdmissionPolicy for ReadyPolicy {
fn admit(&mut self, _request: AdmissionRequest<'_>) -> AdmissionDecision {
AdmissionDecision::Ready(WorkerPlacement::Any)
}
}
struct ZeroIntervalPolicy;
impl PolicyClassAdmissionPolicy for ZeroIntervalPolicy {
fn admit(&mut self, _request: AdmissionRequest<'_>) -> AdmissionDecision {
AdmissionDecision::Bypass
}
fn reconcile_interval(&self) -> Option<Duration> {
Some(Duration::ZERO)
}
}
struct CountingPolicy {
reconciles: Arc<AtomicUsize>,
interval: Duration,
}
impl PolicyClassAdmissionPolicy for CountingPolicy {
fn admit(&mut self, _request: AdmissionRequest<'_>) -> AdmissionDecision {
AdmissionDecision::Bypass
}
fn on_event(&mut self, event: AdmissionEvent) -> Vec<AdmissionAction> {
if event == AdmissionEvent::Reconcile {
self.reconciles.fetch_add(1, AtomicOrdering::Relaxed);
}
Vec::new()
}
fn reconcile_interval(&self) -> Option<Duration> {
Some(self.interval)
}
}
fn profile(yaml: &str) -> PolicyProfile {
RouterPolicyConfig::from_yaml(yaml)
.unwrap()
.resolve_profile(None, None, RouterQueuePolicy::Fcfs)
}
fn admission_profile() -> PolicyProfile {
profile(
r#"
default_policy_family: agents
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: agents
policy_family: agents
cache_bucket: all
queue_policy: fcfs
quantum: 10
"#,
)
}
fn two_class_profile() -> PolicyProfile {
profile(
r#"
default_policy_family: standard
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: standard
policy_family: standard
cache_bucket: all
quantum: 1
- name: agents
quantum: 1
"#,
)
}
#[test]
fn rejects_zero_reconcile_interval() {
let profile = PolicyProfile::synthetic(None, RouterQueuePolicy::Fcfs);
let mut policies = PolicyClassAdmissionPolicies::new();
policies.insert(
profile.default_class().name.clone(),
Box::new(ZeroIntervalPolicy),
);
let error = PolicyQueue::<()>::new_with_admission_policies(
profile,
Duration::from_secs(60),
policies,
)
.err()
.unwrap();
assert!(matches!(error, KvSchedulerError::InitFailed(message) if
message.contains("zero reconcile interval")));
}
#[test]
fn rejects_policy_for_unknown_class() {
let profile = PolicyProfile::synthetic(None, RouterQueuePolicy::Fcfs);
let mut policies = PolicyClassAdmissionPolicies::new();
policies.insert("unknown".to_owned(), Box::new(ReadyPolicy));
let error = PolicyQueue::<()>::new_with_admission_policies(
profile,
Duration::from_secs(60),
policies,
)
.err()
.unwrap();
assert!(matches!(error, KvSchedulerError::InitFailed(message) if
message.contains("unknown policy class") && message.contains("unknown")));
}
#[test]
fn rejects_missing_configured_policy() {
let profile = profile(
r#"
default_policy_family: standard
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: standard
policy_family: standard
cache_bucket: all
quantum: 1
- name: agents
admission:
type: test
quantum: 1
"#,
);
let error = PolicyQueue::<()>::new_with_admission_policies(
profile,
Duration::from_secs(60),
PolicyClassAdmissionPolicies::new(),
)
.err()
.unwrap();
assert!(matches!(error, KvSchedulerError::InitFailed(message) if
message.contains("no admission policy registered") && message.contains("agents")));
}
#[tokio::test(start_paused = true)]
async fn reconciles_only_policies_whose_deadlines_are_due() {
let fast = Arc::new(AtomicUsize::new(0));
let slow = Arc::new(AtomicUsize::new(0));
let mut policies = PolicyClassAdmissionPolicies::new();
policies.insert(
"standard".to_owned(),
Box::new(CountingPolicy {
reconciles: Arc::clone(&fast),
interval: Duration::from_millis(10),
}),
);
policies.insert(
"agents".to_owned(),
Box::new(CountingPolicy {
reconciles: Arc::clone(&slow),
interval: Duration::from_secs(60),
}),
);
let mut queue = PolicyQueue::<()>::new_with_admission_policies(
two_class_profile(),
Duration::from_secs(60),
policies,
)
.unwrap();
tokio::time::advance(Duration::from_millis(10)).await;
queue.reconcile_admission(Instant::now(), false);
assert_eq!(fast.load(AtomicOrdering::Relaxed), 1);
assert_eq!(slow.load(AtomicOrdering::Relaxed), 0);
tokio::time::advance(Duration::from_millis(59_990)).await;
queue.reconcile_admission(Instant::now(), false);
assert_eq!(fast.load(AtomicOrdering::Relaxed), 2);
assert_eq!(slow.load(AtomicOrdering::Relaxed), 1);
}
#[test]
fn admission_ids_are_unique_across_policy_classes() {
let mut policies = PolicyClassAdmissionPolicies::new();
policies.insert("standard".to_owned(), Box::new(ReadyPolicy));
policies.insert("agents".to_owned(), Box::new(ReadyPolicy));
let mut queue = PolicyQueue::<()>::new_with_admission_policies(
two_class_profile(),
Duration::from_secs(60),
policies,
)
.unwrap();
let eligibility = || WorkerEligibility::new(|| WorkerEligibilitySnapshot::new([]));
let first = queue.admit(0, None, 1, eligibility()).unwrap().0.id;
let second = queue.admit(1, None, 1, eligibility()).unwrap().0.id;
assert_ne!(first, second);
}
#[test]
fn deferred_entries_count_toward_limits_without_blocking_ready_work() {
let mut queue = PolicyQueue::new(admission_profile());
let deferred_id = AdmissionId::new(7);
queue
.enqueue_deferred(
0,
1,
QueueSnapshot::new(20, 5),
0.0,
0.0,
0,
deferred_id,
"deferred",
)
.unwrap();
assert!(!queue.has_backlog(0));
queue
.enqueue(
0,
1,
QueueSnapshot::new(10, 0),
1.0,
0.0,
0,
WorkerPlacement::Any,
"ready",
)
.unwrap();
assert!(queue.has_backlog(0));
assert_eq!(queue.pending_count(), 2);
assert_eq!(queue.class_stats(0).requests, 2);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"ready"
);
assert!(queue.pop_next(|_, _, _| true).is_none());
assert!(
queue
.make_ready(0, 0, deferred_id, WorkerPlacement::Any, None)
.is_some()
);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"deferred"
);
assert_eq!(queue.pending_count(), 0);
assert_eq!(queue.class_stats(0), PolicyQueueStats::default());
}
#[test]
fn exact_make_ready_rekeys_wspt_priority() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: agents
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: agents
policy_family: agents
cache_bucket: all
queue_policy: wspt
quantum: 1000
"#,
));
let deferred_id = AdmissionId::new(1);
queue
.enqueue_deferred(
0,
1,
QueueSnapshot::new(100, 99),
0.0,
0.0,
0,
deferred_id,
"rekeyed",
)
.unwrap();
queue
.enqueue(
0,
1,
QueueSnapshot::new(10, 0),
1.0,
0.0,
0,
WorkerPlacement::Any,
"ready",
)
.unwrap();
queue
.make_ready(
0,
0,
deferred_id,
WorkerPlacement::Any,
Some((QueueSnapshot::new(100, 0), 0.0, 0.0)),
)
.unwrap();
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"ready"
);
}
#[test]
fn exact_make_ready_moves_entry_and_accounting_to_target_class() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: agents
uncached_isl_buckets:
- min_tokens: 0
bucket: cached
- min_tokens: 32
bucket: uncached
policy_classes:
- name: agents_cached
policy_family: agents
cache_bucket: cached
queue_policy: fcfs
quantum: 1000
- name: agents_uncached
policy_family: agents
cache_bucket: uncached
queue_policy: wspt
quantum: 1000
"#,
));
let deferred_id = AdmissionId::new(1);
queue
.enqueue_deferred(
0,
1,
QueueSnapshot::new(100, 99),
0.0,
0.0,
0,
deferred_id,
"migrated",
)
.unwrap();
queue
.enqueue(
1,
1,
QueueSnapshot::new(10, 0),
1.0,
0.0,
0,
WorkerPlacement::Any,
"ready",
)
.unwrap();
queue
.make_ready(
0,
1,
deferred_id,
WorkerPlacement::Any,
Some((QueueSnapshot::new(100, 0), 0.0, 0.0)),
)
.unwrap();
assert_eq!(queue.class_stats(0), PolicyQueueStats::default());
assert_eq!(
queue.class_stats(1),
PolicyQueueStats {
requests: 2,
raw_isl_tokens: 110,
cached_tokens: 0,
}
);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"ready"
);
let migrated = queue.pop_next(|_, _, _| true).unwrap();
assert_eq!(migrated.class_index(), 1);
assert_eq!(migrated.into_payload(), "migrated");
}
#[test]
fn per_worker_caps_scale_and_remain_pre_add() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: capped
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: capped
policy_family: capped
cache_bucket: all
quantum: 10
request_queue_limit_per_worker: 1
raw_isl_token_queue_limit_per_worker: 5
cached_token_queue_limit_per_worker: 3
"#,
));
queue
.enqueue(
0,
2,
QueueSnapshot::new(8, 4),
0.0,
0.0,
0,
WorkerPlacement::Any,
"first",
)
.unwrap();
queue
.enqueue(
0,
2,
QueueSnapshot::new(100, 100),
1.0,
0.0,
0,
WorkerPlacement::Any,
"overshoot",
)
.unwrap();
let (rejection, payload) = queue
.enqueue(
0,
2,
QueueSnapshot::new(1, 0),
2.0,
0.0,
0,
WorkerPlacement::Any,
"rejected",
)
.unwrap_err();
assert_eq!(payload, "rejected");
assert_eq!(rejection.limit_kind, QueueLimitKind::Requests);
assert_eq!(rejection.current, 2);
assert_eq!(rejection.limit, 2);
assert_eq!(queue.class_stats(0).raw_isl_tokens, 108);
assert_eq!(queue.class_stats(0).cached_tokens, 104);
}
#[test]
fn retain_removes_payload_and_rebuilds_queue_accounting() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: default
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: default
policy_family: default
cache_bucket: all
quantum: 10
"#,
));
queue
.enqueue(
0,
1,
QueueSnapshot::new(8, 4),
0.0,
0.0,
0,
WorkerPlacement::Any,
"keep",
)
.unwrap();
queue
.enqueue(
0,
1,
QueueSnapshot::new(16, 6),
1.0,
0.0,
0,
WorkerPlacement::Any,
"remove",
)
.unwrap();
queue.retain(|payload| *payload != "remove");
assert_eq!(queue.pending_count(), 1);
assert_eq!(queue.class_stats(0).requests, 1);
assert_eq!(queue.class_stats(0).raw_isl_tokens, 8);
assert_eq!(queue.class_stats(0).cached_tokens, 4);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"keep"
);
}
#[test]
fn blocked_worker_lane_does_not_block_another_lane() {
let mut queue = PolicyQueue::new(admission_profile());
for (worker, payload) in [(1, "blocked"), (2, "ready")] {
queue
.enqueue(
0,
2,
QueueSnapshot::new(1, 0),
worker as f64,
0.0,
0,
WorkerPlacement::Exact(WorkerWithDpRank::new(worker, 0)),
payload,
)
.unwrap();
}
let candidate = queue
.pop_next(|_, _, payload| *payload != "blocked")
.unwrap();
assert_eq!(candidate.into_payload(), "ready");
assert_eq!(queue.pending_count(), 1);
}
#[test]
fn worker_update_rechecks_only_that_blocked_lane() {
let mut queue = PolicyQueue::new(admission_profile());
let worker_1 = WorkerWithDpRank::new(1, 0);
let worker_2 = WorkerWithDpRank::new(2, 0);
for (worker, payload) in [(worker_1, "worker-1"), (worker_2, "worker-2")] {
queue
.enqueue(
0,
2,
QueueSnapshot::new(1, 0),
worker.worker_id as f64,
0.0,
0,
WorkerPlacement::Exact(worker),
payload,
)
.unwrap();
}
assert!(queue.pop_next(|_, _, _| false).is_none());
queue.recheck_worker(worker_1);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"worker-1"
);
assert!(queue.pop_next(|_, _, _| true).is_none());
queue.recheck_worker(worker_2);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"worker-2"
);
}
#[test]
fn exact_lane_index_tracks_a_new_higher_priority_head() {
let mut queue = PolicyQueue::new(admission_profile());
let worker = WorkerWithDpRank::new(1, 0);
queue
.enqueue(
0,
1,
QueueSnapshot::new(1, 0),
0.0,
0.0,
0,
WorkerPlacement::Exact(worker),
"old",
)
.unwrap();
queue
.enqueue(
0,
1,
QueueSnapshot::new(1, 0),
1.0,
10.0,
0,
WorkerPlacement::Exact(worker),
"boosted",
)
.unwrap();
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"boosted"
);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"old"
);
}
#[test]
fn exact_lane_head_index_checks_each_blocked_lane_once() {
const LANES: usize = 10_000;
let mut queue = PolicyQueue::new(admission_profile());
for lane in 0..LANES {
queue
.enqueue(
0,
LANES,
QueueSnapshot::new(1, 0),
lane as f64,
0.0,
0,
WorkerPlacement::Exact(WorkerWithDpRank::new(lane as u64, 0)),
lane,
)
.unwrap();
}
let mut checks = 0;
let mut popped = 0;
while queue
.pop_next(|_, _, lane| {
checks += 1;
!lane.is_multiple_of(2)
})
.is_some()
{
popped += 1;
}
assert_eq!(popped, LANES / 2);
assert_eq!(checks, LANES);
assert_eq!(queue.pending_count(), LANES / 2);
}
#[test]
fn global_recheck_dispatches_workers_as_they_become_available() {
let mut queue = PolicyQueue::new(admission_profile());
for lane in 0..3 {
queue
.enqueue(
0,
3,
QueueSnapshot::new(1, 0),
lane as f64,
0.0,
0,
WorkerPlacement::Exact(WorkerWithDpRank::new(lane, 0)),
lane,
)
.unwrap();
}
assert!(queue.pop_next(|_, _, _| false).is_none());
let mut available = [false, true, false];
queue.recheck_all_workers();
let entry = queue
.pop_next(|_, _, lane| available[*lane as usize])
.expect("newly available worker should dispatch");
assert_eq!(entry.into_payload(), 1);
assert!(
queue
.pop_next(|_, _, lane| available[*lane as usize])
.is_none()
);
available[2] = true;
queue.recheck_all_workers();
let entry = queue
.pop_next(|_, _, lane| available[*lane as usize])
.expect("later worker availability should dispatch on the next recheck");
assert_eq!(entry.into_payload(), 2);
assert_eq!(queue.pending_count(), 1);
}
#[test]
fn global_recheck_preserves_priority_for_multiple_available_workers() {
let mut queue = PolicyQueue::new(admission_profile());
for (worker, strict_priority, payload) in [(1, 0, "low"), (2, 1, "high")] {
queue
.enqueue(
0,
2,
QueueSnapshot::new(1, 0),
0.0,
0.0,
strict_priority,
WorkerPlacement::Exact(WorkerWithDpRank::new(worker, 0)),
payload,
)
.unwrap();
}
assert!(queue.pop_next(|_, _, _| false).is_none());
queue.recheck_all_workers();
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"high"
);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"low"
);
}
#[test]
fn pending_global_recheck_survives_retain_removing_blocked_lane() {
let mut queue = PolicyQueue::new(admission_profile());
for (worker, payload) in [(1, "remove"), (2, "keep")] {
queue
.enqueue(
0,
2,
QueueSnapshot::new(1, 0),
0.0,
0.0,
0,
WorkerPlacement::Exact(WorkerWithDpRank::new(worker, 0)),
payload,
)
.unwrap();
}
assert!(queue.pop_next(|_, _, _| false).is_none());
queue.recheck_all_workers();
queue.retain(|payload| *payload == "keep");
assert_eq!(queue.pending_count(), 1);
assert_eq!(
queue.pop_next(|_, _, _| true).unwrap().into_payload(),
"keep"
);
assert_eq!(queue.pending_count(), 0);
}
#[test]
fn per_worker_token_caps_follow_capacity_without_evicting() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: raw
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: raw
policy_family: raw
cache_bucket: all
quantum: 1
raw_isl_token_queue_limit_per_worker: 10
- name: cached
policy_family: cached
cache_bucket: all
quantum: 1
cached_token_queue_limit_per_worker: 5
- name: zero
policy_family: zero
cache_bucket: all
quantum: 1
request_queue_limit_per_worker: 0
- name: no-workers
policy_family: no-workers
cache_bucket: all
quantum: 1
request_queue_limit_per_worker: 1
"#,
));
queue
.enqueue(
0,
2,
QueueSnapshot::new(11, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"raw-queued",
)
.unwrap();
let (raw_rejection, _) = queue
.enqueue(
0,
1,
QueueSnapshot::new(1, 0),
1.0,
0.0,
0,
WorkerPlacement::Any,
"raw-rejected",
)
.unwrap_err();
assert_eq!(raw_rejection.limit_kind, QueueLimitKind::RawIslTokens);
assert_eq!(raw_rejection.current, 11);
assert_eq!(raw_rejection.limit, 10);
assert_eq!(queue.class_stats(0).raw_isl_tokens, 11);
queue
.enqueue(
0,
2,
QueueSnapshot::new(10, 0),
2.0,
0.0,
0,
WorkerPlacement::Any,
"raw-after-growth",
)
.unwrap();
let (grown_rejection, _) = queue
.enqueue(
0,
2,
QueueSnapshot::new(1, 0),
3.0,
0.0,
0,
WorkerPlacement::Any,
"raw-at-grown-cap",
)
.unwrap_err();
assert_eq!(grown_rejection.current, 21);
assert_eq!(grown_rejection.limit, 20);
queue
.enqueue(
1,
2,
QueueSnapshot::new(8, 6),
0.0,
0.0,
0,
WorkerPlacement::Any,
"cached-queued",
)
.unwrap();
let (cached_rejection, _) = queue
.enqueue(
1,
1,
QueueSnapshot::new(1, 1),
1.0,
0.0,
0,
WorkerPlacement::Any,
"cached-rejected",
)
.unwrap_err();
assert_eq!(cached_rejection.limit_kind, QueueLimitKind::CachedTokens);
assert_eq!(cached_rejection.current, 6);
assert_eq!(cached_rejection.limit, 5);
let (zero_rejection, _) = queue
.enqueue(
2,
4,
QueueSnapshot::new(1, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"zero",
)
.unwrap_err();
assert_eq!(zero_rejection.limit_kind, QueueLimitKind::Requests);
assert_eq!(zero_rejection.limit, 0);
let (no_workers_rejection, _) = queue
.enqueue(
3,
0,
QueueSnapshot::new(1, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"no-workers",
)
.unwrap_err();
assert_eq!(no_workers_rejection.current, 0);
assert_eq!(no_workers_rejection.limit, 0);
}
#[test]
fn per_worker_limit_multiplication_saturates() {
let mut queue = PolicyQueue::new(profile(&format!(
r#"
default_policy_family: capped
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: capped
policy_family: capped
cache_bucket: all
quantum: 1
request_queue_limit_per_worker: {}
"#,
usize::MAX
)));
queue
.enqueue(
0,
2,
QueueSnapshot::new(1, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"queued",
)
.unwrap();
}
#[test]
fn fcfs_and_wspt_order_only_within_each_class() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: fcfs
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: fcfs
policy_family: fcfs
cache_bucket: all
queue_policy: fcfs
quantum: 50
- name: wspt
policy_family: wspt
cache_bucket: all
queue_policy: wspt
quantum: 50
"#,
));
queue
.enqueue(
0,
1,
QueueSnapshot::new(50, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"fcfs-long",
)
.unwrap();
queue
.enqueue(
0,
1,
QueueSnapshot::new(1, 0),
1.0,
0.0,
0,
WorkerPlacement::Any,
"fcfs-short",
)
.unwrap();
queue
.enqueue(
1,
1,
QueueSnapshot::new(50, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"wspt-long",
)
.unwrap();
queue
.enqueue(
1,
1,
QueueSnapshot::new(1, 0),
1.0,
0.0,
0,
WorkerPlacement::Any,
"wspt-short",
)
.unwrap();
let first = queue.pop_next(|_, _, _| true).unwrap();
let second = queue.pop_next(|_, _, _| true).unwrap();
assert_eq!(first.into_payload(), "fcfs-long");
assert_eq!(second.into_payload(), "wspt-short");
}
#[test]
fn drr_weights_progress_and_skips_blocked_classes_without_credit() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: slow
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: slow
policy_family: slow
cache_bucket: all
quantum: 1
- name: fast
policy_family: fast
cache_bucket: all
quantum: 3
"#,
));
for index in 0..6 {
queue
.enqueue(
0,
1,
QueueSnapshot::new(1, 0),
index as f64,
0.0,
0,
WorkerPlacement::Any,
"slow",
)
.unwrap();
queue
.enqueue(
1,
1,
QueueSnapshot::new(1, 0),
index as f64,
0.0,
0,
WorkerPlacement::Any,
"fast",
)
.unwrap();
}
let mut first_six = Vec::new();
for _ in 0..6 {
first_six.push(queue.pop_next(|_, _, _| true).unwrap().into_payload());
}
assert!(first_six.iter().filter(|value| **value == "fast").count() >= 3);
let blocked_deficit = queue.classes[1].deficit;
let slow = queue.pop_next(|class, _, _| class == 0).unwrap();
assert_eq!(slow.into_payload(), "slow");
assert_eq!(queue.classes[1].deficit, blocked_deficit);
}
#[test]
fn drr_carry_uses_next_dispatchable_lane() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: agents
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: agents
policy_family: agents
cache_bucket: all
quantum: 10
- name: batch
policy_family: batch
cache_bucket: all
quantum: 1
"#,
));
queue
.enqueue(
0,
2,
QueueSnapshot::new(5, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"agents-first",
)
.unwrap();
queue
.enqueue(
0,
2,
QueueSnapshot::new(1, 0),
1.0,
0.0,
0,
WorkerPlacement::Exact(WorkerWithDpRank::new(1, 0)),
"agents-blocked",
)
.unwrap();
queue
.enqueue(
0,
2,
QueueSnapshot::new(10, 0),
2.0,
0.0,
0,
WorkerPlacement::Exact(WorkerWithDpRank::new(2, 0)),
"agents-ready",
)
.unwrap();
queue
.enqueue(
1,
2,
QueueSnapshot::new(1, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"batch",
)
.unwrap();
let is_dispatchable =
|_: usize, _: &PolicyClassConfig, payload: &&str| *payload != "agents-blocked";
assert_eq!(
queue.pop_next(is_dispatchable).unwrap().into_payload(),
"agents-first"
);
assert_eq!(queue.classes[0].deficit, 5);
assert_eq!(
queue.pop_next(is_dispatchable).unwrap().into_payload(),
"batch"
);
assert_eq!(queue.classes[0].deficit, 5);
}
#[test]
fn drr_serves_exact_quantum_ratio_for_equal_cost_backlogs() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: one
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: one
policy_family: one
cache_bucket: all
quantum: 1
- name: three
policy_family: three
cache_bucket: all
quantum: 3
"#,
));
for index in 0..20 {
queue
.enqueue(
0,
1,
QueueSnapshot::new(1, 0),
index as f64,
0.0,
0,
WorkerPlacement::Any,
"one",
)
.unwrap();
}
for index in 0..60 {
queue
.enqueue(
1,
1,
QueueSnapshot::new(1, 0),
index as f64,
0.0,
0,
WorkerPlacement::Any,
"three",
)
.unwrap();
}
let dispatches = (0..80)
.map(|_| queue.pop_next(|_, _, _| true).unwrap().into_payload())
.collect::<Vec<_>>();
for round in dispatches.chunks_exact(4) {
assert_eq!(round, ["one", "three", "three", "three"]);
}
}
#[test]
fn fully_blocked_ring_returns_without_accruing_deficit() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: first
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: first
policy_family: first
cache_bucket: all
quantum: 7
- name: second
policy_family: second
cache_bucket: all
quantum: 11
"#,
));
queue
.enqueue(
0,
1,
QueueSnapshot::new(100, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"first",
)
.unwrap();
queue
.enqueue(
1,
1,
QueueSnapshot::new(100, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"second",
)
.unwrap();
for _ in 0..10_000 {
assert!(queue.pop_next(|_, _, _| false).is_none());
}
assert_eq!(queue.classes[0].deficit, 0);
assert_eq!(queue.classes[1].deficit, 0);
}
#[test]
fn oversized_heads_bulk_add_deficit_and_make_progress() {
let mut queue = PolicyQueue::new(profile(
r#"
default_policy_family: large
uncached_isl_buckets:
- min_tokens: 0
bucket: all
policy_classes:
- name: large
policy_family: large
cache_bucket: all
quantum: 4
- name: blocked
policy_family: blocked
cache_bucket: all
quantum: 100
"#,
));
queue
.enqueue(
0,
1,
QueueSnapshot::new(101, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"large",
)
.unwrap();
queue
.enqueue(
1,
1,
QueueSnapshot::new(1, 0),
0.0,
0.0,
0,
WorkerPlacement::Any,
"blocked",
)
.unwrap();
let popped = queue
.pop_next(|class, _, _| class == 0)
.expect("oversized request should make bounded progress");
assert_eq!(popped.into_payload(), "large");
assert_eq!(queue.pending_count(), 1);
}
}