mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-09-02 09:39:37 +00:00
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:
@@ -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
|
// Handlers
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
@@ -476,19 +514,7 @@ pub async fn chat_completions_handler(
|
|||||||
let created = unix_timestamp();
|
let created = unix_timestamp();
|
||||||
|
|
||||||
if has_tools {
|
if has_tools {
|
||||||
let tools = convert_tools(req.tools.as_deref().unwrap_or(&[]));
|
let tool_req = build_tool_request(&req, messages);
|
||||||
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 resp = llm
|
let resp = llm
|
||||||
.complete_with_tools(tool_req)
|
.complete_with_tools(tool_req)
|
||||||
@@ -527,16 +553,7 @@ pub async fn chat_completions_handler(
|
|||||||
|
|
||||||
Ok(Json(response).into_response())
|
Ok(Json(response).into_response())
|
||||||
} else {
|
} else {
|
||||||
let mut comp_req = CompletionRequest::new(messages).with_model(req.model);
|
let comp_req = build_completion_request(&req, messages);
|
||||||
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 resp = llm.complete(comp_req).await.map_err(map_llm_error)?;
|
let resp = llm.complete(comp_req).await.map_err(map_llm_error)?;
|
||||||
let model_name = llm.effective_model_name(Some(requested_model.as_str()));
|
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 llm_result = if has_tools {
|
||||||
let tools = convert_tools(req.tools.as_deref().unwrap_or(&[]));
|
let tool_req = build_tool_request(&req, messages);
|
||||||
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);
|
|
||||||
}
|
|
||||||
LlmResult::WithTools(
|
LlmResult::WithTools(
|
||||||
llm.complete_with_tools(tool_req)
|
llm.complete_with_tools(tool_req)
|
||||||
.await
|
.await
|
||||||
.map_err(map_llm_error)?,
|
.map_err(map_llm_error)?,
|
||||||
)
|
)
|
||||||
} else {
|
} else {
|
||||||
let mut comp_req = CompletionRequest::new(messages).with_model(req.model);
|
let comp_req = build_completion_request(&req, messages);
|
||||||
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);
|
|
||||||
}
|
|
||||||
LlmResult::Simple(llm.complete(comp_req).await.map_err(map_llm_error)?)
|
LlmResult::Simple(llm.complete(comp_req).await.map_err(map_llm_error)?)
|
||||||
};
|
};
|
||||||
let model_name = llm.effective_model_name(Some(requested_model.as_str()));
|
let model_name = llm.effective_model_name(Some(requested_model.as_str()));
|
||||||
|
|||||||
+5
-2
@@ -176,8 +176,11 @@ impl LlmProvider for BedrockProvider {
|
|||||||
builder = builder.tool_config(tc);
|
builder = builder.tool_config(tc);
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(config) = build_inference_config(request.temperature, request.max_tokens, None)
|
if let Some(config) = build_inference_config(
|
||||||
{
|
request.temperature,
|
||||||
|
request.max_tokens,
|
||||||
|
request.stop_sequences.as_deref(),
|
||||||
|
) {
|
||||||
builder = builder.inference_config(config);
|
builder = builder.inference_config(config);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -475,6 +475,7 @@ impl LlmProvider for NearAiChatProvider {
|
|||||||
messages,
|
messages,
|
||||||
temperature: req.temperature,
|
temperature: req.temperature,
|
||||||
max_tokens: req.max_tokens,
|
max_tokens: req.max_tokens,
|
||||||
|
stop: req.stop_sequences,
|
||||||
tools: None,
|
tools: None,
|
||||||
tool_choice: None,
|
tool_choice: None,
|
||||||
};
|
};
|
||||||
@@ -554,6 +555,7 @@ impl LlmProvider for NearAiChatProvider {
|
|||||||
messages,
|
messages,
|
||||||
temperature: req.temperature,
|
temperature: req.temperature,
|
||||||
max_tokens: req.max_tokens,
|
max_tokens: req.max_tokens,
|
||||||
|
stop: req.stop_sequences,
|
||||||
tools: if tools.is_empty() { None } else { Some(tools) },
|
tools: if tools.is_empty() { None } else { Some(tools) },
|
||||||
tool_choice: req.tool_choice,
|
tool_choice: req.tool_choice,
|
||||||
};
|
};
|
||||||
@@ -680,6 +682,8 @@ struct ChatCompletionRequest {
|
|||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
max_tokens: Option<u32>,
|
max_tokens: Option<u32>,
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
stop: Option<Vec<String>>,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
tools: Option<Vec<ChatCompletionTool>>,
|
tools: Option<Vec<ChatCompletionTool>>,
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
tool_choice: Option<String>,
|
tool_choice: Option<String>,
|
||||||
@@ -1666,6 +1670,7 @@ mod tests {
|
|||||||
}],
|
}],
|
||||||
temperature: None,
|
temperature: None,
|
||||||
max_tokens: None,
|
max_tokens: None,
|
||||||
|
stop: None,
|
||||||
tools: None,
|
tools: None,
|
||||||
tool_choice: None,
|
tool_choice: None,
|
||||||
};
|
};
|
||||||
@@ -1687,6 +1692,7 @@ mod tests {
|
|||||||
messages: vec![],
|
messages: vec![],
|
||||||
temperature: Some(0.7),
|
temperature: Some(0.7),
|
||||||
max_tokens: Some(1024),
|
max_tokens: Some(1024),
|
||||||
|
stop: None,
|
||||||
tools: Some(vec![ChatCompletionTool {
|
tools: Some(vec![ChatCompletionTool {
|
||||||
tool_type: "function".to_string(),
|
tool_type: "function".to_string(),
|
||||||
function: ChatCompletionFunction {
|
function: ChatCompletionFunction {
|
||||||
|
|||||||
+24
-3
@@ -251,6 +251,7 @@ pub struct ToolCompletionRequest {
|
|||||||
pub model: Option<String>,
|
pub model: Option<String>,
|
||||||
pub max_tokens: Option<u32>,
|
pub max_tokens: Option<u32>,
|
||||||
pub temperature: Option<f32>,
|
pub temperature: Option<f32>,
|
||||||
|
pub stop_sequences: Option<Vec<String>>,
|
||||||
/// How to handle tool use: "auto", "required", or "none".
|
/// How to handle tool use: "auto", "required", or "none".
|
||||||
pub tool_choice: Option<String>,
|
pub tool_choice: Option<String>,
|
||||||
/// Opaque metadata passed through to the provider (e.g. thread_id for chaining).
|
/// Opaque metadata passed through to the provider (e.g. thread_id for chaining).
|
||||||
@@ -266,6 +267,7 @@ impl ToolCompletionRequest {
|
|||||||
model: None,
|
model: None,
|
||||||
max_tokens: None,
|
max_tokens: None,
|
||||||
temperature: None,
|
temperature: None,
|
||||||
|
stop_sequences: None,
|
||||||
tool_choice: None,
|
tool_choice: None,
|
||||||
metadata: std::collections::HashMap::new(),
|
metadata: std::collections::HashMap::new(),
|
||||||
}
|
}
|
||||||
@@ -289,6 +291,12 @@ impl ToolCompletionRequest {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Set stop sequences.
|
||||||
|
pub fn with_stop_sequences(mut self, stop_sequences: Vec<String>) -> Self {
|
||||||
|
self.stop_sequences = Some(stop_sequences);
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
/// Set tool choice mode.
|
/// Set tool choice mode.
|
||||||
pub fn with_tool_choice(mut self, choice: impl Into<String>) -> Self {
|
pub fn with_tool_choice(mut self, choice: impl Into<String>) -> Self {
|
||||||
self.tool_choice = Some(choice.into());
|
self.tool_choice = Some(choice.into());
|
||||||
@@ -504,8 +512,6 @@ pub fn strip_unsupported_completion_params(
|
|||||||
/// This is the single helper function used by all providers to remove
|
/// This is the single helper function used by all providers to remove
|
||||||
/// parameters they don't support from tool calls, replacing duplicate stringly-typed logic.
|
/// parameters they don't support from tool calls, replacing duplicate stringly-typed logic.
|
||||||
///
|
///
|
||||||
/// Note: Only `Temperature` and `MaxTokens` are supported in `ToolCompletionRequest`.
|
|
||||||
/// `StopSequences` is only available in `CompletionRequest` and is not applicable to tool calls.
|
|
||||||
pub fn strip_unsupported_tool_params(
|
pub fn strip_unsupported_tool_params(
|
||||||
unsupported: &std::collections::HashSet<String>,
|
unsupported: &std::collections::HashSet<String>,
|
||||||
req: &mut ToolCompletionRequest,
|
req: &mut ToolCompletionRequest,
|
||||||
@@ -519,7 +525,9 @@ pub fn strip_unsupported_tool_params(
|
|||||||
if unsupported.contains(UnsupportedParam::MaxTokens.name()) {
|
if unsupported.contains(UnsupportedParam::MaxTokens.name()) {
|
||||||
req.max_tokens = None;
|
req.max_tokens = None;
|
||||||
}
|
}
|
||||||
// Note: StopSequences is not a field in ToolCompletionRequest, so no action needed
|
if unsupported.contains(UnsupportedParam::StopSequences.name()) {
|
||||||
|
req.stop_sequences = None;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -651,4 +659,17 @@ mod tests {
|
|||||||
assert!(messages[2].tool_call_id.is_none());
|
assert!(messages[2].tool_call_id.is_none());
|
||||||
assert!(messages[2].name.is_none());
|
assert!(messages[2].name.is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_strip_unsupported_tool_params_strips_stop_sequences() {
|
||||||
|
let mut unsupported = std::collections::HashSet::new();
|
||||||
|
unsupported.insert(UnsupportedParam::StopSequences.name().to_string());
|
||||||
|
|
||||||
|
let mut req = ToolCompletionRequest::new(vec![ChatMessage::user("hello")], vec![]);
|
||||||
|
req.stop_sequences = Some(vec!["STOP".to_string()]);
|
||||||
|
|
||||||
|
strip_unsupported_tool_params(&unsupported, &mut req);
|
||||||
|
|
||||||
|
assert!(req.stop_sequences.is_none()); // safety: test assertion for explicit strip behavior
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -548,6 +548,7 @@ mod tests {
|
|||||||
model: None,
|
model: None,
|
||||||
max_tokens: None,
|
max_tokens: None,
|
||||||
temperature: None,
|
temperature: None,
|
||||||
|
stop_sequences: None,
|
||||||
tool_choice: None,
|
tool_choice: None,
|
||||||
metadata: Default::default(),
|
metadata: Default::default(),
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -176,6 +176,7 @@ async fn llm_complete_with_tools(
|
|||||||
model: req.model,
|
model: req.model,
|
||||||
max_tokens: req.max_tokens,
|
max_tokens: req.max_tokens,
|
||||||
temperature: req.temperature,
|
temperature: req.temperature,
|
||||||
|
stop_sequences: req.stop_sequences,
|
||||||
tool_choice: req.tool_choice,
|
tool_choice: req.tool_choice,
|
||||||
metadata: std::collections::HashMap::new(),
|
metadata: std::collections::HashMap::new(),
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -65,6 +65,7 @@ pub struct ProxyToolCompletionRequest {
|
|||||||
pub model: Option<String>,
|
pub model: Option<String>,
|
||||||
pub max_tokens: Option<u32>,
|
pub max_tokens: Option<u32>,
|
||||||
pub temperature: Option<f32>,
|
pub temperature: Option<f32>,
|
||||||
|
pub stop_sequences: Option<Vec<String>>,
|
||||||
pub tool_choice: Option<String>,
|
pub tool_choice: Option<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -251,6 +252,7 @@ impl WorkerHttpClient {
|
|||||||
model: request.model.clone(),
|
model: request.model.clone(),
|
||||||
max_tokens: request.max_tokens,
|
max_tokens: request.max_tokens,
|
||||||
temperature: request.temperature,
|
temperature: request.temperature,
|
||||||
|
stop_sequences: request.stop_sequences.clone(),
|
||||||
tool_choice: request.tool_choice.clone(),
|
tool_choice: request.tool_choice.clone(),
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user