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:
co-authored by
Claude Opus 4.8
parent
8857757642
commit
c500203270
+8
-3
@@ -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"]
|
||||||
|
|||||||
@@ -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."""
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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:
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user