codoseo_web/auth/
origin.rs1use axum::extract::{Request, State};
6use axum::http::{Method, header};
7use axum::middleware::Next;
8use axum::response::{IntoResponse, Response};
9use url::Url;
10
11use crate::error::AppError;
12use crate::routes::{api, mcp};
13use crate::state::AppState;
14
15const EXEMPT_PATHS: &[&str] = &["/billing/webhook", mcp::PATH];
20
21const EXEMPT_PREFIXES: &[&str] = &[api::PREFIX];
24
25fn under(path: &str, prefix: &str) -> bool {
27 path.strip_prefix(prefix)
28 .is_some_and(|rest| rest.starts_with('/'))
29}
30
31pub async fn check_origin(State(state): State<AppState>, req: Request, next: Next) -> Response {
32 let unsafe_method = !matches!(*req.method(), Method::GET | Method::HEAD | Method::OPTIONS);
33 let path = req.uri().path();
34 let exempt = EXEMPT_PATHS.contains(&path) || EXEMPT_PREFIXES.iter().any(|p| under(path, p));
35 if unsafe_method && !exempt && !same_origin(&req, &state.config.origin()) {
36 return AppError::Forbidden(
37 "This request came from another site, so it was blocked.".to_owned(),
38 )
39 .into_response();
40 }
41 next.run(req).await
42}
43
44fn same_origin(req: &Request, expected: &str) -> bool {
45 let headers = req.headers();
46 if let Some(origin) = headers.get(header::ORIGIN).and_then(|v| v.to_str().ok()) {
47 return origin == expected;
48 }
49 headers
50 .get(header::REFERER)
51 .and_then(|v| v.to_str().ok())
52 .and_then(|r| Url::parse(r).ok())
53 .is_some_and(|r| r.origin().ascii_serialization() == expected)
54}
55
56#[cfg(test)]
57mod tests {
58 use super::*;
59
60 #[test]
61 fn a_prefix_exemption_needs_the_slash_and_something_after_it() {
62 assert!(under("/api/v1/sites", api::PREFIX));
63 assert!(!under("/api/v1", api::PREFIX));
64 assert!(!under("/api/v10/sites", api::PREFIX));
65 assert!(!under("/api/v2/sites", api::PREFIX));
66 assert!(EXEMPT_PATHS.contains(&"/mcp"));
68 assert!(!EXEMPT_PATHS.contains(&"/mcp/"));
69 }
70}