diff --git a/backend/app/ws.py b/backend/app/ws.py index 6c32990..6f2f17f 100644 --- a/backend/app/ws.py +++ b/backend/app/ws.py @@ -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,10 @@ def serialize_entity(entity: Entity) -> dict: } +def _client_ip(websocket: WebSocket) -> str: + return websocket.client.host if websocket.client else "unknown" + + async def _authenticate(websocket: WebSocket) -> uuid.UUID | None: raw_token = websocket.cookies.get(SESSION_COOKIE_NAME) if raw_token is None: @@ -227,7 +242,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 +286,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 +302,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 +416,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)}) diff --git a/backend/tests/test_ws_session.py b/backend/tests/test_ws_session.py index edc3196..99bb55f 100644 --- a/backend/tests/test_ws_session.py +++ b/backend/tests/test_ws_session.py @@ -7,6 +7,7 @@ import app.ws from app.entities import fallback_profile from app.models.contact_session import ContactSession from app.models.event import Event +from app.rate_limit import RateLimiter class FakeSpiritService: @@ -175,3 +176,59 @@ async def test_same_signature_recontacts_same_entity(sync_client): assert second_frame["entity"]["name"] == first assert second_frame["is_new"] is False assert second_frame["entity"]["contact_count"] == 2 + + +@pytest.mark.asyncio +async def test_summon_rate_limited_per_account(sync_client, monkeypatch): + # Swap in a tight, test-scoped limiter so this doesn't depend on (or + # pollute) the shared module-level budget other tests draw from. + monkeypatch.setattr( + app.ws, "summon_limiter", RateLimiter(max_requests=2, window_seconds=60) + ) + _login(sync_client, "account-limited") + + with _ws_connect(sync_client, sync_client.cookies.get("qm_session")) as ws: + _read_until(ws, "session") + for _ in range(2): + ws.send_json({"type": "summon"}) + _read_until(ws, "entity") + _read_until(ws, "utterance", kind="greeting") + + ws.send_json({"type": "summon"}) + rejection = _read_until(ws, "error") + assert rejection["code"] == "rate_limited" + assert rejection["message"] == ( + "The veil is crowded. The spirits need a moment before another summoning." + ) + + +@pytest.mark.asyncio +async def test_summon_rate_limited_per_ip_even_with_fresh_account( + sync_client, monkeypatch +): + # Starve only the IP bucket; the per-account limiter stays at its + # production default so each account below has plenty of its own budget + # left. This proves the IP limiter alone can reject a request — the + # gap the per-account-only limiters left open. + monkeypatch.setattr( + app.ws, "summon_ip_limiter", RateLimiter(max_requests=1, window_seconds=60) + ) + + _login(sync_client, "ip-limited-a") + with _ws_connect(sync_client, sync_client.cookies.get("qm_session")) as ws: + _read_until(ws, "session") + ws.send_json({"type": "summon"}) + _read_until(ws, "entity") # spends the single per-IP slot + + # A different account — its own per-account budget is untouched — but + # every connection in this test shares the same (simulated) source IP, + # which is already spent. + _login(sync_client, "ip-limited-b") + with _ws_connect(sync_client, sync_client.cookies.get("qm_session")) as ws: + _read_until(ws, "session") + ws.send_json({"type": "summon"}) + rejection = _read_until(ws, "error") + assert rejection["code"] == "rate_limited" + assert rejection["message"] == ( + "The veil is crowded. The spirits need a moment before another summoning." + )