Drop watch websockets on session loss, purge SSO cache on logout

- Successful SSO /auth/api/validate responses are cached per credential
  and perm/renew URL for 10s, so watch websocket re-checks do not hammer
  the auth backend. A POST to the logout endpoint purges all cached
  entries for the request's credentials immediately, so logout/login
  flows are not served stale successes.
- The watch websocket now re-validates auth before each forwarded
  message and every 10s when idle (SSO and built-in sessions alike).
  When the session is gone the client gets an auth error message and
  the socket is closed, instead of streaming updates forever.
- Token-authenticated (API/share token) sockets are exempt from
  re-validation; they are checked once at handshake.
This commit is contained in:
2026-08-12 21:20:12 +00:00
parent bd96b2c7ba
commit 7a0e473fb4
6 changed files with 523 additions and 38 deletions
+211
View File
@@ -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)