use std::net::{IpAddr, Ipv6Addr, SocketAddr};
use std::sync::Arc;
use axum::body::Body;
use axum::extract::{ConnectInfo, Request, State};
use axum::http::{HeaderMap, StatusCode};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use super::admin;
use super::auth::{self, Caller};
use super::respond::{error, Refusal};
use super::AppState;
pub(super) async fn guard(
State(state): State<Arc<AppState>>,
req: Request,
next: Next,
) -> Response {
if let Some(refused) = limit(&state, &req) {
return refused;
}
match authenticate(&state, with_client_ip(&state, req)).await {
Ok(req) => next.run(req).await,
Err(refused) => refused.into_response(),
}
}
pub(super) async fn admin_guard(
State(state): State<Arc<AppState>>,
req: Request,
next: Next,
) -> Response {
if let Some(refused) = limit(&state, &req) {
return refused;
}
let mut req = with_client_ip(&state, req);
let headers = req.headers();
if !headers.contains_key(axum::http::header::AUTHORIZATION)
&& !auth::is_signed(headers)
&& admin::has_session_cookie(headers)
{
return match admin::authenticate_session(&state, &mut req) {
Ok(()) => next.run(req).await,
Err(refused) => refused.into_response(),
};
}
match authenticate(&state, req).await {
Ok(req) => next.run(req).await,
Err(refused) => refused.into_response(),
}
}
fn with_client_ip(state: &AppState, mut req: Request) -> Request {
let ip = client_ip(&req, &state.cfg.trusted_ip_header);
req.extensions_mut().insert(ClientIp(ip));
req
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(super) struct ClientIp(pub(super) String);
pub(super) async fn limited(
State(state): State<Arc<AppState>>,
req: Request,
next: Next,
) -> Response {
unauthenticated(&state, req, next, super::ENROLL_BODY_BYTES).await
}
pub(super) async fn limited_sign_in(
State(state): State<Arc<AppState>>,
req: Request,
next: Next,
) -> Response {
unauthenticated(&state, req, next, super::SIGN_IN_BODY_BYTES).await
}
async fn unauthenticated(state: &AppState, req: Request, next: Next, max_body: usize) -> Response {
if let Some(refused) = limit(state, &req) {
return refused;
}
if declared_length(req.headers()).is_some_and(|n| n > max_body) {
return too_large().into_response();
}
next.run(with_client_ip(state, req)).await
}
fn declared_length(headers: &HeaderMap) -> Option<usize> {
headers
.get(axum::http::header::CONTENT_LENGTH)?
.to_str()
.ok()?
.trim()
.parse()
.ok()
}
pub(super) fn too_large() -> Refusal {
Refusal::new(StatusCode::PAYLOAD_TOO_LARGE, "request body too large")
}
pub(super) async fn admin_only(req: Request, next: Next) -> Response {
match req.extensions().get::<Caller>() {
Some(caller) if caller.is_admin() => next.run(req).await,
_ => error(
StatusCode::FORBIDDEN,
"forbidden: this needs RECALL_TOKEN or a device with the admin scope",
),
}
}
pub(super) async fn not_worker(req: Request, next: Next) -> Response {
match req.extensions().get::<Caller>() {
Some(caller) if caller.is_worker() => error(
StatusCode::FORBIDDEN,
"forbidden: a worker device may only claim jobs and post their results",
),
_ => next.run(req).await,
}
}
pub(super) async fn worker_only(req: Request, next: Next) -> Response {
match req.extensions().get::<Caller>() {
Some(caller) if caller.is_worker() => next.run(req).await,
_ => error(
StatusCode::FORBIDDEN,
"forbidden: this needs a device with the worker scope",
),
}
}
pub(super) fn limit(state: &AppState, req: &Request) -> Option<Response> {
if state
.limiter
.limited(&client_ip(req, &state.cfg.trusted_ip_header))
{
let mut resp = error(
StatusCode::TOO_MANY_REQUESTS,
"rate limit exceeded, try again later",
);
if let Ok(v) = state
.cfg
.rate_limit_window
.as_secs()
.to_string()
.parse::<axum::http::HeaderValue>()
{
resp.headers_mut().insert("retry-after", v);
}
return Some(resp);
}
if let Some(asked) = unsupported_protocol(req.headers()) {
return Some(error(
StatusCode::BAD_REQUEST,
&format!(
"this server speaks Recall protocol {}, and the request asked for {asked}. \
Upgrade whichever side is older; GET {} says what this server supports",
recall_wire::PROTOCOL,
recall_wire::DISCOVERY_PATH
),
));
}
None
}
async fn authenticate(state: &AppState, mut req: Request) -> Result<Request, Refusal> {
if authorized(&state.cfg.token, req.headers()) {
req.extensions_mut().insert(Caller::Operator);
return Ok(req);
}
if !auth::is_signed(req.headers()) {
return Err(Refusal::new(StatusCode::UNAUTHORIZED, "unauthorized"));
}
let (parts, body) = req.into_parts();
let checked = auth::check_headers(state, &parts)?;
if declared_length(&parts.headers).is_some_and(|n| n > super::MAX_BODY_BYTES) {
return Err(too_large());
}
let Ok(bytes) = axum::body::to_bytes(body, super::MAX_BODY_BYTES).await else {
return Err(too_large());
};
let (caller, signed) = auth::finish(state, checked, &bytes)?;
let mut req = Request::from_parts(parts, Body::from(bytes));
req.extensions_mut().insert(caller);
req.extensions_mut().insert(signed);
Ok(req)
}
fn unsupported_protocol(headers: &HeaderMap) -> Option<String> {
let value = headers.get(recall_wire::PROTOCOL_HEADER)?;
let text = value.to_str().unwrap_or("").trim();
match text.parse::<u32>() {
Ok(recall_wire::PROTOCOL) => None,
_ => Some(text.to_string()),
}
}
fn authorized(token: &str, headers: &HeaderMap) -> bool {
let Some(value) = headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "))
else {
return false;
};
!value.is_empty() && constant_time_eq(value.as_bytes(), token.as_bytes())
}
pub(super) fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
let mut diff = 0u8;
for (x, y) in a.iter().zip(b) {
diff |= x ^ y;
}
std::hint::black_box(diff) == 0
}
pub(super) fn client_ip(req: &Request, trusted_header: &str) -> String {
if !trusted_header.is_empty() {
if let Some(ip) = header_str(req.headers(), trusted_header) {
return bucket(ip);
}
}
req.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ConnectInfo(addr)| bucket(&addr.ip().to_string()))
.unwrap_or_else(|| "unknown".to_string())
}
fn bucket(ip: &str) -> String {
match ip.parse::<IpAddr>() {
Ok(IpAddr::V4(v4)) => v4.to_string(),
Ok(IpAddr::V6(v6)) => match v6.to_ipv4_mapped() {
Some(v4) => v4.to_string(),
None => {
let s = v6.segments();
format!("{}/64", Ipv6Addr::new(s[0], s[1], s[2], s[3], 0, 0, 0, 0))
}
},
Err(_) => ip.to_string(),
}
}
fn header_str<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(str::trim)
.filter(|v| !v.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bearer_comparison_rejects_everything_but_the_exact_token() {
let mut h = HeaderMap::new();
assert!(!authorized("secret", &h), "no header");
h.insert("authorization", "secret".parse().unwrap());
assert!(!authorized("secret", &h), "missing Bearer scheme");
h.insert("authorization", "Bearer ".parse().unwrap());
assert!(!authorized("secret", &h), "empty token");
h.insert("authorization", "Bearer secre".parse().unwrap());
assert!(!authorized("secret", &h), "prefix of the token");
h.insert("authorization", "Basic secret".parse().unwrap());
assert!(!authorized("secret", &h), "wrong scheme");
h.insert("authorization", "Bearer secret".parse().unwrap());
assert!(authorized("secret", &h));
}
fn request_with(headers: Vec<(&str, &str)>) -> Request {
let mut req = Request::new(axum::body::Body::empty());
req.extensions_mut()
.insert(ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 1234))));
for (k, v) in headers {
let name = axum::http::HeaderName::from_bytes(k.as_bytes()).unwrap();
req.headers_mut().insert(name, v.parse().unwrap());
}
req
}
#[test]
fn client_ip_reads_the_configured_header_then_the_socket() {
assert_eq!(
client_ip(
&request_with(vec![("cf-connecting-ip", "198.51.100.4")]),
"cf-connecting-ip"
),
"198.51.100.4"
);
assert_eq!(
client_ip(
&request_with(vec![("x-real-ip", "198.51.100.7")]),
"x-real-ip"
),
"198.51.100.7"
);
assert_eq!(
client_ip(&request_with(vec![]), "cf-connecting-ip"),
"127.0.0.1"
);
assert_eq!(
client_ip(
&request_with(vec![("cf-connecting-ip", "198.51.100.4")]),
""
),
"127.0.0.1"
);
}
#[test]
fn an_ipv6_client_is_counted_by_its_64() {
for (ip, want) in [
("2001:db8:1:2::1", "2001:db8:1:2::/64"),
("2001:db8:1:2:ffff:ffff:ffff:ffff", "2001:db8:1:2::/64"),
("2001:DB8:1:2:0:0:0:9", "2001:db8:1:2::/64"),
("2001:db8:1:3::1", "2001:db8:1:3::/64"),
("::ffff:198.51.100.4", "198.51.100.4"),
("198.51.100.4", "198.51.100.4"),
("not an address", "not an address"),
] {
assert_eq!(bucket(ip), want, "{ip}");
}
assert_eq!(
client_ip(
&request_with(vec![("x-real-ip", "2001:db8::abcd")]),
"x-real-ip"
),
"2001:db8::/64"
);
let mut req = request_with(vec![]);
req.extensions_mut().insert(ConnectInfo(SocketAddr::from((
[0x2001, 0xdb8, 0, 7, 1, 2, 3, 4],
1234,
))));
assert_eq!(client_ip(&req, ""), "2001:db8:0:7::/64");
}
#[test]
fn a_header_the_ingress_does_not_set_is_ignored() {
let attacker = request_with(vec![
("cf-connecting-ip", "1.1.1.1"),
("x-forwarded-for", "2.2.2.2"),
("true-client-ip", "3.3.3.3"),
("x-real-ip", "198.51.100.7"),
]);
assert_eq!(
client_ip(&attacker, "x-real-ip"),
"198.51.100.7",
"only the configured header may decide the bucket"
);
assert_eq!(client_ip(&attacker, "cf-connecting-ip"), "1.1.1.1");
}
#[test]
fn forwarded_for_is_no_longer_split_and_trusted() {
let req = request_with(vec![("x-forwarded-for", "203.0.113.9, 10.0.0.1")]);
assert_ne!(
client_ip(&req, "cf-connecting-ip"),
"203.0.113.9",
"x-forwarded-for must not be consulted when it is not the configured header"
);
assert_eq!(client_ip(&req, "cf-connecting-ip"), "127.0.0.1");
}
}