Tiny refactor router test contexts (#16340)
This commit is contained in:
@@ -1,172 +1,25 @@
|
|||||||
mod common;
|
mod common;
|
||||||
|
|
||||||
use std::sync::Arc;
|
|
||||||
|
|
||||||
use axum::{
|
use axum::{
|
||||||
body::Body,
|
body::Body,
|
||||||
extract::Request,
|
extract::Request,
|
||||||
http::{header::CONTENT_TYPE, StatusCode},
|
http::{header::CONTENT_TYPE, StatusCode},
|
||||||
};
|
};
|
||||||
use common::mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType};
|
use common::{
|
||||||
use reqwest::Client;
|
mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType},
|
||||||
use serde_json::json;
|
AppTestContext,
|
||||||
use smg::{
|
|
||||||
app_context::AppContext,
|
|
||||||
config::{RouterConfig, RoutingMode},
|
|
||||||
core::Job,
|
|
||||||
routers::{RouterFactory, RouterTrait},
|
|
||||||
};
|
};
|
||||||
|
use serde_json::json;
|
||||||
|
use smg::{config::RouterConfig, routers::RouterFactory};
|
||||||
use tower::ServiceExt;
|
use tower::ServiceExt;
|
||||||
|
|
||||||
/// Test context that manages mock workers
|
|
||||||
struct TestContext {
|
|
||||||
workers: Vec<MockWorker>,
|
|
||||||
router: Arc<dyn RouterTrait>,
|
|
||||||
_client: Client,
|
|
||||||
_config: RouterConfig,
|
|
||||||
app_context: Arc<AppContext>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl TestContext {
|
|
||||||
async fn new(worker_configs: Vec<MockWorkerConfig>) -> Self {
|
|
||||||
// Create default router config
|
|
||||||
let config = RouterConfig::builder()
|
|
||||||
.regular_mode(vec![])
|
|
||||||
.random_policy()
|
|
||||||
.host("127.0.0.1")
|
|
||||||
.port(3002)
|
|
||||||
.max_payload_size(256 * 1024 * 1024)
|
|
||||||
.request_timeout_secs(600)
|
|
||||||
.worker_startup_timeout_secs(1)
|
|
||||||
.worker_startup_check_interval_secs(1)
|
|
||||||
.max_concurrent_requests(64)
|
|
||||||
.queue_timeout_secs(60)
|
|
||||||
.build_unchecked();
|
|
||||||
|
|
||||||
Self::new_with_config(config, worker_configs).await
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn new_with_config(
|
|
||||||
mut config: RouterConfig,
|
|
||||||
worker_configs: Vec<MockWorkerConfig>,
|
|
||||||
) -> Self {
|
|
||||||
let mut workers = Vec::new();
|
|
||||||
let mut worker_urls = Vec::new();
|
|
||||||
|
|
||||||
// Start mock workers if any
|
|
||||||
for worker_config in worker_configs {
|
|
||||||
let mut worker = MockWorker::new(worker_config);
|
|
||||||
let url = worker.start().await.unwrap();
|
|
||||||
worker_urls.push(url);
|
|
||||||
workers.push(worker);
|
|
||||||
}
|
|
||||||
|
|
||||||
if !workers.is_empty() {
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(200)).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
// Update config with worker URLs if not already set
|
|
||||||
match &mut config.mode {
|
|
||||||
RoutingMode::Regular {
|
|
||||||
worker_urls: ref mut urls,
|
|
||||||
} => {
|
|
||||||
if urls.is_empty() {
|
|
||||||
*urls = worker_urls.clone();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
RoutingMode::OpenAI {
|
|
||||||
worker_urls: ref mut urls,
|
|
||||||
} => {
|
|
||||||
if urls.is_empty() {
|
|
||||||
*urls = worker_urls.clone();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
_ => {} // PrefillDecode mode has its own setup
|
|
||||||
}
|
|
||||||
|
|
||||||
let client = Client::builder()
|
|
||||||
.timeout(std::time::Duration::from_secs(config.request_timeout_secs))
|
|
||||||
.build()
|
|
||||||
.unwrap();
|
|
||||||
|
|
||||||
// Create app context
|
|
||||||
let app_context = common::create_test_context(config.clone()).await;
|
|
||||||
|
|
||||||
// Submit worker initialization job (same as real server does)
|
|
||||||
if !worker_urls.is_empty() {
|
|
||||||
let job_queue = app_context
|
|
||||||
.worker_job_queue
|
|
||||||
.get()
|
|
||||||
.expect("JobQueue should be initialized");
|
|
||||||
let job = Job::InitializeWorkersFromConfig {
|
|
||||||
router_config: Box::new(config.clone()),
|
|
||||||
};
|
|
||||||
job_queue
|
|
||||||
.submit(job)
|
|
||||||
.await
|
|
||||||
.expect("Failed to submit worker initialization job");
|
|
||||||
|
|
||||||
// Poll until all workers are healthy (up to 10 seconds)
|
|
||||||
let expected_count = worker_urls.len();
|
|
||||||
let start = tokio::time::Instant::now();
|
|
||||||
let timeout_duration = tokio::time::Duration::from_secs(10);
|
|
||||||
loop {
|
|
||||||
let healthy_workers = app_context
|
|
||||||
.worker_registry
|
|
||||||
.get_all()
|
|
||||||
.iter()
|
|
||||||
.filter(|w| w.is_healthy())
|
|
||||||
.count();
|
|
||||||
|
|
||||||
if healthy_workers >= expected_count {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
|
|
||||||
if start.elapsed() > timeout_duration {
|
|
||||||
panic!(
|
|
||||||
"Timeout waiting for {} workers to become healthy (only {} ready)",
|
|
||||||
expected_count, healthy_workers
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Create router
|
|
||||||
let router = RouterFactory::create_router(&app_context).await.unwrap();
|
|
||||||
let router = Arc::from(router);
|
|
||||||
|
|
||||||
Self {
|
|
||||||
workers,
|
|
||||||
router,
|
|
||||||
_client: client,
|
|
||||||
_config: config,
|
|
||||||
app_context,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn create_app(&self) -> axum::Router {
|
|
||||||
common::test_app::create_test_app_with_context(
|
|
||||||
Arc::clone(&self.router),
|
|
||||||
Arc::clone(&self.app_context),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn shutdown(mut self) {
|
|
||||||
for worker in &mut self.workers {
|
|
||||||
worker.stop().await;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod health_tests {
|
mod health_tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_liveness_endpoint() {
|
async fn test_liveness_endpoint() {
|
||||||
let ctx = TestContext::new(vec![]).await;
|
let ctx = AppTestContext::new(vec![]).await;
|
||||||
let app = ctx.create_app().await;
|
let app = ctx.create_app().await;
|
||||||
|
|
||||||
let req = Request::builder()
|
let req = Request::builder()
|
||||||
@@ -183,7 +36,7 @@ mod health_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_readiness_with_healthy_workers() {
|
async fn test_readiness_with_healthy_workers() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18001,
|
port: 18001,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -208,7 +61,7 @@ mod health_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_readiness_with_unhealthy_workers() {
|
async fn test_readiness_with_unhealthy_workers() {
|
||||||
let ctx = TestContext::new(vec![]).await;
|
let ctx = AppTestContext::new(vec![]).await;
|
||||||
|
|
||||||
let app = ctx.create_app().await;
|
let app = ctx.create_app().await;
|
||||||
|
|
||||||
@@ -226,7 +79,7 @@ mod health_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_health_endpoint_details() {
|
async fn test_health_endpoint_details() {
|
||||||
let ctx = TestContext::new(vec![
|
let ctx = AppTestContext::new(vec![
|
||||||
MockWorkerConfig {
|
MockWorkerConfig {
|
||||||
port: 18003,
|
port: 18003,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
@@ -260,7 +113,7 @@ mod health_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_health_generate_endpoint() {
|
async fn test_health_generate_endpoint() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18005,
|
port: 18005,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -296,7 +149,7 @@ mod generation_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_generate_success() {
|
async fn test_generate_success() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18101,
|
port: 18101,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -337,7 +190,7 @@ mod generation_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_generate_streaming() {
|
async fn test_generate_streaming() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18102,
|
port: 18102,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -373,7 +226,7 @@ mod generation_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_generate_with_worker_failure() {
|
async fn test_generate_with_worker_failure() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18103,
|
port: 18103,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -404,7 +257,7 @@ mod generation_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_chat_completions_success() {
|
async fn test_v1_chat_completions_success() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18104,
|
port: 18104,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -449,7 +302,7 @@ mod model_info_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_get_server_info() {
|
async fn test_get_server_info() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18201,
|
port: 18201,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -487,7 +340,7 @@ mod model_info_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_get_model_info() {
|
async fn test_get_model_info() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18202,
|
port: 18202,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -532,7 +385,7 @@ mod model_info_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_models() {
|
async fn test_v1_models() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18203,
|
port: 18203,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -588,7 +441,7 @@ mod model_info_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_model_info_with_no_workers() {
|
async fn test_model_info_with_no_workers() {
|
||||||
let ctx = TestContext::new(vec![]).await;
|
let ctx = AppTestContext::new(vec![]).await;
|
||||||
let app = ctx.create_app().await;
|
let app = ctx.create_app().await;
|
||||||
|
|
||||||
let req = Request::builder()
|
let req = Request::builder()
|
||||||
@@ -644,7 +497,7 @@ mod model_info_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_model_info_with_multiple_workers() {
|
async fn test_model_info_with_multiple_workers() {
|
||||||
let ctx = TestContext::new(vec![
|
let ctx = AppTestContext::new(vec![
|
||||||
MockWorkerConfig {
|
MockWorkerConfig {
|
||||||
port: 18204,
|
port: 18204,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
@@ -689,7 +542,7 @@ mod model_info_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_model_info_with_unhealthy_worker() {
|
async fn test_model_info_with_unhealthy_worker() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18206,
|
port: 18206,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -725,7 +578,7 @@ mod router_policy_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_random_policy() {
|
async fn test_random_policy() {
|
||||||
let ctx = TestContext::new(vec![
|
let ctx = AppTestContext::new(vec![
|
||||||
MockWorkerConfig {
|
MockWorkerConfig {
|
||||||
port: 18801,
|
port: 18801,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
@@ -768,7 +621,7 @@ mod router_policy_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_worker_selection() {
|
async fn test_worker_selection() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18207,
|
port: 18207,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -798,7 +651,7 @@ mod responses_endpoint_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_responses_non_streaming() {
|
async fn test_v1_responses_non_streaming() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18950,
|
port: 18950,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -837,7 +690,7 @@ mod responses_endpoint_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_responses_streaming() {
|
async fn test_v1_responses_streaming() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18951,
|
port: 18951,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -878,7 +731,7 @@ mod responses_endpoint_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_responses_get() {
|
async fn test_v1_responses_get() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18952,
|
port: 18952,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -927,7 +780,7 @@ mod responses_endpoint_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_responses_cancel() {
|
async fn test_v1_responses_cancel() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18953,
|
port: 18953,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -976,7 +829,7 @@ mod responses_endpoint_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_responses_delete_not_implemented() {
|
async fn test_v1_responses_delete_not_implemented() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18954,
|
port: 18954,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1019,7 +872,7 @@ mod responses_endpoint_tests {
|
|||||||
.queue_timeout_secs(60)
|
.queue_timeout_secs(60)
|
||||||
.build_unchecked();
|
.build_unchecked();
|
||||||
|
|
||||||
let ctx = TestContext::new_with_config(
|
let ctx = AppTestContext::new_with_config(
|
||||||
config,
|
config,
|
||||||
vec![], // No workers needed
|
vec![], // No workers needed
|
||||||
)
|
)
|
||||||
@@ -1073,7 +926,7 @@ mod responses_endpoint_tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_responses_get_multi_worker_fanout() {
|
async fn test_v1_responses_get_multi_worker_fanout() {
|
||||||
// Start two mock workers
|
// Start two mock workers
|
||||||
let ctx = TestContext::new(vec![
|
let ctx = AppTestContext::new(vec![
|
||||||
MockWorkerConfig {
|
MockWorkerConfig {
|
||||||
port: 18960,
|
port: 18960,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
@@ -1148,7 +1001,7 @@ mod error_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_404_not_found() {
|
async fn test_404_not_found() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18401,
|
port: 18401,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1185,7 +1038,7 @@ mod error_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_method_not_allowed() {
|
async fn test_method_not_allowed() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18402,
|
port: 18402,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1237,7 +1090,7 @@ mod error_tests {
|
|||||||
.queue_timeout_secs(60)
|
.queue_timeout_secs(60)
|
||||||
.build_unchecked();
|
.build_unchecked();
|
||||||
|
|
||||||
let ctx = TestContext::new_with_config(
|
let ctx = AppTestContext::new_with_config(
|
||||||
config,
|
config,
|
||||||
vec![MockWorkerConfig {
|
vec![MockWorkerConfig {
|
||||||
port: 18403,
|
port: 18403,
|
||||||
@@ -1258,7 +1111,7 @@ mod error_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_invalid_json_payload() {
|
async fn test_invalid_json_payload() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18404,
|
port: 18404,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1296,7 +1149,7 @@ mod error_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_invalid_model() {
|
async fn test_invalid_model() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18406,
|
port: 18406,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1334,7 +1187,7 @@ mod cache_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_flush_cache() {
|
async fn test_flush_cache() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18501,
|
port: 18501,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1371,7 +1224,7 @@ mod cache_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_get_loads() {
|
async fn test_get_loads() {
|
||||||
let ctx = TestContext::new(vec![
|
let ctx = AppTestContext::new(vec![
|
||||||
MockWorkerConfig {
|
MockWorkerConfig {
|
||||||
port: 18502,
|
port: 18502,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
@@ -1414,7 +1267,7 @@ mod cache_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_flush_cache_no_workers() {
|
async fn test_flush_cache_no_workers() {
|
||||||
let ctx = TestContext::new(vec![]).await;
|
let ctx = AppTestContext::new(vec![]).await;
|
||||||
|
|
||||||
let app = ctx.create_app().await;
|
let app = ctx.create_app().await;
|
||||||
|
|
||||||
@@ -1441,7 +1294,7 @@ mod load_balancing_tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_request_distribution() {
|
async fn test_request_distribution() {
|
||||||
// Create multiple workers
|
// Create multiple workers
|
||||||
let ctx = TestContext::new(vec![
|
let ctx = AppTestContext::new(vec![
|
||||||
MockWorkerConfig {
|
MockWorkerConfig {
|
||||||
port: 18601,
|
port: 18601,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
@@ -1558,7 +1411,7 @@ mod request_id_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_request_id_generation() {
|
async fn test_request_id_generation() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18901,
|
port: 18901,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1675,7 +1528,7 @@ mod request_id_tests {
|
|||||||
.queue_timeout_secs(60)
|
.queue_timeout_secs(60)
|
||||||
.build_unchecked();
|
.build_unchecked();
|
||||||
|
|
||||||
let ctx = TestContext::new_with_config(
|
let ctx = AppTestContext::new_with_config(
|
||||||
config,
|
config,
|
||||||
vec![MockWorkerConfig {
|
vec![MockWorkerConfig {
|
||||||
port: 18902,
|
port: 18902,
|
||||||
@@ -1720,7 +1573,7 @@ mod rerank_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_rerank_success() {
|
async fn test_rerank_success() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18105,
|
port: 18105,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1772,7 +1625,7 @@ mod rerank_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_rerank_with_top_k() {
|
async fn test_rerank_with_top_k() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18106,
|
port: 18106,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1819,7 +1672,7 @@ mod rerank_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_rerank_without_documents() {
|
async fn test_rerank_without_documents() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18107,
|
port: 18107,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1863,7 +1716,7 @@ mod rerank_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_rerank_worker_failure() {
|
async fn test_rerank_worker_failure() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18108,
|
port: 18108,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1896,7 +1749,7 @@ mod rerank_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_rerank_compatibility() {
|
async fn test_v1_rerank_compatibility() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18110,
|
port: 18110,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -1953,7 +1806,7 @@ mod rerank_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_rerank_invalid_request() {
|
async fn test_rerank_invalid_request() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = AppTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 18111,
|
port: 18111,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
|
|||||||
@@ -13,12 +13,14 @@ use std::{
|
|||||||
sync::{Arc, Mutex, OnceLock},
|
sync::{Arc, Mutex, OnceLock},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
use mock_worker::{MockWorker, MockWorkerConfig};
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
use smg::{
|
use smg::{
|
||||||
app_context::AppContext,
|
app_context::AppContext,
|
||||||
config::{RouterConfig, RoutingMode},
|
config::{RouterConfig, RoutingMode},
|
||||||
core::{
|
core::{
|
||||||
BasicWorkerBuilder, LoadMonitor, ModelCard, RuntimeType, Worker, WorkerRegistry, WorkerType,
|
BasicWorkerBuilder, Job, LoadMonitor, ModelCard, RuntimeType, Worker, WorkerRegistry,
|
||||||
|
WorkerType,
|
||||||
},
|
},
|
||||||
data_connector::{
|
data_connector::{
|
||||||
MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage,
|
MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage,
|
||||||
@@ -27,10 +29,257 @@ use smg::{
|
|||||||
policies::PolicyRegistry,
|
policies::PolicyRegistry,
|
||||||
protocols::common::{Function, Tool},
|
protocols::common::{Function, Tool},
|
||||||
reasoning_parser::ParserFactory as ReasoningParserFactory,
|
reasoning_parser::ParserFactory as ReasoningParserFactory,
|
||||||
|
routers::{RouterFactory, RouterTrait},
|
||||||
tokenizer::registry::TokenizerRegistry,
|
tokenizer::registry::TokenizerRegistry,
|
||||||
tool_parser::ParserFactory as ToolParserFactory,
|
tool_parser::ParserFactory as ToolParserFactory,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
/// Test context for directly testing mock workers without full router setup.
|
||||||
|
pub struct WorkerTestContext {
|
||||||
|
pub workers: Vec<MockWorker>,
|
||||||
|
pub worker_urls: Vec<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl WorkerTestContext {
|
||||||
|
pub async fn new(worker_configs: Vec<MockWorkerConfig>) -> Self {
|
||||||
|
let mut workers = Vec::new();
|
||||||
|
let mut worker_urls = Vec::new();
|
||||||
|
|
||||||
|
for worker_config in worker_configs {
|
||||||
|
let mut worker = MockWorker::new(worker_config);
|
||||||
|
let url = worker.start().await.unwrap();
|
||||||
|
worker_urls.push(url);
|
||||||
|
workers.push(worker);
|
||||||
|
}
|
||||||
|
|
||||||
|
if !workers.is_empty() {
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_millis(200)).await;
|
||||||
|
}
|
||||||
|
|
||||||
|
Self {
|
||||||
|
workers,
|
||||||
|
worker_urls,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn first_worker_url(&self) -> Option<&str> {
|
||||||
|
self.worker_urls.first().map(|s| s.as_str())
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn make_request(
|
||||||
|
&self,
|
||||||
|
endpoint: &str,
|
||||||
|
body: serde_json::Value,
|
||||||
|
) -> Result<serde_json::Value, String> {
|
||||||
|
let client = reqwest::Client::new();
|
||||||
|
let worker_url = self
|
||||||
|
.first_worker_url()
|
||||||
|
.ok_or_else(|| "No workers available".to_string())?;
|
||||||
|
|
||||||
|
let response = client
|
||||||
|
.post(format!("{}{}", worker_url, endpoint))
|
||||||
|
.json(&body)
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.map_err(|e| format!("Request failed: {}", e))?;
|
||||||
|
|
||||||
|
if !response.status().is_success() {
|
||||||
|
return Err(format!("Request failed with status: {}", response.status()));
|
||||||
|
}
|
||||||
|
|
||||||
|
response
|
||||||
|
.json::<serde_json::Value>()
|
||||||
|
.await
|
||||||
|
.map_err(|e| format!("Failed to parse response: {}", e))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn make_streaming_request(
|
||||||
|
&self,
|
||||||
|
endpoint: &str,
|
||||||
|
body: serde_json::Value,
|
||||||
|
) -> Result<Vec<String>, String> {
|
||||||
|
use futures_util::StreamExt;
|
||||||
|
|
||||||
|
let client = reqwest::Client::new();
|
||||||
|
let worker_url = self
|
||||||
|
.first_worker_url()
|
||||||
|
.ok_or_else(|| "No workers available".to_string())?;
|
||||||
|
|
||||||
|
let response = client
|
||||||
|
.post(format!("{}{}", worker_url, endpoint))
|
||||||
|
.json(&body)
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.map_err(|e| format!("Request failed: {}", e))?;
|
||||||
|
|
||||||
|
if !response.status().is_success() {
|
||||||
|
return Err(format!("Request failed with status: {}", response.status()));
|
||||||
|
}
|
||||||
|
|
||||||
|
let content_type = response
|
||||||
|
.headers()
|
||||||
|
.get("content-type")
|
||||||
|
.and_then(|v| v.to_str().ok())
|
||||||
|
.unwrap_or("");
|
||||||
|
|
||||||
|
if !content_type.contains("text/event-stream") {
|
||||||
|
return Err("Response is not a stream".to_string());
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut stream = response.bytes_stream();
|
||||||
|
let mut events = Vec::new();
|
||||||
|
|
||||||
|
while let Some(chunk) = stream.next().await {
|
||||||
|
if let Ok(bytes) = chunk {
|
||||||
|
let text = String::from_utf8_lossy(&bytes);
|
||||||
|
for line in text.lines() {
|
||||||
|
if let Some(stripped) = line.strip_prefix("data: ") {
|
||||||
|
events.push(stripped.to_string());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(events)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn shutdown(mut self) {
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||||
|
for worker in &mut self.workers {
|
||||||
|
worker.stop().await;
|
||||||
|
}
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Test context for integration tests that go through the full axum app stack.
|
||||||
|
pub struct AppTestContext {
|
||||||
|
pub workers: Vec<MockWorker>,
|
||||||
|
pub router: Arc<dyn RouterTrait>,
|
||||||
|
pub config: RouterConfig,
|
||||||
|
pub app_context: Arc<AppContext>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl AppTestContext {
|
||||||
|
pub async fn new(worker_configs: Vec<MockWorkerConfig>) -> Self {
|
||||||
|
let config = RouterConfig::builder()
|
||||||
|
.regular_mode(vec![])
|
||||||
|
.random_policy()
|
||||||
|
.host("127.0.0.1")
|
||||||
|
.port(3002)
|
||||||
|
.max_payload_size(256 * 1024 * 1024)
|
||||||
|
.request_timeout_secs(600)
|
||||||
|
.worker_startup_timeout_secs(1)
|
||||||
|
.worker_startup_check_interval_secs(1)
|
||||||
|
.max_concurrent_requests(64)
|
||||||
|
.queue_timeout_secs(60)
|
||||||
|
.build_unchecked();
|
||||||
|
|
||||||
|
Self::new_with_config(config, worker_configs).await
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn new_with_config(
|
||||||
|
mut config: RouterConfig,
|
||||||
|
worker_configs: Vec<MockWorkerConfig>,
|
||||||
|
) -> Self {
|
||||||
|
let mut workers = Vec::new();
|
||||||
|
let mut worker_urls = Vec::new();
|
||||||
|
|
||||||
|
for worker_config in worker_configs {
|
||||||
|
let mut worker = MockWorker::new(worker_config);
|
||||||
|
let url = worker.start().await.unwrap();
|
||||||
|
worker_urls.push(url);
|
||||||
|
workers.push(worker);
|
||||||
|
}
|
||||||
|
|
||||||
|
if !workers.is_empty() {
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_millis(200)).await;
|
||||||
|
}
|
||||||
|
|
||||||
|
match &mut config.mode {
|
||||||
|
RoutingMode::Regular {
|
||||||
|
worker_urls: ref mut urls,
|
||||||
|
} => {
|
||||||
|
if urls.is_empty() {
|
||||||
|
*urls = worker_urls.clone();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
RoutingMode::OpenAI {
|
||||||
|
worker_urls: ref mut urls,
|
||||||
|
} => {
|
||||||
|
if urls.is_empty() {
|
||||||
|
*urls = worker_urls.clone();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
|
||||||
|
let app_context = create_test_context(config.clone()).await;
|
||||||
|
|
||||||
|
if !worker_urls.is_empty() {
|
||||||
|
let job_queue = app_context
|
||||||
|
.worker_job_queue
|
||||||
|
.get()
|
||||||
|
.expect("JobQueue should be initialized");
|
||||||
|
let job = Job::InitializeWorkersFromConfig {
|
||||||
|
router_config: Box::new(config.clone()),
|
||||||
|
};
|
||||||
|
job_queue
|
||||||
|
.submit(job)
|
||||||
|
.await
|
||||||
|
.expect("Failed to submit worker initialization job");
|
||||||
|
|
||||||
|
let expected_count = worker_urls.len();
|
||||||
|
let start = tokio::time::Instant::now();
|
||||||
|
let timeout_duration = tokio::time::Duration::from_secs(10);
|
||||||
|
loop {
|
||||||
|
let healthy_workers = app_context
|
||||||
|
.worker_registry
|
||||||
|
.get_all()
|
||||||
|
.iter()
|
||||||
|
.filter(|w| w.is_healthy())
|
||||||
|
.count();
|
||||||
|
|
||||||
|
if healthy_workers >= expected_count {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
if start.elapsed() > timeout_duration {
|
||||||
|
panic!(
|
||||||
|
"Timeout waiting for {} workers to become healthy (only {} ready)",
|
||||||
|
expected_count, healthy_workers
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let router = RouterFactory::create_router(&app_context).await.unwrap();
|
||||||
|
let router = Arc::from(router);
|
||||||
|
|
||||||
|
Self {
|
||||||
|
workers,
|
||||||
|
router,
|
||||||
|
config,
|
||||||
|
app_context,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn create_app(&self) -> axum::Router {
|
||||||
|
test_app::create_test_app_with_context(
|
||||||
|
Arc::clone(&self.router),
|
||||||
|
Arc::clone(&self.app_context),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn shutdown(mut self) {
|
||||||
|
for worker in &mut self.workers {
|
||||||
|
worker.stop().await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Helper function to create AppContext for tests
|
/// Helper function to create AppContext for tests
|
||||||
pub async fn create_test_context(config: RouterConfig) -> Arc<AppContext> {
|
pub async fn create_test_context(config: RouterConfig) -> Arc<AppContext> {
|
||||||
let client = reqwest::Client::new();
|
let client = reqwest::Client::new();
|
||||||
|
|||||||
@@ -1,107 +1,10 @@
|
|||||||
mod common;
|
mod common;
|
||||||
|
|
||||||
use std::sync::Arc;
|
use common::{
|
||||||
|
mock_worker::{HealthStatus, MockWorkerConfig, WorkerType},
|
||||||
use common::mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType};
|
WorkerTestContext,
|
||||||
use reqwest::Client;
|
|
||||||
use serde_json::json;
|
|
||||||
use smg::{
|
|
||||||
config::{RouterConfig, RoutingMode},
|
|
||||||
routers::{RouterFactory, RouterTrait},
|
|
||||||
};
|
};
|
||||||
|
use serde_json::json;
|
||||||
/// Test context that manages mock workers
|
|
||||||
struct TestContext {
|
|
||||||
workers: Vec<MockWorker>,
|
|
||||||
_router: Arc<dyn RouterTrait>,
|
|
||||||
worker_urls: Vec<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl TestContext {
|
|
||||||
async fn new(worker_configs: Vec<MockWorkerConfig>) -> Self {
|
|
||||||
let mut config = RouterConfig::builder()
|
|
||||||
.regular_mode(vec![])
|
|
||||||
.port(3003)
|
|
||||||
.worker_startup_timeout_secs(1)
|
|
||||||
.worker_startup_check_interval_secs(1)
|
|
||||||
.build_unchecked();
|
|
||||||
|
|
||||||
let mut workers = Vec::new();
|
|
||||||
let mut worker_urls = Vec::new();
|
|
||||||
|
|
||||||
for worker_config in worker_configs {
|
|
||||||
let mut worker = MockWorker::new(worker_config);
|
|
||||||
let url = worker.start().await.unwrap();
|
|
||||||
worker_urls.push(url);
|
|
||||||
workers.push(worker);
|
|
||||||
}
|
|
||||||
|
|
||||||
if !workers.is_empty() {
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(200)).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
config.mode = RoutingMode::Regular {
|
|
||||||
worker_urls: worker_urls.clone(),
|
|
||||||
};
|
|
||||||
|
|
||||||
let app_context = common::create_test_context(config.clone()).await;
|
|
||||||
|
|
||||||
let router = RouterFactory::create_router(&app_context).await.unwrap();
|
|
||||||
let router = Arc::from(router);
|
|
||||||
|
|
||||||
if !workers.is_empty() {
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(500)).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
Self {
|
|
||||||
workers,
|
|
||||||
_router: router,
|
|
||||||
worker_urls: worker_urls.clone(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn shutdown(mut self) {
|
|
||||||
// Small delay to ensure any pending operations complete
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
|
||||||
|
|
||||||
for worker in &mut self.workers {
|
|
||||||
worker.stop().await;
|
|
||||||
}
|
|
||||||
|
|
||||||
// Another small delay to ensure cleanup completes
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn make_request(
|
|
||||||
&self,
|
|
||||||
endpoint: &str,
|
|
||||||
body: serde_json::Value,
|
|
||||||
) -> Result<serde_json::Value, String> {
|
|
||||||
let client = Client::new();
|
|
||||||
|
|
||||||
// Use the first worker URL from the context
|
|
||||||
let worker_url = self
|
|
||||||
.worker_urls
|
|
||||||
.first()
|
|
||||||
.ok_or_else(|| "No workers available".to_string())?;
|
|
||||||
|
|
||||||
let response = client
|
|
||||||
.post(format!("{}{}", worker_url, endpoint))
|
|
||||||
.json(&body)
|
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Request failed: {}", e))?;
|
|
||||||
|
|
||||||
if !response.status().is_success() {
|
|
||||||
return Err(format!("Request failed with status: {}", response.status()));
|
|
||||||
}
|
|
||||||
|
|
||||||
response
|
|
||||||
.json::<serde_json::Value>()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Failed to parse response: {}", e))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod request_format_tests {
|
mod request_format_tests {
|
||||||
@@ -109,7 +12,7 @@ mod request_format_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_generate_request_formats() {
|
async fn test_generate_request_formats() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 19001,
|
port: 19001,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -156,7 +59,7 @@ mod request_format_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_chat_completions_formats() {
|
async fn test_v1_chat_completions_formats() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 19002,
|
port: 19002,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -204,7 +107,7 @@ mod request_format_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_completions_formats() {
|
async fn test_v1_completions_formats() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 19003,
|
port: 19003,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -256,7 +159,7 @@ mod request_format_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_batch_requests() {
|
async fn test_batch_requests() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 19004,
|
port: 19004,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -290,7 +193,7 @@ mod request_format_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_special_parameters() {
|
async fn test_special_parameters() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 19005,
|
port: 19005,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -338,7 +241,7 @@ mod request_format_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_error_handling() {
|
async fn test_error_handling() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 19006,
|
port: 19006,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
|
|||||||
@@ -1,130 +1,10 @@
|
|||||||
mod common;
|
mod common;
|
||||||
|
|
||||||
use std::sync::Arc;
|
use common::{
|
||||||
|
mock_worker::{HealthStatus, MockWorkerConfig, WorkerType},
|
||||||
use common::mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType};
|
WorkerTestContext,
|
||||||
use futures_util::StreamExt;
|
|
||||||
use reqwest::Client;
|
|
||||||
use serde_json::json;
|
|
||||||
use smg::{
|
|
||||||
config::{RouterConfig, RoutingMode},
|
|
||||||
routers::{RouterFactory, RouterTrait},
|
|
||||||
};
|
};
|
||||||
|
use serde_json::json;
|
||||||
/// Test context that manages mock workers
|
|
||||||
struct TestContext {
|
|
||||||
workers: Vec<MockWorker>,
|
|
||||||
_router: Arc<dyn RouterTrait>,
|
|
||||||
worker_urls: Vec<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl TestContext {
|
|
||||||
async fn new(worker_configs: Vec<MockWorkerConfig>) -> Self {
|
|
||||||
let mut config = RouterConfig::builder()
|
|
||||||
.regular_mode(vec![])
|
|
||||||
.port(3004)
|
|
||||||
.worker_startup_timeout_secs(1)
|
|
||||||
.worker_startup_check_interval_secs(1)
|
|
||||||
.build_unchecked();
|
|
||||||
|
|
||||||
let mut workers = Vec::new();
|
|
||||||
let mut worker_urls = Vec::new();
|
|
||||||
|
|
||||||
for worker_config in worker_configs {
|
|
||||||
let mut worker = MockWorker::new(worker_config);
|
|
||||||
let url = worker.start().await.unwrap();
|
|
||||||
worker_urls.push(url);
|
|
||||||
workers.push(worker);
|
|
||||||
}
|
|
||||||
|
|
||||||
if !workers.is_empty() {
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(200)).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
config.mode = RoutingMode::Regular {
|
|
||||||
worker_urls: worker_urls.clone(),
|
|
||||||
};
|
|
||||||
|
|
||||||
let app_context = common::create_test_context(config.clone()).await;
|
|
||||||
|
|
||||||
let router = RouterFactory::create_router(&app_context).await.unwrap();
|
|
||||||
let router = Arc::from(router);
|
|
||||||
|
|
||||||
if !workers.is_empty() {
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(500)).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
Self {
|
|
||||||
workers,
|
|
||||||
_router: router,
|
|
||||||
worker_urls: worker_urls.clone(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn shutdown(mut self) {
|
|
||||||
// Small delay to ensure any pending operations complete
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
|
||||||
|
|
||||||
for worker in &mut self.workers {
|
|
||||||
worker.stop().await;
|
|
||||||
}
|
|
||||||
|
|
||||||
// Another small delay to ensure cleanup completes
|
|
||||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn make_streaming_request(
|
|
||||||
&self,
|
|
||||||
endpoint: &str,
|
|
||||||
body: serde_json::Value,
|
|
||||||
) -> Result<Vec<String>, String> {
|
|
||||||
let client = Client::new();
|
|
||||||
|
|
||||||
// Use the first worker URL from the context
|
|
||||||
let worker_url = self
|
|
||||||
.worker_urls
|
|
||||||
.first()
|
|
||||||
.ok_or_else(|| "No workers available".to_string())?;
|
|
||||||
|
|
||||||
let response = client
|
|
||||||
.post(format!("{}{}", worker_url, endpoint))
|
|
||||||
.json(&body)
|
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Request failed: {}", e))?;
|
|
||||||
|
|
||||||
if !response.status().is_success() {
|
|
||||||
return Err(format!("Request failed with status: {}", response.status()));
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check if it's a streaming response
|
|
||||||
let content_type = response
|
|
||||||
.headers()
|
|
||||||
.get("content-type")
|
|
||||||
.and_then(|v| v.to_str().ok())
|
|
||||||
.unwrap_or("");
|
|
||||||
|
|
||||||
if !content_type.contains("text/event-stream") {
|
|
||||||
return Err("Response is not a stream".to_string());
|
|
||||||
}
|
|
||||||
|
|
||||||
let mut stream = response.bytes_stream();
|
|
||||||
let mut events = Vec::new();
|
|
||||||
|
|
||||||
while let Some(chunk) = stream.next().await {
|
|
||||||
if let Ok(bytes) = chunk {
|
|
||||||
let text = String::from_utf8_lossy(&bytes);
|
|
||||||
for line in text.lines() {
|
|
||||||
if let Some(stripped) = line.strip_prefix("data: ") {
|
|
||||||
events.push(stripped.to_string());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(events)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod streaming_tests {
|
mod streaming_tests {
|
||||||
@@ -132,7 +12,7 @@ mod streaming_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_generate_streaming() {
|
async fn test_generate_streaming() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 20001,
|
port: 20001,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -154,7 +34,6 @@ mod streaming_tests {
|
|||||||
assert!(result.is_ok());
|
assert!(result.is_ok());
|
||||||
|
|
||||||
let events = result.unwrap();
|
let events = result.unwrap();
|
||||||
// Should have at least one data chunk and [DONE]
|
|
||||||
assert!(events.len() >= 2);
|
assert!(events.len() >= 2);
|
||||||
assert_eq!(events.last().unwrap(), "[DONE]");
|
assert_eq!(events.last().unwrap(), "[DONE]");
|
||||||
|
|
||||||
@@ -163,7 +42,7 @@ mod streaming_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_chat_completions_streaming() {
|
async fn test_v1_chat_completions_streaming() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 20002,
|
port: 20002,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -187,7 +66,7 @@ mod streaming_tests {
|
|||||||
assert!(result.is_ok());
|
assert!(result.is_ok());
|
||||||
|
|
||||||
let events = result.unwrap();
|
let events = result.unwrap();
|
||||||
assert!(events.len() >= 2); // At least one chunk + [DONE]
|
assert!(events.len() >= 2);
|
||||||
|
|
||||||
for event in &events {
|
for event in &events {
|
||||||
if event != "[DONE]" {
|
if event != "[DONE]" {
|
||||||
@@ -207,7 +86,7 @@ mod streaming_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_v1_completions_streaming() {
|
async fn test_v1_completions_streaming() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 20003,
|
port: 20003,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -227,19 +106,19 @@ mod streaming_tests {
|
|||||||
assert!(result.is_ok());
|
assert!(result.is_ok());
|
||||||
|
|
||||||
let events = result.unwrap();
|
let events = result.unwrap();
|
||||||
assert!(events.len() >= 2); // At least one chunk + [DONE]
|
assert!(events.len() >= 2);
|
||||||
|
|
||||||
ctx.shutdown().await;
|
ctx.shutdown().await;
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_streaming_with_error() {
|
async fn test_streaming_with_error() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 20004,
|
port: 20004,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
response_delay_ms: 0,
|
response_delay_ms: 0,
|
||||||
fail_rate: 1.0, // Always fail
|
fail_rate: 1.0,
|
||||||
}])
|
}])
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
@@ -249,7 +128,6 @@ mod streaming_tests {
|
|||||||
});
|
});
|
||||||
|
|
||||||
let result = ctx.make_streaming_request("/generate", payload).await;
|
let result = ctx.make_streaming_request("/generate", payload).await;
|
||||||
// With fail_rate: 1.0, the request should fail
|
|
||||||
assert!(result.is_err());
|
assert!(result.is_err());
|
||||||
|
|
||||||
ctx.shutdown().await;
|
ctx.shutdown().await;
|
||||||
@@ -257,11 +135,11 @@ mod streaming_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_streaming_timeouts() {
|
async fn test_streaming_timeouts() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 20005,
|
port: 20005,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
response_delay_ms: 100, // Slow response
|
response_delay_ms: 100,
|
||||||
fail_rate: 0.0,
|
fail_rate: 0.0,
|
||||||
}])
|
}])
|
||||||
.await;
|
.await;
|
||||||
@@ -280,17 +158,15 @@ mod streaming_tests {
|
|||||||
|
|
||||||
assert!(result.is_ok());
|
assert!(result.is_ok());
|
||||||
let events = result.unwrap();
|
let events = result.unwrap();
|
||||||
|
|
||||||
// Should have received multiple chunks over time
|
|
||||||
assert!(!events.is_empty());
|
assert!(!events.is_empty());
|
||||||
assert!(elapsed.as_millis() >= 100); // At least one delay
|
assert!(elapsed.as_millis() >= 100);
|
||||||
|
|
||||||
ctx.shutdown().await;
|
ctx.shutdown().await;
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_batch_streaming() {
|
async fn test_batch_streaming() {
|
||||||
let ctx = TestContext::new(vec![MockWorkerConfig {
|
let ctx = WorkerTestContext::new(vec![MockWorkerConfig {
|
||||||
port: 20006,
|
port: 20006,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
health_status: HealthStatus::Healthy,
|
health_status: HealthStatus::Healthy,
|
||||||
@@ -299,7 +175,6 @@ mod streaming_tests {
|
|||||||
}])
|
}])
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
// Batch request with streaming
|
|
||||||
let payload = json!({
|
let payload = json!({
|
||||||
"text": ["First", "Second", "Third"],
|
"text": ["First", "Second", "Third"],
|
||||||
"stream": true,
|
"stream": true,
|
||||||
@@ -312,8 +187,7 @@ mod streaming_tests {
|
|||||||
assert!(result.is_ok());
|
assert!(result.is_ok());
|
||||||
|
|
||||||
let events = result.unwrap();
|
let events = result.unwrap();
|
||||||
// Should have multiple events for batch
|
assert!(events.len() >= 4);
|
||||||
assert!(events.len() >= 4); // At least 3 responses + [DONE]
|
|
||||||
|
|
||||||
ctx.shutdown().await;
|
ctx.shutdown().await;
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user