[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},
|
observability::metrics::{metrics_labels, Metrics},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::ChatCompletionRequest,
|
chat::ChatCompletionRequest,
|
||||||
|
classify::ClassifyRequest,
|
||||||
embedding::EmbeddingRequest,
|
embedding::EmbeddingRequest,
|
||||||
generate::GenerateRequest,
|
generate::GenerateRequest,
|
||||||
responses::{ResponsesGetParams, ResponsesRequest},
|
responses::{ResponsesGetParams, ResponsesRequest},
|
||||||
@@ -41,7 +42,8 @@ pub struct GrpcRouter {
|
|||||||
worker_registry: Arc<WorkerRegistry>,
|
worker_registry: Arc<WorkerRegistry>,
|
||||||
pipeline: RequestPipeline,
|
pipeline: RequestPipeline,
|
||||||
harmony_pipeline: RequestPipeline,
|
harmony_pipeline: RequestPipeline,
|
||||||
embedding_pipeline: RequestPipeline, // New field for embedding pipeline
|
embedding_pipeline: RequestPipeline,
|
||||||
|
classify_pipeline: RequestPipeline,
|
||||||
shared_components: Arc<SharedComponents>,
|
shared_components: Arc<SharedComponents>,
|
||||||
responses_context: responses::ResponsesContext,
|
responses_context: responses::ResponsesContext,
|
||||||
harmony_responses_context: responses::ResponsesContext,
|
harmony_responses_context: responses::ResponsesContext,
|
||||||
@@ -99,6 +101,10 @@ impl GrpcRouter {
|
|||||||
let embedding_pipeline =
|
let embedding_pipeline =
|
||||||
RequestPipeline::new_embeddings(worker_registry.clone(), _policy_registry.clone());
|
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
|
// Extract shared dependencies for responses contexts
|
||||||
let mcp_manager = ctx
|
let mcp_manager = ctx
|
||||||
.mcp_manager
|
.mcp_manager
|
||||||
@@ -128,6 +134,7 @@ impl GrpcRouter {
|
|||||||
pipeline,
|
pipeline,
|
||||||
harmony_pipeline,
|
harmony_pipeline,
|
||||||
embedding_pipeline,
|
embedding_pipeline,
|
||||||
|
classify_pipeline,
|
||||||
shared_components,
|
shared_components,
|
||||||
responses_context,
|
responses_context,
|
||||||
harmony_responses_context,
|
harmony_responses_context,
|
||||||
@@ -324,6 +331,25 @@ impl GrpcRouter {
|
|||||||
)
|
)
|
||||||
.await
|
.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 {
|
impl std::fmt::Debug for GrpcRouter {
|
||||||
@@ -390,6 +416,15 @@ impl RouterTrait for GrpcRouter {
|
|||||||
self.route_embeddings_impl(headers, body, model_id).await
|
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 {
|
fn router_type(&self) -> &'static str {
|
||||||
"grpc"
|
"grpc"
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user