Skip to main content

zng_task/
io.rs

1//! IO tasks.
2//!
3//! Most of the types in this module are re-exported from [`futures_lite::io`].
4//!
5//! [`futures_lite::io`]: https://docs.rs/futures-lite/latest/futures_lite/io/index.html
6
7use 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};
23use parking_lot::Mutex;
24use std::io::{Error, Result};
25use zng_time::{DInstant, INSTANT};
26use zng_txt::formatx;
27use zng_unit::{ByteLength, ByteUnits};
28use zng_var::{Var, impl_from_and_into_var, var};
29
30struct MeasureInner {
31    metrics: Var<Metrics>,
32    start_time: DInstant,
33    last_write: DInstant,
34    last_read: DInstant,
35}
36impl MeasureInner {
37    fn new(read_progress: (ByteLength, ByteLength), write_progress: (ByteLength, ByteLength)) -> Self {
38        let now = INSTANT.now();
39        Self {
40            metrics: var(Metrics {
41                read_progress,
42                read_speed: 0.bytes(),
43                write_progress,
44                write_speed: 0.bytes(),
45                total_time: Duration::ZERO,
46            }),
47            start_time: now,
48            last_write: now,
49            last_read: now,
50        }
51    }
52
53    fn on_read(&mut self, bytes: u64) {
54        if bytes == 0 {
55            return;
56        }
57
58        let bytes = bytes.bytes();
59
60        let now = INSTANT.now();
61        let elapsed = now - self.last_read;
62
63        self.last_read = now;
64        let read_speed = bytes_per_sec(bytes, elapsed);
65
66        let total_time = now - self.start_time;
67
68        self.metrics.modify(move |m| {
69            m.read_progress.0 += bytes;
70            m.read_speed = read_speed;
71            m.total_time = total_time;
72        });
73    }
74
75    fn on_write(&mut self, bytes: u64) {
76        if bytes == 0 {
77            return;
78        }
79
80        let bytes = bytes.bytes();
81
82        let now = INSTANT.now();
83        let elapsed = now - self.last_write;
84
85        self.last_write = now;
86        let write_speed = bytes_per_sec(bytes, elapsed);
87
88        let total_time = now - self.start_time;
89
90        self.metrics.modify(move |m| {
91            m.write_progress.0 += bytes;
92            m.write_speed = write_speed;
93            m.total_time = total_time;
94        });
95    }
96}
97
98/// Measure read/write task.
99///
100/// Metrics are updated after each read/write, if you read/write all bytes in one call
101/// the metrics will only update once.
102pub struct Measure<T> {
103    task: T,
104    inner: MeasureInner,
105}
106impl<T> Measure<T> {
107    /// Start measuring a new read/write task.
108    pub fn new(task: T, total_read: ByteLength, total_write: ByteLength) -> Self {
109        Self::new_ongoing(task, (0.bytes(), total_read), (0.bytes(), total_write))
110    }
111
112    /// Continue measuring a read/write task.
113    pub fn new_ongoing(task: T, read_progress: (ByteLength, ByteLength), write_progress: (ByteLength, ByteLength)) -> Self {
114        Measure {
115            task,
116            inner: MeasureInner::new(read_progress, write_progress),
117        }
118    }
119
120    /// Current metrics.
121    ///
122    /// This value is updated after every read/write.
123    pub fn metrics(&self) -> Var<Metrics> {
124        self.inner.metrics.read_only()
125    }
126
127    /// Unwrap the inner task and final metrics.
128    pub fn finish(self) -> (T, Metrics) {
129        let mut metrics = self.inner.metrics.get();
130        metrics.total_time = self.inner.start_time.elapsed();
131        (self.task, metrics)
132    }
133}
134
135fn bytes_per_sec(bytes: ByteLength, elapsed: Duration) -> ByteLength {
136    let bytes_per_sec = bytes.0 as u128 / elapsed.as_nanos() / Duration::from_secs(1).as_nanos();
137    ByteLength(bytes_per_sec as u64)
138}
139
140impl<T: AsyncRead> AsyncRead for Measure<T> {
141    fn poll_read(self: Pin<&mut Self>, cx: &mut task::Context<'_>, buf: &mut [u8]) -> Poll<Result<usize>> {
142        // SAFETY: we don't move anything.
143        let self_ = unsafe { self.get_unchecked_mut() };
144
145        // SAFETY: we don't move task
146        match unsafe { Pin::new_unchecked(&mut self_.task) }.poll_read(cx, buf) {
147            Poll::Ready(Ok(bytes)) => {
148                self_.inner.on_read(bytes as u64);
149                Poll::Ready(Ok(bytes))
150            }
151            p => p,
152        }
153    }
154}
155impl<T: AsyncWrite> AsyncWrite for Measure<T> {
156    fn poll_write(self: Pin<&mut Self>, cx: &mut task::Context<'_>, buf: &[u8]) -> Poll<Result<usize>> {
157        // SAFETY: we don't move anything.
158        let self_ = unsafe { self.get_unchecked_mut() };
159
160        // SAFETY: we don't move task
161        match unsafe { Pin::new_unchecked(&mut self_.task) }.poll_write(cx, buf) {
162            Poll::Ready(Ok(bytes)) => {
163                self_.inner.on_write(bytes as u64);
164                Poll::Ready(Ok(bytes))
165            }
166            p => p,
167        }
168    }
169
170    fn poll_flush(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Result<()>> {
171        // SAFETY: we don't move anything.
172        let self_ = unsafe { self.get_unchecked_mut() };
173
174        // SAFETY: we don't move task
175        unsafe { Pin::new_unchecked(&mut self_.task) }.poll_flush(cx)
176    }
177
178    fn poll_close(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Result<()>> {
179        // SAFETY: we don't move anything.
180        let self_ = unsafe { self.get_unchecked_mut() };
181
182        // SAFETY: we don't move task
183        unsafe { Pin::new_unchecked(&mut self_.task) }.poll_flush(cx)
184    }
185}
186impl<T: AsyncBufRead> AsyncBufRead for Measure<T> {
187    fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Result<&[u8]>> {
188        // SAFETY: we don't move anything.
189        let self_ = unsafe { self.get_unchecked_mut() };
190
191        // SAFETY: we don't move task
192        unsafe { Pin::new_unchecked(&mut self_.task) }.poll_fill_buf(cx)
193    }
194
195    fn consume(self: Pin<&mut Self>, amt: usize) {
196        // SAFETY: we don't move anything.
197        let self_ = unsafe { self.get_unchecked_mut() };
198        // SAFETY: we don't move task
199        unsafe { Pin::new_unchecked(&mut self_.task) }.consume(amt);
200        self_.inner.on_read(amt as u64);
201    }
202}
203impl<T: Read> Read for Measure<T> {
204    fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
205        match self.task.read(buf) {
206            Ok(bytes) => {
207                self.inner.on_read(bytes as u64);
208                Ok(bytes)
209            }
210            r => r,
211        }
212    }
213}
214impl<T: Write> Write for Measure<T> {
215    fn write(&mut self, buf: &[u8]) -> Result<usize> {
216        match self.task.write(buf) {
217            Ok(bytes) => {
218                self.inner.on_write(bytes as u64);
219                Ok(bytes)
220            }
221            r => r,
222        }
223    }
224
225    fn flush(&mut self) -> Result<()> {
226        self.task.flush()
227    }
228}
229impl<T: BufRead> BufRead for Measure<T> {
230    fn fill_buf(&mut self) -> Result<&[u8]> {
231        self.task.fill_buf()
232    }
233
234    fn consume(&mut self, amount: usize) {
235        self.task.consume(amount);
236        self.inner.on_read(amount as u64);
237    }
238}
239
240/// Information about the state of an async IO task.
241///
242/// Read is also called *receive* or *download*. Write is also called *send* or *upload*. The default
243/// display print uses arrows ↓ and ↑ for read and write.
244///
245/// Use [`Measure`] to measure a task.
246#[derive(Debug, Clone, PartialEq, Eq)]
247#[non_exhaustive]
248pub struct Metrics {
249    /// Number of bytes read / estimated total.
250    pub read_progress: (ByteLength, ByteLength),
251
252    /// Average read speed in bytes/second.
253    pub read_speed: ByteLength,
254
255    /// Number of bytes written / estimated total.
256    pub write_progress: (ByteLength, ByteLength),
257
258    /// Average write speed in bytes/second.
259    pub write_speed: ByteLength,
260
261    /// Total time for the entire task. This will continuously increase until
262    /// the task is finished.
263    pub total_time: Duration,
264}
265impl Metrics {
266    /// All zeros.
267    pub fn zero() -> Self {
268        Self {
269            read_progress: (0.bytes(), 0.bytes()),
270            read_speed: 0.bytes(),
271            write_progress: (0.bytes(), 0.bytes()),
272            write_speed: 0.bytes(),
273            total_time: Duration::ZERO,
274        }
275    }
276}
277impl fmt::Display for Metrics {
278    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
279        let mut nl = false;
280        if self.read_progress.1 > 0.bytes() {
281            nl = true;
282            if self.read_progress.0 != self.read_progress.1 {
283                if self.read_progress.0 <= self.read_progress.1 {
284                    write!(f, "↓ {}-{}, {}/s", self.read_progress.0, self.read_progress.1, self.read_speed)?;
285                } else {
286                    write!(f, "↓ {}, {}/s", self.read_progress.0, self.read_speed)?;
287                }
288            } else {
289                write!(f, "↓ {} . {:?}", self.read_progress.0, self.total_time)?;
290            }
291        }
292        if self.write_progress.1 > 0.bytes() {
293            if nl {
294                writeln!(f)?;
295            }
296            if self.write_progress.0 != self.write_progress.1 {
297                if self.write_progress.0 <= self.write_progress.1 {
298                    write!(f, "↑ {}-{}, {}/s", self.write_progress.0, self.write_progress.1, self.write_speed)?;
299                } else {
300                    write!(f, "↑ {}, {}/s", self.write_progress.0, self.write_speed)?;
301                }
302            } else {
303                write!(f, "↑ {} . {:?}", self.write_progress.0, self.total_time)?;
304            }
305        }
306
307        Ok(())
308    }
309}
310impl_from_and_into_var! {
311    fn from(metrics: Metrics) -> Progress {
312        let mut status = Progress::indeterminate();
313        if metrics.read_progress.1 > 0.bytes() {
314            status = Progress::from_n_of(metrics.read_progress.0.0, metrics.read_progress.1.0);
315        }
316        if metrics.write_progress.1 > 0.bytes() {
317            let w_status = Progress::from_n_of(metrics.write_progress.0.0, metrics.write_progress.1.0);
318            if status.is_indeterminate() {
319                status = w_status;
320            } else {
321                status = status.and_fct(w_status.fct());
322            }
323        }
324        status.with_msg(formatx!("{metrics}")).with_meta_mut(|mut m| {
325            m.set(*METRICS_ID, metrics);
326        })
327    }
328}
329
330zng_state_map::static_id! {
331    /// Metrics in a [`Progress::with_meta`] metadata.
332    pub static ref METRICS_ID: zng_state_map::StateId<Metrics>;
333}
334
335/// Extension methods for [`std::io::Error`] to be used with errors returned by [`McBufReader`].
336pub trait McBufErrorExt {
337    /// Returns `true` if this error represents the condition where there are only [`McBufReader::is_lazy`] readers
338    /// left, the buffer is drained and the inner reader is not EOF.
339    ///
340    /// You can recover from this error by turning the reader non-lazy using [`McBufReader::set_lazy`].
341    fn is_only_lazy_left(&self) -> bool;
342}
343impl McBufErrorExt for std::io::Error {
344    fn is_only_lazy_left(&self) -> bool {
345        matches!(self.kind(), ErrorKind::Other) && format!("{self:?}").contains(ONLY_NON_LAZY_ERROR_MSG)
346    }
347}
348const ONLY_NON_LAZY_ERROR_MSG: &str = "no non-lazy readers left to read";
349
350/// Multiple consumer buffered read.
351///
352/// Clone an instance to create a new consumer, already read bytes stay in the buffer until all clones have read it,
353/// clones continue reading from the same offset as the reader they cloned.
354///
355/// A single instance of this reader behaves like a `BufReader`.
356///
357/// # Result
358///
359/// The result is *repeats* ready when `EOF` or an [`Error`] occurs, unfortunately the IO error is not cloneable
360/// so the error is recreated using [`CloneableError`] for subsequent poll attempts.
361///
362/// The inner reader is dropped as soon as it finishes.
363///
364/// # Lazy Clones
365///
366/// You can mark clones as [lazy], lazy clones don't pull from the inner reader, only advance when another clone reads, if
367/// all living clones are lazy they stop reading with an error. You can identify this custom error using the [`McBufErrorExt::is_only_lazy_left`]
368/// extension method.
369///
370/// [lazy]: Self::set_lazy
371pub struct McBufReader<S: AsyncRead> {
372    inner: Arc<Mutex<McBufInner<S>>>,
373    index: usize,
374    lazy: bool,
375}
376struct McBufInner<S: AsyncRead> {
377    source: Option<S>,
378    waker: McWaker,
379    lazy_wakers: Vec<task::Waker>,
380
381    buf: Vec<u8>,
382
383    clones: Vec<usize>,
384    non_lazy_count: usize,
385
386    result: ReadState,
387}
388impl<S: AsyncRead> McBufReader<S> {
389    /// Creates a buffered reader.
390    pub fn new(source: S) -> Self {
391        let mut clones = Vec::with_capacity(2);
392        clones.push(0);
393        McBufReader {
394            inner: Arc::new(Mutex::new(McBufInner {
395                source: Some(source),
396                waker: McWaker::empty(),
397                lazy_wakers: vec![],
398
399                buf: Vec::with_capacity(10.kilobytes().0 as usize),
400
401                clones,
402                non_lazy_count: 1,
403
404                result: ReadState::Running,
405            })),
406            index: 0,
407            lazy: false,
408        }
409    }
410
411    /// Returns `true` if this reader does not pull from the inner reader, only advancing when a non-lazy reader advances.
412    ///
413    /// The initial reader is not lazy, only clones of lazy readers are lazy by default.
414    pub fn is_lazy(&self) -> bool {
415        self.lazy
416    }
417
418    /// Sets [`is_lazy`].
419    ///
420    /// [`is_lazy`]: Self::is_lazy
421    pub fn set_lazy(&mut self, lazy: bool) {
422        if self.lazy != lazy {
423            if lazy {
424                self.inner.lock().non_lazy_count -= 1;
425            } else {
426                self.inner.lock().non_lazy_count += 1;
427            }
428            self.lazy = lazy;
429        }
430    }
431}
432impl<S: AsyncRead> Clone for McBufReader<S> {
433    fn clone(&self) -> Self {
434        let mut inner = self.inner.lock();
435
436        let offset = inner.clones[self.index];
437        let index = inner.clones.len();
438        inner.clones.push(offset);
439
440        if !self.lazy {
441            inner.non_lazy_count += 1;
442        }
443
444        Self {
445            inner: self.inner.clone(),
446            index,
447            lazy: self.lazy,
448        }
449    }
450}
451impl<S: AsyncRead> Drop for McBufReader<S> {
452    fn drop(&mut self) {
453        let mut inner = self.inner.lock();
454        inner.clones[self.index] = usize::MAX;
455        if !self.lazy {
456            inner.non_lazy_count -= 1;
457            if inner.non_lazy_count == 0 {
458                // notify lazy so they get the error.
459                for waker in inner.lazy_wakers.drain(..) {
460                    waker.wake();
461                }
462            }
463        }
464    }
465}
466impl<S: AsyncRead> AsyncRead for McBufReader<S> {
467    fn poll_read(self: Pin<&mut Self>, cx: &mut task::Context<'_>, buf: &mut [u8]) -> Poll<Result<usize>> {
468        let self_ = self.as_ref();
469        let mut inner = self_.inner.lock();
470        let inner = &mut *inner;
471
472        // ready data for this clone.
473        let mut i = inner.clones[self_.index];
474        let mut ready;
475
476        match &inner.result {
477            ReadState::Running => {
478                // source has not finished yet.
479
480                ready = &inner.buf[i..];
481
482                if ready.is_empty() {
483                    if self.lazy {
484                        if inner.non_lazy_count == 0 {
485                            // user can make this reader non-lazy and try again.
486                            return Poll::Ready(Err(Error::other(ONLY_NON_LAZY_ERROR_MSG)));
487                        } else {
488                            // register waker for after non-lazy poll.
489                            inner.lazy_wakers.push(cx.waker().clone());
490
491                            // wait non-lazy to pull.
492                            return Poll::Pending;
493                        }
494                    }
495
496                    // time to poll source.
497
498                    ready = &[];
499
500                    let waker = match inner.waker.push(cx.waker().clone()) {
501                        Some(w) => w,
502                        None => {
503                            // already polling from another clone.
504                            return Poll::Pending;
505                        }
506                    };
507
508                    let min_i = inner.clones.iter().copied().min().unwrap();
509                    if min_i > 0 {
510                        // reuse front.
511                        inner.buf.copy_within(min_i.., 0);
512                        inner.buf.truncate(inner.buf.len() - min_i);
513
514                        i -= min_i;
515                        for i in &mut inner.clones {
516                            *i -= min_i;
517                        }
518                    }
519
520                    let new_start = inner.buf.len();
521
522                    inner.buf.resize(inner.buf.len() + buf.len().max(10.kilobytes().0 as usize), 0);
523
524                    let mut inner_cx = task::Context::from_waker(&waker);
525
526                    // SAFETY: we don't move `source`.
527                    let source = unsafe { Pin::new_unchecked(inner.source.as_mut().unwrap()) };
528                    let result = source.poll_read(&mut inner_cx, &mut inner.buf[new_start..]);
529
530                    match result {
531                        Poll::Ready(result) => {
532                            // notify lazy readers.
533                            for waker in inner.lazy_wakers.drain(..) {
534                                waker.wake();
535                            }
536
537                            match result {
538                                Ok(0) => {
539                                    inner.waker.cancel();
540
541                                    // EOF
542                                    inner.buf.truncate(new_start);
543                                    inner.result = ReadState::Eof;
544                                    inner.source = None;
545
546                                    // continue 'copy ready
547                                }
548                                Ok(read) => {
549                                    inner.waker.cancel();
550
551                                    // Read > 0
552                                    inner.buf.truncate(new_start + read);
553                                    ready = &inner.buf[i..];
554
555                                    // continue 'copy ready
556                                }
557                                Err(e) => {
558                                    inner.waker.cancel();
559
560                                    // Error
561                                    inner.result = ReadState::Err(CloneableError::new(&e));
562                                    inner.buf = vec![];
563                                    inner.source = None;
564
565                                    return Poll::Ready(Err(e));
566                                }
567                            }
568                        }
569
570                        Poll::Pending => {
571                            inner.buf.truncate(new_start);
572                            return Poll::Pending;
573                        }
574                    }
575                }
576            }
577            ReadState::Eof => {
578                ready = &inner.buf[i..];
579
580                // continue 'copy ready
581            }
582            ReadState::Err(e) => return Poll::Ready(e.err()),
583        }
584
585        // 'copy ready
586
587        let max_ready = buf.len().min(ready.len());
588        buf[..max_ready].copy_from_slice(&ready[..max_ready]);
589
590        i += max_ready;
591        inner.clones[self_.index] = i;
592
593        Poll::Ready(Ok(max_ready))
594    }
595}
596
597/// Represents the cloneable parts of an [`Error`].
598///
599/// Unfortunately [`Error`] does not implement clone, this is needed to implemented
600/// IO futures that repeat the ready result after subsequent polls. This type partially
601/// works around the issue by copying enough information to recreate an error that is still useful.
602///
603/// The OS error code, [`ErrorKind`] and display message are preserved. Note that this not an error type,
604/// it must be converted to [`Error`] using `into` or [`err`].
605///
606/// [`err`]: Self::err
607#[derive(Clone)]
608pub struct CloneableError {
609    info: ErrorInfo,
610}
611#[derive(Clone)]
612enum ErrorInfo {
613    OsError(i32),
614    Other(ErrorKind, String),
615}
616impl CloneableError {
617    /// Copy the cloneable information from the [`Error`].
618    pub fn new(e: &Error) -> Self {
619        let info = if let Some(code) = e.raw_os_error() {
620            ErrorInfo::OsError(code)
621        } else {
622            ErrorInfo::Other(e.kind(), format!("{e}"))
623        };
624
625        Self { info }
626    }
627
628    /// Returns an `Err(Error)` generated from the cloneable information.
629    pub fn err<T>(&self) -> Result<T> {
630        Err(self.clone().into())
631    }
632}
633impl From<CloneableError> for Error {
634    fn from(e: CloneableError) -> Self {
635        match e.info {
636            ErrorInfo::OsError(code) => Error::from_raw_os_error(code),
637            ErrorInfo::Other(kind, msg) => Error::new(kind, msg),
638        }
639    }
640}
641
642/// Represents a stream reader that generates an error if the source stream exceeds a limit.
643///
644/// Note that some bytes over the limit may be read once if the source stream is buffered.
645pub struct ReadLimited<S> {
646    source: S,
647    limit: usize,
648    on_limit: fn() -> std::io::Error,
649}
650impl<S> ReadLimited<S> {
651    /// Construct a limited reader.
652    ///
653    /// The `on_limit` closure is called for every read attempt after the limit is reached.
654    pub fn new(source: S, limit: ByteLength, on_limit: fn() -> std::io::Error) -> Self {
655        Self {
656            source,
657            limit: limit.0.try_into().unwrap_or(usize::MAX),
658            on_limit,
659        }
660    }
661
662    /// New with default on limit error.
663    pub fn new_default_err(source: S, limit: ByteLength) -> Self {
664        Self::new(source, limit, || {
665            std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "source exceeded read limit")
666        })
667    }
668}
669impl<S> AsyncRead for ReadLimited<S>
670where
671    S: AsyncRead,
672{
673    fn poll_read(self: Pin<&mut Self>, cx: &mut task::Context<'_>, mut buf: &mut [u8]) -> Poll<Result<usize>> {
674        // SAFETY: we don't move anything.
675        let self_ = unsafe { self.get_unchecked_mut() };
676
677        if self_.limit == 0 {
678            let err = (self_.on_limit)();
679            return Poll::Ready(Err(err));
680        }
681
682        if buf.len() > self_.limit {
683            buf = &mut buf[..self_.limit];
684        }
685
686        // SAFETY: we never move `source`.
687        match unsafe { Pin::new_unchecked(&mut self_.source) }.poll_read(cx, buf) {
688            Poll::Ready(Ok(n)) => {
689                self_.limit = self_.limit.saturating_sub(n);
690                Poll::Ready(Ok(n))
691            }
692            r => r,
693        }
694    }
695}
696impl<S> AsyncBufRead for ReadLimited<S>
697where
698    S: AsyncBufRead,
699{
700    fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> Poll<Result<&[u8]>> {
701        // SAFETY: we don't move anything.
702        let self_ = unsafe { self.get_unchecked_mut() };
703
704        if self_.limit == 0 {
705            let err = (self_.on_limit)();
706            return Poll::Ready(Err(err));
707        }
708
709        // SAFETY: we never move `source`.
710        unsafe { Pin::new_unchecked(&mut self_.source) }.poll_fill_buf(cx)
711    }
712
713    fn consume(self: Pin<&mut Self>, amt: usize) {
714        // SAFETY: we don't move anything.
715        let self_ = unsafe { self.get_unchecked_mut() };
716        // SAFETY: we never move `source`.
717        unsafe { Pin::new_unchecked(&mut self_.source) }.consume(amt);
718        self_.limit = self_.limit.saturating_sub(amt);
719    }
720}
721impl<S> Read for ReadLimited<S>
722where
723    S: Read,
724{
725    fn read(&mut self, mut buf: &mut [u8]) -> Result<usize> {
726        if self.limit == 0 {
727            let err = (self.on_limit)();
728            return Err(err);
729        }
730
731        if buf.len() > self.limit {
732            buf = &mut buf[..self.limit];
733        }
734
735        match self.source.read(buf) {
736            Ok(n) => {
737                self.limit = self.limit.saturating_sub(n);
738                Ok(n)
739            }
740            r => r,
741        }
742    }
743}
744impl<S> BufRead for ReadLimited<S>
745where
746    S: BufRead,
747{
748    fn fill_buf(&mut self) -> Result<&[u8]> {
749        if self.limit == 0 {
750            let err = (self.on_limit)();
751            return Err(err);
752        }
753
754        self.source.fill_buf()
755    }
756
757    fn consume(&mut self, amount: usize) {
758        self.source.consume(amount);
759        self.limit = self.limit.saturating_sub(amount);
760    }
761}
762
763enum ReadState {
764    Running,
765    Eof,
766    Err(CloneableError),
767}
768
769#[cfg(test)]
770mod tests {
771    use super::*;
772    use crate as task;
773    use zng_unit::TimeUnits;
774
775    #[test]
776    pub fn mc_buf_reader_parallel() {
777        let data = Data::new(60.kilobytes().0 as _);
778
779        let mut expected = vec![0; data.len];
780        let _ = data.clone().blocking_read(&mut expected[..]);
781
782        let mut a = McBufReader::new(data);
783        let mut b = a.clone();
784        let mut c = a.clone();
785
786        let (a, b, c) = async_test(async move {
787            let a = task::run(async move {
788                let mut buf = vec![];
789                a.read_to_end(&mut buf).await.unwrap();
790                buf
791            });
792            let b = task::run(async move {
793                let mut buf: Vec<u8> = vec![];
794                b.read_to_end(&mut buf).await.unwrap();
795                buf
796            });
797            let c = task::run(async move {
798                let mut buf: Vec<u8> = vec![];
799                c.read_to_end(&mut buf).await.unwrap();
800                buf
801            });
802
803            task::all!(a, b, c).await
804        });
805
806        crate::assert_vec_eq!(expected, a);
807        crate::assert_vec_eq!(expected, b);
808        crate::assert_vec_eq!(expected, c);
809    }
810
811    #[test]
812    pub fn mc_buf_reader_single() {
813        let data = Data::new(60.kilobytes().0 as _);
814
815        let mut expected = vec![0; data.len];
816        let _ = data.clone().blocking_read(&mut expected[..]);
817
818        let mut a = McBufReader::new(data);
819
820        let a = async_test(async move {
821            let a = task::run(async move {
822                let mut buf = vec![];
823                a.read_to_end(&mut buf).await.unwrap();
824                buf
825            });
826
827            a.await
828        });
829
830        crate::assert_vec_eq!(expected, a);
831    }
832
833    #[test]
834    pub fn mc_buf_reader_sequential() {
835        let data = Data::new(60.kilobytes().0 as _);
836
837        let mut expected = vec![0; data.len];
838        let _ = data.clone().blocking_read(&mut expected[..]);
839
840        let mut clones = vec![McBufReader::new(data)];
841        for _ in 0..5 {
842            clones.push(clones[0].clone());
843        }
844
845        let r = async_test(async move {
846            let mut r = vec![];
847
848            for mut clone in clones {
849                let mut buf = vec![];
850                clone.read_to_end(&mut buf).await.unwrap();
851                r.push(buf);
852            }
853
854            r
855        });
856
857        for r in r {
858            crate::assert_vec_eq!(expected, r);
859        }
860    }
861
862    #[test]
863    pub fn mc_buf_reader_completed() {
864        let data = Data::new(60.kilobytes().0 as _);
865        let mut buf = Vec::with_capacity(data.len);
866        let mut a = McBufReader::new(data);
867
868        let r = async_test(async move {
869            a.read_to_end(&mut buf).await.unwrap();
870
871            let mut b = a.clone();
872            buf.clear();
873
874            b.read_to_end(&mut buf).await.unwrap();
875            buf.len()
876        });
877
878        assert_eq!(0, r);
879    }
880
881    #[test]
882    pub fn mc_buf_reader_error() {
883        let mut data = Data::new(20.kilobytes().0 as _);
884        data.set_error();
885
886        let mut expected = vec![0; data.len];
887        let _ = data.clone().blocking_read(&mut expected[..]);
888
889        let mut a = McBufReader::new(data);
890        let mut b = a.clone();
891
892        let (a, b) = async_test(async move {
893            let a = task::run(async move {
894                let mut buf = vec![];
895                a.read_to_end(&mut buf).await.unwrap_err()
896            });
897            let b = task::run(async move {
898                let mut buf: Vec<u8> = vec![];
899                b.read_to_end(&mut buf).await.unwrap_err()
900            });
901
902            task::all!(a, b).await
903        });
904
905        assert_eq!(ErrorKind::InvalidData, a.kind());
906        assert_eq!(ErrorKind::InvalidData, b.kind());
907    }
908
909    #[test]
910    pub fn mc_buf_reader_error_completed() {
911        let mut data = Data::new(20.kilobytes().0 as _);
912        data.set_error();
913
914        let mut buf = Vec::with_capacity(data.len);
915        let mut a = McBufReader::new(data);
916
917        let (a, b) = async_test(async move {
918            let a_err = a.read_to_end(&mut buf).await.unwrap_err();
919
920            let mut b = a.clone();
921            buf.clear();
922
923            let b_err = b.read_to_end(&mut buf).await.unwrap_err();
924
925            (a_err, b_err)
926        });
927
928        assert_eq!(ErrorKind::InvalidData, a.kind());
929        assert_eq!(ErrorKind::InvalidData, b.kind());
930    }
931
932    #[test]
933    pub fn mc_buf_reader_parallel_with_delay1() {
934        let mut data = Data::new(60.kilobytes().0 as _);
935        data.enable_pending();
936
937        let mut expected = vec![0; data.len];
938        let _ = data.clone().blocking_read(&mut expected[..]);
939
940        let mut a = McBufReader::new(data);
941        let mut b = a.clone();
942        let mut c = a.clone();
943
944        let (a, b, c) = async_test(async move {
945            let a = task::run(async move {
946                let mut buf = vec![];
947                a.read_to_end(&mut buf).await.unwrap();
948                buf
949            });
950            let b = task::run(async move {
951                let mut buf: Vec<u8> = vec![];
952                b.read_to_end(&mut buf).await.unwrap();
953                buf
954            });
955            let c = task::run(async move {
956                let mut buf: Vec<u8> = vec![];
957                c.read_to_end(&mut buf).await.unwrap();
958                buf
959            });
960
961            task::all!(a, b, c).await
962        });
963
964        crate::assert_vec_eq!(expected, a);
965        crate::assert_vec_eq!(expected, b);
966        crate::assert_vec_eq!(expected, c);
967    }
968
969    #[test]
970    pub fn mc_buf_reader_parallel_with_delay2() {
971        let mut data = Data::new(60.kilobytes().0 as _);
972        data.enable_pending();
973
974        let mut expected = vec![0; data.len];
975        let _ = data.clone().blocking_read(&mut expected[..]);
976
977        let mut a = McBufReader::new(data);
978        let mut b = a.clone();
979        let mut c = a.clone();
980
981        let (a, b, c) = async_test(async move {
982            let a = task::run(async move {
983                let mut buf = vec![];
984                a.read_to_end(&mut buf).await.unwrap();
985                buf
986            });
987            let b = task::run(async move {
988                let mut buf: Vec<u8> = vec![];
989                task::deadline(5.ms()).await;
990                b.read_to_end(&mut buf).await.unwrap();
991                buf
992            });
993            let c = task::run(async move {
994                let mut buf: Vec<u8> = vec![];
995                c.read_to_end(&mut buf).await.unwrap();
996                buf
997            });
998
999            task::all!(a, b, c).await
1000        });
1001
1002        crate::assert_vec_eq!(expected, a);
1003        crate::assert_vec_eq!(expected, b);
1004        crate::assert_vec_eq!(expected, c);
1005    }
1006
1007    #[derive(Clone)]
1008    struct Data {
1009        b: u8,
1010        len: usize,
1011        error: Option<CloneableError>,
1012        delay: Duration,
1013        pending: bool,
1014    }
1015    impl Data {
1016        pub fn new(len: usize) -> Self {
1017            Self {
1018                b: 0,
1019                len,
1020                error: None,
1021                delay: 0.ms(),
1022                pending: false,
1023            }
1024        }
1025        pub fn blocking_read(&mut self, buf: &mut [u8]) -> Result<usize> {
1026            let len = self.len;
1027            for b in buf.iter_mut().take(len) {
1028                *b = self.b;
1029                self.len -= 1;
1030                self.b = self.b.wrapping_add(1);
1031            }
1032
1033            if len == 0
1034                && let Some(e) = &self.error
1035            {
1036                return e.err();
1037            }
1038
1039            Ok(buf.len().min(len))
1040        }
1041        pub fn set_error(&mut self) {
1042            self.error = Some(CloneableError::new(&Error::new(ErrorKind::InvalidData, "test error")));
1043        }
1044
1045        pub fn enable_pending(&mut self) {
1046            self.delay = 3.ms();
1047        }
1048    }
1049    impl AsyncRead for Data {
1050        fn poll_read(mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>, buf: &mut [u8]) -> Poll<Result<usize>> {
1051            if self.delay > Duration::ZERO {
1052                self.pending = !self.pending;
1053                if self.pending {
1054                    let waker = cx.waker().clone();
1055                    let delay = self.delay;
1056                    task::spawn(async move {
1057                        task::deadline(delay).await;
1058                        waker.wake();
1059                    });
1060                    return Poll::Pending;
1061                }
1062            }
1063
1064            let r = self.as_mut().blocking_read(buf);
1065            Poll::Ready(r)
1066        }
1067    }
1068
1069    #[track_caller]
1070    fn async_test<F>(test: F) -> F::Output
1071    where
1072        F: Future,
1073    {
1074        task::block_on(task::with_deadline(test, 5.secs())).unwrap()
1075    }
1076
1077    /// Assert vector equality with better error message.
1078    #[macro_export]
1079    macro_rules! assert_vec_eq {
1080        ($a:expr, $b: expr) => {
1081            match (&$a, &$b) {
1082                (ref a, ref b) => {
1083                    let len_not_eq = a.len() != b.len();
1084                    let mut data_not_eq = None;
1085                    for (i, (a, b)) in a.iter().zip(b.iter()).enumerate() {
1086                        if a != b {
1087                            data_not_eq = Some(i);
1088                            break;
1089                        }
1090                    }
1091
1092                    if len_not_eq || data_not_eq.is_some() {
1093                        use std::fmt::*;
1094
1095                        let mut error = format!("`{}` != `{}`", stringify!($a), stringify!($b));
1096                        if len_not_eq {
1097                            let _ = write!(&mut error, "\n  lengths not equal: {} != {}", a.len(), b.len());
1098                        }
1099                        if let Some(i) = data_not_eq {
1100                            let _ = write!(&mut error, "\n  data not equal at index {}: {} != {:?}", i, a[i], b[i]);
1101                        }
1102                        panic!("{error}")
1103                    }
1104                }
1105            }
1106        };
1107    }
1108}