use core::fmt;
pub type Poll<'a> = &'a (dyn Fn() -> bool + Sync);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Stop {
Poll,
Budget,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Interrupted {
pub by: Stop,
pub steps: u64,
}
impl fmt::Display for Interrupted {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.by {
Stop::Poll => write!(
f,
"interrupted by the caller's poll after {} steps",
self.steps
),
Stop::Budget => write!(
f,
"interrupted: the budget of {} steps is spent",
self.steps
),
}
}
}
impl core::error::Error for Interrupted {}
#[derive(Clone, Copy, Default)]
pub struct Control<'a> {
poll: Option<Poll<'a>>,
budget: Option<u64>,
}
impl fmt::Debug for Control<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Control")
.field("poll", &self.poll.map(|_| "…"))
.field("budget", &self.budget)
.finish()
}
}
impl Control<'static> {
pub const NONE: Control<'static> = Control {
poll: None,
budget: None,
};
pub const fn budget(steps: u64) -> Control<'static> {
Control {
poll: None,
budget: Some(steps),
}
}
}
impl<'a> Control<'a> {
pub const fn poll(poll: Poll<'a>) -> Control<'a> {
Control {
poll: Some(poll),
budget: None,
}
}
#[must_use]
pub const fn with_budget(mut self, steps: u64) -> Self {
self.budget = Some(steps);
self
}
#[must_use]
pub const fn with_poll(mut self, poll: Poll<'a>) -> Self {
self.poll = Some(poll);
self
}
}
#[derive(Clone, Copy)]
pub struct Meter<'a> {
poll: Option<Poll<'a>>,
cap: Option<u64>,
steps: u64,
stopped: Option<Interrupted>,
}
impl fmt::Debug for Meter<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Meter")
.field("cap", &self.cap)
.field("steps", &self.steps)
.finish()
}
}
impl Default for Meter<'_> {
fn default() -> Self {
Meter::new(&Control::NONE)
}
}
impl<'a> Meter<'a> {
pub fn new(control: &Control<'a>) -> Self {
Meter {
poll: control.poll,
cap: control.budget,
steps: 0,
stopped: None,
}
}
pub fn tick(&mut self) -> Result<(), Interrupted> {
if self.cap.is_some_and(|cap| self.steps >= cap) {
return Err(self.stopped_by(Stop::Budget));
}
if self.poll.is_some_and(|poll| poll()) {
return Err(self.stopped_by(Stop::Poll));
}
self.steps += 1;
Ok(())
}
pub fn stopped(&self) -> Option<Interrupted> {
self.stopped
}
fn stopped_by(&mut self, by: Stop) -> Interrupted {
let stop = self.stop(by);
*self.stopped.get_or_insert(stop)
}
pub fn steps(&self) -> u64 {
self.steps
}
#[must_use]
pub fn split(&self) -> Meter<'a> {
Meter {
poll: self.poll,
cap: self.cap.map(|cap| cap.saturating_sub(self.steps)),
steps: 0,
stopped: None,
}
}
pub fn charge(&mut self, steps: u64) -> Result<(), Interrupted> {
let total = self.steps.saturating_add(steps);
if let Some(cap) = self.cap
&& total > cap
{
self.steps = cap;
return Err(self.stop(Stop::Budget));
}
self.steps = total;
Ok(())
}
pub fn charge_stop(&mut self, item: Interrupted) -> Interrupted {
self.steps = self.steps.saturating_add(item.steps);
if let (Stop::Budget, Some(cap)) = (item.by, self.cap) {
self.steps = self.steps.min(cap);
}
self.stop(item.by)
}
fn stop(&self, by: Stop) -> Interrupted {
Interrupted {
by,
steps: self.steps,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use core::sync::atomic::{AtomicU64, Ordering};
#[test]
fn none_never_stops() {
let mut m = Meter::new(&Control::NONE);
for _ in 0..1000 {
assert!(m.tick().is_ok());
}
assert_eq!(m.steps(), 1000);
}
#[test]
fn a_meter_remembers_its_first_stop() {
let mut m = Meter::new(&Control::budget(1));
assert_eq!(m.stopped(), None);
m.tick().unwrap();
assert_eq!(m.stopped(), None);
let e = m.tick().unwrap_err();
assert_eq!(m.stopped(), Some(e));
m.tick().unwrap_err();
assert_eq!(m.stopped(), Some(e));
}
#[test]
fn a_budget_of_n_allows_n_steps() {
let mut m = Meter::new(&Control::budget(3));
assert!(m.tick().is_ok() && m.tick().is_ok() && m.tick().is_ok());
let e = m.tick().unwrap_err();
assert_eq!(
e,
Interrupted {
by: Stop::Budget,
steps: 3
}
);
assert_eq!(m.tick().unwrap_err(), e);
assert_eq!(Meter::new(&Control::budget(0)).tick().unwrap_err().steps, 0);
}
#[test]
fn a_poll_is_asked_at_every_step_and_stops_at_its_first_true() {
let calls = AtomicU64::new(0);
let poll = || calls.fetch_add(1, Ordering::Relaxed) + 1 == 4;
let mut m = Meter::new(&Control::poll(&poll));
assert!(m.tick().is_ok() && m.tick().is_ok() && m.tick().is_ok());
let e = m.tick().unwrap_err();
assert_eq!(
e,
Interrupted {
by: Stop::Poll,
steps: 3
}
);
}
#[test]
fn the_budget_is_asked_before_the_poll() {
let calls = AtomicU64::new(0);
let poll = || {
calls.fetch_add(1, Ordering::Relaxed);
false
};
let mut m = Meter::new(&Control::poll(&poll).with_budget(1));
m.tick().unwrap();
assert_eq!(m.tick().unwrap_err().by, Stop::Budget);
assert_eq!(calls.load(Ordering::Relaxed), 1);
}
#[test]
fn a_split_is_capped_at_what_is_left() {
let mut m = Meter::new(&Control::budget(5));
m.tick().unwrap();
m.tick().unwrap();
let mut item = m.split();
for _ in 0..3 {
item.tick().unwrap();
}
assert_eq!(item.tick().unwrap_err().steps, 3);
assert_eq!(Meter::new(&Control::NONE).split().cap, None);
}
#[test]
fn charging_in_order_stops_where_sequential_ticks_stop() {
let items = [2u64, 3, 1, 4];
for budget in 0..=12u64 {
let mut sequential = Meter::new(&Control::budget(budget));
let seq = items
.iter()
.enumerate()
.find_map(|(i, &n)| (0..n).find_map(|_| sequential.tick().err()).map(|e| (i, e)));
let mut joined = Meter::new(&Control::budget(budget));
let par = items.iter().enumerate().find_map(|(i, &n)| {
let mut item = joined.split();
let done = (0..n).try_for_each(|_| item.tick());
match done {
Ok(()) => joined.charge(item.steps()).err().map(|e| (i, e)),
Err(e) => Some((i, joined.charge_stop(e))),
}
});
assert_eq!(par, seq, "budget {budget}");
assert_eq!(joined.steps(), sequential.steps(), "budget {budget}");
}
}
#[test]
fn items_split_together_stop_where_sequential_ticks_stop() {
let items = [2u64, 3, 1, 4];
for budget in 0..=12u64 {
let mut sequential = Meter::new(&Control::budget(budget));
let seq = items
.iter()
.enumerate()
.find_map(|(i, &n)| (0..n).find_map(|_| sequential.tick().err()).map(|e| (i, e)));
let mut joined = Meter::new(&Control::budget(budget));
let base = joined.split();
let par = items.iter().enumerate().find_map(|(i, &n)| {
let mut item = base;
let done = (0..n).try_for_each(|_| item.tick());
match done {
Ok(()) => joined.charge(item.steps()).err().map(|e| (i, e)),
Err(e) => Some((i, joined.charge_stop(e))),
}
});
assert_eq!(par, seq, "budget {budget}");
assert_eq!(joined.steps(), sequential.steps(), "budget {budget}");
}
}
}