mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-30 16:19:21 +00:00
* fix: harden openai-compatible tool flow and local defaults * fix: close approval replay gaps and harden openai-compatible flow * fix: address review feedback and code improvements (takeover #112) - Make ChatCompletionResponse.id Optional<String> to handle providers that omit or null the field - Propagate HTTP client builder errors instead of silently dropping timeout configuration (openai_compatible_chat, nearai_chat) - Add EMBEDDING_DIMENSION env var with smart per-model defaults instead of hardcoding 768/1536 everywhere - Remove duplicated dimension inference logic from main.rs Co-Authored-By: panosAthDBX <[email protected]> Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: harden src/llm/ module from crate audit findings - Replace 9x .expect() on RwLock with graceful poison recovery (nearai.rs: 7, nearai_chat.rs: 2) — eliminates production panics - Propagate HTTP client builder errors in nearai.rs instead of silently dropping timeout config (NearAiProvider::new now returns Result) - Make nearai_chat ChatCompletionResponse.id Optional<String> (mirrors openai_compatible_chat.rs fix for providers that omit id) - Make nearai_chat usage fields optional with defensive parse_usage() helper (was required u32 fields that crash on null/missing) - Truncate error responses to 512 chars in nearai_chat.rs error messages to prevent log bloat and potential data leakage - Delegate 4 missing LlmProvider methods in FailoverProvider (model_metadata, seed_response_chain, get_response_chain_id, calculate_cost) to last-used provider instead of trait defaults Co-Authored-By: Claude Opus 4.6 <[email protected]> * refactor(llm): add RetryProvider, remove openai_compatible_chat, harden decorators - Add composable RetryProvider decorator wrapping any LlmProvider with exponential backoff + jitter, respecting RateLimited retry_after hints - Remove openai_compatible_chat.rs — replaced by rig adapter + RetryProvider - Remove internal retry loop from nearai.rs (was causing double-retry with external RetryProvider, up to 16 attempts instead of 4) - Remove internal retry loop from nearai_chat.rs (same issue) - Wire RetryProvider into main.rs composition chain: each provider gets its own retry wrapper before failover - Move normalize_tool_name to rig_adapter.rs for all rig-based providers - Reconcile is_retryable() vs is_transient() error classification: ModelNotAvailable no longer retryable, Json no longer transient - Fix unchecked Duration subtraction panic in circuit_breaker.rs - Make failover.rs use shared is_retryable() from retry.rs - Remove stale #[allow(dead_code)] on NearAiResponse::id (field is used) Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address PR review feedback — error handling, dimension validation, libSQL warning - Replace response.text().await.unwrap_or_default() with proper error propagation in nearai.rs and nearai_chat.rs (4 call sites). Failures now return LlmError::RequestFailed with context instead of silently proceeding with an empty string. - Add embedding dimension validation in OllamaEmbeddings::embed_batch(): returns EmbeddingError if Ollama returns embeddings with a dimension that doesn't match the configured value. - Add runtime warning when libSQL backend is used with non-1536 embedding dimension, since the libSQL schema uses F32_BLOB(1536) and cannot store different-dimension vectors. Co-Authored-By: Claude Opus 4.6 <[email protected]> * Apply suggestions from code review Co-authored-by: Copilot <[email protected]> --------- Co-authored-by: panosAthDbx <[email protected]> Co-authored-by: panosAthDBX <[email protected]> Co-authored-by: panosAthDBX <[email protected]> Co-authored-by: Claude Opus 4.6 <[email protected]> Co-authored-by: Copilot <[email protected]>
875 lines
31 KiB
Rust
875 lines
31 KiB
Rust
//! Generic adapter that bridges rig-core's `CompletionModel` trait to IronClaw's `LlmProvider`.
|
|
//!
|
|
//! This lets us use any rig-core provider (OpenAI, Anthropic, Ollama, etc.) as an
|
|
//! `Arc<dyn LlmProvider>` without changing any of the agent, reasoning, or tool code.
|
|
|
|
use async_trait::async_trait;
|
|
use rig::OneOrMany;
|
|
use rig::completion::{
|
|
AssistantContent, CompletionModel, CompletionRequest as RigRequest,
|
|
ToolDefinition as RigToolDefinition, Usage as RigUsage,
|
|
};
|
|
use rig::message::{
|
|
Message as RigMessage, ToolChoice as RigToolChoice, ToolFunction, ToolResult as RigToolResult,
|
|
ToolResultContent, UserContent,
|
|
};
|
|
use rust_decimal::Decimal;
|
|
use serde::Serialize;
|
|
use serde::de::DeserializeOwned;
|
|
use serde_json::Value as JsonValue;
|
|
|
|
use std::collections::HashSet;
|
|
|
|
use crate::error::LlmError;
|
|
use crate::llm::costs;
|
|
use crate::llm::provider::{
|
|
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider,
|
|
ToolCall as IronToolCall, ToolCompletionRequest, ToolCompletionResponse,
|
|
ToolDefinition as IronToolDefinition,
|
|
};
|
|
|
|
/// Adapter that wraps a rig-core `CompletionModel` and implements `LlmProvider`.
|
|
pub struct RigAdapter<M: CompletionModel> {
|
|
model: M,
|
|
model_name: String,
|
|
input_cost: Decimal,
|
|
output_cost: Decimal,
|
|
}
|
|
|
|
impl<M: CompletionModel> RigAdapter<M> {
|
|
/// Create a new adapter wrapping the given rig-core model.
|
|
pub fn new(model: M, model_name: impl Into<String>) -> Self {
|
|
let name = model_name.into();
|
|
let (input_cost, output_cost) =
|
|
costs::model_cost(&name).unwrap_or_else(costs::default_cost);
|
|
Self {
|
|
model,
|
|
model_name: name,
|
|
input_cost,
|
|
output_cost,
|
|
}
|
|
}
|
|
}
|
|
|
|
// -- Type conversion helpers --
|
|
|
|
/// Normalize a JSON Schema for OpenAI strict mode compliance.
|
|
///
|
|
/// OpenAI strict function calling requires:
|
|
/// - Every object must have `"additionalProperties": false`
|
|
/// - `"required"` must list ALL property keys
|
|
/// - Optional fields use `"type": ["<original>", "null"]` instead of being omitted from `required`
|
|
/// - Nested objects and array items are recursively normalized
|
|
///
|
|
/// This is applied as a clone-and-transform at the provider boundary so the
|
|
/// original tool definitions remain unchanged for other providers.
|
|
fn normalize_schema_strict(schema: &JsonValue) -> JsonValue {
|
|
let mut schema = schema.clone();
|
|
normalize_schema_recursive(&mut schema);
|
|
schema
|
|
}
|
|
|
|
fn normalize_schema_recursive(schema: &mut JsonValue) {
|
|
let obj = match schema.as_object_mut() {
|
|
Some(o) => o,
|
|
None => return,
|
|
};
|
|
|
|
// Recurse into combinators: anyOf, oneOf, allOf
|
|
for key in &["anyOf", "oneOf", "allOf"] {
|
|
if let Some(JsonValue::Array(variants)) = obj.get_mut(*key) {
|
|
for variant in variants.iter_mut() {
|
|
normalize_schema_recursive(variant);
|
|
}
|
|
}
|
|
}
|
|
|
|
// Recurse into array items
|
|
if let Some(items) = obj.get_mut("items") {
|
|
normalize_schema_recursive(items);
|
|
}
|
|
|
|
// Recurse into `not`, `if`, `then`, `else`
|
|
for key in &["not", "if", "then", "else"] {
|
|
if let Some(sub) = obj.get_mut(*key) {
|
|
normalize_schema_recursive(sub);
|
|
}
|
|
}
|
|
|
|
// Only apply object-level normalization if this schema has "properties"
|
|
// (explicit object schema) or type == "object"
|
|
let is_object = obj
|
|
.get("type")
|
|
.and_then(|t| t.as_str())
|
|
.map(|t| t == "object")
|
|
.unwrap_or(false);
|
|
let has_properties = obj.contains_key("properties");
|
|
|
|
if !is_object && !has_properties {
|
|
return;
|
|
}
|
|
|
|
// Ensure "type": "object" is present
|
|
if !obj.contains_key("type") && has_properties {
|
|
obj.insert("type".to_string(), JsonValue::String("object".to_string()));
|
|
}
|
|
|
|
// Force additionalProperties: false (overwrite any existing value)
|
|
obj.insert("additionalProperties".to_string(), JsonValue::Bool(false));
|
|
|
|
// Ensure "properties" exists
|
|
if !obj.contains_key("properties") {
|
|
obj.insert(
|
|
"properties".to_string(),
|
|
JsonValue::Object(serde_json::Map::new()),
|
|
);
|
|
}
|
|
|
|
// Collect current required set
|
|
let current_required: std::collections::HashSet<String> = obj
|
|
.get("required")
|
|
.and_then(|r| r.as_array())
|
|
.map(|arr| {
|
|
arr.iter()
|
|
.filter_map(|v| v.as_str().map(String::from))
|
|
.collect()
|
|
})
|
|
.unwrap_or_default();
|
|
|
|
// Get all property keys (sorted for deterministic output)
|
|
let all_keys: Vec<String> = obj
|
|
.get("properties")
|
|
.and_then(|p| p.as_object())
|
|
.map(|props| {
|
|
let mut keys: Vec<String> = props.keys().cloned().collect();
|
|
keys.sort();
|
|
keys
|
|
})
|
|
.unwrap_or_default();
|
|
|
|
// For properties NOT in the original required list, make them nullable
|
|
if let Some(JsonValue::Object(props)) = obj.get_mut("properties") {
|
|
for key in &all_keys {
|
|
// Recurse into each property's schema FIRST (before make_nullable,
|
|
// which may change the type to an array and prevent object detection)
|
|
if let Some(prop_schema) = props.get_mut(key) {
|
|
normalize_schema_recursive(prop_schema);
|
|
}
|
|
// Then make originally-optional properties nullable
|
|
if !current_required.contains(key)
|
|
&& let Some(prop_schema) = props.get_mut(key)
|
|
{
|
|
make_nullable(prop_schema);
|
|
}
|
|
}
|
|
}
|
|
|
|
// Set required to ALL property keys
|
|
let required_value: Vec<JsonValue> = all_keys.into_iter().map(JsonValue::String).collect();
|
|
obj.insert("required".to_string(), JsonValue::Array(required_value));
|
|
}
|
|
|
|
/// Make a property schema nullable for OpenAI strict mode.
|
|
///
|
|
/// If it has a simple `"type": "<T>"`, converts to `"type": ["<T>", "null"]`.
|
|
/// If it already has an array type, adds "null" if not present.
|
|
/// Otherwise, wraps with `anyOf: [<existing>, {"type": "null"}]`.
|
|
fn make_nullable(schema: &mut JsonValue) {
|
|
let obj = match schema.as_object_mut() {
|
|
Some(o) => o,
|
|
None => return,
|
|
};
|
|
|
|
if let Some(type_val) = obj.get("type").cloned() {
|
|
match type_val {
|
|
// "type": "string" → "type": ["string", "null"]
|
|
JsonValue::String(ref t) if t != "null" => {
|
|
obj.insert("type".to_string(), serde_json::json!([t, "null"]));
|
|
}
|
|
// "type": ["string", "integer"] → add "null" if missing
|
|
JsonValue::Array(ref arr) => {
|
|
let has_null = arr.iter().any(|v| v.as_str() == Some("null"));
|
|
if !has_null {
|
|
let mut new_arr = arr.clone();
|
|
new_arr.push(JsonValue::String("null".to_string()));
|
|
obj.insert("type".to_string(), JsonValue::Array(new_arr));
|
|
}
|
|
}
|
|
_ => {}
|
|
}
|
|
} else {
|
|
// No "type" key — wrap with anyOf including null
|
|
// (handles enum-only, $ref, or combinator schemas)
|
|
let existing = JsonValue::Object(obj.clone());
|
|
obj.clear();
|
|
obj.insert(
|
|
"anyOf".to_string(),
|
|
serde_json::json!([existing, {"type": "null"}]),
|
|
);
|
|
}
|
|
}
|
|
|
|
/// Convert IronClaw messages to rig-core format.
|
|
///
|
|
/// Returns `(preamble, chat_history)` where preamble is extracted from
|
|
/// any System message and chat_history contains the rest.
|
|
fn convert_messages(messages: &[ChatMessage]) -> (Option<String>, Vec<RigMessage>) {
|
|
let mut preamble: Option<String> = None;
|
|
let mut history = Vec::new();
|
|
|
|
for msg in messages {
|
|
match msg.role {
|
|
crate::llm::Role::System => {
|
|
// Concatenate system messages into preamble
|
|
match preamble {
|
|
Some(ref mut p) => {
|
|
p.push('\n');
|
|
p.push_str(&msg.content);
|
|
}
|
|
None => preamble = Some(msg.content.clone()),
|
|
}
|
|
}
|
|
crate::llm::Role::User => {
|
|
history.push(RigMessage::user(&msg.content));
|
|
}
|
|
crate::llm::Role::Assistant => {
|
|
if let Some(ref tool_calls) = msg.tool_calls {
|
|
// Assistant message with tool calls
|
|
let mut contents: Vec<AssistantContent> = Vec::new();
|
|
if !msg.content.is_empty() {
|
|
contents.push(AssistantContent::text(&msg.content));
|
|
}
|
|
for (idx, tc) in tool_calls.iter().enumerate() {
|
|
let tool_call_id =
|
|
normalized_tool_call_id(Some(tc.id.as_str()), history.len() + idx);
|
|
contents.push(AssistantContent::ToolCall(
|
|
rig::message::ToolCall::new(
|
|
tool_call_id.clone(),
|
|
ToolFunction::new(tc.name.clone(), tc.arguments.clone()),
|
|
)
|
|
.with_call_id(tool_call_id),
|
|
));
|
|
}
|
|
if let Ok(many) = OneOrMany::many(contents) {
|
|
history.push(RigMessage::Assistant {
|
|
id: None,
|
|
content: many,
|
|
});
|
|
} else {
|
|
// Shouldn't happen but fall back to text
|
|
history.push(RigMessage::assistant(&msg.content));
|
|
}
|
|
} else {
|
|
history.push(RigMessage::assistant(&msg.content));
|
|
}
|
|
}
|
|
crate::llm::Role::Tool => {
|
|
// Tool result message: wrap as User { ToolResult }
|
|
let tool_id = normalized_tool_call_id(msg.tool_call_id.as_deref(), history.len());
|
|
history.push(RigMessage::User {
|
|
content: OneOrMany::one(UserContent::ToolResult(RigToolResult {
|
|
id: tool_id.clone(),
|
|
call_id: Some(tool_id),
|
|
content: OneOrMany::one(ToolResultContent::text(&msg.content)),
|
|
})),
|
|
});
|
|
}
|
|
}
|
|
}
|
|
|
|
(preamble, history)
|
|
}
|
|
|
|
/// Responses-style providers require a non-empty tool call ID.
|
|
fn normalized_tool_call_id(raw: Option<&str>, seed: usize) -> String {
|
|
match raw.map(str::trim).filter(|id| !id.is_empty()) {
|
|
Some(id) => id.to_string(),
|
|
None => format!("generated_tool_call_{seed}"),
|
|
}
|
|
}
|
|
|
|
/// Convert IronClaw tool definitions to rig-core format.
|
|
///
|
|
/// Applies OpenAI strict-mode schema normalization to ensure all tool
|
|
/// parameter schemas comply with OpenAI's function calling requirements.
|
|
fn convert_tools(tools: &[IronToolDefinition]) -> Vec<RigToolDefinition> {
|
|
tools
|
|
.iter()
|
|
.map(|t| RigToolDefinition {
|
|
name: t.name.clone(),
|
|
description: t.description.clone(),
|
|
parameters: normalize_schema_strict(&t.parameters),
|
|
})
|
|
.collect()
|
|
}
|
|
|
|
/// Convert IronClaw tool_choice string to rig-core ToolChoice.
|
|
fn convert_tool_choice(choice: Option<&str>) -> Option<RigToolChoice> {
|
|
match choice.map(|s| s.to_lowercase()).as_deref() {
|
|
Some("auto") => Some(RigToolChoice::Auto),
|
|
Some("required") => Some(RigToolChoice::Required),
|
|
Some("none") => Some(RigToolChoice::None),
|
|
_ => None,
|
|
}
|
|
}
|
|
|
|
/// Extract text and tool calls from a rig-core completion response.
|
|
fn extract_response(
|
|
choice: &OneOrMany<AssistantContent>,
|
|
_usage: &RigUsage,
|
|
) -> (Option<String>, Vec<IronToolCall>, FinishReason) {
|
|
let mut text_parts: Vec<String> = Vec::new();
|
|
let mut tool_calls: Vec<IronToolCall> = Vec::new();
|
|
|
|
for content in choice.iter() {
|
|
match content {
|
|
AssistantContent::Text(t) => {
|
|
if !t.text.is_empty() {
|
|
text_parts.push(t.text.clone());
|
|
}
|
|
}
|
|
AssistantContent::ToolCall(tc) => {
|
|
tool_calls.push(IronToolCall {
|
|
id: tc.id.clone(),
|
|
name: tc.function.name.clone(),
|
|
arguments: tc.function.arguments.clone(),
|
|
});
|
|
}
|
|
// Reasoning and Image variants are not mapped to IronClaw types
|
|
_ => {}
|
|
}
|
|
}
|
|
|
|
let text = if text_parts.is_empty() {
|
|
None
|
|
} else {
|
|
Some(text_parts.join(""))
|
|
};
|
|
|
|
let finish = if !tool_calls.is_empty() {
|
|
FinishReason::ToolUse
|
|
} else {
|
|
FinishReason::Stop
|
|
};
|
|
|
|
(text, tool_calls, finish)
|
|
}
|
|
|
|
/// Saturate u64 to u32 for token counts.
|
|
fn saturate_u32(val: u64) -> u32 {
|
|
val.min(u32::MAX as u64) as u32
|
|
}
|
|
|
|
/// Build a rig-core CompletionRequest from our internal types.
|
|
fn build_rig_request(
|
|
preamble: Option<String>,
|
|
mut history: Vec<RigMessage>,
|
|
tools: Vec<RigToolDefinition>,
|
|
tool_choice: Option<RigToolChoice>,
|
|
temperature: Option<f32>,
|
|
max_tokens: Option<u32>,
|
|
) -> Result<RigRequest, LlmError> {
|
|
// rig-core requires at least one message in chat_history
|
|
if history.is_empty() {
|
|
history.push(RigMessage::user("Hello"));
|
|
}
|
|
|
|
let chat_history = OneOrMany::many(history).map_err(|e| LlmError::RequestFailed {
|
|
provider: "rig".to_string(),
|
|
reason: format!("Failed to build chat history: {}", e),
|
|
})?;
|
|
|
|
Ok(RigRequest {
|
|
preamble,
|
|
chat_history,
|
|
documents: Vec::new(),
|
|
tools,
|
|
temperature: temperature.map(|t| t as f64),
|
|
max_tokens: max_tokens.map(|t| t as u64),
|
|
tool_choice,
|
|
additional_params: None,
|
|
})
|
|
}
|
|
|
|
#[async_trait]
|
|
impl<M> LlmProvider for RigAdapter<M>
|
|
where
|
|
M: CompletionModel + Send + Sync + 'static,
|
|
M::Response: Send + Sync + Serialize + DeserializeOwned,
|
|
{
|
|
fn model_name(&self) -> &str {
|
|
&self.model_name
|
|
}
|
|
|
|
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
|
(self.input_cost, self.output_cost)
|
|
}
|
|
|
|
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
|
if let Some(requested_model) = request.model.as_deref()
|
|
&& requested_model != self.model_name.as_str()
|
|
{
|
|
tracing::warn!(
|
|
requested_model = requested_model,
|
|
active_model = %self.model_name,
|
|
"Per-request model override is not supported for this provider; using configured model"
|
|
);
|
|
}
|
|
|
|
let mut messages = request.messages;
|
|
crate::llm::provider::sanitize_tool_messages(&mut messages);
|
|
let (preamble, history) = convert_messages(&messages);
|
|
|
|
let rig_req = build_rig_request(
|
|
preamble,
|
|
history,
|
|
Vec::new(),
|
|
None,
|
|
request.temperature,
|
|
request.max_tokens,
|
|
)?;
|
|
|
|
let response =
|
|
self.model
|
|
.completion(rig_req)
|
|
.await
|
|
.map_err(|e| LlmError::RequestFailed {
|
|
provider: self.model_name.clone(),
|
|
reason: e.to_string(),
|
|
})?;
|
|
|
|
let (text, _tool_calls, finish) = extract_response(&response.choice, &response.usage);
|
|
|
|
Ok(CompletionResponse {
|
|
content: text.unwrap_or_default(),
|
|
input_tokens: saturate_u32(response.usage.input_tokens),
|
|
output_tokens: saturate_u32(response.usage.output_tokens),
|
|
finish_reason: finish,
|
|
response_id: None,
|
|
})
|
|
}
|
|
|
|
async fn complete_with_tools(
|
|
&self,
|
|
request: ToolCompletionRequest,
|
|
) -> Result<ToolCompletionResponse, LlmError> {
|
|
if let Some(requested_model) = request.model.as_deref()
|
|
&& requested_model != self.model_name.as_str()
|
|
{
|
|
tracing::warn!(
|
|
requested_model = requested_model,
|
|
active_model = %self.model_name,
|
|
"Per-request model override is not supported for this provider; using configured model"
|
|
);
|
|
}
|
|
|
|
let known_tool_names: HashSet<String> =
|
|
request.tools.iter().map(|t| t.name.clone()).collect();
|
|
|
|
let mut messages = request.messages;
|
|
crate::llm::provider::sanitize_tool_messages(&mut messages);
|
|
let (preamble, history) = convert_messages(&messages);
|
|
let tools = convert_tools(&request.tools);
|
|
let tool_choice = convert_tool_choice(request.tool_choice.as_deref());
|
|
|
|
let rig_req = build_rig_request(
|
|
preamble,
|
|
history,
|
|
tools,
|
|
tool_choice,
|
|
request.temperature,
|
|
request.max_tokens,
|
|
)?;
|
|
|
|
let response =
|
|
self.model
|
|
.completion(rig_req)
|
|
.await
|
|
.map_err(|e| LlmError::RequestFailed {
|
|
provider: self.model_name.clone(),
|
|
reason: e.to_string(),
|
|
})?;
|
|
|
|
let (text, mut tool_calls, finish) = extract_response(&response.choice, &response.usage);
|
|
|
|
// Normalize tool call names: some proxies prepend "proxy_" prefixes.
|
|
for tc in &mut tool_calls {
|
|
let normalized = normalize_tool_name(&tc.name, &known_tool_names);
|
|
if normalized != tc.name {
|
|
tracing::debug!(
|
|
original = %tc.name,
|
|
normalized = %normalized,
|
|
"Normalized tool call name from provider",
|
|
);
|
|
tc.name = normalized;
|
|
}
|
|
}
|
|
|
|
Ok(ToolCompletionResponse {
|
|
content: text,
|
|
tool_calls,
|
|
input_tokens: saturate_u32(response.usage.input_tokens),
|
|
output_tokens: saturate_u32(response.usage.output_tokens),
|
|
finish_reason: finish,
|
|
response_id: None,
|
|
})
|
|
}
|
|
|
|
fn active_model_name(&self) -> String {
|
|
self.model_name.clone()
|
|
}
|
|
|
|
fn effective_model_name(&self, _requested_model: Option<&str>) -> String {
|
|
self.active_model_name()
|
|
}
|
|
|
|
fn set_model(&self, _model: &str) -> Result<(), LlmError> {
|
|
// rig-core models are baked at construction time.
|
|
// Switching requires creating a new adapter.
|
|
Err(LlmError::RequestFailed {
|
|
provider: self.model_name.clone(),
|
|
reason: "Runtime model switching not supported for rig-core providers. \
|
|
Restart with a different model configured."
|
|
.to_string(),
|
|
})
|
|
}
|
|
}
|
|
|
|
/// Normalize a tool call name returned by an OpenAI-compatible provider.
|
|
///
|
|
/// Some proxies (e.g. VibeProxy) prepend `proxy_` to tool names.
|
|
/// If the returned name doesn't match any known tool but stripping a
|
|
/// `proxy_` prefix yields a match, use the stripped version.
|
|
fn normalize_tool_name(name: &str, known_tools: &HashSet<String>) -> String {
|
|
if known_tools.contains(name) {
|
|
return name.to_string();
|
|
}
|
|
|
|
if let Some(stripped) = name.strip_prefix("proxy_")
|
|
&& known_tools.contains(stripped)
|
|
{
|
|
return stripped.to_string();
|
|
}
|
|
|
|
name.to_string()
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_convert_messages_system_to_preamble() {
|
|
let messages = vec![
|
|
ChatMessage::system("You are a helpful assistant."),
|
|
ChatMessage::user("Hello"),
|
|
];
|
|
let (preamble, history) = convert_messages(&messages);
|
|
assert_eq!(preamble, Some("You are a helpful assistant.".to_string()));
|
|
assert_eq!(history.len(), 1);
|
|
}
|
|
|
|
#[test]
|
|
fn test_convert_messages_multiple_systems_concatenated() {
|
|
let messages = vec![
|
|
ChatMessage::system("System 1"),
|
|
ChatMessage::system("System 2"),
|
|
ChatMessage::user("Hi"),
|
|
];
|
|
let (preamble, history) = convert_messages(&messages);
|
|
assert_eq!(preamble, Some("System 1\nSystem 2".to_string()));
|
|
assert_eq!(history.len(), 1);
|
|
}
|
|
|
|
#[test]
|
|
fn test_convert_messages_tool_result() {
|
|
let messages = vec![ChatMessage::tool_result(
|
|
"call_123",
|
|
"search",
|
|
"result text",
|
|
)];
|
|
let (preamble, history) = convert_messages(&messages);
|
|
assert!(preamble.is_none());
|
|
assert_eq!(history.len(), 1);
|
|
// Tool results become User messages in rig-core
|
|
match &history[0] {
|
|
RigMessage::User { content } => match content.first() {
|
|
UserContent::ToolResult(r) => {
|
|
assert_eq!(r.id, "call_123");
|
|
assert_eq!(r.call_id.as_deref(), Some("call_123"));
|
|
}
|
|
other => panic!("Expected tool result content, got: {:?}", other),
|
|
},
|
|
other => panic!("Expected User message, got: {:?}", other),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_convert_messages_assistant_with_tool_calls() {
|
|
let tc = IronToolCall {
|
|
id: "call_1".to_string(),
|
|
name: "search".to_string(),
|
|
arguments: serde_json::json!({"query": "test"}),
|
|
};
|
|
let msg = ChatMessage::assistant_with_tool_calls(Some("thinking".to_string()), vec![tc]);
|
|
let messages = vec![msg];
|
|
let (_preamble, history) = convert_messages(&messages);
|
|
assert_eq!(history.len(), 1);
|
|
match &history[0] {
|
|
RigMessage::Assistant { content, .. } => {
|
|
// Should have both text and tool call
|
|
assert!(content.iter().count() >= 2);
|
|
for item in content.iter() {
|
|
if let AssistantContent::ToolCall(tc) = item {
|
|
assert_eq!(tc.call_id.as_deref(), Some("call_1"));
|
|
}
|
|
}
|
|
}
|
|
other => panic!("Expected Assistant message, got: {:?}", other),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_convert_messages_tool_result_without_id_gets_fallback() {
|
|
let messages = vec![ChatMessage {
|
|
role: crate::llm::Role::Tool,
|
|
content: "result text".to_string(),
|
|
tool_call_id: None,
|
|
name: Some("search".to_string()),
|
|
tool_calls: None,
|
|
}];
|
|
let (_preamble, history) = convert_messages(&messages);
|
|
match &history[0] {
|
|
RigMessage::User { content } => match content.first() {
|
|
UserContent::ToolResult(r) => {
|
|
assert!(r.id.starts_with("generated_tool_call_"));
|
|
assert_eq!(r.call_id.as_deref(), Some(r.id.as_str()));
|
|
}
|
|
other => panic!("Expected tool result content, got: {:?}", other),
|
|
},
|
|
other => panic!("Expected User message, got: {:?}", other),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_convert_tools() {
|
|
let tools = vec![IronToolDefinition {
|
|
name: "search".to_string(),
|
|
description: "Search the web".to_string(),
|
|
parameters: serde_json::json!({
|
|
"type": "object",
|
|
"properties": {
|
|
"query": {"type": "string"}
|
|
}
|
|
}),
|
|
}];
|
|
let rig_tools = convert_tools(&tools);
|
|
assert_eq!(rig_tools.len(), 1);
|
|
assert_eq!(rig_tools[0].name, "search");
|
|
assert_eq!(rig_tools[0].description, "Search the web");
|
|
}
|
|
|
|
#[test]
|
|
fn test_convert_tool_choice() {
|
|
assert!(matches!(
|
|
convert_tool_choice(Some("auto")),
|
|
Some(RigToolChoice::Auto)
|
|
));
|
|
assert!(matches!(
|
|
convert_tool_choice(Some("required")),
|
|
Some(RigToolChoice::Required)
|
|
));
|
|
assert!(matches!(
|
|
convert_tool_choice(Some("none")),
|
|
Some(RigToolChoice::None)
|
|
));
|
|
assert!(matches!(
|
|
convert_tool_choice(Some("AUTO")),
|
|
Some(RigToolChoice::Auto)
|
|
));
|
|
assert!(convert_tool_choice(None).is_none());
|
|
assert!(convert_tool_choice(Some("unknown")).is_none());
|
|
}
|
|
|
|
#[test]
|
|
fn test_extract_response_text_only() {
|
|
let content = OneOrMany::one(AssistantContent::text("Hello world"));
|
|
let usage = RigUsage::new();
|
|
let (text, calls, finish) = extract_response(&content, &usage);
|
|
assert_eq!(text, Some("Hello world".to_string()));
|
|
assert!(calls.is_empty());
|
|
assert_eq!(finish, FinishReason::Stop);
|
|
}
|
|
|
|
#[test]
|
|
fn test_extract_response_tool_call() {
|
|
let tc = AssistantContent::tool_call("call_1", "search", serde_json::json!({"q": "test"}));
|
|
let content = OneOrMany::one(tc);
|
|
let usage = RigUsage::new();
|
|
let (text, calls, finish) = extract_response(&content, &usage);
|
|
assert!(text.is_none());
|
|
assert_eq!(calls.len(), 1);
|
|
assert_eq!(calls[0].name, "search");
|
|
assert_eq!(finish, FinishReason::ToolUse);
|
|
}
|
|
|
|
#[test]
|
|
fn test_assistant_tool_call_empty_id_gets_generated() {
|
|
let tc = IronToolCall {
|
|
id: "".to_string(),
|
|
name: "search".to_string(),
|
|
arguments: serde_json::json!({"query": "test"}),
|
|
};
|
|
let messages = vec![ChatMessage::assistant_with_tool_calls(None, vec![tc])];
|
|
let (_preamble, history) = convert_messages(&messages);
|
|
|
|
match &history[0] {
|
|
RigMessage::Assistant { content, .. } => {
|
|
let tool_call = content.iter().find_map(|c| match c {
|
|
AssistantContent::ToolCall(tc) => Some(tc),
|
|
_ => None,
|
|
});
|
|
let tc = tool_call.expect("should have a tool call");
|
|
assert!(!tc.id.is_empty(), "tool call id must not be empty");
|
|
assert!(
|
|
tc.id.starts_with("generated_tool_call_"),
|
|
"empty id should be replaced with generated id, got: {}",
|
|
tc.id
|
|
);
|
|
assert_eq!(tc.call_id.as_deref(), Some(tc.id.as_str()));
|
|
}
|
|
other => panic!("Expected Assistant message, got: {:?}", other),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_assistant_tool_call_whitespace_id_gets_generated() {
|
|
let tc = IronToolCall {
|
|
id: " ".to_string(),
|
|
name: "search".to_string(),
|
|
arguments: serde_json::json!({"query": "test"}),
|
|
};
|
|
let messages = vec![ChatMessage::assistant_with_tool_calls(None, vec![tc])];
|
|
let (_preamble, history) = convert_messages(&messages);
|
|
|
|
match &history[0] {
|
|
RigMessage::Assistant { content, .. } => {
|
|
let tool_call = content.iter().find_map(|c| match c {
|
|
AssistantContent::ToolCall(tc) => Some(tc),
|
|
_ => None,
|
|
});
|
|
let tc = tool_call.expect("should have a tool call");
|
|
assert!(
|
|
tc.id.starts_with("generated_tool_call_"),
|
|
"whitespace-only id should be replaced, got: {:?}",
|
|
tc.id
|
|
);
|
|
}
|
|
other => panic!("Expected Assistant message, got: {:?}", other),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_assistant_and_tool_result_missing_ids_share_generated_id() {
|
|
// Simulate: assistant emits a tool call with empty id, then tool
|
|
// result arrives without an id. Both should get deterministic
|
|
// generated ids that match (based on their position in history).
|
|
let tc = IronToolCall {
|
|
id: "".to_string(),
|
|
name: "search".to_string(),
|
|
arguments: serde_json::json!({"query": "test"}),
|
|
};
|
|
let assistant_msg = ChatMessage::assistant_with_tool_calls(None, vec![tc]);
|
|
let tool_result_msg = ChatMessage {
|
|
role: crate::llm::Role::Tool,
|
|
content: "search results here".to_string(),
|
|
tool_call_id: None,
|
|
name: Some("search".to_string()),
|
|
tool_calls: None,
|
|
};
|
|
let messages = vec![assistant_msg, tool_result_msg];
|
|
let (_preamble, history) = convert_messages(&messages);
|
|
|
|
// Extract the generated call_id from the assistant tool call
|
|
let assistant_call_id = match &history[0] {
|
|
RigMessage::Assistant { content, .. } => {
|
|
let tc = content.iter().find_map(|c| match c {
|
|
AssistantContent::ToolCall(tc) => Some(tc),
|
|
_ => None,
|
|
});
|
|
tc.expect("should have tool call").id.clone()
|
|
}
|
|
other => panic!("Expected Assistant message, got: {:?}", other),
|
|
};
|
|
|
|
// Extract the generated call_id from the tool result
|
|
let tool_result_call_id = match &history[1] {
|
|
RigMessage::User { content } => match content.first() {
|
|
UserContent::ToolResult(r) => r
|
|
.call_id
|
|
.clone()
|
|
.expect("tool result call_id must be present"),
|
|
other => panic!("Expected ToolResult, got: {:?}", other),
|
|
},
|
|
other => panic!("Expected User message, got: {:?}", other),
|
|
};
|
|
|
|
assert!(
|
|
!assistant_call_id.is_empty(),
|
|
"assistant call_id must not be empty"
|
|
);
|
|
assert!(
|
|
!tool_result_call_id.is_empty(),
|
|
"tool result call_id must not be empty"
|
|
);
|
|
|
|
// NOTE: With the current seed-based generation, these IDs will differ
|
|
// because the assistant tool call uses seed=0 (history.len() at that
|
|
// point) and the tool result uses seed=1 (history.len() after the
|
|
// assistant message was pushed). This documents the current behavior.
|
|
// A future improvement could thread the assistant's generated ID into
|
|
// the tool result for exact matching.
|
|
assert_ne!(
|
|
assistant_call_id, tool_result_call_id,
|
|
"Current impl generates different IDs for assistant call and tool result \
|
|
because seeds differ; this documents the known limitation"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_saturate_u32() {
|
|
assert_eq!(saturate_u32(100), 100);
|
|
assert_eq!(saturate_u32(u64::MAX), u32::MAX);
|
|
assert_eq!(saturate_u32(u32::MAX as u64), u32::MAX);
|
|
}
|
|
|
|
// -- normalize_tool_name tests --
|
|
|
|
#[test]
|
|
fn test_normalize_tool_name_exact_match() {
|
|
let known = HashSet::from(["echo".to_string(), "list_jobs".to_string()]);
|
|
assert_eq!(normalize_tool_name("echo", &known), "echo");
|
|
}
|
|
|
|
#[test]
|
|
fn test_normalize_tool_name_proxy_prefix_match() {
|
|
let known = HashSet::from(["echo".to_string(), "list_jobs".to_string()]);
|
|
assert_eq!(normalize_tool_name("proxy_echo", &known), "echo");
|
|
}
|
|
|
|
#[test]
|
|
fn test_normalize_tool_name_proxy_prefix_no_match_kept() {
|
|
let known = HashSet::from(["echo".to_string(), "list_jobs".to_string()]);
|
|
assert_eq!(
|
|
normalize_tool_name("proxy_unknown", &known),
|
|
"proxy_unknown"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_normalize_tool_name_unknown_passthrough() {
|
|
let known = HashSet::from(["echo".to_string()]);
|
|
assert_eq!(normalize_tool_name("other_tool", &known), "other_tool");
|
|
}
|
|
}
|