[model-gateway] Wire classify pipeline to gRPC router (#16098)

Co-authored-by: Chang Su <chang.s.su@oracle.com>
This commit is contained in:
Simo Lin
2025-12-29 10:44:45 -08:00
committed by GitHub
parent 684e148e38
commit b2a3f0551a

View File

@@ -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<WorkerRegistry>,
pipeline: RequestPipeline,
harmony_pipeline: RequestPipeline,
embedding_pipeline: RequestPipeline, // New field for embedding pipeline
embedding_pipeline: RequestPipeline,
classify_pipeline: RequestPipeline,
shared_components: Arc<SharedComponents>,
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"
}