ferrum_scheduler/implementations/
fifo.rs1use crate::{
4 BatchHint, BatchPlan, BatchResourceRequirements, PreemptionResult, ScheduledRequest, Scheduler,
5};
6use async_trait::async_trait;
7use ferrum_interfaces::scheduler::SchedulerMetrics;
8use ferrum_types::SchedulerConfig;
9use ferrum_types::{
10 BatchId, InferenceRequest, InferenceResponse, Priority, RequestId, RequestState, Result,
11};
12use parking_lot::RwLock;
13use std::{
14 collections::{HashMap, VecDeque},
15 sync::{
16 atomic::{AtomicU64, Ordering},
17 Arc,
18 },
19 time::Instant,
20};
21use tracing::{debug, info, warn};
22
23pub struct FifoScheduler {
25 config: SchedulerConfig,
27 waiting_queue: RwLock<VecDeque<ScheduledRequest>>,
29 running_requests: RwLock<HashMap<RequestId, ScheduledRequest>>,
31 completed_counter: AtomicU64,
33 failed_counter: AtomicU64,
35 cancelled_counter: AtomicU64,
37 start_time: Instant,
39 metrics_tracker: Arc<MetricsTracker>,
41}
42
43struct MetricsTracker {
45 total_wait_time_ms: AtomicU64,
46 total_execution_time_ms: AtomicU64,
47 request_count: AtomicU64,
48}
49
50impl MetricsTracker {
51 fn new() -> Self {
52 Self {
53 total_wait_time_ms: AtomicU64::new(0),
54 total_execution_time_ms: AtomicU64::new(0),
55 request_count: AtomicU64::new(0),
56 }
57 }
58
59 fn record_completion(&self, wait_time_ms: u64, execution_time_ms: u64) {
60 self.total_wait_time_ms
61 .fetch_add(wait_time_ms, Ordering::Relaxed);
62 self.total_execution_time_ms
63 .fetch_add(execution_time_ms, Ordering::Relaxed);
64 self.request_count.fetch_add(1, Ordering::Relaxed);
65 }
66
67 fn avg_wait_time_ms(&self) -> f64 {
68 let total_wait = self.total_wait_time_ms.load(Ordering::Relaxed) as f64;
69 let count = self.request_count.load(Ordering::Relaxed) as f64;
70 if count > 0.0 {
71 total_wait / count
72 } else {
73 0.0
74 }
75 }
76
77 fn avg_execution_time_ms(&self) -> f64 {
78 let total_exec = self.total_execution_time_ms.load(Ordering::Relaxed) as f64;
79 let count = self.request_count.load(Ordering::Relaxed) as f64;
80 if count > 0.0 {
81 total_exec / count
82 } else {
83 0.0
84 }
85 }
86}
87
88impl FifoScheduler {
89 pub fn new(config: SchedulerConfig) -> Self {
91 info!("Creating FIFO scheduler with config: {:?}", config);
92
93 Self {
94 config,
95 waiting_queue: RwLock::new(VecDeque::new()),
96 running_requests: RwLock::new(HashMap::new()),
97 completed_counter: AtomicU64::new(0),
98 failed_counter: AtomicU64::new(0),
99 cancelled_counter: AtomicU64::new(0),
100 start_time: Instant::now(),
101 metrics_tracker: Arc::new(MetricsTracker::new()),
102 }
103 }
104
105 fn create_batch(&self, hint: BatchHint) -> Option<BatchPlan> {
107 let mut waiting_queue = self.waiting_queue.write();
108 let mut running_requests = self.running_requests.write();
109
110 if waiting_queue.is_empty() {
111 return None;
112 }
113
114 let mut batch_requests = Vec::new();
115 let mut total_tokens = 0;
116 let max_sequence_length = hint.max_tokens.min(2048); while batch_requests.len() < hint.max_batch_size
120 && total_tokens < hint.max_tokens
121 && !waiting_queue.is_empty()
122 {
123 if let Some(mut scheduled_req) = waiting_queue.pop_front() {
124 let request_tokens = scheduled_req.request.sampling_params.max_tokens;
125
126 if total_tokens + request_tokens <= hint.max_tokens {
128 scheduled_req.state = RequestState::Running;
129 scheduled_req.started_at = Some(chrono::Utc::now());
130 scheduled_req.queue_position = None;
131
132 total_tokens += request_tokens;
133
134 let request_id = scheduled_req.request.id.clone();
136 running_requests.insert(request_id, scheduled_req.clone());
137 batch_requests.push(scheduled_req);
138 } else {
139 waiting_queue.push_front(scheduled_req);
141 break;
142 }
143 }
144 }
145
146 if batch_requests.is_empty() {
147 return None;
148 }
149
150 let batch_id = BatchId::new();
151 debug!(
152 "Creating batch {} with {} requests",
153 batch_id,
154 batch_requests.len()
155 );
156
157 Some(BatchPlan {
158 batch_id,
159 requests: batch_requests,
160 max_sequence_length,
161 estimated_time_ms: Some(1000), resource_requirements: BatchResourceRequirements {
163 gpu_memory: (total_tokens * 16) as u64, cpu_memory: (total_tokens * 4) as u64, kv_cache_blocks: total_tokens / 16, recurrent_state_bytes: 0,
167 recurrent_state_slots: 0,
168 compute_units: 1,
169 },
170 created_at: chrono::Utc::now(),
171 })
172 }
173}
174
175#[async_trait]
176impl Scheduler for FifoScheduler {
177 async fn submit(&self, request: InferenceRequest) -> Result<RequestId> {
178 let request_id = request.id.clone();
179 debug!("Submitting request {} to FIFO scheduler", request_id);
180
181 let waiting_queue = self.waiting_queue.read();
183 if waiting_queue.len() >= self.config.max_waiting_requests {
184 warn!("Queue is full, rejecting request {}", request_id);
185 return Err(ferrum_types::FerrumError::scheduler(
186 "Queue is full, cannot accept more requests",
187 ));
188 }
189 drop(waiting_queue);
190
191 let scheduled_request = ScheduledRequest::new(request);
193
194 let mut waiting_queue = self.waiting_queue.write();
196 let queue_position = waiting_queue.len();
197
198 let mut scheduled_req = scheduled_request;
199 scheduled_req.queue_position = Some(queue_position);
200
201 waiting_queue.push_back(scheduled_req);
202
203 info!(
204 "Request {} queued at position {}",
205 request_id, queue_position
206 );
207 Ok(request_id)
208 }
209
210 async fn next_batch(&self, hint: BatchHint) -> Option<BatchPlan> {
211 self.create_batch(hint)
212 }
213
214 async fn complete(&self, request_id: RequestId, response: &InferenceResponse) -> Result<()> {
215 debug!("Completing request {}", request_id);
216
217 let mut running_requests = self.running_requests.write();
218 if let Some(scheduled_req) = running_requests.remove(&request_id) {
219 let wait_time = scheduled_req.age();
221 let execution_time = scheduled_req.processing_time().unwrap_or_default();
222
223 self.metrics_tracker.record_completion(
224 wait_time.as_millis() as u64,
225 execution_time.as_millis() as u64,
226 );
227
228 match response.finish_reason {
229 ferrum_types::FinishReason::EOS
230 | ferrum_types::FinishReason::Stop
231 | ferrum_types::FinishReason::Length => {
232 self.completed_counter.fetch_add(1, Ordering::Relaxed);
233 debug!("Request {} completed successfully", request_id);
234 }
235 _ => {
236 self.failed_counter.fetch_add(1, Ordering::Relaxed);
237 warn!(
238 "Request {} completed with error: {:?}",
239 request_id, response.finish_reason
240 );
241 }
242 }
243
244 Ok(())
245 } else {
246 warn!("Attempted to complete unknown request: {}", request_id);
247 Err(ferrum_types::FerrumError::scheduler(format!(
248 "Request {} not found in running requests",
249 request_id
250 )))
251 }
252 }
253
254 async fn cancel(&self, request_id: RequestId) -> Result<bool> {
255 debug!("Cancelling request {}", request_id);
256
257 let mut waiting_queue = self.waiting_queue.write();
259 if let Some(pos) = waiting_queue
260 .iter()
261 .position(|req| req.request.id == request_id)
262 {
263 waiting_queue.remove(pos);
264 self.cancelled_counter.fetch_add(1, Ordering::Relaxed);
265 info!("Request {} cancelled from waiting queue", request_id);
266 return Ok(true);
267 }
268 drop(waiting_queue);
269
270 let mut running_requests = self.running_requests.write();
272 if running_requests.remove(&request_id).is_some() {
273 self.cancelled_counter.fetch_add(1, Ordering::Relaxed);
274 warn!(
275 "Request {} cancelled while running (may cause issues)",
276 request_id
277 );
278 return Ok(true);
279 }
280
281 warn!("Request {} not found for cancellation", request_id);
282 Ok(false)
283 }
284
285 async fn update_priority(&self, request_id: RequestId, _priority: Priority) -> Result<()> {
286 debug!(
288 "Priority update ignored for request {} in FIFO scheduler",
289 request_id
290 );
291 Ok(())
292 }
293
294 fn metrics(&self) -> SchedulerMetrics {
295 let waiting_queue = self.waiting_queue.read();
296 let running_requests = self.running_requests.read();
297
298 let waiting_count = waiting_queue.len();
299 let running_count = running_requests.len();
300 let completed_count = self.completed_counter.load(Ordering::Relaxed);
301 let failed_count = self.failed_counter.load(Ordering::Relaxed);
302 let cancelled_count = self.cancelled_counter.load(Ordering::Relaxed);
303
304 let uptime_secs = self.start_time.elapsed().as_secs_f64();
305 let throughput = if uptime_secs > 0.0 {
306 completed_count as f64 / uptime_secs
307 } else {
308 0.0
309 };
310
311 let queue_utilization = waiting_count as f32 / self.config.max_waiting_requests as f32;
312
313 ferrum_types::SchedulerStats {
314 waiting_requests: waiting_count,
315 running_requests: running_count,
316 preempted_requests: 0, completed_requests: completed_count,
318 failed_requests: failed_count,
319 cancelled_requests: cancelled_count,
320 avg_wait_time_ms: self.metrics_tracker.avg_wait_time_ms(),
321 avg_execution_time_ms: self.metrics_tracker.avg_execution_time_ms(),
322 throughput_rps: throughput,
323 queue_utilization,
324 }
325 }
326
327 fn config(&self) -> &SchedulerConfig {
328 &self.config
329 }
330
331 fn request_state(&self, request_id: &RequestId) -> Option<RequestState> {
332 if self.running_requests.read().contains_key(request_id) {
333 return Some(RequestState::Running);
334 }
335
336 if self
337 .waiting_queue
338 .read()
339 .iter()
340 .any(|req| req.request.id == *request_id)
341 {
342 return Some(RequestState::Waiting);
343 }
344
345 None
346 }
347
348 async fn preempt(&self, _request_id: RequestId) -> Result<PreemptionResult> {
349 Err(ferrum_types::FerrumError::unsupported(
350 "FIFO scheduler does not support preemption",
351 ))
352 }
353
354 async fn resume(&self, _request_id: RequestId) -> Result<()> {
355 Err(ferrum_types::FerrumError::unsupported(
356 "FIFO scheduler does not support resumption",
357 ))
358 }
359}
360
361#[cfg(test)]
366mod tests {
367 use super::*;
368 use ferrum_types::{ModelId, SamplingParams};
369
370 fn create_test_request(priority: Priority) -> InferenceRequest {
371 InferenceRequest {
372 id: RequestId::new(),
373 prompt: "test".to_string(),
374 model_id: ModelId::new("test-model"),
375 sampling_params: SamplingParams::default(),
376 stream: false,
377 priority,
378 client_id: None,
379 session_id: None,
380 created_at: chrono::Utc::now(),
381 api_request: None,
382 evidence_request: Default::default(),
383 metadata: std::collections::HashMap::new(),
384 }
385 }
386
387 #[tokio::test]
388 async fn test_fifo_scheduler_creation() {
389 let config = SchedulerConfig::default();
390 let scheduler = FifoScheduler::new(config);
391 assert_eq!(scheduler.waiting_queue.read().len(), 0);
392 }
393
394 #[tokio::test]
395 async fn test_fifo_submit_and_batch() {
396 let config = SchedulerConfig::default();
397 let scheduler = FifoScheduler::new(config);
398
399 let _id1 = scheduler
401 .submit(create_test_request(Priority::Normal))
402 .await
403 .unwrap();
404 let _id2 = scheduler
405 .submit(create_test_request(Priority::High))
406 .await
407 .unwrap();
408
409 assert_eq!(scheduler.waiting_queue.read().len(), 2);
411
412 let batch = scheduler.next_batch(BatchHint::simple(5)).await;
414 assert!(batch.is_some());
415 }
416
417 #[tokio::test]
418 async fn test_fifo_cancel() {
419 let config = SchedulerConfig::default();
420 let scheduler = FifoScheduler::new(config);
421
422 let request = create_test_request(Priority::Normal);
423 let id = request.id.clone();
424 scheduler.submit(request).await.unwrap();
425
426 let result = scheduler.cancel(id).await;
427 assert!(result.is_ok());
428 }
429
430 #[tokio::test]
431 async fn test_fifo_request_state_transitions() {
432 let config = SchedulerConfig::default();
433 let scheduler = FifoScheduler::new(config);
434
435 let request = create_test_request(Priority::Normal);
436 let id = request.id.clone();
437 scheduler.submit(request).await.unwrap();
438
439 assert_eq!(scheduler.request_state(&id), Some(RequestState::Waiting));
440
441 let _batch = scheduler.next_batch(BatchHint::simple(1)).await;
442 assert_eq!(scheduler.request_state(&id), Some(RequestState::Running));
443
444 scheduler.cancel(id.clone()).await.unwrap();
445 assert_eq!(scheduler.request_state(&id), None);
446 }
447
448 #[test]
449 fn test_metrics_tracker() {
450 let tracker = MetricsTracker::new();
451 tracker.record_completion(100, 500);
452
453 assert!(tracker.avg_wait_time_ms() > 0.0);
454 assert!(tracker.avg_execution_time_ms() > 0.0);
455 }
456}