Skip to main content

ferrum_scheduler/implementations/
fifo.rs

1//! FIFO (First-In-First-Out) scheduler implementation
2
3use 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
23/// FIFO scheduler that processes requests in first-come-first-served order
24pub struct FifoScheduler {
25    /// Configuration
26    config: SchedulerConfig,
27    /// Waiting queue (FIFO order)
28    waiting_queue: RwLock<VecDeque<ScheduledRequest>>,
29    /// Running requests
30    running_requests: RwLock<HashMap<RequestId, ScheduledRequest>>,
31    /// Completed request counter
32    completed_counter: AtomicU64,
33    /// Failed request counter  
34    failed_counter: AtomicU64,
35    /// Cancelled request counter
36    cancelled_counter: AtomicU64,
37    /// Scheduler start time
38    start_time: Instant,
39    /// Metrics tracking
40    metrics_tracker: Arc<MetricsTracker>,
41}
42
43/// Internal metrics tracker
44struct 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    /// Create new FIFO scheduler
90    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    /// Create batch from waiting queue
106    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); // reasonable default
117
118        // Take requests from front of queue (FIFO)
119        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                // Check if adding this request would exceed limits
127                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                    // Move to running requests
135                    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                    // Put the request back
140                    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), // Simplified estimate
162            resource_requirements: BatchResourceRequirements {
163                gpu_memory: (total_tokens * 16) as u64, // Rough estimate: 16 bytes per token
164                cpu_memory: (total_tokens * 4) as u64,  // Rough estimate: 4 bytes per token
165                kv_cache_blocks: total_tokens / 16,     // Assume 16 tokens per block
166                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        // Check queue capacity
182        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        // Create scheduled request
192        let scheduled_request = ScheduledRequest::new(request);
193
194        // Add to waiting queue
195        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            // Calculate metrics
220            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        // Try to remove from waiting queue first
258        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        // Try to remove from running requests
271        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        // FIFO scheduler ignores priority updates by design
287        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, // FIFO doesn't support preemption
317            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// ============================================================================
362// Unit Tests
363// ============================================================================
364
365#[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        // Submit requests
400        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        // Should have 2 waiting
410        assert_eq!(scheduler.waiting_queue.read().len(), 2);
411
412        // Get batch
413        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}