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
+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 os
import threading
import time
from collections import defaultdict, deque
@@ -37,7 +38,27 @@ _windows: dict[tuple[str, str], deque] = defaultdict(deque)
_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:
"""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"
@@ -62,6 +83,63 @@ def rate_limit(scope: str, request: Request, user: models.User | None = None) ->
_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:
"""Drop callers whose whole window has expired (call with guard held) so
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
from .. import debuglog
from .. import debuglog, netguard
from .base import PromptParts, Provider, ProviderError
# 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]]:
"""Shared SSE plumbing for generate()/chat(): POST a streaming request
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)
received: list[str] = []
try:
+5
View File
@@ -105,13 +105,18 @@ def login(
raise HTTPException(400, "Accounts are disabled in local mode.")
limits.rate_limit("auth", request)
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()
if (
user is None
or not 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.")
limits.note_login_success(email)
_set_session_cookie(response, user.id)
return me_payload(user, db)
+6 -1
View File
@@ -1,8 +1,9 @@
import httpx
from fastapi import APIRouter, Depends, Request
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
router = APIRouter(prefix="/api/settings", tags=["settings"])
@@ -65,6 +66,10 @@ def update_settings(
async def list_endpoint_models(cfg: auth.ProviderConfig) -> dict:
"""GET the endpoint's /models listing. Doubles as a connectivity check, so
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"
headers = {}
if cfg.api_key: