use std::io::{BufRead, BufReader};
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct ChatMessage {
pub role: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub tool_calls: Vec<ToolCall>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
}
impl ChatMessage {
pub fn user(content: impl Into<String>) -> Self {
Self { role: "user".into(), content: Some(content.into()), tool_calls: Vec::new(), tool_call_id: None }
}
pub fn system(content: impl Into<String>) -> Self {
Self { role: "system".into(), content: Some(content.into()), tool_calls: Vec::new(), tool_call_id: None }
}
pub fn tool_result(tool_call_id: impl Into<String>, output: impl Into<String>) -> Self {
Self {
role: "tool".into(),
content: Some(output.into()),
tool_calls: Vec::new(),
tool_call_id: Some(tool_call_id.into()),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct ToolCall {
#[serde(default)]
pub id: String,
pub function: FunctionCall,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct FunctionCall {
pub name: String,
#[serde(default)]
pub arguments: String,
}
#[derive(Debug, Deserialize)]
pub(crate) struct ChatResponse {
#[serde(default)]
pub choices: Vec<Choice>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct Choice {
pub message: ChatMessage,
}
#[derive(Debug, Deserialize)]
pub(crate) struct Usage {
#[serde(default)]
pub prompt_tokens: Option<u64>,
#[serde(default)]
pub completion_tokens: Option<u64>,
#[serde(default)]
pub total_tokens: Option<u64>,
#[serde(default)]
pub prompt_tokens_details: Option<PromptTokensDetails>,
#[serde(default)]
pub cache_read_input_tokens: Option<u64>,
#[serde(default)]
pub cache_creation_input_tokens: Option<u64>,
}
impl Usage {
pub(crate) fn cache_read(&self) -> Option<u64> {
self.cache_read_input_tokens
.or_else(|| self.prompt_tokens_details.as_ref().and_then(|d| d.cached_tokens))
}
pub(crate) fn cache_write(&self) -> Option<u64> {
self.cache_creation_input_tokens
}
}
#[derive(Debug, Deserialize)]
pub(crate) struct PromptTokensDetails {
#[serde(default)]
pub cached_tokens: Option<u64>,
}
const MAX_RETRIES: u32 = 3;
pub(super) fn send_with_retry(url: &str, make: impl Fn() -> Result<ureq::Response, Box<ureq::Error>>) -> Result<ureq::Response, String> {
let mut attempt = 0u32;
loop {
match make() {
Ok(resp) => return Ok(resp),
Err(e) if attempt < MAX_RETRIES && is_retryable(&e) => {
let backoff = retry_after(&e).unwrap_or_else(|| Duration::from_millis(1000 * 2u64.pow(attempt)));
std::thread::sleep(backoff);
attempt += 1;
}
Err(e) => return Err(format!("chat request to {url} failed: {e}")),
}
}
}
fn is_retryable(e: &ureq::Error) -> bool {
match e {
ureq::Error::Status(code, _) => status_is_retryable(*code),
ureq::Error::Transport(_) => true, }
}
fn status_is_retryable(code: u16) -> bool {
code == 429 || (500..=599).contains(&code)
}
fn retry_after(e: &ureq::Error) -> Option<Duration> {
match e {
ureq::Error::Status(_, resp) => {
resp.header("Retry-After").and_then(|v| v.parse::<u64>().ok()).map(Duration::from_secs)
}
ureq::Error::Transport(_) => None,
}
}
pub(crate) fn post_chat(
base: &str,
api_key: Option<&str>,
model: &str,
messages: &[ChatMessage],
tools: &[Value],
) -> Result<ChatResponse, String> {
let url = format!("{base}/v1/chat/completions");
let mut body = json!({
"model": model,
"messages": messages,
"stream": false,
});
if !tools.is_empty() {
body["tools"] = Value::Array(tools.to_vec());
}
let resp = send_with_retry(&url, || {
let mut req = ureq::post(&url);
if let Some(key) = api_key {
req = req.set("Authorization", &format!("Bearer {key}"));
}
req.send_json(body.clone()).map_err(Box::new)
})?;
resp.into_json::<ChatResponse>()
.map_err(|e| format!("decoding chat response from {url}: {e}"))
}
pub(crate) enum Fragment<'a> {
Text(&'a str),
Reasoning(&'a str),
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn post_chat_stream(
base: &str,
api_key: Option<&str>,
model: &str,
messages: &[ChatMessage],
tools: &[Value],
extras: RequestExtras,
cancel: &AtomicBool,
on_delta: impl FnMut(Fragment),
) -> Result<(ChatMessage, Option<Usage>), String> {
let url = format!("{base}/v1/chat/completions");
let mut body = json!({
"model": model,
"messages": messages,
"stream": true,
"stream_options": { "include_usage": true },
});
if !tools.is_empty() {
body["tools"] = Value::Array(tools.to_vec());
}
if let Some(rf) = extras.response_format {
body["response_format"] = rf.clone();
}
if !extras.image_data_uris.is_empty() {
attach_images(&mut body, extras.image_data_uris);
}
let resp = send_with_retry(&url, || {
let mut req = ureq::post(&url);
if let Some(key) = api_key {
req = req.set("Authorization", &format!("Bearer {key}"));
}
req.send_json(body.clone()).map_err(Box::new)
})?;
let reader = BufReader::new(resp.into_reader());
Ok(drain_stream(reader.lines().map_while(Result::ok), extras.reasoning_tag, cancel, on_delta))
}
#[derive(Default)]
pub(crate) struct RequestExtras<'a> {
pub response_format: Option<&'a Value>,
pub image_data_uris: &'a [String],
pub reasoning_tag: Option<&'a str>,
}
fn attach_images(body: &mut Value, uris: &[String]) {
let Some(messages) = body["messages"].as_array_mut() else { return };
let Some(first_user) = messages.iter_mut().find(|m| m["role"] == "user") else { return };
let text = first_user["content"].as_str().unwrap_or_default().to_owned();
let mut parts = vec![json!({ "type": "text", "text": text })];
for uri in uris {
parts.push(json!({ "type": "image_url", "image_url": { "url": uri } }));
}
first_user["content"] = Value::Array(parts);
}
pub(crate) fn image_data_uri(mime: &str, data: &[u8]) -> String {
format!("data:{mime};base64,{}", base64_encode(data))
}
fn base64_encode(data: &[u8]) -> String {
const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
let mut out = String::with_capacity(data.len().div_ceil(3) * 4);
for chunk in data.chunks(3) {
let b1 = *chunk.get(1).unwrap_or(&0);
let b2 = *chunk.get(2).unwrap_or(&0);
let n = ((chunk[0] as u32) << 16) | ((b1 as u32) << 8) | (b2 as u32);
out.push(ALPHABET[((n >> 18) & 63) as usize] as char);
out.push(ALPHABET[((n >> 12) & 63) as usize] as char);
out.push(if chunk.len() > 1 { ALPHABET[((n >> 6) & 63) as usize] as char } else { '=' });
out.push(if chunk.len() > 2 { ALPHABET[(n & 63) as usize] as char } else { '=' });
}
out
}
fn drain_stream(
lines: impl Iterator<Item = String>,
reasoning_tag: Option<&str>,
cancel: &AtomicBool,
mut on_delta: impl FnMut(Fragment),
) -> (ChatMessage, Option<Usage>) {
let mut content = String::new();
let mut tool_calls: Vec<ToolCall> = Vec::new();
let mut usage = None;
let mut think = ThinkSplitter::new(reasoning_tag);
for line in lines {
if cancel.load(Ordering::SeqCst) {
break;
}
let Some(data) = line.strip_prefix("data:").map(str::trim) else {
continue;
};
if data == "[DONE]" {
break;
}
let chunk: StreamChunk = match serde_json::from_str(data) {
Ok(c) => c,
Err(_) => continue, };
if chunk.usage.is_some() {
usage = chunk.usage;
}
if let Some(choice) = chunk.choices.into_iter().next() {
if let Some(text) = choice.delta.content {
if !text.is_empty() {
content.push_str(&think.feed(&text, &mut on_delta));
}
}
if let Some(reasoning) = choice.delta.reasoning_content.or(choice.delta.reasoning) {
if !reasoning.is_empty() {
on_delta(Fragment::Reasoning(&reasoning));
}
}
for delta in choice.delta.tool_calls {
accumulate_tool_call(&mut tool_calls, delta);
}
}
}
content.push_str(&think.finish(&mut on_delta));
let message = ChatMessage {
role: "assistant".to_owned(),
content: (!content.is_empty()).then_some(content),
tool_calls,
tool_call_id: None,
};
(message, usage)
}
#[derive(Default)]
pub(super) struct ThinkSplitter {
open: String,
close: String,
active: bool,
in_think: bool,
carry: String,
}
impl ThinkSplitter {
pub(super) fn new(tag: Option<&str>) -> Self {
match tag {
Some(t) if !t.is_empty() => {
Self { open: format!("<{t}>"), close: format!("</{t}>"), active: true, ..Self::default() }
}
_ => Self::default(),
}
}
pub(super) fn feed(&mut self, piece: &str, on: &mut impl FnMut(Fragment)) -> String {
if !self.active {
on(Fragment::Text(piece));
return piece.to_owned();
}
let mut buf = std::mem::take(&mut self.carry);
buf.push_str(piece);
let mut visible = String::new();
loop {
let tag = if self.in_think { self.close.clone() } else { self.open.clone() };
if let Some(i) = buf.find(&tag) {
let before = &buf[..i];
if !before.is_empty() {
if self.in_think {
on(Fragment::Reasoning(before));
} else {
on(Fragment::Text(before));
visible.push_str(before);
}
}
buf.replace_range(..i + tag.len(), "");
self.in_think = !self.in_think;
} else {
let cut = buf.len() - partial_tag_suffix_len(&buf, &tag);
if cut > 0 {
let run = &buf[..cut];
if self.in_think {
on(Fragment::Reasoning(run));
} else {
on(Fragment::Text(run));
visible.push_str(run);
}
}
self.carry = buf[cut..].to_owned();
break;
}
}
visible
}
pub(super) fn finish(&mut self, on: &mut impl FnMut(Fragment)) -> String {
let rest = std::mem::take(&mut self.carry);
if rest.is_empty() {
return String::new();
}
if self.in_think {
on(Fragment::Reasoning(&rest));
String::new()
} else {
on(Fragment::Text(&rest));
rest
}
}
}
fn partial_tag_suffix_len(buf: &str, tag: &str) -> usize {
let b = buf.as_bytes();
let max = tag.len().min(b.len());
(1..=max).rev().find(|&n| tag.as_bytes().starts_with(&b[b.len() - n..])).unwrap_or(0)
}
fn accumulate_tool_call(calls: &mut Vec<ToolCall>, delta: DeltaToolCall) {
while calls.len() <= delta.index {
calls.push(ToolCall { id: String::new(), function: FunctionCall { name: String::new(), arguments: String::new() } });
}
let call = &mut calls[delta.index];
if let Some(id) = delta.id.filter(|s| !s.is_empty()) {
call.id = id;
}
if let Some(function) = delta.function {
if let Some(name) = function.name.filter(|s| !s.is_empty()) {
call.function.name = name;
}
if let Some(args) = function.arguments {
call.function.arguments.push_str(&args);
}
}
}
#[derive(Deserialize)]
struct StreamChunk {
#[serde(default)]
choices: Vec<StreamChoice>,
#[serde(default)]
usage: Option<Usage>,
}
#[derive(Deserialize)]
struct StreamChoice {
#[serde(default)]
delta: Delta,
}
#[derive(Deserialize, Default)]
struct Delta {
#[serde(default)]
content: Option<String>,
#[serde(default)]
reasoning_content: Option<String>,
#[serde(default)]
reasoning: Option<String>,
#[serde(default)]
tool_calls: Vec<DeltaToolCall>,
}
#[derive(Deserialize)]
struct DeltaToolCall {
#[serde(default)]
index: usize,
#[serde(default)]
id: Option<String>,
#[serde(default)]
function: Option<DeltaFunction>,
}
#[derive(Deserialize)]
struct DeltaFunction {
#[serde(default)]
name: Option<String>,
#[serde(default)]
arguments: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn drain_stream_routes_inline_think_tags_to_reasoning() {
let lines = vec![
r#"data: {"choices":[{"delta":{"content":"Hi <thi"}}]}"#.to_string(),
r#"data: {"choices":[{"delta":{"content":"nk>secret pl"}}]}"#.to_string(),
r#"data: {"choices":[{"delta":{"content":"an</think> answer"}}]}"#.to_string(),
"data: [DONE]".to_string(),
];
let (mut text, mut reasoning) = (String::new(), String::new());
let (msg, _) = drain_stream(lines.into_iter(), Some("think"), &AtomicBool::new(false), |f| match f {
Fragment::Text(t) => text.push_str(t),
Fragment::Reasoning(r) => reasoning.push_str(r),
});
assert_eq!(reasoning, "secret plan");
assert_eq!(text, "Hi answer");
assert_eq!(msg.content.as_deref(), Some("Hi answer"));
}
#[test]
fn drain_stream_passthrough_when_reasoning_disabled() {
let lines = vec![
r#"data: {"choices":[{"delta":{"content":"<think>x</think>hi"}}]}"#.to_string(),
"data: [DONE]".to_string(),
];
let mut text = String::new();
let (msg, _) = drain_stream(lines.into_iter(), None, &AtomicBool::new(false), |f| {
if let Fragment::Text(t) = f {
text.push_str(t);
}
});
assert_eq!(text, "<think>x</think>hi");
assert_eq!(msg.content.as_deref(), Some("<think>x</think>hi"));
}
#[test]
fn base64_and_image_data_uri() {
assert_eq!(base64_encode(b"Man"), "TWFu");
assert_eq!(base64_encode(b"Ma"), "TWE=");
assert_eq!(base64_encode(b"M"), "TQ==");
assert_eq!(image_data_uri("image/png", b"M"), "data:image/png;base64,TQ==");
}
#[test]
fn attach_images_rewrites_first_user_message() {
let mut body = json!({ "messages": [
{ "role": "system", "content": "sys" },
{ "role": "user", "content": "look at this" },
]});
attach_images(&mut body, &["data:image/png;base64,TQ==".to_owned()]);
let user = &body["messages"][1]["content"];
assert_eq!(user[0], json!({ "type": "text", "text": "look at this" }));
assert_eq!(user[1]["type"], "image_url");
assert_eq!(user[1]["image_url"]["url"], "data:image/png;base64,TQ==");
assert_eq!(body["messages"][0]["content"], "sys", "system message untouched");
}
#[test]
fn drain_stream_assembles_text_tool_calls_and_usage() {
let lines = vec![
r#"data: {"choices":[{"delta":{"reasoning_content":"think"}}]}"#.to_string(),
r#"data: {"choices":[{"delta":{"content":"Hel"}}]}"#.to_string(),
r#"data: {"choices":[{"delta":{"content":"lo"}}]}"#.to_string(),
r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"c1","function":{"name":"read","arguments":"{\"pa"}}]}}]}"#.to_string(),
r#"data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"th\":\"x\"}"}}]}}]}"#.to_string(),
r#"data: {"choices":[],"usage":{"prompt_tokens":5,"completion_tokens":3,"total_tokens":8}}"#.to_string(),
"data: [DONE]".to_string(),
];
let mut text = String::new();
let mut reasoning = String::new();
let (msg, usage) = drain_stream(lines.into_iter(), Some("think"), &AtomicBool::new(false), |f| match f {
Fragment::Text(t) => text.push_str(t),
Fragment::Reasoning(r) => reasoning.push_str(r),
});
assert_eq!(reasoning, "think", "reasoning surfaced separately");
assert_eq!(text, "Hello", "text deltas streamed live");
assert_eq!(msg.content.as_deref(), Some("Hello"));
assert_eq!(msg.tool_calls.len(), 1);
assert_eq!(msg.tool_calls[0].id, "c1");
assert_eq!(msg.tool_calls[0].function.name, "read");
assert_eq!(msg.tool_calls[0].function.arguments, r#"{"path":"x"}"#);
assert_eq!(usage.unwrap().total_tokens, Some(8));
}
#[test]
fn drain_stream_ignores_keepalives_and_tolerates_no_done() {
let lines = vec![
": keep-alive".to_string(),
String::new(),
r#"data: {"choices":[{"delta":{"content":"hi"}}]}"#.to_string(),
];
let (msg, usage) = drain_stream(lines.into_iter(), Some("think"), &AtomicBool::new(false), |_| {});
assert_eq!(msg.content.as_deref(), Some("hi"));
assert!(usage.is_none());
}
#[test]
fn usage_reads_cache_tokens_across_shapes() {
let openai: Usage =
serde_json::from_str(r#"{"prompt_tokens":10,"completion_tokens":5,"prompt_tokens_details":{"cached_tokens":7}}"#).unwrap();
assert_eq!(openai.cache_read(), Some(7));
assert_eq!(openai.cache_write(), None);
let anthropic: Usage =
serde_json::from_str(r#"{"cache_read_input_tokens":3,"cache_creation_input_tokens":2}"#).unwrap();
assert_eq!(anthropic.cache_read(), Some(3));
assert_eq!(anthropic.cache_write(), Some(2));
}
#[test]
fn retryable_status_classification() {
assert!(status_is_retryable(429), "rate limit retries");
assert!(status_is_retryable(500) && status_is_retryable(503), "5xx retries");
assert!(!status_is_retryable(400) && !status_is_retryable(401) && !status_is_retryable(404), "4xx is terminal");
}
#[test]
fn drain_stream_stops_within_one_chunk_of_cancel() {
let cancel = AtomicBool::new(false);
let lines: Vec<String> = (0..100)
.map(|i| format!(r#"data: {{"choices":[{{"delta":{{"content":"c{i}"}}}}]}}"#))
.collect();
let mut seen = 0;
let (msg, _) = drain_stream(
lines.into_iter().inspect(|_| {
seen += 1;
cancel.store(true, Ordering::SeqCst);
}),
None,
&cancel,
|_| {},
);
assert_eq!(seen, 1, "the poll after the first pull saw the flag");
assert!(msg.content.as_deref().unwrap_or("").is_empty());
}
}