diff --git a/cista/api.py b/cista/api.py index 59e5d70..051a17a 100644 --- a/cista/api.py +++ b/cista/api.py @@ -5,7 +5,6 @@ import msgspec from mediapreview.office import is_available_cached from sanic import Blueprint, json from sanic.exceptions import BadRequest -from sanic.log import logger from cista import __version__, auth, config, sharefs, sso, watching from cista.auth import ( @@ -15,7 +14,11 @@ from cista.auth import ( list_tokens_handler, ) from cista.fileio import FileServer -from cista.util.apphelpers import websocket_wrapper +from cista.util.apphelpers import ( + get_watch_user_info, + run_auth_checked_watch, + websocket_wrapper, +) bp = Blueprint("api", url_prefix="/api") fileserver = FileServer() @@ -36,29 +39,7 @@ async def stop_fileserver(app): @bp.websocket("watch") @websocket_wrapper async def watch(req, ws): - # Build user info from either built-in auth or SSO - user_info = None - if sso.paskia_enabled(): - # SSO auth: call validation to get user info (don't enforce auth in public mode) - try: - # WebSocket cannot forward Set-Cookie, so ask the auth backend not to - # renew the session here; renewal happens on the HTTP side instead. - await sso.validate_sso_request(req, renew=False) - except Exception as e: - logger.debug("watch SSO validation failed: %s", e) - if sso_user := getattr(req.ctx, "sso_user", None): - ctx = sso_user.get("ctx", {}) - perms = ctx.get("permissions", []) - user_info = { - "username": ctx.get("user", {}).get("display_name", ""), - "privileged": "cista:admin" in perms, - } - elif req.ctx.user: - # Built-in auth: use local user database - user_info = { - "username": req.ctx.username, - "privileged": req.ctx.user.privileged, - } + user_info = await get_watch_user_info(req) await ws.send( msgspec.json.encode( @@ -85,17 +66,7 @@ async def watch(req, ws): await ws.send(root) else: await ws.send(watching.format_root(sharefs.build_virtual_root(share_token))) - # Send updates - while True: - msg = await q.get() - if share_token is None or ( - isinstance(msg, str) and msg.startswith('{"space"') - ): - await ws.send(msg) - else: - await ws.send( - watching.format_root(sharefs.build_virtual_root(share_token)) - ) + await run_auth_checked_watch(req, ws, q, share_token) except RuntimeError as e: if str(e) == "cannot schedule new futures after shutdown": return # Server shutting down, drop the WebSocket diff --git a/cista/app.py b/cista/app.py index 89aa1cc..9f15e9d 100644 --- a/cista/app.py +++ b/cista/app.py @@ -102,6 +102,19 @@ async def forward_sso_cookies(req, res): res.headers.add("set-cookie", cookie) +@app.on_response +async def invalidate_sso_cache_on_logout(req, _res): + """Purge cached SSO validations after a logout request.""" + # Convenience for logout/login flows, not a security feature: cached + # entries expire after 10 seconds anyway if the logout happened elsewhere. + if ( + sso.paskia_enabled() + and req.method == "POST" + and req.path in {"/auth/api/logout", "/auth/logout"} + ): + sso.invalidate_validation_cache(req) + + @app.on_response async def persist_auth_session(req, res): """Persist a session cookie after successful Authorization-based auth.""" diff --git a/cista/sso.py b/cista/sso.py index e64bc7d..b02fc12 100644 --- a/cista/sso.py +++ b/cista/sso.py @@ -10,8 +10,10 @@ Environment variables: """ import asyncio +import hashlib import os import re +from time import time import httpx import websockets @@ -62,6 +64,40 @@ async def close_client(): _client = None +# In-memory cache for successful SSO /auth/api/validate responses. +# Keyed by (credential hash, validation URL) so that entries for different +# perms/renew flags coexist and all entries for a credential can be purged +# on logout. +_VALIDATE_CACHE_TTL = 10 +_validate_cache: dict[tuple[str, str], tuple[float, dict]] = {} + + +def _validate_credential_key(request) -> str: + """Return a stable key for the credential material in the request.""" + cookie = request.headers.get("cookie", "") + authorization = request.headers.get("authorization", "") + return hashlib.sha256(f"{cookie}\x00{authorization}".encode()).hexdigest() + + +def _cleanup_validate_cache() -> None: + """Drop expired cache entries.""" + now = time() + for key, (timestamp, _) in list(_validate_cache.items()): + if now - timestamp >= _VALIDATE_CACHE_TTL: + del _validate_cache[key] + + +def invalidate_validation_cache(request) -> None: + """Remove cached SSO validations for the credentials carried by *request*. + + Called after a logout request so the next request is forced to the + backend instead of being served from a stale success cache. + """ + credential_key = _validate_credential_key(request) + for key in [key for key in _validate_cache if key[0] == credential_key]: + del _validate_cache[key] + + async def validate_sso_request( request, *, perm: str = "cista:login", renew: bool = True ) -> dict | None: @@ -105,6 +141,17 @@ async def validate_sso_request( if not renew: url += "&renew=0" + credential_key = _validate_credential_key(request) + cache_key = (credential_key, url) + + cached = _validate_cache.get(cache_key) + if cached is not None: + timestamp, data = cached + if time() - timestamp < _VALIDATE_CACHE_TTL: + request.ctx.sso_user = data + return data + del _validate_cache[cache_key] + try: response = await client.post( url, @@ -121,6 +168,8 @@ async def validate_sso_request( request.ctx.sso_user = {} return {} else: + _cleanup_validate_cache() + _validate_cache[cache_key] = (time(), data) return data try: diff --git a/cista/util/apphelpers.py b/cista/util/apphelpers.py index e6303a8..923e91a 100644 --- a/cista/util/apphelpers.py +++ b/cista/util/apphelpers.py @@ -1,14 +1,15 @@ +import asyncio import time from functools import wraps import msgspec import websockets.exceptions from sanic import errorpages -from sanic.exceptions import SanicException +from sanic.exceptions import SanicException, Unauthorized from sanic.log import logger from sanic.response import raw, redirect -from cista import auth +from cista import auth, config, session, sharefs, sso, watching from cista.protocol import ErrorMsg from cista.sanic_logging import log_ws_close, log_ws_open @@ -101,3 +102,101 @@ def websocket_wrapper(handler): log_ws_close(ws_id, close_code, duration, extra=close_extra) return wrapper + + +class StopError(Exception): + """Used internally to end a watch websocket's task group cleanly.""" + + +async def get_watch_user_info(request): + """Return the current user info for a watch websocket, re-validating auth. + + Handles all three auth modes: + - Paskia/SSO: re-validates with the auth backend (cache-friendly) + - Built-in: re-reads the local session cookie from the live store + - Public: returns None when no session is present + + Raises Unauthorized/Forbidden in non-public mode when the session is gone. + """ + # Long-lived API/share tokens are validated once at handshake; re-checking + # them on every message would add unnecessary backend calls. + if getattr(request.ctx, "auth_token", None) is not None: + return None + + if sso.paskia_enabled(): + try: + await sso.validate_sso_request(request, renew=False) + except SanicException: + if config.config.public: + return None + raise + sso_user = getattr(request.ctx, "sso_user", None) or {} + if sso_user: + ctx = sso_user.get("ctx", {}) + perms = ctx.get("permissions", []) + return { + "username": ctx.get("user", {}).get("display_name", ""), + "privileged": "cista:admin" in perms, + } + return None + + s = session.get(request) + if s: + user = config.config.users.get(s.get("username")) + if user: + return {"username": s["username"], "privileged": user.privileged} + + if config.config.public: + return None + + raise Unauthorized("Login required", "cookie", quiet=True) + + +async def _check_watch_auth_or_stop(request, ws) -> None: + """Re-validate watch auth; on failure send an error and raise StopError.""" + try: + await get_watch_user_info(request) + except SanicException as exc: + # Match the error format used by websocket_wrapper + message = f"⚠️ {str(exc) or 'Authentication error'}" + await asend( + ws, + ErrorMsg( + {"code": exc.status_code, "message": message, **(exc.context or {})} + ), + ) + raise StopError from None + + +async def run_auth_checked_watch(request, ws, queue, share_token) -> None: + """Run the watch websocket loop with per-message and periodic auth checks. + + Messages are forwarded from *queue* to *ws*. Auth is re-checked before each + message (hitting the SSO cache in the common case) and every 10 seconds when + idle, so a session invalidated on the backend does not stay open forever. + """ + + async def consume() -> None: + while True: + item = await queue.get() + await _check_watch_auth_or_stop(request, ws) + if share_token is None or ( + isinstance(item, str) and item.startswith('{"space"') + ): + await ws.send(item) + else: + await ws.send( + watching.format_root(sharefs.build_virtual_root(share_token)) + ) + + async def idle_checker() -> None: + while True: + await asyncio.sleep(10) + await _check_watch_auth_or_stop(request, ws) + + try: + async with asyncio.TaskGroup() as tg: + tg.create_task(consume()) + tg.create_task(idle_checker()) + except* StopError: + pass diff --git a/tests/test_sso_cache.py b/tests/test_sso_cache.py new file mode 100644 index 0000000..64cd255 --- /dev/null +++ b/tests/test_sso_cache.py @@ -0,0 +1,142 @@ +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import httpx +import pytest +from sanic.exceptions import Unauthorized + +from cista import sso + + +def _make_request(cookie: str = "", authorization: str = ""): + req = SimpleNamespace() + req.headers = {} + if cookie: + req.headers["cookie"] = cookie + if authorization: + req.headers["authorization"] = authorization + req.client_ip = "127.0.0.1" + req.host = "test.local" + req.scheme = "http" + req.ctx = SimpleNamespace() + return req + + +@pytest.fixture(autouse=True) +def _reset_sso_cache_and_client(monkeypatch): + """Clear the SSO validation cache and shared client between tests.""" + sso._validate_cache.clear() + sso._client = None + monkeypatch.setenv("PASKIA_BACKEND_URL", "http://test-paskia.local") + monkeypatch.setattr(sso, "PASKIA_BACKEND_URL", "http://test-paskia.local") + yield + sso._validate_cache.clear() + + +@pytest.fixture +def mock_client(monkeypatch): + client = AsyncMock() + client.is_closed = False + client.headers = {} + monkeypatch.setattr(sso, "_client", client) + return client + + +@pytest.mark.asyncio +async def test_validate_sso_request_caches_successful_responses(mock_client): + req = _make_request(cookie="session=abc123") + mock_client.post.return_value = httpx.Response(200, json={"user": "alice"}) + + data1 = await sso.validate_sso_request(req) + data2 = await sso.validate_sso_request(req) + + assert data1 == {"user": "alice"} + assert data2 == data1 + assert mock_client.post.call_count == 1 + assert req.ctx.sso_user == {"user": "alice"} + + +@pytest.mark.asyncio +async def test_validate_sso_request_does_not_cache_errors(mock_client): + req = _make_request(cookie="session=bad") + mock_client.post.return_value = httpx.Response(401, json={"detail": "nope"}) + + with pytest.raises(Unauthorized): + await sso.validate_sso_request(req) + with pytest.raises(Unauthorized): + await sso.validate_sso_request(req) + + assert mock_client.post.call_count == 2 + + +@pytest.mark.asyncio +async def test_validate_sso_request_cache_is_per_credential(mock_client): + req_alice = _make_request(cookie="session=alice") + req_bob = _make_request(cookie="session=bob") + responses = { + "alice": httpx.Response(200, json={"user": "alice"}), + "bob": httpx.Response(200, json={"user": "bob"}), + } + + def side_effect(*args, **kwargs): + cookie = kwargs.get("headers", {}).get("cookie", "") + if "alice" in cookie: + return responses["alice"] + return responses["bob"] + + mock_client.post.side_effect = side_effect + + assert await sso.validate_sso_request(req_alice) == {"user": "alice"} + assert await sso.validate_sso_request(req_bob) == {"user": "bob"} + assert await sso.validate_sso_request(req_alice) == {"user": "alice"} + + assert mock_client.post.call_count == 2 + + +@pytest.mark.asyncio +async def test_validate_sso_request_cache_is_per_permission(mock_client): + req = _make_request(cookie="session=abc123") + mock_client.post.return_value = httpx.Response(200, json={"user": "alice"}) + + await sso.validate_sso_request(req, perm="cista:login") + await sso.validate_sso_request(req, perm="cista:admin") + + assert mock_client.post.call_count == 2 + + +@pytest.mark.asyncio +async def test_invalidate_validation_cache_forces_backend_call(mock_client): + req = _make_request(cookie="session=abc123") + mock_client.post.return_value = httpx.Response(200, json={"user": "alice"}) + + await sso.validate_sso_request(req) + sso.invalidate_validation_cache(req) + await sso.validate_sso_request(req) + + assert mock_client.post.call_count == 2 + + +@pytest.mark.asyncio +async def test_invalidate_validation_cache_only_affects_same_credentials(mock_client): + alice = _make_request(cookie="session=alice") + bob = _make_request(cookie="session=bob") + responses = { + "alice": httpx.Response(200, json={"user": "alice"}), + "bob": httpx.Response(200, json={"user": "bob"}), + } + + def side_effect(*args, **kwargs): + cookie = kwargs.get("headers", {}).get("cookie", "") + return responses["alice"] if "alice" in cookie else responses["bob"] + + mock_client.post.side_effect = side_effect + + await sso.validate_sso_request(alice) + await sso.validate_sso_request(bob) + sso.invalidate_validation_cache(alice) + + assert await sso.validate_sso_request(alice) == {"user": "alice"} + assert await sso.validate_sso_request(bob) == {"user": "bob"} + + # Alice is re-fetched; bob is still cached. + assert mock_client.post.call_count == 3 diff --git a/tests/test_watch_auth.py b/tests/test_watch_auth.py new file mode 100644 index 0000000..6a8586f --- /dev/null +++ b/tests/test_watch_auth.py @@ -0,0 +1,211 @@ +import asyncio +import contextlib +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from sanic.exceptions import Unauthorized + +from cista import auth, config, session, sso +from cista.util.apphelpers import ( + get_watch_user_info, + run_auth_checked_watch, +) + + +def _make_request(cookie: str = "", auth_token=None): + req = SimpleNamespace() + req.headers = {} + req.cookies = {} + if cookie: + req.headers["cookie"] = cookie + for part in cookie.split(";"): + k, _, v = part.strip().partition("=") + req.cookies[k] = v + req.ctx = SimpleNamespace() + if auth_token: + req.ctx.auth_token = auth_token + return req + + +@pytest.fixture(autouse=True) +def _reset(tmp_path, monkeypatch): + alice = config.User() + auth.set_password(alice, "secret") + admin = config.User(privileged=True) + auth.set_password(admin, "admin-secret") + config.config = config.Config( + path=tmp_path, + listen=":0", + public=False, + users={"alice": alice, "admin": admin}, + ) + session._sessions.clear() + sso._validate_cache.clear() + monkeypatch.setattr(sso, "PASKIA_BACKEND_URL", "") + yield + session._sessions.clear() + sso._validate_cache.clear() + + +@pytest.mark.asyncio +async def test_get_watch_user_info_builtin_valid_session(): + token = "valid-token" + session.put(token, "alice") + req = _make_request(cookie=f"cista={token}") + + info = await get_watch_user_info(req) + + assert info == {"username": "alice", "privileged": False} + + +@pytest.mark.asyncio +async def test_get_watch_user_info_builtin_admin(): + token = "admin-token" + session.put(token, "admin") + req = _make_request(cookie=f"cista={token}") + + info = await get_watch_user_info(req) + + assert info == {"username": "admin", "privileged": True} + + +@pytest.mark.asyncio +async def test_get_watch_user_info_builtin_invalid_session_raises(): + req = _make_request(cookie="cista=bad-token") + + with pytest.raises(Unauthorized): + await get_watch_user_info(req) + + +@pytest.mark.asyncio +async def test_get_watch_user_info_builtin_public_no_session(): + config.config.public = True + req = _make_request() + + info = await get_watch_user_info(req) + + assert info is None + + +@pytest.mark.asyncio +async def test_get_watch_user_info_sso_valid(monkeypatch): + monkeypatch.setattr(sso, "PASKIA_BACKEND_URL", "http://test-paskia.local") + + async def mock_validate(request, *, renew=True): + request.ctx.sso_user = { + "ctx": { + "user": {"display_name": "alice"}, + "permissions": ["cista:login", "cista:admin"], + } + } + + monkeypatch.setattr(sso, "validate_sso_request", mock_validate) + req = _make_request(cookie="session=abc") + + info = await get_watch_user_info(req) + + assert info == {"username": "alice", "privileged": True} + + +@pytest.mark.asyncio +async def test_get_watch_user_info_sso_nonpublic_invalid_raises(monkeypatch): + monkeypatch.setattr(sso, "PASKIA_BACKEND_URL", "http://test-paskia.local") + + async def mock_validate(request, *, renew=True): + raise Unauthorized("Session expired", quiet=True) + + monkeypatch.setattr(sso, "validate_sso_request", mock_validate) + req = _make_request(cookie="session=abc") + + with pytest.raises(Unauthorized): + await get_watch_user_info(req) + + +@pytest.mark.asyncio +async def test_get_watch_user_info_sso_public_invalid_returns_none(monkeypatch): + monkeypatch.setattr(sso, "PASKIA_BACKEND_URL", "http://test-paskia.local") + config.config.public = True + + async def mock_validate(request, *, renew=True): + raise Unauthorized("Session expired", quiet=True) + + monkeypatch.setattr(sso, "validate_sso_request", mock_validate) + req = _make_request(cookie="session=abc") + + info = await get_watch_user_info(req) + + assert info is None + + +@pytest.mark.asyncio +async def test_run_auth_checked_watch_forwards_messages_while_valid(): + token = "valid-token" + session.put(token, "alice") + req = _make_request(cookie=f"cista={token}") + ws = AsyncMock() + q = asyncio.Queue() + + async def producer(): + await q.put('{"space":{}}') + await q.put('{"update":[]}') + # Keep consumer alive briefly, then invalidate. + await asyncio.sleep(0.05) + session._sessions.pop(token, None) + await q.put('{"update":[]}') + + await asyncio.wait_for( + asyncio.gather(producer(), run_auth_checked_watch(req, ws, q, None)), + timeout=1.0, + ) + + calls = [c.args[0] for c in ws.send.call_args_list] + assert calls[0] == '{"space":{}}' + assert calls[1] == '{"update":[]}' + assert '"error"' in calls[2] + assert len(calls) == 3 + + +@pytest.mark.asyncio +async def test_get_watch_user_info_token_auth_skips_revalidation(): + """Token-based auth is considered valid without re-checking the token.""" + token_id = "api-token" + config.config.tokens[token_id] = config.Token( + key=token_id, username="alice", kind="api", mode="rw" + ) + req = _make_request(auth_token=config.config.tokens[token_id]) + + info = await get_watch_user_info(req) + + assert info is None + + +@pytest.mark.asyncio +async def test_run_auth_checked_watch_token_auth_does_not_send_errors(): + """Token-based sockets keep forwarding messages without re-validating.""" + token_id = "api-token" + config.config.tokens[token_id] = config.Token( + key=token_id, username="alice", kind="api", mode="rw" + ) + req = _make_request(auth_token=config.config.tokens[token_id]) + ws = AsyncMock() + q = asyncio.Queue() + + async def producer(): + await q.put('{"space":{}}') + await q.put('{"update":[]}') + # Deleting the token should not affect the already-open websocket. + await asyncio.sleep(0.05) + del config.config.tokens[token_id] + + runner = asyncio.create_task(run_auth_checked_watch(req, ws, q, None)) + await asyncio.wait_for(producer(), timeout=1.0) + await asyncio.sleep(0.05) + runner.cancel() + with contextlib.suppress(asyncio.CancelledError): + await runner + + calls = [c.args[0] for c in ws.send.call_args_list] + assert calls[0] == '{"space":{}}' + assert calls[1] == '{"update":[]}' + assert not any('"error"' in c for c in calls)