use std::sync::{Arc, Mutex};
use schubert::PrincipalId;
use schubert::rate_limit::RateLimiter;
use crate::auth::AuthenticatedPrincipal;
pub type RateLimitState = Arc<Mutex<RateLimiter>>;
pub fn make_rate_limiter(base_tokens_per_second: f64, multiplier: f64) -> RateLimitState {
Arc::new(Mutex::new(RateLimiter::new(
base_tokens_per_second,
multiplier,
)))
}
pub fn consume(
state: &RateLimitState,
principal: &AuthenticatedPrincipal,
) -> schubert::error::Result<()> {
let mut rl = state.lock().expect("rate limiter poisoned");
let pid: PrincipalId = principal.principal.clone();
if rl.capacity(pid.clone()).is_none() {
let weight = principal
.grant
.capabilities
.iter()
.map(|c| c.partition.iter().sum::<usize>() as u64)
.max()
.unwrap_or(1);
rl.configure_principal(pid.clone(), weight);
}
rl.try_consume(pid).map(|_| ())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::auth::IjimaAuth;
use ijima_core::capabilities::{ADMIN, MEMORY_READ, MEMORY_WRITE};
fn auth() -> IjimaAuth {
IjimaAuth::from_embedded_policy().expect("policy must load")
}
fn principal(auth: &IjimaAuth, name: &str, cap: &str) -> AuthenticatedPrincipal {
let bearer = auth.issue_bearer(name, cap).expect("issue");
auth.verify_bearer(&bearer).expect("verify")
}
#[test]
fn first_request_configures_and_consumes() {
let state = make_rate_limiter(1.0, 1.0);
let auth = auth();
let p = principal(&auth, "elliott", MEMORY_READ);
assert!(consume(&state, &p).is_ok()); assert!(consume(&state, &p).is_err()); }
#[test]
fn higher_codimension_gets_more_capacity() {
let state = make_rate_limiter(5.0, 1.0);
let auth = auth();
let reader = principal(&auth, "reader", MEMORY_READ);
let writer = principal(&auth, "writer", MEMORY_WRITE);
for _ in 0..5 {
assert!(consume(&state, &reader).is_ok());
}
assert!(consume(&state, &reader).is_err());
for _ in 0..10 {
assert!(consume(&state, &writer).is_ok());
}
assert!(consume(&state, &writer).is_err());
}
#[test]
fn admin_gets_point_class_capacity() {
let state = make_rate_limiter(1.0, 1.0);
let auth = auth();
let admin = principal(&auth, "admin", ADMIN);
for _ in 0..16 {
assert!(consume(&state, &admin).is_ok());
}
assert!(consume(&state, &admin).is_err());
}
#[test]
fn multiplier_compresses_admin_ratio() {
let state = make_rate_limiter(10.0, 0.1);
let auth = auth();
let reader = principal(&auth, "reader", MEMORY_READ);
let admin = principal(&auth, "admin", ADMIN);
assert!(consume(&state, &reader).is_ok());
assert!(consume(&state, &reader).is_err());
for _ in 0..16 {
assert!(consume(&state, &admin).is_ok());
}
assert!(consume(&state, &admin).is_err());
}
#[test]
fn multi_cap_grant_uses_max_codimension() {
let state = make_rate_limiter(5.0, 1.0);
let auth = auth();
let bearer = auth
.issue_grant_bearer("pi", &[MEMORY_READ, MEMORY_WRITE])
.expect("issue");
let p = auth.verify_bearer(&bearer).expect("verify");
for _ in 0..10 {
assert!(consume(&state, &p).is_ok());
}
assert!(consume(&state, &p).is_err());
}
}