From 5a765518478e3ef732bac755db0988a7328ac47a Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp Date: Fri, 20 Feb 2026 11:51:26 -0400 Subject: [PATCH] Implement spec gaps: audio/document content parts, streaming tool loop improvements - Add Audio/Document content part handling across all providers: Anthropic supports documents natively, Gemini supports both audio and documents, OpenAI and OpenAI-compatible produce text fallbacks - Add stop_when support to streaming tool loops (was only in generate()) - Add retry on initial stream connection (matching generate() behavior) - Add total and per_step timeout support to streaming tool loops Co-Authored-By: Claude Opus 4.6 --- crates/unified-llm/src/generate.rs | 630 +++++++++++++++--- crates/unified-llm/src/providers/anthropic.rs | 77 ++- crates/unified-llm/src/providers/gemini.rs | 120 ++++ crates/unified-llm/src/providers/openai.rs | 67 ++ .../src/providers/openai_compatible.rs | 106 ++- 5 files changed, 919 insertions(+), 81 deletions(-) diff --git a/crates/unified-llm/src/generate.rs b/crates/unified-llm/src/generate.rs index a713a599a..313c84b7b 100644 --- a/crates/unified-llm/src/generate.rs +++ b/crates/unified-llm/src/generate.rs @@ -595,37 +595,19 @@ async fn stream_with_tool_loop(params: GenerateParams) -> Result>(64); let tools = params.tools.clone(); + let retry_policy = RetryPolicy { + max_retries: params.max_retries, + base_delay: 0.001, + jitter: false, + ..Default::default() + }; tokio::spawn(async move { - let mut round = 0u32; + let tool_loop_future = async { + let mut round = 0u32; + let mut steps: Vec = Vec::new(); - 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 { + loop { if let Some(ref token) = abort_signal { if token.is_cancelled() { let _ = tx @@ -637,66 +619,149 @@ async fn stream_with_tool_loop(params: GenerateParams) -> Result 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; } - // Forward the event to the consumer - if tx.send(item).await.is_err() { + // Track step results for stop_when + steps.push(StepResult { + response: response.clone(), + tool_results: tool_results.clone(), + }); + + // Check stop_when condition (Section 4.3) + if params.stop_when.as_ref().is_some_and(|f| f(&steps)) { + // Emit StepFinish but do not continue to next round + let step_finish = StreamEvent::step_finish( + response.finish_reason.clone(), + response.usage.clone(), + response, + tool_calls, + tool_results, + ); + let _ = tx.send(Ok(step_finish)).await; + return; + } + + // Emit StepFinish event between steps + let step_finish = StreamEvent::step_finish( + response.finish_reason.clone(), + response.usage.clone(), + response.clone(), + tool_calls, + tool_results.clone(), + ); + if tx.send(Ok(step_finish)).await.is_err() { return; // Consumer dropped } + + // 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; } + }; - // 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 + // Apply total timeout if configured (Section 4.7) + if let Some(total) = params.timeout.as_ref().and_then(|t| t.total) { + let duration = std::time::Duration::from_secs_f64(total); + if tokio::time::timeout(duration, tool_loop_future) + .await + .is_err() { - return; // No more tool rounds needed + let _ = tx + .send(Err(SdkError::RequestTimeout { + message: format!("Total timeout of {total}s exceeded"), + })) + .await; } - - // 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; - } - - // Emit StepFinish event between steps - let step_finish = StreamEvent::step_finish( - response.finish_reason.clone(), - response.usage.clone(), - response.clone(), - tool_calls, - tool_results.clone(), - ); - if tx.send(Ok(step_finish)).await.is_err() { - return; // Consumer dropped - } - - // 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; + } else { + tool_loop_future.await; } }); @@ -2080,4 +2145,411 @@ mod tests { assert_eq!(step_finish.2.len(), 1); assert_eq!(step_finish.2[0].tool_call_id, "call_1"); } + + #[tokio::test] + async fn stream_stop_when_halts_streaming_tool_loop() { + 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) + .stop_when(|_steps| true) // Stop immediately after first round + .client(client), + ) + .await + .unwrap(); + + let mut events = Vec::new(); + while let Some(item) = result.next().await { + events.push(item); + } + + // stop_when returned true, so only 1 stream call should have been made + assert_eq!(call_count.load(Ordering::SeqCst), 1); + + // Should have a StepFinish event but no second round text + let step_finish_count = events + .iter() + .filter(|e| matches!(e, Ok(StreamEvent::StepFinish { .. }))) + .count(); + assert_eq!(step_finish_count, 1, "Expected StepFinish event from stopped round"); + + // Should NOT have any text deltas (second round never started) + let text_delta_count = events + .iter() + .filter(|e| matches!(e, Ok(StreamEvent::TextDelta { .. }))) + .count(); + assert_eq!(text_delta_count, 0, "Expected no text deltas since loop was stopped"); + } + + /// Mock provider that fails on stream N times then succeeds + struct FailThenStreamProvider { + call_count: Arc, + failures: u32, + } + + #[async_trait::async_trait] + impl ProviderAdapter for FailThenStreamProvider { + 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 < self.failures { + return Err(SdkError::Provider { + kind: crate::error::ProviderErrorKind::Server, + detail: Box::new(crate::error::ProviderErrorDetail { + status_code: Some(500), + ..crate::error::ProviderErrorDetail::new("server error", "mock") + }), + }); + } + + let text = "Hello after retry"; + let response = Response { + id: "resp_1".into(), + model: "mock-model".into(), + provider: "mock".into(), + message: Message::assistant(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, + }; + 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_retry_on_initial_connection() { + let call_count = Arc::new(AtomicU32::new(0)); + let provider: Arc = Arc::new(FailThenStreamProvider { + call_count: call_count.clone(), + failures: 2, // fail twice, succeed on third + }); + + let mut providers: HashMap> = HashMap::new(); + providers.insert("mock".to_string(), provider); + let client = Arc::new(Client::new( + providers, + Some("mock".to_string()), + vec![], + )); + + // Need active tools so the tool loop path (with retry) is used + let mut result = stream( + GenerateParams::new("mock-model") + .prompt("Hi") + .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(1) + .max_retries(3) + .client(client), + ) + .await + .unwrap(); + + let mut events = Vec::new(); + while let Some(item) = result.next().await { + events.push(item); + } + + // Should have called stream 3 times (2 failures + 1 success) + assert_eq!(call_count.load(Ordering::SeqCst), 3); + + // Should have received the text from the successful attempt + 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!["Hello after retry"]); + } + + /// Mock provider that delays before returning stream + struct SlowStreamProvider { + delay: std::time::Duration, + } + + #[async_trait::async_trait] + impl ProviderAdapter for SlowStreamProvider { + 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 { + tokio::time::sleep(self.delay).await; + let text = "Slow response"; + let response = Response { + id: "resp_1".into(), + model: "mock-model".into(), + provider: "mock".into(), + message: Message::assistant(text), + finish_reason: FinishReason::Stop, + usage: Usage::default(), + raw: None, + warnings: vec![], + rate_limit: None, + }; + let events = vec![ + Ok(StreamEvent::text_delta(text, Some("t1".into()))), + Ok(StreamEvent::finish(FinishReason::Stop, Usage::default(), response)), + ]; + Ok(Box::pin(stream::iter(events))) + } + } + + #[tokio::test] + async fn stream_per_step_timeout() { + let provider: Arc = Arc::new(SlowStreamProvider { + delay: std::time::Duration::from_secs(5), + }); + + let mut providers: HashMap> = HashMap::new(); + providers.insert("mock".to_string(), provider); + let client = Arc::new(Client::new( + providers, + Some("mock".to_string()), + vec![], + )); + + // Need active tools so the tool loop path (with timeout) is used + let mut result = stream( + GenerateParams::new("mock-model") + .prompt("Hi") + .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(1) + .timeout(TimeoutConfig { + total: None, + per_step: Some(0.01), // 10ms timeout, provider takes 5s + }) + .max_retries(0) + .client(client), + ) + .await + .unwrap(); + + let mut events = Vec::new(); + while let Some(item) = result.next().await { + events.push(item); + } + + // Should have received a timeout error + let has_timeout = events.iter().any(|e| { + matches!(e, Err(SdkError::RequestTimeout { .. })) + }); + assert!(has_timeout, "Expected a RequestTimeout error"); + } + + #[tokio::test] + async fn stream_total_timeout() { + // Use a streaming tool call provider with a slow tool to trigger total timeout + // across multiple rounds + let call_count = Arc::new(AtomicU32::new(0)); + + /// Provider that always returns tool calls with a delay on the second stream + struct SlowToolCallStreamProvider { + call_count: Arc, + } + + #[async_trait::async_trait] + impl ProviderAdapter for SlowToolCallStreamProvider { + 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 quickly + 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::default(), + raw: None, + warnings: vec![], + rate_limit: None, + }; + let events = vec![ + Ok(StreamEvent::ToolCallEnd { tool_call }), + Ok(StreamEvent::finish( + FinishReason::ToolCalls, + Usage::default(), + response, + )), + ]; + Ok(Box::pin(stream::iter(events))) + } else { + // Second stream: delay long enough to exceed total timeout + tokio::time::sleep(std::time::Duration::from_secs(5)).await; + let text = "Should not arrive"; + let response = Response { + id: "resp_2".into(), + model: "mock-model".into(), + provider: "mock".into(), + message: Message::assistant(text), + finish_reason: FinishReason::Stop, + usage: Usage::default(), + raw: None, + warnings: vec![], + rate_limit: None, + }; + let events = vec![ + Ok(StreamEvent::text_delta(text, Some("t1".into()))), + Ok(StreamEvent::finish(FinishReason::Stop, Usage::default(), response)), + ]; + Ok(Box::pin(stream::iter(events))) + } + } + } + + let provider: Arc = Arc::new(SlowToolCallStreamProvider { + 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(5) + .timeout(TimeoutConfig { + total: Some(0.05), // 50ms total timeout + per_step: None, + }) + .max_retries(0) + .client(client), + ) + .await + .unwrap(); + + let mut events = Vec::new(); + while let Some(item) = result.next().await { + events.push(item); + } + + // Should have received a total timeout error + let has_timeout = events.iter().any(|e| { + matches!(e, Err(SdkError::RequestTimeout { .. })) + }); + assert!(has_timeout, "Expected a RequestTimeout error from total timeout"); + } } diff --git a/crates/unified-llm/src/providers/anthropic.rs b/crates/unified-llm/src/providers/anthropic.rs index eadf7c7a9..27829e1cb 100644 --- a/crates/unified-llm/src/providers/anthropic.rs +++ b/crates/unified-llm/src/providers/anthropic.rs @@ -234,7 +234,29 @@ fn content_part_to_api(part: &ContentPart) -> Option { }) } } - _ => None, + ContentPart::Document(doc) => { + if let Some(url) = &doc.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": "document", + "source": {"type": "base64", "media_type": mime, "data": b64} + })), + Err(_) => None, + }; + } + Some(serde_json::json!({"type": "document", "source": {"type": "url", "url": url}})) + } else { + doc.data.as_ref().map(|data| { + let mime = doc.media_type.as_deref().unwrap_or("application/pdf"); + let b64 = BASE64_STANDARD.encode(data); + serde_json::json!({"type": "document", "source": {"type": "base64", "media_type": mime, "data": b64}}) + }) + } + } + ContentPart::Audio(_) => { + Some(serde_json::json!({"type": "text", "text": "[Audio content not supported by this provider]"})) + } } } @@ -1765,4 +1787,57 @@ mod tests { _ => panic!("expected Finish"), } } + + #[test] + fn document_url_translates_to_url_source() { + let part = ContentPart::Document(crate::types::DocumentData { + url: Some("https://example.com/doc.pdf".to_string()), + data: None, + media_type: None, + file_name: None, + }); + let result = content_part_to_api(&part).expect("should produce JSON"); + assert_eq!(result["type"], "document"); + assert_eq!(result["source"]["type"], "url"); + assert_eq!(result["source"]["url"], "https://example.com/doc.pdf"); + } + + #[test] + fn document_base64_data_translates_to_base64_source() { + let part = ContentPart::Document(crate::types::DocumentData { + url: None, + data: Some(vec![0x25, 0x50, 0x44, 0x46]), + media_type: Some("application/pdf".to_string()), + file_name: Some("test.pdf".to_string()), + }); + let result = content_part_to_api(&part).expect("should produce JSON"); + assert_eq!(result["type"], "document"); + assert_eq!(result["source"]["type"], "base64"); + assert_eq!(result["source"]["media_type"], "application/pdf"); + assert!(result["source"]["data"].as_str().is_some()); + } + + #[test] + fn document_base64_defaults_to_pdf_mime() { + let part = ContentPart::Document(crate::types::DocumentData { + url: None, + data: Some(vec![1, 2, 3]), + media_type: None, + file_name: None, + }); + let result = content_part_to_api(&part).expect("should produce JSON"); + assert_eq!(result["source"]["media_type"], "application/pdf"); + } + + #[test] + fn audio_produces_text_fallback() { + let part = ContentPart::Audio(crate::types::AudioData { + url: Some("https://example.com/audio.wav".to_string()), + data: None, + media_type: None, + }); + let result = content_part_to_api(&part).expect("should produce JSON"); + assert_eq!(result["type"], "text"); + assert_eq!(result["text"], "[Audio content not supported by this provider]"); + } } diff --git a/crates/unified-llm/src/providers/gemini.rs b/crates/unified-llm/src/providers/gemini.rs index 432894c59..381ca0398 100644 --- a/crates/unified-llm/src/providers/gemini.rs +++ b/crates/unified-llm/src/providers/gemini.rs @@ -200,6 +200,7 @@ fn build_tool_call_id_to_name(messages: &[&Message]) -> std::collections::HashMa } /// Translate unified messages to Gemini content format. +#[allow(clippy::too_many_lines)] fn translate_messages(messages: &[&Message]) -> Vec { let id_to_name = build_tool_call_id_to_name(messages); let mut contents: Vec = Vec::new(); @@ -244,6 +245,50 @@ fn translate_messages(messages: &[&Message]) -> Vec { }, ) } + ContentPart::Audio(audio) => { + audio.url.as_ref().map_or_else( + || { + audio.data.as_ref().map(|data| { + let mime = audio.media_type.as_deref().unwrap_or("audio/wav"); + 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 = audio.media_type.as_deref().unwrap_or("audio/wav"); + Some(serde_json::json!({"fileData": {"mimeType": mime, "fileUri": url}})) + } + }, + ) + } + ContentPart::Document(doc) => { + doc.url.as_ref().map_or_else( + || { + doc.data.as_ref().map(|data| { + let mime = doc.media_type.as_deref().unwrap_or("application/pdf"); + 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 = doc.media_type.as_deref().unwrap_or("application/pdf"); + 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. @@ -942,4 +987,79 @@ mod tests { // Should not crash and body should be unchanged assert!(body.get("contents").is_some()); } + + #[test] + fn audio_url_translates_to_file_data() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Audio(crate::types::AudioData { + url: Some("https://example.com/audio.wav".to_string()), + data: None, + media_type: Some("audio/wav".to_string()), + })], + name: None, + tool_call_id: None, + }; + let contents = translate_messages(&[&msg]); + assert_eq!(contents.len(), 1); + let part = &contents[0].parts[0]; + assert_eq!(part["fileData"]["mimeType"], "audio/wav"); + assert_eq!(part["fileData"]["fileUri"], "https://example.com/audio.wav"); + } + + #[test] + fn audio_base64_translates_to_inline_data() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Audio(crate::types::AudioData { + url: None, + data: Some(vec![0xFF, 0xFB, 0x90]), + media_type: None, + })], + name: None, + tool_call_id: None, + }; + let contents = translate_messages(&[&msg]); + let part = &contents[0].parts[0]; + assert_eq!(part["inlineData"]["mimeType"], "audio/wav"); + assert!(part["inlineData"]["data"].as_str().is_some()); + } + + #[test] + fn document_url_translates_to_file_data() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Document(crate::types::DocumentData { + url: Some("https://example.com/doc.pdf".to_string()), + data: None, + media_type: Some("application/pdf".to_string()), + file_name: Some("doc.pdf".to_string()), + })], + name: None, + tool_call_id: None, + }; + let contents = translate_messages(&[&msg]); + let part = &contents[0].parts[0]; + assert_eq!(part["fileData"]["mimeType"], "application/pdf"); + assert_eq!(part["fileData"]["fileUri"], "https://example.com/doc.pdf"); + } + + #[test] + fn document_base64_translates_to_inline_data() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Document(crate::types::DocumentData { + url: None, + data: Some(vec![0x25, 0x50, 0x44, 0x46]), + media_type: None, + file_name: None, + })], + name: None, + tool_call_id: None, + }; + let contents = translate_messages(&[&msg]); + let part = &contents[0].parts[0]; + assert_eq!(part["inlineData"]["mimeType"], "application/pdf"); + assert!(part["inlineData"]["data"].as_str().is_some()); + } } diff --git a/crates/unified-llm/src/providers/openai.rs b/crates/unified-llm/src/providers/openai.rs index 33c5b9bc9..6a0c4d096 100644 --- a/crates/unified-llm/src/providers/openai.rs +++ b/crates/unified-llm/src/providers/openai.rs @@ -159,6 +159,7 @@ fn map_finish_reason(status: Option<&str>, has_tool_calls: bool) -> FinishReason } /// Translate unified messages to Responses API `input` array format. +#[allow(clippy::too_many_lines)] fn translate_input(messages: &[Message]) -> (Option, Vec) { let mut instructions_parts: Vec = Vec::new(); let mut input: Vec = Vec::new(); @@ -197,6 +198,16 @@ fn translate_input(messages: &[Message]) -> (Option, Vec { + Some(serde_json::json!({"type": "input_text", "text": "[Audio content not supported by this provider]"})) + } + ContentPart::Document(doc) => { + let desc = doc.file_name.as_ref().map_or_else( + || "[Document content not supported by this provider]".to_string(), + |name| format!("[Document '{name}': content type not supported by this provider]"), + ); + Some(serde_json::json!({"type": "input_text", "text": desc})) + } _ => None, }) .collect(); @@ -1079,4 +1090,60 @@ mod tests { assert!(adapter.project_id.is_none()); assert!(adapter.default_headers.is_empty()); } + + #[test] + fn audio_content_produces_text_fallback() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Audio(crate::types::AudioData { + url: Some("https://example.com/audio.wav".to_string()), + data: None, + media_type: None, + })], + name: None, + tool_call_id: None, + }; + let (_, input) = translate_input(&[msg]); + let content = input[0]["content"].as_array().expect("content should be array"); + assert_eq!(content[0]["type"], "input_text"); + assert_eq!(content[0]["text"], "[Audio content not supported by this provider]"); + } + + #[test] + fn document_content_produces_text_fallback_with_filename() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Document(crate::types::DocumentData { + url: Some("https://example.com/doc.pdf".to_string()), + data: None, + media_type: None, + file_name: Some("report.pdf".to_string()), + })], + name: None, + tool_call_id: None, + }; + let (_, input) = translate_input(&[msg]); + let content = input[0]["content"].as_array().expect("content should be array"); + assert_eq!(content[0]["type"], "input_text"); + assert_eq!(content[0]["text"], "[Document 'report.pdf': content type not supported by this provider]"); + } + + #[test] + fn document_content_produces_text_fallback_without_filename() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Document(crate::types::DocumentData { + url: None, + data: Some(vec![1, 2, 3]), + media_type: None, + file_name: None, + })], + name: None, + tool_call_id: None, + }; + let (_, input) = translate_input(&[msg]); + let content = input[0]["content"].as_array().expect("content should be array"); + assert_eq!(content[0]["type"], "input_text"); + assert_eq!(content[0]["text"], "[Document content not supported by this provider]"); + } } diff --git a/crates/unified-llm/src/providers/openai_compatible.rs b/crates/unified-llm/src/providers/openai_compatible.rs index 455b9c797..da27fc96b 100644 --- a/crates/unified-llm/src/providers/openai_compatible.rs +++ b/crates/unified-llm/src/providers/openai_compatible.rs @@ -212,6 +212,29 @@ fn map_finish_reason(reason: Option<&str>) -> FinishReason { } } +/// Build the content string from a message's parts, including fallback text +/// for unsupported content types (Audio, Document). +fn content_text_with_fallbacks(parts: &[ContentPart]) -> String { + let mut segments: Vec = Vec::new(); + for part in parts { + match part { + ContentPart::Text(text) => segments.push(text.clone()), + ContentPart::Audio(_) => { + segments.push("[Audio content not supported by this provider]".to_string()); + } + ContentPart::Document(doc) => { + let desc = doc.file_name.as_ref().map_or_else( + || "[Document content not supported by this provider]".to_string(), + |name| format!("[Document '{name}': content type not supported by this provider]"), + ); + segments.push(desc); + } + _ => {} + } + } + segments.join("") +} + fn translate_messages(messages: &[Message]) -> Vec { messages .iter() @@ -243,7 +266,7 @@ fn translate_messages(messages: &[Message]) -> Vec { } } - let text = msg.text(); + let text = content_text_with_fallbacks(&msg.content); let content = if text.is_empty() { None } else { Some(text) }; let tool_calls = if tool_calls.is_empty() { None @@ -1263,4 +1286,85 @@ mod tests { // Should not crash and body should be unchanged assert_eq!(body["model"], "test"); } + + #[test] + fn audio_content_produces_text_fallback() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Audio(crate::types::AudioData { + url: Some("https://example.com/audio.wav".to_string()), + data: None, + media_type: None, + })], + name: None, + tool_call_id: None, + }; + let translated = translate_messages(&[msg]); + assert_eq!( + translated[0].content.as_deref(), + Some("[Audio content not supported by this provider]") + ); + } + + #[test] + fn document_content_produces_text_fallback_with_filename() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Document(crate::types::DocumentData { + url: Some("https://example.com/doc.pdf".to_string()), + data: None, + media_type: None, + file_name: Some("report.pdf".to_string()), + })], + name: None, + tool_call_id: None, + }; + let translated = translate_messages(&[msg]); + assert_eq!( + translated[0].content.as_deref(), + Some("[Document 'report.pdf': content type not supported by this provider]") + ); + } + + #[test] + fn document_content_produces_text_fallback_without_filename() { + let msg = Message { + role: Role::User, + content: vec![ContentPart::Document(crate::types::DocumentData { + url: None, + data: Some(vec![1, 2, 3]), + media_type: None, + file_name: None, + })], + name: None, + tool_call_id: None, + }; + let translated = translate_messages(&[msg]); + assert_eq!( + translated[0].content.as_deref(), + Some("[Document content not supported by this provider]") + ); + } + + #[test] + fn mixed_text_and_audio_content_concatenates() { + let msg = Message { + role: Role::User, + content: vec![ + ContentPart::text("Check this: "), + ContentPart::Audio(crate::types::AudioData { + url: None, + data: Some(vec![1, 2]), + media_type: None, + }), + ], + name: None, + tool_call_id: None, + }; + let translated = translate_messages(&[msg]); + assert_eq!( + translated[0].content.as_deref(), + Some("Check this: [Audio content not supported by this provider]") + ); + } }