diff --git a/docs/rate-limiting.md b/docs/rate-limiting.md index b7b6267..d2eaaf8 100644 --- a/docs/rate-limiting.md +++ b/docs/rate-limiting.md @@ -40,13 +40,25 @@ The `rate_limit` string accepts `/`: By default, rate limiting is tracked per client IP (`rate_limit_key="ip"`). You can customize the key: ### 1. Client IP (default) -Tracks client IP using the `client` tuple or `X-Forwarded-For` / `X-Real-IP` headers when behind reverse proxies: +Tracks the **actual TCP peer address** (`scope.client`). This cannot be spoofed by the client, +and is the safe default: ```python @app.get("/login", rate_limit="5/minute", rate_limit_key="ip") def login() -> dict[str, str]: return {"status": "ok"} ``` +### 1b. Forwarded IP (opt-in, only behind a trusted reverse proxy) +`rate_limit_key="trusted-forwarded-ip"` honors `X-Forwarded-For` / `X-Real-IP` (falling back to +the peer address if neither is present). **Only use this when every request actually passes +through a reverse proxy that overwrites these headers** — otherwise any client can set its own +`X-Forwarded-For` value to dodge the limit, or spoof another client's IP to get it blocked: +```python +@app.get("/login", rate_limit="5/minute", rate_limit_key="trusted-forwarded-ip") +def login() -> dict[str, str]: + return {"status": "ok"} +``` + ### 2. Header-Based (e.g. API Keys or User IDs) Rate limit by custom HTTP header (e.g. `header:X-API-Key` or `header:Authorization`): ```python diff --git a/src/rate_limit.rs b/src/rate_limit.rs index 541a312..e7b1ac0 100644 --- a/src/rate_limit.rs +++ b/src/rate_limit.rs @@ -13,7 +13,14 @@ const MAX_SHARD_ENTRIES: usize = 8192; #[derive(Clone, Debug, PartialEq, Eq)] pub enum RateLimitKeyStrategy { + /// The actual TCP peer address (``scope.client``). Cannot be spoofed by the client — + /// the safe default. Ip, + /// `X-Forwarded-For` / `X-Real-IP`, falling back to ``scope.client``. **Only safe behind + /// a trusted reverse proxy that overwrites these headers** — otherwise any client can pick + /// its own rate-limit bucket, or frame another client's IP, by setting the header itself + /// (issue #207). Opt in explicitly with ``rate_limit_key="trusted-forwarded-ip"``. + TrustedForwardedIp, Header(String), Global, } @@ -23,6 +30,10 @@ impl RateLimitKeyStrategy { let trimmed = s.trim(); if trimmed.eq_ignore_ascii_case("ip") || trimmed.eq_ignore_ascii_case("client") { Self::Ip + } else if trimmed.eq_ignore_ascii_case("trusted-forwarded-ip") + || trimmed.eq_ignore_ascii_case("forwarded-ip") + { + Self::TrustedForwardedIp } else if trimmed.eq_ignore_ascii_case("global") { Self::Global } else if let Some(hdr) = trimmed.strip_prefix("header:") { @@ -192,6 +203,37 @@ impl RateLimiter { } } +/// Read the peer address off ``scope.client``. Granian's RSGI scope exposes this as a +/// ``"host:port"`` string; the ASGI bridge / test scopes expose an ASGI-style ``(host, port)`` +/// tuple. Handles both so the "safe" strategies below get the real peer, not (e.g.) the first +/// character of a ``"host:port"`` string via a stray ``__getitem__``. +fn scope_client_host(scope: &pyo3::Bound<'_, pyo3::PyAny>) -> Option { + let client = scope.getattr("client").ok()?; + if let Ok(s) = client.extract::() { + if s.is_empty() { + return None; + } + // "host:port" (IPv4) or "[::1]:port" (IPv6) — split off the trailing ":port". + if let Some(rest) = s.strip_prefix('[') { + if let Some(end) = rest.find(']') { + return Some(format!("[{}]", &rest[..end])); + } + } + return match s.rsplit_once(':') { + Some((host, port)) if port.chars().all(|c| c.is_ascii_digit()) => { + Some(host.to_string()) + } + _ => Some(s), + }; + } + if let Ok(item) = client.get_item(0) { + if let Ok(ip) = item.extract::() { + return Some(ip); + } + } + None +} + /// Extract rate limit key from Granian / ASGI scope based on strategy. pub fn extract_rate_limit_key( scope: &pyo3::Bound<'_, pyo3::PyAny>, @@ -208,6 +250,12 @@ pub fn extract_rate_limit_key( "unknown".to_string() } RateLimitKeyStrategy::Ip => { + if let Some(ip) = scope_client_host(scope) { + return ip; + } + "127.0.0.1".to_string() + } + RateLimitKeyStrategy::TrustedForwardedIp => { if let Ok(headers) = scope.getattr("headers") { if let Some(xf) = crate::params::header_get_lax(&headers, "x-forwarded-for") { if let Some(first_ip) = xf.split(',').next() { @@ -224,12 +272,8 @@ pub fn extract_rate_limit_key( } } } - if let Ok(client) = scope.getattr("client") { - if let Ok(item) = client.get_item(0) { - if let Ok(ip) = item.extract::() { - return ip; - } - } + if let Some(ip) = scope_client_host(scope) { + return ip; } "127.0.0.1".to_string() } diff --git a/tests/test_rate_limit.py b/tests/test_rate_limit.py index f5c6a95..04e10b4 100644 --- a/tests/test_rate_limit.py +++ b/tests/test_rate_limit.py @@ -45,9 +45,11 @@ async def _run() -> None: assert int(r3.headers.get("ratelimit-reset", "0")) >= 1 assert int(r3.headers.get("retry-after", "0")) >= 1 - # Different IP (via x-forwarded-for) is not blocked - r_other_ip = await c.get("/limited", headers={"x-forwarded-for": "198.51.100.1"}) - assert r_other_ip.status_code == 200 + # Default "ip" strategy keys on the real peer address, not a client-supplied + # X-Forwarded-For (issue #207): the header must NOT let a client escape the + # bucket it's already exhausted. + r_spoofed = await c.get("/limited", headers={"x-forwarded-for": "198.51.100.1"}) + assert r_spoofed.status_code == 429 # 2. Header-keyed rate limiting rh1 = await c.get("/custom-header", headers={"x-api-key": "client-A"}) @@ -70,3 +72,28 @@ async def _run() -> None: assert rr3.status_code == 429 asyncio.run(_run()) + + +def test_trusted_forwarded_ip_opt_in() -> None: + """``rate_limit_key=\"trusted-forwarded-ip\"`` opts into honoring X-Forwarded-For / + X-Real-IP — only safe behind a reverse proxy that overwrites those headers (issue #207).""" + app = App() + + @app.get("/limited", rate_limit="1/minute", rate_limit_key="trusted-forwarded-ip") + def handle_limited() -> dict[str, str]: + return {"status": "ok"} + + async def _run() -> None: + transport = httpx.ASGITransport(app=asgi_test_app(app)) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: + r1 = await c.get("/limited", headers={"x-forwarded-for": "203.0.113.1"}) + assert r1.status_code == 200 + + r2 = await c.get("/limited", headers={"x-forwarded-for": "203.0.113.1"}) + assert r2.status_code == 429 + + # A different forwarded IP gets its own bucket. + r3 = await c.get("/limited", headers={"x-forwarded-for": "203.0.113.2"}) + assert r3.status_code == 200 + + asyncio.run(_run())