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
69 changes: 69 additions & 0 deletions crates/grpc_client/src/sglang_scheduler.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ use openai_protocol::{
chat::ChatCompletionRequest,
common::{ResponseFormat, StringOrArray, ToolChoice, ToolChoiceValue},
generate::GenerateRequest,
messages::CreateMessageRequest,
responses::ResponsesRequest,
sampling_params::SamplingParams as GenerateSamplingParams,
};
Expand Down Expand Up @@ -607,6 +608,74 @@ impl SglangSchedulerClient {
}
}

/// Build a GenerateRequest from CreateMessageRequest (Anthropic Messages API)
#[expect(
clippy::unused_self,
reason = "method receiver kept for consistent public API"
)]
pub fn build_generate_request_from_messages(
&self,
request_id: String,
body: &CreateMessageRequest,
processed_text: String,
token_ids: Vec<u32>,
multimodal_inputs: Option<proto::MultimodalInputs>,
tool_call_constraint: Option<(String, String)>,
) -> Result<proto::GenerateRequest, String> {
let sampling_params =
Self::build_grpc_sampling_params_from_messages(body, tool_call_constraint)?;

let grpc_request = proto::GenerateRequest {
request_id,
tokenized: Some(proto::TokenizedInput {
original_text: processed_text,
input_ids: token_ids,
}),
mm_inputs: multimodal_inputs,
sampling_params: Some(sampling_params),
return_logprob: false,
logprob_start_len: -1,
top_logprobs_num: 0,
return_hidden_states: false,
stream: body.stream.unwrap_or(false),
..Default::default()
};

Ok(grpc_request)
}

/// Build gRPC SamplingParams from CreateMessageRequest
fn build_grpc_sampling_params_from_messages(
request: &CreateMessageRequest,
tool_call_constraint: Option<(String, String)>,
) -> Result<proto::SamplingParams, String> {
let stop_sequences = request.stop_sequences.clone().unwrap_or_default();

// skip_special_tokens: false when tools are present (same logic as chat)
let skip_special_tokens =
tool_call_constraint.is_none() && request.tools.as_ref().is_none_or(|t| t.is_empty());
Comment on lines +655 to +656

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Honor tool_choice none when setting skip_special_tokens

This logic treats any non-empty request.tools as a signal to keep special tokens, even when tool use is effectively disabled (for example tool_choice: none, which yields no tool constraint) or when tools were filtered out earlier for gRPC use. In those cases generation is plain text, but skip_special_tokens is forced to false, which can leak backend control/special tokens into user-visible output; the same pattern is also present in the new vLLM Messages builder.

Useful? React with 👍 / 👎.


Ok(proto::SamplingParams {
temperature: request.temperature.unwrap_or(1.0) as f32,
top_p: request.top_p.unwrap_or(1.0) as f32,
top_k: request.top_k.map(|v| v as i32).unwrap_or(-1),

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Validate top_k before casting to signed backend fields

CreateMessageRequest.top_k is an unsigned integer, but this cast writes it into a signed i32 protobuf field with as; values above i32::MAX will wrap to negative numbers and silently change sampling behavior (including accidental disablement/invalid values). This should be range-checked before conversion; the same lossy cast is also introduced in the new TRT-LLM Messages sampling config path.

Useful? React with 👍 / 👎.

min_p: 0.0,
frequency_penalty: 0.0,
presence_penalty: 0.0,
repetition_penalty: 1.0,
max_new_tokens: Some(request.max_tokens),
stop: stop_sequences,
stop_token_ids: vec![],
skip_special_tokens,
spaces_between_special_tokens: true,
ignore_eos: false,
no_stop_trim: false,
n: 1,
constraint: Self::build_constraint_for_responses(tool_call_constraint)?,
..Default::default()
})
}

fn build_single_constraint_from_plain(
params: &GenerateSamplingParams,
) -> Result<Option<proto::sampling_params::Constraint>, String> {
Expand Down
90 changes: 90 additions & 0 deletions crates/grpc_client/src/trtllm_service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ use openai_protocol::{
chat::ChatCompletionRequest,
common::{ResponseFormat, StringOrArray},
generate::GenerateRequest,
messages::CreateMessageRequest,
responses::ResponsesRequest,
sampling_params::SamplingParams as GenerateSamplingParams,
};
Expand Down Expand Up @@ -622,6 +623,95 @@ impl TrtllmServiceClient {
}
}

