//! Continuous batching scheduler with SLA lanes //! //! This module implements a sophisticated scheduler that forms dynamic batches //! based on SLA requirements, memory constraints, and performance optimization. use serde::{Deserialize, Serialize}; use std::cmp::Ordering; use std::collections::{HashMap, VecDeque}; use std::sync::Arc; use std::time::{Duration, Instant}; use tokio::sync::{Mutex, RwLock}; use tracing::{debug, trace, warn}; use uuid::Uuid; use crate::cache::{AttentionScoreEviction, PageId, PagedKvCache}; use crate::error::{InferenceError, InferenceResult}; use crate::request::{InferenceRequest, RequestId, RequestPriority}; /// Unique identifier for batches #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub struct BatchId(Uuid); impl BatchId { /// Generate a new unique batch ID #[must_use] pub fn new() -> Self { Self(Uuid::new_v4()) } } impl std::fmt::Display for BatchId { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { self.0.fmt(f) } } /// SLA lane configuration for request segregation #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SlaLane { /// Lane name for identification pub name: String, /// Priority level for this lane pub priority: RequestPriority, /// Maximum acceptable latency for requests in this lane pub max_latency: Duration, /// Maximum batch size for this lane pub max_batch_size: usize, /// Memory limit per batch (bytes) pub memory_limit: Option, /// Minimum batch wait time before processing pub min_wait_time: Option, } /// Preemption decision result #[derive(Debug, Clone)] pub struct PreemptionDecision { /// Whether preemption should occur pub should_preempt: bool, /// Target batch to preempt (if any) pub target_batch_id: Option, /// Reason for preemption decision pub reason: String, /// Priority difference that triggered decision pub priority_delta: i32, } /// SLA violation information #[derive(Debug, Clone)] pub struct SlaViolation { /// Request that violated SLA pub request_id: RequestId, /// Lane the request was assigned to pub lane_name: String, /// Type of violation pub violation_type: String, /// Time of violation pub violation_time: Instant, /// Expected vs actual latency pub expected_latency: Duration, pub actual_latency: Duration, } /// Batch of requests ready for processing #[derive(Debug, Clone)] #[allow(dead_code)] pub struct InferenceBatch { /// Unique batch identifier pub id: BatchId, /// Requests in this batch pub requests: Vec, /// Lane this batch belongs to pub lane_name: String, /// Batch creation timestamp pub created_at: Instant, /// Estimated memory usage pub estimated_memory_usage: usize, /// Batch priority (highest among constituent requests) pub priority: RequestPriority, /// Expected processing duration pub estimated_duration: Duration, } /// Lane statistics for monitoring #[derive(Debug, Clone, Serialize, Deserialize)] pub struct LaneStats { /// Number of requests pending in this lane pub pending_count: usize, /// Number of requests currently processing pub processing_count: usize, /// Average latency for completed requests pub average_latency: Duration, /// SLA violation rate pub violation_rate: f64, /// Throughput (requests per second) pub throughput: f64, /// Memory utilization pub memory_utilization: f64, } /// Queue statistics across all lanes #[derive(Debug, Clone, Serialize, Deserialize)] pub struct QueueStats { /// Total pending requests across all lanes pub total_pending: usize, /// Total processing requests pub total_processing: usize, /// Statistics per lane pub lanes: HashMap, /// Overall queue utilization pub queue_utilization: f64, /// Memory pressure level (0.0-1.0) pub memory_pressure: f64, } /// Configuration for batch scheduler #[derive(Debug, Clone)] pub struct BatchSchedulerConfig { /// Maximum batch size across all lanes pub max_batch_size: usize, /// Maximum wait time before forming batch pub max_wait_time: Duration, /// SLA lanes configuration pub sla_lanes: Vec, /// Enable preemption of lower priority batches pub enable_preemption: bool, /// Memory pressure threshold for degradation pub memory_pressure_threshold: f64, /// Maximum memory available for batching pub memory_limit: usize, /// Degradation mode activation threshold pub degradation_mode_threshold: usize, /// Maximum queue size before rejection pub max_queue_size: usize, } impl Default for BatchSchedulerConfig { fn default() -> Self { Self { max_batch_size: 32, max_wait_time: Duration::from_millis(100), sla_lanes: vec![SlaLane { name: "standard".to_string(), priority: RequestPriority::Normal, max_latency: Duration::from_millis(200), max_batch_size: 32, memory_limit: None, min_wait_time: None, }], enable_preemption: false, memory_pressure_threshold: 0.8, memory_limit: 8 * 1024 * 1024 * 1024, // 8GB degradation_mode_threshold: 1000, max_queue_size: 2000, } } } /// Priority-ordered batch request for scheduling #[derive(Debug)] struct SchedulerRequest { request: InferenceRequest, assigned_lane: String, queued_at: Instant, priority_score: i32, /// KV pages that can be reused for the prompt prefix (copy-on-write). /// /// `Some(pages)` when a prefix-cache hit was found in [`PagedKvCache`] at /// submission time. The executing engine should clone these page IDs and /// use them directly for the prefix tokens rather than allocating fresh pages. /// /// `None` means no hit was found (or prefix caching is disabled); the engine /// must allocate pages normally and then call /// [`BatchScheduler::notify_prefill_complete`] to register them. prefix_hit_pages: Option>, /// KV token positions that should be skipped (masked out) during decode. /// /// Populated after prefill completes via /// [`BatchScheduler::notify_prefill_complete`] when an /// [`AttentionScoreEviction`] is attached. An empty `Vec` means nothing is /// evicted (either SnapKV is disabled or no positions qualified). evicted_positions: Vec, } impl PartialEq for SchedulerRequest { fn eq(&self, other: &Self) -> bool { self.priority_score == other.priority_score } } impl Eq for SchedulerRequest {} impl PartialOrd for SchedulerRequest { fn partial_cmp(&self, other: &Self) -> Option { Some(self.cmp(other)) } } impl Ord for SchedulerRequest { fn cmp(&self, other: &Self) -> Ordering { // Higher priority score comes first other.priority_score.cmp(&self.priority_score) } } /// Lane state tracking #[derive(Debug)] struct LaneState { config: SlaLane, pending_requests: VecDeque, processing_batches: HashMap, completed_requests: VecDeque<(Instant, Duration)>, // (completion_time, latency) violation_count: usize, total_processed: usize, } impl LaneState { fn new(config: SlaLane) -> Self { Self { config, pending_requests: VecDeque::new(), processing_batches: HashMap::new(), completed_requests: VecDeque::new(), violation_count: 0, total_processed: 0, } } fn add_request(&mut self, request: SchedulerRequest) { self.pending_requests.push_back(request); } fn has_pending_requests(&self) -> bool { !self.pending_requests.is_empty() } fn pending_count(&self) -> usize { self.pending_requests.len() } fn processing_count(&self) -> usize { self.processing_batches.len() } fn should_form_batch(&self, current_memory_usage: usize, memory_limit: usize) -> bool { if self.pending_requests.is_empty() { return false; } // Always allow at least one request to form a batch let min_batch_size = 1; if self.pending_requests.len() < min_batch_size { return false; } // Check memory constraints let estimated_batch_memory = self.estimate_batch_memory(); if current_memory_usage + estimated_batch_memory > memory_limit { return false; } // Check if we should wait for more requests (but be lenient for batch formation) if let Some(min_wait) = self.config.min_wait_time && let Some(oldest) = self.pending_requests.front() { let wait_time = oldest.queued_at.elapsed(); // Check for urgent deadlines that should bypass min_wait_time let has_urgent_deadline = self.pending_requests.iter().any(|req| { if let Some(deadline) = req.request.deadline { deadline.saturating_duration_since(Instant::now()) < Duration::from_millis(100) } else { false } }); // Be more aggressive about batch formation: // - Always form batch if we only have 1 request (avoid starvation) // - Form batch if we have any requests and have been waiting at least half the min_wait // - Or if we have urgent deadlines // - Or if we have reached desired batch size let half_min_wait = min_wait / 2; let desired_batch_size = std::cmp::max(1, self.config.max_batch_size / 2); let should_wait = wait_time < half_min_wait && self.pending_requests.len() < desired_batch_size && !has_urgent_deadline && self.pending_requests.len() > 1; // Don't wait for single requests if should_wait { return false; } } true } fn form_batch(&mut self) -> Option { if self.pending_requests.is_empty() { return None; } let mut batch_requests = Vec::new(); let mut batch_memory = 0; let max_batch_size = self.config.max_batch_size; let memory_limit = self.config.memory_limit.unwrap_or(usize::MAX); // Collect requests for batch - ensure we don't exceed max_batch_size while !self.pending_requests.is_empty() && batch_requests.len() < max_batch_size { let request = self .pending_requests .pop_front() .expect("checked is_empty before pop_front"); let request_memory = request.request.estimated_memory_bytes(); if batch_memory + request_memory <= memory_limit { batch_memory += request_memory; batch_requests.push(request.request); } else { // Put request back if it doesn't fit self.pending_requests.push_front(request); break; } } if batch_requests.is_empty() { return None; } // Calculate batch priority as highest priority among requests let batch_priority = batch_requests .iter() .map(|r| r.priority) .max() .unwrap_or(RequestPriority::Normal); // Estimate processing duration based on batch size and complexity let estimated_duration = Duration::from_millis(50 + (batch_requests.len() * 10) as u64); let batch = InferenceBatch { id: BatchId::new(), requests: batch_requests, lane_name: self.config.name.clone(), created_at: Instant::now(), estimated_memory_usage: batch_memory, priority: batch_priority, estimated_duration, }; // Track batch as processing self.processing_batches.insert(batch.id, batch.clone()); Some(batch) } fn form_batch_with_limit(&mut self, global_max_batch_size: usize) -> Option { if self.pending_requests.is_empty() { return None; } let mut batch_requests = Vec::new(); let mut batch_memory = 0; // Use the minimum of lane max batch size and global max batch size let max_batch_size = std::cmp::min(self.config.max_batch_size, global_max_batch_size); let memory_limit = self.config.memory_limit.unwrap_or(usize::MAX); // Collect requests for batch - ensure we don't exceed effective max_batch_size while !self.pending_requests.is_empty() && batch_requests.len() < max_batch_size { let request = self .pending_requests .pop_front() .expect("checked is_empty before pop_front"); let request_memory = request.request.estimated_memory_bytes(); if batch_memory + request_memory <= memory_limit { batch_memory += request_memory; batch_requests.push(request.request); } else { // Put request back if it doesn't fit self.pending_requests.push_front(request); break; } } if batch_requests.is_empty() { return None; } // Calculate batch priority as highest priority among requests let batch_priority = batch_requests .iter() .map(|r| r.priority) .max() .unwrap_or(RequestPriority::Normal); // Estimate processing duration based on batch size and complexity let estimated_duration = Duration::from_millis(50 + (batch_requests.len() * 10) as u64); let batch = InferenceBatch { id: BatchId::new(), requests: batch_requests, lane_name: self.config.name.clone(), created_at: Instant::now(), estimated_memory_usage: batch_memory, priority: batch_priority, estimated_duration, }; // Track batch as processing self.processing_batches.insert(batch.id, batch.clone()); Some(batch) } fn estimate_batch_memory(&self) -> usize { self.pending_requests .iter() .take(self.config.max_batch_size) .map(|req| req.request.estimated_memory_bytes()) .sum() } fn calculate_memory_utilization(&self) -> f64 { // Calculate memory used by pending requests let pending_memory = self .pending_requests .iter() .map(|req| req.request.estimated_memory_bytes()) .sum::(); // Calculate memory used by processing batches let processing_memory = self .processing_batches .values() .map(|batch| { batch .requests .iter() .map(super::request::InferenceRequest::estimated_memory_bytes) .sum::() }) .sum::(); let total_memory_used = pending_memory + processing_memory; // Use lane-specific memory limit if available, otherwise use a reasonable default let memory_limit = self.config.memory_limit.unwrap_or(8_000_000_000); // Use configured limit or 8GB default if memory_limit > 0 { (total_memory_used as f64 / memory_limit as f64).min(1.0) } else { 0.0 } } fn get_stats(&self) -> LaneStats { let now = Instant::now(); // If subtraction fails (process younger than 60s), use now as cutoff let recent_cutoff = now .checked_sub(Duration::from_secs(60)) .unwrap_or_else(Instant::now); // Calculate recent metrics let recent_completions: Vec<_> = self .completed_requests .iter() .filter(|(completion_time, _)| *completion_time > recent_cutoff) .collect(); let average_latency = if recent_completions.is_empty() { Duration::ZERO } else { let total_latency: Duration = recent_completions.iter().map(|(_, latency)| *latency).sum(); total_latency / recent_completions.len() as u32 }; let violation_rate = if self.total_processed > 0 { self.violation_count as f64 / self.total_processed as f64 } else { 0.0 }; let throughput = recent_completions.len() as f64 / 60.0; // per second LaneStats { pending_count: self.pending_requests.len(), processing_count: self.processing_batches.len(), average_latency, violation_rate, throughput, memory_utilization: self.calculate_memory_utilization(), } } } /// Continuous batching scheduler pub struct BatchScheduler { config: BatchSchedulerConfig, // Lane management lanes: Arc>>, // Global state current_memory_usage: Arc>, degradation_mode: Arc>, violation_history: Arc>>, // Statistics total_requests_processed: Arc>, last_stats_reset: Arc>, // ── Prefix / SnapKV integration ───────────────────────────────────────── /// Optional paged KV cache used for prefix-cache lookups. /// /// When `Some`, [`submit_request`] checks for a matching prefix in the /// cache before enqueueing the request, storing reusable page IDs in /// [`SchedulerRequest::prefix_hit_pages`]. Set via /// [`BatchScheduler::set_kv_cache`]. kv_cache: Option>>, /// Optional SnapKV attention-score eviction state. /// /// When `Some`, [`notify_prefill_complete`] calls /// [`AttentionScoreEviction::select_evict_positions`] and stores the result /// in the per-request [`SchedulerRequest::evicted_positions`]. Set via /// [`BatchScheduler::set_snapkv_eviction`]. snapkv_eviction: Option, } impl BatchScheduler { /// Create a new batch scheduler pub async fn new(config: BatchSchedulerConfig) -> InferenceResult { if config.sla_lanes.is_empty() { return Err(InferenceError::invalid_request( "At least one SLA lane must be configured", )); } let mut lanes = HashMap::new(); for lane_config in &config.sla_lanes { lanes.insert( lane_config.name.clone(), LaneState::new(lane_config.clone()), ); } Ok(Self { config, lanes: Arc::new(RwLock::new(lanes)), current_memory_usage: Arc::new(RwLock::new(0)), degradation_mode: Arc::new(RwLock::new(false)), violation_history: Arc::new(RwLock::new(Vec::new())), total_requests_processed: Arc::new(RwLock::new(0)), last_stats_reset: Arc::new(RwLock::new(Instant::now())), kv_cache: None, snapkv_eviction: None, }) } /// Submit a request for scheduling pub async fn submit_request(&mut self, request: InferenceRequest) -> InferenceResult<()> { // Check queue size limits first let total_pending = { let lanes = self.lanes.read().await; lanes.values().map(LaneState::pending_count).sum::() }; if total_pending >= self.config.max_queue_size { return Err(InferenceError::queue_full( self.config.max_queue_size, total_pending, )); } // Assign request to appropriate lane let lane_name = self.assign_to_lane(&request).await?; // Calculate priority score let priority_score = self.calculate_priority_score(&request); let request_id = request.id; // ── Prefix-cache lookup ────────────────────────────────────────────── // If a PagedKvCache is attached and prefix caching is enabled, check // whether the prompt tokens are already cached. On a hit we clone the // returned page IDs (copy-on-write: the shared pages are not consumed) // and store them on the request so the engine can skip recomputing the // prefix. On a miss we leave `prefix_hit_pages` as `None`; the engine // will allocate pages normally and should call // `notify_prefill_complete` afterward to register the new pages. // // InferenceRequest::input_tokens is Vec (signed token IDs from the // tokenizer), while PrefixIndex keys are &[u32]. We reinterpret via // bit-cast: token IDs are non-negative in practice so the widening is // lossless; negative values (unlikely in well-formed input) hash to a // distinct key and simply yield a cache miss. let prefix_hit_pages = if let Some(ref cache_arc) = self.kv_cache { // Try a non-blocking lock first so we never stall the hot path // under contention; fall back to a miss if the lock is busy. if let Ok(cache) = cache_arc.try_lock() { if cache.prefix_caching_enabled() { let tokens_u32: Vec = request.input_tokens.iter().map(|&t| t as u32).collect(); let pages = cache.lookup_prefix(&tokens_u32); if pages.is_some() { debug!( "Prefix-cache HIT for request {} ({} prompt tokens)", request_id, request.input_tokens.len() ); } pages } else { None } } else { // Cache is locked; treat as miss to avoid blocking the scheduler. None } } else { None }; let scheduler_request = SchedulerRequest { request, assigned_lane: lane_name.clone(), queued_at: Instant::now(), priority_score, prefix_hit_pages, evicted_positions: Vec::new(), }; // Add to lane { let mut lanes = self.lanes.write().await; if let Some(lane) = lanes.get_mut(&lane_name) { lane.add_request(scheduler_request); debug!("Request {} assigned to lane {}", request_id, lane_name); } else { return Err(InferenceError::internal_error( "Lane assignment", "Lane not found", )); } } // Check if we should enter degradation mode (account for the request we just added) let new_total_pending = total_pending + 1; if new_total_pending >= self.config.degradation_mode_threshold { let mut degradation = self.degradation_mode.write().await; *degradation = true; warn!( "Entering degradation mode: {} pending requests", new_total_pending ); } Ok(()) } /// Get next batch for processing pub async fn get_next_batch(&mut self) -> InferenceResult { let current_memory = *self.current_memory_usage.read().await; let global_max_batch_size = self.config.max_batch_size; let mut best_batch = None; let mut best_priority = RequestPriority::Low; { let mut lanes = self.lanes.write().await; for (_lane_name, lane_state) in lanes.iter_mut() { if lane_state.should_form_batch(current_memory, self.config.memory_limit) && let Some(batch) = lane_state.form_batch_with_limit(global_max_batch_size) { // Choose batch with highest priority if batch.priority >= best_priority { best_priority = batch.priority; best_batch = Some(batch); } } } } best_batch.ok_or_else(|| InferenceError::BatchFormationFailed { reason: "No eligible batches available".to_string(), }) } /// Mark batch as executing pub async fn mark_batch_executing(&mut self, batch_id: BatchId) -> InferenceResult<()> { let lanes = self.lanes.read().await; for lane in lanes.values() { if lane.processing_batches.contains_key(&batch_id) { trace!("Batch {} marked as executing", batch_id); return Ok(()); } } Err(InferenceError::internal_error( "Mark batch executing", "Batch not found", )) } /// Evaluate preemption decisions pub async fn evaluate_preemption(&self) -> InferenceResult { if !self.config.enable_preemption { return Ok(PreemptionDecision { should_preempt: false, target_batch_id: None, reason: "Preemption disabled".to_string(), priority_delta: 0, }); } let lanes = self.lanes.read().await; // Find highest priority pending request let mut highest_pending_priority = RequestPriority::Low; let mut pending_deadline = None; for lane in lanes.values() { if let Some(request) = lane.pending_requests.front() && request.request.priority > highest_pending_priority { highest_pending_priority = request.request.priority; pending_deadline = request.request.deadline; } } // Find lowest priority executing batch let mut lowest_executing_priority = RequestPriority::Critical; let mut target_batch_id = None; for lane in lanes.values() { for batch in lane.processing_batches.values() { if batch.priority < lowest_executing_priority { lowest_executing_priority = batch.priority; target_batch_id = Some(batch.id); } } } // Check if preemption is justified let priority_delta = highest_pending_priority as i32 - lowest_executing_priority as i32; let should_preempt = priority_delta >= 2; // Preempt if 2+ priority levels difference // Also check deadline urgency let urgent_deadline = pending_deadline .is_some_and(|deadline| deadline - Instant::now() < Duration::from_millis(50)); let final_decision = should_preempt || urgent_deadline; Ok(PreemptionDecision { should_preempt: final_decision, target_batch_id, reason: if urgent_deadline { "Urgent deadline approaching".to_string() } else if should_preempt { format!("Priority difference: {priority_delta}") } else { "No preemption needed".to_string() }, priority_delta, }) } /// Check for SLA violations pub async fn check_sla_violations(&self) -> InferenceResult> { let mut violations = Vec::new(); let now = Instant::now(); let lanes = self.lanes.read().await; for lane in lanes.values() { for request in &lane.pending_requests { let age = now - request.queued_at; // Check against lane max latency if age > lane.config.max_latency { violations.push(SlaViolation { request_id: request.request.id, lane_name: lane.config.name.clone(), violation_type: "Queue time exceeded".to_string(), violation_time: now, expected_latency: lane.config.max_latency, actual_latency: age, }); } // Check against request-specific deadline if let Some(deadline) = request.request.deadline && now > deadline { let expected_duration = deadline - request.queued_at; violations.push(SlaViolation { request_id: request.request.id, lane_name: lane.config.name.clone(), violation_type: "Deadline exceeded".to_string(), violation_time: now, expected_latency: expected_duration, actual_latency: age, }); } } } // Store violations in history if !violations.is_empty() { let mut history = self.violation_history.write().await; history.extend(violations.clone()); // Keep only recent violations (last hour) // If subtraction fails (process younger than 1 hour), use now as cutoff let cutoff = now .checked_sub(Duration::from_secs(3600)) .unwrap_or_else(Instant::now); history.retain(|v| v.violation_time > cutoff); } Ok(violations) } /// Set memory pressure level pub async fn set_memory_pressure(&mut self, pressure: f64) -> InferenceResult<()> { let memory_bytes = (pressure * self.config.memory_limit as f64) as usize; let mut current_memory = self.current_memory_usage.write().await; *current_memory = memory_bytes; if pressure > self.config.memory_pressure_threshold { warn!("High memory pressure: {:.1}%", pressure * 100.0); } Ok(()) } /// Check if scheduler is in degradation mode pub async fn is_degradation_mode(&self) -> bool { *self.degradation_mode.read().await } /// Get queue statistics pub async fn queue_stats(&self) -> QueueStats { let lanes_guard = self.lanes.read().await; let mut lane_stats = HashMap::new(); let mut total_pending = 0; let mut total_processing = 0; for (name, lane) in lanes_guard.iter() { let stats = lane.get_stats(); total_pending += stats.pending_count; total_processing += stats.processing_count; lane_stats.insert(name.clone(), stats); } let memory_usage = *self.current_memory_usage.read().await; let memory_pressure = memory_usage as f64 / self.config.memory_limit as f64; QueueStats { total_pending, total_processing, lanes: lane_stats, queue_utilization: total_pending as f64 / self.config.max_queue_size as f64, memory_pressure, } } /// Get scheduler configuration #[must_use] pub fn config(&self) -> &BatchSchedulerConfig { &self.config } // ── Prefix cache / SnapKV public API ────────────────────────────────────── /// Attach a [`PagedKvCache`] for prefix-cache lookups. /// /// When a cache is attached and `enable_prefix_caching` is set on it, /// [`submit_request`] will consult the cache before enqueueing each request. /// The cache must outlive the scheduler; the `Arc>` wrapper /// ensures shared ownership. pub fn set_kv_cache(&mut self, cache: Arc>) { self.kv_cache = Some(cache); } /// Attach a [`AttentionScoreEviction`] for SnapKV position masking. /// /// After prefill completes, call [`notify_prefill_complete`] to trigger /// [`AttentionScoreEviction::select_evict_positions`] and record the /// positions to skip during decode. pub fn set_snapkv_eviction(&mut self, eviction: AttentionScoreEviction) { self.snapkv_eviction = Some(eviction); } /// Notify the scheduler that prefill for a request has completed. /// /// This does two things: /// /// 1. **Prefix registration** — if a KV cache is attached and the request /// did not have a prefix-cache hit at submission, register the newly /// computed pages so future requests sharing the same prompt prefix can /// reuse them. /// /// 2. **SnapKV position selection** — if a [`AttentionScoreEviction`] is /// attached, run [`select_evict_positions`] and store the result in the /// per-request state so the decode loop knows which token positions to /// mask out. /// /// # Arguments /// /// * `request_id` — the request whose prefill has finished. /// * `prompt_tokens` — the token IDs of the prompt (used as the prefix key). /// * `allocated_pages` — the KV pages that were just populated during prefill. /// * `total_kv_positions` — total key positions in the sequence; passed to /// [`AttentionScoreEviction::select_evict_positions`]. pub async fn notify_prefill_complete( &mut self, request_id: RequestId, prompt_tokens: &[u32], allocated_pages: Vec, total_kv_positions: usize, ) -> InferenceResult<()> { // ── 1. Register prefix in the KV cache ─────────────────────────────── if let Some(ref cache_arc) = self.kv_cache { let mut cache = cache_arc.lock().await; if cache.prefix_caching_enabled() { // Only register if there was no hit at submission time // (i.e. we allocated fresh pages). The check is implicit: // if there *was* a hit the caller should not be passing // freshly allocated pages. cache.register_prefix(prompt_tokens, allocated_pages.clone()); debug!( "Prefix registered for request {} ({} tokens, {} pages)", request_id, prompt_tokens.len(), allocated_pages.len() ); } } // ── 2. SnapKV: select evict positions ──────────────────────────────── // // TODO(B2): When an `AttentionScoreEviction` is wired, call // `accumulate_scores` during each prefill attention step (the caller // owns the attention-weight tensors), then call this method once after // the final prefill step. The evicted positions are stored in the // matching `SchedulerRequest::evicted_positions` field so the decode // loop can read them via `get_evicted_positions`. // // For now we only invoke `select_evict_positions` if the eviction state // already has accumulated scores (i.e. the caller drove // `accumulate_scores` externally). This avoids returning a spurious // empty list when no scores were accumulated. let evicted = if let Some(ref eviction) = self.snapkv_eviction { let positions = eviction.select_evict_positions(total_kv_positions); if !positions.is_empty() { debug!( "SnapKV: evicting {} of {} positions for request {}", positions.len(), total_kv_positions, request_id ); } positions } else { Vec::new() }; // Store evicted positions on the matching SchedulerRequest so the // decode path can read them. { let mut lanes = self.lanes.write().await; 'outer: for lane in lanes.values_mut() { for req in &mut lane.pending_requests { if req.request.id == request_id { req.evicted_positions = evicted; break 'outer; } } } } Ok(()) } /// Return the evicted KV positions recorded for a request after prefill. /// /// Returns an empty slice if the request is not found, has not completed /// prefill yet, or if SnapKV is disabled. pub async fn get_evicted_positions(&self, request_id: RequestId) -> Vec { let lanes = self.lanes.read().await; for lane in lanes.values() { for req in &lane.pending_requests { if req.request.id == request_id { return req.evicted_positions.clone(); } } } Vec::new() } /// Return the prefix-hit pages recorded for a request, if any. /// /// Returns `None` if the request was not found or had no prefix-cache hit. pub async fn get_prefix_hit_pages(&self, request_id: RequestId) -> Option> { let lanes = self.lanes.read().await; for lane in lanes.values() { for req in &lane.pending_requests { if req.request.id == request_id { return req.prefix_hit_pages.clone(); } } } None } /// Assign request to appropriate SLA lane async fn assign_to_lane(&self, request: &InferenceRequest) -> InferenceResult { // Find matching lane based on priority for lane_config in &self.config.sla_lanes { if lane_config.priority == request.priority { return Ok(lane_config.name.clone()); } } // Fallback to lane with closest priority let mut best_lane = &self.config.sla_lanes[0]; let mut best_diff = i32::MAX; for lane_config in &self.config.sla_lanes { let diff = (lane_config.priority as i32 - request.priority as i32).abs(); if diff < best_diff { best_diff = diff; best_lane = lane_config; } } Ok(best_lane.name.clone()) } /// Calculate priority score for request ordering fn calculate_priority_score(&self, request: &InferenceRequest) -> i32 { let mut score = request.priority as i32 * 1000; // Boost score for urgent deadlines if let Some(deadline) = request.deadline { let time_remaining = deadline.saturating_duration_since(Instant::now()); if time_remaining < Duration::from_millis(100) { score += 500; // High urgency boost } else if time_remaining < Duration::from_secs(1) { score += 200; // Medium urgency boost } } // Penalize very large requests slightly if request.input_tokens.len() > 1024 { score -= 50; } score } } #[cfg(test)] mod tests { use super::*; #[tokio::test] async fn test_scheduler_creation() { let config = BatchSchedulerConfig::default(); let scheduler = BatchScheduler::new(config).await; assert!(scheduler.is_ok()); } #[tokio::test] async fn test_lane_assignment() { let config = BatchSchedulerConfig { sla_lanes: vec![ SlaLane { name: "high".to_string(), priority: RequestPriority::High, max_latency: Duration::from_millis(50), max_batch_size: 8, memory_limit: None, min_wait_time: None, }, SlaLane { name: "normal".to_string(), priority: RequestPriority::Normal, max_latency: Duration::from_millis(100), max_batch_size: 16, memory_limit: None, min_wait_time: None, }, ], ..Default::default() }; let scheduler = BatchScheduler::new(config).await.unwrap(); let high_request = InferenceRequest { priority: RequestPriority::High, input_tokens: vec![1, 2, 3], ..Default::default() }; let lane = scheduler.assign_to_lane(&high_request).await.unwrap(); assert_eq!(lane, "high"); } #[test] fn test_priority_score_calculation() { let scheduler_config = BatchSchedulerConfig::default(); let scheduler = BatchScheduler { config: scheduler_config, lanes: Arc::new(RwLock::new(HashMap::new())), current_memory_usage: Arc::new(RwLock::new(0)), degradation_mode: Arc::new(RwLock::new(false)), violation_history: Arc::new(RwLock::new(Vec::new())), total_requests_processed: Arc::new(RwLock::new(0)), last_stats_reset: Arc::new(RwLock::new(Instant::now())), kv_cache: None, snapkv_eviction: None, }; let high_priority_request = InferenceRequest { priority: RequestPriority::High, ..Default::default() }; let normal_priority_request = InferenceRequest { priority: RequestPriority::Normal, ..Default::default() }; let high_score = scheduler.calculate_priority_score(&high_priority_request); let normal_score = scheduler.calculate_priority_score(&normal_priority_request); assert!(high_score > normal_score); } // ── Prefix cache / SnapKV tests ─────────────────────────────────────────── #[tokio::test] async fn test_set_kv_cache_does_not_panic() { use crate::cache::{KvCacheConfig, PagedKvCache}; use rtx_tensor::Device; let config = BatchSchedulerConfig::default(); let mut scheduler = BatchScheduler::new(config).await.unwrap(); let kv_config = KvCacheConfig { enable_prefix_caching: true, ..KvCacheConfig::default() }; let cache = PagedKvCache::new(kv_config, Device::cpu()).unwrap(); scheduler.set_kv_cache(Arc::new(Mutex::new(cache))); // If we reach here the field was set without panicking. } #[tokio::test] async fn test_set_snapkv_eviction_does_not_panic() { use crate::cache::AttentionScoreEviction; let config = BatchSchedulerConfig::default(); let mut scheduler = BatchScheduler::new(config).await.unwrap(); scheduler.set_snapkv_eviction(AttentionScoreEviction::new(0.6, 32)); // If we reach here the field was set without panicking. } #[tokio::test] async fn test_submit_request_with_prefix_cache_miss() { use crate::cache::{KvCacheConfig, PagedKvCache}; use rtx_tensor::Device; let config = BatchSchedulerConfig::default(); let mut scheduler = BatchScheduler::new(config).await.unwrap(); let kv_config = KvCacheConfig { enable_prefix_caching: true, ..KvCacheConfig::default() }; let cache = PagedKvCache::new(kv_config, Device::cpu()).unwrap(); scheduler.set_kv_cache(Arc::new(Mutex::new(cache))); let request = InferenceRequest { input_tokens: vec![1, 2, 3, 4], ..Default::default() }; let request_id = request.id; scheduler.submit_request(request).await.unwrap(); // No prefix was registered, so there should be no hit pages. let hit_pages = scheduler.get_prefix_hit_pages(request_id).await; assert!( hit_pages.is_none(), "Expected no prefix-cache hit on empty cache" ); } #[tokio::test] async fn test_notify_prefill_complete_registers_prefix() { use crate::cache::{KvCacheConfig, PagedKvCache}; use rtx_tensor::Device; use uuid::Uuid; let config = BatchSchedulerConfig::default(); let mut scheduler = BatchScheduler::new(config).await.unwrap(); let kv_config = KvCacheConfig { enable_prefix_caching: true, ..KvCacheConfig::default() }; let cache = PagedKvCache::new(kv_config, Device::cpu()).unwrap(); let cache_arc = Arc::new(Mutex::new(cache)); scheduler.set_kv_cache(Arc::clone(&cache_arc)); // input_tokens is Vec; notify_prefill_complete takes &[u32]. // Use i32 tokens and convert to u32 when calling notify_prefill_complete. let prompt_tokens_i32: Vec = vec![10, 20, 30]; let prompt_tokens_u32: Vec = prompt_tokens_i32.iter().map(|&t| t as u32).collect(); let page_id: PageId = Uuid::new_v4(); let request = InferenceRequest { input_tokens: prompt_tokens_i32, ..Default::default() }; let request_id = request.id; scheduler.submit_request(request).await.unwrap(); // Simulate prefill completion — register pages. scheduler .notify_prefill_complete(request_id, &prompt_tokens_u32, vec![page_id], 64) .await .unwrap(); // Now the prefix should be in the cache. let hit = cache_arc.lock().await.lookup_prefix(&prompt_tokens_u32); assert!( hit.is_some(), "Prefix should be registered after notify_prefill_complete" ); assert_eq!(hit.unwrap(), vec![page_id]); } #[tokio::test] async fn test_notify_prefill_complete_snapkv_eviction() { use crate::cache::AttentionScoreEviction; let config = BatchSchedulerConfig::default(); let mut scheduler = BatchScheduler::new(config).await.unwrap(); // Attach an eviction state that already has accumulated scores. let mut eviction = AttentionScoreEviction::new(0.5, 0); // 10 positions, all equal — 5 should be evicted. eviction.accumulate_scores(&[0.1f32; 10]); scheduler.set_snapkv_eviction(eviction); let request = InferenceRequest { input_tokens: vec![1, 2, 3], ..Default::default() }; let request_id = request.id; scheduler.submit_request(request).await.unwrap(); scheduler .notify_prefill_complete(request_id, &[1u32, 2, 3], vec![], 10) .await .unwrap(); let evicted = scheduler.get_evicted_positions(request_id).await; assert_eq!(evicted.len(), 5, "SnapKV should evict 50% of 10 positions"); } }