Skip to main content

ic_testkit/pic/
time.rs

1use pocket_ic::PocketIc;
2use std::{fmt, task::Poll, time::Duration};
3
4/// A caller-owned readiness check failed or exhausted its simulated round budget.
5#[derive(Clone, Debug, Eq, PartialEq)]
6pub enum TickUntilError<E> {
7    /// The predicate was still pending after this many advance/tick rounds.
8    ProgressLimit { rounds: u32 },
9    /// The predicate returned its own failure.
10    Failed(E),
11}
12
13impl<E: fmt::Display> fmt::Display for TickUntilError<E> {
14    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
15        match self {
16            Self::ProgressLimit { rounds } => {
17                write!(
18                    formatter,
19                    "readiness remained pending after {rounds} rounds"
20                )
21            }
22            Self::Failed(error) => write!(formatter, "readiness check failed: {error}"),
23        }
24    }
25}
26
27impl<E: std::error::Error + 'static> std::error::Error for TickUntilError<E> {
28    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
29        match self {
30            Self::ProgressLimit { .. } => None,
31            Self::Failed(error) => Some(error),
32        }
33    }
34}
35
36/// Check a caller-owned predicate, advancing simulated time and ticking while pending.
37///
38/// The predicate runs once before any mutation. Each pending result permits one
39/// round of `advance_time(advance)`, then `tick()`, then another predicate check,
40/// up to `max_rounds`. A zero budget still performs the initial check. A zero
41/// duration permits ticks without advancing time. Completion on the last allowed
42/// round succeeds; predicate failures return immediately without another round.
43///
44/// This bounds simulated progression, not wall-clock time spent inside PocketIC
45/// calls or the predicate. PocketIC operation failures retain upstream panic
46/// behavior. Applications own readiness meaning and predicate error types.
47pub fn tick_until<T, E>(
48    pic: &PocketIc,
49    max_rounds: u32,
50    advance: Duration,
51    mut poll: impl FnMut(&PocketIc) -> Poll<Result<T, E>>,
52) -> Result<T, TickUntilError<E>> {
53    poll_with_rounds(
54        max_rounds,
55        || {
56            pic.advance_time(advance);
57            pic.tick();
58        },
59        || poll(pic),
60    )
61}
62
63fn poll_with_rounds<T, E>(
64    max_rounds: u32,
65    mut round: impl FnMut(),
66    mut poll: impl FnMut() -> Poll<Result<T, E>>,
67) -> Result<T, TickUntilError<E>> {
68    let mut rounds = 0;
69    loop {
70        match poll() {
71            Poll::Ready(result) => return result.map_err(TickUntilError::Failed),
72            Poll::Pending if rounds == max_rounds => {
73                return Err(TickUntilError::ProgressLimit { rounds });
74            }
75            Poll::Pending => {
76                round();
77                rounds += 1;
78            }
79        }
80    }
81}
82
83/// Focused time conversion missing from PocketIC's native API.
84///
85/// All mutation, certified-time, and round operations stay on [`PocketIc`].
86pub trait PocketIcTimeExt {
87    /// Read PocketIC wall-clock time as nanoseconds since the Unix epoch.
88    fn current_time_nanos(&self) -> u64;
89}
90
91impl PocketIcTimeExt for PocketIc {
92    fn current_time_nanos(&self) -> u64 {
93        self.get_time().as_nanos_since_unix_epoch()
94    }
95}
96
97#[cfg(test)]
98mod tests {
99    use super::*;
100    use std::cell::Cell;
101
102    #[test]
103    fn initial_completion_and_failure_do_not_advance() {
104        for budget in [0, 3] {
105            for result in [Ok(42), Err("predicate failure")] {
106                assert_eq!(
107                    poll_with_rounds(
108                        budget,
109                        || panic!("unexpected round"),
110                        || Poll::Ready(result)
111                    ),
112                    result.map_err(TickUntilError::Failed)
113                );
114            }
115        }
116    }
117
118    #[test]
119    fn completion_on_last_round_preserves_poll_order() {
120        let rounds = Cell::new(0);
121        let polls = Cell::new(0);
122        let result = poll_with_rounds(
123            3,
124            || {
125                assert_eq!(polls.get(), rounds.get() + 1);
126                rounds.set(rounds.get() + 1);
127            },
128            || {
129                assert_eq!(polls.get(), rounds.get());
130                polls.set(polls.get() + 1);
131                if rounds.get() == 3 {
132                    Poll::Ready(Ok::<_, ()>(42))
133                } else {
134                    Poll::Pending
135                }
136            },
137        );
138        assert_eq!(result, Ok(42));
139        assert_eq!((rounds.get(), polls.get()), (3, 4));
140    }
141
142    #[test]
143    fn exhaustion_performs_exactly_the_allowed_rounds() {
144        for budget in [0, 1, 3] {
145            let rounds = Cell::new(0);
146            let polls = Cell::new(0);
147            let result = poll_with_rounds(
148                budget,
149                || rounds.set(rounds.get() + 1),
150                || {
151                    polls.set(polls.get() + 1);
152                    Poll::<Result<(), ()>>::Pending
153                },
154            );
155            assert_eq!(
156                result,
157                Err(TickUntilError::ProgressLimit { rounds: budget })
158            );
159            assert_eq!((rounds.get(), polls.get()), (budget, budget + 1));
160        }
161    }
162
163    #[test]
164    fn failure_after_progress_preserves_the_original_cause() {
165        let rounds = Cell::new(0);
166        let result = poll_with_rounds(
167            3,
168            || rounds.set(rounds.get() + 1),
169            || {
170                if rounds.get() == 1 {
171                    Poll::Ready(Err::<(), _>(std::io::Error::from(
172                        std::io::ErrorKind::PermissionDenied,
173                    )))
174                } else {
175                    Poll::Pending
176                }
177            },
178        );
179        let error = result.unwrap_err();
180        let source = std::error::Error::source(&error).unwrap();
181        assert_eq!(
182            source.downcast_ref::<std::io::Error>().unwrap().kind(),
183            std::io::ErrorKind::PermissionDenied
184        );
185        assert_eq!(rounds.get(), 1);
186    }
187}