diskann_disk/build/chunking/continuation/
utils.rs1use std::{error::Error, thread::sleep};
7
8use tracing::info;
9
10use super::continuation_tracker::{ContinuationGrant, ContinuationTrackerTrait};
11use crate::build::chunking::checkpoint::Progress;
12
13pub fn process_while_resource_is_available<Action, ParamIter, Param, E>(
20 mut action: Action,
21 params: ParamIter,
22 continuation_checker: Box<dyn ContinuationTrackerTrait>,
23) -> Result<Progress, E>
24where
25 ParamIter: Iterator<Item = Param>,
26 Action: FnMut(Param) -> Result<(), E>,
27 E: Error,
28{
29 for (idx, param) in params.enumerate() {
30 loop {
31 match continuation_checker.get_continuation_grant() {
32 ContinuationGrant::Continue => {
33 info!("Continue processing.");
34 action(param)?;
35 break;
36 }
37 ContinuationGrant::Yield(duration) => {
38 info!(
39 "Continuation checker asks to yield for {} ms.",
40 duration.as_millis()
41 );
42 sleep(duration);
43 }
44 ContinuationGrant::Stop => {
45 info!("Continuation checker asks to stop. Breaking the loop.");
46 return Ok(Progress::Processed(idx));
47 }
48 }
49 }
50 }
51
52 Ok(Progress::Completed)
53}
54
55pub async fn process_while_resource_is_available_async<Action, ParamIter, Param, Fut, E>(
60 mut action: Action,
61 params: ParamIter,
62 continuation_checker: Box<dyn ContinuationTrackerTrait>,
63) -> Result<Progress, E>
64where
65 ParamIter: Iterator<Item = Param>,
66 Action: FnMut(Param) -> Fut,
67 Fut: core::future::Future<Output = Result<(), E>>,
68 E: Error,
69{
70 for (idx, param) in params.enumerate() {
71 loop {
72 match continuation_checker.get_continuation_grant() {
73 ContinuationGrant::Continue => {
74 info!("Continue processing.");
75 action(param).await?;
76 break;
77 }
78 ContinuationGrant::Yield(duration) => {
79 info!(
80 "Continuation checker asks to yield for {} ms.",
81 duration.as_millis()
82 );
83 sleep(duration);
84 }
85 ContinuationGrant::Stop => {
86 info!("Continuation checker asks to stop. Breaking the loop.");
87 return Ok(Progress::Processed(idx));
88 }
89 }
90 }
91 }
92
93 Ok(Progress::Completed)
94}
95
96#[cfg(test)]
97mod tests {
98 use super::super::continuation_tracker::NaiveContinuationTracker;
99 use super::*;
100 use std::fmt;
101
102 #[derive(Debug)]
103 struct TestError;
104
105 impl fmt::Display for TestError {
106 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
107 write!(f, "TestError")
108 }
109 }
110
111 impl Error for TestError {}
112
113 #[test]
114 fn test_process_while_resource_is_available_completes() {
115 let checker = Box::new(NaiveContinuationTracker::default());
116 let items = vec![1, 2, 3, 4, 5];
117 let mut processed = Vec::new();
118
119 let result = process_while_resource_is_available(
120 |item| {
121 processed.push(item);
122 Ok::<(), TestError>(())
123 },
124 items.into_iter(),
125 checker,
126 );
127
128 assert!(matches!(result.unwrap(), Progress::Completed));
129 assert_eq!(processed, vec![1, 2, 3, 4, 5]);
130 }
131
132 #[test]
133 fn test_process_while_resource_is_available_empty_iter() {
134 let checker = Box::new(NaiveContinuationTracker::default());
135 let items: Vec<i32> = vec![];
136
137 let result = process_while_resource_is_available(
138 |_item| Ok::<(), TestError>(()),
139 items.into_iter(),
140 checker,
141 );
142
143 assert!(matches!(result.unwrap(), Progress::Completed));
144 }
145
146 #[derive(Clone)]
148 struct StopAfterTracker {
149 count: std::sync::Arc<std::sync::Mutex<usize>>,
150 stop_after: usize,
151 }
152
153 impl ContinuationTrackerTrait for StopAfterTracker {
154 fn get_continuation_grant(&self) -> ContinuationGrant {
155 let mut count = self.count.lock().unwrap();
156 if *count >= self.stop_after {
157 ContinuationGrant::Stop
158 } else {
159 *count += 1;
160 ContinuationGrant::Continue
161 }
162 }
163 }
164
165 #[test]
166 fn test_process_while_resource_is_available_stops_early() {
167 let tracker = StopAfterTracker {
168 count: std::sync::Arc::new(std::sync::Mutex::new(0)),
169 stop_after: 3,
170 };
171 let items = vec![10, 20, 30, 40, 50];
172 let mut processed = Vec::new();
173
174 let result = process_while_resource_is_available(
175 |item| {
176 processed.push(item);
177 Ok::<(), TestError>(())
178 },
179 items.into_iter(),
180 Box::new(tracker),
181 );
182
183 assert!(matches!(
185 result.unwrap(),
186 Progress::Processed(processed_count) if processed_count == 3
187 ));
188 assert_eq!(processed, vec![10, 20, 30]);
189 }
190
191 #[derive(Clone)]
193 struct YieldOnceThenContinueTracker {
194 yielded: std::sync::Arc<std::sync::Mutex<bool>>,
195 }
196
197 impl ContinuationTrackerTrait for YieldOnceThenContinueTracker {
198 fn get_continuation_grant(&self) -> ContinuationGrant {
199 let mut yielded = self.yielded.lock().unwrap();
200 if !*yielded {
201 *yielded = true;
202 ContinuationGrant::Yield(std::time::Duration::ZERO)
203 } else {
204 ContinuationGrant::Continue
206 }
207 }
208 }
209
210 #[test]
211 fn test_process_while_resource_is_available_yield_then_continue() {
212 let tracker = YieldOnceThenContinueTracker {
213 yielded: std::sync::Arc::new(std::sync::Mutex::new(false)),
214 };
215 let items = vec![1, 2];
216 let mut processed = Vec::new();
217
218 let result = process_while_resource_is_available(
219 |item| {
220 processed.push(item);
221 Ok::<(), TestError>(())
222 },
223 items.into_iter(),
224 Box::new(tracker),
225 );
226
227 assert!(matches!(result.unwrap(), Progress::Completed));
229 assert_eq!(processed, vec![1, 2]);
230 }
231
232 #[test]
233 fn test_process_while_resource_is_available_action_error() {
234 let checker = Box::new(NaiveContinuationTracker::default());
235 let items = vec![1, 2, 3];
236
237 let result = process_while_resource_is_available(
238 |item| {
239 if item == 2 {
240 Err(TestError)
241 } else {
242 Ok(())
243 }
244 },
245 items.into_iter(),
246 checker,
247 );
248
249 assert!(result.is_err());
250 }
251
252 #[tokio::test]
253 async fn test_process_while_resource_is_available_async_stops_early() {
254 let tracker = StopAfterTracker {
255 count: std::sync::Arc::new(std::sync::Mutex::new(0)),
256 stop_after: 2,
257 };
258 let items = vec![1, 2, 3, 4, 5];
259 let processed = std::sync::Arc::new(tokio::sync::Mutex::new(Vec::new()));
260
261 let result = process_while_resource_is_available_async(
262 |item| {
263 let processed = processed.clone();
264 async move {
265 processed.lock().await.push(item);
266 Ok::<(), TestError>(())
267 }
268 },
269 items.into_iter(),
270 Box::new(tracker),
271 )
272 .await;
273
274 assert!(matches!(
275 result.unwrap(),
276 Progress::Processed(processed_count) if processed_count == 2
277 ));
278 let processed = processed.lock().await;
279 assert_eq!(*processed, vec![1, 2]);
280 }
281
282 #[tokio::test]
283 async fn test_process_while_resource_is_available_async_completes() {
284 let checker = Box::new(NaiveContinuationTracker::default());
285 let items = vec![1, 2, 3];
286 let processed = std::sync::Arc::new(tokio::sync::Mutex::new(Vec::new()));
287
288 let result = process_while_resource_is_available_async(
289 |item| {
290 let processed = processed.clone();
291 async move {
292 processed.lock().await.push(item);
293 Ok::<(), TestError>(())
294 }
295 },
296 items.into_iter(),
297 checker,
298 )
299 .await;
300
301 assert!(matches!(result.unwrap(), Progress::Completed));
302 let processed = processed.lock().await;
303 assert_eq!(*processed, vec![1, 2, 3]);
304 }
305}