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
This commit is contained in:
2026-08-12 17:51:09 +00:00
parent 0889ef0d08
commit 85d8770ba6
4 changed files with 42 additions and 23 deletions
+2 -3
View File
@@ -341,13 +341,12 @@ 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")
+22 -2
View File
@@ -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 <token>")
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:
+2 -5
View File
@@ -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*."""
+13 -10
View File
@@ -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/<id> — 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/<id>
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: