Skip to main content

max / makenotwork

8.2 KB · 247 lines History Blame Raw
1 //! Rate limiting: Cloudflare-aware IP extraction, per-app SyncKit extraction,
2 //! and governor config builders.
3
4 use base64::Engine;
5 use tower_governor::errors::GovernorError;
6 use tower_governor::key_extractor::KeyExtractor;
7
8 use crate::db::SyncAppId;
9
10 /// IP key extractor that prefers `CF-Connecting-IP` (set by Cloudflare, cannot
11 /// be spoofed by clients) over `X-Forwarded-For` (which can be spoofed if the
12 /// proxy chain doesn't strip it). Falls back to `SmartIpKeyExtractor` behavior
13 /// when `CF-Connecting-IP` is absent (e.g., direct/dev access without Cloudflare).
14 #[derive(Debug, Clone, Copy, PartialEq, Eq)]
15 pub struct CloudflareIpKeyExtractor;
16
17 impl KeyExtractor for CloudflareIpKeyExtractor {
18 type Key = std::net::IpAddr;
19
20 fn extract<T>(&self, req: &axum::http::Request<T>) -> Result<Self::Key, GovernorError> {
21 if let Some(ip) = req
22 .headers()
23 .get("cf-connecting-ip")
24 .and_then(|v: &axum::http::HeaderValue| v.to_str().ok())
25 .and_then(|s: &str| s.trim().parse::<std::net::IpAddr>().ok())
26 {
27 return Ok(ip);
28 }
29
30 tower_governor::key_extractor::SmartIpKeyExtractor.extract(req)
31 }
32 }
33
34 /// Per-SyncKit-app key extractor. Decodes the JWT payload from the
35 /// `Authorization: Bearer <token>` header to extract the `app` claim.
36 ///
37 /// This does NOT verify the JWT signature — that's the handler's job via
38 /// `SyncUser`. The rate limiter only needs a consistent key to bucket
39 /// requests. A forged JWT gets a rate-limit bucket but is rejected by the
40 /// handler; the IP-based layer provides defense-in-depth against abuse.
41 #[derive(Debug, Clone, Copy, PartialEq, Eq)]
42 pub struct SyncAppKeyExtractor;
43
44 impl KeyExtractor for SyncAppKeyExtractor {
45 type Key = SyncAppId;
46
47 fn extract<T>(&self, req: &axum::http::Request<T>) -> Result<Self::Key, GovernorError> {
48 let token = match req
49 .headers()
50 .get("authorization")
51 .and_then(|v| v.to_str().ok())
52 .and_then(|s| s.strip_prefix("Bearer "))
53 {
54 Some(t) => t,
55 // No bearer token — use a nil sentinel key so the request passes
56 // through to the handler, where SyncUser will properly return 401.
57 None => return Ok(SyncAppId::nil()),
58 };
59
60 // JWT is header.payload.signature — decode the payload (middle segment)
61 // without verifying the signature. We only need the `app` field.
62 let payload_b64 = token
63 .split('.')
64 .nth(1)
65 .ok_or(GovernorError::UnableToExtractKey)?;
66
67 let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
68 .decode(payload_b64)
69 .or_else(|_| base64::engine::general_purpose::STANDARD.decode(payload_b64))
70 .map_err(|_| GovernorError::UnableToExtractKey)?;
71
72 #[derive(serde::Deserialize)]
73 struct AppClaim {
74 app: SyncAppId,
75 }
76
77 let claims: AppClaim =
78 serde_json::from_slice(&payload_bytes).map_err(|_| GovernorError::UnableToExtractKey)?;
79
80 Ok(claims.app)
81 }
82 }
83
84 // ── Config builders ──
85
86 /// Build an IP-based rate limiter from a per-millisecond interval and burst size.
87 pub fn rate_limiter_ms(
88 ms: u64,
89 burst: u32,
90 ) -> std::sync::Arc<
91 tower_governor::governor::GovernorConfig<
92 CloudflareIpKeyExtractor,
93 ::governor::middleware::StateInformationMiddleware,
94 >,
95 > {
96 std::sync::Arc::new(
97 tower_governor::governor::GovernorConfigBuilder::default()
98 .key_extractor(CloudflareIpKeyExtractor)
99 .per_millisecond(ms)
100 .burst_size(burst)
101 .use_headers()
102 .finish()
103 .expect("rate limiter config"),
104 )
105 }
106
107 /// Build an IP-based rate limiter from a per-second rate and burst size.
108 pub fn rate_limiter_per_sec(
109 per_sec: u64,
110 burst: u32,
111 ) -> std::sync::Arc<
112 tower_governor::governor::GovernorConfig<
113 CloudflareIpKeyExtractor,
114 ::governor::middleware::StateInformationMiddleware,
115 >,
116 > {
117 std::sync::Arc::new(
118 tower_governor::governor::GovernorConfigBuilder::default()
119 .key_extractor(CloudflareIpKeyExtractor)
120 .per_second(per_sec)
121 .burst_size(burst)
122 .use_headers()
123 .finish()
124 .expect("rate limiter config"),
125 )
126 }
127
128 /// Build a per-SyncKit-app rate limiter from a per-millisecond interval and burst size.
129 pub fn synckit_app_rate_limiter_ms(
130 ms: u64,
131 burst: u32,
132 ) -> std::sync::Arc<
133 tower_governor::governor::GovernorConfig<
134 SyncAppKeyExtractor,
135 ::governor::middleware::StateInformationMiddleware,
136 >,
137 > {
138 std::sync::Arc::new(
139 tower_governor::governor::GovernorConfigBuilder::default()
140 .key_extractor(SyncAppKeyExtractor)
141 .per_millisecond(ms)
142 .burst_size(burst)
143 .use_headers()
144 .finish()
145 .expect("synckit app rate limiter config"),
146 )
147 }
148
149 #[cfg(test)]
150 mod tests {
151 use super::*;
152 use axum::http::Request;
153 use tower_governor::key_extractor::KeyExtractor;
154
155 /// Build a fake JWT with the given app ID in the payload (no signature verification).
156 fn fake_jwt(app_id: &SyncAppId) -> String {
157 let header = base64::engine::general_purpose::URL_SAFE_NO_PAD
158 .encode(r#"{"alg":"HS256","typ":"JWT"}"#);
159 let payload_json = serde_json::json!({
160 "sub": "00000000-0000-0000-0000-000000000001",
161 "app": app_id,
162 "iss": "makenotwork-synckit",
163 "exp": 9999999999_i64,
164 "iat": 1000000000_i64,
165 });
166 let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD
167 .encode(payload_json.to_string());
168 let sig = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode("fakesig");
169 format!("{header}.{payload}.{sig}")
170 }
171
172 #[test]
173 fn extracts_app_id_from_jwt() {
174 let app_id = SyncAppId::new();
175 let token = fake_jwt(&app_id);
176
177 let req = Request::builder()
178 .header("authorization", format!("Bearer {token}"))
179 .body(())
180 .unwrap();
181
182 let extracted = SyncAppKeyExtractor.extract(&req).unwrap();
183 assert_eq!(extracted, app_id);
184 }
185
186 #[test]
187 fn missing_auth_header_returns_nil_sentinel() {
188 let req = Request::builder().body(()).unwrap();
189 let key = SyncAppKeyExtractor.extract(&req).unwrap();
190 assert_eq!(key, SyncAppId::nil());
191 }
192
193 #[test]
194 fn non_bearer_auth_returns_nil_sentinel() {
195 let req = Request::builder()
196 .header("authorization", "Basic dXNlcjpwYXNz")
197 .body(())
198 .unwrap();
199 let key = SyncAppKeyExtractor.extract(&req).unwrap();
200 assert_eq!(key, SyncAppId::nil());
201 }
202
203 #[test]
204 fn malformed_jwt_returns_error() {
205 let req = Request::builder()
206 .header("authorization", "Bearer not-a-jwt")
207 .body(())
208 .unwrap();
209 assert!(SyncAppKeyExtractor.extract(&req).is_err());
210 }
211
212 #[test]
213 fn jwt_missing_app_claim_returns_error() {
214 let header = base64::engine::general_purpose::URL_SAFE_NO_PAD
215 .encode(r#"{"alg":"HS256"}"#);
216 let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD
217 .encode(r#"{"sub":"user","iss":"test"}"#);
218 let sig = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode("sig");
219 let token = format!("{header}.{payload}.{sig}");
220
221 let req = Request::builder()
222 .header("authorization", format!("Bearer {token}"))
223 .body(())
224 .unwrap();
225 assert!(SyncAppKeyExtractor.extract(&req).is_err());
226 }
227
228 #[test]
229 fn different_apps_get_different_keys() {
230 let app1 = SyncAppId::new();
231 let app2 = SyncAppId::new();
232
233 let req1 = Request::builder()
234 .header("authorization", format!("Bearer {}", fake_jwt(&app1)))
235 .body(())
236 .unwrap();
237 let req2 = Request::builder()
238 .header("authorization", format!("Bearer {}", fake_jwt(&app2)))
239 .body(())
240 .unwrap();
241
242 let key1 = SyncAppKeyExtractor.extract(&req1).unwrap();
243 let key2 = SyncAppKeyExtractor.extract(&req2).unwrap();
244 assert_ne!(key1, key2);
245 }
246 }
247