Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 13 additions & 1 deletion docs/rate-limiting.md
Original file line number Diff line number Diff line change
Expand Up @@ -40,13 +40,25 @@ The `rate_limit` string accepts `<count>/<window>`:
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
Expand Down
56 changes: 50 additions & 6 deletions src/rate_limit.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
}
Expand All @@ -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:") {
Expand Down Expand Up @@ -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<String> {
let client = scope.getattr("client").ok()?;
if let Ok(s) = client.extract::<String>() {
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::<String>() {
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>,
Expand All @@ -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() {
Expand All @@ -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::<String>() {
return ip;
}
}
if let Some(ip) = scope_client_host(scope) {
return ip;
}
"127.0.0.1".to_string()
}
Expand Down
33 changes: 30 additions & 3 deletions tests/test_rate_limit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"})
Expand All @@ -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())
Loading