use reqwest::StatusCode;
use reqwest::header::HeaderMap;
use tokio::time::Duration;
use crate::client::{ClientError, PError};
pub(crate) fn retry_after_from_headers(headers: &HeaderMap) -> Option<Duration> {
let raw = headers.get(reqwest::header::RETRY_AFTER)?.to_str().ok()?;
let secs: f64 = raw.trim().parse().ok()?;
(secs.is_finite() && secs >= 0.0).then(|| Duration::from_secs_f64(secs))
}
pub(crate) fn classify_status(
status: StatusCode,
retry_after: Option<Duration>,
message: String,
) -> ClientError {
match status.as_u16() {
401 => ClientError::Unauthorized { message },
429 => ClientError::RateLimited { retry_after },
s if (500..=599).contains(&s) => ClientError::RetryableServer { status: s, message },
_ => ClientError::Api(format!("{status}: {message}")),
}
}
pub(crate) async fn map_progenitor_err<E: std::fmt::Debug>(e: PError<E>) -> ClientError {
match e {
PError::CommunicationError(e) => ClientError::Http(e),
PError::InvalidResponsePayload(_, e) => ClientError::Serde(e),
PError::InvalidRequest(s) => ClientError::Api(s),
PError::ErrorResponse(rv) => {
let status = rv.status();
let retry_after = retry_after_from_headers(rv.headers());
let message = format!("{:?}", rv.into_inner());
classify_status(status, retry_after, message)
}
PError::UnexpectedResponse(r) => {
let status = r.status();
let retry_after = retry_after_from_headers(r.headers());
let body = r
.text()
.await
.unwrap_or_else(|e| format!("<body read error: {e}>"));
classify_status(
status,
retry_after,
format!("unexpected response; body: {body}"),
)
}
PError::ResponseBodyError(e) => ClientError::Http(e),
PError::InvalidUpgrade(e) => ClientError::Http(e),
PError::Custom(s) => ClientError::Api(s),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn classify_status_maps_each_class() {
assert!(matches!(
classify_status(StatusCode::UNAUTHORIZED, None, "x".into()),
ClientError::Unauthorized { .. }
));
assert!(matches!(
classify_status(
StatusCode::TOO_MANY_REQUESTS,
Some(Duration::from_secs(2)),
"x".into()
),
ClientError::RateLimited {
retry_after: Some(_)
}
));
assert!(matches!(
classify_status(StatusCode::INTERNAL_SERVER_ERROR, None, "x".into()),
ClientError::RetryableServer { status: 500, .. }
));
assert!(matches!(
classify_status(StatusCode::BAD_REQUEST, None, "x".into()),
ClientError::Api(_)
));
}
#[test]
fn error_predicates() {
assert!(
ClientError::Unauthorized {
message: String::new()
}
.is_unauthorized()
);
assert!(ClientError::RateLimited { retry_after: None }.is_retryable());
assert!(
ClientError::RetryableServer {
status: 503,
message: String::new()
}
.is_retryable()
);
assert!(!ClientError::Api("nope".into()).is_retryable());
assert_eq!(
ClientError::RateLimited {
retry_after: Some(Duration::from_secs(1))
}
.retry_after(),
Some(Duration::from_secs(1))
);
}
#[test]
fn retry_after_header_parses_seconds() {
let mut headers = HeaderMap::new();
headers.insert(reqwest::header::RETRY_AFTER, "3".parse().unwrap());
assert_eq!(
retry_after_from_headers(&headers),
Some(Duration::from_secs(3))
);
let empty = HeaderMap::new();
assert_eq!(retry_after_from_headers(&empty), None);
}
}