diff --git a/sgl-router/src/app_context.rs b/sgl-router/src/app_context.rs new file mode 100644 index 000000000..e8361ca37 --- /dev/null +++ b/sgl-router/src/app_context.rs @@ -0,0 +1,225 @@ +use std::sync::{Arc, OnceLock}; + +use reqwest::Client; + +use crate::{ + config::RouterConfig, + core::{workflow::WorkflowEngine, JobQueue, LoadMonitor, WorkerRegistry}, + data_connector::{ + SharedConversationItemStorage, SharedConversationStorage, SharedResponseStorage, + }, + middleware::TokenBucket, + policies::PolicyRegistry, + reasoning_parser::ParserFactory as ReasoningParserFactory, + routers::router_manager::RouterManager, + tokenizer::traits::Tokenizer, + tool_parser::ParserFactory as ToolParserFactory, +}; + +/// Error type for AppContext builder +#[derive(Debug)] +pub struct AppContextBuildError(&'static str); + +impl std::fmt::Display for AppContextBuildError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "Missing required field: {}", self.0) + } +} + +impl std::error::Error for AppContextBuildError {} + +#[derive(Clone)] +pub struct AppContext { + pub client: Client, + pub router_config: RouterConfig, + pub rate_limiter: Option>, + pub tokenizer: Option>, + pub reasoning_parser_factory: Option, + pub tool_parser_factory: Option, + pub worker_registry: Arc, + pub policy_registry: Arc, + pub router_manager: Option>, + pub response_storage: SharedResponseStorage, + pub conversation_storage: SharedConversationStorage, + pub conversation_item_storage: SharedConversationItemStorage, + pub load_monitor: Option>, + pub configured_reasoning_parser: Option, + pub configured_tool_parser: Option, + pub worker_job_queue: Arc>>, + pub workflow_engine: Arc>>, +} + +pub struct AppContextBuilder { + client: Option, + router_config: Option, + rate_limiter: Option>, + tokenizer: Option>, + reasoning_parser_factory: Option, + tool_parser_factory: Option, + worker_registry: Option>, + policy_registry: Option>, + router_manager: Option>, + response_storage: Option, + conversation_storage: Option, + conversation_item_storage: Option, + load_monitor: Option>, + worker_job_queue: Option>>>, + workflow_engine: Option>>>, +} + +impl AppContext { + pub fn builder() -> AppContextBuilder { + AppContextBuilder::new() + } +} + +impl AppContextBuilder { + pub fn new() -> Self { + Self { + client: None, + router_config: None, + rate_limiter: None, + tokenizer: None, + reasoning_parser_factory: None, + tool_parser_factory: None, + worker_registry: None, + policy_registry: None, + router_manager: None, + response_storage: None, + conversation_storage: None, + conversation_item_storage: None, + load_monitor: None, + worker_job_queue: None, + workflow_engine: None, + } + } + + pub fn client(mut self, client: Client) -> Self { + self.client = Some(client); + self + } + + pub fn router_config(mut self, router_config: RouterConfig) -> Self { + self.router_config = Some(router_config); + self + } + + pub fn rate_limiter(mut self, rate_limiter: Option>) -> Self { + self.rate_limiter = rate_limiter; + self + } + + pub fn tokenizer(mut self, tokenizer: Option>) -> Self { + self.tokenizer = tokenizer; + self + } + + pub fn reasoning_parser_factory( + mut self, + reasoning_parser_factory: Option, + ) -> Self { + self.reasoning_parser_factory = reasoning_parser_factory; + self + } + + pub fn tool_parser_factory(mut self, tool_parser_factory: Option) -> Self { + self.tool_parser_factory = tool_parser_factory; + self + } + + pub fn worker_registry(mut self, worker_registry: Arc) -> Self { + self.worker_registry = Some(worker_registry); + self + } + + pub fn policy_registry(mut self, policy_registry: Arc) -> Self { + self.policy_registry = Some(policy_registry); + self + } + + pub fn router_manager(mut self, router_manager: Option>) -> Self { + self.router_manager = router_manager; + self + } + + pub fn response_storage(mut self, response_storage: SharedResponseStorage) -> Self { + self.response_storage = Some(response_storage); + self + } + + pub fn conversation_storage(mut self, conversation_storage: SharedConversationStorage) -> Self { + self.conversation_storage = Some(conversation_storage); + self + } + + pub fn conversation_item_storage( + mut self, + conversation_item_storage: SharedConversationItemStorage, + ) -> Self { + self.conversation_item_storage = Some(conversation_item_storage); + self + } + + pub fn load_monitor(mut self, load_monitor: Option>) -> Self { + self.load_monitor = load_monitor; + self + } + + pub fn worker_job_queue(mut self, worker_job_queue: Arc>>) -> Self { + self.worker_job_queue = Some(worker_job_queue); + self + } + + pub fn workflow_engine(mut self, workflow_engine: Arc>>) -> Self { + self.workflow_engine = Some(workflow_engine); + self + } + + pub fn build(self) -> Result { + let router_config = self + .router_config + .ok_or(AppContextBuildError("router_config"))?; + let configured_reasoning_parser = router_config.reasoning_parser.clone(); + let configured_tool_parser = router_config.tool_call_parser.clone(); + + Ok(AppContext { + client: self.client.ok_or(AppContextBuildError("client"))?, + router_config, + rate_limiter: self.rate_limiter, + tokenizer: self.tokenizer, + reasoning_parser_factory: self.reasoning_parser_factory, + tool_parser_factory: self.tool_parser_factory, + worker_registry: self + .worker_registry + .ok_or(AppContextBuildError("worker_registry"))?, + policy_registry: self + .policy_registry + .ok_or(AppContextBuildError("policy_registry"))?, + router_manager: self.router_manager, + response_storage: self + .response_storage + .ok_or(AppContextBuildError("response_storage"))?, + conversation_storage: self + .conversation_storage + .ok_or(AppContextBuildError("conversation_storage"))?, + conversation_item_storage: self + .conversation_item_storage + .ok_or(AppContextBuildError("conversation_item_storage"))?, + load_monitor: self.load_monitor, + configured_reasoning_parser, + configured_tool_parser, + worker_job_queue: self + .worker_job_queue + .ok_or(AppContextBuildError("worker_job_queue"))?, + workflow_engine: self + .workflow_engine + .ok_or(AppContextBuildError("workflow_engine"))?, + }) + } +} + +impl Default for AppContextBuilder { + fn default() -> Self { + Self::new() + } +} diff --git a/sgl-router/src/core/job_queue.rs b/sgl-router/src/core/job_queue.rs index 8335420e9..ac396c7ec 100644 --- a/sgl-router/src/core/job_queue.rs +++ b/sgl-router/src/core/job_queue.rs @@ -14,6 +14,7 @@ use tokio::sync::mpsc; use tracing::{debug, error, info, warn}; use crate::{ + app_context::AppContext, config::{RouterConfig, RoutingMode}, core::workflow::{ steps::WorkerRemovalRequest, WorkflowContext, WorkflowEngine, WorkflowId, @@ -21,7 +22,6 @@ use crate::{ }, metrics::RouterMetrics, protocols::worker_spec::{JobStatus, WorkerConfigRequest}, - server::AppContext, }; /// Job types for control plane operations diff --git a/sgl-router/src/core/workflow/steps/worker_registration.rs b/sgl-router/src/core/workflow/steps/worker_registration.rs index 6461459a7..5afd27429 100644 --- a/sgl-router/src/core/workflow/steps/worker_registration.rs +++ b/sgl-router/src/core/workflow/steps/worker_registration.rs @@ -21,13 +21,13 @@ use serde_json::Value; use tracing::{debug, info, warn}; use crate::{ + app_context::AppContext, core::{ workflow::*, BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode, DPAwareWorkerBuilder, HealthConfig, Worker, WorkerType, }, grpc_client::SglangSchedulerClient, protocols::worker_spec::WorkerConfigRequest, - server::AppContext, }; // HTTP client for metadata fetching diff --git a/sgl-router/src/core/workflow/steps/worker_removal.rs b/sgl-router/src/core/workflow/steps/worker_removal.rs index a1cfc351b..a26ab85ed 100644 --- a/sgl-router/src/core/workflow/steps/worker_removal.rs +++ b/sgl-router/src/core/workflow/steps/worker_removal.rs @@ -15,8 +15,8 @@ use async_trait::async_trait; use tracing::{debug, info}; use crate::{ + app_context::AppContext, core::{workflow::*, Worker}, - server::AppContext, }; /// Request structure for worker removal diff --git a/sgl-router/src/lib.rs b/sgl-router/src/lib.rs index 9de44cc03..054386a3f 100644 --- a/sgl-router/src/lib.rs +++ b/sgl-router/src/lib.rs @@ -1,4 +1,5 @@ use pyo3::prelude::*; +pub mod app_context; pub mod config; pub mod logging; use std::collections::HashMap; diff --git a/sgl-router/src/routers/factory.rs b/sgl-router/src/routers/factory.rs index bfb0528f6..54c29eea9 100644 --- a/sgl-router/src/routers/factory.rs +++ b/sgl-router/src/routers/factory.rs @@ -9,10 +9,10 @@ use super::{ RouterTrait, }; use crate::{ + app_context::AppContext, config::{PolicyConfig, RoutingMode}, core::ConnectionMode, policies::PolicyFactory, - server::AppContext, }; /// Factory for creating router instances based on configuration diff --git a/sgl-router/src/routers/grpc/pd_router.rs b/sgl-router/src/routers/grpc/pd_router.rs index 7adeb0fb7..fe62b0dde 100644 --- a/sgl-router/src/routers/grpc/pd_router.rs +++ b/sgl-router/src/routers/grpc/pd_router.rs @@ -13,6 +13,7 @@ use tracing::debug; use super::{context::SharedComponents, pipeline::RequestPipeline}; use crate::{ + app_context::AppContext, config::types::RetryConfig, core::{ConnectionMode, WorkerRegistry, WorkerType}, policies::PolicyRegistry, @@ -27,7 +28,6 @@ use crate::{ }, reasoning_parser::ParserFactory as ReasoningParserFactory, routers::RouterTrait, - server::AppContext, tokenizer::traits::Tokenizer, tool_parser::ParserFactory as ToolParserFactory, }; diff --git a/sgl-router/src/routers/grpc/router.rs b/sgl-router/src/routers/grpc/router.rs index 84733bdfe..24e183d72 100644 --- a/sgl-router/src/routers/grpc/router.rs +++ b/sgl-router/src/routers/grpc/router.rs @@ -18,6 +18,7 @@ use super::{ responses::{self, BackgroundTaskInfo}, }; use crate::{ + app_context::AppContext, config::types::RetryConfig, core::WorkerRegistry, data_connector::{ @@ -35,7 +36,6 @@ use crate::{ }, reasoning_parser::ParserFactory as ReasoningParserFactory, routers::RouterTrait, - server::AppContext, tokenizer::traits::Tokenizer, tool_parser::ParserFactory as ToolParserFactory, }; diff --git a/sgl-router/src/routers/http/pd_router.rs b/sgl-router/src/routers/http/pd_router.rs index c6554c162..a3c16c478 100644 --- a/sgl-router/src/routers/http/pd_router.rs +++ b/sgl-router/src/routers/http/pd_router.rs @@ -127,7 +127,7 @@ impl PDRouter { } } - pub async fn new(ctx: &Arc) -> Result { + pub async fn new(ctx: &Arc) -> Result { Ok(PDRouter { worker_registry: Arc::clone(&ctx.worker_registry), policy_registry: Arc::clone(&ctx.policy_registry), diff --git a/sgl-router/src/routers/http/router.rs b/sgl-router/src/routers/http/router.rs index ca607be40..e9f4dba40 100644 --- a/sgl-router/src/routers/http/router.rs +++ b/sgl-router/src/routers/http/router.rs @@ -48,7 +48,7 @@ pub struct Router { impl Router { /// Create a new router with injected policy and client - pub async fn new(ctx: &Arc) -> Result { + pub async fn new(ctx: &Arc) -> Result { let workers = ctx.worker_registry.get_workers_filtered( None, // any model Some(WorkerType::Regular), diff --git a/sgl-router/src/routers/router_manager.rs b/sgl-router/src/routers/router_manager.rs index 4a12babb8..6535a1d30 100644 --- a/sgl-router/src/routers/router_manager.rs +++ b/sgl-router/src/routers/router_manager.rs @@ -18,6 +18,7 @@ use serde_json::Value; use tracing::{debug, info, warn}; use crate::{ + app_context::AppContext, config::RoutingMode, core::{ConnectionMode, WorkerRegistry, WorkerType}, protocols::{ @@ -30,7 +31,7 @@ use crate::{ responses::{ResponsesGetParams, ResponsesRequest}, }, routers::RouterTrait, - server::{AppContext, ServerConfig}, + server::ServerConfig, }; #[derive(Debug, Clone, Hash, Eq, PartialEq)] diff --git a/sgl-router/src/server.rs b/sgl-router/src/server.rs index 11ab32abf..2baf8896c 100644 --- a/sgl-router/src/server.rs +++ b/sgl-router/src/server.rs @@ -20,6 +20,7 @@ use tokio::{net::TcpListener, signal, spawn}; use tracing::{error, info, warn, Level}; use crate::{ + app_context::AppContext, config::{HistoryBackend, RouterConfig, RoutingMode}, core::{ worker_to_info, @@ -62,72 +63,6 @@ use crate::{ tool_parser::ParserFactory as ToolParserFactory, }; -// - -#[derive(Clone)] -pub struct AppContext { - pub client: Client, - pub router_config: RouterConfig, - pub rate_limiter: Option>, - pub tokenizer: Option>, - pub reasoning_parser_factory: Option, - pub tool_parser_factory: Option, - pub worker_registry: Arc, - pub policy_registry: Arc, - pub router_manager: Option>, - pub response_storage: SharedResponseStorage, - pub conversation_storage: SharedConversationStorage, - pub conversation_item_storage: SharedConversationItemStorage, - pub load_monitor: Option>, - pub configured_reasoning_parser: Option, - pub configured_tool_parser: Option, - pub worker_job_queue: Arc>>, - pub workflow_engine: Arc>>, -} - -impl AppContext { - #[allow(clippy::too_many_arguments)] - pub fn new( - router_config: RouterConfig, - client: Client, - rate_limiter: Option>, - tokenizer: Option>, - reasoning_parser_factory: Option, - tool_parser_factory: Option, - worker_registry: Arc, - policy_registry: Arc, - response_storage: SharedResponseStorage, - conversation_storage: SharedConversationStorage, - conversation_item_storage: SharedConversationItemStorage, - load_monitor: Option>, - worker_job_queue: Arc>>, - workflow_engine: Arc>>, - ) -> Self { - let configured_reasoning_parser = router_config.reasoning_parser.clone(); - let configured_tool_parser = router_config.tool_call_parser.clone(); - - Self { - client, - router_config, - rate_limiter, - tokenizer, - reasoning_parser_factory, - tool_parser_factory, - worker_registry, - policy_registry, - router_manager: None, - response_storage, - conversation_storage, - conversation_item_storage, - load_monitor, - configured_reasoning_parser, - configured_tool_parser, - worker_job_queue, - workflow_engine, - } - } -} - #[derive(Clone)] pub struct AppState { pub router: Arc, @@ -994,26 +929,27 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box Arc { let worker_job_queue = Arc::new(OnceLock::new()); let workflow_engine = Arc::new(OnceLock::new()); - let app_context = Arc::new(AppContext::new( - config, - client, - rate_limiter, - None, // tokenizer - None, // reasoning_parser_factory - None, // tool_parser_factory - worker_registry, - policy_registry, - response_storage, - conversation_storage, - conversation_item_storage, - load_monitor, - worker_job_queue, - workflow_engine, - )); + let app_context = Arc::new( + AppContext::builder() + .router_config(config) + .client(client) + .rate_limiter(rate_limiter) + .tokenizer(None) // tokenizer + .reasoning_parser_factory(None) // reasoning_parser_factory + .tool_parser_factory(None) // tool_parser_factory + .worker_registry(worker_registry) + .policy_registry(policy_registry) + .response_storage(response_storage) + .conversation_storage(conversation_storage) + .conversation_item_storage(conversation_item_storage) + .load_monitor(load_monitor) + .worker_job_queue(worker_job_queue) + .workflow_engine(workflow_engine) + .build() + .unwrap(), + ); // Initialize JobQueue after AppContext is created let weak_context = Arc::downgrade(&app_context); diff --git a/sgl-router/tests/common/test_app.rs b/sgl-router/tests/common/test_app.rs index c93bbeffb..41577b950 100644 --- a/sgl-router/tests/common/test_app.rs +++ b/sgl-router/tests/common/test_app.rs @@ -3,6 +3,7 @@ use std::sync::{Arc, OnceLock}; use axum::Router; use reqwest::Client; use sglang_router_rs::{ + app_context::AppContext, config::RouterConfig, core::{LoadMonitor, WorkerRegistry}, data_connector::{ @@ -11,7 +12,7 @@ use sglang_router_rs::{ middleware::{AuthConfig, TokenBucket}, policies::PolicyRegistry, routers::RouterTrait, - server::{build_app, AppContext, AppState}, + server::{build_app, AppState}, }; /// Create a test Axum application using the actual server's build_app function @@ -57,23 +58,26 @@ pub fn create_test_app( let worker_job_queue = Arc::new(OnceLock::new()); let workflow_engine = Arc::new(OnceLock::new()); - // Create AppContext - let app_context = Arc::new(AppContext::new( - router_config.clone(), - client, - rate_limiter, - None, // tokenizer - None, // reasoning_parser_factory - None, // tool_parser_factory - worker_registry, - policy_registry, - response_storage, - conversation_storage, - conversation_item_storage, - load_monitor, - worker_job_queue, - workflow_engine, - )); + // Create AppContext using builder pattern + let app_context = Arc::new( + AppContext::builder() + .router_config(router_config.clone()) + .client(client) + .rate_limiter(rate_limiter) + .tokenizer(None) // tokenizer + .reasoning_parser_factory(None) // reasoning_parser_factory + .tool_parser_factory(None) // tool_parser_factory + .worker_registry(worker_registry) + .policy_registry(policy_registry) + .response_storage(response_storage) + .conversation_storage(conversation_storage) + .conversation_item_storage(conversation_item_storage) + .load_monitor(load_monitor) + .worker_job_queue(worker_job_queue) + .workflow_engine(workflow_engine) + .build() + .unwrap(), + ); // Create AppState with the test router and context let app_state = Arc::new(AppState { diff --git a/sgl-router/tests/test_pd_routing.rs b/sgl-router/tests/test_pd_routing.rs index d3f3d73bf..a0c1dec0b 100644 --- a/sgl-router/tests/test_pd_routing.rs +++ b/sgl-router/tests/test_pd_routing.rs @@ -2,6 +2,7 @@ mod test_pd_routing { use serde_json::json; use sglang_router_rs::{ + app_context::AppContext, config::{PolicyConfig, RouterConfig, RoutingMode}, core::{BasicWorkerBuilder, Worker, WorkerType}, routers::{http::pd_types::PDSelectionPolicy, RouterFactory}, @@ -221,22 +222,25 @@ mod test_pd_routing { let worker_job_queue = Arc::new(OnceLock::new()); let workflow_engine = Arc::new(OnceLock::new()); - Arc::new(sglang_router_rs::server::AppContext::new( - config, - client, - rate_limiter, - None, // tokenizer - None, // reasoning_parser_factory - None, // tool_parser_factory - worker_registry, - policy_registry, - response_storage, - conversation_storage, - conversation_item_storage, - load_monitor, - worker_job_queue, - workflow_engine, - )) + Arc::new( + AppContext::builder() + .router_config(config) + .client(client) + .rate_limiter(rate_limiter) + .tokenizer(None) // tokenizer + .reasoning_parser_factory(None) // reasoning_parser_factory + .tool_parser_factory(None) // tool_parser_factory + .worker_registry(worker_registry) + .policy_registry(policy_registry) + .response_storage(response_storage) + .conversation_storage(conversation_storage) + .conversation_item_storage(conversation_item_storage) + .load_monitor(load_monitor) + .worker_job_queue(worker_job_queue) + .workflow_engine(workflow_engine) + .build() + .unwrap(), + ) }; let result = RouterFactory::create_router(&app_context).await; assert!(