| 1 |
|
| 2 |
|
| 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 |
|
| 11 |
|
| 12 |
|
| 13 |
|
| 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 |
|
| 35 |
|
| 36 |
|
| 37 |
|
| 38 |
|
| 39 |
|
| 40 |
|
| 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 |
|
| 56 |
|
| 57 |
None => return Ok(SyncAppId::nil()), |
| 58 |
}; |
| 59 |
|
| 60 |
|
| 61 |
|
| 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 |
|
| 85 |
|
| 86 |
|
| 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 |
|
| 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 |
|
| 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 |
|
| 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 |
|