use std::{
collections::HashMap,
time::{SystemTime, UNIX_EPOCH},
};
use chia_sdk_client::{RateLimit, RateLimits, V2_RATE_LIMITS};
use chia_traits::Streamable;
use crate::DigMessage;
#[derive(Debug, Clone)]
pub struct OpcodeRateLimits {
default_settings: RateLimit,
non_tx_frequency: f64,
non_tx_max_total_size: f64,
tx: HashMap<u8, RateLimit>,
other: HashMap<u8, RateLimit>,
}
impl OpcodeRateLimits {
fn from_chia(limits: &RateLimits) -> Self {
let rekey = |map: &HashMap<chia_protocol::ProtocolMessageTypes, RateLimit>| {
map.iter()
.filter_map(|(msg_type, limit)| Some((*msg_type.to_bytes().ok()?.first()?, *limit)))
.collect()
};
Self {
default_settings: limits.default_settings,
non_tx_frequency: limits.non_tx_frequency,
non_tx_max_total_size: limits.non_tx_max_total_size,
tx: rekey(&limits.tx),
other: rekey(&limits.other),
}
}
}
impl Default for OpcodeRateLimits {
fn default() -> Self {
Self::from_chia(&V2_RATE_LIMITS)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Admission {
Admitted,
Deferred,
Unsendable,
}
#[derive(Debug, Clone)]
pub struct OpcodeRateLimiter {
reset_seconds: u64,
period: u64,
limit_factor: f64,
counts: HashMap<u8, f64>,
cumulative_sizes: HashMap<u8, f64>,
non_tx_count: f64,
non_tx_size: f64,
limits: OpcodeRateLimits,
}
impl OpcodeRateLimiter {
#[must_use]
pub fn new(reset_seconds: u64, limit_factor: f64, limits: OpcodeRateLimits) -> Self {
Self {
reset_seconds,
period: now_seconds() / reset_seconds,
limit_factor,
counts: HashMap::new(),
cumulative_sizes: HashMap::new(),
non_tx_count: 0.0,
non_tx_size: 0.0,
limits,
}
}
pub fn allow(&mut self, message: &DigMessage) -> bool {
self.admit(message) == Admission::Admitted
}
pub fn admit(&mut self, message: &DigMessage) -> Admission {
self.roll_window();
let size = f64::from(u32::try_from(message.data.len()).unwrap_or(u32::MAX));
let opcode = message.msg_type;
let mut limit = self.limits.default_settings;
let mut counts_against_non_tx = false;
if let Some(tx_limit) = self.limits.tx.get(&opcode) {
limit = *tx_limit;
} else if let Some(other_limit) = self.limits.other.get(&opcode) {
limit = *other_limit;
counts_against_non_tx = true;
}
let max_total = limit
.max_total_size
.unwrap_or(limit.frequency * limit.max_size);
let fits_an_empty_window = size <= limit.max_size
&& size <= max_total * self.limit_factor
&& 1.0 <= limit.frequency * self.limit_factor
&& (!counts_against_non_tx
|| (1.0 <= self.limits.non_tx_frequency * self.limit_factor
&& size <= self.limits.non_tx_max_total_size * self.limit_factor));
if !fits_an_empty_window {
return Admission::Unsendable;
}
let new_count = self.counts.get(&opcode).unwrap_or(&0.0) + 1.0;
let new_cumulative = self.cumulative_sizes.get(&opcode).unwrap_or(&0.0) + size;
let (new_non_tx_count, new_non_tx_size) = if counts_against_non_tx {
(self.non_tx_count + 1.0, self.non_tx_size + size)
} else {
(self.non_tx_count, self.non_tx_size)
};
let allowed = new_non_tx_count <= self.limits.non_tx_frequency * self.limit_factor
&& new_non_tx_size <= self.limits.non_tx_max_total_size * self.limit_factor
&& new_count <= limit.frequency * self.limit_factor
&& new_cumulative <= max_total * self.limit_factor;
if !allowed {
return Admission::Deferred;
}
self.counts.insert(opcode, new_count);
self.cumulative_sizes.insert(opcode, new_cumulative);
self.non_tx_count = new_non_tx_count;
self.non_tx_size = new_non_tx_size;
Admission::Admitted
}
fn roll_window(&mut self) {
let period = now_seconds() / self.reset_seconds;
if self.period == period {
return;
}
self.period = period;
self.counts.clear();
self.cumulative_sizes.clear();
self.non_tx_count = 0.0;
self.non_tx_size = 0.0;
}
}
fn now_seconds() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("system clock is before the unix epoch")
.as_secs()
}
#[cfg(test)]
mod tests {
use super::{Admission, OpcodeRateLimiter, OpcodeRateLimits};
use crate::{Bytes, DigMessage, DIG_MESSAGE};
use chia_protocol::ProtocolMessageTypes;
use chia_traits::Streamable;
fn message(opcode: u8, payload_len: usize) -> DigMessage {
DigMessage::new(opcode, None, Bytes::new(vec![0u8; payload_len]))
}
#[test]
fn chia_opcodes_keep_their_upstream_limits() {
let limits = OpcodeRateLimits::default();
let handshake = *ProtocolMessageTypes::Handshake
.to_bytes()
.expect("encode")
.first()
.expect("one byte");
let upstream = chia_sdk_client::V2_RATE_LIMITS
.other
.get(&ProtocolMessageTypes::Handshake)
.expect("upstream defines a handshake limit");
let ours = limits
.other
.get(&handshake)
.expect("re-keyed table kept the handshake limit");
assert_eq!(ours.frequency, upstream.frequency);
assert_eq!(ours.max_size, upstream.max_size);
}
#[test]
fn dig_opcodes_fall_back_to_the_default_budget() {
let mut limiter = OpcodeRateLimiter::new(60, 1.0, OpcodeRateLimits::default());
assert!(limiter.allow(&message(DIG_MESSAGE, 16)));
}
#[test]
fn frequency_budget_admits_up_to_the_bound_and_refuses_past_it() {
let limits = OpcodeRateLimits::default();
let allowance = limits.default_settings.frequency as usize;
let mut limiter = OpcodeRateLimiter::new(60, 1.0, limits);
for i in 0..allowance {
assert!(
limiter.allow(&message(DIG_MESSAGE, 1)),
"message {i} refused below the bound"
);
}
assert!(
!limiter.allow(&message(DIG_MESSAGE, 1)),
"one message over the bound was admitted"
);
}
#[test]
fn a_deferrable_refusal_is_distinguished_from_a_permanent_one() {
let limits = OpcodeRateLimits::default();
let allowance = limits.default_settings.frequency as usize;
let max_size = limits.default_settings.max_size as usize;
let mut exhausted = OpcodeRateLimiter::new(60, 1.0, limits);
for _ in 0..allowance {
assert_eq!(
exhausted.admit(&message(DIG_MESSAGE, 1)),
Admission::Admitted
);
}
assert_eq!(
exhausted.admit(&message(DIG_MESSAGE, 1)),
Admission::Deferred,
"an exhausted frequency budget resets on the next window, so waiting can help"
);
let mut fresh = OpcodeRateLimiter::new(60, 1.0, OpcodeRateLimits::default());
assert_eq!(
fresh.admit(&message(DIG_MESSAGE, max_size + 1)),
Admission::Unsendable,
"an oversized message is refused identically in every window"
);
}
#[test]
fn size_cap_is_pinned_from_both_sides() {
let max_size = OpcodeRateLimits::default().default_settings.max_size as usize;
let mut at_bound = OpcodeRateLimiter::new(60, 1.0, OpcodeRateLimits::default());
assert!(at_bound.allow(&message(DIG_MESSAGE, max_size)));
let mut over_bound = OpcodeRateLimiter::new(60, 1.0, OpcodeRateLimits::default());
assert!(!over_bound.allow(&message(DIG_MESSAGE, max_size + 1)));
}
#[test]
fn the_rekeyed_table_pins_upstream_limits_at_absolute_values() {
let limits = OpcodeRateLimits::default();
let handshake = limits
.other
.get(&1)
.expect("opcode 1 (Handshake) kept its entry");
assert_eq!(handshake.frequency, 5.0, "Handshake frequency");
assert_eq!(handshake.max_size, 10.0 * 1024.0, "Handshake max_size");
let tx_ack = limits
.tx
.get(&49)
.expect("opcode 49 (TransactionAck) kept its tx entry");
assert_eq!(tx_ack.frequency, 5000.0, "TransactionAck frequency");
assert_eq!(tx_ack.max_size, 2048.0, "TransactionAck max_size");
let new_tx = limits
.tx
.get(&21)
.expect("opcode 21 (NewTransaction) kept its tx entry");
assert_eq!(new_tx.frequency, 5000.0, "NewTransaction frequency");
assert_eq!(new_tx.max_size, 100.0, "NewTransaction max_size");
assert_eq!(limits.non_tx_frequency, 1000.0);
assert_eq!(limits.non_tx_max_total_size, 100.0 * 1024.0 * 1024.0);
assert_eq!(limits.default_settings.frequency, 100.0);
assert_eq!(limits.default_settings.max_size, 1024.0 * 1024.0);
}
#[test]
fn falling_back_to_the_default_would_be_a_detectable_loosening() {
let limits = OpcodeRateLimits::default();
let handshake = limits
.other
.get(&1)
.expect("opcode 1 (Handshake) kept its entry");
assert!(
handshake.frequency < limits.default_settings.frequency,
"Handshake ({}) is not tighter than default ({}) -- the pin above can no longer distinguish a re-keyed table from a collapsed one",
handshake.frequency,
limits.default_settings.frequency
);
assert!(
handshake.max_size < limits.default_settings.max_size,
"Handshake max_size is not tighter than default"
);
}
#[test]
fn the_rekeyed_table_retains_the_bulk_of_the_upstream_entries() {
let limits = OpcodeRateLimits::default();
assert!(
limits.other.len() >= 30,
"other map holds only {} entries -- the re-key lost most of the table",
limits.other.len()
);
assert!(
limits.tx.len() >= 5,
"tx map holds only {} entries -- the re-key lost most of the table",
limits.tx.len()
);
for opcode in limits.other.keys().chain(limits.tx.keys()) {
assert!(
*opcode < 200,
"opcode {opcode} is outside the chia band -- the re-key is keying off a different ProtocolMessageTypes than the wire uses"
);
}
}
}