sfo_io/
limit_datagram_local.rs1#![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}