1use std::collections::VecDeque;
2use std::io::{IoSlice, Write};
3use std::sync::Arc;
4use std::sync::atomic::{AtomicU64, Ordering};
5
6use bytes::Bytes;
7use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
8use tokio::sync::oneshot;
9use tokio::task::{JoinHandle, JoinSet};
10#[cfg(target_family = "wasm")]
11use tokio_with_wasm::alias as tokio;
12use xet_client::cas_types::FileRange;
13use xet_runtime::core::XetContext;
14use xet_runtime::utils::adjustable_semaphore::AdjustableSemaphorePermit;
15
16use super::super::data_writer::{DataFuture, DataWriter};
17use super::super::run_state::RunState;
18use super::super::{FileReconstructionError, Result};
19use crate::progress_tracking::ItemProgressUpdater;
20
21const WRITEV_MAX_SLICE: usize = 24;
31
32pub(crate) enum SequentialRetrievalItem {
36 Data {
37 receiver: oneshot::Receiver<Bytes>,
38 permit: Option<AdjustableSemaphorePermit>,
39 },
40 Finish,
41}
42
43#[cfg(not(target_family = "wasm"))]
45type PendingWrite = (Bytes, Option<AdjustableSemaphorePermit>);
46
47#[cfg(not(target_family = "wasm"))]
50struct SyncWriterThread {
51 ctx: XetContext,
52 rx: UnboundedReceiver<SequentialRetrievalItem>,
53 bytes_written: Arc<AtomicU64>,
54 progress_updater: Option<Arc<ItemProgressUpdater>>,
55 run_state: Arc<RunState>,
56 pending: Option<SequentialRetrievalItem>,
57 finished: bool,
58}
59
60#[cfg(not(target_family = "wasm"))]
61impl SyncWriterThread {
62 fn new(
63 ctx: XetContext,
64 rx: UnboundedReceiver<SequentialRetrievalItem>,
65 bytes_written: Arc<AtomicU64>,
66 progress_updater: Option<Arc<ItemProgressUpdater>>,
67 run_state: Arc<RunState>,
68 ) -> Self {
69 Self {
70 ctx,
71 rx,
72 bytes_written,
73 progress_updater,
74 run_state,
75 pending: None,
76 finished: false,
77 }
78 }
79
80 #[inline]
87 fn next_write(&mut self, should_block: bool) -> Result<Option<PendingWrite>> {
88 if self.pending.is_none() {
90 self.pending = if should_block {
92 self.rx.blocking_recv()
93 } else {
94 self.rx.try_recv().ok()
95 };
96 }
97
98 match self.pending.take() {
100 Some(SequentialRetrievalItem::Data { mut receiver, permit }) => {
101 if should_block {
102 let data = match receiver.blocking_recv() {
103 Ok(data) => data,
104 Err(_) => {
105 self.run_state.check_error()?;
106 return Err(FileReconstructionError::InternalWriterError(
107 "Data sender was dropped before sending data.".to_string(),
108 ));
109 },
110 };
111 Ok(Some((data, permit)))
112 } else {
113 match receiver.try_recv() {
115 Ok(data) => Ok(Some((data, permit))),
116 Err(oneshot::error::TryRecvError::Empty) => {
117 self.pending = Some(SequentialRetrievalItem::Data { receiver, permit });
119 Ok(None)
120 },
121 Err(oneshot::error::TryRecvError::Closed) => {
122 self.run_state.check_error()?;
123 Err(FileReconstructionError::InternalWriterError(
124 "Data sender was dropped before sending data.".to_string(),
125 ))
126 },
127 }
128 }
129 },
130 Some(SequentialRetrievalItem::Finish) => {
131 self.finished = true;
132 Ok(None)
133 },
134 None => Ok(None),
135 }
136 }
137
138 fn run(mut self, mut writer: impl Write) -> Result<()> {
140 while let Some((data, permit)) = self.next_write(true)? {
141 let len = data.len() as u64;
142 writer.write_all(&data)?;
143 self.bytes_written.fetch_add(len, Ordering::Relaxed);
144 if let Some(ref updater) = self.progress_updater {
145 updater.report_bytes_completed(len);
146 }
147 drop(permit);
148
149 if self.finished {
150 break;
151 }
152
153 self.ctx.check_sigint_shutdown()?;
154 }
155
156 debug_assert!(self.finished);
157
158 writer.flush()?;
159 Ok(())
160 }
161
162 fn run_vectorized(mut self, mut writer: impl Write) -> Result<()> {
164 let mut pending_writes: VecDeque<PendingWrite> = VecDeque::new();
165
166 while !self.finished || !pending_writes.is_empty() {
167 self.ctx.check_sigint_shutdown()?;
168
169 if pending_writes.is_empty() {
171 let Some(write) = self.next_write(true)? else {
172 break;
173 };
174
175 pending_writes.push_back(write);
176 }
177
178 while let Some(write) = self.next_write(false)? {
180 pending_writes.push_back(write);
181 }
182
183 let io_slices: Vec<IoSlice<'_>> = pending_writes
185 .iter()
186 .take(WRITEV_MAX_SLICE)
187 .map(|(data, _)| IoSlice::new(data))
188 .collect();
189
190 let written = match writer.write_vectored(&io_slices) {
192 Ok(0) if !io_slices.is_empty() => {
193 return Err(FileReconstructionError::IoError(Arc::new(std::io::Error::new(
194 std::io::ErrorKind::WriteZero,
195 "write_vectored returned 0 with non-empty buffers",
196 ))));
197 },
198 Ok(n) => n,
199 Err(ref e) if e.kind() == std::io::ErrorKind::Interrupted => continue,
200 Err(e) => return Err(FileReconstructionError::IoError(Arc::new(e))),
201 };
202
203 self.bytes_written.fetch_add(written as u64, Ordering::Relaxed);
204 if let Some(ref updater) = self.progress_updater {
205 updater.report_bytes_completed(written as u64);
206 }
207
208 let mut remaining = written;
210 while remaining > 0 && !pending_writes.is_empty() {
211 let front_len = pending_writes.front().unwrap().0.len();
212 if remaining >= front_len {
213 remaining -= front_len;
214 pending_writes.pop_front();
215 } else {
216 let front = pending_writes.front_mut().unwrap();
217 front.0 = front.0.slice(remaining..);
218 remaining = 0;
219 }
220 }
221 }
222
223 writer.flush()?;
224 Ok(())
225 }
226}
227
228pub struct SequentialWriter {
232 sender: UnboundedSender<SequentialRetrievalItem>,
233 next_position: u64,
234 background_handle: Option<JoinHandle<()>>,
235 run_state: Arc<RunState>,
236 bytes_written: Arc<AtomicU64>,
237 active_tasks: JoinSet<Result<()>>,
238 finished: bool,
239}
240
241impl Drop for SequentialWriter {
242 fn drop(&mut self) {
243 if !self.finished {
244 self.run_state.cancel();
245 }
246 }
247}
248
249#[cfg_attr(not(target_family = "wasm"), async_trait::async_trait)]
250#[cfg_attr(target_family = "wasm", async_trait::async_trait(?Send))]
251impl DataWriter for SequentialWriter {
252 async fn set_next_term_data_source(
256 &mut self,
257 byte_range: FileRange,
258 permit: Option<AdjustableSemaphorePermit>,
259 data_future: DataFuture,
260 ) -> Result<()> {
261 self.run_state.check_error()?;
262
263 while let Some(result) = self.active_tasks.try_join_next() {
264 result.map_err(|e| FileReconstructionError::InternalError(format!("Task join error: {e}")))??;
265 }
266
267 if self.finished {
268 return Err(FileReconstructionError::InternalWriterError("Writer has already finished".to_string()));
269 }
270
271 if byte_range.start != self.next_position {
272 return Err(FileReconstructionError::InternalWriterError(format!(
273 "Byte range not sequential: expected start at {}, got {}",
274 self.next_position, byte_range.start
275 )));
276 }
277
278 let expected_size = byte_range.end - byte_range.start;
279 self.next_position = byte_range.end;
280
281 let (sender, receiver) = oneshot::channel();
282
283 if self.sender.send(SequentialRetrievalItem::Data { receiver, permit }).is_err() {
284 self.run_state.check_error()?;
285 return Err(FileReconstructionError::InternalWriterError("Background writer channel closed".to_string()));
286 }
287
288 let run_state = self.run_state.clone();
289 let task = async move {
290 let result = async {
291 run_state.check_error()?;
292
293 let data = data_future.await?;
294
295 if data.len() as u64 != expected_size {
296 return Err(FileReconstructionError::InternalWriterError(format!(
297 "Data size mismatch: expected {} bytes, got {} bytes",
298 expected_size,
299 data.len()
300 )));
301 }
302
303 if sender.send(data).is_err() {
304 run_state.check_error()?;
305 return Err(FileReconstructionError::InternalWriterError(
306 "Failed to send data: receiver dropped".to_string(),
307 ));
308 }
309
310 Ok(())
311 }
312 .await;
313
314 if let Err(ref e) = result {
315 run_state.set_error(e.clone());
316 }
317 result
318 };
319
320 self.active_tasks.spawn(task);
321
322 Ok(())
323 }
324
325 async fn finish(mut self: Box<Self>) -> Result<u64> {
328 self.run_state.check_error()?;
329
330 if self.finished {
331 return Err(FileReconstructionError::InternalWriterError("Writer has already finished".to_string()));
332 }
333
334 self.finished = true;
335
336 if self.sender.send(SequentialRetrievalItem::Finish).is_err() {
337 self.run_state.check_error()?;
338 return Err(FileReconstructionError::InternalWriterError("Background writer channel closed".to_string()));
339 }
340
341 let expected_bytes = self.next_position;
342
343 while let Some(result) = self.active_tasks.join_next().await {
344 result.map_err(|e| FileReconstructionError::InternalError(format!("Task join error: {e}")))??;
345 }
346
347 match self.background_handle.take() {
348 Some(handle) => {
349 handle.await.map_err(|e| {
350 FileReconstructionError::InternalWriterError(format!("Background writer task failed: {e}"))
351 })?;
352
353 self.run_state.check_error()?;
354
355 let actual_bytes = self.bytes_written.load(Ordering::Relaxed);
356 if actual_bytes != expected_bytes {
357 return Err(FileReconstructionError::InternalWriterError(format!(
358 "Bytes written mismatch: expected {} bytes, but wrote {} bytes",
359 expected_bytes, actual_bytes
360 )));
361 }
362
363 Ok(actual_bytes)
364 },
365 None => {
366 Ok(expected_bytes)
369 },
370 }
371 }
372}
373
374impl SequentialWriter {
375 pub(crate) fn new_streaming(
381 run_state: Arc<RunState>,
382 ) -> (Box<dyn DataWriter>, UnboundedReceiver<SequentialRetrievalItem>) {
383 let (tx, rx) = unbounded_channel::<SequentialRetrievalItem>();
384
385 let writer = Self {
386 sender: tx,
387 next_position: 0,
388 background_handle: None,
389 run_state,
390 bytes_written: Arc::new(AtomicU64::new(0)),
391 active_tasks: JoinSet::new(),
392 finished: false,
393 };
394
395 (Box::new(writer), rx)
396 }
397
398 #[cfg(not(target_family = "wasm"))]
404 #[allow(clippy::new_ret_no_self)]
405 pub(crate) fn new<W: Write + Send + 'static>(
406 ctx: &XetContext,
407 writer: W,
408 use_vectorized: bool,
409 run_state: Arc<RunState>,
410 ) -> Box<dyn DataWriter> {
411 let (tx, rx) = unbounded_channel::<SequentialRetrievalItem>();
412 let bytes_written = Arc::new(AtomicU64::new(0));
413
414 let run_state_clone = run_state.clone();
415 let run_state_thread = run_state.clone();
416 let bytes_written_clone = bytes_written.clone();
417 let progress_updater = run_state.progress_updater().cloned();
418 let ctx_thread = ctx.clone();
419
420 let handle = ctx.runtime.spawn_blocking(move || {
421 let writer_thread =
422 SyncWriterThread::new(ctx_thread, rx, bytes_written_clone, progress_updater, run_state_thread);
423 let result = if use_vectorized {
424 writer_thread.run_vectorized(writer)
425 } else {
426 writer_thread.run(writer)
427 };
428 if let Err(err) = result {
429 run_state_clone.set_error(err);
430 }
431 });
432
433 Box::new(Self {
434 sender: tx,
435 next_position: 0,
436 background_handle: Some(handle),
437 run_state,
438 bytes_written,
439 active_tasks: JoinSet::new(),
440 finished: false,
441 })
442 }
443}
444
445#[cfg(test)]
446mod tests {
447 use std::io;
448 use std::time::Duration;
449
450 use xet_runtime::core::XetContext;
451 use xet_runtime::utils::adjustable_semaphore::AdjustableSemaphore;
452
453 use super::*;
454
455 fn test_context() -> XetContext {
456 XetContext::default().unwrap()
457 }
458
459 struct SharedBuffer(Arc<std::sync::Mutex<Vec<u8>>>);
460
461 impl Write for SharedBuffer {
462 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
463 self.0.lock().unwrap().extend_from_slice(buf);
464 Ok(buf.len())
465 }
466 fn flush(&mut self) -> io::Result<()> {
467 Ok(())
468 }
469 }
470
471 #[derive(Clone, Default)]
473 struct TestWriterConfig {
474 max_write_size: Option<usize>,
476 max_vectored_write_size: Option<usize>,
478 hard_limit_vectored_write_slice: Option<usize>,
481 simulate_interrupts: bool,
483 interrupt_frequency: usize,
485 }
486
487 impl TestWriterConfig {
488 fn vectorized() -> Self {
489 Self::default()
490 }
491
492 fn vectorized_partial(max_size: usize) -> Self {
493 Self {
494 max_vectored_write_size: Some(max_size),
495 ..Default::default()
496 }
497 }
498
499 fn vectorized_hard_limit(max_slice: usize) -> Self {
500 Self {
501 hard_limit_vectored_write_slice: Some(max_slice),
502 ..Default::default()
503 }
504 }
505
506 fn partial(max_size: usize) -> Self {
507 Self {
508 max_write_size: Some(max_size),
509 ..Default::default()
510 }
511 }
512
513 fn vectorized_with_interrupts() -> Self {
514 Self {
515 simulate_interrupts: true,
516 interrupt_frequency: 2,
517 ..Default::default()
518 }
519 }
520 }
521
522 struct TestWriter {
528 buffer: Arc<std::sync::Mutex<Vec<u8>>>,
529 config: TestWriterConfig,
530 write_count: Arc<AtomicU64>,
531 vectored_write_count: Arc<AtomicU64>,
532 interrupt_counter: Arc<AtomicU64>,
533 }
534
535 impl TestWriter {
536 fn new(config: TestWriterConfig) -> Self {
537 Self {
538 buffer: Arc::new(std::sync::Mutex::new(Vec::new())),
539 config,
540 write_count: Arc::new(AtomicU64::new(0)),
541 vectored_write_count: Arc::new(AtomicU64::new(0)),
542 interrupt_counter: Arc::new(AtomicU64::new(0)),
543 }
544 }
545
546 fn should_interrupt(&self) -> bool {
547 if !self.config.simulate_interrupts {
548 return false;
549 }
550 let count = self.interrupt_counter.fetch_add(1, Ordering::Relaxed);
551 count % self.config.interrupt_frequency as u64 == 0
552 }
553 }
554
555 impl Write for TestWriter {
556 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
557 if self.should_interrupt() {
558 return Err(io::Error::new(io::ErrorKind::Interrupted, "simulated interrupt"));
559 }
560
561 self.write_count.fetch_add(1, Ordering::Relaxed);
562
563 let bytes_to_write = match self.config.max_write_size {
564 Some(max) => buf.len().min(max),
565 None => buf.len(),
566 };
567
568 self.buffer.lock().unwrap().extend_from_slice(&buf[..bytes_to_write]);
569 Ok(bytes_to_write)
570 }
571
572 fn write_vectored(&mut self, bufs: &[IoSlice<'_>]) -> io::Result<usize> {
573 if self.should_interrupt() {
574 return Err(io::Error::new(io::ErrorKind::Interrupted, "simulated interrupt"));
575 }
576
577 if let Some(max_slice) = self.config.hard_limit_vectored_write_slice
578 && bufs.len() > max_slice
579 {
580 return Err(io::Error::new(io::ErrorKind::InvalidInput, "simulated iovcnt EINVAL"));
581 }
582
583 self.vectored_write_count.fetch_add(1, Ordering::Relaxed);
584
585 let total_len: usize = bufs.iter().map(|b| b.len()).sum();
586 let max_write = self.config.max_vectored_write_size.unwrap_or(total_len);
587 let bytes_to_write = total_len.min(max_write);
588
589 let mut remaining = bytes_to_write;
590 let mut buffer = self.buffer.lock().unwrap();
591
592 for buf in bufs {
593 if remaining == 0 {
594 break;
595 }
596 let to_write = buf.len().min(remaining);
597 buffer.extend_from_slice(&buf[..to_write]);
598 remaining -= to_write;
599 }
600
601 Ok(bytes_to_write)
602 }
603
604 fn flush(&mut self) -> io::Result<()> {
605 Ok(())
606 }
607 }
608
609 fn immediate_future(data: Bytes) -> DataFuture {
610 Box::pin(async move { Ok(data) })
611 }
612
613 #[tokio::test]
614 async fn test_sequential_writes() {
615 let buffer = Arc::new(std::sync::Mutex::new(Vec::new()));
616 let buffer_clone = buffer.clone();
617
618 let mut writer = SequentialWriter::new(
619 &test_context(),
620 Box::new(SharedBuffer(buffer_clone)),
621 false,
622 RunState::new_for_test(),
623 );
624
625 writer
626 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
627 .await
628 .unwrap();
629 writer
630 .set_next_term_data_source(FileRange::new(5, 6), None, immediate_future(Bytes::from(" ")))
631 .await
632 .unwrap();
633 writer
634 .set_next_term_data_source(FileRange::new(6, 11), None, immediate_future(Bytes::from("World")))
635 .await
636 .unwrap();
637
638 writer.finish().await.unwrap();
639
640 let result = buffer.lock().unwrap();
641 assert_eq!(&*result, b"Hello World");
642 }
643
644 #[tokio::test]
645 async fn test_delayed_future() {
646 let buffer = Arc::new(std::sync::Mutex::new(Vec::new()));
647 let buffer_clone = buffer.clone();
648
649 let mut writer = SequentialWriter::new(
650 &test_context(),
651 Box::new(SharedBuffer(buffer_clone)),
652 false,
653 RunState::new_for_test(),
654 );
655
656 let f0: DataFuture = Box::pin(async {
658 tokio::time::sleep(Duration::from_millis(50)).await;
659 Ok(Bytes::from("Hello"))
660 });
661 let f1: DataFuture = Box::pin(async {
662 tokio::time::sleep(Duration::from_millis(10)).await;
663 Ok(Bytes::from(" "))
664 });
665 let f2: DataFuture = Box::pin(async { Ok(Bytes::from("World")) });
666
667 writer.set_next_term_data_source(FileRange::new(0, 5), None, f0).await.unwrap();
668 writer.set_next_term_data_source(FileRange::new(5, 6), None, f1).await.unwrap();
669 writer.set_next_term_data_source(FileRange::new(6, 11), None, f2).await.unwrap();
670
671 writer.finish().await.unwrap();
672
673 let result = buffer.lock().unwrap();
674 assert_eq!(&*result, b"Hello World");
675 }
676
677 #[tokio::test]
678 async fn test_size_mismatch_error() {
679 let buffer = std::io::Cursor::new(Vec::new());
680 let mut writer = SequentialWriter::new(&test_context(), Box::new(buffer), false, RunState::new_for_test());
681
682 writer
683 .set_next_term_data_source(FileRange::new(0, 10), None, immediate_future(Bytes::from("Hello")))
684 .await
685 .unwrap();
686
687 let result = writer.finish().await;
688 assert!(result.is_err());
689 }
690
691 #[tokio::test]
692 async fn test_background_writer_error_propagates() {
693 struct FailingWriter;
694 impl Write for FailingWriter {
695 fn write(&mut self, _buf: &[u8]) -> io::Result<usize> {
696 Err(io::Error::new(io::ErrorKind::Other, "Simulated write failure"))
697 }
698 fn flush(&mut self) -> io::Result<()> {
699 Ok(())
700 }
701 }
702
703 let mut writer =
704 SequentialWriter::new(&test_context(), Box::new(FailingWriter), false, RunState::new_for_test());
705
706 writer
707 .set_next_term_data_source(FileRange::new(0, 4), None, immediate_future(Bytes::from("Test")))
708 .await
709 .unwrap();
710
711 tokio::time::sleep(Duration::from_millis(200)).await;
712
713 let result = writer
714 .set_next_term_data_source(FileRange::new(4, 8), None, immediate_future(Bytes::from("More")))
715 .await;
716
717 assert!(result.is_err());
718 assert!(matches!(result, Err(FileReconstructionError::IoError(_))));
719 }
720
721 #[tokio::test]
722 async fn test_flush_error_propagates() {
723 struct FlushFailingWriter;
724 impl Write for FlushFailingWriter {
725 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
726 Ok(buf.len())
727 }
728 fn flush(&mut self) -> io::Result<()> {
729 Err(io::Error::new(io::ErrorKind::Other, "Simulated flush failure"))
730 }
731 }
732
733 let writer =
734 SequentialWriter::new(&test_context(), Box::new(FlushFailingWriter), false, RunState::new_for_test());
735 let result = writer.finish().await;
736 assert!(result.is_err());
737 assert!(matches!(result, Err(FileReconstructionError::IoError(_))));
738 }
739
740 #[tokio::test]
741 async fn test_future_error_propagates() {
742 let buffer = Arc::new(std::sync::Mutex::new(Vec::new()));
743 let buffer_clone = buffer.clone();
744
745 let mut writer = SequentialWriter::new(
746 &test_context(),
747 Box::new(SharedBuffer(buffer_clone)),
748 false,
749 RunState::new_for_test(),
750 );
751
752 let failing_future: DataFuture =
753 Box::pin(async { Err(FileReconstructionError::InternalError("Simulated future error".to_string())) });
754
755 writer
756 .set_next_term_data_source(FileRange::new(0, 5), None, failing_future)
757 .await
758 .unwrap();
759
760 let result = writer.finish().await;
761 assert!(result.is_err());
762 }
763
764 #[tokio::test]
765 async fn test_size_mismatch_too_small() {
766 let buffer = std::io::Cursor::new(Vec::new());
767 let mut writer = SequentialWriter::new(&test_context(), Box::new(buffer), false, RunState::new_for_test());
768
769 writer
770 .set_next_term_data_source(FileRange::new(0, 10), None, immediate_future(Bytes::from("Hi")))
771 .await
772 .unwrap();
773
774 let result = writer.finish().await;
775 assert!(result.is_err());
776 }
777
778 #[tokio::test]
779 async fn test_size_mismatch_too_large() {
780 let buffer = std::io::Cursor::new(Vec::new());
781 let mut writer = SequentialWriter::new(&test_context(), Box::new(buffer), false, RunState::new_for_test());
782
783 writer
784 .set_next_term_data_source(FileRange::new(0, 2), None, immediate_future(Bytes::from("Hello World")))
785 .await
786 .unwrap();
787
788 let result = writer.finish().await;
789 assert!(result.is_err());
790 }
791
792 #[tokio::test]
793 async fn test_bytes_written_tracking() {
794 let buffer = Arc::new(std::sync::Mutex::new(Vec::new()));
795 let buffer_clone = buffer.clone();
796
797 let mut writer = SequentialWriter::new(
798 &test_context(),
799 Box::new(SharedBuffer(buffer_clone)),
800 false,
801 RunState::new_for_test(),
802 );
803
804 writer
805 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
806 .await
807 .unwrap();
808 writer
809 .set_next_term_data_source(FileRange::new(5, 11), None, immediate_future(Bytes::from(" World")))
810 .await
811 .unwrap();
812 writer
813 .set_next_term_data_source(FileRange::new(11, 16), None, immediate_future(Bytes::from("!!!!!")))
814 .await
815 .unwrap();
816
817 writer.finish().await.unwrap();
818
819 let result = buffer.lock().unwrap();
820 assert_eq!(&*result, b"Hello World!!!!!");
821 assert_eq!(result.len(), 16);
822 }
823
824 #[tokio::test]
825 async fn test_non_sequential_range_returns_error() {
826 let buffer = std::io::Cursor::new(Vec::new());
827 let mut writer = SequentialWriter::new(&test_context(), Box::new(buffer), false, RunState::new_for_test());
828
829 writer
830 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
831 .await
832 .unwrap();
833
834 let result = writer
835 .set_next_term_data_source(FileRange::new(10, 15), None, immediate_future(Bytes::from("World")))
836 .await;
837 assert!(result.is_err());
838 assert!(matches!(result, Err(FileReconstructionError::InternalWriterError(_))));
839 }
840
841 #[tokio::test]
842 async fn test_first_range_must_start_at_zero() {
843 let buffer = std::io::Cursor::new(Vec::new());
844 let mut writer = SequentialWriter::new(&test_context(), Box::new(buffer), false, RunState::new_for_test());
845
846 let result = writer
847 .set_next_term_data_source(FileRange::new(5, 10), None, immediate_future(Bytes::from("Hello")))
848 .await;
849 assert!(result.is_err());
850 assert!(matches!(result, Err(FileReconstructionError::InternalWriterError(_))));
851 }
852
853 #[tokio::test]
854 async fn test_semaphore_permit_released_after_write() {
855 let buffer = Arc::new(std::sync::Mutex::new(Vec::new()));
856 let buffer_clone = buffer.clone();
857 let semaphore = AdjustableSemaphore::new(2, (0, 2));
858
859 let mut writer = SequentialWriter::new(
860 &test_context(),
861 Box::new(SharedBuffer(buffer_clone)),
862 false,
863 RunState::new_for_test(),
864 );
865
866 let permit1 = semaphore.acquire().await.unwrap();
867 let permit2 = semaphore.acquire().await.unwrap();
868
869 assert_eq!(semaphore.available_permits(), 0);
870
871 writer
872 .set_next_term_data_source(FileRange::new(0, 5), Some(permit1), immediate_future(Bytes::from("Hello")))
873 .await
874 .unwrap();
875
876 tokio::time::sleep(Duration::from_millis(50)).await;
877 assert_eq!(semaphore.available_permits(), 1);
878
879 writer
880 .set_next_term_data_source(FileRange::new(5, 6), Some(permit2), immediate_future(Bytes::from(" ")))
881 .await
882 .unwrap();
883
884 tokio::time::sleep(Duration::from_millis(50)).await;
885 assert_eq!(semaphore.available_permits(), 2);
886
887 writer.finish().await.unwrap();
888
889 let result = buffer.lock().unwrap();
890 assert_eq!(&*result, b"Hello ");
891 }
892
893 #[tokio::test]
896 async fn test_vectorized_basic_writes() {
897 let test_writer = TestWriter::new(TestWriterConfig::vectorized());
898 let buffer = test_writer.buffer.clone();
899 let vectored_count = test_writer.vectored_write_count.clone();
900
901 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
902
903 writer
904 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
905 .await
906 .unwrap();
907 writer
908 .set_next_term_data_source(FileRange::new(5, 6), None, immediate_future(Bytes::from(" ")))
909 .await
910 .unwrap();
911 writer
912 .set_next_term_data_source(FileRange::new(6, 11), None, immediate_future(Bytes::from("World")))
913 .await
914 .unwrap();
915
916 writer.finish().await.unwrap();
917
918 let result = buffer.lock().unwrap();
919 assert_eq!(&*result, b"Hello World");
920 assert!(vectored_count.load(Ordering::Relaxed) > 0);
921 }
922
923 #[tokio::test]
924 async fn test_vectorized_partial_writes() {
925 let test_writer = TestWriter::new(TestWriterConfig::vectorized_partial(3));
926 let buffer = test_writer.buffer.clone();
927
928 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
929
930 writer
931 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
932 .await
933 .unwrap();
934 writer
935 .set_next_term_data_source(FileRange::new(5, 6), None, immediate_future(Bytes::from(" ")))
936 .await
937 .unwrap();
938 writer
939 .set_next_term_data_source(FileRange::new(6, 11), None, immediate_future(Bytes::from("World")))
940 .await
941 .unwrap();
942 writer
943 .set_next_term_data_source(FileRange::new(11, 12), None, immediate_future(Bytes::from("!")))
944 .await
945 .unwrap();
946
947 writer.finish().await.unwrap();
948
949 let result = buffer.lock().unwrap();
950 assert_eq!(&*result, b"Hello World!");
951 }
952
953 #[tokio::test]
954 async fn test_vectorized_with_delays() {
955 let test_writer = TestWriter::new(TestWriterConfig::vectorized());
956 let buffer = test_writer.buffer.clone();
957
958 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
959
960 let f0: DataFuture = Box::pin(async {
962 tokio::time::sleep(Duration::from_millis(30)).await;
963 Ok(Bytes::from("A"))
964 });
965 let f1: DataFuture = Box::pin(async {
966 tokio::time::sleep(Duration::from_millis(10)).await;
967 Ok(Bytes::from("B"))
968 });
969 let f2: DataFuture = Box::pin(async { Ok(Bytes::from("C")) });
970
971 writer.set_next_term_data_source(FileRange::new(0, 1), None, f0).await.unwrap();
972 writer.set_next_term_data_source(FileRange::new(1, 2), None, f1).await.unwrap();
973 writer.set_next_term_data_source(FileRange::new(2, 3), None, f2).await.unwrap();
974
975 writer.finish().await.unwrap();
976
977 let result = buffer.lock().unwrap();
978 assert_eq!(&*result, b"ABC");
979 }
980
981 #[tokio::test]
982 async fn test_vectorized_many_small_writes() {
983 let expected: Vec<u8> = (0..100u8).collect();
984 let test_writer = TestWriter::new(TestWriterConfig::vectorized());
985 let buffer = test_writer.buffer.clone();
986 let vectored_count = test_writer.vectored_write_count.clone();
987
988 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
989
990 for i in 0..100u8 {
992 writer
993 .set_next_term_data_source(
994 FileRange::new(i as u64, i as u64 + 1),
995 None,
996 immediate_future(Bytes::from(vec![i])),
997 )
998 .await
999 .unwrap();
1000 }
1001
1002 writer.finish().await.unwrap();
1003
1004 let result = buffer.lock().unwrap();
1005 assert_eq!(&*result, &expected);
1006
1007 let vectored_calls = vectored_count.load(Ordering::Relaxed);
1009 assert!(vectored_calls < 100);
1010 }
1011
1012 #[tokio::test]
1013 async fn test_vectorized_with_interrupts() {
1014 let test_writer = TestWriter::new(TestWriterConfig::vectorized_with_interrupts());
1015 let buffer = test_writer.buffer.clone();
1016
1017 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
1018
1019 writer
1020 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
1021 .await
1022 .unwrap();
1023 writer
1024 .set_next_term_data_source(FileRange::new(5, 6), None, immediate_future(Bytes::from(" ")))
1025 .await
1026 .unwrap();
1027 writer
1028 .set_next_term_data_source(FileRange::new(6, 11), None, immediate_future(Bytes::from("World")))
1029 .await
1030 .unwrap();
1031
1032 writer.finish().await.unwrap();
1033
1034 let result = buffer.lock().unwrap();
1035 assert_eq!(&*result, b"Hello World");
1036 }
1037
1038 #[tokio::test]
1039 async fn test_vectorized_permit_release() {
1040 let test_writer = TestWriter::new(TestWriterConfig::vectorized());
1041 let buffer = test_writer.buffer.clone();
1042 let semaphore = AdjustableSemaphore::new(2, (0, 2));
1043
1044 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
1045
1046 let permit1 = semaphore.acquire().await.unwrap();
1047 let permit2 = semaphore.acquire().await.unwrap();
1048
1049 assert_eq!(semaphore.available_permits(), 0);
1050
1051 writer
1052 .set_next_term_data_source(FileRange::new(0, 5), Some(permit1), immediate_future(Bytes::from("Hello")))
1053 .await
1054 .unwrap();
1055
1056 tokio::time::sleep(Duration::from_millis(50)).await;
1057 assert_eq!(semaphore.available_permits(), 1);
1058
1059 writer
1060 .set_next_term_data_source(FileRange::new(5, 6), Some(permit2), immediate_future(Bytes::from(" ")))
1061 .await
1062 .unwrap();
1063
1064 tokio::time::sleep(Duration::from_millis(50)).await;
1065 assert_eq!(semaphore.available_permits(), 2);
1066
1067 writer.finish().await.unwrap();
1068
1069 let result = buffer.lock().unwrap();
1070 assert_eq!(&*result, b"Hello ");
1071 }
1072
1073 #[tokio::test]
1074 async fn test_vectorized_partial_permit_release() {
1075 let test_writer = TestWriter::new(TestWriterConfig::vectorized_partial(2));
1076 let buffer = test_writer.buffer.clone();
1077 let semaphore = AdjustableSemaphore::new(3, (0, 3));
1078
1079 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
1080
1081 let permit1 = semaphore.acquire().await.unwrap();
1082 let permit2 = semaphore.acquire().await.unwrap();
1083 let permit3 = semaphore.acquire().await.unwrap();
1084
1085 assert_eq!(semaphore.available_permits(), 0);
1086
1087 writer
1088 .set_next_term_data_source(FileRange::new(0, 5), Some(permit1), immediate_future(Bytes::from("Hello")))
1089 .await
1090 .unwrap();
1091 writer
1092 .set_next_term_data_source(FileRange::new(5, 11), Some(permit2), immediate_future(Bytes::from(" World")))
1093 .await
1094 .unwrap();
1095 writer
1096 .set_next_term_data_source(FileRange::new(11, 12), Some(permit3), immediate_future(Bytes::from("!")))
1097 .await
1098 .unwrap();
1099
1100 writer.finish().await.unwrap();
1101
1102 assert_eq!(semaphore.available_permits(), 3);
1103
1104 let result = buffer.lock().unwrap();
1105 assert_eq!(&*result, b"Hello World!");
1106 }
1107
1108 #[tokio::test]
1109 async fn test_non_vectorized_basic_writes() {
1110 let test_writer = TestWriter::new(TestWriterConfig::default());
1111 let buffer = test_writer.buffer.clone();
1112 let write_count = test_writer.write_count.clone();
1113 let vectored_count = test_writer.vectored_write_count.clone();
1114
1115 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), false, RunState::new_for_test());
1116
1117 writer
1118 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
1119 .await
1120 .unwrap();
1121 writer
1122 .set_next_term_data_source(FileRange::new(5, 6), None, immediate_future(Bytes::from(" ")))
1123 .await
1124 .unwrap();
1125 writer
1126 .set_next_term_data_source(FileRange::new(6, 11), None, immediate_future(Bytes::from("World")))
1127 .await
1128 .unwrap();
1129
1130 writer.finish().await.unwrap();
1131
1132 let result = buffer.lock().unwrap();
1133 assert_eq!(&*result, b"Hello World");
1134 assert!(write_count.load(Ordering::Relaxed) > 0);
1135 assert_eq!(vectored_count.load(Ordering::Relaxed), 0);
1136 }
1137
1138 #[tokio::test]
1139 async fn test_non_vectorized_partial_writes() {
1140 let test_writer = TestWriter::new(TestWriterConfig::partial(3));
1141 let buffer = test_writer.buffer.clone();
1142
1143 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), false, RunState::new_for_test());
1144
1145 writer
1146 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("Hello")))
1147 .await
1148 .unwrap();
1149 writer
1150 .set_next_term_data_source(FileRange::new(5, 6), None, immediate_future(Bytes::from(" ")))
1151 .await
1152 .unwrap();
1153 writer
1154 .set_next_term_data_source(FileRange::new(6, 11), None, immediate_future(Bytes::from("World")))
1155 .await
1156 .unwrap();
1157 writer
1158 .set_next_term_data_source(FileRange::new(11, 12), None, immediate_future(Bytes::from("!")))
1159 .await
1160 .unwrap();
1161
1162 writer.finish().await.unwrap();
1163
1164 let result = buffer.lock().unwrap();
1165 assert_eq!(&*result, b"Hello World!");
1166 }
1167
1168 #[tokio::test]
1169 async fn test_vectorized_single_byte_partial() {
1170 let test_writer = TestWriter::new(TestWriterConfig::vectorized_partial(1));
1171 let buffer = test_writer.buffer.clone();
1172
1173 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
1174
1175 writer
1176 .set_next_term_data_source(FileRange::new(0, 5), None, immediate_future(Bytes::from("ABCDE")))
1177 .await
1178 .unwrap();
1179 writer
1180 .set_next_term_data_source(FileRange::new(5, 10), None, immediate_future(Bytes::from("FGHIJ")))
1181 .await
1182 .unwrap();
1183
1184 writer.finish().await.unwrap();
1185
1186 let result = buffer.lock().unwrap();
1187 assert_eq!(&*result, b"ABCDEFGHIJ");
1188 }
1189
1190 #[tokio::test]
1191 async fn test_vectorized_large_data() {
1192 let expected: Vec<u8> = (0..10000).map(|i| (i % 256) as u8).collect();
1193 let test_writer = TestWriter::new(TestWriterConfig::vectorized());
1194 let buffer = test_writer.buffer.clone();
1195
1196 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
1197
1198 for i in 0..10 {
1200 let start = i * 1000;
1201 let end = start + 1000;
1202 let chunk: Vec<u8> = (start..end).map(|j| (j % 256) as u8).collect();
1203 writer
1204 .set_next_term_data_source(
1205 FileRange::new(start as u64, end as u64),
1206 None,
1207 immediate_future(Bytes::from(chunk)),
1208 )
1209 .await
1210 .unwrap();
1211 }
1212
1213 writer.finish().await.unwrap();
1214
1215 let result = buffer.lock().unwrap();
1216 assert_eq!(&*result, &expected);
1217 }
1218
1219 #[tokio::test]
1220 async fn test_vectorized_large_data_partial() {
1221 let expected: Vec<u8> = (0..5000).map(|i| (i % 256) as u8).collect();
1222 let test_writer = TestWriter::new(TestWriterConfig::vectorized_partial(100));
1223 let buffer = test_writer.buffer.clone();
1224
1225 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test());
1226
1227 for i in 0..10 {
1229 let start = i * 500;
1230 let end = start + 500;
1231 let chunk: Vec<u8> = (start..end).map(|j| (j % 256) as u8).collect();
1232 writer
1233 .set_next_term_data_source(
1234 FileRange::new(start as u64, end as u64),
1235 None,
1236 immediate_future(Bytes::from(chunk)),
1237 )
1238 .await
1239 .unwrap();
1240 }
1241
1242 writer.finish().await.unwrap();
1243
1244 let result = buffer.lock().unwrap();
1245 assert_eq!(&*result, &expected);
1246 }
1247
1248 #[tokio::test]
1249 async fn test_vectorized_exceeded_max_slice() {
1250 let test_writer = TestWriter::new(TestWriterConfig::vectorized_hard_limit(2)); let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test()); for i in 0..1000 {
1256 let start = i * 10;
1257 let end = start + 10;
1258 let chunk: Vec<u8> = (start..end).map(|j| (j % 256) as u8).collect();
1259 if writer
1260 .set_next_term_data_source(
1261 FileRange::new(start as u64, end as u64),
1262 None,
1263 immediate_future(Bytes::from(chunk)),
1264 )
1265 .await
1266 .is_err()
1267 {
1268 break;
1269 }
1270 }
1271
1272 let ret = writer.finish().await;
1273 assert!(ret.is_err());
1274 if let Err(FileReconstructionError::IoError(inner_err)) = ret {
1275 assert_eq!(inner_err.kind(), std::io::ErrorKind::InvalidInput);
1276 };
1277 }
1278
1279 #[tokio::test]
1280 async fn test_vectorized_controlled_max_slice() {
1281 let expected: Vec<u8> = (0..10000).map(|i| (i % 256) as u8).collect();
1282 let test_writer = TestWriter::new(TestWriterConfig::vectorized_hard_limit(40)); let buffer = test_writer.buffer.clone();
1284
1285 let mut writer = SequentialWriter::new(&test_context(), Box::new(test_writer), true, RunState::new_for_test()); for i in 0..1000 {
1289 let start = i * 10;
1290 let end = start + 10;
1291 let chunk: Vec<u8> = (start..end).map(|j| (j % 256) as u8).collect();
1292 writer
1293 .set_next_term_data_source(
1294 FileRange::new(start as u64, end as u64),
1295 None,
1296 immediate_future(Bytes::from(chunk)),
1297 )
1298 .await
1299 .unwrap();
1300 }
1301
1302 writer.finish().await.unwrap();
1303
1304 let result = buffer.lock().unwrap();
1305 assert_eq!(&*result, &expected);
1306 }
1307}