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
16pub(crate) struct CompletedTerm {
20 pub byte_range: FileRange,
21 pub data: Bytes,
22 pub permit: Option<AdjustableSemaphorePermit>,
23}
24
25pub(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
43pub 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 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 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}