1use core::fmt;
12
13pub type Poll<'a> = &'a (dyn Fn() -> bool + Sync);
16
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
19pub enum Stop {
20 Poll,
22 Budget,
24}
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
29pub struct Interrupted {
30 pub by: Stop,
32 pub steps: u64,
36}
37
38impl fmt::Display for Interrupted {
39 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
40 match self.by {
41 Stop::Poll => write!(
42 f,
43 "interrupted by the caller's poll after {} steps",
44 self.steps
45 ),
46 Stop::Budget => write!(
47 f,
48 "interrupted: the budget of {} steps is spent",
49 self.steps
50 ),
51 }
52 }
53}
54
55impl core::error::Error for Interrupted {}
56
57#[derive(Clone, Copy, Default)]
83pub struct Control<'a> {
84 poll: Option<Poll<'a>>,
85 budget: Option<u64>,
86}
87
88impl fmt::Debug for Control<'_> {
89 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
90 f.debug_struct("Control")
91 .field("poll", &self.poll.map(|_| "…"))
92 .field("budget", &self.budget)
93 .finish()
94 }
95}
96
97impl Control<'static> {
98 pub const NONE: Control<'static> = Control {
100 poll: None,
101 budget: None,
102 };
103
104 pub const fn budget(steps: u64) -> Control<'static> {
107 Control {
108 poll: None,
109 budget: Some(steps),
110 }
111 }
112}
113
114impl<'a> Control<'a> {
115 pub const fn poll(poll: Poll<'a>) -> Control<'a> {
118 Control {
119 poll: Some(poll),
120 budget: None,
121 }
122 }
123
124 #[must_use]
126 pub const fn with_budget(mut self, steps: u64) -> Self {
127 self.budget = Some(steps);
128 self
129 }
130
131 #[must_use]
133 pub const fn with_poll(mut self, poll: Poll<'a>) -> Self {
134 self.poll = Some(poll);
135 self
136 }
137}
138
139#[derive(Clone, Copy)]
147pub struct Meter<'a> {
148 poll: Option<Poll<'a>>,
149 cap: Option<u64>,
150 steps: u64,
151 stopped: Option<Interrupted>,
152}
153
154impl fmt::Debug for Meter<'_> {
155 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
156 f.debug_struct("Meter")
157 .field("cap", &self.cap)
158 .field("steps", &self.steps)
159 .finish()
160 }
161}
162
163impl Default for Meter<'_> {
164 fn default() -> Self {
166 Meter::new(&Control::NONE)
167 }
168}
169
170impl<'a> Meter<'a> {
171 pub fn new(control: &Control<'a>) -> Self {
173 Meter {
174 poll: control.poll,
175 cap: control.budget,
176 steps: 0,
177 stopped: None,
178 }
179 }
180
181 pub fn tick(&mut self) -> Result<(), Interrupted> {
185 if self.cap.is_some_and(|cap| self.steps >= cap) {
186 return Err(self.stopped_by(Stop::Budget));
187 }
188 if self.poll.is_some_and(|poll| poll()) {
189 return Err(self.stopped_by(Stop::Poll));
190 }
191 self.steps += 1;
192 Ok(())
193 }
194
195 pub fn stopped(&self) -> Option<Interrupted> {
201 self.stopped
202 }
203
204 fn stopped_by(&mut self, by: Stop) -> Interrupted {
205 let stop = self.stop(by);
206 *self.stopped.get_or_insert(stop)
207 }
208
209 pub fn steps(&self) -> u64 {
211 self.steps
212 }
213
214 #[must_use]
217 pub fn split(&self) -> Meter<'a> {
218 Meter {
219 poll: self.poll,
220 cap: self.cap.map(|cap| cap.saturating_sub(self.steps)),
221 steps: 0,
222 stopped: None,
223 }
224 }
225
226 pub fn charge(&mut self, steps: u64) -> Result<(), Interrupted> {
232 let total = self.steps.saturating_add(steps);
233 if let Some(cap) = self.cap
234 && total > cap
235 {
236 self.steps = cap;
237 return Err(self.stop(Stop::Budget));
238 }
239 self.steps = total;
240 Ok(())
241 }
242
243 pub fn charge_stop(&mut self, item: Interrupted) -> Interrupted {
249 self.steps = self.steps.saturating_add(item.steps);
250 if let (Stop::Budget, Some(cap)) = (item.by, self.cap) {
251 self.steps = self.steps.min(cap);
255 }
256 self.stop(item.by)
257 }
258
259 fn stop(&self, by: Stop) -> Interrupted {
260 Interrupted {
261 by,
262 steps: self.steps,
263 }
264 }
265}
266
267#[cfg(test)]
268mod tests {
269 use super::*;
270 use core::sync::atomic::{AtomicU64, Ordering};
271
272 #[test]
273 fn none_never_stops() {
274 let mut m = Meter::new(&Control::NONE);
275 for _ in 0..1000 {
276 assert!(m.tick().is_ok());
277 }
278 assert_eq!(m.steps(), 1000);
279 }
280
281 #[test]
282 fn a_meter_remembers_its_first_stop() {
283 let mut m = Meter::new(&Control::budget(1));
284 assert_eq!(m.stopped(), None);
285 m.tick().unwrap();
286 assert_eq!(m.stopped(), None);
287 let e = m.tick().unwrap_err();
288 assert_eq!(m.stopped(), Some(e));
289 m.tick().unwrap_err();
290 assert_eq!(m.stopped(), Some(e));
291 }
292
293 #[test]
294 fn a_budget_of_n_allows_n_steps() {
295 let mut m = Meter::new(&Control::budget(3));
296 assert!(m.tick().is_ok() && m.tick().is_ok() && m.tick().is_ok());
297 let e = m.tick().unwrap_err();
298 assert_eq!(
299 e,
300 Interrupted {
301 by: Stop::Budget,
302 steps: 3
303 }
304 );
305 assert_eq!(m.tick().unwrap_err(), e);
307 assert_eq!(Meter::new(&Control::budget(0)).tick().unwrap_err().steps, 0);
308 }
309
310 #[test]
311 fn a_poll_is_asked_at_every_step_and_stops_at_its_first_true() {
312 let calls = AtomicU64::new(0);
313 let poll = || calls.fetch_add(1, Ordering::Relaxed) + 1 == 4;
314 let mut m = Meter::new(&Control::poll(&poll));
315 assert!(m.tick().is_ok() && m.tick().is_ok() && m.tick().is_ok());
316 let e = m.tick().unwrap_err();
317 assert_eq!(
318 e,
319 Interrupted {
320 by: Stop::Poll,
321 steps: 3
322 }
323 );
324 }
325
326 #[test]
327 fn the_budget_is_asked_before_the_poll() {
328 let calls = AtomicU64::new(0);
329 let poll = || {
330 calls.fetch_add(1, Ordering::Relaxed);
331 false
332 };
333 let mut m = Meter::new(&Control::poll(&poll).with_budget(1));
334 m.tick().unwrap();
335 assert_eq!(m.tick().unwrap_err().by, Stop::Budget);
336 assert_eq!(calls.load(Ordering::Relaxed), 1);
337 }
338
339 #[test]
340 fn a_split_is_capped_at_what_is_left() {
341 let mut m = Meter::new(&Control::budget(5));
342 m.tick().unwrap();
343 m.tick().unwrap();
344 let mut item = m.split();
345 for _ in 0..3 {
346 item.tick().unwrap();
347 }
348 assert_eq!(item.tick().unwrap_err().steps, 3);
349 assert_eq!(Meter::new(&Control::NONE).split().cap, None);
350 }
351
352 #[test]
356 fn charging_in_order_stops_where_sequential_ticks_stop() {
357 let items = [2u64, 3, 1, 4];
358 for budget in 0..=12u64 {
359 let mut sequential = Meter::new(&Control::budget(budget));
360 let seq = items
361 .iter()
362 .enumerate()
363 .find_map(|(i, &n)| (0..n).find_map(|_| sequential.tick().err()).map(|e| (i, e)));
364 let mut joined = Meter::new(&Control::budget(budget));
365 let par = items.iter().enumerate().find_map(|(i, &n)| {
366 let mut item = joined.split();
367 let done = (0..n).try_for_each(|_| item.tick());
368 match done {
369 Ok(()) => joined.charge(item.steps()).err().map(|e| (i, e)),
370 Err(e) => Some((i, joined.charge_stop(e))),
371 }
372 });
373 assert_eq!(par, seq, "budget {budget}");
374 assert_eq!(joined.steps(), sequential.steps(), "budget {budget}");
375 }
376 }
377
378 #[test]
382 fn items_split_together_stop_where_sequential_ticks_stop() {
383 let items = [2u64, 3, 1, 4];
384 for budget in 0..=12u64 {
385 let mut sequential = Meter::new(&Control::budget(budget));
386 let seq = items
387 .iter()
388 .enumerate()
389 .find_map(|(i, &n)| (0..n).find_map(|_| sequential.tick().err()).map(|e| (i, e)));
390 let mut joined = Meter::new(&Control::budget(budget));
391 let base = joined.split();
392 let par = items.iter().enumerate().find_map(|(i, &n)| {
393 let mut item = base;
394 let done = (0..n).try_for_each(|_| item.tick());
395 match done {
396 Ok(()) => joined.charge(item.steps()).err().map(|e| (i, e)),
397 Err(e) => Some((i, joined.charge_stop(e))),
398 }
399 });
400 assert_eq!(par, seq, "budget {budget}");
401 assert_eq!(joined.steps(), sequential.steps(), "budget {budget}");
402 }
403 }
404}