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 <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-02-20 11:51:26 -04:00
parent a4978812e7
commit 5a76551847
5 changed files with 919 additions and 81 deletions

View file

@ -595,37 +595,19 @@ async fn stream_with_tool_loop(params: GenerateParams) -> Result<StreamEventStre
let (tx, rx) = tokio::sync::mpsc::channel::<Result<StreamEvent, SdkError>>(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<StepResult> = 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(&params, &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<StreamEventStre
}
}
if let Ok(event) = &item {
accumulator.process(event);
} else {
let _ = tx.send(item).await;
let request = build_request(&params, &messages, tool_definitions.as_deref());
// Retry initial connection (Section 6.6), with optional per_step timeout
let stream_connect = retry(&retry_policy, || {
let c = client.clone();
let r = request.clone();
async move { c.stream(&r).await }
});
let stream_result =
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, stream_connect)
.await
.unwrap_or_else(|_| Err(SdkError::RequestTimeout {
message: format!("Per-step timeout of {per_step}s exceeded"),
}))
} else {
stream_connect.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;
}
// 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<dyn ProviderAdapter> = Arc::new(StreamingToolCallMockProvider {
call_count: call_count.clone(),
});
let mut providers: HashMap<String, Arc<dyn ProviderAdapter>> = 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<AtomicU32>,
failures: u32,
}
#[async_trait::async_trait]
impl ProviderAdapter for FailThenStreamProvider {
fn name(&self) -> &str {
"mock"
}
async fn complete(&self, _request: &Request) -> Result<Response, SdkError> {
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<StreamEventStream, SdkError> {
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<dyn ProviderAdapter> = Arc::new(FailThenStreamProvider {
call_count: call_count.clone(),
failures: 2, // fail twice, succeed on third
});
let mut providers: HashMap<String, Arc<dyn ProviderAdapter>> = 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<Response, SdkError> {
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<StreamEventStream, SdkError> {
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<dyn ProviderAdapter> = Arc::new(SlowStreamProvider {
delay: std::time::Duration::from_secs(5),
});
let mut providers: HashMap<String, Arc<dyn ProviderAdapter>> = 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<AtomicU32>,
}
#[async_trait::async_trait]
impl ProviderAdapter for SlowToolCallStreamProvider {
fn name(&self) -> &str {
"mock"
}
async fn complete(&self, _request: &Request) -> Result<Response, SdkError> {
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<StreamEventStream, SdkError> {
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<dyn ProviderAdapter> = Arc::new(SlowToolCallStreamProvider {
call_count: call_count.clone(),
});
let mut providers: HashMap<String, Arc<dyn ProviderAdapter>> = 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");
}
}

View file

@ -234,7 +234,29 @@ fn content_part_to_api(part: &ContentPart) -> Option<serde_json::Value> {
})
}
}
_ => 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]");
}
}

View file

@ -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<Content> {
let id_to_name = build_tool_call_id_to_name(messages);
let mut contents: Vec<Content> = Vec::new();
@ -244,6 +245,50 @@ fn translate_messages(messages: &[&Message]) -> Vec<Content> {
},
)
}
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());
}
}

View file

@ -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<String>, Vec<serde_json::Value>) {
let mut instructions_parts: Vec<String> = Vec::new();
let mut input: Vec<serde_json::Value> = Vec::new();
@ -197,6 +198,16 @@ fn translate_input(messages: &[Message]) -> (Option<String>, Vec<serde_json::Val
},
)
}
ContentPart::Audio(_) => {
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]");
}
}

View file

@ -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<String> = 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<ChatMessage> {
messages
.iter()
@ -243,7 +266,7 @@ fn translate_messages(messages: &[Message]) -> Vec<ChatMessage> {
}
}
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]")
);
}
}