Fix 8 spec compliance gaps in unified-llm

- Enforce stream_read timeout (30s default) in all 4 providers' streaming code
- Add with_timeout() builder method to all adapter constructors
- Fix ResponseFormatType::JsonObject to serialize as "json" per spec
- Add STEP_FINISH to StreamEventType enum in spec doc
- Add UnsupportedToolChoice error and enforce in all adapters via validate_tool_choice()
- Fix error classification to check status code before message content
- Add stop_sequences support to OpenAI Responses API adapter
- Handle Gemini thought parts (thought: true) in both complete and streaming paths

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-02-20 12:41:41 -04:00
parent 8a19b7be2a
commit 978a477d1e
8 changed files with 319 additions and 71 deletions

View file

@ -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 {

View file

@ -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(())
}

View file

@ -18,6 +18,7 @@ pub struct Adapter {
default_headers: std::collections::HashMap<String, String>,
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<crate::types::RateLimitInfo>,
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<Response, SdkError> {
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<StreamEventStream, SdkError> {
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.

View file

@ -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<String, String>,
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<ContentPart> {
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<crate::types::RateLimitInfo>) -> StreamEventStream {
fn process_sse_stream(http_resp: reqwest::Response, model: String, rate_limit: Option<crate::types::RateLimitInfo>, 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<crate::types::RateLimitInfo>,
stream_read_timeout: std::time::Duration,
}
impl SseStreamState {
fn new(http_resp: reqwest::Response, model: String, rate_limit: Option<crate::types::RateLimitInfo>) -> Self {
fn new(http_resp: reqwest::Response, model: String, rate_limit: Option<crate::types::RateLimitInfo>, 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<ContentPart> = 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<Response, SdkError> {
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<StreamEventStream, SdkError> {
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:?}"),
}
}
}

View file

@ -23,6 +23,7 @@ pub struct Adapter {
default_headers: std::collections::HashMap<String, String>,
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_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
stop: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
metadata: Option<std::collections::HashMap<String, String>>,
#[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<serde_json::Value>,
rate_limit: Option<crate::types::RateLimitInfo>,
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<Response, SdkError> {
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<StreamEventStream, SdkError> {
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());
}
}

View file

@ -24,6 +24,7 @@ pub struct Adapter {
default_headers: std::collections::HashMap<String, String>,
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<Response, SdkError> {
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<StreamEventStream, SdkError> {
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<crate::types::RateLimitInfo>,
stream_read_timeout: std::time::Duration,
}
impl StreamState {
@ -607,6 +628,7 @@ impl StreamState {
provider_name: String,
model: String,
rate_limit: Option<crate::types::RateLimitInfo>,
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] {

View file

@ -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 ---

View file

@ -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.