use std::collections::VecDeque;
use crate::batch::config::BatchConfig;
use crate::batch::sequence::SeqId;
#[derive(Debug, Default)]
pub struct SchedulerDecision {
pub prefill: Vec<(SeqId, usize, usize)>,
pub decode: Vec<SeqId>,
pub evict: Vec<SeqId>,
}
pub trait Scheduler: Send {
fn select_batch(
&mut self,
waiting: &[SeqId],
running: &[SeqId],
kv_free_pages: usize,
gdn_free_slots: usize,
) -> SchedulerDecision;
fn on_preempt(&mut self, seq_id: SeqId);
}
#[derive(Debug)]
pub struct FifoScheduler {
config: BatchConfig,
admission_queue: VecDeque<SeqId>,
}
impl FifoScheduler {
pub fn new(config: BatchConfig) -> Self {
Self {
config,
admission_queue: VecDeque::new(),
}
}
pub fn enqueue(&mut self, seq_id: SeqId) {
self.admission_queue.push_back(seq_id);
}
#[inline]
pub fn waiting_count(&self) -> usize {
self.admission_queue.len()
}
}
impl Scheduler for FifoScheduler {
fn select_batch(
&mut self,
waiting: &[SeqId],
running: &[SeqId],
kv_free_pages: usize,
_gdn_free_slots: usize,
) -> SchedulerDecision {
let mut decision = SchedulerDecision::default();
let active_count = waiting.len() + running.len();
let capacity_remaining = self.config.max_batch_size.saturating_sub(active_count);
let can_admit = kv_free_pages >= self.config.prefill_reserve_pages || active_count == 0;
let admit_limit = if can_admit { capacity_remaining } else { 0 };
for _ in 0..admit_limit {
if let Some(seq_id) = self.admission_queue.pop_front() {
decision.prefill.push((seq_id, 0, self.config.chunk_size));
} else {
break;
}
}
for &seq_id in running {
decision.decode.push(seq_id);
}
let admitted_ids: std::collections::HashSet<SeqId> =
decision.prefill.iter().map(|&(id, _, _)| id).collect();
for &seq_id in waiting {
if !admitted_ids.contains(&seq_id) {
decision
.prefill
.push((seq_id, usize::MAX, self.config.chunk_size));
}
}
decision
}
fn on_preempt(&mut self, _seq_id: SeqId) {
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::batch::config::BatchConfig;
fn default_sched() -> FifoScheduler {
FifoScheduler::new(BatchConfig::default())
}
fn small_sched(max_batch: usize, chunk: usize, reserve: usize) -> FifoScheduler {
FifoScheduler::new(BatchConfig {
max_batch_size: max_batch,
max_seq_len: 4096,
chunk_size: chunk,
prefill_reserve_pages: reserve,
})
}
#[test]
fn empty_state_returns_empty_decision() {
let mut sched = default_sched();
let dec = sched.select_batch(&[], &[], 100, 100);
assert!(dec.prefill.is_empty());
assert!(dec.decode.is_empty());
assert!(dec.evict.is_empty());
}
#[test]
fn new_sequence_admitted_into_prefill() {
let mut sched = default_sched();
sched.enqueue(SeqId(1));
let dec = sched.select_batch(&[], &[], 100, 100);
assert_eq!(dec.prefill.len(), 1);
let (id, start, len) = dec.prefill[0];
assert_eq!(id, SeqId(1));
assert_eq!(start, 0);
assert_eq!(len, 512); }
#[test]
fn admission_respects_max_batch_size() {
let mut sched = small_sched(2, 512, 0);
sched.enqueue(SeqId(1));
sched.enqueue(SeqId(2));
sched.enqueue(SeqId(3)); let dec = sched.select_batch(&[], &[], 100, 100);
assert_eq!(dec.prefill.len(), 2);
assert_eq!(sched.waiting_count(), 1);
}
#[test]
fn admission_refuses_when_active_count_full() {
let mut sched = small_sched(2, 512, 0);
sched.enqueue(SeqId(10));
let dec = sched.select_batch(&[], &[SeqId(1), SeqId(2)], 100, 100);
assert!(dec.prefill.is_empty());
assert_eq!(sched.waiting_count(), 1);
}
#[test]
fn admission_blocked_by_memory_guard() {
let mut sched = small_sched(32, 512, 8);
sched.enqueue(SeqId(1));
let dec = sched.select_batch(&[SeqId(0)], &[], 3, 100);
assert!(dec.prefill.iter().all(|&(id, _, _)| id == SeqId(0)));
assert_eq!(sched.waiting_count(), 1);
}
#[test]
fn admission_allowed_when_batch_empty_despite_low_pages() {
let mut sched = small_sched(32, 512, 8);
sched.enqueue(SeqId(1));
let dec = sched.select_batch(&[], &[], 0, 100);
assert_eq!(dec.prefill.len(), 1);
}
#[test]
fn running_sequences_go_to_decode() {
let mut sched = default_sched();
let dec = sched.select_batch(&[], &[SeqId(5), SeqId(6)], 100, 100);
assert_eq!(dec.decode.len(), 2);
assert!(dec.decode.contains(&SeqId(5)));
assert!(dec.decode.contains(&SeqId(6)));
}
#[test]
fn all_running_sequences_decode_regardless_of_free_gdn_slots() {
let mut sched = default_sched();
let dec = sched.select_batch(&[], &[SeqId(1), SeqId(2), SeqId(3)], 100, 0);
assert_eq!(dec.decode.len(), 3);
assert!(dec.decode.contains(&SeqId(1)));
assert!(dec.decode.contains(&SeqId(2)));
assert!(dec.decode.contains(&SeqId(3)));
}
#[test]
fn waiting_sequences_get_continuation_chunk() {
let mut sched = small_sched(32, 128, 0);
let dec = sched.select_batch(&[SeqId(2)], &[], 100, 100);
let found = dec
.prefill
.iter()
.any(|&(id, start, len)| id == SeqId(2) && start == usize::MAX && len == 128);
assert!(found, "continuation chunk for waiting sequence not found");
}
#[test]
fn no_duplicate_seqid_in_prefill_between_new_and_waiting() {
let mut sched = small_sched(32, 512, 0);
sched.enqueue(SeqId(99));
let dec = sched.select_batch(&[SeqId(10)], &[], 100, 100);
let ids: Vec<SeqId> = dec.prefill.iter().map(|&(id, _, _)| id).collect();
let mut seen = std::collections::HashSet::new();
for id in &ids {
assert!(seen.insert(id), "duplicate SeqId in prefill: {id}");
}
}
#[test]
fn phase1_no_evictions() {
let mut sched = default_sched();
let dec = sched.select_batch(&[], &[SeqId(1)], 100, 100);
assert!(dec.evict.is_empty());
}
#[test]
fn on_preempt_does_not_panic() {
let mut sched = default_sched();
sched.on_preempt(SeqId(42)); }
#[test]
fn chunk_size_propagated_to_new_admissions() {
let mut sched = small_sched(32, 256, 0);
sched.enqueue(SeqId(1));
let dec = sched.select_batch(&[], &[], 100, 100);
let (_, _, len) = dec.prefill[0];
assert_eq!(len, 256);
}
}