- 1
use std::collections::HashMap; - 2
use std::net::IpAddr; - 3
use std::sync::Arc; - 4
use std::time::{Duration, Instant}; - 5
- 6
use axum::Json; - 7
use axum::extract::State; - 8
use axum::http::StatusCode; - 9
use axum::response::{IntoResponse, Response}; - 10
use serde::{Deserialize, Serialize}; - 11
use tokio::sync::RwLock; - 12
- 13
#[derive(Debug, Clone, Deserialize, Serialize)] - 14
pub struct RateLimitConfig { - 15
/// Max requests per window for `POST /gateway/inbound`. - 16
pub inbound_per_min: u32, - 17
/// Max requests per window for `POST /sessions`. - 18
pub sessions_per_min: u32, - 19
/// Max requests per window for `POST /sessions/{id}/run`. - 20
pub runs_per_min: u32, - 21
/// Max requests per window for all other POST endpoints. - 22
pub other_post_per_min: u32, - 23
/// Window duration in seconds. - 24
pub window_secs: u64, - 25
} - 26
- 27
impl Default for RateLimitConfig { - 28
fn default() -> Self { - 29
Self { - 30
inbound_per_min: 30, - 31
sessions_per_min: 5, - 32
runs_per_min: 10, - 33
other_post_per_min: 20, - 34
window_secs: 60, - 35
} - 36
} - 37
} - 38
- 39
impl RateLimitConfig { - 40
pub fn from_settings(s: Option<vak_config::RateLimitSettings>) -> Self { - 41
let d = Self::default(); - 42
match s { - 43
None => d, - 44
Some(s) => Self { - 45
inbound_per_min: s.inbound_per_min.unwrap_or(d.inbound_per_min), - 46
sessions_per_min: s.sessions_per_min.unwrap_or(d.sessions_per_min), - 47
runs_per_min: s.runs_per_min.unwrap_or(d.runs_per_min), - 48
other_post_per_min: s.other_post_per_min.unwrap_or(d.other_post_per_min), - 49
window_secs: s.window_secs.unwrap_or(d.window_secs), - 50
}, - 51
} - 52
} - 53
} - 54
- 55
#[derive(Debug, Clone)] - 56
struct WindowBucket { - 57
count: u32, - 58
window_start: Instant, - 59
} - 60
- 61
#[derive(Debug, Clone, Default)] - 62
struct IpWindow { - 63
buckets: HashMap<String, WindowBucket>, - 64
} - 65
- 66
impl IpWindow { - 67
fn check_and_increment(&mut self, key: &str, max: u32, window: Duration) -> Result<(), u64> { - 68
let now = Instant::now(); - 69
let bucket = self - 70
.buckets - 71
.entry(key.to_string()) - 72
.or_insert_with(|| WindowBucket { - 73
count: 0, - 74
window_start: now, - 75
}); - 76
- 77
if now.duration_since(bucket.window_start) > window { - 78
bucket.count = 0; - 79
bucket.window_start = now; - 80
} - 81
- 82
if bucket.count >= max { - 83
let wait_secs = window - 84
.checked_sub(now.duration_since(bucket.window_start)) - 85
.unwrap_or(Duration::ZERO) - 86
.as_secs() - 87
.max(1); - 88
return Err(wait_secs); - 89
} - 90
- 91
bucket.count += 1; - 92
Ok(()) - 93
} - 94
- 95
fn cleanup(&mut self, window: Duration) { - 96
let now = Instant::now(); - 97
self.buckets - 98
.retain(|_, b| now.duration_since(b.window_start) <= window); - 99
} - 100
} - 101
- 102
#[derive(Clone)] - 103
pub struct RateLimiter { - 104
inner: Arc<RwLock<HashMap<IpAddr, IpWindow>>>, - 105
config: RateLimitConfig, - 106
home: std::path::PathBuf, - 107
} - 108
- 109
impl RateLimiter { - 110
pub fn new(config: RateLimitConfig, home: std::path::PathBuf) -> Self { - 111
Self { - 112
inner: Arc::new(RwLock::new(HashMap::new())), - 113
config, - 114
home, - 115
} - 116
} - 117
- 118
async fn check(&self, ip: IpAddr, key: &str, max: u32) -> Result<(), u64> { - 119
let window = Duration::from_secs(self.config.window_secs); - 120
let mut map = self.inner.write().await; - 121
- 122
// Periodic cleanup: remove expired windows to bound memory - 123
if map.len() > 1000 { - 124
for w in map.values_mut() { - 125
w.cleanup(window); - 126
} - 127
} - 128
- 129
map.entry(ip) - 130
.or_default() - 131
.check_and_increment(key, max, window) - 132
} - 133
- 134
fn limit_for_path(&self, method: &str, path: &str) -> Option<(&str, u32)> { - 135
if method != "POST" { - 136
return None; - 137
} - 138
- 139
if path == "/gateway/inbound" { - 140
Some(("inbound", self.config.inbound_per_min)) - 141
} else if path == "/sessions" { - 142
Some(("sessions", self.config.sessions_per_min)) - 143
} else if path.starts_with("/sessions/") && path.ends_with("/run") { - 144
Some(("runs", self.config.runs_per_min)) - 145
} else { - 146
Some(("other_post", self.config.other_post_per_min)) - 147
} - 148
} - 149
} - 150
- 151
#[derive(Serialize)] - 152
struct RateLimitResponse { - 153
error: String, - 154
retry_after_secs: u64, - 155
} - 156
- 157
#[allow(clippy::result_large_err)] - 158
pub async fn rate_limit_layer( - 159
State(limiter): State<RateLimiter>, - 160
req: axum::extract::Request, - 161
next: axum::middleware::Next, - 162
) -> Result<Response, Response> { - 163
let ip = req - 164
.headers() - 165
.get("x-forwarded-for") - 166
.and_then(|v| v.to_str().ok()) - 167
.and_then(|v| v.split(',').next()) - 168
.and_then(|v| v.trim().parse::<IpAddr>().ok()) - 169
.or_else(|| { - 170
req.headers() - 171
.get("x-real-ip") - 172
.and_then(|v| v.to_str().ok()) - 173
.and_then(|v| v.trim().parse::<IpAddr>().ok()) - 174
}) - 175
.or_else(|| { - 176
req.extensions() - 177
.get::<axum::extract::ConnectInfo<tokio::net::TcpStream>>() - 178
.and_then(|ci| ci.0.peer_addr().ok()) - 179
.map(|addr| addr.ip()) - 180
}) - 181
.unwrap_or(IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)); - 182
- 183
let method = req.method().as_str().to_owned(); - 184
let path = req.uri().path().to_string(); - 185
- 186
if let Some((key, max)) = limiter.limit_for_path(&method, &path) - 187
&& limiter.check(ip, key, max).await.is_err() - 188
{ - 189
let ip_str = ip.to_string(); - 190
vak_core::security_events::record( - 191
&limiter.home, - 192
vak_core::security_events::EventKind::RateLimit, - 193
"rate_limit", - 194
&format!("path={path} limit={key}:{max}/min"), - 195
Some(&ip_str), - 196
); - 197
if let Some(hub) = crate::events::global() { - 198
hub.emit_security("RateLimit", &path); - 199
} - 200
return Err(( - 201
StatusCode::TOO_MANY_REQUESTS, - 202
[("retry-after", "60")], - 203
Json(RateLimitResponse { - 204
error: "rate limit exceeded".into(), - 205
retry_after_secs: 60, - 206
}), - 207
) - 208
.into_response()); - 209
} - 210
- 211
Ok(next.run(req).await) - 212
} - 213
- 214
#[cfg(test)] - 215
#[allow(clippy::unwrap_used, clippy::expect_used)] - 216
mod tests { - 217
use super::*; - 218
- 219
fn test_limiter(config: RateLimitConfig) -> RateLimiter { - 220
let dir = tempfile::tempdir().unwrap(); - 221
RateLimiter::new(config, dir.keep()) - 222
} - 223
- 224
#[tokio::test] - 225
async fn sliding_window_allows_within_limit() { - 226
let limiter = test_limiter(RateLimitConfig { - 227
inbound_per_min: 5, - 228
window_secs: 60, - 229
..Default::default() - 230
}); - 231
let ip = "127.0.0.1".parse().unwrap(); - 232
for _ in 0..5 { - 233
limiter.check(ip, "inbound", 5).await.unwrap(); - 234
} - 235
} - 236
- 237
#[tokio::test] - 238
async fn sliding_window_rejects_over_limit() { - 239
let limiter = test_limiter(RateLimitConfig { - 240
inbound_per_min: 3, - 241
window_secs: 60, - 242
..Default::default() - 243
}); - 244
let ip = "10.0.0.1".parse().unwrap(); - 245
limiter.check(ip, "inbound", 3).await.unwrap(); - 246
limiter.check(ip, "inbound", 3).await.unwrap(); - 247
limiter.check(ip, "inbound", 3).await.unwrap(); - 248
assert!(limiter.check(ip, "inbound", 3).await.is_err()); - 249
} - 250
- 251
#[tokio::test] - 252
async fn different_ips_are_independent() { - 253
let limiter = test_limiter(RateLimitConfig { - 254
inbound_per_min: 1, - 255
window_secs: 60, - 256
..Default::default() - 257
}); - 258
let ip1: IpAddr = "10.0.0.1".parse().unwrap(); - 259
let ip2: IpAddr = "10.0.0.2".parse().unwrap(); - 260
limiter.check(ip1, "inbound", 1).await.unwrap(); - 261
assert!(limiter.check(ip1, "inbound", 1).await.is_err()); - 262
limiter.check(ip2, "inbound", 1).await.unwrap(); - 263
} - 264
- 265
#[tokio::test] - 266
async fn different_keys_are_independent() { - 267
let limiter = test_limiter(RateLimitConfig { - 268
inbound_per_min: 1, - 269
sessions_per_min: 1, - 270
window_secs: 60, - 271
..Default::default() - 272
}); - 273
let ip: IpAddr = "10.0.0.1".parse().unwrap(); - 274
limiter.check(ip, "inbound", 1).await.unwrap(); - 275
assert!(limiter.check(ip, "inbound", 1).await.is_err()); - 276
limiter.check(ip, "sessions", 1).await.unwrap(); - 277
} - 278
} - 279
Indexing the workspace…
Vakyartha documentation is discovering safe artifacts, anchors, and source references.