| 1 |
|
| 2 |
|
| 3 |
|
| 4 |
use tower_governor::errors::GovernorError; |
| 5 |
use tower_governor::key_extractor::KeyExtractor; |
| 6 |
|
| 7 |
use crate::db::SyncAppId; |
| 8 |
|
| 9 |
|
| 10 |
|
| 11 |
|
| 12 |
|
| 13 |
|
| 14 |
|
| 15 |
|
| 16 |
|
| 17 |
|
| 18 |
|
| 19 |
|
| 20 |
|
| 21 |
|
| 22 |
|
| 23 |
|
| 24 |
|
| 25 |
|
| 26 |
|
| 27 |
#[derive(Debug, Clone, Copy, PartialEq, Eq)] |
| 28 |
pub struct CloudflareIpKeyExtractor; |
| 29 |
|
| 30 |
impl KeyExtractor for CloudflareIpKeyExtractor { |
| 31 |
type Key = std::net::IpAddr; |
| 32 |
|
| 33 |
fn extract<T>(&self, req: &axum::http::Request<T>) -> Result<Self::Key, GovernorError> { |
| 34 |
if let Some(ip) = req |
| 35 |
.headers() |
| 36 |
.get("cf-connecting-ip") |
| 37 |
.and_then(|v: &axum::http::HeaderValue| v.to_str().ok()) |
| 38 |
.and_then(|s: &str| s.trim().parse::<std::net::IpAddr>().ok()) |
| 39 |
{ |
| 40 |
return Ok(ip); |
| 41 |
} |
| 42 |
|
| 43 |
tower_governor::key_extractor::PeerIpKeyExtractor.extract(req) |
| 44 |
} |
| 45 |
} |
| 46 |
|
| 47 |
|
| 48 |
|
| 49 |
|
| 50 |
|
| 51 |
|
| 52 |
|
| 53 |
|
| 54 |
|
| 55 |
|
| 56 |
|
| 57 |
|
| 58 |
|
| 59 |
|
| 60 |
|
| 61 |
|
| 62 |
|
| 63 |
#[derive(Debug, Clone)] |
| 64 |
pub struct SyncAppKeyExtractor { |
| 65 |
|
| 66 |
|
| 67 |
|
| 68 |
|
| 69 |
|
| 70 |
secret: Option<std::sync::Arc<String>>, |
| 71 |
} |
| 72 |
|
| 73 |
impl SyncAppKeyExtractor { |
| 74 |
pub fn new(secret: Option<std::sync::Arc<String>>) -> Self { |
| 75 |
Self { secret } |
| 76 |
} |
| 77 |
|
| 78 |
|
| 79 |
|
| 80 |
fn verify_app(secret: &str, token: &str) -> Option<SyncAppId> { |
| 81 |
use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode}; |
| 82 |
|
| 83 |
#[derive(serde::Deserialize)] |
| 84 |
struct AppClaim { |
| 85 |
app: SyncAppId, |
| 86 |
} |
| 87 |
|
| 88 |
let mut validation = Validation::new(Algorithm::HS256); |
| 89 |
|
| 90 |
|
| 91 |
|
| 92 |
validation.validate_exp = false; |
| 93 |
validation.required_spec_claims.clear(); |
| 94 |
|
| 95 |
decode::<AppClaim>( |
| 96 |
token, |
| 97 |
&DecodingKey::from_secret(secret.as_bytes()), |
| 98 |
&validation, |
| 99 |
) |
| 100 |
.ok() |
| 101 |
.map(|data| data.claims.app) |
| 102 |
} |
| 103 |
} |
| 104 |
|
| 105 |
impl KeyExtractor for SyncAppKeyExtractor { |
| 106 |
type Key = SyncAppId; |
| 107 |
|
| 108 |
fn extract<T>(&self, req: &axum::http::Request<T>) -> Result<Self::Key, GovernorError> { |
| 109 |
let Some(token) = req |
| 110 |
.headers() |
| 111 |
.get("authorization") |
| 112 |
.and_then(|v| v.to_str().ok()) |
| 113 |
.and_then(|s| s.strip_prefix("Bearer ")) |
| 114 |
else { |
| 115 |
|
| 116 |
|
| 117 |
return Ok(SyncAppId::nil()); |
| 118 |
}; |
| 119 |
|
| 120 |
let app = match &self.secret { |
| 121 |
Some(secret) => Self::verify_app(secret, token), |
| 122 |
|
| 123 |
|
| 124 |
|
| 125 |
|
| 126 |
|
| 127 |
None => None, |
| 128 |
}; |
| 129 |
|
| 130 |
|
| 131 |
|
| 132 |
Ok(app.unwrap_or_else(SyncAppId::nil)) |
| 133 |
} |
| 134 |
} |
| 135 |
|
| 136 |
|
| 137 |
|
| 138 |
|
| 139 |
|
| 140 |
|
| 141 |
|
| 142 |
|
| 143 |
|
| 144 |
|
| 145 |
|
| 146 |
|
| 147 |
|
| 148 |
|
| 149 |
|
| 150 |
static GOVERNOR_SWEEPERS: std::sync::Mutex<Vec<Box<dyn Fn() -> usize + Send + Sync>>> = |
| 151 |
std::sync::Mutex::new(Vec::new()); |
| 152 |
|
| 153 |
|
| 154 |
|
| 155 |
fn register_for_sweep(hook: impl Fn() -> usize + Send + Sync + 'static) { |
| 156 |
if let Ok(mut hooks) = GOVERNOR_SWEEPERS.lock() { |
| 157 |
hooks.push(Box::new(hook)); |
| 158 |
} |
| 159 |
} |
| 160 |
|
| 161 |
|
| 162 |
|
| 163 |
|
| 164 |
pub fn start_governor_sweeper() { |
| 165 |
static STARTED: std::sync::Once = std::sync::Once::new(); |
| 166 |
STARTED.call_once(|| { |
| 167 |
tokio::spawn(async { |
| 168 |
let interval = |
| 169 |
std::time::Duration::from_secs(crate::constants::GOVERNOR_SWEEP_INTERVAL_SECS); |
| 170 |
loop { |
| 171 |
tokio::time::sleep(interval).await; |
| 172 |
|
| 173 |
let (limiters, retained) = { |
| 174 |
let Ok(hooks) = GOVERNOR_SWEEPERS.lock() else { |
| 175 |
continue; |
| 176 |
}; |
| 177 |
let retained: usize = hooks.iter().map(|hook| hook()).sum(); |
| 178 |
(hooks.len(), retained) |
| 179 |
}; |
| 180 |
tracing::debug!( |
| 181 |
limiters, |
| 182 |
retained_keys = retained, |
| 183 |
"swept governor bucket maps" |
| 184 |
); |
| 185 |
} |
| 186 |
}); |
| 187 |
}); |
| 188 |
} |
| 189 |
|
| 190 |
|
| 191 |
pub fn rate_limiter_ms( |
| 192 |
ms: u64, |
| 193 |
burst: u32, |
| 194 |
) -> std::sync::Arc< |
| 195 |
tower_governor::governor::GovernorConfig< |
| 196 |
CloudflareIpKeyExtractor, |
| 197 |
::governor::middleware::StateInformationMiddleware, |
| 198 |
>, |
| 199 |
> { |
| 200 |
let config = std::sync::Arc::new( |
| 201 |
tower_governor::governor::GovernorConfigBuilder::default() |
| 202 |
.key_extractor(CloudflareIpKeyExtractor) |
| 203 |
.per_millisecond(ms) |
| 204 |
.burst_size(burst) |
| 205 |
.use_headers() |
| 206 |
.finish() |
| 207 |
.expect("rate limiter config"), |
| 208 |
); |
| 209 |
let limiter = config.limiter().clone(); |
| 210 |
register_for_sweep(move || { |
| 211 |
limiter.retain_recent(); |
| 212 |
limiter.len() |
| 213 |
}); |
| 214 |
config |
| 215 |
} |
| 216 |
|
| 217 |
|
| 218 |
pub fn rate_limiter_per_sec( |
| 219 |
per_sec: u64, |
| 220 |
burst: u32, |
| 221 |
) -> std::sync::Arc< |
| 222 |
tower_governor::governor::GovernorConfig< |
| 223 |
CloudflareIpKeyExtractor, |
| 224 |
::governor::middleware::StateInformationMiddleware, |
| 225 |
>, |
| 226 |
> { |
| 227 |
let config = std::sync::Arc::new( |
| 228 |
tower_governor::governor::GovernorConfigBuilder::default() |
| 229 |
.key_extractor(CloudflareIpKeyExtractor) |
| 230 |
.per_second(per_sec) |
| 231 |
.burst_size(burst) |
| 232 |
.use_headers() |
| 233 |
.finish() |
| 234 |
.expect("rate limiter config"), |
| 235 |
); |
| 236 |
let limiter = config.limiter().clone(); |
| 237 |
register_for_sweep(move || { |
| 238 |
limiter.retain_recent(); |
| 239 |
limiter.len() |
| 240 |
}); |
| 241 |
config |
| 242 |
} |
| 243 |
|
| 244 |
|
| 245 |
|
| 246 |
|
| 247 |
|
| 248 |
pub fn synckit_app_rate_limiter_ms( |
| 249 |
secret: Option<std::sync::Arc<String>>, |
| 250 |
ms: u64, |
| 251 |
burst: u32, |
| 252 |
) -> std::sync::Arc< |
| 253 |
tower_governor::governor::GovernorConfig< |
| 254 |
SyncAppKeyExtractor, |
| 255 |
::governor::middleware::StateInformationMiddleware, |
| 256 |
>, |
| 257 |
> { |
| 258 |
let config = std::sync::Arc::new( |
| 259 |
tower_governor::governor::GovernorConfigBuilder::default() |
| 260 |
.key_extractor(SyncAppKeyExtractor::new(secret)) |
| 261 |
.per_millisecond(ms) |
| 262 |
.burst_size(burst) |
| 263 |
.use_headers() |
| 264 |
.finish() |
| 265 |
.expect("synckit app rate limiter config"), |
| 266 |
); |
| 267 |
let limiter = config.limiter().clone(); |
| 268 |
register_for_sweep(move || { |
| 269 |
limiter.retain_recent(); |
| 270 |
limiter.len() |
| 271 |
}); |
| 272 |
config |
| 273 |
} |
| 274 |
|
| 275 |
#[cfg(test)] |
| 276 |
mod tests { |
| 277 |
use super::*; |
| 278 |
use axum::http::Request; |
| 279 |
use base64::Engine; |
| 280 |
use tower_governor::key_extractor::KeyExtractor; |
| 281 |
|
| 282 |
|
| 283 |
fn fake_jwt(app_id: &SyncAppId) -> String { |
| 284 |
let header = base64::engine::general_purpose::URL_SAFE_NO_PAD |
| 285 |
.encode(r#"{"alg":"HS256","typ":"JWT"}"#); |
| 286 |
let payload_json = serde_json::json!({ |
| 287 |
"sub": "00000000-0000-0000-0000-000000000001", |
| 288 |
"app": app_id, |
| 289 |
"iss": "makenotwork-synckit", |
| 290 |
"exp": 9_999_999_999_i64, |
| 291 |
"iat": 1_000_000_000_i64, |
| 292 |
}); |
| 293 |
let payload = |
| 294 |
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(payload_json.to_string()); |
| 295 |
let sig = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode("fakesig"); |
| 296 |
format!("{header}.{payload}.{sig}") |
| 297 |
} |
| 298 |
|
| 299 |
|
| 300 |
fn signed_jwt(secret: &str, app_id: &SyncAppId) -> String { |
| 301 |
use jsonwebtoken::{Algorithm, EncodingKey, Header, encode}; |
| 302 |
let claims = serde_json::json!({ |
| 303 |
"sub": "00000000-0000-0000-0000-000000000001", |
| 304 |
"app": app_id, |
| 305 |
"iss": "makenotwork-synckit", |
| 306 |
"exp": 9_999_999_999_i64, |
| 307 |
"iat": 1_000_000_000_i64, |
| 308 |
}); |
| 309 |
encode( |
| 310 |
&Header::new(Algorithm::HS256), |
| 311 |
&claims, |
| 312 |
&EncodingKey::from_secret(secret.as_bytes()), |
| 313 |
) |
| 314 |
.unwrap() |
| 315 |
} |
| 316 |
|
| 317 |
#[test] |
| 318 |
fn no_secret_collapses_unverified_app_to_nil() { |
| 319 |
|
| 320 |
|
| 321 |
|
| 322 |
let app_id = SyncAppId::new(); |
| 323 |
let token = fake_jwt(&app_id); |
| 324 |
|
| 325 |
let req = Request::builder() |
| 326 |
.header("authorization", format!("Bearer {token}")) |
| 327 |
.body(()) |
| 328 |
.unwrap(); |
| 329 |
|
| 330 |
let extracted = SyncAppKeyExtractor::new(None).extract(&req).unwrap(); |
| 331 |
assert_eq!(extracted, SyncAppId::nil()); |
| 332 |
} |
| 333 |
|
| 334 |
#[test] |
| 335 |
fn verified_extracts_app_id_from_validly_signed_jwt() { |
| 336 |
let secret = "test-secret-key-for-synckit-jwt".to_string(); |
| 337 |
let app_id = SyncAppId::new(); |
| 338 |
let token = signed_jwt(&secret, &app_id); |
| 339 |
|
| 340 |
let req = Request::builder() |
| 341 |
.header("authorization", format!("Bearer {token}")) |
| 342 |
.body(()) |
| 343 |
.unwrap(); |
| 344 |
|
| 345 |
let extractor = SyncAppKeyExtractor::new(Some(std::sync::Arc::new(secret))); |
| 346 |
assert_eq!(extractor.extract(&req).unwrap(), app_id); |
| 347 |
} |
| 348 |
|
| 349 |
#[test] |
| 350 |
fn verified_forged_token_collapses_to_nil_bucket() { |
| 351 |
|
| 352 |
|
| 353 |
|
| 354 |
|
| 355 |
let secret = "test-secret-key-for-synckit-jwt".to_string(); |
| 356 |
let attacker_app = SyncAppId::new(); |
| 357 |
let forged = fake_jwt(&attacker_app); |
| 358 |
|
| 359 |
let req = Request::builder() |
| 360 |
.header("authorization", format!("Bearer {forged}")) |
| 361 |
.body(()) |
| 362 |
.unwrap(); |
| 363 |
|
| 364 |
let extractor = SyncAppKeyExtractor::new(Some(std::sync::Arc::new(secret))); |
| 365 |
assert_eq!(extractor.extract(&req).unwrap(), SyncAppId::nil()); |
| 366 |
} |
| 367 |
|
| 368 |
#[test] |
| 369 |
fn verified_spray_of_forged_apps_all_share_one_bucket() { |
| 370 |
|
| 371 |
|
| 372 |
let secret = std::sync::Arc::new("test-secret-key-for-synckit-jwt".to_string()); |
| 373 |
let extractor = SyncAppKeyExtractor::new(Some(secret)); |
| 374 |
for _ in 0..5 { |
| 375 |
let forged = fake_jwt(&SyncAppId::new()); |
| 376 |
let req = Request::builder() |
| 377 |
.header("authorization", format!("Bearer {forged}")) |
| 378 |
.body(()) |
| 379 |
.unwrap(); |
| 380 |
assert_eq!(extractor.extract(&req).unwrap(), SyncAppId::nil()); |
| 381 |
} |
| 382 |
} |
| 383 |
|
| 384 |
#[test] |
| 385 |
fn missing_auth_header_returns_nil_sentinel() { |
| 386 |
let req = Request::builder().body(()).unwrap(); |
| 387 |
let key = SyncAppKeyExtractor::new(None).extract(&req).unwrap(); |
| 388 |
assert_eq!(key, SyncAppId::nil()); |
| 389 |
} |
| 390 |
|
| 391 |
#[test] |
| 392 |
fn non_bearer_auth_returns_nil_sentinel() { |
| 393 |
let req = Request::builder() |
| 394 |
.header("authorization", "Basic dXNlcjpwYXNz") |
| 395 |
.body(()) |
| 396 |
.unwrap(); |
| 397 |
let key = SyncAppKeyExtractor::new(None).extract(&req).unwrap(); |
| 398 |
assert_eq!(key, SyncAppId::nil()); |
| 399 |
} |
| 400 |
|
| 401 |
#[test] |
| 402 |
fn malformed_jwt_collapses_to_nil_sentinel() { |
| 403 |
|
| 404 |
|
| 405 |
|
| 406 |
let req = Request::builder() |
| 407 |
.header("authorization", "Bearer not-a-jwt") |
| 408 |
.body(()) |
| 409 |
.unwrap(); |
| 410 |
let key = SyncAppKeyExtractor::new(None).extract(&req).unwrap(); |
| 411 |
assert_eq!(key, SyncAppId::nil()); |
| 412 |
} |
| 413 |
|
| 414 |
#[test] |
| 415 |
fn jwt_missing_app_claim_collapses_to_nil_sentinel() { |
| 416 |
let header = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(r#"{"alg":"HS256"}"#); |
| 417 |
let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD |
| 418 |
.encode(r#"{"sub":"user","iss":"test"}"#); |
| 419 |
let sig = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode("sig"); |
| 420 |
let token = format!("{header}.{payload}.{sig}"); |
| 421 |
|
| 422 |
let req = Request::builder() |
| 423 |
.header("authorization", format!("Bearer {token}")) |
| 424 |
.body(()) |
| 425 |
.unwrap(); |
| 426 |
let key = SyncAppKeyExtractor::new(None).extract(&req).unwrap(); |
| 427 |
assert_eq!(key, SyncAppId::nil()); |
| 428 |
} |
| 429 |
|
| 430 |
#[test] |
| 431 |
fn cf_connecting_ip_is_used_when_present() { |
| 432 |
let req = Request::builder() |
| 433 |
.header("cf-connecting-ip", "203.0.113.7") |
| 434 |
.body(()) |
| 435 |
.unwrap(); |
| 436 |
let key = CloudflareIpKeyExtractor.extract(&req).unwrap(); |
| 437 |
assert_eq!(key, "203.0.113.7".parse::<std::net::IpAddr>().unwrap()); |
| 438 |
} |
| 439 |
|
| 440 |
#[test] |
| 441 |
fn forged_x_forwarded_for_is_not_trusted() { |
| 442 |
|
| 443 |
|
| 444 |
|
| 445 |
|
| 446 |
|
| 447 |
let req = Request::builder() |
| 448 |
.header("x-forwarded-for", "1.2.3.4") |
| 449 |
.header("x-real-ip", "1.2.3.4") |
| 450 |
.body(()) |
| 451 |
.unwrap(); |
| 452 |
let result = CloudflareIpKeyExtractor.extract(&req); |
| 453 |
assert!( |
| 454 |
result.is_err(), |
| 455 |
"forged XFF/X-Real-IP must not yield a per-IP bucket; got {result:?}" |
| 456 |
); |
| 457 |
} |
| 458 |
|
| 459 |
#[test] |
| 460 |
fn unsigned_apps_collapse_to_nil_when_no_secret() { |
| 461 |
|
| 462 |
|
| 463 |
|
| 464 |
|
| 465 |
let app1 = SyncAppId::new(); |
| 466 |
let app2 = SyncAppId::new(); |
| 467 |
|
| 468 |
let req1 = Request::builder() |
| 469 |
.header("authorization", format!("Bearer {}", fake_jwt(&app1))) |
| 470 |
.body(()) |
| 471 |
.unwrap(); |
| 472 |
let req2 = Request::builder() |
| 473 |
.header("authorization", format!("Bearer {}", fake_jwt(&app2))) |
| 474 |
.body(()) |
| 475 |
.unwrap(); |
| 476 |
|
| 477 |
let extractor = SyncAppKeyExtractor::new(None); |
| 478 |
let key1 = extractor.extract(&req1).unwrap(); |
| 479 |
let key2 = extractor.extract(&req2).unwrap(); |
| 480 |
assert_eq!(key1, SyncAppId::nil()); |
| 481 |
assert_eq!(key2, SyncAppId::nil()); |
| 482 |
} |
| 483 |
} |
| 484 |
|