Harden auth against a forwarded-header rate-limit bypass, and guard BYOK SSRF

The per-IP rate limits could be bypassed entirely: uvicorn ran with
--forwarded-allow-ips "*", which trusts the leftmost X-Forwarded-For value
(client-controlled), and Render forwards the inbound header rather than
stripping it. Rotating the header handed out a fresh rate-limit bucket per
request, so the login/register limit (10/5min) and guest-minting limit
(30/5min) were no throttle at all — unbounded password guessing and guest-row
creation. Confirmed live: fixed IP -> 429 after 10; rotating spoofed header ->
no 429 across 14 attempts.

Two-layer fix:
- limits._client_ip now derives the client IP from the hop the trusted edge
  appends (rightmost of X-Forwarded-For), which a client can't spoof past;
  tunable via AIDND_TRUSTED_PROXY_HOPS. Dropped --forwarded-allow-ips "*".
- New per-account login throttle (email-keyed, 8 fails / 15 min, cleared on
  success): stops distributed guessing against one account that a per-IP limit
  can't, since it can't be diluted across many source addresses.

Also close an SSRF on the BYOK endpoint_url (hosted mode only): the connection
test and turn/chat streams now refuse a URL that resolves to a non-public
address (private/loopback/link-local metadata/reserved), checked at request
time so it resists a DNS record flipping to a private IP. No-op locally, where
reaching localhost Ollama is intended.

