Skip to main content

xet_data/file_reconstruction/data_writer/
unordered_writer.rs

1use std::sync::Arc;
2use std::sync::atomic::{AtomicU64, Ordering};
3
4use bytes::Bytes;
5use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
6use tokio::task::JoinSet;
7#[cfg(target_family = "wasm")]
8use tokio_with_wasm::alias as tokio;
9use xet_client::cas_types::FileRange;
10use xet_runtime::utils::adjustable_semaphore::AdjustableSemaphorePermit;
11
12use super::super::data_writer::{DataFuture, DataWriter};
13use super::super::run_state::RunState;
14use super::super::{FileReconstructionError, Result};
15
16/// A completed term ready for consumption. Contains the byte range indicating
17/// where this data belongs in the output file, the actual data bytes, and an
18/// optional semaphore permit for backpressure control.
19pub(crate) struct CompletedTerm {
20    pub byte_range: FileRange,
21    pub data: Bytes,
22    pub permit: Option<AdjustableSemaphorePermit>,
23}
24
25/// Atomic progress counters shared between the writer, its spawned tasks,
26/// and the consumer stream. Wrapped in an `Arc` so each party can read/update
27/// counters without holding a reference to the full `UnorderedWriter`.
28pub(crate) struct UnorderedWriterProgress {
29    pub terms_in_progress: AtomicU64,
30    pub bytes_in_progress: AtomicU64,
31}
32
33impl UnorderedWriterProgress {
34    pub fn terms_in_progress(&self) -> u64 {
35        self.terms_in_progress.load(Ordering::Acquire)
36    }
37
38    pub fn bytes_in_progress(&self) -> u64 {
39        self.bytes_in_progress.load(Ordering::Relaxed)
40    }
41}
42
43/// Writer that delivers completed data terms in arbitrary order.
44///
45/// Each call to [`set_next_term_data_source`](DataWriter::set_next_term_data_source)
46/// spawns a task (tracked via a [`JoinSet`]) that resolves the data future and
47/// sends the result through an [`mpsc`](tokio::sync::mpsc) channel. The consumer
48/// (typically an [`UnorderedDownloadStream`](super::unordered_download_stream::UnorderedDownloadStream))
49/// reads from the receiver end and gets items in whatever order tasks complete.
50///
51/// The consumer stream holds only `Arc<UnorderedWriterProgress>`, not the writer
52/// itself, so the writer's channel sender is dropped naturally when the
53/// reconstruction task finishes and consumes the writer via
54/// [`finish()`](DataWriter::finish).
55pub struct UnorderedWriter {
56    result_tx: UnboundedSender<Result<CompletedTerm>>,
57    run_state: Arc<RunState>,
58    progress: Arc<UnorderedWriterProgress>,
59    task_set: JoinSet<Result<u64>>,
60    total_bytes_sent: u64,
61    finished: bool,
62}
63
64impl Drop for UnorderedWriter {
65    fn drop(&mut self) {
66        if !self.finished {
67            self.run_state.cancel();
68        }
69    }
70}
71
72#[cfg_attr(not(target_family = "wasm"), async_trait::async_trait)]
73#[cfg_attr(target_family = "wasm", async_trait::async_trait(?Send))]
74impl DataWriter for UnorderedWriter {
75    async fn set_next_term_data_source(
76        &mut self,
77        byte_range: FileRange,
78        permit: Option<AdjustableSemaphorePermit>,
79        data_future: DataFuture,
80    ) -> Result<()> {
81        self.run_state.check_error()?;
82
83        while let Some(result) = self.task_set.try_join_next() {
84            self.total_bytes_sent +=
85                result.map_err(|e| FileReconstructionError::InternalError(format!("Task join error: {e}")))??;
86        }
87
88        if self.finished {
89            return Err(FileReconstructionError::InternalWriterError("Writer has already finished".to_string()));
90        }
91
92        let expected_size = byte_range.end - byte_range.start;
93        self.progress.terms_in_progress.fetch_add(1, Ordering::Relaxed);
94        self.progress.bytes_in_progress.fetch_add(expected_size, Ordering::Relaxed);
95
96        let result_tx = self.result_tx.clone();
97        let run_state = self.run_state.clone();
98        let progress = self.progress.clone();
99
100        self.task_set.spawn(async move {
101            let result = async {
102                run_state.check_error()?;
103
104                let data = data_future.await?;
105
106                if data.len() as u64 != expected_size {
107                    return Err(FileReconstructionError::InternalWriterError(format!(
108                        "Data size mismatch: expected {} bytes, got {} bytes",
109                        expected_size,
110                        data.len()
111                    )));
112                }
113
114                Ok(CompletedTerm {
115                    byte_range,
116                    data,
117                    permit,
118                })
119            }
120            .await;
121
122            if let Err(ref e) = result {
123                run_state.set_error(e.clone());
124            }
125
126            let completed_bytes = result.as_ref().map(|t| t.data.len() as u64).unwrap_or(0);
127
128            let _ = result_tx.send(result);
129
130            progress.bytes_in_progress.fetch_sub(expected_size, Ordering::Relaxed);
131            progress.terms_in_progress.fetch_sub(1, Ordering::Release);
132
133            if completed_bytes > 0 {
134                Ok(completed_bytes)
135            } else {
136                run_state.check_error()?;
137                Ok(0)
138            }
139        });
140
141        Ok(())
142    }
143
144    async fn finish(mut self: Box<Self>) -> Result<u64> {
145        self.run_state.check_error()?;
146
147        while let Some(result) = self.task_set.join_next().await {
148            self.total_bytes_sent +=
149                result.map_err(|e| FileReconstructionError::InternalError(format!("Task join error: {e}")))??;
150        }
151
152        self.finished = true;
153        Ok(self.total_bytes_sent)
154    }
155}
156
157impl UnorderedWriter {
158    /// Creates an unordered writer for streaming use. Returns the writer (to be
159    /// passed to the reconstruction task as `Box<dyn DataWriter>`), the receiver
160    /// end of the channel, and the shared progress counters for the consumer.
161    ///
162    /// The consumer stream should hold only the `Arc<UnorderedWriterProgress>`,
163    /// **not** the writer itself. This way the channel sender is dropped
164    /// naturally when the reconstruction task finishes (consuming the writer
165    /// via `finish()`), closing the channel without explicit lifetime management.
166    pub(crate) fn new_streaming(
167        run_state: Arc<RunState>,
168    ) -> (Box<dyn DataWriter>, UnboundedReceiver<Result<CompletedTerm>>, Arc<UnorderedWriterProgress>) {
169        let (tx, rx) = unbounded_channel();
170
171        let progress = Arc::new(UnorderedWriterProgress {
172            terms_in_progress: AtomicU64::new(0),
173            bytes_in_progress: AtomicU64::new(0),
174        });
175
176        let writer = Box::new(UnorderedWriter {
177            result_tx: tx,
178            run_state,
179            progress: progress.clone(),
180            task_set: JoinSet::new(),
181            total_bytes_sent: 0,
182            finished: false,
183        });
184
185        (writer, rx, progress)
186    }
187}
188
189#[cfg(test)]
190mod tests {
191    use std::time::Duration;
192
193    use xet_runtime::utils::adjustable_semaphore::AdjustableSemaphore;
194
195    use super::*;
196
197    fn immediate_future(data: Bytes) -> DataFuture {
198        Box::pin(async move { Ok(data) })
199    }
200
201    fn delayed_future(data: Bytes, delay: Duration) -> DataFuture {
202        Box::pin(async move {
203            tokio::time::sleep(delay).await;
204            Ok(data)
205        })
206    }
207
208    /// Drains all results from the receiver, returning data sorted by offset.
209    /// The writer must have been dropped (after calling `finish()`) so that
210    /// the channel closes naturally when all spawned tasks complete.
211    async fn drain_sorted(rx: &mut UnboundedReceiver<Result<CompletedTerm>>) -> Result<Vec<(u64, Bytes)>> {
212        let mut items = Vec::new();
213        while let Some(result) = rx.recv().await {
214            let term = result?;
215            items.push((term.byte_range.start, term.data));
216            drop(term.permit);
217        }
218        items.sort_by_key(|(offset, _)| *offset);
219        Ok(items)
220    }
221
222    #[tokio::test]
223    async fn test_basic_unordered_writes() {
224        let run_state = RunState::new_for_test();
225        let (mut writer, mut rx, _progress) = UnorderedWriter::new_streaming(run_state);
226
227        writer
228            .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
229            .await
230            .unwrap();
231        writer
232            .set_next_term_data_source(FileRange::new(5, 6), None, immediate_future(Bytes::from(" ")))
233            .await
234            .unwrap();
235        writer
236            .set_next_term_data_source(FileRange::new(6, 11), None, immediate_future(Bytes::from("World")))
237            .await
238            .unwrap();
239
240        let total = writer.finish().await.unwrap();
241        assert_eq!(total, 11);
242
243        let items = drain_sorted(&mut rx).await.unwrap();
244        let assembled: Vec<u8> = items.into_iter().flat_map(|(_, data)| data.to_vec()).collect();
245        assert_eq!(&assembled, b"Hello World");
246    }
247
248    #[tokio::test]
249    async fn test_delayed_futures_complete_out_of_order() {
250        let run_state = RunState::new_for_test();
251        let (mut writer, mut rx, _progress) = UnorderedWriter::new_streaming(run_state);
252
253        writer
254            .set_next_term_data_source(
255                FileRange::new(0, 5),
256                None,
257                delayed_future(Bytes::from("Hello"), Duration::from_millis(80)),
258            )
259            .await
260            .unwrap();
261        writer
262            .set_next_term_data_source(
263                FileRange::new(5, 6),
264                None,
265                delayed_future(Bytes::from(" "), Duration::from_millis(40)),
266            )
267            .await
268            .unwrap();
269        writer
270            .set_next_term_data_source(FileRange::new(6, 11), None, immediate_future(Bytes::from("World")))
271            .await
272            .unwrap();
273
274        let total = writer.finish().await.unwrap();
275        assert_eq!(total, 11);
276
277        let items = drain_sorted(&mut rx).await.unwrap();
278        let assembled: Vec<u8> = items.into_iter().flat_map(|(_, data)| data.to_vec()).collect();
279        assert_eq!(&assembled, b"Hello World");
280    }
281
282    #[tokio::test]
283    async fn test_size_mismatch_error() {
284        let run_state = RunState::new_for_test();
285        let (mut writer, mut rx, _progress) = UnorderedWriter::new_streaming(run_state);
286
287        writer
288            .set_next_term_data_source(FileRange::new(0, 10), None, immediate_future(Bytes::from("Hello")))
289            .await
290            .unwrap();
291
292        let result = writer.finish().await;
293        assert!(result.is_err());
294
295        let result = rx.recv().await.unwrap();
296        assert!(result.is_err());
297        assert!(matches!(result, Err(FileReconstructionError::InternalWriterError(_))));
298    }
299
300    #[tokio::test]
301    async fn test_future_error_propagates() {
302        let run_state = RunState::new_for_test();
303        let (mut writer, mut rx, _progress) = UnorderedWriter::new_streaming(run_state);
304
305        let failing_future: DataFuture =
306            Box::pin(async { Err(FileReconstructionError::InternalError("Simulated error".to_string())) });
307
308        writer
309            .set_next_term_data_source(FileRange::new(0, 5), None, failing_future)
310            .await
311            .unwrap();
312
313        let result = writer.finish().await;
314        assert!(result.is_err());
315
316        let result = rx.recv().await.unwrap();
317        assert!(result.is_err());
318    }
319
320    #[tokio::test]
321    async fn test_semaphore_permit_released_after_consumption() {
322        let run_state = RunState::new_for_test();
323        let (mut writer, mut rx, _progress) = UnorderedWriter::new_streaming(run_state);
324        let semaphore = AdjustableSemaphore::new(2, (0, 2));
325
326        let permit1 = semaphore.acquire().await.unwrap();
327        let permit2 = semaphore.acquire().await.unwrap();
328        assert_eq!(semaphore.available_permits(), 0);
329
330        writer
331            .set_next_term_data_source(FileRange::new(0, 5), Some(permit1), immediate_future(Bytes::from("Hello")))
332            .await
333            .unwrap();
334        writer
335            .set_next_term_data_source(FileRange::new(5, 6), Some(permit2), immediate_future(Bytes::from(" ")))
336            .await
337            .unwrap();
338
339        writer.finish().await.unwrap();
340
341        let items = drain_sorted(&mut rx).await.unwrap();
342        drop(items);
343
344        assert_eq!(semaphore.available_permits(), 2);
345    }
346
347    #[tokio::test]
348    async fn test_counter_accuracy() {
349        let run_state = RunState::new_for_test();
350        let (mut writer, mut rx, progress) = UnorderedWriter::new_streaming(run_state);
351
352        writer
353            .set_next_term_data_source(
354                FileRange::new(0, 5),
355                None,
356                delayed_future(Bytes::from("Hello"), Duration::from_millis(50)),
357            )
358            .await
359            .unwrap();
360        writer
361            .set_next_term_data_source(
362                FileRange::new(5, 11),
363                None,
364                delayed_future(Bytes::from(" World"), Duration::from_millis(50)),
365            )
366            .await
367            .unwrap();
368
369        let total = writer.finish().await.unwrap();
370        assert_eq!(total, 11);
371
372        let _items = drain_sorted(&mut rx).await.unwrap();
373
374        assert_eq!(progress.bytes_in_progress(), 0);
375        assert_eq!(progress.terms_in_progress(), 0);
376    }
377
378    #[tokio::test]
379    async fn test_finish_returns_total_bytes() {
380        let run_state = RunState::new_for_test();
381        let (mut writer, mut rx, _progress) = UnorderedWriter::new_streaming(run_state);
382
383        writer
384            .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
385            .await
386            .unwrap();
387        writer
388            .set_next_term_data_source(FileRange::new(5, 11), None, immediate_future(Bytes::from(" World")))
389            .await
390            .unwrap();
391
392        let total = writer.finish().await.unwrap();
393        assert_eq!(total, 11);
394
395        let _items = drain_sorted(&mut rx).await.unwrap();
396    }
397
398    #[tokio::test]
399    async fn test_error_propagation_prevents_subsequent_writes() {
400        let run_state = RunState::new_for_test();
401        let (mut writer, mut _rx, _progress) = UnorderedWriter::new_streaming(run_state.clone());
402
403        let failing_future: DataFuture =
404            Box::pin(async { Err(FileReconstructionError::InternalError("fail".to_string())) });
405
406        writer
407            .set_next_term_data_source(FileRange::new(0, 5), None, failing_future)
408            .await
409            .unwrap();
410
411        let wait_for_error = tokio::time::timeout(Duration::from_secs(1), async {
412            loop {
413                if run_state.check_error().is_err() {
414                    break;
415                }
416                tokio::task::yield_now().await;
417            }
418        })
419        .await;
420        assert!(wait_for_error.is_ok());
421
422        let result = writer
423            .set_next_term_data_source(FileRange::new(5, 10), None, immediate_future(Bytes::from("World")))
424            .await;
425        assert!(result.is_err());
426    }
427
428    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
429    async fn stress_test_many_concurrent_terms() {
430        let run_state = RunState::new_for_test();
431        let (mut writer, mut rx, _progress) = UnorderedWriter::new_streaming(run_state);
432
433        let num_terms: usize = 100;
434        let mut expected: Vec<(u64, Vec<u8>)> = Vec::new();
435        let mut offset = 0u64;
436
437        for i in 0..num_terms {
438            let size = 100 + (i % 50) * 10;
439            let data: Vec<u8> = (0..size).map(|j| ((i * 7 + j * 13) % 256) as u8).collect();
440            let bytes = Bytes::from(data.clone());
441            expected.push((offset, data));
442
443            let delay = Duration::from_micros((i % 10) as u64 * 100);
444            writer
445                .set_next_term_data_source(
446                    FileRange::new(offset, offset + size as u64),
447                    None,
448                    delayed_future(bytes, delay),
449                )
450                .await
451                .unwrap();
452
453            offset += size as u64;
454        }
455
456        let total = writer.finish().await.unwrap();
457        assert_eq!(total, offset);
458
459        let items = drain_sorted(&mut rx).await.unwrap();
460        assert_eq!(items.len(), num_terms);
461
462        for ((exp_offset, exp_data), (act_offset, act_data)) in expected.iter().zip(items.iter()) {
463            assert_eq!(*exp_offset, *act_offset);
464            assert_eq!(exp_data.as_slice(), act_data.as_ref());
465        }
466    }
467
468    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
469    async fn stress_test_rapid_finish_after_writes() {
470        for _ in 0..50 {
471            let run_state = RunState::new_for_test();
472            let (mut writer, mut rx, _progress) = UnorderedWriter::new_streaming(run_state);
473
474            for i in 0..10u64 {
475                let data = Bytes::from(vec![i as u8; 100]);
476                writer
477                    .set_next_term_data_source(FileRange::new(i * 100, (i + 1) * 100), None, immediate_future(data))
478                    .await
479                    .unwrap();
480            }
481
482            let total = writer.finish().await.unwrap();
483            assert_eq!(total, 1000);
484
485            let items = drain_sorted(&mut rx).await.unwrap();
486            assert_eq!(items.len(), 10);
487
488            let total_bytes: usize = items.iter().map(|(_, data)| data.len()).sum();
489            assert_eq!(total_bytes, 1000);
490        }
491    }
492
493    #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
494    async fn stress_test_mixed_immediate_and_delayed() {
495        for _ in 0..20 {
496            let run_state = RunState::new_for_test();
497            let (mut writer, mut rx, progress) = UnorderedWriter::new_streaming(run_state);
498
499            let mut offset = 0u64;
500            let mut total_size = 0u64;
501            let num_terms = 30usize;
502
503            for i in 0..num_terms {
504                let size = ((i + 1) * 50) as u64;
505                let data = Bytes::from(vec![(i % 256) as u8; size as usize]);
506                total_size += size;
507
508                let future = if i % 3 == 0 {
509                    delayed_future(data, Duration::from_millis((i % 5) as u64))
510                } else {
511                    immediate_future(data)
512                };
513
514                writer
515                    .set_next_term_data_source(FileRange::new(offset, offset + size), None, future)
516                    .await
517                    .unwrap();
518                offset += size;
519            }
520
521            let total = writer.finish().await.unwrap();
522            assert_eq!(total, total_size);
523
524            let items = drain_sorted(&mut rx).await.unwrap();
525            assert_eq!(items.len(), num_terms);
526
527            let received_bytes: u64 = items.iter().map(|(_, data)| data.len() as u64).sum();
528            assert_eq!(received_bytes, total_size);
529            assert_eq!(progress.terms_in_progress(), 0);
530        }
531    }
532}