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};
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
102/// Measure read/write task.
103///
104/// Metrics are updated after each read/write, if you read/write all bytes in one call
105/// the metrics will only update once.
106pub struct Measure<T> {
107    task: T,
108    inner: MeasureInner,
109}
110impl<T> Measure<T> {
111    /// Start measuring a new read/write task.
112    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    /// Continue measuring a read/write task.
117    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    /// Current metrics.
125    ///
126    /// This value is updated after every read/write.
127    pub fn metrics(&self) -> Var<Metrics> {
128        self.inner.metrics.read_only()
129    }
130
131    /// Unwrap the inner task and final metrics.
132    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        // SAFETY: we don't move anything.
147        let self_ = unsafe { self.get_unchecked_mut() };
148
149        // SAFETY: we don't move task
150        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        // SAFETY: we don't move anything.
162        let self_ = unsafe { self.get_unchecked_mut() };
163
164        // SAFETY: we don't move task
165        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        // SAFETY: we don't move anything.
176        let self_ = unsafe { self.get_unchecked_mut() };
177
178        // SAFETY: we don't move task
179        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        // SAFETY: we don't move anything.
184        let self_ = unsafe { self.get_unchecked_mut() };
185
186        // SAFETY: we don't move task
187        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        // SAFETY: we don't move anything.
193        let self_ = unsafe { self.get_unchecked_mut() };
194
195        // SAFETY: we don't move task
196        unsafe { Pin::new_unchecked(&mut self_.task) }.poll_fill_buf(cx)
197    }
198
199    fn consume(self: Pin<&mut Self>, amt: usize) {
200        // SAFETY: we don't move anything.
201        let self_ = unsafe { self.get_unchecked_mut() };
202        // SAFETY: we don't move task
203        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/// Information about the state of an async IO task.
245///
246/// Read is also called *receive* or *download*. Write is also called *send* or *upload*. The default
247/// display print uses arrows ↓ and ↑ for read and write.
248///
249/// Use [`Measure`] to measure a task.
250#[derive(Debug, Clone, PartialEq, Eq)]
251#[non_exhaustive]
252pub struct Metrics {
253    /// Number of bytes read / estimated total.
254    pub read_progress: (ByteLength, ByteLength),
255
256    /// Average read speed in bytes/second.
257    pub read_speed: ByteLength,
258
259    /// Number of bytes written / estimated total.
260    pub write_progress: (ByteLength, ByteLength),
261
262    /// Average write speed in bytes/second.
263    pub write_speed: ByteLength,
264
265    /// Total time for the entire task. This will continuously increase until
266    /// the task is finished.
267    pub total_time: Duration,
268}
269impl Metrics {
270    /// All zeros.
271    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    /// Metrics in a [`Progress::with_meta`] metadata.
336    pub static ref METRICS_ID: zng_state_map::StateId<Metrics>;
337}
338
339/// Extension methods for [`std::io::Error`] to be used with errors returned by [`McBufReader`].
340pub trait McBufErrorExt {
341    /// Returns `true` if this error represents the condition where there are only [`McBufReader::is_lazy`] readers
342    /// left, the buffer is drained and the inner reader is not EOF.
343    ///
344    /// You can recover from this error by turning the reader non-lazy using [`McBufReader::set_lazy`].
345    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
354/// Multiple consumer buffered read.
355///
356/// Clone an instance to create a new consumer, already read bytes stay in the buffer until all clones have read it,
357/// clones continue reading from the same offset as the reader they cloned.
358///
359/// A single instance of this reader behaves like a `BufReader`.
360///
361/// # Result
362///
363/// The result is *repeats* ready when `EOF` or an [`Error`] occurs, unfortunately the IO error is not cloneable
364/// so the error is recreated using [`CloneableError`] for subsequent poll attempts.
365///
366/// The inner reader is dropped as soon as it finishes.
367///
368/// # Lazy Clones
369///
370/// You can mark clones as [lazy], lazy clones don't pull from the inner reader, only advance when another clone reads, if
371/// all living clones are lazy they stop reading with an error. You can identify this custom error using the [`McBufErrorExt::is_only_lazy_left`]
372/// extension method.
373///
374/// [lazy]: Self::set_lazy
375pub 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    /// Creates a buffered reader.
394    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    /// Returns `true` if this reader does not pull from the inner reader, only advancing when a non-lazy reader advances.
416    ///
417    /// The initial reader is not lazy, only clones of lazy readers are lazy by default.
418    pub fn is_lazy(&self) -> bool {
419        self.lazy
420    }
421
422    /// Sets [`is_lazy`].
423    ///
424    /// [`is_lazy`]: Self::is_lazy
425    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                // notify lazy so they get the error.
463                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        // ready data for this clone.
477        let mut i = inner.clones[self_.index];
478        let mut ready;
479
480        match &inner.result {
481            ReadState::Running => {
482                // source has not finished yet.
483
484                ready = &inner.buf[i..];
485
486                if ready.is_empty() {
487                    if self.lazy {
488                        if inner.non_lazy_count == 0 {
489                            // user can make this reader non-lazy and try again.
490                            return Poll::Ready(Err(Error::other(ONLY_NON_LAZY_ERROR_MSG)));
491                        } else {
492                            // register waker for after non-lazy poll.
493                            inner.lazy_wakers.push(cx.waker().clone());
494
495                            // wait non-lazy to pull.
496                            return Poll::Pending;
497                        }
498                    }
499
500                    // time to poll source.
501
502                    ready = &[];
503
504                    let waker = match inner.waker.push(cx.waker().clone()) {
505                        Some(w) => w,
506                        None => {
507                            // already polling from another clone.
508                            return Poll::Pending;
509                        }
510                    };
511
512                    let min_i = inner.clones.iter().copied().min().unwrap();
513                    if min_i > 0 {
514                        // reuse front.
515                        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                    // SAFETY: we don't move `source`.
531                    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                            // notify lazy readers.
537                            for waker in inner.lazy_wakers.drain(..) {
538                                waker.wake();
539                            }
540
541                            match result {
542                                Ok(0) => {
543                                    inner.waker.cancel();
544
545                                    // EOF
546                                    inner.buf.truncate(new_start);
547                                    inner.result = ReadState::Eof;
548                                    inner.source = None;
549
550                                    // continue 'copy ready
551                                }
552                                Ok(read) => {
553                                    inner.waker.cancel();
554
555                                    // Read > 0
556                                    inner.buf.truncate(new_start + read);
557                                    ready = &inner.buf[i..];
558
559                                    // continue 'copy ready
560                                }
561                                Err(e) => {
562                                    inner.waker.cancel();
563
564                                    // Error
565                                    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                // continue 'copy ready
585            }
586            ReadState::Err(e) => return Poll::Ready(e.err()),
587        }
588
589        // 'copy ready
590
591        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/// Represents the cloneable parts of an [`Error`].
602///
603/// Unfortunately [`Error`] does not implement clone, this is needed to implemented
604/// IO futures that repeat the ready result after subsequent polls. This type partially
605/// works around the issue by copying enough information to recreate an error that is still useful.
606///
607/// The OS error code, [`ErrorKind`] and display message are preserved. Note that this not an error type,
608/// it must be converted to [`Error`] using `into` or [`err`].
609///
610/// [`err`]: Self::err
611#[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    /// Copy the cloneable information from the [`Error`].
622    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    /// Returns an `Err(Error)` generated from the cloneable information.
633    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
646/// Represents a stream reader that generates an error if the source stream exceeds a limit.
647///
648/// Note that some bytes over the limit may be read once if the source stream is buffered.
649pub struct ReadLimited<S> {
650    source: S,
651    limit: usize,
652    on_limit: fn() -> std::io::Error,
653}
654impl<S> ReadLimited<S> {
655    /// Construct a limited reader.
656    ///
657    /// The `on_limit` closure is called for every read attempt after the limit is reached.
658    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    /// New with default on limit error.
667    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        // SAFETY: we don't move anything.
679        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        // SAFETY: we never move `source`.
691        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        // SAFETY: we don't move anything.
706        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        // SAFETY: we never move `source`.
714        unsafe { Pin::new_unchecked(&mut self_.source) }.poll_fill_buf(cx)
715    }
716
717    fn consume(self: Pin<&mut Self>, amt: usize) {
718        // SAFETY: we don't move anything.
719        let self_ = unsafe { self.get_unchecked_mut() };
720        // SAFETY: we never move `source`.
721        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    /// Assert vector equality with better error message.
1082    #[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}