Merge gap-b-per-ip-limit: per-IP WS rate limiting via CF-Connecting-IP
This commit is contained in:
@@ -1,5 +1,24 @@
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from typing import Mapping
|
||||
|
||||
|
||||
def resolve_client_ip(headers: Mapping[str, str], direct_host: str | None) -> str:
|
||||
"""Resolves the real visitor IP for per-IP rate limiting.
|
||||
|
||||
The App CT sits behind a Cloudflare Tunnel that runs on a separate
|
||||
machine (see README Architecture) — every internet-facing connection's
|
||||
raw TCP peer is that tunnel machine, not the visitor, which would
|
||||
collapse per-IP limiting to a single shared bucket for all remote
|
||||
traffic. Cloudflare's edge sets `CF-Connecting-IP` itself, stripping any
|
||||
client-supplied value first, so it's safe to trust here. Direct
|
||||
LAN/local access (no Cloudflare in front, e.g. local dev) has no such
|
||||
header and falls back to the raw socket peer.
|
||||
"""
|
||||
forwarded = headers.get("cf-connecting-ip")
|
||||
if forwarded:
|
||||
return forwarded
|
||||
return direct_host or "unknown"
|
||||
|
||||
|
||||
class RateLimiter:
|
||||
|
||||
@@ -36,7 +36,7 @@ from app.models.contact_session import ContactSession
|
||||
from app.models.entity import Entity
|
||||
from app.models.entity_sighting import EntitySighting
|
||||
from app.models.event import Event
|
||||
from app.rate_limit import RateLimiter
|
||||
from app.rate_limit import RateLimiter, resolve_client_ip
|
||||
from app.telemetry import detect_wire_spike, sample_network
|
||||
from app.tts.piper import synthesize_spirit_voice
|
||||
from app.tts.voices import pick_voice
|
||||
@@ -53,6 +53,16 @@ fragment_limiter = RateLimiter(max_requests=30, window_seconds=60)
|
||||
question_limiter = RateLimiter(max_requests=6, window_seconds=60)
|
||||
summon_limiter = RateLimiter(max_requests=4, window_seconds=60)
|
||||
|
||||
# Per-IP limiters for the same trigger points (spec §5). Ollama is a shared,
|
||||
# single-instance, CPU-only resource — per-account limits alone don't stop
|
||||
# one compromised/scripted account from hammering it across many source IPs,
|
||||
# nor do they protect the resource from many accounts sharing one IP. IP
|
||||
# thresholds are looser than the per-account ones since a single IP can
|
||||
# legitimately host multiple accounts on a shared network.
|
||||
fragment_ip_limiter = RateLimiter(max_requests=60, window_seconds=60)
|
||||
question_ip_limiter = RateLimiter(max_requests=12, window_seconds=60)
|
||||
summon_ip_limiter = RateLimiter(max_requests=8, window_seconds=60)
|
||||
|
||||
AUDIO_DIR = Path(settings.data_dir) / "audio"
|
||||
|
||||
|
||||
@@ -65,6 +75,7 @@ def _audio_dir() -> Path:
|
||||
class SeanceState:
|
||||
user_id: uuid.UUID
|
||||
session_id: uuid.UUID
|
||||
client_ip: str
|
||||
send_queue: asyncio.Queue = field(default_factory=asyncio.Queue)
|
||||
mode: str = "unknown"
|
||||
language: str = "en"
|
||||
@@ -91,6 +102,11 @@ def serialize_entity(entity: Entity) -> dict:
|
||||
}
|
||||
|
||||
|
||||
def _client_ip(websocket: WebSocket) -> str:
|
||||
host = websocket.client.host if websocket.client else None
|
||||
return resolve_client_ip(websocket.headers, host)
|
||||
|
||||
|
||||
async def _authenticate(websocket: WebSocket) -> uuid.UUID | None:
|
||||
raw_token = websocket.cookies.get(SESSION_COOKIE_NAME)
|
||||
if raw_token is None:
|
||||
@@ -227,7 +243,12 @@ async def _summon(state: SeanceState, channel: str) -> tuple[Entity, bool]:
|
||||
|
||||
|
||||
async def _handle_summon(state: SeanceState) -> None:
|
||||
if not summon_limiter.allow(str(state.user_id)):
|
||||
# Short-circuits: an account already over its own cap never gets far
|
||||
# enough to spend from the IP budget too.
|
||||
if not (
|
||||
summon_limiter.allow(str(state.user_id))
|
||||
and summon_ip_limiter.allow(state.client_ip)
|
||||
):
|
||||
await state.send_queue.put(
|
||||
{
|
||||
"type": "error",
|
||||
@@ -266,7 +287,10 @@ async def _handle_anomaly(state: SeanceState, message: dict) -> None:
|
||||
await state.send_queue.put({"type": "status", "state": "attuning"})
|
||||
return
|
||||
|
||||
if not fragment_limiter.allow(str(state.user_id)):
|
||||
if not (
|
||||
fragment_limiter.allow(str(state.user_id))
|
||||
and fragment_ip_limiter.allow(state.client_ip)
|
||||
):
|
||||
return # anomalies during a crowded veil just pass unheard
|
||||
|
||||
try:
|
||||
@@ -279,7 +303,10 @@ async def _handle_anomaly(state: SeanceState, message: dict) -> None:
|
||||
|
||||
|
||||
async def _handle_question(state: SeanceState, text: str) -> None:
|
||||
if not question_limiter.allow(str(state.user_id)):
|
||||
if not (
|
||||
question_limiter.allow(str(state.user_id))
|
||||
and question_ip_limiter.allow(state.client_ip)
|
||||
):
|
||||
await state.send_queue.put(
|
||||
{
|
||||
"type": "error",
|
||||
@@ -390,7 +417,11 @@ async def session_socket(websocket: WebSocket) -> None:
|
||||
await db.commit()
|
||||
await db.refresh(contact_session)
|
||||
|
||||
state = SeanceState(user_id=user_id, session_id=contact_session.id)
|
||||
state = SeanceState(
|
||||
user_id=user_id,
|
||||
session_id=contact_session.id,
|
||||
client_ip=_client_ip(websocket),
|
||||
)
|
||||
sender = asyncio.create_task(_sender(state, websocket))
|
||||
await state.send_queue.put({"type": "session", "id": str(contact_session.id)})
|
||||
|
||||
|
||||
Reference in New Issue
Block a user