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
8 changes: 4 additions & 4 deletions model_gateway/src/routers/grpc/context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ use super::{
EncodeItemBootstrapInfo, ProtoEmbedComplete, ProtoEmbedRequest, ProtoGenerateRequest,
ProtoRequest, ProtoStream,
},
utils::ParserResolver,
};
use crate::{
middleware::TenantRequestMeta,
Expand Down Expand Up @@ -144,10 +145,9 @@ pub(crate) struct SharedComponents {
pub worker_registry: Arc<WorkerRegistry>,
pub tool_parser_factory: ToolParserFactory,
pub reasoning_parser_factory: ReasoningParserFactory,
/// Configured tool parser name (from CLI `--tool-call-parser`)
pub configured_tool_parser: Option<String>,
/// Configured reasoning parser name (from CLI `--reasoning-parser`)
pub configured_reasoning_parser: Option<String>,
/// Per-request parser-name resolution (model-card override → configured
/// CLI `--tool-call-parser`/`--reasoning-parser` names).
pub parser_resolver: ParserResolver,
/// Multimodal processing components (initialized at router creation)
pub multimodal: Option<Arc<MultimodalComponents>>,
}
Expand Down
24 changes: 12 additions & 12 deletions model_gateway/src/routers/grpc/pipeline.rs
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@ use super::{
},
streaming,
},
utils,
utils::error_type_from_status,
};
use crate::{
Expand Down Expand Up @@ -138,17 +139,20 @@ impl PipelineDeps {
processor::ResponseProcessor,
Arc<streaming::StreamingProcessor>,
) {
let parser_resolver = utils::ParserResolver::new(
self.worker_registry.clone(),
self.configured_tool_parser.clone(),
self.configured_reasoning_parser.clone(),
);
let processor = processor::ResponseProcessor::new(
self.tool_parser_factory.clone(),
self.reasoning_parser_factory.clone(),
self.configured_tool_parser.clone(),
self.configured_reasoning_parser.clone(),
parser_resolver.clone(),
);
let streaming_processor = Arc::new(streaming::StreamingProcessor::new(
self.tool_parser_factory.clone(),
self.reasoning_parser_factory.clone(),
self.configured_tool_parser.clone(),
self.configured_reasoning_parser.clone(),
parser_resolver,
backend,
));
(processor, streaming_processor)
Expand All @@ -165,14 +169,12 @@ impl PipelineDeps {
let processor = processor::ResponseProcessor::new(
ToolParserFactory::default(),
ReasoningParserFactory::default(),
None,
None,
utils::ParserResolver::disabled(),
);
let streaming_processor = Arc::new(streaming::StreamingProcessor::new(
ToolParserFactory::default(),
ReasoningParserFactory::default(),
None,
None,
utils::ParserResolver::disabled(),
backend,
));
(processor, streaming_processor)
Expand Down Expand Up @@ -1570,8 +1572,7 @@ mod alias_pipeline_tests {
worker_registry,
tool_parser_factory: ToolParserFactory::default(),
reasoning_parser_factory: ReasoningParserFactory::default(),
configured_tool_parser: None,
configured_reasoning_parser: None,
parser_resolver: utils::ParserResolver::disabled(),
multimodal: None,
});
let request: GenerateRequest = serde_json::from_value(json!({
Expand Down Expand Up @@ -1644,8 +1645,7 @@ mod rate_limit_reserve_tests {
worker_registry,
tool_parser_factory: ToolParserFactory::default(),
reasoning_parser_factory: ReasoningParserFactory::default(),
configured_tool_parser: None,
configured_reasoning_parser: None,
parser_resolver: utils::ParserResolver::disabled(),
multimodal: None,
})
}
Expand Down
52 changes: 32 additions & 20 deletions model_gateway/src/routers/grpc/regular/processor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,22 +38,20 @@ use crate::routers::{
pub(crate) struct ResponseProcessor {
pub tool_parser_factory: ToolParserFactory,
pub reasoning_parser_factory: ReasoningParserFactory,
pub configured_tool_parser: Option<String>,
pub configured_reasoning_parser: Option<String>,
/// Per-request parser-name resolution (model-card override → configured).
pub parser_resolver: utils::ParserResolver,
}

impl ResponseProcessor {
pub fn new(
tool_parser_factory: ToolParserFactory,
reasoning_parser_factory: ReasoningParserFactory,
configured_tool_parser: Option<String>,
configured_reasoning_parser: Option<String>,
parser_resolver: utils::ParserResolver,
) -> Self {
Self {
tool_parser_factory,
reasoning_parser_factory,
configured_tool_parser,
configured_reasoning_parser,
parser_resolver,
}
}

Expand All @@ -69,6 +67,11 @@ impl ResponseProcessor {
history_tool_calls_count: usize,
reasoning_parser_available: bool,
tool_parser_available: bool,
// Resolved once per request by the caller: keeps every choice of one
// request on the same parser even if the worker registry changes
// between availability check and parsing.
reasoning_parser_name: Option<&str>,
tool_parser_name: Option<&str>,
) -> Result<ChatChoice, String> {
stop_decoder.reset();
// Decode tokens
Expand Down Expand Up @@ -109,7 +112,7 @@ impl ResponseProcessor {
// across requests, so avoid serializing on the shared pooled mutex.
if let Some(mut parser) = utils::create_reasoning_parser(
&self.reasoning_parser_factory,
self.configured_reasoning_parser.as_deref(),
reasoning_parser_name,
&original_request.model,
) {
// If the template injected `<think>` in the prefill (thinking toggle
Expand Down Expand Up @@ -151,7 +154,7 @@ impl ResponseProcessor {
let has_structural_tag = self
.tool_parser_factory
.registry()
.has_structural_tag_for_parser(self.configured_tool_parser.as_deref());
.has_structural_tag_for_parser(tool_parser_name);
let used_json_schema = if has_structural_tag {
false
} else {
Expand All @@ -175,6 +178,7 @@ impl ResponseProcessor {
.parse_tool_calls(
&processed_text,
&original_request.model,
tool_parser_name,
original_request.tools.as_deref().unwrap_or(&[]),
history_tool_calls_count,
)
Expand Down Expand Up @@ -241,6 +245,8 @@ impl ResponseProcessor {
stop_decoder: &mut StopSequenceDecoder,
request_logprobs: bool,
) -> Result<ChatCompletionResponse, axum::response::Response> {
let reasoning_parser_name = self.parser_resolver.reasoning_parser(&chat_request.model);
let tool_parser_name = self.parser_resolver.tool_parser(&chat_request.model);
// Collect all responses from the execution result
let all_responses =
response_collection::collect_responses(execution_result, request_logprobs).await?;
Expand All @@ -251,7 +257,7 @@ impl ResponseProcessor {
let reasoning_parser_available = chat_request.separate_reasoning
&& utils::check_reasoning_parser_availability(
&self.reasoning_parser_factory,
self.configured_reasoning_parser.as_deref(),
reasoning_parser_name.as_deref(),
&chat_request.model,
);

Expand All @@ -264,7 +270,7 @@ impl ResponseProcessor {
&& chat_request.tools.is_some()
&& utils::check_tool_parser_availability(
&self.tool_parser_factory,
self.configured_tool_parser.as_deref(),
tool_parser_name.as_deref(),
&chat_request.model,
);

Expand Down Expand Up @@ -296,6 +302,8 @@ impl ResponseProcessor {
history_tool_calls_count,
reasoning_parser_available,
tool_parser_available,
reasoning_parser_name.as_deref(),
tool_parser_name.as_deref(),
)
.await
{
Expand Down Expand Up @@ -328,15 +336,14 @@ impl ResponseProcessor {
&self,
processed_text: &str,
model: &str,
// Resolved once per request by the caller (see process_single_choice).
tool_parser_name: Option<&str>,
tools: &[Tool],
history_tool_calls_count: usize,
) -> (Option<Vec<ToolCall>>, String) {
// Get pooled parser for this model
let pooled_parser = utils::get_tool_parser(
&self.tool_parser_factory,
self.configured_tool_parser.as_deref(),
model,
);
let pooled_parser =
utils::get_tool_parser(&self.tool_parser_factory, tool_parser_name, model);

// Try parsing directly (parser will handle detection internally). Pass the
// tool schemas so schema-aware parsers coerce argument types by their
Expand Down Expand Up @@ -502,6 +509,10 @@ impl ResponseProcessor {
tokenizer: Arc<dyn Tokenizer>,
stop_decoder: &mut StopSequenceDecoder,
) -> Result<Message, axum::response::Response> {
let reasoning_parser_name = self
.parser_resolver
.reasoning_parser(&messages_request.model);
let tool_parser_name = self.parser_resolver.tool_parser(&messages_request.model);
// Collect all responses (no logprobs for Messages API)
let all_responses = response_collection::collect_responses(execution_result, false).await?;

Expand Down Expand Up @@ -537,7 +548,7 @@ impl ResponseProcessor {
// or when the selected parser needs structural special tokens (e.g. Inkling).
let reasoning_requires_special_tokens = utils::reasoning_parser_requires_special_tokens(
&self.reasoning_parser_factory,
self.configured_reasoning_parser.as_deref(),
reasoning_parser_name.as_deref(),
&messages_request.model,
);
let separate_reasoning = reasoning_requires_special_tokens
Expand All @@ -551,7 +562,7 @@ impl ResponseProcessor {
let reasoning_parser_available = separate_reasoning
&& utils::check_reasoning_parser_availability(
&self.reasoning_parser_factory,
self.configured_reasoning_parser.as_deref(),
reasoning_parser_name.as_deref(),
&messages_request.model,
);

Expand All @@ -564,7 +575,7 @@ impl ResponseProcessor {
&& messages_request.tools.is_some()
&& utils::check_tool_parser_availability(
&self.tool_parser_factory,
self.configured_tool_parser.as_deref(),
tool_parser_name.as_deref(),
&messages_request.model,
);

Expand Down Expand Up @@ -624,7 +635,7 @@ impl ResponseProcessor {
// across requests, so avoid serializing on the shared pooled mutex.
if let Some(mut parser) = utils::create_reasoning_parser(
&self.reasoning_parser_factory,
self.configured_reasoning_parser.as_deref(),
reasoning_parser_name.as_deref(),
&messages_request.model,
) {
// If thinking is effectively ON and template has a toggle, start in reasoning mode.
Expand Down Expand Up @@ -664,7 +675,7 @@ impl ResponseProcessor {
let has_structural_tag = self
.tool_parser_factory
.registry()
.has_structural_tag_for_parser(self.configured_tool_parser.as_deref());
.has_structural_tag_for_parser(tool_parser_name.as_deref());
let used_json_schema = !has_structural_tag
&& matches!(
&messages_request.tool_choice,
Expand Down Expand Up @@ -694,6 +705,7 @@ impl ResponseProcessor {
.parse_tool_calls(
&processed_text,
&messages_request.model,
tool_parser_name.as_deref(),
&chat_tools,
utils::message_utils::get_history_tool_calls_count_messages(
&messages_request,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -239,7 +239,10 @@ impl ChatPreparationStage {
.tool_parser_factory
.registry()
.generate_tool_constraint(
ctx.components.configured_tool_parser.as_deref(),
ctx.components
.parser_resolver
.tool_parser(&request.model)
.as_deref(),
tools,
tool_choice,
)
Expand All @@ -257,7 +260,10 @@ impl ChatPreparationStage {
let preserve_reasoning_special_tokens = request.separate_reasoning
&& utils::reasoning_parser_requires_special_tokens(
&ctx.components.reasoning_parser_factory,
ctx.components.configured_reasoning_parser.as_deref(),
ctx.components
.parser_resolver
.reasoning_parser(&request.model)
.as_deref(),
&request.model,
);

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -267,7 +267,10 @@ impl MessagePreparationStage {
.tool_parser_factory
.registry()
.generate_tool_constraint(
ctx.components.configured_tool_parser.as_deref(),
ctx.components
.parser_resolver
.tool_parser(&request.model)
.as_deref(),
&filtered_tools,
tool_choice,
)
Expand All @@ -290,7 +293,10 @@ impl MessagePreparationStage {

let preserve_reasoning_special_tokens = utils::reasoning_parser_requires_special_tokens(
&ctx.components.reasoning_parser_factory,
ctx.components.configured_reasoning_parser.as_deref(),
ctx.components
.parser_resolver
.reasoning_parser(&request.model)
.as_deref(),
&request.model,
);

Expand Down
Loading
Loading