use crate::message::{ContentBlock, Message, ToolCall};
use crate::provider::{
ChatRequest, ChatResponse, FinishReason, ModelOptions, Provider, ProviderError, StreamEvent,
TimeoutStage, Usage,
};
use crate::tool::ToolSchema;
use async_trait::async_trait;
use base64::Engine;
use futures::stream::{BoxStream, Stream, StreamExt, unfold};
use serde::Serialize;
use std::collections::BTreeMap;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
#[derive(Clone)]
pub struct OpenAiProvider {
base_url: String,
api_key: String,
model: String,
client: reqwest::Client,
request_timeout: Duration,
idle_timeout: Duration,
stream_timeout: Duration,
structured_mode: StructuredOutputMode,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum StructuredOutputMode {
#[default]
Native,
JsonObject,
Off,
}
impl std::fmt::Debug for OpenAiProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OpenAiProvider")
.field("base_url", &self.base_url)
.field("api_key", &"***")
.field("model", &self.model)
.field("request_timeout", &self.request_timeout)
.field("idle_timeout", &self.idle_timeout)
.field("stream_timeout", &self.stream_timeout)
.field("structured_mode", &self.structured_mode)
.finish_non_exhaustive()
}
}
impl OpenAiProvider {
pub fn new(
base_url: impl Into<String>,
api_key: impl Into<String>,
model: impl Into<String>,
) -> Self {
Self {
base_url: base_url.into().trim_end_matches('/').to_string(),
api_key: api_key.into(),
model: model.into(),
client: build_client(Duration::from_secs(30)),
request_timeout: Duration::from_secs(600),
idle_timeout: Duration::from_secs(60),
stream_timeout: Duration::from_secs(1800),
structured_mode: StructuredOutputMode::default(),
}
}
pub fn with_connect_timeout(mut self, timeout: Duration) -> Self {
self.client = build_client(timeout);
self
}
pub fn with_request_timeout(mut self, timeout: Duration) -> Self {
self.request_timeout = timeout;
self
}
pub fn with_idle_timeout(mut self, timeout: Duration) -> Self {
self.idle_timeout = timeout;
self
}
pub fn with_stream_timeout(mut self, timeout: Duration) -> Self {
self.stream_timeout = timeout;
self
}
pub fn with_structured_output_mode(mut self, mode: StructuredOutputMode) -> Self {
self.structured_mode = mode;
self
}
fn wire_response_format(
&self,
schema: Option<&serde_json::Value>,
) -> Option<OpenAiResponseFormat> {
match (self.structured_mode, schema) {
(StructuredOutputMode::Off, _) | (_, None) => None,
(StructuredOutputMode::Native, Some(schema)) => {
Some(OpenAiResponseFormat::from_schema(schema))
}
(StructuredOutputMode::JsonObject, Some(_)) => {
Some(OpenAiResponseFormat::json_object())
}
}
}
}
fn build_client(connect_timeout: Duration) -> reqwest::Client {
reqwest::Client::builder()
.connect_timeout(connect_timeout)
.build()
.expect("reqwest Client::builder cannot fail without TLS config")
}
#[async_trait]
impl Provider for OpenAiProvider {
async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, ProviderError> {
let wire = OpenAiChatRequest {
model: self.model.clone(),
messages: request
.messages
.iter()
.map(OpenAiMessage::from_message)
.collect(),
tools: request
.tools
.iter()
.map(OpenAiToolDef::from_schema)
.collect(),
stream: None,
stream_options: None,
temperature: request.options.temperature,
max_tokens: request.options.max_tokens,
response_format: self.wire_response_format(request.options.structured.as_ref()),
extra: wire_extra(&request.options),
};
let url = format!("{}/chat/completions", self.base_url);
let mut req = self.client.post(url).json(&wire);
if !self.api_key.is_empty() {
req = req.bearer_auth(&self.api_key);
}
req = req.timeout(self.request_timeout);
let resp = req.send().await.map_err(|e| {
if e.is_timeout() {
ProviderError::Timeout(TimeoutStage::Request)
} else {
map_network_error(e)
}
})?;
let status = resp.status();
let headers = resp.headers().clone();
let body = read_response_body(resp).await?;
if !status.is_success() {
return Err(map_status_error(status.as_u16(), &headers, &body));
}
parse_response(&body)
}
async fn stream_chat(
&self,
request: ChatRequest,
) -> Result<BoxStream<'static, Result<StreamEvent, ProviderError>>, ProviderError> {
let wire = OpenAiChatRequest {
model: self.model.clone(),
messages: request
.messages
.iter()
.map(OpenAiMessage::from_message)
.collect(),
tools: request
.tools
.iter()
.map(OpenAiToolDef::from_schema)
.collect(),
stream: Some(true),
stream_options: Some(StreamOptions {
include_usage: true,
}),
temperature: request.options.temperature,
max_tokens: request.options.max_tokens,
response_format: self.wire_response_format(request.options.structured.as_ref()),
extra: wire_extra(&request.options),
};
let url = format!("{}/chat/completions", self.base_url);
let mut req = self.client.post(url).json(&wire);
if !self.api_key.is_empty() {
req = req.bearer_auth(&self.api_key);
}
let resp = match tokio::time::timeout(self.request_timeout, req.send()).await {
Ok(result) => result.map_err(map_network_error)?,
Err(_) => return Err(ProviderError::Timeout(TimeoutStage::Request)),
};
let status = resp.status();
let headers = resp.headers().clone();
if !status.is_success() {
let body = read_response_body(resp).await?;
return Err(map_status_error(status.as_u16(), &headers, &body));
}
let aggregator = Arc::new(Mutex::new(ToolCallAggregator::default()));
Ok(Box::pin(sse_stream(
resp.bytes_stream(),
self.idle_timeout,
self.stream_timeout,
aggregator,
)))
}
}
async fn read_response_body(resp: reqwest::Response) -> Result<String, ProviderError> {
let result = tokio::time::timeout(RESPONSE_BODY_TIMEOUT, async {
if let Some(len) = resp.content_length()
&& len > MAX_RESPONSE_BODY as u64
{
return Err(limit_error("response body exceeds size limit"));
}
let mut stream = resp.bytes_stream();
let mut buf = Vec::new();
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(map_network_error)?;
if buf.len() + chunk.len() > MAX_RESPONSE_BODY {
return Err(limit_error("response body exceeds size limit"));
}
buf.extend_from_slice(&chunk);
}
Ok(String::from_utf8_lossy(&buf).into_owned())
})
.await;
match result {
Ok(r) => r,
Err(_) => Err(ProviderError::Timeout(TimeoutStage::ResponseBody)),
}
}
const RESPONSE_BODY_TIMEOUT: Duration = Duration::from_secs(60);
fn sse_stream(
chunks: impl Stream<Item = Result<impl AsRef<[u8]>, reqwest::Error>> + Unpin,
idle_timeout: Duration,
stream_timeout: Duration,
aggregator: Arc<Mutex<ToolCallAggregator>>,
) -> impl Stream<Item = Result<StreamEvent, ProviderError>> {
let tail = aggregator.clone();
sse_lines(chunks, idle_timeout, stream_timeout)
.map(move |line| match line {
Err(e) => {
let mut agg = aggregator.lock().expect("SSE aggregator lock poisoned");
if agg.finished {
return Vec::new();
}
agg.set_errored();
vec![Err(e)]
}
Ok(line) => {
let mut agg = aggregator.lock().expect("SSE aggregator lock poisoned");
if agg.finished {
Vec::new()
} else {
parse_sse_line(&line, &mut agg)
}
}
})
.flat_map(futures::stream::iter)
.chain(
futures::stream::once(async move {
tail.lock()
.expect("SSE aggregator lock poisoned")
.flush_done()
})
.flat_map(futures::stream::iter),
)
}
fn sse_lines(
chunks: impl Stream<Item = Result<impl AsRef<[u8]>, reqwest::Error>> + Unpin,
idle_timeout: Duration,
stream_timeout: Duration,
) -> impl Stream<Item = Result<String, ProviderError>> {
let deadline = Instant::now() + stream_timeout;
unfold(
(chunks, Vec::new(), false, deadline),
move |(mut chunks, mut buf, terminated, deadline)| async move {
if terminated {
return None;
}
loop {
if let Some(line) = pop_line(&mut buf) {
return Some((Ok(line), (chunks, buf, false, deadline)));
}
if Instant::now() >= deadline {
return Some((
Err(ProviderError::Timeout(TimeoutStage::StreamTotal)),
(chunks, buf, true, deadline),
));
}
let next = match tokio::time::timeout(idle_timeout, chunks.next()).await {
Ok(next) => next,
Err(_) => {
return Some((
Err(ProviderError::Timeout(TimeoutStage::Idle)),
(chunks, buf, true, deadline),
));
}
};
match next {
Some(Ok(bytes)) => {
if buf.len() + bytes.as_ref().len() > MAX_SSE_LINE {
return Some((
Err(ProviderError::Api {
status: 0,
message: format!(
"stream line exceeds size limit ({} bytes)",
MAX_SSE_LINE
),
}),
(chunks, buf, true, deadline),
));
}
buf.extend_from_slice(bytes.as_ref());
}
Some(Err(e)) => {
return Some((Err(map_network_error(e)), (chunks, buf, true, deadline)));
}
None => {
if buf.is_empty() {
return None;
}
return Some((
Ok(to_line(std::mem::take(&mut buf))),
(chunks, buf, false, deadline),
));
}
}
}
},
)
}
const MAX_SSE_LINE: usize = 1 << 20;
const MAX_RESPONSE_BODY: usize = 16 << 20;
fn pop_line(buf: &mut Vec<u8>) -> Option<String> {
let pos = buf.iter().position(|&b| b == b'\n')?;
let rest = buf.split_off(pos + 1);
let line = std::mem::replace(buf, rest);
Some(to_line(line))
}
fn to_line(mut bytes: Vec<u8>) -> String {
while matches!(bytes.last(), Some(b'\r' | b'\n')) {
bytes.pop();
}
String::from_utf8_lossy(&bytes).into_owned()
}
fn parse_sse_line(
line: &str,
tool_calls: &mut ToolCallAggregator,
) -> Vec<Result<StreamEvent, ProviderError>> {
let mut events = Vec::new();
if line.is_empty() || line.starts_with(':') {
return events;
}
let Some(data) = line.strip_prefix("data:") else {
return events;
};
let data = data.trim_start();
if data == "[DONE]" {
let events = tool_calls.flush_done();
tool_calls.finished = true;
return events;
}
let chunk = match serde_json::from_str::<OpenAiStreamChunk>(data) {
Ok(chunk) => chunk,
Err(_) => {
return line_error(
tool_calls,
ProviderError::Api {
status: 0,
message: format!("invalid stream event: {}", extract_error_message(data)),
},
);
}
};
if let Some(error) = chunk.error {
return line_error(
tool_calls,
ProviderError::Api {
status: 0,
message: format!("provider stream error: {}", error.message),
},
);
}
if let Some(usage) = chunk.usage {
tool_calls.push_usage(usage);
}
let Some(choice) = chunk.choices.into_iter().next() else {
return events;
};
if let Some(content) = choice.delta.content
&& !tool_calls.errored
{
events.push(Ok(StreamEvent::Delta(content)));
}
if let Some(reasoning) = choice.delta.reasoning_content
&& !tool_calls.errored
{
events.push(Ok(StreamEvent::Reasoning(reasoning)));
}
if let Some(calls) = choice.delta.tool_calls {
for call in calls {
if let Err(e) = tool_calls.push_chunk(call) {
return line_error(tool_calls, e);
}
}
}
if let Some(finish_reason) = choice.finish_reason {
if !tool_calls.errored {
for call in tool_calls.take_all() {
events.push(Ok(StreamEvent::ToolCall {
id: call.id,
name: call.name,
arguments: call.arguments,
}));
}
}
tool_calls.set_done_reason(map_finish_reason(&finish_reason));
}
events
}
fn map_finish_reason(reason: &str) -> FinishReason {
match reason {
"stop" | "tool_calls" => FinishReason::Stop,
"length" => FinishReason::Length,
other => FinishReason::Other(other.to_string()),
}
}
#[derive(Default)]
struct ToolCallAggregator {
calls: Vec<AccumulatedCall>,
done_reason: Option<FinishReason>,
usage: Option<Usage>,
errored: bool,
finished: bool,
}
const MAX_TOOL_CALLS: usize = 64;
const MAX_ACCUMULATED_TEXT: usize = 1 << 20;
fn limit_error(message: &str) -> ProviderError {
ProviderError::Api {
status: 0,
message: message.into(),
}
}
fn line_error(
tool_calls: &mut ToolCallAggregator,
e: ProviderError,
) -> Vec<Result<StreamEvent, ProviderError>> {
tool_calls.set_errored();
vec![Err(e)]
}
struct AccumulatedCall {
index: usize,
id: String,
name: String,
arguments: String,
}
impl ToolCallAggregator {
fn push_chunk(&mut self, chunk: OpenAiToolCallChunk) -> Result<(), ProviderError> {
let name = chunk
.function
.as_ref()
.and_then(|f| f.name.clone())
.unwrap_or_default();
let arguments = chunk
.function
.as_ref()
.and_then(|f| f.arguments.clone())
.unwrap_or_default();
if let Some(call) = self.calls.iter_mut().find(|c| c.index == chunk.index) {
if call.id.is_empty() {
call.id = chunk.id.unwrap_or_default();
}
if call.name.is_empty() {
call.name = name;
}
if call.arguments.len() + arguments.len() > MAX_ACCUMULATED_TEXT {
return Err(limit_error("tool arguments exceed size limit"));
}
call.arguments.push_str(&arguments);
} else {
if self.calls.len() >= MAX_TOOL_CALLS {
return Err(limit_error("tool call count exceeds limit"));
}
if arguments.len() > MAX_ACCUMULATED_TEXT {
return Err(limit_error("tool arguments exceed size limit"));
}
self.calls.push(AccumulatedCall {
index: chunk.index,
id: chunk.id.unwrap_or_default(),
name,
arguments,
});
}
Ok(())
}
fn take_all(&mut self) -> Vec<AccumulatedCall> {
std::mem::take(&mut self.calls)
}
fn set_done_reason(&mut self, reason: FinishReason) {
if !self.errored {
self.done_reason = Some(reason);
}
}
fn push_usage(&mut self, usage: OpenAiUsage) {
if !self.errored {
self.usage = Some(map_usage(usage));
}
}
fn flush_done(&mut self) -> Vec<Result<StreamEvent, ProviderError>> {
if self.errored || self.finished {
return Vec::new();
}
let Some(reason) = self.done_reason.take() else {
return vec![Err(ProviderError::Api {
status: 0,
message: "stream ended without finish_reason".into(),
})];
};
let usage = self.usage.take();
vec![Ok(StreamEvent::Done { reason, usage })]
}
fn set_errored(&mut self) {
self.errored = true;
self.done_reason = None;
}
}
fn map_usage(wire: OpenAiUsage) -> Usage {
Usage {
prompt_tokens: wire.prompt_tokens,
completion_tokens: wire.completion_tokens,
total_tokens: wire.total_tokens,
}
}
fn map_network_error(e: reqwest::Error) -> ProviderError {
if e.is_timeout() {
ProviderError::Timeout(TimeoutStage::Transport)
} else {
ProviderError::Network(e.to_string())
}
}
fn map_status_error(
status: u16,
headers: &reqwest::header::HeaderMap,
body: &str,
) -> ProviderError {
let message = extract_error_message(body);
if status == 429 && !message.contains("insufficient_quota") {
ProviderError::RateLimited {
retry_after: extract_retry_after(headers),
}
} else {
ProviderError::Api { status, message }
}
}
fn extract_retry_after(headers: &reqwest::header::HeaderMap) -> Option<Duration> {
headers
.get(reqwest::header::RETRY_AFTER)?
.to_str()
.ok()?
.trim()
.parse::<u64>()
.ok()
.map(Duration::from_secs)
}
fn extract_error_message(body: &str) -> String {
serde_json::from_str::<OpenAiErrorBody>(body)
.ok()
.and_then(|body| body.error)
.map(|error| error.message)
.unwrap_or_else(|| body.to_string())
}
fn parse_response(body: &str) -> Result<ChatResponse, ProviderError> {
let parsed: OpenAiChatResponse =
serde_json::from_str(body).map_err(|e| ProviderError::Api {
status: 0,
message: format!("invalid provider response: {e}"),
})?;
let choice = parsed
.choices
.into_iter()
.next()
.ok_or_else(|| ProviderError::Api {
status: 0,
message: "provider returned empty choices".to_string(),
})?;
let wire = choice.message;
let content = match wire.content {
Some(serde_json::Value::String(s)) => s,
Some(other) => other.to_string(),
None => String::new(),
};
let raw_tool_calls = wire.tool_calls.unwrap_or_default();
if raw_tool_calls.len() > MAX_TOOL_CALLS {
return Err(ProviderError::Api {
status: 0,
message: format!("too many tool calls ({})", raw_tool_calls.len()),
});
}
let tool_calls: Vec<ToolCall> = raw_tool_calls
.into_iter()
.map(|call| ToolCall {
id: call.id,
name: call.function.name,
arguments: call.function.arguments,
})
.collect();
Ok(ChatResponse {
message: Message::Assistant {
content,
reasoning: wire.reasoning_content,
tool_calls,
},
finish_reason: map_finish_reason(&choice.finish_reason),
usage: parsed.usage.map(map_usage).unwrap_or_default(),
})
}
#[derive(serde::Serialize)]
struct OpenAiChatRequest {
model: String,
messages: Vec<OpenAiMessage>,
#[serde(skip_serializing_if = "Vec::is_empty")]
tools: Vec<OpenAiToolDef>,
#[serde(skip_serializing_if = "Option::is_none")]
stream: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
stream_options: Option<StreamOptions>,
#[serde(skip_serializing_if = "Option::is_none")]
temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
response_format: Option<OpenAiResponseFormat>,
#[serde(flatten)]
extra: BTreeMap<String, serde_json::Value>,
}
#[derive(Serialize)]
struct OpenAiResponseFormat {
#[serde(rename = "type")]
kind: &'static str,
#[serde(skip_serializing_if = "Option::is_none")]
strict: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
json_schema: Option<OpenAiJsonSchema>,
}
#[derive(Serialize)]
struct OpenAiJsonSchema {
name: &'static str,
schema: serde_json::Value,
}
impl OpenAiResponseFormat {
fn from_schema(schema: &serde_json::Value) -> Self {
if schema.get("type").and_then(|t| t.as_str()) == Some("object")
&& let (Some(properties), Some(required)) = (
schema.get("properties").and_then(|p| p.as_object()),
schema.get("required").and_then(|r| r.as_array()),
)
&& required.len() == properties.len()
&& required.iter().all(|k| k.is_string())
{
let mut strict_schema = schema.clone();
strict_schema["additionalProperties"] = serde_json::Value::Bool(false);
Self {
kind: "json_schema",
strict: Some(true),
json_schema: Some(OpenAiJsonSchema {
name: "output",
schema: strict_schema,
}),
}
} else {
Self {
kind: "json_schema",
strict: None,
json_schema: Some(OpenAiJsonSchema {
name: "output",
schema: schema.clone(),
}),
}
}
}
fn json_object() -> Self {
Self {
kind: "json_object",
strict: None,
json_schema: None,
}
}
}
fn wire_extra(options: &ModelOptions) -> BTreeMap<String, serde_json::Value> {
options
.extra
.iter()
.filter(|(key, _)| {
!matches!(
key.as_str(),
"model"
| "messages"
| "tools"
| "stream"
| "stream_options"
| "temperature"
| "max_tokens"
)
})
.map(|(key, value)| (key.clone(), value.clone()))
.collect()
}
#[derive(serde::Serialize)]
struct StreamOptions {
include_usage: bool,
}
#[derive(serde::Serialize, serde::Deserialize)]
struct OpenAiMessage {
role: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
content: Option<serde_json::Value>,
#[serde(default, skip_serializing_if = "Option::is_none", alias = "reasoning")]
reasoning_content: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
tool_calls: Option<Vec<OpenAiToolCall>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
tool_call_id: Option<String>,
}
impl OpenAiMessage {
fn from_message(message: &Message) -> Self {
match message {
Message::System(content) => Self::text("system", content),
Message::User(blocks) => Self {
role: "user".to_string(),
content: user_content(blocks),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
},
Message::Assistant {
content,
reasoning,
tool_calls,
} => Self {
role: "assistant".to_string(),
content: (!content.is_empty()).then(|| serde_json::Value::String(content.clone())),
reasoning_content: reasoning.clone(),
tool_calls: (!tool_calls.is_empty()).then(|| {
tool_calls
.iter()
.map(|call| OpenAiToolCall {
id: call.id.clone(),
kind: "function".to_string(),
function: OpenAiFunctionCall {
name: call.name.clone(),
arguments: call.arguments.clone(),
},
})
.collect()
}),
tool_call_id: None,
},
Message::ToolResult { id, content } => Self {
role: "tool".to_string(),
content: Some(serde_json::Value::String(content.clone())),
reasoning_content: None,
tool_calls: None,
tool_call_id: Some(id.clone()),
},
}
}
fn text(role: &str, content: &str) -> Self {
Self {
role: role.to_string(),
content: Some(serde_json::Value::String(content.to_string())),
reasoning_content: None,
tool_calls: None,
tool_call_id: None,
}
}
}
fn user_content(blocks: &[ContentBlock]) -> Option<serde_json::Value> {
match blocks {
[ContentBlock::Text(text)] => Some(serde_json::Value::String(text.clone())),
blocks => {
let parts: Vec<serde_json::Value> = blocks.iter().map(content_block_wire).collect();
(!parts.is_empty()).then_some(serde_json::Value::Array(parts))
}
}
}
fn content_block_wire(block: &ContentBlock) -> serde_json::Value {
match block {
ContentBlock::Text(text) => serde_json::json!({"type": "text", "text": text}),
ContentBlock::Image(image) => serde_json::json!({
"type": "image_url",
"image_url": {
"url": format!(
"data:{};base64,{}",
image.mime_type,
base64::engine::general_purpose::STANDARD.encode(&image.data)
)
}
}),
ContentBlock::Wire(value) => value.clone(),
}
}
#[derive(serde::Serialize, serde::Deserialize)]
struct OpenAiToolCall {
id: String,
#[serde(rename = "type")]
kind: String,
function: OpenAiFunctionCall,
}
#[derive(serde::Serialize, serde::Deserialize)]
struct OpenAiFunctionCall {
name: String,
arguments: String,
}
#[derive(serde::Serialize)]
struct OpenAiToolDef {
#[serde(rename = "type")]
kind: &'static str,
function: OpenAiFunctionDef,
}
#[derive(serde::Serialize)]
struct OpenAiFunctionDef {
name: String,
description: String,
parameters: serde_json::Value,
}
impl OpenAiToolDef {
fn from_schema(schema: &ToolSchema) -> Self {
Self {
kind: "function",
function: OpenAiFunctionDef {
name: schema.name.clone(),
description: schema.description.clone(),
parameters: schema.parameters.clone(),
},
}
}
}
#[derive(serde::Deserialize)]
struct OpenAiUsage {
prompt_tokens: u32,
completion_tokens: u32,
total_tokens: u32,
}
#[derive(serde::Deserialize)]
struct OpenAiChatResponse {
choices: Vec<OpenAiChoice>,
#[serde(default)]
usage: Option<OpenAiUsage>,
}
#[derive(serde::Deserialize)]
struct OpenAiChoice {
message: OpenAiMessage,
finish_reason: String,
}
#[derive(serde::Deserialize)]
struct OpenAiStreamChunk {
#[serde(default)]
choices: Vec<OpenAiStreamChoice>,
#[serde(default)]
usage: Option<OpenAiUsage>,
#[serde(default)]
error: Option<OpenAiStreamError>,
}
#[derive(serde::Deserialize)]
struct OpenAiStreamError {
message: String,
}
#[derive(serde::Deserialize)]
struct OpenAiStreamChoice {
delta: OpenAiStreamDelta,
#[serde(default)]
finish_reason: Option<String>,
}
#[derive(serde::Deserialize)]
struct OpenAiStreamDelta {
#[serde(default)]
content: Option<String>,
#[serde(default, alias = "reasoning")]
reasoning_content: Option<String>,
#[serde(default)]
tool_calls: Option<Vec<OpenAiToolCallChunk>>,
}
#[derive(serde::Deserialize)]
struct OpenAiToolCallChunk {
index: usize,
#[serde(default)]
id: Option<String>,
#[serde(default)]
function: Option<OpenAiFunctionChunk>,
}
#[derive(serde::Deserialize)]
struct OpenAiFunctionChunk {
#[serde(default)]
name: Option<String>,
#[serde(default)]
arguments: Option<String>,
}
#[derive(serde::Deserialize)]
struct OpenAiErrorBody {
error: Option<OpenAiErrorDetail>,
}
#[derive(serde::Deserialize)]
struct OpenAiErrorDetail {
message: String,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::ImageContent;
use reqwest::header::{HeaderMap, RETRY_AFTER};
use serde_json::json;
#[test]
fn status_429_quota_goes_to_api() {
let body = r#"{"error":{"message":"You exceeded your current quota: insufficient_quota"}}"#;
let err = map_status_error(429, &HeaderMap::new(), body);
assert!(matches!(err, ProviderError::Api { status: 429, .. }));
}
#[test]
fn status_429_rate_limited_with_retry_after() {
let mut headers = HeaderMap::new();
headers.insert(RETRY_AFTER, "5".parse().unwrap());
let body = r#"{"error":{"message":"rate limit exceeded"}}"#;
let err = map_status_error(429, &headers, body);
assert!(matches!(
err,
ProviderError::RateLimited { retry_after: Some(d) } if d == Duration::from_secs(5)
));
}
#[test]
fn status_429_without_retry_after_header() {
let body = r#"{"error":{"message":"rate limit exceeded"}}"#;
let err = map_status_error(429, &HeaderMap::new(), body);
assert!(matches!(
err,
ProviderError::RateLimited { retry_after: None }
));
}
#[test]
fn status_5xx_goes_to_api_with_status() {
let body = r#"{"error":{"message":"server overloaded"}}"#;
let err = map_status_error(503, &HeaderMap::new(), body);
assert!(matches!(err, ProviderError::Api { status: 503, .. }));
}
#[test]
fn maps_message_to_wire() {
let message = Message::user("hi");
let wire = OpenAiMessage::from_message(&message);
assert_eq!(
serde_json::to_value(&wire).unwrap(),
json!({"role": "user", "content": "hi"})
);
}
#[test]
fn maps_multi_block_user_content_to_wire_parts() {
let message = Message::user_blocks(vec![
ContentBlock::Text("take a look:".into()),
ContentBlock::Text("see the attachment".into()),
]);
assert_eq!(
serde_json::to_value(OpenAiMessage::from_message(&message)).unwrap(),
json!({
"role": "user",
"content": [
{"type": "text", "text": "take a look:"},
{"type": "text", "text": "see the attachment"}
]
})
);
}
#[test]
fn maps_image_block_to_wire_data_url() {
let message = Message::user_blocks(vec![ContentBlock::Image(ImageContent::new(
"image/png",
b"fake-png-bytes".to_vec(),
))]);
let expected_b64 = base64::engine::general_purpose::STANDARD.encode(b"fake-png-bytes");
assert_eq!(
serde_json::to_value(OpenAiMessage::from_message(&message)).unwrap(),
json!({
"role": "user",
"content": [{
"type": "image_url",
"image_url": {"url": format!("data:image/png;base64,{expected_b64}")}
}]
})
);
}
#[test]
fn passes_through_wire_blocks() {
let message = Message::user_blocks(vec![
ContentBlock::Wire(json!({
"type": "input_audio",
"input_audio": {
"data": "UklGRi4A",
"format": "wav"
}
})),
ContentBlock::Text("transcribe this".into()),
]);
assert_eq!(
serde_json::to_value(OpenAiMessage::from_message(&message)).unwrap(),
json!({
"role": "user",
"content": [
{
"type": "input_audio",
"input_audio": {"data": "UklGRi4A", "format": "wav"}
},
{"type": "text", "text": "transcribe this"}
]
})
);
}
#[test]
fn maps_mixed_text_and_image_blocks() {
let message = Message::user_blocks(vec![
ContentBlock::Text("what is this?".into()),
ContentBlock::Image(ImageContent::new("image/jpeg", vec![1, 2, 3])),
]);
assert_eq!(
serde_json::to_value(OpenAiMessage::from_message(&message)).unwrap(),
json!({
"role": "user",
"content": [
{"type": "text", "text": "what is this?"},
{"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,AQID"}}
]
})
);
}
#[test]
fn maps_system_and_assistant_roles() {
assert_eq!(
OpenAiMessage::from_message(&Message::system("s")).role,
"system"
);
assert_eq!(
OpenAiMessage::from_message(&Message::assistant("a")).role,
"assistant"
);
}
#[test]
fn maps_tool_call_to_wire() {
let message = Message::Assistant {
content: String::new(),
reasoning: None,
tool_calls: vec![ToolCall {
id: "call_1".into(),
name: "calculator".into(),
arguments: r#"{"expression":"1+1"}"#.into(),
}],
};
assert_eq!(
serde_json::to_value(OpenAiMessage::from_message(&message)).unwrap(),
json!({
"role": "assistant",
"tool_calls": [{
"id": "call_1",
"type": "function",
"function": {"name": "calculator", "arguments": "{\"expression\":\"1+1\"}"}
}]
})
);
}
#[test]
fn maps_tool_result_to_wire() {
let message = Message::ToolResult {
id: "call_1".into(),
content: "2".into(),
};
assert_eq!(
serde_json::to_value(OpenAiMessage::from_message(&message)).unwrap(),
json!({"role": "tool", "tool_call_id": "call_1", "content": "2"})
);
}
#[test]
fn maps_reasoning_to_wire() {
let message = Message::assistant_with_reasoning("2", "thinking process");
assert_eq!(
serde_json::to_value(OpenAiMessage::from_message(&message)).unwrap(),
json!({
"role": "assistant",
"content": "2",
"reasoning_content": "thinking process"
})
);
}
#[test]
fn omits_empty_reasoning_in_wire() {
let message = Message::assistant("2");
assert_eq!(
serde_json::to_value(OpenAiMessage::from_message(&message)).unwrap(),
json!({"role": "assistant", "content": "2"})
);
}
#[test]
fn parses_reasoning_content() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":"2","reasoning_content":"thinking process"},"finish_reason":"stop"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(
response.message,
Message::assistant_with_reasoning("2", "thinking process")
);
}
#[test]
fn parses_reasoning_with_tool_call() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":null,"reasoning_content":"decided","tool_calls":[{"id":"call_1","type":"function","function":{"name":"calculator","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(
response.message,
Message::Assistant {
content: String::new(),
reasoning: Some("decided".into()),
tool_calls: vec![ToolCall {
id: "call_1".into(),
name: "calculator".into(),
arguments: "{}".into(),
}],
}
);
}
#[test]
fn accepts_ollama_reasoning_field() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":"2","reasoning":"thinking"},"finish_reason":"stop"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(
response.message,
Message::assistant_with_reasoning("2", "thinking")
);
}
#[test]
fn serializes_tool_defs() {
let schema = ToolSchema {
name: "calculator".into(),
description: "Evaluate expressions".into(),
parameters: json!({"type": "object"}),
};
let def = OpenAiToolDef::from_schema(&schema);
assert_eq!(
serde_json::to_value(&def).unwrap(),
json!({
"type": "function",
"function": {
"name": "calculator",
"description": "Evaluate expressions",
"parameters": {"type": "object"}
}
})
);
}
#[test]
fn omits_empty_tools_in_request() {
let request = OpenAiChatRequest {
model: "m".into(),
messages: vec![OpenAiMessage::from_message(&Message::user("hi"))],
tools: vec![],
stream: None,
stream_options: None,
temperature: None,
max_tokens: None,
response_format: None,
extra: BTreeMap::new(),
};
assert_eq!(
serde_json::to_value(&request).unwrap(),
json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]})
);
}
#[test]
fn serializes_structured_output_as_response_format() {
let schema = json!({
"type": "object",
"properties": { "city": { "type": "string" } },
"required": ["city"],
});
let options = ModelOptions {
structured: Some(schema.clone()),
..Default::default()
};
let request = OpenAiChatRequest {
model: "m".into(),
messages: vec![OpenAiMessage::from_message(&Message::user("hi"))],
tools: vec![],
stream: None,
stream_options: None,
temperature: None,
max_tokens: None,
response_format: options
.structured
.as_ref()
.map(OpenAiResponseFormat::from_schema),
extra: BTreeMap::new(),
};
let mut expected_schema = schema.clone();
expected_schema["additionalProperties"] = json!(false);
assert_eq!(
serde_json::to_value(&request).unwrap(),
json!({
"model": "m",
"messages": [{"role": "user", "content": "hi"}],
"response_format": {
"type": "json_schema",
"strict": true,
"json_schema": { "name": "output", "schema": expected_schema }
}
})
);
}
#[test]
fn strict_mode_skipped_when_required_incomplete() {
let schema = json!({
"type": "object",
"properties": {
"city": { "type": "string" },
"country": { "type": "string" },
},
"required": ["city"],
});
let request = OpenAiChatRequest {
model: "m".into(),
messages: vec![OpenAiMessage::from_message(&Message::user("hi"))],
tools: vec![],
stream: None,
stream_options: None,
temperature: None,
max_tokens: None,
response_format: Some(OpenAiResponseFormat::from_schema(&schema)),
extra: BTreeMap::new(),
};
assert_eq!(
serde_json::to_value(&request).unwrap(),
json!({
"model": "m",
"messages": [{"role": "user", "content": "hi"}],
"response_format": {
"type": "json_schema",
"json_schema": { "name": "output", "schema": schema }
}
})
);
}
#[test]
fn json_object_mode_omits_schema() {
let wire = OpenAiResponseFormat::json_object();
assert_eq!(
serde_json::to_value(&wire).unwrap(),
json!({ "type": "json_object" })
);
}
#[test]
fn response_format_mode_gating() {
let schema = json!({ "type": "object" });
let provider = OpenAiProvider::new("http://x", "k", "m");
assert!(provider.wire_response_format(None).is_none());
assert!(matches!(
provider.wire_response_format(Some(&schema)),
Some(OpenAiResponseFormat {
kind: "json_schema",
..
})
));
let off = OpenAiProvider::new("http://x", "k", "m")
.with_structured_output_mode(StructuredOutputMode::Off);
assert!(off.wire_response_format(Some(&schema)).is_none());
let json_object = OpenAiProvider::new("http://x", "k", "m")
.with_structured_output_mode(StructuredOutputMode::JsonObject);
assert!(matches!(
json_object.wire_response_format(Some(&schema)),
Some(OpenAiResponseFormat {
kind: "json_object",
..
})
));
}
#[test]
fn serializes_model_options_in_request() {
let request = OpenAiChatRequest {
model: "m".into(),
messages: vec![OpenAiMessage::from_message(&Message::user("hi"))],
tools: vec![],
stream: None,
stream_options: None,
temperature: Some(0.5), max_tokens: Some(128),
response_format: None,
extra: BTreeMap::new(),
};
assert_eq!(
serde_json::to_value(&request).unwrap(),
json!({
"model": "m",
"messages": [{"role": "user", "content": "hi"}],
"temperature": 0.5,
"max_tokens": 128
})
);
}
#[test]
fn passes_through_extra_options() {
let mut extra = BTreeMap::new();
extra.insert("top_p".into(), serde_json::json!(0.9));
extra.insert(
"response_format".into(),
serde_json::json!({"type": "json_object"}),
);
let request = OpenAiChatRequest {
model: "m".into(),
messages: vec![OpenAiMessage::from_message(&Message::user("hi"))],
tools: vec![],
stream: None,
stream_options: None,
temperature: None,
max_tokens: None,
response_format: None,
extra,
};
assert_eq!(
serde_json::to_value(&request).unwrap(),
json!({
"model": "m",
"messages": [{"role": "user", "content": "hi"}],
"top_p": 0.9,
"response_format": {"type": "json_object"}
})
);
}
#[test]
fn extra_keys_colliding_with_managed_fields_are_filtered() {
let mut extra = BTreeMap::new();
extra.insert("temperature".into(), serde_json::json!(1.0));
extra.insert("top_p".into(), serde_json::json!(0.9));
let options = ModelOptions {
temperature: Some(0.5),
max_tokens: None,
extra,
structured: None,
};
let request = OpenAiChatRequest {
model: "m".into(),
messages: vec![OpenAiMessage::from_message(&Message::user("hi"))],
tools: vec![],
stream: None,
stream_options: None,
temperature: options.temperature,
max_tokens: options.max_tokens,
response_format: options
.structured
.as_ref()
.map(OpenAiResponseFormat::from_schema),
extra: wire_extra(&options),
};
assert_eq!(
serde_json::to_value(&request).unwrap(),
json!({
"model": "m",
"messages": [{"role": "user", "content": "hi"}],
"temperature": 0.5,
"top_p": 0.9
})
);
}
#[test]
fn parses_chat_response() {
let body = r#"{"id":"chatcmpl-1","object":"chat.completion","created":1,"model":"gpt-4o-mini","choices":[{"index":0,"message":{"role":"assistant","content":"hi!"},"finish_reason":"stop"}],"usage":{"prompt_tokens":5,"completion_tokens":3,"total_tokens":8}}"#;
let response = parse_response(body).unwrap();
assert_eq!(response.message, Message::assistant("hi!"));
assert_eq!(response.finish_reason, FinishReason::Stop);
assert_eq!(response.usage, Usage::new(5, 3));
}
#[test]
fn parses_tool_call_response() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":null,"tool_calls":[{"id":"call_1","type":"function","function":{"name":"calculator","arguments":"{\"expression\": \"1+1\"}"}}]},"finish_reason":"tool_calls"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(
response.message,
Message::Assistant {
content: String::new(),
reasoning: None,
tool_calls: vec![ToolCall {
id: "call_1".into(),
name: "calculator".into(),
arguments: r#"{"expression": "1+1"}"#.into(),
}],
}
);
assert_eq!(response.finish_reason, FinishReason::Stop);
}
#[test]
fn parses_text_with_tool_calls() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":"let me compute","tool_calls":[{"id":"call_1","type":"function","function":{"name":"calculator","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(
response.message,
Message::Assistant {
content: "let me compute".into(),
reasoning: None,
tool_calls: vec![ToolCall {
id: "call_1".into(),
name: "calculator".into(),
arguments: "{}".into(),
}],
}
);
}
#[test]
fn parses_multiple_tool_calls_in_one_message() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":null,"tool_calls":[{"id":"a","type":"function","function":{"name":"calculator","arguments":"{}"}},{"id":"b","type":"function","function":{"name":"calculator","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(
response.message,
Message::Assistant {
content: String::new(),
reasoning: None,
tool_calls: vec![
ToolCall {
id: "a".into(),
name: "calculator".into(),
arguments: "{}".into(),
},
ToolCall {
id: "b".into(),
name: "calculator".into(),
arguments: "{}".into(),
},
],
}
);
}
#[test]
fn maps_length_finish_reason() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":""},"finish_reason":"length"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(response.message, Message::assistant(""));
assert_eq!(response.finish_reason, FinishReason::Length);
}
#[test]
fn maps_unknown_finish_reason_to_other() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":""},"finish_reason":"content_filter"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(
response.finish_reason,
FinishReason::Other("content_filter".into())
);
}
#[test]
fn rejects_empty_choices() {
let body = r#"{"choices":[]}"#;
assert!(parse_response(body).is_err());
}
#[test]
fn missing_usage_defaults_to_zero() {
let body = r#"{"choices":[{"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}]}"#;
let response = parse_response(body).unwrap();
assert_eq!(response.usage, Usage::default());
}
#[test]
fn extracts_error_message() {
let body = r#"{"error":{"message":"Incorrect API key","type":"invalid_request_error"}}"#;
assert_eq!(extract_error_message(body), "Incorrect API key");
}
#[test]
fn falls_back_to_raw_body() {
let body = "not json";
assert_eq!(extract_error_message(body), "not json");
}
fn parse_events(line: &str, aggregator: &mut ToolCallAggregator) -> Vec<StreamEvent> {
parse_sse_line(line, aggregator)
.into_iter()
.map(|event| event.expect("test input must not produce an error"))
.collect()
}
#[test]
fn stream_parses_delta() {
let line = r#"data: {"choices":[{"delta":{"content":"hi"},"finish_reason":null}]}"#;
assert_eq!(
parse_events(line, &mut ToolCallAggregator::default()),
vec![StreamEvent::Delta("hi".to_string())]
);
}
#[test]
fn stream_parses_finish_reason() {
let line = r#"data: {"choices":[{"delta":{},"finish_reason":"stop"}]}"#;
let mut aggregator = ToolCallAggregator::default();
assert!(parse_events(line, &mut aggregator).is_empty());
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: None,
})]
);
}
#[test]
fn stream_maps_length_finish_reason() {
let line = r#"data: {"choices":[{"delta":{},"finish_reason":"length"}]}"#;
let mut aggregator = ToolCallAggregator::default();
assert!(parse_events(line, &mut aggregator).is_empty());
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Length,
usage: None,
})]
);
}
#[test]
fn stream_maps_tool_calls_finish_reason() {
let line = r#"data: {"choices":[{"delta":{},"finish_reason":"tool_calls"}]}"#;
let mut aggregator = ToolCallAggregator::default();
assert!(parse_events(line, &mut aggregator).is_empty());
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: None,
})]
);
}
#[test]
fn stream_parses_usage_only_tail_chunk() {
let mut aggregator = ToolCallAggregator::default();
let done_line = r#"data: {"choices":[{"delta":{"content":"ok"},"finish_reason":"stop"}]}"#;
let usage_line = r#"data: {"choices":[],"usage":{"prompt_tokens":5,"completion_tokens":3,"total_tokens":8}}"#;
assert_eq!(
parse_events(done_line, &mut aggregator),
vec![StreamEvent::Delta("ok".to_string())]
);
assert!(parse_events(usage_line, &mut aggregator).is_empty());
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(Usage::new(5, 3)),
})]
);
}
#[test]
fn stream_parses_usage_in_finish_chunk() {
let line = r#"data: {"choices":[{"delta":{},"finish_reason":"stop"}],"usage":{"prompt_tokens":5,"completion_tokens":3,"total_tokens":8}}"#;
let mut aggregator = ToolCallAggregator::default();
assert!(parse_events(line, &mut aggregator).is_empty());
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(Usage::new(5, 3)),
})]
);
}
#[test]
fn stream_skips_irrelevant_lines() {
let mut aggregator = ToolCallAggregator::default();
assert!(parse_sse_line("", &mut aggregator).is_empty());
assert!(parse_sse_line(": keep-alive", &mut aggregator).is_empty());
assert!(parse_sse_line("event: message", &mut aggregator).is_empty());
let events = parse_sse_line("data: [DONE]", &mut aggregator);
assert_eq!(events.len(), 1);
assert!(matches!(
events[0],
Err(ProviderError::Api { status: 0, .. })
));
assert!(parse_sse_line(r#"data: {"choices":[]}"#, &mut aggregator).is_empty());
assert!(
parse_sse_line(
r#"data: {"choices":[{"delta":{"role":"assistant"},"finish_reason":null}]}"#,
&mut aggregator
)
.is_empty()
);
}
#[test]
fn stream_rejects_invalid_json() {
let line = "data: not-json";
let events = parse_sse_line(line, &mut ToolCallAggregator::default());
match events.as_slice() {
[Err(ProviderError::Api { status: 0, message })] => {
assert_eq!(message, "invalid stream event: not-json")
}
_ => panic!("expected single api error, got {events:?}"),
}
}
#[test]
fn stream_extracts_provider_error() {
let line = r#"data: {"error":{"message":"rate limited"}}"#;
let events = parse_sse_line(line, &mut ToolCallAggregator::default());
let err = events.into_iter().next().unwrap().unwrap_err();
assert!(err.to_string().contains("rate limited"));
}
#[test]
fn stream_forwards_reasoning_chunks() {
let mut aggregator = ToolCallAggregator::default();
let line1 =
r#"data: {"choices":[{"delta":{"reasoning_content":"first"},"finish_reason":null}]}"#;
let line2 = r#"data: {"choices":[{"delta":{"reasoning_content":" multiply"},"finish_reason":null}]}"#;
let line3 = r#"data: {"choices":[{"delta":{"content":"answer"},"finish_reason":"stop"}]}"#;
assert_eq!(
parse_events(line1, &mut aggregator),
vec![StreamEvent::Reasoning("first".to_string())]
);
assert_eq!(
parse_events(line2, &mut aggregator),
vec![StreamEvent::Reasoning(" multiply".to_string())]
);
assert_eq!(
parse_events(line3, &mut aggregator),
vec![StreamEvent::Delta("answer".to_string())]
);
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: None,
})]
);
}
#[test]
fn stream_reasoning_with_tool_call() {
let mut aggregator = ToolCallAggregator::default();
let line1 =
r#"data: {"choices":[{"delta":{"reasoning_content":"all"},"finish_reason":null}]}"#;
let line2 =
r#"data: {"choices":[{"delta":{"reasoning_content":" set"},"finish_reason":null}]}"#;
let line3 = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","function":{"name":"calculator","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}"#;
assert_eq!(
parse_events(line1, &mut aggregator),
vec![StreamEvent::Reasoning("all".to_string())]
);
assert_eq!(
parse_events(line2, &mut aggregator),
vec![StreamEvent::Reasoning(" set".to_string())]
);
assert_eq!(
parse_events(line3, &mut aggregator),
vec![StreamEvent::ToolCall {
id: "call_1".to_string(),
name: "calculator".to_string(),
arguments: "{}".to_string(),
},]
);
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: None,
})]
);
}
#[test]
fn stream_aggregates_tool_call_chunks() {
let mut aggregator = ToolCallAggregator::default();
let line1 = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"calculator","arguments":"{\"expression\": \""}}]},"finish_reason":null}]}"#;
let line2 = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"1+1"}}]},"finish_reason":null}]}"#;
let line3 = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"}"}}]},"finish_reason":"tool_calls"}]}"#;
assert!(parse_events(line1, &mut aggregator).is_empty());
assert!(parse_events(line2, &mut aggregator).is_empty());
assert_eq!(
parse_events(line3, &mut aggregator),
vec![StreamEvent::ToolCall {
id: "call_1".to_string(),
name: "calculator".to_string(),
arguments: "{\"expression\": \"1+1\"}".to_string(),
},]
);
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: None,
})]
);
}
#[test]
fn stream_handles_multiple_interleaved_tool_calls() {
let mut aggregator = ToolCallAggregator::default();
let line1 = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","function":{"name":"now","arguments":"{}"}},{"index":1,"id":"call_2","function":{"name":"calculator","arguments":"{\"exp"}}]},"finish_reason":null}]}"#;
let line2 = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":1,"function":{"arguments":"ression\":\"1\"}"}}]},"finish_reason":"tool_calls"}]}"#;
assert!(parse_events(line1, &mut aggregator).is_empty());
assert_eq!(
parse_events(line2, &mut aggregator),
vec![
StreamEvent::ToolCall {
id: "call_1".to_string(),
name: "now".to_string(),
arguments: "{}".to_string(),
},
StreamEvent::ToolCall {
id: "call_2".to_string(),
name: "calculator".to_string(),
arguments: "{\"expression\":\"1\"}".to_string(),
},
]
);
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: None,
})]
);
}
#[test]
fn stream_keeps_delta_and_tool_call_order() {
let mut aggregator = ToolCallAggregator::default();
let line1 =
r#"data: {"choices":[{"delta":{"content":"let me compute"},"finish_reason":null}]}"#;
let line2 = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","function":{"name":"calculator","arguments":"{}"}}]},"finish_reason":"tool_calls"}]}"#;
assert_eq!(
parse_events(line1, &mut aggregator),
vec![StreamEvent::Delta("let me compute".to_string())]
);
assert_eq!(
parse_events(line2, &mut aggregator),
vec![StreamEvent::ToolCall {
id: "call_1".to_string(),
name: "calculator".to_string(),
arguments: "{}".to_string(),
},]
);
assert_eq!(
aggregator.flush_done(),
vec![Ok(StreamEvent::Done {
reason: FinishReason::Stop,
usage: None,
})]
);
}
#[tokio::test]
async fn sse_lines_joins_chunks_and_strips_terminators() {
let stream = sse_lines(
futures::stream::iter([
Ok(b"data: a\r".to_vec()),
Ok(b"\ndata: b\ndata: c".to_vec()),
]),
Duration::from_secs(60),
Duration::from_secs(60),
);
let lines: Vec<_> = stream.map(|line| line.unwrap()).collect().await;
assert_eq!(
lines,
vec![
"data: a".to_string(),
"data: b".to_string(),
"data: c".to_string()
]
);
}
#[tokio::test]
async fn sse_lines_emits_tail_without_newline() {
let stream = sse_lines(
futures::stream::iter([Ok(b"data: x".to_vec())]),
Duration::from_secs(60),
Duration::from_secs(60),
);
let lines: Vec<_> = stream.map(|line| line.unwrap()).collect().await;
assert_eq!(lines, vec!["data: x".to_string()]);
}
#[tokio::test]
async fn sse_lines_idle_timeout_emits_timeout() {
let mut stream = Box::pin(sse_lines(
futures::stream::pending::<Result<Vec<u8>, reqwest::Error>>(),
Duration::from_millis(10),
Duration::from_secs(60),
));
let first = stream.next().await.unwrap();
assert!(matches!(first, Err(ProviderError::Timeout(_))));
assert!(stream.next().await.is_none());
}
#[tokio::test]
async fn sse_lines_total_deadline_kills_keepalive_stream() {
let mut stream = Box::pin(sse_lines(
futures::stream::iter(std::iter::repeat_with(|| Ok(b": keep-alive\n".to_vec()))),
Duration::from_secs(60),
Duration::from_millis(50),
));
let mut lines = 0usize;
loop {
match stream.next().await {
Some(Ok(_)) => lines += 1,
Some(Err(ProviderError::Timeout(_))) => break,
other => panic!("expected line or Timeout, got {other:?}"),
}
}
assert!(lines > 0, "lines must be produced before the deadline");
assert!(
stream.next().await.is_none(),
"stream must terminate after an error"
);
}
#[tokio::test]
async fn sse_lines_keepalive_lines_delivered_within_deadline() {
let mut stream = Box::pin(sse_lines(
futures::stream::iter(std::iter::repeat_with(|| Ok(b": keep-alive\n".to_vec()))),
Duration::from_secs(60),
Duration::from_secs(60),
));
let first = stream.next().await.unwrap().unwrap();
assert_eq!(first, ": keep-alive");
}
#[tokio::test]
async fn sse_lines_over_limit_terminates_stream() {
let mut stream = Box::pin(sse_lines(
futures::stream::iter([Ok(vec![b'x'; MAX_SSE_LINE + 1])]),
Duration::from_secs(60),
Duration::from_secs(60),
));
let first = stream.next().await.unwrap();
assert!(matches!(first, Err(ProviderError::Api { status: 0, .. })));
assert!(
stream.next().await.is_none(),
"stream must terminate after exceeding the limit"
);
}
#[test]
fn bad_line_discards_pending_done() {
let mut aggregator = ToolCallAggregator::default();
let finish = r#"data: {"choices":[{"delta":{},"finish_reason":"stop"}]}"#;
assert!(parse_sse_line(finish, &mut aggregator).is_empty());
let bad = "data: not-json";
let events = parse_sse_line(bad, &mut aggregator);
assert_eq!(events.len(), 1);
assert!(matches!(
events[0],
Err(ProviderError::Api { status: 0, .. })
));
assert!(parse_sse_line("data: [DONE]", &mut aggregator).is_empty());
}
#[test]
fn tool_call_count_limit_only_for_new_calls() {
let mut aggregator = ToolCallAggregator::default();
for i in 0..MAX_TOOL_CALLS {
let line = format!(
r#"data: {{"choices":[{{"delta":{{"tool_calls":[{{"index":{i},"id":"c{i}","function":{{"name":"f","arguments":"{{}}"}}}}]}},"finish_reason":null}}]}}"#
);
assert!(
parse_sse_line(&line, &mut aggregator).is_empty(),
"tool call {i} must be accepted"
);
}
let over = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":64,"id":"c64","function":{"name":"f","arguments":"{}"}}]},"finish_reason":null}]}"#;
let events = parse_sse_line(over, &mut aggregator);
assert_eq!(events.len(), 1);
assert!(matches!(
events[0],
Err(ProviderError::Api { status: 0, .. })
));
let cont = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"more\":true}"}}]},"finish_reason":null}]}"#;
assert!(parse_sse_line(cont, &mut aggregator).is_empty());
}
#[test]
fn tool_arguments_size_limit() {
let mut aggregator = ToolCallAggregator::default();
let huge = "x".repeat(MAX_ACCUMULATED_TEXT + 1);
let line = format!(
r#"data: {{"choices":[{{"delta":{{"tool_calls":[{{"index":0,"id":"c0","function":{{"name":"f","arguments":"{huge}"}}}}]}},"finish_reason":null}}]}}"#
);
let events = parse_sse_line(&line, &mut aggregator);
assert_eq!(events.len(), 1);
assert!(matches!(
events[0],
Err(ProviderError::Api { status: 0, .. })
));
}
#[test]
fn debug_masks_api_key() {
let provider = OpenAiProvider::new("https://api.example.com", "sk-super-secret", "model");
let debug = format!("{provider:?}");
assert!(
!debug.contains("sk-super-secret"),
"Debug leaked api_key: {debug}"
);
assert!(
debug.contains("\"***\""),
"mask placeholder missing: {debug}"
);
}
#[test]
fn tool_arguments_continuation_size_limit() {
let mut aggregator = ToolCallAggregator::default();
let first = "x".repeat(MAX_ACCUMULATED_TEXT - 1);
let line1 = format!(
r#"data: {{"choices":[{{"delta":{{"tool_calls":[{{"index":0,"id":"c0","function":{{"name":"f","arguments":"{first}"}}}}]}},"finish_reason":null}}]}}"#
);
assert!(parse_sse_line(&line1, &mut aggregator).is_empty());
let cont = r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"yy"}}]},"finish_reason":null}]}"#;
let events = parse_sse_line(cont, &mut aggregator);
assert_eq!(events.len(), 1);
assert!(matches!(
events[0],
Err(ProviderError::Api { status: 0, .. })
));
}
#[test]
fn errored_latch_blocks_resurrected_done() {
let mut aggregator = ToolCallAggregator::default();
let events = parse_sse_line("data: not-json", &mut aggregator);
assert_eq!(events.len(), 1);
assert!(matches!(
events[0],
Err(ProviderError::Api { status: 0, .. })
));
let finish = r#"data: {"choices":[{"delta":{},"finish_reason":"stop"}]}"#;
assert!(parse_sse_line(finish, &mut aggregator).is_empty());
assert!(parse_sse_line("data: [DONE]", &mut aggregator).is_empty());
}
#[tokio::test]
async fn sse_stream_eof_flushes_pending_done() {
let aggregator = Arc::new(Mutex::new(ToolCallAggregator::default()));
let chunks = futures::stream::iter([
Ok(b"data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}]}\n".to_vec()),
Ok(b"data: {\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":2,\"total_tokens\":3}}\n".to_vec()),
]);
let mut stream = Box::pin(sse_stream(
chunks,
Duration::from_secs(60),
Duration::from_secs(60),
aggregator,
));
let events: Vec<_> = stream.by_ref().map(|e| e.unwrap()).collect().await;
assert_eq!(
events,
vec![StreamEvent::Done {
reason: FinishReason::Stop,
usage: Some(Usage::new(1, 2)),
}]
);
}
#[tokio::test]
async fn sse_stream_done_marker_flushes() {
let aggregator = Arc::new(Mutex::new(ToolCallAggregator::default()));
let chunks = futures::stream::iter([
Ok(b"data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}]}\n".to_vec()),
Ok(b"data: [DONE]\n".to_vec()),
]);
let mut stream = Box::pin(sse_stream(
chunks,
Duration::from_secs(60),
Duration::from_secs(60),
aggregator,
));
let events: Vec<_> = stream.by_ref().map(|e| e.unwrap()).collect().await;
assert_eq!(
events,
vec![StreamEvent::Done {
reason: FinishReason::Stop,
usage: None
}]
);
}
#[tokio::test]
async fn sse_stream_error_then_no_done() {
let aggregator = Arc::new(Mutex::new(ToolCallAggregator::default()));
let chunks = futures::stream::iter([
Ok(b"data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}]}".to_vec()),
Ok(b"data: not-json".to_vec()),
Ok(b"data: [DONE]".to_vec()),
]);
let mut stream = Box::pin(sse_stream(
chunks,
Duration::from_secs(60),
Duration::from_secs(60),
aggregator,
));
let events: Vec<Result<StreamEvent, ProviderError>> = stream.by_ref().collect().await;
assert_eq!(events.len(), 1);
assert!(matches!(
events[0],
Err(ProviderError::Api { status: 0, .. })
));
}
#[tokio::test]
async fn chat_over_real_http_wire() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buf = [0u8; 8192];
let n = socket.read(&mut buf).await.unwrap();
let req = String::from_utf8_lossy(&buf[..n]).to_string();
assert!(
req.starts_with("POST /chat/completions HTTP/1.1"),
"method/path: {req}"
);
assert!(
req.contains("authorization: Bearer sk-test"),
"auth header: {req}"
);
assert!(
req.contains(r#""model":"deepseek-chat""#),
"body model: {req}"
);
let body = r#"{"choices":[{"message":{"role":"assistant","content":"hi","reasoning_content":null,"tool_calls":null},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":2,"total_tokens":3}}"#;
let resp = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\n\r\n{}",
body.len(),
body
);
socket.write_all(resp.as_bytes()).await.unwrap();
});
let provider = OpenAiProvider::new(format!("http://{addr}"), "sk-test", "deepseek-chat");
let resp = provider
.chat(ChatRequest {
messages: vec![crate::message::Message::user("hi")],
tools: vec![],
options: Default::default(),
})
.await
.unwrap();
assert_eq!(resp.message, crate::message::Message::assistant("hi"));
assert_eq!(resp.usage, Usage::new(1, 2));
server.await.unwrap();
}
#[tokio::test]
async fn chat_request_timeout_fires() {
use tokio::io::AsyncReadExt;
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buf = [0u8; 8192];
let _ = socket.read(&mut buf).await;
std::future::pending::<()>().await;
});
let provider = OpenAiProvider::new(format!("http://{addr}"), "sk-test", "m")
.with_request_timeout(Duration::from_millis(100));
let err = provider
.chat(ChatRequest {
messages: vec![crate::message::Message::user("hi")],
tools: vec![],
options: Default::default(),
})
.await
.unwrap_err();
assert!(matches!(err, ProviderError::Timeout(TimeoutStage::Request)));
server.abort();
}
#[tokio::test]
async fn chat_connect_timeout_fires() {
let provider = OpenAiProvider::new("http://192.0.2.1/mcp", "sk-test", "m")
.with_connect_timeout(Duration::from_millis(200));
let err = provider
.chat(ChatRequest {
messages: vec![crate::message::Message::user("hi")],
tools: vec![],
options: Default::default(),
})
.await
.unwrap_err();
assert!(
matches!(&err, ProviderError::Timeout(_) | ProviderError::Network(_)),
"connection failure should classify as Timeout or Network, got: {err:?}"
);
}
#[tokio::test]
async fn sse_stream_empty_stream_is_error() {
let aggregator = Arc::new(Mutex::new(ToolCallAggregator::default()));
let mut stream = Box::pin(sse_stream(
futures::stream::iter([Ok(b"\n".to_vec())]),
Duration::from_secs(60),
Duration::from_secs(60),
aggregator,
));
let events: Vec<Result<StreamEvent, ProviderError>> = stream.by_ref().collect().await;
assert_eq!(events.len(), 1);
assert!(matches!(
events[0],
Err(ProviderError::Api { status: 0, ref message }) if message.contains("without finish_reason")
));
}
#[tokio::test]
async fn sse_stream_drops_data_after_done_marker() {
let aggregator = Arc::new(Mutex::new(ToolCallAggregator::default()));
let chunks = futures::stream::iter([
Ok(b"data: {\"choices\":[{\"delta\":{\"content\":\"ok\"},\"finish_reason\":\"stop\"}]}\n".to_vec()),
Ok(b"data: [DONE]\n".to_vec()),
Ok(b"data: {\"choices\":[{\"delta\":{\"content\":\"x\"}}]}\n".to_vec()),
]);
let mut stream = Box::pin(sse_stream(
chunks,
Duration::from_secs(60),
Duration::from_secs(60),
aggregator,
));
let events: Vec<_> = stream.by_ref().map(|e| e.unwrap()).collect().await;
assert_eq!(
events,
vec![
StreamEvent::Delta("ok".into()),
StreamEvent::Done {
reason: FinishReason::Stop,
usage: None
},
]
);
}
#[tokio::test]
async fn network_error_terminates_and_drops_buffered_tail() {
let conn_err = reqwest::Client::new()
.get("http://[::1")
.build()
.expect_err("invalid URL must fail to build");
let mut stream = Box::pin(sse_lines(
futures::stream::iter([Ok(b"data: partial".to_vec()), Err(conn_err)]),
Duration::from_secs(60),
Duration::from_secs(60),
));
let first = stream.next().await.unwrap();
assert!(matches!(first, Err(ProviderError::Network(_))));
assert!(stream.next().await.is_none());
}
}