1use axum::{
16 http::{HeaderMap, HeaderName, HeaderValue, StatusCode},
17 response::{IntoResponse, Response},
18 Json,
19};
20use openkind_engine::EngineError;
21use serde_json::json;
22use std::time::Duration;
23
24const RETRY_AFTER_MS: HeaderName = HeaderName::from_static("retry-after-ms");
25const RETRY_AFTER: HeaderName = HeaderName::from_static("retry-after");
26
27#[derive(Debug, thiserror::Error)]
31pub enum ApiError {
32 #[error("invalid request body: {0}")]
34 InvalidBody(String),
35
36 #[error("invalid JSON: {0}")]
38 BadJson(String),
39
40 #[error("payload too large: {0}")]
42 PayloadTooLarge(String),
43
44 #[error("missing or invalid API key")]
46 Unauthorized,
47
48 #[error("rate limited; retry after {retry_after_ms} ms")]
50 RateLimited {
51 retry_after_ms: u64,
53 },
54
55 #[error("server overloaded; retry after {retry_after_ms} ms")]
57 Overloaded {
58 retry_after_ms: u64,
60 },
61
62 #[error("engine error: {0}")]
64 Engine(#[from] EngineError),
65
66 #[error("bad gateway: {0}")]
69 BadGateway(String),
70
71 #[error("internal error: {0}")]
73 Internal(String),
74}
75
76impl ApiError {
77 fn status_and_code(&self) -> (StatusCode, &'static str) {
78 match self {
79 ApiError::InvalidBody(_)
80 | ApiError::Engine(EngineError::Invalid(_))
81 | ApiError::Engine(EngineError::Unsupported { .. }) => {
82 (StatusCode::UNPROCESSABLE_ENTITY, "invalid_body")
83 }
84 ApiError::BadJson(_) => (StatusCode::BAD_REQUEST, "bad_json"),
85 ApiError::PayloadTooLarge(_) => (StatusCode::PAYLOAD_TOO_LARGE, "payload_too_large"),
86 ApiError::Unauthorized => (StatusCode::UNAUTHORIZED, "unauthorized"),
87 ApiError::RateLimited { .. } => (StatusCode::TOO_MANY_REQUESTS, "rate_limited"),
88 ApiError::Overloaded { .. } => (StatusCode::from_u16(529).unwrap(), "overloaded"),
89 ApiError::Engine(EngineError::UnknownModel(_)) => {
90 (StatusCode::NOT_FOUND, "unknown_model")
91 }
92 ApiError::Engine(EngineError::Overloaded { .. }) => {
93 (StatusCode::from_u16(529).unwrap(), "overloaded")
94 }
95 ApiError::Engine(EngineError::DeadlineExceeded { .. }) => {
96 (StatusCode::GATEWAY_TIMEOUT, "deadline_exceeded")
97 }
98 ApiError::Engine(EngineError::Backend { .. }) => {
99 (StatusCode::INTERNAL_SERVER_ERROR, "backend_error")
100 }
101 ApiError::BadGateway(_) => (StatusCode::BAD_GATEWAY, "bad_gateway"),
102 ApiError::Internal(_) => (StatusCode::INTERNAL_SERVER_ERROR, "internal_error"),
103 }
104 }
105
106 fn retry_after_ms(&self) -> Option<u64> {
108 match self {
109 ApiError::RateLimited { retry_after_ms } | ApiError::Overloaded { retry_after_ms } => {
110 Some(*retry_after_ms)
111 }
112 ApiError::Engine(EngineError::Overloaded { retry_after_ms, .. }) => {
113 Some(*retry_after_ms)
114 }
115 _ => None,
116 }
117 }
118}
119
120impl IntoResponse for ApiError {
121 fn into_response(self) -> Response {
122 let (status, code) = self.status_and_code();
123
124 let mut headers = HeaderMap::new();
125 if let Some(ms) = self.retry_after_ms() {
126 headers.insert(RETRY_AFTER_MS, HeaderValue::from(ms));
128 let secs = ms.div_ceil(1000);
129 if let Ok(v) = HeaderValue::from_str(&secs.to_string()) {
130 headers.insert(RETRY_AFTER, v);
131 }
132 }
133 if matches!(self, ApiError::Unauthorized) {
136 if let Ok(v) = HeaderValue::from_str("Bearer") {
137 headers.insert(axum::http::header::WWW_AUTHENTICATE, v);
138 }
139 }
140
141 let body = Json(json!({
142 "error": {
143 "code": code,
144 "message": self.to_string(),
145 }
146 }));
147
148 (status, headers, body).into_response()
149 }
150}
151
152impl From<serde_json::Error> for ApiError {
153 fn from(e: serde_json::Error) -> Self {
154 ApiError::BadJson(e.to_string())
155 }
156}
157
158impl From<Duration> for ApiError {
160 fn from(d: Duration) -> Self {
161 ApiError::RateLimited {
162 retry_after_ms: d.as_millis().min(u64::MAX as u128) as u64,
163 }
164 }
165}
166
167#[cfg(test)]
168mod tests {
169 use super::*;
170 use axum::http::StatusCode;
171 use http_body_util::BodyExt;
172 use openkind_core::ValidationError;
173
174 async fn extract_body_json(resp: Response) -> (StatusCode, HeaderMap, serde_json::Value) {
175 let status = resp.status();
176 let headers = resp.headers().clone();
177 let bytes = resp.into_body().collect().await.unwrap().to_bytes();
178 let val: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
179 (status, headers, val)
180 }
181
182 #[tokio::test]
183 async fn rate_limited_error_into_response() {
184 let err = ApiError::RateLimited {
185 retry_after_ms: 1500,
186 };
187 let (status, headers, body) = extract_body_json(err.into_response()).await;
188 assert_eq!(status, StatusCode::TOO_MANY_REQUESTS);
189 assert_eq!(headers.get("retry-after-ms").unwrap(), "1500");
190 assert_eq!(headers.get("retry-after").unwrap(), "2"); assert_eq!(body["error"]["code"], "rate_limited");
192 }
193
194 #[tokio::test]
195 async fn bad_gateway_error_into_response() {
196 let err = ApiError::BadGateway("upstream unreachable".into());
197 let (status, _headers, body) = extract_body_json(err.into_response()).await;
198 assert_eq!(status, StatusCode::BAD_GATEWAY);
199 assert_eq!(body["error"]["code"], "bad_gateway");
200 assert!(
201 body["error"]["message"]
202 .as_str()
203 .unwrap()
204 .contains("upstream unreachable"),
205 "{body}"
206 );
207 }
208
209 #[tokio::test]
210 async fn overloaded_error_into_response() {
211 let err = ApiError::Overloaded {
212 retry_after_ms: 500,
213 };
214 let (status, headers, body) = extract_body_json(err.into_response()).await;
215 assert_eq!(status, StatusCode::from_u16(529).unwrap());
216 assert_eq!(headers.get("retry-after-ms").unwrap(), "500");
217 assert_eq!(headers.get("retry-after").unwrap(), "1");
218 assert_eq!(body["error"]["code"], "overloaded");
219 }
220
221 #[tokio::test]
222 async fn unauthorized_error_into_response() {
223 let err = ApiError::Unauthorized;
224 let (status, headers, body) = extract_body_json(err.into_response()).await;
225 assert_eq!(status, StatusCode::UNAUTHORIZED);
226 assert_eq!(headers.get("www-authenticate").unwrap(), "Bearer");
227 assert_eq!(body["error"]["code"], "unauthorized");
228 }
229
230 #[tokio::test]
231 async fn internal_error_into_response() {
232 let err = ApiError::Internal("db crashed".into());
233 let (status, _, body) = extract_body_json(err.into_response()).await;
234 assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR);
235 assert_eq!(body["error"]["code"], "internal_error");
236 }
237
238 #[tokio::test]
239 async fn bad_json_and_model_not_found_into_response() {
240 let err_json = ApiError::BadJson("syntax error".into());
241 let (status, _, body) = extract_body_json(err_json.into_response()).await;
242 assert_eq!(status, StatusCode::BAD_REQUEST);
243 assert_eq!(body["error"]["code"], "bad_json");
244
245 let err_model = ApiError::Engine(EngineError::UnknownModel("gpt-5".into()));
246 let (status, _, body) = extract_body_json(err_model.into_response()).await;
247 assert_eq!(status, StatusCode::NOT_FOUND);
248 assert_eq!(body["error"]["code"], "unknown_model");
249 }
250
251 #[tokio::test]
252 async fn validation_and_backend_errors_into_response() {
253 let err_val = ApiError::Engine(EngineError::Invalid(ValidationError::NoQuestions));
254 let (status, _, body) = extract_body_json(err_val.into_response()).await;
255 assert_eq!(status, StatusCode::UNPROCESSABLE_ENTITY);
256 assert_eq!(body["error"]["code"], "invalid_body");
257
258 let err_backend = ApiError::Engine(EngineError::Backend {
259 backend: "mock".into(),
260 message: "simulated failure".into(),
261 });
262 let (status, _, body) = extract_body_json(err_backend.into_response()).await;
263 assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR);
264 assert_eq!(body["error"]["code"], "backend_error");
265 }
266
267 #[tokio::test]
268 async fn unsupported_and_engine_overload_keep_transport_semantics() {
269 let unsupported = ApiError::Engine(EngineError::Unsupported {
270 backend: "native".into(),
271 message: "explicit semantic none required".into(),
272 });
273 let (status, _, body) = extract_body_json(unsupported.into_response()).await;
274 assert_eq!(status, StatusCode::UNPROCESSABLE_ENTITY);
275 assert_eq!(body["error"]["code"], "invalid_body");
276
277 let overloaded = ApiError::Engine(EngineError::Overloaded {
278 backend: "native".into(),
279 retry_after_ms: 750,
280 });
281 let (status, headers, body) = extract_body_json(overloaded.into_response()).await;
282 assert_eq!(status, StatusCode::from_u16(529).unwrap());
283 assert_eq!(headers["retry-after-ms"], "750");
284 assert_eq!(body["error"]["code"], "overloaded");
285
286 let deadline = ApiError::Engine(EngineError::DeadlineExceeded {
287 backend: "native".into(),
288 timeout_ms: 30_000,
289 });
290 let (status, _, body) = extract_body_json(deadline.into_response()).await;
291 assert_eq!(status, StatusCode::GATEWAY_TIMEOUT);
292 assert_eq!(body["error"]["code"], "deadline_exceeded");
293 }
294
295 #[test]
296 fn duration_conversion_to_rate_limited() {
297 let err: ApiError = Duration::from_millis(2500).into();
298 match err {
299 ApiError::RateLimited { retry_after_ms } => assert_eq!(retry_after_ms, 2500),
300 _ => panic!("expected RateLimited"),
301 }
302 }
303
304 #[test]
305 fn duration_conversion_saturates_instead_of_wrapping() {
306 let err: ApiError = Duration::from_secs(u64::MAX).into();
307 assert!(matches!(
308 err,
309 ApiError::RateLimited {
310 retry_after_ms: u64::MAX
311 }
312 ));
313 }
314}