Skip to content
Closed
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
2 changes: 1 addition & 1 deletion crates/data_connector/src/memory.rs
Original file line number Diff line number Diff line change
Expand Up @@ -435,7 +435,7 @@ impl ResponseStorage for MemoryResponseStorage {
.collect();

// Sort by creation time (newest first)
responses_with_time.sort_by(|a, b| b.0.cmp(&a.0));
responses_with_time.sort_by_key(|(created_at, _)| std::cmp::Reverse(*created_at));

// Apply limit and collect the actual responses
let limit = limit.unwrap_or(responses_with_time.len());
Expand Down
2 changes: 1 addition & 1 deletion crates/mcp/src/core/metrics.rs
Original file line number Diff line number Diff line change
Expand Up @@ -206,7 +206,7 @@ impl LatencyStats {

LatencySnapshot {
count,
avg_ms: if count > 0 { total / count } else { 0 },
avg_ms: total.checked_div(count).unwrap_or(0),
min_ms: if min == u64::MAX { 0 } else { min },
max_ms: max,
}
Expand Down
16 changes: 8 additions & 8 deletions crates/protocols/src/builders/responses/response.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ pub struct ResponsesResponseBuilder {
store: bool,
temperature: Option<f32>,
text: Option<TextConfig>,
tool_choice: String,
tool_choice: Value,
tools: Vec<ResponseTool>,
top_p: Option<f32>,
truncation: Option<String>,
Expand Down Expand Up @@ -64,7 +64,7 @@ impl ResponsesResponseBuilder {
store: true,
temperature: None,
text: None,
tool_choice: "auto".to_string(),
tool_choice: serde_json::json!("auto"),
tools: Vec::new(),
top_p: None,
truncation: None,
Expand All @@ -91,11 +91,11 @@ impl ResponsesResponseBuilder {
.clone_from(&request.previous_response_id);
self.store = request.store.unwrap_or(true);
self.temperature = request.temperature;
self.tool_choice = if let Some(ref tc) = request.tool_choice {
serde_json::to_string(tc).unwrap_or_else(|_| "auto".to_string())
} else {
"auto".to_string()
};
self.tool_choice = request
.tool_choice
.as_ref()
.map(responses_tool_choice_value)
.unwrap_or_else(|| serde_json::json!("auto"));
self.tools = request.tools.clone().unwrap_or_default();
self.top_p = request.top_p;
self.user.clone_from(&request.user);
Expand Down Expand Up @@ -196,7 +196,7 @@ impl ResponsesResponseBuilder {
}

/// Set tool choice setting
pub fn tool_choice(mut self, tool_choice: impl Into<String>) -> Self {
pub fn tool_choice(mut self, tool_choice: impl Into<Value>) -> Self {
self.tool_choice = tool_choice.into();
self
}
Expand Down
77 changes: 71 additions & 6 deletions crates/protocols/src/responses.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,9 @@ use validator::{Validate, ValidationError};

use super::{
common::{
default_model, default_true, validate_stop, ChatLogProbs, Function, GenerationRequest,
PromptTokenUsageInfo, StringOrArray, ToolChoice, ToolChoiceValue, ToolReference, UsageInfo,
default_model, default_true, validate_stop, ChatLogProbs, Function, FunctionChoice,
GenerationRequest, PromptTokenUsageInfo, StringOrArray, ToolChoice, ToolChoiceValue,
ToolReference, UsageInfo,
},
sampling_params::{validate_top_k_value, validate_top_p_value},
};
Expand Down Expand Up @@ -708,7 +709,12 @@ pub struct ResponsesRequest {
pub temperature: Option<f32>,

/// Tool choice behavior
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(
default,
deserialize_with = "deserialize_responses_tool_choice",
serialize_with = "serialize_responses_tool_choice",
skip_serializing_if = "Option::is_none"
)]
pub tool_choice: Option<ToolChoice>,

/// Available tools
Expand Down Expand Up @@ -934,6 +940,65 @@ impl GenerationRequest for ResponsesRequest {
}
}

pub fn responses_tool_choice_value(tool_choice: &ToolChoice) -> Value {
match tool_choice {
ToolChoice::Function { function, .. } => serde_json::json!({
"type": "function",
"name": function.name,
}),
_ => serde_json::to_value(tool_choice).unwrap_or_else(|_| serde_json::json!("auto")),
}
}

#[expect(
clippy::ref_option,
reason = "serde serialize_with passes the field as &Option<T>"
)]
fn serialize_responses_tool_choice<S>(
tool_choice: &Option<ToolChoice>,
serializer: S,
) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match tool_choice {
Some(tool_choice) => responses_tool_choice_value(tool_choice).serialize(serializer),
None => serializer.serialize_none(),
}
}

fn deserialize_responses_tool_choice<'de, D>(
deserializer: D,
) -> Result<Option<ToolChoice>, D::Error>
where
D: serde::Deserializer<'de>,
{
let Some(value) = Option::<Value>::deserialize(deserializer)? else {
return Ok(None);
};

if value.is_null() {
return Ok(None);
}

if let Some(name) = value
.as_object()
.filter(|obj| obj.get("type").and_then(Value::as_str) == Some("function"))
.and_then(|obj| obj.get("name").and_then(Value::as_str))
{
return Ok(Some(ToolChoice::Function {
tool_type: "function".to_string(),
function: FunctionChoice {
name: name.to_string(),
},
}));
}

serde_json::from_value(value)
.map(Some)
.map_err(serde::de::Error::custom)
}

/// Validate conversation ID format
pub fn validate_conversation_id(conv_id: &str) -> Result<(), ValidationError> {
if !conv_id.starts_with("conv_") {
Expand Down Expand Up @@ -1338,7 +1403,7 @@ pub struct ResponsesResponse {

/// Tool choice setting
#[serde(default = "default_tool_choice")]
pub tool_choice: String,
pub tool_choice: Value,

/// Available tools
#[serde(default)]
Expand Down Expand Up @@ -1368,8 +1433,8 @@ fn default_object_type() -> String {
"response".to_string()
}

fn default_tool_choice() -> String {
"auto".to_string()
fn default_tool_choice() -> Value {
serde_json::json!("auto")
}

impl ResponsesResponse {
Expand Down
6 changes: 1 addition & 5 deletions crates/tokenizer/src/cache/fingerprint.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,11 +42,7 @@ impl TokenizerFingerprint {

// Sample up to 1000 tokens for speed
let sample_size = vocab_size.min(1000);
let step = if sample_size > 0 {
vocab_size / sample_size
} else {
1
};
let step = vocab_size.checked_div(sample_size).unwrap_or(1);

for i in (0..vocab_size).step_by(step.max(1)) {
if let Some(token) = tokenizer.id_to_token(i as u32) {
Expand Down
18 changes: 8 additions & 10 deletions crates/tokenizer/src/chat_template.rs
Original file line number Diff line number Diff line change
Expand Up @@ -209,12 +209,11 @@ impl<'a> Detector<'a> {
self.inspect_expr_for_structure(&e.expr);
}
// {% set content = message.content %}
Stmt::Set(s) => {
Stmt::Set(s)
if Self::is_var_access(&s.target, "content")
&& self.is_any_scope_var_content(&s.expr)
{
self.flags.saw_assignment = true;
}
&& self.is_any_scope_var_content(&s.expr) =>
{
self.flags.saw_assignment = true;
}
Stmt::Macro(m) => {
// Heuristic: macro that checks type (via `is` test) and also has any loop
Expand All @@ -236,13 +235,12 @@ impl<'a> Detector<'a> {

match expr {
// content[0] or message.content[0]
Expr::GetItem(gi) => {
Expr::GetItem(gi)
if (matches!(&gi.expr, Expr::Var(v) if v.id == "content")
|| self.is_any_scope_var_content(&gi.expr))
&& Self::is_numeric_const(&gi.subscript_expr)
{
self.flags.saw_structure = true;
}
&& Self::is_numeric_const(&gi.subscript_expr) =>
{
self.flags.saw_structure = true;
}
// content|length or message.content|length
Expr::Filter(f) => {
Expand Down
2 changes: 1 addition & 1 deletion model_gateway/src/core/metrics_aggregator.rs
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ pub fn aggregate_metrics(metric_packs: Vec<MetricPack>) -> anyhow::Result<String
expositions.push(exposition);
}

let text = try_reduce(expositions.into_iter(), merge_exposition)?
let text = try_reduce(expositions, merge_exposition)?
.map(|x| format!("{x}"))
.unwrap_or_default();
Ok(text)
Expand Down
7 changes: 3 additions & 4 deletions model_gateway/src/policies/bucket.rs
Original file line number Diff line number Diff line change
Expand Up @@ -363,10 +363,7 @@ impl Bucket {
}

let worker_cnt = bucket_cnt;
let boundary = if worker_cnt == 0 {
Vec::new()
} else {
let gap = self.l_max / worker_cnt;
let boundary = if let Some(gap) = self.l_max.checked_div(worker_cnt) {
self.l_max = usize::MAX;
prefill_worker_urls
.iter()
Expand All @@ -381,6 +378,8 @@ impl Bucket {
Boundary::new(url.clone(), [min, max])
})
.collect()
} else {
Vec::new()
};

self.boundary = boundary;
Expand Down
49 changes: 47 additions & 2 deletions model_gateway/src/routers/grpc/common/responses/streaming.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,8 @@ use openai_protocol::{
McpEvent, OutputItemEvent, OutputTextEvent, ResponseEvent, WebSearchCallEvent,
},
responses::{
ResponseOutputItem, ResponseStatus, ResponsesRequest, ResponsesResponse, ResponsesUsage,
responses_tool_choice_value, ResponseOutputItem, ResponseStatus, ResponsesRequest,
ResponsesResponse, ResponsesUsage,
},
};
use serde_json::json;
Expand Down Expand Up @@ -340,7 +341,7 @@ impl ResponseStreamEventEmitter {

// tool_choice: serialize if present, otherwise use "auto"
if let Some(ref tc) = req.tool_choice {
response_obj["tool_choice"] = json!(tc);
response_obj["tool_choice"] = responses_tool_choice_value(tc);
} else {
response_obj["tool_choice"] = json!("auto");
}
Expand Down Expand Up @@ -1010,3 +1011,47 @@ pub(crate) fn attach_mcp_server_label(
item["server_label"] = json!(label);
}
}

#[cfg(test)]
mod tests {
use openai_protocol::{
common::{Function, FunctionChoice, ToolChoice},
responses::{FunctionTool, ResponseInput, ResponseTool},
};

use super::*;

#[test]
fn completed_event_serializes_flat_responses_function_tool_choice() {
let mut emitter =
ResponseStreamEventEmitter::new("resp_test".to_string(), "mock-model".to_string(), 1);
emitter.set_original_request(ResponsesRequest {
input: ResponseInput::Text("test".to_string()),
tools: Some(vec![ResponseTool::Function(FunctionTool {
function: Function {
name: "lookup_city".to_string(),
description: None,
parameters: json!({}),
strict: None,
},
})]),
tool_choice: Some(ToolChoice::Function {
tool_type: "function".to_string(),
function: FunctionChoice {
name: "lookup_city".to_string(),
},
}),
..Default::default()
});

let event = emitter.emit_completed(None);

assert_eq!(
event["response"]["tool_choice"],
json!({
"type": "function",
"name": "lookup_city"
})
);
}
}
8 changes: 3 additions & 5 deletions model_gateway/src/routers/openai/responses/accumulator.rs
Original file line number Diff line number Diff line change
Expand Up @@ -104,11 +104,9 @@ impl StreamingResponseAccumulator {
};

match get_event_type(event_name, &parsed) {
ResponseEvent::CREATED => {
if self.initial_response.is_none() {
if let Some(response) = parsed.get("response") {
self.initial_response = Some(response.clone());
}
ResponseEvent::CREATED if self.initial_response.is_none() => {
if let Some(response) = parsed.get("response") {
self.initial_response = Some(response.clone());
}
}
ResponseEvent::COMPLETED => {
Expand Down
55 changes: 55 additions & 0 deletions model_gateway/tests/api/api_endpoints_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -688,6 +688,61 @@ mod responses_endpoint_tests {
ctx.shutdown().await;
}

#[tokio::test]
async fn test_v1_responses_accepts_flat_forced_function_tool_choice() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 0,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}])
.await;

let app = ctx.create_app();

let payload = json!({
"input": "Run lookup_city.",
"model": "mock-model",
"stream": false,
"tools": [
{
"type": "function",
"name": "lookup_city",
"description": "Look up a city",
"parameters": {
"type": "object",
"properties": {
"city": { "type": "string" }
},
"required": ["city"]
}
}
],
"tool_choice": {
"type": "function",
"name": "lookup_city"
}
});

let req = Request::builder()
.method("POST")
.uri("/v1/responses")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();

let resp = app.clone().oneshot(req).await.unwrap();
assert_ne!(
resp.status(),
StatusCode::BAD_REQUEST,
"Responses flat function tool_choice should not be rejected by request validation"
);
assert_eq!(resp.status(), StatusCode::OK);

ctx.shutdown().await;
}

#[tokio::test]
async fn test_v1_responses_streaming() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
Expand Down
Loading