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,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct CompletionRequest {
pub system: String,
pub user: String,
pub tools: Vec<ToolSpec>,
}
#[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"
}
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}")
}
fn request() -> CompletionRequest {
CompletionRequest {
system: "s".into(),
user: "u".into(),
tools: Vec::new(),
}
}
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}"
);
}
}
}
}