diff --git a/crates/unified-llm/src/error.rs b/crates/unified-llm/src/error.rs index 7aac6453c..29ea84d4e 100644 --- a/crates/unified-llm/src/error.rs +++ b/crates/unified-llm/src/error.rs @@ -80,6 +80,9 @@ pub enum SdkError { #[error("Configuration error: {message}")] Configuration { message: String }, + + #[error("Unsupported tool choice: {message}")] + UnsupportedToolChoice { message: String }, } impl SdkError { @@ -99,7 +102,8 @@ impl SdkError { Self::InvalidToolCall { .. } | Self::NoObjectGenerated { .. } | Self::Abort { .. } - | Self::Configuration { .. } => false, + | Self::Configuration { .. } + | Self::UnsupportedToolChoice { .. } => false, _ => true, } } @@ -148,35 +152,8 @@ pub fn error_from_status_code( raw, }; - // First check message-based classification for ambiguous cases - let lower_msg = detail.message.to_lowercase(); - if lower_msg.contains("not found") || lower_msg.contains("does not exist") { - return SdkError::Provider { - kind: ProviderErrorKind::NotFound, - detail: Box::new(detail), - }; - } - if lower_msg.contains("unauthorized") || lower_msg.contains("invalid key") { - return SdkError::Provider { - kind: ProviderErrorKind::Authentication, - detail: Box::new(detail), - }; - } - if lower_msg.contains("context length") || lower_msg.contains("too many tokens") { - return SdkError::Provider { - kind: ProviderErrorKind::ContextLength, - detail: Box::new(detail), - }; - } - if lower_msg.contains("content filter") || lower_msg.contains("safety") { - return SdkError::Provider { - kind: ProviderErrorKind::ContentFilter, - detail: Box::new(detail), - }; - } - + // Check specific status codes first -- these always map to their designated error types let kind = match status_code { - 400 | 422 => ProviderErrorKind::InvalidRequest, 401 => ProviderErrorKind::Authentication, 403 => ProviderErrorKind::AccessDenied, 404 => ProviderErrorKind::NotFound, @@ -187,7 +164,24 @@ pub fn error_from_status_code( } 413 => ProviderErrorKind::ContextLength, 429 => ProviderErrorKind::RateLimit, - _ => ProviderErrorKind::Server, + 500..=504 => ProviderErrorKind::Server, + // For ambiguous status codes (400, 422, etc.), use message-based classification + _ => { + let lower_msg = detail.message.to_lowercase(); + if lower_msg.contains("not found") || lower_msg.contains("does not exist") { + ProviderErrorKind::NotFound + } else if lower_msg.contains("unauthorized") || lower_msg.contains("invalid key") { + ProviderErrorKind::Authentication + } else if lower_msg.contains("context length") + || lower_msg.contains("too many tokens") + { + ProviderErrorKind::ContextLength + } else if lower_msg.contains("content filter") || lower_msg.contains("safety") { + ProviderErrorKind::ContentFilter + } else { + ProviderErrorKind::InvalidRequest + } + } }; SdkError::Provider { diff --git a/crates/unified-llm/src/provider.rs b/crates/unified-llm/src/provider.rs index 771062683..1e12947c3 100644 --- a/crates/unified-llm/src/provider.rs +++ b/crates/unified-llm/src/provider.rs @@ -1,5 +1,5 @@ use crate::error::SdkError; -use crate::types::{Request, Response, StreamEvent}; +use crate::types::{Request, Response, StreamEvent, ToolChoice}; use futures::Stream; use std::pin::Pin; @@ -34,3 +34,23 @@ pub trait ProviderAdapter: Send + Sync { true } } + +/// Validate that the adapter supports the requested tool choice mode. +/// +/// Returns `Err(SdkError::UnsupportedToolChoice)` if the adapter does not +/// support the given mode. +pub fn validate_tool_choice( + adapter: &dyn ProviderAdapter, + tool_choice: &ToolChoice, +) -> Result<(), SdkError> { + let mode = tool_choice.mode_str(); + if !adapter.supports_tool_choice(mode) { + return Err(SdkError::UnsupportedToolChoice { + message: format!( + "provider '{}' does not support tool_choice mode '{mode}'", + adapter.name() + ), + }); + } + Ok(()) +} diff --git a/crates/unified-llm/src/providers/anthropic.rs b/crates/unified-llm/src/providers/anthropic.rs index d87008cfc..70b087c42 100644 --- a/crates/unified-llm/src/providers/anthropic.rs +++ b/crates/unified-llm/src/providers/anthropic.rs @@ -18,6 +18,7 @@ pub struct Adapter { default_headers: std::collections::HashMap, client: reqwest::Client, request_timeout: std::time::Duration, + stream_read_timeout: std::time::Duration, } impl Adapter { @@ -34,6 +35,7 @@ impl Adapter { default_headers: std::collections::HashMap::new(), client, request_timeout: std::time::Duration::from_secs_f64(timeout.request), + stream_read_timeout: std::time::Duration::from_secs_f64(timeout.stream_read), } } @@ -49,6 +51,17 @@ impl Adapter { self } + #[must_use] + pub fn with_timeout(mut self, timeout: crate::types::AdapterTimeout) -> Self { + self.client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs_f64(timeout.connect)) + .build() + .unwrap_or_default(); + self.request_timeout = std::time::Duration::from_secs_f64(timeout.request); + self.stream_read_timeout = std::time::Duration::from_secs_f64(timeout.stream_read); + self + } + fn messages_url(&self) -> String { format!("{}/messages", self.base_url) } @@ -900,6 +913,7 @@ struct SseReaderState { done: bool, /// When true, `tool_use` events for the synthetic tool are converted to text events. json_schema_mode: bool, + stream_read_timeout: std::time::Duration, } impl SseReaderState { @@ -909,6 +923,7 @@ impl SseReaderState { + 'static, rate_limit: Option, json_schema_mode: bool, + stream_read_timeout: std::time::Duration, ) -> Self { use futures::StreamExt; Self { @@ -918,6 +933,7 @@ impl SseReaderState { pending_events: std::collections::VecDeque::new(), done: false, json_schema_mode, + stream_read_timeout, } } @@ -939,17 +955,17 @@ impl SseReaderState { } // Read more bytes from the stream. - match self.byte_stream.next().await { - Some(Ok(chunk)) => { + match tokio::time::timeout(self.stream_read_timeout, self.byte_stream.next()).await { + Ok(Some(Ok(chunk))) => { let text = String::from_utf8_lossy(&chunk); self.buffer.push_str(&text); } - Some(Err(e)) => { + Ok(Some(Err(e))) => { return SseResult::Error(SdkError::Stream { message: e.to_string(), }); } - None => { + Ok(None) => { self.done = true; // Try one more time to parse any remaining data. if let Some(result) = self.try_parse_event() { @@ -957,6 +973,11 @@ impl SseReaderState { } return SseResult::Done; } + Err(_) => { + return SseResult::Error(SdkError::Stream { + message: "stream read timed out waiting for next event".to_string(), + }); + } } } } @@ -1081,6 +1102,9 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { + if let Some(tc) = &request.tool_choice { + crate::provider::validate_tool_choice(self, tc)?; + } let (_api_request, req_builder) = build_api_request(self, request, false); let (body, headers) = @@ -1141,6 +1165,9 @@ impl ProviderAdapter for Adapter { } async fn stream(&self, request: &Request) -> Result { + if let Some(tc) = &request.tool_choice { + crate::provider::validate_tool_choice(self, tc)?; + } let (_api_request, req_builder) = build_api_request(self, request, true); let http_resp = req_builder.send().await.map_err(|e| SdkError::Network { @@ -1167,9 +1194,10 @@ impl ProviderAdapter for Adapter { let rate_limit = parse_rate_limit_headers(http_resp.headers()); let byte_stream = http_resp.bytes_stream(); let json_schema_mode = uses_json_schema_format(request); + let stream_read_timeout = self.stream_read_timeout; let stream = futures::stream::unfold( - SseReaderState::new(byte_stream, rate_limit, json_schema_mode), + SseReaderState::new(byte_stream, rate_limit, json_schema_mode, stream_read_timeout), |mut state| async move { loop { // Drain any buffered events first. diff --git a/crates/unified-llm/src/providers/gemini.rs b/crates/unified-llm/src/providers/gemini.rs index ace9f86c3..e0c1f55dc 100644 --- a/crates/unified-llm/src/providers/gemini.rs +++ b/crates/unified-llm/src/providers/gemini.rs @@ -8,7 +8,7 @@ use crate::providers::common::{ }; use crate::types::{ ContentPart, FinishReason, Message, Request, Response, ResponseFormat, ResponseFormatType, - Role, StreamEvent, ToolCall, ToolChoice, ToolDefinition, Usage, + Role, StreamEvent, ThinkingData, ToolCall, ToolChoice, ToolDefinition, Usage, }; const DEFAULT_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta"; @@ -20,6 +20,7 @@ pub struct Adapter { default_headers: std::collections::HashMap, client: reqwest::Client, request_timeout: std::time::Duration, + stream_read_timeout: std::time::Duration, } impl Adapter { @@ -36,6 +37,7 @@ impl Adapter { default_headers: std::collections::HashMap::new(), client, request_timeout: std::time::Duration::from_secs_f64(timeout.request), + stream_read_timeout: std::time::Duration::from_secs_f64(timeout.stream_read), } } @@ -50,6 +52,17 @@ impl Adapter { self.default_headers = headers; self } + + #[must_use] + pub fn with_timeout(mut self, timeout: crate::types::AdapterTimeout) -> Self { + self.client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs_f64(timeout.connect)) + .build() + .unwrap_or_default(); + self.request_timeout = std::time::Duration::from_secs_f64(timeout.request); + self.stream_read_timeout = std::time::Duration::from_secs_f64(timeout.stream_read); + self + } } // --- Request types --- @@ -157,6 +170,17 @@ fn map_finish_reason(reason: Option<&str>, has_function_calls: bool) -> FinishRe fn parse_part(part: &serde_json::Value) -> Option { if let Some(text) = part.get("text").and_then(serde_json::Value::as_str) { + let is_thought = part + .get("thought") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false); + if is_thought { + return Some(ContentPart::Thinking(ThinkingData { + text: text.to_string(), + signature: None, + redacted: false, + })); + } return Some(ContentPart::text(text)); } if let Some(fc) = part.get("functionCall") { @@ -537,9 +561,9 @@ async fn send_streaming_request( /// Process a stream of SSE chunks from the Gemini `streamGenerateContent` endpoint /// and yield `StreamEvent` values. -fn process_sse_stream(http_resp: reqwest::Response, model: String, rate_limit: Option) -> StreamEventStream { +fn process_sse_stream(http_resp: reqwest::Response, model: String, rate_limit: Option, stream_read_timeout: std::time::Duration) -> StreamEventStream { Box::pin(stream::unfold( - SseStreamState::new(http_resp, model, rate_limit), + SseStreamState::new(http_resp, model, rate_limit, stream_read_timeout), |mut state| async move { // If we have buffered events, yield them first. if let Some(event) = state.pending_events.pop_front() { @@ -628,6 +652,10 @@ struct SseStreamState { stream_started: bool, /// Whether we have emitted a `TextStart` event. text_started: bool, + /// Whether we are currently inside a reasoning (thought) segment. + reasoning_started: bool, + /// Accumulated thinking text across all chunks. + accumulated_thinking: String, /// Accumulated text across all chunks. accumulated_text: String, /// Accumulated tool calls across all chunks. @@ -642,10 +670,11 @@ struct SseStreamState { finished: bool, /// Rate limit info parsed from HTTP response headers. rate_limit: Option, + stream_read_timeout: std::time::Duration, } impl SseStreamState { - fn new(http_resp: reqwest::Response, model: String, rate_limit: Option) -> Self { + fn new(http_resp: reqwest::Response, model: String, rate_limit: Option, stream_read_timeout: std::time::Duration) -> Self { Self { http_resp, model, @@ -653,6 +682,8 @@ impl SseStreamState { pending_events: std::collections::VecDeque::new(), stream_started: false, text_started: false, + reasoning_started: false, + accumulated_thinking: String::new(), accumulated_text: String::new(), accumulated_tool_calls: Vec::new(), text_id: uuid::Uuid::new_v4().to_string(), @@ -660,6 +691,7 @@ impl SseStreamState { finish_reason_str: None, finished: false, rate_limit, + stream_read_timeout, } } @@ -678,12 +710,12 @@ impl SseStreamState { } // Read more bytes from the HTTP response. - match self.http_resp.chunk().await { - Ok(Some(bytes)) => { + match tokio::time::timeout(self.stream_read_timeout, self.http_resp.chunk()).await { + Ok(Ok(Some(bytes))) => { let text = String::from_utf8_lossy(&bytes); self.line_buffer.push_str(&text); } - Ok(None) => { + Ok(Ok(None)) => { // Stream ended. Return any remaining buffered content. if self.line_buffer.is_empty() { return Ok(None); @@ -695,11 +727,16 @@ impl SseStreamState { } return Ok(Some(line)); } - Err(e) => { + Ok(Err(e)) => { return Err(SdkError::Stream { message: format!("error reading Gemini stream: {e}"), }); } + Err(_) => { + return Err(SdkError::Stream { + message: "stream read timed out waiting for next event".to_string(), + }); + } } } } @@ -723,16 +760,40 @@ impl SseStreamState { }; for part in parts { + let is_thought = part + .get("thought") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false); + if let Some(text) = part.get("text").and_then(serde_json::Value::as_str) { - if !self.text_started { - self.text_started = true; - self.pending_events.push_back(StreamEvent::TextStart { - text_id: Some(self.text_id.clone()), - }); + if is_thought { + if !self.reasoning_started { + self.reasoning_started = true; + self.pending_events + .push_back(StreamEvent::ReasoningStart); + } + self.accumulated_thinking.push_str(text); + self.pending_events + .push_back(StreamEvent::ReasoningDelta { + delta: text.to_string(), + }); + } else { + // Transition from reasoning to text: close reasoning segment. + if self.reasoning_started { + self.reasoning_started = false; + self.pending_events + .push_back(StreamEvent::ReasoningEnd); + } + if !self.text_started { + self.text_started = true; + self.pending_events.push_back(StreamEvent::TextStart { + text_id: Some(self.text_id.clone()), + }); + } + self.accumulated_text.push_str(text); + self.pending_events + .push_back(StreamEvent::text_delta(text, Some(self.text_id.clone()))); } - self.accumulated_text.push_str(text); - self.pending_events - .push_back(StreamEvent::text_delta(text, Some(self.text_id.clone()))); } else if let Some(fc) = part.get("functionCall") { let name = fc .get("name") @@ -765,10 +826,17 @@ impl SseStreamState { .and_then(|c| c.finish_reason.as_ref()) .is_some(); - if has_finish_reason && self.text_started { - self.pending_events.push_back(StreamEvent::TextEnd { - text_id: Some(self.text_id.clone()), - }); + if has_finish_reason { + if self.reasoning_started { + self.reasoning_started = false; + self.pending_events + .push_back(StreamEvent::ReasoningEnd); + } + if self.text_started { + self.pending_events.push_back(StreamEvent::TextEnd { + text_id: Some(self.text_id.clone()), + }); + } } } @@ -779,6 +847,13 @@ impl SseStreamState { map_finish_reason(self.finish_reason_str.as_deref(), has_tool_calls); let mut content_parts: Vec = Vec::new(); + if !self.accumulated_thinking.is_empty() { + content_parts.push(ContentPart::Thinking(ThinkingData { + text: self.accumulated_thinking.clone(), + signature: None, + redacted: false, + })); + } if !self.accumulated_text.is_empty() { content_parts.push(ContentPart::text(&self.accumulated_text)); } @@ -815,6 +890,9 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { + if let Some(tc) = &request.tool_choice { + crate::provider::validate_tool_choice(self, tc)?; + } let api_body = build_api_request(request); let url = format!( @@ -880,6 +958,9 @@ impl ProviderAdapter for Adapter { } async fn stream(&self, request: &Request) -> Result { + if let Some(tc) = &request.tool_choice { + crate::provider::validate_tool_choice(self, tc)?; + } let api_body = build_api_request(request); let url = format!( @@ -894,7 +975,7 @@ impl ProviderAdapter for Adapter { let http_resp = send_streaming_request(req.json(&api_body)).await?; let rate_limit = parse_rate_limit_headers(http_resp.headers()); - Ok(process_sse_stream(http_resp, request.model.clone(), rate_limit)) + Ok(process_sse_stream(http_resp, request.model.clone(), rate_limit, self.stream_read_timeout)) } } @@ -1136,4 +1217,38 @@ mod tests { let err = gemini_error(500, "internal".into(), None, None, None); assert!(matches!(err, SdkError::Provider { kind: ProviderErrorKind::Server, .. })); } + + #[test] + fn parse_part_handles_thought_text() { + let part = serde_json::json!({"text": "Let me think about this...", "thought": true}); + let result = parse_part(&part).expect("should parse thought part"); + match result { + ContentPart::Thinking(td) => { + assert_eq!(td.text, "Let me think about this..."); + assert!(td.signature.is_none()); + assert!(!td.redacted); + } + other => panic!("expected Thinking, got {other:?}"), + } + } + + #[test] + fn parse_part_text_without_thought_flag() { + let part = serde_json::json!({"text": "Hello world"}); + let result = parse_part(&part).expect("should parse text part"); + match result { + ContentPart::Text(text) => assert_eq!(text, "Hello world"), + other => panic!("expected Text, got {other:?}"), + } + } + + #[test] + fn parse_part_thought_false_is_regular_text() { + let part = serde_json::json!({"text": "Regular text", "thought": false}); + let result = parse_part(&part).expect("should parse text part"); + match result { + ContentPart::Text(text) => assert_eq!(text, "Regular text"), + other => panic!("expected Text, got {other:?}"), + } + } } diff --git a/crates/unified-llm/src/providers/openai.rs b/crates/unified-llm/src/providers/openai.rs index 29163a376..0340e151e 100644 --- a/crates/unified-llm/src/providers/openai.rs +++ b/crates/unified-llm/src/providers/openai.rs @@ -23,6 +23,7 @@ pub struct Adapter { default_headers: std::collections::HashMap, client: reqwest::Client, request_timeout: std::time::Duration, + stream_read_timeout: std::time::Duration, } impl Adapter { @@ -41,6 +42,7 @@ impl Adapter { default_headers: std::collections::HashMap::new(), client, request_timeout: std::time::Duration::from_secs_f64(timeout.request), + stream_read_timeout: std::time::Duration::from_secs_f64(timeout.stream_read), } } @@ -68,6 +70,17 @@ impl Adapter { self } + #[must_use] + pub fn with_timeout(mut self, timeout: crate::types::AdapterTimeout) -> Self { + self.client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs_f64(timeout.connect)) + .build() + .unwrap_or_default(); + self.request_timeout = std::time::Duration::from_secs_f64(timeout.request); + self.stream_read_timeout = std::time::Duration::from_secs_f64(timeout.stream_read); + self + } + /// Build a `reqwest::RequestBuilder` with default headers, org/project headers, and auth. fn build_request(&self, url: &str) -> reqwest::RequestBuilder { let mut req = self.client.post(url); @@ -109,6 +122,8 @@ struct ApiRequest { #[serde(skip_serializing_if = "Option::is_none")] text: Option, #[serde(skip_serializing_if = "Option::is_none")] + stop: Option>, + #[serde(skip_serializing_if = "Option::is_none")] metadata: Option>, #[serde(skip_serializing_if = "std::ops::Not::not")] stream: bool, @@ -352,6 +367,7 @@ fn build_api_request(request: &Request, stream: bool) -> ApiRequest { tool_choice, reasoning, text, + stop: request.stop_sequences.clone(), metadata: request.metadata.clone(), stream, } @@ -449,6 +465,7 @@ struct SseStreamState { emitted_text_start: bool, raw_response: Option, rate_limit: Option, + stream_read_timeout: std::time::Duration, } /// Extract complete SSE messages from the buffer. @@ -513,17 +530,17 @@ async fn process_next_sse_events( return Ok(dispatch_sse_messages(state, messages)); } - match state.byte_stream.next().await { - Some(Ok(bytes)) => { + match tokio::time::timeout(state.stream_read_timeout, state.byte_stream.next()).await { + Ok(Some(Ok(bytes))) => { let text = String::from_utf8_lossy(&bytes); state.buffer.push_str(&text); } - Some(Err(e)) => { + Ok(Some(Err(e))) => { return Err(SdkError::Stream { message: e.to_string(), }); } - None => { + Ok(None) => { // Stream ended. Process any remaining data in the buffer. if !state.buffer.is_empty() { state.buffer.push_str("\n\n"); @@ -532,6 +549,11 @@ async fn process_next_sse_events( } return Ok(vec![]); } + Err(_) => { + return Err(SdkError::Stream { + message: "stream read timed out waiting for next event".to_string(), + }); + } } } } @@ -823,6 +845,9 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { + if let Some(tc) = &request.tool_choice { + crate::provider::validate_tool_choice(self, tc)?; + } let request_body = build_request_body(request, false); let url = format!("{}/responses", self.base_url); @@ -880,6 +905,9 @@ impl ProviderAdapter for Adapter { } async fn stream(&self, request: &Request) -> Result { + if let Some(tc) = &request.tool_choice { + crate::provider::validate_tool_choice(self, tc)?; + } let request_body = build_request_body(request, true); let url = format!("{}/responses", self.base_url); @@ -913,6 +941,7 @@ impl ProviderAdapter for Adapter { let rate_limit = parse_rate_limit_headers(http_resp.headers()); let byte_stream = http_resp.bytes_stream(); + let stream_read_timeout = self.stream_read_timeout; let state = SseStreamState { byte_stream: Box::pin(byte_stream), buffer: String::new(), @@ -927,6 +956,7 @@ impl ProviderAdapter for Adapter { emitted_text_start: false, raw_response: None, rate_limit, + stream_read_timeout, }; let stream = futures::stream::unfold(state, |mut state| async move { @@ -1150,4 +1180,24 @@ mod tests { assert_eq!(content[0]["type"], "input_text"); assert_eq!(content[0]["text"], "[Document content not supported by this provider]"); } + + #[test] + fn build_request_body_includes_stop_sequences() { + let mut request = minimal_request(); + request.stop_sequences = Some(vec!["END".to_string(), "STOP".to_string()]); + + let body = build_request_body(&request, false); + let stop = body.get("stop").expect("stop should be present"); + let arr = stop.as_array().expect("stop should be an array"); + assert_eq!(arr.len(), 2); + assert_eq!(arr[0], "END"); + assert_eq!(arr[1], "STOP"); + } + + #[test] + fn build_request_body_omits_stop_when_none() { + let request = minimal_request(); + let body = build_request_body(&request, false); + assert!(body.get("stop").is_none()); + } } diff --git a/crates/unified-llm/src/providers/openai_compatible.rs b/crates/unified-llm/src/providers/openai_compatible.rs index b6a470601..88ca70c20 100644 --- a/crates/unified-llm/src/providers/openai_compatible.rs +++ b/crates/unified-llm/src/providers/openai_compatible.rs @@ -24,6 +24,7 @@ pub struct Adapter { default_headers: std::collections::HashMap, client: reqwest::Client, request_timeout: std::time::Duration, + stream_read_timeout: std::time::Duration, } impl Adapter { @@ -41,6 +42,7 @@ impl Adapter { default_headers: std::collections::HashMap::new(), client, request_timeout: std::time::Duration::from_secs_f64(timeout.request), + stream_read_timeout: std::time::Duration::from_secs_f64(timeout.stream_read), } } @@ -56,6 +58,17 @@ impl Adapter { self } + #[must_use] + pub fn with_timeout(mut self, timeout: crate::types::AdapterTimeout) -> Self { + self.client = reqwest::Client::builder() + .connect_timeout(std::time::Duration::from_secs_f64(timeout.connect)) + .build() + .unwrap_or_default(); + self.request_timeout = std::time::Duration::from_secs_f64(timeout.request); + self.stream_read_timeout = std::time::Duration::from_secs_f64(timeout.stream_read); + self + } + /// Build a `reqwest::RequestBuilder` with default headers and auth. fn build_request(&self, url: &str) -> reqwest::RequestBuilder { let mut req = self.client.post(url); @@ -395,6 +408,9 @@ impl ProviderAdapter for Adapter { } async fn complete(&self, request: &Request) -> Result { + if let Some(tc) = &request.tool_choice { + crate::provider::validate_tool_choice(self, tc)?; + } let api_body = build_api_request(request, None, &self.provider_name); let url = format!("{}/chat/completions", self.base_url); @@ -467,6 +483,9 @@ impl ProviderAdapter for Adapter { } async fn stream(&self, request: &Request) -> Result { + if let Some(tc) = &request.tool_choice { + crate::provider::validate_tool_choice(self, tc)?; + } let api_body = build_api_request(request, Some(true), &self.provider_name); let url = format!("{}/chat/completions", self.base_url); @@ -502,9 +521,10 @@ impl ProviderAdapter for Adapter { let provider_name = self.provider_name.clone(); let model = request.model.clone(); let rate_limit = parse_rate_limit_headers(http_resp.headers()); + let stream_read_timeout = self.stream_read_timeout; let stream = futures::stream::unfold( - StreamState::new(http_resp, provider_name, model, rate_limit), + StreamState::new(http_resp, provider_name, model, rate_limit, stream_read_timeout), |mut state| async move { loop { let line = match state.next_line().await { @@ -599,6 +619,7 @@ struct StreamState { text_started: bool, done: bool, rate_limit: Option, + stream_read_timeout: std::time::Duration, } impl StreamState { @@ -607,6 +628,7 @@ impl StreamState { provider_name: String, model: String, rate_limit: Option, + stream_read_timeout: std::time::Duration, ) -> Self { Self { response, @@ -622,6 +644,7 @@ impl StreamState { text_started: false, done: false, rate_limit, + stream_read_timeout, } } @@ -638,12 +661,12 @@ impl StreamState { return Ok(Some(line)); } - match self.response.chunk().await { - Ok(Some(bytes)) => { + match tokio::time::timeout(self.stream_read_timeout, self.response.chunk()).await { + Ok(Ok(Some(bytes))) => { let text = String::from_utf8_lossy(&bytes); self.buffer.push_str(&text); } - Ok(None) => { + Ok(Ok(None)) => { self.done = true; if self.buffer.is_empty() { return Ok(None); @@ -651,11 +674,16 @@ impl StreamState { let remaining = std::mem::take(&mut self.buffer); return Ok(Some(remaining)); } - Err(e) => { + Ok(Err(e)) => { return Err(SdkError::Stream { message: e.to_string(), }); } + Err(_) => { + return Err(SdkError::Stream { + message: "stream read timed out waiting for next event".to_string(), + }); + } } } } @@ -888,7 +916,7 @@ mod tests { .body("") .unwrap(), ); - let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None); + let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None, std::time::Duration::from_secs(30)); // First text chunk should emit TextStart + TextDelta. let chunk1: StreamChunk = serde_json::from_str( @@ -918,7 +946,7 @@ mod tests { .body("") .unwrap(), ); - let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None); + let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None, std::time::Duration::from_secs(30)); // First tool call chunk (has id and name) -> ToolCallStart. let chunk1: StreamChunk = serde_json::from_str( @@ -947,7 +975,7 @@ mod tests { .body("") .unwrap(), ); - let mut state = StreamState::new(http_resp, "test-provider".into(), "test-model".into(), None); + let mut state = StreamState::new(http_resp, "test-provider".into(), "test-model".into(), None, std::time::Duration::from_secs(30)); state.response_id = "resp-1".into(); state.response_model = "gpt-4".into(); state.accumulated_text = "Hello world".into(); @@ -989,7 +1017,7 @@ mod tests { .body("") .unwrap(), ); - let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None); + let mut state = StreamState::new(http_resp, "test".into(), "model".into(), None, std::time::Duration::from_secs(30)); state.response_id = "resp-1".into(); state.tool_calls.push(AccumulatedToolCall { id: "call_1".into(), @@ -1032,7 +1060,7 @@ mod tests { .body("") .unwrap(), ); - let mut state = StreamState::new(http_resp, "test".into(), "fallback-model".into(), None); + let mut state = StreamState::new(http_resp, "test".into(), "fallback-model".into(), None, std::time::Duration::from_secs(30)); // response_model is empty, so finish_events should use the request model. let events = state.finish_events(); match &events[0] { diff --git a/crates/unified-llm/src/types.rs b/crates/unified-llm/src/types.rs index 0145b428d..c47523c65 100644 --- a/crates/unified-llm/src/types.rs +++ b/crates/unified-llm/src/types.rs @@ -356,6 +356,7 @@ impl std::ops::Add for Usage { #[serde(rename_all = "snake_case")] pub enum ResponseFormatType { Text, + #[serde(rename = "json")] JsonObject, JsonSchema, } @@ -433,6 +434,17 @@ impl ToolChoice { tool_name: name.into(), } } + + /// Return the mode string used by `ProviderAdapter::supports_tool_choice`. + #[must_use] + pub fn mode_str(&self) -> &'static str { + match self { + Self::Auto => "auto", + Self::None => "none", + Self::Required => "required", + Self::Named { .. } => "named", + } + } } // --- 3.7 Response --- diff --git a/docs/specs/unified-llm-spec.md b/docs/specs/unified-llm-spec.md index 5676e8870..046d1e176 100644 --- a/docs/specs/unified-llm-spec.md +++ b/docs/specs/unified-llm-spec.md @@ -786,6 +786,7 @@ ENUM StreamEventType: TOOL_CALL_START -- A tool call has begun. Includes tool name and call ID. TOOL_CALL_DELTA -- Incremental tool call arguments (partial JSON). TOOL_CALL_END -- Tool call is fully formed and ready for execution. + STEP_FINISH -- A tool execution step completed. Includes tool calls and results from this step. FINISH -- Generation complete. Includes finish_reason, usage, response. ERROR -- An error occurred during streaming. PROVIDER_EVENT -- Raw provider event not mapped to the unified model.