use std::fmt;
use std::time::Duration;
use serde_json::Value;
use crate::REQUEST_ID_HEADER;
const MAX_RAW_BODY_IN_MESSAGE: usize = 200;
#[derive(Debug, Clone)]
pub struct ApiError {
pub status: u16,
pub kind: ApiErrorKind,
pub message: String,
pub body: Option<Value>,
pub headers: reqwest::header::HeaderMap,
pub request_id: Option<String>,
pub endpoint: String,
pub retry_after: Option<Duration>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum ApiErrorKind {
BadRequest,
Authentication,
PermissionDenied,
NotFound,
UnprocessableEntity,
RateLimit,
InternalServer,
Other,
}
impl ApiErrorKind {
pub fn from_status(status: u16) -> Self {
match status {
400 => Self::BadRequest,
401 => Self::Authentication,
403 => Self::PermissionDenied,
404 => Self::NotFound,
422 => Self::UnprocessableEntity,
429 => Self::RateLimit,
s if (500..=599).contains(&s) => Self::InternalServer,
_ => Self::Other,
}
}
}
impl fmt::Display for ApiErrorKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let name = match self {
Self::BadRequest => "bad request",
Self::Authentication => "authentication",
Self::PermissionDenied => "permission denied",
Self::NotFound => "not found",
Self::UnprocessableEntity => "unprocessable entity",
Self::RateLimit => "rate limit",
Self::InternalServer => "internal server",
Self::Other => "other",
};
f.write_str(name)
}
}
impl fmt::Display for ApiError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}: {} {}", self.endpoint, self.status, self.message)?;
if let Some(request_id) = &self.request_id {
write!(f, " (request_id={request_id})")?;
}
Ok(())
}
}
impl std::error::Error for ApiError {}
pub(super) fn extract_message(body: &Value) -> Option<String> {
match body {
Value::String(s) => {
if s.is_empty() {
return None;
}
Some(truncate_chars(s))
}
Value::Object(map) => {
let error = map.get("error");
if let Some(Value::String(error)) = error {
return Some(error.clone());
}
if let Some(Value::Object(error)) = error {
if let Some(Value::String(message)) = error.get("message") {
return Some(message.clone());
}
}
if let Some(Value::String(message)) = map.get("message") {
return Some(message.clone());
}
let detail = map.get("detail");
if let Some(Value::String(detail)) = detail {
return Some(detail.clone());
}
if let Some(Value::Object(detail)) = detail {
if let Some(Value::String(message)) = detail.get("message") {
return Some(message.clone());
}
}
if let Some(Value::Array(entries)) = detail {
let mut parts = Vec::new();
for entry in entries {
let Value::Object(entry) = entry else {
continue;
};
let Some(Value::String(msg)) = entry.get("msg") else {
continue;
};
let path = match entry.get("loc") {
Some(Value::Array(location)) => location
.iter()
.filter(|item| **item != Value::String("body".into()))
.map(|item| match item {
Value::String(s) => s.clone(),
other => other.to_string(),
})
.collect::<Vec<_>>()
.join("."),
_ => String::new(),
};
parts.push(if path.is_empty() {
msg.clone()
} else {
format!("{path}: {msg}")
});
}
if !parts.is_empty() {
return Some(parts.join("; "));
}
}
None
}
_ => None,
}
}
fn truncate_chars(raw: &str) -> String {
if raw.chars().count() > MAX_RAW_BODY_IN_MESSAGE {
let mut truncated: String = raw.chars().take(MAX_RAW_BODY_IN_MESSAGE).collect();
truncated.push('…');
truncated
} else {
raw.to_owned()
}
}
pub(super) fn truncated_body(body: &Value) -> String {
let raw = match body {
Value::String(s) => s.clone(),
other => other.to_string(),
};
truncate_chars(&raw)
}
pub(super) fn api_error_message(body: &Option<Value>) -> String {
if let Some(body) = body {
if let Some(message) = extract_message(body) {
return message;
}
if matches!(body, Value::String(s) if s.is_empty()) || (body.is_null() && false) {
return "status code (no body)".into();
}
return truncated_body(body);
}
"status code (no body)".into()
}
#[derive(Debug)]
#[non_exhaustive]
pub enum Error {
Config(String),
InvalidRequest(String),
Api(Box<ApiError>),
Connection {
message: String,
source: Option<Box<dyn std::error::Error + Send + Sync>>,
},
Timeout {
timeout: Duration,
},
ResponseValidation {
status: u16,
field_path: String,
request_id: Option<String>,
endpoint: String,
},
}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Config(message) => write!(f, "configuration error: {message}"),
Self::InvalidRequest(message) => write!(f, "invalid request: {message}"),
Self::Api(error) => write!(f, "{error}"),
Self::Connection { message, .. } => write!(f, "connection error: {message}"),
Self::Timeout { timeout } => {
write!(f, "request timed out (timeout={:?})", timeout)
}
Self::ResponseValidation { field_path, .. } => {
write!(f, "invalid response data at '{field_path}'")
}
}
}
}
impl std::error::Error for Error {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Api(error) => Some(error.as_ref()),
Self::Connection {
source: Some(source),
..
} => Some(source.as_ref()),
_ => None,
}
}
}
impl From<ApiError> for Error {
fn from(error: ApiError) -> Self {
Self::Api(Box::new(error))
}
}
impl Error {
pub fn is_timeout(&self) -> bool {
matches!(self, Self::Timeout { .. })
}
pub fn is_connection(&self) -> bool {
matches!(self, Self::Connection { .. } | Self::Timeout { .. })
}
pub fn status(&self) -> Option<u16> {
match self {
Self::Api(error) => Some(error.status),
Self::ResponseValidation { status, .. } => Some(*status),
_ => None,
}
}
pub fn request_id(&self) -> Option<&str> {
match self {
Self::Api(error) => error.request_id.as_deref(),
Self::ResponseValidation { request_id, .. } => request_id.as_deref(),
_ => None,
}
}
pub fn as_api(&self) -> Option<&ApiError> {
match self {
Self::Api(error) => Some(error),
_ => None,
}
}
}
pub(super) fn request_id_of(headers: &reqwest::header::HeaderMap) -> Option<String> {
headers
.get(REQUEST_ID_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::to_owned)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn api_error(status: u16, body: Value, request_id: Option<&str>) -> ApiError {
let mut headers = reqwest::header::HeaderMap::new();
if let Some(request_id) = request_id {
headers.insert(
REQUEST_ID_HEADER,
reqwest::header::HeaderValue::from_str(request_id).unwrap(),
);
}
let body = if body.is_null() { None } else { Some(body) };
ApiError {
status,
kind: ApiErrorKind::from_status(status),
message: api_error_message(&body),
body,
headers,
request_id: request_id.map(str::to_owned),
endpoint: "POST https://api.typesafe.ai/v1/systemone".into(),
retry_after: None,
}
}
#[test]
fn kind_per_status() {
for (status, kind) in [
(400, ApiErrorKind::BadRequest),
(401, ApiErrorKind::Authentication),
(403, ApiErrorKind::PermissionDenied),
(404, ApiErrorKind::NotFound),
(422, ApiErrorKind::UnprocessableEntity),
(429, ApiErrorKind::RateLimit),
(500, ApiErrorKind::InternalServer),
(503, ApiErrorKind::InternalServer),
(599, ApiErrorKind::InternalServer),
(418, ApiErrorKind::Other),
] {
assert_eq!(ApiErrorKind::from_status(status), kind, "{status}");
}
}
#[test]
fn message_extraction_first_match_wins() {
let cases = [
(json!("plain string"), "plain string"),
(json!({"error": "an error"}), "an error"),
(
json!({"error": {"message": "nested"}, "message": "flat"}),
"nested",
),
(json!({"message": "flat"}), "flat"),
(json!({"detail": "details"}), "details"),
(
json!({"detail": {"message": "detail message"}}),
"detail message",
),
(
json!({"detail": [
{"loc": ["body", "questions", "urgency", "score", "criteria"], "msg": "Field required"},
{"loc": ["body", "state"], "msg": "Bad state"}
]}),
"questions.urgency.score.criteria: Field required; state: Bad state",
),
];
for (body, expected) in cases {
assert_eq!(extract_message(&body).as_deref(), Some(expected), "{body}");
}
}
#[test]
fn plain_string_body_is_truncated_to_200_chars() {
let long = "é".repeat(250);
let message = extract_message(&Value::String(long.clone())).unwrap();
assert_eq!(message.chars().count(), 201);
assert!(message.ends_with('…'));
assert_eq!(message.chars().filter(|c| *c == 'é').count(), 200);
assert_eq!(message.len(), 200 * 2 + 3);
let exact = "é".repeat(200);
let message = extract_message(&Value::String(exact.clone())).unwrap();
assert_eq!(message, exact);
let raw_html = format!("<html>{}</html>", "z".repeat(300));
let message = extract_message(&Value::String(raw_html)).unwrap();
assert_eq!(message.chars().count(), 201);
assert!(message.ends_with('…'));
}
#[test]
fn fastapi_detail_renders_loc_without_body() {
let error = api_error(
422,
json!({"detail": [
{"loc": ["body", "questions", "urgency", "score", "criteria"], "msg": "Field required", "type": "missing"}
]}),
None,
);
assert_eq!(
error.message,
"questions.urgency.score.criteria: Field required"
);
}
#[test]
fn message_falls_back_to_raw_body_truncated() {
let long = "x".repeat(300);
let error = api_error(500, json!({"unexpected": long.clone()}), None);
assert_eq!(error.message.chars().count(), 201);
assert!(error.message.ends_with('…'));
assert!(error.message.starts_with("{\"unexpected\":\"xxx"));
let exact = "y".repeat(183);
let error = api_error(500, json!({"unexpected": exact}), None);
assert_eq!(
error.message,
"{\"unexpected\":\"".to_owned() + &"y".repeat(183) + "\"}"
);
let raw = "z".repeat(250);
let error = api_error(500, Value::String(raw.clone()), None);
assert_eq!(error.message, "z".repeat(200) + "…");
}
#[test]
fn empty_body_message() {
let error = api_error(404, Value::Null, None);
assert_eq!(error.message, "status code (no body)");
let error = api_error(404, json!("whoops"), None);
assert_eq!(error.message, "whoops");
}
#[test]
fn display_format() {
let error = api_error(429, json!({"error": "too fast"}), Some("req-123"));
assert_eq!(
error.to_string(),
"POST https://api.typesafe.ai/v1/systemone: 429 too fast (request_id=req-123)"
);
let error = api_error(500, json!("boom"), None);
assert_eq!(
error.to_string(),
"POST https://api.typesafe.ai/v1/systemone: 500 boom"
);
}
#[test]
fn error_helpers() {
let error = Error::from(api_error(429, json!("rl"), Some("req-1")));
assert_eq!(error.status(), Some(429));
assert_eq!(error.request_id(), Some("req-1"));
assert!(error.as_api().is_some());
assert!(!error.is_connection());
assert!(!error.is_timeout());
let error = Error::Timeout {
timeout: Duration::from_secs(2),
};
assert!(error.is_timeout());
assert!(error.is_connection());
assert_eq!(error.status(), None);
let error = Error::Connection {
message: "reset".into(),
source: None,
};
assert!(error.is_connection());
assert!(!error.is_timeout());
}
}