mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-08 03:10:26 +00:00
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:
parent
8a19b7be2a
commit
978a477d1e
8 changed files with 319 additions and 71 deletions
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:?}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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] {
|
||||
|
|
|
|||
|
|
@ -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 ---
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue