Files
optimclaw/src/llm/provider.rs
T
41ed0a0f98 feat(agent): thread per-tool reasoning through provider, session, and all surfaces (#1513)
* 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]>
2026-03-25 08:35:41 -07:00

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
}
}