1use crate::StatusCode;
8use async_lock::Semaphore;
9use async_std::{
10 net::TcpListener,
11 sync::Arc,
12 task::{sleep, spawn},
13};
14use async_trait::async_trait;
15use derivative::Derivative;
16use futures::stream::StreamExt;
17use std::{
18 fmt::{self, Display, Formatter},
19 io::{self, ErrorKind},
20 net::SocketAddr,
21 time::Duration,
22};
23use tide::{
24 Server, http,
25 listener::{ListenInfo, Listener, ToListener},
26};
27
28#[derive(Derivative)]
35#[derivative(Debug(bound = "State: Send + Sync + 'static"))]
36pub struct RateLimitListener<State> {
37 addr: SocketAddr,
38 listener: Option<TcpListener>,
39 server: Option<Server<State>>,
40 info: Option<ListenInfo>,
41 permit: Arc<Semaphore>,
42}
43
44impl<State> RateLimitListener<State> {
45 pub fn new(addr: SocketAddr, limit: usize) -> Self {
47 Self {
48 addr,
49 listener: None,
50 server: None,
51 info: None,
52 permit: Arc::new(Semaphore::new(limit)),
53 }
54 }
55
56 pub fn with_port(port: u16, limit: usize) -> Self {
58 Self::new(([0, 0, 0, 0], port).into(), limit)
59 }
60}
61
62#[async_trait]
63impl<State> Listener<State> for RateLimitListener<State>
64where
65 State: Clone + Send + Sync + 'static,
66{
67 async fn bind(&mut self, app: Server<State>) -> io::Result<()> {
68 if self.server.is_some() {
69 return Err(io::Error::new(
70 ErrorKind::AlreadyExists,
71 "`bind` should only be called once",
72 ));
73 }
74 self.server = Some(app);
75 self.listener = Some(TcpListener::bind(&[self.addr][..]).await?);
76
77 let conn_string = format!("{}", self);
79 let transport = "tcp".to_owned();
80 let tls = false;
81 self.info = Some(ListenInfo::new(conn_string, transport, tls));
82
83 Ok(())
84 }
85
86 async fn accept(&mut self) -> io::Result<()> {
87 let server = self.server.take().ok_or_else(|| {
88 io::Error::other("`Listener::bind` must be called before `Listener::accept`")
89 })?;
90 let listener = self.listener.take().ok_or_else(|| {
91 io::Error::other("`Listener::bind` must be called before `Listener::accept`")
92 })?;
93
94 let mut incoming = listener.incoming();
95 while let Some(stream) = incoming.next().await {
96 match stream {
97 Err(err) if is_transient_error(&err) => continue,
98 Err(err) => {
99 tracing::warn!(%err, "TCP error");
100 sleep(Duration::from_millis(500)).await;
101 continue;
102 }
103 Ok(stream) => {
104 let app = server.clone();
105 let permit = self.permit.clone();
106 spawn(async move {
107 let local_addr = stream.local_addr().ok();
108 let peer_addr = stream.peer_addr().ok();
109
110 let fut = async_h1::accept(stream, |mut req| async {
111 if let Some(_guard) = permit.try_acquire() {
113 req.set_local_addr(local_addr);
114 req.set_peer_addr(peer_addr);
115 app.respond(req).await
116 } else {
117 Ok(http::Response::new(StatusCode::TOO_MANY_REQUESTS))
120 }
121 });
122
123 if let Err(error) = fut.await {
124 tracing::error!(%error, "HTTP error");
125 }
126 });
127 }
128 };
129 }
130 Ok(())
131 }
132
133 fn info(&self) -> Vec<ListenInfo> {
134 match &self.info {
135 Some(info) => vec![info.clone()],
136 None => vec![],
137 }
138 }
139}
140
141impl<State> ToListener<State> for RateLimitListener<State>
142where
143 State: Clone + Send + Sync + 'static,
144{
145 type Listener = Self;
146
147 fn to_listener(self) -> io::Result<Self::Listener> {
148 Ok(self)
149 }
150}
151
152impl<State> Display for RateLimitListener<State> {
153 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
154 match &self.listener {
155 Some(listener) => {
156 let addr = listener.local_addr().expect("Could not get local addr");
157 write!(f, "http://{}", addr)
158 }
159 None => write!(f, "http://{}", self.addr),
160 }
161 }
162}
163
164fn is_transient_error(e: &io::Error) -> bool {
165 matches!(
166 e.kind(),
167 ErrorKind::ConnectionRefused | ErrorKind::ConnectionAborted | ErrorKind::ConnectionReset
168 )
169}
170
171#[cfg(test)]
172mod test {
173 use super::*;
174 use crate::{
175 App,
176 error::ServerError,
177 testing::{Client, setup_test},
178 };
179 use futures::future::{FutureExt, try_join_all};
180 use portpicker::pick_unused_port;
181 use toml::toml;
182 use vbs::version::{StaticVersion, StaticVersionType};
183
184 type StaticVer01 = StaticVersion<0, 1>;
185
186 #[async_std::test]
187 async fn test_rate_limiting() {
188 setup_test();
189
190 let mut app = App::<_, ServerError>::with_state(());
191 let api_toml = toml! {
192 [route.test]
193 PATH = ["/test"]
194 METHOD = "GET"
195 };
196 {
197 let mut api = app
198 .module::<ServerError, StaticVer01>("mod", api_toml)
199 .unwrap();
200 api.get("test", |_req, _state| {
201 async move {
202 sleep(Duration::from_secs(30)).await;
204 Ok(())
205 }
206 .boxed()
207 })
208 .unwrap();
209 }
210
211 let limit = 10;
212 let port = pick_unused_port().unwrap();
213 spawn(app.serve(
214 RateLimitListener::with_port(port, limit),
215 StaticVer01::instance(),
216 ));
217 let client = Client::new(format!("http://localhost:{port}").parse().unwrap()).await;
218
219 let reqs = (0..limit)
221 .map(|_| spawn(client.get("mod/test").send()))
222 .collect::<Vec<_>>();
223
224 sleep(Duration::from_secs(5)).await;
226
227 let res = client.get("mod/test").send().await.unwrap();
229 assert_eq!(StatusCode::TOO_MANY_REQUESTS, res.status());
230
231 for res in try_join_all(reqs).await.unwrap() {
233 assert_eq!(StatusCode::OK, res.status());
234 }
235 }
236}