1use core::num::NonZeroU64;
4
5use super::DeliveryOutcomes;
6
7use behavior::{
8 Actions, Address, Behavior, BehaviorActed, BehaviorBase, Never, NoBirths, Protocol,
9 SendEffects, User,
10};
11use thiserror::Error;
12
13use crate::DeliveryRoute;
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
17pub struct TokenCount(pub NonZeroU64);
18impl TokenCount {
19 #[must_use]
21 pub const fn new(value: NonZeroU64) -> Self {
22 Self(value)
23 }
24 #[must_use]
26 pub const fn get(self) -> u64 {
27 self.0.get()
28 }
29}
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub struct RateLimiterState {
34 pub capacity: TokenCount,
36 available: u64,
37}
38
39impl RateLimiterState {
40 #[must_use]
42 pub fn available(&self) -> u64 {
43 self.available
44 }
45}
46
47#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
49pub enum RateLimitRejection {
50 ExceedsCapacity,
52 InsufficientTokens,
54}
55
56#[derive(Debug, PartialEq, Eq)]
58pub enum RateLimiterOutcome<T> {
59 Admitted {
61 remaining: u64,
63 },
64 Rejected {
66 cost: TokenCount,
68 value: T,
70 reason: RateLimitRejection,
72 },
73}
74
75pub enum RateLimiterMessage<T, TargetRoute, ReplyRoute> {
77 Acquire {
79 cost: TokenCount,
81 value: T,
83 to: TargetRoute,
85 reply_to: ReplyRoute,
87 },
88 Refill {
90 tokens: TokenCount,
92 },
93}
94
95#[derive(Debug, Error, Clone, Copy, PartialEq, Eq)]
97pub enum RateLimiterConfigError {
98 #[error("initial rate-limit tokens exceed capacity")]
100 InitialExceedsCapacity {
101 capacity: TokenCount,
103 initial: u64,
105 },
106}
107
108pub struct RateLimiter<
119 A: Address,
120 T,
121 TargetRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = T>>,
122 ReplyRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = RateLimiterOutcome<T>>>,
123> {
124 capacity: TokenCount,
125 available: u64,
126 marker: core::marker::PhantomData<fn() -> (A, TargetRoute, ReplyRoute)>,
127}
128type RateActions<A, TargetSends, OutcomeSends> =
129 Actions<A, Never, DeliveryOutcomes<TargetSends, OutcomeSends>, NoBirths>;
130impl<A, T, TargetRoute, ReplyRoute> RateLimiter<A, T, TargetRoute, ReplyRoute>
131where
132 A: Address,
133 TargetRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = T>>,
134 ReplyRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = RateLimiterOutcome<T>>>,
135{
136 pub fn new(capacity: TokenCount, initial: u64) -> Result<Self, RateLimiterConfigError> {
140 if initial > capacity.get() {
141 return Err(RateLimiterConfigError::InitialExceedsCapacity { capacity, initial });
142 }
143 Ok(Self {
144 capacity,
145 available: initial,
146 marker: core::marker::PhantomData,
147 })
148 }
149 #[must_use]
151 pub const fn state(&self) -> RateLimiterState {
152 RateLimiterState {
153 capacity: self.capacity,
154 available: self.available,
155 }
156 }
157 fn result(
158 deliveries: TargetRoute::Sends,
159 reply_to: ReplyRoute,
160 outcome: RateLimiterOutcome<T>,
161 ) -> RateActions<A, TargetRoute::Sends, ReplyRoute::Sends> {
162 Actions::send(DeliveryOutcomes {
163 deliveries,
164 outcomes: reply_to.deliver(outcome),
165 })
166 }
167 fn acquire(
168 &mut self,
169 cost: TokenCount,
170 value: T,
171 to: TargetRoute,
172 reply_to: ReplyRoute,
173 ) -> RateActions<A, TargetRoute::Sends, ReplyRoute::Sends> {
174 let required = cost.get();
175 if required > self.capacity.get() {
176 return Self::result(
177 TargetRoute::Sends::empty(),
178 reply_to,
179 RateLimiterOutcome::Rejected {
180 cost,
181 value,
182 reason: RateLimitRejection::ExceedsCapacity,
183 },
184 );
185 }
186 if required > self.available {
187 return Self::result(
188 TargetRoute::Sends::empty(),
189 reply_to,
190 RateLimiterOutcome::Rejected {
191 cost,
192 value,
193 reason: RateLimitRejection::InsufficientTokens,
194 },
195 );
196 }
197 self.available -= required;
198 Self::result(
199 to.deliver(value),
200 reply_to,
201 RateLimiterOutcome::Admitted {
202 remaining: self.available,
203 },
204 )
205 }
206}
207impl<A, T, TargetRoute, ReplyRoute> BehaviorBase for RateLimiter<A, T, TargetRoute, ReplyRoute>
208where
209 A: Address,
210 TargetRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = T>>,
211 ReplyRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = RateLimiterOutcome<T>>>,
212{
213 type Base = Self;
214 fn base(&self) -> &Self {
215 self
216 }
217}
218impl<A, T, TargetRoute, ReplyRoute> behavior::Protocol
219 for RateLimiter<A, T, TargetRoute, ReplyRoute>
220where
221 A: Address,
222 TargetRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = T>>,
223 ReplyRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = RateLimiterOutcome<T>>>,
224{
225 type Addr = A;
226 type Msg = RateLimiterMessage<T, TargetRoute, ReplyRoute>;
227}
228
229impl<A, T, TargetRoute, ReplyRoute> Behavior for RateLimiter<A, T, TargetRoute, ReplyRoute>
230where
231 A: Address,
232 TargetRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = T>>,
233 ReplyRoute: DeliveryRoute<Protocol: Protocol<Addr = A, Msg = RateLimiterOutcome<T>>>,
234 TargetRoute::Sends: behavior::SendsFor<User<A, RateLimiterMessage<T, TargetRoute, ReplyRoute>>>,
235 ReplyRoute::Sends: behavior::SendsFor<User<A, RateLimiterMessage<T, TargetRoute, ReplyRoute>>>,
236{
237 type Protocol = Self;
238 type Event = User<A, behavior::BehaviorMessage<Self>>;
239 type Sends = DeliveryOutcomes<TargetRoute::Sends, ReplyRoute::Sends>;
240 type Ph = Never;
241 type Error = Never;
242 type Birth = NoBirths;
243 fn transition(&mut self, _: behavior::ActiveTurn, event: Self::Event) -> BehaviorActed<Self> {
244 Ok(match event.message {
245 RateLimiterMessage::Acquire {
246 cost,
247 value,
248 to,
249 reply_to,
250 } => self.acquire(cost, value, to, reply_to),
251 RateLimiterMessage::Refill { tokens } => {
252 self.available = self
253 .available
254 .saturating_add(tokens.get())
255 .min(self.capacity.get());
256 Actions::cont()
257 }
258 })
259 }
260}
261
262#[cfg(test)]
263mod tests {
264 use super::*;
265 use crate::Activate as _;
266 use behavior::{Delivery, MailAddr, Recipient};
267 struct Target;
268 struct Reply;
269 impl behavior::Protocol for Target {
270 type Addr = MailAddr;
271 type Msg = u8;
272 }
273
274 impl Behavior for Target {
275 type Protocol = Self;
276 type Event = User<MailAddr, u8>;
277 type Sends = Vec<Never>;
278 type Ph = Never;
279 type Error = Never;
280 type Birth = NoBirths;
281 fn transition(&mut self, _: behavior::ActiveTurn, _: Self::Event) -> BehaviorActed<Self> {
282 Ok(Actions::cont())
283 }
284 }
285 impl behavior::Protocol for Reply {
286 type Addr = MailAddr;
287 type Msg = RateLimiterOutcome<u8>;
288 }
289
290 impl Behavior for Reply {
291 type Protocol = Self;
292 type Event = User<MailAddr, behavior::BehaviorMessage<Self>>;
293 type Sends = Vec<Never>;
294 type Ph = Never;
295 type Error = Never;
296 type Birth = NoBirths;
297 fn transition(&mut self, _: behavior::ActiveTurn, _: Self::Event) -> BehaviorActed<Self> {
298 Ok(Actions::cont())
299 }
300 }
301 type Subject = RateLimiter<MailAddr, u8, Recipient<Target>, Recipient<Reply>>;
302 fn tokens(n: u64) -> TokenCount {
303 TokenCount::new(NonZeroU64::new(n).unwrap())
304 }
305 fn acquire(
306 s: &mut crate::Active<Subject>,
307 cost: u64,
308 value: u8,
309 ) -> RateActions<MailAddr, Vec<Delivery<Target>>, Vec<Delivery<Reply>>> {
310 s.receive(
311 MailAddr(0),
312 RateLimiterMessage::Acquire {
313 cost: tokens(cost),
314 value,
315 to: Recipient::global(MailAddr(1)),
316 reply_to: Recipient::global(MailAddr(2)),
317 },
318 )
319 .unwrap()
320 }
321 #[test]
322 fn admission_and_rejections_preserve_tokens_and_ownership() {
323 let mut s = (Subject::new(tokens(5), 3).unwrap())
324 .initialize()
325 .unwrap()
326 .behavior;
327 let admitted = acquire(&mut s, 2, 7);
328 assert_eq!(admitted.sends.deliveries[0].message, 7);
329 let insufficient = acquire(&mut s, 2, 8);
330 assert!(matches!(
331 insufficient.sends.outcomes[0].message,
332 RateLimiterOutcome::Rejected {
333 cost,
334 value: 8,
335 reason: RateLimitRejection::InsufficientTokens
336 } if cost == tokens(2)
337 ));
338 let over_capacity = acquire(&mut s, 6, 9);
339 assert!(matches!(
340 over_capacity.sends.outcomes[0].message,
341 RateLimiterOutcome::Rejected {
342 cost,
343 value: 9,
344 reason: RateLimitRejection::ExceedsCapacity
345 } if cost == tokens(6)
346 ));
347 assert_eq!(s.state().available(), 1);
348 }
349 #[test]
350 fn refill_saturates_without_overflow() {
351 let mut s = (Subject::new(tokens(5), 0).unwrap())
352 .initialize()
353 .unwrap()
354 .behavior;
355 let refilled = s
356 .receive(
357 MailAddr(0),
358 RateLimiterMessage::Refill {
359 tokens: tokens(u64::MAX),
360 },
361 )
362 .unwrap();
363 assert!(refilled.sends.deliveries.is_empty());
364 assert!(refilled.sends.outcomes.is_empty());
365 assert!(refilled.creates.is_empty());
366 assert_eq!(refilled.become_, behavior::Step::Continue);
367 assert_eq!(s.state().available(), 5);
368 }
369}