Skip to main content

diskann_disk/build/chunking/continuation/
utils.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6use std::{error::Error, thread::sleep};
7
8use tracing::info;
9
10use super::continuation_tracker::{ContinuationGrant, ContinuationTrackerTrait};
11use crate::build::chunking::checkpoint::Progress;
12
13/// This takes an operation with an iterator of oprands,
14/// and processes the oprands using the operation in a loop,
15/// until the continuation_checker asks it to stop.
16/// The continuation_checker is used to get continuation grants between processing each operation.
17/// The clean_up function is called after the loop is broken and before exit.
18/// The function returns a Progress enum, which indicates the number of operations executed.
19pub 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
55/// Asynchronous version of [`process_while_resource_is_available`].
56///
57/// Takes an async operation with an iterator of operands and processes them in a loop
58/// until the continuation_checker signals to stop.
59pub 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    /// A tracker that returns Stop after `stop_after` Continue grants.
147    #[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        // `Processed(n)` reports the number of items processed before the stop grant.
184        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    /// A tracker that yields once (with a tiny duration), then continues.
192    #[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                // After yielding once, always continue
205                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        // After yielding, it should have continued and processed all items
228        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}