fix(llm): add stop_sequences parity for tool completions (#1170)

* fix(llm): add stop_sequences parity for tool completions

* refactor(web-openai): dedupe request builders and satisfy no-panics gate

* test(llm): mark multiline assert with safety comment for CI gate

* test(llm): make safety-marked assert formatting-stable
This commit is contained in:
Nige
2026-03-14 13:06:48 -07:00
committed by GitHub
parent cc52a046c1
commit ffe384b66e
7 changed files with 81 additions and 51 deletions
+42 -46
View File
@@ -419,6 +419,44 @@ fn parse_stop(val: &serde_json::Value) -> Option<Vec<String>> {
}
}
fn build_completion_request(
req: &OpenAiChatRequest,
messages: Vec<ChatMessage>,
) -> CompletionRequest {
let mut comp_req = CompletionRequest::new(messages).with_model(req.model.clone());
if let Some(t) = req.temperature {
comp_req = comp_req.with_temperature(t);
}
if let Some(mt) = req.max_tokens {
comp_req = comp_req.with_max_tokens(mt);
}
if let Some(stops) = req.stop.as_ref().and_then(parse_stop) {
comp_req.stop_sequences = Some(stops);
}
comp_req
}
fn build_tool_request(
req: &OpenAiChatRequest,
messages: Vec<ChatMessage>,
) -> ToolCompletionRequest {
let tools = convert_tools(req.tools.as_deref().unwrap_or(&[]));
let mut tool_req = ToolCompletionRequest::new(messages, tools).with_model(req.model.clone());
if let Some(t) = req.temperature {
tool_req = tool_req.with_temperature(t);
}
if let Some(mt) = req.max_tokens {
tool_req = tool_req.with_max_tokens(mt);
}
if let Some(stops) = req.stop.as_ref().and_then(parse_stop) {
tool_req = tool_req.with_stop_sequences(stops);
}
if let Some(choice) = req.tool_choice.as_ref().and_then(normalize_tool_choice) {
tool_req = tool_req.with_tool_choice(choice);
}
tool_req
}
// ---------------------------------------------------------------------------
// Handlers
// ---------------------------------------------------------------------------
@@ -476,19 +514,7 @@ pub async fn chat_completions_handler(
let created = unix_timestamp();
if has_tools {
let tools = convert_tools(req.tools.as_deref().unwrap_or(&[]));
let mut tool_req = ToolCompletionRequest::new(messages, tools).with_model(req.model);
if let Some(t) = req.temperature {
tool_req = tool_req.with_temperature(t);
}
if let Some(mt) = req.max_tokens {
tool_req = tool_req.with_max_tokens(mt);
}
if let Some(ref tc) = req.tool_choice
&& let Some(choice) = normalize_tool_choice(tc)
{
tool_req = tool_req.with_tool_choice(choice);
}
let tool_req = build_tool_request(&req, messages);
let resp = llm
.complete_with_tools(tool_req)
@@ -527,16 +553,7 @@ pub async fn chat_completions_handler(
Ok(Json(response).into_response())
} else {
let mut comp_req = CompletionRequest::new(messages).with_model(req.model);
if let Some(t) = req.temperature {
comp_req = comp_req.with_temperature(t);
}
if let Some(mt) = req.max_tokens {
comp_req = comp_req.with_max_tokens(mt);
}
if let Some(ref stop_val) = req.stop {
comp_req.stop_sequences = parse_stop(stop_val);
}
let comp_req = build_completion_request(&req, messages);
let resp = llm.complete(comp_req).await.map_err(map_llm_error)?;
let model_name = llm.effective_model_name(Some(requested_model.as_str()));
@@ -596,35 +613,14 @@ async fn handle_streaming(
}
let llm_result = if has_tools {
let tools = convert_tools(req.tools.as_deref().unwrap_or(&[]));
let mut tool_req = ToolCompletionRequest::new(messages, tools).with_model(req.model);
if let Some(t) = req.temperature {
tool_req = tool_req.with_temperature(t);
}
if let Some(mt) = req.max_tokens {
tool_req = tool_req.with_max_tokens(mt);
}
if let Some(ref tc) = req.tool_choice
&& let Some(choice) = normalize_tool_choice(tc)
{
tool_req = tool_req.with_tool_choice(choice);
}
let tool_req = build_tool_request(&req, messages);
LlmResult::WithTools(
llm.complete_with_tools(tool_req)
.await
.map_err(map_llm_error)?,
)
} else {
let mut comp_req = CompletionRequest::new(messages).with_model(req.model);
if let Some(t) = req.temperature {
comp_req = comp_req.with_temperature(t);
}
if let Some(mt) = req.max_tokens {
comp_req = comp_req.with_max_tokens(mt);
}
if let Some(ref stop_val) = req.stop {
comp_req.stop_sequences = parse_stop(stop_val);
}
let comp_req = build_completion_request(&req, messages);
LlmResult::Simple(llm.complete(comp_req).await.map_err(map_llm_error)?)
};
let model_name = llm.effective_model_name(Some(requested_model.as_str()));