Files
rustytorch/crates/production/rtx-inference/tests/request_management.rs
T
2026-03-04 00:08:42 +00:00

443 lines
13 KiB
Rust

//! Tests for request management system
//!
//! This module tests request lifecycle, priority handling, and SLA tracking.
use rtx_inference::{
error::InferenceError,
request::{
InferenceRequest, OverflowStrategy, RequestId, RequestManager, RequestManagerConfig,
RequestMetrics, RequestPriority, RequestStatus,
},
};
use std::time::{Duration, Instant};
use tokio::time::sleep;
/// Test request creation and validation
#[tokio::test]
async fn test_request_creation() {
let request = InferenceRequest {
id: RequestId::new(),
model_name: "llama-7b".to_string(),
input_tokens: vec![1, 2, 3, 4, 5],
max_new_tokens: 100,
temperature: 0.8,
top_p: 0.9,
priority: RequestPriority::Normal,
deadline: Some(Instant::now() + Duration::from_secs(10)),
..Default::default()
};
assert!(!request.id.is_nil());
assert_eq!(request.model_name, "llama-7b");
assert_eq!(request.input_tokens.len(), 5);
assert_eq!(request.max_new_tokens, 100);
assert!(request.deadline.is_some());
// Test request validation
let validation_result = request.validate();
assert!(validation_result.is_ok());
}
/// Test request validation with invalid parameters
#[tokio::test]
async fn test_request_validation() {
// Empty input tokens should fail
let invalid_request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![],
max_new_tokens: 100,
..Default::default()
};
let result = invalid_request.validate();
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
InferenceError::InvalidRequest { .. }
));
// Zero max_new_tokens should fail
let invalid_request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![1, 2, 3],
max_new_tokens: 0,
..Default::default()
};
let result = invalid_request.validate();
assert!(result.is_err());
// Invalid temperature should fail
let invalid_request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![1, 2, 3],
max_new_tokens: 100,
temperature: -1.0, // Invalid negative temperature
..Default::default()
};
let result = invalid_request.validate();
assert!(result.is_err());
}
/// Test request manager creation and configuration
#[tokio::test]
async fn test_request_manager_creation() {
let config = RequestManagerConfig {
max_concurrent_requests: 100,
max_queue_size: 1000,
default_timeout: Duration::from_secs(30),
metrics_enabled: true,
..Default::default()
};
let manager = RequestManager::new(config).await;
assert!(manager.is_ok());
let manager = manager.unwrap();
assert_eq!(manager.config().max_concurrent_requests, 100);
assert_eq!(manager.config().max_queue_size, 1000);
assert!(manager.config().metrics_enabled);
}
/// Test request submission and tracking
#[tokio::test]
async fn test_request_submission_tracking() {
let config = RequestManagerConfig::default();
let mut manager = RequestManager::new(config).await.unwrap();
let request = InferenceRequest {
id: RequestId::new(),
model_name: "gpt-3.5-turbo".to_string(),
input_tokens: vec![1, 2, 3, 4],
max_new_tokens: 50,
priority: RequestPriority::High,
..Default::default()
};
let request_id = request.id;
let result = manager.submit_request(request).await;
assert!(result.is_ok());
// Request should be tracked
let status = manager.get_request_status(request_id).await.unwrap();
assert_eq!(status, RequestStatus::Queued);
// Should appear in active requests
let active_requests = manager.get_active_requests().await;
assert!(active_requests.contains_key(&request_id));
}
/// Test request priority ordering
#[tokio::test]
async fn test_request_priority_ordering() {
let config = RequestManagerConfig::default();
let mut manager = RequestManager::new(config).await.unwrap();
// Submit requests with different priorities
let low_priority = InferenceRequest {
id: RequestId::new(),
priority: RequestPriority::Low,
input_tokens: vec![1; 10],
..Default::default()
};
let low_id = low_priority.id;
let high_priority = InferenceRequest {
id: RequestId::new(),
priority: RequestPriority::High,
input_tokens: vec![2; 10],
..Default::default()
};
let high_id = high_priority.id;
let normal_priority = InferenceRequest {
id: RequestId::new(),
priority: RequestPriority::Normal,
input_tokens: vec![3; 10],
..Default::default()
};
let normal_id = normal_priority.id;
// Submit in low -> high -> normal order
manager.submit_request(low_priority).await.unwrap();
manager.submit_request(high_priority).await.unwrap();
manager.submit_request(normal_priority).await.unwrap();
// Get next requests - should be in priority order: high -> normal -> low
let next_requests = manager.get_next_requests(3).await.unwrap();
assert_eq!(next_requests.len(), 3);
assert_eq!(next_requests[0].id, high_id);
assert_eq!(next_requests[1].id, normal_id);
assert_eq!(next_requests[2].id, low_id);
}
/// Test request lifecycle state transitions
#[tokio::test]
async fn test_request_lifecycle() {
let config = RequestManagerConfig::default();
let mut manager = RequestManager::new(config).await.unwrap();
let request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![1, 2, 3],
max_new_tokens: 10,
..Default::default()
};
let request_id = request.id;
manager.submit_request(request).await.unwrap();
// Initial state: Queued
assert_eq!(
manager.get_request_status(request_id).await.unwrap(),
RequestStatus::Queued
);
// Transition to Processing
manager.mark_request_processing(request_id).await.unwrap();
assert_eq!(
manager.get_request_status(request_id).await.unwrap(),
RequestStatus::Processing
);
// Transition to Generating
manager.mark_request_generating(request_id).await.unwrap();
assert_eq!(
manager.get_request_status(request_id).await.unwrap(),
RequestStatus::Generating
);
// Complete successfully
let output_tokens = vec![10, 11, 12, 13, 14];
manager
.complete_request(request_id, output_tokens.clone())
.await
.unwrap();
assert_eq!(
manager.get_request_status(request_id).await.unwrap(),
RequestStatus::Completed
);
// Verify output
let result = manager.get_request_result(request_id).await.unwrap();
assert_eq!(result.output_tokens, output_tokens);
assert!(result.completion_time.is_some());
}
/// Test request timeout handling
#[tokio::test]
async fn test_request_timeout() {
let config = RequestManagerConfig {
default_timeout: Duration::from_millis(100),
..Default::default()
};
let mut manager = RequestManager::new(config).await.unwrap();
let request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![1, 2, 3],
deadline: Some(Instant::now() + Duration::from_millis(50)),
..Default::default()
};
let request_id = request.id;
manager.submit_request(request).await.unwrap();
// Mark as processing but don't complete
manager.mark_request_processing(request_id).await.unwrap();
// Wait for timeout
sleep(Duration::from_millis(150)).await;
// Check for timeouts
let timed_out = manager.check_timeouts().await.unwrap();
assert!(!timed_out.is_empty());
assert!(timed_out.contains(&request_id));
// Status should be timeout
assert_eq!(
manager.get_request_status(request_id).await.unwrap(),
RequestStatus::TimedOut
);
}
/// Test request cancellation
#[tokio::test]
async fn test_request_cancellation() {
let config = RequestManagerConfig::default();
let mut manager = RequestManager::new(config).await.unwrap();
let request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![1, 2, 3],
..Default::default()
};
let request_id = request.id;
manager.submit_request(request).await.unwrap();
// Cancel the request
let result = manager.cancel_request(request_id).await;
assert!(result.is_ok());
// Status should be cancelled
assert_eq!(
manager.get_request_status(request_id).await.unwrap(),
RequestStatus::Cancelled
);
// Should not appear in next requests
let next_requests = manager.get_next_requests(10).await.unwrap();
assert!(!next_requests.iter().any(|r| r.id == request_id));
}
/// Test request metrics collection
#[tokio::test]
async fn test_request_metrics() {
let config = RequestManagerConfig {
metrics_enabled: true,
..Default::default()
};
let mut manager = RequestManager::new(config).await.unwrap();
let request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![1, 2, 3, 4, 5],
max_new_tokens: 20,
..Default::default()
};
let request_id = request.id;
manager.submit_request(request).await.unwrap();
// Process through lifecycle
manager.mark_request_processing(request_id).await.unwrap();
sleep(Duration::from_millis(10)).await;
manager.mark_request_generating(request_id).await.unwrap();
sleep(Duration::from_millis(50)).await;
let output = vec![10, 11, 12];
manager.complete_request(request_id, output).await.unwrap();
// Get metrics
let metrics = manager.get_request_metrics(request_id).await.unwrap();
assert!(metrics.queue_time > Duration::ZERO);
assert!(metrics.processing_time > Duration::ZERO);
assert!(metrics.generation_time > Duration::ZERO);
assert!(metrics.total_time > Duration::ZERO);
assert_eq!(metrics.input_token_count, 5);
assert_eq!(metrics.output_token_count, 3);
assert!(metrics.tokens_per_second > 0.0);
}
/// Test concurrent request handling
#[tokio::test]
async fn test_concurrent_request_handling() {
let config = RequestManagerConfig {
max_concurrent_requests: 5,
..Default::default()
};
let mut manager = RequestManager::new(config).await.unwrap();
// Submit more requests than max concurrent
let mut request_ids = Vec::new();
for i in 0..8 {
let request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![i; 10],
..Default::default()
};
request_ids.push(request.id);
manager.submit_request(request).await.unwrap();
}
// Get next batch - should respect max concurrent limit
let next_batch = manager.get_next_requests(10).await.unwrap();
assert!(next_batch.len() <= 5);
// Mark some as processing
for (i, request) in next_batch.iter().enumerate() {
if i < 3 {
manager.mark_request_processing(request.id).await.unwrap();
}
}
// Should be able to get more requests up to the limit
let additional_batch = manager.get_next_requests(10).await.unwrap();
let total_processing = manager.count_processing_requests().await;
assert!(total_processing <= 5);
}
/// Test queue capacity and overflow handling
#[tokio::test]
async fn test_queue_capacity_overflow() {
let config = RequestManagerConfig {
max_queue_size: 5,
overflow_strategy: OverflowStrategy::RejectNew,
..Default::default()
};
let mut manager = RequestManager::new(config).await.unwrap();
// Fill queue to capacity
for i in 0..5 {
let request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![i; 5],
..Default::default()
};
manager.submit_request(request).await.unwrap();
}
// Additional request should be rejected
let overflow_request = InferenceRequest {
id: RequestId::new(),
input_tokens: vec![99; 5],
..Default::default()
};
let result = manager.submit_request(overflow_request).await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
InferenceError::QueueFull { .. }
));
}
/// Test queue statistics and monitoring
#[tokio::test]
async fn test_queue_statistics() {
let config = RequestManagerConfig::default();
let mut manager = RequestManager::new(config).await.unwrap();
// Submit requests with different priorities
for i in 0..3 {
let request = InferenceRequest {
id: RequestId::new(),
priority: RequestPriority::High,
input_tokens: vec![i; 5],
..Default::default()
};
manager.submit_request(request).await.unwrap();
}
for i in 0..5 {
let request = InferenceRequest {
id: RequestId::new(),
priority: RequestPriority::Normal,
input_tokens: vec![i; 5],
..Default::default()
};
manager.submit_request(request).await.unwrap();
}
let stats = manager.get_queue_stats().await;
assert_eq!(stats.total_queued, 8);
assert_eq!(stats.high_priority_count, 3);
assert_eq!(stats.normal_priority_count, 5);
assert_eq!(stats.low_priority_count, 0);
assert!(stats.average_queue_time >= Duration::ZERO);
assert_eq!(stats.total_processing, 0);
}