1use std::{
24 collections::HashMap,
25 time::{SystemTime, UNIX_EPOCH},
26};
27
28use chia_sdk_client::{RateLimit, RateLimits, V2_RATE_LIMITS};
29use chia_traits::Streamable;
30
31use crate::DigMessage;
32
33#[derive(Debug, Clone)]
38pub struct OpcodeRateLimits {
39 default_settings: RateLimit,
40 non_tx_frequency: f64,
41 non_tx_max_total_size: f64,
42 tx: HashMap<u8, RateLimit>,
43 other: HashMap<u8, RateLimit>,
44}
45
46impl OpcodeRateLimits {
47 fn from_chia(limits: &RateLimits) -> Self {
49 let rekey = |map: &HashMap<chia_protocol::ProtocolMessageTypes, RateLimit>| {
52 map.iter()
53 .filter_map(|(msg_type, limit)| Some((*msg_type.to_bytes().ok()?.first()?, *limit)))
54 .collect()
55 };
56
57 Self {
58 default_settings: limits.default_settings,
59 non_tx_frequency: limits.non_tx_frequency,
60 non_tx_max_total_size: limits.non_tx_max_total_size,
61 tx: rekey(&limits.tx),
62 other: rekey(&limits.other),
63 }
64 }
65}
66
67impl Default for OpcodeRateLimits {
68 fn default() -> Self {
69 Self::from_chia(&V2_RATE_LIMITS)
70 }
71}
72
73#[derive(Debug, Clone, Copy, PartialEq, Eq)]
78pub enum Admission {
79 Admitted,
81 Deferred,
83 Unsendable,
86}
87
88#[derive(Debug, Clone)]
93pub struct OpcodeRateLimiter {
94 reset_seconds: u64,
95 period: u64,
96 limit_factor: f64,
97 counts: HashMap<u8, f64>,
98 cumulative_sizes: HashMap<u8, f64>,
99 non_tx_count: f64,
100 non_tx_size: f64,
101 limits: OpcodeRateLimits,
102}
103
104impl OpcodeRateLimiter {
105 #[must_use]
110 pub fn new(reset_seconds: u64, limit_factor: f64, limits: OpcodeRateLimits) -> Self {
111 Self {
112 reset_seconds,
113 period: now_seconds() / reset_seconds,
114 limit_factor,
115 counts: HashMap::new(),
116 cumulative_sizes: HashMap::new(),
117 non_tx_count: 0.0,
118 non_tx_size: 0.0,
119 limits,
120 }
121 }
122
123 pub fn allow(&mut self, message: &DigMessage) -> bool {
131 self.admit(message) == Admission::Admitted
132 }
133
134 pub fn admit(&mut self, message: &DigMessage) -> Admission {
141 self.roll_window();
142
143 let size = f64::from(u32::try_from(message.data.len()).unwrap_or(u32::MAX));
144 let opcode = message.msg_type;
145
146 let mut limit = self.limits.default_settings;
147 let mut counts_against_non_tx = false;
148 if let Some(tx_limit) = self.limits.tx.get(&opcode) {
149 limit = *tx_limit;
150 } else if let Some(other_limit) = self.limits.other.get(&opcode) {
151 limit = *other_limit;
152 counts_against_non_tx = true;
153 }
154
155 let max_total = limit
156 .max_total_size
157 .unwrap_or(limit.frequency * limit.max_size);
158
159 let fits_an_empty_window = size <= limit.max_size
162 && size <= max_total * self.limit_factor
163 && 1.0 <= limit.frequency * self.limit_factor
164 && (!counts_against_non_tx
165 || (1.0 <= self.limits.non_tx_frequency * self.limit_factor
166 && size <= self.limits.non_tx_max_total_size * self.limit_factor));
167 if !fits_an_empty_window {
168 return Admission::Unsendable;
169 }
170
171 let new_count = self.counts.get(&opcode).unwrap_or(&0.0) + 1.0;
172 let new_cumulative = self.cumulative_sizes.get(&opcode).unwrap_or(&0.0) + size;
173 let (new_non_tx_count, new_non_tx_size) = if counts_against_non_tx {
174 (self.non_tx_count + 1.0, self.non_tx_size + size)
175 } else {
176 (self.non_tx_count, self.non_tx_size)
177 };
178
179 let allowed = new_non_tx_count <= self.limits.non_tx_frequency * self.limit_factor
180 && new_non_tx_size <= self.limits.non_tx_max_total_size * self.limit_factor
181 && new_count <= limit.frequency * self.limit_factor
182 && new_cumulative <= max_total * self.limit_factor;
183
184 if !allowed {
185 return Admission::Deferred;
186 }
187
188 self.counts.insert(opcode, new_count);
189 self.cumulative_sizes.insert(opcode, new_cumulative);
190 self.non_tx_count = new_non_tx_count;
191 self.non_tx_size = new_non_tx_size;
192 Admission::Admitted
193 }
194
195 fn roll_window(&mut self) {
197 let period = now_seconds() / self.reset_seconds;
198 if self.period == period {
199 return;
200 }
201 self.period = period;
202 self.counts.clear();
203 self.cumulative_sizes.clear();
204 self.non_tx_count = 0.0;
205 self.non_tx_size = 0.0;
206 }
207}
208
209fn now_seconds() -> u64 {
210 SystemTime::now()
211 .duration_since(UNIX_EPOCH)
212 .expect("system clock is before the unix epoch")
213 .as_secs()
214}
215
216#[cfg(test)]
217mod tests {
218 use super::{Admission, OpcodeRateLimiter, OpcodeRateLimits};
219 use crate::{DigMessage, DIG_MESSAGE};
220 use chia_protocol::{Bytes, ProtocolMessageTypes};
221 use chia_traits::Streamable;
222
223 fn message(opcode: u8, payload_len: usize) -> DigMessage {
224 DigMessage::new(opcode, None, Bytes::new(vec![0u8; payload_len]))
225 }
226
227 #[test]
232 fn chia_opcodes_keep_their_upstream_limits() {
233 let limits = OpcodeRateLimits::default();
234 let handshake = *ProtocolMessageTypes::Handshake
235 .to_bytes()
236 .expect("encode")
237 .first()
238 .expect("one byte");
239
240 let upstream = chia_sdk_client::V2_RATE_LIMITS
241 .other
242 .get(&ProtocolMessageTypes::Handshake)
243 .expect("upstream defines a handshake limit");
244 let ours = limits
245 .other
246 .get(&handshake)
247 .expect("re-keyed table kept the handshake limit");
248
249 assert_eq!(ours.frequency, upstream.frequency);
250 assert_eq!(ours.max_size, upstream.max_size);
251 }
252
253 #[test]
256 fn dig_opcodes_fall_back_to_the_default_budget() {
257 let mut limiter = OpcodeRateLimiter::new(60, 1.0, OpcodeRateLimits::default());
258 assert!(limiter.allow(&message(DIG_MESSAGE, 16)));
259 }
260
261 #[test]
264 fn frequency_budget_admits_up_to_the_bound_and_refuses_past_it() {
265 let limits = OpcodeRateLimits::default();
266 let allowance = limits.default_settings.frequency as usize;
267 let mut limiter = OpcodeRateLimiter::new(60, 1.0, limits);
268
269 for i in 0..allowance {
270 assert!(
271 limiter.allow(&message(DIG_MESSAGE, 1)),
272 "message {i} refused below the bound"
273 );
274 }
275 assert!(
276 !limiter.allow(&message(DIG_MESSAGE, 1)),
277 "one message over the bound was admitted"
278 );
279 }
280
281 #[test]
288 fn a_deferrable_refusal_is_distinguished_from_a_permanent_one() {
289 let limits = OpcodeRateLimits::default();
290 let allowance = limits.default_settings.frequency as usize;
291 let max_size = limits.default_settings.max_size as usize;
292
293 let mut exhausted = OpcodeRateLimiter::new(60, 1.0, limits);
294 for _ in 0..allowance {
295 assert_eq!(
296 exhausted.admit(&message(DIG_MESSAGE, 1)),
297 Admission::Admitted
298 );
299 }
300 assert_eq!(
301 exhausted.admit(&message(DIG_MESSAGE, 1)),
302 Admission::Deferred,
303 "an exhausted frequency budget resets on the next window, so waiting can help"
304 );
305
306 let mut fresh = OpcodeRateLimiter::new(60, 1.0, OpcodeRateLimits::default());
307 assert_eq!(
308 fresh.admit(&message(DIG_MESSAGE, max_size + 1)),
309 Admission::Unsendable,
310 "an oversized message is refused identically in every window"
311 );
312 }
313
314 #[test]
321 fn size_cap_is_pinned_from_both_sides() {
322 let max_size = OpcodeRateLimits::default().default_settings.max_size as usize;
323
324 let mut at_bound = OpcodeRateLimiter::new(60, 1.0, OpcodeRateLimits::default());
325 assert!(at_bound.allow(&message(DIG_MESSAGE, max_size)));
326
327 let mut over_bound = OpcodeRateLimiter::new(60, 1.0, OpcodeRateLimits::default());
328 assert!(!over_bound.allow(&message(DIG_MESSAGE, max_size + 1)));
329 }
330}