//! Rate limiting: Cloudflare-aware IP extraction, per-app SyncKit extraction, //! and governor config builders. use base64::Engine; use tower_governor::errors::GovernorError; use tower_governor::key_extractor::KeyExtractor; use crate::db::SyncAppId; /// IP key extractor that prefers `CF-Connecting-IP` (set by Cloudflare, cannot /// be spoofed by clients) over `X-Forwarded-For` (which can be spoofed if the /// proxy chain doesn't strip it). Falls back to `SmartIpKeyExtractor` behavior /// when `CF-Connecting-IP` is absent (e.g., direct/dev access without Cloudflare). #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct CloudflareIpKeyExtractor; impl KeyExtractor for CloudflareIpKeyExtractor { type Key = std::net::IpAddr; fn extract(&self, req: &axum::http::Request) -> Result { if let Some(ip) = req .headers() .get("cf-connecting-ip") .and_then(|v: &axum::http::HeaderValue| v.to_str().ok()) .and_then(|s: &str| s.trim().parse::().ok()) { return Ok(ip); } tower_governor::key_extractor::SmartIpKeyExtractor.extract(req) } } /// Per-SyncKit-app key extractor. Decodes the JWT payload from the /// `Authorization: Bearer ` header to extract the `app` claim. /// /// This does NOT verify the JWT signature — that's the handler's job via /// `SyncUser`. The rate limiter only needs a consistent key to bucket /// requests. A forged JWT gets a rate-limit bucket but is rejected by the /// handler; the IP-based layer provides defense-in-depth against abuse. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct SyncAppKeyExtractor; impl KeyExtractor for SyncAppKeyExtractor { type Key = SyncAppId; fn extract(&self, req: &axum::http::Request) -> Result { let token = match req .headers() .get("authorization") .and_then(|v| v.to_str().ok()) .and_then(|s| s.strip_prefix("Bearer ")) { Some(t) => t, // No bearer token — use a nil sentinel key so the request passes // through to the handler, where SyncUser will properly return 401. None => return Ok(SyncAppId::nil()), }; // JWT is header.payload.signature — decode the payload (middle segment) // without verifying the signature. We only need the `app` field. let payload_b64 = token .split('.') .nth(1) .ok_or(GovernorError::UnableToExtractKey)?; let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD .decode(payload_b64) .or_else(|_| base64::engine::general_purpose::STANDARD.decode(payload_b64)) .map_err(|_| GovernorError::UnableToExtractKey)?; #[derive(serde::Deserialize)] struct AppClaim { app: SyncAppId, } let claims: AppClaim = serde_json::from_slice(&payload_bytes).map_err(|_| GovernorError::UnableToExtractKey)?; Ok(claims.app) } } // ── Config builders ── /// Build an IP-based rate limiter from a per-millisecond interval and burst size. pub fn rate_limiter_ms( ms: u64, burst: u32, ) -> std::sync::Arc< tower_governor::governor::GovernorConfig< CloudflareIpKeyExtractor, ::governor::middleware::StateInformationMiddleware, >, > { std::sync::Arc::new( tower_governor::governor::GovernorConfigBuilder::default() .key_extractor(CloudflareIpKeyExtractor) .per_millisecond(ms) .burst_size(burst) .use_headers() .finish() .expect("rate limiter config"), ) } /// Build an IP-based rate limiter from a per-second rate and burst size. pub fn rate_limiter_per_sec( per_sec: u64, burst: u32, ) -> std::sync::Arc< tower_governor::governor::GovernorConfig< CloudflareIpKeyExtractor, ::governor::middleware::StateInformationMiddleware, >, > { std::sync::Arc::new( tower_governor::governor::GovernorConfigBuilder::default() .key_extractor(CloudflareIpKeyExtractor) .per_second(per_sec) .burst_size(burst) .use_headers() .finish() .expect("rate limiter config"), ) } /// Build a per-SyncKit-app rate limiter from a per-millisecond interval and burst size. pub fn synckit_app_rate_limiter_ms( ms: u64, burst: u32, ) -> std::sync::Arc< tower_governor::governor::GovernorConfig< SyncAppKeyExtractor, ::governor::middleware::StateInformationMiddleware, >, > { std::sync::Arc::new( tower_governor::governor::GovernorConfigBuilder::default() .key_extractor(SyncAppKeyExtractor) .per_millisecond(ms) .burst_size(burst) .use_headers() .finish() .expect("synckit app rate limiter config"), ) } #[cfg(test)] mod tests { use super::*; use axum::http::Request; use tower_governor::key_extractor::KeyExtractor; /// Build a fake JWT with the given app ID in the payload (no signature verification). fn fake_jwt(app_id: &SyncAppId) -> String { let header = base64::engine::general_purpose::URL_SAFE_NO_PAD .encode(r#"{"alg":"HS256","typ":"JWT"}"#); let payload_json = serde_json::json!({ "sub": "00000000-0000-0000-0000-000000000001", "app": app_id, "iss": "makenotwork-synckit", "exp": 9999999999_i64, "iat": 1000000000_i64, }); let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD .encode(payload_json.to_string()); let sig = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode("fakesig"); format!("{header}.{payload}.{sig}") } #[test] fn extracts_app_id_from_jwt() { let app_id = SyncAppId::new(); let token = fake_jwt(&app_id); let req = Request::builder() .header("authorization", format!("Bearer {token}")) .body(()) .unwrap(); let extracted = SyncAppKeyExtractor.extract(&req).unwrap(); assert_eq!(extracted, app_id); } #[test] fn missing_auth_header_returns_nil_sentinel() { let req = Request::builder().body(()).unwrap(); let key = SyncAppKeyExtractor.extract(&req).unwrap(); assert_eq!(key, SyncAppId::nil()); } #[test] fn non_bearer_auth_returns_nil_sentinel() { let req = Request::builder() .header("authorization", "Basic dXNlcjpwYXNz") .body(()) .unwrap(); let key = SyncAppKeyExtractor.extract(&req).unwrap(); assert_eq!(key, SyncAppId::nil()); } #[test] fn malformed_jwt_returns_error() { let req = Request::builder() .header("authorization", "Bearer not-a-jwt") .body(()) .unwrap(); assert!(SyncAppKeyExtractor.extract(&req).is_err()); } #[test] fn jwt_missing_app_claim_returns_error() { let header = base64::engine::general_purpose::URL_SAFE_NO_PAD .encode(r#"{"alg":"HS256"}"#); let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD .encode(r#"{"sub":"user","iss":"test"}"#); let sig = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode("sig"); let token = format!("{header}.{payload}.{sig}"); let req = Request::builder() .header("authorization", format!("Bearer {token}")) .body(()) .unwrap(); assert!(SyncAppKeyExtractor.extract(&req).is_err()); } #[test] fn different_apps_get_different_keys() { let app1 = SyncAppId::new(); let app2 = SyncAppId::new(); let req1 = Request::builder() .header("authorization", format!("Bearer {}", fake_jwt(&app1))) .body(()) .unwrap(); let req2 = Request::builder() .header("authorization", format!("Bearer {}", fake_jwt(&app2))) .body(()) .unwrap(); let key1 = SyncAppKeyExtractor.extract(&req1).unwrap(); let key2 = SyncAppKeyExtractor.extract(&req2).unwrap(); assert_ne!(key1, key2); } }