[model-gateway] Wire classify pipeline to gRPC router (#16098)
Co-authored-by: Chang Su <chang.s.su@oracle.com>
This commit is contained in:
@@ -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"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user