1use 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 priority_queue::PriorityQueue;
14use std::{
15 collections::HashMap,
16 sync::{
17 atomic::{AtomicU64, Ordering},
18 Arc,
19 },
20 time::Instant,
21};
22use tracing::{debug, info, warn};
23
24pub struct PriorityScheduler {
26 config: SchedulerConfig,
28 waiting_queue: RwLock<PriorityQueue<RequestId, RequestPriority>>,
30 request_map: RwLock<HashMap<RequestId, ScheduledRequest>>,
32 running_requests: RwLock<HashMap<RequestId, ScheduledRequest>>,
34 completed_counter: AtomicU64,
36 failed_counter: AtomicU64,
38 cancelled_counter: AtomicU64,
40 start_time: Instant,
42 metrics_tracker: Arc<MetricsTracker>,
44}
45
46#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
48struct RequestPriority {
49 priority: i32,
51 submission_time_nanos: i64,
53}
54
55impl RequestPriority {
56 fn new(priority: Priority, submitted_at: chrono::DateTime<chrono::Utc>) -> Self {
57 let priority_value = match priority {
58 Priority::Critical => 100,
59 Priority::High => 75,
60 Priority::Normal => 50,
61 Priority::Low => 25,
62 };
63
64 let submission_time_nanos = -submitted_at.timestamp_nanos_opt().unwrap_or(0);
66
67 Self {
68 priority: priority_value,
69 submission_time_nanos,
70 }
71 }
72}
73
74struct MetricsTracker {
76 total_wait_time_ms: AtomicU64,
77 total_execution_time_ms: AtomicU64,
78 request_count: AtomicU64,
79 priority_stats: parking_lot::RwLock<HashMap<Priority, (u64, u64)>>, }
81
82impl MetricsTracker {
83 fn new() -> Self {
84 Self {
85 total_wait_time_ms: AtomicU64::new(0),
86 total_execution_time_ms: AtomicU64::new(0),
87 request_count: AtomicU64::new(0),
88 priority_stats: parking_lot::RwLock::new(HashMap::new()),
89 }
90 }
91
92 fn record_completion(&self, wait_time_ms: u64, execution_time_ms: u64, priority: Priority) {
93 self.total_wait_time_ms
94 .fetch_add(wait_time_ms, Ordering::Relaxed);
95 self.total_execution_time_ms
96 .fetch_add(execution_time_ms, Ordering::Relaxed);
97 self.request_count.fetch_add(1, Ordering::Relaxed);
98
99 let mut priority_stats = self.priority_stats.write();
101 let (count, total_wait) = priority_stats.entry(priority).or_insert((0, 0));
102 *count += 1;
103 *total_wait += wait_time_ms;
104 }
105
106 fn avg_wait_time_ms(&self) -> f64 {
107 let total_wait = self.total_wait_time_ms.load(Ordering::Relaxed) as f64;
108 let count = self.request_count.load(Ordering::Relaxed) as f64;
109 if count > 0.0 {
110 total_wait / count
111 } else {
112 0.0
113 }
114 }
115
116 fn avg_execution_time_ms(&self) -> f64 {
117 let total_exec = self.total_execution_time_ms.load(Ordering::Relaxed) as f64;
118 let count = self.request_count.load(Ordering::Relaxed) as f64;
119 if count > 0.0 {
120 total_exec / count
121 } else {
122 0.0
123 }
124 }
125
126 fn priority_wait_time(&self, priority: Priority) -> f64 {
127 let priority_stats = self.priority_stats.read();
128 if let Some((count, total_wait)) = priority_stats.get(&priority) {
129 if *count > 0 {
130 *total_wait as f64 / *count as f64
131 } else {
132 0.0
133 }
134 } else {
135 0.0
136 }
137 }
138}
139
140impl PriorityScheduler {
141 pub fn new(config: SchedulerConfig) -> Self {
143 info!("Creating Priority scheduler with config: {:?}", config);
144
145 Self {
146 config,
147 waiting_queue: RwLock::new(PriorityQueue::new()),
148 request_map: RwLock::new(HashMap::new()),
149 running_requests: RwLock::new(HashMap::new()),
150 completed_counter: AtomicU64::new(0),
151 failed_counter: AtomicU64::new(0),
152 cancelled_counter: AtomicU64::new(0),
153 start_time: Instant::now(),
154 metrics_tracker: Arc::new(MetricsTracker::new()),
155 }
156 }
157
158 fn create_batch(&self, hint: BatchHint) -> Option<BatchPlan> {
160 let mut waiting_queue = self.waiting_queue.write();
161 let mut request_map = self.request_map.write();
162 let mut running_requests = self.running_requests.write();
163
164 if waiting_queue.is_empty() {
165 return None;
166 }
167
168 let mut batch_requests = Vec::new();
169 let mut total_tokens = 0;
170 let max_sequence_length = hint.max_tokens.min(2048);
171
172 let mut requests_to_readd = Vec::new();
174
175 while batch_requests.len() < hint.max_batch_size
176 && total_tokens < hint.max_tokens
177 && !waiting_queue.is_empty()
178 {
179 if let Some((request_id, _priority)) = waiting_queue.pop() {
180 if let Some(mut scheduled_req) = request_map.remove(&request_id) {
181 let request_tokens = scheduled_req.request.sampling_params.max_tokens;
182
183 if total_tokens + request_tokens <= hint.max_tokens {
185 scheduled_req.state = RequestState::Running;
186 scheduled_req.started_at = Some(chrono::Utc::now());
187 scheduled_req.queue_position = None;
188
189 total_tokens += request_tokens;
190
191 running_requests.insert(request_id.clone(), scheduled_req.clone());
193 batch_requests.push(scheduled_req);
194 } else {
195 let priority = RequestPriority::new(
197 scheduled_req.request.priority,
198 scheduled_req.submitted_at,
199 );
200 requests_to_readd.push((request_id, scheduled_req, priority));
201 break;
202 }
203 }
204 }
205 }
206
207 for (request_id, scheduled_req, priority) in requests_to_readd {
209 waiting_queue.push(request_id.clone(), priority);
210 request_map.insert(request_id, scheduled_req);
211 }
212
213 if batch_requests.is_empty() {
214 return None;
215 }
216
217 let batch_id = BatchId::new();
218 debug!(
219 "Creating priority batch {} with {} requests",
220 batch_id,
221 batch_requests.len()
222 );
223
224 let gpu_memory = (total_tokens * 16) as u64; let cpu_memory = (total_tokens * 4) as u64; let kv_cache_blocks = total_tokens / 16; let highest_priority = batch_requests
231 .iter()
232 .map(|req| req.request.priority)
233 .max()
234 .unwrap_or(Priority::Low);
235
236 let estimated_time_ms = match highest_priority {
237 Priority::Critical => 500, Priority::High => 750,
239 Priority::Normal => 1000,
240 Priority::Low => 1500,
241 };
242
243 Some(BatchPlan {
244 batch_id,
245 requests: batch_requests,
246 max_sequence_length,
247 estimated_time_ms: Some(estimated_time_ms),
248 resource_requirements: BatchResourceRequirements {
249 gpu_memory,
250 cpu_memory,
251 kv_cache_blocks,
252 recurrent_state_bytes: 0,
253 recurrent_state_slots: 0,
254 compute_units: 1,
255 },
256 created_at: chrono::Utc::now(),
257 })
258 }
259}
260
261#[async_trait]
262impl Scheduler for PriorityScheduler {
263 async fn submit(&self, request: InferenceRequest) -> Result<RequestId> {
264 let request_id = request.id.clone();
265 let priority = request.priority;
266 debug!(
267 "Submitting request {} with priority {:?} to priority scheduler",
268 request_id, priority
269 );
270
271 let waiting_queue = self.waiting_queue.read();
273 if waiting_queue.len() >= self.config.max_waiting_requests {
274 warn!("Queue is full, rejecting request {}", request_id);
275 return Err(ferrum_types::FerrumError::scheduler(
276 "Queue is full, cannot accept more requests",
277 ));
278 }
279 drop(waiting_queue);
280
281 let scheduled_request = ScheduledRequest::new(request);
283 let request_priority = RequestPriority::new(priority, scheduled_request.submitted_at);
284
285 let mut waiting_queue = self.waiting_queue.write();
287 let mut request_map = self.request_map.write();
288
289 let queue_position = waiting_queue.len();
290 let mut scheduled_req = scheduled_request;
291 scheduled_req.queue_position = Some(queue_position);
292
293 waiting_queue.push(request_id.clone(), request_priority);
294 request_map.insert(request_id.clone(), scheduled_req);
295
296 info!(
297 "Request {} queued with priority {:?} at position {}",
298 request_id, priority, queue_position
299 );
300 Ok(request_id)
301 }
302
303 async fn next_batch(&self, hint: BatchHint) -> Option<BatchPlan> {
304 self.create_batch(hint)
305 }
306
307 async fn complete(&self, request_id: RequestId, response: &InferenceResponse) -> Result<()> {
308 debug!("Completing request {}", request_id);
309
310 let mut running_requests = self.running_requests.write();
311 if let Some(scheduled_req) = running_requests.remove(&request_id) {
312 let wait_time = scheduled_req.age();
314 let execution_time = scheduled_req.processing_time().unwrap_or_default();
315 let priority = scheduled_req.request.priority;
316
317 self.metrics_tracker.record_completion(
318 wait_time.as_millis() as u64,
319 execution_time.as_millis() as u64,
320 priority,
321 );
322
323 match response.finish_reason {
324 ferrum_types::FinishReason::EOS
325 | ferrum_types::FinishReason::Stop
326 | ferrum_types::FinishReason::Length => {
327 self.completed_counter.fetch_add(1, Ordering::Relaxed);
328 debug!(
329 "Request {} (priority {:?}) completed successfully",
330 request_id, priority
331 );
332 }
333 _ => {
334 self.failed_counter.fetch_add(1, Ordering::Relaxed);
335 warn!(
336 "Request {} (priority {:?}) completed with error: {:?}",
337 request_id, priority, response.finish_reason
338 );
339 }
340 }
341
342 Ok(())
343 } else {
344 warn!("Attempted to complete unknown request: {}", request_id);
345 Err(ferrum_types::FerrumError::scheduler(format!(
346 "Request {} not found in running requests",
347 request_id
348 )))
349 }
350 }
351
352 async fn cancel(&self, request_id: RequestId) -> Result<bool> {
353 debug!("Cancelling request {}", request_id);
354
355 let mut waiting_queue = self.waiting_queue.write();
357 let mut request_map = self.request_map.write();
358
359 if waiting_queue.remove(&request_id).is_some() {
360 request_map.remove(&request_id);
361 self.cancelled_counter.fetch_add(1, Ordering::Relaxed);
362 info!("Request {} cancelled from waiting queue", request_id);
363 return Ok(true);
364 }
365 drop((waiting_queue, request_map));
366
367 let mut running_requests = self.running_requests.write();
369 if running_requests.remove(&request_id).is_some() {
370 self.cancelled_counter.fetch_add(1, Ordering::Relaxed);
371 warn!(
372 "Request {} cancelled while running (may cause issues)",
373 request_id
374 );
375 return Ok(true);
376 }
377
378 warn!("Request {} not found for cancellation", request_id);
379 Ok(false)
380 }
381
382 async fn update_priority(&self, request_id: RequestId, new_priority: Priority) -> Result<()> {
383 debug!(
384 "Updating priority for request {} to {:?}",
385 request_id, new_priority
386 );
387
388 let mut waiting_queue = self.waiting_queue.write();
389 let mut request_map = self.request_map.write();
390
391 if waiting_queue
393 .change_priority(
394 &request_id,
395 RequestPriority::new(new_priority, chrono::Utc::now()),
396 )
397 .is_some()
398 {
399 if let Some(scheduled_req) = request_map.get_mut(&request_id) {
401 scheduled_req.request.priority = new_priority;
402 info!(
403 "Updated priority for request {} to {:?}",
404 request_id, new_priority
405 );
406 return Ok(());
407 }
408 }
409
410 let mut running_requests = self.running_requests.write();
412 if let Some(scheduled_req) = running_requests.get_mut(&request_id) {
413 let old_priority = scheduled_req.request.priority;
414 scheduled_req.request.priority = new_priority;
415 warn!("Updated priority for running request {} from {:?} to {:?} (no effect on scheduling)",
416 request_id, old_priority, new_priority);
417 return Ok(());
418 }
419
420 warn!("Request {} not found for priority update", request_id);
421 Err(ferrum_types::FerrumError::scheduler(format!(
422 "Request {} not found",
423 request_id
424 )))
425 }
426
427 fn metrics(&self) -> SchedulerMetrics {
428 let waiting_queue = self.waiting_queue.read();
429 let running_requests = self.running_requests.read();
430
431 let waiting_count = waiting_queue.len();
432 let running_count = running_requests.len();
433 let completed_count = self.completed_counter.load(Ordering::Relaxed);
434 let failed_count = self.failed_counter.load(Ordering::Relaxed);
435 let cancelled_count = self.cancelled_counter.load(Ordering::Relaxed);
436
437 let uptime_secs = self.start_time.elapsed().as_secs_f64();
438 let throughput = if uptime_secs > 0.0 {
439 completed_count as f64 / uptime_secs
440 } else {
441 0.0
442 };
443
444 let queue_utilization = waiting_count as f32 / self.config.max_waiting_requests as f32;
445
446 let critical_wait = self.metrics_tracker.priority_wait_time(Priority::Critical);
448 let high_wait = self.metrics_tracker.priority_wait_time(Priority::High);
449 let avg_wait = self.metrics_tracker.avg_wait_time_ms();
450
451 let _p95_wait_estimate = if critical_wait > 0.0 || high_wait > 0.0 {
452 (critical_wait + high_wait) / 2.0 * 1.2 } else {
454 avg_wait * 1.5 };
456
457 ferrum_types::SchedulerStats {
458 waiting_requests: waiting_count,
459 running_requests: running_count,
460 preempted_requests: 0, completed_requests: completed_count,
462 failed_requests: failed_count,
463 cancelled_requests: cancelled_count,
464 avg_wait_time_ms: avg_wait,
465 avg_execution_time_ms: self.metrics_tracker.avg_execution_time_ms(),
466 throughput_rps: throughput,
467 queue_utilization,
468 }
469 }
470
471 fn config(&self) -> &SchedulerConfig {
472 &self.config
473 }
474
475 fn request_state(&self, request_id: &RequestId) -> Option<RequestState> {
476 if self.running_requests.read().contains_key(request_id) {
477 return Some(RequestState::Running);
478 }
479
480 if let Some(req) = self.request_map.read().get(request_id) {
481 return Some(req.state);
482 }
483
484 if self.waiting_queue.read().get_priority(request_id).is_some() {
485 return Some(RequestState::Waiting);
486 }
487
488 None
489 }
490
491 async fn preempt(&self, request_id: RequestId) -> Result<PreemptionResult> {
492 let mut running_requests = self.running_requests.write();
494
495 if let Some(scheduled_req) = running_requests.get(&request_id) {
496 let priority = scheduled_req.request.priority;
497
498 if matches!(priority, Priority::Low | Priority::Normal) {
500 if let Some(removed_req) = running_requests.remove(&request_id) {
501 let mut waiting_queue = self.waiting_queue.write();
503 let mut request_map = self.request_map.write();
504
505 let mut preempted_req = removed_req;
506 preempted_req.state = RequestState::Waiting;
507 preempted_req.started_at = None;
508
509 let request_priority =
510 RequestPriority::new(priority, preempted_req.submitted_at);
511 waiting_queue.push(request_id.clone(), request_priority);
512 request_map.insert(request_id.clone(), preempted_req);
513
514 warn!(
515 "Preempted request {} with priority {:?}",
516 request_id, priority
517 );
518
519 return Ok(PreemptionResult {
520 success: true,
521 saved_state: None, freed_resources: Default::default(),
523 });
524 }
525 }
526 }
527
528 Err(ferrum_types::FerrumError::scheduler(format!(
529 "Cannot preempt request {} (not found or high priority)",
530 request_id
531 )))
532 }
533
534 async fn resume(&self, _request_id: RequestId) -> Result<()> {
535 Ok(())
537 }
538}
539
540#[cfg(test)]
545mod tests {
546 use super::*;
547 use ferrum_types::{ModelId, SamplingParams};
548
549 fn create_test_request_with_priority(priority: Priority) -> InferenceRequest {
550 InferenceRequest {
551 id: RequestId::new(),
552 prompt: "test prompt".to_string(),
553 model_id: ModelId::new("test-model"),
554 sampling_params: SamplingParams::default(),
555 stream: false,
556 priority,
557 client_id: None,
558 session_id: None,
559 created_at: chrono::Utc::now(),
560 api_request: None,
561 evidence_request: Default::default(),
562 metadata: std::collections::HashMap::new(),
563 }
564 }
565
566 #[test]
567 fn test_request_priority_ordering() {
568 let now = chrono::Utc::now();
569
570 let critical = RequestPriority::new(Priority::Critical, now);
571 let high = RequestPriority::new(Priority::High, now);
572 let normal = RequestPriority::new(Priority::Normal, now);
573 let low = RequestPriority::new(Priority::Low, now);
574
575 assert!(critical > high);
577 assert!(high > normal);
578 assert!(normal > low);
579 }
580
581 #[test]
582 fn test_request_priority_fifo_within_same_level() {
583 let now = chrono::Utc::now();
584 let later = now + chrono::Duration::seconds(1);
585
586 let high1 = RequestPriority::new(Priority::High, now);
587 let high2 = RequestPriority::new(Priority::High, later);
588
589 assert!(high1 > high2);
591 }
592
593 #[tokio::test]
594 async fn test_priority_scheduler_creation() {
595 let config = SchedulerConfig::default();
596 let scheduler = PriorityScheduler::new(config.clone());
597
598 assert_eq!(
599 scheduler.config().max_waiting_requests,
600 config.max_waiting_requests
601 );
602 }
603
604 #[tokio::test]
605 async fn test_priority_scheduler_submit() {
606 let config = SchedulerConfig::default();
607 let scheduler = PriorityScheduler::new(config);
608
609 let request = create_test_request_with_priority(Priority::Normal);
610 let request_id = request.id.clone();
611
612 let result = scheduler.submit(request).await;
613 assert!(result.is_ok());
614 assert_eq!(result.unwrap(), request_id);
615
616 let waiting_queue = scheduler.waiting_queue.read();
618 assert_eq!(waiting_queue.len(), 1);
619 }
620
621 #[tokio::test]
622 async fn test_priority_ordering_in_queue() {
623 let config = SchedulerConfig::default();
624 let scheduler = PriorityScheduler::new(config);
625
626 let low_req = create_test_request_with_priority(Priority::Low);
628 let normal_req = create_test_request_with_priority(Priority::Normal);
629 let high_req = create_test_request_with_priority(Priority::High);
630 let critical_req = create_test_request_with_priority(Priority::Critical);
631
632 scheduler.submit(low_req).await.unwrap();
633 scheduler.submit(normal_req).await.unwrap();
634 scheduler.submit(high_req).await.unwrap();
635 scheduler.submit(critical_req.clone()).await.unwrap();
636
637 let batch = scheduler.next_batch(BatchHint::simple(10)).await;
639 assert!(batch.is_some());
640
641 let batch = batch.unwrap();
642 assert_eq!(batch.requests[0].request.priority, Priority::Critical);
644 }
645
646 #[tokio::test]
647 async fn test_priority_scheduler_cancel() {
648 let config = SchedulerConfig::default();
649 let scheduler = PriorityScheduler::new(config);
650
651 let request = create_test_request_with_priority(Priority::Normal);
652 let request_id = request.id.clone();
653
654 scheduler.submit(request).await.unwrap();
655
656 let result = scheduler.cancel(request_id).await;
658 assert!(result.is_ok());
659 assert!(result.unwrap());
660
661 let waiting_queue = scheduler.waiting_queue.read();
663 assert_eq!(waiting_queue.len(), 0);
664 }
665
666 #[tokio::test]
667 async fn test_priority_request_state_transitions() {
668 let config = SchedulerConfig::default();
669 let scheduler = PriorityScheduler::new(config);
670
671 let request = create_test_request_with_priority(Priority::Normal);
672 let request_id = request.id.clone();
673 scheduler.submit(request).await.unwrap();
674
675 assert_eq!(
676 scheduler.request_state(&request_id),
677 Some(RequestState::Waiting)
678 );
679
680 let _batch = scheduler.next_batch(BatchHint::simple(1)).await;
681 assert_eq!(
682 scheduler.request_state(&request_id),
683 Some(RequestState::Running)
684 );
685
686 scheduler.cancel(request_id.clone()).await.unwrap();
687 assert_eq!(scheduler.request_state(&request_id), None);
688 }
689
690 #[tokio::test]
691 async fn test_priority_update() {
692 let config = SchedulerConfig::default();
693 let scheduler = PriorityScheduler::new(config);
694
695 let request = create_test_request_with_priority(Priority::Low);
696 let request_id = request.id.clone();
697
698 scheduler.submit(request).await.unwrap();
699
700 let result = scheduler
702 .update_priority(request_id.clone(), Priority::High)
703 .await;
704 assert!(result.is_ok());
705
706 let request_map = scheduler.request_map.read();
708 if let Some(scheduled_req) = request_map.get(&request_id) {
709 assert_eq!(scheduled_req.request.priority, Priority::High);
710 }
711 }
712
713 #[tokio::test]
714 async fn test_metrics_tracking() {
715 let config = SchedulerConfig::default();
716 let scheduler = PriorityScheduler::new(config);
717
718 let request = create_test_request_with_priority(Priority::Normal);
719 scheduler.submit(request).await.unwrap();
720
721 let metrics = scheduler.metrics();
722 assert_eq!(metrics.waiting_requests, 1);
723 assert_eq!(metrics.running_requests, 0);
724 assert_eq!(metrics.completed_requests, 0);
725 }
726
727 #[tokio::test]
728 async fn test_batch_creation() {
729 let config = SchedulerConfig::default();
730 let scheduler = PriorityScheduler::new(config);
731
732 for i in 0..5 {
734 let priority = if i % 2 == 0 {
735 Priority::High
736 } else {
737 Priority::Normal
738 };
739 let request = create_test_request_with_priority(priority);
740 scheduler.submit(request).await.unwrap();
741 }
742
743 let batch = scheduler.next_batch(BatchHint::simple(3)).await;
745 assert!(batch.is_some());
746
747 let batch = batch.unwrap();
748 assert!(batch.requests.len() <= 3);
749 assert!(!batch.requests.is_empty());
750 }
751
752 #[tokio::test]
753 async fn test_preemption_low_priority() {
754 let config = SchedulerConfig::default();
755 let scheduler = PriorityScheduler::new(config);
756
757 let low_request = create_test_request_with_priority(Priority::Low);
758 let request_id = low_request.id.clone();
759
760 scheduler.submit(low_request).await.unwrap();
761
762 let _batch = scheduler.next_batch(BatchHint::simple(10)).await;
764
765 let result = scheduler.preempt(request_id).await;
767 assert!(result.is_ok());
768
769 let preemption_result = result.unwrap();
770 assert!(preemption_result.success);
771 }
772
773 #[tokio::test]
774 async fn test_cannot_preempt_high_priority() {
775 let config = SchedulerConfig::default();
776 let scheduler = PriorityScheduler::new(config);
777
778 let high_request = create_test_request_with_priority(Priority::High);
779 let request_id = high_request.id.clone();
780
781 scheduler.submit(high_request).await.unwrap();
782
783 let _batch = scheduler.next_batch(BatchHint::simple(10)).await;
785
786 let result = scheduler.preempt(request_id).await;
788 assert!(result.is_err());
789 }
790
791 #[tokio::test]
792 async fn test_queue_full() {
793 let mut config = SchedulerConfig::default();
794 config.max_waiting_requests = 2;
795 let scheduler = PriorityScheduler::new(config);
796
797 scheduler
799 .submit(create_test_request_with_priority(Priority::Normal))
800 .await
801 .unwrap();
802 scheduler
803 .submit(create_test_request_with_priority(Priority::Normal))
804 .await
805 .unwrap();
806
807 let result = scheduler
809 .submit(create_test_request_with_priority(Priority::Normal))
810 .await;
811 assert!(result.is_err());
812 }
813
814 #[test]
815 fn test_metrics_tracker_priority_stats() {
816 let tracker = MetricsTracker::new();
817
818 tracker.record_completion(100, 500, Priority::High);
819 tracker.record_completion(200, 600, Priority::High);
820 tracker.record_completion(300, 700, Priority::Normal);
821
822 let high_wait = tracker.priority_wait_time(Priority::High);
823 assert_eq!(high_wait, 150.0); let avg_wait = tracker.avg_wait_time_ms();
826 assert_eq!(avg_wait, 200.0); }
828
829 #[tokio::test]
830 async fn test_complete_request() {
831 let config = SchedulerConfig::default();
832 let scheduler = PriorityScheduler::new(config);
833
834 let request = create_test_request_with_priority(Priority::Normal);
835 let request_id = request.id.clone();
836
837 scheduler.submit(request).await.unwrap();
838
839 let _batch = scheduler.next_batch(BatchHint::simple(10)).await;
841
842 let response = InferenceResponse {
844 request_id: request_id.clone(),
845 text: "test".to_string(),
846 tokens: vec![],
847 finish_reason: ferrum_types::FinishReason::EOS,
848 usage: ferrum_types::TokenUsage::new(10, 5),
849 latency_ms: 100,
850 created_at: chrono::Utc::now(),
851 metadata: std::collections::HashMap::new(),
852 api_response: None,
853 execution_evidence: None,
854 };
855
856 let result = scheduler.complete(request_id, &response).await;
857 assert!(result.is_ok());
858
859 let metrics = scheduler.metrics();
861 assert_eq!(metrics.completed_requests, 1);
862 }
863
864 #[tokio::test]
865 async fn test_resume_request() {
866 let config = SchedulerConfig::default();
867 let scheduler = PriorityScheduler::new(config);
868
869 let request_id = RequestId::new();
870
871 let result = scheduler.resume(request_id).await;
873 assert!(result.is_ok());
874 }
875
876 #[tokio::test]
877 async fn test_multiple_priority_levels() {
878 let config = SchedulerConfig::default();
879 let scheduler = PriorityScheduler::new(config);
880
881 let priorities = vec![
883 Priority::Low,
884 Priority::Normal,
885 Priority::High,
886 Priority::Critical,
887 Priority::Normal,
888 Priority::Low,
889 ];
890
891 for priority in priorities {
892 scheduler
893 .submit(create_test_request_with_priority(priority))
894 .await
895 .unwrap();
896 }
897
898 let metrics = scheduler.metrics();
899 assert_eq!(metrics.waiting_requests, 6);
900
901 let batch = scheduler.next_batch(BatchHint::simple(10)).await;
903 assert!(batch.is_some());
904
905 let batch = batch.unwrap();
907 if batch.requests.len() >= 2 {
908 assert_eq!(batch.requests[0].request.priority, Priority::Critical);
910 }
911 }
912}