use std::{fmt, num::NonZeroU64, time::Duration};
use sha2::{Digest, Sha256};
use thiserror::Error;
use crate::{PolicyId, ScopeId};
const FINGERPRINT_DOMAIN: &[u8] = b"runlimit/fixed-window-policy/v1\0";
const MAX_EXACT_DOUBLE_INTEGER: u64 = 1_u64 << f64::MANTISSA_DIGITS;
pub const MAX_LIMIT: u64 = i64::MAX as u64;
pub const MAX_WINDOW_MILLIS: u64 = MAX_EXACT_DOUBLE_INTEGER / 1_000;
pub const MAX_WINDOW: Duration = Duration::from_millis(MAX_WINDOW_MILLIS);
#[derive(Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct PolicyFingerprint([u8; 32]);
impl PolicyFingerprint {
pub const fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
pub const fn into_bytes(self) -> [u8; 32] {
self.0
}
}
impl fmt::Debug for PolicyFingerprint {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("PolicyFingerprint(")?;
write_hex(formatter, &self.0)?;
formatter.write_str(")")
}
}
impl fmt::Display for PolicyFingerprint {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
write_hex(formatter, &self.0)
}
}
fn write_hex(formatter: &mut fmt::Formatter<'_>, bytes: &[u8]) -> fmt::Result {
for byte in bytes {
write!(formatter, "{byte:02x}")?;
}
Ok(())
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub struct FixedWindowPolicy {
id: PolicyId,
scope: ScopeId,
limit: NonZeroU64,
window_millis: NonZeroU64,
fingerprint: PolicyFingerprint,
}
impl FixedWindowPolicy {
pub fn new(
id: PolicyId,
scope: ScopeId,
limit: u64,
window: Duration,
) -> Result<Self, PolicyError> {
let limit = NonZeroU64::new(limit).ok_or(PolicyError::ZeroLimit)?;
if limit.get() > MAX_LIMIT {
return Err(PolicyError::LimitTooLarge {
actual: limit.get(),
maximum: MAX_LIMIT,
});
}
let window_millis = validate_window(window)?;
let fingerprint = fingerprint(&id, &scope, limit, window_millis);
Ok(Self {
id,
scope,
limit,
window_millis,
fingerprint,
})
}
pub const fn id(&self) -> &PolicyId {
&self.id
}
pub const fn scope(&self) -> &ScopeId {
&self.scope
}
pub const fn limit(&self) -> u64 {
self.limit.get()
}
pub const fn window(&self) -> Duration {
Duration::from_millis(self.window_millis.get())
}
pub const fn window_millis(&self) -> u64 {
self.window_millis.get()
}
pub const fn fingerprint(&self) -> PolicyFingerprint {
self.fingerprint
}
}
fn validate_window(window: Duration) -> Result<NonZeroU64, PolicyError> {
if window.is_zero() {
return Err(PolicyError::ZeroWindow);
}
if !window.subsec_nanos().is_multiple_of(1_000_000) {
return Err(PolicyError::WindowNotWholeMilliseconds);
}
if window > MAX_WINDOW {
return Err(PolicyError::WindowTooLarge {
actual: window,
maximum: MAX_WINDOW,
});
}
let millis =
u64::try_from(window.as_millis()).expect("the portable window maximum fits in u64");
NonZeroU64::new(millis).ok_or(PolicyError::ZeroWindow)
}
fn fingerprint(
id: &PolicyId,
scope: &ScopeId,
limit: NonZeroU64,
window_millis: NonZeroU64,
) -> PolicyFingerprint {
let mut digest = Sha256::new();
digest.update(FINGERPRINT_DOMAIN);
digest.update(id.as_str().as_bytes());
digest.update([0]);
digest.update(scope.as_str().as_bytes());
digest.update([0]);
digest.update(limit.get().to_be_bytes());
digest.update(window_millis.get().to_be_bytes());
PolicyFingerprint(digest.finalize().into())
}
#[derive(Clone, Copy, Debug, Error, Eq, PartialEq)]
pub enum PolicyError {
#[error("fixed-window limit must be greater than zero")]
ZeroLimit,
#[error("fixed-window limit {actual} exceeds portable maximum {maximum}")]
LimitTooLarge {
actual: u64,
maximum: u64,
},
#[error("fixed-window duration must be greater than zero")]
ZeroWindow,
#[error("fixed-window duration must be an exact whole number of milliseconds")]
WindowNotWholeMilliseconds,
#[error("fixed-window duration {actual:?} exceeds portable maximum {maximum:?}")]
WindowTooLarge {
actual: Duration,
maximum: Duration,
},
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::{FixedWindowPolicy, MAX_LIMIT, MAX_WINDOW, MAX_WINDOW_MILLIS, PolicyError};
use crate::{PolicyId, ScopeId};
fn policy(limit: u64, window: Duration) -> Result<FixedWindowPolicy, PolicyError> {
FixedWindowPolicy::new(
PolicyId::new("auth.login").unwrap(),
ScopeId::new("client").unwrap(),
limit,
window,
)
}
#[test]
fn accepts_nonzero_whole_millisecond_windows() {
let policy = policy(8, Duration::from_millis(60_001)).unwrap();
assert_eq!(policy.limit(), 8);
assert_eq!(policy.window(), Duration::from_millis(60_001));
assert_eq!(policy.window_millis(), 60_001);
assert_eq!(policy.id().as_str(), "auth.login");
assert_eq!(policy.scope().as_str(), "client");
}
#[test]
fn rejects_zero_limit_and_window() {
assert_eq!(
policy(0, Duration::from_secs(1)),
Err(PolicyError::ZeroLimit)
);
assert_eq!(policy(1, Duration::ZERO), Err(PolicyError::ZeroWindow));
}
#[test]
fn rejects_sub_millisecond_and_fractional_millisecond_windows() {
assert_eq!(
policy(1, Duration::from_nanos(1)),
Err(PolicyError::WindowNotWholeMilliseconds)
);
assert_eq!(
policy(1, Duration::from_micros(1_500)),
Err(PolicyError::WindowNotWholeMilliseconds)
);
}
#[test]
fn accepts_portable_upper_bounds() {
let policy = policy(MAX_LIMIT, MAX_WINDOW).unwrap();
assert_eq!(policy.limit(), MAX_LIMIT);
assert_eq!(policy.window(), MAX_WINDOW);
assert_eq!(policy.window_millis(), MAX_WINDOW_MILLIS);
}
#[test]
fn rejects_limit_above_portable_maximum() {
assert_eq!(
policy(MAX_LIMIT + 1, Duration::from_secs(1)),
Err(PolicyError::LimitTooLarge {
actual: MAX_LIMIT + 1,
maximum: MAX_LIMIT,
})
);
}
#[test]
fn rejects_window_above_portable_maximum() {
let actual = MAX_WINDOW + Duration::from_millis(1);
assert_eq!(
policy(1, actual),
Err(PolicyError::WindowTooLarge {
actual,
maximum: MAX_WINDOW,
})
);
}
#[test]
fn rejects_windows_far_beyond_portable_maximum() {
let actual = Duration::from_secs(u64::MAX);
assert_eq!(
policy(1, actual),
Err(PolicyError::WindowTooLarge {
actual,
maximum: MAX_WINDOW,
})
);
}
#[test]
fn fingerprint_is_deterministic() {
let first = policy(8, Duration::from_secs(60)).unwrap();
let second = policy(8, Duration::from_secs(60)).unwrap();
assert_eq!(first.fingerprint(), second.fingerprint());
assert_eq!(first.fingerprint().as_bytes().len(), 32);
assert_eq!(first.fingerprint().to_string().len(), 64);
}
#[test]
fn fingerprint_changes_with_every_storage_relevant_field() {
let baseline = policy(8, Duration::from_secs(60)).unwrap();
let different_limit = policy(9, Duration::from_secs(60)).unwrap();
let different_window = policy(8, Duration::from_secs(61)).unwrap();
let different_id = FixedWindowPolicy::new(
PolicyId::new("auth.signup").unwrap(),
ScopeId::new("client").unwrap(),
8,
Duration::from_secs(60),
)
.unwrap();
let different_scope = FixedWindowPolicy::new(
PolicyId::new("auth.login").unwrap(),
ScopeId::new("identity").unwrap(),
8,
Duration::from_secs(60),
)
.unwrap();
assert_ne!(baseline.fingerprint(), different_limit.fingerprint());
assert_ne!(baseline.fingerprint(), different_window.fingerprint());
assert_ne!(baseline.fingerprint(), different_id.fingerprint());
assert_ne!(baseline.fingerprint(), different_scope.fingerprint());
}
}