use std::time::Duration;
use thiserror::Error;
use time::{OffsetDateTime, PrimitiveDateTime};
#[derive(Debug, Error)]
pub enum ProviderError {
#[error("provider transport error: {0}")]
Transport(#[source] reqwest::Error),
#[error("retryable provider error: {message}")]
Retryable {
message: String,
delay: Option<Duration>,
},
#[error("provider does not support capability: {0}")]
UnsupportedCapability(String),
#[error("{message}", message = provider_http_error(.status, .body))]
Http {
status: reqwest::StatusCode,
body: String,
retry_after: Option<Duration>,
},
#[error("failed to decode provider response: {0}")]
Decode(#[source] reqwest::Error),
#[error("failed to serialize provider request: {0}")]
Serialize(#[source] serde_json::Error),
#[error("failed to deserialize provider payload: {0}")]
Deserialize(#[source] serde_json::Error),
#[error("invalid provider request: {0}")]
InvalidRequest(String),
#[error("invalid provider response: {0}")]
InvalidResponse(String),
#[error("malformed provider stream: {0}")]
MalformedStream(String),
}
impl ProviderError {
pub async fn from_http_response(response: reqwest::Response) -> Self {
let status = response.status();
let retry_after = retry_after_from_headers(response.headers());
Self::Http {
status,
body: response.text().await.unwrap_or_default(),
retry_after,
}
}
pub fn retry_after(&self) -> Option<Duration> {
match self {
Self::Http { retry_after, .. } => *retry_after,
Self::Retryable { delay, .. } => *delay,
_ => None,
}
}
}
fn provider_http_error(status: &reqwest::StatusCode, body: &str) -> String {
if body.trim().is_empty() {
format!("provider returned HTTP {status}")
} else {
format!("provider returned HTTP {status}: {body}")
}
}
pub(crate) fn retry_after_from_headers(headers: &reqwest::header::HeaderMap) -> Option<Duration> {
retry_after_from_header_value(headers.get(reqwest::header::RETRY_AFTER)?.to_str().ok()?)
}
pub(crate) fn retry_after_from_header_value(value: &str) -> Option<Duration> {
parse_retry_after(value, OffsetDateTime::now_utc())
}
fn parse_retry_after(value: &str, now: OffsetDateTime) -> Option<Duration> {
let value = value.trim();
if value.is_empty() {
return None;
}
if let Ok(seconds) = value.parse::<u64>() {
return Some(Duration::from_secs(seconds));
}
let deadline = parse_http_date(value)?;
Some((deadline - now).try_into().unwrap_or(Duration::ZERO))
}
fn parse_http_date(value: &str) -> Option<OffsetDateTime> {
let format = time::format_description::parse_borrowed::<2>(
"[weekday repr:short], [day] [month repr:short] [year] [hour]:[minute]:[second] GMT",
)
.ok()?;
PrimitiveDateTime::parse(value, format.as_slice())
.ok()
.map(PrimitiveDateTime::assume_utc)
}
#[cfg(test)]
mod tests {
use super::*;
fn now() -> OffsetDateTime {
time::Date::from_calendar_date(1994, time::Month::November, 6)
.expect("a real date")
.with_hms(8, 49, 37)
.expect("a real time")
.assume_utc()
}
#[test]
fn a_delay_in_seconds_is_read_as_seconds() {
assert_eq!(
parse_retry_after("30", now()),
Some(Duration::from_secs(30))
);
assert_eq!(
parse_retry_after(" 30 ", now()),
Some(Duration::from_secs(30)),
"surrounding whitespace is not part of the value"
);
}
#[test]
fn an_http_date_is_read_as_the_wait_until_it() {
assert_eq!(
parse_retry_after("Sun, 06 Nov 1994 08:50:37 GMT", now()),
Some(Duration::from_secs(60))
);
}
#[test]
fn an_http_date_that_has_passed_means_retry_now() {
assert_eq!(
parse_retry_after("Sun, 06 Nov 1994 08:49:00 GMT", now()),
Some(Duration::ZERO),
"the server answered; the answer is that the wait is over"
);
}
#[test]
fn an_unparseable_value_is_no_hint_at_all() {
assert_eq!(parse_retry_after("", now()), None);
assert_eq!(parse_retry_after("soon", now()), None);
assert_eq!(parse_retry_after("-5", now()), None);
}
#[test]
fn a_rate_limited_response_reports_the_interval_it_named() {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(reqwest::header::RETRY_AFTER, "42".parse().expect("valid"));
assert_eq!(
retry_after_from_headers(&headers),
Some(Duration::from_secs(42))
);
}
#[tokio::test]
async fn a_429_response_becomes_an_error_that_still_knows_the_interval() {
let response = http::Response::builder()
.status(429)
.header("retry-after", "60")
.body("rate limit exceeded")
.expect("a response");
let error = ProviderError::from_http_response(reqwest::Response::from(response)).await;
let ProviderError::Http {
status,
body,
retry_after,
} = error
else {
panic!("an unsuccessful response is an Http error");
};
assert_eq!(status, reqwest::StatusCode::TOO_MANY_REQUESTS);
assert_eq!(body, "rate limit exceeded");
assert_eq!(retry_after, Some(Duration::from_secs(60)));
}
#[tokio::test]
async fn a_response_without_the_header_asks_for_nothing() {
let response = http::Response::builder()
.status(503)
.body("upstream is restarting")
.expect("a response");
let error = ProviderError::from_http_response(reqwest::Response::from(response)).await;
assert_eq!(error.retry_after(), None);
}
#[test]
fn both_retryable_shapes_answer_the_same_question() {
let http = ProviderError::Http {
status: reqwest::StatusCode::TOO_MANY_REQUESTS,
body: String::new(),
retry_after: Some(Duration::from_secs(20)),
};
let retryable = ProviderError::Retryable {
message: "connection closed".to_string(),
delay: Some(Duration::from_millis(750)),
};
let silent = ProviderError::InvalidRequest("bad model".to_string());
assert_eq!(http.retry_after(), Some(Duration::from_secs(20)));
assert_eq!(retryable.retry_after(), Some(Duration::from_millis(750)));
assert_eq!(silent.retry_after(), None);
}
}