Skip to main content

tide_disco/
listener.rs

1// Copyright (c) 2022 Espresso Systems (espressosys.com)
2// This file is part of the tide-disco library.
3
4// You should have received a copy of the MIT License
5// along with the tide-disco library. If not, see <https://mit-license.org/>.
6
7use 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/// TCP listener which accepts only a limited number of connections at a time.
29///
30/// This listener is based on `tide::listener::TcpListener` and should match the semantics of that
31/// listener in every way, accept that when there are more simultaneous outstanding requests than
32/// the configured limit, excess requests will fail immediately with error code 429 (Too Many
33/// Requests).
34#[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    /// Listen at the given address.
46    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    /// Listen at the given port on all interfaces.
57    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        // Format the listen information.
78        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                            // Handle the request if we can get a permit.
112                            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                                // Otherwise, we are rate limited. Respond immediately with an
118                                // error.
119                                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                    // Make a really slow endpoint so we can have many simultaneous requests.
203                    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        // Start the maximum number of simultaneous requests.
220        let reqs = (0..limit)
221            .map(|_| spawn(client.get("mod/test").send()))
222            .collect::<Vec<_>>();
223
224        // Wait a bit for those requests to get accepted.
225        sleep(Duration::from_secs(5)).await;
226
227        // The next request gets rate limited.
228        let res = client.get("mod/test").send().await.unwrap();
229        assert_eq!(StatusCode::TOO_MANY_REQUESTS, res.status());
230
231        // The other requests eventually complete successfully.
232        for res in try_join_all(reqs).await.unwrap() {
233            assert_eq!(StatusCode::OK, res.status());
234        }
235    }
236}