1use std::{
4 future::Future,
5 pin::Pin,
6 task::{Context, Poll},
7};
8
9use axum::{Json, body::Body, body::to_bytes, extract::Request, response::IntoResponse};
10use http::StatusCode;
11use serde::{Deserialize, Serialize};
12use tower::{Layer, Service};
13
14use crate::context::RequestContext;
15
16#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
29#[error("{message}")]
30pub struct HttpError {
31 status: StatusCode,
32 code: &'static str,
33 message: String,
34}
35
36impl HttpError {
37 pub fn new(status: StatusCode, code: &'static str, message: impl Into<String>) -> Self {
43 Self {
44 status,
45 code,
46 message: message.into(),
47 }
48 }
49
50 pub fn bad_request(message: impl Into<String>) -> Self {
52 Self::new(StatusCode::BAD_REQUEST, "bad_request", message)
53 }
54
55 pub fn unauthorized(message: impl Into<String>) -> Self {
57 Self::new(StatusCode::UNAUTHORIZED, "unauthorized", message)
58 }
59
60 pub fn forbidden(message: impl Into<String>) -> Self {
62 Self::new(StatusCode::FORBIDDEN, "forbidden", message)
63 }
64
65 pub fn not_found(message: impl Into<String>) -> Self {
67 Self::new(StatusCode::NOT_FOUND, "not_found", message)
68 }
69
70 pub fn conflict(message: impl Into<String>) -> Self {
72 Self::new(StatusCode::CONFLICT, "conflict", message)
73 }
74
75 pub fn too_many_requests(message: impl Into<String>) -> Self {
77 Self::new(StatusCode::TOO_MANY_REQUESTS, "too_many_requests", message)
78 }
79
80 pub fn unprocessable_entity(message: impl Into<String>) -> Self {
82 Self::new(
83 StatusCode::UNPROCESSABLE_ENTITY,
84 "unprocessable_entity",
85 message,
86 )
87 }
88
89 pub fn internal_server_error() -> Self {
94 Self::new(
95 StatusCode::INTERNAL_SERVER_ERROR,
96 "internal_server_error",
97 "internal server error",
98 )
99 }
100
101 pub fn status(&self) -> StatusCode {
103 self.status
104 }
105
106 pub fn code(&self) -> &'static str {
108 self.code
109 }
110
111 pub fn message(&self) -> &str {
113 &self.message
114 }
115}
116
117impl IntoResponse for HttpError {
118 fn into_response(self) -> axum::response::Response {
119 let status = self.status;
120 let code = self.code;
121 let message = self.message;
122
123 if status.is_server_error() {
124 tracing::error!(
125 http.status = status.as_u16(),
126 error.code = code,
127 error.message = %message,
128 "http error response"
129 );
130 } else {
131 tracing::warn!(
132 http.status = status.as_u16(),
133 error.code = code,
134 error.message = %message,
135 "http error response"
136 );
137 }
138
139 let body = Json(ErrorBody {
140 error: ErrorDetails { code, message },
141 });
142 (status, body).into_response()
143 }
144}
145
146pub async fn not_found_fallback() -> HttpError {
152 HttpError::not_found("route not found")
153}
154
155#[derive(Debug, Serialize)]
156struct ErrorBody {
157 error: ErrorDetails,
158}
159
160#[derive(Debug, Serialize)]
161struct ErrorDetails {
162 code: &'static str,
163 message: String,
164}
165
166#[derive(Clone, Copy, Debug, Default)]
200pub struct ErrorEnvelopeLayer;
201
202impl ErrorEnvelopeLayer {
203 pub fn new() -> Self {
205 Self
206 }
207}
208
209impl<S> Layer<S> for ErrorEnvelopeLayer {
210 type Service = ErrorEnvelopeService<S>;
211
212 fn layer(&self, inner: S) -> Self::Service {
213 ErrorEnvelopeService { inner }
214 }
215}
216
217#[derive(Clone, Debug)]
219pub struct ErrorEnvelopeService<S> {
220 inner: S,
221}
222
223impl<S> Service<Request> for ErrorEnvelopeService<S>
224where
225 S: Service<Request, Response = axum::response::Response> + Send + 'static,
226 S::Future: Send + 'static,
227 S::Error: Send + 'static,
228{
229 type Response = axum::response::Response;
230 type Error = S::Error;
231 type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
232
233 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
234 self.inner.poll_ready(cx)
235 }
236
237 fn call(&mut self, request: Request) -> Self::Future {
238 let path = request.uri().path().to_owned();
239 let request_id = request
242 .extensions()
243 .get::<RequestContext>()
244 .map(|context| context.request_id().to_owned());
245 let future = self.inner.call(request);
246
247 Box::pin(async move {
248 let response = future.await?;
249 if !response.status().is_client_error() && !response.status().is_server_error() {
250 return Ok(response);
251 }
252 Ok(envelope_response(response, request_id, path).await)
253 })
254 }
255}
256
257async fn envelope_response(
258 response: axum::response::Response,
259 request_id: Option<String>,
260 path: String,
261) -> axum::response::Response {
262 let (mut parts, body) = response.into_parts();
263 let status = parts.status;
264 let extracted = read_legacy_error_body(body).await;
265 let mut code = extracted
266 .as_ref()
267 .map(|body| body.error.code.clone())
268 .unwrap_or_else(|| default_code(status).to_owned());
269 let mut message = extracted
270 .as_ref()
271 .map(|body| body.error.message.clone())
272 .unwrap_or_else(|| status.canonical_reason().unwrap_or("error").to_owned());
273 let mut details = extracted
274 .map(|body| {
275 if body.error.details.is_empty() {
276 serde_json::Value::Null
277 } else {
278 serde_json::Value::Object(body.error.details)
279 }
280 })
281 .unwrap_or(serde_json::Value::Null);
282 if status.is_server_error() {
283 tracing::error!(
284 http.status = status.as_u16(),
285 error.code = %code,
286 request.id = request_id.as_deref().unwrap_or(""),
287 http.path = %path,
288 "http error envelope"
289 );
290 message = "internal server error".to_owned();
293 details = serde_json::Value::Null;
294 code = default_code(status).to_owned();
295 }
296
297 let envelope = ProductionErrorBody {
298 error: ProductionErrorDetails {
299 status_code: status.as_u16(),
300 code,
301 message,
302 details,
303 timestamp: timestamp_now(),
304 path,
305 request_id: request_id.unwrap_or_default(),
306 },
307 };
308 let body = serde_json::to_vec(&envelope).expect("error envelope should serialize");
309 parts.headers.insert(
310 http::header::CONTENT_TYPE,
311 http::HeaderValue::from_static("application/json"),
312 );
313 axum::response::Response::from_parts(parts, Body::from(body))
314}
315
316const MAX_ERROR_ENVELOPE_BODY_BYTES: usize = 64 * 1024;
317
318async fn read_legacy_error_body(body: Body) -> Option<LegacyErrorBody> {
319 let bytes = to_bytes(body, MAX_ERROR_ENVELOPE_BODY_BYTES).await.ok()?;
320 serde_json::from_slice::<LegacyErrorBody>(&bytes).ok()
321}
322
323pub(crate) fn timestamp_now() -> String {
325 time::OffsetDateTime::now_utc()
326 .format(&time::format_description::well_known::Rfc3339)
327 .expect("UTC timestamp should format as RFC3339")
328}
329
330fn default_code(status: StatusCode) -> &'static str {
331 match status {
332 StatusCode::BAD_REQUEST => "bad_request",
333 StatusCode::UNAUTHORIZED => "unauthorized",
334 StatusCode::FORBIDDEN => "forbidden",
335 StatusCode::NOT_FOUND => "not_found",
336 StatusCode::CONFLICT => "conflict",
337 StatusCode::UNPROCESSABLE_ENTITY => "unprocessable_entity",
338 StatusCode::TOO_MANY_REQUESTS => "too_many_requests",
339 status if status.is_server_error() => "internal_server_error",
340 _ => "http_error",
341 }
342}
343
344#[derive(Debug, Deserialize)]
345struct LegacyErrorBody {
346 error: LegacyErrorDetails,
347}
348
349#[derive(Debug, Deserialize)]
350struct LegacyErrorDetails {
351 code: String,
352 message: String,
353 #[serde(flatten)]
354 details: serde_json::Map<String, serde_json::Value>,
355}
356
357#[derive(Debug, Serialize)]
358struct ProductionErrorBody {
359 error: ProductionErrorDetails,
360}
361
362#[derive(Debug, Serialize)]
363#[serde(rename_all = "camelCase")]
364struct ProductionErrorDetails {
365 status_code: u16,
366 code: String,
367 message: String,
368 details: serde_json::Value,
369 timestamp: String,
370 path: String,
371 request_id: String,
372}
373
374#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
376#[error("route path `{path}` contains a parameter segment without a name after ':'")]
377pub struct RoutePathError {
378 path: String,
379}
380
381impl RoutePathError {
382 pub fn empty_parameter(path: impl Into<String>) -> Self {
384 Self { path: path.into() }
385 }
386
387 pub fn path(&self) -> &str {
389 &self.path
390 }
391}