use axum::extract::Request;
use axum::http::StatusCode;
use axum::middleware::{from_fn, Next};
use axum::response::IntoResponse;
use axum::Router;
pub const EDGE_SECRET_HEADER: &str = "x-m4a-edge-secret";
fn ct_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
}
pub fn require_edge_secret(router: Router, secret: String) -> Router {
let secret: std::sync::Arc<str> = secret.into();
router.layer(from_fn(move |req: Request, next: Next| {
let secret = std::sync::Arc::clone(&secret);
async move {
let ok = req
.headers()
.get(EDGE_SECRET_HEADER)
.and_then(|v| v.to_str().ok())
.is_some_and(|got| ct_eq(got.as_bytes(), secret.as_bytes()));
if ok {
next.run(req).await
} else {
(StatusCode::UNAUTHORIZED, axum::Json(serde_json::json!({"errcode": "M_UNKNOWN_TOKEN", "error": "edge secret required"}))).into_response()
}
}
}))
}
pub fn require_link_token(router: Router, token: String) -> Router {
let token: std::sync::Arc<str> = token.into();
router.layer(from_fn(move |mut req: Request, next: Next| {
let token = std::sync::Arc::clone(&token);
async move {
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);
req.headers_mut().remove(m4a_seam::LINK_TOKEN_HEADER);
if ok {
next.run(req).await
} else {
(StatusCode::UNAUTHORIZED, axum::Json(serde_json::json!({"errcode": "M4A_LINK_TOKEN_REQUIRED", "error": "link token required"}))).into_response()
}
}
}))
}
#[cfg(test)]
mod tests {
#[test]
fn constant_time_compare() {
assert!(super::ct_eq(b"abc", b"abc"));
assert!(!super::ct_eq(b"abc", b"abd"));
assert!(!super::ct_eq(b"abc", b"ab"));
}
}