net_backend_server/http/
mod.rs1pub mod call;
19pub(crate) mod client_ip;
20pub(crate) mod middleware;
21pub(crate) mod routes;
22
23pub use client_ip::{ClientIp, FORWARDED_FOR_HEADER};
24
25use std::fmt;
26use std::sync::atomic::{AtomicU64, Ordering};
27use std::sync::{Arc, OnceLock};
28
29use axum::extract::rejection::JsonRejection;
30use axum::extract::{FromRequest, FromRequestParts, Request};
31use axum::response::{IntoResponse, Response};
32use http::request::Parts;
33use net_backend_protocol::codes;
34use serde::de::DeserializeOwned;
35use serde::Serialize;
36
37pub use axum::extract::DefaultBodyLimit;
38
39use crate::error::AppError;
40use crate::state::AppState;
41
42pub const REQUEST_ID_HEADER: &str = "x-request-id";
44
45pub fn body_limit(bytes: usize) -> DefaultBodyLimit {
49 DefaultBodyLimit::max(bytes)
50}
51
52#[derive(Clone, PartialEq, Eq, Hash)]
55pub struct RequestId(Arc<str>);
56
57impl RequestId {
58 pub fn generate() -> Self {
60 static COUNTER: AtomicU64 = AtomicU64::new(1);
61 static PREFIX: OnceLock<String> = OnceLock::new();
62 let prefix = PREFIX.get_or_init(|| format!("{:x}", net_backend_protocol::UnixMillis::now().get().max(0)));
63 let n = COUNTER.fetch_add(1, Ordering::Relaxed);
64 Self(Arc::from(format!("{prefix}-{n:x}")))
65 }
66
67 pub fn from_client(value: &str) -> Option<Self> {
69 let ok = (1..=64).contains(&value.len()) && value.bytes().all(|b| b.is_ascii_alphanumeric() || matches!(b, b'.' | b'_' | b'-'));
70 ok.then(|| Self(Arc::from(value)))
71 }
72
73 pub fn as_str(&self) -> &str {
75 &self.0
76 }
77}
78
79impl fmt::Debug for RequestId {
80 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
81 write!(f, "RequestId({})", self.0)
82 }
83}
84
85impl fmt::Display for RequestId {
86 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
87 f.write_str(&self.0)
88 }
89}
90
91impl<S: Send + Sync> FromRequestParts<S> for RequestId {
92 type Rejection = std::convert::Infallible;
93
94 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
95 Ok(parts.extensions.get::<RequestId>().cloned().unwrap_or_else(RequestId::generate))
96 }
97}
98
99#[derive(Clone, Copy, Debug, Default)]
104pub struct ApiJson<T>(pub T);
105
106impl<T, S> FromRequest<S> for ApiJson<T>
107where
108 T: DeserializeOwned,
109 S: Send + Sync,
110{
111 type Rejection = AppError;
112
113 async fn from_request(req: Request, state: &S) -> Result<Self, Self::Rejection> {
114 match axum::Json::<T>::from_request(req, state).await {
115 Ok(axum::Json(value)) => Ok(ApiJson(value)),
116 Err(rejection) => Err(json_rejection(rejection)),
117 }
118 }
119}
120
121fn json_rejection(rejection: JsonRejection) -> AppError {
122 let status = rejection.status();
123 match status.as_u16() {
124 413 => AppError::payload_too_large("the request body is too large"),
125 415 => AppError::with_status(status, "unsupported_media_type", "the request body must be JSON (content-type: application/json)"),
126 _ => AppError::new(codes::BAD_REQUEST, rejection.body_text()),
128 }
129}
130
131impl<T: Serialize> IntoResponse for ApiJson<T> {
132 fn into_response(self) -> Response {
133 axum::Json(self.0).into_response()
134 }
135}
136
137#[derive(Debug)]
140pub struct Ext<T>(pub Arc<T>);
141
142impl<T> std::ops::Deref for Ext<T> {
143 type Target = T;
144 fn deref(&self) -> &T {
145 &self.0
146 }
147}
148
149impl<T: Send + Sync + 'static> FromRequestParts<AppState> for Ext<T> {
150 type Rejection = AppError;
151
152 async fn from_request_parts(_parts: &mut Parts, state: &AppState) -> Result<Self, Self::Rejection> {
153 state
154 .get::<T>()
155 .map(Ext)
156 .ok_or_else(|| AppError::internal(std::io::Error::other(format!("no state of type {} was registered", std::any::type_name::<T>()))))
157 }
158}
159
160#[cfg(test)]
161mod tests {
162 use super::*;
163
164 #[test]
165 fn request_ids() {
166 let a = RequestId::generate();
167 let b = RequestId::generate();
168 assert_ne!(a, b);
169 assert!(RequestId::from_client(a.as_str()).is_some());
170 assert!(RequestId::from_client("abc-DEF_1.2").is_some());
171 for bad in ["", "a b", "x\ny", "<script>", &"x".repeat(65)] {
172 assert!(RequestId::from_client(bad).is_none(), "{bad}");
173 }
174 }
175}