1001 lines
31 KiB
Rust
1001 lines
31 KiB
Rust
//! Request prioritization and queue management
|
|
//!
|
|
//! Provides comprehensive queue management including:
|
|
//! - Priority-based request scheduling with configurable levels
|
|
//! - Resource-aware load balancing across model instances
|
|
//! - Queue management with fair scheduling and starvation prevention
|
|
//! - Deadline-sensitive scheduling for time-critical requests
|
|
//! - Admission control with capacity planning
|
|
//! - Request lifecycle tracking and SLA monitoring
|
|
|
|
use anyhow::{Result, anyhow};
|
|
use chrono::{DateTime, Utc};
|
|
use dashmap::DashMap;
|
|
use parking_lot::{Mutex, RwLock};
|
|
use serde::{Deserialize, Serialize};
|
|
use std::{
|
|
cmp::Ordering,
|
|
collections::{BinaryHeap, HashMap},
|
|
sync::Arc,
|
|
time::Duration,
|
|
};
|
|
use tokio::sync::{Semaphore, oneshot};
|
|
use uuid::Uuid;
|
|
|
|
/// Request priority levels
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
|
|
pub enum RequestPriority {
|
|
Low = 1,
|
|
Normal = 2,
|
|
High = 3,
|
|
Critical = 4,
|
|
Emergency = 5,
|
|
}
|
|
|
|
impl RequestPriority {
|
|
/// Get priority multiplier for cost calculations
|
|
#[must_use]
|
|
pub fn cost_multiplier(self) -> f64 {
|
|
match self {
|
|
Self::Low => 0.8,
|
|
Self::Normal => 1.0,
|
|
Self::High => 1.5,
|
|
Self::Critical => 2.0,
|
|
Self::Emergency => 3.0,
|
|
}
|
|
}
|
|
|
|
/// Get queue jump allowance (how many requests can be skipped)
|
|
#[must_use]
|
|
pub fn queue_jump_allowance(self) -> usize {
|
|
match self {
|
|
Self::Low => 0,
|
|
Self::Normal => 0,
|
|
Self::High => 5,
|
|
Self::Critical => 20,
|
|
Self::Emergency => 100,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Request metadata for queue management
|
|
#[derive(Debug)]
|
|
pub struct QueuedRequest {
|
|
pub request_id: String,
|
|
pub user_id: String,
|
|
pub organization_id: Option<String>,
|
|
pub priority: RequestPriority,
|
|
pub submitted_at: DateTime<Utc>,
|
|
pub deadline: Option<DateTime<Utc>>,
|
|
pub estimated_duration: Duration,
|
|
pub estimated_tokens: u64,
|
|
pub model_name: String,
|
|
pub queue_position: usize,
|
|
pub wait_time: Duration,
|
|
pub retry_count: u32,
|
|
pub max_retries: u32,
|
|
pub callback_channel: Option<oneshot::Sender<QueueResult>>,
|
|
}
|
|
|
|
impl Clone for QueuedRequest {
|
|
fn clone(&self) -> Self {
|
|
Self {
|
|
request_id: self.request_id.clone(),
|
|
user_id: self.user_id.clone(),
|
|
organization_id: self.organization_id.clone(),
|
|
priority: self.priority,
|
|
submitted_at: self.submitted_at,
|
|
deadline: self.deadline,
|
|
estimated_duration: self.estimated_duration,
|
|
estimated_tokens: self.estimated_tokens,
|
|
model_name: self.model_name.clone(),
|
|
queue_position: self.queue_position,
|
|
wait_time: self.wait_time,
|
|
retry_count: self.retry_count,
|
|
max_retries: self.max_retries,
|
|
callback_channel: None, // Cannot clone oneshot::Sender
|
|
}
|
|
}
|
|
}
|
|
|
|
impl QueuedRequest {
|
|
/// Create new queued request
|
|
#[must_use]
|
|
pub fn new(
|
|
user_id: String,
|
|
organization_id: Option<String>,
|
|
priority: RequestPriority,
|
|
deadline: Option<DateTime<Utc>>,
|
|
estimated_duration: Duration,
|
|
estimated_tokens: u64,
|
|
model_name: String,
|
|
) -> (Self, oneshot::Receiver<QueueResult>) {
|
|
let (tx, rx) = oneshot::channel();
|
|
let request = Self {
|
|
request_id: Uuid::new_v4().to_string(),
|
|
user_id,
|
|
organization_id,
|
|
priority,
|
|
submitted_at: Utc::now(),
|
|
deadline,
|
|
estimated_duration,
|
|
estimated_tokens,
|
|
model_name,
|
|
queue_position: 0,
|
|
wait_time: Duration::ZERO,
|
|
retry_count: 0,
|
|
max_retries: 3,
|
|
callback_channel: Some(tx),
|
|
};
|
|
|
|
(request, rx)
|
|
}
|
|
|
|
/// Check if request has expired based on deadline
|
|
#[must_use]
|
|
pub fn is_expired(&self) -> bool {
|
|
if let Some(deadline) = self.deadline {
|
|
Utc::now() > deadline
|
|
} else {
|
|
false
|
|
}
|
|
}
|
|
|
|
/// Get urgency score (higher is more urgent)
|
|
#[must_use]
|
|
pub fn urgency_score(&self) -> f64 {
|
|
let base_score = f64::from(self.priority as u32) * 1000.0;
|
|
let wait_penalty = self.wait_time.as_secs() as f64 * 0.1;
|
|
let deadline_penalty = if let Some(deadline) = self.deadline {
|
|
let time_to_deadline = deadline.signed_duration_since(Utc::now()).num_seconds() as f64;
|
|
if time_to_deadline > 0.0 {
|
|
1000.0 / time_to_deadline
|
|
} else {
|
|
10000.0 // Very urgent if past deadline
|
|
}
|
|
} else {
|
|
0.0
|
|
};
|
|
|
|
base_score + wait_penalty + deadline_penalty
|
|
}
|
|
|
|
/// Check if request can be retried
|
|
#[must_use]
|
|
pub fn can_retry(&self) -> bool {
|
|
self.retry_count < self.max_retries
|
|
}
|
|
}
|
|
|
|
impl PartialEq for QueuedRequest {
|
|
fn eq(&self, other: &Self) -> bool {
|
|
self.request_id == other.request_id
|
|
}
|
|
}
|
|
|
|
impl Eq for QueuedRequest {}
|
|
|
|
impl PartialOrd for QueuedRequest {
|
|
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
|
|
Some(self.cmp(other))
|
|
}
|
|
}
|
|
|
|
impl Ord for QueuedRequest {
|
|
fn cmp(&self, other: &Self) -> Ordering {
|
|
// Higher urgency score first
|
|
other
|
|
.urgency_score()
|
|
.partial_cmp(&self.urgency_score())
|
|
.unwrap_or(Ordering::Equal)
|
|
.then_with(|| self.submitted_at.cmp(&other.submitted_at))
|
|
}
|
|
}
|
|
|
|
/// Queue result
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub enum QueueResult {
|
|
Accepted {
|
|
estimated_wait: Duration,
|
|
queue_position: usize,
|
|
},
|
|
Processing {
|
|
started_at: DateTime<Utc>,
|
|
},
|
|
Completed {
|
|
processing_time: Duration,
|
|
total_wait_time: Duration,
|
|
},
|
|
Failed {
|
|
error: String,
|
|
retry_possible: bool,
|
|
},
|
|
Expired {
|
|
reason: String,
|
|
},
|
|
Cancelled {
|
|
reason: String,
|
|
},
|
|
}
|
|
|
|
/// Queue statistics
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct QueueStats {
|
|
pub total_queued: usize,
|
|
pub total_processing: usize,
|
|
pub average_wait_time: Duration,
|
|
pub priority_breakdown: HashMap<RequestPriority, usize>,
|
|
pub model_breakdown: HashMap<String, usize>,
|
|
pub throughput_per_minute: f64,
|
|
pub success_rate: f64,
|
|
pub sla_compliance: f64,
|
|
}
|
|
|
|
/// Resource requirements for a request
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct ResourceRequirements {
|
|
pub gpu_memory_mb: u64,
|
|
pub system_memory_mb: u64,
|
|
pub gpu_compute_units: u32,
|
|
pub cpu_cores: u32,
|
|
pub estimated_duration: Duration,
|
|
}
|
|
|
|
/// Available system resources
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct SystemResources {
|
|
pub available_gpu_memory_mb: u64,
|
|
pub available_system_memory_mb: u64,
|
|
pub available_gpu_compute_units: u32,
|
|
pub available_cpu_cores: u32,
|
|
pub load_factor: f64,
|
|
}
|
|
|
|
impl SystemResources {
|
|
/// Check if resources can satisfy requirements
|
|
#[must_use]
|
|
pub fn can_satisfy(&self, requirements: &ResourceRequirements) -> bool {
|
|
self.available_gpu_memory_mb >= requirements.gpu_memory_mb
|
|
&& self.available_system_memory_mb >= requirements.system_memory_mb
|
|
&& self.available_gpu_compute_units >= requirements.gpu_compute_units
|
|
&& self.available_cpu_cores >= requirements.cpu_cores
|
|
}
|
|
|
|
/// Reserve resources
|
|
pub fn reserve(&mut self, requirements: &ResourceRequirements) -> Result<()> {
|
|
if !self.can_satisfy(requirements) {
|
|
return Err(anyhow!("Insufficient resources"));
|
|
}
|
|
|
|
self.available_gpu_memory_mb -= requirements.gpu_memory_mb;
|
|
self.available_system_memory_mb -= requirements.system_memory_mb;
|
|
self.available_gpu_compute_units -= requirements.gpu_compute_units;
|
|
self.available_cpu_cores -= requirements.cpu_cores;
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Release reserved resources
|
|
pub fn release(&mut self, requirements: &ResourceRequirements) {
|
|
self.available_gpu_memory_mb += requirements.gpu_memory_mb;
|
|
self.available_system_memory_mb += requirements.system_memory_mb;
|
|
self.available_gpu_compute_units += requirements.gpu_compute_units;
|
|
self.available_cpu_cores += requirements.cpu_cores;
|
|
}
|
|
}
|
|
|
|
/// Fair scheduling state to prevent starvation
|
|
#[derive(Debug)]
|
|
struct FairSchedulingState {
|
|
user_last_served: HashMap<String, DateTime<Utc>>,
|
|
organization_last_served: HashMap<String, DateTime<Utc>>,
|
|
priority_counters: HashMap<RequestPriority, u64>,
|
|
starvation_threshold: Duration,
|
|
}
|
|
|
|
impl Default for FairSchedulingState {
|
|
fn default() -> Self {
|
|
Self {
|
|
user_last_served: HashMap::new(),
|
|
organization_last_served: HashMap::new(),
|
|
priority_counters: HashMap::new(),
|
|
starvation_threshold: Duration::from_secs(300), // 5 minutes
|
|
}
|
|
}
|
|
}
|
|
|
|
impl FairSchedulingState {
|
|
/// Check if user/org is being starved
|
|
pub fn is_starved(&self, user_id: &str, organization_id: Option<&str>) -> bool {
|
|
let now = Utc::now();
|
|
|
|
// Check user starvation
|
|
if let Some(&last_served) = self.user_last_served.get(user_id)
|
|
&& now
|
|
.signed_duration_since(last_served)
|
|
.to_std()
|
|
.unwrap_or(Duration::ZERO)
|
|
> self.starvation_threshold
|
|
{
|
|
return true;
|
|
}
|
|
|
|
// Check organization starvation
|
|
if let Some(org_id) = organization_id
|
|
&& let Some(&last_served) = self.organization_last_served.get(org_id)
|
|
&& now
|
|
.signed_duration_since(last_served)
|
|
.to_std()
|
|
.unwrap_or(Duration::ZERO)
|
|
> self.starvation_threshold
|
|
{
|
|
return true;
|
|
}
|
|
|
|
false
|
|
}
|
|
|
|
/// Update last served time
|
|
pub fn update_served(
|
|
&mut self,
|
|
user_id: &str,
|
|
organization_id: Option<&str>,
|
|
priority: RequestPriority,
|
|
) {
|
|
let now = Utc::now();
|
|
self.user_last_served.insert(user_id.to_string(), now);
|
|
|
|
if let Some(org_id) = organization_id {
|
|
self.organization_last_served
|
|
.insert(org_id.to_string(), now);
|
|
}
|
|
|
|
*self.priority_counters.entry(priority).or_insert(0) += 1;
|
|
}
|
|
}
|
|
|
|
/// SLA configuration
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct SlaConfig {
|
|
pub max_wait_time: HashMap<RequestPriority, Duration>,
|
|
pub max_processing_time: HashMap<String, Duration>, // Per model
|
|
pub target_success_rate: f64,
|
|
pub target_availability: f64,
|
|
}
|
|
|
|
impl Default for SlaConfig {
|
|
fn default() -> Self {
|
|
let mut max_wait_time = HashMap::new();
|
|
max_wait_time.insert(RequestPriority::Low, Duration::from_secs(300));
|
|
max_wait_time.insert(RequestPriority::Normal, Duration::from_secs(120));
|
|
max_wait_time.insert(RequestPriority::High, Duration::from_secs(60));
|
|
max_wait_time.insert(RequestPriority::Critical, Duration::from_secs(30));
|
|
max_wait_time.insert(RequestPriority::Emergency, Duration::from_secs(10));
|
|
|
|
Self {
|
|
max_wait_time,
|
|
max_processing_time: HashMap::new(),
|
|
target_success_rate: 0.99,
|
|
target_availability: 0.999,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Request queue manager
|
|
pub struct QueueManager {
|
|
queue: Arc<Mutex<BinaryHeap<QueuedRequest>>>,
|
|
processing: Arc<DashMap<String, QueuedRequest>>,
|
|
completed: Arc<DashMap<String, QueueResult>>,
|
|
resources: Arc<RwLock<SystemResources>>,
|
|
fair_scheduler: Arc<Mutex<FairSchedulingState>>,
|
|
sla_config: Arc<RwLock<SlaConfig>>,
|
|
semaphore: Arc<Semaphore>,
|
|
stats: Arc<RwLock<QueueStats>>,
|
|
admission_controller: Arc<AdmissionController>,
|
|
}
|
|
|
|
/// Admission controller for managing system capacity
|
|
#[derive(Debug)]
|
|
pub struct AdmissionController {
|
|
max_queue_size: usize,
|
|
max_concurrent_requests: usize,
|
|
load_shedding_threshold: f64,
|
|
current_load: Arc<RwLock<f64>>,
|
|
}
|
|
|
|
impl AdmissionController {
|
|
/// Create new admission controller
|
|
#[must_use]
|
|
pub fn new(max_queue_size: usize, max_concurrent_requests: usize) -> Self {
|
|
Self {
|
|
max_queue_size,
|
|
max_concurrent_requests,
|
|
load_shedding_threshold: 0.9,
|
|
current_load: Arc::new(RwLock::new(0.0)),
|
|
}
|
|
}
|
|
|
|
/// Check if request should be admitted
|
|
#[must_use]
|
|
pub fn should_admit(
|
|
&self,
|
|
request: &QueuedRequest,
|
|
queue_size: usize,
|
|
processing_count: usize,
|
|
) -> bool {
|
|
// Check queue capacity
|
|
if queue_size >= self.max_queue_size {
|
|
// Only admit high priority requests if queue is full
|
|
return request.priority >= RequestPriority::High;
|
|
}
|
|
|
|
// Check processing capacity
|
|
if processing_count >= self.max_concurrent_requests {
|
|
return false;
|
|
}
|
|
|
|
// Check system load
|
|
let current_load = *self.current_load.read();
|
|
if current_load >= self.load_shedding_threshold {
|
|
// Only admit critical requests during high load
|
|
return request.priority >= RequestPriority::Critical;
|
|
}
|
|
|
|
true
|
|
}
|
|
|
|
/// Update system load
|
|
pub fn update_load(&self, load: f64) {
|
|
*self.current_load.write() = load.clamp(0.0, 1.0);
|
|
}
|
|
}
|
|
|
|
impl QueueManager {
|
|
/// Create new queue manager
|
|
#[must_use]
|
|
pub fn new(max_concurrent_requests: usize, max_queue_size: usize) -> Self {
|
|
let resources = SystemResources {
|
|
available_gpu_memory_mb: 16384, // 16GB
|
|
available_system_memory_mb: 32768, // 32GB
|
|
available_gpu_compute_units: 108, // Example GPU specs
|
|
available_cpu_cores: 16,
|
|
load_factor: 1.0,
|
|
};
|
|
|
|
Self {
|
|
queue: Arc::new(Mutex::new(BinaryHeap::new())),
|
|
processing: Arc::new(DashMap::new()),
|
|
completed: Arc::new(DashMap::new()),
|
|
resources: Arc::new(RwLock::new(resources)),
|
|
fair_scheduler: Arc::new(Mutex::new(FairSchedulingState::default())),
|
|
sla_config: Arc::new(RwLock::new(SlaConfig::default())),
|
|
semaphore: Arc::new(Semaphore::new(max_concurrent_requests)),
|
|
stats: Arc::new(RwLock::new(QueueStats {
|
|
total_queued: 0,
|
|
total_processing: 0,
|
|
average_wait_time: Duration::ZERO,
|
|
priority_breakdown: HashMap::new(),
|
|
model_breakdown: HashMap::new(),
|
|
throughput_per_minute: 0.0,
|
|
success_rate: 0.0,
|
|
sla_compliance: 0.0,
|
|
})),
|
|
admission_controller: Arc::new(AdmissionController::new(
|
|
max_queue_size,
|
|
max_concurrent_requests,
|
|
)),
|
|
}
|
|
}
|
|
|
|
/// Submit request to queue
|
|
pub async fn submit_request(
|
|
&self,
|
|
mut request: QueuedRequest,
|
|
) -> Result<oneshot::Receiver<QueueResult>> {
|
|
// Check admission control
|
|
let queue_size = self.queue.lock().len();
|
|
let processing_count = self.processing.len();
|
|
|
|
if !self
|
|
.admission_controller
|
|
.should_admit(&request, queue_size, processing_count)
|
|
{
|
|
let (tx, rx) = oneshot::channel();
|
|
let _ = tx.send(QueueResult::Failed {
|
|
error: "Request rejected by admission control".to_string(),
|
|
retry_possible: false,
|
|
});
|
|
return Ok(rx);
|
|
}
|
|
|
|
// Update queue position
|
|
request.queue_position = queue_size + 1;
|
|
|
|
let (callback_tx, callback_rx) = oneshot::channel();
|
|
request.callback_channel = Some(callback_tx);
|
|
|
|
// Add to queue
|
|
{
|
|
let mut queue = self.queue.lock();
|
|
queue.push(request.clone());
|
|
}
|
|
|
|
// Update stats
|
|
self.update_queue_stats().await;
|
|
|
|
Ok(callback_rx)
|
|
}
|
|
|
|
/// Process next request from queue
|
|
pub async fn process_next_request(&self) -> Option<QueuedRequest> {
|
|
// Try to acquire semaphore permit
|
|
let permit = self.semaphore.try_acquire();
|
|
if permit.is_err() {
|
|
return None;
|
|
}
|
|
|
|
let mut selected_request = None;
|
|
|
|
// Select request using fair scheduling
|
|
{
|
|
let mut queue = self.queue.lock();
|
|
let fair_scheduler = self.fair_scheduler.lock();
|
|
|
|
// Convert heap to vector for processing
|
|
let mut requests: Vec<_> = queue.drain().collect();
|
|
requests.sort(); // Sort by urgency/priority
|
|
|
|
// Apply fair scheduling
|
|
for (i, request) in requests.iter().enumerate() {
|
|
// Check if expired
|
|
if request.is_expired() {
|
|
// Skip expired requests
|
|
continue;
|
|
}
|
|
|
|
// Check starvation prevention
|
|
if fair_scheduler.is_starved(&request.user_id, request.organization_id.as_deref()) {
|
|
selected_request = Some(requests.remove(i));
|
|
break;
|
|
}
|
|
|
|
// Check resource availability
|
|
let resource_requirements = self
|
|
.estimate_resource_requirements(&request.model_name, request.estimated_tokens);
|
|
if self.resources.read().can_satisfy(&resource_requirements) {
|
|
selected_request = Some(requests.remove(i));
|
|
break;
|
|
}
|
|
}
|
|
|
|
// Put remaining requests back in queue
|
|
for request in requests {
|
|
queue.push(request);
|
|
}
|
|
}
|
|
|
|
if let Some(mut request) = selected_request {
|
|
// Reserve resources
|
|
let resource_requirements =
|
|
self.estimate_resource_requirements(&request.model_name, request.estimated_tokens);
|
|
if let Err(_) = self.resources.write().reserve(&resource_requirements) {
|
|
// Resource reservation failed, put request back
|
|
let mut queue = self.queue.lock();
|
|
queue.push(request);
|
|
return None;
|
|
}
|
|
|
|
// Update fair scheduler
|
|
{
|
|
let mut fair_scheduler = self.fair_scheduler.lock();
|
|
fair_scheduler.update_served(
|
|
&request.user_id,
|
|
request.organization_id.as_deref(),
|
|
request.priority,
|
|
);
|
|
}
|
|
|
|
// Move to processing
|
|
request.wait_time = Utc::now()
|
|
.signed_duration_since(request.submitted_at)
|
|
.to_std()
|
|
.unwrap_or(Duration::ZERO);
|
|
self.processing
|
|
.insert(request.request_id.clone(), request.clone());
|
|
|
|
// Notify request started
|
|
if let Some(callback) = request.callback_channel.take() {
|
|
let _ = callback.send(QueueResult::Processing {
|
|
started_at: Utc::now(),
|
|
});
|
|
}
|
|
|
|
// Update stats
|
|
self.update_processing_stats().await;
|
|
|
|
Some(request)
|
|
} else {
|
|
None
|
|
}
|
|
}
|
|
|
|
/// Complete request processing
|
|
pub async fn complete_request(&self, request_id: &str, result: QueueResult) -> Result<()> {
|
|
if let Some((_, mut request)) = self.processing.remove(request_id) {
|
|
// Release resources
|
|
let resource_requirements =
|
|
self.estimate_resource_requirements(&request.model_name, request.estimated_tokens);
|
|
self.resources.write().release(&resource_requirements);
|
|
|
|
// Release semaphore permit
|
|
self.semaphore.add_permits(1);
|
|
|
|
// Store result
|
|
self.completed
|
|
.insert(request_id.to_string(), result.clone());
|
|
|
|
// Notify completion
|
|
if let Some(callback) = request.callback_channel.take() {
|
|
let _ = callback.send(result);
|
|
}
|
|
|
|
// Update stats
|
|
self.update_completion_stats().await;
|
|
|
|
Ok(())
|
|
} else {
|
|
Err(anyhow!("Request {request_id} not found in processing"))
|
|
}
|
|
}
|
|
|
|
/// Cancel request
|
|
pub async fn cancel_request(&self, request_id: &str, reason: String) -> Result<()> {
|
|
// Try to remove from queue first
|
|
{
|
|
let mut queue = self.queue.lock();
|
|
let mut requests: Vec<_> = queue.drain().collect();
|
|
|
|
if let Some(pos) = requests.iter().position(|r| r.request_id == request_id) {
|
|
let mut request = requests.remove(pos);
|
|
|
|
// Notify cancellation
|
|
if let Some(callback) = request.callback_channel.take() {
|
|
let _ = callback.send(QueueResult::Cancelled { reason });
|
|
}
|
|
|
|
// Put remaining requests back
|
|
for req in requests {
|
|
queue.push(req);
|
|
}
|
|
|
|
return Ok(());
|
|
}
|
|
// Put all requests back
|
|
for req in requests {
|
|
queue.push(req);
|
|
}
|
|
}
|
|
|
|
// Try to remove from processing
|
|
if let Some((_, mut request)) = self.processing.remove(request_id) {
|
|
// Release resources
|
|
let resource_requirements =
|
|
self.estimate_resource_requirements(&request.model_name, request.estimated_tokens);
|
|
self.resources.write().release(&resource_requirements);
|
|
|
|
// Release semaphore permit
|
|
self.semaphore.add_permits(1);
|
|
|
|
// Notify cancellation
|
|
if let Some(callback) = request.callback_channel.take() {
|
|
let _ = callback.send(QueueResult::Cancelled { reason });
|
|
}
|
|
|
|
return Ok(());
|
|
}
|
|
|
|
Err(anyhow!("Request {request_id} not found"))
|
|
}
|
|
|
|
/// Get queue statistics
|
|
#[must_use]
|
|
pub fn get_queue_stats(&self) -> QueueStats {
|
|
self.stats.read().clone()
|
|
}
|
|
|
|
/// Get request status
|
|
#[must_use]
|
|
pub fn get_request_status(&self, request_id: &str) -> Option<QueueResult> {
|
|
// Check if completed
|
|
if let Some(result) = self.completed.get(request_id) {
|
|
return Some(result.clone());
|
|
}
|
|
|
|
// Check if processing
|
|
if let Some(_request) = self.processing.get(request_id) {
|
|
return Some(QueueResult::Processing {
|
|
started_at: Utc::now(), // Approximate
|
|
});
|
|
}
|
|
|
|
// Check if in queue
|
|
{
|
|
let queue = self.queue.lock();
|
|
for (i, request) in queue.iter().enumerate() {
|
|
if request.request_id == request_id {
|
|
return Some(QueueResult::Accepted {
|
|
estimated_wait: Duration::from_secs(i as u64 * 30), // Rough estimate
|
|
queue_position: i + 1,
|
|
});
|
|
}
|
|
}
|
|
}
|
|
|
|
None
|
|
}
|
|
|
|
/// Estimate resource requirements for a model/token combination
|
|
fn estimate_resource_requirements(
|
|
&self,
|
|
model_name: &str,
|
|
estimated_tokens: u64,
|
|
) -> ResourceRequirements {
|
|
// Simple estimation - in production this would be more sophisticated
|
|
let base_memory = match model_name {
|
|
name if name.contains("gpt-4") => 8192, // 8GB
|
|
name if name.contains("gpt-3.5") => 4096, // 4GB
|
|
_ => 2048, // 2GB default
|
|
};
|
|
|
|
let token_memory = (estimated_tokens / 1000) * 10; // 10MB per 1K tokens
|
|
|
|
ResourceRequirements {
|
|
gpu_memory_mb: base_memory + token_memory,
|
|
system_memory_mb: u64::midpoint(base_memory, token_memory),
|
|
gpu_compute_units: if estimated_tokens > 4000 { 4 } else { 2 },
|
|
cpu_cores: 2,
|
|
estimated_duration: Duration::from_millis(estimated_tokens * 10), // 10ms per token
|
|
}
|
|
}
|
|
|
|
/// Update queue statistics
|
|
async fn update_queue_stats(&self) {
|
|
let queue_size = self.queue.lock().len();
|
|
let processing_count = self.processing.len();
|
|
|
|
let mut stats = self.stats.write();
|
|
stats.total_queued = queue_size;
|
|
stats.total_processing = processing_count;
|
|
|
|
// Update priority breakdown
|
|
stats.priority_breakdown.clear();
|
|
{
|
|
let queue = self.queue.lock();
|
|
for request in queue.iter() {
|
|
*stats
|
|
.priority_breakdown
|
|
.entry(request.priority)
|
|
.or_insert(0) += 1;
|
|
}
|
|
}
|
|
|
|
// Update model breakdown
|
|
stats.model_breakdown.clear();
|
|
{
|
|
let queue = self.queue.lock();
|
|
for request in queue.iter() {
|
|
*stats
|
|
.model_breakdown
|
|
.entry(request.model_name.clone())
|
|
.or_insert(0) += 1;
|
|
}
|
|
}
|
|
|
|
for request_ref in self.processing.iter() {
|
|
*stats
|
|
.model_breakdown
|
|
.entry(request_ref.model_name.clone())
|
|
.or_insert(0) += 1;
|
|
}
|
|
}
|
|
|
|
/// Update processing statistics
|
|
async fn update_processing_stats(&self) {
|
|
// Implementation for processing stats
|
|
}
|
|
|
|
/// Update completion statistics
|
|
async fn update_completion_stats(&self) {
|
|
// Implementation for completion stats
|
|
}
|
|
|
|
/// Cleanup expired requests
|
|
pub async fn cleanup_expired_requests(&self) {
|
|
let mut expired_requests = Vec::new();
|
|
|
|
// Check queue for expired requests
|
|
{
|
|
let mut queue = self.queue.lock();
|
|
let requests: Vec<_> = queue.drain().collect();
|
|
|
|
for request in requests {
|
|
if request.is_expired() {
|
|
expired_requests.push(request);
|
|
} else {
|
|
queue.push(request);
|
|
}
|
|
}
|
|
}
|
|
|
|
// Notify expired requests
|
|
for mut request in expired_requests {
|
|
if let Some(callback) = request.callback_channel.take() {
|
|
let _ = callback.send(QueueResult::Expired {
|
|
reason: "Request deadline exceeded".to_string(),
|
|
});
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use tokio::time::sleep;
|
|
|
|
#[test]
|
|
fn test_request_priority_ordering() {
|
|
assert!(RequestPriority::Emergency > RequestPriority::Critical);
|
|
assert!(RequestPriority::Critical > RequestPriority::High);
|
|
assert!(RequestPriority::High > RequestPriority::Normal);
|
|
assert!(RequestPriority::Normal > RequestPriority::Low);
|
|
}
|
|
|
|
#[test]
|
|
fn test_queued_request_urgency() {
|
|
let (request, _) = QueuedRequest::new(
|
|
"user1".to_string(),
|
|
None,
|
|
RequestPriority::High,
|
|
None,
|
|
Duration::from_secs(30),
|
|
100,
|
|
"gpt-4".to_string(),
|
|
);
|
|
|
|
let urgency = request.urgency_score();
|
|
assert!(urgency > 0.0);
|
|
assert!(urgency >= 3000.0); // High priority base score
|
|
}
|
|
|
|
#[test]
|
|
fn test_system_resources() {
|
|
let mut resources = SystemResources {
|
|
available_gpu_memory_mb: 8192,
|
|
available_system_memory_mb: 16384,
|
|
available_gpu_compute_units: 10,
|
|
available_cpu_cores: 8,
|
|
load_factor: 1.0,
|
|
};
|
|
|
|
let requirements = ResourceRequirements {
|
|
gpu_memory_mb: 4096,
|
|
system_memory_mb: 8192,
|
|
gpu_compute_units: 4,
|
|
cpu_cores: 4,
|
|
estimated_duration: Duration::from_secs(30),
|
|
};
|
|
|
|
assert!(resources.can_satisfy(&requirements));
|
|
|
|
resources.reserve(&requirements).unwrap();
|
|
assert_eq!(resources.available_gpu_memory_mb, 4096);
|
|
assert_eq!(resources.available_system_memory_mb, 8192);
|
|
|
|
resources.release(&requirements);
|
|
assert_eq!(resources.available_gpu_memory_mb, 8192);
|
|
assert_eq!(resources.available_system_memory_mb, 16384);
|
|
}
|
|
|
|
#[test]
|
|
fn test_admission_controller() {
|
|
let controller = AdmissionController::new(100, 10);
|
|
|
|
let (request, _) = QueuedRequest::new(
|
|
"user1".to_string(),
|
|
None,
|
|
RequestPriority::Normal,
|
|
None,
|
|
Duration::from_secs(30),
|
|
100,
|
|
"gpt-4".to_string(),
|
|
);
|
|
|
|
// Should admit under normal conditions
|
|
assert!(controller.should_admit(&request, 50, 5));
|
|
|
|
// Should reject when queue is full (unless high priority)
|
|
assert!(!controller.should_admit(&request, 100, 5));
|
|
|
|
let (high_priority_request, _) = QueuedRequest::new(
|
|
"user1".to_string(),
|
|
None,
|
|
RequestPriority::High,
|
|
None,
|
|
Duration::from_secs(30),
|
|
100,
|
|
"gpt-4".to_string(),
|
|
);
|
|
|
|
assert!(controller.should_admit(&high_priority_request, 100, 5));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_queue_manager() {
|
|
let manager = QueueManager::new(5, 100);
|
|
|
|
let (request, callback) = QueuedRequest::new(
|
|
"user1".to_string(),
|
|
None,
|
|
RequestPriority::Normal,
|
|
None,
|
|
Duration::from_secs(30),
|
|
100,
|
|
"gpt-4".to_string(),
|
|
);
|
|
|
|
let receiver = manager.submit_request(request).await.unwrap();
|
|
|
|
// Should be able to get next request
|
|
let next_request = manager.process_next_request().await;
|
|
assert!(next_request.is_some());
|
|
|
|
let request = next_request.unwrap();
|
|
|
|
// Complete the request
|
|
manager
|
|
.complete_request(
|
|
&request.request_id,
|
|
QueueResult::Completed {
|
|
processing_time: Duration::from_secs(30),
|
|
total_wait_time: Duration::from_secs(5),
|
|
},
|
|
)
|
|
.await
|
|
.unwrap();
|
|
|
|
// Check stats
|
|
let stats = manager.get_queue_stats();
|
|
assert_eq!(stats.total_processing, 0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_fair_scheduling_starvation_detection() {
|
|
let mut scheduler = FairSchedulingState::default();
|
|
|
|
// Initially not starved
|
|
assert!(!scheduler.is_starved("user1", None));
|
|
|
|
// Update served time to past
|
|
let past_time = Utc::now() - chrono::Duration::seconds(600); // 10 minutes ago
|
|
scheduler
|
|
.user_last_served
|
|
.insert("user1".to_string(), past_time);
|
|
|
|
// Should be starved now
|
|
assert!(scheduler.is_starved("user1", None));
|
|
}
|
|
|
|
#[test]
|
|
fn test_sla_config() {
|
|
let sla = SlaConfig::default();
|
|
|
|
assert!(sla.max_wait_time.contains_key(&RequestPriority::Emergency));
|
|
assert!(
|
|
sla.max_wait_time[&RequestPriority::Emergency]
|
|
< sla.max_wait_time[&RequestPriority::Low]
|
|
);
|
|
assert!(sla.target_success_rate > 0.0);
|
|
}
|
|
}
|