/// Build a GenerateRequest from CreateMessageRequest (Anthropic Messages API)
#[expect(
clippy::unused_self,
reason = "method receiver kept for consistent public API"
)]
pub fn build_generate_request_from_messages(
&self,
request_id: String,
body: &CreateMessageRequest,
processed_text: String,
token_ids: Vec<u32>,
multimodal_input: Option<proto::MultimodalInput>,
tool_call_constraint: Option<(String, String)>,
) -> Result<proto::GenerateRequest, String> {
let sampling_config = Self::build_sampling_config_from_messages(body);
let output_config = proto::OutputConfig {
logprobs: None,
prompt_logprobs: None,
return_context_logits: false,
return_generation_logits: false,
exclude_input_from_output: true,
return_encoder_output: false,
return_perf_metrics: false,
};

let guided_decoding = Self::build_guided_decoding_from_responses(tool_call_constraint)?;

let stop = body.stop_sequences.clone().unwrap_or_default();
let max_tokens = body.max_tokens;

let grpc_request = proto::GenerateRequest {
request_id,
tokenized: Some(proto::TokenizedInput {
original_text: processed_text,
input_token_ids: token_ids,
query_token_ids: vec![],
}),
sampling_config: Some(sampling_config),
output_config: Some(output_config),
max_tokens,
streaming: body.stream.unwrap_or(false),
stop,
stop_token_ids: vec![],
ignore_eos: false,
bad: vec![],
bad_token_ids: vec![],
guided_decoding,
embedding_bias: vec![],
lora_config: None,
prompt_tuning_config: None,
multimodal_input,
kv_cache_retention: None,
disaggregated_params: None,
lookahead_config: None,
cache_salt_id: None,
arrival_time: None,
};

Ok(grpc_request)
}

/// Build SamplingConfig from CreateMessageRequest
fn build_sampling_config_from_messages(
request: &CreateMessageRequest,
) -> proto::SamplingConfig {
proto::SamplingConfig {
beam_width: 1,
num_return_sequences: 1,
top_k: request.top_k.map(|v| v as i32),
top_p: Some(request.top_p.unwrap_or(1.0) as f32),
top_p_min: None,
top_p_reset_ids: None,
top_p_decay: None,
seed: None,
temperature: Some(request.temperature.unwrap_or(1.0) as f32),
min_tokens: None,
beam_search_diversity_rate: None,
repetition_penalty: Some(1.0),
presence_penalty: None,
frequency_penalty: None,
prompt_ignore_length: None,
length_penalty: None,
early_stopping: None,
no_repeat_ngram_size: None,
min_p: None,
beam_width_array: vec![],
}
}

fn build_sampling_config_from_plain(
params: Option<&GenerateSamplingParams>,
) -> proto::SamplingConfig {
Expand Down
67 changes: 67 additions & 0 deletions crates/grpc_client/src/vllm_engine.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ use openai_protocol::{
chat::ChatCompletionRequest,
common::{ResponseFormat, StringOrArray, ToolChoice, ToolChoiceValue},
generate::GenerateRequest,
messages::CreateMessageRequest,
responses::ResponsesRequest,
sampling_params::SamplingParams as GenerateSamplingParams,
};
Expand Down Expand Up @@ -540,6 +541,72 @@ impl VllmEngineClient {
}
}

/// Build a GenerateRequest from CreateMessageRequest (Anthropic Messages API)
#[expect(
clippy::unused_self,
reason = "method receiver kept for consistent public API across gRPC backends"
)]
pub fn build_generate_request_from_messages(
&self,
request_id: String,
body: &CreateMessageRequest,
processed_text: String,
token_ids: Vec<u32>,
multimodal_inputs: Option<proto::MultimodalInputs>,
tool_call_constraint: Option<(String, String)>,
) -> Result<proto::GenerateRequest, String> {
let sampling_params =
Self::build_grpc_sampling_params_from_messages(body, tool_call_constraint)?;

let grpc_request = proto::GenerateRequest {
request_id,
input: Some(proto::generate_request::Input::Tokenized(
proto::TokenizedInput {
original_text: processed_text,
input_ids: token_ids,
},
)),
sampling_params: Some(sampling_params),
stream: body.stream.unwrap_or(false),
kv_transfer_params: None,
mm_inputs: multimodal_inputs,
};

Ok(grpc_request)
}

