diff --git a/Cargo.lock b/Cargo.lock index c219a3c..c9c69c1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -35,6 +35,8 @@ dependencies = [ "chrono", "clap", "futures", + "governor", + "http-body-util", "log", "poem", "reqwest", @@ -65,6 +67,12 @@ dependencies = [ "alloc-no-stdlib", ] +[[package]] +name = "allocator-api2" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" + [[package]] name = "android_system_properties" version = "0.1.5" @@ -386,6 +394,12 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "crossbeam-utils" +version = "0.8.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" + [[package]] name = "crypto-common" version = "0.1.6" @@ -396,6 +410,20 @@ dependencies = [ "typenum", ] +[[package]] +name = "dashmap" +version = "6.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5041cc499144891f3790297212f32a74fb938e5136a14943f338ef9e0ae276cf" +dependencies = [ + "cfg-if", + "crossbeam-utils", + "hashbrown 0.14.5", + "lock_api", + "once_cell", + "parking_lot_core 0.9.11", +] + [[package]] name = "digest" version = "0.10.7" @@ -477,6 +505,12 @@ version = "1.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + [[package]] name = "foreign-types" version = "0.3.2" @@ -572,6 +606,12 @@ version = "0.3.31" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f90f7dce0722e95104fcb095585910c0977252f286e354b5e3bd38902cd99988" +[[package]] +name = "futures-timer" +version = "3.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f288b0a4f20f9a56b5d1da57e2227c661b7b16168e2f72365f57b63326e29b24" + [[package]] name = "futures-util" version = "0.3.31" @@ -620,9 +660,11 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "26145e563e54f2cadc477553f1ec5ee650b00862f0a58bcd12cbdc5f0ea2d2f4" dependencies = [ "cfg-if", + "js-sys", "libc", "r-efi", "wasi 0.14.5+wasi-0.2.4", + "wasm-bindgen", ] [[package]] @@ -631,6 +673,29 @@ version = "0.31.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "07e28edb80900c19c28f1072f2e8aeca7fa06b23cd4169cefe1af5aa3260783f" +[[package]] +name = "governor" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "444405bbb1a762387aa22dd569429533b54a1d8759d35d3b64cb39b0293eaa19" +dependencies = [ + "cfg-if", + "dashmap", + "futures-sink", + "futures-timer", + "futures-util", + "getrandom 0.3.3", + "hashbrown 0.15.5", + "nonzero_ext", + "parking_lot 0.12.4", + "portable-atomic", + "quanta", + "rand 0.9.2", + "smallvec", + "spinning_top", + "web-time", +] + [[package]] name = "h2" version = "0.4.12" @@ -650,11 +715,22 @@ dependencies = [ "tracing", ] +[[package]] +name = "hashbrown" +version = "0.14.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1" + [[package]] name = "hashbrown" version = "0.15.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" +dependencies = [ + "allocator-api2", + "equivalent", + "foldhash", +] [[package]] name = "headers" @@ -960,7 +1036,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4b0f83760fb341a774ed326568e19f5a863af4a952def8c39f9ab92fd95b88e5" dependencies = [ "equivalent", - "hashbrown", + "hashbrown 0.15.5", ] [[package]] @@ -1165,6 +1241,12 @@ dependencies = [ "libc", ] +[[package]] +name = "nonzero_ext" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38bf9645c8b145698bb0b18a4637dcacbc421ea49bef2317e4fd8065a387cf21" + [[package]] name = "nu-ansi-term" version = "0.50.1" @@ -1384,6 +1466,12 @@ dependencies = [ "syn", ] +[[package]] +name = "portable-atomic" +version = "1.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f84267b20a16ea918e43c6a88433c2d54fa145c92a811b5b047ccbe153674483" + [[package]] name = "postgres-protocol" version = "0.6.8" @@ -1452,6 +1540,21 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "quanta" +version = "0.12.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f3ab5a9d756f0d97bdc89019bd2e4ea098cf9cde50ee7564dde6b81ccc8f06c7" +dependencies = [ + "crossbeam-utils", + "libc", + "once_cell", + "raw-cpuid", + "wasi 0.11.1+wasi-snapshot-preview1", + "web-sys", + "winapi", +] + [[package]] name = "quote" version = "1.0.40" @@ -1526,6 +1629,15 @@ dependencies = [ "getrandom 0.3.3", ] +[[package]] +name = "raw-cpuid" +version = "11.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "498cd0dc59d73224351ee52a95fee0f1a617a2eae0e7d9d720cc622c73a54186" +dependencies = [ + "bitflags 2.9.4", +] + [[package]] name = "redox_syscall" version = "0.2.16" @@ -1925,6 +2037,15 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "spinning_top" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d96d2d1d716fb500937168cc09353ffdc7a012be8475ac7308e1bdf0e3923300" +dependencies = [ + "lock_api", +] + [[package]] name = "stable_deref_trait" version = "1.2.0" @@ -2576,6 +2697,16 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "web-time" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a6580f308b1fad9207618087a65c04e7a10bc77e02c8e84e9b00dd4b12fa0bb" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + [[package]] name = "whoami" version = "1.6.1" diff --git a/Cargo.toml b/Cargo.toml index 4d0750b..9d9cc15 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -12,6 +12,8 @@ async-compression = { version = "0.4.30", features = ["futures-io", "tokio", "gz chrono = { version = "0.4.42", features = ["serde"] } clap = { version = "4.5.47", features = ["derive", "env"] } futures = "0.3.31" +governor = "0.10.1" +http-body-util = "0.1.3" log = "0.4.28" poem = { version = "3.1.12", features = ["compression"] } reqwest = { version = "0.12.23", features = ["stream"] } diff --git a/src/lib.rs b/src/lib.rs index 5f5f43f..9b3dfdb 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -5,6 +5,7 @@ mod client; mod mirror; mod plc_pg; mod poll; +mod ratelimit; mod weekly; pub use backfill::backfill; @@ -12,6 +13,7 @@ pub use client::CLIENT; pub use mirror::serve; pub use plc_pg::{Db, backfill_to_pg, pages_to_pg}; pub use poll::{PageBoundaryState, get_page, poll_upstream}; +pub use ratelimit::GovernorMiddleware; pub use weekly::{BundleSource, FolderSource, HttpSource, Week, pages_to_weeks, week_to_pages}; pub type Dt = chrono::DateTime; diff --git a/src/mirror.rs b/src/mirror.rs index 0f2aba9..da0d264 100644 --- a/src/mirror.rs +++ b/src/mirror.rs @@ -1,4 +1,5 @@ -use crate::logo; +use crate::{GovernorMiddleware, logo}; +use futures::TryStreamExt; use poem::{ EndpointExt, Error, IntoResponse, Request, Response, Result, Route, Server, get, handler, http::{StatusCode, Uri}, @@ -7,8 +8,7 @@ use poem::{ web::Data, }; use reqwest::{Client, Url}; -use std::net::SocketAddr; -use std::time::Duration; +use std::{net::SocketAddr, time::Duration}; #[derive(Debug, Clone)] struct State { @@ -73,14 +73,24 @@ async fn proxy(req: &Request, Data(state): Data<&State>) -> Result = upstream_res.into(); + let (parts, reqw_body) = http_res.into_parts(); + + let parts = poem::ResponseParts { + status: parts.status, + version: parts.version, + headers: parts.headers, + extensions: parts.extensions, + }; + + let body = http_body_util::BodyDataStream::new(reqw_body) + .map_err(|e| std::io::Error::other(Box::new(e))); + + Ok(Response::from_parts( + parts, + poem::Body::from_bytes_stream(body), + )) } #[handler] @@ -107,6 +117,7 @@ pub async fn serve(upstream: &Url, plc: Url, bind: SocketAddr) -> std::io::Resul .with(AddData::new(state)) .with(Cors::new().allow_credentials(false)) .with(Compression::new()) + .with(GovernorMiddleware::per_minute(3000).unwrap()) .with(CatchPanic::new()) .with(Tracing); diff --git a/src/ratelimit.rs b/src/ratelimit.rs new file mode 100644 index 0000000..2fb140b --- /dev/null +++ b/src/ratelimit.rs @@ -0,0 +1,141 @@ +use crate::logo; + +use governor::{ + Quota, RateLimiter, + clock::{Clock, DefaultClock}, + state::keyed::DefaultKeyedStateStore, +}; +use poem::{Endpoint, Middleware, Request, Response, Result, http::StatusCode}; +use std::{ + convert::TryInto, error::Error, net::IpAddr, num::NonZeroU32, sync::Arc, sync::LazyLock, + time::Duration, +}; + +static CLOCK: LazyLock = LazyLock::new(DefaultClock::default); + +/// Once the rate limit has been reached, the middleware will respond with +/// status code 429 (too many requests) and a `Retry-After` header with the amount +/// of time that needs to pass before another request will be allowed. +#[derive(Debug, Clone)] +pub struct GovernorMiddleware { + limiter: Arc, DefaultClock>>, +} + +impl GovernorMiddleware { + /// Constructs a rate-limiting middleware from a [`Duration`] that allows one request in the given time interval. + /// + /// If the time interval is zero, returns `None`. + #[must_use] + pub fn with_period(duration: Duration) -> Option { + Some(Self { + limiter: Arc::new(RateLimiter::::keyed(Quota::with_period( + duration, + )?)), + }) + } + + /// Constructs a rate-limiting middleware that allows a specified number of requests every second. + /// + /// Returns an error if `times` can't be converted into a [`NonZeroU32`]. + pub fn per_second(times: T) -> Result + where + T: TryInto, + T::Error: Error + Send + Sync + 'static, + { + Ok(Self { + limiter: Arc::new(RateLimiter::::keyed(Quota::per_second( + times.try_into().unwrap(), // TODO + ))), + }) + } + + /// Constructs a rate-limiting middleware that allows a specified number of requests every minute. + /// + /// Returns an error if `times` can't be converted into a [`NonZeroU32`]. + pub fn per_minute(times: T) -> Result + where + T: TryInto, + T::Error: Error + Send + Sync + 'static, + { + Ok(Self { + limiter: Arc::new(RateLimiter::::keyed(Quota::per_minute( + times.try_into().unwrap(), // TODO + ))), + }) + } + + /// Constructs a rate-limiting middleware that allows a specified number of requests every hour. + /// + /// Returns an error if `times` can't be converted into a [`NonZeroU32`]. + pub fn per_hour(times: T) -> Result + where + T: TryInto, + T::Error: Error + Send + Sync + 'static, + { + Ok(Self { + limiter: Arc::new(RateLimiter::::keyed(Quota::per_hour( + times.try_into().unwrap(), // TODO + ))), + }) + } +} + +impl Middleware for GovernorMiddleware { + type Output = GovernorMiddlewareImpl; + fn transform(&self, ep: E) -> Self::Output { + GovernorMiddlewareImpl { + ep, + limiter: self.limiter.clone(), + } + } +} + +pub struct GovernorMiddlewareImpl { + ep: E, + limiter: Arc, DefaultClock>>, +} + +impl Endpoint for GovernorMiddlewareImpl { + type Output = E::Output; + + async fn call(&self, req: Request) -> Result { + let remote = req + .remote_addr() + .as_socket_addr() + .unwrap_or_else(|| panic!("failed to get request's remote addr")) // TODO + .ip(); + + log::trace!("remote: {remote}"); + + match self.limiter.check_key(&remote) { + Ok(_) => { + log::debug!("allowing remote {remote}"); + self.ep.call(req).await + } + Err(negative) => { + let wait_time = negative.wait_time_from(CLOCK.now()).as_secs(); + + log::debug!("rate limit exceeded for {remote}, quota reset in {wait_time}s"); + + let res = Response::builder() + .status(StatusCode::TOO_MANY_REQUESTS) + .header("x-ratelimit-after", wait_time) + .header("retry-after", wait_time) + .body(booo()); + Err(poem::Error::from_response(res)) + } + } + } +} + +fn booo() -> String { + format!( + r#"{} + +You're going a bit too fast. + +Tip: check out the `x-ratelimit-after` response header. +"#, + logo("mirror 429") + ) +}