mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
* feat(agent): thread per-tool reasoning from LLM through to REPL, HTTP, SSE, and DB Add end-to-end agent reasoning summaries so users can see *why* the agent chose specific tools, not just what it did. - Add `reasoning: Option<String>` to `ToolCall` (all providers) - Populate from LLM response content in `Reasoning::respond_with_tools` and `select_tools`, with per-tool override when providers supply it - Extend `Turn` with `narrative` and `TurnToolCall` with `rationale` + `tool_call_id` for identity-based result matching - Persist reasoning in DB via existing tool_calls JSON (no migration) - Add `StatusUpdate::ReasoningUpdate` and `SseEvent::ReasoningUpdate` + `SseEvent::JobReasoning` for real-time streaming - Emit reasoning events in both chat dispatcher and worker job path - Add `/reasoning [N|all]` command for inspecting turn reasoning - Surface `narrative` and `rationale` in HTTP `/api/chat/history` Based on the design from #361 and #456, reconstructed cleanly with Option<String> to minimize blast radius (vs mandatory String that broke compilation in #456). Closes #456 Co-Authored-By: panosAthDBX <[email protected]> Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: address PR review feedback from Gemini and Copilot - Fix `_ => Ok(None)` in agent_loop.rs to avoid accidental shutdown - Fix fallback in record_tool_result_for/record_tool_error_for to use first pending call instead of last_mut (parallel execution safety) - Include per-tool decisions in WASM channel reasoning messages - Apply truncate_at_tool_tags + clean_response to shared_reasoning in select_tools (parity with respond_with_tools) - Persist turn-level narrative to DB in tool_calls JSON wrapper - Parse both old (array) and new (object) tool_calls formats in build_turns_from_db_messages for backward compatibility - Populate reasoning from action.reasoning in execute_plan ToolCalls [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: address second round of review comments + merge fixes - Add reasoning: None to new github_copilot.rs ToolCall sites (from staging merge) - Run cargo fmt on 4 files with formatting diffs - Truncate narrative to 1000 chars before DB persistence - Clone turn data and drop session lock in /reasoning command - Extract ToolDecisionDto::from_json_array shared helper (deduplicate worker/job.rs and orchestrator/api.rs) - Add unit tests for wrapped tool_calls JSON format with narrative [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: address third round of review comments (Copilot + serrrfirat) - Reword ToolCall.reasoning docstring to reflect provider-supplied or fallback contract - Sanitize narrative through SafetyLayer before storage/emission - Clean per-tool reasoning via truncate_at_tool_tags + clean_response in select_tools (parity with shared reasoning) - Convert 4 approval-path recording sites in thread_ops.rs to identity-based record_tool_result_for/record_tool_error_for - Preserve tool_call_id and reasoning through restore_from_messages - Fix has_result/has_error to reject JSON null values - Truncate tool_call_id to 128 chars before DB persistence - Add 4 unit tests for record_tool_result_for/error_for edge cases Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: address zmanian review — sanitize JobDelegate reasoning + warn on dropped results - Sanitize narrative and per-tool rationale through SafetyLayer in JobDelegate reasoning events (parity with ChatDelegate) - Add tracing::warn when record_tool_result_for/error_for drops a result because no matching or pending tool call exists - Add 3 unit tests for reasoning normalization (thinking tags, tool tags, empty-after-cleaning) Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: address 4 remaining unreplied review comments - Clean per-tool reasoning in respond_with_tools via truncate_at_tool_tags + clean_response (parity with select_tools) - Handle wrapped JSON format in rebuild_chat_messages_from_db so cold hydration works after persist_tool_calls format change - Update persist_tool_calls doc comment to describe new JSON shape - Sanitize per-tool rationale through SafetyLayer in ChatDelegate before emission and storage (parity with JobDelegate) Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: address zmanian review round 2 - Add tracing::debug on fallback-to-pending path in record_tool_result_for and record_tool_error_for (item 1) - Add comment explaining why /reasoning is special-cased in agent_loop.rs (item 4) - Items 2 (narrative persistence), 3 (rationale sanitization), and 5 (catch-all fix) were already addressed in prior commits Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> --------- Co-authored-by: panosAthDBX <[email protected]> Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]>
781 lines
26 KiB
Rust
781 lines
26 KiB
Rust
//! LLM provider trait and types.
|
|
|
|
use async_trait::async_trait;
|
|
use rust_decimal::Decimal;
|
|
use serde::{Deserialize, Serialize};
|
|
|
|
use crate::llm::error::LlmError;
|
|
|
|
/// Role in a conversation.
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
|
#[serde(rename_all = "lowercase")]
|
|
pub enum Role {
|
|
System,
|
|
User,
|
|
Assistant,
|
|
Tool,
|
|
}
|
|
|
|
/// A part of multimodal message content (OpenAI Chat Completions format).
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
#[serde(tag = "type")]
|
|
pub enum ContentPart {
|
|
/// Text content part.
|
|
#[serde(rename = "text")]
|
|
Text { text: String },
|
|
/// Image URL content part (supports data: URLs for inline base64 images).
|
|
#[serde(rename = "image_url")]
|
|
ImageUrl { image_url: ImageUrl },
|
|
}
|
|
|
|
/// Image URL reference for multimodal content.
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct ImageUrl {
|
|
/// URL or data: URI (e.g., "data:image/jpeg;base64,...").
|
|
pub url: String,
|
|
/// Detail level hint: "auto", "low", or "high".
|
|
#[serde(skip_serializing_if = "Option::is_none")]
|
|
pub detail: Option<String>,
|
|
}
|
|
|
|
/// A message in a conversation.
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct ChatMessage {
|
|
pub role: Role,
|
|
pub content: String,
|
|
/// Multimodal content parts (images, etc.).
|
|
/// When non-empty, providers serialize content as an array of parts
|
|
/// (with `content` included as a text part) instead of a plain string.
|
|
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
|
pub content_parts: Vec<ContentPart>,
|
|
/// Tool call ID if this is a tool result message.
|
|
#[serde(skip_serializing_if = "Option::is_none")]
|
|
pub tool_call_id: Option<String>,
|
|
/// Name of the tool for tool results.
|
|
#[serde(skip_serializing_if = "Option::is_none")]
|
|
pub name: Option<String>,
|
|
/// Tool calls made by the assistant (OpenAI protocol requires these
|
|
/// to appear on the assistant message preceding tool result messages).
|
|
#[serde(skip_serializing_if = "Option::is_none")]
|
|
pub tool_calls: Option<Vec<ToolCall>>,
|
|
}
|
|
|
|
impl ChatMessage {
|
|
/// Create a system message.
|
|
pub fn system(content: impl Into<String>) -> Self {
|
|
Self {
|
|
role: Role::System,
|
|
content: content.into(),
|
|
content_parts: Vec::new(),
|
|
tool_call_id: None,
|
|
name: None,
|
|
tool_calls: None,
|
|
}
|
|
}
|
|
|
|
/// Create a user message.
|
|
pub fn user(content: impl Into<String>) -> Self {
|
|
Self {
|
|
role: Role::User,
|
|
content: content.into(),
|
|
content_parts: Vec::new(),
|
|
tool_call_id: None,
|
|
name: None,
|
|
tool_calls: None,
|
|
}
|
|
}
|
|
|
|
/// Create a user message with multimodal content parts (e.g., images).
|
|
///
|
|
/// The text `content` is included as the primary text alongside the parts.
|
|
pub fn user_with_parts(content: impl Into<String>, parts: Vec<ContentPart>) -> Self {
|
|
Self {
|
|
role: Role::User,
|
|
content: content.into(),
|
|
content_parts: parts,
|
|
tool_call_id: None,
|
|
name: None,
|
|
tool_calls: None,
|
|
}
|
|
}
|
|
|
|
/// Create an assistant message.
|
|
pub fn assistant(content: impl Into<String>) -> Self {
|
|
Self {
|
|
role: Role::Assistant,
|
|
content: content.into(),
|
|
content_parts: Vec::new(),
|
|
tool_call_id: None,
|
|
name: None,
|
|
tool_calls: None,
|
|
}
|
|
}
|
|
|
|
/// Create an assistant message that includes tool calls.
|
|
///
|
|
/// Per the OpenAI protocol, an assistant message with tool_calls must
|
|
/// precede the corresponding tool result messages in the conversation.
|
|
pub fn assistant_with_tool_calls(content: Option<String>, tool_calls: Vec<ToolCall>) -> Self {
|
|
Self {
|
|
role: Role::Assistant,
|
|
content: content.unwrap_or_default(),
|
|
content_parts: Vec::new(),
|
|
tool_call_id: None,
|
|
name: None,
|
|
tool_calls: if tool_calls.is_empty() {
|
|
None
|
|
} else {
|
|
Some(tool_calls)
|
|
},
|
|
}
|
|
}
|
|
|
|
/// Create a tool result message.
|
|
pub fn tool_result(
|
|
tool_call_id: impl Into<String>,
|
|
name: impl Into<String>,
|
|
content: impl Into<String>,
|
|
) -> Self {
|
|
Self {
|
|
role: Role::Tool,
|
|
content: content.into(),
|
|
content_parts: Vec::new(),
|
|
tool_call_id: Some(tool_call_id.into()),
|
|
name: Some(name.into()),
|
|
tool_calls: None,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Request for a chat completion.
|
|
#[derive(Debug, Clone)]
|
|
pub struct CompletionRequest {
|
|
pub messages: Vec<ChatMessage>,
|
|
/// Optional per-request model override.
|
|
pub model: Option<String>,
|
|
pub max_tokens: Option<u32>,
|
|
pub temperature: Option<f32>,
|
|
pub stop_sequences: Option<Vec<String>>,
|
|
/// Opaque metadata passed through to the provider (e.g. thread_id for chaining).
|
|
pub metadata: std::collections::HashMap<String, String>,
|
|
}
|
|
|
|
impl CompletionRequest {
|
|
/// Create a new completion request.
|
|
pub fn new(messages: Vec<ChatMessage>) -> Self {
|
|
Self {
|
|
messages,
|
|
model: None,
|
|
max_tokens: None,
|
|
temperature: None,
|
|
stop_sequences: None,
|
|
metadata: std::collections::HashMap::new(),
|
|
}
|
|
}
|
|
|
|
/// Set model override.
|
|
pub fn with_model(mut self, model: impl Into<String>) -> Self {
|
|
self.model = Some(model.into());
|
|
self
|
|
}
|
|
|
|
/// Set max tokens.
|
|
pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
|
|
self.max_tokens = Some(max_tokens);
|
|
self
|
|
}
|
|
|
|
/// Set temperature.
|
|
pub fn with_temperature(mut self, temperature: f32) -> Self {
|
|
self.temperature = Some(temperature);
|
|
self
|
|
}
|
|
}
|
|
|
|
/// Response from a chat completion.
|
|
#[derive(Debug, Clone)]
|
|
pub struct CompletionResponse {
|
|
pub content: String,
|
|
pub input_tokens: u32,
|
|
pub output_tokens: u32,
|
|
pub finish_reason: FinishReason,
|
|
/// Tokens read from the provider's server-side prompt cache (Anthropic).
|
|
/// Zero when caching is not supported or on a cache miss.
|
|
pub cache_read_input_tokens: u32,
|
|
/// Tokens written to the provider's server-side prompt cache (Anthropic).
|
|
/// Zero when caching is not supported or no new prefix was cached.
|
|
pub cache_creation_input_tokens: u32,
|
|
}
|
|
|
|
/// Why the completion finished.
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
pub enum FinishReason {
|
|
Stop,
|
|
Length,
|
|
ToolUse,
|
|
ContentFilter,
|
|
Unknown,
|
|
}
|
|
|
|
/// Definition of a tool for the LLM.
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct ToolDefinition {
|
|
pub name: String,
|
|
pub description: String,
|
|
pub parameters: serde_json::Value,
|
|
}
|
|
|
|
/// A tool call requested by the LLM.
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct ToolCall {
|
|
pub id: String,
|
|
pub name: String,
|
|
pub arguments: serde_json::Value,
|
|
/// Optional reasoning for why this tool was chosen — supplied by the provider
|
|
/// or derived from the shared response content as a fallback.
|
|
#[serde(default, skip_serializing_if = "Option::is_none")]
|
|
pub reasoning: Option<String>,
|
|
}
|
|
|
|
/// Generate a tool-call ID that satisfies all providers.
|
|
///
|
|
/// Mistral requires exactly 9 alphanumeric characters (`[a-zA-Z0-9]{9}`).
|
|
/// Other providers accept any non-empty string. By default we produce a
|
|
/// 9-char base-62 string derived from two seed values so the ID is both
|
|
/// deterministic (for replayed history) and provider-compatible.
|
|
pub fn generate_tool_call_id(seed_a: usize, seed_b: usize) -> String {
|
|
// Mix the two seeds into a single u64 using a simple hash-like combine.
|
|
let combined = (seed_a as u64)
|
|
.wrapping_mul(6364136223846793005)
|
|
.wrapping_add(seed_b as u64);
|
|
// Format as 9-char zero-padded base-62 (0-9, a-z, A-Z).
|
|
let mut buf = [b'0'; 9];
|
|
let mut val = combined;
|
|
for b in buf.iter_mut().rev() {
|
|
let digit = (val % 62) as u8;
|
|
*b = match digit {
|
|
0..=9 => b'0' + digit,
|
|
10..=35 => b'a' + (digit - 10),
|
|
_ => b'A' + (digit - 36),
|
|
};
|
|
val /= 62;
|
|
}
|
|
buf.iter().map(|&b| b as char).collect::<String>()
|
|
}
|
|
|
|
/// Result of a tool execution to send back to the LLM.
|
|
#[derive(Debug, Clone)]
|
|
pub struct ToolResult {
|
|
pub tool_call_id: String,
|
|
pub name: String,
|
|
pub content: String,
|
|
pub is_error: bool,
|
|
}
|
|
|
|
/// Request for a completion with tool use.
|
|
#[derive(Debug, Clone)]
|
|
pub struct ToolCompletionRequest {
|
|
pub messages: Vec<ChatMessage>,
|
|
pub tools: Vec<ToolDefinition>,
|
|
/// Optional per-request model override.
|
|
pub model: Option<String>,
|
|
pub max_tokens: Option<u32>,
|
|
pub temperature: Option<f32>,
|
|
pub stop_sequences: Option<Vec<String>>,
|
|
/// How to handle tool use: "auto", "required", or "none".
|
|
pub tool_choice: Option<String>,
|
|
/// Opaque metadata passed through to the provider (e.g. thread_id for chaining).
|
|
pub metadata: std::collections::HashMap<String, String>,
|
|
}
|
|
|
|
impl ToolCompletionRequest {
|
|
/// Create a new tool completion request.
|
|
pub fn new(messages: Vec<ChatMessage>, tools: Vec<ToolDefinition>) -> Self {
|
|
Self {
|
|
messages,
|
|
tools,
|
|
model: None,
|
|
max_tokens: None,
|
|
temperature: None,
|
|
stop_sequences: None,
|
|
tool_choice: None,
|
|
metadata: std::collections::HashMap::new(),
|
|
}
|
|
}
|
|
|
|
/// Set model override.
|
|
pub fn with_model(mut self, model: impl Into<String>) -> Self {
|
|
self.model = Some(model.into());
|
|
self
|
|
}
|
|
|
|
/// Set max tokens.
|
|
pub fn with_max_tokens(mut self, max_tokens: u32) -> Self {
|
|
self.max_tokens = Some(max_tokens);
|
|
self
|
|
}
|
|
|
|
/// Set temperature.
|
|
pub fn with_temperature(mut self, temperature: f32) -> Self {
|
|
self.temperature = Some(temperature);
|
|
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.
|
|
pub fn with_tool_choice(mut self, choice: impl Into<String>) -> Self {
|
|
self.tool_choice = Some(choice.into());
|
|
self
|
|
}
|
|
}
|
|
|
|
/// Response from a completion with potential tool calls.
|
|
#[derive(Debug, Clone)]
|
|
pub struct ToolCompletionResponse {
|
|
/// Text content (may be empty if tool calls are present).
|
|
pub content: Option<String>,
|
|
/// Tool calls requested by the model.
|
|
pub tool_calls: Vec<ToolCall>,
|
|
pub input_tokens: u32,
|
|
pub output_tokens: u32,
|
|
pub finish_reason: FinishReason,
|
|
/// Tokens read from the provider's server-side prompt cache (Anthropic).
|
|
pub cache_read_input_tokens: u32,
|
|
/// Tokens written to the provider's server-side prompt cache (Anthropic).
|
|
pub cache_creation_input_tokens: u32,
|
|
}
|
|
|
|
/// Metadata about a model returned by the provider's API.
|
|
#[derive(Debug, Clone)]
|
|
pub struct ModelMetadata {
|
|
pub id: String,
|
|
/// Total context window size in tokens.
|
|
pub context_length: Option<u32>,
|
|
}
|
|
|
|
/// Trait for LLM providers.
|
|
#[async_trait]
|
|
pub trait LlmProvider: Send + Sync {
|
|
/// Get the model name.
|
|
fn model_name(&self) -> &str;
|
|
|
|
/// Get cost per token (input, output).
|
|
fn cost_per_token(&self) -> (Decimal, Decimal);
|
|
|
|
/// Complete a chat conversation.
|
|
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError>;
|
|
|
|
/// Complete with tool use support.
|
|
async fn complete_with_tools(
|
|
&self,
|
|
request: ToolCompletionRequest,
|
|
) -> Result<ToolCompletionResponse, LlmError>;
|
|
|
|
/// List available models from the provider.
|
|
/// Default implementation returns empty list.
|
|
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
|
|
Ok(Vec::new())
|
|
}
|
|
|
|
/// Fetch metadata for the current model (context length, etc.).
|
|
/// Default returns the model name with no size info.
|
|
async fn model_metadata(&self) -> Result<ModelMetadata, LlmError> {
|
|
Ok(ModelMetadata {
|
|
id: self.model_name().to_string(),
|
|
context_length: None,
|
|
})
|
|
}
|
|
|
|
/// Resolve which model should be reported for a given request.
|
|
///
|
|
/// Providers that ignore per-request model overrides should override this
|
|
/// and return `active_model_name()`.
|
|
fn effective_model_name(&self, requested_model: Option<&str>) -> String {
|
|
requested_model
|
|
.map(std::borrow::ToOwned::to_owned)
|
|
.unwrap_or_else(|| self.active_model_name())
|
|
}
|
|
|
|
/// Get the currently active model name.
|
|
///
|
|
/// May differ from `model_name()` if the model was switched at runtime
|
|
/// via `set_model()`. Default returns `model_name()`.
|
|
fn active_model_name(&self) -> String {
|
|
self.model_name().to_string()
|
|
}
|
|
|
|
/// Switch the active model at runtime. Not all providers support this.
|
|
fn set_model(&self, _model: &str) -> Result<(), LlmError> {
|
|
Err(LlmError::RequestFailed {
|
|
provider: "unknown".to_string(),
|
|
reason: "Runtime model switching not supported by this provider".to_string(),
|
|
})
|
|
}
|
|
|
|
/// Calculate cost for a completion.
|
|
fn calculate_cost(&self, input_tokens: u32, output_tokens: u32) -> Decimal {
|
|
let (input_cost, output_cost) = self.cost_per_token();
|
|
input_cost * Decimal::from(input_tokens) + output_cost * Decimal::from(output_tokens)
|
|
}
|
|
|
|
/// Cost multiplier for cache-creation tokens (Anthropic prompt caching).
|
|
///
|
|
/// Returns `1.0` by default (no surcharge). Anthropic providers return
|
|
/// `1.25` for 5-minute TTL or `2.0` for 1-hour TTL.
|
|
fn cache_write_multiplier(&self) -> Decimal {
|
|
Decimal::ONE
|
|
}
|
|
|
|
/// Discount divisor for cache-read tokens.
|
|
///
|
|
/// Cached-read cost = `input_rate / cache_read_discount()`.
|
|
/// Returns `1` by default (no discount). Anthropic returns `10` (90% off),
|
|
/// OpenAI would return `2` (50% off).
|
|
fn cache_read_discount(&self) -> Decimal {
|
|
Decimal::ONE
|
|
}
|
|
}
|
|
|
|
/// Sanitize a message list to ensure tool_use / tool_result integrity.
|
|
///
|
|
/// LLM APIs (especially Anthropic) require every tool_result to reference a
|
|
/// tool_call_id that exists in an immediately preceding assistant message's
|
|
/// tool_calls. Orphaned tool_results cause HTTP 400 errors.
|
|
///
|
|
/// This function:
|
|
/// 1. Tracks all tool_call_ids emitted by assistant messages.
|
|
/// 2. Rewrites orphaned tool_result messages (whose tool_call_id has no
|
|
/// matching assistant tool_call) as user messages so the content is
|
|
/// preserved without violating the protocol.
|
|
///
|
|
/// Call this before sending messages to any LLM provider.
|
|
pub fn sanitize_tool_messages(messages: &mut [ChatMessage]) {
|
|
use std::collections::HashSet;
|
|
|
|
// Collect all tool_call_ids from assistant messages with tool_calls.
|
|
let mut known_ids: HashSet<String> = HashSet::new();
|
|
for msg in messages.iter() {
|
|
if msg.role == Role::Assistant
|
|
&& let Some(ref calls) = msg.tool_calls
|
|
{
|
|
for tc in calls {
|
|
known_ids.insert(tc.id.clone());
|
|
}
|
|
}
|
|
}
|
|
|
|
// Rewrite orphaned tool_result messages as user messages.
|
|
for msg in messages.iter_mut() {
|
|
if msg.role != Role::Tool {
|
|
continue;
|
|
}
|
|
let is_orphaned = match &msg.tool_call_id {
|
|
Some(id) => !known_ids.contains(id),
|
|
None => true,
|
|
};
|
|
if is_orphaned {
|
|
let tool_name = msg.name.as_deref().unwrap_or("unknown");
|
|
tracing::debug!(
|
|
tool_call_id = ?msg.tool_call_id,
|
|
tool_name,
|
|
"Rewriting orphaned tool_result as user message",
|
|
);
|
|
msg.role = Role::User;
|
|
msg.content = format!("[Tool `{}` returned: {}]", tool_name, msg.content);
|
|
msg.tool_call_id = None;
|
|
msg.name = None;
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Represents a request parameter that may not be supported by all LLM providers.
|
|
///
|
|
/// This typed enum replaces stringly-typed parameter names across the codebase,
|
|
/// providing type safety and single-point-of-maintenance for parameter handling.
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
|
|
pub enum UnsupportedParam {
|
|
Temperature,
|
|
MaxTokens,
|
|
StopSequences,
|
|
}
|
|
|
|
impl UnsupportedParam {
|
|
/// Get the string name of this parameter for config/error messages.
|
|
pub fn name(&self) -> &'static str {
|
|
match self {
|
|
UnsupportedParam::Temperature => "temperature",
|
|
UnsupportedParam::MaxTokens => "max_tokens",
|
|
UnsupportedParam::StopSequences => "stop_sequences",
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Strip unsupported parameters from a `CompletionRequest` in place.
|
|
///
|
|
/// This is the single helper function used by all providers to remove
|
|
/// parameters they don't support, replacing duplicate stringly-typed logic.
|
|
pub fn strip_unsupported_completion_params(
|
|
unsupported: &std::collections::HashSet<String>,
|
|
req: &mut CompletionRequest,
|
|
) {
|
|
if unsupported.is_empty() {
|
|
return;
|
|
}
|
|
if unsupported.contains(UnsupportedParam::Temperature.name()) {
|
|
req.temperature = None;
|
|
}
|
|
if unsupported.contains(UnsupportedParam::MaxTokens.name()) {
|
|
req.max_tokens = None;
|
|
}
|
|
if unsupported.contains(UnsupportedParam::StopSequences.name()) {
|
|
req.stop_sequences = None;
|
|
}
|
|
}
|
|
|
|
/// Strip unsupported parameters from a `ToolCompletionRequest` in place.
|
|
///
|
|
/// 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.
|
|
///
|
|
pub fn strip_unsupported_tool_params(
|
|
unsupported: &std::collections::HashSet<String>,
|
|
req: &mut ToolCompletionRequest,
|
|
) {
|
|
if unsupported.is_empty() {
|
|
return;
|
|
}
|
|
if unsupported.contains(UnsupportedParam::Temperature.name()) {
|
|
req.temperature = None;
|
|
}
|
|
if unsupported.contains(UnsupportedParam::MaxTokens.name()) {
|
|
req.max_tokens = None;
|
|
}
|
|
if unsupported.contains(UnsupportedParam::StopSequences.name()) {
|
|
req.stop_sequences = None;
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use std::collections::HashSet;
|
|
|
|
#[test]
|
|
fn generate_tool_call_id_has_valid_format() {
|
|
let samples = [
|
|
(0usize, 0usize),
|
|
(1usize, 2usize),
|
|
(42usize, 999usize),
|
|
(usize::MAX, usize::MAX),
|
|
];
|
|
|
|
for (a, b) in samples {
|
|
let id = generate_tool_call_id(a, b);
|
|
assert_eq!(
|
|
id.len(),
|
|
9,
|
|
"tool-call ID must be exactly 9 characters for seeds ({a}, {b})"
|
|
);
|
|
assert!(
|
|
id.chars().all(|c| c.is_ascii_alphanumeric()),
|
|
"tool-call ID must be ASCII alphanumeric for seeds ({a}, {b}), got: {id}"
|
|
);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn generate_tool_call_id_is_deterministic_for_same_seeds() {
|
|
let pairs = [
|
|
(0usize, 0usize),
|
|
(1usize, 2usize),
|
|
(123usize, 456usize),
|
|
(usize::MAX, 0usize),
|
|
];
|
|
|
|
for (a, b) in pairs {
|
|
let id1 = generate_tool_call_id(a, b);
|
|
let id2 = generate_tool_call_id(a, b);
|
|
let id3 = generate_tool_call_id(a, b);
|
|
assert_eq!(
|
|
id1, id2,
|
|
"tool-call ID must be deterministic for seeds ({a}, {b})"
|
|
);
|
|
assert_eq!(
|
|
id2, id3,
|
|
"tool-call ID must be deterministic across multiple calls for seeds ({a}, {b})"
|
|
);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn generate_tool_call_id_differs_for_different_seeds_in_small_sample() {
|
|
let seed_pairs = [
|
|
(0usize, 1usize),
|
|
(1usize, 0usize),
|
|
(1usize, 2usize),
|
|
(2usize, 3usize),
|
|
(10usize, 20usize),
|
|
(100usize, 200usize),
|
|
];
|
|
|
|
let mut ids = HashSet::new();
|
|
for (a, b) in seed_pairs {
|
|
let id = generate_tool_call_id(a, b);
|
|
let inserted = ids.insert(id.clone());
|
|
assert!(
|
|
inserted,
|
|
"expected distinct tool-call IDs for different seeds, \
|
|
but duplicate ID '{id}' found for seeds ({a}, {b})"
|
|
);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_sanitize_preserves_valid_pairs() {
|
|
let tc = ToolCall {
|
|
id: "call_1".to_string(),
|
|
name: "echo".to_string(),
|
|
arguments: serde_json::json!({}),
|
|
reasoning: None,
|
|
};
|
|
let mut messages = vec![
|
|
ChatMessage::user("hello"),
|
|
ChatMessage::assistant_with_tool_calls(None, vec![tc]),
|
|
ChatMessage::tool_result("call_1", "echo", "result"),
|
|
];
|
|
sanitize_tool_messages(&mut messages);
|
|
assert_eq!(messages[2].role, Role::Tool);
|
|
assert_eq!(messages[2].tool_call_id, Some("call_1".to_string()));
|
|
}
|
|
|
|
#[test]
|
|
fn test_sanitize_rewrites_orphaned_tool_result() {
|
|
let mut messages = vec![
|
|
ChatMessage::user("hello"),
|
|
ChatMessage::assistant("I'll use a tool"),
|
|
ChatMessage::tool_result("call_missing", "search", "some result"),
|
|
];
|
|
sanitize_tool_messages(&mut messages);
|
|
assert_eq!(messages[2].role, Role::User);
|
|
assert!(messages[2].content.contains("[Tool `search` returned:"));
|
|
assert!(messages[2].tool_call_id.is_none());
|
|
assert!(messages[2].name.is_none());
|
|
}
|
|
|
|
#[test]
|
|
fn test_sanitize_handles_no_tool_messages() {
|
|
let mut messages = vec![
|
|
ChatMessage::system("prompt"),
|
|
ChatMessage::user("hello"),
|
|
ChatMessage::assistant("hi"),
|
|
];
|
|
let original_len = messages.len();
|
|
sanitize_tool_messages(&mut messages);
|
|
assert_eq!(messages.len(), original_len);
|
|
}
|
|
|
|
#[test]
|
|
fn test_sanitize_multiple_orphaned() {
|
|
let tc = ToolCall {
|
|
id: "call_1".to_string(),
|
|
name: "echo".to_string(),
|
|
arguments: serde_json::json!({}),
|
|
reasoning: None,
|
|
};
|
|
let mut messages = vec![
|
|
ChatMessage::user("test"),
|
|
ChatMessage::assistant_with_tool_calls(None, vec![tc]),
|
|
ChatMessage::tool_result("call_1", "echo", "ok"),
|
|
// These are orphaned (call_2 and call_3 have no matching assistant message)
|
|
ChatMessage::tool_result("call_2", "search", "orphan 1"),
|
|
ChatMessage::tool_result("call_3", "http", "orphan 2"),
|
|
];
|
|
sanitize_tool_messages(&mut messages);
|
|
assert_eq!(messages[2].role, Role::Tool); // call_1 is valid
|
|
assert_eq!(messages[3].role, Role::User); // call_2 orphaned
|
|
assert_eq!(messages[4].role, Role::User); // call_3 orphaned
|
|
}
|
|
|
|
/// Regression: worker's select_tools/execute_plan now emit
|
|
/// assistant_with_tool_calls before tool_result messages.
|
|
/// Verify sanitize_tool_messages preserves all tool_results when
|
|
/// each has a matching assistant tool_call.
|
|
#[test]
|
|
fn test_sanitize_preserves_tool_results_with_matching_assistant() {
|
|
let tc1 = ToolCall {
|
|
id: "call_sel_1".to_string(),
|
|
name: "search".to_string(),
|
|
arguments: serde_json::json!({"q": "test"}),
|
|
reasoning: None,
|
|
};
|
|
let tc2 = ToolCall {
|
|
id: "call_sel_2".to_string(),
|
|
name: "http".to_string(),
|
|
arguments: serde_json::json!({"url": "https://example.com"}),
|
|
reasoning: None,
|
|
};
|
|
let mut messages = vec![
|
|
ChatMessage::system("You are a helpful assistant."),
|
|
ChatMessage::assistant_with_tool_calls(None, vec![tc1, tc2]),
|
|
ChatMessage::tool_result("call_sel_1", "search", "found 3 results"),
|
|
ChatMessage::tool_result("call_sel_2", "http", "200 OK"),
|
|
];
|
|
sanitize_tool_messages(&mut messages);
|
|
|
|
// All tool_results must keep Role::Tool -- none should be rewritten.
|
|
assert_eq!(messages[2].role, Role::Tool);
|
|
assert_eq!(messages[2].tool_call_id, Some("call_sel_1".to_string()));
|
|
assert_eq!(messages[2].content, "found 3 results");
|
|
|
|
assert_eq!(messages[3].role, Role::Tool);
|
|
assert_eq!(messages[3].tool_call_id, Some("call_sel_2".to_string()));
|
|
assert_eq!(messages[3].content, "200 OK");
|
|
}
|
|
|
|
/// Regression: the OLD buggy worker code pushed tool_result messages
|
|
/// without a preceding assistant_with_tool_calls, causing
|
|
/// sanitize_tool_messages to rewrite them as orphaned user messages.
|
|
/// This test reproduces that buggy sequence and confirms the rewrite.
|
|
#[test]
|
|
fn test_sanitize_rewrites_orphaned_tool_results() {
|
|
let mut messages = vec![
|
|
ChatMessage::system("You are a helpful assistant."),
|
|
// No assistant_with_tool_calls -- mimics the old bug.
|
|
ChatMessage::tool_result("call_bug_1", "search", "found 3 results"),
|
|
ChatMessage::tool_result("call_bug_2", "http", "200 OK"),
|
|
];
|
|
sanitize_tool_messages(&mut messages);
|
|
|
|
// Both tool_results must be rewritten to Role::User.
|
|
assert_eq!(messages[1].role, Role::User);
|
|
assert!(messages[1].content.contains("[Tool `search` returned:"));
|
|
assert!(messages[1].content.contains("found 3 results"));
|
|
assert!(messages[1].tool_call_id.is_none());
|
|
assert!(messages[1].name.is_none());
|
|
|
|
assert_eq!(messages[2].role, Role::User);
|
|
assert!(messages[2].content.contains("[Tool `http` returned:"));
|
|
assert!(messages[2].content.contains("200 OK"));
|
|
assert!(messages[2].tool_call_id.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
|
|
}
|
|
}
|