Skip to main content

mail4agent_server/http/
edge_auth.rs

1//! Core-side gate for the edge/core split. When the server runs with
2//! `--role core`, every request must carry the shared secret the edge adds
3//! (`X-M4A-Edge-Secret`). The tunnel and a firewall allowlist are the first
4//! line; this header is the second, so a stray process on the tunnel network
5//! cannot talk to the core. The secret comes from the environment and is never
6//! logged.
7
8use axum::extract::Request;
9use axum::http::StatusCode;
10use axum::middleware::{from_fn, Next};
11use axum::response::IntoResponse;
12use axum::Router;
13
14/// Header name carrying the shared secret edge -> core.
15pub const EDGE_SECRET_HEADER: &str = "x-m4a-edge-secret";
16
17fn ct_eq(a: &[u8], b: &[u8]) -> bool {
18    if a.len() != b.len() {
19        return false;
20    }
21    a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
22}
23
24/// Wraps `router` so requests without the right secret get a Matrix-shaped 401.
25pub fn require_edge_secret(router: Router, secret: String) -> Router {
26    let secret: std::sync::Arc<str> = secret.into();
27    router.layer(from_fn(move |req: Request, next: Next| {
28        let secret = std::sync::Arc::clone(&secret);
29        async move {
30            let ok = req
31                .headers()
32                .get(EDGE_SECRET_HEADER)
33                .and_then(|v| v.to_str().ok())
34                .is_some_and(|got| ct_eq(got.as_bytes(), secret.as_bytes()));
35            if ok {
36                next.run(req).await
37            } else {
38                (StatusCode::UNAUTHORIZED, axum::Json(serde_json::json!({"errcode": "M_UNKNOWN_TOKEN", "error": "edge secret required"}))).into_response()
39            }
40        }
41    }))
42}
43
44/// Barrier token gate for a standalone core reached directly by a product server
45/// (`M4A_LINK_TOKEN`): every request outside the public protocol surfaces must carry
46/// `x-m4a-link-token`. The header is removed before routing.
47pub fn require_link_token(router: Router, token: String) -> Router {
48    let token: std::sync::Arc<str> = token.into();
49    router.layer(from_fn(move |mut req: Request, next: Next| {
50        let token = std::sync::Arc::clone(&token);
51        async move {
52            let ok = m4a_seam::link_path_is_open(req.uri().path()) || m4a_seam::link_token_ok(req.headers().get(m4a_seam::LINK_TOKEN_HEADER).and_then(|v| v.to_str().ok()), &token);
53            req.headers_mut().remove(m4a_seam::LINK_TOKEN_HEADER);
54            if ok {
55                next.run(req).await
56            } else {
57                (StatusCode::UNAUTHORIZED, axum::Json(serde_json::json!({"errcode": "M4A_LINK_TOKEN_REQUIRED", "error": "link token required"}))).into_response()
58            }
59        }
60    }))
61}
62
63#[cfg(test)]
64mod tests {
65    #[test]
66    fn constant_time_compare() {
67        assert!(super::ct_eq(b"abc", b"abc"));
68        assert!(!super::ct_eq(b"abc", b"abd"));
69        assert!(!super::ct_eq(b"abc", b"ab"));
70    }
71}