443 lines
13 KiB
Rust
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);
|
|
}
|