diff --git a/Cargo.lock b/Cargo.lock index d5d3ac56c..2a702e07f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1401,6 +1401,7 @@ dependencies = [ "thiserror", "tokio", "tokio-stream", + "tokio-util", "uuid", ] diff --git a/Cargo.toml b/Cargo.toml index ff9923ea0..ed773b712 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -29,3 +29,4 @@ tokio-stream = "0.1" async-trait = "0.1" base64 = "0.22" bytes = "1" +tokio-util = "0.7" diff --git a/crates/unified-llm/Cargo.toml b/crates/unified-llm/Cargo.toml index 9a208f539..191b01ae4 100644 --- a/crates/unified-llm/Cargo.toml +++ b/crates/unified-llm/Cargo.toml @@ -23,6 +23,7 @@ async-trait.workspace = true reqwest.workspace = true base64.workspace = true bytes.workspace = true +tokio-util.workspace = true [dev-dependencies] dotenvy.workspace = true diff --git a/crates/unified-llm/src/generate.rs b/crates/unified-llm/src/generate.rs index 65c29fcbd..7682f9a2b 100644 --- a/crates/unified-llm/src/generate.rs +++ b/crates/unified-llm/src/generate.rs @@ -8,10 +8,12 @@ use crate::types::{ ResponseFormatType, RetryPolicy, StepResult, StreamEvent, TimeoutConfig, ToolCall, ToolChoice, ToolDefinition, Usage, }; -use futures::StreamExt; +use futures::{Stream, StreamExt}; use std::pin::Pin; use std::sync::Arc; +use std::task::{Context, Poll}; use tokio::sync::OnceCell; +use tokio_util::sync::CancellationToken; /// Module-level default client (Section 2.5). static DEFAULT_CLIENT: OnceCell> = OnceCell::const_new(); @@ -95,6 +97,7 @@ fn build_generate_result(steps: Vec, total_usage: Usage) -> Generate /// # Panics /// /// Panics if a tool's `execute` handler is `None` when matched during tool execution. +#[allow(clippy::too_many_lines)] pub async fn generate(params: GenerateParams) -> Result { let client = params.client.clone().unwrap_or_else(get_default_client); let retry_policy = RetryPolicy { @@ -112,12 +115,22 @@ pub async fn generate(params: GenerateParams) -> Result = Vec::new(); let mut total_usage = Usage::default(); let mut round = 0u32; loop { + if let Some(ref token) = abort_signal { + if token.is_cancelled() { + return Err(SdkError::Abort { + message: "Generation aborted by cancellation token".into(), + }); + } + } + let request = build_request(¶ms, &messages, tool_definitions.as_deref()); let client_ref = client.clone(); @@ -148,6 +161,7 @@ pub async fn generate(params: GenerateParams) -> Result 0 { let tools = params.tools.as_ref().expect("checked above"); if tools.iter().any(|t| t.is_active()) { @@ -175,6 +189,14 @@ pub async fn generate(params: GenerateParams) -> Result, pub client: Option>, + /// Cancellation token to abort generation (Section 4.8). + pub abort_signal: Option, /// Custom stop condition checked after each tool round (Section 4.3). pub stop_when: Option, } @@ -254,6 +278,7 @@ impl GenerateParams { max_retries: 2, timeout: None, client: None, + abort_signal: None, stop_when: None, } } @@ -366,6 +391,12 @@ impl GenerateParams { self } + #[must_use] + pub fn abort_signal(mut self, token: CancellationToken) -> Self { + self.abort_signal = Some(token); + self + } + /// Set a custom stop condition for the tool loop (Section 4.3). /// /// The callback receives the accumulated steps so far and returns `true` @@ -453,15 +484,237 @@ impl Default for StreamAccumulator { } } +/// Wraps a streaming response with an internal `StreamAccumulator` and convenience methods. +/// +/// Implements `Stream>` so it can be used +/// as a drop-in replacement for `StreamEventStream`. Also supports multi-step +/// tool loops when active tools are provided. +pub struct StreamResult { + inner: StreamEventStream, + accumulator: StreamAccumulator, +} + +impl StreamResult { + fn new(inner: StreamEventStream) -> Self { + Self { + inner, + accumulator: StreamAccumulator::new(), + } + } + + /// Returns the accumulated response after the stream has ended. + #[must_use] + pub const fn response(&self) -> Option<&Response> { + self.accumulator.response() + } + + /// Returns the current partially accumulated response state. + #[must_use] + pub const fn partial_response(&self) -> Option<&Response> { + self.accumulator.response() + } + + /// Returns a stream that yields only text delta strings. + #[must_use] + pub fn text_stream(self) -> Pin> + Send>> { + Box::pin(self.filter_map(|result| { + futures::future::ready(match result { + Ok(StreamEvent::TextDelta { delta, .. }) => Some(Ok(delta)), + Err(e) => Some(Err(e)), + _ => None, + }) + })) + } +} + +impl Stream for StreamResult { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let inner = self.inner.as_mut(); + match inner.poll_next(cx) { + Poll::Ready(Some(Ok(event))) => { + self.accumulator.process(&event); + Poll::Ready(Some(Ok(event))) + } + other => other, + } + } +} + /// High-level streaming generation (Section 4.4). -/// Returns a `StreamEventStream` that the caller can iterate over. +/// Returns a `StreamResult` that the caller can iterate over. +/// Supports multi-step tool loops when active tools are provided. /// /// # Errors /// /// 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 +pub async fn stream(params: GenerateParams) -> Result { + let inner = stream_with_tool_loop(params).await?; + Ok(StreamResult::new(inner)) +} + +/// Streaming generation with multi-step tool loop support. +/// +/// When active tools are provided and the model returns tool calls: +/// - Collects the stream to get the complete first response +/// - Executes tools concurrently +/// - Starts a new stream with updated conversation +/// - Yields all events from all rounds seamlessly +/// - Continues until no more tool calls or `max_tool_rounds` reached +/// +/// # Errors +/// +/// Returns `SdkError::Configuration` if both `prompt` and `messages` are set, +/// or any provider error encountered during streaming setup. +#[allow(clippy::too_many_lines)] +async fn stream_with_tool_loop(params: GenerateParams) -> Result { + let client = params.client.clone().unwrap_or_else(get_default_client); + let mut messages = build_initial_messages(¶ms)?; + let tool_definitions: Option> = params + .tools + .as_ref() + .map(|tools| tools.iter().map(|t| t.definition.clone()).collect()); + let abort_signal = params.abort_signal.clone(); + let max_tool_rounds = params.max_tool_rounds; + + let has_active_tools = max_tool_rounds > 0 + && params + .tools + .as_ref() + .is_some_and(|tools| tools.iter().any(|t| t.is_active())); + + if !has_active_tools { + // No tool loop needed, just stream directly + return stream_generate_raw(&client, ¶ms, &messages, tool_definitions.as_deref()) + .await; + } + + // Tool loop: collect events from each round, execute tools, continue + let (tx, rx) = tokio::sync::mpsc::channel::>(64); + + let tools = params.tools.clone(); + + tokio::spawn(async move { + let mut round = 0u32; + + loop { + if let Some(ref token) = abort_signal { + if token.is_cancelled() { + let _ = tx + .send(Err(SdkError::Abort { + message: "Stream aborted by cancellation token".into(), + })) + .await; + return; + } + } + + let request = build_request(¶ms, &messages, tool_definitions.as_deref()); + let stream_result = client.stream(&request).await; + + let mut inner_stream = match stream_result { + Ok(s) => s, + Err(e) => { + let _ = tx.send(Err(e)).await; + return; + } + }; + + // Collect stream and forward events, accumulating for tool call detection + let mut accumulator = StreamAccumulator::new(); + + while let Some(item) = inner_stream.next().await { + if let Some(ref token) = abort_signal { + if token.is_cancelled() { + let _ = tx + .send(Err(SdkError::Abort { + message: "Stream aborted by cancellation token".into(), + })) + .await; + return; + } + } + + if let Ok(event) = &item { + accumulator.process(event); + } else { + let _ = tx.send(item).await; + return; + } + + // Forward the event to the consumer + if tx.send(item).await.is_err() { + return; // Consumer dropped + } + } + + // Check if we should continue with tool calls + let response = match accumulator.response() { + Some(r) => r.clone(), + None => return, // No response accumulated, stream ended + }; + + let tool_calls = response.tool_calls(); + if tool_calls.is_empty() + || response.finish_reason != FinishReason::ToolCalls + || round >= max_tool_rounds + { + return; // No more tool rounds needed + } + + // Execute tools + let Some(tool_list) = &tools else { return }; + + let tool_refs: Vec<&Tool> = tool_list.iter().map(std::convert::AsRef::as_ref).collect(); + let tool_results = execute_all_tools(&tool_refs, &tool_calls).await; + + if tool_results.is_empty() { + return; + } + + // Append assistant message and tool results to conversation + messages.push(response.message.clone()); + for result in &tool_results { + messages.push(Message::tool_result( + &result.tool_call_id, + result.content.to_string(), + result.is_error, + )); + } + + round += 1; + } + }); + + Ok(Box::pin(tokio_stream::wrappers::ReceiverStream::new(rx))) +} + +/// Internal single-round streaming (no tool loop). Used by `stream_object()`. +async fn stream_generate_raw( + client: &Arc, + params: &GenerateParams, + messages: &[Message], + tool_definitions: Option<&[ToolDefinition]>, +) -> Result { + let request = build_request(params, messages, tool_definitions); + let inner_stream = client.stream(&request).await?; + + if let Some(ref token) = params.abort_signal { + let token = token.clone(); + let mapped = inner_stream.map(move |item| { + if token.is_cancelled() { + return Err(SdkError::Abort { + message: "Stream aborted by cancellation token".into(), + }); + } + item + }); + Ok(Box::pin(mapped)) + } else { + Ok(inner_stream) + } } /// High-level streaming generation (Section 4.4). @@ -481,8 +734,7 @@ pub async fn stream_generate(params: GenerateParams) -> Result, + cancel_token: CancellationToken, + } + + #[async_trait::async_trait] + impl ProviderAdapter for AlwaysToolCallProvider { + fn name(&self) -> &str { + "mock" + } + + async fn complete(&self, _request: &Request) -> Result { + let count = self.call_count.fetch_add(1, Ordering::SeqCst); + // Cancel after first call completes + if count == 0 { + self.cancel_token.cancel(); + } + Ok(Response { + id: format!("resp_{count}"), + model: "mock-model".into(), + provider: "mock".into(), + message: Message { + role: Role::Assistant, + content: vec![ContentPart::ToolCall(ToolCall::new( + format!("call_{count}"), + "get_weather", + serde_json::json!({"city": "SF"}), + ))], + name: None, + tool_call_id: None, + }, + finish_reason: FinishReason::ToolCalls, + usage: Usage::default(), + raw: None, + warnings: vec![], + rate_limit: None, + }) + } + + async fn stream( + &self, + _request: &Request, + ) -> Result { + Ok(Box::pin(stream::empty())) + } + } + + let provider: Arc = Arc::new(AlwaysToolCallProvider { + call_count: call_count.clone(), + cancel_token: token_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?") + .tools(vec![Tool::active( + "get_weather", + "Get weather", + serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), + |_args| async { Ok(serde_json::json!("72F")) }, + )]) + .max_tool_rounds(10) + .abort_signal(token) + .client(client), + ) + .await; + + assert!(result.is_err()); + assert!(matches!(result.unwrap_err(), SdkError::Abort { .. })); + // Should have only made 1 call before aborting + assert_eq!(call_count.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn stream_abort_signal_terminates_stream() { + let token = CancellationToken::new(); + let token_clone = token.clone(); + + // Create a mock that produces events, but cancel after stream starts + let client = mock_client("Hello stream!"); + token_clone.cancel(); + + let mut stream_result = stream( + GenerateParams::new("mock-model") + .prompt("Hi") + .client(client) + .abort_signal(token), + ) + .await + .unwrap(); + + let first = stream_result.next().await.unwrap(); + assert!(first.is_err()); + assert!(matches!(first.unwrap_err(), SdkError::Abort { .. })); + } + + #[tokio::test] + async fn generate_max_tool_rounds_zero_skips_tool_execution() { + 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 tool_executed = Arc::new(AtomicU32::new(0)); + let tool_executed_clone = tool_executed.clone(); + + 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"}}}), + move |_args| { + let counter = tool_executed_clone.clone(); + async move { + counter.fetch_add(1, Ordering::SeqCst); + Ok(serde_json::json!("72F")) + } + }, + )]) + .max_tool_rounds(0) + .client(client), + ) + .await + .unwrap(); + + // Should return after first LLM call without executing any tools + assert_eq!(result.steps.len(), 1); + assert_eq!(call_count.load(Ordering::SeqCst), 1); + assert_eq!(tool_executed.load(Ordering::SeqCst), 0); + // The tool results should be empty since tools were not executed + assert!(result.tool_results.is_empty()); + } + + #[test] + fn generate_params_abort_signal_builder() { + let token = CancellationToken::new(); + let params = GenerateParams::new("test-model") + .abort_signal(token); + assert!(params.abort_signal.is_some()); + } + + #[tokio::test] + async fn stream_result_accumulates_response() { + let client = mock_client("Hello!"); + let mut result = stream( + GenerateParams::new("mock-model") + .prompt("Hi") + .client(client), + ) + .await + .unwrap(); + + assert!(result.response().is_none()); + assert!(result.partial_response().is_none()); + + // Consume all events + while result.next().await.is_some() {} + + assert!(result.response().is_some()); + assert_eq!(result.response().unwrap().text(), "Hello!"); + } + + #[tokio::test] + async fn stream_result_text_stream() { + let client = streaming_json_mock_client(vec!["Hello", " ", "world"]); + let result = stream( + GenerateParams::new("mock-model") + .prompt("Hi") + .client(client), + ) + .await + .unwrap(); + + let texts: Vec = result + .text_stream() + .filter_map(|r| futures::future::ready(r.ok())) + .collect() + .await; + + assert_eq!(texts, vec!["Hello", " ", "world"]); + } + + /// Mock provider that streams tool calls then text on second stream + struct StreamingToolCallMockProvider { + call_count: Arc, + } + + #[async_trait::async_trait] + impl ProviderAdapter for StreamingToolCallMockProvider { + 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("fallback"), + finish_reason: FinishReason::Stop, + usage: Usage::default(), + raw: None, + warnings: vec![], + rate_limit: None, + }) + } + + async fn stream( + &self, + _request: &Request, + ) -> Result { + let count = self.call_count.fetch_add(1, Ordering::SeqCst); + + if count == 0 { + // First stream: return tool call + let tool_call = ToolCall::new("call_1", "get_weather", serde_json::json!({"city": "SF"})); + let response = Response { + id: "resp_1".into(), + model: "mock-model".into(), + provider: "mock".into(), + message: Message { + role: Role::Assistant, + content: vec![ContentPart::ToolCall(tool_call.clone())], + name: None, + tool_call_id: None, + }, + finish_reason: FinishReason::ToolCalls, + usage: Usage { + input_tokens: 10, + output_tokens: 5, + total_tokens: 15, + ..Default::default() + }, + raw: None, + warnings: vec![], + rate_limit: None, + }; + let events = vec![ + Ok(StreamEvent::ToolCallEnd { tool_call }), + Ok(StreamEvent::finish( + FinishReason::ToolCalls, + response.usage.clone(), + response, + )), + ]; + Ok(Box::pin(stream::iter(events))) + } else { + // Second stream: return text + let text = "The weather in SF is 72F"; + let response = Response { + id: "resp_2".into(), + model: "mock-model".into(), + provider: "mock".into(), + message: Message::assistant(text), + finish_reason: FinishReason::Stop, + usage: Usage { + input_tokens: 20, + output_tokens: 10, + total_tokens: 30, + ..Default::default() + }, + raw: None, + warnings: vec![], + rate_limit: None, + }; + let events = vec![ + Ok(StreamEvent::text_delta(text, Some("t1".into()))), + Ok(StreamEvent::finish( + FinishReason::Stop, + response.usage.clone(), + response, + )), + ]; + Ok(Box::pin(stream::iter(events))) + } + } + } + + #[tokio::test] + async fn stream_with_tool_loop_executes_tools() { + let call_count = Arc::new(AtomicU32::new(0)); + let provider: Arc = Arc::new(StreamingToolCallMockProvider { + 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 mut result = stream( + 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 { Ok(serde_json::json!("72F")) }, + )]) + .max_tool_rounds(5) + .client(client), + ) + .await + .unwrap(); + + // Collect all events + let mut events = Vec::new(); + while let Some(item) = result.next().await { + events.push(item); + } + + // Should have events from both rounds + assert_eq!(call_count.load(Ordering::SeqCst), 2); + + // Should have text deltas from the second round + let text_deltas: Vec<_> = events + .iter() + .filter_map(|e| match e { + Ok(StreamEvent::TextDelta { delta, .. }) => Some(delta.as_str()), + _ => None, + }) + .collect(); + assert_eq!(text_deltas, vec!["The weather in SF is 72F"]); + + // The final response should be the text response + assert!(result.response().is_some()); + assert_eq!(result.response().unwrap().text(), "The weather in SF is 72F"); + } + + #[tokio::test] + async fn stream_no_tool_loop_when_max_rounds_zero() { + let call_count = Arc::new(AtomicU32::new(0)); + let provider: Arc = Arc::new(StreamingToolCallMockProvider { + 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 mut result = stream( + GenerateParams::new("mock-model") + .prompt("What's the weather?") + .tools(vec![Tool::active( + "get_weather", + "Get weather", + serde_json::json!({"type": "object", "properties": {"city": {"type": "string"}}}), + |_args| async { Ok(serde_json::json!("72F")) }, + )]) + .max_tool_rounds(0) + .client(client), + ) + .await + .unwrap(); + + // Consume all events + while result.next().await.is_some() {} + + // Only one stream call, no tool execution + assert_eq!(call_count.load(Ordering::SeqCst), 1); + } } diff --git a/crates/unified-llm/src/lib.rs b/crates/unified-llm/src/lib.rs index d75e74379..af92eeaf3 100644 --- a/crates/unified-llm/src/lib.rs +++ b/crates/unified-llm/src/lib.rs @@ -8,3 +8,5 @@ pub mod retry; pub mod generate; pub mod catalog; pub mod providers; + +pub use tokio_util::sync::CancellationToken; diff --git a/crates/unified-llm/src/providers/gemini.rs b/crates/unified-llm/src/providers/gemini.rs index e60ef50e3..f9f95fd96 100644 --- a/crates/unified-llm/src/providers/gemini.rs +++ b/crates/unified-llm/src/providers/gemini.rs @@ -328,7 +328,10 @@ fn translate_response_format( } /// Build the Gemini API request body from a unified `Request`. -fn build_api_request(request: &Request) -> ApiRequest { +/// +/// Returns a `serde_json::Value` so that `provider_options.gemini` fields can be +/// merged into the request before sending. +fn build_api_request(request: &Request) -> serde_json::Value { let (system_text, other_messages) = extract_system_prompt(&request.messages); let system_instruction = system_text.map(|text| SystemInstruction { @@ -354,12 +357,37 @@ fn build_api_request(request: &Request) -> ApiRequest { 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 { + let api_request = ApiRequest { contents, system_instruction, generation_config: Some(generation_config), tools: api_tools, tool_config, + }; + + let mut body = serde_json::to_value(&api_request).unwrap_or_default(); + merge_provider_options(&mut body, request.provider_options.as_ref()); + body +} + +/// Merge `provider_options.gemini` fields into the serialized API request body. +/// +/// Known fields like `safety_settings` and `cached_content` are set directly. +/// Any other fields are merged at the top level, allowing pass-through of +/// Gemini-specific options not covered by the unified schema. +fn merge_provider_options(body: &mut serde_json::Value, provider_options: Option<&serde_json::Value>) { + let Some(gemini_opts) = provider_options.and_then(|opts| opts.get("gemini")) else { + return; + }; + let Some(body_map) = body.as_object_mut() else { + return; + }; + let Some(gemini_map) = gemini_opts.as_object() else { + return; + }; + + for (key, value) in gemini_map { + body_map.insert(key.clone(), value.clone()); } } @@ -688,7 +716,7 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { - let api_request = build_api_request(request); + let api_body = build_api_request(request); let url = format!( "{}/models/{}:generateContent?key={}", @@ -696,7 +724,7 @@ impl ProviderAdapter for Adapter { ); let body = send_and_read_body( - self.client.post(&url).json(&api_request).timeout(self.request_timeout), + self.client.post(&url).json(&api_body).timeout(self.request_timeout), "gemini", "status", ) @@ -751,7 +779,7 @@ impl ProviderAdapter for Adapter { } async fn stream(&self, request: &Request) -> Result { - let api_request = build_api_request(request); + let api_body = build_api_request(request); let url = format!( "{}/models/{}:streamGenerateContent?alt=sse&key={}", @@ -759,8 +787,139 @@ impl ProviderAdapter for Adapter { ); let http_resp = - send_streaming_request(self.client.post(&url).json(&api_request)).await?; + send_streaming_request(self.client.post(&url).json(&api_body)).await?; Ok(process_sse_stream(http_resp, request.model.clone())) } } + +#[cfg(test)] +mod tests { + use super::*; + + fn minimal_request() -> Request { + Request { + model: "gemini-2.0-flash".to_string(), + messages: vec![Message::user("Hello")], + provider: None, + tools: None, + tool_choice: None, + response_format: None, + temperature: None, + top_p: None, + max_tokens: None, + stop_sequences: None, + reasoning_effort: None, + metadata: None, + provider_options: None, + } + } + + #[test] + fn provider_options_none_produces_standard_body() { + let request = minimal_request(); + let body = build_api_request(&request); + assert!(body.get("safetySettings").is_none()); + assert!(body.get("cachedContent").is_none()); + } + + #[test] + fn provider_options_gemini_safety_settings_merged() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "gemini": { + "safetySettings": [ + {"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"} + ] + } + })); + + let body = build_api_request(&request); + let safety = body.get("safetySettings").expect("safetySettings should be present"); + let arr = safety.as_array().expect("should be an array"); + assert_eq!(arr.len(), 1); + assert_eq!(arr[0]["category"], "HARM_CATEGORY_HARASSMENT"); + } + + #[test] + fn provider_options_gemini_cached_content_merged() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "gemini": { + "cachedContent": "projects/my-project/cachedContents/abc123" + } + })); + + let body = build_api_request(&request); + assert_eq!( + body.get("cachedContent").and_then(serde_json::Value::as_str), + Some("projects/my-project/cachedContents/abc123") + ); + } + + #[test] + fn provider_options_gemini_multiple_fields_merged() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "gemini": { + "safetySettings": [{"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_LOW_AND_ABOVE"}], + "cachedContent": "cache-id", + "customField": "custom-value" + } + })); + + let body = build_api_request(&request); + assert!(body.get("safetySettings").is_some()); + assert_eq!( + body.get("cachedContent").and_then(serde_json::Value::as_str), + Some("cache-id") + ); + assert_eq!( + body.get("customField").and_then(serde_json::Value::as_str), + Some("custom-value") + ); + } + + #[test] + fn provider_options_other_provider_ignored() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "anthropic": { + "auto_cache": false + } + })); + + let body = build_api_request(&request); + assert!(body.get("auto_cache").is_none()); + } + + #[test] + fn provider_options_gemini_preserves_standard_fields() { + let mut request = minimal_request(); + request.temperature = Some(0.5); + request.max_tokens = Some(100); + request.provider_options = Some(serde_json::json!({ + "gemini": { + "cachedContent": "cache-id" + } + })); + + let body = build_api_request(&request); + let gen_config = body.get("generationConfig").expect("generationConfig should exist"); + assert_eq!(gen_config.get("temperature").and_then(serde_json::Value::as_f64), Some(0.5)); + assert_eq!(gen_config.get("maxOutputTokens").and_then(serde_json::Value::as_i64), Some(100)); + assert_eq!( + body.get("cachedContent").and_then(serde_json::Value::as_str), + Some("cache-id") + ); + } + + #[test] + fn merge_provider_options_with_non_object_gemini_value() { + let mut body = serde_json::json!({"contents": []}); + let opts = serde_json::json!({"gemini": "not-an-object"}); + merge_provider_options(&mut body, Some(&opts)); + // Should not crash and body should be unchanged + assert!(body.get("contents").is_some()); + } +} diff --git a/crates/unified-llm/src/providers/openai.rs b/crates/unified-llm/src/providers/openai.rs index c9400d20d..76d925902 100644 --- a/crates/unified-llm/src/providers/openai.rs +++ b/crates/unified-llm/src/providers/openai.rs @@ -67,6 +67,8 @@ struct ApiRequest { reasoning: Option, #[serde(skip_serializing_if = "Option::is_none")] text: Option, + #[serde(skip_serializing_if = "Option::is_none")] + metadata: Option>, #[serde(skip_serializing_if = "std::ops::Not::not")] stream: bool, } @@ -294,10 +296,31 @@ fn build_api_request(request: &Request, stream: bool) -> ApiRequest { tool_choice, reasoning, text, + metadata: request.metadata.clone(), stream, } } +/// Serialize an `ApiRequest` to JSON and merge any `provider_options.openai` keys into it. +fn build_request_body(request: &Request, stream: bool) -> serde_json::Value { + let api_request = build_api_request(request, stream); + let mut body = serde_json::to_value(&api_request).unwrap_or_else(|_| serde_json::json!({})); + + if let Some(openai_opts) = request + .provider_options + .as_ref() + .and_then(|opts| opts.get("openai")) + { + if let (Some(base), Some(overrides)) = (body.as_object_mut(), openai_opts.as_object()) { + for (key, value) in overrides { + base.insert(key.clone(), value.clone()); + } + } + } + + body +} + /// Parse output items from the Responses API into content parts. fn parse_output(output: &[serde_json::Value]) -> (Vec, bool) { let mut parts = Vec::new(); @@ -744,14 +767,14 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { - let api_request = build_api_request(request, false); + let request_body = build_request_body(request, false); let url = format!("{}/responses", self.base_url); let (body, headers) = send_and_read_response( self.client .post(&url) .bearer_auth(&self.api_key) - .json(&api_request) + .json(&request_body) .timeout(self.request_timeout), "openai", "type", @@ -803,14 +826,14 @@ impl ProviderAdapter for Adapter { } async fn stream(&self, request: &Request) -> Result { - let api_request = build_api_request(request, true); + let request_body = build_request_body(request, true); let url = format!("{}/responses", self.base_url); let http_resp = self .client .post(&url) .bearer_auth(&self.api_key) - .json(&api_request) + .json(&request_body) .send() .await .map_err(|e| SdkError::Network { @@ -868,3 +891,127 @@ impl ProviderAdapter for Adapter { Ok(Box::pin(stream)) } } + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + fn minimal_request() -> Request { + Request { + model: "gpt-4o".to_string(), + messages: vec![Message::user("Hello")], + provider: None, + tools: None, + tool_choice: None, + response_format: None, + temperature: None, + top_p: None, + max_tokens: None, + stop_sequences: None, + reasoning_effort: None, + metadata: None, + provider_options: None, + } + } + + #[test] + fn build_request_body_includes_metadata() { + let mut metadata = HashMap::new(); + metadata.insert("user_id".to_string(), "u123".to_string()); + metadata.insert("session".to_string(), "s456".to_string()); + + let mut request = minimal_request(); + request.metadata = Some(metadata); + + let body = build_request_body(&request, false); + let meta = body.get("metadata").expect("metadata should be present"); + assert_eq!(meta["user_id"], "u123"); + assert_eq!(meta["session"], "s456"); + } + + #[test] + fn build_request_body_omits_metadata_when_none() { + let request = minimal_request(); + let body = build_request_body(&request, false); + assert!(body.get("metadata").is_none()); + } + + #[test] + fn build_request_body_merges_provider_options_openai() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "openai": { + "store": true, + "previous_response_id": "resp_abc123" + } + })); + + let body = build_request_body(&request, false); + assert_eq!(body["store"], true); + assert_eq!(body["previous_response_id"], "resp_abc123"); + } + + #[test] + fn build_request_body_provider_options_override_fields() { + let mut request = minimal_request(); + request.temperature = Some(0.5); + request.provider_options = Some(serde_json::json!({ + "openai": { + "temperature": 0.9 + } + })); + + let body = build_request_body(&request, false); + // provider_options should override the base field + assert_eq!(body["temperature"], 0.9); + } + + #[test] + fn build_request_body_ignores_non_openai_provider_options() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "anthropic": { + "thinking": {"type": "enabled", "budget_tokens": 10000} + } + })); + + let body = build_request_body(&request, false); + // anthropic options should not leak into the OpenAI request + assert!(body.get("thinking").is_none()); + } + + #[test] + fn build_request_body_no_provider_options() { + let request = minimal_request(); + let body = build_request_body(&request, false); + assert_eq!(body["model"], "gpt-4o"); + // stream field is omitted when false (skip_serializing_if) + assert!(body.get("stream").is_none()); + } + + #[test] + fn build_request_body_stream_flag() { + let request = minimal_request(); + let body = build_request_body(&request, true); + assert!(body["stream"].as_bool().unwrap_or(false)); + } + + #[test] + fn build_request_body_metadata_and_provider_options_together() { + let mut metadata = HashMap::new(); + metadata.insert("trace_id".to_string(), "t789".to_string()); + + let mut request = minimal_request(); + request.metadata = Some(metadata); + request.provider_options = Some(serde_json::json!({ + "openai": { + "store": true + } + })); + + let body = build_request_body(&request, false); + assert_eq!(body["metadata"]["trace_id"], "t789"); + assert_eq!(body["store"], true); + } +} diff --git a/crates/unified-llm/src/providers/openai_compatible.rs b/crates/unified-llm/src/providers/openai_compatible.rs index 8d0408592..63196cc0b 100644 --- a/crates/unified-llm/src/providers/openai_compatible.rs +++ b/crates/unified-llm/src/providers/openai_compatible.rs @@ -291,8 +291,11 @@ fn translate_response_format(format: &ResponseFormat) -> serde_json::Value { } } -/// Build an `ApiRequest` from a unified `Request`. -fn build_api_request(request: &Request, stream: Option) -> ApiRequest { +/// Build the API request body from a unified `Request`. +/// +/// Returns a `serde_json::Value` so that `provider_options.` fields +/// can be merged into the request before sending. +fn build_api_request(request: &Request, stream: Option, provider_name: &str) -> serde_json::Value { 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); @@ -301,7 +304,7 @@ fn build_api_request(request: &Request, stream: Option) -> ApiRequest { .as_ref() .map(translate_response_format); - ApiRequest { + let api_request = ApiRequest { model: request.model.clone(), messages: chat_messages, temperature: request.temperature, @@ -312,6 +315,34 @@ fn build_api_request(request: &Request, stream: Option) -> ApiRequest { tool_choice, response_format, stream, + }; + + let mut body = serde_json::to_value(&api_request).unwrap_or_default(); + merge_provider_options(&mut body, request.provider_options.as_ref(), provider_name); + body +} + +/// Merge `provider_options.` fields into the serialized API request body. +/// +/// The provider name is configurable (e.g. "groq", "together", "openai-compatible"), +/// allowing each instance to have its own namespace in `provider_options`. +fn merge_provider_options( + body: &mut serde_json::Value, + provider_options: Option<&serde_json::Value>, + provider_name: &str, +) { + let Some(opts) = provider_options.and_then(|opts| opts.get(provider_name)) else { + return; + }; + let Some(body_map) = body.as_object_mut() else { + return; + }; + let Some(opts_map) = opts.as_object() else { + return; + }; + + for (key, value) in opts_map { + body_map.insert(key.clone(), value.clone()); } } @@ -323,14 +354,14 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { - let api_request = build_api_request(request, None); + let api_body = build_api_request(request, None, &self.provider_name); 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) + .json(&api_body) .timeout(self.request_timeout), &self.provider_name, "type", @@ -397,14 +428,14 @@ impl ProviderAdapter for Adapter { } async fn stream(&self, request: &Request) -> Result { - let api_request = build_api_request(request, Some(true)); + let api_body = build_api_request(request, Some(true), &self.provider_name); let url = format!("{}/chat/completions", self.base_url); let http_resp = self .client .post(&url) .bearer_auth(&self.api_key) - .json(&api_request) + .json(&api_body) .send() .await .map_err(|e| SdkError::Network { @@ -1111,4 +1142,111 @@ mod tests { assert_eq!(tool_calls[0]["id"], "call_1"); assert_eq!(tool_calls[0]["function"]["name"], "get_weather"); } + + fn minimal_request() -> Request { + Request { + model: "llama-3.1-70b".to_string(), + messages: vec![Message::user("Hello")], + provider: None, + tools: None, + tool_choice: None, + response_format: None, + temperature: None, + top_p: None, + max_tokens: None, + stop_sequences: None, + reasoning_effort: None, + metadata: None, + provider_options: None, + } + } + + #[test] + fn provider_options_none_produces_standard_body() { + let request = minimal_request(); + let body = build_api_request(&request, None, "groq"); + assert_eq!(body["model"], "llama-3.1-70b"); + assert!(body.get("stream").is_none()); + } + + #[test] + fn provider_options_matching_name_merged() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "groq": { + "frequency_penalty": 0.5, + "presence_penalty": 0.3 + } + })); + + let body = build_api_request(&request, None, "groq"); + assert_eq!(body["frequency_penalty"], 0.5); + assert_eq!(body["presence_penalty"], 0.3); + } + + #[test] + fn provider_options_different_name_ignored() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "together": { + "repetition_penalty": 1.2 + } + })); + + let body = build_api_request(&request, None, "groq"); + assert!(body.get("repetition_penalty").is_none()); + } + + #[test] + fn provider_options_uses_adapter_name() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "together": { + "repetition_penalty": 1.2 + } + })); + + let body = build_api_request(&request, None, "together"); + assert_eq!(body["repetition_penalty"], 1.2); + } + + #[test] + fn provider_options_preserves_standard_fields() { + let mut request = minimal_request(); + request.temperature = Some(0.7); + request.max_tokens = Some(200); + request.provider_options = Some(serde_json::json!({ + "groq": { + "frequency_penalty": 0.5 + } + })); + + let body = build_api_request(&request, Some(true), "groq"); + assert_eq!(body["temperature"], 0.7); + assert_eq!(body["max_tokens"], 200); + assert_eq!(body["stream"], true); + assert_eq!(body["frequency_penalty"], 0.5); + } + + #[test] + fn provider_options_can_override_model() { + let mut request = minimal_request(); + request.provider_options = Some(serde_json::json!({ + "groq": { + "model": "custom-model" + } + })); + + let body = build_api_request(&request, None, "groq"); + assert_eq!(body["model"], "custom-model"); + } + + #[test] + fn merge_provider_options_with_non_object_value() { + let mut body = serde_json::json!({"model": "test"}); + let opts = serde_json::json!({"groq": "not-an-object"}); + merge_provider_options(&mut body, Some(&opts), "groq"); + // Should not crash and body should be unchanged + assert_eq!(body["model"], "test"); + } }