1use async_trait::async_trait;
7use ferrum_types::{
8 BatchId, InferenceRequest, InferenceResponse, Priority, RequestId, RequestState, Result,
9 SchedulerConfig as TypesSchedulerConfig, SchedulerStats,
10};
11use serde::{Deserialize, Serialize};
12use std::{collections::HashMap, time::Duration};
13
14mod prefix_restore;
15pub use prefix_restore::PreparedPrefixRestore;
16
17#[async_trait]
19pub trait Scheduler: Send + Sync {
20 async fn submit(&self, request: InferenceRequest) -> Result<RequestId>;
22
23 async fn next_batch(&self, hint: BatchHint) -> Option<BatchPlan>;
25
26 async fn complete(&self, request_id: RequestId, response: &InferenceResponse) -> Result<()>;
28
29 async fn cancel(&self, request_id: RequestId) -> Result<bool>;
31
32 async fn update_priority(&self, request_id: RequestId, priority: Priority) -> Result<()>;
34
35 fn metrics(&self) -> SchedulerMetrics;
37
38 fn config(&self) -> &TypesSchedulerConfig;
40
41 fn request_state(&self, request_id: &RequestId) -> Option<RequestState> {
43 let _ = request_id;
44 None
45 }
46
47 fn prepare_prefix_restore(
52 &self,
53 _request_id: &RequestId,
54 _expected_offset: usize,
55 _prompt_tokens: usize,
56 ) -> Result<Option<PreparedPrefixRestore>> {
57 Ok(None)
58 }
59
60 fn commit_prefix_restored(
65 &self,
66 _prepared: PreparedPrefixRestore,
67 _restored_boundary: usize,
68 ) -> Result<()> {
69 Err(ferrum_types::FerrumError::unsupported(
70 "Prefix state restoration is not supported by this scheduler",
71 ))
72 }
73
74 async fn preempt(&self, _request_id: RequestId) -> Result<PreemptionResult> {
76 Err(ferrum_types::FerrumError::unsupported(
78 "Preemption not supported",
79 ))
80 }
81
82 async fn resume(&self, _request_id: RequestId) -> Result<()> {
84 Err(ferrum_types::FerrumError::unsupported(
86 "Resumption not supported",
87 ))
88 }
89}
90
91#[derive(Debug, Clone)]
93pub struct BatchHint {
94 pub max_batch_size: usize,
96 pub max_tokens: usize,
98 pub target_latency_ms: Option<u64>,
100 pub available_memory: Option<u64>,
102 pub resource_constraints: ResourceConstraints,
104}
105
106impl BatchHint {
107 pub fn simple(max_batch_size: usize) -> Self {
109 Self {
110 max_batch_size,
111 max_tokens: max_batch_size * 2048, target_latency_ms: None,
113 available_memory: None,
114 resource_constraints: ResourceConstraints::default(),
115 }
116 }
117}
118
119#[derive(Debug, Clone, Serialize, Deserialize, Default)]
121pub struct ResourceConstraints {
122 pub max_gpu_memory: Option<u64>,
124 pub max_cpu_memory: Option<u64>,
126 pub max_recurrent_state_bytes: Option<u64>,
128 pub max_recurrent_state_slots: Option<usize>,
130 pub max_compute_units: Option<usize>,
132 pub required_devices: Vec<ferrum_types::Device>,
134}
135
136#[derive(Debug, Clone)]
138pub struct BatchPlan {
139 pub batch_id: BatchId,
141 pub requests: Vec<ScheduledRequest>,
143 pub max_sequence_length: usize,
145 pub estimated_time_ms: Option<u64>,
147 pub resource_requirements: BatchResourceRequirements,
149 pub created_at: chrono::DateTime<chrono::Utc>,
151}
152
153impl BatchPlan {
154 pub fn total_tokens(&self) -> usize {
156 self.requests
157 .iter()
158 .map(|req| {
159 req.tokens_to_process
160 .unwrap_or(req.request.sampling_params.max_tokens)
161 })
162 .sum()
163 }
164
165 pub fn size(&self) -> usize {
167 self.requests.len()
168 }
169
170 pub fn is_empty(&self) -> bool {
172 self.requests.is_empty()
173 }
174
175 pub fn max_priority(&self) -> Priority {
177 self.requests
178 .iter()
179 .map(|req| req.request.priority)
180 .max()
181 .unwrap_or(Priority::Low)
182 }
183}
184
185#[derive(Debug, Clone)]
187pub struct ScheduledRequest {
188 pub request: InferenceRequest,
190 pub state: RequestState,
192 pub queue_position: Option<usize>,
194 pub estimated_wait_time: Option<Duration>,
196 pub tokens_processed: usize,
198 pub tokens_to_process: Option<usize>,
203 pub allocated_resources: AllocatedResources,
205 pub submitted_at: chrono::DateTime<chrono::Utc>,
207 pub started_at: Option<chrono::DateTime<chrono::Utc>>,
209}
210
211impl ScheduledRequest {
212 pub fn new(request: InferenceRequest) -> Self {
214 Self {
215 request,
216 state: RequestState::Waiting,
217 queue_position: None,
218 estimated_wait_time: None,
219 tokens_processed: 0,
220 tokens_to_process: None,
221 allocated_resources: AllocatedResources::default(),
222 submitted_at: chrono::Utc::now(),
223 started_at: None,
224 }
225 }
226
227 pub fn age(&self) -> Duration {
229 (chrono::Utc::now() - self.submitted_at)
230 .to_std()
231 .unwrap_or_default()
232 }
233
234 pub fn processing_time(&self) -> Option<Duration> {
236 self.started_at
237 .map(|start| (chrono::Utc::now() - start).to_std().unwrap_or_default())
238 }
239}
240
241#[derive(Debug, Clone, Default)]
243pub struct AllocatedResources {
244 pub kv_cache_blocks: Vec<ferrum_types::BlockId>,
246 pub gpu_memory: u64,
248 pub cpu_memory: u64,
250 pub recurrent_state_bytes: u64,
252 pub recurrent_state_slots: usize,
254 pub compute_units: usize,
256}
257
258#[derive(Debug, Clone, Default)]
260pub struct BatchResourceRequirements {
261 pub gpu_memory: u64,
263 pub cpu_memory: u64,
265 pub kv_cache_blocks: usize,
267 pub recurrent_state_bytes: u64,
269 pub recurrent_state_slots: usize,
271 pub compute_units: usize,
273}
274
275#[derive(Debug, Clone)]
277pub struct PreemptionResult {
278 pub success: bool,
280 pub saved_state: Option<PreemptionState>,
282 pub freed_resources: AllocatedResources,
284}
285
286#[derive(Debug, Clone)]
288pub struct PreemptionState {
289 pub kv_cache_checkpoint: Vec<u8>,
291 pub tokens_processed: usize,
293 pub generation_state: HashMap<String, serde_json::Value>,
295}
296
297pub type SchedulerConfig = TypesSchedulerConfig;
299
300#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
302pub enum SchedulingPolicy {
303 FCFS,
305 Priority,
307 FairShare,
309 SJF,
311 ResourceAware,
313 SlaAware,
315}
316
317#[derive(Debug, Clone, Serialize, Deserialize)]
319pub struct FairShareConfig {
320 pub client_shares: HashMap<String, f32>,
322 pub default_share: f32,
324 pub enforcement_strictness: f32,
326}
327
328#[derive(Debug, Clone, Serialize, Deserialize)]
330pub struct SlaConfig {
331 pub enabled: bool,
333 pub default_sla: SlaRequirements,
335 pub client_slas: HashMap<String, SlaRequirements>,
337}
338
339#[derive(Debug, Clone, Serialize, Deserialize)]
341pub struct SlaRequirements {
342 pub max_latency_p95_ms: u64,
344 pub max_latency_p99_ms: u64,
346 pub min_throughput_rps: f32,
348 pub availability_percent: f32,
350}
351
352#[derive(Debug, Clone, Serialize, Deserialize, Default)]
354pub struct ResourceLimits {
355 pub max_gpu_memory: Option<u64>,
357 pub max_cpu_memory: Option<u64>,
359 pub max_kv_cache_blocks: Option<usize>,
361 pub max_recurrent_state_bytes: Option<u64>,
363 pub max_recurrent_state_slots: Option<usize>,
365 pub per_client_limits: HashMap<String, ClientResourceLimits>,
367}
368
369#[derive(Debug, Clone, Serialize, Deserialize)]
371pub struct ClientResourceLimits {
372 pub max_concurrent_requests: usize,
374 pub max_gpu_memory: Option<u64>,
376 pub max_recurrent_state_bytes: Option<u64>,
378 pub max_requests_per_minute: Option<u32>,
380}
381
382pub type SchedulerMetrics = SchedulerStats;
383
384#[async_trait]
386pub trait AdvancedScheduler: Scheduler {
387 async fn enable_resource_awareness(&mut self, config: ResourceAwarenessConfig) -> Result<()>;
389
390 async fn set_admission_policy(&mut self, policy: Box<dyn AdmissionPolicy>) -> Result<()>;
392
393 async fn configure_dynamic_batching(&mut self, config: DynamicBatchingConfig) -> Result<()>;
395
396 fn queue_analysis(&self) -> QueueAnalysis;
398
399 async fn simulate_load(
401 &self,
402 workload: &SimulatedWorkload,
403 ) -> Result<SchedulingSimulationResult>;
404}
405
406#[derive(Debug, Clone, Serialize, Deserialize)]
408pub struct ResourceAwarenessConfig {
409 pub enable_memory_awareness: bool,
411 pub enable_compute_awareness: bool,
413 pub prediction_horizon_ms: u64,
415 pub safety_margin: f32,
417}
418
419pub trait AdmissionPolicy: Send + Sync {
421 fn should_admit(
423 &self,
424 request: &InferenceRequest,
425 current_metrics: &SchedulerMetrics,
426 ) -> AdmissionDecision;
427
428 fn name(&self) -> &str;
430}
431
432#[derive(Debug, Clone)]
434pub enum AdmissionDecision {
435 Accept,
437 Reject(String),
439 AcceptWithDelay(Duration),
441}
442
443#[derive(Debug, Clone, Serialize, Deserialize)]
445pub struct DynamicBatchingConfig {
446 pub min_batch_size: usize,
448 pub max_batch_size: usize,
450 pub batch_timeout_ms: u64,
452 pub enable_adaptive_sizing: bool,
454 pub target_utilization: f32,
456}
457
458#[derive(Debug, Clone)]
460pub struct QueueAnalysis {
461 pub queue_depth_history: Vec<(chrono::DateTime<chrono::Utc>, usize)>,
463 pub wait_time_distribution: WaitTimeDistribution,
465 pub request_patterns: RequestPatternAnalysis,
467 pub bottlenecks: Vec<BottleneckAnalysis>,
469}
470
471#[derive(Debug, Clone)]
473pub struct WaitTimeDistribution {
474 pub p50_ms: f64,
476 pub p95_ms: f64,
478 pub p99_ms: f64,
480 pub max_ms: f64,
482 pub mean_ms: f64,
484}
485
486#[derive(Debug, Clone)]
488pub struct RequestPatternAnalysis {
489 pub peak_times: Vec<chrono::DateTime<chrono::Utc>>,
491 pub rate_trend: RateTrend,
493 pub seasonality: SeasonalityPattern,
495}
496
497#[derive(Debug, Clone, Copy)]
499pub enum RateTrend {
500 Increasing,
501 Decreasing,
502 Stable,
503 Volatile,
504}
505
506#[derive(Debug, Clone)]
508pub struct SeasonalityPattern {
509 pub hourly_pattern: Vec<f32>,
511 pub daily_pattern: Vec<f32>,
513 pub weekly_pattern: Vec<f32>,
515}
516
517#[derive(Debug, Clone)]
519pub struct BottleneckAnalysis {
520 pub bottleneck_type: BottleneckType,
522 pub severity: f32,
524 pub description: String,
526 pub mitigation: String,
528}
529
530#[derive(Debug, Clone, Copy)]
532pub enum BottleneckType {
533 Memory,
535 Compute,
537 IO,
539 Scheduling,
541 Network,
543}
544
545#[derive(Debug, Clone)]
547pub struct SimulatedWorkload {
548 pub arrival_pattern: ArrivalPattern,
550 pub size_distribution: SizeDistribution,
552 pub duration_seconds: u64,
554}
555
556#[derive(Debug, Clone)]
558pub enum ArrivalPattern {
559 Constant { rate_rps: f32 },
561 Poisson { lambda: f32 },
563 Bursty {
565 burst_rate: f32,
566 quiet_rate: f32,
567 burst_duration_s: f32,
568 },
569 Seasonal {
571 base_rate: f32,
572 peaks: Vec<(f32, f32)>,
573 }, }
575
576#[derive(Debug, Clone)]
578pub enum SizeDistribution {
579 Fixed { tokens: usize },
581 Uniform {
583 min_tokens: usize,
584 max_tokens: usize,
585 },
586 Normal { mean: f32, std_dev: f32 },
588 LogNormal { mu: f32, sigma: f32 },
590}
591
592#[derive(Debug, Clone)]
594pub struct SchedulingSimulationResult {
595 pub total_requests: u64,
597 pub successful_requests: u64,
599 pub failed_requests: u64,
601 pub avg_latency_ms: f64,
603 pub p95_latency_ms: f64,
605 pub p99_latency_ms: f64,
607 pub throughput_rps: f32,
609 pub resource_utilization: Option<ResourceStats>,
611 pub bottlenecks: Vec<BottleneckAnalysis>,
613}
614
615#[derive(Debug, Clone, Serialize, Deserialize, Default)]
616pub struct ResourceStats {
617 pub gpu_memory_bytes: Option<u64>,
618 pub cpu_memory_bytes: Option<u64>,
619 pub compute_utilization: Option<f32>,
620}