diff --git a/crates/unified-llm/src/catalog.json b/crates/unified-llm/src/catalog.json new file mode 100644 index 000000000..c837a2e21 --- /dev/null +++ b/crates/unified-llm/src/catalog.json @@ -0,0 +1,93 @@ +[ + { + "id": "claude-opus-4-6", + "provider": "anthropic", + "display_name": "Claude Opus 4.6", + "context_window": 200000, + "max_output": null, + "supports_tools": true, + "supports_vision": true, + "supports_reasoning": true, + "input_cost_per_million": 15.0, + "output_cost_per_million": 75.0, + "aliases": ["opus", "claude-opus"] + }, + { + "id": "claude-sonnet-4-5", + "provider": "anthropic", + "display_name": "Claude Sonnet 4.5", + "context_window": 200000, + "max_output": null, + "supports_tools": true, + "supports_vision": true, + "supports_reasoning": true, + "input_cost_per_million": 3.0, + "output_cost_per_million": 15.0, + "aliases": ["sonnet", "claude-sonnet"] + }, + { + "id": "gpt-5.2", + "provider": "openai", + "display_name": "GPT-5.2", + "context_window": 1047576, + "max_output": null, + "supports_tools": true, + "supports_vision": true, + "supports_reasoning": true, + "input_cost_per_million": null, + "output_cost_per_million": null, + "aliases": ["gpt5"] + }, + { + "id": "gpt-5.2-mini", + "provider": "openai", + "display_name": "GPT-5.2 Mini", + "context_window": 1047576, + "max_output": null, + "supports_tools": true, + "supports_vision": true, + "supports_reasoning": true, + "input_cost_per_million": null, + "output_cost_per_million": null, + "aliases": [] + }, + { + "id": "gpt-5.2-codex", + "provider": "openai", + "display_name": "GPT-5.2 Codex", + "context_window": 1047576, + "max_output": null, + "supports_tools": true, + "supports_vision": true, + "supports_reasoning": true, + "input_cost_per_million": null, + "output_cost_per_million": null, + "aliases": ["codex"] + }, + { + "id": "gemini-3-pro-preview", + "provider": "gemini", + "display_name": "Gemini 3 Pro (Preview)", + "context_window": 1048576, + "max_output": null, + "supports_tools": true, + "supports_vision": true, + "supports_reasoning": true, + "input_cost_per_million": null, + "output_cost_per_million": null, + "aliases": ["gemini-pro"] + }, + { + "id": "gemini-3-flash-preview", + "provider": "gemini", + "display_name": "Gemini 3 Flash (Preview)", + "context_window": 1048576, + "max_output": null, + "supports_tools": true, + "supports_vision": true, + "supports_reasoning": true, + "input_cost_per_million": null, + "output_cost_per_million": null, + "aliases": ["gemini-flash"] + } +] diff --git a/crates/unified-llm/src/catalog.rs b/crates/unified-llm/src/catalog.rs index 759711abc..c01b896c7 100644 --- a/crates/unified-llm/src/catalog.rs +++ b/crates/unified-llm/src/catalog.rs @@ -1,105 +1,11 @@ use crate::types::ModelInfo; use std::sync::LazyLock; -/// Built-in model catalog (Section 2.9). +/// Built-in model catalog loaded from catalog.json (Section 2.9). /// The catalog is advisory, not restrictive -- unknown model strings pass through. static BUILT_IN_MODELS: LazyLock> = LazyLock::new(|| { - vec![ - // === Anthropic === - ModelInfo { - id: "claude-opus-4-6".into(), - provider: "anthropic".into(), - display_name: "Claude Opus 4.6".into(), - context_window: 200_000, - max_output: None, - supports_tools: true, - supports_vision: true, - supports_reasoning: true, - input_cost_per_million: Some(15.0), - output_cost_per_million: Some(75.0), - aliases: vec!["opus".into(), "claude-opus".into()], - }, - ModelInfo { - id: "claude-sonnet-4-5".into(), - provider: "anthropic".into(), - display_name: "Claude Sonnet 4.5".into(), - context_window: 200_000, - max_output: None, - supports_tools: true, - supports_vision: true, - supports_reasoning: true, - input_cost_per_million: Some(3.0), - output_cost_per_million: Some(15.0), - aliases: vec!["sonnet".into(), "claude-sonnet".into()], - }, - // === OpenAI === - ModelInfo { - id: "gpt-5.2".into(), - provider: "openai".into(), - display_name: "GPT-5.2".into(), - context_window: 1_047_576, - max_output: None, - supports_tools: true, - supports_vision: true, - supports_reasoning: true, - input_cost_per_million: None, - output_cost_per_million: None, - aliases: vec!["gpt5".into()], - }, - ModelInfo { - id: "gpt-5.2-mini".into(), - provider: "openai".into(), - display_name: "GPT-5.2 Mini".into(), - context_window: 1_047_576, - max_output: None, - supports_tools: true, - supports_vision: true, - supports_reasoning: true, - input_cost_per_million: None, - output_cost_per_million: None, - aliases: vec![], - }, - ModelInfo { - id: "gpt-5.2-codex".into(), - provider: "openai".into(), - display_name: "GPT-5.2 Codex".into(), - context_window: 1_047_576, - max_output: None, - supports_tools: true, - supports_vision: true, - supports_reasoning: true, - input_cost_per_million: None, - output_cost_per_million: None, - aliases: vec!["codex".into()], - }, - // === Gemini === - ModelInfo { - id: "gemini-3-pro-preview".into(), - provider: "gemini".into(), - display_name: "Gemini 3 Pro (Preview)".into(), - context_window: 1_048_576, - max_output: None, - supports_tools: true, - supports_vision: true, - supports_reasoning: true, - input_cost_per_million: None, - output_cost_per_million: None, - aliases: vec!["gemini-pro".into()], - }, - ModelInfo { - id: "gemini-3-flash-preview".into(), - provider: "gemini".into(), - display_name: "Gemini 3 Flash (Preview)".into(), - context_window: 1_048_576, - max_output: None, - supports_tools: true, - supports_vision: true, - supports_reasoning: true, - input_cost_per_million: None, - output_cost_per_million: None, - aliases: vec!["gemini-flash".into()], - }, - ] + serde_json::from_str(include_str!("catalog.json")) + .expect("embedded catalog.json must be valid") }); /// Get model info by model ID (Section 2.9). diff --git a/crates/unified-llm/src/client.rs b/crates/unified-llm/src/client.rs index 5d4e00c37..6e699d788 100644 --- a/crates/unified-llm/src/client.rs +++ b/crates/unified-llm/src/client.rs @@ -31,8 +31,11 @@ impl Client { /// Create a Client from environment variables (Section 2.2). /// Registers providers whose API keys are present in the environment. /// The first registered provider becomes the default. - #[must_use] - pub fn from_env() -> Self { + /// + /// # Errors + /// + /// Returns `SdkError` if any provider adapter fails to initialize. + pub async fn from_env() -> Result { let mut client = Self { providers: HashMap::new(), default_provider: None, @@ -46,7 +49,7 @@ impl Client { if let Ok(base_url) = std::env::var("ANTHROPIC_BASE_URL") { adapter = adapter.with_base_url(base_url); } - client.register_provider(Arc::new(adapter)); + client.register_provider(Arc::new(adapter)).await?; } if let Ok(key) = std::env::var("OPENAI_API_KEY") { let mut adapter = providers::OpenAiAdapter::new(key); @@ -59,7 +62,7 @@ impl Client { if let Ok(project_id) = std::env::var("OPENAI_PROJECT_ID") { adapter = adapter.with_project_id(project_id); } - client.register_provider(Arc::new(adapter)); + client.register_provider(Arc::new(adapter)).await?; } if let Ok(key) = std::env::var("GEMINI_API_KEY").or_else(|_| std::env::var("GOOGLE_API_KEY")) { @@ -67,19 +70,28 @@ impl Client { if let Ok(base_url) = std::env::var("GEMINI_BASE_URL") { adapter = adapter.with_base_url(base_url); } - client.register_provider(Arc::new(adapter)); + client.register_provider(Arc::new(adapter)).await?; } - client + Ok(client) } - /// Register a provider adapter. - pub fn register_provider(&mut self, adapter: Arc) { + /// Register a provider adapter. Calls `initialize()` on the adapter (Section 2.4). + /// + /// # Errors + /// + /// Returns `SdkError` if the adapter's `initialize()` method fails. + pub async fn register_provider( + &mut self, + adapter: Arc, + ) -> Result<(), SdkError> { + adapter.initialize().await?; let name = adapter.name().to_string(); if self.default_provider.is_none() { self.default_provider = Some(name.clone()); } self.providers.insert(name, adapter); + Ok(()) } /// Add middleware. @@ -292,7 +304,7 @@ mod tests { #[tokio::test] async fn complete_routes_to_default_provider() { let mut client = Client::new(HashMap::new(), None, vec![]); - client.register_provider(Arc::new(MockProvider::new("test", "Hello!"))); + client.register_provider(Arc::new(MockProvider::new("test", "Hello!"))).await.unwrap(); let response = client.complete(&test_request()).await.unwrap(); assert_eq!(response.text(), "Hello!"); @@ -302,8 +314,8 @@ mod tests { #[tokio::test] async fn complete_routes_to_named_provider() { let mut client = Client::new(HashMap::new(), None, vec![]); - client.register_provider(Arc::new(MockProvider::new("provider_a", "from A"))); - client.register_provider(Arc::new(MockProvider::new("provider_b", "from B"))); + client.register_provider(Arc::new(MockProvider::new("provider_a", "from A"))).await.unwrap(); + client.register_provider(Arc::new(MockProvider::new("provider_b", "from B"))).await.unwrap(); let mut req = test_request(); req.provider = Some("provider_b".into()); @@ -325,7 +337,7 @@ mod tests { #[tokio::test] async fn complete_errors_on_unknown_provider() { let mut client = Client::new(HashMap::new(), None, vec![]); - client.register_provider(Arc::new(MockProvider::new("test", "Hello"))); + client.register_provider(Arc::new(MockProvider::new("test", "Hello"))).await.unwrap(); let mut req = test_request(); req.provider = Some("nonexistent".into()); @@ -342,10 +354,10 @@ mod tests { let mut client = Client::new(HashMap::new(), None, vec![]); assert_eq!(client.default_provider(), None); - client.register_provider(Arc::new(MockProvider::new("first", "1"))); + client.register_provider(Arc::new(MockProvider::new("first", "1"))).await.unwrap(); assert_eq!(client.default_provider(), Some("first")); - client.register_provider(Arc::new(MockProvider::new("second", "2"))); + client.register_provider(Arc::new(MockProvider::new("second", "2"))).await.unwrap(); assert_eq!(client.default_provider(), Some("first")); } @@ -354,7 +366,7 @@ mod tests { use futures::StreamExt; let mut client = Client::new(HashMap::new(), None, vec![]); - client.register_provider(Arc::new(MockProvider::new("test", "streamed"))); + client.register_provider(Arc::new(MockProvider::new("test", "streamed"))).await.unwrap(); let mut stream = client.stream(&test_request()).await.unwrap(); let first = stream.next().await.unwrap().unwrap(); @@ -367,8 +379,8 @@ mod tests { #[tokio::test] async fn provider_names_returns_registered() { let mut client = Client::new(HashMap::new(), None, vec![]); - client.register_provider(Arc::new(MockProvider::new("alpha", ""))); - client.register_provider(Arc::new(MockProvider::new("beta", ""))); + client.register_provider(Arc::new(MockProvider::new("alpha", ""))).await.unwrap(); + client.register_provider(Arc::new(MockProvider::new("beta", ""))).await.unwrap(); let mut names = client.provider_names(); names.sort_unstable(); assert_eq!(names, vec!["alpha", "beta"]); @@ -402,7 +414,7 @@ mod tests { #[tokio::test] async fn middleware_wraps_complete() { let mut client = Client::new(HashMap::new(), None, vec![]); - client.register_provider(Arc::new(MockProvider::new("test", "hello"))); + client.register_provider(Arc::new(MockProvider::new("test", "hello"))).await.unwrap(); client.add_middleware(Arc::new(UppercaseMiddleware)); let response = client.complete(&test_request()).await.unwrap(); diff --git a/crates/unified-llm/src/error.rs b/crates/unified-llm/src/error.rs index b68781fc7..cd878c95c 100644 --- a/crates/unified-llm/src/error.rs +++ b/crates/unified-llm/src/error.rs @@ -1,4 +1,5 @@ -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[serde(rename_all = "snake_case")] pub enum ProviderErrorKind { Authentication, AccessDenied, @@ -27,7 +28,7 @@ impl std::fmt::Display for ProviderErrorKind { } } -#[derive(Debug, Clone)] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct ProviderErrorDetail { pub message: String, pub provider: String, @@ -50,7 +51,8 @@ impl ProviderErrorDetail { } } -#[derive(Debug, thiserror::Error)] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, thiserror::Error)] +#[serde(tag = "type", rename_all = "snake_case")] pub enum SdkError { #[error("{kind} {}: {}", .detail.provider, .detail.message)] Provider { diff --git a/crates/unified-llm/src/generate.rs b/crates/unified-llm/src/generate.rs index 313c84b7b..1e4f8603e 100644 --- a/crates/unified-llm/src/generate.rs +++ b/crates/unified-llm/src/generate.rs @@ -24,11 +24,13 @@ pub fn set_default_client(client: Client) { } /// Get the default client, lazily initialized from env. -fn get_default_client() -> Arc { - DEFAULT_CLIENT - .get() - .cloned() - .unwrap_or_else(|| Arc::new(Client::from_env())) +async fn get_default_client() -> Result, SdkError> { + if let Some(client) = DEFAULT_CLIENT.get() { + return Ok(client.clone()); + } + let client = Arc::new(Client::from_env().await?); + let _ = DEFAULT_CLIENT.set(client.clone()); + Ok(client) } fn build_initial_messages(params: &GenerateParams) -> Result, SdkError> { @@ -99,7 +101,10 @@ fn build_generate_result(steps: Vec, total_usage: Usage) -> Generate /// Panics if a tool's `execute` handler is `None` when matched during tool execution. #[allow(clippy::too_many_lines)] pub async fn generate(params: GenerateParams) -> Result { - let client = params.client.clone().unwrap_or_else(get_default_client); + let client = match params.client.clone() { + Some(c) => c, + None => get_default_client().await?, + }; let retry_policy = RetryPolicy { max_retries: params.max_retries, base_delay: 0.001, @@ -167,7 +172,7 @@ pub async fn generate(params: GenerateParams) -> Result = tools.iter().map(std::convert::AsRef::as_ref).collect(); - tool_results = execute_all_tools(&tool_refs, &tool_calls).await; + tool_results = execute_all_tools(&tool_refs, &tool_calls, &messages, abort_signal.as_ref()).await; } } @@ -570,7 +575,10 @@ pub async fn stream(params: GenerateParams) -> Result { /// or any provider error encountered during streaming setup. #[allow(clippy::too_many_lines)] async fn stream_with_tool_loop(params: GenerateParams) -> Result { - let client = params.client.clone().unwrap_or_else(get_default_client); + let client = match params.client.clone() { + Some(c) => c, + None => get_default_client().await?, + }; let mut messages = build_initial_messages(¶ms)?; let tool_definitions: Option> = params .tools @@ -695,7 +703,7 @@ async fn stream_with_tool_loop(params: GenerateParams) -> Result = tool_list.iter().map(std::convert::AsRef::as_ref).collect(); - let tool_results = execute_all_tools(&tool_refs, &tool_calls).await; + let tool_results = execute_all_tools(&tool_refs, &tool_calls, &messages, abort_signal.as_ref()).await; if tool_results.is_empty() { return; @@ -804,7 +812,10 @@ async fn stream_generate_raw( /// Returns `SdkError::Configuration` if both `prompt` and `messages` are set, /// or any provider error encountered during streaming setup. pub async fn stream_generate(params: GenerateParams) -> Result { - let client = params.client.clone().unwrap_or_else(get_default_client); + let client = match params.client.clone() { + Some(c) => c, + None => get_default_client().await?, + }; let messages = build_initial_messages(¶ms)?; let tool_definitions: Option> = params .tools @@ -1188,7 +1199,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |args| async move { + |args, _ctx| async move { let city = args["city"].as_str().unwrap_or("unknown"); Ok(serde_json::json!(format!("72F in {}", city))) }, @@ -1348,7 +1359,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |args| async move { + |args, _ctx| async move { let city = args["city"].as_str().unwrap_or("unknown"); Ok(serde_json::json!(format!("72F in {}", city))) }, @@ -1714,7 +1725,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |_args| async { Ok(serde_json::json!("72F")) }, + |_args, _ctx| async { Ok(serde_json::json!("72F")) }, )]) .max_tool_rounds(10) .abort_signal(token) @@ -1776,7 +1787,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - move |_args| { + move |_args, _ctx| { let counter = tool_executed_clone.clone(); async move { counter.fetch_add(1, Ordering::SeqCst); @@ -1965,7 +1976,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |_args| async { Ok(serde_json::json!("72F")) }, + |_args, _ctx| async { Ok(serde_json::json!("72F")) }, )]) .max_tool_rounds(5) .client(client), @@ -2019,7 +2030,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |_args| async { Ok(serde_json::json!("72F")) }, + |_args, _ctx| async { Ok(serde_json::json!("72F")) }, )]) .max_tool_rounds(0) .client(client), @@ -2105,7 +2116,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |_args| async { Ok(serde_json::json!("72F")) }, + |_args, _ctx| async { Ok(serde_json::json!("72F")) }, )]) .max_tool_rounds(5) .client(client), @@ -2168,7 +2179,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |_args| async { Ok(serde_json::json!("72F")) }, + |_args, _ctx| async { Ok(serde_json::json!("72F")) }, )]) .max_tool_rounds(5) .stop_when(|_steps| true) // Stop immediately after first round @@ -2295,7 +2306,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |_args| async { Ok(serde_json::json!("72F")) }, + |_args, _ctx| async { Ok(serde_json::json!("72F")) }, )]) .max_tool_rounds(1) .max_retries(3) @@ -2395,7 +2406,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |_args| async { Ok(serde_json::json!("72F")) }, + |_args, _ctx| async { Ok(serde_json::json!("72F")) }, )]) .max_tool_rounds(1) .timeout(TimeoutConfig { @@ -2528,7 +2539,7 @@ mod tests { "get_weather", "Get weather", serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), - |_args| async { Ok(serde_json::json!("72F")) }, + |_args, _ctx| async { Ok(serde_json::json!("72F")) }, )]) .max_tool_rounds(5) .timeout(TimeoutConfig { diff --git a/crates/unified-llm/src/providers/anthropic.rs b/crates/unified-llm/src/providers/anthropic.rs index 27829e1cb..721c25e3c 100644 --- a/crates/unified-llm/src/providers/anthropic.rs +++ b/crates/unified-llm/src/providers/anthropic.rs @@ -140,6 +140,24 @@ struct ApiUsage { cache_creation_input_tokens: Option, } +/// Estimate reasoning tokens from thinking content blocks. +/// Anthropic does not provide a separate reasoning token count, +/// so we estimate by dividing the character count of thinking text by 4. +fn estimate_reasoning_tokens(content_parts: &[ContentPart]) -> Option { + let total_chars: usize = content_parts + .iter() + .filter_map(|part| match part { + ContentPart::Thinking(td) => Some(td.text.len()), + _ => None, + }) + .sum(); + if total_chars > 0 { + Some((total_chars / 4).max(1) as i64) + } else { + None + } +} + fn map_finish_reason(stop_reason: Option<&str>) -> FinishReason { match stop_reason { Some("end_turn" | "stop_sequence") | None => FinishReason::Stop, @@ -257,6 +275,7 @@ fn content_part_to_api(part: &ContentPart) -> Option { ContentPart::Audio(_) => { Some(serde_json::json!({"type": "text", "text": "[Audio content not supported by this provider]"})) } + ContentPart::Other { .. } => None, } } @@ -834,6 +853,7 @@ impl StreamAccumulator { } fn handle_message_stop(&mut self) -> Vec { + self.usage.reasoning_tokens = estimate_reasoning_tokens(&self.content_parts); let response = self.take_response(); vec![StreamEvent::Finish { finish_reason: response.finish_reason.clone(), @@ -1091,6 +1111,7 @@ impl ProviderAdapter for Adapter { map_finish_reason(api_resp.stop_reason.as_deref()) }; let total = api_resp.usage.input_tokens + api_resp.usage.output_tokens; + let reasoning_tokens = estimate_reasoning_tokens(&content_parts); Ok(Response { id: api_resp.id, @@ -1107,6 +1128,7 @@ impl ProviderAdapter for Adapter { input_tokens: api_resp.usage.input_tokens, output_tokens: api_resp.usage.output_tokens, total_tokens: total, + reasoning_tokens, cache_read_tokens: api_resp.usage.cache_read_input_tokens, cache_write_tokens: api_resp.usage.cache_creation_input_tokens, ..Usage::default() diff --git a/crates/unified-llm/src/providers/gemini.rs b/crates/unified-llm/src/providers/gemini.rs index 381ca0398..ace9f86c3 100644 --- a/crates/unified-llm/src/providers/gemini.rs +++ b/crates/unified-llm/src/providers/gemini.rs @@ -1,11 +1,10 @@ use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine}; use futures::stream; -use crate::error::{error_from_status_code, ProviderErrorDetail, ProviderErrorKind, SdkError}; +use crate::error::{error_from_grpc_status, error_from_status_code, ProviderErrorDetail, ProviderErrorKind, SdkError}; use crate::provider::{ProviderAdapter, StreamEventStream}; use crate::providers::common::{ extract_system_prompt, parse_error_body, parse_rate_limit_headers, parse_retry_after, - send_and_read_response, }; use crate::types::{ ContentPart, FinishReason, Message, Request, Response, ResponseFormat, ResponseFormatType, @@ -462,10 +461,60 @@ fn parse_usage(metadata: Option<&UsageMetadata>) -> Usage { }) } +/// Send an HTTP request and read the Gemini response body. +/// +/// Like `send_and_read_response` but uses gRPC status code mapping when available. +async fn send_gemini_response( + request: reqwest::RequestBuilder, +) -> Result<(String, reqwest::header::HeaderMap), SdkError> { + let http_resp = request.send().await.map_err(|e| { + if e.is_timeout() { + SdkError::RequestTimeout { + message: format!("gemini: {e}"), + } + } else { + SdkError::Network { + message: e.to_string(), + } + } + })?; + + let status = http_resp.status(); + let retry_after = parse_retry_after(http_resp.headers()); + let headers = http_resp.headers().clone(); + let body = http_resp + .text() + .await + .map_err(|e| SdkError::Network { + message: e.to_string(), + })?; + + if !status.is_success() { + let (msg, code, raw) = parse_error_body(&body, "status"); + return Err(gemini_error(status.as_u16(), msg, code, raw, retry_after)); + } + + Ok((body, headers)) +} + +/// Map Gemini error response using gRPC status when available, falling back to HTTP status. +fn gemini_error( + status_code: u16, + msg: String, + grpc_status: Option, + raw: Option, + retry_after: Option, +) -> SdkError { + match grpc_status { + Some(grpc_code) => error_from_grpc_status(&grpc_code, msg, "gemini".to_string(), Some(grpc_code.clone()), raw, retry_after), + None => error_from_status_code(status_code, msg, "gemini".to_string(), None, raw, retry_after), + } +} + /// Send an HTTP request for streaming and return the `reqwest::Response`. /// /// Checks for HTTP errors before returning. On error, reads the body and -/// maps it to `SdkError` using the same logic as `send_and_read_body`. +/// maps it to `SdkError` using gRPC status code mapping when available. async fn send_streaming_request( request: reqwest::RequestBuilder, ) -> Result { @@ -480,14 +529,7 @@ async fn send_streaming_request( message: e.to_string(), })?; let (msg, code, raw) = parse_error_body(&body, "status"); - return Err(error_from_status_code( - status.as_u16(), - msg, - "gemini".to_string(), - code, - raw, - retry_after, - )); + return Err(gemini_error(status.as_u16(), msg, code, raw, retry_after)); } Ok(http_resp) @@ -784,10 +826,8 @@ impl ProviderAdapter for Adapter { for (key, value) in &self.default_headers { req = req.header(key, value); } - let (body, headers) = send_and_read_response( + let (body, headers) = send_gemini_response( req.json(&api_body).timeout(self.request_timeout), - "gemini", - "status", ) .await?; @@ -1062,4 +1102,38 @@ mod tests { assert_eq!(part["inlineData"]["mimeType"], "application/pdf"); assert!(part["inlineData"]["data"].as_str().is_some()); } + + #[test] + fn gemini_error_uses_grpc_status_when_available() { + use crate::error::ProviderErrorKind; + + let err = gemini_error(400, "model not found".into(), Some("NOT_FOUND".into()), None, None); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::NotFound, .. })); + + let err = gemini_error(400, "bad args".into(), Some("INVALID_ARGUMENT".into()), None, None); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::InvalidRequest, .. })); + + let err = gemini_error(429, "rate limited".into(), Some("RESOURCE_EXHAUSTED".into()), None, None); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::RateLimit, .. })); + + let err = gemini_error(401, "bad key".into(), Some("UNAUTHENTICATED".into()), None, None); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Authentication, .. })); + + let err = gemini_error(403, "denied".into(), Some("PERMISSION_DENIED".into()), None, None); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::AccessDenied, .. })); + + let err = gemini_error(504, "timeout".into(), Some("DEADLINE_EXCEEDED".into()), None, None); + assert!(matches!(err, SdkError::RequestTimeout { .. })); + } + + #[test] + fn gemini_error_falls_back_to_http_status_without_grpc() { + use crate::error::ProviderErrorKind; + + let err = gemini_error(429, "rate limited".into(), None, None, None); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::RateLimit, .. })); + + let err = gemini_error(500, "internal".into(), None, None, None); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Server, .. })); + } } diff --git a/crates/unified-llm/src/tools.rs b/crates/unified-llm/src/tools.rs index 590977b3e..f6cef8f98 100644 --- a/crates/unified-llm/src/tools.rs +++ b/crates/unified-llm/src/tools.rs @@ -1,11 +1,20 @@ -use crate::types::{ToolCall, ToolDefinition, ToolResult}; +use crate::types::{Message, ToolCall, ToolDefinition, ToolResult}; use std::future::Future; use std::pin::Pin; use std::sync::Arc; +use tokio_util::sync::CancellationToken; + +/// Context passed to tool execute handlers (Section 5.2). +#[derive(Clone)] +pub struct ToolContext { + pub tool_call_id: String, + pub messages: Vec, + pub abort_signal: Option, +} /// An execute handler for a tool. pub type ExecuteHandler = Arc< - dyn Fn(serde_json::Value) -> Pin> + Send>> + dyn Fn(serde_json::Value, ToolContext) -> Pin> + Send>> + Send + Sync, >; @@ -51,7 +60,7 @@ impl Tool { handler: F, ) -> Self where - F: Fn(serde_json::Value) -> Fut + Send + Sync + 'static, + F: Fn(serde_json::Value, ToolContext) -> Fut + Send + Sync + 'static, Fut: Future> + Send + 'static, { if let Err(e) = validate_tool_name(name) { @@ -63,7 +72,7 @@ impl Tool { description: description.to_string(), parameters, }, - execute: Some(Arc::new(move |args| Box::pin(handler(args)))), + execute: Some(Arc::new(move |args, ctx| Box::pin(handler(args, ctx)))), } } @@ -115,6 +124,8 @@ pub fn validate_tool_name(name: &str) -> Result<(), String> { pub async fn execute_all_tools( tools: &[&Tool], tool_calls: &[ToolCall], + messages: &[Message], + abort_signal: Option<&CancellationToken>, ) -> Vec { use futures::future::join_all; @@ -125,12 +136,17 @@ pub async fn execute_all_tools( let call_id = call.id.clone(); let call_name = call.name.clone(); let args = call.arguments.clone(); + let ctx = ToolContext { + tool_call_id: call_id.clone(), + messages: messages.to_vec(), + abort_signal: abort_signal.cloned(), + }; async move { match tool { Some(t) if t.execute.is_some() => { let handler = t.execute.as_ref().unwrap(); - match handler(args).await { + match handler(args, ctx).await { Ok(result) => ToolResult { tool_call_id: call_id, content: result, @@ -164,6 +180,162 @@ pub async fn execute_all_tools( join_all(futures).await } +/// A callback to repair invalid tool call arguments (Section 5.8). +/// Receives the tool call and the validation error message, returns repaired arguments +/// or an error if repair is not possible. +pub type RepairToolCallFn = Arc< + dyn Fn( + ToolCall, + String, + ) -> Pin> + Send>> + + Send + + Sync, +>; + +/// Validate tool call arguments against the tool's parameter schema. +/// Performs a lightweight structural check: verifies that when the schema +/// specifies `"type": "object"`, the arguments are a JSON object, and that +/// required properties are present. +fn validate_tool_args(args: &serde_json::Value, schema: &serde_json::Value) -> Result<(), String> { + let schema_type = schema.get("type").and_then(serde_json::Value::as_str); + if schema_type == Some("object") && !args.is_object() { + return Err(format!( + "Expected object arguments, got {}", + args_type_name(args) + )); + } + if let (Some(obj), Some(required)) = ( + args.as_object(), + schema.get("required").and_then(serde_json::Value::as_array), + ) { + let missing: Vec<&str> = required + .iter() + .filter_map(serde_json::Value::as_str) + .filter(|key| !obj.contains_key(*key)) + .collect(); + if !missing.is_empty() { + return Err(format!("Missing required properties: {}", missing.join(", "))); + } + } + Ok(()) +} + +fn args_type_name(value: &serde_json::Value) -> &'static str { + match value { + serde_json::Value::Null => "null", + serde_json::Value::Bool(_) => "boolean", + serde_json::Value::Number(_) => "number", + serde_json::Value::String(_) => "string", + serde_json::Value::Array(_) => "array", + serde_json::Value::Object(_) => "object", + } +} + +/// Execute all tool calls with optional schema validation and repair (Section 5.8). +/// +/// Before calling a tool's execute handler, validates the arguments against the +/// tool's parameter schema. If validation fails and a `repair` callback is provided, +/// calls it to attempt repair. If repair succeeds, uses the repaired arguments. +/// If repair fails or is not configured, returns an error `ToolResult`. +pub async fn execute_all_tools_with_repair( + tools: &[&Tool], + tool_calls: &[ToolCall], + messages: &[Message], + abort_signal: Option<&CancellationToken>, + repair: Option<&RepairToolCallFn>, +) -> Vec { + use futures::future::join_all; + + let futures: Vec<_> = tool_calls + .iter() + .map(|call| { + let tool = tools.iter().find(|t| t.definition.name == call.name).copied(); + let call_id = call.id.clone(); + let call_name = call.name.clone(); + let args = call.arguments.clone(); + let call_clone = call.clone(); + let ctx = ToolContext { + tool_call_id: call_id.clone(), + messages: messages.to_vec(), + abort_signal: abort_signal.cloned(), + }; + + async move { + let Some(t) = tool else { + return ToolResult { + tool_call_id: call_id, + content: serde_json::Value::String(format!("Unknown tool: {call_name}")), + is_error: true, + image_data: None, + image_media_type: None, + }; + }; + + let Some(handler) = &t.execute else { + return ToolResult { + tool_call_id: call_id, + content: serde_json::Value::String(format!("Unknown tool: {call_name}")), + is_error: true, + image_data: None, + image_media_type: None, + }; + }; + + let validated_args = match validate_tool_args(&args, &t.definition.parameters) { + Ok(()) => args, + Err(validation_error) => { + if let Some(repair_fn) = repair { + match repair_fn(call_clone, validation_error).await { + Ok(repaired) => repaired, + Err(repair_error) => { + return ToolResult { + tool_call_id: call_id, + content: serde_json::Value::String(format!( + "Tool call validation failed and repair failed: {repair_error}" + )), + is_error: true, + image_data: None, + image_media_type: None, + }; + } + } + } else { + return ToolResult { + tool_call_id: call_id, + content: serde_json::Value::String(format!( + "Tool call validation failed: {validation_error}" + )), + is_error: true, + image_data: None, + image_media_type: None, + }; + } + } + }; + + match handler(validated_args, ctx).await { + Ok(result) => ToolResult { + tool_call_id: call_id, + content: result, + is_error: false, + image_data: None, + image_media_type: None, + }, + Err(err_msg) => ToolResult { + tool_call_id: call_id, + content: serde_json::Value::String(err_msg), + is_error: true, + image_data: None, + image_media_type: None, + }, + } + } + }) + .collect(); + + join_all(futures).await +} + #[cfg(test)] mod tests { use super::*; @@ -224,7 +396,7 @@ mod tests { "test", "test tool", serde_json::json!({"type": "object", "properties": {}}), - |_args| async { Ok(serde_json::json!("result")) }, + |_args, _ctx| async { Ok(serde_json::json!("result")) }, ); assert!(tool.is_active()); } @@ -235,7 +407,7 @@ mod tests { "greet", "Greet someone", serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}}), - |args| async move { + |args, _ctx| async move { let name = args["name"].as_str().unwrap_or("world"); Ok(serde_json::json!(format!("Hello, {}!", name))) }, @@ -248,7 +420,7 @@ mod tests { )]; let tool_refs: Vec<&Tool> = tools.iter().collect(); - let results = execute_all_tools(&tool_refs, &calls).await; + let results = execute_all_tools(&tool_refs, &calls, &[], None).await; assert_eq!(results.len(), 1); assert_eq!(results[0].tool_call_id, "call_1"); assert!(!results[0].is_error); @@ -266,7 +438,7 @@ mod tests { )]; let tool_refs: Vec<&Tool> = tools.iter().collect(); - let results = execute_all_tools(&tool_refs, &calls).await; + let results = execute_all_tools(&tool_refs, &calls, &[], None).await; assert_eq!(results.len(), 1); assert!(results[0].is_error); assert!(results[0] @@ -282,7 +454,7 @@ mod tests { "fail", "Always fails", serde_json::json!({"type": "object", "properties": {}}), - |_args| async { Err("something went wrong".to_string()) }, + |_args, _ctx| async { Err("something went wrong".to_string()) }, )]; let calls = vec![ToolCall::new( @@ -292,7 +464,7 @@ mod tests { )]; let tool_refs: Vec<&Tool> = tools.iter().collect(); - let results = execute_all_tools(&tool_refs, &calls).await; + let results = execute_all_tools(&tool_refs, &calls, &[], None).await; assert_eq!(results.len(), 1); assert!(results[0].is_error); assert_eq!( @@ -308,13 +480,13 @@ mod tests { "tool_a", "Tool A", serde_json::json!({"type": "object", "properties": {}}), - |_args| async { Ok(serde_json::json!("result_a")) }, + |_args, _ctx| async { Ok(serde_json::json!("result_a")) }, ), Tool::active( "tool_b", "Tool B", serde_json::json!({"type": "object", "properties": {}}), - |_args| async { Ok(serde_json::json!("result_b")) }, + |_args, _ctx| async { Ok(serde_json::json!("result_b")) }, ), ]; @@ -324,7 +496,7 @@ mod tests { ]; let tool_refs: Vec<&Tool> = tools.iter().collect(); - let results = execute_all_tools(&tool_refs, &calls).await; + let results = execute_all_tools(&tool_refs, &calls, &[], None).await; assert_eq!(results.len(), 2); assert_eq!(results[0].tool_call_id, "call_1"); assert_eq!(results[0].content, serde_json::json!("result_a")); @@ -339,13 +511,13 @@ mod tests { "succeed", "Succeeds", serde_json::json!({"type": "object", "properties": {}}), - |_args| async { Ok(serde_json::json!("ok")) }, + |_args, _ctx| async { Ok(serde_json::json!("ok")) }, ), Tool::active( "fail", "Fails", serde_json::json!({"type": "object", "properties": {}}), - |_args| async { Err("boom".to_string()) }, + |_args, _ctx| async { Err("boom".to_string()) }, ), ]; @@ -355,7 +527,7 @@ mod tests { ]; let tool_refs: Vec<&Tool> = tools.iter().collect(); - let results = execute_all_tools(&tool_refs, &calls).await; + let results = execute_all_tools(&tool_refs, &calls, &[], None).await; assert_eq!(results.len(), 2); assert!(!results[0].is_error); assert!(results[1].is_error); @@ -378,7 +550,129 @@ mod tests { "my-tool", "bad name", serde_json::json!({"type": "object"}), - |_args| async { Ok(serde_json::json!("result")) }, + |_args, _ctx| async { Ok(serde_json::json!("result")) }, ); } + + #[test] + fn validate_tool_args_valid_object() { + let schema = serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}}); + let args = serde_json::json!({"name": "Alice"}); + assert!(validate_tool_args(&args, &schema).is_ok()); + } + + #[test] + fn validate_tool_args_non_object_when_object_expected() { + let schema = serde_json::json!({"type": "object", "properties": {}}); + let args = serde_json::json!("not an object"); + let result = validate_tool_args(&args, &schema); + assert!(result.is_err()); + assert!(result.unwrap_err().contains("Expected object")); + } + + #[test] + fn validate_tool_args_missing_required_properties() { + let schema = serde_json::json!({ + "type": "object", + "properties": {"name": {"type": "string"}, "age": {"type": "number"}}, + "required": ["name", "age"] + }); + let args = serde_json::json!({"name": "Alice"}); + let result = validate_tool_args(&args, &schema); + assert!(result.is_err()); + assert!(result.unwrap_err().contains("age")); + } + + #[test] + fn validate_tool_args_no_schema_type_passes() { + let schema = serde_json::json!({}); + let args = serde_json::json!("anything"); + assert!(validate_tool_args(&args, &schema).is_ok()); + } + + #[tokio::test] + async fn execute_with_repair_valid_args_no_repair_needed() { + let tools = vec![Tool::active( + "greet", + "Greet someone", + serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}}), + |args, _ctx| async move { + let name = args["name"].as_str().unwrap_or("world"); + Ok(serde_json::json!(format!("Hello, {}!", name))) + }, + )]; + let calls = vec![ToolCall::new("call_1", "greet", serde_json::json!({"name": "Alice"}))]; + let tool_refs: Vec<&Tool> = tools.iter().collect(); + + let results = execute_all_tools_with_repair(&tool_refs, &calls, &[], None, None).await; + assert_eq!(results.len(), 1); + assert!(!results[0].is_error); + assert_eq!(results[0].content, serde_json::json!("Hello, Alice!")); + } + + #[tokio::test] + async fn execute_with_repair_invalid_args_no_repair_fn() { + let tools = vec![Tool::active( + "greet", + "Greet someone", + serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}), + |args, _ctx| async move { + let name = args["name"].as_str().unwrap_or("world"); + Ok(serde_json::json!(format!("Hello, {}!", name))) + }, + )]; + let calls = vec![ToolCall::new("call_1", "greet", serde_json::json!({}))]; + let tool_refs: Vec<&Tool> = tools.iter().collect(); + + let results = execute_all_tools_with_repair(&tool_refs, &calls, &[], None, None).await; + assert_eq!(results.len(), 1); + assert!(results[0].is_error); + assert!(results[0].content.as_str().unwrap().contains("validation failed")); + } + + #[tokio::test] + async fn execute_with_repair_invalid_args_repair_succeeds() { + let tools = vec![Tool::active( + "greet", + "Greet someone", + serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}), + |args, _ctx| async move { + let name = args["name"].as_str().unwrap_or("world"); + Ok(serde_json::json!(format!("Hello, {}!", name))) + }, + )]; + let calls = vec![ToolCall::new("call_1", "greet", serde_json::json!({}))]; + let tool_refs: Vec<&Tool> = tools.iter().collect(); + + let repair: RepairToolCallFn = Arc::new(|_call, _error| { + Box::pin(async { Ok(serde_json::json!({"name": "Repaired"})) }) + }); + let results = execute_all_tools_with_repair(&tool_refs, &calls, &[], None, Some(&repair)).await; + assert_eq!(results.len(), 1); + assert!(!results[0].is_error); + assert_eq!(results[0].content, serde_json::json!("Hello, Repaired!")); + } + + #[tokio::test] + async fn execute_with_repair_invalid_args_repair_fails() { + let tools = vec![Tool::active( + "greet", + "Greet someone", + serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}), + |args, _ctx| async move { + let name = args["name"].as_str().unwrap_or("world"); + Ok(serde_json::json!(format!("Hello, {}!", name))) + }, + )]; + let calls = vec![ToolCall::new("call_1", "greet", serde_json::json!({}))]; + let tool_refs: Vec<&Tool> = tools.iter().collect(); + + let repair: RepairToolCallFn = Arc::new(|_call, _error| { + Box::pin(async { Err("cannot repair".to_string()) }) + }); + let results = execute_all_tools_with_repair(&tool_refs, &calls, &[], None, Some(&repair)).await; + assert_eq!(results.len(), 1); + assert!(results[0].is_error); + assert!(results[0].content.as_str().unwrap().contains("repair failed")); + } } diff --git a/crates/unified-llm/src/types.rs b/crates/unified-llm/src/types.rs index 26c2ecc9b..f738f4e5e 100644 --- a/crates/unified-llm/src/types.rs +++ b/crates/unified-llm/src/types.rs @@ -85,8 +85,7 @@ pub struct ToolResult { // --- 3.3 ContentPart --- -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -#[serde(tag = "kind", content = "data", rename_all = "snake_case")] +#[derive(Debug, Clone, PartialEq, Eq)] pub enum ContentPart { Text(String), Image(ImageData), @@ -96,6 +95,97 @@ pub enum ContentPart { ToolResult(ToolResult), Thinking(ThinkingData), RedactedThinking(ThinkingData), + Other { + kind: String, + data: serde_json::Value, + }, +} + +impl Serialize for ContentPart { + fn serialize(&self, serializer: S) -> Result { + use serde::ser::SerializeMap; + let mut map = serializer.serialize_map(Some(2))?; + match self { + Self::Text(v) => { + map.serialize_entry("kind", "text")?; + map.serialize_entry("data", v)?; + } + Self::Image(v) => { + map.serialize_entry("kind", "image")?; + map.serialize_entry("data", v)?; + } + Self::Audio(v) => { + map.serialize_entry("kind", "audio")?; + map.serialize_entry("data", v)?; + } + Self::Document(v) => { + map.serialize_entry("kind", "document")?; + map.serialize_entry("data", v)?; + } + Self::ToolCall(v) => { + map.serialize_entry("kind", "tool_call")?; + map.serialize_entry("data", v)?; + } + Self::ToolResult(v) => { + map.serialize_entry("kind", "tool_result")?; + map.serialize_entry("data", v)?; + } + Self::Thinking(v) => { + map.serialize_entry("kind", "thinking")?; + map.serialize_entry("data", v)?; + } + Self::RedactedThinking(v) => { + map.serialize_entry("kind", "redacted_thinking")?; + map.serialize_entry("data", v)?; + } + Self::Other { kind, data } => { + map.serialize_entry("kind", kind)?; + map.serialize_entry("data", data)?; + } + } + map.end() + } +} + +impl<'de> Deserialize<'de> for ContentPart { + fn deserialize>(deserializer: D) -> Result { + let value = serde_json::Value::deserialize(deserializer)?; + let kind = value + .get("kind") + .and_then(serde_json::Value::as_str) + .ok_or_else(|| serde::de::Error::missing_field("kind"))?; + let data = value.get("data").cloned().unwrap_or(serde_json::Value::Null); + match kind { + "text" => serde_json::from_value(data) + .map(Self::Text) + .map_err(serde::de::Error::custom), + "image" => serde_json::from_value(data) + .map(Self::Image) + .map_err(serde::de::Error::custom), + "audio" => serde_json::from_value(data) + .map(Self::Audio) + .map_err(serde::de::Error::custom), + "document" => serde_json::from_value(data) + .map(Self::Document) + .map_err(serde::de::Error::custom), + "tool_call" => serde_json::from_value(data) + .map(Self::ToolCall) + .map_err(serde::de::Error::custom), + "tool_result" => serde_json::from_value(data) + .map(Self::ToolResult) + .map_err(serde::de::Error::custom), + "thinking" => serde_json::from_value(data) + .map(Self::Thinking) + .map_err(serde::de::Error::custom), + "redacted_thinking" => serde_json::from_value(data) + .map(Self::RedactedThinking) + .map_err(serde::de::Error::custom), + other => Ok(Self::Other { + kind: other.to_string(), + data, + }), + } + } } impl ContentPart { @@ -441,7 +531,7 @@ pub enum StreamEvent { response: Box, }, Error { - error: String, + error: SdkError, raw: Option, }, ProviderEvent { @@ -483,11 +573,8 @@ impl StreamEvent { } } - pub fn error(message: impl Into) -> Self { - Self::Error { - error: message.into(), - raw: None, - } + pub fn error(error: SdkError) -> Self { + Self::Error { error, raw: None } } } @@ -955,10 +1042,12 @@ mod tests { #[test] fn stream_event_error() { - let event = StreamEvent::error("something went wrong"); + let event = StreamEvent::error(SdkError::Stream { + message: "something went wrong".into(), + }); match &event { StreamEvent::Error { error, .. } => { - assert_eq!(error, "something went wrong"); + assert_eq!(error.to_string(), "Stream error: something went wrong"); } other => panic!("Expected Error, got {other:?}"), }