Tests: test_ratelimit_hardening.py (8), test_netguard.py (13). 172 pass.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_015CYEJKobJ2Re4Dv7qUoSA7
This commit is contained in:
parththakkar106
2026-08-15 13:47:14 +05:30
co-authored by Claude Opus 4.8
parent 8857757642
commit c500203270
8 changed files with 331 additions and 5 deletions
+8 -3
View File
@@ -33,9 +33,14 @@ VOLUME /data
EXPOSE 8000 EXPOSE 8000
WORKDIR /app/backend WORKDIR /app/backend
# --proxy-headers: behind a reverse proxy (any hosted deploy), trust # --proxy-headers lets uvicorn fix up the request scheme (https) behind the
# X-Forwarded-For so per-IP rate limits key on the client, not the proxy. # platform's edge. We deliberately do NOT pass --forwarded-allow-ips "*": that
# made uvicorn trust the LEFTMOST X-Forwarded-For value, which the client fully
# controls, so anyone could rotate the header to dodge the per-IP rate limits.
# The client IP used for rate limiting is derived in limits._client_ip from the
# hop the edge appends (rightmost), which a client cannot spoof past; tune with
# AIDND_TRUSTED_PROXY_HOPS if the platform adds more proxy hops.
# Single worker on purpose: the turn lock, rate limiter, and debug log are # Single worker on purpose: the turn lock, rate limiter, and debug log are
# in-process state. # in-process state.
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", \ CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", \
"--proxy-headers", "--forwarded-allow-ips", "*"] "--proxy-headers"]
+78
View File
@@ -7,6 +7,7 @@ hostile visitor can't burn the demo key, peg the CPU, or bloat the database.
""" """
import json import json
import os
import threading import threading
import time import time
from collections import defaultdict, deque from collections import defaultdict, deque
@@ -37,7 +38,27 @@ _windows: dict[tuple[str, str], deque] = defaultdict(deque)
_windows_guard = threading.Lock() _windows_guard = threading.Lock()
# How many proxy hops sit between the app and the real client. On Render (and
# most PaaS) that's one: the platform's edge appends the connecting IP to the
# RIGHT of X-Forwarded-For. A client can prepend anything it likes to the left,
# but it cannot push a value past the edge's own append — so the trustworthy
# client IP is the (hops)-th entry from the right, NOT uvicorn's leftmost pick.
# Trusting the leftmost let anyone rotate X-Forwarded-For to mint a fresh
# rate-limit bucket per request and bypass the auth/guest limits entirely.
# Override with AIDND_TRUSTED_PROXY_HOPS if the deployment adds more hops.
TRUSTED_PROXY_HOPS = max(1, int(os.environ.get("AIDND_TRUSTED_PROXY_HOPS", "1") or 1))
def _client_ip(request: Request) -> str: def _client_ip(request: Request) -> str:
"""The real client IP for rate-limit keying, resistant to a spoofed
X-Forwarded-For. Takes the hop the trusted edge appended (rightmost minus
any extra trusted hops); falls back to the socket peer when no forwarded
header is present (local/dev, or a direct connection)."""
forwarded = request.headers.get("x-forwarded-for")
if forwarded:
parts = [p.strip() for p in forwarded.split(",") if p.strip()]
if parts:
return parts[-min(TRUSTED_PROXY_HOPS, len(parts))]
return request.client.host if request.client else "unknown" return request.client.host if request.client else "unknown"
@@ -62,6 +83,63 @@ def rate_limit(scope: str, request: Request, user: models.User | None = None) ->
_prune(now) _prune(now)
# ---------- Per-account login throttle ----------
# Defense in depth beside the per-IP `auth` limit: that one can be diluted by a
# botnet (many real source IPs, one bucket each), so it can't by itself stop a
# distributed guessing run against a single account. This cap keys on the target
# email instead of the caller, so guessing ONE account's password stays
# expensive regardless of how many addresses it comes from. Failures only — a
# correct password clears the record — and it's a short sliding window, not a
# hard lock, so a user mistyping a few times recovers on their own in minutes.
# Tradeoff: an attacker can keep a known account throttled (a nuisance), which
# is strictly preferable to letting it be brute-forced.
LOGIN_FAIL_LIMIT = 8 # failed attempts per account...
LOGIN_FAIL_WINDOW = 900 # ...within this many seconds (15 min)
_login_fails: dict[str, deque] = defaultdict(deque)
_login_guard = threading.Lock()
def check_login_allowed(email: str) -> None:
"""429 when an account has too many recent failed logins. Call before
verifying the password so guesses don't even reach the hash."""
if not auth.MULTI_USER:
return
now = time.time()
with _login_guard:
window = _login_fails[email]
while window and window[0] < now - LOGIN_FAIL_WINDOW:
window.popleft()
if len(window) >= LOGIN_FAIL_LIMIT:
raise HTTPException(
429,
"Too many failed sign-in attempts for this account — "
"wait a few minutes and try again.",
)
def note_login_failure(email: str) -> None:
"""Record one failed attempt against `email`."""
if not auth.MULTI_USER:
return
now = time.time()
with _login_guard:
_login_fails[email].append(now)
if len(_login_fails) > 10_000: # bound the map on a flood of unique emails
stale = [
key for key, window in _login_fails.items()
if not window or window[-1] < now - LOGIN_FAIL_WINDOW
]
for key in stale:
del _login_fails[key]
def note_login_success(email: str) -> None:
"""A correct password wipes the account's failure streak."""
with _login_guard:
_login_fails.pop(email, None)
def _prune(now: float) -> None: def _prune(now: float) -> None:
"""Drop callers whose whole window has expired (call with guard held) so """Drop callers whose whole window has expired (call with guard held) so
the per-IP dict can't grow without bound.""" the per-IP dict can't grow without bound."""
+51
View File
@@ -0,0 +1,51 @@
"""SSRF guard for the one place the server makes an outbound request to a
user-supplied address: the BYOK `endpoint_url` (connection test + turns/chat).
Without this, a hosted user could point endpoint_url at an internal service or
the cloud metadata endpoint (169.254.169.254) and have the server fetch it —
the connection test even echoes part of the response back. We refuse any URL
that resolves to a non-public address.
No-op in local mode: a local install talking to http://localhost:11434 (Ollama)
is the normal, intended case — the guard only applies to the hosted, multi-user
deployment where the endpoint comes from an untrusted visitor.
"""
import ipaddress
import socket
from urllib.parse import urlparse
from . import auth
def endpoint_block_reason(url: str) -> str | None:
"""A human-readable reason this URL must NOT be fetched server-side, or None
if it's allowed. Resolves the host and rejects it if any resulting address
is non-public (private, loopback, link-local/metadata, reserved, …).
Checking at request time (not just on save) is deliberate: it resists a DNS
record that flips to a private IP after the value was stored.
"""
if not auth.MULTI_USER:
return None
parsed = urlparse(url)
if parsed.scheme not in ("http", "https"):
return "the endpoint URL must start with http:// or https://"
host = parsed.hostname
if not host:
return "the endpoint URL has no host"
port = parsed.port or (443 if parsed.scheme == "https" else 80)
try:
infos = socket.getaddrinfo(host, port, type=socket.SOCK_STREAM)
except socket.gaierror:
return "the endpoint host could not be resolved"
for info in infos:
try:
ip = ipaddress.ip_address(info[4][0])
except ValueError:
return "the endpoint host resolved to an unrecognized address"
# is_global is the strict allowlist: private/loopback/link-local/CGNAT
# all report False, so this one check covers the metadata IP too.
if not ip.is_global or ip.is_multicast or ip.is_reserved:
return "the endpoint URL resolves to a non-public address"
return None
+6 -1
View File
@@ -3,7 +3,7 @@ from typing import AsyncIterator
import httpx import httpx
from .. import debuglog from .. import debuglog, netguard
from .base import PromptParts, Provider, ProviderError from .base import PromptParts, Provider, ProviderError
# Framing appended after the story text in chat mode, so chat-tuned models keep # Framing appended after the story text in chat mode, so chat-tuned models keep
@@ -168,6 +168,11 @@ class OpenAICompatibleProvider(Provider):
async def _stream(self, url: str, body: dict) -> AsyncIterator[tuple[str, str]]: async def _stream(self, url: str, body: dict) -> AsyncIterator[tuple[str, str]]:
"""Shared SSE plumbing for generate()/chat(): POST a streaming request """Shared SSE plumbing for generate()/chat(): POST a streaming request
and yield ("text" | "reasoning", chunk) pairs, logging the exchange.""" and yield ("text" | "reasoning", chunk) pairs, logging the exchange."""
# SSRF guard (hosted mode): a user-supplied endpoint_url must not point
# at an internal/metadata address. No-op for local installs.
reason = netguard.endpoint_block_reason(url)
if reason:
raise ProviderError(f"This endpoint can't be used — {reason}.")
log = debuglog.start_entry(url, self.model, body) log = debuglog.start_entry(url, self.model, body)
received: list[str] = [] received: list[str] = []
try: try:
+5
View File
@@ -105,13 +105,18 @@ def login(
raise HTTPException(400, "Accounts are disabled in local mode.") raise HTTPException(400, "Accounts are disabled in local mode.")
limits.rate_limit("auth", request) limits.rate_limit("auth", request)
email = payload.email.strip().lower() email = payload.email.strip().lower()
# Per-account throttle: stops distributed guessing against one email even
# when the per-IP limit above is diluted across many source addresses.
limits.check_login_allowed(email)
user = db.query(models.User).filter(models.User.email == email).first() user = db.query(models.User).filter(models.User.email == email).first()
if ( if (
user is None user is None
or not user.password_hash or not user.password_hash
or not security.verify_password(payload.password, user.password_hash) or not security.verify_password(payload.password, user.password_hash)
): ):
limits.note_login_failure(email)
raise HTTPException(401, "Incorrect email or password.") raise HTTPException(401, "Incorrect email or password.")
limits.note_login_success(email)
_set_session_cookie(response, user.id) _set_session_cookie(response, user.id)
return me_payload(user, db) return me_payload(user, db)
+6 -1
View File
@@ -1,8 +1,9 @@
import httpx import httpx
from fastapi import APIRouter, Depends, Request from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from starlette.concurrency import run_in_threadpool
from .. import auth, limits, models, schemas, security from .. import auth, limits, models, netguard, schemas, security
from ..database import get_db from ..database import get_db
router = APIRouter(prefix="/api/settings", tags=["settings"]) router = APIRouter(prefix="/api/settings", tags=["settings"])
@@ -65,6 +66,10 @@ def update_settings(
async def list_endpoint_models(cfg: auth.ProviderConfig) -> dict: async def list_endpoint_models(cfg: auth.ProviderConfig) -> dict:
"""GET the endpoint's /models listing. Doubles as a connectivity check, so """GET the endpoint's /models listing. Doubles as a connectivity check, so
failures come back as {"ok": False, "detail": ...} rather than raising.""" failures come back as {"ok": False, "detail": ...} rather than raising."""
# SSRF guard: never probe a non-public address the user pointed us at.
reason = await run_in_threadpool(netguard.endpoint_block_reason, cfg.endpoint_url)
if reason:
return {"ok": False, "detail": f"Can't reach that endpoint — {reason}."}
url = cfg.endpoint_url.rstrip("/") + "/models" url = cfg.endpoint_url.rstrip("/") + "/models"
headers = {} headers = {}
if cfg.api_key: if cfg.api_key:
+69
View File
@@ -0,0 +1,69 @@
"""Tests for the SSRF guard on the user-supplied BYOK endpoint_url.
python -m pytest tests/test_netguard.py -v
"""
import os
import tempfile
_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
_tmp.close()
os.environ["AIDND_DB_PATH"] = _tmp.name
os.environ.pop("AIDND_DATABASE_URL", None)
os.environ.pop("DATABASE_URL", None)
import pytest
from app import auth, netguard
@pytest.fixture
def hosted(monkeypatch):
monkeypatch.setattr(auth, "MULTI_USER", True)
def _resolves_to(monkeypatch, ip: str):
"""Pin getaddrinfo so we test the address decision, not real DNS."""
monkeypatch.setattr(
netguard.socket, "getaddrinfo",
lambda *a, **k: [(2, 1, 6, "", (ip, 443))],
)
@pytest.mark.parametrize("ip", [
"127.0.0.1", # loopback
"169.254.169.254", # cloud metadata (link-local)
"10.0.0.5", # RFC1918
"192.168.1.1", # RFC1918
"172.16.0.9", # RFC1918
"0.0.0.0", # unspecified
"100.64.0.1", # carrier-grade NAT
"::1", # IPv6 loopback
"fd00::1", # IPv6 unique-local
])
def test_blocks_non_public_addresses(hosted, monkeypatch, ip):
_resolves_to(monkeypatch, ip)
assert netguard.endpoint_block_reason("https://evil.example.com/v1") is not None
def test_allows_public_address(hosted, monkeypatch):
_resolves_to(monkeypatch, "104.18.0.1") # a public IP
assert netguard.endpoint_block_reason("https://openrouter.ai/api/v1") is None
def test_rejects_non_http_scheme(hosted):
assert netguard.endpoint_block_reason("file:///etc/passwd") is not None
assert netguard.endpoint_block_reason("gopher://x/") is not None
def test_unresolvable_host_is_blocked(hosted, monkeypatch):
def boom(*a, **k):
raise netguard.socket.gaierror("no such host")
monkeypatch.setattr(netguard.socket, "getaddrinfo", boom)
assert netguard.endpoint_block_reason("https://nope.invalid/v1") is not None
def test_noop_in_local_mode(monkeypatch):
monkeypatch.setattr(auth, "MULTI_USER", False)
# Local installs legitimately reach localhost (Ollama) — never blocked.
assert netguard.endpoint_block_reason("http://localhost:11434/v1") is None
assert netguard.endpoint_block_reason("http://127.0.0.1:11434/v1") is None
+108
View File
@@ -0,0 +1,108 @@
"""Regression tests for the X-Forwarded-For rate-limit bypass and the
per-account login throttle added to close it.
Background: uvicorn's --forwarded-allow-ips "*" trusted the LEFTMOST
X-Forwarded-For entry, which the client controls, so rotating the header
handed out a fresh rate-limit bucket per request. _client_ip now reads the
hop the trusted edge appends (rightmost), and login has an email-keyed throttle
that no IP trick can dilute.
python -m pytest tests/test_ratelimit_hardening.py -v
"""
import os
import tempfile
_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
_tmp.close()
os.environ["AIDND_DB_PATH"] = _tmp.name
os.environ.pop("AIDND_DATABASE_URL", None)
os.environ.pop("DATABASE_URL", None)
import pytest
from app import auth, limits
class _Req:
"""Minimal stand-in for starlette's Request: a header lookup and a peer."""
def __init__(self, xff: str | None, peer: str | None = "10.0.0.1"):
self.headers = {} if xff is None else {"x-forwarded-for": xff}
self.client = None if peer is None else type("C", (), {"host": peer})()
# ---------- _client_ip: the spoof-resistant hop ----------
def test_client_ip_takes_appended_rightmost_hop(monkeypatch):
monkeypatch.setattr(limits, "TRUSTED_PROXY_HOPS", 1)
# Attacker prepends a fake IP; the edge appends the real one on the right.
req = _Req("203.0.113.9, 198.51.100.77")
assert limits._client_ip(req) == "198.51.100.77"
def test_client_ip_ignores_spoofed_leftmost(monkeypatch):
monkeypatch.setattr(limits, "TRUSTED_PROXY_HOPS", 1)
# Whatever the client stuffs to the left, the keyed IP stays the real hop —
# so rotating it no longer mints a new bucket.
a = limits._client_ip(_Req("1.1.1.1, 198.51.100.77"))
b = limits._client_ip(_Req("2.2.2.2, 198.51.100.77"))
c = limits._client_ip(_Req("evil, junk, 198.51.100.77"))
assert a == b == c == "198.51.100.77"
def test_client_ip_honours_extra_trusted_hops(monkeypatch):
monkeypatch.setattr(limits, "TRUSTED_PROXY_HOPS", 2)
# Two trusted hops: real client is second from the right.
req = _Req("9.9.9.9, 203.0.113.5, 198.51.100.77")
assert limits._client_ip(req) == "203.0.113.5"
def test_client_ip_falls_back_to_socket_peer():
assert limits._client_ip(_Req(None, peer="172.16.0.4")) == "172.16.0.4"
assert limits._client_ip(_Req(None, peer=None)) == "unknown"
# ---------- per-account login throttle ----------
@pytest.fixture(autouse=True)
def _multi_user(monkeypatch):
monkeypatch.setattr(auth, "MULTI_USER", True)
# Isolate the module-level failure map for each test.
from collections import defaultdict, deque
monkeypatch.setattr(limits, "_login_fails", defaultdict(deque))
def test_login_throttle_blocks_after_limit():
email = "victim@example.com"
# Up to the limit: allowed, each a recorded failure.
for _ in range(limits.LOGIN_FAIL_LIMIT):
limits.check_login_allowed(email) # does not raise
limits.note_login_failure(email)
# One more crosses the line.
with pytest.raises(limits.HTTPException) as exc:
limits.check_login_allowed(email)
assert exc.value.status_code == 429
def test_login_throttle_is_per_account():
for _ in range(limits.LOGIN_FAIL_LIMIT):
limits.note_login_failure("a@example.com")
with pytest.raises(limits.HTTPException):
limits.check_login_allowed("a@example.com")
# A different account is unaffected — this is not an IP bucket.
limits.check_login_allowed("b@example.com") # must not raise
def test_successful_login_clears_the_streak():
email = "typo@example.com"
for _ in range(limits.LOGIN_FAIL_LIMIT):
limits.note_login_failure(email)
limits.note_login_success(email)
limits.check_login_allowed(email) # streak wiped — must not raise
def test_throttle_is_noop_in_local_mode(monkeypatch):
monkeypatch.setattr(auth, "MULTI_USER", False)
for _ in range(limits.LOGIN_FAIL_LIMIT * 3):
limits.note_login_failure("solo@example.com")
limits.check_login_allowed("solo@example.com") # never throttled locally