Skip to main content

ironflow_api/
rate_limit.rs

1//! Identity-aware rate limiting middleware.
2//!
3//! Requests are keyed by authenticated identity (API key ID or user ID) when
4//! available, falling back to the client IP address. Each key gets a
5//! fixed-window counter that resets every 60 seconds.
6//!
7//! Responses carry standard rate-limit headers:
8//! - `X-RateLimit-Limit` -- maximum requests allowed per window
9//! - `X-RateLimit-Remaining` -- requests left in the current window
10//! - `X-RateLimit-Reset` -- epoch timestamp when the window resets
11//!
12//! API keys with a `rate_limit_override` use that value instead of the
13//! server-wide default.
14
15use std::net::{IpAddr, Ipv4Addr};
16use std::sync::Arc;
17use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
18use std::time::{SystemTime, UNIX_EPOCH};
19
20use axum::extract::Request;
21use axum::http::{HeaderValue, StatusCode};
22use axum::middleware::Next;
23use axum::response::{IntoResponse, Response};
24use dashmap::DashMap;
25use serde_json::json;
26use uuid::Uuid;
27
28use ironflow_auth::extractor::{API_KEY_PREFIX, API_KEY_SUFFIX_LEN};
29use ironflow_auth::jwt::{AccessToken, JwtConfig};
30use ironflow_store::store::Store;
31
32/// Rate limit key: identity-based when authenticated, IP-based otherwise.
33///
34/// # Examples
35///
36/// ```
37/// use std::net::{IpAddr, Ipv4Addr};
38/// use uuid::Uuid;
39/// use ironflow_api::rate_limit::RateLimitKey;
40///
41/// let ip_key = RateLimitKey::Ip(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)));
42/// let user_key = RateLimitKey::User(Uuid::nil());
43/// let api_key = RateLimitKey::ApiKey(Uuid::nil());
44/// ```
45#[derive(Debug, Clone, PartialEq, Eq, Hash)]
46pub enum RateLimitKey {
47    /// Authenticated via API key.
48    ApiKey(Uuid),
49    /// Authenticated via JWT.
50    User(Uuid),
51    /// Unauthenticated, identified by client IP.
52    Ip(IpAddr),
53}
54
55struct WindowEntry {
56    count: AtomicU32,
57    window_start: AtomicU64,
58}
59
60/// Shared rate limiter state with fixed-window counters.
61///
62/// # Examples
63///
64/// ```
65/// use ironflow_api::rate_limit::per_minute;
66///
67/// let limiter = per_minute(60);
68/// ```
69#[derive(Clone)]
70pub struct RateLimitState {
71    counters: Arc<DashMap<RateLimitKey, WindowEntry>>,
72    burst: u32,
73}
74
75/// Context needed by the rate limit middleware.
76///
77/// Combines the limiter state with auth dependencies for identity extraction.
78///
79/// # Examples
80///
81/// ```no_run
82/// use std::sync::Arc;
83/// use ironflow_api::rate_limit::{RateLimitContext, per_minute};
84/// use ironflow_auth::jwt::JwtConfig;
85/// use ironflow_store::memory::InMemoryStore;
86/// use ironflow_store::store::Store;
87///
88/// let store: Arc<dyn Store> = Arc::new(InMemoryStore::new());
89/// let jwt_config = Arc::new(JwtConfig {
90///     secret: "secret".to_string(),
91///     access_token_ttl_secs: 900,
92///     refresh_token_ttl_secs: 604800,
93///     cookie_domain: None,
94///     cookie_secure: false,
95/// });
96/// let ctx = RateLimitContext {
97///     store,
98///     jwt_config,
99///     limiter: per_minute(60),
100/// };
101/// ```
102#[derive(Clone)]
103pub struct RateLimitContext {
104    /// Backing store for API key prefix lookups.
105    pub store: Arc<dyn Store>,
106    /// JWT config for token decoding.
107    pub jwt_config: Arc<JwtConfig>,
108    /// The rate limiter itself.
109    pub limiter: RateLimitState,
110}
111
112struct RateLimitResult {
113    limit: u32,
114    remaining: u32,
115    reset: u64,
116    allowed: bool,
117}
118
119/// Build a per-identity rate limiter allowing `requests_per_minute` requests
120/// per minute.
121///
122/// # Panics
123///
124/// Panics if `requests_per_minute` is 0.
125///
126/// # Examples
127///
128/// ```
129/// use ironflow_api::rate_limit::per_minute;
130///
131/// let auth_limiter = per_minute(10);
132/// let general_limiter = per_minute(60);
133/// ```
134pub fn per_minute(requests_per_minute: u32) -> RateLimitState {
135    assert!(requests_per_minute > 0, "burst must be > 0");
136    RateLimitState {
137        counters: Arc::new(DashMap::new()),
138        burst: requests_per_minute,
139    }
140}
141
142const WINDOW_SECS: u64 = 60;
143
144fn now_epoch() -> u64 {
145    SystemTime::now()
146        .duration_since(UNIX_EPOCH)
147        .expect("system clock before UNIX epoch")
148        .as_secs()
149}
150
151fn check_rate_limit(limiter: &RateLimitState, key: RateLimitKey, limit: u32) -> RateLimitResult {
152    let now = now_epoch();
153
154    let entry = limiter.counters.entry(key).or_insert_with(|| WindowEntry {
155        count: AtomicU32::new(0),
156        window_start: AtomicU64::new(now),
157    });
158
159    let window_start = entry.window_start.load(Ordering::Acquire);
160    if now >= window_start + WINDOW_SECS
161        && entry
162            .window_start
163            .compare_exchange(window_start, now, Ordering::AcqRel, Ordering::Relaxed)
164            .is_ok()
165    {
166        entry.count.store(0, Ordering::Release);
167    }
168
169    let count = entry.count.fetch_add(1, Ordering::AcqRel) + 1;
170    let reset = entry.window_start.load(Ordering::Acquire) + WINDOW_SECS;
171
172    if count > limit {
173        entry.count.fetch_sub(1, Ordering::AcqRel);
174        RateLimitResult {
175            limit,
176            remaining: 0,
177            reset,
178            allowed: false,
179        }
180    } else {
181        RateLimitResult {
182            limit,
183            remaining: limit.saturating_sub(count),
184            reset,
185            allowed: true,
186        }
187    }
188}
189
190/// Axum middleware that enforces identity-aware rate limiting.
191///
192/// Extracts the caller's identity from the request:
193/// - API key (`irfl_*` Bearer token) -> `RateLimitKey::ApiKey(id)`
194/// - JWT (Bearer token or cookie) -> `RateLimitKey::User(id)`
195/// - Fallback -> `RateLimitKey::Ip(addr)`
196///
197/// API keys with a `rate_limit_override` use that value instead of
198/// the default burst.
199pub async fn rate_limit(mut req: Request, next: Next) -> Response {
200    let ctx = req.extensions_mut().remove::<RateLimitContext>();
201    let Some(ctx) = ctx else {
202        return next.run(req).await;
203    };
204    let (key, override_limit) = extract_rate_limit_key_from_headers(req.headers(), &ctx).await;
205    let limit = override_limit.unwrap_or(ctx.limiter.burst);
206
207    if limit == 0 {
208        return next.run(req).await;
209    }
210
211    let result = check_rate_limit(&ctx.limiter, key, limit);
212
213    if result.allowed {
214        let mut resp = next.run(req).await;
215        insert_rate_limit_headers(resp.headers_mut(), &result);
216        resp
217    } else {
218        let retry_after = result.reset.saturating_sub(now_epoch()).max(1);
219
220        let body = json!({
221            "error": {
222                "code": "RATE_LIMIT_EXCEEDED",
223                "message": "Too many requests, please try again later",
224                "retry_after_secs": retry_after,
225            }
226        });
227
228        let mut resp = (StatusCode::TOO_MANY_REQUESTS, axum::Json(body)).into_response();
229        resp.headers_mut()
230            .insert("retry-after", HeaderValue::from(retry_after));
231        insert_rate_limit_headers(resp.headers_mut(), &result);
232        resp
233    }
234}
235
236fn insert_rate_limit_headers(headers: &mut axum::http::HeaderMap, result: &RateLimitResult) {
237    headers.insert("x-ratelimit-limit", HeaderValue::from(result.limit));
238    headers.insert("x-ratelimit-remaining", HeaderValue::from(result.remaining));
239    headers.insert("x-ratelimit-reset", HeaderValue::from(result.reset));
240}
241
242/// Extract the rate limit key and optional override from the request.
243///
244/// Attempts lightweight identity extraction without full auth verification:
245/// - API key prefix lookup (cheap index scan, no argon2)
246/// - JWT decode (local crypto, no DB)
247/// - Fallback to client IP
248async fn extract_rate_limit_key_from_headers(
249    headers: &axum::http::HeaderMap,
250    ctx: &RateLimitContext,
251) -> (RateLimitKey, Option<u32>) {
252    let bearer = headers
253        .get("authorization")
254        .and_then(|v| v.to_str().ok())
255        .and_then(|v| v.strip_prefix("Bearer "))
256        .map(|s| s.to_string());
257
258    if let Some(ref token) = bearer {
259        if token.starts_with(API_KEY_PREFIX) {
260            if let Some((key, override_limit)) = try_api_key_identity(token, ctx).await {
261                return (key, override_limit);
262            }
263        } else if let Some(key) = try_jwt_identity(token, ctx) {
264            return (key, None);
265        }
266    }
267
268    // Try JWT from cookie
269    if let Some(cookie_header) = headers.get("cookie").and_then(|v| v.to_str().ok()) {
270        for part in cookie_header.split(';') {
271            let part = part.trim();
272            if let Some(value) = part.strip_prefix("ironflow_session=")
273                && let Some(key) = try_jwt_identity(value, ctx)
274            {
275                return (key, None);
276            }
277        }
278    }
279
280    let ip = headers
281        .get("x-forwarded-for")
282        .and_then(|v| v.to_str().ok())
283        .and_then(|v| v.split(',').next())
284        .and_then(|s| s.trim().parse::<IpAddr>().ok())
285        .or_else(|| {
286            headers
287                .get("x-real-ip")
288                .and_then(|v| v.to_str().ok())
289                .and_then(|s| s.trim().parse::<IpAddr>().ok())
290        })
291        .unwrap_or(IpAddr::V4(Ipv4Addr::UNSPECIFIED));
292
293    (RateLimitKey::Ip(ip), None)
294}
295
296async fn try_api_key_identity(
297    token: &str,
298    ctx: &RateLimitContext,
299) -> Option<(RateLimitKey, Option<u32>)> {
300    let suffix_len = (token.len() - API_KEY_PREFIX.len()).min(API_KEY_SUFFIX_LEN);
301    let prefix = &token[..API_KEY_PREFIX.len() + suffix_len];
302
303    let api_key = ctx.store.find_api_key_by_prefix(prefix).await.ok()??;
304
305    let override_limit = api_key.rate_limit_override;
306    Some((RateLimitKey::ApiKey(api_key.id), override_limit))
307}
308
309fn try_jwt_identity(token: &str, ctx: &RateLimitContext) -> Option<RateLimitKey> {
310    let claims = AccessToken::decode(token, &ctx.jwt_config).ok()?;
311    Some(RateLimitKey::User(claims.user_id))
312}
313
314#[cfg(test)]
315mod tests {
316    use std::net::SocketAddr;
317
318    use axum::extract::ConnectInfo;
319
320    fn extract_client_ip(req: &Request<Body>) -> IpAddr {
321        if let Some(forwarded) = req
322            .headers()
323            .get("x-forwarded-for")
324            .and_then(|v| v.to_str().ok())
325            && let Some(first) = forwarded.split(',').next()
326            && let Ok(ip) = first.trim().parse::<IpAddr>()
327        {
328            return ip;
329        }
330
331        if let Some(real_ip) = req.headers().get("x-real-ip").and_then(|v| v.to_str().ok())
332            && let Ok(ip) = real_ip.trim().parse::<IpAddr>()
333        {
334            return ip;
335        }
336
337        req.extensions()
338            .get::<ConnectInfo<SocketAddr>>()
339            .map(|ci| ci.0.ip())
340            .unwrap_or(IpAddr::V4(Ipv4Addr::UNSPECIFIED))
341    }
342
343    use axum::Extension;
344    use axum::Router;
345    use axum::body::Body;
346    use axum::http::{Request, StatusCode};
347    use axum::middleware as axum_mw;
348    use axum::routing::get;
349    use http_body_util::BodyExt;
350    use ironflow_auth::password;
351    use ironflow_store::entities::{ApiKeyScope, NewApiKey, NewUser};
352    use ironflow_store::memory::InMemoryStore;
353    use serde_json::Value as JsonValue;
354    use tower::ServiceExt;
355
356    use super::*;
357
358    async fn ok_handler() -> &'static str {
359        "ok"
360    }
361
362    fn test_jwt_config() -> Arc<JwtConfig> {
363        Arc::new(JwtConfig {
364            secret: "test-secret-for-rate-limit".to_string(),
365            access_token_ttl_secs: 900,
366            refresh_token_ttl_secs: 604800,
367            cookie_domain: None,
368            cookie_secure: false,
369        })
370    }
371
372    fn test_ctx(burst: u32) -> RateLimitContext {
373        let store: Arc<dyn Store> = Arc::new(InMemoryStore::new());
374        RateLimitContext {
375            store,
376            jwt_config: test_jwt_config(),
377            limiter: per_minute(burst),
378        }
379    }
380
381    fn test_app(ctx: RateLimitContext) -> Router {
382        Router::new()
383            .route("/test", get(ok_handler))
384            .layer(axum_mw::from_fn(rate_limit))
385            .layer(Extension(ctx))
386    }
387
388    fn ip_request(ip: &str) -> Request<Body> {
389        Request::builder()
390            .uri("/test")
391            .header("x-forwarded-for", ip)
392            .body(Body::empty())
393            .unwrap()
394    }
395
396    async fn setup_api_key_in_store(
397        store: &Arc<dyn Store>,
398        rate_limit_override: Option<u32>,
399    ) -> (Uuid, String) {
400        let user = store
401            .create_user(NewUser {
402                email: "rl-test@test.com".to_string(),
403                username: "rl-test".to_string(),
404                password_hash: password::hash("pass").unwrap(),
405                is_admin: Some(false),
406            })
407            .await
408            .unwrap();
409
410        let raw_key = "irfl_abcdef12rest-of-secret-key";
411        let key_hash = password::hash(raw_key).unwrap();
412        let prefix =
413            &raw_key[..API_KEY_PREFIX.len() + ironflow_auth::extractor::API_KEY_SUFFIX_LEN];
414
415        let api_key = store
416            .create_api_key(NewApiKey {
417                user_id: user.id,
418                name: "rl-test-key".to_string(),
419                key_hash,
420                key_prefix: prefix.to_string(),
421                scopes: vec![ApiKeyScope::RunsRead],
422                expires_at: None,
423                rate_limit_override,
424            })
425            .await
426            .unwrap();
427
428        (api_key.id, raw_key.to_string())
429    }
430
431    #[tokio::test]
432    async fn unauthenticated_uses_ip_bucket() {
433        let ctx = test_ctx(2);
434        let app = test_app(ctx.clone());
435
436        let resp = app.oneshot(ip_request("1.2.3.4")).await.unwrap();
437        assert_eq!(resp.status(), StatusCode::OK);
438
439        let app = test_app(ctx.clone());
440        let resp = app.oneshot(ip_request("1.2.3.4")).await.unwrap();
441        assert_eq!(resp.status(), StatusCode::OK);
442
443        let app = test_app(ctx);
444        let resp = app.oneshot(ip_request("1.2.3.4")).await.unwrap();
445        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
446    }
447
448    #[tokio::test]
449    async fn api_key_auth_uses_separate_bucket() {
450        let store: Arc<dyn Store> = Arc::new(InMemoryStore::new());
451        let (_api_key_id, raw_key) = setup_api_key_in_store(&store, None).await;
452
453        let ctx = RateLimitContext {
454            store,
455            jwt_config: test_jwt_config(),
456            limiter: per_minute(2),
457        };
458
459        // Use up both IP requests from 1.2.3.4
460        let app = test_app(ctx.clone());
461        let _ = app.oneshot(ip_request("1.2.3.4")).await;
462        let app = test_app(ctx.clone());
463        let _ = app.oneshot(ip_request("1.2.3.4")).await;
464
465        // IP 1.2.3.4 is now exhausted
466        let app = test_app(ctx.clone());
467        let resp = app.oneshot(ip_request("1.2.3.4")).await.unwrap();
468        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
469
470        // But API key from same IP still works (separate bucket)
471        let app = test_app(ctx);
472        let req = Request::builder()
473            .uri("/test")
474            .header("x-forwarded-for", "1.2.3.4")
475            .header("authorization", format!("Bearer {raw_key}"))
476            .body(Body::empty())
477            .unwrap();
478        let resp = app.oneshot(req).await.unwrap();
479        assert_eq!(resp.status(), StatusCode::OK);
480    }
481
482    #[tokio::test]
483    async fn jwt_auth_uses_user_bucket() {
484        let ctx = test_ctx(2);
485        let user_id = Uuid::now_v7();
486        let token = AccessToken::for_user(user_id, "alice", false, &ctx.jwt_config).unwrap();
487
488        // Use up both IP requests
489        let app = test_app(ctx.clone());
490        let _ = app.oneshot(ip_request("10.0.0.1")).await;
491        let app = test_app(ctx.clone());
492        let _ = app.oneshot(ip_request("10.0.0.1")).await;
493
494        // IP exhausted
495        let app = test_app(ctx.clone());
496        let resp = app.oneshot(ip_request("10.0.0.1")).await.unwrap();
497        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
498
499        // JWT from same IP still works (separate User bucket)
500        let app = test_app(ctx);
501        let req = Request::builder()
502            .uri("/test")
503            .header("x-forwarded-for", "10.0.0.1")
504            .header("authorization", format!("Bearer {}", token.0))
505            .body(Body::empty())
506            .unwrap();
507        let resp = app.oneshot(req).await.unwrap();
508        assert_eq!(resp.status(), StatusCode::OK);
509    }
510
511    #[tokio::test]
512    async fn remaining_header_decrements() {
513        let ctx = test_ctx(3);
514
515        let app = test_app(ctx.clone());
516        let resp = app.oneshot(ip_request("2.2.2.2")).await.unwrap();
517        assert_eq!(resp.status(), StatusCode::OK);
518        let remaining: u32 = resp
519            .headers()
520            .get("x-ratelimit-remaining")
521            .unwrap()
522            .to_str()
523            .unwrap()
524            .parse()
525            .unwrap();
526        assert_eq!(remaining, 2);
527
528        let app = test_app(ctx.clone());
529        let resp = app.oneshot(ip_request("2.2.2.2")).await.unwrap();
530        let remaining: u32 = resp
531            .headers()
532            .get("x-ratelimit-remaining")
533            .unwrap()
534            .to_str()
535            .unwrap()
536            .parse()
537            .unwrap();
538        assert_eq!(remaining, 1);
539
540        let app = test_app(ctx);
541        let resp = app.oneshot(ip_request("2.2.2.2")).await.unwrap();
542        let remaining: u32 = resp
543            .headers()
544            .get("x-ratelimit-remaining")
545            .unwrap()
546            .to_str()
547            .unwrap()
548            .parse()
549            .unwrap();
550        assert_eq!(remaining, 0);
551    }
552
553    #[tokio::test]
554    async fn reset_header_present() {
555        let ctx = test_ctx(5);
556        let app = test_app(ctx);
557
558        let resp = app.oneshot(ip_request("3.3.3.3")).await.unwrap();
559        assert_eq!(resp.status(), StatusCode::OK);
560
561        let reset: u64 = resp
562            .headers()
563            .get("x-ratelimit-reset")
564            .unwrap()
565            .to_str()
566            .unwrap()
567            .parse()
568            .unwrap();
569
570        let now = now_epoch();
571        assert!(
572            reset > now,
573            "reset {reset} should be in the future (now {now})"
574        );
575        assert!(
576            reset <= now + WINDOW_SECS,
577            "reset {reset} should be within one window of now {now}"
578        );
579    }
580
581    #[tokio::test]
582    async fn api_key_override_uses_custom_limit() {
583        let store: Arc<dyn Store> = Arc::new(InMemoryStore::new());
584        let (_api_key_id, raw_key) = setup_api_key_in_store(&store, Some(1)).await;
585
586        let ctx = RateLimitContext {
587            store,
588            jwt_config: test_jwt_config(),
589            limiter: per_minute(100),
590        };
591
592        let app = test_app(ctx.clone());
593        let req = Request::builder()
594            .uri("/test")
595            .header("authorization", format!("Bearer {raw_key}"))
596            .body(Body::empty())
597            .unwrap();
598        let resp = app.oneshot(req).await.unwrap();
599        assert_eq!(resp.status(), StatusCode::OK);
600        assert_eq!(
601            resp.headers()
602                .get("x-ratelimit-limit")
603                .unwrap()
604                .to_str()
605                .unwrap(),
606            "1"
607        );
608
609        // Second request with override=1 should be rejected
610        let app = test_app(ctx);
611        let req = Request::builder()
612            .uri("/test")
613            .header("authorization", format!("Bearer {raw_key}"))
614            .body(Body::empty())
615            .unwrap();
616        let resp = app.oneshot(req).await.unwrap();
617        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
618    }
619
620    #[tokio::test]
621    async fn rate_limited_response_includes_all_headers() {
622        let ctx = test_ctx(1);
623
624        let app = test_app(ctx.clone());
625        let _ = app.oneshot(ip_request("5.5.5.5")).await;
626
627        let app = test_app(ctx);
628        let resp = app.oneshot(ip_request("5.5.5.5")).await.unwrap();
629        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
630
631        assert!(resp.headers().contains_key("retry-after"));
632        assert!(resp.headers().contains_key("x-ratelimit-limit"));
633        assert!(resp.headers().contains_key("x-ratelimit-remaining"));
634        assert!(resp.headers().contains_key("x-ratelimit-reset"));
635
636        let remaining: u32 = resp
637            .headers()
638            .get("x-ratelimit-remaining")
639            .unwrap()
640            .to_str()
641            .unwrap()
642            .parse()
643            .unwrap();
644        assert_eq!(remaining, 0);
645
646        let body = resp.into_body().collect().await.unwrap().to_bytes();
647        let json_val: JsonValue = serde_json::from_slice(&body).unwrap();
648        assert_eq!(json_val["error"]["code"], "RATE_LIMIT_EXCEEDED");
649    }
650
651    #[tokio::test]
652    async fn different_ips_have_separate_limits() {
653        let ctx = test_ctx(1);
654
655        let app = test_app(ctx.clone());
656        let resp = app.oneshot(ip_request("10.0.0.1")).await.unwrap();
657        assert_eq!(resp.status(), StatusCode::OK);
658
659        let app = test_app(ctx);
660        let resp = app.oneshot(ip_request("10.0.0.2")).await.unwrap();
661        assert_eq!(resp.status(), StatusCode::OK);
662    }
663
664    #[tokio::test]
665    async fn allows_requests_within_limit() {
666        let ctx = test_ctx(5);
667        let app = test_app(ctx);
668
669        let resp = app.oneshot(ip_request("1.2.3.4")).await.unwrap();
670        assert_eq!(resp.status(), StatusCode::OK);
671        assert!(resp.headers().contains_key("x-ratelimit-limit"));
672    }
673
674    #[tokio::test]
675    async fn rejects_when_limit_exceeded() {
676        let ctx = test_ctx(2);
677
678        let app = test_app(ctx.clone());
679        let _ = app.oneshot(ip_request("1.2.3.4")).await;
680        let app = test_app(ctx.clone());
681        let _ = app.oneshot(ip_request("1.2.3.4")).await;
682
683        let app = test_app(ctx);
684        let resp = app.oneshot(ip_request("1.2.3.4")).await.unwrap();
685        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
686
687        let body = resp.into_body().collect().await.unwrap().to_bytes();
688        let json_val: JsonValue = serde_json::from_slice(&body).unwrap();
689        assert_eq!(json_val["error"]["code"], "RATE_LIMIT_EXCEEDED");
690    }
691
692    #[tokio::test]
693    async fn includes_retry_after_header() {
694        let ctx = test_ctx(1);
695
696        let app = test_app(ctx.clone());
697        let _ = app.oneshot(ip_request("1.2.3.4")).await;
698
699        let app = test_app(ctx);
700        let resp = app.oneshot(ip_request("1.2.3.4")).await.unwrap();
701        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
702        assert!(resp.headers().contains_key("retry-after"));
703    }
704
705    #[tokio::test]
706    async fn extracts_ip_from_x_real_ip() {
707        let ctx = test_ctx(1);
708
709        let app = test_app(ctx.clone());
710        let req = Request::builder()
711            .uri("/test")
712            .header("x-real-ip", "192.168.1.1")
713            .body(Body::empty())
714            .unwrap();
715        let resp = app.oneshot(req).await.unwrap();
716        assert_eq!(resp.status(), StatusCode::OK);
717
718        let app = test_app(ctx);
719        let req = Request::builder()
720            .uri("/test")
721            .header("x-real-ip", "192.168.1.1")
722            .body(Body::empty())
723            .unwrap();
724        let resp = app.oneshot(req).await.unwrap();
725        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
726    }
727
728    #[test]
729    fn extract_ip_x_forwarded_for_first_ip() {
730        let req = Request::builder()
731            .uri("/test")
732            .header("x-forwarded-for", "1.2.3.4, 5.6.7.8")
733            .body(Body::empty())
734            .unwrap();
735        let ip = extract_client_ip(&req);
736        assert_eq!(ip, "1.2.3.4".parse::<IpAddr>().unwrap());
737    }
738
739    #[test]
740    fn extract_ip_fallback_to_unspecified() {
741        let req = Request::builder().uri("/test").body(Body::empty()).unwrap();
742        let ip = extract_client_ip(&req);
743        assert_eq!(ip, IpAddr::V4(Ipv4Addr::UNSPECIFIED));
744    }
745
746    #[tokio::test]
747    async fn general_limiter_allows_more_requests() {
748        let ctx = test_ctx(60);
749
750        for _ in 0..10 {
751            let app = test_app(ctx.clone());
752            let resp = app.oneshot(ip_request("1.2.3.4")).await.unwrap();
753            assert_eq!(resp.status(), StatusCode::OK);
754        }
755    }
756
757    #[tokio::test]
758    async fn jwt_cookie_uses_user_bucket() {
759        let ctx = test_ctx(1);
760        let user_id = Uuid::now_v7();
761        let token = AccessToken::for_user(user_id, "cookie-user", false, &ctx.jwt_config).unwrap();
762
763        // Exhaust IP bucket
764        let app = test_app(ctx.clone());
765        let _ = app.oneshot(ip_request("20.0.0.1")).await;
766
767        // IP exhausted
768        let app = test_app(ctx.clone());
769        let resp = app.oneshot(ip_request("20.0.0.1")).await.unwrap();
770        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
771
772        // JWT in cookie from same IP still works (separate User bucket)
773        let app = test_app(ctx);
774        let req = Request::builder()
775            .uri("/test")
776            .header("x-forwarded-for", "20.0.0.1")
777            .header("cookie", format!("ironflow_session={}", token.0))
778            .body(Body::empty())
779            .unwrap();
780        let resp = app.oneshot(req).await.unwrap();
781        assert_eq!(resp.status(), StatusCode::OK);
782    }
783
784    #[tokio::test]
785    async fn invalid_bearer_falls_back_to_ip() {
786        let ctx = test_ctx(1);
787
788        // First request with garbage Bearer uses IP bucket
789        let app = test_app(ctx.clone());
790        let req = Request::builder()
791            .uri("/test")
792            .header("x-forwarded-for", "30.0.0.1")
793            .header("authorization", "Bearer not_a_valid_token")
794            .body(Body::empty())
795            .unwrap();
796        let resp = app.oneshot(req).await.unwrap();
797        assert_eq!(resp.status(), StatusCode::OK);
798
799        // Second request from same IP is rate limited (same bucket)
800        let app = test_app(ctx);
801        let req = Request::builder()
802            .uri("/test")
803            .header("x-forwarded-for", "30.0.0.1")
804            .header("authorization", "Bearer not_a_valid_token")
805            .body(Body::empty())
806            .unwrap();
807        let resp = app.oneshot(req).await.unwrap();
808        assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS);
809    }
810
811    #[tokio::test]
812    async fn api_key_override_zero_disables_rate_limiting() {
813        let store: Arc<dyn Store> = Arc::new(InMemoryStore::new());
814        let (_api_key_id, raw_key) = setup_api_key_in_store(&store, Some(0)).await;
815
816        let ctx = RateLimitContext {
817            store,
818            jwt_config: test_jwt_config(),
819            limiter: per_minute(1),
820        };
821
822        // override=0 means unlimited: many requests should all pass
823        for _ in 0..10 {
824            let app = test_app(ctx.clone());
825            let req = Request::builder()
826                .uri("/test")
827                .header("authorization", format!("Bearer {raw_key}"))
828                .body(Body::empty())
829                .unwrap();
830            let resp = app.oneshot(req).await.unwrap();
831            assert_eq!(resp.status(), StatusCode::OK);
832        }
833    }
834}