1use serde::de::DeserializeOwned;
32use std::sync::Arc;
33use warp::http::{Method, StatusCode};
34use warp::{Filter, Rejection};
35
36pub const MAX_JSON_BODY_BYTES: u64 = 1024 * 1024;
38
39#[derive(Debug, Clone, Copy, PartialEq, Eq)]
41pub enum RequestRefused {
42 CrossOrigin,
44 UnsupportedMediaType,
46}
47
48impl warp::reject::Reject for RequestRefused {}
49
50impl RequestRefused {
51 pub fn message(&self) -> &'static str {
53 match self {
54 Self::CrossOrigin => "Cross-origin request refused",
55 Self::UnsupportedMediaType => "Content-Type must be application/json",
56 }
57 }
58
59 pub fn status(&self) -> StatusCode {
61 match self {
62 Self::CrossOrigin => StatusCode::FORBIDDEN,
63 Self::UnsupportedMediaType => StatusCode::UNSUPPORTED_MEDIA_TYPE,
64 }
65 }
66}
67
68pub fn normalize_origin(origin: &str) -> Option<String> {
72 let origin = origin.trim().to_ascii_lowercase();
73 let origin = origin.strip_suffix('/').unwrap_or(&origin);
74 let (scheme, authority) = origin.split_once("://")?;
75 if !matches!(scheme, "http" | "https") || authority.is_empty() {
76 return None;
77 }
78 if authority.contains(['/', '?', '#', '@', ' ']) {
79 return None;
80 }
81 let host = match authority.rsplit_once(':') {
82 Some((host, port)) if !port.ends_with(']') => {
84 if port.is_empty() || !port.bytes().all(|b| b.is_ascii_digit()) {
85 return None;
86 }
87 port.parse::<u16>().ok()?;
88 host
89 }
90 _ => authority,
91 };
92 if host.is_empty() {
93 return None;
94 }
95 Some(format!("{scheme}://{authority}"))
96}
97
98#[derive(Debug, Clone, Default)]
101pub struct AllowedOrigins(Arc<Vec<String>>);
102
103impl AllowedOrigins {
104 pub fn new<I, S>(origins: I) -> crate::Result<Self>
106 where
107 I: IntoIterator<Item = S>,
108 S: AsRef<str>,
109 {
110 let origins = origins
111 .into_iter()
112 .map(|origin| {
113 let origin = origin.as_ref();
114 normalize_origin(origin).ok_or_else(|| {
115 anyhow::anyhow!(
116 "invalid origin {origin:?}: expected scheme://host[:port], for example \
117 https://ops.example.com"
118 )
119 })
120 })
121 .collect::<crate::Result<Vec<_>>>()?;
122 Ok(Self(Arc::new(origins)))
123 }
124
125 pub fn contains(&self, origin: &str) -> bool {
127 normalize_origin(origin).is_some_and(|origin| self.0.contains(&origin))
128 }
129
130 pub fn as_slice(&self) -> &[String] {
132 &self.0
133 }
134}
135
136fn without_default_port(scheme: &str, authority: &str) -> String {
138 let authority = authority.to_ascii_lowercase();
139 let default = match scheme {
140 "https" => ":443",
141 _ => ":80",
142 };
143 authority
144 .strip_suffix(default)
145 .map(str::to_string)
146 .unwrap_or(authority)
147}
148
149pub fn request_allowed(
155 origin: Option<&str>,
156 fetch_site: Option<&str>,
157 host: Option<&str>,
158 allowed: &AllowedOrigins,
159) -> bool {
160 if origin.is_some_and(|origin| allowed.contains(origin)) {
161 return true;
162 }
163 if let Some(site) = fetch_site {
164 return matches!(
165 site.trim().to_ascii_lowercase().as_str(),
166 "same-origin" | "none"
167 );
168 }
169 let Some(origin) = origin else {
170 return true;
171 };
172 let Some(normalized) = normalize_origin(origin) else {
173 return false;
175 };
176 let Some((scheme, authority)) = normalized.split_once("://") else {
177 return false;
178 };
179 host.is_some_and(|host| {
180 without_default_port(scheme, host) == without_default_port(scheme, authority)
181 })
182}
183
184fn origin_headers()
186-> impl Filter<Extract = (Method, Option<String>, Option<String>, Option<String>), Error = Rejection>
187+ Clone {
188 warp::method()
189 .and(warp::header::optional::<String>("origin"))
190 .and(warp::header::optional::<String>("sec-fetch-site"))
191 .and(
192 warp::host::optional().map(|authority: Option<warp::host::Authority>| {
193 authority.map(|authority| authority.as_str().to_string())
194 }),
195 )
196}
197
198fn guard(
199 allowed: AllowedOrigins,
200 writes_only: bool,
201) -> impl Filter<Extract = (), Error = Rejection> + Clone {
202 origin_headers()
203 .and_then(
204 move |method: Method,
205 origin: Option<String>,
206 fetch_site: Option<String>,
207 host: Option<String>| {
208 let allowed = allowed.clone();
209 async move {
210 let safe = matches!(method, Method::GET | Method::HEAD | Method::OPTIONS);
211 if (writes_only && safe)
212 || request_allowed(
213 origin.as_deref(),
214 fetch_site.as_deref(),
215 host.as_deref(),
216 &allowed,
217 )
218 {
219 Ok(())
220 } else {
221 tracing::warn!(
222 origin = origin.as_deref().unwrap_or("-"),
223 sec_fetch_site = fetch_site.as_deref().unwrap_or("-"),
224 %method,
225 "Refused a cross-origin request"
226 );
227 Err(warp::reject::custom(RequestRefused::CrossOrigin))
228 }
229 }
230 },
231 )
232 .untuple_one()
233}
234
235pub fn same_origin_writes(
238 allowed: AllowedOrigins,
239) -> impl Filter<Extract = (), Error = Rejection> + Clone {
240 guard(allowed, true)
241}
242
243pub fn same_origin(
246 allowed: AllowedOrigins,
247) -> impl Filter<Extract = (), Error = Rejection> + Clone {
248 guard(allowed, false)
249}
250
251fn is_json(content_type: &str) -> bool {
253 content_type
254 .split(';')
255 .next()
256 .is_some_and(|mime| mime.trim().eq_ignore_ascii_case("application/json"))
257}
258
259pub fn json_body<T>() -> impl Filter<Extract = (T,), Error = Rejection> + Clone
262where
263 T: DeserializeOwned + Send,
264{
265 warp::header::optional::<String>("content-type")
266 .and_then(|content_type: Option<String>| async move {
267 if content_type.as_deref().is_some_and(is_json) {
268 Ok(())
269 } else {
270 Err(warp::reject::custom(RequestRefused::UnsupportedMediaType))
271 }
272 })
273 .untuple_one()
274 .and(warp::body::content_length_limit(MAX_JSON_BODY_BYTES))
275 .and(warp::body::json())
276}
277
278#[cfg(test)]
279mod tests {
280 use super::*;
281
282 fn allowed(origins: &[&str]) -> AllowedOrigins {
283 AllowedOrigins::new(origins).unwrap()
284 }
285
286 #[test]
287 fn origins_are_normalized_and_validated() {
288 for (input, expected) in [
289 ("http://localhost:8080", Some("http://localhost:8080")),
290 ("HTTPS://Ops.Example.COM/", Some("https://ops.example.com")),
291 ("http://[::1]", Some("http://[::1]")),
292 ("http://[::1]:9000", Some("http://[::1]:9000")),
293 (" https://a.b ", Some("https://a.b")),
294 ("*", None),
295 ("null", None),
296 ("ftp://example.com", None),
297 ("https://", None),
298 ("https://example.com/path", None),
299 ("https://example.com?x", None),
300 ("https://user@example.com", None),
301 ("https://example.com:", None),
302 ("https://example.com:http", None),
303 ("https://example.com:99999", None),
304 ("https://:8080", None),
305 ("example.com", None),
306 ] {
307 assert_eq!(normalize_origin(input).as_deref(), expected, "{input}");
308 }
309 let err = AllowedOrigins::new(["https://ok.example", "*"]).unwrap_err();
310 assert!(err.to_string().contains("invalid origin \"*\""), "{err}");
311 let list = allowed(&["https://Ops.Example.com/"]);
312 assert_eq!(list.as_slice(), ["https://ops.example.com"]);
313 assert!(list.contains("https://ops.example.com"));
314 assert!(!list.contains("null"));
315 assert!(AllowedOrigins::default().as_slice().is_empty());
316 }
317
318 #[test]
319 fn browsers_on_other_origins_are_refused() {
320 let none = AllowedOrigins::default();
321 let host = Some("127.0.0.1:8080");
322 assert!(request_allowed(None, None, host, &none));
324 assert!(request_allowed(None, None, None, &none));
325 assert!(request_allowed(None, Some("same-origin"), host, &none));
327 assert!(request_allowed(None, Some("none"), host, &none));
328 assert!(!request_allowed(None, Some("cross-site"), host, &none));
329 assert!(
330 !request_allowed(
331 Some("http://127.0.0.1:9999"),
332 Some("same-site"),
333 host,
334 &none
335 ),
336 "another port on the same host is another origin"
337 );
338 assert!(
339 request_allowed(
340 Some("https://public.example"),
341 Some("same-origin"),
342 Some("10.0.0.5:8080"),
343 &none
344 ),
345 "behind a proxy the browser's same-origin verdict wins"
346 );
347 assert!(request_allowed(
349 Some("http://127.0.0.1:8080"),
350 None,
351 host,
352 &none
353 ));
354 assert!(request_allowed(
355 Some("http://Dash.Example"),
356 None,
357 Some("dash.example:80"),
358 &none
359 ));
360 assert!(request_allowed(
361 Some("https://dash.example"),
362 None,
363 Some("dash.example:443"),
364 &none
365 ));
366 assert!(!request_allowed(
367 Some("http://evil.example"),
368 None,
369 host,
370 &none
371 ));
372 assert!(!request_allowed(Some("null"), None, host, &none));
373 assert!(!request_allowed(
374 Some("http://127.0.0.1:8080"),
375 None,
376 None,
377 &none
378 ));
379 let ops = allowed(&["https://ops.example"]);
381 assert!(request_allowed(
382 Some("https://ops.example"),
383 Some("cross-site"),
384 host,
385 &ops
386 ));
387 assert!(!request_allowed(
388 Some("https://evil.example"),
389 Some("cross-site"),
390 host,
391 &ops
392 ));
393 }
394
395 #[test]
396 fn content_types() {
397 assert!(is_json("application/json"));
398 assert!(is_json("Application/JSON; charset=utf-8"));
399 assert!(!is_json("text/plain"));
400 assert!(!is_json("application/x-www-form-urlencoded"));
401 assert!(!is_json("multipart/form-data; boundary=x"));
402 assert!(!is_json(""));
403 }
404
405 fn write_route() -> impl Filter<Extract = (String,), Error = Rejection> + Clone {
406 same_origin_writes(AllowedOrigins::default())
407 .and(json_body::<serde_json::Value>())
408 .map(|body: serde_json::Value| body.to_string())
409 }
410
411 async fn status(request: warp::test::RequestBuilder) -> u16 {
412 let route = write_route().recover(crate::auth::handle_auth_rejection);
413 request.reply(&route).await.status().as_u16()
414 }
415
416 fn post(body: &str) -> warp::test::RequestBuilder {
417 warp::test::request()
418 .method("POST")
419 .path("/")
420 .header("host", "127.0.0.1:8080")
421 .header("content-length", body.len().to_string())
422 .body(body)
423 }
424
425 #[tokio::test]
426 async fn json_bodies_need_the_json_content_type_and_a_bounded_length() {
427 let ok = post("{}").header("content-type", "application/json");
428 assert_eq!(status(ok).await, 200);
429 assert_eq!(status(post("{}")).await, 415);
431 let form = post("{}").header("content-type", "text/plain");
432 assert_eq!(status(form).await, 415);
433 let big = format!("\"{}\"", "x".repeat(MAX_JSON_BODY_BYTES as usize));
435 let big = post(&big).header("content-type", "application/json");
436 assert_eq!(status(big).await, 413);
437 let chunked = warp::test::request()
439 .method("POST")
440 .path("/")
441 .header("content-type", "application/json");
442 assert_eq!(status(chunked).await, 411);
443 }
444
445 #[tokio::test]
446 async fn cross_origin_writes_are_refused_and_reads_are_not() {
447 let cross = post("{}")
448 .header("content-type", "application/json")
449 .header("origin", "http://evil.example")
450 .header("sec-fetch-site", "cross-site");
451 assert_eq!(status(cross).await, 403);
452 let same = post("{}")
453 .header("content-type", "application/json")
454 .header("origin", "http://127.0.0.1:8080");
455 assert_eq!(status(same).await, 200);
456
457 let read = same_origin_writes(AllowedOrigins::default()).map(|| "read");
458 let get = warp::test::request()
459 .path("/")
460 .header("origin", "http://evil.example")
461 .reply(&read)
462 .await;
463 assert_eq!(
464 get.status(),
465 200,
466 "reads are left to the same-origin policy"
467 );
468
469 let any = same_origin(AllowedOrigins::default())
470 .map(|| "ws")
471 .recover(crate::auth::handle_auth_rejection);
472 let get = warp::test::request()
473 .path("/")
474 .header("origin", "http://evil.example")
475 .reply(&any)
476 .await;
477 assert_eq!(get.status(), 403, "same_origin checks every method");
478 }
479}