1use axum::{
18 Router,
19 body::Body,
20 extract::Request,
21 http::{HeaderMap, HeaderName, HeaderValue, Method, StatusCode, header},
22 middleware::{self, Next},
23 response::{IntoResponse, Response},
24};
25use tower_http::cors::{AllowOrigin, CorsLayer};
26
27use crate::session::{Keys, Session};
28
29pub const ALLOWED_ORIGINS: &str = "ALLOWED_ORIGINS";
55
56pub const ALLOWED_HOSTS: &str = "ALLOWED_HOSTS";
82
83pub(crate) struct Config {
85 pub keys: Result<Keys, String>,
86 pub allowed_origins: Vec<HeaderValue>,
87 pub allowed_hosts: Vec<String>,
88}
89
90pub(crate) fn parse_origins(value: Option<String>) -> Vec<HeaderValue> {
93 value
94 .unwrap_or_default()
95 .split(',')
96 .map(|origin| origin.trim().trim_end_matches('/'))
97 .filter(|origin| !origin.is_empty())
98 .filter_map(|origin| HeaderValue::from_str(origin).ok())
99 .collect()
100}
101
102pub(crate) fn parse_hosts(value: Option<String>) -> Vec<String> {
104 let value = value.unwrap_or_default();
105 value.split(',').map(|host| host.trim().to_ascii_lowercase()).filter(|host| !host.is_empty()).collect()
106}
107
108pub(crate) fn host_allowed(host: Option<&str>, allowed: &[String]) -> bool {
110 if allowed.is_empty() {
111 return true;
112 }
113 let Some(host) = host else { return false };
114 let host = host.to_ascii_lowercase();
115 let name = match host.rsplit_once(':') {
117 Some((name, port)) if !name.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => name,
118 _ => host.as_str(),
119 };
120 if matches!(name, "localhost" | "127.0.0.1" | "[::1]") {
121 return true;
122 }
123 allowed.iter().any(|entry| match entry.strip_prefix('.') {
124 Some(domain) => name == domain || name.strip_suffix(domain).is_some_and(|sub| sub.ends_with('.')),
125 None => name == entry,
126 })
127}
128
129pub(crate) fn wrap(router: Router, config: Config) -> Router {
131 let Config { keys, allowed_origins, allowed_hosts } = config;
132 let trusted = allowed_origins.clone();
133 let mut router = router
134 .layer(middleware::from_fn(move |req: Request, next: Next| {
135 let keys = keys.clone();
136 async move { session(keys, req, next).await }
137 }))
138 .layer(middleware::from_fn(move |req: Request, next: Next| {
139 let refused = cross_origin(req.method(), req.headers(), &trusted);
140 async move {
141 match refused {
142 Some(reason) => (StatusCode::FORBIDDEN, reason).into_response(),
143 None => next.run(req).await,
144 }
145 }
146 }));
147 if !allowed_origins.is_empty() {
148 router = router.layer(
149 CorsLayer::new()
150 .allow_origin(AllowOrigin::list(allowed_origins))
151 .allow_methods([Method::GET, Method::POST, Method::PUT, Method::PATCH, Method::DELETE])
152 .allow_headers([header::CONTENT_TYPE, header::AUTHORIZATION, header::ACCEPT])
153 .allow_credentials(true),
154 );
155 }
156 if !allowed_hosts.is_empty() {
157 router = router.layer(middleware::from_fn(move |req: Request, next: Next| {
158 let header = req.headers().get(header::HOST).and_then(|host| host.to_str().ok());
159 let allowed = host_allowed(req.uri().host().or(header), &allowed_hosts);
160 async move {
161 if allowed {
162 next.run(req).await
163 } else {
164 (StatusCode::FORBIDDEN, "Forbidden: blocked host. Add it to ALLOWED_HOSTS to allow it.")
165 .into_response()
166 }
167 }
168 }));
169 }
170 router.layer(middleware::from_fn(security_headers))
171}
172
173async fn session(keys: Result<Keys, String>, mut req: Request, next: Next) -> Response {
174 let secure = req.uri().scheme_str() == Some("https");
175 let cookies = crate::cookies::Cookies::from_headers(req.headers(), keys.clone(), secure);
176 let session = Session::from_headers(req.headers(), keys, secure);
177 req.extensions_mut().insert(session.clone());
178 req.extensions_mut().insert(cookies.clone());
179 let mut response = next.run(req).await;
180 for cookie in cookies.set_cookies() {
181 response.headers_mut().append(header::SET_COOKIE, cookie);
182 }
183 match session.set_cookie() {
184 Ok(Some(cookie)) => {
185 response.headers_mut().append(header::SET_COOKIE, cookie);
186 response
187 }
188 Ok(None) => response,
189 #[cfg(feature = "html")]
190 Err(err) => err.into_response(),
191 #[cfg(not(feature = "html"))]
192 Err(err) => crate::ApiError::from(err).into_response(),
193 }
194}
195
196pub(crate) fn is_websocket_upgrade(headers: &HeaderMap) -> bool {
198 headers.get(header::UPGRADE).is_some_and(|value| value.as_bytes().eq_ignore_ascii_case(b"websocket"))
199}
200
201pub(crate) fn cross_origin(method: &Method, headers: &HeaderMap, trusted: &[HeaderValue]) -> Option<&'static str> {
205 if matches!(*method, Method::GET | Method::HEAD | Method::OPTIONS) && !is_websocket_upgrade(headers) {
206 return None;
207 }
208 let origin = headers.get(header::ORIGIN);
209 if origin.is_some_and(|origin| trusted.contains(origin)) {
210 return None;
211 }
212 match headers.get("sec-fetch-site").map(HeaderValue::as_bytes) {
213 Some(b"same-origin" | b"none") => return None,
214 Some(_) => return Some("Forbidden: cross-site request. Add the origin to ALLOWED_ORIGINS to allow it."),
215 None => {}
216 }
217 let (Some(origin), Some(host)) = (origin, headers.get(header::HOST)) else {
219 return None;
220 };
221 let origin_host = origin.to_str().ok().and_then(|origin| origin.split_once("://")).map(|(_, host)| host);
222 if origin_host.is_some_and(|origin_host| origin_host.as_bytes() == host.as_bytes()) {
223 None
224 } else {
225 Some("Forbidden: cross-origin request. Add the origin to ALLOWED_ORIGINS to allow it.")
226 }
227}
228
229async fn security_headers(req: Request<Body>, next: Next) -> Response {
232 let https = req.uri().scheme_str() == Some("https");
233 let mut response = next.run(req).await;
234 let headers = response.headers_mut();
235 let defaults: [(HeaderName, &str); 5] = [
236 (header::X_CONTENT_TYPE_OPTIONS, "nosniff"),
237 (header::X_FRAME_OPTIONS, "SAMEORIGIN"),
238 (header::REFERRER_POLICY, "strict-origin-when-cross-origin"),
239 (header::X_XSS_PROTECTION, "0"),
240 (HeaderName::from_static("x-permitted-cross-domain-policies"), "none"),
241 ];
242 for (name, value) in defaults {
243 headers.entry(name).or_insert(HeaderValue::from_static(value));
244 }
245 if https {
246 headers.entry(header::STRICT_TRANSPORT_SECURITY).or_insert(HeaderValue::from_static("max-age=63072000"));
247 }
248 response
249}
250
251#[cfg(test)]
252#[path = "../tests/protect.rs"]
253mod tests;