1use pocket_ic::PocketIc;
2use std::{fmt, task::Poll, time::Duration};
3
4#[derive(Clone, Debug, Eq, PartialEq)]
6pub enum TickUntilError<E> {
7 ProgressLimit { rounds: u32 },
9 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
36pub 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
83pub trait PocketIcTimeExt {
87 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}