[model-gateway] Add classify pipeline stages and protocol types (#16094)

Co-authored-by: Chang Su <chang.s.su@oracle.com>
This commit is contained in:
Simo Lin
2025-12-29 10:23:36 -08:00
committed by GitHub
parent 7380ec9d55
commit f44c4b3715
15 changed files with 632 additions and 56 deletions

View File

@@ -375,6 +375,7 @@ pub mod metrics_labels {
pub const ENDPOINT_COMPLETIONS: &str = "completions";
pub const ENDPOINT_RERANK: &str = "rerank";
pub const ENDPOINT_EMBEDDINGS: &str = "embeddings";
pub const ENDPOINT_CLASSIFY: &str = "classify";
// Worker types
pub const WORKER_REGULAR: &str = "regular";

View File

@@ -1,41 +1,52 @@
//! Classify API protocol definitions.
//!
//! This module defines the request and response types for the `/v1/classify` API,
//! which is compatible with vLLM's classification endpoint.
//!
//! Classification reuses the embedding backend - the scheduler returns logits as
//! "embeddings", and the classify layer applies softmax + label mapping.
use serde::{Deserialize, Serialize};
use serde_json::Value;
use super::common::GenerationRequest;
use super::common::{GenerationRequest, UsageInfo};
// ============================================================================
// Embedding API
// Classify API
// ============================================================================
/// Classification request - compatible with vLLM's /v1/classify API
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct ClassifyRequest {
/// ID of the model to use
pub model: String,
/// Input can be a string, array of strings, tokens, or batch inputs
/// Input can be a string, array of strings, or token IDs
/// - Single string: "text to classify"
/// - Array of strings: ["text1", "text2"]
/// - Token IDs: [1, 2, 3] (advanced usage)
pub input: Value,
/// Optional encoding format (e.g., "float", "base64")
#[serde(skip_serializing_if = "Option::is_none")]
pub encoding_format: Option<String>,
/// Optional user identifier
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
/// Optional number of dimensions for the embedding
#[serde(skip_serializing_if = "Option::is_none")]
pub dimensions: Option<u32>,
/// SGLang extension: request id for tracking
#[serde(skip_serializing_if = "Option::is_none")]
pub rid: Option<String>,
/// SGLang extension: request priority
#[serde(skip_serializing_if = "Option::is_none")]
pub priority: Option<i32>,
/// SGLang extension: enable/disable logging of metrics
#[serde(skip_serializing_if = "Option::is_none")]
pub log_metrics: Option<bool>,
}
impl GenerationRequest for ClassifyRequest {
fn is_stream(&self) -> bool {
// Embeddings are non-streaming
false
false // Classification is always non-streaming
}
fn get_model(&self) -> Option<&str> {
@@ -43,7 +54,6 @@ impl GenerationRequest for ClassifyRequest {
}
fn extract_text_for_routing(&self) -> String {
// Best effort: extract text content for routing decisions
match &self.input {
Value::String(s) => s.clone(),
Value::Array(arr) => arr
@@ -55,3 +65,57 @@ impl GenerationRequest for ClassifyRequest {
}
}
}
// ============================================================================
// Classify Response
// ============================================================================
/// Single classification result
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClassifyData {
/// Index of this result (for batch requests)
pub index: u32,
/// Predicted class label (from id2label mapping)
pub label: String,
/// Probability distribution over all classes (softmax of logits)
pub probs: Vec<f32>,
/// Number of classes
pub num_classes: u32,
}
/// Classification response
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClassifyResponse {
/// Unique request ID (format: "classify-{uuid}")
pub id: String,
/// Always "list"
pub object: String,
/// Unix timestamp (seconds since epoch)
pub created: u64,
/// Model name
pub model: String,
/// Classification results (one per input in batch)
pub data: Vec<ClassifyData>,
/// Token usage info
pub usage: UsageInfo,
}
impl ClassifyResponse {
/// Create a new ClassifyResponse with the given data
pub fn new(
id: String,
model: String,
created: u64,
data: Vec<ClassifyData>,
usage: UsageInfo,
) -> Self {
Self {
id,
object: "list".to_string(),
created,
model,
data,
usage,
}
}
}

View File

@@ -42,6 +42,7 @@ impl PipelineStage for DispatchMetadataStage {
}
RequestType::Responses(req) => req.model.clone(),
RequestType::Embedding(req) => req.model.clone(),
RequestType::Classify(req) => req.model.clone(),
};
let weight_version = ctx

View File

@@ -17,6 +17,7 @@ use crate::{
core::{attach_guards_to_response, Worker, WorkerLoadGuard},
protocols::{
chat::{ChatCompletionRequest, ChatCompletionResponse},
classify::{ClassifyRequest, ClassifyResponse},
embedding::{EmbeddingRequest, EmbeddingResponse},
generate::{GenerateRequest, GenerateResponse},
responses::ResponsesRequest,
@@ -51,6 +52,7 @@ pub enum RequestType {
Generate(Arc<GenerateRequest>),
Responses(Arc<ResponsesRequest>),
Embedding(Arc<EmbeddingRequest>),
Classify(Arc<ClassifyRequest>),
}
/// Shared components (injected once at creation)
@@ -320,6 +322,24 @@ impl RequestContext {
}
}
/// Create context for classify request
pub fn for_classify(
request: Arc<ClassifyRequest>,
headers: Option<HeaderMap>,
model_id: Option<String>,
components: Arc<SharedComponents>,
) -> Self {
Self {
input: RequestInput {
request_type: RequestType::Classify(request),
headers,
model_id,
},
components,
state: ProcessingState::default(),
}
}
/// Get reference to original request (type-safe)
pub fn request(&self) -> &RequestType {
&self.input.request_type
@@ -389,6 +409,22 @@ impl RequestContext {
}
}
/// Get classify request (panics if not classify)
pub fn classify_request(&self) -> &ClassifyRequest {
match &self.input.request_type {
RequestType::Classify(req) => req.as_ref(),
_ => panic!("Expected classify request"),
}
}
/// Get Arc clone of classify request (panics if not classify)
pub fn classify_request_arc(&self) -> Arc<ClassifyRequest> {
match &self.input.request_type {
RequestType::Classify(req) => Arc::clone(req),
_ => panic!("Expected classify request"),
}
}
/// Check if request is streaming
pub fn is_streaming(&self) -> bool {
match &self.input.request_type {
@@ -396,6 +432,7 @@ impl RequestContext {
RequestType::Generate(req) => req.stream,
RequestType::Responses(req) => req.stream.unwrap_or(false),
RequestType::Embedding(_) => false, // Embeddings are never streaming
RequestType::Classify(_) => false, // Classification is never streaming
}
}
@@ -548,4 +585,6 @@ pub enum FinalResponse {
Generate(Vec<GenerateResponse>),
/// Embedding response
Embedding(EmbeddingResponse),
/// Classification response
Classify(ClassifyResponse),
}

View File

@@ -90,6 +90,16 @@ impl PipelineStage for HarmonyRequestBuildingStage {
"Embedding requests are not supported with Harmony models".to_string(),
));
}
RequestType::Classify(_) => {
error!(
function = "HarmonyRequestBuildingStage::execute",
"Classify requests not supported for Harmony models"
);
return Err(error::bad_request(
"harmony_classify_not_supported",
"Classify requests are not supported with Harmony models".to_string(),
));
}
};
// Build gRPC request using token_ids directly (Harmony encoding already handled message rendering)

View File

@@ -135,14 +135,14 @@ impl PipelineStage for HarmonyResponseProcessingStage {
ctx.state.response.responses_iteration_result = Some(iteration_result);
Ok(None)
}
RequestType::Generate(_) | RequestType::Embedding(_) => {
RequestType::Generate(_) | RequestType::Embedding(_) | RequestType::Classify(_) => {
error!(
function = "HarmonyResponseProcessingStage::execute",
"Generate/Embedding request type not supported in Harmony pipeline"
"Generate/Embedding/Classify request type not supported in Harmony pipeline"
);
Err(error::internal_error(
"requests_not_supported_in_harmony",
"Generate/Embedding requests not supported in Harmony pipeline",
"Generate/Embedding/Classify requests not supported in Harmony pipeline",
))
}
}

View File

@@ -8,8 +8,8 @@ use std::{sync::Arc, time::Instant};
use axum::response::{IntoResponse, Response};
use tracing::{debug, error};
// Import embedding-specific stages
use super::regular::stages::embedding::preparation::EmbeddingPreparationStage;
// Import embedding-specific and classify-specific stages
use super::regular::stages::classify::ClassifyResponseProcessingStage;
use super::{
common::stages::*,
context::*,
@@ -18,6 +18,7 @@ use super::{
processor,
stages::{
embedding::{
preparation::EmbeddingPreparationStage,
request_building::EmbeddingRequestBuildingStage,
response_processing::EmbeddingResponseProcessingStage,
},
@@ -33,6 +34,7 @@ use crate::{
policies::PolicyRegistry,
protocols::{
chat::{ChatCompletionRequest, ChatCompletionResponse},
classify::ClassifyRequest,
embedding::EmbeddingRequest,
generate::GenerateRequest,
},
@@ -224,6 +226,34 @@ impl RequestPipeline {
}
}
/// Create a classify pipeline
///
/// Classify reuses embedding stages for preparation and request building,
/// but uses its own response processing for softmax + label mapping.
pub fn new_classify(
worker_registry: Arc<WorkerRegistry>,
policy_registry: Arc<PolicyRegistry>,
) -> Self {
let stages: Vec<Box<dyn PipelineStage>> = vec![
Box::new(EmbeddingPreparationStage::new()),
Box::new(WorkerSelectionStage::new(
worker_registry,
policy_registry,
WorkerSelectionMode::Regular, // Classify is always single worker
)),
Box::new(ClientAcquisitionStage),
Box::new(EmbeddingRequestBuildingStage::new()),
Box::new(DispatchMetadataStage),
Box::new(RequestExecutionStage::new(ExecutionMode::Single)),
Box::new(ClassifyResponseProcessingStage::new()),
];
Self {
stages: Arc::new(stages),
backend_type: metrics_labels::BACKEND_REGULAR,
}
}
/// Execute the complete pipeline for a chat request
pub async fn execute_chat(
&self,
@@ -295,10 +325,12 @@ impl RequestPipeline {
);
axum::Json(response).into_response()
}
Some(FinalResponse::Generate(_)) | Some(FinalResponse::Embedding(_)) => {
Some(FinalResponse::Generate(_))
| Some(FinalResponse::Embedding(_))
| Some(FinalResponse::Classify(_)) => {
error!(
function = "execute_chat",
"Wrong response type: expected Chat, got Generate/Embedding"
"Wrong response type: expected Chat, got Generate/Embedding/Classify"
);
Metrics::record_router_error(
metrics_labels::ROUTER_GRPC,
@@ -399,10 +431,12 @@ impl RequestPipeline {
);
axum::Json(response).into_response()
}
Some(FinalResponse::Chat(_)) | Some(FinalResponse::Embedding(_)) => {
Some(FinalResponse::Chat(_))
| Some(FinalResponse::Embedding(_))
| Some(FinalResponse::Classify(_)) => {
error!(
function = "execute_generate",
"Wrong response type: expected Generate, got Chat/Embedding"
"Wrong response type: expected Generate, got Chat/Embedding/Classify"
);
Metrics::record_router_error(
metrics_labels::ROUTER_GRPC,
@@ -538,6 +572,110 @@ impl RequestPipeline {
}
}
/// Execute the complete pipeline for a classify request
pub async fn execute_classify(
&self,
request: Arc<ClassifyRequest>,
headers: Option<http::HeaderMap>,
model_id: Option<String>,
components: Arc<SharedComponents>,
) -> Response {
debug!(
"execute_classify: Starting execution for model: {:?}",
model_id
);
let start = Instant::now();
// Record request start
Metrics::record_router_request(
metrics_labels::ROUTER_GRPC,
self.backend_type,
metrics_labels::CONNECTION_GRPC,
model_id.as_deref().unwrap_or("unknown"),
metrics_labels::ENDPOINT_CLASSIFY,
bool_to_static_str(false), // Classify is never streaming
);
let mut ctx = RequestContext::for_classify(request, headers, model_id.clone(), components);
for stage in self.stages.iter() {
debug!("execute_classify: Executing stage: {}", stage.name());
match stage.execute(&mut ctx).await {
Ok(Some(response)) => {
debug!(
"execute_classify: Stage {} returned final response.",
stage.name()
);
Metrics::record_router_duration(
metrics_labels::ROUTER_GRPC,
self.backend_type,
metrics_labels::CONNECTION_GRPC,
model_id.as_deref().unwrap_or("unknown"),
metrics_labels::ENDPOINT_CLASSIFY,
start.elapsed(),
);
return response;
}
Ok(None) => {
debug!(
"execute_classify: Stage {} completed, continuing to next stage.",
stage.name()
);
continue;
}
Err(response) => {
error!(
"execute_classify: Stage {} failed with status {:?}, returning error response.",
stage.name(),
response.status()
);
Metrics::record_router_error(
metrics_labels::ROUTER_GRPC,
self.backend_type,
metrics_labels::CONNECTION_GRPC,
model_id.as_deref().unwrap_or("unknown"),
metrics_labels::ENDPOINT_CLASSIFY,
error_type_from_status(response.status()),
);
return response;
}
}
}
debug!(
"execute_classify: Pipeline finished, processing final_response. Current state: {:?}",
ctx.state.response.final_response
);
match ctx.state.response.final_response {
Some(FinalResponse::Classify(_)) => {
error!("execute_classify: Classify FinalResponse found, but pipeline finished without returning response directly. This should be handled by the last stage.");
Metrics::record_router_duration(
metrics_labels::ROUTER_GRPC,
self.backend_type,
metrics_labels::CONNECTION_GRPC,
model_id.as_deref().unwrap_or("unknown"),
metrics_labels::ENDPOINT_CLASSIFY,
start.elapsed(),
);
error::internal_error(
"pipeline_fallthrough",
"Pipeline finished without returning response",
)
}
Some(_) => {
error!(function = "execute_classify", "Wrong response type");
error::internal_error("wrong_response_type", "Internal error: wrong response type")
}
None => {
error!(
function = "execute_classify",
"No final response produced by pipeline."
);
error::internal_error("no_response_produced", "No response produced")
}
}
}
/// Execute chat pipeline for responses endpoint
///
/// Used by ALL non-streaming /v1/responses requests.
@@ -584,10 +722,12 @@ impl RequestPipeline {
match ctx.state.response.final_response {
Some(FinalResponse::Chat(response)) => Ok(response),
Some(FinalResponse::Generate(_)) | Some(FinalResponse::Embedding(_)) => {
Some(FinalResponse::Generate(_))
| Some(FinalResponse::Embedding(_))
| Some(FinalResponse::Classify(_)) => {
error!(
function = "execute_chat_for_responses",
"Wrong response type: expected Chat, got Generate/Embedding"
"Wrong response type: expected Chat, got Generate/Embedding/Classify"
);
Err(error::internal_error(
"wrong_response_type",

View File

@@ -0,0 +1,9 @@
//! Pipeline stages for classify requests.
//!
//! Classify reuses embedding stages for preparation and request building,
//! as the scheduler treats classify as an embedding request and returns logits.
//! Only response processing is classify-specific (softmax + label mapping).
pub mod response_processing;
pub use response_processing::ClassifyResponseProcessingStage;

View File

@@ -0,0 +1,304 @@
//! Response processing stage for classify requests.
//!
//! Key responsibilities:
//! 1. Extract embedding (logits) from EmbedComplete response
//! 2. Apply softmax to convert logits to probabilities
//! 3. Find predicted class (argmax)
//! 4. Map class index to label (from id2label or generic LABEL_N)
//! 5. Build ClassifyResponse
use std::collections::HashMap;
use async_trait::async_trait;
use axum::{
response::{IntoResponse, Response},
Json,
};
use tracing::error;
use crate::{
protocols::{
classify::{ClassifyData, ClassifyResponse},
common::UsageInfo,
},
routers::{
error,
grpc::{
common::stages::PipelineStage,
context::{ExecutionResult, FinalResponse, RequestContext, WorkerSelection},
},
},
};
/// Response processing stage for classify requests.
///
/// Takes the logits from the embedding response and converts them to
/// classification results with probabilities and labels.
///
/// The stage is stateless - id2label mapping is obtained from the
/// selected worker's model card at runtime.
pub struct ClassifyResponseProcessingStage;
impl ClassifyResponseProcessingStage {
/// Create a new classify response processing stage.
pub fn new() -> Self {
Self
}
/// Apply softmax to logits to get probability distribution.
///
/// Uses the numerically stable formula: softmax(x)_i = exp(x_i - max(x)) / sum(exp(x - max(x)))
fn softmax(logits: &[f32]) -> Vec<f32> {
if logits.is_empty() {
return vec![];
}
// Find max for numerical stability
let max_logit = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
// Compute exp(x - max) for each element
let exp_vals: Vec<f32> = logits.iter().map(|&x| (x - max_logit).exp()).collect();
// Sum of exponentials
let sum: f32 = exp_vals.iter().sum();
// Normalize to get probabilities
if sum == 0.0 {
// Avoid division by zero - return uniform distribution
let n = exp_vals.len();
return vec![1.0 / n as f32; n];
}
exp_vals.iter().map(|&x| x / sum).collect()
}
/// Find the index of the maximum value (argmax).
fn argmax(probs: &[f32]) -> u32 {
probs
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(idx, _)| idx as u32)
.unwrap_or(0)
}
/// Get label for a class index.
///
/// Returns the label from id2label if available, otherwise returns generic "LABEL_N".
fn get_label(id2label: &HashMap<u32, String>, class_idx: u32) -> String {
id2label
.get(&class_idx)
.cloned()
.unwrap_or_else(|| format!("LABEL_{}", class_idx))
}
/// Extract id2label mapping from the selected worker's model card.
fn get_id2label_from_context(ctx: &RequestContext) -> HashMap<u32, String> {
// Get the selected worker
let worker = match ctx.state.workers.as_ref() {
Some(WorkerSelection::Single { worker }) => worker,
Some(WorkerSelection::Dual { prefill, .. }) => prefill, // Use prefill worker for model info
None => return HashMap::new(),
};
// Get id2label from the first model card
worker
.metadata()
.models
.first()
.map(|model| model.id2label.clone())
.unwrap_or_default()
}
}
impl Default for ClassifyResponseProcessingStage {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl PipelineStage for ClassifyResponseProcessingStage {
async fn execute(&self, ctx: &mut RequestContext) -> Result<Option<Response>, Response> {
// Extract execution result
let execution_result = ctx.state.response.execution_result.take().ok_or_else(|| {
error!(
function = "ClassifyResponseProcessingStage::execute",
"Execution result missing"
);
error::internal_error("execution_result_missing", "Execution result missing")
})?;
// Expect Embedding result variant (classify uses embed backend)
let proto_response = if let ExecutionResult::Embedding { response } = execution_result {
response
} else {
error!(
function = "ClassifyResponseProcessingStage::execute",
"Invalid execution result: expected Embedding"
);
return Err(error::internal_error(
"invalid_execution_result",
"Expected Embedding result for classify",
));
};
// Get logits from embedding response
let logits = proto_response.embedding();
if logits.is_empty() {
error!(
function = "ClassifyResponseProcessingStage::execute",
"Empty logits received from scheduler"
);
return Err(error::internal_error(
"empty_logits",
"Empty logits received from scheduler",
));
}
// Get id2label from the worker's model card
let id2label = Self::get_id2label_from_context(ctx);
// Apply softmax to get probabilities
let probs = Self::softmax(logits);
// Get predicted class (argmax)
let predicted_class = Self::argmax(&probs);
// Get label for predicted class
let label = Self::get_label(&id2label, predicted_class);
// Build classify data
let classify_data = ClassifyData {
index: 0,
label,
probs: probs.clone(),
num_classes: probs.len() as u32,
};
// Get dispatch metadata
let dispatch = ctx.state.dispatch.as_ref().ok_or_else(|| {
error!(
function = "ClassifyResponseProcessingStage::execute",
"Dispatch metadata missing"
);
error::internal_error("dispatch_missing", "Dispatch metadata missing")
})?;
// Build usage info
let prompt_tokens = proto_response.prompt_tokens().max(0) as u32;
let usage = UsageInfo {
prompt_tokens,
total_tokens: prompt_tokens,
completion_tokens: 0,
prompt_tokens_details: None,
reasoning_tokens: None,
};
// Build response
let response = ClassifyResponse::new(
format!("classify-{}", dispatch.request_id),
dispatch.model.clone(),
dispatch.created,
vec![classify_data],
usage,
);
// Store in context
ctx.state.response.final_response = Some(FinalResponse::Classify(response.clone()));
// Return HTTP response
Ok(Some(Json(response).into_response()))
}
fn name(&self) -> &'static str {
"ClassifyResponseProcessing"
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_softmax_basic() {
let logits = vec![1.0, 2.0, 3.0];
let probs = ClassifyResponseProcessingStage::softmax(&logits);
// Probabilities should sum to 1
let sum: f32 = probs.iter().sum();
assert!((sum - 1.0).abs() < 1e-6);
// Highest logit should have highest probability
assert!(probs[2] > probs[1]);
assert!(probs[1] > probs[0]);
}
#[test]
fn test_softmax_empty() {
let probs = ClassifyResponseProcessingStage::softmax(&[]);
assert!(probs.is_empty());
}
#[test]
fn test_softmax_single() {
let probs = ClassifyResponseProcessingStage::softmax(&[5.0]);
assert_eq!(probs.len(), 1);
assert!((probs[0] - 1.0).abs() < 1e-6);
}
#[test]
fn test_softmax_numerical_stability() {
// Large values that would overflow without max subtraction
let logits = vec![1000.0, 1001.0, 1002.0];
let probs = ClassifyResponseProcessingStage::softmax(&logits);
let sum: f32 = probs.iter().sum();
assert!((sum - 1.0).abs() < 1e-6);
assert!(probs[2] > probs[1]);
}
#[test]
fn test_argmax() {
assert_eq!(ClassifyResponseProcessingStage::argmax(&[0.1, 0.7, 0.2]), 1);
assert_eq!(
ClassifyResponseProcessingStage::argmax(&[0.9, 0.05, 0.05]),
0
);
assert_eq!(ClassifyResponseProcessingStage::argmax(&[0.1, 0.1, 0.8]), 2);
}
#[test]
fn test_get_label_with_mapping() {
let mut id2label = HashMap::new();
id2label.insert(0, "negative".to_string());
id2label.insert(1, "positive".to_string());
assert_eq!(
ClassifyResponseProcessingStage::get_label(&id2label, 0),
"negative"
);
assert_eq!(
ClassifyResponseProcessingStage::get_label(&id2label, 1),
"positive"
);
assert_eq!(
ClassifyResponseProcessingStage::get_label(&id2label, 2),
"LABEL_2"
); // Fallback for unknown
}
#[test]
fn test_get_label_without_mapping() {
let id2label = HashMap::new();
assert_eq!(
ClassifyResponseProcessingStage::get_label(&id2label, 0),
"LABEL_0"
);
assert_eq!(
ClassifyResponseProcessingStage::get_label(&id2label, 5),
"LABEL_5"
);
}
}

View File

@@ -33,22 +33,21 @@ impl Default for EmbeddingPreparationStage {
#[async_trait]
impl PipelineStage for EmbeddingPreparationStage {
async fn execute(&self, ctx: &mut RequestContext) -> Result<Option<Response>, Response> {
// Extract embedding request
let request = if let RequestType::Embedding(req) = &ctx.input.request_type {
req
} else {
error!(
function = "EmbeddingPreparationStage::execute",
"Invalid request type: expected Embedding"
);
return Err(error::internal_error(
"invalid_request_type",
"Expected Embedding request",
));
// Extract text from embedding or classify request (both use same preparation)
let text = match &ctx.input.request_type {
RequestType::Embedding(req) => req.extract_text_for_routing(),
RequestType::Classify(req) => req.extract_text_for_routing(),
_ => {
error!(
function = "EmbeddingPreparationStage::execute",
"Invalid request type: expected Embedding or Classify"
);
return Err(error::internal_error(
"invalid_request_type",
"Expected Embedding or Classify request",
));
}
};
// Extract text from request
let text = request.extract_text_for_routing();
if text.is_empty() {
return Err(error::bad_request(
"empty_input",

View File

@@ -31,19 +31,21 @@ impl Default for EmbeddingRequestBuildingStage {
#[async_trait]
impl PipelineStage for EmbeddingRequestBuildingStage {
async fn execute(&self, ctx: &mut RequestContext) -> Result<Option<Response>, Response> {
// Check if the request is of type Embedding
if let RequestType::Embedding(_) = &ctx.input.request_type {
// Proceed as expected
} else {
error!(
function = "EmbeddingRequestBuildingStage::execute",
"Invalid request type: expected Embedding"
);
return Err(error::internal_error(
"invalid_request_type",
"Expected Embedding request",
));
}
// Extract log_metrics from embedding or classify request (both use same backend)
let log_metrics = match &ctx.input.request_type {
RequestType::Embedding(req) => req.log_metrics,
RequestType::Classify(req) => req.log_metrics,
_ => {
error!(
function = "EmbeddingRequestBuildingStage::execute",
"Invalid request type: expected Embedding or Classify"
);
return Err(error::internal_error(
"invalid_request_type",
"Expected Embedding or Classify request",
));
}
};
// Preparation output should have tokenized input
let prep_output = ctx.state.preparation.as_ref().ok_or_else(|| {
@@ -82,13 +84,12 @@ impl PipelineStage for EmbeddingRequestBuildingStage {
// Use backend-specific builder to create ProtoEmbedRequest
// Currently only SGLang supports embedding via gRPC
let sglang_client = client.as_sglang();
let embedding_request = ctx.embedding_request();
let sglang_req = sglang_client.build_embed_request(
request_id.clone(),
original_text,
prep_output.token_ids.clone(),
embedding_request.log_metrics,
log_metrics,
);
let proto_req = ProtoEmbedRequest::Sglang(Box::new(sglang_req));

View File

@@ -3,6 +3,7 @@
//! This module defines stages specific to regular tokenizer-based models.
pub mod chat;
pub mod classify;
pub mod embedding;
pub mod generate;
mod preparation;
@@ -10,6 +11,7 @@ mod request_building;
mod response_processing;
pub use chat::{ChatPreparationStage, ChatRequestBuildingStage, ChatResponseProcessingStage};
pub use classify::ClassifyResponseProcessingStage;
pub use generate::{
GeneratePreparationStage, GenerateRequestBuildingStage, GenerateResponseProcessingStage,
};

View File

@@ -49,6 +49,8 @@ impl PipelineStage for PreparationStage {
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
RequestType::Embedding(_) => self.embedding_stage.execute(ctx).await,
// Classify reuses the embedding preparation (tokenization)
RequestType::Classify(_) => self.embedding_stage.execute(ctx).await,
RequestType::Responses(_) => {
error!(
function = "PreparationStage::execute",

View File

@@ -40,6 +40,7 @@ impl PipelineStage for RequestBuildingStage {
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
RequestType::Embedding(_) => self.embedding_stage.execute(ctx).await,
RequestType::Classify(_) => self.embedding_stage.execute(ctx).await,
RequestType::Responses(_request) => {
error!(
function = "RequestBuildingStage::execute",

View File

@@ -7,7 +7,7 @@ use axum::response::Response;
use tracing::error;
use super::{
chat::ChatResponseProcessingStage,
chat::ChatResponseProcessingStage, classify::ClassifyResponseProcessingStage,
embedding::response_processing::EmbeddingResponseProcessingStage,
generate::GenerateResponseProcessingStage,
};
@@ -25,6 +25,7 @@ pub struct ResponseProcessingStage {
chat_stage: ChatResponseProcessingStage,
generate_stage: GenerateResponseProcessingStage,
embedding_stage: EmbeddingResponseProcessingStage,
classify_stage: ClassifyResponseProcessingStage,
}
impl ResponseProcessingStage {
@@ -39,6 +40,7 @@ impl ResponseProcessingStage {
),
generate_stage: GenerateResponseProcessingStage::new(processor, streaming_processor),
embedding_stage: EmbeddingResponseProcessingStage::new(),
classify_stage: ClassifyResponseProcessingStage::new(),
}
}
}
@@ -50,6 +52,7 @@ impl PipelineStage for ResponseProcessingStage {
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
RequestType::Embedding(_) => self.embedding_stage.execute(ctx).await,
RequestType::Classify(_) => self.classify_stage.execute(ctx).await,
RequestType::Responses(_) => {
error!(
function = "ResponseProcessingStage::execute",