Skip to main content

ferrum_scheduler/implementations/
priority.rs

1//! Priority-based 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 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
24/// Priority scheduler that processes requests based on priority and submission time
25pub struct PriorityScheduler {
26    /// Configuration
27    config: SchedulerConfig,
28    /// Priority queue for waiting requests (higher priority first, then FIFO within priority)
29    waiting_queue: RwLock<PriorityQueue<RequestId, RequestPriority>>,
30    /// Map from request ID to scheduled request
31    request_map: RwLock<HashMap<RequestId, ScheduledRequest>>,
32    /// Running requests
33    running_requests: RwLock<HashMap<RequestId, ScheduledRequest>>,
34    /// Completed request counter
35    completed_counter: AtomicU64,
36    /// Failed request counter
37    failed_counter: AtomicU64,
38    /// Cancelled request counter
39    cancelled_counter: AtomicU64,
40    /// Scheduler start time
41    start_time: Instant,
42    /// Metrics tracking
43    metrics_tracker: Arc<MetricsTracker>,
44}
45
46/// Priority wrapper for the priority queue (higher values = higher priority)
47#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
48struct RequestPriority {
49    /// Priority level (higher = more important)
50    priority: i32,
51    /// Submission time as negative nanoseconds (older = higher priority within same level)
52    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        // Use negative timestamp so older requests get higher priority within same level
65        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
74/// Internal metrics tracker
75struct 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)>>, // (count, total_wait_time)
80}
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        // Track per-priority stats
100        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    /// Create new priority scheduler
142    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    /// Create batch from waiting queue based on priority
159    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        // Process requests in priority order
173        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                    // Check if adding this request would exceed limits
184                    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                        // Move to running requests
192                        running_requests.insert(request_id.clone(), scheduled_req.clone());
193                        batch_requests.push(scheduled_req);
194                    } else {
195                        // Put the request back
196                        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        // Re-add requests that didn't fit
208        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        // Calculate resource requirements based on batch composition
225        let gpu_memory = (total_tokens * 16) as u64; // Estimate: 16 bytes per token
226        let cpu_memory = (total_tokens * 4) as u64; // Estimate: 4 bytes per token
227        let kv_cache_blocks = total_tokens / 16; // Assume 16 tokens per block
228
229        // Estimate execution time based on highest priority in batch
230        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, // Fast lane for critical requests
238            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        // Check queue capacity
272        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        // Create scheduled request
282        let scheduled_request = ScheduledRequest::new(request);
283        let request_priority = RequestPriority::new(priority, scheduled_request.submitted_at);
284
285        // Add to priority queue and request map
286        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            // Calculate metrics
313            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        // Try to remove from waiting queue first
356        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        // Try to remove from running requests
368        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        // Check if request is in waiting queue
392        if waiting_queue
393            .change_priority(
394                &request_id,
395                RequestPriority::new(new_priority, chrono::Utc::now()),
396            )
397            .is_some()
398        {
399            // Update the request priority in the request map
400            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        // Check running requests (can't change priority of running request, but update for logging)
411        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        // Calculate priority-based P95 estimates
447        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 // Priority requests should have better P95
453        } else {
454            avg_wait * 1.5 // Fallback to simple estimate
455        };
456
457        ferrum_types::SchedulerStats {
458            waiting_requests: waiting_count,
459            running_requests: running_count,
460            preempted_requests: 0, // MVP: no preemption tracking
461            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        // Simple preemption: can preempt lower priority running requests
493        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            // Only allow preempting Low and Medium priority requests
499            if matches!(priority, Priority::Low | Priority::Normal) {
500                if let Some(removed_req) = running_requests.remove(&request_id) {
501                    // Move back to waiting queue with updated priority
502                    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, // Simplified - no state saving in this implementation
522                        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        // In this implementation, preempted requests are automatically re-queued
536        Ok(())
537    }
538}
539
540// ============================================================================
541// 内联单元测试
542// ============================================================================
543
544#[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        // 优先级应该是:Critical > High > Normal > Low
576        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        // 相同优先级内,早提交的请求优先级更高(负时间戳)
590        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        // 验证请求在队列中
617        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        // 提交不同优先级的请求
627        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        // 获取批次,应该优先处理高优先级请求
638        let batch = scheduler.next_batch(BatchHint::simple(10)).await;
639        assert!(batch.is_some());
640
641        let batch = batch.unwrap();
642        // 第一个请求应该是 Critical 优先级
643        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        // 取消请求
657        let result = scheduler.cancel(request_id).await;
658        assert!(result.is_ok());
659        assert!(result.unwrap());
660
661        // 验证请求不在队列中
662        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        // 更新优先级
701        let result = scheduler
702            .update_priority(request_id.clone(), Priority::High)
703            .await;
704        assert!(result.is_ok());
705
706        // 验证优先级已更新
707        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        // 提交多个请求
733        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        // 创建批次
744        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        // 获取批次,将请求移到运行队列
763        let _batch = scheduler.next_batch(BatchHint::simple(10)).await;
764
765        // 尝试抢占
766        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        // 获取批次
784        let _batch = scheduler.next_batch(BatchHint::simple(10)).await;
785
786        // 尝试抢占高优先级请求应该失败
787        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        // 填满队列
798        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        // 第三个请求应该失败
808        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); // (100 + 200) / 2
824
825        let avg_wait = tracker.avg_wait_time_ms();
826        assert_eq!(avg_wait, 200.0); // (100 + 200 + 300) / 3
827    }
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        // 获取批次,移到运行队列
840        let _batch = scheduler.next_batch(BatchHint::simple(10)).await;
841
842        // 完成请求
843        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        // 验证已完成计数增加
860        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        // Resume 应该总是成功(在这个实现中)
872        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        // 提交各种优先级的请求
882        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        // 获取批次
902        let batch = scheduler.next_batch(BatchHint::simple(10)).await;
903        assert!(batch.is_some());
904
905        // 批次中的请求应该按优先级排序
906        let batch = batch.unwrap();
907        if batch.requests.len() >= 2 {
908            // 第一个应该是最高优先级
909            assert_eq!(batch.requests[0].request.priority, Priority::Critical);
910        }
911    }
912}