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
4 changes: 0 additions & 4 deletions model_gateway/src/routers/grpc/context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
13 changes: 13 additions & 0 deletions model_gateway/src/routers/grpc/regular/stages/completion/mod.rs
Original file line number Diff line number Diff line change
@@ -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;
Original file line number Diff line number Diff line change
@@ -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<Option<Response>, 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"
}
}
1 change: 1 addition & 0 deletions model_gateway/src/routers/grpc/regular/stages/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
10 changes: 8 additions & 2 deletions model_gateway/src/routers/grpc/regular/stages/preparation.rs
Original file line number Diff line number Diff line change
@@ -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::{
Expand All @@ -20,13 +23,15 @@ use crate::routers::{
pub(crate) struct PreparationStage {
chat_stage: ChatPreparationStage,
generate_stage: GeneratePreparationStage,
completion_stage: CompletionPreparationStage,
}

impl PreparationStage {
pub fn new() -> Self {
Self {
chat_stage: ChatPreparationStage,
generate_stage: GeneratePreparationStage,
completion_stage: CompletionPreparationStage,
}
}
}
Expand All @@ -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",
Expand Down
Loading