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:
+7
-36
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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:
|
||||
|
||||
+101
-2
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user