diff --git a/model_gateway/src/routers/grpc/context.rs b/model_gateway/src/routers/grpc/context.rs index 3b9dea48ae..2ba56cd646 100644 --- a/model_gateway/src/routers/grpc/context.rs +++ b/model_gateway/src/routers/grpc/context.rs @@ -425,10 +425,6 @@ impl RequestContext { } /// Get Arc clone of completion request (panics if not completion) - #[expect( - dead_code, - reason = "Arc accessor is introduced before later stacked PRs use it from completion stages" - )] #[expect( clippy::panic, reason = "typed accessor: caller guarantees variant via RequestType construction" diff --git a/model_gateway/src/routers/grpc/regular/stages/completion/mod.rs b/model_gateway/src/routers/grpc/regular/stages/completion/mod.rs new file mode 100644 index 0000000000..04d8746132 --- /dev/null +++ b/model_gateway/src/routers/grpc/regular/stages/completion/mod.rs @@ -0,0 +1,13 @@ +//! Completion API endpoint pipeline stages +//! +//! This module is the next stacked step after the native completion typing work +//! landed in PR #840. That PR introduced `RequestType::Completion`, +//! `FinalResponse::Completion`, and `execute_completion()`. This branch begins +//! the endpoint-specific stage stack with `CompletionPreparationStage`. +//! +//! Later follow-up PRs can add completion-specific request building and +//! response processing here, similar to the Messages API pipeline. + +mod preparation; + +pub(crate) use preparation::CompletionPreparationStage; diff --git a/model_gateway/src/routers/grpc/regular/stages/completion/preparation.rs b/model_gateway/src/routers/grpc/regular/stages/completion/preparation.rs new file mode 100644 index 0000000000..8ed7cfb392 --- /dev/null +++ b/model_gateway/src/routers/grpc/regular/stages/completion/preparation.rs @@ -0,0 +1,79 @@ +//! Completion preparation stage: resolve prompt, tokenize, create stop decoder. +//! +//! This is the `/v1/completions` Stage 1 equivalent. It intentionally builds on top of +//! the native completion pipeline typing introduced in PR #840. It keeps +//! `CompletionRequest` native in the request context instead of laundering it +//! through `GenerateRequest`. + +use async_trait::async_trait; +use axum::response::Response; +use openai_protocol::common::StringOrArray; +use tracing::error; + +use crate::routers::{ + error, + grpc::{ + common::stages::PipelineStage, + context::{PreparationOutput, RequestContext}, + utils, + }, +}; + +pub(crate) struct CompletionPreparationStage; + +#[async_trait] +impl PipelineStage for CompletionPreparationStage { + async fn execute(&self, ctx: &mut RequestContext) -> Result, Response> { + let request = ctx.completion_request_arc(); + + let tokenizer = + utils::resolve_tokenizer(ctx, "CompletionPreparationStage::execute").map_err(|e| *e)?; + + let prompt_text = match &request.prompt { + StringOrArray::String(text) => text.clone(), + StringOrArray::Array(_) => { + return Err(error::bad_request( + "batch_prompts_not_supported", + "Batched prompt arrays are not supported for gRPC /v1/completions yet", + )); + } + }; + + let encoding = tokenizer.encode(&prompt_text, false).map_err(|e| { + error!( + function = "CompletionPreparationStage::execute", + error = %e, + "Tokenization failed" + ); + error::bad_request("tokenization_failed", format!("Tokenization failed: {e}")) + })?; + + let stop_decoder = utils::create_stop_decoder( + &tokenizer, + request.stop.as_ref(), + request.stop_token_ids.as_ref(), + request.skip_special_tokens, + request.no_stop_trim, + ); + + ctx.state.preparation = Some(PreparationOutput { + original_text: Some(prompt_text), + token_ids: encoding.token_ids().to_vec(), + processed_messages: None, + tool_constraints: None, + filtered_request: None, + harmony_mode: false, + selection_text: None, + harmony_messages: None, + harmony_stop_ids: None, + }); + + ctx.state.response.stop_decoder = Some(stop_decoder); + + Ok(None) + } + + fn name(&self) -> &'static str { + "CompletionPreparation" + } +} diff --git a/model_gateway/src/routers/grpc/regular/stages/mod.rs b/model_gateway/src/routers/grpc/regular/stages/mod.rs index 2c8f43e404..c7f6f1841d 100644 --- a/model_gateway/src/routers/grpc/regular/stages/mod.rs +++ b/model_gateway/src/routers/grpc/regular/stages/mod.rs @@ -4,6 +4,7 @@ pub(crate) mod chat; pub(crate) mod classify; +pub(crate) mod completion; pub(crate) mod embedding; pub(crate) mod generate; pub(crate) mod messages; diff --git a/model_gateway/src/routers/grpc/regular/stages/preparation.rs b/model_gateway/src/routers/grpc/regular/stages/preparation.rs index 1823f11276..95a9679372 100644 --- a/model_gateway/src/routers/grpc/regular/stages/preparation.rs +++ b/model_gateway/src/routers/grpc/regular/stages/preparation.rs @@ -1,13 +1,16 @@ //! Preparation stage that delegates to endpoint-specific implementations //! //! This stage checks RequestType at runtime and delegates to the appropriate -//! endpoint-specific stage (ChatPreparationStage or GeneratePreparationStage). +//! endpoint-specific stage. (ChatPreparationStage, CompletionPreparationStage or GeneratePreparationStage). use async_trait::async_trait; use axum::response::Response; use tracing::error; -use super::{chat::ChatPreparationStage, generate::GeneratePreparationStage}; +use super::{ + chat::ChatPreparationStage, completion::CompletionPreparationStage, + generate::GeneratePreparationStage, +}; use crate::routers::{ error as grpc_error, grpc::{ @@ -20,6 +23,7 @@ use crate::routers::{ pub(crate) struct PreparationStage { chat_stage: ChatPreparationStage, generate_stage: GeneratePreparationStage, + completion_stage: CompletionPreparationStage, } impl PreparationStage { @@ -27,6 +31,7 @@ impl PreparationStage { Self { chat_stage: ChatPreparationStage, generate_stage: GeneratePreparationStage, + completion_stage: CompletionPreparationStage, } } } @@ -43,6 +48,7 @@ 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::Completion(_) => self.completion_stage.execute(ctx).await, request_type => { error!( function = "PreparationStage::execute",