From b2a3f0551af0a9f5c53ce8eaafa9ca4f5a0e43fc Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Mon, 29 Dec 2025 10:44:45 -0800 Subject: [PATCH] [model-gateway] Wire classify pipeline to gRPC router (#16098) Co-authored-by: Chang Su --- sgl-model-gateway/src/routers/grpc/router.rs | 37 +++++++++++++++++++- 1 file changed, 36 insertions(+), 1 deletion(-) diff --git a/sgl-model-gateway/src/routers/grpc/router.rs b/sgl-model-gateway/src/routers/grpc/router.rs index fc8831f98..50ed9cfd4 100644 --- a/sgl-model-gateway/src/routers/grpc/router.rs +++ b/sgl-model-gateway/src/routers/grpc/router.rs @@ -27,6 +27,7 @@ use crate::{ observability::metrics::{metrics_labels, Metrics}, protocols::{ chat::ChatCompletionRequest, + classify::ClassifyRequest, embedding::EmbeddingRequest, generate::GenerateRequest, responses::{ResponsesGetParams, ResponsesRequest}, @@ -41,7 +42,8 @@ pub struct GrpcRouter { worker_registry: Arc, pipeline: RequestPipeline, harmony_pipeline: RequestPipeline, - embedding_pipeline: RequestPipeline, // New field for embedding pipeline + embedding_pipeline: RequestPipeline, + classify_pipeline: RequestPipeline, shared_components: Arc, responses_context: responses::ResponsesContext, harmony_responses_context: responses::ResponsesContext, @@ -99,6 +101,10 @@ impl GrpcRouter { let embedding_pipeline = RequestPipeline::new_embeddings(worker_registry.clone(), _policy_registry.clone()); + // Create Classify pipeline + let classify_pipeline = + RequestPipeline::new_classify(worker_registry.clone(), _policy_registry.clone()); + // Extract shared dependencies for responses contexts let mcp_manager = ctx .mcp_manager @@ -128,6 +134,7 @@ impl GrpcRouter { pipeline, harmony_pipeline, embedding_pipeline, + classify_pipeline, shared_components, responses_context, harmony_responses_context, @@ -324,6 +331,25 @@ impl GrpcRouter { ) .await } + + /// Main route_classify implementation + async fn route_classify_impl( + &self, + headers: Option<&HeaderMap>, + body: &ClassifyRequest, + model_id: Option<&str>, + ) -> Response { + debug!("Processing classify request for model: {:?}", model_id); + + self.classify_pipeline + .execute_classify( + Arc::new(body.clone()), + headers.cloned(), + model_id.map(|s| s.to_string()), + self.shared_components.clone(), + ) + .await + } } impl std::fmt::Debug for GrpcRouter { @@ -390,6 +416,15 @@ impl RouterTrait for GrpcRouter { self.route_embeddings_impl(headers, body, model_id).await } + async fn route_classify( + &self, + headers: Option<&HeaderMap>, + body: &ClassifyRequest, + model_id: Option<&str>, + ) -> Response { + self.route_classify_impl(headers, body, model_id).await + } + fn router_type(&self) -> &'static str { "grpc" }