pub mod anthropic;
pub mod openai;
pub(crate) mod openai_wire;
pub mod openrouter;
pub mod fallback;
pub mod record;
pub mod replay;
pub use anthropic::Anthropic;
pub use fallback::Fallback;
pub use openai::OpenAi;
pub use openrouter::OpenRouter;
pub use record::Record;
pub use replay::Replay;
use futures_util::StreamExt;
use crate::error::{Error, Result};
pub(crate) async fn ensure_success(resp: reqwest::Response) -> Result<reqwest::Response> {
let status = resp.status();
if status.is_success() {
return Ok(resp);
}
let retry_after = crate::net::retry_after(resp.headers());
let detail = resp.text().await.unwrap_or_default();
let detail = detail.trim();
Err(Error::provider_status(
status.as_u16(),
retry_after,
if detail.is_empty() {
status.canonical_reason().unwrap_or("no detail").to_string()
} else {
detail.to_string()
},
))
}
pub(crate) fn ensure_parsed(response: CompletionResponse) -> Result<CompletionResponse> {
if response.text.is_none() && response.tool_calls.is_empty() && response.usage.is_none() {
return Err(Error::provider_malformed(
"the response stream yielded no text, no tool call and no usage",
));
}
Ok(response)
}
pub(crate) async fn read_sse<F>(resp: reqwest::Response, mut ingest: F) -> Result<()>
where
F: FnMut(&str) -> bool,
{
let mut stream = resp.bytes_stream();
let mut buf = String::new();
while let Some(chunk) = stream.next().await {
let chunk = chunk?;
buf.push_str(&String::from_utf8_lossy(&chunk));
while let Some(nl) = buf.find('\n') {
let line = buf[..nl].trim_end_matches('\r').to_string();
buf.drain(..=nl);
let Some(data) = line.strip_prefix("data:") else {
continue;
};
let data = data.trim();
if data.is_empty() {
continue;
}
if ingest(data) {
return Ok(());
}
}
}
Ok(())
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ToolSpec {
pub name: String,
pub description: String,
pub parameters: serde_json::Value,
}
#[cfg(feature = "media")]
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct Media {
pub media_type: String,
pub base64: String,
}
#[cfg(feature = "media")]
pub const IMAGE_MEDIA_TYPES: [&str; 4] = ["image/jpeg", "image/png", "image/gif", "image/webp"];
#[cfg(feature = "media")]
pub const MAX_IMAGE_BYTES: usize = 5 * 1024 * 1024;
#[cfg(feature = "media")]
pub const MAX_REQUEST_IMAGE_BYTES: usize = 20 * 1024 * 1024;
#[cfg(feature = "media")]
impl Media {
pub fn image(media_type: impl Into<String>, bytes: &[u8]) -> Result<Self> {
use base64::Engine as _;
let media_type = media_type.into();
if !IMAGE_MEDIA_TYPES.contains(&media_type.as_str()) {
return Err(Error::Config(format!(
"unsupported image media type {media_type:?}: expected one of {}",
IMAGE_MEDIA_TYPES.join(", ")
)));
}
if bytes.len() > MAX_IMAGE_BYTES {
return Err(Error::Config(format!(
"image is {} bytes, over the {MAX_IMAGE_BYTES}-byte per-image bound; \
resize it before attaching rather than sending it truncated",
bytes.len()
)));
}
Ok(Self {
media_type,
base64: base64::engine::general_purpose::STANDARD.encode(bytes),
})
}
pub fn media_type_for(path: &str) -> Option<&'static str> {
let ext = path.rsplit('.').next()?.to_ascii_lowercase();
Some(match ext.as_str() {
"jpg" | "jpeg" => "image/jpeg",
"png" => "image/png",
"gif" => "image/gif",
"webp" => "image/webp",
_ => return None,
})
}
pub fn byte_len(&self) -> usize {
let pad = self.base64.bytes().rev().take_while(|b| *b == b'=').count();
(self.base64.len() / 4 * 3).saturating_sub(pad)
}
pub fn digest(&self) -> String {
use std::hash::{Hash as _, Hasher as _};
let mut h = std::collections::hash_map::DefaultHasher::new();
self.media_type.hash(&mut h);
self.base64.hash(&mut h);
format!("{:016x}", h.finish())
}
}
#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct CompletionRequest {
pub system: String,
pub user: String,
pub tools: Vec<ToolSpec>,
#[cfg(feature = "media")]
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub media: Vec<Media>,
}
#[cfg(feature = "media")]
pub(crate) fn ensure_media_accepted(
name: &str,
accepts: bool,
request: &CompletionRequest,
) -> Result<()> {
if request.media.is_empty() || accepts {
return Ok(());
}
Err(Error::Config(format!(
"provider {name:?} does not accept image input, and the request carries {} image(s); \
no request was sent",
request.media.len()
)))
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ToolCall {
pub name: String,
pub arguments: serde_json::Value,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct Usage {
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub total_tokens: u64,
}
#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct CompletionResponse {
pub text: Option<String>,
pub tool_calls: Vec<ToolCall>,
pub usage: Option<Usage>,
}
pub trait Provider {
fn complete(
&self,
request: CompletionRequest,
) -> impl std::future::Future<Output = Result<CompletionResponse>> + Send;
fn name(&self) -> &str {
"provider"
}
#[cfg(feature = "media")]
fn accepts_images(&self) -> bool {
false
}
fn endpoint(&self) -> Option<&str> {
None
}
fn endpoints(&self) -> Vec<&str> {
self.endpoint().into_iter().collect()
}
fn last_served(&self) -> Option<String> {
None
}
}
#[cfg(test)]
mod failures {
use std::io::{Read, Write};
use std::net::TcpListener;
use std::time::Duration;
use super::*;
use crate::error::ProviderErrorKind as Kind;
use crate::net::{http_date, unix_now};
fn serve(response: String) -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(mut stream) = stream else { break };
drain_request(&mut stream);
let _ = stream.write_all(response.as_bytes());
let _ = stream.flush();
}
});
url
}
fn drain_request(stream: &mut std::net::TcpStream) {
let mut seen = Vec::new();
let mut byte = [0u8; 1];
while stream.read(&mut byte).unwrap_or(0) == 1 {
seen.push(byte[0]);
if seen.ends_with(b"\r\n\r\n") {
break;
}
}
let head = String::from_utf8_lossy(&seen).to_ascii_lowercase();
let len: usize = head
.lines()
.find_map(|l| l.strip_prefix("content-length:"))
.and_then(|v| v.trim().parse().ok())
.unwrap_or(0);
let mut body = vec![0u8; len];
let _ = stream.read_exact(&mut body);
}
fn status_response(status: &str, extra: &[&str]) -> String {
let body = "{\"error\":\"nope\"}";
let mut head = format!("HTTP/1.1 {status}\r\nContent-Length: {}\r\n", body.len());
for line in extra {
head.push_str(line);
head.push_str("\r\n");
}
format!("{head}\r\n{body}")
}
fn stream_response(events: &str) -> String {
format!("HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n{events}")
}
#[allow(clippy::needless_update)] fn request() -> CompletionRequest {
CompletionRequest {
system: "s".into(),
user: "u".into(),
tools: Vec::new(),
..Default::default()
}
}
fn failure(result: Result<CompletionResponse>) -> (Kind, Option<u16>, Option<Duration>) {
match result {
Err(Error::Provider {
kind,
status,
retry_after,
..
}) => (kind, status, retry_after),
other => panic!("expected a provider error, got {other:?}"),
}
}
fn openrouter(url: &str) -> OpenRouter {
OpenRouter::at(url, Duration::from_secs(1))
}
#[tokio::test]
async fn a_refused_connection_is_transport() {
let dead = TcpListener::bind("127.0.0.1:0").unwrap();
let url = format!("http://{}/v1", dead.local_addr().unwrap());
drop(dead);
let (kind, status, _) = failure(openrouter(&url).complete(request()).await);
assert!(
matches!(kind, Kind::Transport | Kind::Timeout),
"a connection that never opened must be Transport or Timeout, got {kind:?}"
);
assert!(
kind.is_retryable(),
"a connection that never opened is worth another attempt"
);
assert_eq!(
status, None,
"a connection that never happened has no status"
);
assert!(kind.is_retryable());
}
#[tokio::test]
async fn a_server_that_accepts_and_never_answers_ends_as_a_timeout() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
std::thread::spawn(move || {
let held: Vec<_> = listener.incoming().filter_map(|s| s.ok()).collect();
std::thread::sleep(Duration::from_secs(30));
drop(held);
});
let started = std::time::Instant::now();
let (kind, status, _) = failure(
OpenRouter::at(&url, Duration::from_millis(300))
.complete(request())
.await,
);
assert_eq!(kind, Kind::Timeout);
assert_eq!(status, None);
assert!(kind.is_retryable());
assert!(
started.elapsed() < Duration::from_secs(10),
"the deadline, not the server, ended the call"
);
}
#[tokio::test]
async fn a_rate_limit_without_a_retry_after_is_rate_limited_and_carries_no_wait() {
let url = serve(status_response("429 Too Many Requests", &[]));
let (kind, status, retry_after) = failure(openrouter(&url).complete(request()).await);
assert_eq!(kind, Kind::RateLimited);
assert_eq!(status, Some(429));
assert_eq!(retry_after, None);
}
#[tokio::test]
async fn a_rate_limit_with_delta_seconds_keeps_the_wait_the_server_asked_for() {
let url = serve(status_response(
"429 Too Many Requests",
&["Retry-After: 11"],
));
let (kind, status, retry_after) = failure(openrouter(&url).complete(request()).await);
assert_eq!(kind, Kind::RateLimited);
assert_eq!(status, Some(429));
assert_eq!(retry_after, Some(Duration::from_secs(11)));
}
#[tokio::test]
async fn a_rate_limit_with_an_http_date_keeps_the_wait_until_that_date() {
let header = format!("Retry-After: {}", http_date(unix_now() + 45));
let url = serve(status_response("429 Too Many Requests", &[&header]));
let (kind, status, retry_after) = failure(openrouter(&url).complete(request()).await);
assert_eq!(kind, Kind::RateLimited);
assert_eq!(status, Some(429));
let waited = retry_after.expect("the date is a wait");
assert!(
waited > Duration::from_secs(40) && waited <= Duration::from_secs(45),
"{waited:?}"
);
}
#[tokio::test]
async fn a_server_error_is_server_and_retryable() {
let url = serve(status_response("503 Service Unavailable", &[]));
let (kind, status, _) = failure(openrouter(&url).complete(request()).await);
assert_eq!(kind, Kind::Server);
assert_eq!(status, Some(503));
assert!(kind.is_retryable());
}
#[tokio::test]
async fn a_rejected_key_is_auth_and_not_retryable() {
let url = serve(status_response("401 Unauthorized", &[]));
let (kind, status, _) = failure(openrouter(&url).complete(request()).await);
assert_eq!(kind, Kind::Auth);
assert_eq!(status, Some(401));
assert!(!kind.is_retryable(), "a wrong key stays wrong");
}
#[tokio::test]
async fn a_bad_request_is_request_and_not_retryable() {
let url = serve(status_response("400 Bad Request", &[]));
let (kind, status, _) = failure(openrouter(&url).complete(request()).await);
assert_eq!(kind, Kind::Request);
assert_eq!(status, Some(400));
assert!(!kind.is_retryable(), "the same request fails the same way");
}
#[tokio::test]
async fn a_stream_that_parses_to_nothing_is_malformed_not_an_empty_answer() {
let url = serve(stream_response(
"data: not json at all\n\ndata: {\"unterminated\n\n",
));
let (kind, status, _) = failure(openrouter(&url).complete(request()).await);
assert_eq!(kind, Kind::Malformed);
assert_eq!(status, None, "the status was fine; the body was not");
assert!(kind.is_retryable(), "re-asking is cheap");
}
#[tokio::test]
async fn a_stream_with_text_and_no_tool_call_stays_a_quiet_success() {
let url = serve(stream_response(
"data: {\"choices\":[{\"delta\":{\"content\":\"done\"}}]}\n\ndata: [DONE]\n\n",
));
let out = openrouter(&url).complete(request()).await.unwrap();
assert_eq!(out.text.as_deref(), Some("done"));
assert!(out.tool_calls.is_empty());
}
#[tokio::test]
async fn a_stream_carrying_only_usage_is_not_malformed() {
let url = serve(stream_response(
"data: {\"choices\":[],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":0,\"total_tokens\":3}}\n\ndata: [DONE]\n\n",
));
let out = openrouter(&url).complete(request()).await.unwrap();
assert_eq!(out.usage.unwrap().total_tokens, 3);
}
#[tokio::test]
async fn anthropics_own_stream_shape_is_held_to_the_same_two_meanings() {
let empty = serve(stream_response("data: {\"type\":\"whatever\"}\n\n"));
let (kind, _, _) = failure(
Anthropic::at(&empty, Duration::from_secs(1))
.complete(request())
.await,
);
assert_eq!(kind, Kind::Malformed);
let quiet = serve(stream_response(
"data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hi\"}}\n\ndata: {\"type\":\"message_stop\"}\n\n",
));
let out = Anthropic::at(&quiet, Duration::from_secs(1))
.complete(request())
.await
.unwrap();
assert_eq!(out.text.as_deref(), Some("hi"));
assert!(out.tool_calls.is_empty());
}
#[tokio::test]
async fn all_three_providers_map_a_status_to_the_same_kind() {
for (status, want) in [
("429 Too Many Requests", Kind::RateLimited),
("503 Service Unavailable", Kind::Server),
("403 Forbidden", Kind::Auth),
("404 Not Found", Kind::Request),
] {
let url = serve(status_response(status, &["Retry-After: 3"]));
let code: u16 = status[..3].parse().unwrap();
let timeout = Duration::from_secs(1);
let seen = [
failure(OpenRouter::at(&url, timeout).complete(request()).await),
failure(OpenAi::at(&url, timeout).complete(request()).await),
failure(Anthropic::at(&url, timeout).complete(request()).await),
];
for observed in &seen {
assert_eq!(
*observed,
(want, Some(code), Some(Duration::from_secs(3))),
"{status}"
);
}
}
}
}
#[cfg(all(test, feature = "media"))]
mod media_tests {
use super::*;
const PNG: &[u8] = &[
0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44,
0x52,
];
struct Blind;
impl Provider for Blind {
async fn complete(&self, _r: CompletionRequest) -> Result<CompletionResponse> {
unreachable!("the refusal must happen before the request is sent")
}
fn name(&self) -> &str {
"blind"
}
}
struct Seeing;
impl Provider for Seeing {
async fn complete(&self, _r: CompletionRequest) -> Result<CompletionResponse> {
Ok(CompletionResponse::default())
}
fn name(&self) -> &str {
"seeing"
}
fn accepts_images(&self) -> bool {
true
}
}
#[allow(clippy::needless_update)] fn with_image() -> CompletionRequest {
CompletionRequest {
user: "what is in this picture".into(),
media: vec![Media::image("image/png", PNG).unwrap()],
..Default::default()
}
}
#[test]
fn an_unsupported_media_type_is_refused_at_construction() {
let err = Media::image("image/tiff", PNG).unwrap_err();
assert!(
matches!(&err, Error::Config(m) if m.contains("image/tiff")),
"{err:?}"
);
}
#[test]
fn every_documented_image_type_is_accepted() {
for t in IMAGE_MEDIA_TYPES {
assert!(Media::image(t, PNG).is_ok(), "{t} should be constructible");
}
}
#[test]
fn a_provider_that_does_not_accept_images_refuses_before_the_request_is_sent() {
let err =
ensure_media_accepted("blind", Blind.accepts_images(), &with_image()).unwrap_err();
let Error::Config(message) = &err else {
panic!("expected a configuration error, got {err:?}");
};
assert!(message.contains("does not accept image input"), "{message}");
assert!(
message.contains("no request was sent"),
"the caller must be told nothing was spent: {message}"
);
}
#[test]
fn a_provider_that_accepts_images_is_not_refused() {
assert!(ensure_media_accepted("seeing", Seeing.accepts_images(), &with_image()).is_ok());
}
#[test]
fn a_text_only_request_is_never_refused_even_by_a_blind_provider() {
#[allow(clippy::needless_update)] let text_only = CompletionRequest {
user: "no picture here".into(),
..Default::default()
};
assert!(ensure_media_accepted("blind", Blind.accepts_images(), &text_only).is_ok());
}
#[test]
fn byte_len_reports_the_decoded_size_not_the_encoded_one() {
let m = Media::image("image/png", PNG).unwrap();
assert_eq!(m.byte_len(), PNG.len());
assert!(m.base64.len() > PNG.len(), "base64 grows the payload");
}
#[test]
fn the_digest_distinguishes_images_and_is_stable() {
let a = Media::image("image/png", PNG).unwrap();
let b = Media::image("image/png", PNG).unwrap();
let c = Media::image("image/png", &[0xff, 0x00]).unwrap();
assert_eq!(a.digest(), b.digest(), "same bytes, same digest");
assert_ne!(a.digest(), c.digest(), "different bytes, different digest");
assert_eq!(a.digest().len(), 16);
}
#[test]
fn a_path_maps_to_the_media_type_its_extension_names() {
assert_eq!(Media::media_type_for("shot.PNG"), Some("image/png"));
assert_eq!(Media::media_type_for("a/b/photo.jpeg"), Some("image/jpeg"));
assert_eq!(Media::media_type_for("scan.webp"), Some("image/webp"));
assert_eq!(Media::media_type_for("report.pdf"), None);
assert_eq!(Media::media_type_for("clip.mp4"), None);
assert_eq!(Media::media_type_for("noextension"), None);
}
}