Skip to main content

sfo_io/
limit_datagram_local.rs

1#![cfg_attr(coverage_nightly, feature(coverage_attribute))]
2
3use crate::SpeedLimitSession;
4
5#[async_trait::async_trait(?Send)]
6pub trait LocalDatagramSend {
7    type Error;
8    async fn send_to(&mut self, buf: &[u8]) -> Result<usize, Self::Error>;
9}
10
11#[async_trait::async_trait(?Send)]
12pub trait LocalDatagramRecv {
13    type Error;
14    async fn recv_from(&mut self, buf: &mut [u8]) -> Result<usize, Self::Error>;
15}
16
17enum ReadState {
18    Idle,
19    Reading((usize, usize)),
20}
21
22enum WriteState {
23    Idle,
24    Writing((usize, usize)),
25}
26
27pub struct LocalLimitDatagramSend<S: LocalDatagramSend> {
28    inner: S,
29    write_limiter: SpeedLimitSession,
30    write_state: WriteState,
31}
32
33impl<S: LocalDatagramSend> LocalLimitDatagramSend<S> {
34    pub fn new(inner: S, write_limiter: SpeedLimitSession) -> Self {
35        Self {
36            inner,
37            write_limiter,
38            write_state: WriteState::Idle,
39        }
40    }
41}
42
43#[async_trait::async_trait(?Send)]
44impl<S: LocalDatagramSend> LocalDatagramSend for LocalLimitDatagramSend<S> {
45    type Error = S::Error;
46
47    async fn send_to(&mut self, buf: &[u8]) -> Result<usize, Self::Error> {
48        match &mut self.write_state {
49            WriteState::Idle => {
50                let write_len = self.write_limiter.until_ready().await;
51                self.inner.send_to(buf).await?;
52                if buf.len() > write_len {
53                    self.write_state = WriteState::Idle;
54                } else {
55                    self.write_state = WriteState::Writing((write_len, buf.len()));
56                }
57                Ok(buf.len())
58            }
59            WriteState::Writing((write_len, written_len)) => {
60                self.inner.send_to(buf).await?;
61                if *written_len + buf.len() >= *write_len {
62                    self.write_state = WriteState::Idle;
63                    Ok(buf.len())
64                } else {
65                    self.write_state = WriteState::Writing((*write_len, *written_len + buf.len()));
66                    Ok(buf.len())
67                }
68            }
69        }
70    }
71}
72
73pub struct LocalLimitDatagramRecv<R: LocalDatagramRecv> {
74    inner: R,
75    read_limiter: SpeedLimitSession,
76    read_state: ReadState,
77}
78
79impl<R: LocalDatagramRecv> LocalLimitDatagramRecv<R> {
80    pub fn new(inner: R, read_limiter: SpeedLimitSession) -> Self {
81        Self {
82            inner,
83            read_limiter,
84            read_state: ReadState::Idle,
85        }
86    }
87}
88
89#[async_trait::async_trait(?Send)]
90impl<R: LocalDatagramRecv> LocalDatagramRecv for LocalLimitDatagramRecv<R> {
91    type Error = R::Error;
92
93    async fn recv_from(&mut self, buf: &mut [u8]) -> Result<usize, Self::Error> {
94        match &mut self.read_state {
95            ReadState::Idle => {
96                let read_len = self.read_limiter.until_ready().await;
97                let len = self.inner.recv_from(buf).await?;
98                if len > read_len {
99                    self.read_state = ReadState::Idle;
100                    Ok(len)
101                } else {
102                    self.read_state = ReadState::Reading((read_len, len));
103                    Ok(len)
104                }
105            }
106            ReadState::Reading((read_len, readded_len)) => {
107                let len = self.inner.recv_from(buf).await?;
108                if *readded_len + len >= *read_len {
109                    self.read_state = ReadState::Idle;
110                } else {
111                    self.read_state = ReadState::Reading((*read_len, *readded_len + len));
112                }
113                Ok(len)
114            }
115        }
116    }
117}
118
119#[async_trait::async_trait(?Send)]
120pub trait LocalDatagram {
121    type Error;
122    async fn send_to(&mut self, buf: &[u8]) -> Result<usize, Self::Error>;
123    async fn recv_from(&mut self, buf: &mut [u8]) -> Result<usize, Self::Error>;
124}
125
126pub struct LocalLimitDatagram<D: LocalDatagram> {
127    inner: D,
128    write_limiter: SpeedLimitSession,
129    read_limiter: SpeedLimitSession,
130    read_state: ReadState,
131    write_state: WriteState,
132}
133
134impl<D: LocalDatagram> LocalLimitDatagram<D> {
135    pub fn new(inner: D, read_limit: SpeedLimitSession, write_limit: SpeedLimitSession) -> Self {
136        Self {
137            inner,
138            write_limiter: write_limit,
139            read_limiter: read_limit,
140            read_state: ReadState::Idle,
141            write_state: WriteState::Idle,
142        }
143    }
144}
145
146#[async_trait::async_trait(?Send)]
147impl<D: LocalDatagram> LocalDatagram for LocalLimitDatagram<D> {
148    type Error = D::Error;
149
150    async fn send_to(&mut self, buf: &[u8]) -> Result<usize, Self::Error> {
151        match &mut self.write_state {
152            WriteState::Idle => {
153                let write_len = self.write_limiter.until_ready().await;
154                self.inner.send_to(buf).await?;
155                if buf.len() > write_len {
156                    self.write_state = WriteState::Idle;
157                } else {
158                    self.write_state = WriteState::Writing((write_len, buf.len()));
159                }
160                Ok(buf.len())
161            }
162            WriteState::Writing((write_len, written_len)) => {
163                self.inner.send_to(buf).await?;
164                if *written_len + buf.len() >= *write_len {
165                    self.write_state = WriteState::Idle;
166                    Ok(buf.len())
167                } else {
168                    self.write_state = WriteState::Writing((*write_len, *written_len + buf.len()));
169                    Ok(buf.len())
170                }
171            }
172        }
173    }
174
175    async fn recv_from(&mut self, buf: &mut [u8]) -> Result<usize, Self::Error> {
176        match &mut self.read_state {
177            ReadState::Idle => {
178                let read_len = self.read_limiter.until_ready().await;
179                let len = self.inner.recv_from(buf).await?;
180                if len > read_len {
181                    self.read_state = ReadState::Idle;
182                    Ok(len)
183                } else {
184                    self.read_state = ReadState::Reading((read_len, len));
185                    Ok(len)
186                }
187            }
188            ReadState::Reading((read_len, readded_len)) => {
189                let len = self.inner.recv_from(buf).await?;
190                if *readded_len + len >= *read_len {
191                    self.read_state = ReadState::Idle;
192                } else {
193                    self.read_state = ReadState::Reading((*read_len, *readded_len + len));
194                }
195                Ok(len)
196            }
197        }
198    }
199}
200
201#[cfg_attr(coverage_nightly, coverage(off))]
202#[cfg(test)]
203mod tests {
204    use super::*;
205    use crate::SpeedLimiter;
206    use std::cell::RefCell;
207    use std::num::NonZeroU32;
208    use std::rc::Rc;
209
210    struct LocalMockDatagram {
211        calls: Rc<RefCell<Vec<&'static str>>>,
212    }
213
214    #[async_trait::async_trait(?Send)]
215    impl LocalDatagram for LocalMockDatagram {
216        type Error = &'static str;
217
218        async fn send_to(&mut self, buf: &[u8]) -> Result<usize, Self::Error> {
219            self.calls.borrow_mut().push("send_to");
220            Ok(buf.len())
221        }
222
223        async fn recv_from(&mut self, buf: &mut [u8]) -> Result<usize, Self::Error> {
224            self.calls.borrow_mut().push("recv_from");
225            Ok(buf.len())
226        }
227    }
228
229    #[tokio::test(flavor = "current_thread")]
230    async fn local_limit_datagram_accepts_non_send_inner() {
231        let calls = Rc::new(RefCell::new(Vec::new()));
232        let mock = LocalMockDatagram {
233            calls: calls.clone(),
234        };
235        let read_limit = SpeedLimiter::new(None, None, NonZeroU32::new(64)).new_limit_session();
236        let write_limit = SpeedLimiter::new(None, None, NonZeroU32::new(64)).new_limit_session();
237        let mut datagram = LocalLimitDatagram::new(mock, read_limit, write_limit);
238
239        let mut recv_buf = [0; 8];
240        assert_eq!(datagram.send_to(&[1, 2, 3]).await, Ok(3));
241        assert_eq!(datagram.recv_from(&mut recv_buf).await, Ok(8));
242        assert_eq!(&*calls.borrow(), &["send_to", "recv_from"]);
243    }
244}