use std::collections::HashMap;
use std::hash::Hash;
use std::sync::{Arc, Mutex, MutexGuard};
use crate::frame::VerifiedRequest;
use super::LinkError;
const DEADLINE_PAST_TOLERANCE_MS: i64 = 5 * 60_000;
const DEADLINE_AHEAD_MAX_MS: i64 = 10 * 60_000;
const KEPT_PAST_DEADLINE_MS: i64 = 5 * 60_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AdmissionLimits {
pub caller_quota: usize,
pub share: usize,
pub cap: usize,
pub reply_bytes: usize,
pub reply_bytes_total: usize,
pub sessions_per_caller: usize,
pub sessions: usize,
pub inbox_bytes_per_caller: usize,
pub inbox_bytes: usize,
}
impl Default for AdmissionLimits {
fn default() -> Self {
AdmissionLimits {
caller_quota: 256,
share: 1024,
cap: 1024,
reply_bytes: 256 * 1024,
reply_bytes_total: 16 * 1024 * 1024,
sessions_per_caller: 16,
sessions: 1000,
inbox_bytes_per_caller: 16 * 1024 * 1024,
inbox_bytes: 256 * 1024 * 1024,
}
}
}
impl AdmissionLimits {
pub fn validate(&self) -> Result<(), LinkError> {
let l = self;
let valid = l.caller_quota > 0
&& l.share > 0
&& l.cap > 0
&& l.reply_bytes > 0
&& l.reply_bytes_total > 0
&& l.caller_quota <= l.share
&& l.reply_bytes <= l.reply_bytes_total
&& l.sessions_per_caller > 0
&& l.sessions >= l.sessions_per_caller
&& l.inbox_bytes_per_caller > 0
&& l.inbox_bytes >= l.inbox_bytes_per_caller;
if valid {
Ok(())
} else {
Err(LinkError::InvalidConfig(
"admission limits must be positive, each per-caller bound within its total".into(),
))
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(super) enum Verdict {
Refused(&'static str),
Copy(Option<Vec<u8>>),
New,
}
struct Entry {
hash: [u8; 48],
expires_at: i64,
share: String,
answered: bool,
reply: Option<Vec<u8>>,
}
#[derive(Default)]
struct Held {
entries: HashMap<([u8; 32], [u8; 16]), Entry>,
callers: HashMap<[u8; 32], usize>,
shares: HashMap<String, usize>,
reply_bytes: HashMap<[u8; 32], usize>,
reply_total: usize,
sessions: HashMap<[u8; 32], usize>,
sessions_total: usize,
inbox: HashMap<[u8; 32], usize>,
inbox_total: usize,
}
pub struct Admission {
limits: AdmissionLimits,
held: Mutex<Held>,
}
impl Admission {
pub fn new(limits: AdmissionLimits) -> Admission {
Admission {
limits,
held: Mutex::new(Held::default()),
}
}
pub fn limits(&self) -> AdmissionLimits {
self.limits
}
fn lock(&self) -> MutexGuard<'_, Held> {
self.held.lock().unwrap_or_else(|p| p.into_inner())
}
pub(super) fn admit(&self, request: &VerifiedRequest, share: &str, now_ms: i64) -> Verdict {
let deadline = request.deadline as i64;
if deadline < now_ms - DEADLINE_PAST_TOLERANCE_MS {
return Verdict::Refused("expired");
}
if deadline > now_ms + DEADLINE_AHEAD_MAX_MS {
return Verdict::Refused("not_yet_valid");
}
let mut held = self.lock();
held.sweep(now_ms);
let key = (request.caller, request.request_id);
if let Some(entry) = held.entries.get(&key) {
if entry.hash != request.request_hash {
return Verdict::Refused("request_id_reused");
}
if entry.answered && entry.reply.is_none() {
return Verdict::Refused("reply_not_kept");
}
return Verdict::Copy(entry.reply.clone());
}
if held.callers.get(&request.caller).copied().unwrap_or(0) >= self.limits.caller_quota {
return Verdict::Refused("caller_quota");
}
if held.shares.get(share).copied().unwrap_or(0) >= self.limits.share {
return Verdict::Refused("share_full");
}
if held.entries.len() >= self.limits.cap {
return Verdict::Refused("admission_full");
}
held.entries.insert(
key,
Entry {
hash: request.request_hash,
expires_at: deadline + KEPT_PAST_DEADLINE_MS,
share: share.to_string(),
answered: false,
reply: None,
},
);
*held.callers.entry(request.caller).or_default() += 1;
*held.shares.entry(share.to_string()).or_default() += 1;
Verdict::New
}
pub(super) fn store(&self, request: &VerifiedRequest, reply: Vec<u8>) {
let mut held = self.lock();
let held = &mut *held;
let Some(entry) = held.entries.get_mut(&(request.caller, request.request_id)) else {
return;
};
if entry.answered || entry.hash != request.request_hash {
return;
}
entry.answered = true;
let caller_bytes = held.reply_bytes.get(&request.caller).copied().unwrap_or(0);
if caller_bytes + reply.len() > self.limits.reply_bytes
|| held.reply_total + reply.len() > self.limits.reply_bytes_total
{
return;
}
*held.reply_bytes.entry(request.caller).or_default() += reply.len();
held.reply_total += reply.len();
entry.reply = Some(reply);
}
pub(super) fn open_session(self: &Arc<Self>, caller: [u8; 32]) -> Option<SessionPlace> {
let mut held = self.lock();
if held.sessions.get(&caller).copied().unwrap_or(0) >= self.limits.sessions_per_caller
|| held.sessions_total >= self.limits.sessions
{
return None;
}
*held.sessions.entry(caller).or_default() += 1;
held.sessions_total += 1;
Some(SessionPlace {
admission: self.clone(),
caller,
})
}
pub(super) fn charge_inbox(&self, caller: [u8; 32], n: usize) -> bool {
let mut held = self.lock();
if held.inbox.get(&caller).copied().unwrap_or(0) + n > self.limits.inbox_bytes_per_caller
|| held.inbox_total + n > self.limits.inbox_bytes
{
return false;
}
*held.inbox.entry(caller).or_default() += n;
held.inbox_total += n;
true
}
pub(super) fn release_inbox(&self, caller: [u8; 32], n: usize) {
let mut held = self.lock();
decrement(&mut held.inbox, caller, n);
held.inbox_total = held.inbox_total.saturating_sub(n);
}
}
pub(super) struct SessionPlace {
admission: Arc<Admission>,
caller: [u8; 32],
}
impl Drop for SessionPlace {
fn drop(&mut self) {
let mut held = self.admission.lock();
decrement(&mut held.sessions, self.caller, 1);
held.sessions_total = held.sessions_total.saturating_sub(1);
}
}
impl Held {
fn sweep(&mut self, now_ms: i64) {
let expired: Vec<_> = self
.entries
.iter()
.filter(|(_, e)| e.expires_at < now_ms)
.map(|(k, _)| *k)
.collect();
for key in expired {
let Some(entry) = self.entries.remove(&key) else {
continue;
};
decrement(&mut self.callers, key.0, 1);
decrement(&mut self.shares, entry.share, 1);
if let Some(reply) = entry.reply {
decrement(&mut self.reply_bytes, key.0, reply.len());
self.reply_total = self.reply_total.saturating_sub(reply.len());
}
}
}
}
fn decrement<K: Eq + Hash>(m: &mut HashMap<K, usize>, k: K, n: usize) {
if let Some(v) = m.get_mut(&k) {
*v = v.saturating_sub(n);
if *v == 0 {
m.remove(&k);
}
}
}