From 85d8770ba633dd30134251a3fce37176b063f75b Mon Sep 17 00:00:00 2001 From: Mike Teehan Date: Wed, 12 Aug 2026 17:51:09 +0000 Subject: [PATCH] fix: address auth subsystem issues from ws-debug review - lib/auth: make RateLimiter.is_allowed read-only (no dict mutation on read) - daemon/server: add periodic blacklist_expired cleanup to poll loop (60s interval) - daemon/server: negotiate only matched Bearer subprotocol on WebSocket connect - webui/server: rewrite _is_personal_auth with path-prefix matching, cover WebAuthn register routes - daemon/handlers/auth: eliminate redundant get_user call in auth_update_user --- daemon/handlers/auth.py | 11 +++++------ daemon/server.py | 24 ++++++++++++++++++++++-- lib/auth.py | 7 ++----- webui/server.py | 23 +++++++++++++---------- 4 files changed, 42 insertions(+), 23 deletions(-) diff --git a/daemon/handlers/auth.py b/daemon/handlers/auth.py index b291ecc..0cfb306 100644 --- a/daemon/handlers/auth.py +++ b/daemon/handlers/auth.py @@ -341,16 +341,15 @@ def auth_update_user(_request: Any, body: Any) -> dict[str, Any]: if not username: raise ValueError("username is required") - existing = get_user(username) - if existing is None: + user = get_user(username) + if user is None: raise NotFoundError(f"User {username!r} not found") if "permissions" in body: update_permissions(username, body["permissions"]) - - user = get_user(username) - if user is None: - raise NotFoundError(f"User {username!r} not found") + user = get_user(username) + if user is None: + raise NotFoundError(f"User {username!r} not found") return { "id": user["id"], diff --git a/daemon/server.py b/daemon/server.py index 0648e2e..6c1c523 100644 --- a/daemon/server.py +++ b/daemon/server.py @@ -10,6 +10,7 @@ import json import logging import os import signal +import time from collections.abc import Callable from pathlib import Path from typing import Any @@ -17,6 +18,7 @@ from typing import Any from aiohttp import web from daemon.iface import PathLike +from lib.auth import blacklist_expired from lib.state import _DEFAULT_POLL_INTERVALS from lib.state import state as state_store @@ -361,6 +363,8 @@ def create_app() -> web.Application: _ws_subscribers: set[web.WebSocketResponse] = set() _ws_tasks: set[asyncio.Task[None]] = set() _poll_tasks: set[asyncio.Task[None]] = set() +_last_blacklist_cleanup: float = 0 +_last_blacklist_cleanup_lock: asyncio.Lock | None = None async def _handle_ws(request: web.Request) -> web.Response: @@ -376,12 +380,14 @@ async def _handle_ws(request: web.Request) -> web.Response: from lib.auth import validate_token token_param = None + matched_proto = None # Prefer subprotocol header (client JS sends "Bearer ") subprotocols = request.get_subprotocols() for proto in subprotocols or []: if proto and proto.startswith("Bearer "): token_param = proto[7:] + matched_proto = proto break if token_param is None: @@ -401,8 +407,10 @@ async def _handle_ws(request: web.Request) -> web.Response: if payload is None: return web.json_response({"ok": False, "error": "unauthorized"}, status=401) - # Negotiate the subprotocol the client sent - ws = web.WebSocketResponse(protocols=request.get_subprotocols()) + # Negotiate only the matched auth subprotocol (or all if token came from header) + ws = web.WebSocketResponse( + protocols=[matched_proto] if matched_proto else subprotocols + ) await ws.prepare(request) _ws_subscribers.add(ws) @@ -453,6 +461,12 @@ async def broadcast_tick(subsystems: list[str]) -> None: async def _poll_loop(subsystem: str, interval: int) -> None: """Periodically poll a subsystem for state changes and broadcast as needed.""" + global _last_blacklist_cleanup, _last_blacklist_cleanup_lock + + # Lazy-init lock (requires running event loop) + if _last_blacklist_cleanup_lock is None: + _last_blacklist_cleanup_lock = asyncio.Lock() + offset = int(hashlib.md5(subsystem.encode()).hexdigest(), 16) % interval await asyncio.sleep(offset) while True: @@ -463,6 +477,12 @@ async def _poll_loop(subsystem: str, interval: int) -> None: await broadcast_versions() elif volatile: await broadcast_tick([subsystem]) + # Periodic blacklist cleanup — coordinated across all poll loops + async with _last_blacklist_cleanup_lock: + now = time.time() + if now - _last_blacklist_cleanup >= 60: + blacklist_expired() + _last_blacklist_cleanup = now except asyncio.CancelledError: raise except Exception: diff --git a/lib/auth.py b/lib/auth.py index c56d289..f8f5067 100644 --- a/lib/auth.py +++ b/lib/auth.py @@ -379,11 +379,8 @@ class RateLimiter: now = time.time() cutoff = now - self.window timestamps = self.failures.get(key, []) - - # Clean old entries - self.failures[key] = [t for t in timestamps if t > cutoff] - - return len(self.failures[key]) < self.max_attempts + clean = [t for t in timestamps if t > cutoff] + return len(clean) < self.max_attempts def record_failure(self, key: str) -> None: """Record a failed attempt for *key*.""" diff --git a/webui/server.py b/webui/server.py index 89bb39f..6a368ba 100644 --- a/webui/server.py +++ b/webui/server.py @@ -128,14 +128,18 @@ _AUTH_EXEMPT = { # ── Personal auth routes (operates on own account, no subsystem permission needed) ── # These routes require a valid JWT but do NOT require an "auth" permission entry. # A user with only "firewall:read" can still view session, change password, logout, etc. -_AUTH_PERSONAL = { - ("GET", "/api/auth/session"), - ("POST", "/api/auth/password"), - ("POST", "/api/auth/logout"), - ("GET", "/api/auth/webauthn/credentials"), -} +# Method-agnostic — covers all HTTP methods for future-proofing. +_AUTH_PERSONAL_PATHS = ( + "/api/auth/session", + "/api/auth/password", + "/api/auth/logout", +) -# Pattern: DELETE /api/auth/webauthn/creds/ — match prefix only +_AUTH_PERSONAL_PREFIXES = ( + "/api/auth/webauthn/register-", + "/api/auth/webauthn/credentials", + "/api/auth/webauthn/creds/", +) def _subsystem_from_path(path: str) -> str | None: @@ -150,10 +154,9 @@ def _subsystem_from_path(path: str) -> str | None: def _is_personal_auth(method: str, path: str) -> bool: """Check if route is a personal auth operation (no subsystem permission needed).""" - if (method, path) in _AUTH_PERSONAL: + if path in _AUTH_PERSONAL_PATHS: return True - # Personal credential deletion: DELETE /api/auth/webauthn/creds/ - return method == "DELETE" and path.startswith("/api/auth/webauthn/creds/") + return any(path.startswith(prefix) for prefix in _AUTH_PERSONAL_PREFIXES) def _has_permission(perms: dict, subsystem: str, method: str) -> bool: