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}