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("provider context length exceeded: {message}", message = provider_http_error(.status, .body))]
ContextLengthExceeded {
status: reqwest::StatusCode,
body: String,
},
#[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());
let body = response.text().await.unwrap_or_default();
if is_context_overflow(status, &body) {
return Self::ContextLengthExceeded { status, body };
}
Self::Http {
status,
body,
retry_after,
}
}
pub fn is_context_length_exceeded(&self) -> bool {
matches!(self, Self::ContextLengthExceeded { .. })
}
pub fn retry_after(&self) -> Option<Duration> {
match self {
Self::Http { retry_after, .. } => *retry_after,
Self::Retryable { delay, .. } => *delay,
_ => None,
}
}
}
const CONTEXT_OVERFLOW_MARKERS: &[&str] = &[
"context_length_exceeded",
"maximum context length",
"context length exceeded",
"prompt is too long",
"exceed context limit",
"exceeds the maximum number of tokens",
"reduce the length of the messages",
"please reduce the length",
];
fn is_context_overflow(status: reqwest::StatusCode, body: &str) -> bool {
if !matches!(
status,
reqwest::StatusCode::BAD_REQUEST | reqwest::StatusCode::PAYLOAD_TOO_LARGE
) {
return false;
}
let body = body.to_ascii_lowercase();
CONTEXT_OVERFLOW_MARKERS
.iter()
.any(|marker| body.contains(marker))
}
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::{ProviderError, is_context_overflow};
#[test]
fn each_providers_way_of_saying_too_long_is_recognized() {
let bad_request = reqwest::StatusCode::BAD_REQUEST;
for body in [
r#"{"error":{"code":"context_length_exceeded","message":"This model's maximum context length is 128000 tokens"}}"#,
r#"{"type":"error","error":{"type":"invalid_request_error","message":"prompt is too long: 215000 tokens > 200000 maximum"}}"#,
r#"{"error":{"message":"The input token count exceeds the maximum number of tokens allowed"}}"#,
r#"{"object":"error","message":"This model's maximum context length is 8192 tokens. Please reduce the length of the messages."}"#,
] {
assert!(is_context_overflow(bad_request, body), "{body}");
}
}
#[test]
fn an_ordinary_bad_request_is_not_mistaken_for_an_overflow() {
assert!(!is_context_overflow(
reqwest::StatusCode::BAD_REQUEST,
r#"{"error":{"message":"unknown field `temperatur`"}}"#
));
assert!(!is_context_overflow(
reqwest::StatusCode::BAD_REQUEST,
r#"{"error":{"message":"tool `read` has an invalid input_schema"}}"#
));
}
#[test]
fn an_overflow_message_on_another_status_is_left_alone() {
assert!(!is_context_overflow(
reqwest::StatusCode::INTERNAL_SERVER_ERROR,
"maximum context length"
));
}
#[test]
fn an_overflow_error_answers_the_predicate_and_nothing_else_does() {
let overflow = ProviderError::ContextLengthExceeded {
status: reqwest::StatusCode::BAD_REQUEST,
body: "prompt is too long".to_string(),
};
let other = ProviderError::InvalidRequest("nope".to_string());
assert!(overflow.is_context_length_exceeded());
assert!(!other.is_context_length_exceeded());
}
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);
}
}