From ed07d43335b0f3763bba7f69e542d818b2df9395 Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp Date: Fri, 20 Feb 2026 11:05:10 -0400 Subject: [PATCH] Implement spec gaps: rate limit headers, error classification, total timeout, metadata, stream_object - Parse x-ratelimit-* headers into RateLimitInfo for Anthropic, OpenAI, and OpenAI-compatible providers (previously hardcoded to None) - Add "not found"/"does not exist" and "unauthorized"/"invalid key" error message classification patterns for ambiguous HTTP status codes - Apply TimeoutConfig.total to wrap the entire multi-step generate() loop (previously only per_step was used) - Add metadata field to GenerateParams with builder method, pass through to Request instead of hardcoding None - Implement stream_object() for streaming structured output with incremental JSON parsing via new ObjectStreamEvent type (Partial/Delta/Complete variants) - Add OpenAI-compatible Chat Completions adapter for third-party endpoints Co-Authored-By: Claude Opus 4.6 --- Cargo.lock | 19 + Cargo.toml | 4 +- crates/unified-llm/Cargo.toml | 3 + crates/unified-llm/src/client.rs | 32 +- crates/unified-llm/src/error.rs | 64 + crates/unified-llm/src/generate.rs | 612 +++++++- crates/unified-llm/src/providers/anthropic.rs | 1283 ++++++++++++++++- crates/unified-llm/src/providers/common.rs | 292 +++- crates/unified-llm/src/providers/gemini.rs | 676 ++++++++- crates/unified-llm/src/providers/mod.rs | 2 + crates/unified-llm/src/providers/openai.rs | 888 ++++++++++-- .../src/providers/openai_compatible.rs | 1114 ++++++++++++++ crates/unified-llm/src/retry.rs | 57 +- crates/unified-llm/src/tools.rs | 43 +- crates/unified-llm/src/types.rs | 61 +- docs/specs/unified-llm-spec.md | 24 +- 16 files changed, 4880 insertions(+), 294 deletions(-) create mode 100644 crates/unified-llm/src/providers/openai_compatible.rs diff --git a/Cargo.lock b/Cargo.lock index d264da939..d5d3ac56c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -880,6 +880,7 @@ dependencies = [ "bytes", "encoding_rs", "futures-core", + "futures-util", "h2", "http", "http-body", @@ -901,12 +902,14 @@ dependencies = [ "sync_wrapper", "tokio", "tokio-native-tls", + "tokio-util", "tower", "tower-http", "tower-service", "url", "wasm-bindgen", "wasm-bindgen-futures", + "wasm-streams", "web-sys", ] @@ -1386,8 +1389,11 @@ version = "0.1.0" dependencies = [ "anyhow", "async-trait", + "base64", + "bytes", "dotenvy", "futures", + "http", "rand", "reqwest", "serde", @@ -1553,6 +1559,19 @@ dependencies = [ "wasmparser", ] +[[package]] +name = "wasm-streams" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "15053d8d85c7eccdbefef60f06769760a563c7f0a9d6902a13d35c7800b0ad65" +dependencies = [ + "futures-util", + "js-sys", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + [[package]] name = "wasmparser" version = "0.244.0" diff --git a/Cargo.toml b/Cargo.toml index f083f6b3b..ff9923ea0 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -20,10 +20,12 @@ thiserror = "2" serde = { version = "1", features = ["derive"] } serde_json = "1" tokio = { version = "1", features = ["full"] } -reqwest = { version = "0.12", features = ["json"] } +reqwest = { version = "0.12", features = ["json", "stream"] } uuid = { version = "1", features = ["v4"] } rand = "0.8" dotenvy = "0.15" futures = "0.3" tokio-stream = "0.1" async-trait = "0.1" +base64 = "0.22" +bytes = "1" diff --git a/crates/unified-llm/Cargo.toml b/crates/unified-llm/Cargo.toml index 1c2553c65..9a208f539 100644 --- a/crates/unified-llm/Cargo.toml +++ b/crates/unified-llm/Cargo.toml @@ -21,9 +21,12 @@ futures.workspace = true tokio-stream.workspace = true async-trait.workspace = true reqwest.workspace = true +base64.workspace = true +bytes.workspace = true [dev-dependencies] dotenvy.workspace = true +http = "1" tokio = { workspace = true, features = ["test-util", "macros"] } [lints] diff --git a/crates/unified-llm/src/client.rs b/crates/unified-llm/src/client.rs index 763af55c0..74c7afe98 100644 --- a/crates/unified-llm/src/client.rs +++ b/crates/unified-llm/src/client.rs @@ -1,6 +1,7 @@ use crate::error::SdkError; use crate::middleware::{Middleware, NextFn, NextStreamFn}; use crate::provider::{ProviderAdapter, StreamEventStream}; +use crate::providers; use crate::types::{Request, Response}; use std::collections::HashMap; use std::sync::Arc; @@ -32,13 +33,38 @@ impl Client { /// The first registered provider becomes the default. #[must_use] pub fn from_env() -> Self { - // In a real implementation, this would check for OPENAI_API_KEY, ANTHROPIC_API_KEY, etc. - // and register the appropriate adapters. For now, return an empty client. - Self { + let mut client = Self { providers: HashMap::new(), default_provider: None, middleware: Vec::new(), + }; + + // Register providers whose API keys are present in the environment. + // Order determines which becomes the default provider. + if let Ok(key) = std::env::var("ANTHROPIC_API_KEY") { + let mut adapter = providers::AnthropicAdapter::new(key); + if let Ok(base_url) = std::env::var("ANTHROPIC_BASE_URL") { + adapter = adapter.with_base_url(base_url); + } + client.register_provider(Arc::new(adapter)); } + if let Ok(key) = std::env::var("OPENAI_API_KEY") { + let mut adapter = providers::OpenAiAdapter::new(key); + if let Ok(base_url) = std::env::var("OPENAI_BASE_URL") { + adapter = adapter.with_base_url(base_url); + } + client.register_provider(Arc::new(adapter)); + } + if let Ok(key) = std::env::var("GEMINI_API_KEY").or_else(|_| std::env::var("GOOGLE_API_KEY")) + { + let mut adapter = providers::GeminiAdapter::new(key); + 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 a provider adapter. diff --git a/crates/unified-llm/src/error.rs b/crates/unified-llm/src/error.rs index 05030ae1b..b68781fc7 100644 --- a/crates/unified-llm/src/error.rs +++ b/crates/unified-llm/src/error.rs @@ -140,6 +140,18 @@ pub fn error_from_status_code( // First check message-based classification for ambiguous cases let lower_msg = detail.message.to_lowercase(); + if lower_msg.contains("not found") || lower_msg.contains("does not exist") { + return SdkError::Provider { + kind: ProviderErrorKind::NotFound, + detail: Box::new(detail), + }; + } + if lower_msg.contains("unauthorized") || lower_msg.contains("invalid key") { + return SdkError::Provider { + kind: ProviderErrorKind::Authentication, + detail: Box::new(detail), + }; + } if lower_msg.contains("context length") || lower_msg.contains("too many tokens") { return SdkError::Provider { kind: ProviderErrorKind::ContextLength, @@ -372,6 +384,58 @@ mod tests { assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::ContentFilter, .. })); } + #[test] + fn error_message_classification_not_found() { + let err = error_from_status_code( + 400, + "The model gpt-5 was not found".into(), + "openai".into(), + None, + None, + None, + ); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::NotFound, .. })); + } + + #[test] + fn error_message_classification_does_not_exist() { + let err = error_from_status_code( + 400, + "The resource does not exist".into(), + "openai".into(), + None, + None, + None, + ); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::NotFound, .. })); + } + + #[test] + fn error_message_classification_unauthorized() { + let err = error_from_status_code( + 400, + "Request unauthorized for this resource".into(), + "openai".into(), + None, + None, + None, + ); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Authentication, .. })); + } + + #[test] + fn error_message_classification_invalid_key() { + let err = error_from_status_code( + 400, + "Provided invalid key for authentication".into(), + "openai".into(), + None, + None, + None, + ); + assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Authentication, .. })); + } + #[test] fn grpc_status_mapping() { let err = error_from_grpc_status("NOT_FOUND", "model not found".into(), "gemini".into(), None, None, None); diff --git a/crates/unified-llm/src/generate.rs b/crates/unified-llm/src/generate.rs index e755ce8e6..65c29fcbd 100644 --- a/crates/unified-llm/src/generate.rs +++ b/crates/unified-llm/src/generate.rs @@ -4,9 +4,12 @@ use crate::provider::StreamEventStream; use crate::retry::retry; use crate::tools::{execute_all_tools, Tool}; use crate::types::{ - FinishReason, GenerateResult, Message, Request, Response, ResponseFormat, ResponseFormatType, - RetryPolicy, StepResult, StreamEvent, ToolCall, ToolChoice, ToolDefinition, Usage, + FinishReason, GenerateResult, Message, ObjectStreamEvent, Request, Response, ResponseFormat, + ResponseFormatType, RetryPolicy, StepResult, StreamEvent, TimeoutConfig, ToolCall, ToolChoice, + ToolDefinition, Usage, }; +use futures::StreamExt; +use std::pin::Pin; use std::sync::Arc; use tokio::sync::OnceCell; @@ -61,7 +64,7 @@ fn build_request( max_tokens: params.max_tokens, stop_sequences: params.stop_sequences.clone(), reasoning_effort: params.reasoning_effort.clone(), - metadata: None, + metadata: params.metadata.clone(), provider_options: params.provider_options.clone(), } } @@ -108,69 +111,101 @@ pub async fn generate(params: GenerateParams) -> Result = Vec::new(); - let mut total_usage = Usage::default(); - let mut round = 0u32; - loop { - let request = build_request(¶ms, &messages, tool_definitions.as_deref()); + let generate_future = async { + let mut steps: Vec = Vec::new(); + let mut total_usage = Usage::default(); - let client_ref = client.clone(); - let response = retry(&retry_policy, || { - let c = client_ref.clone(); - let r = request.clone(); - async move { c.complete(&r).await } - }) - .await?; + let mut round = 0u32; + loop { + let request = build_request(¶ms, &messages, tool_definitions.as_deref()); - let tool_calls = response.tool_calls(); - let mut tool_results = Vec::new(); + let client_ref = client.clone(); + let response = + if let Some(per_step) = params.timeout.as_ref().and_then(|t| t.per_step) { + let duration = std::time::Duration::from_secs_f64(per_step); + tokio::time::timeout(duration, retry(&retry_policy, || { + let c = client_ref.clone(); + let r = request.clone(); + async move { c.complete(&r).await } + })) + .await + .map_err(|_| SdkError::RequestTimeout { + message: format!("Per-step timeout of {per_step}s exceeded"), + })? + } else { + retry(&retry_policy, || { + let c = client_ref.clone(); + let r = request.clone(); + async move { c.complete(&r).await } + }) + .await + }?; - if !tool_calls.is_empty() - && response.finish_reason == FinishReason::ToolCalls - && params.tools.is_some() - { - let tools = params.tools.as_ref().expect("checked above"); - if tools.iter().any(|t| t.is_active()) { - let tool_refs: Vec<&Tool> = - tools.iter().map(std::convert::AsRef::as_ref).collect(); - tool_results = execute_all_tools(&tool_refs, &tool_calls).await; + let tool_calls = response.tool_calls(); + let mut tool_results = Vec::new(); + + if !tool_calls.is_empty() + && response.finish_reason == FinishReason::ToolCalls + && params.tools.is_some() + { + let tools = params.tools.as_ref().expect("checked above"); + if tools.iter().any(|t| t.is_active()) { + let tool_refs: Vec<&Tool> = + tools.iter().map(std::convert::AsRef::as_ref).collect(); + tool_results = execute_all_tools(&tool_refs, &tool_calls).await; + } } - } - total_usage = total_usage + response.usage.clone(); + total_usage = total_usage + response.usage.clone(); - let should_continue = !tool_calls.is_empty() - && response.finish_reason == FinishReason::ToolCalls - && round < max_tool_rounds - && !tool_results.is_empty(); + steps.push(StepResult { + response, + tool_results, + }); - if should_continue { - messages.push(response.message.clone()); - for result in &tool_results { + let last = steps.last().expect("just pushed"); + let should_continue = !tool_calls.is_empty() + && last.response.finish_reason == FinishReason::ToolCalls + && round < max_tool_rounds + && !last.tool_results.is_empty() + && !params.stop_when.as_ref().is_some_and(|f| f(&steps)); + + if !should_continue { + break; + } + + let last = steps.last().expect("just pushed"); + messages.push(last.response.message.clone()); + for result in &last.tool_results { messages.push(Message::tool_result( &result.tool_call_id, result.content.to_string(), result.is_error, )); } + + round += 1; } - steps.push(StepResult { - response, - tool_results, - }); + Ok(build_generate_result(steps, total_usage)) + }; - if !should_continue { - break; - } - - round += 1; + if let Some(total) = params.timeout.as_ref().and_then(|t| t.total) { + let duration = std::time::Duration::from_secs_f64(total); + tokio::time::timeout(duration, generate_future) + .await + .map_err(|_| SdkError::RequestTimeout { + message: format!("Total timeout of {total}s exceeded"), + })? + } else { + generate_future.await } - - Ok(build_generate_result(steps, total_usage)) } +/// Callback type for custom stop conditions in the tool loop. +pub type StopCondition = Arc bool + Send + Sync>; + /// Parameters for `generate()` (Section 4.3). #[derive(Clone)] pub struct GenerateParams { @@ -189,8 +224,12 @@ pub struct GenerateParams { pub reasoning_effort: Option, pub provider: Option, pub provider_options: Option, + pub metadata: Option>, pub max_retries: u32, + pub timeout: Option, pub client: Option>, + /// Custom stop condition checked after each tool round (Section 4.3). + pub stop_when: Option, } impl GenerateParams { @@ -211,8 +250,11 @@ impl GenerateParams { reasoning_effort: None, provider: None, provider_options: None, + metadata: None, max_retries: 2, + timeout: None, client: None, + stop_when: None, } } @@ -257,6 +299,85 @@ impl GenerateParams { self.provider = Some(provider.into()); self } + + #[must_use] + pub fn tool_choice(mut self, tool_choice: ToolChoice) -> Self { + self.tool_choice = Some(tool_choice); + self + } + + #[must_use] + pub fn response_format(mut self, response_format: ResponseFormat) -> Self { + self.response_format = Some(response_format); + self + } + + #[must_use] + pub const fn temperature(mut self, temperature: f64) -> Self { + self.temperature = Some(temperature); + self + } + + #[must_use] + pub const fn top_p(mut self, top_p: f64) -> Self { + self.top_p = Some(top_p); + self + } + + #[must_use] + pub const fn max_tokens(mut self, max_tokens: i64) -> Self { + self.max_tokens = Some(max_tokens); + self + } + + #[must_use] + pub fn stop_sequences(mut self, stop_sequences: Vec) -> Self { + self.stop_sequences = Some(stop_sequences); + self + } + + #[must_use] + pub fn reasoning_effort(mut self, reasoning_effort: impl Into) -> Self { + self.reasoning_effort = Some(reasoning_effort.into()); + self + } + + #[must_use] + pub fn provider_options(mut self, provider_options: serde_json::Value) -> Self { + self.provider_options = Some(provider_options); + self + } + + #[must_use] + pub fn metadata(mut self, metadata: std::collections::HashMap) -> Self { + self.metadata = Some(metadata); + self + } + + #[must_use] + pub const fn max_retries(mut self, max_retries: u32) -> Self { + self.max_retries = max_retries; + self + } + + #[must_use] + pub const fn timeout(mut self, timeout: TimeoutConfig) -> Self { + self.timeout = Some(timeout); + self + } + + /// Set a custom stop condition for the tool loop (Section 4.3). + /// + /// The callback receives the accumulated steps so far and returns `true` + /// to stop the tool loop early. + #[must_use] + pub fn stop_when( + mut self, + f: impl Fn(&[StepResult]) -> bool + Send + Sync + 'static, + ) -> Self { + self.stop_when = Some(Arc::new(f)); + self + } } /// `StreamAccumulator` collects stream events into a complete Response (Section 4.4). @@ -339,6 +460,19 @@ impl Default for StreamAccumulator { /// /// Returns `SdkError::Configuration` if both `prompt` and `messages` are set, /// or any provider error encountered during streaming setup. +pub async fn stream(params: GenerateParams) -> Result { + stream_generate(params).await +} + +/// High-level streaming generation (Section 4.4). +/// Returns a `StreamEventStream` that the caller can iterate over. +/// +/// Alias: prefer [`stream()`] for consistency with the spec. +/// +/// # Errors +/// +/// 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 messages = build_initial_messages(¶ms)?; @@ -384,6 +518,94 @@ pub async fn generate_object( } } +/// Stream type for `stream_object()`. +pub type ObjectStream = + Pin> + Send>>; + +/// Streaming structured output with incremental JSON parsing (Section 4.6). +/// +/// Combines streaming with structured output: sets `response_format` to `json_schema`, +/// streams the response, and attempts to parse the accumulated text as JSON on each +/// text delta. Yields `ObjectStreamEvent::Partial` when a new valid partial parse is +/// obtained, `ObjectStreamEvent::Delta` for every raw stream event, and +/// `ObjectStreamEvent::Complete` when the stream finishes with the final parsed object. +/// +/// # Errors +/// +/// Returns `SdkError::Configuration` if both `prompt` and `messages` are set, +/// `SdkError::NoObjectGenerated` if the final accumulated text is not valid JSON, +/// or any provider error encountered during streaming. +pub async fn stream_object( + params: GenerateParams, + schema: serde_json::Value, +) -> Result { + let params = GenerateParams { + response_format: Some(ResponseFormat { + kind: ResponseFormatType::JsonSchema, + json_schema: Some(schema), + strict: true, + }), + ..params + }; + + let inner_stream = stream(params).await?; + + let mapped = inner_stream.scan( + (String::new(), Option::::None), + |(accumulated_text, last_parsed), event| { + let mut events: Vec> = Vec::new(); + + match &event { + Ok(stream_event) => { + // Accumulate text from TextDelta events + if let StreamEvent::TextDelta { delta, .. } = stream_event { + accumulated_text.push_str(delta); + + // Try incremental JSON parse + if let Ok(parsed) = serde_json::from_str::(accumulated_text) { + if last_parsed.as_ref() != Some(&parsed) { + *last_parsed = Some(parsed.clone()); + events.push(Ok(ObjectStreamEvent::Partial { object: parsed })); + } + } + } + + // On Finish, yield the Complete event with final parsed object + if let StreamEvent::Finish { response, .. } = stream_event { + match serde_json::from_str::(accumulated_text) { + Ok(final_object) => { + events.push(Ok(ObjectStreamEvent::Complete { + object: final_object, + response: response.clone(), + })); + } + Err(e) => { + events.push(Err(SdkError::NoObjectGenerated { + message: format!("Failed to parse final response as JSON: {e}"), + })); + } + } + } else { + // Yield the raw delta event + events.push(Ok(ObjectStreamEvent::Delta { + event: stream_event.clone(), + })); + } + } + Err(e) => { + events.push(Err(SdkError::Stream { + message: format!("{e}"), + })); + } + } + + futures::future::ready(Some(futures::stream::iter(events))) + }, + ); + + Ok(Box::pin(mapped.flatten())) +} + #[cfg(test)] mod tests { use super::*; @@ -774,4 +996,302 @@ mod tests { SdkError::NoObjectGenerated { .. } )); } + + #[tokio::test] + async fn generate_stop_when_halts_tool_loop() { + let call_count = Arc::new(AtomicU32::new(0)); + let provider: Arc = Arc::new(ToolCallMockProvider { + call_count: call_count.clone(), + }); + + let mut providers: HashMap> = HashMap::new(); + providers.insert("mock".to_string(), provider); + let client = Arc::new(Client::new( + providers, + Some("mock".to_string()), + vec![], + )); + + let result = generate( + GenerateParams::new("mock-model") + .prompt("What's the weather in SF?") + .tools(vec![Tool::active( + "get_weather", + "Get weather", + serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), + |args| async move { + let city = args["city"].as_str().unwrap_or("unknown"); + Ok(serde_json::json!(format!("72F in {}", city))) + }, + )]) + .max_tool_rounds(5) + .stop_when(|_steps| true) // Stop immediately after first round + .client(client), + ) + .await + .unwrap(); + + // stop_when returned true, so the tool loop should stop after 1 step + assert_eq!(result.steps.len(), 1); + assert_eq!(call_count.load(Ordering::SeqCst), 1); + } + + #[test] + fn generate_params_builder_methods() { + let params = GenerateParams::new("test-model") + .prompt("hello") + .system("you are helpful") + .temperature(0.7) + .top_p(0.9) + .max_tokens(100) + .stop_sequences(vec!["STOP".to_string()]) + .reasoning_effort("high") + .provider("anthropic") + .provider_options(serde_json::json!({"key": "value"})) + .max_retries(5) + .tool_choice(ToolChoice::Required) + .response_format(ResponseFormat { + kind: ResponseFormatType::JsonObject, + json_schema: None, + strict: false, + }) + .max_tool_rounds(3); + + assert_eq!(params.model, "test-model"); + assert_eq!(params.prompt.as_deref(), Some("hello")); + assert_eq!(params.system.as_deref(), Some("you are helpful")); + assert_eq!(params.temperature, Some(0.7)); + assert_eq!(params.top_p, Some(0.9)); + assert_eq!(params.max_tokens, Some(100)); + assert_eq!( + params.stop_sequences, + Some(vec!["STOP".to_string()]) + ); + assert_eq!(params.reasoning_effort.as_deref(), Some("high")); + assert_eq!(params.provider.as_deref(), Some("anthropic")); + assert!(params.provider_options.is_some()); + assert_eq!(params.max_retries, 5); + assert_eq!(params.tool_choice, Some(ToolChoice::Required)); + assert!(params.response_format.is_some()); + assert_eq!(params.max_tool_rounds, 3); + } + + #[test] + fn generate_params_timeout_builder() { + let params = GenerateParams::new("test-model") + .timeout(TimeoutConfig { + total: Some(30.0), + per_step: Some(10.0), + }); + assert!(params.timeout.is_some()); + let t = params.timeout.unwrap(); + assert_eq!(t.total, Some(30.0)); + assert_eq!(t.per_step, Some(10.0)); + } + + /// Mock provider that streams JSON tokens incrementally. + struct StreamingJsonMockProvider { + deltas: Vec, + full_text: String, + } + + impl StreamingJsonMockProvider { + fn new(deltas: Vec<&str>) -> Self { + let full_text: String = deltas.iter().copied().collect(); + Self { + deltas: deltas.into_iter().map(String::from).collect(), + full_text, + } + } + } + + #[async_trait::async_trait] + impl ProviderAdapter for StreamingJsonMockProvider { + fn name(&self) -> &str { + "mock" + } + + async fn complete(&self, _request: &Request) -> Result { + Ok(Response { + id: "resp_1".into(), + model: "mock-model".into(), + provider: "mock".into(), + message: Message::assistant(&self.full_text), + finish_reason: FinishReason::Stop, + usage: Usage::default(), + raw: None, + warnings: vec![], + rate_limit: None, + }) + } + + async fn stream( + &self, + _request: &Request, + ) -> Result { + let mut events: Vec> = self + .deltas + .iter() + .map(|d| Ok(StreamEvent::text_delta(d.as_str(), Some("t1".into())))) + .collect(); + + events.push(Ok(StreamEvent::finish( + FinishReason::Stop, + Usage { + input_tokens: 10, + output_tokens: 20, + total_tokens: 30, + ..Default::default() + }, + Response { + id: "resp_1".into(), + model: "mock-model".into(), + provider: "mock".into(), + message: Message::assistant(&self.full_text), + finish_reason: FinishReason::Stop, + usage: Usage { + input_tokens: 10, + output_tokens: 20, + total_tokens: 30, + ..Default::default() + }, + raw: None, + warnings: vec![], + rate_limit: None, + }, + ))); + + Ok(Box::pin(stream::iter(events))) + } + } + + fn streaming_json_mock_client(deltas: Vec<&str>) -> Arc { + let mut providers: HashMap> = HashMap::new(); + providers.insert( + "mock".to_string(), + Arc::new(StreamingJsonMockProvider::new(deltas)), + ); + Arc::new(Client::new( + providers, + Some("mock".to_string()), + vec![], + )) + } + + #[tokio::test] + async fn stream_object_yields_complete_event() { + let client = streaming_json_mock_client(vec![r#"{"name": "Alice", "age": 30}"#]); + + let schema = serde_json::json!({ + "type": "object", + "properties": { + "name": {"type": "string"}, + "age": {"type": "integer"} + }, + "required": ["name", "age"] + }); + + let obj_stream = stream_object( + GenerateParams::new("mock-model") + .prompt("Extract info") + .client(client), + schema, + ) + .await + .unwrap(); + + let events: Vec = obj_stream + .filter_map(|r| futures::future::ready(r.ok())) + .collect() + .await; + + let complete = events + .iter() + .find(|e| matches!(e, ObjectStreamEvent::Complete { .. })); + assert!(complete.is_some(), "Expected a Complete event"); + + if let ObjectStreamEvent::Complete { object, .. } = complete.unwrap() { + assert_eq!(object["name"], "Alice"); + assert_eq!(object["age"], 30); + } + } + + #[tokio::test] + async fn stream_object_yields_partial_events_incrementally() { + let client = streaming_json_mock_client(vec![ + r#"{"name""#, + r#": "Bob""#, + r#", "age": 25}"#, + ]); + + let schema = serde_json::json!({ + "type": "object", + "properties": { + "name": {"type": "string"}, + "age": {"type": "integer"} + } + }); + + let obj_stream = stream_object( + GenerateParams::new("mock-model") + .prompt("Extract info") + .client(client), + schema, + ) + .await + .unwrap(); + + let events: Vec = obj_stream + .filter_map(|r| futures::future::ready(r.ok())) + .collect() + .await; + + let partial_count = events + .iter() + .filter(|e| matches!(e, ObjectStreamEvent::Partial { .. })) + .count(); + + assert!( + partial_count >= 1, + "Expected at least one Partial event, got {partial_count}" + ); + + let delta_count = events + .iter() + .filter(|e| matches!(e, ObjectStreamEvent::Delta { .. })) + .count(); + + assert_eq!(delta_count, 3); + + let last_complete = events + .iter() + .rev() + .find(|e| matches!(e, ObjectStreamEvent::Complete { .. })); + assert!(last_complete.is_some(), "Expected a Complete event"); + if let ObjectStreamEvent::Complete { object, .. } = last_complete.unwrap() { + assert_eq!(object["name"], "Bob"); + assert_eq!(object["age"], 25); + } + } + + #[tokio::test] + async fn stream_object_errors_on_invalid_final_json() { + let client = streaming_json_mock_client(vec![r#"{"name": "Alice"#]); + + let schema = serde_json::json!({"type": "object"}); + + let obj_stream = stream_object( + GenerateParams::new("mock-model") + .prompt("Extract info") + .client(client), + schema, + ) + .await + .unwrap(); + + let results: Vec> = obj_stream.collect().await; + + let has_error = results.iter().any(|r| r.is_err()); + assert!(has_error, "Expected an error for invalid final JSON"); + } } diff --git a/crates/unified-llm/src/providers/anthropic.rs b/crates/unified-llm/src/providers/anthropic.rs index 242c6d0aa..41d3a8b66 100644 --- a/crates/unified-llm/src/providers/anthropic.rs +++ b/crates/unified-llm/src/providers/anthropic.rs @@ -1,26 +1,52 @@ +use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine}; + use crate::error::SdkError; use crate::provider::{ProviderAdapter, StreamEventStream}; -use crate::providers::common::{extract_system_prompt, send_and_read_body, ApiMessage}; +use crate::providers::common::{ + extract_system_prompt, parse_error_body, parse_rate_limit_headers, send_and_read_response, +}; use crate::types::{ - ContentPart, FinishReason, Message, Request, Response, Role, ToolCall, Usage, + ContentPart, FinishReason, Message, Request, Response, Role, StreamEvent, ThinkingData, + ToolCall, ToolChoice, ToolDefinition, Usage, }; /// Provider adapter for the Anthropic Messages API. pub struct Adapter { api_key: String, + base_url: String, client: reqwest::Client, + request_timeout: std::time::Duration, } impl Adapter { #[must_use] pub fn new(api_key: impl Into) -> Self { + let timeout = crate::types::AdapterTimeout::default(); + let client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs_f64(timeout.connect)) + .build() + .unwrap_or_default(); Self { api_key: api_key.into(), - client: reqwest::Client::new(), + base_url: DEFAULT_BASE_URL.to_string(), + client, + request_timeout: std::time::Duration::from_secs_f64(timeout.request), } } + + #[must_use] + pub fn with_base_url(mut self, base_url: impl Into) -> Self { + self.base_url = base_url.into(); + self + } + + fn messages_url(&self) -> String { + format!("{}/messages", self.base_url) + } } +const DEFAULT_BASE_URL: &str = "https://api.anthropic.com/v1"; + // --- Request types --- #[derive(serde::Serialize)] @@ -28,14 +54,58 @@ struct ApiRequest { model: String, messages: Vec, max_tokens: i64, + /// System prompt: either a plain string or an array of content blocks + /// (with optional `cache_control` annotations for prompt caching). #[serde(skip_serializing_if = "Option::is_none")] - system: Option, + system: Option, #[serde(skip_serializing_if = "Option::is_none")] temperature: Option, #[serde(skip_serializing_if = "Option::is_none")] top_p: Option, #[serde(skip_serializing_if = "Option::is_none")] stop_sequences: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + tool_choice: Option, + /// Extended thinking configuration (e.g. `{"type": "enabled", "budget_tokens": 10000}`). + /// Passed through from `provider_options.anthropic.thinking`. + #[serde(skip_serializing_if = "Option::is_none")] + thinking: Option, + #[serde(skip_serializing_if = "std::ops::Not::not")] + stream: bool, +} + +/// Anthropic messages use structured content blocks, not plain strings. +#[derive(serde::Serialize)] +struct ApiMessage { + role: String, + content: Vec, +} + +/// Anthropic tool definition format. +#[derive(serde::Serialize)] +struct ApiToolDef { + name: String, + description: String, + input_schema: serde_json::Value, + #[serde(skip_serializing_if = "Option::is_none")] + cache_control: Option, +} + +/// Anthropic `cache_control` annotation. +#[derive(serde::Serialize, Clone)] +struct CacheControl { + #[serde(rename = "type")] + kind: String, +} + +impl CacheControl { + fn ephemeral() -> Self { + Self { + kind: "ephemeral".to_string(), + } + } } // --- Response types --- @@ -50,14 +120,19 @@ struct ApiResponse { } #[derive(serde::Deserialize)] +#[allow(clippy::struct_field_names)] struct ApiUsage { input_tokens: i64, output_tokens: i64, + #[serde(default)] + cache_read_input_tokens: Option, + #[serde(default)] + cache_creation_input_tokens: Option, } fn map_finish_reason(stop_reason: Option<&str>) -> FinishReason { match stop_reason { - Some("end_turn") | None => FinishReason::Stop, + Some("end_turn" | "stop_sequence") | None => FinishReason::Stop, Some("max_tokens") => FinishReason::Length, Some("tool_use") => FinishReason::ToolCalls, Some(other) => FinishReason::Other(other.to_string()), @@ -72,10 +147,734 @@ fn parse_content_block(block: &serde_json::Value) -> Option { block.get("name")?.as_str()?, block.get("input")?.clone(), ))), + "thinking" => Some(ContentPart::Thinking(ThinkingData { + text: block.get("thinking")?.as_str()?.to_string(), + signature: block + .get("signature") + .and_then(serde_json::Value::as_str) + .map(String::from), + redacted: false, + })), + "redacted_thinking" => Some(ContentPart::RedactedThinking(ThinkingData { + text: block + .get("data") + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(), + signature: None, + redacted: true, + })), _ => None, } } +/// Translate a unified `ContentPart` to an Anthropic content block JSON value. +fn content_part_to_api(part: &ContentPart) -> Option { + match part { + ContentPart::Text(text) => Some(serde_json::json!({"type": "text", "text": text})), + ContentPart::ToolCall(tc) => Some(serde_json::json!({ + "type": "tool_use", + "id": tc.id, + "name": tc.name, + "input": tc.arguments, + })), + ContentPart::ToolResult(tr) => { + let content = tr + .content + .as_str() + .map_or_else(|| tr.content.to_string(), str::to_string); + Some(serde_json::json!({ + "type": "tool_result", + "tool_use_id": tr.tool_call_id, + "content": content, + "is_error": tr.is_error, + })) + } + ContentPart::Thinking(td) => { + let mut block = serde_json::json!({ + "type": "thinking", + "thinking": td.text, + }); + if let Some(sig) = &td.signature { + block["signature"] = serde_json::Value::String(sig.clone()); + } + Some(block) + } + ContentPart::RedactedThinking(td) => Some(serde_json::json!({ + "type": "redacted_thinking", + "data": td.text, + })), + ContentPart::Image(img) => { + if let Some(url) = &img.url { + if crate::providers::common::is_file_path(url) { + return match crate::providers::common::load_file_as_base64(url) { + Ok((b64, mime)) => Some(serde_json::json!({ + "type": "image", + "source": {"type": "base64", "media_type": mime, "data": b64} + })), + Err(_) => None, + }; + } + Some(serde_json::json!({"type": "image", "source": {"type": "url", "url": url}})) + } else { + img.data.as_ref().map(|data| { + let mime = img.media_type.as_deref().unwrap_or("image/png"); + let b64 = BASE64_STANDARD.encode(data); + serde_json::json!({"type": "image", "source": {"type": "base64", "media_type": mime, "data": b64}}) + }) + } + } + _ => None, + } +} + +/// Convert unified messages to Anthropic API messages. +/// +/// Handles: role mapping, content block translation, strict alternation +/// (merging consecutive same-role messages), and tool results in user messages. +fn translate_messages(messages: &[&Message]) -> Vec { + let mut api_messages: Vec = Vec::new(); + + for msg in messages { + let role = match msg.role { + Role::Assistant => "assistant", + // Tool results go in user messages for Anthropic + Role::User | Role::Tool => "user", + // System and Developer are extracted separately + Role::System | Role::Developer => continue, + }; + + let content: Vec = msg + .content + .iter() + .filter_map(content_part_to_api) + .collect(); + + if content.is_empty() { + continue; + } + + // Enforce strict user/assistant alternation by merging consecutive same-role messages + if let Some(last) = api_messages.last_mut() { + if last.role == role { + last.content.extend(content); + continue; + } + } + + api_messages.push(ApiMessage { + role: role.to_string(), + content, + }); + } + + api_messages +} + +/// Translate unified `ToolDefinition` to Anthropic format. +fn translate_tools(tools: &[ToolDefinition]) -> Vec { + tools + .iter() + .map(|t| ApiToolDef { + name: t.name.clone(), + description: t.description.clone(), + input_schema: t.parameters.clone(), + cache_control: None, + }) + .collect() +} + +/// Translate unified `ToolChoice` to Anthropic's `tool_choice` JSON. +fn translate_tool_choice(choice: &ToolChoice) -> Option { + match choice { + ToolChoice::Auto => Some(serde_json::json!({"type": "auto"})), + // Anthropic does not support tool_choice none with tools present. + // The caller should omit tools from the request instead. + ToolChoice::None => None, + ToolChoice::Required => Some(serde_json::json!({"type": "any"})), + ToolChoice::Named { tool_name } => { + Some(serde_json::json!({"type": "tool", "name": tool_name})) + } + } +} + +// --- Prompt caching helpers --- + +const CACHE_BETA_HEADER: &str = "prompt-caching-2024-07-31"; + +/// Check whether auto-caching is disabled via `provider_options`. +/// +/// Returns `true` if caching should be applied (the default). +/// Only returns `false` if `provider_options.anthropic.auto_cache` is explicitly `false`. +/// Extract the `thinking` configuration from `provider_options.anthropic.thinking`. +fn extract_thinking_config(provider_options: Option<&serde_json::Value>) -> Option { + provider_options + .and_then(|opts| opts.get("anthropic")) + .and_then(|anthropic| anthropic.get("thinking")) + .cloned() +} + +fn is_auto_cache_enabled(provider_options: Option<&serde_json::Value>) -> bool { + provider_options + .and_then(|opts| opts.get("anthropic")) + .and_then(|anthropic| anthropic.get("auto_cache")) + .and_then(serde_json::Value::as_bool) + .unwrap_or(true) +} + +/// Wrap a system prompt string as an array of content blocks with `cache_control` +/// on the last block. +fn system_with_cache_control(system: &str) -> serde_json::Value { + serde_json::json!([{ + "type": "text", + "text": system, + "cache_control": {"type": "ephemeral"} + }]) +} + +/// Add `cache_control` to the last tool definition. +fn apply_cache_control_to_last_tool(tools: &mut [ApiToolDef]) { + if let Some(last) = tools.last_mut() { + last.cache_control = Some(CacheControl::ephemeral()); + } +} + +/// Add `cache_control` to the last content block of the second-to-last user message. +/// +/// In a multi-turn conversation, the conversation prefix (everything before the latest +/// user turn) is stable and benefits from caching. We find the last user message before +/// the final one and annotate its last content block. +fn apply_cache_control_to_conversation_prefix(messages: &mut [ApiMessage]) { + // Find all user message indices + let user_indices: Vec = messages + .iter() + .enumerate() + .filter(|(_, m)| m.role == "user") + .map(|(i, _)| i) + .collect(); + + // We need at least 2 user messages to have a "prefix" user message + if user_indices.len() < 2 { + return; + } + + // The second-to-last user message is the one to cache + let target_idx = user_indices[user_indices.len() - 2]; + if let Some(serde_json::Value::Object(map)) = messages[target_idx].content.last_mut() { + map.insert( + "cache_control".to_string(), + serde_json::json!({"type": "ephemeral"}), + ); + } +} + +/// Collect beta headers from `provider_options` and merge with the caching header +/// when auto-caching is active. +fn build_beta_header( + provider_options: Option<&serde_json::Value>, + include_cache_header: bool, +) -> Option { + let mut headers: Vec = Vec::new(); + + // Add user-provided beta headers + if let Some(beta_array) = provider_options + .and_then(|opts| opts.get("anthropic")) + .and_then(|anthropic| anthropic.get("beta_headers")) + .and_then(serde_json::Value::as_array) + { + headers.extend( + beta_array + .iter() + .filter_map(serde_json::Value::as_str) + .map(String::from), + ); + } + + // Add prompt-caching header if caching is active and not already present + if include_cache_header && !headers.iter().any(|h| h == CACHE_BETA_HEADER) { + headers.push(CACHE_BETA_HEADER.to_string()); + } + + if headers.is_empty() { + None + } else { + Some(headers.join(",")) + } +} + +// --- Streaming types and helpers --- + +/// The type of the current content block being streamed. +#[derive(Clone)] +enum ContentBlockKind { + Text, + ToolUse { id: String, name: String }, + Thinking { signature: Option }, +} + +/// Accumulated state across SSE events during streaming. +struct StreamAccumulator { + id: String, + model: String, + content_parts: Vec, + usage: Usage, + finish_reason: FinishReason, + /// The kind of the current content block, set by `content_block_start`. + current_block: Option, + /// Accumulated text for the current text block. + current_text: String, + /// Accumulated thinking text for the current thinking block. + current_thinking: String, + /// Accumulated raw JSON arguments for the current `tool_use` block. + current_tool_args: String, + /// Rate limit info parsed from the initial HTTP response headers. + rate_limit: Option, +} + +impl StreamAccumulator { + fn new(rate_limit: Option) -> Self { + Self { + id: String::new(), + model: String::new(), + content_parts: Vec::new(), + usage: Usage::default(), + finish_reason: FinishReason::Stop, + current_block: None, + current_text: String::new(), + current_thinking: String::new(), + current_tool_args: String::new(), + rate_limit, + } + } + + /// Build the final `Response` from accumulated state, consuming content parts. + fn take_response(&mut self) -> Response { + let content_parts = std::mem::take(&mut self.content_parts); + Response { + id: self.id.clone(), + model: self.model.clone(), + provider: "anthropic".to_string(), + message: Message { + role: Role::Assistant, + content: content_parts, + name: None, + tool_call_id: None, + }, + finish_reason: self.finish_reason.clone(), + usage: self.usage.clone(), + raw: None, + warnings: vec![], + rate_limit: self.rate_limit.clone(), + } + } +} + +impl StreamAccumulator { + fn handle_message_start(&mut self, data: &serde_json::Value) -> Vec { + if let Some(message) = data.get("message") { + if let Some(id) = message.get("id").and_then(serde_json::Value::as_str) { + self.id = id.to_string(); + } + if let Some(model) = message.get("model").and_then(serde_json::Value::as_str) { + self.model = model.to_string(); + } + if let Some(usage) = message.get("usage") { + self.usage.input_tokens = usage + .get("input_tokens") + .and_then(serde_json::Value::as_i64) + .unwrap_or(0); + self.usage.cache_read_tokens = usage + .get("cache_read_input_tokens") + .and_then(serde_json::Value::as_i64); + self.usage.cache_write_tokens = usage + .get("cache_creation_input_tokens") + .and_then(serde_json::Value::as_i64); + } + } + vec![StreamEvent::StreamStart] + } + + fn handle_content_block_start(&mut self, data: &serde_json::Value) -> Vec { + let block_type = data + .get("content_block") + .and_then(|b| b.get("type")) + .and_then(serde_json::Value::as_str) + .unwrap_or(""); + + let index = data + .get("index") + .and_then(serde_json::Value::as_u64) + .unwrap_or(0); + let text_id = Some(format!("block_{index}")); + + match block_type { + "text" => { + self.current_block = Some(ContentBlockKind::Text); + self.current_text.clear(); + vec![StreamEvent::TextStart { text_id }] + } + "tool_use" => { + let content_block = data.get("content_block"); + let id = content_block + .and_then(|b| b.get("id")) + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let name = content_block + .and_then(|b| b.get("name")) + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + self.current_block = Some(ContentBlockKind::ToolUse { + id: id.clone(), + name: name.clone(), + }); + self.current_tool_args.clear(); + vec![StreamEvent::ToolCallStart { + tool_call: ToolCall::new(id, name, serde_json::json!({})), + }] + } + "thinking" => { + let signature = data + .get("content_block") + .and_then(|b| b.get("signature")) + .and_then(serde_json::Value::as_str) + .map(String::from); + self.current_block = Some(ContentBlockKind::Thinking { signature }); + self.current_thinking.clear(); + vec![StreamEvent::ReasoningStart] + } + _ => vec![], + } + } + + fn handle_content_block_delta(&mut self, data: &serde_json::Value) -> Vec { + let delta = data.get("delta"); + let delta_type = delta + .and_then(|d| d.get("type")) + .and_then(serde_json::Value::as_str) + .unwrap_or(""); + + match delta_type { + "text_delta" => { + let text = delta + .and_then(|d| d.get("text")) + .and_then(serde_json::Value::as_str) + .unwrap_or(""); + self.current_text.push_str(text); + + let index = data + .get("index") + .and_then(serde_json::Value::as_u64) + .unwrap_or(0); + + vec![StreamEvent::TextDelta { + delta: text.to_string(), + text_id: Some(format!("block_{index}")), + }] + } + "input_json_delta" => { + let partial_json = delta + .and_then(|d| d.get("partial_json")) + .and_then(serde_json::Value::as_str) + .unwrap_or(""); + self.current_tool_args.push_str(partial_json); + + if let Some(ContentBlockKind::ToolUse { id, name }) = &self.current_block { + vec![StreamEvent::ToolCallDelta { + tool_call: ToolCall::new( + id.clone(), + name.clone(), + serde_json::json!(partial_json), + ), + }] + } else { + vec![] + } + } + "thinking_delta" => { + let thinking = delta + .and_then(|d| d.get("thinking")) + .and_then(serde_json::Value::as_str) + .unwrap_or(""); + self.current_thinking.push_str(thinking); + vec![StreamEvent::ReasoningDelta { + delta: thinking.to_string(), + }] + } + _ => vec![], + } + } + + fn handle_content_block_stop(&mut self, data: &serde_json::Value) -> Vec { + let current_block = self.current_block.take(); + match current_block { + Some(ContentBlockKind::Text) => { + let text = std::mem::take(&mut self.current_text); + self.content_parts.push(ContentPart::text(&text)); + + let index = data + .get("index") + .and_then(serde_json::Value::as_u64) + .unwrap_or(0); + + vec![StreamEvent::TextEnd { + text_id: Some(format!("block_{index}")), + }] + } + Some(ContentBlockKind::ToolUse { id, name }) => { + let raw_args = std::mem::take(&mut self.current_tool_args); + let arguments = serde_json::from_str(&raw_args) + .unwrap_or_else(|_| serde_json::json!({})); + let mut tool_call = ToolCall::new(id, name, arguments); + tool_call.raw_arguments = Some(raw_args); + self.content_parts + .push(ContentPart::ToolCall(tool_call.clone())); + vec![StreamEvent::ToolCallEnd { tool_call }] + } + Some(ContentBlockKind::Thinking { signature }) => { + let thinking_text = std::mem::take(&mut self.current_thinking); + // Prefer signature from content_block_stop if available, + // fall back to one captured at content_block_start. + let stop_signature = data + .get("content_block") + .and_then(|b| b.get("signature")) + .and_then(serde_json::Value::as_str) + .map(String::from); + self.content_parts.push(ContentPart::Thinking(ThinkingData { + text: thinking_text, + signature: stop_signature.or(signature), + redacted: false, + })); + vec![StreamEvent::ReasoningEnd] + } + None => vec![], + } + } + + fn handle_message_delta(&mut self, data: &serde_json::Value) { + if let Some(delta) = data.get("delta") { + let stop_reason = delta + .get("stop_reason") + .and_then(serde_json::Value::as_str); + self.finish_reason = map_finish_reason(stop_reason); + } + if let Some(usage) = data.get("usage") { + self.usage.output_tokens = usage + .get("output_tokens") + .and_then(serde_json::Value::as_i64) + .unwrap_or(0); + self.usage.total_tokens = self.usage.input_tokens + self.usage.output_tokens; + } + } + + fn handle_message_stop(&mut self) -> Vec { + let response = self.take_response(); + vec![StreamEvent::Finish { + finish_reason: response.finish_reason.clone(), + usage: response.usage.clone(), + response: Box::new(response), + }] + } +} + +/// Process a single SSE event and return zero or more `StreamEvent`s. +fn process_sse_event( + event_type: &str, + data: &serde_json::Value, + acc: &mut StreamAccumulator, +) -> Vec { + match event_type { + "message_start" => acc.handle_message_start(data), + "content_block_start" => acc.handle_content_block_start(data), + "content_block_delta" => acc.handle_content_block_delta(data), + "content_block_stop" => acc.handle_content_block_stop(data), + "message_delta" => { + acc.handle_message_delta(data); + vec![] + } + "message_stop" => acc.handle_message_stop(), + _ => vec![], + } +} + +// --- SSE reader --- + +enum SseResult { + Event { event_type: String, data: String }, + Done, + Error(SdkError), +} + +struct SseReaderState { + byte_stream: futures::stream::BoxStream<'static, Result>, + buffer: String, + accumulator: StreamAccumulator, + pending_events: std::collections::VecDeque, + done: bool, +} + +impl SseReaderState { + fn new( + byte_stream: impl futures::Stream> + + Send + + 'static, + rate_limit: Option, + ) -> Self { + use futures::StreamExt; + Self { + byte_stream: byte_stream.boxed(), + buffer: String::new(), + accumulator: StreamAccumulator::new(rate_limit), + pending_events: std::collections::VecDeque::new(), + done: false, + } + } + + /// Read the next complete SSE event from the byte stream. + /// + /// SSE events are separated by double newlines. Each event has optional + /// `event:` and `data:` lines. + async fn next_sse_event(&mut self) -> SseResult { + use futures::StreamExt; + + loop { + // Try to extract a complete SSE event from the buffer. + if let Some(result) = self.try_parse_event() { + return result; + } + + if self.done { + return SseResult::Done; + } + + // Read more bytes from the stream. + match self.byte_stream.next().await { + Some(Ok(chunk)) => { + let text = String::from_utf8_lossy(&chunk); + self.buffer.push_str(&text); + } + Some(Err(e)) => { + return SseResult::Error(SdkError::Stream { + message: e.to_string(), + }); + } + None => { + self.done = true; + // Try one more time to parse any remaining data. + if let Some(result) = self.try_parse_event() { + return result; + } + return SseResult::Done; + } + } + } + } + + /// Attempt to parse one complete SSE event from the buffer. + /// + /// Returns `None` if no complete event is available yet. + fn try_parse_event(&mut self) -> Option { + // SSE events are terminated by a blank line (double newline). + let separator = self.buffer.find("\n\n")?; + let event_block = self.buffer[..separator].to_string(); + self.buffer = self.buffer[separator + 2..].to_string(); + + let mut event_type = String::new(); + let mut data_parts: Vec = Vec::new(); + + for line in event_block.lines() { + if let Some(rest) = line.strip_prefix("event:") { + event_type = rest.trim().to_string(); + } else if let Some(rest) = line.strip_prefix("data:") { + data_parts.push(rest.trim().to_string()); + } + // Ignore other SSE fields (id:, retry:, comments starting with :) + } + + // Skip events with no data (e.g. heartbeat comments). + if data_parts.is_empty() { + return None; + } + + let data = data_parts.join("\n"); + Some(SseResult::Event { event_type, data }) + } +} + +/// Build the streaming request body and headers, mirroring `complete()` logic +/// but with `stream: true`. +fn build_stream_request( + adapter: &Adapter, + request: &Request, +) -> (ApiRequest, reqwest::RequestBuilder) { + let (system, other_messages) = extract_system_prompt(&request.messages); + let mut api_messages = translate_messages(&other_messages); + + let mut omit_tools = false; + let tool_choice_json = request.tool_choice.as_ref().and_then(|tc| { + if matches!(tc, ToolChoice::None) { + omit_tools = true; + None + } else { + translate_tool_choice(tc) + } + }); + + let mut api_tools = if omit_tools { + None + } else { + request.tools.as_ref().map(|t| translate_tools(t)) + }; + + let auto_cache = is_auto_cache_enabled(request.provider_options.as_ref()); + + let system_value = system.map(|s| { + if auto_cache { + system_with_cache_control(&s) + } else { + serde_json::Value::String(s) + } + }); + + if auto_cache { + if let Some(ref mut tools) = api_tools { + apply_cache_control_to_last_tool(tools); + } + apply_cache_control_to_conversation_prefix(&mut api_messages); + } + + let thinking = extract_thinking_config(request.provider_options.as_ref()); + + let api_request = ApiRequest { + model: request.model.clone(), + messages: api_messages, + max_tokens: request.max_tokens.unwrap_or(4096), + system: system_value, + temperature: request.temperature, + top_p: request.top_p, + stop_sequences: request.stop_sequences.clone(), + tools: api_tools, + tool_choice: tool_choice_json, + thinking, + stream: true, + }; + + let url = adapter.messages_url(); + let mut req_builder = adapter + .client + .post(&url) + .header("x-api-key", &adapter.api_key) + .header("anthropic-version", "2023-06-01"); + + if let Some(beta_str) = build_beta_header(request.provider_options.as_ref(), auto_cache) { + req_builder = req_builder.header("anthropic-beta", beta_str); + } + + let req_builder = req_builder.json(&api_request); + (api_request, req_builder) +} + #[allow(clippy::unnecessary_literal_bound)] #[async_trait::async_trait] impl ProviderAdapter for Adapter { @@ -85,41 +884,78 @@ impl ProviderAdapter for Adapter { async fn complete(&self, request: &Request) -> Result { let (system, other_messages) = extract_system_prompt(&request.messages); + let mut api_messages = translate_messages(&other_messages); - let api_messages: Vec = other_messages - .iter() - .map(|msg| { - let role = match msg.role { - Role::Assistant => "assistant", - Role::System | Role::User | Role::Tool | Role::Developer => "user", - }; - ApiMessage { - role: role.to_string(), - content: msg.text(), - } - }) - .collect(); + // Handle tools and tool_choice + let mut omit_tools = false; + let tool_choice_json = request.tool_choice.as_ref().and_then(|tc| { + if matches!(tc, ToolChoice::None) { + omit_tools = true; + None + } else { + translate_tool_choice(tc) + } + }); + + let mut api_tools = if omit_tools { + None + } else { + request.tools.as_ref().map(|t| translate_tools(t)) + }; + + // Determine if auto-caching is enabled + let auto_cache = is_auto_cache_enabled(request.provider_options.as_ref()); + + // Build system prompt value, optionally with cache_control + let system_value = system.map(|s| { + if auto_cache { + system_with_cache_control(&s) + } else { + serde_json::Value::String(s) + } + }); + + // Apply cache_control breakpoints when auto-caching is enabled + if auto_cache { + if let Some(ref mut tools) = api_tools { + apply_cache_control_to_last_tool(tools); + } + apply_cache_control_to_conversation_prefix(&mut api_messages); + } + + let thinking = extract_thinking_config(request.provider_options.as_ref()); let api_request = ApiRequest { model: request.model.clone(), messages: api_messages, - max_tokens: request.max_tokens.unwrap_or(1024), - system, + max_tokens: request.max_tokens.unwrap_or(4096), + system: system_value, temperature: request.temperature, top_p: request.top_p, stop_sequences: request.stop_sequences.clone(), + tools: api_tools, + tool_choice: tool_choice_json, + thinking, + stream: false, }; - let body = send_and_read_body( - self.client - .post("https://api.anthropic.com/v1/messages") - .header("x-api-key", &self.api_key) - .header("anthropic-version", "2023-06-01") - .json(&api_request), - "anthropic", - "type", - ) - .await?; + // Build headers + let url = self.messages_url(); + let mut req_builder = self + .client + .post(&url) + .header("x-api-key", &self.api_key) + .header("anthropic-version", "2023-06-01"); + + // Build beta header: merge user-provided headers with cache header when active + if let Some(beta_str) = + build_beta_header(request.provider_options.as_ref(), auto_cache) + { + req_builder = req_builder.header("anthropic-beta", beta_str); + } + + let (body, headers) = + send_and_read_response(req_builder.json(&api_request).timeout(self.request_timeout), "anthropic", "type").await?; let api_resp: ApiResponse = serde_json::from_str(&body).map_err(|e| SdkError::Network { @@ -150,17 +986,394 @@ impl ProviderAdapter for Adapter { input_tokens: api_resp.usage.input_tokens, output_tokens: api_resp.usage.output_tokens, total_tokens: total, + cache_read_tokens: api_resp.usage.cache_read_input_tokens, + cache_write_tokens: api_resp.usage.cache_creation_input_tokens, ..Usage::default() }, raw: serde_json::from_str(&body).ok(), warnings: vec![], - rate_limit: None, + rate_limit: parse_rate_limit_headers(&headers), }) } - async fn stream(&self, _request: &Request) -> Result { - Err(SdkError::Configuration { - message: "streaming not yet implemented".to_string(), - }) + async fn stream(&self, request: &Request) -> Result { + let (_api_request, req_builder) = build_stream_request(self, request); + + let http_resp = req_builder.send().await.map_err(|e| SdkError::Network { + message: e.to_string(), + })?; + + let status = http_resp.status(); + if !status.is_success() { + let body = http_resp.text().await.map_err(|e| SdkError::Network { + message: e.to_string(), + })?; + let (msg, code, raw) = parse_error_body(&body, "type"); + return Err(crate::error::error_from_status_code( + status.as_u16(), + msg, + "anthropic".to_string(), + code, + raw, + None, + )); + } + + let rate_limit = parse_rate_limit_headers(http_resp.headers()); + let byte_stream = http_resp.bytes_stream(); + + let stream = futures::stream::unfold( + SseReaderState::new(byte_stream, rate_limit), + |mut state| async move { + loop { + // Drain any buffered events first. + if let Some(event) = state.pending_events.pop_front() { + return Some((Ok(event), state)); + } + + // Read more SSE data from the byte stream. + match state.next_sse_event().await { + SseResult::Event { event_type, data } => { + let parsed: serde_json::Value = match serde_json::from_str(&data) { + Ok(v) => v, + Err(e) => { + return Some(( + Err(SdkError::Stream { + message: format!("failed to parse SSE data: {e}"), + }), + state, + )); + } + }; + let events = + process_sse_event(&event_type, &parsed, &mut state.accumulator); + state.pending_events.extend(events); + // Loop to drain from pending_events. + } + SseResult::Done => return None, + SseResult::Error(err) => return Some((Err(err), state)), + } + } + }, + ); + + Ok(Box::pin(stream)) + } + + fn supports_tool_choice(&self, mode: &str) -> bool { + matches!(mode, "auto" | "none" | "required" | "named") + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn auto_cache_enabled_by_default() { + assert!(is_auto_cache_enabled(None)); + } + + #[test] + fn auto_cache_enabled_when_true() { + let opts = serde_json::json!({"anthropic": {"auto_cache": true}}); + assert!(is_auto_cache_enabled(Some(&opts))); + } + + #[test] + fn auto_cache_disabled_when_false() { + let opts = serde_json::json!({"anthropic": {"auto_cache": false}}); + assert!(!is_auto_cache_enabled(Some(&opts))); + } + + #[test] + fn auto_cache_enabled_when_key_missing() { + let opts = serde_json::json!({"anthropic": {}}); + assert!(is_auto_cache_enabled(Some(&opts))); + } + + #[test] + fn auto_cache_enabled_when_anthropic_missing() { + let opts = serde_json::json!({"openai": {}}); + assert!(is_auto_cache_enabled(Some(&opts))); + } + + #[test] + fn system_prompt_cache_control_wraps_as_array() { + let result = system_with_cache_control("You are helpful."); + let arr = result.as_array().expect("should be an array"); + assert_eq!(arr.len(), 1); + assert_eq!(arr[0]["type"], "text"); + assert_eq!(arr[0]["text"], "You are helpful."); + assert_eq!(arr[0]["cache_control"]["type"], "ephemeral"); + } + + #[test] + fn tool_cache_control_applied_to_last_tool() { + let mut tools = vec![ + ApiToolDef { + name: "tool_a".to_string(), + description: "first".to_string(), + input_schema: serde_json::json!({}), + cache_control: None, + }, + ApiToolDef { + name: "tool_b".to_string(), + description: "second".to_string(), + input_schema: serde_json::json!({}), + cache_control: None, + }, + ]; + apply_cache_control_to_last_tool(&mut tools); + + assert!(tools[0].cache_control.is_none()); + assert!(tools[1].cache_control.is_some()); + assert_eq!(tools[1].cache_control.as_ref().unwrap().kind, "ephemeral"); + } + + #[test] + fn tool_cache_control_empty_slice() { + let mut tools: Vec = vec![]; + apply_cache_control_to_last_tool(&mut tools); + assert!(tools.is_empty()); + } + + #[test] + fn tool_cache_control_single_tool() { + let mut tools = vec![ApiToolDef { + name: "only_tool".to_string(), + description: "the one".to_string(), + input_schema: serde_json::json!({}), + cache_control: None, + }]; + apply_cache_control_to_last_tool(&mut tools); + assert!(tools[0].cache_control.is_some()); + } + + #[test] + fn conversation_prefix_cache_control_with_two_user_messages() { + let mut messages = vec![ + ApiMessage { + role: "user".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Hello"})], + }, + ApiMessage { + role: "assistant".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Hi there"})], + }, + ApiMessage { + role: "user".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "How are you?"})], + }, + ]; + + apply_cache_control_to_conversation_prefix(&mut messages); + + // First user message should have cache_control + assert_eq!( + messages[0].content[0]["cache_control"]["type"], + "ephemeral" + ); + // Last user message should NOT have cache_control + assert!(messages[2].content[0].get("cache_control").is_none()); + // Assistant message should NOT have cache_control + assert!(messages[1].content[0].get("cache_control").is_none()); + } + + #[test] + fn conversation_prefix_cache_control_with_multiple_content_blocks() { + let mut messages = vec![ + ApiMessage { + role: "user".to_string(), + content: vec![ + serde_json::json!({"type": "text", "text": "Part 1"}), + serde_json::json!({"type": "text", "text": "Part 2"}), + ], + }, + ApiMessage { + role: "assistant".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Reply"})], + }, + ApiMessage { + role: "user".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Follow up"})], + }, + ]; + + apply_cache_control_to_conversation_prefix(&mut messages); + + // Only the LAST content block of the first user message should have cache_control + assert!(messages[0].content[0].get("cache_control").is_none()); + assert_eq!( + messages[0].content[1]["cache_control"]["type"], + "ephemeral" + ); + } + + #[test] + fn conversation_prefix_cache_control_single_user_message() { + let mut messages = vec![ApiMessage { + role: "user".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Hello"})], + }]; + + apply_cache_control_to_conversation_prefix(&mut messages); + + // With only one user message, no cache_control should be added + assert!(messages[0].content[0].get("cache_control").is_none()); + } + + #[test] + fn conversation_prefix_cache_control_no_user_messages() { + let mut messages: Vec = vec![]; + // Should not panic on empty messages + apply_cache_control_to_conversation_prefix(&mut messages); + } + + #[test] + fn conversation_prefix_cache_control_three_user_messages() { + let mut messages = vec![ + ApiMessage { + role: "user".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "First"})], + }, + ApiMessage { + role: "assistant".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Reply 1"})], + }, + ApiMessage { + role: "user".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Second"})], + }, + ApiMessage { + role: "assistant".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Reply 2"})], + }, + ApiMessage { + role: "user".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Third"})], + }, + ]; + + apply_cache_control_to_conversation_prefix(&mut messages); + + // Only the second-to-last user message (index 2) should get cache_control + assert!(messages[0].content[0].get("cache_control").is_none()); + assert_eq!( + messages[2].content[0]["cache_control"]["type"], + "ephemeral" + ); + assert!(messages[4].content[0].get("cache_control").is_none()); + } + + #[test] + fn beta_header_includes_cache_header() { + let result = build_beta_header(None, true); + assert_eq!(result, Some(CACHE_BETA_HEADER.to_string())); + } + + #[test] + fn beta_header_no_cache_no_user_headers() { + let result = build_beta_header(None, false); + assert_eq!(result, None); + } + + #[test] + fn beta_header_merges_user_headers_with_cache() { + let opts = serde_json::json!({ + "anthropic": { + "beta_headers": ["interleaved-thinking-2025-05-14"] + } + }); + let result = build_beta_header(Some(&opts), true); + assert_eq!( + result, + Some(format!( + "interleaved-thinking-2025-05-14,{CACHE_BETA_HEADER}" + )) + ); + } + + #[test] + fn beta_header_no_duplicate_cache_header() { + let opts = serde_json::json!({ + "anthropic": { + "beta_headers": [CACHE_BETA_HEADER] + } + }); + let result = build_beta_header(Some(&opts), true); + // Should not duplicate the header + assert_eq!(result, Some(CACHE_BETA_HEADER.to_string())); + } + + #[test] + fn beta_header_user_headers_only_when_cache_disabled() { + let opts = serde_json::json!({ + "anthropic": { + "beta_headers": ["interleaved-thinking-2025-05-14"] + } + }); + let result = build_beta_header(Some(&opts), false); + assert_eq!( + result, + Some("interleaved-thinking-2025-05-14".to_string()) + ); + } + + #[test] + fn tool_serialization_includes_cache_control() { + let tool = ApiToolDef { + name: "test_tool".to_string(), + description: "A test tool".to_string(), + input_schema: serde_json::json!({"type": "object"}), + cache_control: Some(CacheControl::ephemeral()), + }; + let json = serde_json::to_value(&tool).expect("should serialize"); + assert_eq!(json["cache_control"]["type"], "ephemeral"); + } + + #[test] + fn tool_serialization_omits_cache_control_when_none() { + let tool = ApiToolDef { + name: "test_tool".to_string(), + description: "A test tool".to_string(), + input_schema: serde_json::json!({"type": "object"}), + cache_control: None, + }; + let json = serde_json::to_value(&tool).expect("should serialize"); + assert!(json.get("cache_control").is_none()); + } + + #[test] + fn system_prompt_as_string_when_cache_disabled() { + let system = "You are helpful.".to_string(); + let value = serde_json::Value::String(system); + assert_eq!(value.as_str(), Some("You are helpful.")); + } + + #[test] + fn api_request_serialization_with_cached_system() { + let api_request = ApiRequest { + model: "claude-sonnet-4-20250514".to_string(), + messages: vec![ApiMessage { + role: "user".to_string(), + content: vec![serde_json::json!({"type": "text", "text": "Hello"})], + }], + max_tokens: 4096, + system: Some(system_with_cache_control("You are helpful.")), + temperature: None, + top_p: None, + stop_sequences: None, + tools: None, + tool_choice: None, + thinking: None, + stream: false, + }; + + let json = serde_json::to_value(&api_request).expect("should serialize"); + let system = json.get("system").expect("system should be present"); + let arr = system.as_array().expect("system should be an array"); + assert_eq!(arr.len(), 1); + assert_eq!(arr[0]["cache_control"]["type"], "ephemeral"); } } diff --git a/crates/unified-llm/src/providers/common.rs b/crates/unified-llm/src/providers/common.rs index 3c9b70a96..662bd6659 100644 --- a/crates/unified-llm/src/providers/common.rs +++ b/crates/unified-llm/src/providers/common.rs @@ -1,5 +1,7 @@ +use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine}; + use crate::error::{error_from_status_code, SdkError}; -use crate::types::{Message, Role}; +use crate::types::{Message, RateLimitInfo, Role}; #[derive(serde::Serialize)] pub struct ApiMessage { @@ -8,6 +10,7 @@ pub struct ApiMessage { } /// Parse an error response body, extracting the message and error code. +/// /// `error_code_field` is the JSON field name for the error code (e.g. "type" or "status"). #[must_use] pub fn parse_error_body( @@ -33,7 +36,9 @@ pub fn parse_error_body( ) } -/// Send an HTTP request and read the response body, returning an error on non-success status. +/// Send an HTTP request and read the response body. +/// +/// Returns an error on non-success status. /// /// # Errors /// @@ -43,14 +48,20 @@ pub async fn send_and_read_body( provider: &str, error_code_field: &str, ) -> Result { - let http_resp = request - .send() - .await - .map_err(|e| SdkError::Network { - message: e.to_string(), - })?; + let http_resp = request.send().await.map_err(|e| { + if e.is_timeout() { + SdkError::RequestTimeout { + message: format!("{provider}: {e}"), + } + } else { + SdkError::Network { + message: e.to_string(), + } + } + })?; let status = http_resp.status(); + let retry_after = parse_retry_after(http_resp.headers()); let body = http_resp .text() .await @@ -66,21 +77,24 @@ pub async fn send_and_read_body( provider.to_string(), code, raw, - None, + retry_after, )); } Ok(body) } -/// Extract system messages from a message list, returning the joined system prompt -/// and the remaining non-system messages. +/// Extract system and developer messages from a message list. +/// +/// Returns the joined system prompt and the remaining messages. +/// Per spec, Developer role messages are merged with system messages +/// for Anthropic and Gemini. #[must_use] pub fn extract_system_prompt(messages: &[Message]) -> (Option, Vec<&Message>) { let mut system_parts = Vec::new(); let mut other = Vec::new(); for msg in messages { - if msg.role == Role::System { + if msg.role == Role::System || msg.role == Role::Developer { system_parts.push(msg.text()); } else { other.push(msg); @@ -93,3 +107,257 @@ pub fn extract_system_prompt(messages: &[Message]) -> (Option, Vec<&Mess }; (system, other) } + +/// Check if a URL string looks like a local file path. +#[must_use] +pub fn is_file_path(url: &str) -> bool { + url.starts_with('/') || url.starts_with("./") || url.starts_with("~/") +} + +/// Infer MIME type from a file extension. +#[must_use] +pub fn mime_from_extension(path: &str) -> &str { + match path.rsplit('.').next().map(str::to_lowercase).as_deref() { + Some("png") => "image/png", + Some("jpg" | "jpeg") => "image/jpeg", + Some("gif") => "image/gif", + Some("webp") => "image/webp", + Some("heic") => "image/heic", + Some("heif") => "image/heif", + Some("pdf") => "application/pdf", + Some("wav") => "audio/wav", + Some("mp3") => "audio/mp3", + _ => "application/octet-stream", + } +} + +/// Load a local file, returning (`base64_data`, `mime_type`). +/// Expands ~ to home directory. +/// +/// # Errors +/// Returns an error if the file cannot be read. +pub fn load_file_as_base64(path: &str) -> Result<(String, String), std::io::Error> { + let expanded = path.strip_prefix("~/").map_or_else( + || path.to_string(), + |rest| { + let home = std::env::var("HOME").unwrap_or_else(|_| "/".to_string()); + format!("{home}/{rest}") + }, + ); + let data = std::fs::read(&expanded)?; + let mime = mime_from_extension(&expanded).to_string(); + let b64 = BASE64_STANDARD.encode(&data); + Ok((b64, mime)) +} + +/// Extract the `Retry-After` header value from an HTTP response as seconds. +#[must_use] +pub fn parse_retry_after(headers: &reqwest::header::HeaderMap) -> Option { + headers + .get("retry-after") + .and_then(|v| v.to_str().ok()) + .and_then(|s| s.parse::().ok()) +} + +/// Parse `x-ratelimit-*` headers into a `RateLimitInfo`. +/// +/// Returns `None` if no rate limit headers are present. +#[must_use] +pub fn parse_rate_limit_headers(headers: &reqwest::header::HeaderMap) -> Option { + fn header_i64(headers: &reqwest::header::HeaderMap, name: &str) -> Option { + headers + .get(name) + .and_then(|v| v.to_str().ok()) + .and_then(|s| s.parse::().ok()) + } + + fn header_str(headers: &reqwest::header::HeaderMap, name: &str) -> Option { + headers + .get(name) + .and_then(|v| v.to_str().ok()) + .map(String::from) + } + + let requests_remaining = header_i64(headers, "x-ratelimit-remaining-requests"); + let requests_limit = header_i64(headers, "x-ratelimit-limit-requests"); + let tokens_remaining = header_i64(headers, "x-ratelimit-remaining-tokens"); + let tokens_limit = header_i64(headers, "x-ratelimit-limit-tokens"); + let reset_at = header_str(headers, "x-ratelimit-reset-requests") + .or_else(|| header_str(headers, "x-ratelimit-reset-tokens")); + + if requests_remaining.is_none() + && requests_limit.is_none() + && tokens_remaining.is_none() + && tokens_limit.is_none() + && reset_at.is_none() + { + return None; + } + + Some(RateLimitInfo { + requests_remaining, + requests_limit, + tokens_remaining, + tokens_limit, + reset_at, + }) +} + +/// Send an HTTP request, read the response body, and return it along with the response headers. +/// +/// Returns an error on non-success status. +/// +/// # Errors +/// +/// Returns `SdkError::Network` on connection failure or `SdkError::Provider` on non-success status. +pub async fn send_and_read_response( + request: reqwest::RequestBuilder, + provider: &str, + error_code_field: &str, +) -> Result<(String, reqwest::header::HeaderMap), SdkError> { + let http_resp = request.send().await.map_err(|e| { + if e.is_timeout() { + SdkError::RequestTimeout { + message: format!("{provider}: {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, error_code_field); + return Err(error_from_status_code( + status.as_u16(), + msg, + provider.to_string(), + code, + raw, + retry_after, + )); + } + + Ok((body, headers)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn is_file_path_absolute() { + assert!(is_file_path("/tmp/image.png")); + assert!(is_file_path("/home/user/photo.jpg")); + } + + #[test] + fn is_file_path_relative() { + assert!(is_file_path("./image.png")); + assert!(is_file_path("./subdir/photo.jpg")); + } + + #[test] + fn is_file_path_tilde() { + assert!(is_file_path("~/image.png")); + assert!(is_file_path("~/Documents/photo.jpg")); + } + + #[test] + fn is_file_path_url() { + assert!(!is_file_path("https://example.com/image.png")); + assert!(!is_file_path("http://example.com/image.png")); + assert!(!is_file_path("data:image/png;base64,abc")); + } + + #[test] + fn mime_from_extension_known() { + assert_eq!(mime_from_extension("photo.png"), "image/png"); + assert_eq!(mime_from_extension("photo.jpg"), "image/jpeg"); + assert_eq!(mime_from_extension("photo.jpeg"), "image/jpeg"); + assert_eq!(mime_from_extension("photo.gif"), "image/gif"); + assert_eq!(mime_from_extension("photo.webp"), "image/webp"); + assert_eq!(mime_from_extension("doc.pdf"), "application/pdf"); + } + + #[test] + fn mime_from_extension_unknown() { + assert_eq!(mime_from_extension("file.xyz"), "application/octet-stream"); + assert_eq!(mime_from_extension("noext"), "application/octet-stream"); + } + + #[test] + fn parse_rate_limit_headers_all_present() { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("x-ratelimit-remaining-requests", "99".parse().unwrap()); + headers.insert("x-ratelimit-limit-requests", "100".parse().unwrap()); + headers.insert("x-ratelimit-remaining-tokens", "9000".parse().unwrap()); + headers.insert("x-ratelimit-limit-tokens", "10000".parse().unwrap()); + headers.insert( + "x-ratelimit-reset-requests", + "2024-01-01T00:00:00Z".parse().unwrap(), + ); + + let info = parse_rate_limit_headers(&headers).unwrap(); + assert_eq!(info.requests_remaining, Some(99)); + assert_eq!(info.requests_limit, Some(100)); + assert_eq!(info.tokens_remaining, Some(9000)); + assert_eq!(info.tokens_limit, Some(10000)); + assert_eq!(info.reset_at.as_deref(), Some("2024-01-01T00:00:00Z")); + } + + #[test] + fn parse_rate_limit_headers_none_present() { + let headers = reqwest::header::HeaderMap::new(); + assert!(parse_rate_limit_headers(&headers).is_none()); + } + + #[test] + fn parse_rate_limit_headers_partial() { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("x-ratelimit-remaining-requests", "50".parse().unwrap()); + + let info = parse_rate_limit_headers(&headers).unwrap(); + assert_eq!(info.requests_remaining, Some(50)); + assert_eq!(info.requests_limit, None); + assert_eq!(info.tokens_remaining, None); + assert_eq!(info.tokens_limit, None); + assert_eq!(info.reset_at, None); + } + + #[test] + fn parse_rate_limit_headers_reset_tokens_fallback() { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("x-ratelimit-limit-tokens", "5000".parse().unwrap()); + headers.insert( + "x-ratelimit-reset-tokens", + "2024-06-01T12:00:00Z".parse().unwrap(), + ); + + let info = parse_rate_limit_headers(&headers).unwrap(); + assert_eq!(info.tokens_limit, Some(5000)); + assert_eq!(info.reset_at.as_deref(), Some("2024-06-01T12:00:00Z")); + } + + #[test] + fn parse_rate_limit_headers_invalid_values_ignored() { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("x-ratelimit-remaining-requests", "not-a-number".parse().unwrap()); + headers.insert("x-ratelimit-limit-tokens", "10000".parse().unwrap()); + + let info = parse_rate_limit_headers(&headers).unwrap(); + assert_eq!(info.requests_remaining, None); + assert_eq!(info.tokens_limit, Some(10000)); + } +} diff --git a/crates/unified-llm/src/providers/gemini.rs b/crates/unified-llm/src/providers/gemini.rs index 334ad68b8..e60ef50e3 100644 --- a/crates/unified-llm/src/providers/gemini.rs +++ b/crates/unified-llm/src/providers/gemini.rs @@ -1,25 +1,47 @@ -use crate::error::SdkError; +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::provider::{ProviderAdapter, StreamEventStream}; -use crate::error::{ProviderErrorDetail, ProviderErrorKind}; -use crate::providers::common::{extract_system_prompt, send_and_read_body}; -use crate::types::{ - ContentPart, FinishReason, Message, Request, Response, Role, ToolCall, Usage, +use crate::providers::common::{ + extract_system_prompt, parse_error_body, parse_retry_after, send_and_read_body, }; +use crate::types::{ + ContentPart, FinishReason, Message, Request, Response, ResponseFormat, ResponseFormatType, + Role, StreamEvent, ToolCall, ToolChoice, ToolDefinition, Usage, +}; + +const DEFAULT_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta"; /// Provider adapter for the Google Gemini `generateContent` API. pub struct Adapter { api_key: String, + base_url: String, client: reqwest::Client, + request_timeout: std::time::Duration, } impl Adapter { #[must_use] pub fn new(api_key: impl Into) -> Self { + let timeout = crate::types::AdapterTimeout::default(); + let client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs_f64(timeout.connect)) + .build() + .unwrap_or_default(); Self { api_key: api_key.into(), - client: reqwest::Client::new(), + base_url: DEFAULT_BASE_URL.to_string(), + client, + request_timeout: std::time::Duration::from_secs_f64(timeout.request), } } + + #[must_use] + pub fn with_base_url(mut self, base_url: impl Into) -> Self { + self.base_url = base_url.into(); + self + } } // --- Request types --- @@ -32,22 +54,21 @@ struct ApiRequest { system_instruction: Option, #[serde(skip_serializing_if = "Option::is_none")] generation_config: Option, + #[serde(skip_serializing_if = "Option::is_none")] + tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + tool_config: Option, } #[derive(serde::Serialize)] struct Content { role: String, - parts: Vec, + parts: Vec, } #[derive(serde::Serialize)] struct SystemInstruction { - parts: Vec, -} - -#[derive(serde::Serialize)] -struct Part { - text: String, + parts: Vec, } #[derive(serde::Serialize)] @@ -61,6 +82,24 @@ struct GenerationConfig { top_p: Option, #[serde(skip_serializing_if = "Option::is_none")] stop_sequences: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + response_mime_type: Option, + #[serde(skip_serializing_if = "Option::is_none")] + response_schema: Option, +} + +/// Gemini groups function declarations under a `tools` array. +#[derive(serde::Serialize)] +#[serde(rename_all = "camelCase")] +struct GeminiToolGroup { + function_declarations: Vec, +} + +#[derive(serde::Serialize)] +struct GeminiFunctionDecl { + name: String, + description: String, + parameters: serde_json::Value, } // --- Response types --- @@ -91,9 +130,15 @@ struct UsageMetadata { prompt_token_count: Option, candidates_token_count: Option, total_token_count: Option, + thoughts_token_count: Option, + cached_content_token_count: Option, } -fn map_finish_reason(reason: Option<&str>) -> FinishReason { +/// Map Gemini's finish reason, inferring `ToolCalls` from content when needed. +fn map_finish_reason(reason: Option<&str>, has_function_calls: bool) -> FinishReason { + if has_function_calls { + return FinishReason::ToolCalls; + } match reason { Some("STOP") | None => FinishReason::Stop, Some("MAX_TOKENS") => FinishReason::Length, @@ -121,6 +166,520 @@ fn parse_part(part: &serde_json::Value) -> Option { None } +/// Check if any parts contain function calls. +fn parts_have_function_calls(parts: &[serde_json::Value]) -> bool { + parts.iter().any(|p| p.get("functionCall").is_some()) +} + +/// Build a mapping from tool call ID to function name by scanning assistant messages. +/// +/// Gemini uses function names (not call IDs) in `functionResponse`. Since the adapter +/// generates synthetic UUIDs as tool call IDs, we need this mapping to recover the +/// original function name when sending tool results back. +fn build_tool_call_id_to_name(messages: &[&Message]) -> std::collections::HashMap { + let mut map = std::collections::HashMap::new(); + for msg in messages { + if msg.role == Role::Assistant { + for part in &msg.content { + if let ContentPart::ToolCall(tc) = part { + map.insert(tc.id.clone(), tc.name.clone()); + } + } + } + } + map +} + +/// Translate unified messages to Gemini content format. +fn translate_messages(messages: &[&Message]) -> Vec { + let id_to_name = build_tool_call_id_to_name(messages); + let mut contents: Vec = Vec::new(); + + for msg in messages { + let role = match msg.role { + Role::Assistant => "model", + Role::User | Role::Tool => "user", + Role::System | Role::Developer => continue, + }; + + let parts: Vec = msg + .content + .iter() + .filter_map(|part| match part { + ContentPart::Text(text) => Some(serde_json::json!({"text": text})), + ContentPart::ToolCall(tc) => Some(serde_json::json!({ + "functionCall": { + "name": tc.name, + "args": tc.arguments, + } + })), + ContentPart::Image(img) => { + img.url.as_ref().map_or_else( + || { + img.data.as_ref().map(|data| { + let mime = img.media_type.as_deref().unwrap_or("image/png"); + let b64 = BASE64_STANDARD.encode(data); + serde_json::json!({"inlineData": {"mimeType": mime, "data": b64}}) + }) + }, + |url| { + if crate::providers::common::is_file_path(url) { + match crate::providers::common::load_file_as_base64(url) { + Ok((b64, mime)) => Some(serde_json::json!({"inlineData": {"mimeType": mime, "data": b64}})), + Err(_) => None, + } + } else { + let mime = img.media_type.as_deref().unwrap_or("image/png"); + Some(serde_json::json!({"fileData": {"mimeType": mime, "fileUri": url}})) + } + }, + ) + } + ContentPart::ToolResult(tr) => { + // Gemini's functionResponse uses the function *name*, not the call ID. + // Look up the original function name from the tool call mapping. + let function_name = id_to_name + .get(&tr.tool_call_id) + .cloned() + .unwrap_or_else(|| tr.tool_call_id.clone()); + let response = tr.content.as_str().map_or_else( + || { + if tr.content.is_object() { + tr.content.clone() + } else { + serde_json::json!({"result": tr.content.to_string()}) + } + }, + |s| serde_json::json!({"result": s}), + ); + Some(serde_json::json!({ + "functionResponse": { + "name": function_name, + "response": response, + } + })) + } + _ => None, + }) + .collect(); + + if parts.is_empty() { + continue; + } + + contents.push(Content { + role: role.to_string(), + parts, + }); + } + + contents +} + +/// Translate unified tool definitions to Gemini's format. +fn translate_tools(tools: &[ToolDefinition]) -> Vec { + vec![GeminiToolGroup { + function_declarations: tools + .iter() + .map(|t| GeminiFunctionDecl { + name: t.name.clone(), + description: t.description.clone(), + parameters: t.parameters.clone(), + }) + .collect(), + }] +} + +/// Translate unified `ToolChoice` to Gemini's `toolConfig`. +fn translate_tool_choice(choice: &ToolChoice) -> serde_json::Value { + match choice { + ToolChoice::Auto => serde_json::json!({ + "functionCallingConfig": {"mode": "AUTO"} + }), + ToolChoice::None => serde_json::json!({ + "functionCallingConfig": {"mode": "NONE"} + }), + ToolChoice::Required => serde_json::json!({ + "functionCallingConfig": {"mode": "ANY"} + }), + ToolChoice::Named { tool_name } => serde_json::json!({ + "functionCallingConfig": { + "mode": "ANY", + "allowedFunctionNames": [tool_name], + } + }), + } +} + +/// Translate unified `ResponseFormat` to Gemini generation config fields. +/// +/// Returns `(response_mime_type, response_schema)`. +fn translate_response_format( + format: &ResponseFormat, +) -> (Option, Option) { + match format.kind { + ResponseFormatType::Text => (None, None), + ResponseFormatType::JsonObject => (Some("application/json".to_string()), None), + ResponseFormatType::JsonSchema => ( + Some("application/json".to_string()), + format.json_schema.clone(), + ), + } +} + +/// Build the Gemini API request body from a unified `Request`. +fn build_api_request(request: &Request) -> ApiRequest { + let (system_text, other_messages) = extract_system_prompt(&request.messages); + + let system_instruction = system_text.map(|text| SystemInstruction { + parts: vec![serde_json::json!({"text": text})], + }); + + let contents = translate_messages(&other_messages); + + let (response_mime_type, response_schema) = request + .response_format + .as_ref() + .map_or((None, None), translate_response_format); + + let generation_config = GenerationConfig { + temperature: request.temperature, + max_output_tokens: request.max_tokens, + top_p: request.top_p, + stop_sequences: request.stop_sequences.clone(), + response_mime_type, + response_schema, + }; + + let api_tools = request.tools.as_ref().map(|t| translate_tools(t)); + let tool_config = request.tool_choice.as_ref().map(translate_tool_choice); + + ApiRequest { + contents, + system_instruction, + generation_config: Some(generation_config), + tools: api_tools, + tool_config, + } +} + +/// Convert `UsageMetadata` from the Gemini API into a unified `Usage`. +fn parse_usage(metadata: Option<&UsageMetadata>) -> Usage { + metadata.map_or_else(Usage::default, |u| { + let input = u.prompt_token_count.unwrap_or(0); + let output = u.candidates_token_count.unwrap_or(0); + let total = u.total_token_count.unwrap_or(input + output); + Usage { + input_tokens: input, + output_tokens: output, + total_tokens: total, + reasoning_tokens: u.thoughts_token_count, + cache_read_tokens: u.cached_content_token_count, + ..Usage::default() + } + }) +} + +/// 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`. +async fn send_streaming_request( + request: reqwest::RequestBuilder, +) -> Result { + let http_resp = request.send().await.map_err(|e| SdkError::Network { + message: e.to_string(), + })?; + + let status = http_resp.status(); + if !status.is_success() { + let retry_after = parse_retry_after(http_resp.headers()); + let body = http_resp.text().await.map_err(|e| SdkError::Network { + 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, + )); + } + + Ok(http_resp) +} + +/// Process a stream of SSE chunks from the Gemini `streamGenerateContent` endpoint +/// and yield `StreamEvent` values. +fn process_sse_stream(http_resp: reqwest::Response, model: String) -> StreamEventStream { + Box::pin(stream::unfold( + SseStreamState::new(http_resp, model), + |mut state| async move { + // If we have buffered events, yield them first. + if let Some(event) = state.pending_events.pop_front() { + return Some((Ok(event), state)); + } + + // Read SSE lines until we get a data payload or the stream ends. + loop { + let line = match state.read_line().await { + Ok(Some(line)) => line, + Ok(None) => { + // Stream ended. Emit Finish if we haven't yet. + if !state.finished { + state.finished = true; + let event = state.build_finish_event(); + return Some((Ok(event), state)); + } + return None; + } + Err(e) => return Some((Err(e), state)), + }; + + // SSE format: lines starting with "data:" carry the payload. + let data = if let Some(stripped) = line.strip_prefix("data:") { + stripped.trim() + } else { + // Ignore non-data lines (empty lines, comments, event: lines). + continue; + }; + + // Skip empty data lines. + if data.is_empty() { + continue; + } + + // Parse the JSON chunk. + let chunk: ApiResponse = match serde_json::from_str(data) { + Ok(v) => v, + Err(e) => { + return Some(( + Err(SdkError::Stream { + message: format!("failed to parse Gemini SSE chunk: {e}"), + }), + state, + )); + } + }; + + // Extract events from this chunk. + state.process_chunk(&chunk); + + // Track usage from every chunk; the final one will have the totals. + if let Some(ref usage_meta) = chunk.usage_metadata { + state.usage = parse_usage(Some(usage_meta)); + } + + // Extract finish reason from the candidate if present. + let candidate_finish = chunk + .candidates + .as_ref() + .and_then(|c| c.first()) + .and_then(|c| c.finish_reason.clone()); + if let Some(reason) = candidate_finish { + state.finish_reason_str = Some(reason); + } + + // Yield the first buffered event if any were produced. + if let Some(event) = state.pending_events.pop_front() { + return Some((Ok(event), state)); + } + // If no events were produced from this chunk, continue reading. + } + }, + )) +} + +/// Internal state for the SSE stream processor. +struct SseStreamState { + http_resp: reqwest::Response, + model: String, + /// Buffered SSE text not yet split into complete lines. + line_buffer: String, + /// Events extracted from a chunk but not yet yielded. + pending_events: std::collections::VecDeque, + /// Whether we have emitted a `StreamStart` event. + stream_started: bool, + /// Whether we have emitted a `TextStart` event. + text_started: bool, + /// Accumulated text across all chunks. + accumulated_text: String, + /// Accumulated tool calls across all chunks. + accumulated_tool_calls: Vec, + /// The `text_id` used for `TextStart`/`TextDelta`/`TextEnd`. + text_id: String, + /// Latest usage metadata (updated per chunk; final chunk has totals). + usage: Usage, + /// The finish reason string from the candidate, if received. + finish_reason_str: Option, + /// Whether we have emitted the `Finish` event. + finished: bool, +} + +impl SseStreamState { + fn new(http_resp: reqwest::Response, model: String) -> Self { + Self { + http_resp, + model, + line_buffer: String::new(), + pending_events: std::collections::VecDeque::new(), + stream_started: false, + text_started: false, + accumulated_text: String::new(), + accumulated_tool_calls: Vec::new(), + text_id: uuid::Uuid::new_v4().to_string(), + usage: Usage::default(), + finish_reason_str: None, + finished: false, + } + } + + /// Read the next complete line from the HTTP byte stream. + /// + /// Returns `Ok(None)` when the stream is exhausted. + async fn read_line(&mut self) -> Result, SdkError> { + loop { + // Check if we already have a complete line in the buffer. + if let Some(newline_pos) = self.line_buffer.find('\n') { + let line = self.line_buffer[..newline_pos] + .trim_end_matches('\r') + .to_string(); + self.line_buffer = self.line_buffer[newline_pos + 1..].to_string(); + return Ok(Some(line)); + } + + // Read more bytes from the HTTP response. + match self.http_resp.chunk().await { + Ok(Some(bytes)) => { + let text = String::from_utf8_lossy(&bytes); + self.line_buffer.push_str(&text); + } + Ok(None) => { + // Stream ended. Return any remaining buffered content. + if self.line_buffer.is_empty() { + return Ok(None); + } + let remaining = std::mem::take(&mut self.line_buffer); + let line = remaining.trim_end_matches('\r').to_string(); + if line.is_empty() { + return Ok(None); + } + return Ok(Some(line)); + } + Err(e) => { + return Err(SdkError::Stream { + message: format!("error reading Gemini stream: {e}"), + }); + } + } + } + } + + /// Extract stream events from a parsed SSE chunk and buffer them. + fn process_chunk(&mut self, chunk: &ApiResponse) { + if !self.stream_started { + self.stream_started = true; + self.pending_events.push_back(StreamEvent::StreamStart); + } + + let parts = chunk + .candidates + .as_ref() + .and_then(|c| c.first()) + .and_then(|c| c.content.as_ref()) + .and_then(|c| c.parts.as_ref()); + + let Some(parts) = parts else { + return; + }; + + for part in parts { + if let Some(text) = part.get("text").and_then(serde_json::Value::as_str) { + if !self.text_started { + self.text_started = true; + self.pending_events.push_back(StreamEvent::TextStart { + text_id: Some(self.text_id.clone()), + }); + } + self.accumulated_text.push_str(text); + self.pending_events + .push_back(StreamEvent::text_delta(text, Some(self.text_id.clone()))); + } else if let Some(fc) = part.get("functionCall") { + let name = fc + .get("name") + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let args = fc + .get("args") + .cloned() + .unwrap_or_else(|| serde_json::Value::Object(serde_json::Map::new())); + let tool_call = ToolCall::new(uuid::Uuid::new_v4().to_string(), name, args); + + // Gemini delivers function calls as complete objects in a single chunk. + self.pending_events + .push_back(StreamEvent::ToolCallStart { + tool_call: tool_call.clone(), + }); + self.pending_events.push_back(StreamEvent::ToolCallEnd { + tool_call: tool_call.clone(), + }); + self.accumulated_tool_calls.push(tool_call); + } + } + + // If a finish reason is present on this chunk's candidate, emit TextEnd. + let has_finish_reason = chunk + .candidates + .as_ref() + .and_then(|c| c.first()) + .and_then(|c| c.finish_reason.as_ref()) + .is_some(); + + if has_finish_reason && self.text_started { + self.pending_events.push_back(StreamEvent::TextEnd { + text_id: Some(self.text_id.clone()), + }); + } + } + + /// Build the final `Finish` event from accumulated state. + fn build_finish_event(&self) -> StreamEvent { + let has_tool_calls = !self.accumulated_tool_calls.is_empty(); + let finish_reason = + map_finish_reason(self.finish_reason_str.as_deref(), has_tool_calls); + + let mut content_parts: Vec = Vec::new(); + if !self.accumulated_text.is_empty() { + content_parts.push(ContentPart::text(&self.accumulated_text)); + } + for tc in &self.accumulated_tool_calls { + content_parts.push(ContentPart::ToolCall(tc.clone())); + } + + let response = Response { + id: uuid::Uuid::new_v4().to_string(), + model: self.model.clone(), + provider: "gemini".to_string(), + message: Message { + role: Role::Assistant, + content: content_parts, + name: None, + tool_call_id: None, + }, + finish_reason: finish_reason.clone(), + usage: self.usage.clone(), + raw: None, + warnings: vec![], + rate_limit: None, + }; + + StreamEvent::finish(finish_reason, self.usage.clone(), response) + } +} + #[allow(clippy::unnecessary_literal_bound)] #[async_trait::async_trait] impl ProviderAdapter for Adapter { @@ -129,48 +688,15 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { - let (system_text, other_messages) = extract_system_prompt(&request.messages); - - let system_instruction = system_text.map(|text| SystemInstruction { - parts: vec![Part { text }], - }); - - let contents: Vec = other_messages - .iter() - .map(|msg| { - let role = match msg.role { - Role::Assistant => "model", - Role::System | Role::User | Role::Tool | Role::Developer => "user", - }; - Content { - role: role.to_string(), - parts: vec![Part { - text: msg.text(), - }], - } - }) - .collect(); - - let generation_config = GenerationConfig { - temperature: request.temperature, - max_output_tokens: request.max_tokens, - top_p: request.top_p, - stop_sequences: request.stop_sequences.clone(), - }; - - let api_request = ApiRequest { - contents, - system_instruction, - generation_config: Some(generation_config), - }; + let api_request = build_api_request(request); let url = format!( - "https://generativelanguage.googleapis.com/v1beta/models/{}:generateContent?key={}", - request.model, self.api_key + "{}/models/{}:generateContent?key={}", + self.base_url, request.model, self.api_key ); let body = send_and_read_body( - self.client.post(&url).json(&api_request), + self.client.post(&url).json(&api_request).timeout(self.request_timeout), "gemini", "status", ) @@ -187,32 +713,24 @@ impl ProviderAdapter for Adapter { .and_then(|c| c.first()) .ok_or_else(|| SdkError::Provider { kind: ProviderErrorKind::Server, - detail: Box::new(ProviderErrorDetail::new("no candidates in Gemini response", "gemini")), + detail: Box::new(ProviderErrorDetail::new( + "no candidates in Gemini response", + "gemini", + )), })?; - let content_parts: Vec = candidate - .content - .as_ref() - .and_then(|c| c.parts.as_ref()) + let raw_parts = candidate.content.as_ref().and_then(|c| c.parts.as_ref()); + + let content_parts: Vec = raw_parts .map(|parts| parts.iter().filter_map(parse_part).collect()) .unwrap_or_default(); - let finish_reason = map_finish_reason(candidate.finish_reason.as_deref()); + // Gemini has no dedicated tool_calls finish reason; infer from parts + let has_tool_calls = raw_parts.is_some_and(|p| parts_have_function_calls(p)); + let finish_reason = + map_finish_reason(candidate.finish_reason.as_deref(), has_tool_calls); - let usage = api_resp - .usage_metadata - .as_ref() - .map_or_else(Usage::default, |u| { - let input = u.prompt_token_count.unwrap_or(0); - let output = u.candidates_token_count.unwrap_or(0); - let total = u.total_token_count.unwrap_or(input + output); - Usage { - input_tokens: input, - output_tokens: output, - total_tokens: total, - ..Usage::default() - } - }); + let usage = parse_usage(api_resp.usage_metadata.as_ref()); Ok(Response { id: uuid::Uuid::new_v4().to_string(), @@ -232,9 +750,17 @@ impl ProviderAdapter for Adapter { }) } - async fn stream(&self, _request: &Request) -> Result { - Err(SdkError::Configuration { - message: "streaming not yet implemented".to_string(), - }) + async fn stream(&self, request: &Request) -> Result { + let api_request = build_api_request(request); + + let url = format!( + "{}/models/{}:streamGenerateContent?alt=sse&key={}", + self.base_url, request.model, self.api_key + ); + + let http_resp = + send_streaming_request(self.client.post(&url).json(&api_request)).await?; + + Ok(process_sse_stream(http_resp, request.model.clone())) } } diff --git a/crates/unified-llm/src/providers/mod.rs b/crates/unified-llm/src/providers/mod.rs index fe19107de..ff0421985 100644 --- a/crates/unified-llm/src/providers/mod.rs +++ b/crates/unified-llm/src/providers/mod.rs @@ -2,7 +2,9 @@ pub mod anthropic; pub mod common; pub mod gemini; pub mod openai; +pub mod openai_compatible; pub use anthropic::Adapter as AnthropicAdapter; pub use gemini::Adapter as GeminiAdapter; pub use openai::Adapter as OpenAiAdapter; +pub use openai_compatible::Adapter as OpenAiCompatibleAdapter; diff --git a/crates/unified-llm/src/providers/openai.rs b/crates/unified-llm/src/providers/openai.rs index 58360ccd1..c9400d20d 100644 --- a/crates/unified-llm/src/providers/openai.rs +++ b/crates/unified-llm/src/providers/openai.rs @@ -1,95 +1,741 @@ +use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine}; +use futures::StreamExt; + use crate::error::SdkError; use crate::provider::{ProviderAdapter, StreamEventStream}; -use crate::error::{ProviderErrorDetail, ProviderErrorKind}; -use crate::providers::common::{send_and_read_body, ApiMessage}; +use crate::providers::common::{ + parse_error_body, parse_rate_limit_headers, parse_retry_after, send_and_read_response, +}; use crate::types::{ - ContentPart, FinishReason, Message, Request, Response, Role, ToolCall, Usage, + ContentPart, FinishReason, Message, Request, Response, ResponseFormat, ResponseFormatType, + Role, StreamEvent, ToolCall, ToolChoice, ToolDefinition, Usage, }; -/// Provider adapter for the `OpenAI` Chat Completions API. +/// Provider adapter for the `OpenAI` Responses API (`/v1/responses`). +/// +/// Per spec Section 2.7, this adapter uses the Responses API (not Chat Completions) +/// to properly surface reasoning tokens, built-in tools, and server-side state. pub struct Adapter { api_key: String, + base_url: String, client: reqwest::Client, + request_timeout: std::time::Duration, } impl Adapter { #[must_use] pub fn new(api_key: impl Into) -> Self { + let timeout = crate::types::AdapterTimeout::default(); + let client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs_f64(timeout.connect)) + .build() + .unwrap_or_default(); Self { api_key: api_key.into(), - client: reqwest::Client::new(), + base_url: "https://api.openai.com/v1".to_string(), + client, + request_timeout: std::time::Duration::from_secs_f64(timeout.request), } } + + #[must_use] + pub fn with_base_url(mut self, base_url: impl Into) -> Self { + self.base_url = base_url.into(); + self + } } -// --- Request types --- +// --- Request types (Responses API format) --- #[derive(serde::Serialize)] struct ApiRequest { model: String, - messages: Vec, + input: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + instructions: Option, #[serde(skip_serializing_if = "Option::is_none")] temperature: Option, #[serde(skip_serializing_if = "Option::is_none")] - max_tokens: Option, + max_output_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] top_p: Option, #[serde(skip_serializing_if = "Option::is_none")] - stop: Option>, + tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + tool_choice: Option, + #[serde(skip_serializing_if = "Option::is_none")] + reasoning: Option, + #[serde(skip_serializing_if = "Option::is_none")] + text: Option, + #[serde(skip_serializing_if = "std::ops::Not::not")] + stream: bool, } -// --- Response types --- +// --- Response types (Responses API format) --- #[derive(serde::Deserialize)] struct ApiResponse { id: String, - model: String, - choices: Vec, + model: Option, + output: Vec, + status: Option, usage: Option, } -#[derive(serde::Deserialize)] -struct ApiChoice { - message: ApiChoiceMessage, - finish_reason: Option, -} - -#[derive(serde::Deserialize)] -struct ApiChoiceMessage { - content: Option, - tool_calls: Option>, -} - -#[derive(serde::Deserialize)] -struct ApiToolCall { - id: String, - function: ApiFunction, -} - -#[derive(serde::Deserialize)] -struct ApiFunction { - name: String, - arguments: String, -} - #[derive(serde::Deserialize)] #[allow(clippy::struct_field_names)] struct ApiUsage { - prompt_tokens: i64, - completion_tokens: i64, - total_tokens: i64, + input_tokens: i64, + output_tokens: i64, + total_tokens: Option, + output_tokens_details: Option, + input_tokens_details: Option, } -fn map_finish_reason(reason: Option<&str>) -> FinishReason { - match reason { - Some("stop") | None => FinishReason::Stop, - Some("length") => FinishReason::Length, - Some("tool_calls") => FinishReason::ToolCalls, - Some("content_filter") => FinishReason::ContentFilter, +#[derive(serde::Deserialize)] +struct OutputTokenDetails { + reasoning_tokens: Option, +} + +#[derive(serde::Deserialize)] +struct InputTokenDetails { + cached_tokens: Option, +} + +/// Map the Responses API status to a `FinishReason`. +fn map_finish_reason(status: Option<&str>, has_tool_calls: bool) -> FinishReason { + if has_tool_calls { + return FinishReason::ToolCalls; + } + match status { + Some("completed") | None => FinishReason::Stop, + Some("incomplete") => FinishReason::Length, + Some("failed") => FinishReason::Error, Some(other) => FinishReason::Other(other.to_string()), } } +/// Translate unified messages to Responses API `input` array format. +fn translate_input(messages: &[Message]) -> (Option, Vec) { + let mut instructions_parts: Vec = Vec::new(); + let mut input: Vec = Vec::new(); + + for msg in messages { + match msg.role { + Role::System | Role::Developer => { + instructions_parts.push(msg.text()); + } + Role::User => { + let content: Vec = msg + .content + .iter() + .filter_map(|part| match part { + ContentPart::Text(text) => { + Some(serde_json::json!({"type": "input_text", "text": text})) + } + ContentPart::Image(img) => { + img.url.as_ref().map_or_else( + || { + img.data.as_ref().map(|data| { + let mime = img.media_type.as_deref().unwrap_or("image/png"); + let b64 = BASE64_STANDARD.encode(data); + serde_json::json!({"type": "input_image", "image_url": format!("data:{mime};base64,{b64}")}) + }) + }, + |url| { + if crate::providers::common::is_file_path(url) { + match crate::providers::common::load_file_as_base64(url) { + Ok((b64, mime)) => Some(serde_json::json!({"type": "input_image", "image_url": format!("data:{mime};base64,{b64}")})), + Err(_) => None, + } + } else { + Some(serde_json::json!({"type": "input_image", "image_url": url})) + } + }, + ) + } + _ => None, + }) + .collect(); + if !content.is_empty() { + input.push(serde_json::json!({ + "type": "message", + "role": "user", + "content": content, + })); + } + } + Role::Assistant => { + for part in &msg.content { + match part { + ContentPart::Text(text) => { + input.push(serde_json::json!({ + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": text}], + })); + } + ContentPart::ToolCall(tc) => { + let args = tc + .raw_arguments + .as_ref() + .map_or_else(|| tc.arguments.to_string(), Clone::clone); + input.push(serde_json::json!({ + "type": "function_call", + "id": tc.id, + "call_id": tc.id, + "name": tc.name, + "arguments": args, + })); + } + _ => {} + } + } + } + Role::Tool => { + for part in &msg.content { + if let ContentPart::ToolResult(tr) = part { + let output = tr + .content + .as_str() + .map_or_else(|| tr.content.to_string(), str::to_string); + input.push(serde_json::json!({ + "type": "function_call_output", + "call_id": tr.tool_call_id, + "output": output, + })); + } + } + } + } + } + + let instructions = if instructions_parts.is_empty() { + None + } else { + Some(instructions_parts.join("\n")) + }; + + (instructions, input) +} + +/// Translate unified tool definitions to Responses API tool format. +fn translate_tools(tools: &[ToolDefinition]) -> Vec { + tools + .iter() + .map(|t| { + serde_json::json!({ + "type": "function", + "name": t.name, + "description": t.description, + "parameters": t.parameters, + }) + }) + .collect() +} + +/// Translate unified `ToolChoice` to Responses API format. +fn translate_tool_choice(choice: &ToolChoice) -> serde_json::Value { + match choice { + ToolChoice::Auto => serde_json::json!("auto"), + ToolChoice::None => serde_json::json!("none"), + ToolChoice::Required => serde_json::json!("required"), + ToolChoice::Named { tool_name } => { + serde_json::json!({"type": "function", "name": tool_name}) + } + } +} + +/// Translate unified `ResponseFormat` to Responses API `text` field. +/// +/// The Responses API uses `"text": {"format": {...}}` for structured output. +fn translate_response_format(format: &ResponseFormat) -> Option { + match format.kind { + ResponseFormatType::Text => None, + ResponseFormatType::JsonObject => { + Some(serde_json::json!({"format": {"type": "json_object"}})) + } + ResponseFormatType::JsonSchema => { + let mut schema_obj = serde_json::json!({ + "type": "json_schema", + "name": "response", + "strict": format.strict, + }); + if let Some(schema) = &format.json_schema { + schema_obj["schema"] = schema.clone(); + } + Some(serde_json::json!({"format": schema_obj})) + } + } +} + +/// Build an `ApiRequest` from a unified `Request`. +fn build_api_request(request: &Request, stream: bool) -> ApiRequest { + let (instructions, input) = translate_input(&request.messages); + let api_tools = request.tools.as_ref().map(|t| translate_tools(t)); + let tool_choice = request.tool_choice.as_ref().map(translate_tool_choice); + let reasoning = request + .reasoning_effort + .as_ref() + .map(|effort| serde_json::json!({"effort": effort})); + let text = request + .response_format + .as_ref() + .and_then(translate_response_format); + + ApiRequest { + model: request.model.clone(), + input, + instructions, + temperature: request.temperature, + max_output_tokens: request.max_tokens, + top_p: request.top_p, + tools: api_tools, + tool_choice, + reasoning, + text, + stream, + } +} + +/// Parse output items from the Responses API into content parts. +fn parse_output(output: &[serde_json::Value]) -> (Vec, bool) { + let mut parts = Vec::new(); + let mut has_tool_calls = false; + + for item in output { + let item_type = item.get("type").and_then(serde_json::Value::as_str); + match item_type { + Some("message") => { + if let Some(content) = item.get("content").and_then(|c| c.as_array()) { + for block in content { + if block.get("type").and_then(serde_json::Value::as_str) + == Some("output_text") + { + if let Some(text) = + block.get("text").and_then(serde_json::Value::as_str) + { + parts.push(ContentPart::text(text)); + } + } + } + } + } + Some("function_call") => { + has_tool_calls = true; + let id = item + .get("call_id") + .or_else(|| item.get("id")) + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let name = item + .get("name") + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let args_str = item + .get("arguments") + .and_then(serde_json::Value::as_str) + .unwrap_or("{}"); + let arguments = serde_json::from_str(args_str) + .unwrap_or_else(|_| serde_json::json!({})); + let mut tc = ToolCall::new(id, name, arguments); + tc.raw_arguments = Some(args_str.to_string()); + parts.push(ContentPart::ToolCall(tc)); + } + _ => {} + } + } + + (parts, has_tool_calls) +} + +// --- SSE streaming support --- + +/// Mutable state carried through SSE stream processing. +struct SseStreamState { + byte_stream: std::pin::Pin< + Box> + Send>, + >, + buffer: String, + model: String, + response_id: String, + response_model: String, + accumulated_text: String, + tool_calls: Vec, + usage: Usage, + finish_reason: FinishReason, + emitted_start: bool, + emitted_text_start: bool, + raw_response: Option, + rate_limit: Option, +} + +/// Extract complete SSE messages from the buffer. +/// +/// Each SSE message consists of one or more lines (`event:` and `data:` prefixed) +/// terminated by a blank line. Returns parsed (`event_type`, data) pairs. +fn extract_sse_messages(buffer: &mut String) -> Vec<(Option, String)> { + let mut messages = Vec::new(); + + while let Some(pos) = buffer.find("\n\n") { + let message_block = buffer[..pos].to_string(); + *buffer = buffer[pos + 2..].to_string(); + + let mut current_event: Option = None; + let mut current_data = String::new(); + + for line in message_block.lines() { + if let Some(stripped) = line.strip_prefix("event: ") { + current_event = Some(stripped.to_string()); + } else if let Some(stripped) = line.strip_prefix("event:") { + current_event = Some(stripped.trim().to_string()); + } else if let Some(stripped) = line.strip_prefix("data: ") { + if !current_data.is_empty() { + current_data.push('\n'); + } + current_data.push_str(stripped); + } else if let Some(stripped) = line.strip_prefix("data:") { + if !current_data.is_empty() { + current_data.push('\n'); + } + current_data.push_str(stripped.trim()); + } + } + + if !current_data.is_empty() { + messages.push((current_event, current_data)); + } + } + + messages +} + +/// Dispatch SSE messages from the buffer and return the resulting `StreamEvent`s. +fn dispatch_sse_messages( + state: &mut SseStreamState, + messages: Vec<(Option, String)>, +) -> Vec { + let mut events = Vec::new(); + for (event_type, data) in messages { + events.extend(process_sse_event(state, event_type.as_deref(), &data)); + } + events +} + +/// Process the next chunk(s) from the byte stream and return `StreamEvent`s. +async fn process_next_sse_events( + state: &mut SseStreamState, +) -> Result, SdkError> { + loop { + let messages = extract_sse_messages(&mut state.buffer); + if !messages.is_empty() { + return Ok(dispatch_sse_messages(state, messages)); + } + + match state.byte_stream.next().await { + Some(Ok(bytes)) => { + let text = String::from_utf8_lossy(&bytes); + state.buffer.push_str(&text); + } + Some(Err(e)) => { + return Err(SdkError::Stream { + message: e.to_string(), + }); + } + None => { + // Stream ended. Process any remaining data in the buffer. + if !state.buffer.is_empty() { + state.buffer.push_str("\n\n"); + let messages = extract_sse_messages(&mut state.buffer); + return Ok(dispatch_sse_messages(state, messages)); + } + return Ok(vec![]); + } + } + } +} + +/// Process a single SSE event and return the corresponding `StreamEvent`(s). +fn process_sse_event( + state: &mut SseStreamState, + event_type: Option<&str>, + data: &str, +) -> Vec { + let mut events = Vec::new(); + + if !state.emitted_start { + state.emitted_start = true; + events.push(StreamEvent::StreamStart); + } + + let json: serde_json::Value = match serde_json::from_str(data) { + Ok(v) => v, + Err(_) => return events, + }; + + // Resolve event type from the `event:` SSE line or from the JSON `type` field. + let resolved_type = event_type + .map(str::to_string) + .or_else(|| { + json.get("type") + .and_then(serde_json::Value::as_str) + .map(str::to_string) + }) + .unwrap_or_default(); + + match resolved_type.as_str() { + "response.created" => handle_response_created(state, &json), + "response.output_text.delta" => handle_text_delta(state, &json, &mut events), + "response.function_call_arguments.delta" => { + handle_tool_call_delta(state, &json, &mut events); + } + "response.output_item.done" => handle_output_item_done(state, &json, &mut events), + "response.completed" => handle_response_completed(state, &json, &mut events), + _ => {} + } + + events +} + +/// Handle `response.created` by extracting the response ID and model. +fn handle_response_created(state: &mut SseStreamState, json: &serde_json::Value) { + if let Some(id) = json + .get("response") + .and_then(|r| r.get("id")) + .and_then(serde_json::Value::as_str) + { + state.response_id = id.to_string(); + } + if let Some(model) = json + .get("response") + .and_then(|r| r.get("model")) + .and_then(serde_json::Value::as_str) + { + state.response_model = model.to_string(); + } +} + +/// Handle `response.output_text.delta` by accumulating text and emitting events. +fn handle_text_delta( + state: &mut SseStreamState, + json: &serde_json::Value, + events: &mut Vec, +) { + if let Some(delta) = json.get("delta").and_then(serde_json::Value::as_str) { + if !state.emitted_text_start { + state.emitted_text_start = true; + events.push(StreamEvent::TextStart { text_id: None }); + } + state.accumulated_text.push_str(delta); + events.push(StreamEvent::text_delta(delta, None)); + } +} + +/// Handle `response.function_call_arguments.delta` by accumulating args and emitting events. +fn handle_tool_call_delta( + state: &mut SseStreamState, + json: &serde_json::Value, + events: &mut Vec, +) { + let Some(delta) = json.get("delta").and_then(serde_json::Value::as_str) else { + return; + }; + + let call_id = json + .get("call_id") + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let item_id = json + .get("item_id") + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let lookup_id = if call_id.is_empty() { + &item_id + } else { + &call_id + }; + + let tc_index = state + .tool_calls + .iter() + .position(|tc| tc.id == *lookup_id); + + if let Some(idx) = tc_index { + if let Some(ref mut raw) = state.tool_calls[idx].raw_arguments { + raw.push_str(delta); + } + } else { + let name = json + .get("name") + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let mut tc = ToolCall::new(lookup_id, name, serde_json::json!({})); + tc.raw_arguments = Some(delta.to_string()); + state.tool_calls.push(tc.clone()); + events.push(StreamEvent::ToolCallStart { tool_call: tc }); + } + + let current_tc = state + .tool_calls + .iter() + .find(|tc| tc.id == *lookup_id) + .cloned() + .unwrap_or_else(|| ToolCall::new("", "", serde_json::json!({}))); + + events.push(StreamEvent::ToolCallDelta { + tool_call: ToolCall { + raw_arguments: Some(delta.to_string()), + ..current_tc + }, + }); +} + +/// Handle `response.output_item.done` for text and function call items. +fn handle_output_item_done( + state: &mut SseStreamState, + json: &serde_json::Value, + events: &mut Vec, +) { + let item_type = json + .get("item") + .and_then(|i| i.get("type")) + .and_then(serde_json::Value::as_str); + + match item_type { + Some("message") => { + if state.emitted_text_start { + events.push(StreamEvent::TextEnd { text_id: None }); + state.emitted_text_start = false; + } + } + Some("function_call") => { + let item = json.get("item").unwrap_or(json); + let call_id = item + .get("call_id") + .or_else(|| item.get("id")) + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let name = item + .get("name") + .and_then(serde_json::Value::as_str) + .unwrap_or("") + .to_string(); + let args_str = item + .get("arguments") + .and_then(serde_json::Value::as_str) + .unwrap_or("{}"); + let arguments = + serde_json::from_str(args_str).unwrap_or_else(|_| serde_json::json!({})); + + let mut tc = ToolCall::new(&call_id, &name, arguments); + tc.raw_arguments = Some(args_str.to_string()); + + if let Some(existing) = state.tool_calls.iter_mut().find(|t| t.id == call_id) { + existing.name.clone_from(&name); + existing.arguments = tc.arguments.clone(); + existing.raw_arguments.clone_from(&tc.raw_arguments); + } else { + state.tool_calls.push(tc.clone()); + } + + events.push(StreamEvent::ToolCallEnd { tool_call: tc }); + } + _ => {} + } +} + +/// Handle `response.completed` by extracting usage and building the final response. +fn handle_response_completed( + state: &mut SseStreamState, + json: &serde_json::Value, + events: &mut Vec, +) { + let response_data = json.get("response").unwrap_or(json); + + if let Some(usage_data) = response_data.get("usage") { + if let Ok(u) = serde_json::from_value::(usage_data.clone()) { + state.usage = Usage { + input_tokens: u.input_tokens, + output_tokens: u.output_tokens, + total_tokens: u.total_tokens.unwrap_or(u.input_tokens + u.output_tokens), + reasoning_tokens: u + .output_tokens_details + .as_ref() + .and_then(|d| d.reasoning_tokens), + cache_read_tokens: u + .input_tokens_details + .as_ref() + .and_then(|d| d.cached_tokens), + ..Usage::default() + }; + } + } + + if let Some(id) = response_data + .get("id") + .and_then(serde_json::Value::as_str) + { + state.response_id = id.to_string(); + } + if let Some(model) = response_data + .get("model") + .and_then(serde_json::Value::as_str) + { + state.response_model = model.to_string(); + } + + let status = response_data + .get("status") + .and_then(serde_json::Value::as_str); + let has_tool_calls = !state.tool_calls.is_empty(); + state.finish_reason = map_finish_reason(status, has_tool_calls); + + state.raw_response = Some(response_data.clone()); + + let mut content_parts = Vec::new(); + if !state.accumulated_text.is_empty() { + content_parts.push(ContentPart::text(&state.accumulated_text)); + } + for tc in &state.tool_calls { + content_parts.push(ContentPart::ToolCall(tc.clone())); + } + + let model = if state.response_model.is_empty() { + state.model.clone() + } else { + state.response_model.clone() + }; + + let response = Response { + id: state.response_id.clone(), + model, + provider: "openai".to_string(), + message: Message { + role: Role::Assistant, + content: content_parts, + name: None, + tool_call_id: None, + }, + finish_reason: state.finish_reason.clone(), + usage: state.usage.clone(), + raw: state.raw_response.clone(), + warnings: vec![], + rate_limit: state.rate_limit.clone(), + }; + + events.push(StreamEvent::finish( + state.finish_reason.clone(), + state.usage.clone(), + response, + )); +} + #[allow(clippy::unnecessary_literal_bound)] #[async_trait::async_trait] impl ProviderAdapter for Adapter { @@ -98,36 +744,15 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { - let api_messages: Vec = request - .messages - .iter() - .map(|msg| { - let role = match msg.role { - Role::System | Role::Developer => "system", - Role::User | Role::Tool => "user", - Role::Assistant => "assistant", - }; - ApiMessage { - role: role.to_string(), - content: msg.text(), - } - }) - .collect(); + let api_request = build_api_request(request, false); + let url = format!("{}/responses", self.base_url); - let api_request = ApiRequest { - model: request.model.clone(), - messages: api_messages, - temperature: request.temperature, - max_tokens: request.max_tokens, - top_p: request.top_p, - stop: request.stop_sequences.clone(), - }; - - let body = send_and_read_body( + let (body, headers) = send_and_read_response( self.client - .post("https://api.openai.com/v1/chat/completions") + .post(&url) .bearer_auth(&self.api_key) - .json(&api_request), + .json(&api_request) + .timeout(self.request_timeout), "openai", "type", ) @@ -138,41 +763,30 @@ impl ProviderAdapter for Adapter { message: format!("failed to parse OpenAI response: {e}"), })?; - let choice = api_resp.choices.first().ok_or_else(|| SdkError::Provider { - kind: ProviderErrorKind::Server, - detail: Box::new(ProviderErrorDetail::new("no choices in OpenAI response", "openai")), - })?; + let (content_parts, has_tool_calls) = parse_output(&api_resp.output); + let finish_reason = map_finish_reason(api_resp.status.as_deref(), has_tool_calls); - let mut content_parts = Vec::new(); - if let Some(text) = &choice.message.content { - if !text.is_empty() { - content_parts.push(ContentPart::text(text)); - } - } - if let Some(tool_calls) = &choice.message.tool_calls { - for tc in tool_calls { - let arguments = serde_json::from_str(&tc.function.arguments) - .unwrap_or_else(|_| serde_json::json!({})); - content_parts.push(ContentPart::ToolCall(ToolCall::new( - &tc.id, - &tc.function.name, - arguments, - ))); - } - } - - let finish_reason = map_finish_reason(choice.finish_reason.as_deref()); - - let usage = api_resp.usage.as_ref().map_or_else(Usage::default, |u| Usage { - input_tokens: u.prompt_tokens, - output_tokens: u.completion_tokens, - total_tokens: u.total_tokens, - ..Usage::default() - }); + let usage = api_resp + .usage + .as_ref() + .map_or_else(Usage::default, |u| Usage { + input_tokens: u.input_tokens, + output_tokens: u.output_tokens, + total_tokens: u.total_tokens.unwrap_or(u.input_tokens + u.output_tokens), + reasoning_tokens: u + .output_tokens_details + .as_ref() + .and_then(|d| d.reasoning_tokens), + cache_read_tokens: u + .input_tokens_details + .as_ref() + .and_then(|d| d.cached_tokens), + ..Usage::default() + }); Ok(Response { id: api_resp.id, - model: api_resp.model, + model: api_resp.model.unwrap_or_else(|| request.model.clone()), provider: "openai".to_string(), message: Message { role: Role::Assistant, @@ -184,13 +798,73 @@ impl ProviderAdapter for Adapter { usage, raw: serde_json::from_str(&body).ok(), warnings: vec![], - rate_limit: None, + rate_limit: parse_rate_limit_headers(&headers), }) } - async fn stream(&self, _request: &Request) -> Result { - Err(SdkError::Configuration { - message: "streaming not yet implemented".to_string(), + async fn stream(&self, request: &Request) -> Result { + let api_request = build_api_request(request, true); + let url = format!("{}/responses", self.base_url); + + let http_resp = self + .client + .post(&url) + .bearer_auth(&self.api_key) + .json(&api_request) + .send() + .await + .map_err(|e| SdkError::Network { + message: e.to_string(), + })?; + + let status = http_resp.status(); + if !status.is_success() { + let retry_after = parse_retry_after(http_resp.headers()); + let body = http_resp.text().await.map_err(|e| SdkError::Network { + message: e.to_string(), + })?; + let (msg, code, raw) = parse_error_body(&body, "type"); + return Err(crate::error::error_from_status_code( + status.as_u16(), + msg, + "openai".to_string(), + code, + raw, + retry_after, + )); + } + + let model = request.model.clone(); + let rate_limit = parse_rate_limit_headers(http_resp.headers()); + let byte_stream = http_resp.bytes_stream(); + + let state = SseStreamState { + byte_stream: Box::pin(byte_stream), + buffer: String::new(), + model, + response_id: String::new(), + response_model: String::new(), + accumulated_text: String::new(), + tool_calls: Vec::new(), + usage: Usage::default(), + finish_reason: FinishReason::Stop, + emitted_start: false, + emitted_text_start: false, + raw_response: None, + rate_limit, + }; + + let stream = futures::stream::unfold(state, |mut state| async move { + let events = process_next_sse_events(&mut state).await; + let items: Vec> = match events { + Ok(events) if events.is_empty() => return None, + Ok(events) => events.into_iter().map(Ok).collect(), + Err(e) => vec![Err(e)], + }; + Some((futures::stream::iter(items), state)) }) + .flatten(); + + Ok(Box::pin(stream)) } } diff --git a/crates/unified-llm/src/providers/openai_compatible.rs b/crates/unified-llm/src/providers/openai_compatible.rs new file mode 100644 index 000000000..8d0408592 --- /dev/null +++ b/crates/unified-llm/src/providers/openai_compatible.rs @@ -0,0 +1,1114 @@ +use futures::StreamExt; + +use crate::error::{error_from_status_code, ProviderErrorDetail, ProviderErrorKind, SdkError}; +use crate::provider::{ProviderAdapter, StreamEventStream}; +use crate::providers::common::{ + parse_error_body, parse_rate_limit_headers, parse_retry_after, send_and_read_response, +}; +use crate::types::{ + ContentPart, FinishReason, Message, Request, Response, ResponseFormat, ResponseFormatType, + Role, StreamEvent, ToolCall, ToolChoice, ToolDefinition, Usage, +}; + +/// `OpenAI`-compatible Chat Completions adapter (Section 7.10). +/// +/// Use this for third-party services (vLLM, Ollama, Together AI, Groq, etc.) +/// that implement the `OpenAI` Chat Completions API (`/v1/chat/completions`). +/// +/// Does NOT support reasoning tokens, built-in tools, or other Responses API +/// features. Use the primary `OpenAiAdapter` for `OpenAI`'s own API. +pub struct Adapter { + api_key: String, + base_url: String, + provider_name: String, + client: reqwest::Client, + request_timeout: std::time::Duration, +} + +impl Adapter { + #[must_use] + pub fn new(api_key: impl Into, base_url: impl Into) -> Self { + let timeout = crate::types::AdapterTimeout::default(); + let client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs_f64(timeout.connect)) + .build() + .unwrap_or_default(); + Self { + api_key: api_key.into(), + base_url: base_url.into(), + provider_name: "openai-compatible".to_string(), + client, + request_timeout: std::time::Duration::from_secs_f64(timeout.request), + } + } + + #[must_use] + pub fn with_name(mut self, name: impl Into) -> Self { + self.provider_name = name.into(); + self + } +} + +// --- Request types (Chat Completions format) --- + +#[derive(serde::Serialize)] +struct ApiRequest { + model: String, + messages: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + max_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + top_p: Option, + #[serde(skip_serializing_if = "Option::is_none")] + stop: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + tool_choice: Option, + #[serde(skip_serializing_if = "Option::is_none")] + response_format: Option, + #[serde(skip_serializing_if = "Option::is_none")] + stream: Option, +} + +#[derive(serde::Serialize)] +struct ChatMessage { + role: String, + #[serde(skip_serializing_if = "Option::is_none")] + content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + tool_call_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + tool_calls: Option>, +} + +#[derive(serde::Serialize)] +struct ChatToolCall { + id: String, + #[serde(rename = "type")] + kind: String, + function: ChatFunction, +} + +#[derive(serde::Serialize)] +struct ChatFunction { + name: String, + arguments: String, +} + +// --- Response types (non-streaming) --- + +#[derive(serde::Deserialize)] +struct ApiResponse { + id: String, + model: String, + choices: Vec, + usage: Option, +} + +#[derive(serde::Deserialize)] +struct ApiChoice { + message: ApiChoiceMessage, + finish_reason: Option, +} + +#[derive(serde::Deserialize)] +struct ApiChoiceMessage { + content: Option, + tool_calls: Option>, +} + +#[derive(serde::Deserialize)] +struct ApiToolCall { + id: String, + function: ApiFunction, +} + +#[derive(serde::Deserialize)] +struct ApiFunction { + name: String, + arguments: String, +} + +#[derive(serde::Deserialize)] +#[allow(clippy::struct_field_names)] +struct ApiUsage { + prompt_tokens: i64, + completion_tokens: i64, + total_tokens: i64, +} + +// --- Streaming response types --- + +#[derive(serde::Deserialize)] +struct StreamChunk { + id: Option, + model: Option, + choices: Option>, + usage: Option, +} + +#[derive(serde::Deserialize)] +struct StreamChoice { + delta: Option, + finish_reason: Option, +} + +#[derive(serde::Deserialize)] +struct StreamDelta { + content: Option, + tool_calls: Option>, +} + +#[derive(serde::Deserialize)] +struct StreamToolCall { + index: usize, + id: Option, + function: Option, +} + +#[derive(serde::Deserialize)] +struct StreamFunction { + name: Option, + arguments: Option, +} + +// --- Accumulated tool call state for streaming --- + +struct AccumulatedToolCall { + id: String, + name: String, + arguments: String, + started: bool, +} + +fn map_finish_reason(reason: Option<&str>) -> FinishReason { + match reason { + Some("stop") | None => FinishReason::Stop, + Some("length") => FinishReason::Length, + Some("tool_calls") => FinishReason::ToolCalls, + Some("content_filter") => FinishReason::ContentFilter, + Some(other) => FinishReason::Other(other.to_string()), + } +} + +fn translate_messages(messages: &[Message]) -> Vec { + messages + .iter() + .map(|msg| { + let role = match msg.role { + Role::System | Role::Developer => "system", + Role::User => "user", + Role::Assistant => "assistant", + Role::Tool => "tool", + }; + + let mut tool_calls: Vec = Vec::new(); + if msg.role == Role::Assistant { + for part in &msg.content { + if let ContentPart::ToolCall(tc) = part { + let arguments = tc + .raw_arguments + .clone() + .unwrap_or_else(|| tc.arguments.to_string()); + tool_calls.push(ChatToolCall { + id: tc.id.clone(), + kind: "function".to_string(), + function: ChatFunction { + name: tc.name.clone(), + arguments, + }, + }); + } + } + } + + let text = msg.text(); + let content = if text.is_empty() { None } else { Some(text) }; + let tool_calls = if tool_calls.is_empty() { + None + } else { + Some(tool_calls) + }; + + ChatMessage { + role: role.to_string(), + content, + tool_call_id: msg.tool_call_id.clone(), + tool_calls, + } + }) + .collect() +} + +fn translate_tools(tools: &[ToolDefinition]) -> Vec { + tools + .iter() + .map(|t| { + serde_json::json!({ + "type": "function", + "function": { + "name": t.name, + "description": t.description, + "parameters": t.parameters, + } + }) + }) + .collect() +} + +fn translate_tool_choice(choice: &ToolChoice) -> serde_json::Value { + match choice { + ToolChoice::Auto => serde_json::json!("auto"), + ToolChoice::None => serde_json::json!("none"), + ToolChoice::Required => serde_json::json!("required"), + ToolChoice::Named { tool_name } => { + serde_json::json!({"type": "function", "function": {"name": tool_name}}) + } + } +} + +/// Translate unified `ResponseFormat` to Chat Completions `response_format`. +fn translate_response_format(format: &ResponseFormat) -> serde_json::Value { + match format.kind { + ResponseFormatType::Text => serde_json::json!({"type": "text"}), + ResponseFormatType::JsonObject => serde_json::json!({"type": "json_object"}), + ResponseFormatType::JsonSchema => { + let mut json_schema = serde_json::json!({ + "name": "response", + "strict": format.strict, + }); + if let Some(schema) = &format.json_schema { + json_schema["schema"] = schema.clone(); + } + serde_json::json!({ + "type": "json_schema", + "json_schema": json_schema, + }) + } + } +} + +/// Build an `ApiRequest` from a unified `Request`. +fn build_api_request(request: &Request, stream: Option) -> ApiRequest { + let chat_messages = translate_messages(&request.messages); + let tools = request.tools.as_ref().map(|t| translate_tools(t)); + let tool_choice = request.tool_choice.as_ref().map(translate_tool_choice); + let response_format = request + .response_format + .as_ref() + .map(translate_response_format); + + ApiRequest { + model: request.model.clone(), + messages: chat_messages, + temperature: request.temperature, + max_tokens: request.max_tokens, + top_p: request.top_p, + stop: request.stop_sequences.clone(), + tools, + tool_choice, + response_format, + stream, + } +} + +#[allow(clippy::unnecessary_literal_bound)] +#[async_trait::async_trait] +impl ProviderAdapter for Adapter { + fn name(&self) -> &str { + &self.provider_name + } + + async fn complete(&self, request: &Request) -> Result { + let api_request = build_api_request(request, None); + let url = format!("{}/chat/completions", self.base_url); + + let (body, headers) = send_and_read_response( + self.client + .post(&url) + .bearer_auth(&self.api_key) + .json(&api_request) + .timeout(self.request_timeout), + &self.provider_name, + "type", + ) + .await?; + + let api_resp: ApiResponse = + serde_json::from_str(&body).map_err(|e| SdkError::Network { + message: format!("failed to parse response: {e}"), + })?; + + let choice = api_resp.choices.first().ok_or_else(|| SdkError::Provider { + kind: ProviderErrorKind::Server, + detail: Box::new(ProviderErrorDetail::new( + "no choices in response", + &self.provider_name, + )), + })?; + + let mut content_parts = Vec::new(); + if let Some(text) = &choice.message.content { + if !text.is_empty() { + content_parts.push(ContentPart::text(text)); + } + } + if let Some(tool_calls) = &choice.message.tool_calls { + for tc in tool_calls { + let arguments = serde_json::from_str(&tc.function.arguments) + .unwrap_or_else(|_| serde_json::json!({})); + let mut tool_call = ToolCall::new(&tc.id, &tc.function.name, arguments); + tool_call.raw_arguments = Some(tc.function.arguments.clone()); + content_parts.push(ContentPart::ToolCall(tool_call)); + } + } + + let finish_reason = map_finish_reason(choice.finish_reason.as_deref()); + + let usage = api_resp + .usage + .as_ref() + .map_or_else(Usage::default, |u| Usage { + input_tokens: u.prompt_tokens, + output_tokens: u.completion_tokens, + total_tokens: u.total_tokens, + ..Usage::default() + }); + + Ok(Response { + id: api_resp.id, + model: api_resp.model, + provider: self.provider_name.clone(), + message: Message { + role: Role::Assistant, + content: content_parts, + name: None, + tool_call_id: None, + }, + finish_reason, + usage, + raw: serde_json::from_str(&body).ok(), + warnings: vec![], + rate_limit: parse_rate_limit_headers(&headers), + }) + } + + async fn stream(&self, request: &Request) -> Result { + let api_request = build_api_request(request, Some(true)); + let url = format!("{}/chat/completions", self.base_url); + + let http_resp = self + .client + .post(&url) + .bearer_auth(&self.api_key) + .json(&api_request) + .send() + .await + .map_err(|e| SdkError::Network { + message: e.to_string(), + })?; + + let status = http_resp.status(); + if !status.is_success() { + let retry_after = parse_retry_after(http_resp.headers()); + let body = http_resp + .text() + .await + .map_err(|e| SdkError::Network { + message: e.to_string(), + })?; + let (msg, code, raw) = parse_error_body(&body, "type"); + return Err(error_from_status_code( + status.as_u16(), + msg, + self.provider_name.clone(), + code, + raw, + retry_after, + )); + } + + let provider_name = self.provider_name.clone(); + let model = request.model.clone(); + let rate_limit = parse_rate_limit_headers(http_resp.headers()); + + let stream = futures::stream::unfold( + StreamState::new(http_resp, provider_name, model, rate_limit), + |mut state| async move { + loop { + let line = match state.next_line().await { + Ok(Some(line)) => line, + Ok(None) => return None, + Err(e) => return Some((Err(e), state)), + }; + + let line = line.trim(); + if line.is_empty() || line.starts_with(':') { + continue; + } + + let data = match line.strip_prefix("data:") { + Some(d) => d.trim(), + None => continue, + }; + + if data == "[DONE]" { + let events = state.finish_events(); + return Some((Ok(events), state)); + } + + let chunk: StreamChunk = match serde_json::from_str(data) { + Ok(c) => c, + Err(e) => { + return Some(( + Err(SdkError::Stream { + message: format!("failed to parse SSE chunk: {e}"), + }), + state, + )); + } + }; + + if let Some(events) = state.process_chunk(&chunk) { + return Some((Ok(events), state)); + } + } + }, + ); + + // Flatten batched events into individual stream events. + let flat_stream = futures::stream::unfold( + FlattenState { + inner: Box::pin(stream), + pending: Vec::new(), + }, + |mut flatten_state| async { + loop { + if let Some(event) = flatten_state.pending.pop() { + return Some((Ok(event), flatten_state)); + } + + match flatten_state.inner.next().await { + Some(Ok(mut events)) => { + // Reverse so we can pop from the end in order. + events.reverse(); + flatten_state.pending = events; + } + Some(Err(e)) => return Some((Err(e), flatten_state)), + None => return None, + } + } + }, + ); + + Ok(Box::pin(flat_stream)) + } +} + +/// State for flattening batched events into individual stream events. +struct FlattenState { + inner: std::pin::Pin< + Box, SdkError>> + Send>, + >, + pending: Vec, +} + +/// Accumulated state while processing the SSE stream. +struct StreamState { + response: reqwest::Response, + buffer: String, + provider_name: String, + model: String, + response_id: String, + response_model: String, + accumulated_text: String, + tool_calls: Vec, + usage: Usage, + finish_reason: FinishReason, + text_started: bool, + done: bool, + rate_limit: Option, +} + +impl StreamState { + fn new( + response: reqwest::Response, + provider_name: String, + model: String, + rate_limit: Option, + ) -> Self { + Self { + response, + buffer: String::new(), + provider_name, + model, + response_id: String::new(), + response_model: String::new(), + accumulated_text: String::new(), + tool_calls: Vec::new(), + usage: Usage::default(), + finish_reason: FinishReason::Stop, + text_started: false, + done: false, + rate_limit, + } + } + + /// Read the next complete line from the SSE byte stream. + async fn next_line(&mut self) -> Result, SdkError> { + if self.done { + return Ok(None); + } + + loop { + if let Some(newline_pos) = self.buffer.find('\n') { + let line = self.buffer[..newline_pos].to_string(); + self.buffer = self.buffer[newline_pos + 1..].to_string(); + return Ok(Some(line)); + } + + match self.response.chunk().await { + Ok(Some(bytes)) => { + let text = String::from_utf8_lossy(&bytes); + self.buffer.push_str(&text); + } + Ok(None) => { + self.done = true; + if self.buffer.is_empty() { + return Ok(None); + } + let remaining = std::mem::take(&mut self.buffer); + return Ok(Some(remaining)); + } + Err(e) => { + return Err(SdkError::Stream { + message: e.to_string(), + }); + } + } + } + } + + /// Process a parsed SSE chunk and return events to emit, if any. + fn process_chunk(&mut self, chunk: &StreamChunk) -> Option> { + // Capture response metadata from the first chunk. + if let Some(id) = &chunk.id { + if self.response_id.is_empty() { + self.response_id.clone_from(id); + } + } + if let Some(model) = &chunk.model { + if self.response_model.is_empty() { + self.response_model.clone_from(model); + } + } + + // Capture usage if present (often in a dedicated chunk). + if let Some(usage) = &chunk.usage { + self.usage = Usage { + input_tokens: usage.prompt_tokens, + output_tokens: usage.completion_tokens, + total_tokens: usage.total_tokens, + ..Usage::default() + }; + } + + let choices = chunk.choices.as_ref()?; + let choice = choices.first()?; + + let mut events = Vec::new(); + + // Check for finish_reason. + if let Some(reason) = &choice.finish_reason { + self.finish_reason = map_finish_reason(Some(reason.as_str())); + } + + let delta = choice.delta.as_ref()?; + + // Handle text content delta. + if let Some(content) = &delta.content { + if !content.is_empty() { + if !self.text_started { + self.text_started = true; + events.push(StreamEvent::TextStart { text_id: None }); + } + self.accumulated_text.push_str(content); + events.push(StreamEvent::text_delta(content, None)); + } + } + + // Handle tool call deltas. + if let Some(tool_calls) = &delta.tool_calls { + for tc in tool_calls { + let index = tc.index; + + // Grow the accumulated tool calls vector if needed. + while self.tool_calls.len() <= index { + self.tool_calls.push(AccumulatedToolCall { + id: String::new(), + name: String::new(), + arguments: String::new(), + started: false, + }); + } + + let accumulated = &mut self.tool_calls[index]; + + // First chunk for this tool call carries id and name. + if let Some(id) = &tc.id { + accumulated.id.clone_from(id); + } + if let Some(func) = &tc.function { + if let Some(name) = &func.name { + accumulated.name.clone_from(name); + } + if let Some(args) = &func.arguments { + accumulated.arguments.push_str(args); + } + } + + let partial_tool_call = ToolCall::new( + &accumulated.id, + &accumulated.name, + serde_json::json!(null), + ); + + if accumulated.started { + events.push(StreamEvent::ToolCallDelta { + tool_call: partial_tool_call, + }); + } else { + accumulated.started = true; + events.push(StreamEvent::ToolCallStart { + tool_call: partial_tool_call, + }); + } + } + } + + if events.is_empty() { + None + } else { + Some(events) + } + } + + /// Generate the final events when `[DONE]` is received. + fn finish_events(&mut self) -> Vec { + let mut events = Vec::new(); + + // End text segment if it was started. + if self.text_started { + events.push(StreamEvent::TextEnd { text_id: None }); + } + + // End all tool calls with complete data. + let mut content_parts = Vec::new(); + + if !self.accumulated_text.is_empty() { + content_parts.push(ContentPart::text(&self.accumulated_text)); + } + + for accumulated in &self.tool_calls { + let arguments = serde_json::from_str(&accumulated.arguments) + .unwrap_or_else(|_| serde_json::json!({})); + let mut tool_call = + ToolCall::new(&accumulated.id, &accumulated.name, arguments); + tool_call.raw_arguments = Some(accumulated.arguments.clone()); + + events.push(StreamEvent::ToolCallEnd { + tool_call: tool_call.clone(), + }); + content_parts.push(ContentPart::ToolCall(tool_call)); + } + + // Infer finish reason from tool calls if not explicitly set. + if !self.tool_calls.is_empty() && self.finish_reason == FinishReason::Stop { + self.finish_reason = FinishReason::ToolCalls; + } + + let response_model = if self.response_model.is_empty() { + self.model.clone() + } else { + self.response_model.clone() + }; + + let response = Response { + id: self.response_id.clone(), + model: response_model, + provider: self.provider_name.clone(), + message: Message { + role: Role::Assistant, + content: content_parts, + name: None, + tool_call_id: None, + }, + finish_reason: self.finish_reason.clone(), + usage: self.usage.clone(), + raw: None, + warnings: vec![], + rate_limit: self.rate_limit.clone(), + }; + + events.push(StreamEvent::finish( + self.finish_reason.clone(), + self.usage.clone(), + response, + )); + + events + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn stream_chunk_text_delta_parsing() { + let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[{"delta":{"content":"Hello"},"finish_reason":null}]}"#; + let chunk: StreamChunk = serde_json::from_str(json).unwrap(); + assert_eq!(chunk.id.as_deref(), Some("chatcmpl-1")); + assert_eq!(chunk.model.as_deref(), Some("gpt-4")); + let choices = chunk.choices.unwrap(); + assert_eq!(choices.len(), 1); + let delta = choices[0].delta.as_ref().unwrap(); + assert_eq!(delta.content.as_deref(), Some("Hello")); + assert!(choices[0].finish_reason.is_none()); + } + + #[test] + fn stream_chunk_tool_call_parsing() { + let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","function":{"name":"get_weather","arguments":"{\"ci"}}]},"finish_reason":null}]}"#; + let chunk: StreamChunk = serde_json::from_str(json).unwrap(); + let choices = chunk.choices.unwrap(); + let delta = choices[0].delta.as_ref().unwrap(); + let tc = &delta.tool_calls.as_ref().unwrap()[0]; + assert_eq!(tc.index, 0); + assert_eq!(tc.id.as_deref(), Some("call_1")); + let func = tc.function.as_ref().unwrap(); + assert_eq!(func.name.as_deref(), Some("get_weather")); + assert_eq!(func.arguments.as_deref(), Some("{\"ci")); + } + + #[test] + fn stream_chunk_usage_parsing() { + let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[],"usage":{"prompt_tokens":10,"completion_tokens":20,"total_tokens":30}}"#; + let chunk: StreamChunk = serde_json::from_str(json).unwrap(); + let usage = chunk.usage.unwrap(); + assert_eq!(usage.prompt_tokens, 10); + assert_eq!(usage.completion_tokens, 20); + assert_eq!(usage.total_tokens, 30); + } + + #[test] + fn stream_chunk_finish_reason_parsing() { + let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[{"delta":{},"finish_reason":"stop"}]}"#; + let chunk: StreamChunk = serde_json::from_str(json).unwrap(); + let choices = chunk.choices.unwrap(); + assert_eq!(choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[test] + fn stream_state_process_text_chunks() { + let http_resp = reqwest::Response::from( + http::Response::builder() + .status(200) + .body("") + .unwrap(), + ); + let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None); + + // First text chunk should emit TextStart + TextDelta. + let chunk1: StreamChunk = serde_json::from_str( + r#"{"id":"c1","model":"m1","choices":[{"delta":{"content":"Hello"},"finish_reason":null}]}"#, + ).unwrap(); + let events1 = state.process_chunk(&chunk1).unwrap(); + assert_eq!(events1.len(), 2); + assert!(matches!(events1[0], StreamEvent::TextStart { .. })); + assert!(matches!(events1[1], StreamEvent::TextDelta { .. })); + + // Second text chunk should emit only TextDelta (no second TextStart). + let chunk2: StreamChunk = serde_json::from_str( + r#"{"id":"c1","model":"m1","choices":[{"delta":{"content":" world"},"finish_reason":null}]}"#, + ).unwrap(); + let events2 = state.process_chunk(&chunk2).unwrap(); + assert_eq!(events2.len(), 1); + assert!(matches!(events2[0], StreamEvent::TextDelta { .. })); + + assert_eq!(state.accumulated_text, "Hello world"); + } + + #[test] + fn stream_state_process_tool_call_chunks() { + let http_resp = reqwest::Response::from( + http::Response::builder() + .status(200) + .body("") + .unwrap(), + ); + let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None); + + // First tool call chunk (has id and name) -> ToolCallStart. + let chunk1: StreamChunk = serde_json::from_str( + r#"{"id":"c1","model":"m1","choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","function":{"name":"fn1","arguments":"{\"k"}}]},"finish_reason":null}]}"#, + ).unwrap(); + let events1 = state.process_chunk(&chunk1).unwrap(); + assert_eq!(events1.len(), 1); + assert!(matches!(events1[0], StreamEvent::ToolCallStart { .. })); + + // Subsequent chunk (more arguments) -> ToolCallDelta. + let chunk2: StreamChunk = serde_json::from_str( + r#"{"id":"c1","model":"m1","choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"ey\"}"}}]},"finish_reason":null}]}"#, + ).unwrap(); + let events2 = state.process_chunk(&chunk2).unwrap(); + assert_eq!(events2.len(), 1); + assert!(matches!(events2[0], StreamEvent::ToolCallDelta { .. })); + + assert_eq!(state.tool_calls[0].arguments, r#"{"key"}"#); + } + + #[test] + fn stream_state_finish_events_text_only() { + let http_resp = reqwest::Response::from( + http::Response::builder() + .status(200) + .body("") + .unwrap(), + ); + let mut state = StreamState::new(http_resp, "test-provider".into(), "test-model".into(), None); + state.response_id = "resp-1".into(); + state.response_model = "gpt-4".into(); + state.accumulated_text = "Hello world".into(); + state.text_started = true; + state.usage = Usage { + input_tokens: 5, + output_tokens: 10, + total_tokens: 15, + ..Usage::default() + }; + + let events = state.finish_events(); + // TextEnd + Finish + assert_eq!(events.len(), 2); + assert!(matches!(events[0], StreamEvent::TextEnd { .. })); + match &events[1] { + StreamEvent::Finish { + finish_reason, + usage, + response, + } => { + assert_eq!(*finish_reason, FinishReason::Stop); + assert_eq!(usage.input_tokens, 5); + assert_eq!(usage.output_tokens, 10); + assert_eq!(response.text(), "Hello world"); + assert_eq!(response.id, "resp-1"); + assert_eq!(response.model, "gpt-4"); + assert_eq!(response.provider, "test-provider"); + } + other => panic!("Expected Finish, got {other:?}"), + } + } + + #[test] + fn stream_state_finish_events_with_tool_calls() { + let http_resp = reqwest::Response::from( + http::Response::builder() + .status(200) + .body("") + .unwrap(), + ); + let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None); + state.response_id = "resp-1".into(); + state.tool_calls.push(AccumulatedToolCall { + id: "call_1".into(), + name: "get_weather".into(), + arguments: r#"{"city":"SF"}"#.into(), + started: true, + }); + + let events = state.finish_events(); + // ToolCallEnd + Finish (no TextEnd since text_started is false) + assert_eq!(events.len(), 2); + match &events[0] { + StreamEvent::ToolCallEnd { tool_call } => { + assert_eq!(tool_call.id, "call_1"); + assert_eq!(tool_call.name, "get_weather"); + assert_eq!(tool_call.raw_arguments.as_deref(), Some(r#"{"city":"SF"}"#)); + } + other => panic!("Expected ToolCallEnd, got {other:?}"), + } + match &events[1] { + StreamEvent::Finish { + finish_reason, + response, + .. + } => { + assert_eq!(*finish_reason, FinishReason::ToolCalls); + let calls = response.tool_calls(); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "get_weather"); + } + other => panic!("Expected Finish, got {other:?}"), + } + } + + #[test] + fn stream_state_uses_request_model_as_fallback() { + let http_resp = reqwest::Response::from( + http::Response::builder() + .status(200) + .body("") + .unwrap(), + ); + let mut state = StreamState::new(http_resp, "test".into(), "fallback-model".into(), None); + // response_model is empty, so finish_events should use the request model. + let events = state.finish_events(); + match &events[0] { + StreamEvent::Finish { response, .. } => { + assert_eq!(response.model, "fallback-model"); + } + other => panic!("Expected Finish, got {other:?}"), + } + } + + #[test] + fn api_request_stream_field_serialization() { + let req = ApiRequest { + model: "test".into(), + messages: vec![], + temperature: None, + max_tokens: None, + top_p: None, + stop: None, + tools: None, + tool_choice: None, + response_format: None, + stream: Some(true), + }; + let json = serde_json::to_value(&req).unwrap(); + assert_eq!(json["stream"], true); + + // When stream is None, it should be omitted. + let req_no_stream = ApiRequest { + model: "test".into(), + messages: vec![], + temperature: None, + max_tokens: None, + top_p: None, + stop: None, + tools: None, + tool_choice: None, + response_format: None, + stream: None, + }; + let json_no_stream = serde_json::to_value(&req_no_stream).unwrap(); + assert!(json_no_stream.get("stream").is_none()); + } + + #[test] + fn translate_assistant_message_with_tool_calls_only() { + let msg = Message { + role: Role::Assistant, + content: vec![ContentPart::ToolCall(ToolCall::new( + "call_1", + "get_weather", + serde_json::json!({"city": "SF"}), + ))], + name: None, + tool_call_id: None, + }; + let translated = translate_messages(&[msg]); + assert_eq!(translated.len(), 1); + assert_eq!(translated[0].role, "assistant"); + assert!(translated[0].content.is_none()); + let tool_calls = translated[0].tool_calls.as_ref().unwrap(); + assert_eq!(tool_calls.len(), 1); + assert_eq!(tool_calls[0].id, "call_1"); + assert_eq!(tool_calls[0].kind, "function"); + assert_eq!(tool_calls[0].function.name, "get_weather"); + assert_eq!(tool_calls[0].function.arguments, r#"{"city":"SF"}"#); + } + + #[test] + fn translate_assistant_message_with_text_and_tool_calls() { + let msg = Message { + role: Role::Assistant, + content: vec![ + ContentPart::text("Let me check the weather"), + ContentPart::ToolCall(ToolCall::new( + "call_2", + "get_weather", + serde_json::json!({"city": "NYC"}), + )), + ], + name: None, + tool_call_id: None, + }; + let translated = translate_messages(&[msg]); + assert_eq!(translated[0].content.as_deref(), Some("Let me check the weather")); + let tool_calls = translated[0].tool_calls.as_ref().unwrap(); + assert_eq!(tool_calls.len(), 1); + assert_eq!(tool_calls[0].function.name, "get_weather"); + } + + #[test] + fn translate_assistant_message_with_raw_arguments() { + let mut tc = ToolCall::new("call_3", "search", serde_json::json!({"q": "rust"})); + tc.raw_arguments = Some(r#"{"q": "rust"}"#.to_string()); + let msg = Message { + role: Role::Assistant, + content: vec![ContentPart::ToolCall(tc)], + name: None, + tool_call_id: None, + }; + let translated = translate_messages(&[msg]); + let tool_calls = translated[0].tool_calls.as_ref().unwrap(); + // Should prefer raw_arguments over serializing arguments + assert_eq!(tool_calls[0].function.arguments, r#"{"q": "rust"}"#); + } + + #[test] + fn translate_tool_message_has_tool_call_id() { + let msg = Message::tool_result("call_1", "72F and sunny", false); + let translated = translate_messages(&[msg]); + assert_eq!(translated[0].role, "tool"); + assert_eq!(translated[0].tool_call_id.as_deref(), Some("call_1")); + assert!(translated[0].tool_calls.is_none()); + } + + #[test] + fn translate_user_message_has_no_tool_calls() { + let msg = Message::user("Hello"); + let translated = translate_messages(&[msg]); + assert_eq!(translated[0].role, "user"); + assert_eq!(translated[0].content.as_deref(), Some("Hello")); + assert!(translated[0].tool_calls.is_none()); + } + + #[test] + fn assistant_tool_calls_serialize_correctly() { + let msg = Message { + role: Role::Assistant, + content: vec![ContentPart::ToolCall(ToolCall::new( + "call_1", + "get_weather", + serde_json::json!({"city": "SF"}), + ))], + name: None, + tool_call_id: None, + }; + let translated = translate_messages(&[msg]); + let json = serde_json::to_value(&translated[0]).unwrap(); + assert!(json.get("content").is_none()); + assert!(json.get("tool_call_id").is_none()); + let tool_calls = json["tool_calls"].as_array().unwrap(); + assert_eq!(tool_calls.len(), 1); + assert_eq!(tool_calls[0]["type"], "function"); + assert_eq!(tool_calls[0]["id"], "call_1"); + assert_eq!(tool_calls[0]["function"]["name"], "get_weather"); + } +} diff --git a/crates/unified-llm/src/retry.rs b/crates/unified-llm/src/retry.rs index 2288b7d38..33783567f 100644 --- a/crates/unified-llm/src/retry.rs +++ b/crates/unified-llm/src/retry.rs @@ -27,16 +27,21 @@ where } // Check Retry-After - if let Some(retry_after) = err.retry_after() { + let delay = if let Some(retry_after) = err.retry_after() { if retry_after > policy.max_delay { return Err(err); } - tokio::time::sleep(std::time::Duration::from_secs_f64(retry_after)).await; + retry_after } else { - let delay = policy.delay_for_attempt(attempt); - tokio::time::sleep(std::time::Duration::from_secs_f64(delay)).await; + policy.delay_for_attempt(attempt) + }; + + if let Some(ref on_retry) = policy.on_retry { + on_retry(&err, attempt, delay); } + tokio::time::sleep(std::time::Duration::from_secs_f64(delay)).await; + attempt += 1; } } @@ -46,6 +51,7 @@ where #[cfg(test)] mod tests { use super::*; + use crate::types::RetryPolicy; use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::Arc; @@ -243,4 +249,47 @@ mod tests { // Should have waited ~0.01s, not ~10s assert!(elapsed.as_secs_f64() < 1.0); } + + #[tokio::test] + async fn retry_invokes_on_retry_callback() { + let retry_attempts = Arc::new(AtomicU32::new(0)); + let retry_attempts_clone = retry_attempts.clone(); + + let policy = RetryPolicy { + max_retries: 2, + base_delay: 0.001, + jitter: false, + on_retry: Some(Arc::new(move |_err, _attempt, _delay| { + retry_attempts_clone.fetch_add(1, Ordering::SeqCst); + })), + ..Default::default() + }; + + let call_count = Arc::new(AtomicU32::new(0)); + let cc = call_count.clone(); + + let result = retry(&policy, || { + let cc = cc.clone(); + async move { + let count = cc.fetch_add(1, Ordering::SeqCst); + if count < 2 { + Err(SdkError::Provider { + kind: crate::error::ProviderErrorKind::Server, + detail: Box::new(crate::error::ProviderErrorDetail { + status_code: Some(500), + ..crate::error::ProviderErrorDetail::new("error", "test") + }), + }) + } else { + Ok(99) + } + } + }) + .await; + + assert_eq!(result.unwrap(), 99); + assert_eq!(call_count.load(Ordering::SeqCst), 3); + // on_retry should have been called twice (before each retry) + assert_eq!(retry_attempts.load(Ordering::SeqCst), 2); + } } diff --git a/crates/unified-llm/src/tools.rs b/crates/unified-llm/src/tools.rs index ba47ef831..590977b3e 100644 --- a/crates/unified-llm/src/tools.rs +++ b/crates/unified-llm/src/tools.rs @@ -20,8 +20,15 @@ pub struct Tool { impl Tool { /// Create a passive tool (no execute handler). - #[must_use] + /// + /// # Panics + /// + /// Panics if the tool name is invalid (see [`validate_tool_name`]). + #[must_use] pub fn passive(name: &str, description: &str, parameters: serde_json::Value) -> Self { + if let Err(e) = validate_tool_name(name) { + panic!("Invalid tool name: {e}"); + } Self { definition: ToolDefinition { name: name.to_string(), @@ -33,6 +40,10 @@ impl Tool { } /// Create an active tool with an execute handler. + /// + /// # Panics + /// + /// Panics if the tool name is invalid (see [`validate_tool_name`]). pub fn active( name: &str, description: &str, @@ -43,6 +54,9 @@ impl Tool { F: Fn(serde_json::Value) -> Fut + Send + Sync + 'static, Fut: Future> + Send + 'static, { + if let Err(e) = validate_tool_name(name) { + panic!("Invalid tool name: {e}"); + } Self { definition: ToolDefinition { name: name.to_string(), @@ -121,11 +135,15 @@ pub async fn execute_all_tools( 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, }, } } @@ -135,6 +153,8 @@ pub async fn execute_all_tools( "Unknown tool: {call_name}" )), is_error: true, + image_data: None, + image_media_type: None, }, } } @@ -340,4 +360,25 @@ mod tests { assert!(!results[0].is_error); assert!(results[1].is_error); } + + #[test] + #[should_panic(expected = "Invalid tool name")] + fn passive_tool_panics_on_invalid_name() { + let _ = Tool::passive( + "1invalid", + "bad name", + serde_json::json!({"type": "object"}), + ); + } + + #[test] + #[should_panic(expected = "Invalid tool name")] + fn active_tool_panics_on_invalid_name() { + Tool::active( + "my-tool", + "bad name", + serde_json::json!({"type": "object"}), + |_args| async { Ok(serde_json::json!("result")) }, + ); + } } diff --git a/crates/unified-llm/src/types.rs b/crates/unified-llm/src/types.rs index 9c7a388e2..26d35905e 100644 --- a/crates/unified-llm/src/types.rs +++ b/crates/unified-llm/src/types.rs @@ -1,5 +1,7 @@ +use crate::error::SdkError; use serde::{Deserialize, Serialize}; use std::collections::HashMap; +use std::sync::Arc; // --- 3.2 Role --- @@ -75,6 +77,10 @@ pub struct ToolResult { pub tool_call_id: String, pub content: serde_json::Value, pub is_error: bool, + #[serde(skip_serializing_if = "Option::is_none")] + pub image_data: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub image_media_type: Option, } // --- 3.3 ContentPart --- @@ -149,6 +155,8 @@ impl Message { tool_call_id: id.clone(), content: serde_json::Value::String(content.into()), is_error, + image_data: None, + image_media_type: None, })], name: None, tool_call_id: Some(id), @@ -503,13 +511,31 @@ impl Default for AdapterTimeout { // --- 6.6 RetryPolicy --- -#[derive(Debug, Clone)] +/// Callback invoked before each retry attempt with (error, attempt, delay in seconds). +pub type OnRetryCallback = Arc; + +#[derive(Clone)] pub struct RetryPolicy { pub max_retries: u32, pub base_delay: f64, pub max_delay: f64, pub backoff_multiplier: f64, pub jitter: bool, + /// Called before each retry with (error, attempt number, delay in seconds). + pub on_retry: Option, +} + +impl std::fmt::Debug for RetryPolicy { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RetryPolicy") + .field("max_retries", &self.max_retries) + .field("base_delay", &self.base_delay) + .field("max_delay", &self.max_delay) + .field("backoff_multiplier", &self.backoff_multiplier) + .field("jitter", &self.jitter) + .field("on_retry", &self.on_retry.as_ref().map(|_| "...")) + .finish() + } } impl Default for RetryPolicy { @@ -520,6 +546,7 @@ impl Default for RetryPolicy { max_delay: 60.0, backoff_multiplier: 2.0, jitter: true, + on_retry: None, } } } @@ -540,6 +567,22 @@ impl RetryPolicy { } } +// --- 4.6 ObjectStreamEvent --- + +/// Events yielded by `stream_object()` for streaming structured output. +#[derive(Debug, Clone)] +pub enum ObjectStreamEvent { + /// A new partial parse of the accumulated JSON text. + Partial { object: serde_json::Value }, + /// A raw stream event from the underlying provider stream. + Delta { event: StreamEvent }, + /// The stream completed with a fully parsed object and response. + Complete { + object: serde_json::Value, + response: Box, + }, +} + // --- 4.3 GenerateResult / StepResult --- #[derive(Debug, Clone)] @@ -905,6 +948,7 @@ mod tests { max_delay: 60.0, backoff_multiplier: 2.0, jitter: false, + ..Default::default() }; assert!((policy.delay_for_attempt(0) - 1.0).abs() < f64::EPSILON); assert!((policy.delay_for_attempt(1) - 2.0).abs() < f64::EPSILON); @@ -920,6 +964,7 @@ mod tests { max_delay: 5.0, backoff_multiplier: 2.0, jitter: false, + ..Default::default() }; assert!((policy.delay_for_attempt(5) - 5.0).abs() < f64::EPSILON); } @@ -932,6 +977,7 @@ mod tests { max_delay: 60.0, backoff_multiplier: 2.0, jitter: true, + ..Default::default() }; let delay = policy.delay_for_attempt(0); // base * 0.5 to base * 1.5 => 0.5 to 1.5 @@ -972,6 +1018,19 @@ mod tests { assert_eq!(deserialized, tc); } + #[test] + fn tool_result_with_image_data() { + let result = ToolResult { + tool_call_id: "call_1".into(), + content: serde_json::json!("screenshot taken"), + is_error: false, + image_data: Some(vec![0x89, 0x50, 0x4E, 0x47]), + image_media_type: Some("image/png".into()), + }; + assert!(result.image_data.is_some()); + assert_eq!(result.image_media_type.as_deref(), Some("image/png")); + } + #[test] fn tool_call_new_constructor() { let tc = ToolCall::new("c1", "test", serde_json::json!({})); diff --git a/docs/specs/unified-llm-spec.md b/docs/specs/unified-llm-spec.md index 488809cee..5676e8870 100644 --- a/docs/specs/unified-llm-spec.md +++ b/docs/specs/unified-llm-spec.md @@ -510,10 +510,10 @@ RECORD DocumentData: ``` RECORD ToolCallData: - id : String -- unique identifier for this call (provider-assigned) - name : String -- tool name - arguments : Dict | String -- parsed JSON arguments or raw argument string - type : String -- "function" (default) or "custom" + id : String -- unique identifier for this call (provider-assigned) + name : String -- tool name + arguments : Dict -- parsed JSON arguments + raw_arguments : String | None -- raw argument string before parsing (for debugging) ``` The `id` field is assigned by the provider and is required for linking tool results back to calls. For providers that do not assign unique IDs (e.g., Gemini), the adapter must generate synthetic unique IDs (e.g., `"call_" + random_uuid()`) and maintain a mapping to the function name. @@ -617,6 +617,8 @@ RECORD FinishReason: raw : String | None -- the provider's native finish reason string ``` +**Note for statically-typed languages:** An enum (discriminated union) with variants `Stop`, `Length`, `ToolCalls`, `ContentFilter`, `Error`, and `Other(String)` is an acceptable representation. The provider's raw finish reason string is available in `Response.raw` (the full provider response). A separate `raw` field on FinishReason itself is optional in typed implementations. + Unified reason values: | Value | Meaning | @@ -632,10 +634,14 @@ Provider finish reason mapping: | Provider | Provider Value | Unified Value | |-----------|-------------------|------------------| -| OpenAI | stop | stop | -| OpenAI | length | length | -| OpenAI | tool_calls | tool_calls | -| OpenAI | content_filter | content_filter | +| OpenAI (Responses API) | completed | stop | +| OpenAI (Responses API) | incomplete | length | +| OpenAI (Responses API) | failed | error | +| OpenAI (Responses API) | (has function_call items) | tool_calls | +| OpenAI (Chat Completions) | stop | stop | +| OpenAI (Chat Completions) | length | length | +| OpenAI (Chat Completions) | tool_calls | tool_calls | +| OpenAI (Chat Completions) | content_filter | content_filter | | Anthropic | end_turn | stop | | Anthropic | stop_sequence | stop | | Anthropic | max_tokens | length | @@ -646,7 +652,7 @@ Provider finish reason mapping: | Gemini | RECITATION | content_filter | | Gemini | (has tool calls) | tool_calls | -Note: Gemini does not have a dedicated "tool_calls" finish reason. The adapter infers it from the presence of `functionCall` parts in the response. +Note: Gemini does not have a dedicated "tool_calls" finish reason. The adapter infers it from the presence of `functionCall` parts in the response. Similarly, the OpenAI Responses API uses `completed`/`incomplete`/`failed` status rather than Chat Completions-style finish reasons, and tool calls are inferred from the presence of `function_call` output items. ### 3.9 Usage