/// Build gRPC SamplingParams from CreateMessageRequest
fn build_grpc_sampling_params_from_messages(
request: &CreateMessageRequest,
tool_call_constraint: Option<(String, String)>,
) -> Result<proto::SamplingParams, String> {
let stop_sequences = request.stop_sequences.clone().unwrap_or_default();

// skip_special_tokens: false when tools are present (same logic as chat)
let skip_special_tokens =
tool_call_constraint.is_none() && request.tools.as_ref().is_none_or(|t| t.is_empty());

Ok(proto::SamplingParams {
temperature: Some(request.temperature.unwrap_or(1.0) as f32),
top_p: request.top_p.unwrap_or(1.0) as f32,
top_k: request.top_k.unwrap_or(0), // 0 means disabled in vLLM
min_p: 0.0,
frequency_penalty: 0.0,
presence_penalty: 0.0,
repetition_penalty: 1.0,
max_tokens: Some(request.max_tokens),
stop: stop_sequences,
stop_token_ids: vec![],
skip_special_tokens,
spaces_between_special_tokens: true,
ignore_eos: false,
n: 1,
logprobs: None,
constraint: Self::build_constraint_for_responses(tool_call_constraint)?,
..Default::default()
})
}

fn build_single_constraint_from_plain(
params: &GenerateSamplingParams,
) -> Result<Option<proto::sampling_params::Constraint>, String> {
Expand Down
65 changes: 64 additions & 1 deletion model_gateway/src/routers/grpc/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,8 @@
use std::collections::HashMap;

use openai_protocol::{
chat::ChatCompletionRequest, generate::GenerateRequest, worker::WorkerLoadResponse,
chat::ChatCompletionRequest, generate::GenerateRequest, messages::CreateMessageRequest,
worker::WorkerLoadResponse,
};
use smg_grpc_client::{
tokenizer_bundle, tokenizer_bundle::StreamBundle, SglangSchedulerClient, TrtllmServiceClient,
Expand Down Expand Up @@ -321,6 +322,68 @@ impl GrpcClient {
}
}

#[expect(
clippy::unreachable,
reason = "assembly stage guarantees matching MultimodalData variant for each backend"
)]
pub fn build_messages_request(
&self,
request_id: String,
body: &CreateMessageRequest,
processed_text: String,
token_ids: Vec<u32>,
multimodal_inputs: Option<MultimodalData>,
tool_constraints: Option<(String, String)>,
) -> Result<ProtoGenerateRequest, String> {
match self {
Self::Sglang(client) => {
let sglang_mm = multimodal_inputs.map(|mm| match mm {
MultimodalData::Sglang(data) => data.into_proto(),
_ => unreachable!("caller guarantees matching variant"),
});
let req = client.build_generate_request_from_messages(
request_id,
body,
processed_text,
token_ids,
sglang_mm,
tool_constraints,
)?;
Ok(ProtoGenerateRequest::Sglang(Box::new(req)))
}
Self::Vllm(client) => {
let vllm_mm = multimodal_inputs.map(|mm| match mm {
MultimodalData::Vllm(data) => data.into_proto(),
_ => unreachable!("caller guarantees matching variant"),
});
let req = client.build_generate_request_from_messages(
request_id,
body,
processed_text,
token_ids,
vllm_mm,
tool_constraints,
)?;
Ok(ProtoGenerateRequest::Vllm(Box::new(req)))
}
Self::Trtllm(client) => {
let trtllm_mm = multimodal_inputs.map(|mm| match mm {
MultimodalData::Trtllm(data) => data.into_proto(),
_ => unreachable!("caller guarantees matching variant"),
});
let req = client.build_generate_request_from_messages(
request_id,
body,
processed_text,
token_ids,
trtllm_mm,
tool_constraints,
)?;
Ok(ProtoGenerateRequest::Trtllm(Box::new(req)))
}
}
}
Comment thread
slin1237 marked this conversation as resolved.

pub fn build_generate_request(
&self,
request_id: String,
Expand Down
Original file line number Diff line number Diff line change
@@ -1,9 +1,12 @@
//! Messages API endpoint pipeline stages
//!
//! These stages handle Messages API-specific preprocessing.
//! Request building and response processing will be added in follow-up PRs.
//! These stages handle Messages API-specific preprocessing and request building.
//! Response processing will be added in a follow-up PR.

mod preparation;
mod request_building;

#[expect(unused_imports, reason = "wired in follow-up PR (pipeline factory)")]
pub(crate) use preparation::MessagePreparationStage;
#[expect(unused_imports, reason = "wired in follow-up PR (pipeline factory)")]
pub(crate) use request_building::MessageRequestBuildingStage;
Loading
Loading