From 877c8e3a962c244195d02c192d078d33983784f1 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Sun, 4 Jan 2026 01:44:10 +0800 Subject: [PATCH] Tiny refactor router test contexts (#16340) --- sgl-model-gateway/tests/api_endpoints_test.rs | 239 ++++------------- sgl-model-gateway/tests/common/mod.rs | 251 +++++++++++++++++- .../tests/request_formats_test.rs | 117 +------- sgl-model-gateway/tests/streaming_tests.rs | 158 ++--------- 4 files changed, 322 insertions(+), 443 deletions(-) diff --git a/sgl-model-gateway/tests/api_endpoints_test.rs b/sgl-model-gateway/tests/api_endpoints_test.rs index 9a6bd6fa2..a16663a87 100644 --- a/sgl-model-gateway/tests/api_endpoints_test.rs +++ b/sgl-model-gateway/tests/api_endpoints_test.rs @@ -1,172 +1,25 @@ mod common; -use std::sync::Arc; - use axum::{ body::Body, extract::Request, http::{header::CONTENT_TYPE, StatusCode}, }; -use common::mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType}; -use reqwest::Client; -use serde_json::json; -use smg::{ - app_context::AppContext, - config::{RouterConfig, RoutingMode}, - core::Job, - routers::{RouterFactory, RouterTrait}, +use common::{ + mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType}, + AppTestContext, }; +use serde_json::json; +use smg::{config::RouterConfig, routers::RouterFactory}; use tower::ServiceExt; -/// Test context that manages mock workers -struct TestContext { - workers: Vec, - router: Arc, - _client: Client, - _config: RouterConfig, - app_context: Arc, -} - -impl TestContext { - async fn new(worker_configs: Vec) -> 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, - ) -> 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)] mod health_tests { use super::*; #[tokio::test] async fn test_liveness_endpoint() { - let ctx = TestContext::new(vec![]).await; + let ctx = AppTestContext::new(vec![]).await; let app = ctx.create_app().await; let req = Request::builder() @@ -183,7 +36,7 @@ mod health_tests { #[tokio::test] async fn test_readiness_with_healthy_workers() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18001, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -208,7 +61,7 @@ mod health_tests { #[tokio::test] 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; @@ -226,7 +79,7 @@ mod health_tests { #[tokio::test] async fn test_health_endpoint_details() { - let ctx = TestContext::new(vec![ + let ctx = AppTestContext::new(vec![ MockWorkerConfig { port: 18003, worker_type: WorkerType::Regular, @@ -260,7 +113,7 @@ mod health_tests { #[tokio::test] async fn test_health_generate_endpoint() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18005, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -296,7 +149,7 @@ mod generation_tests { #[tokio::test] async fn test_generate_success() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18101, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -337,7 +190,7 @@ mod generation_tests { #[tokio::test] async fn test_generate_streaming() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18102, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -373,7 +226,7 @@ mod generation_tests { #[tokio::test] async fn test_generate_with_worker_failure() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18103, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -404,7 +257,7 @@ mod generation_tests { #[tokio::test] async fn test_v1_chat_completions_success() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18104, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -449,7 +302,7 @@ mod model_info_tests { #[tokio::test] async fn test_get_server_info() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18201, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -487,7 +340,7 @@ mod model_info_tests { #[tokio::test] async fn test_get_model_info() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18202, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -532,7 +385,7 @@ mod model_info_tests { #[tokio::test] async fn test_v1_models() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18203, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -588,7 +441,7 @@ mod model_info_tests { #[tokio::test] 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 req = Request::builder() @@ -644,7 +497,7 @@ mod model_info_tests { #[tokio::test] async fn test_model_info_with_multiple_workers() { - let ctx = TestContext::new(vec![ + let ctx = AppTestContext::new(vec![ MockWorkerConfig { port: 18204, worker_type: WorkerType::Regular, @@ -689,7 +542,7 @@ mod model_info_tests { #[tokio::test] async fn test_model_info_with_unhealthy_worker() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18206, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -725,7 +578,7 @@ mod router_policy_tests { #[tokio::test] async fn test_random_policy() { - let ctx = TestContext::new(vec![ + let ctx = AppTestContext::new(vec![ MockWorkerConfig { port: 18801, worker_type: WorkerType::Regular, @@ -768,7 +621,7 @@ mod router_policy_tests { #[tokio::test] async fn test_worker_selection() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18207, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -798,7 +651,7 @@ mod responses_endpoint_tests { #[tokio::test] async fn test_v1_responses_non_streaming() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18950, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -837,7 +690,7 @@ mod responses_endpoint_tests { #[tokio::test] async fn test_v1_responses_streaming() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18951, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -878,7 +731,7 @@ mod responses_endpoint_tests { #[tokio::test] async fn test_v1_responses_get() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18952, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -927,7 +780,7 @@ mod responses_endpoint_tests { #[tokio::test] async fn test_v1_responses_cancel() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18953, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -976,7 +829,7 @@ mod responses_endpoint_tests { #[tokio::test] async fn test_v1_responses_delete_not_implemented() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18954, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1019,7 +872,7 @@ mod responses_endpoint_tests { .queue_timeout_secs(60) .build_unchecked(); - let ctx = TestContext::new_with_config( + let ctx = AppTestContext::new_with_config( config, vec![], // No workers needed ) @@ -1073,7 +926,7 @@ mod responses_endpoint_tests { #[tokio::test] async fn test_v1_responses_get_multi_worker_fanout() { // Start two mock workers - let ctx = TestContext::new(vec![ + let ctx = AppTestContext::new(vec![ MockWorkerConfig { port: 18960, worker_type: WorkerType::Regular, @@ -1148,7 +1001,7 @@ mod error_tests { #[tokio::test] async fn test_404_not_found() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18401, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1185,7 +1038,7 @@ mod error_tests { #[tokio::test] async fn test_method_not_allowed() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18402, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1237,7 +1090,7 @@ mod error_tests { .queue_timeout_secs(60) .build_unchecked(); - let ctx = TestContext::new_with_config( + let ctx = AppTestContext::new_with_config( config, vec![MockWorkerConfig { port: 18403, @@ -1258,7 +1111,7 @@ mod error_tests { #[tokio::test] async fn test_invalid_json_payload() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18404, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1296,7 +1149,7 @@ mod error_tests { #[tokio::test] async fn test_invalid_model() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18406, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1334,7 +1187,7 @@ mod cache_tests { #[tokio::test] async fn test_flush_cache() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18501, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1371,7 +1224,7 @@ mod cache_tests { #[tokio::test] async fn test_get_loads() { - let ctx = TestContext::new(vec![ + let ctx = AppTestContext::new(vec![ MockWorkerConfig { port: 18502, worker_type: WorkerType::Regular, @@ -1414,7 +1267,7 @@ mod cache_tests { #[tokio::test] 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; @@ -1441,7 +1294,7 @@ mod load_balancing_tests { #[tokio::test] async fn test_request_distribution() { // Create multiple workers - let ctx = TestContext::new(vec![ + let ctx = AppTestContext::new(vec![ MockWorkerConfig { port: 18601, worker_type: WorkerType::Regular, @@ -1558,7 +1411,7 @@ mod request_id_tests { #[tokio::test] async fn test_request_id_generation() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18901, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1675,7 +1528,7 @@ mod request_id_tests { .queue_timeout_secs(60) .build_unchecked(); - let ctx = TestContext::new_with_config( + let ctx = AppTestContext::new_with_config( config, vec![MockWorkerConfig { port: 18902, @@ -1720,7 +1573,7 @@ mod rerank_tests { #[tokio::test] async fn test_rerank_success() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18105, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1772,7 +1625,7 @@ mod rerank_tests { #[tokio::test] async fn test_rerank_with_top_k() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18106, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1819,7 +1672,7 @@ mod rerank_tests { #[tokio::test] async fn test_rerank_without_documents() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18107, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1863,7 +1716,7 @@ mod rerank_tests { #[tokio::test] async fn test_rerank_worker_failure() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18108, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1896,7 +1749,7 @@ mod rerank_tests { #[tokio::test] async fn test_v1_rerank_compatibility() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18110, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -1953,7 +1806,7 @@ mod rerank_tests { #[tokio::test] async fn test_rerank_invalid_request() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = AppTestContext::new(vec![MockWorkerConfig { port: 18111, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, diff --git a/sgl-model-gateway/tests/common/mod.rs b/sgl-model-gateway/tests/common/mod.rs index 893b10243..48dbeedda 100644 --- a/sgl-model-gateway/tests/common/mod.rs +++ b/sgl-model-gateway/tests/common/mod.rs @@ -13,12 +13,14 @@ use std::{ sync::{Arc, Mutex, OnceLock}, }; +use mock_worker::{MockWorker, MockWorkerConfig}; use serde_json::json; use smg::{ app_context::AppContext, config::{RouterConfig, RoutingMode}, core::{ - BasicWorkerBuilder, LoadMonitor, ModelCard, RuntimeType, Worker, WorkerRegistry, WorkerType, + BasicWorkerBuilder, Job, LoadMonitor, ModelCard, RuntimeType, Worker, WorkerRegistry, + WorkerType, }, data_connector::{ MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage, @@ -27,10 +29,257 @@ use smg::{ policies::PolicyRegistry, protocols::common::{Function, Tool}, reasoning_parser::ParserFactory as ReasoningParserFactory, + routers::{RouterFactory, RouterTrait}, tokenizer::registry::TokenizerRegistry, tool_parser::ParserFactory as ToolParserFactory, }; +/// Test context for directly testing mock workers without full router setup. +pub struct WorkerTestContext { + pub workers: Vec, + pub worker_urls: Vec, +} + +impl WorkerTestContext { + pub async fn new(worker_configs: Vec) -> 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 { + 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::() + .await + .map_err(|e| format!("Failed to parse response: {}", e)) + } + + pub async fn make_streaming_request( + &self, + endpoint: &str, + body: serde_json::Value, + ) -> Result, 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, + pub router: Arc, + pub config: RouterConfig, + pub app_context: Arc, +} + +impl AppTestContext { + pub async fn new(worker_configs: Vec) -> 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, + ) -> 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 pub async fn create_test_context(config: RouterConfig) -> Arc { let client = reqwest::Client::new(); diff --git a/sgl-model-gateway/tests/request_formats_test.rs b/sgl-model-gateway/tests/request_formats_test.rs index e8201b291..2e83fe805 100644 --- a/sgl-model-gateway/tests/request_formats_test.rs +++ b/sgl-model-gateway/tests/request_formats_test.rs @@ -1,107 +1,10 @@ mod common; -use std::sync::Arc; - -use common::mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType}; -use reqwest::Client; -use serde_json::json; -use smg::{ - config::{RouterConfig, RoutingMode}, - routers::{RouterFactory, RouterTrait}, +use common::{ + mock_worker::{HealthStatus, MockWorkerConfig, WorkerType}, + WorkerTestContext, }; - -/// Test context that manages mock workers -struct TestContext { - workers: Vec, - _router: Arc, - worker_urls: Vec, -} - -impl TestContext { - async fn new(worker_configs: Vec) -> 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 { - 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::() - .await - .map_err(|e| format!("Failed to parse response: {}", e)) - } -} +use serde_json::json; #[cfg(test)] mod request_format_tests { @@ -109,7 +12,7 @@ mod request_format_tests { #[tokio::test] async fn test_generate_request_formats() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 19001, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -156,7 +59,7 @@ mod request_format_tests { #[tokio::test] async fn test_v1_chat_completions_formats() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 19002, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -204,7 +107,7 @@ mod request_format_tests { #[tokio::test] async fn test_v1_completions_formats() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 19003, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -256,7 +159,7 @@ mod request_format_tests { #[tokio::test] async fn test_batch_requests() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 19004, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -290,7 +193,7 @@ mod request_format_tests { #[tokio::test] async fn test_special_parameters() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 19005, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -338,7 +241,7 @@ mod request_format_tests { #[tokio::test] async fn test_error_handling() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 19006, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, diff --git a/sgl-model-gateway/tests/streaming_tests.rs b/sgl-model-gateway/tests/streaming_tests.rs index eaad35204..bcfd7099c 100644 --- a/sgl-model-gateway/tests/streaming_tests.rs +++ b/sgl-model-gateway/tests/streaming_tests.rs @@ -1,130 +1,10 @@ mod common; -use std::sync::Arc; - -use common::mock_worker::{HealthStatus, MockWorker, MockWorkerConfig, WorkerType}; -use futures_util::StreamExt; -use reqwest::Client; -use serde_json::json; -use smg::{ - config::{RouterConfig, RoutingMode}, - routers::{RouterFactory, RouterTrait}, +use common::{ + mock_worker::{HealthStatus, MockWorkerConfig, WorkerType}, + WorkerTestContext, }; - -/// Test context that manages mock workers -struct TestContext { - workers: Vec, - _router: Arc, - worker_urls: Vec, -} - -impl TestContext { - async fn new(worker_configs: Vec) -> 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, 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) - } -} +use serde_json::json; #[cfg(test)] mod streaming_tests { @@ -132,7 +12,7 @@ mod streaming_tests { #[tokio::test] async fn test_generate_streaming() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 20001, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -154,7 +34,6 @@ mod streaming_tests { assert!(result.is_ok()); let events = result.unwrap(); - // Should have at least one data chunk and [DONE] assert!(events.len() >= 2); assert_eq!(events.last().unwrap(), "[DONE]"); @@ -163,7 +42,7 @@ mod streaming_tests { #[tokio::test] async fn test_v1_chat_completions_streaming() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 20002, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -187,7 +66,7 @@ mod streaming_tests { assert!(result.is_ok()); let events = result.unwrap(); - assert!(events.len() >= 2); // At least one chunk + [DONE] + assert!(events.len() >= 2); for event in &events { if event != "[DONE]" { @@ -207,7 +86,7 @@ mod streaming_tests { #[tokio::test] async fn test_v1_completions_streaming() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 20003, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -227,19 +106,19 @@ mod streaming_tests { assert!(result.is_ok()); let events = result.unwrap(); - assert!(events.len() >= 2); // At least one chunk + [DONE] + assert!(events.len() >= 2); ctx.shutdown().await; } #[tokio::test] async fn test_streaming_with_error() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 20004, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, response_delay_ms: 0, - fail_rate: 1.0, // Always fail + fail_rate: 1.0, }]) .await; @@ -249,7 +128,6 @@ mod streaming_tests { }); let result = ctx.make_streaming_request("/generate", payload).await; - // With fail_rate: 1.0, the request should fail assert!(result.is_err()); ctx.shutdown().await; @@ -257,11 +135,11 @@ mod streaming_tests { #[tokio::test] async fn test_streaming_timeouts() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 20005, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, - response_delay_ms: 100, // Slow response + response_delay_ms: 100, fail_rate: 0.0, }]) .await; @@ -280,17 +158,15 @@ mod streaming_tests { assert!(result.is_ok()); let events = result.unwrap(); - - // Should have received multiple chunks over time assert!(!events.is_empty()); - assert!(elapsed.as_millis() >= 100); // At least one delay + assert!(elapsed.as_millis() >= 100); ctx.shutdown().await; } #[tokio::test] async fn test_batch_streaming() { - let ctx = TestContext::new(vec![MockWorkerConfig { + let ctx = WorkerTestContext::new(vec![MockWorkerConfig { port: 20006, worker_type: WorkerType::Regular, health_status: HealthStatus::Healthy, @@ -299,7 +175,6 @@ mod streaming_tests { }]) .await; - // Batch request with streaming let payload = json!({ "text": ["First", "Second", "Third"], "stream": true, @@ -312,8 +187,7 @@ mod streaming_tests { assert!(result.is_ok()); let events = result.unwrap(); - // Should have multiple events for batch - assert!(events.len() >= 4); // At least 3 responses + [DONE] + assert!(events.len() >= 4); ctx.shutdown().await; }