Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c445061451 | ||
|
|
43fd7df098 | ||
|
|
768ccd7739 | ||
|
|
321a86acf8 | ||
|
|
b5f4f67813 | ||
|
|
697a9416e8 | ||
|
|
d005ae1d88 | ||
|
|
3d8e20de8f | ||
|
|
54d7129bfc | ||
|
|
0c9fe7638c | ||
|
|
226f96c477 | ||
|
|
29e816cb53 | ||
|
|
7a0e473fb4 |
+2
-1
@@ -1,5 +1,6 @@
|
|||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
from importlib.metadata import version as pkg_version
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from docopt import docopt
|
from docopt import docopt
|
||||||
@@ -30,7 +31,7 @@ def create_startup_box(
|
|||||||
*, folder, url, unix=None, dev=False, paskia_url=None, public=False
|
*, folder, url, unix=None, dev=False, paskia_url=None, public=False
|
||||||
):
|
):
|
||||||
"""Create a framed startup box with server information."""
|
"""Create a framed startup box with server information."""
|
||||||
title = f"Cista {cista.__version__}"
|
title = f"Cista {cista.__version__} (mediapreview {pkg_version('mediapreview')})"
|
||||||
listen = unix or url
|
listen = unix or url
|
||||||
location = f"{folder} @ {listen}"
|
location = f"{folder} @ {listen}"
|
||||||
lines = [title, location]
|
lines = [title, location]
|
||||||
|
|||||||
+7
-36
@@ -5,7 +5,6 @@ import msgspec
|
|||||||
from mediapreview.office import is_available_cached
|
from mediapreview.office import is_available_cached
|
||||||
from sanic import Blueprint, json
|
from sanic import Blueprint, json
|
||||||
from sanic.exceptions import BadRequest
|
from sanic.exceptions import BadRequest
|
||||||
from sanic.log import logger
|
|
||||||
|
|
||||||
from cista import __version__, auth, config, sharefs, sso, watching
|
from cista import __version__, auth, config, sharefs, sso, watching
|
||||||
from cista.auth import (
|
from cista.auth import (
|
||||||
@@ -15,7 +14,11 @@ from cista.auth import (
|
|||||||
list_tokens_handler,
|
list_tokens_handler,
|
||||||
)
|
)
|
||||||
from cista.fileio import FileServer
|
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")
|
bp = Blueprint("api", url_prefix="/api")
|
||||||
fileserver = FileServer()
|
fileserver = FileServer()
|
||||||
@@ -36,29 +39,7 @@ async def stop_fileserver(app):
|
|||||||
@bp.websocket("watch")
|
@bp.websocket("watch")
|
||||||
@websocket_wrapper
|
@websocket_wrapper
|
||||||
async def watch(req, ws):
|
async def watch(req, ws):
|
||||||
# Build user info from either built-in auth or SSO
|
user_info = await get_watch_user_info(req)
|
||||||
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,
|
|
||||||
}
|
|
||||||
|
|
||||||
await ws.send(
|
await ws.send(
|
||||||
msgspec.json.encode(
|
msgspec.json.encode(
|
||||||
@@ -85,17 +66,7 @@ async def watch(req, ws):
|
|||||||
await ws.send(root)
|
await ws.send(root)
|
||||||
else:
|
else:
|
||||||
await ws.send(watching.format_root(sharefs.build_virtual_root(share_token)))
|
await ws.send(watching.format_root(sharefs.build_virtual_root(share_token)))
|
||||||
# Send updates
|
await run_auth_checked_watch(req, ws, q, share_token)
|
||||||
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))
|
|
||||||
)
|
|
||||||
except RuntimeError as e:
|
except RuntimeError as e:
|
||||||
if str(e) == "cannot schedule new futures after shutdown":
|
if str(e) == "cannot schedule new futures after shutdown":
|
||||||
return # Server shutting down, drop the WebSocket
|
return # Server shutting down, drop the WebSocket
|
||||||
|
|||||||
@@ -102,6 +102,19 @@ async def forward_sso_cookies(req, res):
|
|||||||
res.headers.add("set-cookie", cookie)
|
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
|
@app.on_response
|
||||||
async def persist_auth_session(req, res):
|
async def persist_auth_session(req, res):
|
||||||
"""Persist a session cookie after successful Authorization-based auth."""
|
"""Persist a session cookie after successful Authorization-based auth."""
|
||||||
|
|||||||
+1
-1
@@ -226,7 +226,7 @@ def hydrate_request_auth_context(request, *, source: str) -> None:
|
|||||||
|
|
||||||
|
|
||||||
_AUTH_REALM = "cista"
|
_AUTH_REALM = "cista"
|
||||||
_AUTH_CACHE_TTL = 10
|
_AUTH_CACHE_TTL = 300
|
||||||
_auth_cache: dict[str, tuple[float, config.User]] = {}
|
_auth_cache: dict[str, tuple[float, config.User]] = {}
|
||||||
_WINDOWS_UA_HINTS = (
|
_WINDOWS_UA_HINTS = (
|
||||||
"windows",
|
"windows",
|
||||||
|
|||||||
+10
-1
@@ -6,6 +6,7 @@ only bridges cista's config-derived JWT secret into it and wires the
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import mediapreview.office
|
import mediapreview.office
|
||||||
@@ -36,4 +37,12 @@ def setup_docker(confdir: Path | None = None) -> int:
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
configure()
|
configure()
|
||||||
return mediapreview.office.setup_docker()
|
try:
|
||||||
|
mediapreview.office.setup_docker()
|
||||||
|
finally:
|
||||||
|
# Print regardless of build outcome: the secret is deterministic
|
||||||
|
# (derived from the config).
|
||||||
|
sys.stdout.write(
|
||||||
|
f"ONLYOFFICE_JWT_SECRET={os.environ['ONLYOFFICE_JWT_SECRET']}\n"
|
||||||
|
)
|
||||||
|
return 0
|
||||||
|
|||||||
+17
-59
@@ -6,22 +6,21 @@ Sanic with auth, etag negotiation and the in-memory response cache.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import re
|
|
||||||
import urllib.parse
|
import urllib.parse
|
||||||
from pathlib import PurePosixPath
|
from pathlib import PurePosixPath
|
||||||
from urllib.parse import unquote
|
from urllib.parse import unquote
|
||||||
from wsgiref.handlers import format_date_time
|
from wsgiref.handlers import format_date_time
|
||||||
|
|
||||||
import httpx
|
|
||||||
from mediapreview import CachedPreview, PreviewCache, is_previewable_path
|
from mediapreview import CachedPreview, PreviewCache, is_previewable_path
|
||||||
|
from mediapreview.exceptions import (
|
||||||
|
PreviewBackendError,
|
||||||
|
PreviewCancelledError,
|
||||||
|
PreviewError,
|
||||||
|
)
|
||||||
from mediapreview.formats import OFFICE_PREVIEW_SUFFIXES
|
from mediapreview.formats import OFFICE_PREVIEW_SUFFIXES
|
||||||
from mediapreview.formats import expected_backend as _expected_preview_backend
|
from mediapreview.formats import expected_backend as _expected_preview_backend
|
||||||
from mediapreview.office import onlyoffice_error_short_text
|
|
||||||
from mediapreview.pool import (
|
from mediapreview.pool import (
|
||||||
PREVIEW_TIMEOUT,
|
PREVIEW_TIMEOUT,
|
||||||
PreviewError,
|
|
||||||
PreviewPoolClosedError,
|
|
||||||
PreviewTimeoutError,
|
|
||||||
generate_office_preview,
|
generate_office_preview,
|
||||||
run_preview,
|
run_preview,
|
||||||
)
|
)
|
||||||
@@ -39,21 +38,6 @@ bp = Blueprint("preview", url_prefix="/preview")
|
|||||||
_preview_cache = PreviewCache(capacity=500)
|
_preview_cache = PreviewCache(capacity=500)
|
||||||
|
|
||||||
|
|
||||||
def _shorten_error(detail: str) -> str:
|
|
||||||
"""Shorten an upstream backend error for the single-line access log.
|
|
||||||
|
|
||||||
Backend errors arrive verbatim from ffmpeg/pyvips/pymupdf and often carry
|
|
||||||
an '[Errno N]' prefix, the input file path, and multi-line library noise —
|
|
||||||
all redundant with the URL already in the log line. Keep the first line,
|
|
||||||
drop the bracket prefix, and cut at the first ': ' separator.
|
|
||||||
"""
|
|
||||||
lines = detail.splitlines()
|
|
||||||
if not lines:
|
|
||||||
return ""
|
|
||||||
first_line = re.sub(r"^\[[^\]]*\]\s*", "", lines[0].strip())
|
|
||||||
return first_line.split(": ", 1)[0]
|
|
||||||
|
|
||||||
|
|
||||||
@bp.on_request
|
@bp.on_request
|
||||||
async def verify_preview(request):
|
async def verify_preview(request):
|
||||||
"""Verify access to preview routes."""
|
"""Verify access to preview routes."""
|
||||||
@@ -98,7 +82,9 @@ async def preview(req, path):
|
|||||||
logger.debug(f"Preview cache hit: {rel}")
|
logger.debug(f"Preview cache hit: {rel}")
|
||||||
return raw(cached.body, headers=cached.headers)
|
return raw(cached.body, headers=cached.headers)
|
||||||
|
|
||||||
# Generate preview
|
# Generate preview. The outer deadline is strict: pool internals have
|
||||||
|
# their own timeouts, but queueing (workers, the OnlyOffice semaphore)
|
||||||
|
# must not let a request exceed PREVIEW_TIMEOUT.
|
||||||
try:
|
try:
|
||||||
if filepath.suffix.lower() in OFFICE_PREVIEW_SUFFIXES:
|
if filepath.suffix.lower() in OFFICE_PREVIEW_SUFFIXES:
|
||||||
img, preview_resp = await asyncio.wait_for(
|
img, preview_resp = await asyncio.wait_for(
|
||||||
@@ -113,45 +99,17 @@ async def preview(req, path):
|
|||||||
except TimeoutError:
|
except TimeoutError:
|
||||||
req.ctx.log_extra = f"{_expected_preview_backend(filepath)} timeout"
|
req.ctx.log_extra = f"{_expected_preview_backend(filepath)} timeout"
|
||||||
return empty(503)
|
return empty(503)
|
||||||
except PreviewTimeoutError as e:
|
|
||||||
req.ctx.log_extra = (
|
|
||||||
f"{(e.backend or _expected_preview_backend(filepath))} timeout"
|
|
||||||
)
|
|
||||||
return empty(503)
|
|
||||||
except httpx.HTTPStatusError:
|
|
||||||
req.ctx.log_extra = "onlyoffice N/A"
|
|
||||||
return empty(503)
|
|
||||||
except httpx.RequestError:
|
|
||||||
req.ctx.log_extra = "onlyoffice N/A"
|
|
||||||
return empty(503)
|
|
||||||
except RuntimeError as e:
|
|
||||||
detail = str(e)
|
|
||||||
if detail.startswith("OnlyOffice"):
|
|
||||||
req.ctx.log_extra = onlyoffice_error_short_text(detail)
|
|
||||||
return empty(503)
|
|
||||||
raise
|
|
||||||
except PreviewPoolClosedError:
|
|
||||||
# Server is shutting down; not an error, just a cancelled preview.
|
|
||||||
req.ctx.log_extra = "preview cancelled"
|
|
||||||
return empty(503)
|
|
||||||
except PreviewError as e:
|
except PreviewError as e:
|
||||||
detail = str(e)
|
# mediapreview is responsible for backend-specific diagnostics; cista only
|
||||||
if detail == "preview worker error" and e.stderr:
|
# needs the backend name, a short access-log reason, and a response status.
|
||||||
captured = e.stderr.strip()
|
if isinstance(e, PreviewCancelledError):
|
||||||
if captured:
|
req.ctx.log_extra = e.short or "preview cancelled"
|
||||||
detail = captured.splitlines()[0]
|
raise asyncio.CancelledError from e
|
||||||
# The worker already logged the failure (with traceback where the
|
status = 422 if isinstance(e, PreviewBackendError) else 503
|
||||||
# error occurred) — annotate the access log instead of re-logging,
|
req.ctx.log_extra = f"{e.backend}: {e.short}" if e.backend else e.short
|
||||||
# with a shortened reason. In dev mode, print the full error too.
|
|
||||||
backend = e.backend or _expected_preview_backend(filepath)
|
|
||||||
if req.app.debug:
|
if req.app.debug:
|
||||||
full = detail
|
logger.warning("%s", str(e))
|
||||||
if e.stderr and e.stderr.strip() not in detail:
|
return empty(status)
|
||||||
full = f"{detail}\n{e.stderr.strip()}"
|
|
||||||
logger.warning("[%s] preview failed: %s", backend, full.strip())
|
|
||||||
short = _shorten_error(detail)
|
|
||||||
req.ctx.log_extra = f"{backend}: {short}" if short else backend
|
|
||||||
return empty(422)
|
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
# Server shutdown or client disconnect: the connection is being torn
|
# Server shutdown or client disconnect: the connection is being torn
|
||||||
# down, so responding is impossible — just annotate the access log.
|
# down, so responding is impossible — just annotate the access log.
|
||||||
|
|||||||
+37
-5
@@ -274,11 +274,44 @@ def configure_access_logging() -> None:
|
|||||||
|
|
||||||
|
|
||||||
def configure_main_logging() -> None:
|
def configure_main_logging() -> None:
|
||||||
"""Replace Sanic's verbose 'Main yyyy-mm-dd INFO:' prefix with emoji-only format.
|
"""Replace Sanic's verbose 'Main yyyy-mm-dd INFO:' prefix with emoji-only format
|
||||||
|
|
||||||
Patches LOGGING_CONFIG_DEFAULTS so the formatter survives every dictConfig
|
and make sure the root logger catches unhandled loggers instead of falling back
|
||||||
call Sanic makes during serve_single() / serve().
|
to logging.lastResort (which prints a bare message with no level prefix).
|
||||||
|
|
||||||
|
Patches LOGGING_CONFIG_DEFAULTS so the formatter and root logger survive every
|
||||||
|
dictConfig call Sanic makes during serve_single() / serve().
|
||||||
"""
|
"""
|
||||||
|
# Give the root logger a real handler so third-party warnings (e.g.
|
||||||
|
# mediapreview.office) are formatted with the emoji prefix instead of being
|
||||||
|
# printed plain by logging.lastResort.
|
||||||
|
root = logging.getLogger()
|
||||||
|
root.setLevel(logging.WARNING)
|
||||||
|
if not root.handlers:
|
||||||
|
root_handler = ReentrantSafeStreamHandler(sys.stderr)
|
||||||
|
root_handler.setFormatter(_EmojiFormatter())
|
||||||
|
root.addHandler(root_handler)
|
||||||
|
|
||||||
|
# Ensure future dictConfig calls keep a root logger so unhandled loggers still
|
||||||
|
# get the emoji formatter rather than falling back to logging.lastResort.
|
||||||
|
LOGGING_CONFIG_DEFAULTS["root"] = {
|
||||||
|
"level": "WARNING",
|
||||||
|
"handlers": ["error_console"],
|
||||||
|
}
|
||||||
|
|
||||||
|
# Sanic's loggers already have their own handlers; stop them from bubbling up
|
||||||
|
# to the root handler we just added so messages are not duplicated.
|
||||||
|
for name in (
|
||||||
|
"sanic.root",
|
||||||
|
"sanic.error",
|
||||||
|
"sanic.access",
|
||||||
|
"sanic.server",
|
||||||
|
"sanic.websockets",
|
||||||
|
):
|
||||||
|
logging.getLogger(name).propagate = False
|
||||||
|
if name in LOGGING_CONFIG_DEFAULTS["loggers"]:
|
||||||
|
LOGGING_CONFIG_DEFAULTS["loggers"][name]["propagate"] = False
|
||||||
|
|
||||||
for handler_name in ("console", "error_console", "access_console"):
|
for handler_name in ("console", "error_console", "access_console"):
|
||||||
LOGGING_CONFIG_DEFAULTS["handlers"][handler_name]["class"] = (
|
LOGGING_CONFIG_DEFAULTS["handlers"][handler_name]["class"] = (
|
||||||
"cista.sanic_logging.ReentrantSafeStreamHandler"
|
"cista.sanic_logging.ReentrantSafeStreamHandler"
|
||||||
@@ -295,8 +328,7 @@ def configure_main_logging() -> None:
|
|||||||
LOGGING_CONFIG_DEFAULTS["loggers"]["sanic.websockets"]["level"] = "ERROR"
|
LOGGING_CONFIG_DEFAULTS["loggers"]["sanic.websockets"]["level"] = "ERROR"
|
||||||
logging.getLogger("sanic.websockets").setLevel(logging.ERROR)
|
logging.getLogger("sanic.websockets").setLevel(logging.ERROR)
|
||||||
# Preview worker timeouts are already annotated in the access log extra;
|
# Preview worker timeouts are already annotated in the access log extra;
|
||||||
# the pool's WARNING would otherwise fall to logging.lastResort, printing
|
# keep the pool's own warnings quiet so they are not logged twice.
|
||||||
# a bare message with no level prefix.
|
|
||||||
logging.getLogger("mediapreview.pool").setLevel(logging.ERROR)
|
logging.getLogger("mediapreview.pool").setLevel(logging.ERROR)
|
||||||
# Also reformat any handlers already attached (covers the initial Sanic() call)
|
# Also reformat any handlers already attached (covers the initial Sanic() call)
|
||||||
for name in ("sanic.root", "sanic.error", "sanic.server", "sanic.websockets"):
|
for name in ("sanic.root", "sanic.error", "sanic.server", "sanic.websockets"):
|
||||||
|
|||||||
@@ -10,8 +10,10 @@ Environment variables:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import hashlib
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
from time import time
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import websockets
|
import websockets
|
||||||
@@ -62,6 +64,40 @@ async def close_client():
|
|||||||
_client = None
|
_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(
|
async def validate_sso_request(
|
||||||
request, *, perm: str = "cista:login", renew: bool = True
|
request, *, perm: str = "cista:login", renew: bool = True
|
||||||
) -> dict | None:
|
) -> dict | None:
|
||||||
@@ -105,6 +141,17 @@ async def validate_sso_request(
|
|||||||
if not renew:
|
if not renew:
|
||||||
url += "&renew=0"
|
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:
|
try:
|
||||||
response = await client.post(
|
response = await client.post(
|
||||||
url,
|
url,
|
||||||
@@ -121,6 +168,8 @@ async def validate_sso_request(
|
|||||||
request.ctx.sso_user = {}
|
request.ctx.sso_user = {}
|
||||||
return {}
|
return {}
|
||||||
else:
|
else:
|
||||||
|
_cleanup_validate_cache()
|
||||||
|
_validate_cache[cache_key] = (time(), data)
|
||||||
return data
|
return data
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
+101
-2
@@ -1,14 +1,15 @@
|
|||||||
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
|
|
||||||
import msgspec
|
import msgspec
|
||||||
import websockets.exceptions
|
import websockets.exceptions
|
||||||
from sanic import errorpages
|
from sanic import errorpages
|
||||||
from sanic.exceptions import SanicException
|
from sanic.exceptions import SanicException, Unauthorized
|
||||||
from sanic.log import logger
|
from sanic.log import logger
|
||||||
from sanic.response import raw, redirect
|
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.protocol import ErrorMsg
|
||||||
from cista.sanic_logging import log_ws_close, log_ws_open
|
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)
|
log_ws_close(ws_id, close_code, duration, extra=close_extra)
|
||||||
|
|
||||||
return wrapper
|
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
|
||||||
|
|||||||
@@ -1,28 +0,0 @@
|
|||||||
services:
|
|
||||||
onlyoffice:
|
|
||||||
build:
|
|
||||||
context: ./mediapreview/mediapreview/docker
|
|
||||||
args:
|
|
||||||
ONLYOFFICE_VERSION: "9.3.1"
|
|
||||||
container_name: onlyoffice
|
|
||||||
ports:
|
|
||||||
- "8080:80"
|
|
||||||
environment:
|
|
||||||
# Number of converter workers (default 8).
|
|
||||||
# Set to your CPU count or slightly below.
|
|
||||||
- WORKERS
|
|
||||||
# JWT secret shared with Cista.
|
|
||||||
# OnlyOffice reads it as JWT_SECRET; Cista reads it as ONLYOFFICE_JWT_SECRET.
|
|
||||||
# We use ONLYOFFICE_JWT_SECRET as the canonical name so you only set one variable.
|
|
||||||
- JWT_SECRET=${ONLYOFFICE_JWT_SECRET}
|
|
||||||
- JWT_ENABLED=true
|
|
||||||
- JWT_HEADER=Authorization
|
|
||||||
volumes:
|
|
||||||
# Persist fonts and generated caches across restarts
|
|
||||||
- onlyoffice-data:/var/www/onlyoffice/Data
|
|
||||||
- onlyoffice-lib:/var/lib/onlyoffice
|
|
||||||
restart: unless-stopped
|
|
||||||
|
|
||||||
volumes:
|
|
||||||
onlyoffice-data:
|
|
||||||
onlyoffice-lib:
|
|
||||||
@@ -28,9 +28,9 @@
|
|||||||
</template>
|
</template>
|
||||||
<div v-if="!props.editorMode && showSortHints" class="sort-hints">
|
<div v-if="!props.editorMode && showSortHints" class="sort-hints">
|
||||||
<span class="sort-label">Order</span>
|
<span class="sort-label">Order</span>
|
||||||
<span class="keycap">1</span>
|
<button type="button" class="keycap" aria-label="Alphabetical order" @click="store.sort('name')">1</button>
|
||||||
<span class="keycap">2</span>
|
<button type="button" class="keycap" aria-label="Newest first" @click="store.sort('modified')">2</button>
|
||||||
<span class="keycap">3</span>
|
<button type="button" class="keycap" aria-label="Largest first" @click="store.sort('size')">3</button>
|
||||||
</div>
|
</div>
|
||||||
<SvgButton
|
<SvgButton
|
||||||
v-if="props.editorMode"
|
v-if="props.editorMode"
|
||||||
@@ -314,6 +314,14 @@ onUnmounted(() => {
|
|||||||
border-radius: 0.3em;
|
border-radius: 0.3em;
|
||||||
padding: 0 0.45em;
|
padding: 0 0.45em;
|
||||||
line-height: 1.4;
|
line-height: 1.4;
|
||||||
|
cursor: pointer;
|
||||||
|
transition: all 0.2s ease;
|
||||||
|
}
|
||||||
|
.keycap:hover,
|
||||||
|
.keycap:focus {
|
||||||
|
background: #e6e6e6;
|
||||||
|
border-color: #aaa;
|
||||||
|
transform: scale(1.05);
|
||||||
}
|
}
|
||||||
@media screen and (min-width: 800px) {
|
@media screen and (min-width: 800px) {
|
||||||
.sort-hints {
|
.sort-hints {
|
||||||
|
|||||||
@@ -79,6 +79,7 @@ export const useMainStore = defineStore('main', {
|
|||||||
connected: false,
|
connected: false,
|
||||||
authInProgress: false,
|
authInProgress: false,
|
||||||
cursor: '' as string,
|
cursor: '' as string,
|
||||||
|
lastSearchLoc: '' as string,
|
||||||
server: {} as Record<string, any> & {
|
server: {} as Record<string, any> & {
|
||||||
public?: boolean
|
public?: boolean
|
||||||
paskia?: boolean
|
paskia?: boolean
|
||||||
@@ -159,6 +160,10 @@ export const useMainStore = defineStore('main', {
|
|||||||
this.docVersion++
|
this.docVersion++
|
||||||
// Sync documents to search worker
|
// Sync documents to search worker
|
||||||
this.syncSearchWorker()
|
this.syncSearchWorker()
|
||||||
|
// Re-run the current search against the updated file list
|
||||||
|
if (this.query) {
|
||||||
|
this.search(this.query, this.lastSearchLoc)
|
||||||
|
}
|
||||||
},
|
},
|
||||||
/** Patch aspect ratios on existing docs from a server ar update message */
|
/** Patch aspect ratios on existing docs from a server ar update message */
|
||||||
updateAr(arMap: Record<string, number>) {
|
updateAr(arMap: Record<string, number>) {
|
||||||
@@ -268,6 +273,7 @@ export const useMainStore = defineStore('main', {
|
|||||||
|
|
||||||
// Update query immediately so watchers know we're handling this
|
// Update query immediately so watchers know we're handling this
|
||||||
this.query = query
|
this.query = query
|
||||||
|
this.lastSearchLoc = loc
|
||||||
|
|
||||||
// Cancel pending timers
|
// Cancel pending timers
|
||||||
if (loadingTimer) {
|
if (loadingTimer) {
|
||||||
|
|||||||
+1
-1
@@ -33,7 +33,7 @@ dependencies = [
|
|||||||
"html5tagger>=1.3.0",
|
"html5tagger>=1.3.0",
|
||||||
"httpx>=0.28.0",
|
"httpx>=0.28.0",
|
||||||
"inotify>=0.2.12",
|
"inotify>=0.2.12",
|
||||||
"mediapreview[standard]",
|
"mediapreview[standard]>=0.2.3",
|
||||||
"msgspec>=0.19.0",
|
"msgspec>=0.19.0",
|
||||||
"natsort>=8.4.0",
|
"natsort>=8.4.0",
|
||||||
"numpy>=2.3.2",
|
"numpy>=2.3.2",
|
||||||
|
|||||||
@@ -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