Skip to main content

sfo_io/
limit_stream_local.rs

1#![cfg_attr(coverage_nightly, feature(coverage_attribute))]
2
3use crate::SpeedLimitSession;
4use pin_project::pin_project;
5use std::future::Future;
6use std::io::Error;
7use std::pin::Pin;
8use std::task::{Context, Poll};
9use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
10
11enum ReadState {
12    Idle,
13    Waiting(Option<(Pin<Box<dyn Future<Output = usize> + 'static>>, usize)>),
14    Reading(Option<(usize, usize)>),
15}
16
17enum WriteState {
18    Idle,
19    Waiting(Option<(Pin<Box<dyn Future<Output = usize> + 'static>>, usize)>),
20    Writing(Option<(usize, usize)>),
21}
22
23#[pin_project]
24pub struct LocalLimitStream<S: AsyncRead + AsyncWrite + Unpin> {
25    #[pin]
26    read: LocalLimitRead<sfo_split::ReadHalf<S>>,
27    #[pin]
28    write: LocalLimitWrite<sfo_split::WriteHalf<S>>,
29}
30
31impl<S: AsyncRead + AsyncWrite + Unpin> LocalLimitStream<S> {
32    pub fn new(stream: S, read_limit: SpeedLimitSession, write_limit: SpeedLimitSession) -> Self {
33        let (read, write) = sfo_split::split(stream);
34        let limit_read = LocalLimitRead::new(read, read_limit);
35        let limit_write = LocalLimitWrite::new(write, write_limit);
36        LocalLimitStream {
37            read: limit_read,
38            write: limit_write,
39        }
40    }
41
42    pub fn with_lock_raw_stream<R>(&mut self, f: impl FnOnce(Pin<&mut S>) -> R) -> R {
43        self.read.raw_read().with_lock(f)
44    }
45}
46
47impl<S: AsyncRead + AsyncWrite + Unpin> AsyncWrite for LocalLimitStream<S> {
48    fn poll_write(
49        self: Pin<&mut Self>,
50        cx: &mut Context<'_>,
51        buf: &[u8],
52    ) -> Poll<Result<usize, Error>> {
53        self.project().write.poll_write(cx, buf)
54    }
55
56    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
57        self.project().write.poll_flush(cx)
58    }
59
60    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
61        self.project().write.poll_shutdown(cx)
62    }
63}
64
65impl<S: AsyncRead + AsyncWrite + Unpin> AsyncRead for LocalLimitStream<S> {
66    fn poll_read(
67        self: Pin<&mut Self>,
68        cx: &mut Context<'_>,
69        buf: &mut ReadBuf<'_>,
70    ) -> Poll<std::io::Result<()>> {
71        self.project().read.poll_read(cx, buf)
72    }
73}
74
75#[pin_project]
76pub struct LocalLimitRead<S: AsyncRead + Unpin> {
77    #[pin]
78    read: S,
79    read_limit: SpeedLimitSession,
80    read_state: ReadState,
81}
82
83impl<S: AsyncRead + Unpin> LocalLimitRead<S> {
84    pub fn new(read: S, read_limit: SpeedLimitSession) -> Self {
85        LocalLimitRead {
86            read,
87            read_limit,
88            read_state: ReadState::Idle,
89        }
90    }
91
92    pub fn raw_read_mut(&mut self) -> &mut S {
93        &mut self.read
94    }
95
96    pub fn raw_read(&self) -> &S {
97        &self.read
98    }
99
100    pub fn into_raw_read(self) -> S {
101        self.read
102    }
103}
104
105impl<S: AsyncRead + Unpin> AsyncRead for LocalLimitRead<S> {
106    fn poll_read(
107        self: Pin<&mut Self>,
108        cx: &mut Context<'_>,
109        buf: &mut ReadBuf<'_>,
110    ) -> Poll<std::io::Result<()>> {
111        let this = self.project();
112        buf.initialize_unfilled();
113        match this.read_state {
114            ReadState::Idle => {
115                let mut readded_len = 0;
116                let read_limit: &'static mut SpeedLimitSession =
117                    unsafe { std::mem::transmute(this.read_limit) };
118                let mut waiting_future = Box::pin(read_limit.until_ready());
119                match Pin::new(&mut waiting_future).poll(cx) {
120                    Poll::Ready(read_len) => {
121                        let mut read_buf = if read_len <= buf.remaining() {
122                            buf.take(read_len)
123                        } else {
124                            buf.take(buf.remaining())
125                        };
126                        match this.read.poll_read(cx, &mut read_buf) {
127                            Poll::Ready(Ok(())) => {
128                                let len = read_buf.filled().len();
129                                readded_len += len;
130                                buf.advance(len);
131                                if readded_len >= read_len {
132                                    *this.read_state = ReadState::Idle;
133                                } else {
134                                    *this.read_state =
135                                        ReadState::Reading(Some((read_len, readded_len)));
136                                }
137                                Poll::Ready(Ok(()))
138                            }
139                            Poll::Ready(Err(e)) => {
140                                *this.read_state = ReadState::Idle;
141                                Poll::Ready(Err(e))
142                            }
143                            Poll::Pending => {
144                                *this.read_state =
145                                    ReadState::Reading(Some((read_len, readded_len)));
146                                Poll::Pending
147                            }
148                        }
149                    }
150                    Poll::Pending => {
151                        *this.read_state = ReadState::Waiting(Some((waiting_future, readded_len)));
152                        Poll::Pending
153                    }
154                }
155            }
156            ReadState::Waiting(state) => {
157                let (mut rx, mut readded_len) = state.take().unwrap();
158                match Pin::new(&mut rx).poll(cx) {
159                    Poll::Ready(read_len) => {
160                        let mut read_buf = if (read_len - readded_len) <= buf.remaining() {
161                            buf.take(read_len - readded_len)
162                        } else {
163                            buf.take(buf.remaining())
164                        };
165                        match this.read.poll_read(cx, &mut read_buf) {
166                            Poll::Ready(Ok(())) => {
167                                let len = read_buf.filled().len();
168                                readded_len += len;
169                                buf.advance(len);
170                                if readded_len >= read_len {
171                                    *this.read_state = ReadState::Idle;
172                                } else {
173                                    *this.read_state =
174                                        ReadState::Reading(Some((read_len, readded_len)));
175                                }
176                                Poll::Ready(Ok(()))
177                            }
178                            Poll::Ready(Err(e)) => {
179                                *this.read_state = ReadState::Idle;
180                                Poll::Ready(Err(e))
181                            }
182                            Poll::Pending => {
183                                *this.read_state =
184                                    ReadState::Reading(Some((read_len, readded_len)));
185                                Poll::Pending
186                            }
187                        }
188                    }
189                    Poll::Pending => {
190                        *this.read_state = ReadState::Waiting(Some((rx, readded_len)));
191                        Poll::Pending
192                    }
193                }
194            }
195            ReadState::Reading(state) => match state.take() {
196                Some((read_len, mut readded_len)) => {
197                    let mut read_buf = if (read_len - readded_len) <= buf.remaining() {
198                        buf.take(read_len - readded_len)
199                    } else {
200                        buf.take(buf.remaining())
201                    };
202                    match this.read.poll_read(cx, &mut read_buf) {
203                        Poll::Ready(Ok(())) => {
204                            let len = read_buf.filled().len();
205                            readded_len += len;
206                            buf.advance(len);
207                            if readded_len >= read_len {
208                                *this.read_state = ReadState::Idle;
209                            } else {
210                                *this.read_state =
211                                    ReadState::Reading(Some((read_len, readded_len)));
212                            }
213                            Poll::Ready(Ok(()))
214                        }
215                        Poll::Ready(Err(e)) => {
216                            *this.read_state = ReadState::Idle;
217                            Poll::Ready(Err(e))
218                        }
219                        Poll::Pending => {
220                            *this.read_state = ReadState::Reading(Some((read_len, readded_len)));
221                            Poll::Pending
222                        }
223                    }
224                }
225                None => match this.read.poll_read(cx, buf) {
226                    Poll::Ready(Ok(())) => {
227                        *this.read_state = ReadState::Idle;
228                        Poll::Ready(Ok(()))
229                    }
230                    Poll::Ready(Err(e)) => {
231                        *this.read_state = ReadState::Idle;
232                        Poll::Ready(Err(e))
233                    }
234                    Poll::Pending => {
235                        *this.read_state = ReadState::Reading(None);
236                        Poll::Pending
237                    }
238                },
239            },
240        }
241    }
242}
243
244#[pin_project]
245pub struct LocalLimitWrite<S: AsyncWrite + Unpin> {
246    #[pin]
247    write: S,
248    write_limit: SpeedLimitSession,
249    write_state: WriteState,
250}
251
252impl<S: AsyncWrite + Unpin> LocalLimitWrite<S> {
253    pub fn new(write: S, write_limit: SpeedLimitSession) -> Self {
254        LocalLimitWrite {
255            write,
256            write_limit,
257            write_state: WriteState::Idle,
258        }
259    }
260
261    pub fn raw_write_mut(&mut self) -> &mut S {
262        &mut self.write
263    }
264
265    pub fn raw_write(&self) -> &S {
266        &self.write
267    }
268
269    pub fn into_raw_write(self) -> S {
270        self.write
271    }
272}
273
274impl<S: AsyncWrite + Unpin> AsyncWrite for LocalLimitWrite<S> {
275    fn poll_write(
276        self: Pin<&mut Self>,
277        cx: &mut Context<'_>,
278        buf: &[u8],
279    ) -> Poll<Result<usize, Error>> {
280        let this = self.project();
281        match this.write_state {
282            WriteState::Idle => {
283                let mut written_len = 0;
284                let write_limiter: &'static mut SpeedLimitSession =
285                    unsafe { std::mem::transmute(this.write_limit) };
286                let mut waiting_future = Box::pin(write_limiter.until_ready());
287                match Pin::new(&mut waiting_future).poll(cx) {
288                    Poll::Ready(write_len) => {
289                        let write_buf = if write_len <= buf.len() {
290                            &buf[..write_len]
291                        } else {
292                            buf
293                        };
294                        match this.write.poll_write(cx, write_buf) {
295                            Poll::Ready(Ok(len)) => {
296                                written_len += len;
297                                if written_len >= write_len {
298                                    *this.write_state = WriteState::Idle;
299                                } else {
300                                    *this.write_state =
301                                        WriteState::Writing(Some((write_len, written_len)));
302                                }
303                                Poll::Ready(Ok(written_len))
304                            }
305                            Poll::Ready(Err(e)) => {
306                                *this.write_state = WriteState::Idle;
307                                Poll::Ready(Err(e))
308                            }
309                            Poll::Pending => {
310                                *this.write_state =
311                                    WriteState::Writing(Some((write_len, written_len)));
312                                Poll::Pending
313                            }
314                        }
315                    }
316                    Poll::Pending => {
317                        *this.write_state =
318                            WriteState::Waiting(Some((waiting_future, written_len)));
319                        Poll::Pending
320                    }
321                }
322            }
323            WriteState::Waiting(state) => {
324                let (mut waiting_future, mut written_len) = state.take().unwrap();
325                match Pin::new(&mut waiting_future).poll(cx) {
326                    Poll::Ready(write_len) => {
327                        let write_buf = if write_len - written_len <= buf.len() {
328                            &buf[..(write_len - written_len)]
329                        } else {
330                            buf
331                        };
332                        match this.write.poll_write(cx, write_buf) {
333                            Poll::Ready(Ok(len)) => {
334                                written_len += len;
335                                if written_len >= write_len {
336                                    *this.write_state = WriteState::Idle;
337                                } else {
338                                    *this.write_state =
339                                        WriteState::Writing(Some((write_len, written_len)));
340                                }
341                                Poll::Ready(Ok(len))
342                            }
343                            Poll::Ready(Err(e)) => {
344                                *this.write_state = WriteState::Idle;
345                                Poll::Ready(Err(e))
346                            }
347                            Poll::Pending => {
348                                *this.write_state =
349                                    WriteState::Writing(Some((write_len, written_len)));
350                                Poll::Pending
351                            }
352                        }
353                    }
354                    Poll::Pending => {
355                        *this.write_state =
356                            WriteState::Waiting(Some((waiting_future, written_len)));
357                        Poll::Pending
358                    }
359                }
360            }
361            WriteState::Writing(state) => match state.take() {
362                Some((write_len, mut written_len)) => {
363                    let write_buf = if write_len - written_len <= buf.len() {
364                        &buf[..(write_len - written_len)]
365                    } else {
366                        buf
367                    };
368                    match this.write.poll_write(cx, write_buf) {
369                        Poll::Ready(Ok(len)) => {
370                            written_len += len;
371                            if written_len >= write_len {
372                                *this.write_state = WriteState::Idle;
373                            } else {
374                                *this.write_state =
375                                    WriteState::Writing(Some((write_len, written_len)));
376                            }
377                            Poll::Ready(Ok(len))
378                        }
379                        Poll::Ready(Err(e)) => {
380                            *this.write_state = WriteState::Idle;
381                            Poll::Ready(Err(e))
382                        }
383                        Poll::Pending => {
384                            *this.write_state = WriteState::Writing(Some((write_len, written_len)));
385                            Poll::Pending
386                        }
387                    }
388                }
389                None => match this.write.poll_write(cx, buf) {
390                    Poll::Ready(Ok(len)) => {
391                        *this.write_state = WriteState::Idle;
392                        Poll::Ready(Ok(len))
393                    }
394                    Poll::Ready(Err(e)) => {
395                        *this.write_state = WriteState::Idle;
396                        Poll::Ready(Err(e))
397                    }
398                    Poll::Pending => {
399                        *this.write_state = WriteState::Writing(None);
400                        Poll::Pending
401                    }
402                },
403            },
404        }
405    }
406
407    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
408        self.project().write.poll_flush(cx)
409    }
410
411    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
412        self.project().write.poll_shutdown(cx)
413    }
414}
415
416#[cfg_attr(coverage_nightly, coverage(off))]
417#[cfg(test)]
418mod tests {
419    use super::*;
420    use crate::SpeedLimiter;
421    use std::cell::RefCell;
422    use std::io;
423    use std::num::NonZeroU32;
424    use std::rc::Rc;
425    use tokio::io::{AsyncReadExt, AsyncWriteExt};
426
427    struct LocalMockStream {
428        read_data: Rc<RefCell<Vec<u8>>>,
429        written_data: Rc<RefCell<Vec<u8>>>,
430    }
431
432    impl AsyncRead for LocalMockStream {
433        fn poll_read(
434            self: Pin<&mut Self>,
435            _cx: &mut Context<'_>,
436            buf: &mut ReadBuf<'_>,
437        ) -> Poll<io::Result<()>> {
438            let mut read_data = self.read_data.borrow_mut();
439            let len = read_data.len().min(buf.remaining());
440            buf.put_slice(&read_data[..len]);
441            read_data.drain(..len);
442            Poll::Ready(Ok(()))
443        }
444    }
445
446    impl AsyncWrite for LocalMockStream {
447        fn poll_write(
448            self: Pin<&mut Self>,
449            _cx: &mut Context<'_>,
450            buf: &[u8],
451        ) -> Poll<io::Result<usize>> {
452            self.written_data.borrow_mut().extend_from_slice(buf);
453            Poll::Ready(Ok(buf.len()))
454        }
455
456        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
457            Poll::Ready(Ok(()))
458        }
459
460        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
461            Poll::Ready(Ok(()))
462        }
463    }
464
465    #[tokio::test(flavor = "current_thread")]
466    async fn local_limit_stream_accepts_non_send_inner() {
467        let read_data = Rc::new(RefCell::new(vec![1, 2, 3]));
468        let written_data = Rc::new(RefCell::new(Vec::new()));
469        let mock = LocalMockStream {
470            read_data,
471            written_data: written_data.clone(),
472        };
473        let read_limit = SpeedLimiter::new(None, None, NonZeroU32::new(64)).new_limit_session();
474        let write_limit = SpeedLimiter::new(None, None, NonZeroU32::new(64)).new_limit_session();
475        let mut stream = LocalLimitStream::new(mock, read_limit, write_limit);
476
477        let mut read_buf = [0; 3];
478        stream.read_exact(&mut read_buf).await.unwrap();
479        stream.write_all(&[4, 5, 6]).await.unwrap();
480
481        assert_eq!(read_buf, [1, 2, 3]);
482        assert_eq!(&*written_data.borrow(), &[4, 5, 6]);
483    }
484}