1use std::fmt;
4use std::sync::Arc;
5use std::time::Duration;
6
7use serde_json::Value;
8
9#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
11#[non_exhaustive]
12pub enum ErrorKind {
13 NoApiKey,
15 InvalidRequest,
18 BadRequest,
20 Authentication,
22 PermissionDenied,
24 NotFound,
26 UnprocessableEntity,
28 RateLimited,
30 Overloaded,
32 ServerError,
34 HttpError,
36 Timeout,
38 Connection,
40 InvalidResponse,
42}
43
44impl ErrorKind {
45 pub(crate) fn for_status(status: u16) -> ErrorKind {
46 match status {
47 400 => ErrorKind::BadRequest,
48 401 => ErrorKind::Authentication,
49 403 => ErrorKind::PermissionDenied,
50 404 => ErrorKind::NotFound,
51 422 => ErrorKind::UnprocessableEntity,
52 429 => ErrorKind::RateLimited,
53 529 => ErrorKind::Overloaded,
54 500..=599 => ErrorKind::ServerError,
55 _ => ErrorKind::HttpError,
56 }
57 }
58}
59
60#[derive(Clone, Debug)]
65pub struct Error {
66 kind: ErrorKind,
67 status: Option<u16>,
68 message: String,
69 body: Option<Value>,
70 request_id: Option<String>,
71 retry_after: Option<Duration>,
72 source: Option<Arc<dyn std::error::Error + Send + Sync>>,
73}
74
75impl Error {
76 pub fn kind(&self) -> ErrorKind {
77 self.kind
78 }
79
80 pub fn status(&self) -> Option<u16> {
82 self.status
83 }
84
85 pub fn message(&self) -> &str {
87 &self.message
88 }
89
90 pub fn body(&self) -> Option<&Value> {
92 self.body.as_ref()
93 }
94
95 pub fn request_id(&self) -> Option<&str> {
97 self.request_id.as_deref()
98 }
99
100 pub fn retry_after(&self) -> Option<Duration> {
103 self.retry_after
104 }
105
106 fn new(kind: ErrorKind, message: impl Into<String>) -> Self {
107 Error {
108 kind,
109 status: None,
110 message: message.into(),
111 body: None,
112 request_id: None,
113 retry_after: None,
114 source: None,
115 }
116 }
117
118 pub(crate) fn no_api_key() -> Self {
119 Error::new(
120 ErrorKind::NoApiKey,
121 "no API key: pass one to Config::api_key or set TYPESAFE_API_KEY",
122 )
123 }
124
125 pub(crate) fn invalid_request(message: impl Into<String>) -> Self {
126 Error::new(ErrorKind::InvalidRequest, message)
127 }
128
129 pub(crate) fn invalid_response(
130 message: impl Into<String>,
131 status: u16,
132 body: Option<Value>,
133 request_id: Option<String>,
134 ) -> Self {
135 Error {
136 status: Some(status),
137 body,
138 request_id,
139 ..Error::new(ErrorKind::InvalidResponse, message)
140 }
141 }
142
143 pub(crate) fn from_response(
144 status: u16,
145 body: &[u8],
146 request_id: Option<String>,
147 retry_after: Option<Duration>,
148 ) -> Self {
149 let (body, message) = read_error_body(body);
150 Error {
151 status: Some(status),
152 body,
153 request_id,
154 retry_after,
155 ..Error::new(ErrorKind::for_status(status), message)
156 }
157 }
158
159 pub(crate) fn from_reqwest(error: reqwest::Error) -> Self {
160 let kind = if error.is_timeout() {
161 ErrorKind::Timeout
162 } else if error.is_builder() {
163 ErrorKind::InvalidRequest
164 } else {
165 ErrorKind::Connection
166 };
167 Error {
168 source: Some(Arc::new(error.without_url())),
169 ..Error::new(kind, describe_reqwest(kind))
170 }
171 }
172}
173
174fn describe_reqwest(kind: ErrorKind) -> &'static str {
175 match kind {
176 ErrorKind::Timeout => "the request timed out",
177 ErrorKind::InvalidRequest => "the request could not be built",
178 _ => "the connection failed",
179 }
180}
181
182impl fmt::Display for Error {
183 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
184 if let Some(status) = self.status {
185 write!(f, "{status} ")?;
186 }
187 f.write_str(&self.message)?;
188 if let Some(source) = &self.source {
189 write!(f, ": {source}")?;
190 }
191 if let Some(id) = &self.request_id {
192 write!(f, " (request_id={id})")?;
193 }
194 Ok(())
195 }
196}
197
198impl std::error::Error for Error {
199 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
200 self.source
201 .as_deref()
202 .map(|e| e as &(dyn std::error::Error + 'static))
203 }
204}
205
206const MAX_BODY_IN_MESSAGE: usize = 200;
207
208fn read_error_body(bytes: &[u8]) -> (Option<Value>, String) {
210 let text = String::from_utf8_lossy(bytes);
211 if text.trim().is_empty() {
212 return (None, "(no body)".into());
213 }
214 match serde_json::from_slice::<Value>(bytes) {
215 Ok(json) => {
216 let message = extract_message(&json).unwrap_or_else(|| truncate(&json.to_string()));
217 (Some(json), message)
218 }
219 Err(_) => (Some(Value::String(text.to_string())), truncate(text.trim())),
220 }
221}
222
223fn extract_message(body: &Value) -> Option<String> {
226 let str_at = |value: Option<&Value>| value.and_then(Value::as_str).map(str::to_string);
227 if let Value::String(s) = body {
228 return (!s.is_empty()).then(|| s.clone());
229 }
230 let error = body.get("error");
231 let detail = body.get("detail");
232 str_at(error)
233 .or_else(|| str_at(error.and_then(|e| e.get("message"))))
234 .or_else(|| str_at(body.get("message")))
235 .or_else(|| str_at(detail))
236 .or_else(|| str_at(detail.and_then(|d| d.get("message"))))
237 .or_else(|| {
238 detail
239 .and_then(Value::as_array)
240 .and_then(|d| validation_message(d))
241 })
242}
243
244fn validation_message(entries: &[Value]) -> Option<String> {
245 let parts: Vec<String> = entries
246 .iter()
247 .filter_map(|entry| {
248 let msg = entry.get("msg")?.as_str()?;
249 let path = match entry.get("loc") {
250 Some(Value::Array(loc)) => loc
251 .iter()
252 .filter(|part| part.as_str() != Some("body"))
253 .map(|part| match part {
254 Value::String(s) => s.clone(),
255 other => other.to_string(),
256 })
257 .collect::<Vec<_>>()
258 .join("."),
259 Some(Value::String(s)) if s != "body" => s.clone(),
260 _ => String::new(),
261 };
262 Some(if path.is_empty() {
263 msg.to_string()
264 } else {
265 format!("{path}: {msg}")
266 })
267 })
268 .collect();
269 (!parts.is_empty()).then(|| parts.join("; "))
270}
271
272fn truncate(raw: &str) -> String {
273 match raw.char_indices().nth(MAX_BODY_IN_MESSAGE) {
274 Some((cut, _)) => format!("{}…", &raw[..cut]),
275 None => raw.to_string(),
276 }
277}
278
279#[cfg(test)]
280mod tests {
281 use super::*;
282 use serde_json::json;
283
284 fn message(body: Value) -> String {
285 Error::from_response(400, body.to_string().as_bytes(), None, None).message
286 }
287
288 #[test]
289 fn maps_statuses() {
290 let kinds: Vec<ErrorKind> = [400, 401, 403, 404, 422, 429, 529, 500, 503, 418]
291 .into_iter()
292 .map(ErrorKind::for_status)
293 .collect();
294 assert_eq!(
295 kinds,
296 [
297 ErrorKind::BadRequest,
298 ErrorKind::Authentication,
299 ErrorKind::PermissionDenied,
300 ErrorKind::NotFound,
301 ErrorKind::UnprocessableEntity,
302 ErrorKind::RateLimited,
303 ErrorKind::Overloaded,
304 ErrorKind::ServerError,
305 ErrorKind::ServerError,
306 ErrorKind::HttpError,
307 ]
308 );
309 }
310
311 #[test]
312 fn reads_the_message_shapes() {
313 assert_eq!(message(json!({"error": "bad key"})), "bad key");
314 assert_eq!(message(json!({"error": {"message": "nested"}})), "nested");
315 assert_eq!(message(json!({"message": "plain"})), "plain");
316 assert_eq!(message(json!({"detail": "detailed"})), "detailed");
317 assert_eq!(message(json!({"detail": {"message": "deep"}})), "deep");
318 assert_eq!(
319 message(json!({"detail": [
320 {"loc": ["body", "questions", "tone", "criteria"], "msg": "field required"},
321 {"loc": ["body", "state"], "msg": "must not be empty"},
322 {"loc": [], "msg": "and one more"}
323 ]})),
324 "questions.tone.criteria: field required; state: must not be empty; and one more"
325 );
326 assert_eq!(message(json!({"other": 1})), r#"{"other":1}"#);
327 }
328
329 #[test]
330 fn falls_back_to_the_body() {
331 let error = Error::from_response(502, b"<html>Bad gateway</html>", None, None);
332 assert_eq!(error.message(), "<html>Bad gateway</html>");
333 assert_eq!(error.body(), Some(&json!("<html>Bad gateway</html>")));
334
335 let empty = Error::from_response(503, b"", None, None);
336 assert_eq!(empty.message(), "(no body)");
337 assert_eq!(empty.body(), None);
338
339 let long = "x".repeat(300);
340 let error = Error::from_response(500, long.as_bytes(), None, None);
341 assert_eq!(error.message().chars().count(), MAX_BODY_IN_MESSAGE + 1);
342 assert!(error.message().ends_with('…'));
343 }
344
345 #[test]
346 fn displays_status_and_request_id() {
347 let error = Error::from_response(
348 422,
349 br#"{"detail": "bad"}"#,
350 Some("req_9".into()),
351 Some(Duration::from_millis(5)),
352 );
353 assert_eq!(error.to_string(), "422 bad (request_id=req_9)");
354 assert_eq!(error.kind(), ErrorKind::UnprocessableEntity);
355 assert_eq!(error.retry_after(), Some(Duration::from_millis(5)));
356 assert_eq!(Error::no_api_key().status(), None);
357 }
358}