diff --git a/sgl-router/benches/request_processing.rs b/sgl-router/benches/request_processing.rs index 87380a796..54d704512 100644 --- a/sgl-router/benches/request_processing.rs +++ b/sgl-router/benches/request_processing.rs @@ -33,6 +33,7 @@ fn get_bootstrap_info(worker: &BasicWorker) -> (String, Option) { fn default_generate_request() -> GenerateRequest { GenerateRequest { text: None, + model: None, input_ids: None, input_embeds: None, image_data: None, diff --git a/sgl-router/src/protocols/generate.rs b/sgl-router/src/protocols/generate.rs index 4f3c1301a..d5819095a 100644 --- a/sgl-router/src/protocols/generate.rs +++ b/sgl-router/src/protocols/generate.rs @@ -21,6 +21,8 @@ pub struct GenerateRequest { #[serde(skip_serializing_if = "Option::is_none")] pub text: Option, + pub model: Option, + /// Input IDs for tokenized input #[serde(skip_serializing_if = "Option::is_none")] pub input_ids: Option, @@ -201,8 +203,12 @@ impl GenerationRequest for GenerateRequest { } fn get_model(&self) -> Option<&str> { - // Generate requests typically don't have a model field - None + // Generate requests have an optional model field + if let Some(s) = &self.model { + Some(s.as_str()) + } else { + None + } } fn extract_text_for_routing(&self) -> String { diff --git a/sgl-router/src/routers/router_manager.rs b/sgl-router/src/routers/router_manager.rs index 6535a1d30..a4b5554ef 100644 --- a/sgl-router/src/routers/router_manager.rs +++ b/sgl-router/src/routers/router_manager.rs @@ -350,12 +350,12 @@ impl RouterTrait for RouterManager { &self, headers: Option<&HeaderMap>, body: &GenerateRequest, - _model_id: Option<&str>, + model_id: Option<&str>, ) -> Response { - let router = self.select_router_for_request(headers, None); + let router = self.select_router_for_request(headers, model_id); if let Some(router) = router { - router.route_generate(headers, body, None).await + router.route_generate(headers, body, model_id).await } else { ( StatusCode::NOT_FOUND, @@ -369,12 +369,12 @@ impl RouterTrait for RouterManager { &self, headers: Option<&HeaderMap>, body: &ChatCompletionRequest, - _model_id: Option<&str>, + model_id: Option<&str>, ) -> Response { - let router = self.select_router_for_request(headers, Some(&body.model)); + let router = self.select_router_for_request(headers, model_id); if let Some(router) = router { - router.route_chat(headers, body, Some(&body.model)).await + router.route_chat(headers, body, model_id).await } else { ( StatusCode::NOT_FOUND, @@ -388,14 +388,12 @@ impl RouterTrait for RouterManager { &self, headers: Option<&HeaderMap>, body: &CompletionRequest, - _model_id: Option<&str>, + model_id: Option<&str>, ) -> Response { - let router = self.select_router_for_request(headers, Some(&body.model)); + let router = self.select_router_for_request(headers, model_id); if let Some(router) = router { - router - .route_completion(headers, body, Some(&body.model)) - .await + router.route_completion(headers, body, model_id).await } else { ( StatusCode::NOT_FOUND, @@ -487,14 +485,12 @@ impl RouterTrait for RouterManager { &self, headers: Option<&HeaderMap>, body: &EmbeddingRequest, - _model_id: Option<&str>, + model_id: Option<&str>, ) -> Response { - let router = self.select_router_for_request(headers, Some(&body.model)); + let router = self.select_router_for_request(headers, model_id); if let Some(router) = router { - router - .route_embeddings(headers, body, Some(&body.model)) - .await + router.route_embeddings(headers, body, model_id).await } else { ( StatusCode::NOT_FOUND, @@ -510,7 +506,7 @@ impl RouterTrait for RouterManager { body: &RerankRequest, model_id: Option<&str>, ) -> Response { - let router = self.select_router_for_request(headers, None); + let router = self.select_router_for_request(headers, model_id); if let Some(router) = router { router.route_rerank(headers, body, model_id).await @@ -529,7 +525,7 @@ impl RouterTrait for RouterManager { body: &ClassifyRequest, model_id: Option<&str>, ) -> Response { - let router = self.select_router_for_request(headers, Some(&body.model)); + let router = self.select_router_for_request(headers, model_id); if let Some(router) = router { router.route_classify(headers, body, model_id).await diff --git a/sgl-router/src/server.rs b/sgl-router/src/server.rs index 24032cc50..5db429788 100644 --- a/sgl-router/src/server.rs +++ b/sgl-router/src/server.rs @@ -136,9 +136,10 @@ async fn generate( headers: http::HeaderMap, Json(body): Json, ) -> Response { + let model_id = body.model.as_deref(); state .router - .route_generate(Some(&headers), &body, None) + .route_generate(Some(&headers), &body, model_id) .await } @@ -147,7 +148,10 @@ async fn v1_chat_completions( headers: http::HeaderMap, ValidatedJson(body): ValidatedJson, ) -> Response { - state.router.route_chat(Some(&headers), &body, None).await + state + .router + .route_chat(Some(&headers), &body, Some(&body.model)) + .await } async fn v1_completions( @@ -157,7 +161,7 @@ async fn v1_completions( ) -> Response { state .router - .route_completion(Some(&headers), &body, None) + .route_completion(Some(&headers), &body, Some(&body.model)) .await } @@ -166,7 +170,10 @@ async fn rerank( headers: http::HeaderMap, ValidatedJson(body): ValidatedJson, ) -> Response { - state.router.route_rerank(Some(&headers), &body, None).await + state + .router + .route_rerank(Some(&headers), &body, Some(&body.model)) + .await } async fn v1_rerank( @@ -174,9 +181,10 @@ async fn v1_rerank( headers: http::HeaderMap, Json(body): Json, ) -> Response { + let rerank_body = &body.into(); state .router - .route_rerank(Some(&headers), &body.into(), None) + .route_rerank(Some(&headers), rerank_body, Some(&rerank_body.model)) .await } @@ -187,7 +195,7 @@ async fn v1_responses( ) -> Response { state .router - .route_responses(Some(&headers), &body, None) + .route_responses(Some(&headers), &body, Some(&body.model)) .await } @@ -198,7 +206,7 @@ async fn v1_embeddings( ) -> Response { state .router - .route_embeddings(Some(&headers), &body, None) + .route_embeddings(Some(&headers), &body, Some(&body.model)) .await } @@ -209,7 +217,7 @@ async fn v1_classify( ) -> Response { state .router - .route_classify(Some(&headers), &body, None) + .route_classify(Some(&headers), &body, Some(&body.model)) .await } diff --git a/sgl-router/tests/test_openai_routing.rs b/sgl-router/tests/test_openai_routing.rs index 7bfdadc4a..8e7756011 100644 --- a/sgl-router/tests/test_openai_routing.rs +++ b/sgl-router/tests/test_openai_routing.rs @@ -602,6 +602,7 @@ async fn test_unsupported_endpoints() { let generate_request = GenerateRequest { text: Some("Hello world".to_string()), + model: None, input_ids: None, input_embeds: None, image_data: None,