1use std::{
8 fmt,
9 io::{BufRead, ErrorKind, Read, Write},
10 pin::Pin,
11 sync::Arc,
12 task::{self, Poll},
13 time::Duration,
14};
15
16use crate::{McWaker, Progress};
17
18#[doc(no_inline)]
19pub use futures_lite::io::{
20 AsyncBufRead, AsyncBufReadExt, AsyncRead, AsyncReadExt, AsyncSeek, AsyncSeekExt, AsyncWrite, AsyncWriteExt, BoxedReader, BoxedWriter,
21 BufReader, BufWriter, Cursor, ReadHalf, WriteHalf, copy, empty, repeat, sink, split,
22};
23
24#[doc(no_inline)]
25pub use blocking::Unblock;
26
27use parking_lot::Mutex;
28use std::io::{Error, Result};
29use zng_time::{DInstant, INSTANT};
30use zng_txt::formatx;
31use zng_unit::{ByteLength, ByteUnits};
32use zng_var::{Var, impl_from_and_into_var, var};
33
34struct MeasureInner {
35 metrics: Var<Metrics>,
36 start_time: DInstant,
37 last_write: DInstant,
38 last_read: DInstant,
39}
40impl MeasureInner {
41 fn new(read_progress: (ByteLength, ByteLength), write_progress: (ByteLength, ByteLength)) -> Self {
42 let now = INSTANT.now();
43 Self {
44 metrics: var(Metrics {
45 read_progress,
46 read_speed: 0.bytes(),
47 write_progress,
48 write_speed: 0.bytes(),
49 total_time: Duration::ZERO,
50 }),
51 start_time: now,
52 last_write: now,
53 last_read: now,
54 }
55 }
56
57 fn on_read(&mut self, bytes: u64) {
58 if bytes == 0 {
59 return;
60 }
61
62 let bytes = bytes.bytes();
63
64 let now = INSTANT.now();
65 let elapsed = now - self.last_read;
66
67 self.last_read = now;
68 let read_speed = bytes_per_sec(bytes, elapsed);
69
70 let total_time = now - self.start_time;
71
72 self.metrics.modify(move |m| {
73 m.read_progress.0 += bytes;
74 m.read_speed = read_speed;
75 m.total_time = total_time;
76 });
77 }
78
79 fn on_write(&mut self, bytes: u64) {
80 if bytes == 0 {
81 return;
82 }
83
84 let bytes = bytes.bytes();
85
86 let now = INSTANT.now();
87 let elapsed = now - self.last_write;
88
89 self.last_write = now;
90 let write_speed = bytes_per_sec(bytes, elapsed);
91
92 let total_time = now - self.start_time;
93
94 self.metrics.modify(move |m| {
95 m.write_progress.0 += bytes;
96 m.write_speed = write_speed;
97 m.total_time = total_time;
98 });
99 }
100}
101
102pub struct Measure<T> {
107 task: T,
108 inner: MeasureInner,
109}
110impl<T> Measure<T> {
111 pub fn new(task: T, total_read: ByteLength, total_write: ByteLength) -> Self {
113 Self::new_ongoing(task, (0.bytes(), total_read), (0.bytes(), total_write))
114 }
115
116 pub fn new_ongoing(task: T, read_progress: (ByteLength, ByteLength), write_progress: (ByteLength, ByteLength)) -> Self {
118 Measure {
119 task,
120 inner: MeasureInner::new(read_progress, write_progress),
121 }
122 }
123
124 pub fn metrics(&self) -> Var<Metrics> {
128 self.inner.metrics.read_only()
129 }
130
131 pub fn finish(self) -> (T, Metrics) {
133 let mut metrics = self.inner.metrics.get();
134 metrics.total_time = self.inner.start_time.elapsed();
135 (self.task, metrics)
136 }
137}
138
139fn bytes_per_sec(bytes: ByteLength, elapsed: Duration) -> ByteLength {
140 let bytes_per_sec = bytes.0 as u128 / elapsed.as_nanos() / Duration::from_secs(1).as_nanos();
141 ByteLength(bytes_per_sec as u64)
142}
143
144impl<T: AsyncRead> AsyncRead for Measure<T> {
145 fn poll_read(self: Pin<&mut Self>, cx: &mut task::Context<'_>, buf: &mut [u8]) -> Poll<Result<usize>> {
146 let self_ = unsafe { self.get_unchecked_mut() };
148
149 match unsafe { Pin::new_unchecked(&mut self_.task) }.poll_read(cx, buf) {
151 Poll::Ready(Ok(bytes)) => {
152 self_.inner.on_read(bytes as u64);
153 Poll::Ready(Ok(bytes))
154 }
155 p => p,
156 }
157 }
158}
159impl<T: AsyncWrite> AsyncWrite for Measure<T> {
160 fn poll_write(self: Pin<&mut Self>, cx: &mut task::Context<'_>, buf: &[u8]) -> Poll<Result<usize>> {
161 let self_ = unsafe { self.get_unchecked_mut() };
163
164 match unsafe { Pin::new_unchecked(&mut self_.task) }.poll_write(cx, buf) {
166 Poll::Ready(Ok(bytes)) => {
167 self_.inner.on_write(bytes as u64);
168 Poll::Ready(Ok(bytes))
169 }
170 p => p,
171 }
172 }
173
174 fn poll_flush(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Result<()>> {
175 let self_ = unsafe { self.get_unchecked_mut() };
177
178 unsafe { Pin::new_unchecked(&mut self_.task) }.poll_flush(cx)
180 }
181
182 fn poll_close(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Result<()>> {
183 let self_ = unsafe { self.get_unchecked_mut() };
185
186 unsafe { Pin::new_unchecked(&mut self_.task) }.poll_flush(cx)
188 }
189}
190impl<T: AsyncBufRead> AsyncBufRead for Measure<T> {
191 fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Result<&[u8]>> {
192 let self_ = unsafe { self.get_unchecked_mut() };
194
195 unsafe { Pin::new_unchecked(&mut self_.task) }.poll_fill_buf(cx)
197 }
198
199 fn consume(self: Pin<&mut Self>, amt: usize) {
200 let self_ = unsafe { self.get_unchecked_mut() };
202 unsafe { Pin::new_unchecked(&mut self_.task) }.consume(amt);
204 self_.inner.on_read(amt as u64);
205 }
206}
207impl<T: Read> Read for Measure<T> {
208 fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
209 match self.task.read(buf) {
210 Ok(bytes) => {
211 self.inner.on_read(bytes as u64);
212 Ok(bytes)
213 }
214 r => r,
215 }
216 }
217}
218impl<T: Write> Write for Measure<T> {
219 fn write(&mut self, buf: &[u8]) -> Result<usize> {
220 match self.task.write(buf) {
221 Ok(bytes) => {
222 self.inner.on_write(bytes as u64);
223 Ok(bytes)
224 }
225 r => r,
226 }
227 }
228
229 fn flush(&mut self) -> Result<()> {
230 self.task.flush()
231 }
232}
233impl<T: BufRead> BufRead for Measure<T> {
234 fn fill_buf(&mut self) -> Result<&[u8]> {
235 self.task.fill_buf()
236 }
237
238 fn consume(&mut self, amount: usize) {
239 self.task.consume(amount);
240 self.inner.on_read(amount as u64);
241 }
242}
243
244#[derive(Debug, Clone, PartialEq, Eq)]
251#[non_exhaustive]
252pub struct Metrics {
253 pub read_progress: (ByteLength, ByteLength),
255
256 pub read_speed: ByteLength,
258
259 pub write_progress: (ByteLength, ByteLength),
261
262 pub write_speed: ByteLength,
264
265 pub total_time: Duration,
268}
269impl Metrics {
270 pub fn zero() -> Self {
272 Self {
273 read_progress: (0.bytes(), 0.bytes()),
274 read_speed: 0.bytes(),
275 write_progress: (0.bytes(), 0.bytes()),
276 write_speed: 0.bytes(),
277 total_time: Duration::ZERO,
278 }
279 }
280}
281impl fmt::Display for Metrics {
282 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
283 let mut nl = false;
284 if self.read_progress.1 > 0.bytes() {
285 nl = true;
286 if self.read_progress.0 != self.read_progress.1 {
287 if self.read_progress.0 <= self.read_progress.1 {
288 write!(f, "↓ {}-{}, {}/s", self.read_progress.0, self.read_progress.1, self.read_speed)?;
289 } else {
290 write!(f, "↓ {}, {}/s", self.read_progress.0, self.read_speed)?;
291 }
292 } else {
293 write!(f, "↓ {} . {:?}", self.read_progress.0, self.total_time)?;
294 }
295 }
296 if self.write_progress.1 > 0.bytes() {
297 if nl {
298 writeln!(f)?;
299 }
300 if self.write_progress.0 != self.write_progress.1 {
301 if self.write_progress.0 <= self.write_progress.1 {
302 write!(f, "↑ {}-{}, {}/s", self.write_progress.0, self.write_progress.1, self.write_speed)?;
303 } else {
304 write!(f, "↑ {}, {}/s", self.write_progress.0, self.write_speed)?;
305 }
306 } else {
307 write!(f, "↑ {} . {:?}", self.write_progress.0, self.total_time)?;
308 }
309 }
310
311 Ok(())
312 }
313}
314impl_from_and_into_var! {
315 fn from(metrics: Metrics) -> Progress {
316 let mut status = Progress::indeterminate();
317 if metrics.read_progress.1 > 0.bytes() {
318 status = Progress::from_n_of(metrics.read_progress.0.0, metrics.read_progress.1.0);
319 }
320 if metrics.write_progress.1 > 0.bytes() {
321 let w_status = Progress::from_n_of(metrics.write_progress.0.0, metrics.write_progress.1.0);
322 if status.is_indeterminate() {
323 status = w_status;
324 } else {
325 status = status.and_fct(w_status.fct());
326 }
327 }
328 status.with_msg(formatx!("{metrics}")).with_meta_mut(|mut m| {
329 m.set(*METRICS_ID, metrics);
330 })
331 }
332}
333
334zng_state_map::static_id! {
335 pub static ref METRICS_ID: zng_state_map::StateId<Metrics>;
337}
338
339pub trait McBufErrorExt {
341 fn is_only_lazy_left(&self) -> bool;
346}
347impl McBufErrorExt for std::io::Error {
348 fn is_only_lazy_left(&self) -> bool {
349 matches!(self.kind(), ErrorKind::Other) && format!("{self:?}").contains(ONLY_NON_LAZY_ERROR_MSG)
350 }
351}
352const ONLY_NON_LAZY_ERROR_MSG: &str = "no non-lazy readers left to read";
353
354pub struct McBufReader<S: AsyncRead> {
376 inner: Arc<Mutex<McBufInner<S>>>,
377 index: usize,
378 lazy: bool,
379}
380struct McBufInner<S: AsyncRead> {
381 source: Option<S>,
382 waker: McWaker,
383 lazy_wakers: Vec<task::Waker>,
384
385 buf: Vec<u8>,
386
387 clones: Vec<usize>,
388 non_lazy_count: usize,
389
390 result: ReadState,
391}
392impl<S: AsyncRead> McBufReader<S> {
393 pub fn new(source: S) -> Self {
395 let mut clones = Vec::with_capacity(2);
396 clones.push(0);
397 McBufReader {
398 inner: Arc::new(Mutex::new(McBufInner {
399 source: Some(source),
400 waker: McWaker::empty(),
401 lazy_wakers: vec![],
402
403 buf: Vec::with_capacity(10.kilobytes().0 as usize),
404
405 clones,
406 non_lazy_count: 1,
407
408 result: ReadState::Running,
409 })),
410 index: 0,
411 lazy: false,
412 }
413 }
414
415 pub fn is_lazy(&self) -> bool {
419 self.lazy
420 }
421
422 pub fn set_lazy(&mut self, lazy: bool) {
426 if self.lazy != lazy {
427 if lazy {
428 self.inner.lock().non_lazy_count -= 1;
429 } else {
430 self.inner.lock().non_lazy_count += 1;
431 }
432 self.lazy = lazy;
433 }
434 }
435}
436impl<S: AsyncRead> Clone for McBufReader<S> {
437 fn clone(&self) -> Self {
438 let mut inner = self.inner.lock();
439
440 let offset = inner.clones[self.index];
441 let index = inner.clones.len();
442 inner.clones.push(offset);
443
444 if !self.lazy {
445 inner.non_lazy_count += 1;
446 }
447
448 Self {
449 inner: self.inner.clone(),
450 index,
451 lazy: self.lazy,
452 }
453 }
454}
455impl<S: AsyncRead> Drop for McBufReader<S> {
456 fn drop(&mut self) {
457 let mut inner = self.inner.lock();
458 inner.clones[self.index] = usize::MAX;
459 if !self.lazy {
460 inner.non_lazy_count -= 1;
461 if inner.non_lazy_count == 0 {
462 for waker in inner.lazy_wakers.drain(..) {
464 waker.wake();
465 }
466 }
467 }
468 }
469}
470impl<S: AsyncRead> AsyncRead for McBufReader<S> {
471 fn poll_read(self: Pin<&mut Self>, cx: &mut task::Context<'_>, buf: &mut [u8]) -> Poll<Result<usize>> {
472 let self_ = self.as_ref();
473 let mut inner = self_.inner.lock();
474 let inner = &mut *inner;
475
476 let mut i = inner.clones[self_.index];
478 let mut ready;
479
480 match &inner.result {
481 ReadState::Running => {
482 ready = &inner.buf[i..];
485
486 if ready.is_empty() {
487 if self.lazy {
488 if inner.non_lazy_count == 0 {
489 return Poll::Ready(Err(Error::other(ONLY_NON_LAZY_ERROR_MSG)));
491 } else {
492 inner.lazy_wakers.push(cx.waker().clone());
494
495 return Poll::Pending;
497 }
498 }
499
500 ready = &[];
503
504 let waker = match inner.waker.push(cx.waker().clone()) {
505 Some(w) => w,
506 None => {
507 return Poll::Pending;
509 }
510 };
511
512 let min_i = inner.clones.iter().copied().min().unwrap();
513 if min_i > 0 {
514 inner.buf.copy_within(min_i.., 0);
516 inner.buf.truncate(inner.buf.len() - min_i);
517
518 i -= min_i;
519 for i in &mut inner.clones {
520 *i -= min_i;
521 }
522 }
523
524 let new_start = inner.buf.len();
525
526 inner.buf.resize(inner.buf.len() + buf.len().max(10.kilobytes().0 as usize), 0);
527
528 let mut inner_cx = task::Context::from_waker(&waker);
529
530 let source = unsafe { Pin::new_unchecked(inner.source.as_mut().unwrap()) };
532 let result = source.poll_read(&mut inner_cx, &mut inner.buf[new_start..]);
533
534 match result {
535 Poll::Ready(result) => {
536 for waker in inner.lazy_wakers.drain(..) {
538 waker.wake();
539 }
540
541 match result {
542 Ok(0) => {
543 inner.waker.cancel();
544
545 inner.buf.truncate(new_start);
547 inner.result = ReadState::Eof;
548 inner.source = None;
549
550 }
552 Ok(read) => {
553 inner.waker.cancel();
554
555 inner.buf.truncate(new_start + read);
557 ready = &inner.buf[i..];
558
559 }
561 Err(e) => {
562 inner.waker.cancel();
563
564 inner.result = ReadState::Err(CloneableError::new(&e));
566 inner.buf = vec![];
567 inner.source = None;
568
569 return Poll::Ready(Err(e));
570 }
571 }
572 }
573
574 Poll::Pending => {
575 inner.buf.truncate(new_start);
576 return Poll::Pending;
577 }
578 }
579 }
580 }
581 ReadState::Eof => {
582 ready = &inner.buf[i..];
583
584 }
586 ReadState::Err(e) => return Poll::Ready(e.err()),
587 }
588
589 let max_ready = buf.len().min(ready.len());
592 buf[..max_ready].copy_from_slice(&ready[..max_ready]);
593
594 i += max_ready;
595 inner.clones[self_.index] = i;
596
597 Poll::Ready(Ok(max_ready))
598 }
599}
600
601#[derive(Clone)]
612pub struct CloneableError {
613 info: ErrorInfo,
614}
615#[derive(Clone)]
616enum ErrorInfo {
617 OsError(i32),
618 Other(ErrorKind, String),
619}
620impl CloneableError {
621 pub fn new(e: &Error) -> Self {
623 let info = if let Some(code) = e.raw_os_error() {
624 ErrorInfo::OsError(code)
625 } else {
626 ErrorInfo::Other(e.kind(), format!("{e}"))
627 };
628
629 Self { info }
630 }
631
632 pub fn err<T>(&self) -> Result<T> {
634 Err(self.clone().into())
635 }
636}
637impl From<CloneableError> for Error {
638 fn from(e: CloneableError) -> Self {
639 match e.info {
640 ErrorInfo::OsError(code) => Error::from_raw_os_error(code),
641 ErrorInfo::Other(kind, msg) => Error::new(kind, msg),
642 }
643 }
644}
645
646pub struct ReadLimited<S> {
650 source: S,
651 limit: usize,
652 on_limit: fn() -> std::io::Error,
653}
654impl<S> ReadLimited<S> {
655 pub fn new(source: S, limit: ByteLength, on_limit: fn() -> std::io::Error) -> Self {
659 Self {
660 source,
661 limit: limit.0.try_into().unwrap_or(usize::MAX),
662 on_limit,
663 }
664 }
665
666 pub fn new_default_err(source: S, limit: ByteLength) -> Self {
668 Self::new(source, limit, || {
669 std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "source exceeded read limit")
670 })
671 }
672}
673impl<S> AsyncRead for ReadLimited<S>
674where
675 S: AsyncRead,
676{
677 fn poll_read(self: Pin<&mut Self>, cx: &mut task::Context<'_>, mut buf: &mut [u8]) -> Poll<Result<usize>> {
678 let self_ = unsafe { self.get_unchecked_mut() };
680
681 if self_.limit == 0 {
682 let err = (self_.on_limit)();
683 return Poll::Ready(Err(err));
684 }
685
686 if buf.len() > self_.limit {
687 buf = &mut buf[..self_.limit];
688 }
689
690 match unsafe { Pin::new_unchecked(&mut self_.source) }.poll_read(cx, buf) {
692 Poll::Ready(Ok(n)) => {
693 self_.limit = self_.limit.saturating_sub(n);
694 Poll::Ready(Ok(n))
695 }
696 r => r,
697 }
698 }
699}
700impl<S> AsyncBufRead for ReadLimited<S>
701where
702 S: AsyncBufRead,
703{
704 fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Result<&[u8]>> {
705 let self_ = unsafe { self.get_unchecked_mut() };
707
708 if self_.limit == 0 {
709 let err = (self_.on_limit)();
710 return Poll::Ready(Err(err));
711 }
712
713 unsafe { Pin::new_unchecked(&mut self_.source) }.poll_fill_buf(cx)
715 }
716
717 fn consume(self: Pin<&mut Self>, amt: usize) {
718 let self_ = unsafe { self.get_unchecked_mut() };
720 unsafe { Pin::new_unchecked(&mut self_.source) }.consume(amt);
722 self_.limit = self_.limit.saturating_sub(amt);
723 }
724}
725impl<S> Read for ReadLimited<S>
726where
727 S: Read,
728{
729 fn read(&mut self, mut buf: &mut [u8]) -> Result<usize> {
730 if self.limit == 0 {
731 let err = (self.on_limit)();
732 return Err(err);
733 }
734
735 if buf.len() > self.limit {
736 buf = &mut buf[..self.limit];
737 }
738
739 match self.source.read(buf) {
740 Ok(n) => {
741 self.limit = self.limit.saturating_sub(n);
742 Ok(n)
743 }
744 r => r,
745 }
746 }
747}
748impl<S> BufRead for ReadLimited<S>
749where
750 S: BufRead,
751{
752 fn fill_buf(&mut self) -> Result<&[u8]> {
753 if self.limit == 0 {
754 let err = (self.on_limit)();
755 return Err(err);
756 }
757
758 self.source.fill_buf()
759 }
760
761 fn consume(&mut self, amount: usize) {
762 self.source.consume(amount);
763 self.limit = self.limit.saturating_sub(amount);
764 }
765}
766
767enum ReadState {
768 Running,
769 Eof,
770 Err(CloneableError),
771}
772
773#[cfg(test)]
774mod tests {
775 use super::*;
776 use crate as task;
777 use zng_unit::TimeUnits;
778
779 #[test]
780 pub fn mc_buf_reader_parallel() {
781 let data = Data::new(60.kilobytes().0 as _);
782
783 let mut expected = vec![0; data.len];
784 let _ = data.clone().blocking_read(&mut expected[..]);
785
786 let mut a = McBufReader::new(data);
787 let mut b = a.clone();
788 let mut c = a.clone();
789
790 let (a, b, c) = async_test(async move {
791 let a = task::run(async move {
792 let mut buf = vec![];
793 a.read_to_end(&mut buf).await.unwrap();
794 buf
795 });
796 let b = task::run(async move {
797 let mut buf: Vec<u8> = vec![];
798 b.read_to_end(&mut buf).await.unwrap();
799 buf
800 });
801 let c = task::run(async move {
802 let mut buf: Vec<u8> = vec![];
803 c.read_to_end(&mut buf).await.unwrap();
804 buf
805 });
806
807 task::all!(a, b, c).await
808 });
809
810 crate::assert_vec_eq!(expected, a);
811 crate::assert_vec_eq!(expected, b);
812 crate::assert_vec_eq!(expected, c);
813 }
814
815 #[test]
816 pub fn mc_buf_reader_single() {
817 let data = Data::new(60.kilobytes().0 as _);
818
819 let mut expected = vec![0; data.len];
820 let _ = data.clone().blocking_read(&mut expected[..]);
821
822 let mut a = McBufReader::new(data);
823
824 let a = async_test(async move {
825 let a = task::run(async move {
826 let mut buf = vec![];
827 a.read_to_end(&mut buf).await.unwrap();
828 buf
829 });
830
831 a.await
832 });
833
834 crate::assert_vec_eq!(expected, a);
835 }
836
837 #[test]
838 pub fn mc_buf_reader_sequential() {
839 let data = Data::new(60.kilobytes().0 as _);
840
841 let mut expected = vec![0; data.len];
842 let _ = data.clone().blocking_read(&mut expected[..]);
843
844 let mut clones = vec![McBufReader::new(data)];
845 for _ in 0..5 {
846 clones.push(clones[0].clone());
847 }
848
849 let r = async_test(async move {
850 let mut r = vec![];
851
852 for mut clone in clones {
853 let mut buf = vec![];
854 clone.read_to_end(&mut buf).await.unwrap();
855 r.push(buf);
856 }
857
858 r
859 });
860
861 for r in r {
862 crate::assert_vec_eq!(expected, r);
863 }
864 }
865
866 #[test]
867 pub fn mc_buf_reader_completed() {
868 let data = Data::new(60.kilobytes().0 as _);
869 let mut buf = Vec::with_capacity(data.len);
870 let mut a = McBufReader::new(data);
871
872 let r = async_test(async move {
873 a.read_to_end(&mut buf).await.unwrap();
874
875 let mut b = a.clone();
876 buf.clear();
877
878 b.read_to_end(&mut buf).await.unwrap();
879 buf.len()
880 });
881
882 assert_eq!(0, r);
883 }
884
885 #[test]
886 pub fn mc_buf_reader_error() {
887 let mut data = Data::new(20.kilobytes().0 as _);
888 data.set_error();
889
890 let mut expected = vec![0; data.len];
891 let _ = data.clone().blocking_read(&mut expected[..]);
892
893 let mut a = McBufReader::new(data);
894 let mut b = a.clone();
895
896 let (a, b) = async_test(async move {
897 let a = task::run(async move {
898 let mut buf = vec![];
899 a.read_to_end(&mut buf).await.unwrap_err()
900 });
901 let b = task::run(async move {
902 let mut buf: Vec<u8> = vec![];
903 b.read_to_end(&mut buf).await.unwrap_err()
904 });
905
906 task::all!(a, b).await
907 });
908
909 assert_eq!(ErrorKind::InvalidData, a.kind());
910 assert_eq!(ErrorKind::InvalidData, b.kind());
911 }
912
913 #[test]
914 pub fn mc_buf_reader_error_completed() {
915 let mut data = Data::new(20.kilobytes().0 as _);
916 data.set_error();
917
918 let mut buf = Vec::with_capacity(data.len);
919 let mut a = McBufReader::new(data);
920
921 let (a, b) = async_test(async move {
922 let a_err = a.read_to_end(&mut buf).await.unwrap_err();
923
924 let mut b = a.clone();
925 buf.clear();
926
927 let b_err = b.read_to_end(&mut buf).await.unwrap_err();
928
929 (a_err, b_err)
930 });
931
932 assert_eq!(ErrorKind::InvalidData, a.kind());
933 assert_eq!(ErrorKind::InvalidData, b.kind());
934 }
935
936 #[test]
937 pub fn mc_buf_reader_parallel_with_delay1() {
938 let mut data = Data::new(60.kilobytes().0 as _);
939 data.enable_pending();
940
941 let mut expected = vec![0; data.len];
942 let _ = data.clone().blocking_read(&mut expected[..]);
943
944 let mut a = McBufReader::new(data);
945 let mut b = a.clone();
946 let mut c = a.clone();
947
948 let (a, b, c) = async_test(async move {
949 let a = task::run(async move {
950 let mut buf = vec![];
951 a.read_to_end(&mut buf).await.unwrap();
952 buf
953 });
954 let b = task::run(async move {
955 let mut buf: Vec<u8> = vec![];
956 b.read_to_end(&mut buf).await.unwrap();
957 buf
958 });
959 let c = task::run(async move {
960 let mut buf: Vec<u8> = vec![];
961 c.read_to_end(&mut buf).await.unwrap();
962 buf
963 });
964
965 task::all!(a, b, c).await
966 });
967
968 crate::assert_vec_eq!(expected, a);
969 crate::assert_vec_eq!(expected, b);
970 crate::assert_vec_eq!(expected, c);
971 }
972
973 #[test]
974 pub fn mc_buf_reader_parallel_with_delay2() {
975 let mut data = Data::new(60.kilobytes().0 as _);
976 data.enable_pending();
977
978 let mut expected = vec![0; data.len];
979 let _ = data.clone().blocking_read(&mut expected[..]);
980
981 let mut a = McBufReader::new(data);
982 let mut b = a.clone();
983 let mut c = a.clone();
984
985 let (a, b, c) = async_test(async move {
986 let a = task::run(async move {
987 let mut buf = vec![];
988 a.read_to_end(&mut buf).await.unwrap();
989 buf
990 });
991 let b = task::run(async move {
992 let mut buf: Vec<u8> = vec![];
993 task::deadline(5.ms()).await;
994 b.read_to_end(&mut buf).await.unwrap();
995 buf
996 });
997 let c = task::run(async move {
998 let mut buf: Vec<u8> = vec![];
999 c.read_to_end(&mut buf).await.unwrap();
1000 buf
1001 });
1002
1003 task::all!(a, b, c).await
1004 });
1005
1006 crate::assert_vec_eq!(expected, a);
1007 crate::assert_vec_eq!(expected, b);
1008 crate::assert_vec_eq!(expected, c);
1009 }
1010
1011 #[derive(Clone)]
1012 struct Data {
1013 b: u8,
1014 len: usize,
1015 error: Option<CloneableError>,
1016 delay: Duration,
1017 pending: bool,
1018 }
1019 impl Data {
1020 pub fn new(len: usize) -> Self {
1021 Self {
1022 b: 0,
1023 len,
1024 error: None,
1025 delay: 0.ms(),
1026 pending: false,
1027 }
1028 }
1029 pub fn blocking_read(&mut self, buf: &mut [u8]) -> Result<usize> {
1030 let len = self.len;
1031 for b in buf.iter_mut().take(len) {
1032 *b = self.b;
1033 self.len -= 1;
1034 self.b = self.b.wrapping_add(1);
1035 }
1036
1037 if len == 0
1038 && let Some(e) = &self.error
1039 {
1040 return e.err();
1041 }
1042
1043 Ok(buf.len().min(len))
1044 }
1045 pub fn set_error(&mut self) {
1046 self.error = Some(CloneableError::new(&Error::new(ErrorKind::InvalidData, "test error")));
1047 }
1048
1049 pub fn enable_pending(&mut self) {
1050 self.delay = 3.ms();
1051 }
1052 }
1053 impl AsyncRead for Data {
1054 fn poll_read(mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>, buf: &mut [u8]) -> Poll<Result<usize>> {
1055 if self.delay > Duration::ZERO {
1056 self.pending = !self.pending;
1057 if self.pending {
1058 let waker = cx.waker().clone();
1059 let delay = self.delay;
1060 task::spawn(async move {
1061 task::deadline(delay).await;
1062 waker.wake();
1063 });
1064 return Poll::Pending;
1065 }
1066 }
1067
1068 let r = self.as_mut().blocking_read(buf);
1069 Poll::Ready(r)
1070 }
1071 }
1072
1073 #[track_caller]
1074 fn async_test<F>(test: F) -> F::Output
1075 where
1076 F: Future,
1077 {
1078 task::block_on(task::with_deadline(test, 5.secs())).unwrap()
1079 }
1080
1081 #[macro_export]
1083 macro_rules! assert_vec_eq {
1084 ($a:expr, $b: expr) => {
1085 match (&$a, &$b) {
1086 (ref a, ref b) => {
1087 let len_not_eq = a.len() != b.len();
1088 let mut data_not_eq = None;
1089 for (i, (a, b)) in a.iter().zip(b.iter()).enumerate() {
1090 if a != b {
1091 data_not_eq = Some(i);
1092 break;
1093 }
1094 }
1095
1096 if len_not_eq || data_not_eq.is_some() {
1097 use std::fmt::*;
1098
1099 let mut error = format!("`{}` != `{}`", stringify!($a), stringify!($b));
1100 if len_not_eq {
1101 let _ = write!(&mut error, "\n lengths not equal: {} != {}", a.len(), b.len());
1102 }
1103 if let Some(i) = data_not_eq {
1104 let _ = write!(&mut error, "\n data not equal at index {}: {} != {:?}", i, a[i], b[i]);
1105 }
1106 panic!("{error}")
1107 }
1108 }
1109 }
1110 };
1111 }
1112}