Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -22,12 +22,6 @@ impl EmbeddingPreparationStage {
}
}

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> {
Expand Down
24 changes: 11 additions & 13 deletions model_gateway/src/routers/grpc/regular/stages/preparation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,7 @@ use async_trait::async_trait;
use axum::response::Response;
use tracing::error;

use super::{
chat::ChatPreparationStage, embedding::preparation::EmbeddingPreparationStage,
generate::GeneratePreparationStage,
};
use super::{chat::ChatPreparationStage, generate::GeneratePreparationStage};
use crate::routers::{
error as grpc_error,
grpc::{
Expand All @@ -23,15 +20,13 @@ use crate::routers::{
pub(crate) struct PreparationStage {
chat_stage: ChatPreparationStage,
generate_stage: GeneratePreparationStage,
embedding_stage: EmbeddingPreparationStage,
}

impl PreparationStage {
pub fn new() -> Self {
Self {
chat_stage: ChatPreparationStage,
generate_stage: GeneratePreparationStage,
embedding_stage: EmbeddingPreparationStage::new(),
}
}
}
Expand All @@ -48,17 +43,20 @@ impl PipelineStage for PreparationStage {
match &ctx.input.request_type {
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(_) => {
other => {
let type_name = match other {
RequestType::Embedding(_) => "Embedding",
RequestType::Classify(_) => "Classify",
RequestType::Responses(_) => "Responses",
_ => "Unknown",
};
error!(
function = "PreparationStage::execute",
"RequestType::Responses reached regular preparation stage"
"RequestType::{type_name} reached regular preparation stage"
);
Err(grpc_error::internal_error(
"responses_in_wrong_pipeline",
"RequestType::Responses reached regular preparation stage",
"wrong_pipeline",
format!("RequestType::{type_name} should use its dedicated pipeline"),
))
}
}
Expand Down
Loading