[model-gateway]: add gRPC router embeddings endpoint implementation (#15273)
This commit is contained in:
@@ -455,10 +455,22 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer)
|
|||||||
input_text = grpc_req.tokenized.original_text
|
input_text = grpc_req.tokenized.original_text
|
||||||
input_ids = list(grpc_req.tokenized.input_ids)
|
input_ids = list(grpc_req.tokenized.input_ids)
|
||||||
|
|
||||||
|
# Convert sampling params
|
||||||
|
sampling_params = self._convert_sampling_params(grpc_req.sampling_params)
|
||||||
|
|
||||||
|
# For embedding requests, max_new_tokens should be 0.
|
||||||
|
# The scheduler logic expects an integer, not None.
|
||||||
|
sampling_params.max_new_tokens = 0
|
||||||
|
|
||||||
|
sampling_params.normalize(tokenizer=None)
|
||||||
|
|
||||||
return TokenizedEmbeddingReqInput(
|
return TokenizedEmbeddingReqInput(
|
||||||
rid=grpc_req.request_id,
|
rid=grpc_req.request_id,
|
||||||
input_text=input_text,
|
input_text=input_text,
|
||||||
input_ids=input_ids,
|
input_ids=input_ids,
|
||||||
|
image_inputs={"mm_items": []},
|
||||||
|
token_type_ids=list(grpc_req.token_type_ids),
|
||||||
|
sampling_params=sampling_params,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _convert_sampling_params(
|
def _convert_sampling_params(
|
||||||
|
|||||||
@@ -688,7 +688,9 @@ class GrpcRequestManager:
|
|||||||
batch_out.prompt_tokens[i] if batch_out.prompt_tokens else 0
|
batch_out.prompt_tokens[i] if batch_out.prompt_tokens else 0
|
||||||
),
|
),
|
||||||
"finish_reason": (
|
"finish_reason": (
|
||||||
batch_out.finish_reason[i] if batch_out.finish_reason else None
|
batch_out.finished_reasons[i]
|
||||||
|
if batch_out.finished_reasons
|
||||||
|
else None
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1710,9 +1710,14 @@ class Scheduler(
|
|||||||
if recv_req.image_inputs is not None:
|
if recv_req.image_inputs is not None:
|
||||||
image_inputs = self._get_multimodal_inputs(recv_req.image_inputs)
|
image_inputs = self._get_multimodal_inputs(recv_req.image_inputs)
|
||||||
# Expand a single image token into multiple dummy tokens for receiving image embeddings
|
# Expand a single image token into multiple dummy tokens for receiving image embeddings
|
||||||
req.origin_input_ids = self.pad_input_ids_func(
|
# The `pad_input_ids_func` is model-specific and may be None for
|
||||||
req.origin_input_ids, image_inputs
|
# embedding models or models not requiring special padding.
|
||||||
)
|
# If None, `req.origin_input_ids` is expected to be correctly populated already.
|
||||||
|
if self.pad_input_ids_func:
|
||||||
|
req.origin_input_ids = self.pad_input_ids_func(
|
||||||
|
req.origin_input_ids, image_inputs
|
||||||
|
)
|
||||||
|
|
||||||
req.extend_image_inputs(image_inputs)
|
req.extend_image_inputs(image_inputs)
|
||||||
|
|
||||||
if len(req.origin_input_ids) >= self.max_req_input_len:
|
if len(req.origin_input_ids) >= self.max_req_input_len:
|
||||||
|
|||||||
@@ -0,0 +1,110 @@
|
|||||||
|
"""
|
||||||
|
gRPC Router E2E Test - Embedding Server
|
||||||
|
|
||||||
|
Test the embedding functionality of the gRPC router.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import sys
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import openai
|
||||||
|
|
||||||
|
_TEST_DIR = Path(__file__).parent
|
||||||
|
sys.path.insert(0, str(_TEST_DIR.parent))
|
||||||
|
from fixtures import popen_launch_workers_and_router
|
||||||
|
from util import (
|
||||||
|
DEFAULT_EMBEDDING_MODEL_PATH,
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
kill_process_tree,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestEmbeddingServer(CustomTestCase):
|
||||||
|
"""
|
||||||
|
Test Embedding API through gRPC router.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = DEFAULT_EMBEDDING_MODEL_PATH
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
cls.api_key = "sk-123456"
|
||||||
|
|
||||||
|
# Launch workers with --is-embedding flag
|
||||||
|
cls.cluster = popen_launch_workers_and_router(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
num_workers=1,
|
||||||
|
tp_size=1,
|
||||||
|
policy="round_robin",
|
||||||
|
api_key=cls.api_key,
|
||||||
|
worker_args=["--is-embedding"],
|
||||||
|
)
|
||||||
|
|
||||||
|
cls.base_url += "/v1"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
# Cleanup router and workers
|
||||||
|
kill_process_tree(cls.cluster["router"].pid)
|
||||||
|
for worker in cls.cluster.get("workers", []):
|
||||||
|
kill_process_tree(worker.pid)
|
||||||
|
|
||||||
|
def test_embedding(self):
|
||||||
|
client = openai.Client(api_key=self.api_key, base_url=self.base_url)
|
||||||
|
|
||||||
|
input_text = "Hello world"
|
||||||
|
response = client.embeddings.create(
|
||||||
|
model=self.model,
|
||||||
|
input=input_text,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.object == "list"
|
||||||
|
assert len(response.data) == 1
|
||||||
|
embedding = response.data[0]
|
||||||
|
assert embedding.object == "embedding"
|
||||||
|
assert embedding.index == 0
|
||||||
|
assert len(embedding.embedding) > 0
|
||||||
|
assert isinstance(embedding.embedding[0], float)
|
||||||
|
|
||||||
|
# Verify usage statistics
|
||||||
|
assert response.usage.prompt_tokens > 0
|
||||||
|
assert response.usage.total_tokens == response.usage.prompt_tokens
|
||||||
|
|
||||||
|
def test_embedding_batch(self):
|
||||||
|
client = openai.Client(api_key=self.api_key, base_url=self.base_url)
|
||||||
|
|
||||||
|
input_texts = ["Hello world", "SGLang is fast"]
|
||||||
|
response = client.embeddings.create(
|
||||||
|
model=self.model,
|
||||||
|
input=input_texts,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(response.data) == 1
|
||||||
|
assert response.data[0].index == 0
|
||||||
|
assert len(response.data[0].embedding) > 0
|
||||||
|
|
||||||
|
def test_embedding_dimensions(self):
|
||||||
|
client = openai.Client(api_key=self.api_key, base_url=self.base_url)
|
||||||
|
|
||||||
|
response1 = client.embeddings.create(
|
||||||
|
model=self.model,
|
||||||
|
input="A short text",
|
||||||
|
)
|
||||||
|
dim1 = len(response1.data[0].embedding)
|
||||||
|
|
||||||
|
response2 = client.embeddings.create(
|
||||||
|
model=self.model,
|
||||||
|
input="A much longer text to ensure dimensions match",
|
||||||
|
)
|
||||||
|
dim2 = len(response2.data[0].embedding)
|
||||||
|
|
||||||
|
assert dim1 == dim2
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -92,6 +92,9 @@ DEFAULT_MISTRAL_FUNCTION_CALLING_MODEL_PATH = _get_model_path(
|
|||||||
# GPT-OSS models
|
# GPT-OSS models
|
||||||
DEFAULT_GPT_OSS_MODEL_PATH = _get_model_path("openai/gpt-oss-20b")
|
DEFAULT_GPT_OSS_MODEL_PATH = _get_model_path("openai/gpt-oss-20b")
|
||||||
|
|
||||||
|
# Embedding models
|
||||||
|
DEFAULT_EMBEDDING_MODEL_PATH = _get_model_path("intfloat/e5-mistral-7b-instruct")
|
||||||
|
|
||||||
|
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
# Process Management
|
# Process Management
|
||||||
|
|||||||
@@ -180,6 +180,21 @@ impl SglangSchedulerClient {
|
|||||||
))
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Submit an embedding request
|
||||||
|
pub async fn embed(
|
||||||
|
&self,
|
||||||
|
req: proto::EmbedRequest,
|
||||||
|
) -> Result<proto::EmbedResponse, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
let mut client = self.client.clone();
|
||||||
|
let mut request = Request::new(req);
|
||||||
|
|
||||||
|
// Inject W3C trace context into gRPC metadata
|
||||||
|
inject_trace_context_grpc(request.metadata_mut());
|
||||||
|
|
||||||
|
let response = client.embed(request).await?;
|
||||||
|
Ok(response.into_inner())
|
||||||
|
}
|
||||||
|
|
||||||
/// Perform health check
|
/// Perform health check
|
||||||
pub async fn health_check(
|
pub async fn health_check(
|
||||||
&self,
|
&self,
|
||||||
@@ -246,6 +261,25 @@ impl SglangSchedulerClient {
|
|||||||
Ok(response.into_inner())
|
Ok(response.into_inner())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Build a single SGLang EmbedRequest
|
||||||
|
pub fn build_embed_request(
|
||||||
|
&self,
|
||||||
|
request_id: String,
|
||||||
|
original_text: Option<String>,
|
||||||
|
token_ids: Vec<u32>,
|
||||||
|
log_metrics: Option<bool>,
|
||||||
|
) -> proto::EmbedRequest {
|
||||||
|
proto::EmbedRequest {
|
||||||
|
request_id,
|
||||||
|
tokenized: Some(proto::TokenizedInput {
|
||||||
|
original_text: original_text.unwrap_or_default(),
|
||||||
|
input_ids: token_ids,
|
||||||
|
}),
|
||||||
|
log_metrics: log_metrics.unwrap_or(false), // Default to false if not specified
|
||||||
|
..Default::default()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Build a single SGLang GenerateRequest from OpenAI ChatCompletionRequest
|
/// Build a single SGLang GenerateRequest from OpenAI ChatCompletionRequest
|
||||||
pub fn build_generate_request_from_chat(
|
pub fn build_generate_request_from_chat(
|
||||||
&self,
|
&self,
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
use super::common::GenerationRequest;
|
use super::common::{GenerationRequest, UsageInfo};
|
||||||
|
|
||||||
// ============================================================================
|
// ============================================================================
|
||||||
// Embedding API
|
// Embedding API
|
||||||
@@ -30,6 +30,10 @@ pub struct EmbeddingRequest {
|
|||||||
/// SGLang extension: request id for tracking
|
/// SGLang extension: request id for tracking
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub rid: Option<String>,
|
pub rid: Option<String>,
|
||||||
|
|
||||||
|
/// SGLang extension: enable/disable logging of metrics for this request
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
pub log_metrics: Option<bool>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl GenerationRequest for EmbeddingRequest {
|
impl GenerationRequest for EmbeddingRequest {
|
||||||
@@ -55,3 +59,18 @@ impl GenerationRequest for EmbeddingRequest {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
|
pub struct EmbeddingObject {
|
||||||
|
pub object: String, // "embedding"
|
||||||
|
pub embedding: Vec<f32>,
|
||||||
|
pub index: u32,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
|
pub struct EmbeddingResponse {
|
||||||
|
pub object: String, // "list"
|
||||||
|
pub data: Vec<EmbeddingObject>,
|
||||||
|
pub model: String,
|
||||||
|
pub usage: UsageInfo,
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,7 +2,9 @@
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
grpc_client::{SglangSchedulerClient, VllmEngineClient},
|
grpc_client::{SglangSchedulerClient, VllmEngineClient},
|
||||||
routers::grpc::proto_wrapper::{ProtoGenerateRequest, ProtoStream},
|
routers::grpc::proto_wrapper::{
|
||||||
|
ProtoEmbedRequest, ProtoEmbedResponse, ProtoGenerateRequest, ProtoStream,
|
||||||
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
/// Health check response (common across backends)
|
/// Health check response (common across backends)
|
||||||
@@ -131,6 +133,20 @@ impl GrpcClient {
|
|||||||
_ => panic!("Mismatched client and request types"),
|
_ => panic!("Mismatched client and request types"),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Submit an embedding request
|
||||||
|
pub async fn embed(
|
||||||
|
&mut self,
|
||||||
|
req: ProtoEmbedRequest,
|
||||||
|
) -> Result<ProtoEmbedResponse, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
match (self, req) {
|
||||||
|
(Self::Sglang(client), ProtoEmbedRequest::Sglang(boxed_req)) => {
|
||||||
|
let resp = client.embed(*boxed_req).await?;
|
||||||
|
Ok(ProtoEmbedResponse::Sglang(resp))
|
||||||
|
}
|
||||||
|
_ => panic!("Mismatched client and request types or unsupported embedding backend"),
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Unified ModelInfo wrapper
|
/// Unified ModelInfo wrapper
|
||||||
|
|||||||
@@ -55,6 +55,13 @@ pub async fn collect_responses(
|
|||||||
|
|
||||||
decode_responses
|
decode_responses
|
||||||
}
|
}
|
||||||
|
ExecutionResult::Embedding { .. } => {
|
||||||
|
// Embeddings do not support this path (no generate complete response)
|
||||||
|
return Err(error::internal_error(
|
||||||
|
"invalid_execution_mode",
|
||||||
|
"Embedding result encountered in response collection",
|
||||||
|
));
|
||||||
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
if all_responses.is_empty() {
|
if all_responses.is_empty() {
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ impl PipelineStage for DispatchMetadataStage {
|
|||||||
.unwrap_or_else(|| "unknown".to_string())
|
.unwrap_or_else(|| "unknown".to_string())
|
||||||
}
|
}
|
||||||
RequestType::Responses(req) => req.model.clone(),
|
RequestType::Responses(req) => req.model.clone(),
|
||||||
|
RequestType::Embedding(req) => req.model.clone(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let weight_version = ctx
|
let weight_version = ctx
|
||||||
|
|||||||
@@ -9,7 +9,10 @@ use crate::routers::{
|
|||||||
error,
|
error,
|
||||||
grpc::{
|
grpc::{
|
||||||
context::{ClientSelection, ExecutionResult, LoadGuards, RequestContext, WorkerSelection},
|
context::{ClientSelection, ExecutionResult, LoadGuards, RequestContext, WorkerSelection},
|
||||||
proto_wrapper::{ProtoGenerateRequest, ProtoStream},
|
proto_wrapper::{
|
||||||
|
ProtoEmbedRequest, ProtoEmbedResponseVariant, ProtoGenerateRequest, ProtoRequest,
|
||||||
|
ProtoStream,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -95,12 +98,14 @@ impl PipelineStage for RequestExecutionStage {
|
|||||||
);
|
);
|
||||||
|
|
||||||
let result = async {
|
let result = async {
|
||||||
match self.mode {
|
match proto_request {
|
||||||
ExecutionMode::Single => self.execute_single(proto_request, clients, workers).await,
|
ProtoRequest::Generate(req) => match self.mode {
|
||||||
ExecutionMode::DualDispatch => {
|
ExecutionMode::Single => self.execute_single(req, clients, workers).await,
|
||||||
self.execute_dual_dispatch(proto_request, clients, workers)
|
ExecutionMode::DualDispatch => {
|
||||||
.await
|
self.execute_dual_dispatch(req, clients, workers).await
|
||||||
}
|
}
|
||||||
|
},
|
||||||
|
ProtoRequest::Embed(req) => self.execute_single_embed(req, clients).await,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
.instrument(span)
|
.instrument(span)
|
||||||
@@ -154,6 +159,62 @@ impl RequestExecutionStage {
|
|||||||
Ok(ExecutionResult::Single { stream })
|
Ok(ExecutionResult::Single { stream })
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn execute_single_embed(
|
||||||
|
&self,
|
||||||
|
proto_request: ProtoEmbedRequest,
|
||||||
|
clients: &mut ClientSelection,
|
||||||
|
) -> Result<ExecutionResult, Response> {
|
||||||
|
let client = clients.single_mut().ok_or_else(|| {
|
||||||
|
error!(
|
||||||
|
function = "execute_single_embed",
|
||||||
|
"Expected single client but got dual"
|
||||||
|
);
|
||||||
|
error::internal_error(
|
||||||
|
"expected_single_client_got_dual",
|
||||||
|
"Expected single client but got dual",
|
||||||
|
)
|
||||||
|
})?;
|
||||||
|
|
||||||
|
let response = client.embed(proto_request).await.map_err(|e| {
|
||||||
|
error!(
|
||||||
|
function = "execute_single_embed",
|
||||||
|
error = %e,
|
||||||
|
"Failed to start embedding"
|
||||||
|
);
|
||||||
|
error::internal_error(
|
||||||
|
"start_embedding_failed",
|
||||||
|
format!("Failed to start embedding: {}", e),
|
||||||
|
)
|
||||||
|
})?;
|
||||||
|
|
||||||
|
match response.into_response() {
|
||||||
|
ProtoEmbedResponseVariant::Complete(complete) => {
|
||||||
|
Ok(ExecutionResult::Embedding { response: complete })
|
||||||
|
}
|
||||||
|
ProtoEmbedResponseVariant::Error(e) => {
|
||||||
|
error!(
|
||||||
|
function = "execute_single_embed",
|
||||||
|
error = %e.message(),
|
||||||
|
"Embedding execution failed"
|
||||||
|
);
|
||||||
|
Err(error::internal_error(
|
||||||
|
"embedding_execution_failed",
|
||||||
|
e.message().to_string(),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
ProtoEmbedResponseVariant::None => {
|
||||||
|
error!(
|
||||||
|
function = "execute_single_embed",
|
||||||
|
"Embedding execution returned no response"
|
||||||
|
);
|
||||||
|
Err(error::internal_error(
|
||||||
|
"embedding_no_response",
|
||||||
|
"Embedding execution returned no response",
|
||||||
|
))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
async fn execute_dual_dispatch(
|
async fn execute_dual_dispatch(
|
||||||
&self,
|
&self,
|
||||||
proto_request: ProtoGenerateRequest,
|
proto_request: ProtoGenerateRequest,
|
||||||
|
|||||||
@@ -11,12 +11,13 @@ use serde_json::Value;
|
|||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
client::GrpcClient,
|
client::GrpcClient,
|
||||||
proto_wrapper::{ProtoGenerateComplete, ProtoGenerateRequest, ProtoStream},
|
proto_wrapper::{ProtoEmbedComplete, ProtoGenerateComplete, ProtoRequest, ProtoStream},
|
||||||
};
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
core::{attach_guards_to_response, Worker, WorkerLoadGuard},
|
core::{attach_guards_to_response, Worker, WorkerLoadGuard},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatCompletionRequest, ChatCompletionResponse},
|
chat::{ChatCompletionRequest, ChatCompletionResponse},
|
||||||
|
embedding::{EmbeddingRequest, EmbeddingResponse},
|
||||||
generate::{GenerateRequest, GenerateResponse},
|
generate::{GenerateRequest, GenerateResponse},
|
||||||
responses::ResponsesRequest,
|
responses::ResponsesRequest,
|
||||||
},
|
},
|
||||||
@@ -49,6 +50,7 @@ pub enum RequestType {
|
|||||||
Chat(Arc<ChatCompletionRequest>),
|
Chat(Arc<ChatCompletionRequest>),
|
||||||
Generate(Arc<GenerateRequest>),
|
Generate(Arc<GenerateRequest>),
|
||||||
Responses(Arc<ResponsesRequest>),
|
Responses(Arc<ResponsesRequest>),
|
||||||
|
Embedding(Arc<EmbeddingRequest>),
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Shared components (injected once at creation)
|
/// Shared components (injected once at creation)
|
||||||
@@ -71,7 +73,7 @@ pub struct ProcessingState {
|
|||||||
pub clients: Option<ClientSelection>,
|
pub clients: Option<ClientSelection>,
|
||||||
|
|
||||||
// Stage 4: Request building outputs
|
// Stage 4: Request building outputs
|
||||||
pub proto_request: Option<ProtoGenerateRequest>,
|
pub proto_request: Option<ProtoRequest>,
|
||||||
|
|
||||||
// Stage 5: Dispatch metadata
|
// Stage 5: Dispatch metadata
|
||||||
pub dispatch: Option<DispatchMetadata>,
|
pub dispatch: Option<DispatchMetadata>,
|
||||||
@@ -202,6 +204,9 @@ pub struct ResponseState {
|
|||||||
/// Collected responses (non-streaming)
|
/// Collected responses (non-streaming)
|
||||||
pub collected: Option<Vec<ProtoGenerateComplete>>,
|
pub collected: Option<Vec<ProtoGenerateComplete>>,
|
||||||
|
|
||||||
|
/// Collected embeddings (non-streaming)
|
||||||
|
pub collected_embeddings: Option<Vec<ProtoEmbedComplete>>,
|
||||||
|
|
||||||
/// Execution result (streams from workers)
|
/// Execution result (streams from workers)
|
||||||
pub execution_result: Option<ExecutionResult>,
|
pub execution_result: Option<ExecutionResult>,
|
||||||
|
|
||||||
@@ -293,6 +298,24 @@ impl RequestContext {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Create context for embedding request
|
||||||
|
pub fn for_embedding(
|
||||||
|
request: Arc<EmbeddingRequest>,
|
||||||
|
headers: Option<HeaderMap>,
|
||||||
|
model_id: Option<String>,
|
||||||
|
components: Arc<SharedComponents>,
|
||||||
|
) -> Self {
|
||||||
|
Self {
|
||||||
|
input: RequestInput {
|
||||||
|
request_type: RequestType::Embedding(request),
|
||||||
|
headers,
|
||||||
|
model_id,
|
||||||
|
},
|
||||||
|
components,
|
||||||
|
state: ProcessingState::default(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Get reference to original request (type-safe)
|
/// Get reference to original request (type-safe)
|
||||||
pub fn request(&self) -> &RequestType {
|
pub fn request(&self) -> &RequestType {
|
||||||
&self.input.request_type
|
&self.input.request_type
|
||||||
@@ -346,12 +369,29 @@ impl RequestContext {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Get embedding request (panics if not embedding)
|
||||||
|
pub fn embedding_request(&self) -> &EmbeddingRequest {
|
||||||
|
match &self.input.request_type {
|
||||||
|
RequestType::Embedding(req) => req.as_ref(),
|
||||||
|
_ => panic!("Expected embedding request"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get Arc clone of embedding request (panics if not embedding)
|
||||||
|
pub fn embedding_request_arc(&self) -> Arc<EmbeddingRequest> {
|
||||||
|
match &self.input.request_type {
|
||||||
|
RequestType::Embedding(req) => Arc::clone(req),
|
||||||
|
_ => panic!("Expected embedding request"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Check if request is streaming
|
/// Check if request is streaming
|
||||||
pub fn is_streaming(&self) -> bool {
|
pub fn is_streaming(&self) -> bool {
|
||||||
match &self.input.request_type {
|
match &self.input.request_type {
|
||||||
RequestType::Chat(req) => req.stream,
|
RequestType::Chat(req) => req.stream,
|
||||||
RequestType::Generate(req) => req.stream,
|
RequestType::Generate(req) => req.stream,
|
||||||
RequestType::Responses(req) => req.stream.unwrap_or(false),
|
RequestType::Responses(req) => req.stream.unwrap_or(false),
|
||||||
|
RequestType::Embedding(_) => false, // Embeddings are never streaming
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -482,11 +522,18 @@ pub enum ExecutionResult {
|
|||||||
prefill: ProtoStream,
|
prefill: ProtoStream,
|
||||||
decode: Box<ProtoStream>,
|
decode: Box<ProtoStream>,
|
||||||
},
|
},
|
||||||
|
/// Embedding requests return a single response, not a stream
|
||||||
|
Embedding {
|
||||||
|
response: ProtoEmbedComplete,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Final processed response
|
/// Final processed response
|
||||||
|
#[derive(Debug)]
|
||||||
pub enum FinalResponse {
|
pub enum FinalResponse {
|
||||||
Chat(ChatCompletionResponse),
|
Chat(ChatCompletionResponse),
|
||||||
/// Generate response is a Vec of GenerateResponse (n=1 returns single item, n>1 returns multiple)
|
/// Generate response is a Vec of GenerateResponse (n=1 returns single item, n>1 returns multiple)
|
||||||
Generate(Vec<GenerateResponse>),
|
Generate(Vec<GenerateResponse>),
|
||||||
|
/// Embedding response
|
||||||
|
Embedding(EmbeddingResponse),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ use crate::routers::{
|
|||||||
grpc::{
|
grpc::{
|
||||||
common::stages::{helpers, PipelineStage},
|
common::stages::{helpers, PipelineStage},
|
||||||
context::{ClientSelection, RequestContext, RequestType, WorkerSelection},
|
context::{ClientSelection, RequestContext, RequestType, WorkerSelection},
|
||||||
proto_wrapper::ProtoGenerateRequest,
|
proto_wrapper::{ProtoGenerateRequest, ProtoRequest},
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -80,6 +80,16 @@ impl PipelineStage for HarmonyRequestBuildingStage {
|
|||||||
"Generate requests are not supported with Harmony models".to_string(),
|
"Generate requests are not supported with Harmony models".to_string(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
RequestType::Embedding(_) => {
|
||||||
|
error!(
|
||||||
|
function = "HarmonyRequestBuildingStage::execute",
|
||||||
|
"Embedding requests not supported for Harmony models"
|
||||||
|
);
|
||||||
|
return Err(error::bad_request(
|
||||||
|
"harmony_embedding_not_supported",
|
||||||
|
"Embedding requests are not supported with Harmony models".to_string(),
|
||||||
|
));
|
||||||
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
// Build gRPC request using token_ids directly (Harmony encoding already handled message rendering)
|
// Build gRPC request using token_ids directly (Harmony encoding already handled message rendering)
|
||||||
@@ -133,7 +143,17 @@ impl PipelineStage for HarmonyRequestBuildingStage {
|
|||||||
format!("Invalid request parameters: {}", e),
|
format!("Invalid request parameters: {}", e),
|
||||||
)
|
)
|
||||||
})?,
|
})?,
|
||||||
_ => unreachable!(),
|
RequestType::Embedding(_) => {
|
||||||
|
error!(
|
||||||
|
function = "HarmonyRequestBuildingStage::execute",
|
||||||
|
"Embedding requests not supported for Harmony models"
|
||||||
|
);
|
||||||
|
return Err(error::bad_request(
|
||||||
|
"harmony_embedding_not_supported",
|
||||||
|
"Embedding requests are not supported with Harmony models".to_string(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
_ => unreachable!(), // All other request types should be handled above
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut proto_request = ProtoGenerateRequest::Sglang(Box::new(proto_request_inner));
|
let mut proto_request = ProtoGenerateRequest::Sglang(Box::new(proto_request_inner));
|
||||||
@@ -159,7 +179,7 @@ impl PipelineStage for HarmonyRequestBuildingStage {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx.state.proto_request = Some(proto_request);
|
ctx.state.proto_request = Some(ProtoRequest::Generate(proto_request));
|
||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -135,14 +135,14 @@ impl PipelineStage for HarmonyResponseProcessingStage {
|
|||||||
ctx.state.response.responses_iteration_result = Some(iteration_result);
|
ctx.state.response.responses_iteration_result = Some(iteration_result);
|
||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
RequestType::Generate(_) => {
|
RequestType::Generate(_) | RequestType::Embedding(_) => {
|
||||||
error!(
|
error!(
|
||||||
function = "HarmonyResponseProcessingStage::execute",
|
function = "HarmonyResponseProcessingStage::execute",
|
||||||
"Generate request type not supported in Harmony pipeline"
|
"Generate/Embedding request type not supported in Harmony pipeline"
|
||||||
);
|
);
|
||||||
Err(error::internal_error(
|
Err(error::internal_error(
|
||||||
"generate_requests_not_supported_in_harmony",
|
"requests_not_supported_in_harmony",
|
||||||
"Generate requests not supported in Harmony pipeline",
|
"Generate/Embedding requests not supported in Harmony pipeline",
|
||||||
))
|
))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -175,6 +175,19 @@ impl HarmonyStreamingProcessor {
|
|||||||
let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n")));
|
let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n")));
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
context::ExecutionResult::Embedding { .. } => {
|
||||||
|
error!("Harmony streaming not supported for embeddings");
|
||||||
|
let error_chunk = format!(
|
||||||
|
"data: {}\n\n",
|
||||||
|
json!({
|
||||||
|
"error": {
|
||||||
|
"message": "Embeddings not supported in Harmony streaming",
|
||||||
|
"type": "invalid_request_error"
|
||||||
|
}
|
||||||
|
})
|
||||||
|
);
|
||||||
|
let _ = tx.send(Ok(Bytes::from(error_chunk)));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Return SSE response
|
// Return SSE response
|
||||||
@@ -644,6 +657,9 @@ impl HarmonyStreamingProcessor {
|
|||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
context::ExecutionResult::Embedding { .. } => {
|
||||||
|
Err("Embeddings not supported in Responses API streaming".to_string())
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -6,13 +6,25 @@
|
|||||||
use std::{sync::Arc, time::Instant};
|
use std::{sync::Arc, time::Instant};
|
||||||
|
|
||||||
use axum::response::{IntoResponse, Response};
|
use axum::response::{IntoResponse, Response};
|
||||||
use tracing::error;
|
use tracing::{debug, error};
|
||||||
|
|
||||||
|
// Import embedding-specific stages
|
||||||
|
use super::regular::stages::embedding::preparation::EmbeddingPreparationStage;
|
||||||
use super::{
|
use super::{
|
||||||
common::stages::*,
|
common::stages::*,
|
||||||
context::*,
|
context::*,
|
||||||
harmony,
|
harmony,
|
||||||
regular::{processor, stages::*, streaming},
|
regular::{
|
||||||
|
processor,
|
||||||
|
stages::{
|
||||||
|
embedding::{
|
||||||
|
request_building::EmbeddingRequestBuildingStage,
|
||||||
|
response_processing::EmbeddingResponseProcessingStage,
|
||||||
|
},
|
||||||
|
*,
|
||||||
|
},
|
||||||
|
streaming,
|
||||||
|
},
|
||||||
utils::error_type_from_status,
|
utils::error_type_from_status,
|
||||||
};
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
@@ -21,6 +33,7 @@ use crate::{
|
|||||||
policies::PolicyRegistry,
|
policies::PolicyRegistry,
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatCompletionRequest, ChatCompletionResponse},
|
chat::{ChatCompletionRequest, ChatCompletionResponse},
|
||||||
|
embedding::EmbeddingRequest,
|
||||||
generate::GenerateRequest,
|
generate::GenerateRequest,
|
||||||
},
|
},
|
||||||
reasoning_parser::ParserFactory as ReasoningParserFactory,
|
reasoning_parser::ParserFactory as ReasoningParserFactory,
|
||||||
@@ -186,6 +199,32 @@ impl RequestPipeline {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Create an embeddings pipeline
|
||||||
|
pub fn new_embeddings(
|
||||||
|
worker_registry: Arc<WorkerRegistry>,
|
||||||
|
policy_registry: Arc<PolicyRegistry>,
|
||||||
|
_tokenizer: Arc<dyn Tokenizer>,
|
||||||
|
) -> Self {
|
||||||
|
let stages: Vec<Box<dyn PipelineStage>> = vec![
|
||||||
|
Box::new(EmbeddingPreparationStage::new()),
|
||||||
|
Box::new(WorkerSelectionStage::new(
|
||||||
|
worker_registry,
|
||||||
|
policy_registry,
|
||||||
|
WorkerSelectionMode::Regular, // Embeddings are always single
|
||||||
|
)),
|
||||||
|
Box::new(ClientAcquisitionStage),
|
||||||
|
Box::new(EmbeddingRequestBuildingStage::new()),
|
||||||
|
Box::new(DispatchMetadataStage),
|
||||||
|
Box::new(RequestExecutionStage::new(ExecutionMode::Single)),
|
||||||
|
Box::new(EmbeddingResponseProcessingStage::new()),
|
||||||
|
];
|
||||||
|
|
||||||
|
Self {
|
||||||
|
stages: Arc::new(stages),
|
||||||
|
backend_type: metrics_labels::BACKEND_REGULAR, // Embeddings are regular for now
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Execute the complete pipeline for a chat request
|
/// Execute the complete pipeline for a chat request
|
||||||
pub async fn execute_chat(
|
pub async fn execute_chat(
|
||||||
&self,
|
&self,
|
||||||
@@ -257,10 +296,10 @@ impl RequestPipeline {
|
|||||||
);
|
);
|
||||||
axum::Json(response).into_response()
|
axum::Json(response).into_response()
|
||||||
}
|
}
|
||||||
Some(FinalResponse::Generate(_)) => {
|
Some(FinalResponse::Generate(_)) | Some(FinalResponse::Embedding(_)) => {
|
||||||
error!(
|
error!(
|
||||||
function = "execute_chat",
|
function = "execute_chat",
|
||||||
"Wrong response type: expected Chat, got Generate"
|
"Wrong response type: expected Chat, got Generate/Embedding"
|
||||||
);
|
);
|
||||||
Metrics::record_router_error(
|
Metrics::record_router_error(
|
||||||
metrics_labels::ROUTER_GRPC,
|
metrics_labels::ROUTER_GRPC,
|
||||||
@@ -361,10 +400,10 @@ impl RequestPipeline {
|
|||||||
);
|
);
|
||||||
axum::Json(response).into_response()
|
axum::Json(response).into_response()
|
||||||
}
|
}
|
||||||
Some(FinalResponse::Chat(_)) => {
|
Some(FinalResponse::Chat(_)) | Some(FinalResponse::Embedding(_)) => {
|
||||||
error!(
|
error!(
|
||||||
function = "execute_generate",
|
function = "execute_generate",
|
||||||
"Wrong response type: expected Generate, got Chat"
|
"Wrong response type: expected Generate, got Chat/Embedding"
|
||||||
);
|
);
|
||||||
Metrics::record_router_error(
|
Metrics::record_router_error(
|
||||||
metrics_labels::ROUTER_GRPC,
|
metrics_labels::ROUTER_GRPC,
|
||||||
@@ -394,6 +433,112 @@ impl RequestPipeline {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Execute the complete pipeline for an embedding request
|
||||||
|
pub async fn execute_embeddings(
|
||||||
|
&self,
|
||||||
|
request: Arc<EmbeddingRequest>,
|
||||||
|
headers: Option<http::HeaderMap>,
|
||||||
|
model_id: Option<String>,
|
||||||
|
components: Arc<SharedComponents>,
|
||||||
|
) -> Response {
|
||||||
|
debug!(
|
||||||
|
"execute_embeddings: 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_EMBEDDINGS,
|
||||||
|
bool_to_static_str(false),
|
||||||
|
);
|
||||||
|
|
||||||
|
let mut ctx = RequestContext::for_embedding(request, headers, model_id.clone(), components);
|
||||||
|
|
||||||
|
for stage in self.stages.iter() {
|
||||||
|
debug!("execute_embeddings: Executing stage: {}", stage.name());
|
||||||
|
match stage.execute(&mut ctx).await {
|
||||||
|
Ok(Some(response)) => {
|
||||||
|
debug!(
|
||||||
|
"execute_embeddings: 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_EMBEDDINGS,
|
||||||
|
start.elapsed(),
|
||||||
|
);
|
||||||
|
return response;
|
||||||
|
}
|
||||||
|
Ok(None) => {
|
||||||
|
debug!(
|
||||||
|
"execute_embeddings: Stage {} completed, continuing to next stage.",
|
||||||
|
stage.name()
|
||||||
|
);
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
Err(response) => {
|
||||||
|
error!(
|
||||||
|
"execute_embeddings: 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_EMBEDDINGS,
|
||||||
|
error_type_from_status(response.status()),
|
||||||
|
);
|
||||||
|
return response;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
debug!(
|
||||||
|
"execute_embeddings: Pipeline finished, processing final_response. Current state: {:?}",
|
||||||
|
ctx.state.response.final_response
|
||||||
|
);
|
||||||
|
match ctx.state.response.final_response {
|
||||||
|
Some(FinalResponse::Embedding(_)) => {
|
||||||
|
error!("execute_embeddings: Embedding FinalResponse found, but pipeline finished without returning response directly. This should be handled by the last stage.");
|
||||||
|
// Already handled in ResponseProcessingStage, but just in case
|
||||||
|
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_EMBEDDINGS,
|
||||||
|
start.elapsed(),
|
||||||
|
);
|
||||||
|
// The response should have been returned by the last stage
|
||||||
|
error::internal_error(
|
||||||
|
"pipeline_fallthrough",
|
||||||
|
"Pipeline finished without returning response",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
Some(_) => {
|
||||||
|
error!(function = "execute_embeddings", "Wrong response type");
|
||||||
|
error::internal_error("wrong_response_type", "Internal error: wrong response type")
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
error!(
|
||||||
|
function = "execute_embeddings",
|
||||||
|
"No final response produced by pipeline."
|
||||||
|
);
|
||||||
|
error::internal_error("no_response_produced", "No response produced")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Execute chat pipeline for responses endpoint
|
/// Execute chat pipeline for responses endpoint
|
||||||
///
|
///
|
||||||
/// Used by ALL non-streaming /v1/responses requests.
|
/// Used by ALL non-streaming /v1/responses requests.
|
||||||
@@ -440,10 +585,10 @@ impl RequestPipeline {
|
|||||||
|
|
||||||
match ctx.state.response.final_response {
|
match ctx.state.response.final_response {
|
||||||
Some(FinalResponse::Chat(response)) => Ok(response),
|
Some(FinalResponse::Chat(response)) => Ok(response),
|
||||||
Some(FinalResponse::Generate(_)) => {
|
Some(FinalResponse::Generate(_)) | Some(FinalResponse::Embedding(_)) => {
|
||||||
error!(
|
error!(
|
||||||
function = "execute_chat_for_responses",
|
function = "execute_chat_for_responses",
|
||||||
"Wrong response type: expected Chat, got Generate"
|
"Wrong response type: expected Chat, got Generate/Embedding"
|
||||||
);
|
);
|
||||||
Err(error::internal_error(
|
Err(error::internal_error(
|
||||||
"wrong_response_type",
|
"wrong_response_type",
|
||||||
|
|||||||
@@ -12,6 +12,36 @@ use crate::grpc_client::{
|
|||||||
vllm_proto as vllm,
|
vllm_proto as vllm,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
/// Unified ProtoRequest
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum ProtoRequest {
|
||||||
|
Generate(ProtoGenerateRequest),
|
||||||
|
Embed(ProtoEmbedRequest),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoRequest {
|
||||||
|
pub fn as_generate(&self) -> &ProtoGenerateRequest {
|
||||||
|
match self {
|
||||||
|
Self::Generate(req) => req,
|
||||||
|
_ => panic!("Expected Generate request"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn as_embed(&self) -> &ProtoEmbedRequest {
|
||||||
|
match self {
|
||||||
|
Self::Embed(req) => req,
|
||||||
|
_ => panic!("Expected Embed request"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn request_id(&self) -> &str {
|
||||||
|
match self {
|
||||||
|
Self::Generate(req) => req.request_id(),
|
||||||
|
Self::Embed(req) => req.request_id(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Unified GenerateRequest that works with both backends
|
/// Unified GenerateRequest that works with both backends
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
pub enum ProtoGenerateRequest {
|
pub enum ProtoGenerateRequest {
|
||||||
@@ -379,3 +409,129 @@ impl ProtoStream {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Unified EmbedRequest that works with both backends
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum ProtoEmbedRequest {
|
||||||
|
Sglang(Box<sglang::EmbedRequest>),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoEmbedRequest {
|
||||||
|
/// Get SGLang variant
|
||||||
|
pub fn as_sglang(&self) -> &sglang::EmbedRequest {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(req) => req,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get mutable SGLang variant
|
||||||
|
pub fn as_sglang_mut(&mut self) -> &mut sglang::EmbedRequest {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(req) => req,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is SGLang
|
||||||
|
pub fn is_sglang(&self) -> bool {
|
||||||
|
matches!(self, Self::Sglang(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Clone the inner request (for passing to embed())
|
||||||
|
pub fn clone_inner(&self) -> Self {
|
||||||
|
self.clone()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get request ID
|
||||||
|
pub fn request_id(&self) -> &str {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(req) => &req.request_id,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified EmbedResponse
|
||||||
|
pub enum ProtoEmbedResponse {
|
||||||
|
Sglang(sglang::EmbedResponse),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoEmbedResponse {
|
||||||
|
/// Get the response variant (complete or error)
|
||||||
|
pub fn into_response(self) -> ProtoEmbedResponseVariant {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(resp) => match resp.response {
|
||||||
|
Some(sglang::embed_response::Response::Complete(complete)) => {
|
||||||
|
ProtoEmbedResponseVariant::Complete(ProtoEmbedComplete::Sglang(complete))
|
||||||
|
}
|
||||||
|
Some(sglang::embed_response::Response::Error(error)) => {
|
||||||
|
ProtoEmbedResponseVariant::Error(ProtoEmbedError::Sglang(error))
|
||||||
|
}
|
||||||
|
None => ProtoEmbedResponseVariant::None,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Response variant extracted from EmbedResponse
|
||||||
|
pub enum ProtoEmbedResponseVariant {
|
||||||
|
Complete(ProtoEmbedComplete),
|
||||||
|
Error(ProtoEmbedError),
|
||||||
|
None,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified EmbedComplete response
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum ProtoEmbedComplete {
|
||||||
|
Sglang(sglang::EmbedComplete),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoEmbedComplete {
|
||||||
|
/// Get embeddings
|
||||||
|
pub fn embedding(&self) -> &[f32] {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => &c.embedding,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get prompt tokens
|
||||||
|
pub fn prompt_tokens(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.prompt_tokens,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get cached tokens
|
||||||
|
pub fn cached_tokens(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.cached_tokens,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get embedding dimension
|
||||||
|
pub fn embedding_dim(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.embedding_dim,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified EmbedError
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum ProtoEmbedError {
|
||||||
|
Sglang(sglang::EmbedError),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoEmbedError {
|
||||||
|
/// Get error message
|
||||||
|
pub fn message(&self) -> &str {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(e) => &e.message,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get error code
|
||||||
|
pub fn code(&self) -> &str {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(e) => &e.code,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -108,7 +108,9 @@ impl PipelineStage for ChatRequestBuildingStage {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx.state.proto_request = Some(proto_request);
|
ctx.state.proto_request = Some(
|
||||||
|
crate::routers::grpc::proto_wrapper::ProtoRequest::Generate(proto_request),
|
||||||
|
);
|
||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
pub mod preparation;
|
||||||
|
pub mod request_building;
|
||||||
|
pub mod response_processing;
|
||||||
@@ -0,0 +1,93 @@
|
|||||||
|
//! Preparation stage for embedding requests
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use axum::response::Response;
|
||||||
|
use tracing::error;
|
||||||
|
|
||||||
|
use crate::{
|
||||||
|
protocols::common::GenerationRequest,
|
||||||
|
routers::{
|
||||||
|
error,
|
||||||
|
grpc::{
|
||||||
|
common::stages::PipelineStage,
|
||||||
|
context::{PreparationOutput, RequestContext, RequestType},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
pub struct EmbeddingPreparationStage;
|
||||||
|
|
||||||
|
impl EmbeddingPreparationStage {
|
||||||
|
pub fn new() -> Self {
|
||||||
|
Self
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for EmbeddingPreparationStage {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[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 request
|
||||||
|
let text = request.extract_text_for_routing();
|
||||||
|
if text.is_empty() {
|
||||||
|
return Err(error::bad_request(
|
||||||
|
"empty_input",
|
||||||
|
"Input text cannot be empty",
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tokenize
|
||||||
|
let token_ids = ctx
|
||||||
|
.components
|
||||||
|
.tokenizer
|
||||||
|
.encode(&text)
|
||||||
|
.map_err(|e| {
|
||||||
|
error!(
|
||||||
|
function = "EmbeddingPreparationStage::execute",
|
||||||
|
error = %e,
|
||||||
|
"Tokenization failed"
|
||||||
|
);
|
||||||
|
error::bad_request("tokenization_failed", format!("Tokenization failed: {}", e))
|
||||||
|
})?
|
||||||
|
.token_ids()
|
||||||
|
.to_vec();
|
||||||
|
|
||||||
|
// Store preparation output
|
||||||
|
ctx.state.preparation = Some(PreparationOutput {
|
||||||
|
original_text: Some(text),
|
||||||
|
token_ids,
|
||||||
|
processed_messages: None,
|
||||||
|
tool_constraints: None,
|
||||||
|
filtered_request: None,
|
||||||
|
harmony_mode: false,
|
||||||
|
selection_text: None,
|
||||||
|
harmony_messages: None,
|
||||||
|
harmony_stop_ids: None,
|
||||||
|
});
|
||||||
|
|
||||||
|
Ok(None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn name(&self) -> &'static str {
|
||||||
|
"EmbeddingPreparation"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
//! Request building stage for embedding requests
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use axum::response::Response;
|
||||||
|
use tracing::error;
|
||||||
|
|
||||||
|
use crate::routers::{
|
||||||
|
error,
|
||||||
|
grpc::{
|
||||||
|
common::stages::PipelineStage,
|
||||||
|
context::{RequestContext, RequestType},
|
||||||
|
proto_wrapper::{ProtoEmbedRequest, ProtoRequest},
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
/// Request building stage for embedding requests
|
||||||
|
pub struct EmbeddingRequestBuildingStage;
|
||||||
|
|
||||||
|
impl EmbeddingRequestBuildingStage {
|
||||||
|
pub fn new() -> Self {
|
||||||
|
Self
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for EmbeddingRequestBuildingStage {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[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",
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Preparation output should have tokenized input
|
||||||
|
let prep_output = ctx.state.preparation.as_ref().ok_or_else(|| {
|
||||||
|
error!(
|
||||||
|
function = "EmbeddingRequestBuildingStage::execute",
|
||||||
|
"Preparation output missing"
|
||||||
|
);
|
||||||
|
error::internal_error("preparation_missing", "Preparation output missing")
|
||||||
|
})?;
|
||||||
|
|
||||||
|
// Extract client
|
||||||
|
let client = ctx
|
||||||
|
.state
|
||||||
|
.clients
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|c| c.single())
|
||||||
|
.ok_or_else(|| {
|
||||||
|
error!(
|
||||||
|
function = "EmbeddingRequestBuildingStage::execute",
|
||||||
|
"Client not selected"
|
||||||
|
);
|
||||||
|
error::internal_error("client_missing", "Client not selected")
|
||||||
|
})?;
|
||||||
|
|
||||||
|
// Extract request ID
|
||||||
|
let request_id = ctx
|
||||||
|
.state
|
||||||
|
.dispatch
|
||||||
|
.as_ref()
|
||||||
|
.map(|d| d.request_id.clone())
|
||||||
|
.unwrap_or_else(|| "unknown".to_string());
|
||||||
|
|
||||||
|
// Extract original text
|
||||||
|
let original_text = prep_output.original_text.clone();
|
||||||
|
|
||||||
|
// 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,
|
||||||
|
);
|
||||||
|
|
||||||
|
let proto_req = ProtoEmbedRequest::Sglang(Box::new(sglang_req));
|
||||||
|
|
||||||
|
ctx.state.proto_request = Some(ProtoRequest::Embed(proto_req));
|
||||||
|
Ok(None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn name(&self) -> &'static str {
|
||||||
|
"EmbeddingRequestBuilding"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
//! Response processing stage for embedding requests
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use axum::{
|
||||||
|
response::{IntoResponse, Response},
|
||||||
|
Json,
|
||||||
|
};
|
||||||
|
use tracing::error;
|
||||||
|
|
||||||
|
use crate::{
|
||||||
|
protocols::embedding::{EmbeddingObject, EmbeddingResponse},
|
||||||
|
routers::{
|
||||||
|
error,
|
||||||
|
grpc::{
|
||||||
|
common::stages::PipelineStage,
|
||||||
|
context::{ExecutionResult, FinalResponse, RequestContext},
|
||||||
|
proto_wrapper::ProtoEmbedComplete,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
/// Response processing stage for embedding requests
|
||||||
|
pub struct EmbeddingResponseProcessingStage;
|
||||||
|
|
||||||
|
impl EmbeddingResponseProcessingStage {
|
||||||
|
pub fn new() -> Self {
|
||||||
|
Self
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for EmbeddingResponseProcessingStage {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl PipelineStage for EmbeddingResponseProcessingStage {
|
||||||
|
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 = "EmbeddingResponseProcessingStage::execute",
|
||||||
|
"Execution result missing"
|
||||||
|
);
|
||||||
|
error::internal_error("execution_result_missing", "Execution result missing")
|
||||||
|
})?;
|
||||||
|
|
||||||
|
// Expect Embedding result variant
|
||||||
|
let proto_response = if let ExecutionResult::Embedding { response } = execution_result {
|
||||||
|
response
|
||||||
|
} else {
|
||||||
|
error!(
|
||||||
|
function = "EmbeddingResponseProcessingStage::execute",
|
||||||
|
"Invalid execution result: expected Embedding"
|
||||||
|
);
|
||||||
|
return Err(error::internal_error(
|
||||||
|
"invalid_execution_result",
|
||||||
|
"Expected Embedding result",
|
||||||
|
));
|
||||||
|
};
|
||||||
|
|
||||||
|
// Convert proto response to HTTP response
|
||||||
|
let embedding_response = self
|
||||||
|
.convert_response(ctx, proto_response)
|
||||||
|
.map_err(|boxed_err| *boxed_err)?;
|
||||||
|
|
||||||
|
// Store in context
|
||||||
|
ctx.state.response.final_response =
|
||||||
|
Some(FinalResponse::Embedding(embedding_response.clone()));
|
||||||
|
|
||||||
|
// Return the HTTP response directly
|
||||||
|
Ok(Some(Json(embedding_response).into_response()))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn name(&self) -> &'static str {
|
||||||
|
"EmbeddingResponseProcessing"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl EmbeddingResponseProcessingStage {
|
||||||
|
fn convert_response(
|
||||||
|
&self,
|
||||||
|
ctx: &RequestContext,
|
||||||
|
proto: ProtoEmbedComplete,
|
||||||
|
) -> Result<EmbeddingResponse, Box<Response>> {
|
||||||
|
let dispatch = ctx.state.dispatch.as_ref().ok_or_else(|| {
|
||||||
|
error!(
|
||||||
|
function = "EmbeddingResponseProcessingStage::convert_response",
|
||||||
|
"Dispatch metadata missing in context"
|
||||||
|
);
|
||||||
|
error::internal_error("dispatch_missing", "Dispatch metadata missing")
|
||||||
|
})?;
|
||||||
|
|
||||||
|
let model = dispatch.model.clone();
|
||||||
|
|
||||||
|
// Convert flat embedding vector to response
|
||||||
|
// single input -> single embedding object
|
||||||
|
|
||||||
|
let embedding_data = EmbeddingObject {
|
||||||
|
object: "embedding".to_string(),
|
||||||
|
embedding: proto.embedding().to_vec(),
|
||||||
|
index: 0,
|
||||||
|
};
|
||||||
|
|
||||||
|
// Casting i32 to u32 for usage stats
|
||||||
|
let prompt_tokens = proto.prompt_tokens().max(0) as u32;
|
||||||
|
|
||||||
|
let usage = crate::protocols::common::UsageInfo {
|
||||||
|
prompt_tokens,
|
||||||
|
total_tokens: prompt_tokens, // Embedding has no completion tokens
|
||||||
|
completion_tokens: 0,
|
||||||
|
prompt_tokens_details: None,
|
||||||
|
reasoning_tokens: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
Ok(EmbeddingResponse {
|
||||||
|
object: "list".to_string(),
|
||||||
|
data: vec![embedding_data],
|
||||||
|
model,
|
||||||
|
usage,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -103,7 +103,9 @@ impl PipelineStage for GenerateRequestBuildingStage {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx.state.proto_request = Some(proto_request);
|
ctx.state.proto_request = Some(
|
||||||
|
crate::routers::grpc::proto_wrapper::ProtoRequest::Generate(proto_request),
|
||||||
|
);
|
||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,7 @@
|
|||||||
//! This module defines stages specific to regular tokenizer-based models.
|
//! This module defines stages specific to regular tokenizer-based models.
|
||||||
|
|
||||||
pub mod chat;
|
pub mod chat;
|
||||||
|
pub mod embedding;
|
||||||
pub mod generate;
|
pub mod generate;
|
||||||
mod preparation;
|
mod preparation;
|
||||||
mod request_building;
|
mod request_building;
|
||||||
|
|||||||
@@ -7,7 +7,10 @@ use async_trait::async_trait;
|
|||||||
use axum::response::Response;
|
use axum::response::Response;
|
||||||
use tracing::error;
|
use tracing::error;
|
||||||
|
|
||||||
use super::{chat::ChatPreparationStage, generate::GeneratePreparationStage};
|
use super::{
|
||||||
|
chat::ChatPreparationStage, embedding::preparation::EmbeddingPreparationStage,
|
||||||
|
generate::GeneratePreparationStage,
|
||||||
|
};
|
||||||
use crate::routers::{
|
use crate::routers::{
|
||||||
error as grpc_error,
|
error as grpc_error,
|
||||||
grpc::{
|
grpc::{
|
||||||
@@ -20,6 +23,7 @@ use crate::routers::{
|
|||||||
pub struct PreparationStage {
|
pub struct PreparationStage {
|
||||||
chat_stage: ChatPreparationStage,
|
chat_stage: ChatPreparationStage,
|
||||||
generate_stage: GeneratePreparationStage,
|
generate_stage: GeneratePreparationStage,
|
||||||
|
embedding_stage: EmbeddingPreparationStage,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl PreparationStage {
|
impl PreparationStage {
|
||||||
@@ -27,6 +31,7 @@ impl PreparationStage {
|
|||||||
Self {
|
Self {
|
||||||
chat_stage: ChatPreparationStage,
|
chat_stage: ChatPreparationStage,
|
||||||
generate_stage: GeneratePreparationStage,
|
generate_stage: GeneratePreparationStage,
|
||||||
|
embedding_stage: EmbeddingPreparationStage::new(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -43,6 +48,7 @@ impl PipelineStage for PreparationStage {
|
|||||||
match &ctx.input.request_type {
|
match &ctx.input.request_type {
|
||||||
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
||||||
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
||||||
|
RequestType::Embedding(_) => self.embedding_stage.execute(ctx).await,
|
||||||
RequestType::Responses(_) => {
|
RequestType::Responses(_) => {
|
||||||
error!(
|
error!(
|
||||||
function = "PreparationStage::execute",
|
function = "PreparationStage::execute",
|
||||||
|
|||||||
@@ -4,7 +4,10 @@ use async_trait::async_trait;
|
|||||||
use axum::response::Response;
|
use axum::response::Response;
|
||||||
use tracing::error;
|
use tracing::error;
|
||||||
|
|
||||||
use super::{chat::ChatRequestBuildingStage, generate::GenerateRequestBuildingStage};
|
use super::{
|
||||||
|
chat::ChatRequestBuildingStage, embedding::request_building::EmbeddingRequestBuildingStage,
|
||||||
|
generate::GenerateRequestBuildingStage,
|
||||||
|
};
|
||||||
use crate::routers::{
|
use crate::routers::{
|
||||||
error as grpc_error,
|
error as grpc_error,
|
||||||
grpc::{
|
grpc::{
|
||||||
@@ -17,6 +20,7 @@ use crate::routers::{
|
|||||||
pub struct RequestBuildingStage {
|
pub struct RequestBuildingStage {
|
||||||
chat_stage: ChatRequestBuildingStage,
|
chat_stage: ChatRequestBuildingStage,
|
||||||
generate_stage: GenerateRequestBuildingStage,
|
generate_stage: GenerateRequestBuildingStage,
|
||||||
|
embedding_stage: EmbeddingRequestBuildingStage,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl RequestBuildingStage {
|
impl RequestBuildingStage {
|
||||||
@@ -24,6 +28,7 @@ impl RequestBuildingStage {
|
|||||||
Self {
|
Self {
|
||||||
chat_stage: ChatRequestBuildingStage::new(inject_pd_metadata),
|
chat_stage: ChatRequestBuildingStage::new(inject_pd_metadata),
|
||||||
generate_stage: GenerateRequestBuildingStage::new(inject_pd_metadata),
|
generate_stage: GenerateRequestBuildingStage::new(inject_pd_metadata),
|
||||||
|
embedding_stage: EmbeddingRequestBuildingStage::new(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -34,6 +39,7 @@ impl PipelineStage for RequestBuildingStage {
|
|||||||
match &ctx.input.request_type {
|
match &ctx.input.request_type {
|
||||||
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
||||||
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
||||||
|
RequestType::Embedding(_) => self.embedding_stage.execute(ctx).await,
|
||||||
RequestType::Responses(_request) => {
|
RequestType::Responses(_request) => {
|
||||||
error!(
|
error!(
|
||||||
function = "RequestBuildingStage::execute",
|
function = "RequestBuildingStage::execute",
|
||||||
|
|||||||
@@ -6,7 +6,11 @@ use async_trait::async_trait;
|
|||||||
use axum::response::Response;
|
use axum::response::Response;
|
||||||
use tracing::error;
|
use tracing::error;
|
||||||
|
|
||||||
use super::{chat::ChatResponseProcessingStage, generate::GenerateResponseProcessingStage};
|
use super::{
|
||||||
|
chat::ChatResponseProcessingStage,
|
||||||
|
embedding::response_processing::EmbeddingResponseProcessingStage,
|
||||||
|
generate::GenerateResponseProcessingStage,
|
||||||
|
};
|
||||||
use crate::routers::{
|
use crate::routers::{
|
||||||
error,
|
error,
|
||||||
grpc::{
|
grpc::{
|
||||||
@@ -20,6 +24,7 @@ use crate::routers::{
|
|||||||
pub struct ResponseProcessingStage {
|
pub struct ResponseProcessingStage {
|
||||||
chat_stage: ChatResponseProcessingStage,
|
chat_stage: ChatResponseProcessingStage,
|
||||||
generate_stage: GenerateResponseProcessingStage,
|
generate_stage: GenerateResponseProcessingStage,
|
||||||
|
embedding_stage: EmbeddingResponseProcessingStage,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ResponseProcessingStage {
|
impl ResponseProcessingStage {
|
||||||
@@ -33,6 +38,7 @@ impl ResponseProcessingStage {
|
|||||||
streaming_processor.clone(),
|
streaming_processor.clone(),
|
||||||
),
|
),
|
||||||
generate_stage: GenerateResponseProcessingStage::new(processor, streaming_processor),
|
generate_stage: GenerateResponseProcessingStage::new(processor, streaming_processor),
|
||||||
|
embedding_stage: EmbeddingResponseProcessingStage::new(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -43,6 +49,7 @@ impl PipelineStage for ResponseProcessingStage {
|
|||||||
match &ctx.input.request_type {
|
match &ctx.input.request_type {
|
||||||
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
||||||
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
||||||
|
RequestType::Embedding(_) => self.embedding_stage.execute(ctx).await,
|
||||||
RequestType::Responses(_) => {
|
RequestType::Responses(_) => {
|
||||||
error!(
|
error!(
|
||||||
function = "ResponseProcessingStage::execute",
|
function = "ResponseProcessingStage::execute",
|
||||||
|
|||||||
@@ -167,6 +167,18 @@ impl StreamingProcessor {
|
|||||||
let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n")));
|
let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n")));
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
context::ExecutionResult::Embedding { .. } => {
|
||||||
|
let error_chunk = format!(
|
||||||
|
"data: {}\n\n",
|
||||||
|
json!({
|
||||||
|
"error": {
|
||||||
|
"message": "Embeddings not supported in streaming mode",
|
||||||
|
"type": "invalid_request_error"
|
||||||
|
}
|
||||||
|
})
|
||||||
|
);
|
||||||
|
let _ = tx.send(Ok(Bytes::from(error_chunk)));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Return SSE response
|
// Return SSE response
|
||||||
@@ -703,6 +715,11 @@ impl StreamingProcessor {
|
|||||||
let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n")));
|
let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n")));
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
context::ExecutionResult::Embedding { .. } => {
|
||||||
|
let error_chunk =
|
||||||
|
"data: {\"error\": \"Embeddings not supported in streaming generate\"}\n\n";
|
||||||
|
let _ = tx.send(Ok(Bytes::from(error_chunk)));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Return SSE response
|
// Return SSE response
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ use crate::{
|
|||||||
observability::metrics::{metrics_labels, Metrics},
|
observability::metrics::{metrics_labels, Metrics},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::ChatCompletionRequest,
|
chat::ChatCompletionRequest,
|
||||||
|
embedding::EmbeddingRequest,
|
||||||
generate::GenerateRequest,
|
generate::GenerateRequest,
|
||||||
responses::{ResponsesGetParams, ResponsesRequest},
|
responses::{ResponsesGetParams, ResponsesRequest},
|
||||||
},
|
},
|
||||||
@@ -40,6 +41,7 @@ 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
|
||||||
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,
|
||||||
@@ -93,6 +95,13 @@ impl GrpcRouter {
|
|||||||
ctx.configured_reasoning_parser.clone(),
|
ctx.configured_reasoning_parser.clone(),
|
||||||
);
|
);
|
||||||
|
|
||||||
|
// Create Embedding pipeline
|
||||||
|
let embedding_pipeline = RequestPipeline::new_embeddings(
|
||||||
|
worker_registry.clone(),
|
||||||
|
_policy_registry.clone(),
|
||||||
|
tokenizer.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
|
||||||
@@ -121,6 +130,7 @@ impl GrpcRouter {
|
|||||||
worker_registry,
|
worker_registry,
|
||||||
pipeline,
|
pipeline,
|
||||||
harmony_pipeline,
|
harmony_pipeline,
|
||||||
|
embedding_pipeline,
|
||||||
shared_components,
|
shared_components,
|
||||||
responses_context,
|
responses_context,
|
||||||
harmony_responses_context,
|
harmony_responses_context,
|
||||||
@@ -298,6 +308,25 @@ impl GrpcRouter {
|
|||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Main route_embeddings implementation
|
||||||
|
async fn route_embeddings_impl(
|
||||||
|
&self,
|
||||||
|
headers: Option<&HeaderMap>,
|
||||||
|
body: &EmbeddingRequest,
|
||||||
|
model_id: Option<&str>,
|
||||||
|
) -> Response {
|
||||||
|
debug!("Processing embedding request for model: {:?}", model_id);
|
||||||
|
|
||||||
|
self.embedding_pipeline
|
||||||
|
.execute_embeddings(
|
||||||
|
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 {
|
||||||
@@ -355,6 +384,15 @@ impl RouterTrait for GrpcRouter {
|
|||||||
cancel_response_impl(&self.responses_context, response_id).await
|
cancel_response_impl(&self.responses_context, response_id).await
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn route_embeddings(
|
||||||
|
&self,
|
||||||
|
headers: Option<&HeaderMap>,
|
||||||
|
body: &EmbeddingRequest,
|
||||||
|
model_id: Option<&str>,
|
||||||
|
) -> Response {
|
||||||
|
self.route_embeddings_impl(headers, body, model_id).await
|
||||||
|
}
|
||||||
|
|
||||||
fn router_type(&self) -> &'static str {
|
fn router_type(&self) -> &'static str {
|
||||||
"grpc"
|
"grpc"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ fn test_embedding_request_serialization_string_input() {
|
|||||||
user: Some("user-1".to_string()),
|
user: Some("user-1".to_string()),
|
||||||
dimensions: Some(128),
|
dimensions: Some(128),
|
||||||
rid: Some("rid-123".to_string()),
|
rid: Some("rid-123".to_string()),
|
||||||
|
log_metrics: None,
|
||||||
};
|
};
|
||||||
|
|
||||||
let serialized = to_string(&req).unwrap();
|
let serialized = to_string(&req).unwrap();
|
||||||
@@ -32,6 +33,7 @@ fn test_embedding_request_serialization_array_input() {
|
|||||||
user: None,
|
user: None,
|
||||||
dimensions: None,
|
dimensions: None,
|
||||||
rid: None,
|
rid: None,
|
||||||
|
log_metrics: None,
|
||||||
};
|
};
|
||||||
|
|
||||||
let serialized = to_string(&req).unwrap();
|
let serialized = to_string(&req).unwrap();
|
||||||
@@ -49,6 +51,7 @@ fn test_embedding_generation_request_trait_string() {
|
|||||||
user: None,
|
user: None,
|
||||||
dimensions: None,
|
dimensions: None,
|
||||||
rid: None,
|
rid: None,
|
||||||
|
log_metrics: None,
|
||||||
};
|
};
|
||||||
assert!(!req.is_stream());
|
assert!(!req.is_stream());
|
||||||
assert_eq!(req.get_model(), Some("emb-model"));
|
assert_eq!(req.get_model(), Some("emb-model"));
|
||||||
@@ -64,6 +67,7 @@ fn test_embedding_generation_request_trait_array() {
|
|||||||
user: None,
|
user: None,
|
||||||
dimensions: None,
|
dimensions: None,
|
||||||
rid: None,
|
rid: None,
|
||||||
|
log_metrics: None,
|
||||||
};
|
};
|
||||||
assert_eq!(req.extract_text_for_routing(), "hello world");
|
assert_eq!(req.extract_text_for_routing(), "hello world");
|
||||||
}
|
}
|
||||||
@@ -77,6 +81,7 @@ fn test_embedding_generation_request_trait_non_text() {
|
|||||||
user: None,
|
user: None,
|
||||||
dimensions: None,
|
dimensions: None,
|
||||||
rid: None,
|
rid: None,
|
||||||
|
log_metrics: None,
|
||||||
};
|
};
|
||||||
assert_eq!(req.extract_text_for_routing(), "");
|
assert_eq!(req.extract_text_for_routing(), "");
|
||||||
}
|
}
|
||||||
@@ -90,6 +95,7 @@ fn test_embedding_generation_request_trait_mixed_array_ignores_nested() {
|
|||||||
user: None,
|
user: None,
|
||||||
dimensions: None,
|
dimensions: None,
|
||||||
rid: None,
|
rid: None,
|
||||||
|
log_metrics: None,
|
||||||
};
|
};
|
||||||
// Only top-level string elements are extracted
|
// Only top-level string elements are extracted
|
||||||
assert_eq!(req.extract_text_for_routing(), "a");
|
assert_eq!(req.extract_text_for_routing(), "a");
|
||||||
|
|||||||
Reference in New Issue
Block a user