use std::collections::HashMap;
use std::sync::{Mutex, PoisonError};
use std::time::Duration;
use axum::http::request::Parts;
use axum::http::{HeaderMap, StatusCode};
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine;
use recall_wire::devices::{SCOPE_ADMIN, SCOPE_SYNC, SCOPE_WORKER};
use recall_wire::signature::{
self, Received, SignatureInput, Target, LABEL, MAX_AHEAD_SECONDS, SIGNATURE_HEADER,
SIGNATURE_INPUT_HEADER, WINDOW_SECONDS,
};
use recall_wire::Device;
use super::respond::Refusal;
use super::AppState;
use crate::{now, parse_timestamp};
pub(super) const WINDOW: u64 = WINDOW_SECONDS;
pub(super) const NONCE_LIFETIME: u64 = WINDOW + MAX_AHEAD_SECONDS;
#[derive(Debug, Clone, PartialEq, Eq)]
pub(super) enum Caller {
Operator,
Device {
id: String,
name: String,
scope: String,
ephemeral: bool,
agent: String,
},
Owner {
credential_id: String,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(super) struct SignedRequestInfo {
pub(super) body_sha256: String,
pub(super) signature_base: String,
pub(super) signature: String,
}
impl Caller {
pub(super) fn is_admin(&self) -> bool {
match self {
Caller::Operator | Caller::Owner { .. } => true,
Caller::Device { scope, .. } => scope == SCOPE_ADMIN,
}
}
pub(super) fn is_worker(&self) -> bool {
matches!(self, Caller::Device { scope, .. } if scope == SCOPE_WORKER)
}
}
pub(super) fn is_signed(headers: &HeaderMap) -> bool {
headers.contains_key(SIGNATURE_INPUT_HEADER) || headers.contains_key(SIGNATURE_HEADER)
}
const REPLAY_CAPACITY: usize = 16_384;
pub(super) const MAX_NONCES_PER_DEVICE: usize = REPLAY_CAPACITY / 64;
const LAST_SEEN_EVERY: Duration = Duration::from_secs(60);
pub(super) struct Checked {
device: Device,
digest: String,
nonce: String,
created: i64,
signature_base: String,
signature: Vec<u8>,
}
fn known_scope(scope: &str) -> bool {
scope == SCOPE_SYNC || scope == SCOPE_ADMIN || scope == SCOPE_WORKER
}
fn rejected(why: &dyn std::fmt::Display) -> Refusal {
Refusal::new(StatusCode::UNAUTHORIZED, format!("unauthorized: {why}"))
}
pub(super) fn check_headers(state: &AppState, parts: &Parts) -> Result<Checked, Refusal> {
let field = |name: &str| joined(&parts.headers, name);
let (Some(input_field), Some(signature_field)) =
(field(SIGNATURE_INPUT_HEADER), field(SIGNATURE_HEADER))
else {
return Err(rejected(
&"a signed request needs both signature-input and signature",
));
};
let input = SignatureInput::parse(&input_field, LABEL).map_err(|e| rejected(&e))?;
let sig = signature::parse_signature(&signature_field, LABEL).map_err(|e| rejected(&e))?;
let Some(keyid) = input.keyid() else {
return Err(rejected(&signature::SignatureError::MissingParameter(
"keyid",
)));
};
let device = match state.store.device(keyid) {
Ok(Some(device)) => device,
Ok(None) => return Err(rejected(&"unknown device")),
Err(e) => return Err(Refusal::internal(e)),
};
if device.revoked_at.is_some() {
return Err(rejected(&"this device has been revoked"));
}
if !known_scope(&device.scope) {
return Err(Refusal::new(
StatusCode::FORBIDDEN,
format!(
"forbidden: this server does not know the scope {:?}; revoke the device \
or run the server version that enrolled it",
device.scope
),
));
}
let key = signature::parse_public_key(&device.public_key).map_err(|e| {
Refusal::internal(anyhow::anyhow!("device {} has a bad key: {e}", device.id))
})?;
let host = parts
.headers
.get(axum::http::header::HOST)
.and_then(|v| v.to_str().ok())
.or_else(|| parts.uri.authority().map(|a| a.as_str()))
.unwrap_or("");
let authority = signature::normalize_authority(host);
let received = Received {
input: &input,
signature: &sig,
target: Target {
method: parts.method.as_str(),
authority: &authority,
path: parts.uri.path(),
query: parts.uri.query(),
},
field: &field,
};
let now = state.now();
let signature_base = signature::verify_headers_and_base(&received, &key, now, WINDOW)
.map_err(|e| rejected(&e))?;
let (nonce, created) = (input.nonce().unwrap_or(""), input.created().unwrap_or(0));
let settled = state.started().saturating_add(MAX_AHEAD_SECONDS as i64);
if created <= settled {
return Err(rejected(
&"signature created before this server started, or too soon after; \
sign the request again in a few seconds",
));
}
if state.replay.seen(keyid, nonce, now) {
return Err(rejected(&"this request was already received once"));
}
Ok(Checked {
digest: field(signature::CONTENT_DIGEST_HEADER).unwrap_or_default(),
nonce: nonce.to_string(),
created,
signature_base,
signature: sig,
device,
})
}
pub(super) fn finish(
state: &AppState,
checked: Checked,
body: &[u8],
) -> Result<(Caller, SignedRequestInfo), Refusal> {
let Checked {
device,
digest,
nonce,
created,
signature_base,
signature,
} = checked;
signature::check_content_digest(&digest, body).map_err(|e| rejected(&e))?;
match state
.replay
.first_use(&device.id, &nonce, created, &|| state.now())
{
Recorded::Fresh => {}
Recorded::Replayed => return Err(rejected(&"this request was already received once")),
Recorded::Stale => {
return Err(rejected(
&"the signature's window closed while the request was read; sign it again",
))
}
Recorded::DeviceFull => {
return Err(Refusal::new(
StatusCode::TOO_MANY_REQUESTS,
"too many signed requests from this device, try again later",
))
}
Recorded::Full => {
return Err(Refusal::new(
StatusCode::SERVICE_UNAVAILABLE,
"too many signed requests at once, try again later",
))
}
}
let stale = device
.last_seen
.as_deref()
.and_then(parse_timestamp)
.is_none_or(|seen| time::OffsetDateTime::now_utc() - seen >= LAST_SEEN_EVERY);
if stale {
if let Err(e) = state.store.touch_device(&device.id, &now()) {
eprintln!("recording last_seen for {}: {e:#}", device.id);
}
}
let info = SignedRequestInfo {
body_sha256: signature::content_digest_base64(&digest).unwrap_or_default(),
signature_base,
signature: BASE64_STANDARD.encode(&signature),
};
Ok((
Caller::Device {
id: device.id,
name: device.name,
scope: device.scope,
ephemeral: device.ephemeral,
agent: device.agent,
},
info,
))
}
fn joined(headers: &HeaderMap, name: &str) -> Option<String> {
let mut out: Option<String> = None;
for value in headers.get_all(name) {
let value = value.to_str().ok()?.trim();
match &mut out {
Some(s) => {
s.push_str(", ");
s.push_str(value);
}
None => out = Some(value.to_string()),
}
}
out
}
pub(super) struct ReplayCache {
window: i64,
capacity: usize,
per_device: usize,
state: Mutex<ReplayState>,
}
struct ReplayState {
seen: HashMap<(String, String), i64>,
per_device: HashMap<String, usize>,
last_sweep: i64,
#[cfg(test)]
sweeps: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum Recorded {
Fresh,
Replayed,
Stale,
DeviceFull,
Full,
}
impl ReplayCache {
pub(super) fn new(window: u64, per_device: usize) -> Self {
Self::with_capacity(
window,
REPLAY_CAPACITY,
per_device.min(MAX_NONCES_PER_DEVICE),
)
}
fn with_capacity(window: u64, capacity: usize, per_device: usize) -> Self {
Self {
window: window as i64,
capacity,
per_device: per_device.max(1),
state: Mutex::new(ReplayState {
seen: HashMap::new(),
per_device: HashMap::new(),
last_sweep: 0,
#[cfg(test)]
sweeps: 0,
}),
}
}
fn lock(&self) -> std::sync::MutexGuard<'_, ReplayState> {
self.state.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(super) fn seen(&self, keyid: &str, nonce: &str, now: i64) -> bool {
self.lock()
.seen
.get(&(keyid.to_string(), nonce.to_string()))
.is_some_and(|until| *until >= now)
}
pub(super) fn first_use(
&self,
keyid: &str,
nonce: &str,
created: i64,
clock: &dyn Fn() -> i64,
) -> Recorded {
let until = created.saturating_add(self.window);
let mut state = self.lock();
let now = clock().max(state.last_sweep);
let device_count = state.per_device.get(keyid).copied().unwrap_or(0);
let crowded = state.seen.len() >= self.capacity || device_count >= self.per_device;
if now - state.last_sweep >= self.window || (crowded && now > state.last_sweep) {
sweep(&mut state, now);
}
if until < now {
return Recorded::Stale;
}
let key = (keyid.to_string(), nonce.to_string());
if state.seen.get(&key).is_some_and(|u| *u >= now) {
return Recorded::Replayed;
}
if state.per_device.get(keyid).copied().unwrap_or(0) >= self.per_device {
return Recorded::DeviceFull;
}
if state.seen.len() >= self.capacity {
return Recorded::Full;
}
state.seen.insert(key, until);
*state.per_device.entry(keyid.to_string()).or_insert(0) += 1;
Recorded::Fresh
}
}
fn sweep(state: &mut ReplayState, now: i64) {
#[cfg(test)]
{
state.sweeps += 1;
}
state.last_sweep = now;
state.seen.retain(|_, until| *until >= now);
state.per_device.clear();
for (keyid, _) in state.seen.keys() {
*state.per_device.entry(keyid.clone()).or_insert(0) += 1;
}
}
#[cfg(test)]
mod tests {
use super::*;
fn at(t: i64) -> impl Fn() -> i64 {
move || t
}
#[test]
fn a_nonce_is_accepted_once_inside_its_window() {
let cache = ReplayCache::new(60, 100);
assert_eq!(
cache.first_use("dev_a", "n1", 1000, &at(1000)),
Recorded::Fresh
);
assert!(cache.seen("dev_a", "n1", 1030));
assert_eq!(
cache.first_use("dev_a", "n1", 1000, &at(1030)),
Recorded::Replayed
);
assert_eq!(
cache.first_use("dev_b", "n1", 1000, &at(1030)),
Recorded::Fresh
);
assert_eq!(
cache.first_use("dev_a", "n2", 1000, &at(1030)),
Recorded::Fresh
);
}
#[test]
fn nonces_are_forgotten_once_their_window_has_passed() {
let cache = ReplayCache::with_capacity(60, 2, 100);
assert_eq!(cache.first_use("d", "a", 1000, &at(1000)), Recorded::Fresh);
assert_eq!(cache.first_use("e", "b", 1000, &at(1000)), Recorded::Fresh);
assert_eq!(
cache.first_use("f", "c", 1000, &at(1000)),
Recorded::Full,
"full of live nonces: refuse, never forget one"
);
assert_eq!(cache.first_use("f", "c", 1100, &at(1100)), Recorded::Fresh);
assert_eq!(cache.lock().seen.len(), 1, "expired nonces were swept");
}
#[test]
fn one_device_cannot_fill_the_cache_for_the_others() {
let cache = ReplayCache::with_capacity(60, 100, 3);
for n in ["a", "b", "c"] {
assert_eq!(
cache.first_use("greedy", n, 1000, &at(1000)),
Recorded::Fresh
);
}
assert_eq!(
cache.first_use("greedy", "d", 1000, &at(1000)),
Recorded::DeviceFull
);
assert_eq!(
cache.first_use("other", "a", 1000, &at(1000)),
Recorded::Fresh
);
assert_eq!(
cache.first_use("greedy", "d", 1100, &at(1100)),
Recorded::Fresh
);
}
#[test]
fn a_devices_share_is_capped_however_high_the_rate_limit() {
assert_eq!(ReplayCache::new(60, usize::MAX).per_device, 256);
assert_eq!(ReplayCache::new(60, 120).per_device, 120);
let cache = ReplayCache::new(60, usize::MAX);
for n in 0..256 {
assert_eq!(
cache.first_use("greedy", &n.to_string(), 1000, &at(1000)),
Recorded::Fresh
);
}
assert_eq!(
cache.first_use("greedy", "one more", 1000, &at(1000)),
Recorded::DeviceFull
);
}
#[test]
fn a_nonce_swept_away_is_not_recorded_afresh() {
let cache = ReplayCache::new(60, 100);
assert_eq!(cache.first_use("d", "n", 40, &at(100)), Recorded::Fresh);
assert_eq!(cache.first_use("e", "m", 101, &at(161)), Recorded::Fresh);
assert!(!cache.seen("d", "n", 101));
assert_eq!(cache.first_use("d", "n", 40, &at(101)), Recorded::Stale);
}
#[test]
fn a_clock_stepped_back_does_not_bring_a_swept_nonce_back() {
let cache = ReplayCache::new(60, 100);
assert_eq!(cache.first_use("d", "n", 40, &at(100)), Recorded::Fresh);
assert_eq!(cache.first_use("e", "m", 101, &at(161)), Recorded::Fresh);
assert_eq!(cache.first_use("d", "n", 40, &at(90)), Recorded::Stale);
}
#[test]
fn a_crowded_cache_is_swept_at_most_once_a_second() {
let cache = ReplayCache::new(60, 2);
for n in ["a", "b"] {
assert_eq!(cache.first_use("d", n, 1000, &at(1000)), Recorded::Fresh);
}
let swept = cache.lock().sweeps;
for n in 0..100 {
assert_eq!(
cache.first_use("d", &n.to_string(), 1000, &at(1000)),
Recorded::DeviceFull
);
}
assert_eq!(cache.lock().sweeps, swept, "no sweep in the same second");
assert_eq!(
cache.first_use("d", "c", 1001, &at(1001)),
Recorded::DeviceFull
);
assert_eq!(cache.lock().sweeps, swept + 1, "one in the next");
}
#[test]
fn repeated_headers_are_joined_as_rfc9110_combines_them() {
let mut h = HeaderMap::new();
h.append("content-digest", "sha-256=:a:".parse().unwrap());
h.append("content-digest", " sha-512=:b: ".parse().unwrap());
assert_eq!(
joined(&h, "content-digest").as_deref(),
Some("sha-256=:a:, sha-512=:b:")
);
assert_eq!(joined(&h, "signature"), None);
}
#[test]
fn only_the_operator_and_admin_devices_are_admins() {
let device = |scope: &str| Caller::Device {
id: "dev_a".into(),
name: "laptop".into(),
scope: scope.into(),
ephemeral: false,
agent: String::new(),
};
assert!(Caller::Operator.is_admin());
assert!(Caller::Owner {
credential_id: "c".into()
}
.is_admin());
assert!(device("admin").is_admin());
assert!(!device("sync").is_admin());
assert!(!device("worker").is_admin());
}
#[test]
fn only_a_worker_device_is_a_worker() {
let device = |scope: &str| Caller::Device {
id: "dev_a".into(),
name: "worker".into(),
scope: scope.into(),
ephemeral: false,
agent: String::new(),
};
assert!(device("worker").is_worker());
assert!(!device("admin").is_worker());
assert!(!device("sync").is_worker());
assert!(!Caller::Operator.is_worker());
}
#[test]
fn a_scope_this_server_does_not_know_is_not_read_as_sync() {
assert!(known_scope("sync"));
assert!(known_scope("admin"));
assert!(known_scope("worker"));
for later in ["evaluator", "Worker", "Sync", "", "sync "] {
assert!(!known_scope(later), "{later:?}");
}
}
}