Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5c4965a86b | ||
|
|
bd7291e9ef | ||
|
|
2fa52229cc | ||
|
|
2dea459d8f | ||
|
|
922069c603 | ||
|
|
5c7c7343ad | ||
|
|
593d16d8c5 | ||
|
|
20d8d317fa | ||
|
|
8d89c397a4 | ||
|
|
1bec73f4cd | ||
|
|
4dd1d4c7e6 | ||
|
|
84ef91a360 | ||
|
|
421d90e9c5 | ||
|
|
8b4e622aef | ||
|
|
13f32c57ab | ||
|
|
25a2a5f20c |
+35
-31
@@ -5,7 +5,8 @@ from pathlib import Path
|
||||
from docopt import docopt
|
||||
|
||||
import cista
|
||||
from cista import app, config, droppy, serve, server80
|
||||
from cista import app, config, droppy, onlyoffice, serve, server80
|
||||
from cista.sso import PASKIA_BACKEND_URL
|
||||
from cista.util import pwgen
|
||||
|
||||
del app, server80.app # Only import needed, for Sanic multiprocessing
|
||||
@@ -30,7 +31,7 @@ def create_startup_box(
|
||||
):
|
||||
"""Create a framed startup box with server information."""
|
||||
title = f"Cista {cista.__version__}"
|
||||
listen = unix if unix else url
|
||||
listen = unix or url
|
||||
location = f"{folder} @ {listen}"
|
||||
lines = [title, location]
|
||||
# Auth line: Paskia <url> or Password, with optional Public suffix
|
||||
@@ -53,40 +54,39 @@ def create_startup_box(
|
||||
|
||||
banner = create_banner()
|
||||
|
||||
doc = """\
|
||||
_default_confdir = (
|
||||
(Path(os.environ["XDG_CONFIG_HOME"]) / "cista").as_posix()
|
||||
if os.environ.get("XDG_CONFIG_HOME")
|
||||
else (Path.home() / ".config/cista").as_posix()
|
||||
)
|
||||
|
||||
doc = f"""\
|
||||
Usage:
|
||||
cista [-c <confdir>] [-l <host>] [--import-droppy] [--dev] [<path>]
|
||||
cista [-c <confdir>] --user <name> [--privileged] [--password]
|
||||
cista [-c <confdir>] --oosetup
|
||||
cista --version
|
||||
|
||||
Options:
|
||||
-c CONFDIR Custom config directory
|
||||
-l, --listen LISTEN-ADDR
|
||||
Listen on
|
||||
:8989 (localhost port, plain http)
|
||||
<addr>:3000 (bind another address, port)
|
||||
/path/to/unix.sock (unix socket)
|
||||
example.com (run on 80 and 443 with LetsEncrypt)
|
||||
--import-droppy Import Droppy config from ~/.droppy/config
|
||||
--dev Developer mode (reloads, friendlier crashes, more logs)
|
||||
|
||||
Listen address and path are preserved in config,
|
||||
and only config dir and dev mode need to be specified on subsequent runs.
|
||||
|
||||
User management:
|
||||
--user NAME Create or modify user
|
||||
--privileged Give the user full admin rights
|
||||
--password Reset password
|
||||
-c CONFDIR Config directory [{_default_confdir}]
|
||||
-l, --listen ADDR Listen on address (port, :port, /socket or domain for https)
|
||||
--import-droppy Import Droppy config from ~/.droppy/config
|
||||
--dev Developer mode (reloads, friendlier crashes, more logs)
|
||||
--user NAME Create or modify a user account (when server is not running)
|
||||
--privileged Grant admin rights
|
||||
--password Reset password
|
||||
--oosetup Build and run OnlyOffice in Docker for document previews
|
||||
|
||||
Environment:
|
||||
PASKIA_BACKEND_URL Paskia single sign-on (e.g. http://localhost:4401)
|
||||
https://git.zi.fi/leovasanko/paskia
|
||||
PASKIA_BACKEND_URL Paskia single sign-on (e.g. http://localhost:4401)
|
||||
https://git.zi.fi/leovasanko/paskia
|
||||
ONLYOFFICE_CISTA_URL, ONLYOFFICE_JWT_SECRET, ONLYOFFICE_CALLBACK_HOST (if needed)
|
||||
"""
|
||||
|
||||
first_time_help = """\
|
||||
No config file found! Get started with:
|
||||
cista --user yourname --privileged # If you want user accounts
|
||||
cista -l :8989 /path/to/files # Run the server on localhost:8989
|
||||
cista --user yourname --privileged # If you want user accounts
|
||||
cista -l :8989 /path/to/files # Run the server on localhost:8989
|
||||
|
||||
See cista --help for other options!
|
||||
"""
|
||||
@@ -115,6 +115,8 @@ def _main():
|
||||
args = docopt(doc)
|
||||
if args["--user"]:
|
||||
return _user(args)
|
||||
if args["--oosetup"]:
|
||||
return onlyoffice.setup_docker(_resolve_confdir(args))
|
||||
listen = args["--listen"]
|
||||
# Validate arguments first
|
||||
if args["<path>"]:
|
||||
@@ -153,9 +155,6 @@ def _main():
|
||||
if not config.config.path.is_dir():
|
||||
raise ValueError(f"No such directory: {config.config.path}")
|
||||
dev = args["--dev"]
|
||||
# Check for Paskia SSO
|
||||
from cista.sso import PASKIA_BACKEND_URL
|
||||
|
||||
# Print startup box
|
||||
startup_box = create_startup_box(
|
||||
folder=config.config.path,
|
||||
@@ -171,17 +170,22 @@ def _main():
|
||||
return 0
|
||||
|
||||
|
||||
def _confdir(args):
|
||||
def _resolve_confdir(args):
|
||||
confdir = None
|
||||
if args["-c"]:
|
||||
# Custom config directory
|
||||
confdir = Path(args["-c"]).resolve()
|
||||
if confdir.exists() and not confdir.is_dir():
|
||||
if confdir.name != config.conffile.name:
|
||||
if confdir.name != "db.toml":
|
||||
raise ValueError("Config path is not a directory")
|
||||
# Accidentally pointed to the db.toml, use parent
|
||||
confdir = confdir.parent
|
||||
os.environ["CISTA_HOME"] = confdir.as_posix()
|
||||
config.init_confdir() # Uses environ if available
|
||||
return confdir
|
||||
|
||||
|
||||
def _confdir(args):
|
||||
confdir = _resolve_confdir(args)
|
||||
config.init_confdir(confdir)
|
||||
|
||||
|
||||
def _user(args):
|
||||
|
||||
+9
-9
@@ -6,7 +6,7 @@ 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 import __version__, auth, config, onlyoffice, sharefs, sso, watching
|
||||
from cista.auth import (
|
||||
create_share_token_handler,
|
||||
create_token_handler,
|
||||
@@ -22,11 +22,13 @@ fileserver = FileServer()
|
||||
|
||||
@bp.before_server_start
|
||||
async def start_fileserver(app):
|
||||
_ = app
|
||||
await fileserver.start()
|
||||
|
||||
|
||||
@bp.after_server_stop
|
||||
async def stop_fileserver(app):
|
||||
_ = app
|
||||
await fileserver.stop()
|
||||
|
||||
|
||||
@@ -63,6 +65,7 @@ async def watch(req, ws):
|
||||
"version": __version__,
|
||||
"public": config.config.public,
|
||||
"paskia": sso.paskia_enabled(),
|
||||
"office_previews": await onlyoffice.is_available_cached(),
|
||||
},
|
||||
"user": user_info,
|
||||
}
|
||||
@@ -99,6 +102,7 @@ async def watch(req, ws):
|
||||
|
||||
|
||||
def subscribe(uuid, ws):
|
||||
_ = ws
|
||||
with watching.state.lock:
|
||||
q = watching.pubsub[uuid] = asyncio.Queue()
|
||||
# Init with disk usage and full tree
|
||||
@@ -125,12 +129,10 @@ async def update_public(request):
|
||||
await auth.verify(request, privileged=True)
|
||||
try:
|
||||
public = request.json["public"]
|
||||
if not isinstance(public, bool):
|
||||
raise ValueError("public must be a boolean")
|
||||
except KeyError:
|
||||
raise BadRequest("Missing public field") from None
|
||||
except ValueError as e:
|
||||
raise BadRequest(str(e)) from None
|
||||
if not isinstance(public, bool):
|
||||
raise BadRequest("public must be a boolean")
|
||||
config.update_config({"public": public})
|
||||
return json({"message": "Public access setting updated", "public": public})
|
||||
|
||||
@@ -140,12 +142,10 @@ async def update_name(request):
|
||||
await auth.verify(request, privileged=True)
|
||||
try:
|
||||
name = request.json["name"]
|
||||
if not isinstance(name, str):
|
||||
raise ValueError("name must be a string")
|
||||
except KeyError:
|
||||
raise BadRequest("Missing name field") from None
|
||||
except ValueError as e:
|
||||
raise BadRequest(str(e)) from None
|
||||
if not isinstance(name, str):
|
||||
raise BadRequest("name must be a string")
|
||||
config.update_config({"name": name})
|
||||
# Return the effective name (fallback to path.name if empty)
|
||||
effective_name = name or config.config.path.name
|
||||
|
||||
+27
-6
@@ -8,6 +8,7 @@ from stat import S_IFDIR, S_IFREG
|
||||
from urllib.parse import unquote
|
||||
from wsgiref.handlers import format_date_time
|
||||
|
||||
import tracerite
|
||||
from blake3 import blake3
|
||||
from sanic import Sanic, empty, raw, redirect
|
||||
from sanic.exceptions import Forbidden, NotFound
|
||||
@@ -16,7 +17,17 @@ from setproctitle import setproctitle
|
||||
from stream_zip import ZIP_AUTO, stream_zip
|
||||
from zstandard import ZstdCompressor
|
||||
|
||||
from cista import auth, config, fileserver, preview, session, sharefs, sso, watching
|
||||
from cista import (
|
||||
auth,
|
||||
config,
|
||||
fileserver,
|
||||
onlyoffice,
|
||||
preview,
|
||||
session,
|
||||
sharefs,
|
||||
sso,
|
||||
watching,
|
||||
)
|
||||
from cista.api import bp
|
||||
from cista.preview import shutdown_preview_workers, start_preview_workers
|
||||
from cista.sanic_logging import (
|
||||
@@ -27,8 +38,10 @@ from cista.sanic_logging import (
|
||||
from cista.sanic_logging import logger as access_logger
|
||||
from cista.util.apphelpers import handle_sanic_exception
|
||||
|
||||
tracerite.load()
|
||||
configure_access_logging()
|
||||
|
||||
|
||||
app = Sanic("cista", strict_slashes=True)
|
||||
app.router.ALLOWED_METHODS = (
|
||||
*app.router.ALLOWED_METHODS,
|
||||
@@ -43,8 +56,8 @@ configure_main_logging()
|
||||
|
||||
@app.on_request
|
||||
async def use_session(req):
|
||||
req.ctx._log_start = time.perf_counter()
|
||||
req.ctx._auth_flow = ["session: start"]
|
||||
req.ctx.log_start = time.perf_counter()
|
||||
req.ctx.auth_flow = ["session: start"]
|
||||
auth.hydrate_request_auth_context(req, source="app.on_request")
|
||||
# CSRF protection
|
||||
if req.method == "GET" and req.headers.upgrade != "websocket":
|
||||
@@ -61,7 +74,7 @@ async def log_access(req, res):
|
||||
"""Log HTTP access in a clean single-line format."""
|
||||
if req.headers.get("upgrade", "").lower() == "websocket":
|
||||
return res
|
||||
start = getattr(req.ctx, "_log_start", None)
|
||||
start = getattr(req.ctx, "log_start", None)
|
||||
duration_ms = (time.perf_counter() - start) * 1000 if start is not None else 0.0
|
||||
client = req.client_ip or "-"
|
||||
host = req.host or "-"
|
||||
@@ -71,7 +84,7 @@ async def log_access(req, res):
|
||||
if isinstance(qs, bytes):
|
||||
qs = qs.decode(errors="replace")
|
||||
path = f"{path}?{qs}"
|
||||
extra = getattr(req.ctx, "_log_extra", None)
|
||||
extra = getattr(req.ctx, "log_extra", None)
|
||||
line = format_access_log(
|
||||
client, res.status, req.method, host, path, duration_ms, extra=extra
|
||||
)
|
||||
@@ -90,7 +103,7 @@ async def forward_sso_cookies(req, res):
|
||||
@app.on_response
|
||||
async def persist_auth_session(req, res):
|
||||
"""Persist a session cookie after successful Authorization-based auth."""
|
||||
username = getattr(req.ctx, "_create_session_username", None)
|
||||
username = getattr(req.ctx, "create_session_username", None)
|
||||
if not username or res.status >= 400:
|
||||
return
|
||||
existing = getattr(req.ctx, "session", None)
|
||||
@@ -126,10 +139,17 @@ async def main_start(app):
|
||||
watching.start(app)
|
||||
|
||||
|
||||
@app.after_server_start
|
||||
async def main_after_start(app):
|
||||
_ = app
|
||||
onlyoffice.log_reachable_info()
|
||||
|
||||
|
||||
# Sanic sometimes fails to execute after_server_stop, so we do it before instead (potentially interrupting handlers)
|
||||
@app.before_server_stop
|
||||
async def main_stop(app):
|
||||
watching.stop(app)
|
||||
await onlyoffice.close_oo_client()
|
||||
await shutdown_preview_workers()
|
||||
app.ctx.threadexec.shutdown()
|
||||
app.ctx.zipexec.shutdown(cancel_futures=True)
|
||||
@@ -236,6 +256,7 @@ async def wwwroot(req, path=""):
|
||||
|
||||
@app.route("/favicon.ico", methods=["GET", "HEAD"])
|
||||
async def favicon(req):
|
||||
_ = req
|
||||
# Browsers keep asking for it when viewing files (not HTML with icon link)
|
||||
return redirect("/assets/logo-ctv8tVwU.svg", status=308)
|
||||
|
||||
|
||||
+39
-69
@@ -2,22 +2,21 @@ import base64
|
||||
import binascii
|
||||
import hashlib
|
||||
import hmac
|
||||
import re
|
||||
import secrets
|
||||
import struct
|
||||
from pathlib import PurePosixPath
|
||||
from time import time
|
||||
from unicodedata import normalize
|
||||
|
||||
import argon2
|
||||
import msgspec
|
||||
from Crypto.Hash import MD4
|
||||
from html5tagger import Document
|
||||
from sanic import Blueprint, html, json, redirect
|
||||
from sanic.exceptions import BadRequest, Forbidden, Unauthorized
|
||||
from sanic.log import logger
|
||||
|
||||
from cista import config, session, sharefs
|
||||
from cista.util import pwgen
|
||||
from cista import sso as _sso_module
|
||||
from cista.util import pwgen, pwhash
|
||||
from cista.util.filename import sanitize
|
||||
|
||||
_LOGIN_PAGE_CSS = """\
|
||||
@@ -175,16 +174,8 @@ form.onsubmit = async (e) => {
|
||||
};
|
||||
"""
|
||||
|
||||
# Import for SSO validation (lazily loaded to avoid circular imports)
|
||||
_sso_module = None
|
||||
|
||||
|
||||
def _get_sso():
|
||||
global _sso_module
|
||||
if _sso_module is None:
|
||||
from cista import sso
|
||||
|
||||
_sso_module = sso
|
||||
return _sso_module
|
||||
|
||||
|
||||
@@ -202,13 +193,13 @@ def _set_auth_failure_log(request, auth_flow: list[str]) -> None:
|
||||
value = request.headers.get(header)
|
||||
if value:
|
||||
parts.append(f"{label}={value}")
|
||||
request.ctx._log_extra = " | ".join(parts)
|
||||
request.ctx.log_extra = " | ".join(parts)
|
||||
|
||||
|
||||
def hydrate_request_auth_context(request, *, source: str) -> None:
|
||||
auth_flow = getattr(request.ctx, "_auth_flow", None)
|
||||
auth_flow = getattr(request.ctx, "auth_flow", None)
|
||||
if auth_flow is None:
|
||||
auth_flow = request.ctx._auth_flow = []
|
||||
auth_flow = request.ctx.auth_flow = []
|
||||
|
||||
if hasattr(request.ctx, "session"):
|
||||
# Already hydrated by an earlier caller (e.g., use_session middleware)
|
||||
@@ -234,9 +225,6 @@ def hydrate_request_auth_context(request, *, source: str) -> None:
|
||||
auth_flow.append(f"session:{source}(bad-jwt)")
|
||||
|
||||
|
||||
_argon = argon2.PasswordHasher()
|
||||
_droppyhash = re.compile(r"^([a-f0-9]{64})\$([a-f0-9]{8})$")
|
||||
|
||||
_AUTH_REALM = "cista"
|
||||
_AUTH_CACHE_TTL = 10
|
||||
_auth_cache: dict[str, tuple[float, config.User]] = {}
|
||||
@@ -280,6 +268,7 @@ def _log_webdav_user_agent_once(request, user_agent: str):
|
||||
|
||||
|
||||
def _build_ua_auth_headers(request, *, include_hint=False) -> dict[str, str]:
|
||||
_ = include_hint
|
||||
user_agent = request.headers.get("user-agent", "")
|
||||
_log_webdav_user_agent_once(request, user_agent)
|
||||
if _is_windows_auth_client(user_agent):
|
||||
@@ -448,12 +437,6 @@ def _ntlmv2_verify(
|
||||
nt_response: bytes,
|
||||
) -> bool:
|
||||
"""Verify an NTLMv2 response using the plaintext token secret as the password."""
|
||||
try:
|
||||
from Crypto.Hash import MD4
|
||||
except ImportError:
|
||||
logger.error("pycryptodome MD4 not available, cannot verify NTLM")
|
||||
return False
|
||||
|
||||
if len(nt_response) < 16:
|
||||
return False
|
||||
|
||||
@@ -517,47 +500,30 @@ def _ntlmv2_verify(
|
||||
return False
|
||||
|
||||
|
||||
def _pwnorm(password):
|
||||
return normalize("NFC", password).strip().encode()
|
||||
|
||||
|
||||
def _cache_key(username: str, password: str) -> str:
|
||||
return hashlib.sha256(f"{username}\x00{password}".encode()).hexdigest()
|
||||
|
||||
|
||||
def login(username: str, password: str):
|
||||
normalized_username = pwhash.normalize_secret(username).decode()
|
||||
cache_key = _cache_key(username, password)
|
||||
cached = _auth_cache.get(cache_key)
|
||||
if cached:
|
||||
ts, user = cached
|
||||
if time() - ts < _AUTH_CACHE_TTL:
|
||||
return user
|
||||
current = config.config.users.get(normalized_username)
|
||||
if current and current.hash == user.hash:
|
||||
return current
|
||||
del _auth_cache[cache_key]
|
||||
|
||||
un = _pwnorm(username)
|
||||
pw = _pwnorm(password)
|
||||
try:
|
||||
u = config.config.users[un.decode()]
|
||||
u = config.config.users[normalized_username]
|
||||
except KeyError:
|
||||
raise ValueError("Invalid username") from None
|
||||
# Verify password
|
||||
need_rehash = False
|
||||
if not u.hash:
|
||||
raise ValueError("Account disabled")
|
||||
if (m := _droppyhash.match(u.hash)) is not None:
|
||||
h, s = m.groups()
|
||||
h2 = hmac.digest(pw + s.encode() + un, b"", "sha256").hex()
|
||||
if not hmac.compare_digest(h, h2):
|
||||
raise ValueError("Invalid password")
|
||||
# Droppy hashes are weak, do a hash update
|
||||
need_rehash = True
|
||||
else:
|
||||
try:
|
||||
_argon.verify(u.hash, pw)
|
||||
except Exception:
|
||||
raise ValueError("Invalid password") from None
|
||||
if _argon.check_needs_rehash(u.hash):
|
||||
need_rehash = True
|
||||
need_rehash = pwhash.verify_hash(
|
||||
u.hash, username=normalized_username, password=password
|
||||
)
|
||||
# Login successful
|
||||
if need_rehash:
|
||||
set_password(u, password)
|
||||
@@ -568,7 +534,7 @@ def login(username: str, password: str):
|
||||
|
||||
|
||||
def set_password(user: config.User, password: str):
|
||||
user.hash = _argon.hash(_pwnorm(password))
|
||||
pwhash.set_password(user, password)
|
||||
_auth_cache.clear()
|
||||
|
||||
|
||||
@@ -670,11 +636,12 @@ async def _token_auth_login(request, *, privileged=False):
|
||||
ctx = data.get("ctx", {}) if isinstance(data, dict) else {}
|
||||
user_info = ctx.get("user", {}) if isinstance(ctx, dict) else {}
|
||||
request.ctx.username = user_info.get("display_name", "")
|
||||
return True
|
||||
except Forbidden:
|
||||
raise
|
||||
except Exception:
|
||||
return False
|
||||
else:
|
||||
return True
|
||||
|
||||
if token.username:
|
||||
user = config.config.users.get(token.username)
|
||||
@@ -840,12 +807,13 @@ async def _ntlm_auth_login(request, *, privileged=False):
|
||||
token.sso_user_id,
|
||||
tid[:8],
|
||||
)
|
||||
return True
|
||||
except Forbidden:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.warning("NTLM SSO check failed: %s", e)
|
||||
continue
|
||||
else:
|
||||
return True
|
||||
|
||||
if token.username:
|
||||
user = config.config.users.get(token.username)
|
||||
@@ -865,7 +833,7 @@ async def _ntlm_auth_login(request, *, privileged=False):
|
||||
request.ctx.user = user
|
||||
request.ctx.auth_token_id = tid
|
||||
request.ctx.auth_token = token
|
||||
request.ctx._create_session_username = token.username
|
||||
request.ctx.create_session_username = token.username
|
||||
logger.debug(
|
||||
"NTLM auth success for local user %s (token=%s...)",
|
||||
token.username,
|
||||
@@ -914,7 +882,7 @@ async def verify(request, *, privileged=False):
|
||||
scheme = auth_header.split()[0].lower() if has_auth_header else None
|
||||
|
||||
# Concise auth flow for diagnostics (populated by use_session + verify)
|
||||
auth_flow = list(getattr(request.ctx, "_auth_flow", ["session:skipped"]))
|
||||
auth_flow = list(getattr(request.ctx, "auth_flow", ["session:skipped"]))
|
||||
tried: list[str] = []
|
||||
|
||||
sso = _get_sso()
|
||||
@@ -927,7 +895,6 @@ async def verify(request, *, privileged=False):
|
||||
try:
|
||||
perm = "cista:admin" if privileged else "cista:login"
|
||||
await sso.validate_sso_request(request, perm=perm)
|
||||
return
|
||||
except Unauthorized as e:
|
||||
auth_flow.append(f"tried={','.join(tried)} result=failed")
|
||||
_set_auth_failure_log(request, auth_flow)
|
||||
@@ -936,6 +903,8 @@ async def verify(request, *, privileged=False):
|
||||
headers=_build_ua_auth_headers(request),
|
||||
quiet=True,
|
||||
) from e
|
||||
else:
|
||||
return
|
||||
tried.append("sso")
|
||||
perm = "cista:admin" if privileged else "cista:login"
|
||||
await sso.validate_sso_request(request, perm=perm)
|
||||
@@ -986,10 +955,10 @@ async def verify(request, *, privileged=False):
|
||||
user = None
|
||||
else:
|
||||
if user is not None:
|
||||
if getattr(request.ctx, "_create_session_username", None) is None:
|
||||
if getattr(request.ctx, "create_session_username", None) is None:
|
||||
username = getattr(request.ctx, "username", None)
|
||||
if username:
|
||||
request.ctx._create_session_username = username
|
||||
request.ctx.create_session_username = username
|
||||
return
|
||||
# Auth header present but invalid → try session fallback
|
||||
tried.append("session")
|
||||
@@ -1113,6 +1082,7 @@ async def login_page(request):
|
||||
|
||||
def _login_success_page(username: str) -> str:
|
||||
"""Minimal page that signals auth-success to parent iframe."""
|
||||
_ = username
|
||||
return str(
|
||||
Document().script_("window.parent.postMessage({type:'auth-success'},'*')")
|
||||
)
|
||||
@@ -1127,13 +1097,16 @@ async def login_post(request):
|
||||
else:
|
||||
username = request.form["username"][0]
|
||||
password = request.form["password"][0]
|
||||
if not username or not password:
|
||||
raise KeyError
|
||||
except KeyError:
|
||||
raise BadRequest(
|
||||
"Missing username or password",
|
||||
context={"redirect": "/login"},
|
||||
) from None
|
||||
if not username or not password:
|
||||
raise BadRequest(
|
||||
"Missing username or password",
|
||||
context={"redirect": "/login"},
|
||||
)
|
||||
try:
|
||||
user = login(username, password)
|
||||
except ValueError as e:
|
||||
@@ -1172,12 +1145,12 @@ async def change_password(request):
|
||||
username = request.form["username"][0]
|
||||
pwchange = request.form["passwordChange"][0]
|
||||
password = request.form["password"][0]
|
||||
if not username or not password:
|
||||
raise KeyError
|
||||
except KeyError:
|
||||
raise BadRequest(
|
||||
"Missing username, passwordChange or password",
|
||||
) from None
|
||||
if not username or not password:
|
||||
raise BadRequest("Missing username, passwordChange or password")
|
||||
try:
|
||||
user = login(username, password)
|
||||
set_password(user, pwchange)
|
||||
@@ -1220,16 +1193,15 @@ async def create_user(request):
|
||||
username = request.form["username"][0]
|
||||
password = request.form.get("password", [None])[0]
|
||||
privileged = request.form.get("privileged", ["false"])[0].lower() == "true"
|
||||
if not username or not username.isidentifier():
|
||||
raise ValueError("Invalid username")
|
||||
except (KeyError, ValueError) as e:
|
||||
raise BadRequest(str(e)) from e
|
||||
except KeyError as e:
|
||||
raise BadRequest("Missing fields") from e
|
||||
if not username or not username.isidentifier():
|
||||
raise BadRequest("Invalid username")
|
||||
if username in config.config.users:
|
||||
raise BadRequest("User already exists")
|
||||
if not password:
|
||||
password = pwgen.generate()
|
||||
changes = {"privileged": privileged}
|
||||
changes["hash"] = _argon.hash(_pwnorm(password))
|
||||
changes = {"privileged": privileged, "password": password}
|
||||
try:
|
||||
config.update_user(username, changes)
|
||||
except Exception as e:
|
||||
@@ -1256,8 +1228,6 @@ async def update_user(request, username):
|
||||
if changes["password"] == "":
|
||||
changes["password"] = pwgen.generate()
|
||||
password_response = changes["password"]
|
||||
changes["hash"] = _argon.hash(_pwnorm(changes["password"]))
|
||||
del changes["password"]
|
||||
if not changes:
|
||||
return json({"message": "No changes"})
|
||||
try:
|
||||
|
||||
+6
-6
@@ -14,6 +14,8 @@ from typing import Concatenate, Literal, ParamSpec
|
||||
import msgspec
|
||||
import msgspec.toml
|
||||
|
||||
from .util import pwhash
|
||||
|
||||
|
||||
class Config(msgspec.Struct):
|
||||
path: Path
|
||||
@@ -61,10 +63,10 @@ config: Config
|
||||
conffile: Path
|
||||
|
||||
|
||||
def init_confdir() -> None:
|
||||
def init_confdir(confdir: Path | str | None = None) -> None:
|
||||
global conffile
|
||||
if p := os.environ.get("CISTA_HOME"):
|
||||
home = Path(p)
|
||||
if confdir is not None:
|
||||
home = Path(confdir).expanduser()
|
||||
else:
|
||||
xdg = os.environ.get("XDG_CONFIG_HOME")
|
||||
home = (
|
||||
@@ -199,9 +201,7 @@ def update_user(conf: Config, name: str, changes: dict) -> Config:
|
||||
except KeyError:
|
||||
u = User()
|
||||
if "password" in changes:
|
||||
from . import auth
|
||||
|
||||
auth.set_password(u, changes["password"])
|
||||
pwhash.set_password(u, changes["password"])
|
||||
del changes["password"]
|
||||
udict = msgspec.to_builtins(u, enc_hook=enc_hook)
|
||||
udict.update(changes)
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
# Patched OnlyOffice Document Server with configurable converter worker count.
|
||||
#
|
||||
# The Community Edition hardcodes the document converter to 1 worker,
|
||||
# which creates a severe bottleneck under concurrent load.
|
||||
# This image patches the open-source license.js to spawn a configurable
|
||||
# number of converter workers (default 8).
|
||||
#
|
||||
# Build:
|
||||
# docker build -t onlyoffice-cista docker/onlyoffice-converter-patch
|
||||
#
|
||||
# Run:
|
||||
# docker run -d -p 8988:80 \
|
||||
# -e WORKERS=16 \
|
||||
# -e JWT_SECRET=your-strong-secret \
|
||||
# --name onlyoffice onlyoffice-cista
|
||||
#
|
||||
# JWT:
|
||||
# Set JWT_SECRET to the same value you pass to Cista as ONLYOFFICE_JWT_SECRET.
|
||||
# OnlyOffice will enable token validation automatically.
|
||||
#
|
||||
# The ONLYOFFICE_VERSION build arg lets you target a specific release.
|
||||
|
||||
ARG ONLYOFFICE_VERSION=9.3.1
|
||||
|
||||
FROM onlyoffice/documentserver:${ONLYOFFICE_VERSION}
|
||||
|
||||
# Prevent interactive apt prompts
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Install Node.js, npm, and git so we can run the FileConverter from source.
|
||||
RUN apt-get update -qq && \
|
||||
apt-get install -y -qq --no-install-recommends \
|
||||
nodejs \
|
||||
npm \
|
||||
git \
|
||||
ca-certificates && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Clone the open-source server components (shallow, ~15 MB).
|
||||
# The master branch is used because the Linux/web tags are not published
|
||||
# in the server repo; the license.js file has been stable for years.
|
||||
RUN git clone --depth 1 https://github.com/ONLYOFFICE/server.git /opt/oo-server
|
||||
|
||||
# Patch license.js so the converter worker count is read from an env var
|
||||
# instead of being hardcoded to 1.
|
||||
RUN sed -i \
|
||||
's/count: 1,/count: parseInt(process.env.WORKERS, 10) || 8,/' \
|
||||
/opt/oo-server/Common/sources/license.js
|
||||
|
||||
# Install npm dependencies for the modules the FileConverter touches.
|
||||
# DocService deps are also needed because converter.js pulls in baseConnector.
|
||||
RUN cd /opt/oo-server/Common && npm ci --no-audit --no-fund
|
||||
RUN cd /opt/oo-server/FileConverter && npm ci --no-audit --no-fund
|
||||
RUN cd /opt/oo-server/DocService && npm ci --no-audit --no-fund
|
||||
|
||||
# Back up the compiled pkg binary and replace it with our wrapper.
|
||||
RUN mv /var/www/onlyoffice/documentserver/server/FileConverter/converter \
|
||||
/var/www/onlyoffice/documentserver/server/FileConverter/converter.orig
|
||||
|
||||
COPY converter-wrapper.sh /var/www/onlyoffice/documentserver/server/FileConverter/converter
|
||||
RUN chmod +x /var/www/onlyoffice/documentserver/server/FileConverter/converter
|
||||
|
||||
# Default worker count (override at runtime with -e WORKERS=16).
|
||||
ENV WORKERS=8
|
||||
|
||||
# Use our custom entrypoint to persist the env var to a file that the
|
||||
# non-root converter process (user=ds) can read.
|
||||
COPY entrypoint.sh /app/ds/run-document-server-patched.sh
|
||||
RUN chmod +x /app/ds/run-document-server-patched.sh
|
||||
ENTRYPOINT ["/app/ds/run-document-server-patched.sh"]
|
||||
@@ -0,0 +1,19 @@
|
||||
#!/bin/bash
|
||||
# Wrapper that runs the OnlyOffice FileConverter from patched Node.js source.
|
||||
# Replaces the compiled pkg binary shipped with the Community Edition.
|
||||
|
||||
# The env var is not passed through supervisor to the 'ds' user, so we read
|
||||
# it from a file written by the custom entrypoint.
|
||||
if [ -z "${WORKERS}" ] && [ -r /tmp/oo-converter-workers.txt ]; then
|
||||
export WORKERS=$(cat /tmp/oo-converter-workers.txt)
|
||||
fi
|
||||
|
||||
cd /opt/oo-server/FileConverter || exit 1
|
||||
|
||||
export NODE_ENV=production-linux
|
||||
export NODE_CONFIG_DIR=/etc/onlyoffice/documentserver
|
||||
export NODE_DISABLE_COLORS=1
|
||||
export APPLICATION_NAME=onlyoffice
|
||||
export LD_LIBRARY_PATH=/var/www/onlyoffice/documentserver/server/FileConverter/bin
|
||||
|
||||
exec node sources/convertermaster.js "$@"
|
||||
@@ -0,0 +1,8 @@
|
||||
#!/bin/bash
|
||||
# Custom entrypoint that persists WORKERS to a file readable by
|
||||
# the non-root user that supervisor uses to run the converter.
|
||||
|
||||
echo "${WORKERS:-8}" > /tmp/oo-converter-workers.txt
|
||||
chmod 644 /tmp/oo-converter-workers.txt
|
||||
|
||||
exec /app/ds/run-document-server.sh "$@"
|
||||
+6
-24
@@ -76,7 +76,7 @@ async def upload_file_chunk(request, name):
|
||||
size_after = upload_info.get("size_after")
|
||||
if size_before is not None and size_after is not None and size_before != size_after:
|
||||
extras.append("resized")
|
||||
request.ctx._log_extra = " ".join(extras) if extras else None
|
||||
request.ctx.log_extra = " ".join(extras) if extras else None
|
||||
real_rel = PurePosixPath(path.relative_to(config.config.path.resolve()).as_posix())
|
||||
watching.notify_change(real_rel, *real_rel.parents)
|
||||
return json(
|
||||
@@ -197,38 +197,18 @@ async def copy_or_move(request, name=""):
|
||||
|
||||
def _apply():
|
||||
for op_name, op_keys in (("cp", cp_keys), ("mv", mv_keys)):
|
||||
op_multi = len(op_keys) > 1
|
||||
for key in op_keys:
|
||||
try:
|
||||
src_rel = key_paths[key]
|
||||
src_abs = _resolve_from_relpath(src_rel, request=request)
|
||||
|
||||
if op_multi:
|
||||
if not dst_is_dir:
|
||||
raise BadRequest(
|
||||
"Destination must be an existing directory for multiple keys"
|
||||
)
|
||||
dst_item_rel = (
|
||||
dst_rel / src_rel.name
|
||||
if dst_rel.parts
|
||||
else PurePosixPath(src_rel.name)
|
||||
)
|
||||
elif dst_is_dir:
|
||||
if dst_is_dir:
|
||||
dst_item_rel = (
|
||||
dst_rel / src_rel.name
|
||||
if dst_rel.parts
|
||||
else PurePosixPath(src_rel.name)
|
||||
)
|
||||
else:
|
||||
if not dst_rel.parts:
|
||||
raise BadRequest("Destination file path is required")
|
||||
parent_abs = dst_abs.parent
|
||||
if not parent_abs.is_dir():
|
||||
raise BadRequest("Destination parent folder does not exist")
|
||||
if src_abs.is_dir() and dst_exists and dst_abs.is_file():
|
||||
raise BadRequest(
|
||||
"Cannot move/copy a directory to an existing file"
|
||||
)
|
||||
dst_item_rel = dst_rel
|
||||
|
||||
dst_item_abs = _resolve_from_relpath(dst_item_rel, request=request)
|
||||
@@ -301,6 +281,8 @@ async def head_file(request, name=""):
|
||||
@bp.route("/", methods=["OPTIONS"], name="options_root", strict_slashes=False)
|
||||
@bp.route("/<name:path>", methods=["OPTIONS"], name="options_path")
|
||||
async def dav_options(request, name=""):
|
||||
_ = request
|
||||
_ = name
|
||||
return HTTPResponse(
|
||||
status=200,
|
||||
headers={
|
||||
@@ -362,7 +344,7 @@ async def dav_copy(request, name=""):
|
||||
dst_rel, dst_abs = _parse_webdav_destination(dest_header, request=request)
|
||||
if auth.request_share_token(request) is not None and not dst_rel.parts:
|
||||
raise BadRequest("Destination cannot be virtual root")
|
||||
request.ctx._log_extra = f"→ {dst_rel}"
|
||||
request.ctx.log_extra = f"→ {dst_rel}"
|
||||
if not src_abs.exists():
|
||||
raise NotFound(f"Source not found: {name}")
|
||||
if src_abs == dst_abs:
|
||||
@@ -401,7 +383,7 @@ async def dav_move(request, name=""):
|
||||
dst_rel, dst_abs = _parse_webdav_destination(dest_header, request=request)
|
||||
if auth.request_share_token(request) is not None and not dst_rel.parts:
|
||||
raise BadRequest("Destination cannot be virtual root")
|
||||
request.ctx._log_extra = f"→ {dst_rel}"
|
||||
request.ctx.log_extra = f"→ {dst_rel}"
|
||||
if not src_abs.exists():
|
||||
raise NotFound(f"Source not found: {name}")
|
||||
if src_abs == dst_abs:
|
||||
|
||||
+157
-24
@@ -10,12 +10,14 @@ Environment requirements:
|
||||
reachable from the container (usually the docker bridge IP).
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import socketserver
|
||||
import subprocess
|
||||
import threading
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from functools import partial
|
||||
from http.server import SimpleHTTPRequestHandler
|
||||
@@ -23,20 +25,29 @@ from pathlib import Path
|
||||
from time import perf_counter
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
import jwt
|
||||
from sanic.log import logger
|
||||
|
||||
from cista import config
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Configuration helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
_httpx_client: httpx.AsyncClient | None = None
|
||||
|
||||
|
||||
def _get_onlyoffice_url() -> str:
|
||||
return os.environ.get("ONLYOFFICE_URL", "http://localhost:8080")
|
||||
return os.environ.get("ONLYOFFICE_CISTA_URL", "http://localhost:8988")
|
||||
|
||||
|
||||
def _get_jwt_secret() -> str | None:
|
||||
return os.environ.get("ONLYOFFICE_JWT_SECRET") or None
|
||||
def _get_jwt_secret() -> str:
|
||||
return (
|
||||
os.environ.get("ONLYOFFICE_JWT_SECRET")
|
||||
or config.derived_secret("onlyoffice", size=16).hex()
|
||||
)
|
||||
|
||||
|
||||
def _get_callback_host() -> str:
|
||||
@@ -62,19 +73,140 @@ def _get_callback_host() -> str:
|
||||
return "127.0.0.1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Async HTTP client
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def get_httpx_client() -> httpx.AsyncClient:
|
||||
"""Return the shared async HTTP client for OnlyOffice requests."""
|
||||
global _httpx_client
|
||||
if _httpx_client is None:
|
||||
_httpx_client = httpx.AsyncClient()
|
||||
return _httpx_client
|
||||
|
||||
|
||||
async def close_oo_client() -> None:
|
||||
"""Close the shared async HTTP client."""
|
||||
global _httpx_client
|
||||
if _httpx_client is not None:
|
||||
await _httpx_client.aclose()
|
||||
_httpx_client = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Availability check
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def is_available() -> bool:
|
||||
"""Return True if the configured OnlyOffice Document Server is reachable."""
|
||||
url = _get_onlyoffice_url()
|
||||
def _probe_status() -> tuple[bool, bool, str | None]:
|
||||
"""Return (ok, responded, detail) for a lightweight reachability probe."""
|
||||
url = _get_onlyoffice_url().rstrip("/") + "/ConvertService.ashx"
|
||||
try:
|
||||
with urllib.request.urlopen(url, timeout=3) as resp: # noqa: S310
|
||||
return resp.status == 200
|
||||
with urllib.request.urlopen(url, timeout=2) as resp: # noqa: S310
|
||||
status = resp.status
|
||||
except urllib.error.HTTPError as e:
|
||||
status = e.code
|
||||
except Exception:
|
||||
return False, False, None
|
||||
|
||||
if status in (200, 405):
|
||||
return True, True, None
|
||||
if status >= 500:
|
||||
return False, True, f"HTTP {status}"
|
||||
return False, True, f"HTTP {status}"
|
||||
|
||||
|
||||
def log_reachable_info() -> None:
|
||||
"""Log info on success, warning on responded probe errors, silent on no-response."""
|
||||
ok, responded, detail = _probe_status()
|
||||
if ok:
|
||||
logger.info("Using OnlyOffice document server at %s", _get_onlyoffice_url())
|
||||
elif responded:
|
||||
suffix = f": {detail}" if detail else ""
|
||||
logger.warning("OnlyOffice probe failed%s", suffix)
|
||||
|
||||
|
||||
def setup_docker(confdir: Path | None = None) -> int:
|
||||
"""Build and run the patched OnlyOffice Docker image."""
|
||||
config.init_confdir(confdir)
|
||||
if config.conffile.exists():
|
||||
config.load_config()
|
||||
else:
|
||||
config.update_config(
|
||||
{
|
||||
"listen": ":8989",
|
||||
"path": Path.home() / "Downloads",
|
||||
"public": False,
|
||||
}
|
||||
)
|
||||
|
||||
secret = config.derived_secret("onlyoffice", size=16).hex()
|
||||
docker_dir = Path(__file__).parent / "docker"
|
||||
if not docker_dir.is_dir():
|
||||
raise FileNotFoundError(
|
||||
f"Docker files not found at {docker_dir}. Is the package installed correctly?"
|
||||
)
|
||||
|
||||
logger.info("Building OnlyOffice image")
|
||||
build_cmd = ["docker", "build", "-t", "onlyoffice-cista", str(docker_dir)]
|
||||
logger.info("%s", " ".join(build_cmd))
|
||||
result = subprocess.run(build_cmd, check=False, shell=False) # noqa: S603
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError("Failed to build OnlyOffice image")
|
||||
|
||||
logger.info("Starting OnlyOffice container")
|
||||
run_cmd = [
|
||||
"docker",
|
||||
"run",
|
||||
"-d",
|
||||
"-p",
|
||||
"8988:80",
|
||||
"-e",
|
||||
f"JWT_SECRET={secret}",
|
||||
"-e",
|
||||
"WORKERS=8",
|
||||
"--name",
|
||||
"onlyoffice-cista",
|
||||
"--restart",
|
||||
"unless-stopped",
|
||||
"onlyoffice-cista",
|
||||
]
|
||||
logger.info("%s", " ".join(run_cmd))
|
||||
result = subprocess.run(run_cmd, check=False, shell=False) # noqa: S603
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError("Failed to start OnlyOffice container")
|
||||
logger.info("OnlyOffice is running on http://localhost:8988")
|
||||
return 0
|
||||
|
||||
|
||||
async def is_available_async(request_timeout: float = 2.0) -> bool:
|
||||
"""Return True if the configured OnlyOffice Document Server is reachable."""
|
||||
url = _get_onlyoffice_url().rstrip("/") + "/ConvertService.ashx"
|
||||
client = get_httpx_client()
|
||||
try:
|
||||
response = await client.get(url, timeout=request_timeout)
|
||||
except Exception:
|
||||
return False
|
||||
else:
|
||||
return response.status_code in (200, 405)
|
||||
|
||||
|
||||
_oo_available_cache: tuple[bool, float] | None = None
|
||||
OO_AVAILABILITY_CACHE_TTL = 30.0
|
||||
|
||||
|
||||
async def is_available_cached() -> bool:
|
||||
"""Return cached OnlyOffice availability, refreshed every 30 seconds."""
|
||||
global _oo_available_cache
|
||||
now = perf_counter()
|
||||
if _oo_available_cache is not None:
|
||||
result, timestamp = _oo_available_cache
|
||||
if now - timestamp < OO_AVAILABILITY_CACHE_TTL:
|
||||
return result
|
||||
result = await is_available_async()
|
||||
_oo_available_cache = (result, now)
|
||||
return result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -121,22 +253,23 @@ def _build_jwt_token(payload: dict) -> str | None:
|
||||
return jwt.encode(payload, secret, algorithm="HS256")
|
||||
|
||||
|
||||
def convert_to_png(file_path: Path, timeout: float = 30.0) -> bytes:
|
||||
"""Convert *file_path* to PNG using OnlyOffice Document Server.
|
||||
async def convert_to_png_async(file_path: Path, request_timeout: float = 5.0) -> bytes:
|
||||
"""Convert *file_path* to PNG using OnlyOffice Document Server (async).
|
||||
|
||||
Returns the PNG bytes. Raises RuntimeError on failure.
|
||||
"""
|
||||
oo_url = _get_onlyoffice_url().rstrip("/")
|
||||
convert_url = f"{oo_url}/ConvertService.ashx"
|
||||
client = get_httpx_client()
|
||||
|
||||
# Start temporary HTTP server so OnlyOffice can fetch the file
|
||||
doc_url, httpd = _serve_file_temporarily(file_path)
|
||||
doc_url, httpd = await asyncio.to_thread(_serve_file_temporarily, file_path)
|
||||
try:
|
||||
suffix = file_path.suffix.lstrip(".").lower()
|
||||
payload = {
|
||||
"async": False,
|
||||
"filetype": suffix,
|
||||
"key": f"cista_{file_path.stat().st_mtime_ns}",
|
||||
"key": f"cista_{(await asyncio.to_thread(file_path.stat)).st_mtime_ns}",
|
||||
"outputtype": "png",
|
||||
"title": file_path.name,
|
||||
"url": doc_url,
|
||||
@@ -149,16 +282,15 @@ def convert_to_png(file_path: Path, timeout: float = 30.0) -> bytes:
|
||||
payload["token"] = token
|
||||
headers["Authorization"] = token
|
||||
|
||||
req = urllib.request.Request( # noqa: S310
|
||||
convert_url,
|
||||
data=json.dumps(payload).encode(),
|
||||
headers=headers,
|
||||
method="POST",
|
||||
)
|
||||
|
||||
t_start = perf_counter()
|
||||
with urllib.request.urlopen(req, timeout=timeout) as resp: # noqa: S310
|
||||
body = resp.read()
|
||||
response = await client.post(
|
||||
convert_url,
|
||||
content=json.dumps(payload).encode(),
|
||||
headers=headers,
|
||||
timeout=request_timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
body = response.content
|
||||
t_end = perf_counter()
|
||||
|
||||
# Parse XML response
|
||||
@@ -178,7 +310,8 @@ def convert_to_png(file_path: Path, timeout: float = 30.0) -> bytes:
|
||||
logger.debug("OnlyOffice converted in %.2fs: %s", t_end - t_start, file_url)
|
||||
|
||||
# Download converted PNG
|
||||
with urllib.request.urlopen(file_url, timeout=timeout) as png_resp: # noqa: S310
|
||||
return png_resp.read()
|
||||
png_response = await client.get(file_url, timeout=request_timeout)
|
||||
png_response.raise_for_status()
|
||||
return png_response.content
|
||||
finally:
|
||||
httpd.shutdown()
|
||||
await asyncio.to_thread(httpd.shutdown)
|
||||
|
||||
+228
-330
@@ -1,7 +1,5 @@
|
||||
import asyncio
|
||||
import contextlib
|
||||
import gc
|
||||
import io
|
||||
import mimetypes
|
||||
import struct
|
||||
import sys
|
||||
@@ -10,41 +8,27 @@ import urllib.parse
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass
|
||||
from multiprocessing import cpu_count
|
||||
from pathlib import PurePosixPath
|
||||
from pathlib import Path, PurePosixPath
|
||||
from time import perf_counter
|
||||
from urllib.parse import unquote
|
||||
from wsgiref.handlers import format_date_time
|
||||
|
||||
import av
|
||||
import fitz # PyMuPDF
|
||||
import httpx
|
||||
import msgspec
|
||||
import numpy as np
|
||||
import pyvips
|
||||
from blake3 import blake3
|
||||
from sanic import Blueprint, empty, raw, redirect
|
||||
from sanic.exceptions import NotFound
|
||||
from sanic.log import logger
|
||||
|
||||
from cista import auth, config, sharefs
|
||||
from cista.preview_worker import PreviewRequest, PreviewResponse
|
||||
from cista import auth, config, onlyoffice, sharefs
|
||||
from cista.preview_worker import (
|
||||
DOC_PREVIEW_SUFFIXES,
|
||||
OFFICE_PREVIEW_SUFFIXES,
|
||||
PreviewRequest,
|
||||
PreviewResponse,
|
||||
)
|
||||
from cista.util.filename import sanitize
|
||||
|
||||
# OnlyOffice integration is loaded lazily; availability is checked at runtime.
|
||||
_onlyoffice = None
|
||||
|
||||
|
||||
def _get_onlyoffice():
|
||||
global _onlyoffice
|
||||
if _onlyoffice is None:
|
||||
try:
|
||||
from cista import onlyoffice as oo
|
||||
|
||||
_onlyoffice = oo
|
||||
except Exception:
|
||||
_onlyoffice = False
|
||||
return _onlyoffice
|
||||
|
||||
|
||||
bp = Blueprint("preview", url_prefix="/preview")
|
||||
|
||||
|
||||
@@ -112,24 +96,30 @@ class _PreviewWorker:
|
||||
def __init__(self, proc: asyncio.subprocess.Process):
|
||||
self.proc = proc
|
||||
|
||||
async def request(self, filepath, quality: int, maxsize: int, maxzoom: float):
|
||||
async def request(
|
||||
self,
|
||||
filepath,
|
||||
quality: int,
|
||||
maxsize: int,
|
||||
maxzoom: float,
|
||||
data: bytes | None = None,
|
||||
):
|
||||
if self.proc.returncode is not None:
|
||||
raise WorkerProtocolError("worker already exited")
|
||||
if self.proc.stdin is None or self.proc.stdout is None:
|
||||
raise WorkerProtocolError("worker streams not available")
|
||||
|
||||
line = (
|
||||
msgspec.json.encode(
|
||||
PreviewRequest(
|
||||
path=str(filepath),
|
||||
quality=quality,
|
||||
maxsize=maxsize,
|
||||
maxzoom=maxzoom,
|
||||
)
|
||||
meta = msgspec.json.encode(
|
||||
PreviewRequest(
|
||||
path=str(filepath),
|
||||
quality=quality,
|
||||
maxsize=maxsize,
|
||||
maxzoom=maxzoom,
|
||||
)
|
||||
+ b"\n"
|
||||
)
|
||||
self.proc.stdin.write(line)
|
||||
payload = data or b""
|
||||
packet = struct.pack("<II", len(meta), len(payload)) + meta + payload
|
||||
self.proc.stdin.write(packet)
|
||||
await self.proc.stdin.drain()
|
||||
|
||||
checksum = await self.proc.stdout.readexactly(WORKER_CHECKSUM_BYTES)
|
||||
@@ -172,6 +162,14 @@ class _PreviewWorkerPool:
|
||||
self._seq = 0
|
||||
self._closed = False
|
||||
|
||||
async def _read_startup_stderr(self, proc: asyncio.subprocess.Process) -> str:
|
||||
if proc.stderr is None:
|
||||
return ""
|
||||
with contextlib.suppress(TimeoutError):
|
||||
data = await asyncio.wait_for(proc.stderr.read(), timeout=0.5)
|
||||
return data.decode(errors="replace").strip()
|
||||
return ""
|
||||
|
||||
async def _spawn_worker(self) -> _PreviewWorker:
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
sys.executable,
|
||||
@@ -179,10 +177,35 @@ class _PreviewWorkerPool:
|
||||
"cista.preview_worker",
|
||||
stdin=asyncio.subprocess.PIPE,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.DEVNULL,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
start_new_session=True,
|
||||
)
|
||||
_active_procs.add(proc)
|
||||
try:
|
||||
ready = await asyncio.wait_for(proc.stdout.readexactly(1), timeout=30.0)
|
||||
except TimeoutError as err:
|
||||
with contextlib.suppress(ProcessLookupError):
|
||||
proc.kill()
|
||||
with contextlib.suppress(Exception):
|
||||
await proc.wait()
|
||||
stderr = await self._read_startup_stderr(proc)
|
||||
if stderr:
|
||||
raise WorkerProtocolError(
|
||||
"preview worker failed to become ready: " + stderr.splitlines()[-1]
|
||||
) from err
|
||||
raise WorkerProtocolError("preview worker failed to become ready") from err
|
||||
except asyncio.IncompleteReadError as err:
|
||||
stderr = await self._read_startup_stderr(proc)
|
||||
if stderr:
|
||||
raise WorkerProtocolError(
|
||||
"preview worker exited before signalling readiness: "
|
||||
+ stderr.splitlines()[-1]
|
||||
) from err
|
||||
raise WorkerProtocolError(
|
||||
"preview worker exited before signalling readiness"
|
||||
) from err
|
||||
if ready != b"\x01":
|
||||
raise WorkerProtocolError(f"preview worker ready signal invalid: {ready!r}")
|
||||
return _PreviewWorker(proc)
|
||||
|
||||
async def _add_worker(self) -> None:
|
||||
@@ -210,7 +233,20 @@ class _PreviewWorkerPool:
|
||||
if future.cancelled():
|
||||
continue
|
||||
|
||||
worker = await self._idle.get()
|
||||
try:
|
||||
worker = await asyncio.wait_for(
|
||||
self._idle.get(), timeout=PREVIEW_TIMEOUT
|
||||
)
|
||||
except TimeoutError:
|
||||
logger.warning(
|
||||
"Preview worker unavailable (%ds) for %s",
|
||||
int(PREVIEW_TIMEOUT),
|
||||
args[0].name,
|
||||
)
|
||||
if not future.done():
|
||||
future.set_exception(PreviewTimeoutError(args[0].name))
|
||||
continue
|
||||
|
||||
filepath = args[0]
|
||||
replace = False
|
||||
try:
|
||||
@@ -256,6 +292,15 @@ class _PreviewWorkerPool:
|
||||
f"worker protocol failure for {filepath.name}: {e}"
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
replace = True
|
||||
logger.exception(
|
||||
"Unexpected preview worker error for %s", filepath.name
|
||||
)
|
||||
if not future.done():
|
||||
future.set_exception(
|
||||
PreviewError(f"unexpected worker error for {filepath.name}")
|
||||
)
|
||||
finally:
|
||||
if replace:
|
||||
await self._replace_worker(worker)
|
||||
@@ -265,12 +310,23 @@ class _PreviewWorkerPool:
|
||||
await self._replace_worker(worker)
|
||||
|
||||
async def start(self) -> None:
|
||||
for _ in range(self.size):
|
||||
await self._add_worker()
|
||||
workers = await asyncio.gather(
|
||||
*(self._spawn_worker() for _ in range(self.size))
|
||||
)
|
||||
for worker in workers:
|
||||
self._workers.add(worker)
|
||||
await self._idle.put(worker)
|
||||
for _ in range(self.size):
|
||||
self._dispatchers.append(asyncio.create_task(self._dispatch_loop()))
|
||||
|
||||
async def run(self, filepath, quality: int, maxsize: int, maxzoom: float):
|
||||
async def run(
|
||||
self,
|
||||
filepath,
|
||||
quality: int,
|
||||
maxsize: int,
|
||||
maxzoom: float,
|
||||
data: bytes | None = None,
|
||||
):
|
||||
if self._closed:
|
||||
raise PreviewError("preview worker pool closed")
|
||||
loop = asyncio.get_running_loop()
|
||||
@@ -281,7 +337,7 @@ class _PreviewWorkerPool:
|
||||
_preview_job_priority(filepath),
|
||||
self._seq,
|
||||
future,
|
||||
(filepath, quality, maxsize, maxzoom),
|
||||
(filepath, quality, maxsize, maxzoom, data),
|
||||
)
|
||||
)
|
||||
return await future
|
||||
@@ -370,58 +426,108 @@ class PreviewError(Exception):
|
||||
self.backend = backend
|
||||
|
||||
|
||||
# Max concurrent OnlyOffice conversion requests. OO has its own queue;
|
||||
# we must not flood it. This is intentionally small.
|
||||
OO_MAX_CONCURRENT = PREVIEW_WORKERS
|
||||
|
||||
|
||||
class OOConversionManager:
|
||||
"""Manages async OnlyOffice conversions with deduplication and concurrency limits."""
|
||||
|
||||
def __init__(self, max_concurrent: int = OO_MAX_CONCURRENT):
|
||||
self._semaphore = asyncio.Semaphore(max_concurrent)
|
||||
self._in_flight: dict[str, asyncio.Future[bytes]] = {}
|
||||
self._tasks: set[asyncio.Task[None]] = set()
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def convert(self, filepath: Path) -> bytes:
|
||||
"""Return PNG bytes for *filepath*, deduplicating concurrent requests."""
|
||||
stat = await asyncio.to_thread(filepath.stat)
|
||||
key = f"{filepath}:{stat.st_mtime_ns}"
|
||||
|
||||
async with self._lock:
|
||||
if key in self._in_flight:
|
||||
future = self._in_flight[key]
|
||||
else:
|
||||
future = asyncio.get_running_loop().create_future()
|
||||
self._in_flight[key] = future
|
||||
task = asyncio.create_task(self._do_convert(filepath, key, future))
|
||||
self._tasks.add(task)
|
||||
task.add_done_callback(self._tasks.discard)
|
||||
|
||||
return await future
|
||||
|
||||
async def _do_convert(
|
||||
self, filepath: Path, key: str, future: asyncio.Future[bytes]
|
||||
) -> None:
|
||||
try:
|
||||
async with self._semaphore:
|
||||
png_bytes = await onlyoffice.convert_to_png_async(
|
||||
filepath, request_timeout=5.0
|
||||
)
|
||||
except Exception as e:
|
||||
future.set_exception(e)
|
||||
async with self._lock:
|
||||
self._in_flight.pop(key, None)
|
||||
else:
|
||||
future.set_result(png_bytes)
|
||||
async with self._lock:
|
||||
self._in_flight.pop(key, None)
|
||||
|
||||
|
||||
_oo_manager: OOConversionManager | None = None
|
||||
|
||||
|
||||
def get_oo_manager() -> OOConversionManager:
|
||||
"""Return the singleton OOConversionManager."""
|
||||
global _oo_manager
|
||||
if _oo_manager is None:
|
||||
_oo_manager = OOConversionManager(max_concurrent=OO_MAX_CONCURRENT)
|
||||
return _oo_manager
|
||||
|
||||
|
||||
async def _generate_office_preview(
|
||||
filepath: Path, quality: int, maxsize: int, maxzoom: float
|
||||
) -> tuple[bytes | None, PreviewResponse | None]:
|
||||
"""Generate a preview for an office file using OnlyOffice + worker AVIF conversion."""
|
||||
manager = get_oo_manager()
|
||||
t_oo_start = perf_counter()
|
||||
png_bytes = await manager.convert(filepath)
|
||||
t_oo_end = perf_counter()
|
||||
|
||||
img, resp = await _run_preview_process(
|
||||
filepath, quality, maxsize, maxzoom, data=png_bytes
|
||||
)
|
||||
|
||||
if resp is not None:
|
||||
resp.backend = "onlyoffice+" + (resp.backend or "pyvips")
|
||||
if resp.timings:
|
||||
resp.timings = [round((t_oo_end - t_oo_start) * 1000, 1), *resp.timings]
|
||||
return img, resp
|
||||
|
||||
|
||||
async def _run_preview_process(
|
||||
filepath, quality: int, maxsize: int, maxzoom: float
|
||||
filepath, quality: int, maxsize: int, maxzoom: float, data: bytes | None = None
|
||||
) -> tuple[bytes | None, PreviewResponse | None]:
|
||||
"""Run preview request in a persistent worker process."""
|
||||
await start_preview_workers()
|
||||
if _preview_pool is None:
|
||||
raise PreviewError(f"preview worker pool unavailable for {filepath.name}")
|
||||
return await _preview_pool.run(filepath, quality, maxsize, maxzoom)
|
||||
return await _preview_pool.run(filepath, quality, maxsize, maxzoom, data)
|
||||
|
||||
|
||||
DOC_PREVIEW_SUFFIXES = {".pdf", ".xps", ".epub", ".mobi"}
|
||||
|
||||
OFFICE_PREVIEW_SUFFIXES = {
|
||||
".doc",
|
||||
".dot",
|
||||
".docx",
|
||||
".docm",
|
||||
".dotx",
|
||||
".dotm",
|
||||
".rtf",
|
||||
".odt",
|
||||
".ott",
|
||||
".txt",
|
||||
".md",
|
||||
".mhtml",
|
||||
".mht",
|
||||
".html",
|
||||
".htm",
|
||||
".xml",
|
||||
".wps",
|
||||
".wri",
|
||||
# Spreadsheets
|
||||
".xls",
|
||||
".xlsx",
|
||||
".xlsm",
|
||||
".xlsb",
|
||||
".xltx",
|
||||
".xltm",
|
||||
".ods",
|
||||
".ots",
|
||||
".csv",
|
||||
# Presentations
|
||||
".ppt",
|
||||
".pptx",
|
||||
".pptm",
|
||||
".pps",
|
||||
".ppsx",
|
||||
".pot",
|
||||
".potx",
|
||||
".odp",
|
||||
".otp",
|
||||
}
|
||||
def _onlyoffice_error_short_text(detail: str) -> str:
|
||||
if detail.startswith("OnlyOffice conversion error:"):
|
||||
code = detail.rsplit(":", 1)[-1].strip()
|
||||
return {
|
||||
"-8": "onlyoffice jwt error",
|
||||
"-4": "onlyoffice input error",
|
||||
"-2": "onlyoffice timeout error",
|
||||
"-1": "onlyoffice unknown error",
|
||||
}.get(code, f"onlyoffice {code} error")
|
||||
if "OnlyOffice response did not contain FileUrl" in detail:
|
||||
return "onlyoffice no-fileurl error"
|
||||
return "onlyoffice error"
|
||||
|
||||
|
||||
def _preview_job_priority(path) -> int:
|
||||
@@ -492,14 +598,37 @@ async def preview(req, path):
|
||||
|
||||
# Generate preview
|
||||
try:
|
||||
img, preview_resp = await _run_preview_process(
|
||||
filepath, quality, maxsize, maxzoom
|
||||
)
|
||||
if filepath.suffix.lower() in OFFICE_PREVIEW_SUFFIXES:
|
||||
img, preview_resp = await asyncio.wait_for(
|
||||
_generate_office_preview(filepath, quality, maxsize, maxzoom),
|
||||
timeout=PREVIEW_TIMEOUT,
|
||||
)
|
||||
else:
|
||||
img, preview_resp = await asyncio.wait_for(
|
||||
_run_preview_process(filepath, quality, maxsize, maxzoom),
|
||||
timeout=PREVIEW_TIMEOUT,
|
||||
)
|
||||
except TimeoutError:
|
||||
logger.warning("Preview timeout for %s", filepath)
|
||||
return empty(503)
|
||||
except PreviewTimeoutError:
|
||||
return empty(504)
|
||||
logger.warning("Preview worker timeout for %s", filepath)
|
||||
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 PreviewError as e:
|
||||
if e.backend:
|
||||
req.ctx._log_extra = e.backend
|
||||
req.ctx.log_extra = e.backend
|
||||
detail = str(e)
|
||||
if detail == "preview worker error" and e.stderr:
|
||||
captured = e.stderr.strip()
|
||||
@@ -507,14 +636,20 @@ async def preview(req, path):
|
||||
detail = captured.splitlines()[0]
|
||||
logger.error("%s preview: %s", filepath, detail)
|
||||
return empty(422)
|
||||
except asyncio.CancelledError:
|
||||
req.ctx.log_extra = "preview cancelled"
|
||||
return empty(503)
|
||||
except Exception:
|
||||
logger.exception("Unhandled preview error for %s", filepath)
|
||||
return empty(500)
|
||||
if preview_resp and preview_resp.backend:
|
||||
if preview_resp.timings:
|
||||
timing_detail = "/".join(
|
||||
str(round(value)) for value in preview_resp.timings
|
||||
)
|
||||
req.ctx._log_extra = f"{preview_resp.backend} {timing_detail} ➛"
|
||||
req.ctx.log_extra = f"{preview_resp.backend} {timing_detail} ➛"
|
||||
else:
|
||||
req.ctx._log_extra = preview_resp.backend
|
||||
req.ctx.log_extra = preview_resp.backend
|
||||
if not img:
|
||||
# Preview generation failed, redirect to the file itself
|
||||
return redirect(f"/files/{path}", status=303)
|
||||
@@ -537,240 +672,3 @@ async def preview(req, path):
|
||||
_preview_cache.set(etag, CachedPreview(headers=headers, body=img))
|
||||
|
||||
return raw(img, headers=headers)
|
||||
|
||||
|
||||
def dispatch(path, quality, maxsize, maxzoom):
|
||||
backend = "unknown"
|
||||
try:
|
||||
suffix = path.suffix.lower()
|
||||
if suffix in DOC_PREVIEW_SUFFIXES:
|
||||
backend = "pdf"
|
||||
return process_pdf(path, quality=quality, maxsize=maxsize, maxzoom=maxzoom)
|
||||
if suffix in OFFICE_PREVIEW_SUFFIXES:
|
||||
backend = "onlyoffice"
|
||||
return process_office(
|
||||
path, quality=quality, maxsize=maxsize, maxzoom=maxzoom
|
||||
)
|
||||
mime_type, _ = mimetypes.guess_type(path.name)
|
||||
if mime_type and mime_type.startswith("video/"):
|
||||
backend = "video"
|
||||
return process_video(path, quality=quality, maxsize=maxsize)
|
||||
if mime_type and mime_type.startswith("image/"):
|
||||
backend = "pyvips"
|
||||
return process_image(path, quality=quality, maxsize=maxsize)
|
||||
except ValueError as e:
|
||||
return None, PreviewResponse(ok=False, backend=backend, error=str(e))
|
||||
except Exception as e:
|
||||
return None, PreviewResponse(ok=False, backend=backend, error=str(e))
|
||||
return None, PreviewResponse(ok=False, backend=backend, error="preview unsupported")
|
||||
|
||||
|
||||
def process_image(path, *, maxsize, quality):
|
||||
return process_image_pyvips(path, maxsize=maxsize, quality=quality)
|
||||
|
||||
|
||||
def process_image_pyvips(path, *, maxsize, quality):
|
||||
t_start = perf_counter()
|
||||
img = pyvips.Image.new_from_file(str(path), access="sequential")
|
||||
img = img.autorot()
|
||||
scale = min(maxsize / img.width, maxsize / img.height, 1.0)
|
||||
if scale < 1.0:
|
||||
img = img.resize(scale)
|
||||
ret = img.write_to_buffer(
|
||||
".avif",
|
||||
Q=quality,
|
||||
effort=AVIF_FAST_EFFORT,
|
||||
strip=True,
|
||||
)
|
||||
t_end = perf_counter()
|
||||
|
||||
return ret, PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend="pyvips",
|
||||
timings=[round((t_end - t_start) * 1000, 1)],
|
||||
)
|
||||
|
||||
|
||||
def process_pdf(path, *, maxsize, maxzoom, quality, page_number=0):
|
||||
t_load_start = perf_counter()
|
||||
pdf = fitz.open(path)
|
||||
page = pdf.load_page(page_number)
|
||||
w, h = page.rect[2:4]
|
||||
zoom = min(maxsize / w, maxsize / h, maxzoom)
|
||||
mat = fitz.Matrix(zoom, zoom)
|
||||
pix = page.get_pixmap(matrix=mat)
|
||||
t_load_end = perf_counter()
|
||||
|
||||
t_save_start = perf_counter()
|
||||
img = pyvips.Image.new_from_memory(
|
||||
pix.samples_mv, pix.width, pix.height, pix.n, "uchar"
|
||||
)
|
||||
ret = img.write_to_buffer(".avif", Q=quality, effort=AVIF_FAST_EFFORT, strip=True)
|
||||
backend = "pdf+pyvips"
|
||||
t_save_end = perf_counter()
|
||||
|
||||
return ret, PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend=backend,
|
||||
timings=[
|
||||
round((t_load_end - t_load_start) * 1000, 1),
|
||||
round((t_save_end - t_save_start) * 1000, 1),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def process_office(path, *, quality, maxsize, maxzoom):
|
||||
t_load_start = perf_counter()
|
||||
oo = _get_onlyoffice()
|
||||
if oo is False:
|
||||
raise RuntimeError("OnlyOffice is not installed")
|
||||
if not oo.is_available():
|
||||
raise RuntimeError("OnlyOffice Document Server is not reachable")
|
||||
png_bytes = oo.convert_to_png(path)
|
||||
t_load_end = perf_counter()
|
||||
|
||||
t_save_start = perf_counter()
|
||||
img = pyvips.Image.new_from_buffer(png_bytes, "")
|
||||
scale = min(maxsize / img.width, maxsize / img.height, 1.0)
|
||||
if scale < 1.0:
|
||||
img = img.resize(scale)
|
||||
ret = img.write_to_buffer(".avif", Q=quality, effort=AVIF_FAST_EFFORT, strip=True)
|
||||
backend = "onlyoffice+pyvips"
|
||||
t_save_end = perf_counter()
|
||||
|
||||
return ret, PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend=backend,
|
||||
timings=[
|
||||
round((t_load_end - t_load_start) * 1000, 1),
|
||||
round((t_save_end - t_save_start) * 1000, 1),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def process_video(path, *, maxsize, quality):
|
||||
frame = None
|
||||
imgdata = io.BytesIO()
|
||||
istream = ostream = icc = occ = frame = None
|
||||
t_load_start = perf_counter()
|
||||
# Initialize to avoid "possibly unbound" in static analysis when exceptions occur
|
||||
t_load_end = t_load_start
|
||||
t_save_start = t_load_start
|
||||
t_save_end = t_load_start
|
||||
with (
|
||||
av.open(
|
||||
str(path),
|
||||
options={
|
||||
"analyzeduration": "1000000", # 1 second (in microseconds)
|
||||
"fflags": "fastseek",
|
||||
},
|
||||
) as icontainer,
|
||||
av.open(imgdata, "w", format="avif") as ocontainer,
|
||||
):
|
||||
istream = icontainer.streams.video[0]
|
||||
istream.codec_context.skip_frame = "NONKEY"
|
||||
icontainer.seek((icontainer.duration or 0) // 8)
|
||||
for frame in icontainer.decode(istream):
|
||||
if frame.dts is not None:
|
||||
break
|
||||
else:
|
||||
raise RuntimeError("No frames found in video")
|
||||
|
||||
# Resize frame to thumbnail size
|
||||
if frame.width > maxsize or frame.height > maxsize:
|
||||
scale_factor = min(maxsize / frame.width, maxsize / frame.height)
|
||||
new_width = int(frame.width * scale_factor)
|
||||
new_height = int(frame.height * scale_factor)
|
||||
frame = frame.reformat(width=new_width, height=new_height)
|
||||
|
||||
# Apply EXIF rotation if present
|
||||
if frame.rotation:
|
||||
# frame.rotation indicates clockwise rotation needed to display correctly
|
||||
# np.rot90 rotates counter-clockwise, so we negate k
|
||||
k = (frame.rotation // 90) % 4 # Convert to counter-clockwise rotations
|
||||
if k == 2:
|
||||
# 180° rotation can be done in YUV420p, preserving HDR
|
||||
try:
|
||||
fplanes = frame.to_ndarray()
|
||||
# Split into Y, U, V planes of proper dimensions
|
||||
planes = [
|
||||
fplanes[: frame.height],
|
||||
fplanes[
|
||||
frame.height : frame.height + frame.height // 4
|
||||
].reshape(frame.height // 2, frame.width // 2),
|
||||
fplanes[frame.height + frame.height // 4 :].reshape(
|
||||
frame.height // 2, frame.width // 2
|
||||
),
|
||||
]
|
||||
# Rotate each plane by 180°
|
||||
planes = [np.rot90(p, 2) for p in planes]
|
||||
# Restore PyAV format
|
||||
planes = np.hstack([p.flat for p in planes]).reshape(
|
||||
-1, planes[0].shape[1]
|
||||
)
|
||||
frame = av.VideoFrame.from_ndarray(planes, format=frame.format.name)
|
||||
del planes, fplanes
|
||||
except Exception as e:
|
||||
logger.exception(f"Error rotating video frame by 180°: {e}")
|
||||
elif k in (1, 3):
|
||||
# 90° or 270° rotation requires RGB conversion (loses HDR)
|
||||
try:
|
||||
rgb = frame.to_ndarray(format="rgb24")
|
||||
rgb = np.rot90(rgb, k)
|
||||
frame = av.VideoFrame.from_ndarray(rgb, format="rgb24")
|
||||
frame = frame.reformat(
|
||||
format="yuv420p"
|
||||
) # Convert back for encoding
|
||||
del rgb
|
||||
except Exception as e:
|
||||
logger.exception(
|
||||
f"Error rotating video frame by {frame.rotation}°: {e}"
|
||||
)
|
||||
t_load_end = perf_counter()
|
||||
|
||||
t_save_start = perf_counter()
|
||||
crf = str(int(63 * (1 - quality / 100) ** 2)) # Closely matching PIL quality-%
|
||||
ostream = ocontainer.add_stream(
|
||||
"av1",
|
||||
options={
|
||||
"crf": crf,
|
||||
"usage": "realtime",
|
||||
"cpu-used": "8",
|
||||
"threads": "1",
|
||||
},
|
||||
)
|
||||
if not isinstance(ostream, av.VideoStream):
|
||||
raise PreviewError("failed to initialize AV1 video stream")
|
||||
ostream.width = frame.width
|
||||
ostream.height = frame.height
|
||||
ostream.pix_fmt = frame.format.name
|
||||
icc = istream.codec_context
|
||||
occ = ostream.codec_context
|
||||
|
||||
# Copy HDR metadata from input video stream
|
||||
occ.color_primaries = icc.color_primaries
|
||||
occ.color_trc = icc.color_trc
|
||||
occ.colorspace = icc.colorspace
|
||||
occ.color_range = icc.color_range
|
||||
|
||||
ocontainer.mux(ostream.encode(frame))
|
||||
ocontainer.mux(ostream.encode(None)) # Flush the stream
|
||||
t_save_end = perf_counter()
|
||||
|
||||
# Capture result before cleanup
|
||||
ret = imgdata.getvalue()
|
||||
resp = PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend="video",
|
||||
timings=[
|
||||
round((t_load_end - t_load_start) * 1000, 1),
|
||||
round((t_save_end - t_save_start) * 1000, 1),
|
||||
],
|
||||
)
|
||||
del imgdata, istream, ostream, icc, occ, frame
|
||||
gc.collect()
|
||||
return ret, resp
|
||||
|
||||
+420
-11
@@ -1,24 +1,85 @@
|
||||
"""Preview generation worker subprocess.
|
||||
"""Preview generation worker subprocess and synchronous preview engine.
|
||||
|
||||
Two modes are supported:
|
||||
1) Legacy one-shot mode: argv has path/quality/maxsize/maxzoom.
|
||||
2) Long-lived mode: read JSONL commands from stdin and write framed responses.
|
||||
2) Long-lived mode: read framed requests from stdin and write framed responses.
|
||||
|
||||
Framed response format:
|
||||
Framed request format (stdin):
|
||||
(uint32 json size)(uint32 data size)(json)(binary data)
|
||||
|
||||
Framed response format (stdout):
|
||||
(blake3(packet))(uint32 json size)(uint32 payload size)(json)(binary payload)
|
||||
where packet = (uint32 json size)(uint32 payload size)(json)(binary payload).
|
||||
"""
|
||||
|
||||
import contextlib
|
||||
import gc
|
||||
import io
|
||||
import logging
|
||||
import mimetypes
|
||||
import struct
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from time import perf_counter
|
||||
|
||||
import av
|
||||
import fitz # PyMuPDF
|
||||
import msgspec
|
||||
import numpy as np
|
||||
import pyvips
|
||||
from blake3 import blake3
|
||||
|
||||
from cista import config
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
AVIF_FAST_EFFORT = 0
|
||||
|
||||
DOC_PREVIEW_SUFFIXES = {".pdf", ".xps", ".epub", ".mobi"}
|
||||
|
||||
OFFICE_PREVIEW_SUFFIXES = {
|
||||
".doc",
|
||||
".dot",
|
||||
".docx",
|
||||
".docm",
|
||||
".dotx",
|
||||
".dotm",
|
||||
".rtf",
|
||||
".odt",
|
||||
".ott",
|
||||
".txt",
|
||||
".md",
|
||||
".mhtml",
|
||||
".mht",
|
||||
".html",
|
||||
".htm",
|
||||
".xml",
|
||||
".wps",
|
||||
".wri",
|
||||
# Spreadsheets
|
||||
".xls",
|
||||
".xlsx",
|
||||
".xlsm",
|
||||
".xlsb",
|
||||
".xltx",
|
||||
".xltm",
|
||||
".ods",
|
||||
".ots",
|
||||
".csv",
|
||||
# Presentations
|
||||
".ppt",
|
||||
".pptx",
|
||||
".pptm",
|
||||
".pps",
|
||||
".ppsx",
|
||||
".pot",
|
||||
".potx",
|
||||
".odp",
|
||||
".otp",
|
||||
}
|
||||
|
||||
|
||||
class PreviewRequest(msgspec.Struct, omit_defaults=True):
|
||||
path: str
|
||||
@@ -40,6 +101,30 @@ _enc = msgspec.json.Encoder()
|
||||
_dec_req = msgspec.json.Decoder(PreviewRequest)
|
||||
|
||||
|
||||
def _read_exactly(f, n: int) -> bytes:
|
||||
buf = b""
|
||||
while len(buf) < n:
|
||||
chunk = f.read(n - len(buf))
|
||||
if not chunk:
|
||||
raise EOFError
|
||||
buf += chunk
|
||||
return buf
|
||||
|
||||
|
||||
def _read_request() -> tuple[PreviewRequest, bytes] | None:
|
||||
try:
|
||||
header = _read_exactly(sys.stdin.buffer, 8)
|
||||
except EOFError:
|
||||
return None
|
||||
json_size, data_size = struct.unpack("<II", header)
|
||||
meta_raw = _read_exactly(sys.stdin.buffer, json_size)
|
||||
data = b""
|
||||
if data_size:
|
||||
data = _read_exactly(sys.stdin.buffer, data_size)
|
||||
req = _dec_req.decode(meta_raw)
|
||||
return req, data
|
||||
|
||||
|
||||
def _write_response(resp: PreviewResponse, payload: bytes) -> None:
|
||||
meta_bytes = _enc.encode(resp)
|
||||
packet = struct.pack("<II", len(meta_bytes), len(payload)) + meta_bytes + payload
|
||||
@@ -49,13 +134,325 @@ def _write_response(resp: PreviewResponse, payload: bytes) -> None:
|
||||
sys.stdout.buffer.flush()
|
||||
|
||||
|
||||
def dispatch(path, quality, maxsize, maxzoom, data=None):
|
||||
backend = "unknown"
|
||||
try:
|
||||
if data:
|
||||
backend = "pyvips"
|
||||
return process_image_buffer(
|
||||
data, quality=quality, maxsize=maxsize, maxzoom=maxzoom
|
||||
)
|
||||
suffix = path.suffix.lower()
|
||||
if suffix in DOC_PREVIEW_SUFFIXES:
|
||||
backend = "pdf"
|
||||
return process_pdf(path, quality=quality, maxsize=maxsize, maxzoom=maxzoom)
|
||||
mime_type, _ = mimetypes.guess_type(path.name)
|
||||
if mime_type and mime_type.startswith("video/"):
|
||||
backend = "video"
|
||||
return process_video(path, quality=quality, maxsize=maxsize)
|
||||
if mime_type and mime_type.startswith("image/"):
|
||||
backend = "pyvips"
|
||||
return process_image(path, quality=quality, maxsize=maxsize)
|
||||
except ValueError as e:
|
||||
return None, PreviewResponse(ok=False, backend=backend, error=str(e))
|
||||
except Exception as e:
|
||||
logger.exception("Preview dispatch failed for %s", path)
|
||||
return None, PreviewResponse(ok=False, backend=backend, error=str(e))
|
||||
return None, PreviewResponse(ok=False, backend=backend, error="preview unsupported")
|
||||
|
||||
|
||||
def process_image(path, *, maxsize, quality):
|
||||
return process_image_pyvips(path, maxsize=maxsize, quality=quality)
|
||||
|
||||
|
||||
def _get_image_dimensions(path: Path) -> tuple[int, int] | None:
|
||||
"""Probe image dimensions.
|
||||
|
||||
pyvips can read the header of most formats (including HEIC) without
|
||||
fully decoding the image.
|
||||
"""
|
||||
try:
|
||||
img = pyvips.Image.new_from_file(str(path))
|
||||
except pyvips.error.Error:
|
||||
return None
|
||||
else:
|
||||
return img.width, img.height
|
||||
|
||||
|
||||
def _image_via_ffmpeg(path: Path, maxsize: int, quality: int) -> bytes:
|
||||
"""Convert any image to AVIF using ffmpeg CLI.
|
||||
|
||||
ffmpeg handles HEIC tile assembly, EXIF rotation, HDR metadata and
|
||||
ICC profile embedding automatically.
|
||||
"""
|
||||
dims = _get_image_dimensions(path)
|
||||
crf = int(63 * (1 - quality / 100) ** 2)
|
||||
with tempfile.NamedTemporaryFile(suffix=".avif", delete=False) as tmp_f:
|
||||
tmp_path = tmp_f.name
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(path),
|
||||
"-frames:v",
|
||||
"1",
|
||||
"-c:v",
|
||||
"av1",
|
||||
"-crf",
|
||||
str(crf),
|
||||
"-cpu-used",
|
||||
"8",
|
||||
tmp_path,
|
||||
]
|
||||
if dims is not None:
|
||||
w, h = dims
|
||||
if max(w, h) > maxsize:
|
||||
scale = min(maxsize / w, maxsize / h)
|
||||
new_w = int(w * scale)
|
||||
new_h = int(h * scale)
|
||||
# insert -s <wxh> right after the input file
|
||||
cmd.insert(4, "-s")
|
||||
cmd.insert(5, f"{new_w}x{new_h}")
|
||||
try:
|
||||
subprocess.run(cmd, capture_output=True, check=True, shell=False) # noqa: S603
|
||||
with Path(tmp_path).open("rb") as f:
|
||||
return f.read()
|
||||
finally:
|
||||
Path(tmp_path).unlink(missing_ok=True)
|
||||
|
||||
|
||||
def process_image_pyvips(path, *, maxsize, quality):
|
||||
t_start = perf_counter()
|
||||
suffix = path.suffix.lower()
|
||||
|
||||
# HEIC/HEIF: ffmpeg handles tile assembly and HDR correctly;
|
||||
# skip pyvips entirely.
|
||||
if suffix in (".heic", ".heif"):
|
||||
ret = _image_via_ffmpeg(path, maxsize, quality)
|
||||
t_end = perf_counter()
|
||||
return ret, PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend="ffmpeg",
|
||||
timings=[round((t_end - t_start) * 1000, 1)],
|
||||
)
|
||||
|
||||
# Other image formats: pyvips first, ffmpeg fallback.
|
||||
load_opts = {"access": "sequential"}
|
||||
try:
|
||||
img = pyvips.Image.new_from_file(str(path), **load_opts)
|
||||
img = img.autorot()
|
||||
scale = min(maxsize / img.width, maxsize / img.height, 1.0)
|
||||
if scale < 1.0:
|
||||
img = img.resize(scale)
|
||||
ret = img.write_to_buffer(
|
||||
".avif",
|
||||
Q=quality,
|
||||
effort=AVIF_FAST_EFFORT,
|
||||
strip=True,
|
||||
)
|
||||
backend = "pyvips"
|
||||
except pyvips.error.Error:
|
||||
ret = _image_via_ffmpeg(path, maxsize, quality)
|
||||
backend = "ffmpeg"
|
||||
t_end = perf_counter()
|
||||
|
||||
return ret, PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend=backend,
|
||||
timings=[round((t_end - t_start) * 1000, 1)],
|
||||
)
|
||||
|
||||
|
||||
def process_image_buffer(data: bytes, *, quality, maxsize, maxzoom):
|
||||
_ = maxzoom
|
||||
t_start = perf_counter()
|
||||
img = pyvips.Image.new_from_buffer(data, "")
|
||||
img = img.autorot()
|
||||
scale = min(maxsize / img.width, maxsize / img.height, 1.0)
|
||||
if scale < 1.0:
|
||||
img = img.resize(scale)
|
||||
ret = img.write_to_buffer(
|
||||
".avif",
|
||||
Q=quality,
|
||||
effort=AVIF_FAST_EFFORT,
|
||||
strip=True,
|
||||
)
|
||||
t_end = perf_counter()
|
||||
|
||||
return ret, PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend="pyvips",
|
||||
timings=[round((t_end - t_start) * 1000, 1)],
|
||||
)
|
||||
|
||||
|
||||
def process_pdf(path, *, maxsize, maxzoom, quality, page_number=0):
|
||||
t_load_start = perf_counter()
|
||||
pdf = fitz.open(path)
|
||||
page = pdf.load_page(page_number)
|
||||
w, h = page.rect[2:4]
|
||||
zoom = min(maxsize / w, maxsize / h, maxzoom)
|
||||
mat = fitz.Matrix(zoom, zoom)
|
||||
pix = page.get_pixmap(matrix=mat)
|
||||
t_load_end = perf_counter()
|
||||
|
||||
t_save_start = perf_counter()
|
||||
img = pyvips.Image.new_from_memory(
|
||||
pix.samples_mv, pix.width, pix.height, pix.n, "uchar"
|
||||
)
|
||||
ret = img.write_to_buffer(".avif", Q=quality, effort=AVIF_FAST_EFFORT, strip=True)
|
||||
backend = "pdf+pyvips"
|
||||
t_save_end = perf_counter()
|
||||
|
||||
return ret, PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend=backend,
|
||||
timings=[
|
||||
round((t_load_end - t_load_start) * 1000, 1),
|
||||
round((t_save_end - t_save_start) * 1000, 1),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def process_video(path, *, maxsize, quality):
|
||||
frame = None
|
||||
imgdata = io.BytesIO()
|
||||
istream = ostream = icc = occ = frame = None
|
||||
t_load_start = perf_counter()
|
||||
# Initialize to avoid "possibly unbound" in static analysis when exceptions occur
|
||||
t_load_end = t_load_start
|
||||
t_save_start = t_load_start
|
||||
t_save_end = t_load_start
|
||||
with (
|
||||
av.open(
|
||||
str(path),
|
||||
options={
|
||||
"analyzeduration": "1000000", # 1 second (in microseconds)
|
||||
"fflags": "fastseek",
|
||||
},
|
||||
) as icontainer,
|
||||
av.open(imgdata, "w", format="avif") as ocontainer,
|
||||
):
|
||||
istream = icontainer.streams.video[0]
|
||||
istream.codec_context.skip_frame = "NONKEY"
|
||||
icontainer.seek((icontainer.duration or 0) // 8)
|
||||
for frame in icontainer.decode(istream):
|
||||
if frame.dts is not None:
|
||||
break
|
||||
else:
|
||||
raise RuntimeError("No frames found in video")
|
||||
|
||||
# Resize frame to thumbnail size
|
||||
if frame.width > maxsize or frame.height > maxsize:
|
||||
scale_factor = min(maxsize / frame.width, maxsize / frame.height)
|
||||
new_width = int(frame.width * scale_factor)
|
||||
new_height = int(frame.height * scale_factor)
|
||||
frame = frame.reformat(width=new_width, height=new_height)
|
||||
|
||||
# Apply EXIF rotation if present
|
||||
if frame.rotation:
|
||||
# frame.rotation indicates clockwise rotation needed to display correctly
|
||||
# np.rot90 rotates counter-clockwise, so we negate k
|
||||
k = (frame.rotation // 90) % 4 # Convert to counter-clockwise rotations
|
||||
if k == 2:
|
||||
# 180° rotation can be done in YUV420p, preserving HDR
|
||||
try:
|
||||
fplanes = frame.to_ndarray()
|
||||
# Split into Y, U, V planes of proper dimensions
|
||||
planes = [
|
||||
fplanes[: frame.height],
|
||||
fplanes[
|
||||
frame.height : frame.height + frame.height // 4
|
||||
].reshape(frame.height // 2, frame.width // 2),
|
||||
fplanes[frame.height + frame.height // 4 :].reshape(
|
||||
frame.height // 2, frame.width // 2
|
||||
),
|
||||
]
|
||||
# Rotate each plane by 180°
|
||||
planes = [np.rot90(p, 2) for p in planes]
|
||||
# Restore PyAV format
|
||||
planes = np.hstack([p.flat for p in planes]).reshape(
|
||||
-1, planes[0].shape[1]
|
||||
)
|
||||
frame = av.VideoFrame.from_ndarray(planes, format=frame.format.name)
|
||||
del planes, fplanes
|
||||
except Exception:
|
||||
logger.exception("Error rotating video frame by 180°")
|
||||
elif k in (1, 3):
|
||||
# 90° or 270° rotation requires RGB conversion (loses HDR)
|
||||
try:
|
||||
rgb = frame.to_ndarray(format="rgb24")
|
||||
rgb = np.rot90(rgb, k)
|
||||
frame = av.VideoFrame.from_ndarray(rgb, format="rgb24")
|
||||
frame = frame.reformat(
|
||||
format="yuv420p"
|
||||
) # Convert back for encoding
|
||||
del rgb
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Error rotating video frame by %s°", frame.rotation
|
||||
)
|
||||
|
||||
# libsvtav1 rejects full-range JPEG-style YUV pixel formats such as
|
||||
# yuvj420p, so normalize them before opening the encoder.
|
||||
if frame.format.name.startswith("yuvj"):
|
||||
frame = frame.reformat(format="yuv420p")
|
||||
t_load_end = perf_counter()
|
||||
|
||||
t_save_start = perf_counter()
|
||||
crf = str(int(63 * (1 - quality / 100) ** 2)) # Closely matching PIL quality-%
|
||||
ostream = ocontainer.add_stream(
|
||||
"av1",
|
||||
options={
|
||||
"crf": crf,
|
||||
"usage": "realtime",
|
||||
"cpu-used": "8",
|
||||
"threads": "1",
|
||||
},
|
||||
)
|
||||
if not isinstance(ostream, av.VideoStream):
|
||||
raise TypeError("failed to initialize AV1 video stream")
|
||||
ostream.width = frame.width
|
||||
ostream.height = frame.height
|
||||
ostream.pix_fmt = frame.format.name
|
||||
icc = istream.codec_context
|
||||
occ = ostream.codec_context
|
||||
|
||||
# Copy HDR metadata from input video stream
|
||||
occ.color_primaries = icc.color_primaries
|
||||
occ.color_trc = icc.color_trc
|
||||
occ.colorspace = icc.colorspace
|
||||
occ.color_range = icc.color_range
|
||||
|
||||
ocontainer.mux(ostream.encode(frame))
|
||||
ocontainer.mux(ostream.encode(None)) # Flush the stream
|
||||
t_save_end = perf_counter()
|
||||
|
||||
# Capture result before cleanup
|
||||
ret = imgdata.getvalue()
|
||||
resp = PreviewResponse(
|
||||
ok=True,
|
||||
mime="image/avif",
|
||||
backend="video",
|
||||
timings=[
|
||||
round((t_load_end - t_load_start) * 1000, 1),
|
||||
round((t_save_end - t_save_start) * 1000, 1),
|
||||
],
|
||||
)
|
||||
del imgdata, istream, ostream, icc, occ, frame
|
||||
gc.collect()
|
||||
return ret, resp
|
||||
|
||||
|
||||
def _run_once() -> None:
|
||||
if len(sys.argv) != 5:
|
||||
sys.stderr.write(f"Usage: {sys.argv[0]} <path> <quality> <maxsize> <maxzoom>\n")
|
||||
sys.exit(1)
|
||||
|
||||
from cista.preview import dispatch
|
||||
|
||||
path = Path(sys.argv[1])
|
||||
quality = int(sys.argv[2])
|
||||
maxsize = int(sys.argv[3])
|
||||
@@ -67,21 +464,19 @@ def _run_once() -> None:
|
||||
|
||||
|
||||
def _run_loop() -> None:
|
||||
from cista.preview import dispatch
|
||||
|
||||
while True:
|
||||
line = sys.stdin.buffer.readline()
|
||||
if not line:
|
||||
result = _read_request()
|
||||
if result is None:
|
||||
return
|
||||
req, data = result
|
||||
stderr_capture = io.StringIO()
|
||||
handler = logging.StreamHandler(stderr_capture)
|
||||
root_logger = logging.getLogger()
|
||||
root_logger.addHandler(handler)
|
||||
try:
|
||||
with contextlib.redirect_stderr(stderr_capture):
|
||||
req = _dec_req.decode(line)
|
||||
result, resp = dispatch(
|
||||
Path(req.path), req.quality, req.maxsize, req.maxzoom
|
||||
Path(req.path), req.quality, req.maxsize, req.maxzoom, data
|
||||
)
|
||||
if not resp.ok:
|
||||
captured = stderr_capture.getvalue().strip()
|
||||
@@ -94,6 +489,7 @@ def _run_loop() -> None:
|
||||
)
|
||||
_write_response(resp, result or b"")
|
||||
except Exception as e:
|
||||
logger.exception("Preview worker error for %s", req.path)
|
||||
captured = stderr_capture.getvalue().strip()
|
||||
_write_response(
|
||||
PreviewResponse(ok=False, error=str(e), stderr=captured or None), b""
|
||||
@@ -106,9 +502,22 @@ def _run_loop() -> None:
|
||||
def main() -> None:
|
||||
# Configure all log output to stderr before any imports that may emit logs.
|
||||
logging.basicConfig(stream=sys.stderr, level=logging.INFO)
|
||||
try:
|
||||
config.load_config()
|
||||
logger.warning(
|
||||
"preview-worker config=%s master_secret=%s",
|
||||
config.conffile,
|
||||
config.config.secret,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("preview-worker failed to load config at startup")
|
||||
if len(sys.argv) > 1:
|
||||
_run_once()
|
||||
return
|
||||
# Eagerly import heavy modules before signalling readiness so the parent
|
||||
# does not hand us a request while we are still initialising.
|
||||
sys.stdout.buffer.write(b"\x01")
|
||||
sys.stdout.buffer.flush()
|
||||
_run_loop()
|
||||
|
||||
|
||||
|
||||
+15
-10
@@ -5,6 +5,8 @@ import sys
|
||||
import unicodedata
|
||||
from ipaddress import IPv6Address
|
||||
|
||||
from sanic.log import LOGGING_CONFIG_DEFAULTS
|
||||
|
||||
logger = logging.getLogger("cista.access")
|
||||
|
||||
_RESET = "\033[0m"
|
||||
@@ -96,10 +98,11 @@ def format_duration_ms(duration_ms: float) -> str:
|
||||
|
||||
|
||||
def _display_width(text: str) -> int:
|
||||
width = 0
|
||||
for char in text:
|
||||
width += 2 if unicodedata.east_asian_width(char) in {"F", "W"} else 1
|
||||
return width
|
||||
return sum(
|
||||
1 + (unicodedata.east_asian_width(c) in "FW")
|
||||
for c in text
|
||||
if unicodedata.category(c) != "Mn"
|
||||
)
|
||||
|
||||
|
||||
def _format_left(label: str) -> str:
|
||||
@@ -242,20 +245,24 @@ def configure_access_logging() -> None:
|
||||
|
||||
_LEVEL_EMOJI = {
|
||||
logging.DEBUG: "🔍",
|
||||
logging.INFO: "i",
|
||||
logging.INFO: "ℹ️", # noqa: RUF001
|
||||
logging.WARNING: "⚠️",
|
||||
logging.ERROR: "🛑",
|
||||
logging.CRITICAL: "🛑",
|
||||
}
|
||||
|
||||
|
||||
def _format_level_prefix(levelno: int) -> str:
|
||||
emoji = _LEVEL_EMOJI.get(levelno, "▪️")
|
||||
prefix = f"{emoji} "
|
||||
return prefix + (" " * max(0, 3 - _display_width(prefix)))
|
||||
|
||||
|
||||
class _EmojiFormatter(logging.Formatter):
|
||||
"""Compact formatter: emoji + message, no timestamp/level text/logger name."""
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
emoji = _LEVEL_EMOJI.get(record.levelno, "▪️")
|
||||
sep = " " if record.levelno in (logging.INFO, logging.WARNING) else " "
|
||||
return f"{emoji}{sep}{record.getMessage()}"
|
||||
return _format_level_prefix(record.levelno) + record.getMessage()
|
||||
|
||||
|
||||
def configure_main_logging() -> None:
|
||||
@@ -264,8 +271,6 @@ def configure_main_logging() -> None:
|
||||
Patches LOGGING_CONFIG_DEFAULTS so the formatter survives every dictConfig
|
||||
call Sanic makes during serve_single() / serve().
|
||||
"""
|
||||
from sanic.log import LOGGING_CONFIG_DEFAULTS
|
||||
|
||||
LOGGING_CONFIG_DEFAULTS["formatters"]["generic"] = {
|
||||
"class": "cista.sanic_logging._EmojiFormatter",
|
||||
}
|
||||
|
||||
+2
-2
@@ -6,12 +6,11 @@ from fastapi_vue.hostutil import parse_endpoint
|
||||
from sanic import Sanic
|
||||
|
||||
from cista import config, server80
|
||||
from cista.app import app
|
||||
|
||||
|
||||
def run(*, dev=False):
|
||||
"""Run Sanic main process that spawns worker processes to serve HTTP requests."""
|
||||
from .app import app
|
||||
|
||||
_url, opts = parse_listen(config.config.listen)
|
||||
# Silence Sanic's warning about running in production rather than debug
|
||||
os.environ["SANIC_IGNORE_PRODUCTION_WARNING"] = "1"
|
||||
@@ -36,6 +35,7 @@ def run(*, dev=False):
|
||||
|
||||
|
||||
def check_cert(certdir, domain):
|
||||
_ = domain
|
||||
if (certdir / "privkey.pem").exist() and (certdir / "fullchain.pem").exists():
|
||||
return
|
||||
# Certificate provisioning is external; files must exist before startup.
|
||||
|
||||
@@ -6,6 +6,7 @@ app = Sanic("server80")
|
||||
# Send all HTTP users to HTTPS
|
||||
@app.exception(exceptions.NotFound, exceptions.MethodNotSupported)
|
||||
def redirect_everything_else(request, exception):
|
||||
_ = exception
|
||||
server, path = request.server_name, request.path
|
||||
if server and path.startswith("/"):
|
||||
return response.redirect(f"https://{server}{path}", status=308)
|
||||
@@ -15,6 +16,7 @@ def redirect_everything_else(request, exception):
|
||||
# ACME challenge for LetsEncrypt
|
||||
@app.get("/.well-known/acme-challenge/<challenge>")
|
||||
async def letsencrypt(request, challenge):
|
||||
_ = request
|
||||
try:
|
||||
return response.text(acme_challenges[challenge])
|
||||
except KeyError:
|
||||
|
||||
+8
-1
@@ -36,7 +36,7 @@ def get(request):
|
||||
def create(request, res, username, **kwargs):
|
||||
_purge_expired()
|
||||
token = _token()
|
||||
_sessions[token] = {"exp": int(time()) + max_age, "username": username, **kwargs}
|
||||
put(token, username, **kwargs)
|
||||
secure = request.scheme == "https"
|
||||
res.cookies.add_cookie(
|
||||
SESSION_COOKIE_NAME,
|
||||
@@ -49,10 +49,17 @@ def create(request, res, username, **kwargs):
|
||||
|
||||
|
||||
def delete(request, res):
|
||||
token = request.cookies.get(SESSION_COOKIE_NAME)
|
||||
if token is not None:
|
||||
_sessions.pop(token, None)
|
||||
secure = request.scheme == "https"
|
||||
res.cookies.delete_cookie(SESSION_COOKIE_NAME, host_prefix=secure)
|
||||
|
||||
|
||||
def put(token: str, username: str, **kwargs) -> None:
|
||||
_sessions[token] = {"exp": int(time()) + max_age, "username": username, **kwargs}
|
||||
|
||||
|
||||
def flash(res, message: str | None):
|
||||
if message is None:
|
||||
res.cookies.delete_cookie("message")
|
||||
|
||||
+5
-2
@@ -107,10 +107,11 @@ async def validate_sso_request(request, *, perm: str = "cista:login") -> dict |
|
||||
request.ctx.sso_user = data
|
||||
if "set-cookie" in response.headers:
|
||||
request.ctx.sso_cookies = response.headers.get_list("set-cookie")
|
||||
return data
|
||||
except Exception:
|
||||
request.ctx.sso_user = {}
|
||||
return {}
|
||||
else:
|
||||
return data
|
||||
|
||||
try:
|
||||
error_data = response.json()
|
||||
@@ -257,7 +258,7 @@ async def proxy_auth_request(request):
|
||||
method=request.method,
|
||||
url=url,
|
||||
headers=headers,
|
||||
content=request.body if request.body else None,
|
||||
content=request.body or None,
|
||||
) as response:
|
||||
raw_content = b"".join([chunk async for chunk in response.aiter_raw()])
|
||||
|
||||
@@ -348,6 +349,7 @@ bp = Blueprint("sso", url_prefix="/auth")
|
||||
@bp.websocket("/ws/<path:path>")
|
||||
async def auth_websocket_proxy(request, ws, path=""):
|
||||
"""Proxy WebSocket connections to the auth backend."""
|
||||
_ = path
|
||||
await proxy_auth_websocket(request, ws)
|
||||
|
||||
|
||||
@@ -362,6 +364,7 @@ async def auth_websocket_proxy_root(request, ws):
|
||||
)
|
||||
async def auth_proxy(request, path=""):
|
||||
"""Proxy all auth requests to the auth backend."""
|
||||
_ = path
|
||||
return await proxy_auth_request(request)
|
||||
|
||||
|
||||
|
||||
@@ -60,7 +60,7 @@ def websocket_wrapper(handler):
|
||||
@wraps(handler)
|
||||
async def wrapper(request, ws, *args, **kwargs):
|
||||
username = getattr(request.ctx, "username", None)
|
||||
extra = username if username else None
|
||||
extra = username or None
|
||||
start = time.perf_counter()
|
||||
ws_id = log_ws_open(request, extra=extra)
|
||||
close_extra = None
|
||||
|
||||
@@ -23,7 +23,7 @@ class AsyncLink:
|
||||
@property
|
||||
def to_sync(self):
|
||||
"""Yield SyncRequests from async caller when called from worker thread."""
|
||||
while (req := self._await(self._get())) is not None:
|
||||
while (req := self.await_sync(self._get())) is not None:
|
||||
yield SyncRequest(self, req)
|
||||
|
||||
async def _get(self):
|
||||
@@ -33,7 +33,7 @@ class AsyncLink:
|
||||
self.queue.task_done()
|
||||
return ret
|
||||
|
||||
def _await(self, coro):
|
||||
def await_sync(self, coro):
|
||||
"""Run coroutine in main thread and return result; called from worker."""
|
||||
return asyncio.run_coroutine_threadsafe(coro, self.loop).result()
|
||||
|
||||
@@ -87,9 +87,9 @@ class SyncRequest:
|
||||
def set_result(self, value):
|
||||
"""Set result value; mark as done."""
|
||||
self.done = True
|
||||
self.alink._await(set_result(self.future, value))
|
||||
self.alink.await_sync(set_result(self.future, value))
|
||||
|
||||
def set_exception(self, exc):
|
||||
"""Set exception; mark as done."""
|
||||
self.done = True
|
||||
self.alink._await(set_result(self.future, exception=exc))
|
||||
self.alink.await_sync(set_result(self.future, exception=exc))
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
import hmac
|
||||
import re
|
||||
from typing import Protocol
|
||||
from unicodedata import normalize
|
||||
|
||||
import argon2
|
||||
|
||||
_argon = argon2.PasswordHasher()
|
||||
_droppyhash = re.compile(r"^([a-f0-9]{64})\$([a-f0-9]{8})$")
|
||||
|
||||
|
||||
class SupportsHash(Protocol):
|
||||
hash: str
|
||||
|
||||
|
||||
def normalize_secret(value: str) -> bytes:
|
||||
return normalize("NFC", value).strip().encode()
|
||||
|
||||
|
||||
def verify_hash(user_hash: str, *, username: str, password: str) -> bool:
|
||||
"""Verify password hash and return whether the stored hash should be upgraded."""
|
||||
if not user_hash:
|
||||
raise ValueError("Account disabled")
|
||||
|
||||
normalized_username = normalize_secret(username)
|
||||
normalized_password = normalize_secret(password)
|
||||
|
||||
if (match := _droppyhash.match(user_hash)) is not None:
|
||||
expected_hash, salt = match.groups()
|
||||
computed_hash = hmac.digest(
|
||||
normalized_password + salt.encode() + normalized_username,
|
||||
b"",
|
||||
"sha256",
|
||||
).hex()
|
||||
if not hmac.compare_digest(expected_hash, computed_hash):
|
||||
raise ValueError("Invalid password")
|
||||
return True
|
||||
|
||||
try:
|
||||
_argon.verify(user_hash, normalized_password)
|
||||
except Exception:
|
||||
raise ValueError("Invalid password") from None
|
||||
return _argon.check_needs_rehash(user_hash)
|
||||
|
||||
|
||||
def set_password(user: SupportsHash, password: str) -> None:
|
||||
user.hash = _argon.hash(normalize_secret(password))
|
||||
+9
-5
@@ -9,6 +9,7 @@ from os import stat_result
|
||||
from pathlib import Path, PurePosixPath
|
||||
from stat import S_ISDIR, S_ISREG
|
||||
|
||||
import inotify.adapters
|
||||
import msgspec
|
||||
from natsort import humansorted, natsort_keygen, ns
|
||||
from sanic.log import logger
|
||||
@@ -30,6 +31,7 @@ if sys.platform == "win32":
|
||||
|
||||
def get_allocated_size(path: Path, st: stat_result) -> int:
|
||||
"""Get actual disk allocation on Windows using GetCompressedFileSizeW."""
|
||||
_ = st
|
||||
high = wintypes.DWORD()
|
||||
low = GetCompressedFileSizeW(str(path), ctypes.byref(high))
|
||||
if low == INVALID_FILE_SIZE and ctypes.get_last_error() != 0:
|
||||
@@ -40,6 +42,7 @@ else:
|
||||
|
||||
def get_allocated_size(path: Path, st: stat_result) -> int:
|
||||
"""Get actual disk allocation on Unix using st_blocks."""
|
||||
_ = path
|
||||
# st_blocks is in 512-byte units
|
||||
return st.st_blocks * 512
|
||||
|
||||
@@ -48,6 +51,10 @@ pubsub = {}
|
||||
sortkey = natsort_keygen(alg=ns.LOCALE)
|
||||
|
||||
|
||||
class FormatUpdateLoopError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class State:
|
||||
def __init__(self):
|
||||
self.lock = threading.RLock()
|
||||
@@ -244,6 +251,7 @@ def update_root(loop):
|
||||
def update_path(rootmod: list[FileEntry], relpath: PurePosixPath, loop):
|
||||
"""Called on FS updates, check the filesystem and broadcast any changes."""
|
||||
new = walk(relpath)
|
||||
_ = loop
|
||||
obegin, old = treeget(rootmod, relpath)
|
||||
|
||||
if old == new:
|
||||
@@ -300,7 +308,7 @@ def format_update(old, new):
|
||||
logger.error(
|
||||
f"format_update potential infinite loop! iteration={iteration_count}, oidx={oidx}, nidx={nidx}"
|
||||
)
|
||||
raise Exception(
|
||||
raise FormatUpdateLoopError(
|
||||
f"format_update infinite loop detected at iteration {iteration_count}"
|
||||
)
|
||||
|
||||
@@ -641,8 +649,6 @@ def watcher(loop):
|
||||
modified_flags = frozenset()
|
||||
|
||||
if use_inotify:
|
||||
import inotify.adapters
|
||||
|
||||
modified_flags = frozenset(
|
||||
(
|
||||
"IN_CREATE",
|
||||
@@ -657,8 +663,6 @@ def watcher(loop):
|
||||
|
||||
while not stop_event.is_set():
|
||||
if use_inotify:
|
||||
import inotify.adapters
|
||||
|
||||
inotify_tree = inotify.adapters.InotifyTree(rootpath.as_posix())
|
||||
|
||||
# Initialize the tree from filesystem
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
services:
|
||||
onlyoffice:
|
||||
build:
|
||||
context: ./docker/onlyoffice-converter-patch
|
||||
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:
|
||||
@@ -0,0 +1,70 @@
|
||||
# Patched OnlyOffice Document Server with configurable converter worker count.
|
||||
#
|
||||
# The Community Edition hardcodes the document converter to 1 worker,
|
||||
# which creates a severe bottleneck under concurrent load.
|
||||
# This image patches the open-source license.js to spawn a configurable
|
||||
# number of converter workers (default 8).
|
||||
#
|
||||
# Build:
|
||||
# docker build -t onlyoffice-cista docker/onlyoffice-converter-patch
|
||||
#
|
||||
# Run:
|
||||
# docker run -d -p 8988:80 \
|
||||
# -e WORKERS=16 \
|
||||
# -e JWT_SECRET=your-strong-secret \
|
||||
# --name onlyoffice onlyoffice-cista
|
||||
#
|
||||
# JWT:
|
||||
# Set JWT_SECRET to the same value you pass to Cista as ONLYOFFICE_JWT_SECRET.
|
||||
# OnlyOffice will enable token validation automatically.
|
||||
#
|
||||
# The ONLYOFFICE_VERSION build arg lets you target a specific release.
|
||||
|
||||
ARG ONLYOFFICE_VERSION=9.3.1
|
||||
|
||||
FROM onlyoffice/documentserver:${ONLYOFFICE_VERSION}
|
||||
|
||||
# Prevent interactive apt prompts
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Install Node.js, npm, and git so we can run the FileConverter from source.
|
||||
RUN apt-get update -qq && \
|
||||
apt-get install -y -qq --no-install-recommends \
|
||||
nodejs \
|
||||
npm \
|
||||
git \
|
||||
ca-certificates && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Clone the open-source server components (shallow, ~15 MB).
|
||||
# The master branch is used because the Linux/web tags are not published
|
||||
# in the server repo; the license.js file has been stable for years.
|
||||
RUN git clone --depth 1 https://github.com/ONLYOFFICE/server.git /opt/oo-server
|
||||
|
||||
# Patch license.js so the converter worker count is read from an env var
|
||||
# instead of being hardcoded to 1.
|
||||
RUN sed -i \
|
||||
's/count: 1,/count: parseInt(process.env.WORKERS, 10) || 8,/' \
|
||||
/opt/oo-server/Common/sources/license.js
|
||||
|
||||
# Install npm dependencies for the modules the FileConverter touches.
|
||||
# DocService deps are also needed because converter.js pulls in baseConnector.
|
||||
RUN cd /opt/oo-server/Common && npm ci --no-audit --no-fund
|
||||
RUN cd /opt/oo-server/FileConverter && npm ci --no-audit --no-fund
|
||||
RUN cd /opt/oo-server/DocService && npm ci --no-audit --no-fund
|
||||
|
||||
# Back up the compiled pkg binary and replace it with our wrapper.
|
||||
RUN mv /var/www/onlyoffice/documentserver/server/FileConverter/converter \
|
||||
/var/www/onlyoffice/documentserver/server/FileConverter/converter.orig
|
||||
|
||||
COPY converter-wrapper.sh /var/www/onlyoffice/documentserver/server/FileConverter/converter
|
||||
RUN chmod +x /var/www/onlyoffice/documentserver/server/FileConverter/converter
|
||||
|
||||
# Default worker count (override at runtime with -e WORKERS=16).
|
||||
ENV WORKERS=8
|
||||
|
||||
# Use our custom entrypoint to persist the env var to a file that the
|
||||
# non-root converter process (user=ds) can read.
|
||||
COPY entrypoint.sh /app/ds/run-document-server-patched.sh
|
||||
RUN chmod +x /app/ds/run-document-server-patched.sh
|
||||
ENTRYPOINT ["/app/ds/run-document-server-patched.sh"]
|
||||
@@ -0,0 +1,19 @@
|
||||
#!/bin/bash
|
||||
# Wrapper that runs the OnlyOffice FileConverter from patched Node.js source.
|
||||
# Replaces the compiled pkg binary shipped with the Community Edition.
|
||||
|
||||
# The env var is not passed through supervisor to the 'ds' user, so we read
|
||||
# it from a file written by the custom entrypoint.
|
||||
if [ -z "${WORKERS}" ] && [ -r /tmp/oo-converter-workers.txt ]; then
|
||||
export WORKERS=$(cat /tmp/oo-converter-workers.txt)
|
||||
fi
|
||||
|
||||
cd /opt/oo-server/FileConverter || exit 1
|
||||
|
||||
export NODE_ENV=production-linux
|
||||
export NODE_CONFIG_DIR=/etc/onlyoffice/documentserver
|
||||
export NODE_DISABLE_COLORS=1
|
||||
export APPLICATION_NAME=onlyoffice
|
||||
export LD_LIBRARY_PATH=/var/www/onlyoffice/documentserver/server/FileConverter/bin
|
||||
|
||||
exec node sources/convertermaster.js "$@"
|
||||
@@ -0,0 +1,8 @@
|
||||
#!/bin/bash
|
||||
# Custom entrypoint that persists WORKERS to a file readable by
|
||||
# the non-root user that supervisor uses to run the converter.
|
||||
|
||||
echo "${WORKERS:-8}" > /tmp/oo-converter-workers.txt
|
||||
chmod 644 /tmp/oo-converter-workers.txt
|
||||
|
||||
exec /app/ds/run-document-server.sh "$@"
|
||||
@@ -20,6 +20,7 @@
|
||||
<script setup lang="ts">
|
||||
import { Play as PlayIcon, Spinner as SpinnerIcon } from '@/assets/svg'
|
||||
import type { Doc } from '@/repositories/Document'
|
||||
import { useMainStore } from '@/stores/main'
|
||||
import { computed, ref } from 'vue'
|
||||
|
||||
const aud = ref<HTMLAudioElement | null>(null)
|
||||
@@ -125,64 +126,22 @@ defineExpose({
|
||||
media
|
||||
})
|
||||
|
||||
const video = () => ['mkv', 'mp4', 'webm', 'mov', 'avi'].includes(props.doc.ext)
|
||||
const audio = () => ['mp3', 'flac', 'ogg', 'aac'].includes(props.doc.ext)
|
||||
const archive = () =>
|
||||
['zip', 'tar', 'gz', 'bz2', 'xz', '7z', 'rar'].includes(props.doc.ext)
|
||||
const video = () => props.doc.video
|
||||
const audio = () => props.doc.audio
|
||||
const archive = () => props.doc.archive
|
||||
const docs = () => props.doc.document
|
||||
// image = requires server-side preview (browsers cannot display it natively)
|
||||
// img = browser-viewable image that can be used directly in an <img> tag
|
||||
const image = () => props.doc.image
|
||||
const print = () => props.doc.print
|
||||
const showProgress = () => !props.doc.complete && (preview() || props.doc.img)
|
||||
const preview = () =>
|
||||
[
|
||||
'bmp',
|
||||
'ico',
|
||||
'tif',
|
||||
'tiff',
|
||||
'heic',
|
||||
'heif',
|
||||
'pdf',
|
||||
'epub',
|
||||
'mobi',
|
||||
// Documents
|
||||
'doc',
|
||||
'dot',
|
||||
'docx',
|
||||
'docm',
|
||||
'dotx',
|
||||
'dotm',
|
||||
'rtf',
|
||||
'odt',
|
||||
'ott',
|
||||
'txt',
|
||||
'md',
|
||||
'mhtml',
|
||||
'mht',
|
||||
'html',
|
||||
'htm',
|
||||
'xml',
|
||||
'wps',
|
||||
'wri',
|
||||
// Spreadsheets
|
||||
'xls',
|
||||
'xlsx',
|
||||
'xlsm',
|
||||
'xlsb',
|
||||
'xltx',
|
||||
'xltm',
|
||||
'ods',
|
||||
'ots',
|
||||
'csv',
|
||||
// Presentations
|
||||
'ppt',
|
||||
'pptx',
|
||||
'pptm',
|
||||
'pps',
|
||||
'ppsx',
|
||||
'pot',
|
||||
'potx',
|
||||
'odp',
|
||||
'otp'
|
||||
].includes(props.doc.ext) ||
|
||||
(props.doc.size > 500000 &&
|
||||
['avif', 'webp', 'png', 'jpg', 'jpeg'].includes(props.doc.ext))
|
||||
const preview = () => {
|
||||
const store = useMainStore()
|
||||
return (
|
||||
!(store.server.office_previews === false && docs()) &&
|
||||
(image() || print() || (props.doc.img && props.doc.size > 500000))
|
||||
)
|
||||
}
|
||||
</script>
|
||||
|
||||
<style scoped>
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import { formatSize, formatUnixDate } from '@/utils'
|
||||
import { useMainStore } from '@/stores/main'
|
||||
import { FILE_TYPES, formatSize, formatUnixDate } from '@/utils'
|
||||
|
||||
export type FUID = string
|
||||
|
||||
@@ -63,86 +64,53 @@ export class Doc {
|
||||
return this.url.replace(/^\/#/, '')
|
||||
}
|
||||
get img(): boolean {
|
||||
// Folders cannot be images
|
||||
if (this.dir) return false
|
||||
return [
|
||||
'jpg',
|
||||
'jpeg',
|
||||
'png',
|
||||
'gif',
|
||||
'webp',
|
||||
'avif',
|
||||
'heic',
|
||||
'heif',
|
||||
'svg'
|
||||
].includes(this.ext)
|
||||
return (
|
||||
!this.dir && (FILE_TYPES.imageBrowser as readonly string[]).includes(this.ext)
|
||||
)
|
||||
}
|
||||
get video(): boolean {
|
||||
return (FILE_TYPES.video as readonly string[]).includes(this.ext)
|
||||
}
|
||||
get audio(): boolean {
|
||||
return (FILE_TYPES.audio as readonly string[]).includes(this.ext)
|
||||
}
|
||||
get archive(): boolean {
|
||||
return (FILE_TYPES.archive as readonly string[]).includes(this.ext)
|
||||
}
|
||||
get document(): boolean {
|
||||
return (FILE_TYPES.document as readonly string[]).includes(this.ext)
|
||||
}
|
||||
// Images that require server-side preview (browsers cannot display them natively)
|
||||
get image(): boolean {
|
||||
return (FILE_TYPES.image as readonly string[]).includes(this.ext)
|
||||
}
|
||||
get print(): boolean {
|
||||
return (FILE_TYPES.print as readonly string[]).includes(this.ext)
|
||||
}
|
||||
get complete(): boolean {
|
||||
return !this.ghost && (this.dir || this.size <= this.allocated)
|
||||
}
|
||||
get previewable(): boolean {
|
||||
// Folders cannot be previewable
|
||||
if (this.dir) return false
|
||||
if (this.img) return true
|
||||
// Not a comprehensive list, but good enough for now
|
||||
return [
|
||||
'mp4',
|
||||
'mkv',
|
||||
'webm',
|
||||
'ogg',
|
||||
'mp3',
|
||||
'flac',
|
||||
'aac',
|
||||
'pdf',
|
||||
// Documents
|
||||
'doc',
|
||||
'dot',
|
||||
'docx',
|
||||
'docm',
|
||||
'dotx',
|
||||
'dotm',
|
||||
'rtf',
|
||||
'odt',
|
||||
'ott',
|
||||
'txt',
|
||||
'md',
|
||||
'mhtml',
|
||||
'mht',
|
||||
'html',
|
||||
'htm',
|
||||
'xml',
|
||||
'wps',
|
||||
'wri',
|
||||
// Spreadsheets
|
||||
'xls',
|
||||
'xlsx',
|
||||
'xlsm',
|
||||
'xlsb',
|
||||
'xltx',
|
||||
'xltm',
|
||||
'ods',
|
||||
'ots',
|
||||
'csv',
|
||||
// Presentations
|
||||
'ppt',
|
||||
'pptx',
|
||||
'pptm',
|
||||
'pps',
|
||||
'ppsx',
|
||||
'pot',
|
||||
'potx',
|
||||
'odp',
|
||||
'otp'
|
||||
].includes(this.ext)
|
||||
return (
|
||||
this.img ||
|
||||
this.video ||
|
||||
this.audio ||
|
||||
this.image ||
|
||||
this.print ||
|
||||
(this.document && useMainStore().server.office_previews !== false)
|
||||
)
|
||||
}
|
||||
get previewurl(): string {
|
||||
if (!this.complete || !this.previewable) return ''
|
||||
return this.url.replace(/^\/files/, '/preview')
|
||||
return !this.complete || !this.previewable
|
||||
? ''
|
||||
: this.url.replace(/^\/files/, '/preview')
|
||||
}
|
||||
get ext(): string {
|
||||
const dotIndex = this.name.lastIndexOf('.')
|
||||
if (dotIndex === -1 || dotIndex === this.name.length - 1) return ''
|
||||
return this.name.slice(dotIndex + 1).toLowerCase()
|
||||
return dotIndex === -1 || dotIndex === this.name.length - 1
|
||||
? ''
|
||||
: this.name.slice(dotIndex + 1).toLowerCase()
|
||||
}
|
||||
}
|
||||
export type errorEvent = {
|
||||
|
||||
@@ -79,7 +79,11 @@ export const useMainStore = defineStore('main', {
|
||||
connected: false,
|
||||
authInProgress: false,
|
||||
cursor: '' as string,
|
||||
server: {} as Record<string, any> & { public?: boolean; paskia?: boolean },
|
||||
server: {} as Record<string, any> & {
|
||||
public?: boolean
|
||||
paskia?: boolean
|
||||
office_previews?: boolean
|
||||
},
|
||||
dialog: '' as '' | 'settings' | 'usermgmt' | 'accessdenied' | 'tokens',
|
||||
uprogress: {} as any,
|
||||
dprogress: {} as any,
|
||||
|
||||
+18
-12
@@ -69,23 +69,29 @@ export function getFileExtension(filename: string) {
|
||||
}
|
||||
return filename.slice(dotIndex + 1)
|
||||
}
|
||||
interface FileTypes {
|
||||
[key: string]: string[]
|
||||
}
|
||||
|
||||
const filetypes: FileTypes = {
|
||||
export const FILE_TYPES = {
|
||||
video: ['avi', 'mkv', 'mov', 'mp4', 'webm'],
|
||||
image: ['avif', 'gif', 'jpg', 'jpeg', 'png', 'webp', 'svg'],
|
||||
pdf: ['pdf']
|
||||
}
|
||||
audio: ['mp3', 'flac', 'ogg', 'aac'],
|
||||
archive: ['zip', 'tar', 'gz', 'bz2', 'xz', '7z', 'rar'],
|
||||
document: ['doc', 'docx', 'xls', 'xlsx', 'ppt', 'pptx', 'odt', 'ods', 'odp', 'rtf'],
|
||||
imageBrowser: ['avif', 'gif', 'jpg', 'jpeg', 'png', 'webp', 'svg'],
|
||||
// Images that require server-side preview (browsers cannot display them natively)
|
||||
image: ['bmp', 'heic', 'heif', 'ico', 'tif', 'tiff'],
|
||||
print: ['epub', 'mobi', 'pdf']
|
||||
} as const
|
||||
|
||||
export function getFileType(name: string): string {
|
||||
export type FileCategory = keyof typeof FILE_TYPES
|
||||
|
||||
export function getFileType(name: string): FileCategory | 'unknown' {
|
||||
const dotIndex = name.lastIndexOf('.')
|
||||
if (dotIndex === -1 || dotIndex === name.length - 1) return 'unknown'
|
||||
const ext = name.slice(dotIndex + 1).toLowerCase()
|
||||
return (
|
||||
Object.keys(filetypes).find(type => filetypes[type]!.includes(ext)) || 'unknown'
|
||||
)
|
||||
for (const category of Object.keys(FILE_TYPES) as FileCategory[]) {
|
||||
if ((FILE_TYPES[category] as readonly string[]).includes(ext)) {
|
||||
return category
|
||||
}
|
||||
}
|
||||
return 'unknown'
|
||||
}
|
||||
|
||||
// Prebuilt for fast & consistent sorting
|
||||
|
||||
+4
-11
@@ -77,8 +77,8 @@ docs = [
|
||||
source = "vcs"
|
||||
|
||||
[tool.hatch.build]
|
||||
artifacts = ["cista/frontend-build"]
|
||||
targets.sdist.hooks.custom.path = "scripts/fastapi-vue/build-frontend.py"
|
||||
artifacts = ["cista/frontend-build", "cista/docker"]
|
||||
targets.sdist.hooks.custom.path = "scripts/fastapi-vue/buildhook.py"
|
||||
targets.sdist.include = [
|
||||
"/cista",
|
||||
]
|
||||
@@ -130,7 +130,6 @@ ignore = [
|
||||
"ANN202", # legacy codebase: no full runtime annotation coverage yet
|
||||
"ANN204", # legacy codebase: no full runtime annotation coverage yet
|
||||
"ANN205", # legacy codebase: no full runtime annotation coverage yet
|
||||
"ARG001", # framework and callback signatures commonly require unused args
|
||||
"BLE001", # broad catch remains in boundary/proxy/error-handling paths
|
||||
"C901", # legacy complexity; keep other correctness rules enabled
|
||||
"D100", # legacy docs not yet standardized
|
||||
@@ -152,22 +151,16 @@ ignore = [
|
||||
"EM101", # exception-message style; low signal for this project
|
||||
"EM102", # exception-message style; low signal for this project
|
||||
"INP001", # scripts folder intentionally lacks package markers
|
||||
"PLC0415", # lazy imports used to avoid startup/circular import issues
|
||||
"PLR0911", # legacy complexity; keep other correctness rules enabled
|
||||
"PLR0912", # legacy complexity; keep other correctness rules enabled
|
||||
"PLR0913", # legacy complexity; keep other correctness rules enabled
|
||||
"PLR0915", # legacy complexity; keep other correctness rules enabled
|
||||
"PLR2004", # legacy comparisons use inline constants
|
||||
"PLR2004", # we like magic numbers (don't remove this suppression)
|
||||
"PLW0603", # module-level shared state exists in server runtime code
|
||||
"SLF001", # cohesive modules occasionally need private-member access
|
||||
"TRY002", # exception-class strictness too noisy on legacy handlers
|
||||
"TRY003", # exception-message strictness too noisy on legacy handlers
|
||||
"TRY004", # type-check strictness too noisy on legacy handlers
|
||||
"TRY300", # stylistic try/else preference
|
||||
"TRY301", # stylistic raise-in-try preference
|
||||
]
|
||||
isort.known-first-party = ["cista"]
|
||||
per-file-ignores."tests/*" = ["S", "ANN", "D", "INP", "PLR2004"]
|
||||
per-file-ignores."tests/*" = ["S", "ANN", "D", "INP", "PLR2004", "ARG001"]
|
||||
per-file-ignores."scripts/*" = ["T20"]
|
||||
|
||||
[dependency-groups]
|
||||
|
||||
+10
-9
@@ -16,15 +16,16 @@ Environment:
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import contextlib
|
||||
import os
|
||||
import sys
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
|
||||
# Import devutil from scripts/fastapi-vue (not a package, so we adjust sys.path)
|
||||
sys.path.insert(0, str(Path(__file__).with_name("fastapi-vue")))
|
||||
from devutil import ( # type: ignore[import-not-found]
|
||||
ProcessGroup,
|
||||
check_ports_free,
|
||||
logger,
|
||||
ready,
|
||||
setup_vite,
|
||||
@@ -33,7 +34,9 @@ from devutil import ( # type: ignore[import-not-found]
|
||||
from cista import config
|
||||
from cista.serve import parse_listen
|
||||
|
||||
DEFAULT_VITE_PORT = 8989
|
||||
DEFAULT_BACKEND_PORT = 8999
|
||||
HEALTH = "/api/health?from=devserver.py"
|
||||
|
||||
|
||||
def setup_sanic_backend(
|
||||
@@ -64,7 +67,7 @@ async def run_devserver(
|
||||
logger.warning("Frontend source not found at %s", front)
|
||||
raise SystemExit(1)
|
||||
|
||||
_frontend_url, npm_install, vite = setup_vite(frontend or "")
|
||||
frontend_url, npm_install, vite = setup_vite(frontend or "", DEFAULT_VITE_PORT)
|
||||
backend_url, sanic_cmd = setup_sanic_backend(backend, extra_args)
|
||||
|
||||
# Tell vite where to proxy API requests
|
||||
@@ -72,19 +75,17 @@ async def run_devserver(
|
||||
|
||||
async with ProcessGroup() as pg:
|
||||
install_proc = await pg.spawn(*npm_install, cwd=str(front))
|
||||
await asyncio.sleep(0.2) # reduce message overlap
|
||||
await check_ports_free(frontend_url, backend_url)
|
||||
await pg.spawn(*sanic_cmd, cwd=str(reporoot))
|
||||
|
||||
# Wait for both install and backend to be ready
|
||||
async with asyncio.TaskGroup() as tg:
|
||||
tg.create_task(pg.wait(install_proc))
|
||||
tg.create_task(ready(backend_url, path="/api/health?from=devserver.py"))
|
||||
# Wait for dependencies to be installed and backend to accept requests
|
||||
await pg.wait(install_proc, ready(backend_url, path=HEALTH))
|
||||
|
||||
# Start Vite dev server (ProcessGroup waits for any exit, then terminates others)
|
||||
await pg.spawn(*vite, cwd=str(front))
|
||||
|
||||
|
||||
def main():
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Run Vite and Cista (Sanic) development servers",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
@@ -102,7 +103,7 @@ def main():
|
||||
help="Cista backend endpoint (default: from config, or :8999)",
|
||||
)
|
||||
args, unknown = parser.parse_known_args()
|
||||
with contextlib.suppress(KeyboardInterrupt):
|
||||
with suppress(KeyboardInterrupt):
|
||||
asyncio.run(run_devserver(args.listen, args.backend, unknown))
|
||||
|
||||
|
||||
|
||||
@@ -1,17 +0,0 @@
|
||||
"""Hatch build hook for building Vue frontend during package build."""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from hatchling.builders.hooks.plugin.interface import (
|
||||
BuildHookInterface, # type: ignore[import-not-found]
|
||||
)
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
from buildutil import build
|
||||
|
||||
|
||||
class CustomBuildHook(BuildHookInterface):
|
||||
def initialize(self, version, build_data):
|
||||
super().initialize(version, build_data)
|
||||
build("frontend")
|
||||
@@ -0,0 +1,18 @@
|
||||
"""Hatch build hook for building Vue frontend during package build."""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from hatchling.builders.hooks.plugin.interface import BuildHookInterface
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
from buildutil import build
|
||||
|
||||
|
||||
class CustomBuildHook(BuildHookInterface): # type: ignore[misc]
|
||||
"""Hatch build hook that builds Vue frontend during package build."""
|
||||
|
||||
def initialize(self, version: str, build_data: dict) -> None: # type: ignore[override]
|
||||
"""Build frontend before package is built."""
|
||||
super().initialize(version, build_data)
|
||||
build("frontend")
|
||||
@@ -7,13 +7,15 @@ import shutil
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
MIN_NODE_VERSION = 20
|
||||
|
||||
|
||||
class _PrefixFormatter(logging.Formatter):
|
||||
"""Formatter that adds prefix based on log level."""
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
if record.levelno >= logging.WARNING:
|
||||
return f"┃ ⚠️ {record.getMessage()}"
|
||||
return f"⚠️ {record.getMessage()}"
|
||||
return record.getMessage()
|
||||
|
||||
|
||||
@@ -41,74 +43,108 @@ def _check_node_version(node_path: str) -> None:
|
||||
match = re.match(r"v(\d+)", version_str)
|
||||
if match:
|
||||
major_version = int(match.group(1))
|
||||
if major_version >= 20:
|
||||
if major_version >= MIN_NODE_VERSION:
|
||||
return
|
||||
raise RuntimeError(
|
||||
f"Node.js {version_str} found, but v20+ required (install with nvm)"
|
||||
)
|
||||
msg = f"Node.js {version_str} found, but v{MIN_NODE_VERSION}+ required"
|
||||
raise RuntimeError(msg)
|
||||
except (subprocess.CalledProcessError, FileNotFoundError, ValueError):
|
||||
pass
|
||||
raise RuntimeError("Could not determine Node.js version")
|
||||
msg = "Could not determine Node.js version"
|
||||
raise RuntimeError(msg)
|
||||
|
||||
|
||||
def _validate_npm_runtime(tool: str) -> bool:
|
||||
"""Validate npm runtime by checking Node.js version. Returns True if valid."""
|
||||
node_path = shutil.which("node", path=str(Path(tool).parent))
|
||||
if node_path is None:
|
||||
return False
|
||||
try:
|
||||
_check_node_version(node_path)
|
||||
except RuntimeError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _find_runtime_from_env(options: list[str]) -> tuple[str, str] | None:
|
||||
"""Find runtime specified by JS_RUNTIME environment variable."""
|
||||
js_runtime_env = os.environ.get("JS_RUNTIME")
|
||||
if not js_runtime_env:
|
||||
return None
|
||||
|
||||
js_runtime = js_runtime_env
|
||||
js_path = Path(js_runtime)
|
||||
runtime_name = js_path.name
|
||||
|
||||
# Map node to npm
|
||||
if runtime_name == "node":
|
||||
runtime_name = "npm"
|
||||
js_runtime = str(js_path.parent / "npm") if js_path.parent.name else "npm"
|
||||
|
||||
for option in options:
|
||||
if option != runtime_name and not runtime_name.startswith(option):
|
||||
continue
|
||||
|
||||
tool = shutil.which(js_runtime)
|
||||
if tool is None:
|
||||
msg = f"JS_RUNTIME={js_runtime_env}: {option} not found"
|
||||
raise RuntimeError(msg)
|
||||
|
||||
if option == "npm":
|
||||
node_path = shutil.which("node", path=str(Path(tool).parent))
|
||||
if node_path is None:
|
||||
msg = f"JS_RUNTIME={js_runtime_env}: node not found"
|
||||
raise RuntimeError(msg)
|
||||
_check_node_version(node_path)
|
||||
|
||||
return tool, option
|
||||
|
||||
msg = f"JS_RUNTIME={js_runtime_env} not recognized"
|
||||
raise RuntimeError(msg)
|
||||
|
||||
|
||||
def _auto_detect_runtime(options: list[str]) -> tuple[str, str]:
|
||||
"""Auto-detect JavaScript runtime from available options."""
|
||||
node_version_error: RuntimeError | None = None
|
||||
|
||||
for option in options:
|
||||
tool = shutil.which(option)
|
||||
if not tool:
|
||||
continue
|
||||
|
||||
if option == "npm" and not _validate_npm_runtime(tool):
|
||||
try:
|
||||
node_path = shutil.which("node", path=str(Path(tool).parent))
|
||||
if node_path:
|
||||
_check_node_version(node_path)
|
||||
except RuntimeError as e:
|
||||
node_version_error = e
|
||||
continue
|
||||
|
||||
return tool, option
|
||||
|
||||
if node_version_error:
|
||||
raise node_version_error
|
||||
msg = "Node.js (v20+), Deno or Bun is required but none was found"
|
||||
raise RuntimeError(msg)
|
||||
|
||||
|
||||
def find_js_runtime() -> tuple[str, str]:
|
||||
"""Find a JavaScript runtime from JS_RUNTIME env or auto-detect.
|
||||
|
||||
Returns (tool_path, tool_name) where tool_name is "deno", "npm", or "bun".
|
||||
Raises JSRuntimeError if no suitable runtime is found.
|
||||
Raises RuntimeError if no suitable runtime is found.
|
||||
"""
|
||||
options = ["npm", "deno", "bun"]
|
||||
node_version_error: RuntimeError | None = None
|
||||
|
||||
# Check for JS_RUNTIME environment variable
|
||||
if js_runtime_env := os.environ.get("JS_RUNTIME"):
|
||||
js_runtime = js_runtime_env
|
||||
js_path = Path(js_runtime)
|
||||
runtime_name = js_path.name
|
||||
# Map node to npm
|
||||
if runtime_name == "node":
|
||||
runtime_name = "npm"
|
||||
js_runtime = str(js_path.parent / "npm") if js_path.parent.name else "npm"
|
||||
for option in options:
|
||||
if option == runtime_name or runtime_name.startswith(option):
|
||||
tool = shutil.which(js_runtime)
|
||||
if tool is None:
|
||||
raise RuntimeError(
|
||||
f"JS_RUNTIME={js_runtime_env}: {option} not found"
|
||||
)
|
||||
# Check Node.js version if using npm
|
||||
if option == "npm":
|
||||
node_path = shutil.which("node", path=str(Path(tool).parent))
|
||||
if node_path is None:
|
||||
raise RuntimeError(
|
||||
f"JS_RUNTIME={js_runtime_env}: node not found"
|
||||
)
|
||||
_check_node_version(node_path) # Raises on failure
|
||||
return tool, option
|
||||
raise RuntimeError(f"JS_RUNTIME={js_runtime_env} not recognized")
|
||||
if result := _find_runtime_from_env(options):
|
||||
return result
|
||||
|
||||
# Auto-detect
|
||||
for option in options:
|
||||
if tool := shutil.which(option):
|
||||
# Check Node.js version if using npm
|
||||
if option == "npm":
|
||||
node_path = shutil.which("node", path=str(Path(tool).parent))
|
||||
if node_path is None:
|
||||
continue
|
||||
try:
|
||||
_check_node_version(node_path)
|
||||
except RuntimeError as e:
|
||||
node_version_error = e
|
||||
continue # Try next runtime
|
||||
return tool, option
|
||||
|
||||
# No runtime found - provide helpful error
|
||||
if node_version_error:
|
||||
raise node_version_error
|
||||
raise RuntimeError("Node.js (v20+), Deno or Bun is required but none was found")
|
||||
return _auto_detect_runtime(options)
|
||||
|
||||
|
||||
def find_build_tool():
|
||||
def find_build_tool() -> tuple[list[str], list[str]]:
|
||||
"""Find JavaScript runtime and construct install/build commands.
|
||||
|
||||
Returns (install_cmd, build_cmd) tuples of command lists.
|
||||
@@ -146,7 +182,7 @@ def find_dev_tool() -> list[str]:
|
||||
|
||||
if name == "bun":
|
||||
logger.warning(
|
||||
"Bun has a bug in WS proxying (https://github.com/oven-sh/bun/issues/9882). Consider using npm instead."
|
||||
"Bun has a WS proxy bug (github.com/oven-sh/bun/issues/9882). Consider npm.",
|
||||
)
|
||||
|
||||
return [tool, *dev_args[name]]
|
||||
@@ -179,10 +215,10 @@ def build(folder: str = "frontend") -> None:
|
||||
install_cmd, build_cmd = find_build_tool()
|
||||
except RuntimeError as e:
|
||||
logger.warning(e)
|
||||
raise SystemExit(1) from e
|
||||
raise SystemExit(1) from None
|
||||
|
||||
def run(cmd):
|
||||
display_cmd = [Path(cmd[0]).name, *cmd[1:]]
|
||||
def run(cmd: list[str]) -> None:
|
||||
display_cmd = [Path(cmd[0]).stem, *cmd[1:]]
|
||||
logger.info("### %s", " ".join(display_cmd))
|
||||
subprocess.run(cmd, check=True, cwd=folder) # noqa: S603
|
||||
|
||||
@@ -190,5 +226,5 @@ def build(folder: str = "frontend") -> None:
|
||||
run(install_cmd)
|
||||
logger.info("")
|
||||
run(build_cmd)
|
||||
except subprocess.CalledProcessError as e:
|
||||
raise SystemExit(1) from e
|
||||
except subprocess.CalledProcessError:
|
||||
raise SystemExit(1) from None
|
||||
|
||||
+134
-63
@@ -1,111 +1,156 @@
|
||||
"""Utilities meant for devserver script, used only in source repository with dev deps."""
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import subprocess
|
||||
import sys
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Self
|
||||
|
||||
import httpx
|
||||
from buildutil import find_dev_tool, find_install_tool, logger
|
||||
from fastapi_vue.hostutil import parse_endpoint
|
||||
|
||||
DEFAULT_VITE_PORT = 8989
|
||||
DEFAULT_BACKEND_PORT = 8999
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Coroutine
|
||||
|
||||
|
||||
class ProcessGroup:
|
||||
"""Manage async subprocesses with automatic cleanup, like TaskGroup for processes."""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self) -> None:
|
||||
"""Initialize empty process tracking."""
|
||||
self._procs: list[asyncio.subprocess.Process] = []
|
||||
self._cmds: dict[int, str] = {} # pid -> command name
|
||||
|
||||
async def spawn(
|
||||
self, *cmd: str, cwd: str | None = None
|
||||
self,
|
||||
*cmd: str,
|
||||
cwd: str | None = None,
|
||||
) -> asyncio.subprocess.Process:
|
||||
"""Spawn a subprocess and track it."""
|
||||
logger.info(">>> %s", " ".join([Path(cmd[0]).name, *cmd[1:]]))
|
||||
cmd_name = Path(cmd[0]).stem
|
||||
logger.info(">>> %s", " ".join([cmd_name, *cmd[1:]]))
|
||||
proc = await asyncio.create_subprocess_exec(*cmd, cwd=cwd)
|
||||
self._procs.append(proc)
|
||||
self._cmds[proc.pid] = cmd_name
|
||||
return proc
|
||||
|
||||
async def wait(self, proc: asyncio.subprocess.Process) -> None:
|
||||
"""Wait for a process to complete, raise SystemExit(1) on failure."""
|
||||
if await proc.wait() != 0:
|
||||
logger.warning("Command failed")
|
||||
raise SystemExit(1)
|
||||
async def wait(
|
||||
self,
|
||||
*waitables: "asyncio.subprocess.Process | Coroutine[Any, Any, Any]",
|
||||
) -> None:
|
||||
"""Wait for processes/coroutines to complete, raise SystemExit on failure."""
|
||||
|
||||
async def __aenter__(self):
|
||||
async def wait_proc(proc: asyncio.subprocess.Process) -> None:
|
||||
returncode = await proc.wait()
|
||||
if returncode != 0:
|
||||
cmd_name = self._cmds.get(proc.pid, "unknown")
|
||||
raise subprocess.CalledProcessError(returncode, cmd_name)
|
||||
|
||||
tasks = [
|
||||
wait_proc(w) if isinstance(w, asyncio.subprocess.Process) else w
|
||||
for w in waitables
|
||||
]
|
||||
try:
|
||||
await asyncio.gather(*tasks)
|
||||
except subprocess.CalledProcessError as e:
|
||||
logger.warning("%s failed with exit status %d", e.cmd, e.returncode)
|
||||
raise SystemExit(1) from None
|
||||
|
||||
async def __aenter__(self) -> Self:
|
||||
"""Enter the async context manager."""
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, *_):
|
||||
async def __aexit__(self, exc_type: type[BaseException] | None, *_: object) -> None:
|
||||
"""Wait for one process to exit, terminate others, then wait for all."""
|
||||
cleanup_task = asyncio.create_task(self._cleanup())
|
||||
try:
|
||||
await asyncio.shield(cleanup_task)
|
||||
except asyncio.CancelledError:
|
||||
# Shield was cancelled but cleanup_task continues - wait for it
|
||||
await cleanup_task
|
||||
await self._cleanup(immediate=exc_type is not None)
|
||||
|
||||
async def _cleanup(self):
|
||||
async def _cleanup(self, *, immediate: bool = False) -> None:
|
||||
running = [p for p in self._procs if p.returncode is None]
|
||||
if not running:
|
||||
return
|
||||
|
||||
# Wait for any one process to exit
|
||||
await asyncio.wait(
|
||||
[asyncio.create_task(p.wait()) for p in running],
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
if not immediate:
|
||||
# Wait for any one process to exit
|
||||
with suppress(asyncio.CancelledError):
|
||||
await asyncio.wait(
|
||||
[asyncio.create_task(p.wait()) for p in running],
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
|
||||
# Terminate remaining processes
|
||||
for p in self._procs:
|
||||
if p.returncode is None:
|
||||
with contextlib.suppress(ProcessLookupError):
|
||||
with suppress(ProcessLookupError):
|
||||
p.terminate()
|
||||
|
||||
# Wait for all to finish (with overall timeout)
|
||||
# Wait for all to finish (with overall timeout), shielded from cancellation
|
||||
still_running = [p for p in self._procs if p.returncode is None]
|
||||
if still_running:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(*[p.wait() for p in still_running]),
|
||||
timeout=10,
|
||||
)
|
||||
except TimeoutError:
|
||||
for p in self._procs:
|
||||
if p.returncode is None:
|
||||
with contextlib.suppress(ProcessLookupError):
|
||||
p.kill()
|
||||
await p.wait()
|
||||
with suppress(asyncio.CancelledError):
|
||||
try:
|
||||
await asyncio.shield(
|
||||
asyncio.wait_for(
|
||||
asyncio.gather(*[p.wait() for p in still_running]),
|
||||
timeout=10,
|
||||
),
|
||||
)
|
||||
except TimeoutError:
|
||||
for p in self._procs:
|
||||
if p.returncode is None:
|
||||
with suppress(ProcessLookupError):
|
||||
p.kill()
|
||||
await p.wait()
|
||||
|
||||
|
||||
async def ready(url: str, path: str = "") -> None:
|
||||
async def check_ports_free(*urls: str) -> None:
|
||||
"""Verify URLs are not responding (ports are free). Raise SystemExit if any respond."""
|
||||
|
||||
async def check(client: httpx.AsyncClient, url: str) -> None:
|
||||
with suppress(httpx.RequestError):
|
||||
res = await client.get(url, timeout=0.1)
|
||||
server = res.headers.get("server", "server")
|
||||
logger.warning("Conflicting %s already running at %s", server, url)
|
||||
raise SystemExit(1)
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
await asyncio.gather(*[check(client, url) for url in urls])
|
||||
|
||||
|
||||
async def ready(url: str, path: str = "", max_attempts: int = 50) -> None:
|
||||
"""Wait for the server to be ready by polling an endpoint.
|
||||
|
||||
Use empty path to disable the check and make this return immediately.
|
||||
Raises SystemExit(1) if server doesn't start in time.
|
||||
"""
|
||||
max_attempts = 50
|
||||
full_url = f"{url}{path}"
|
||||
if not path:
|
||||
return
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
for attempt in range(max_attempts):
|
||||
try:
|
||||
await client.get(full_url, timeout=1.0)
|
||||
logger.info("✓ Backend ready!")
|
||||
return
|
||||
except httpx.RequestError as e:
|
||||
await client.get(f"{url}{path}", timeout=1.0)
|
||||
except httpx.RequestError:
|
||||
if attempt == max_attempts - 1:
|
||||
logger.warning("Backend didn't start in time")
|
||||
raise SystemExit(1) from e
|
||||
raise SystemExit(1) from None
|
||||
await asyncio.sleep(0.1)
|
||||
else:
|
||||
logger.info("✓ Backend ready!")
|
||||
return
|
||||
|
||||
|
||||
def setup_vite(endpoint: str) -> tuple[str, list[str], list[str]]:
|
||||
def setup_vite(
|
||||
endpoint: str,
|
||||
default_port: int = 5173,
|
||||
) -> tuple[str, list[str], list[str]]:
|
||||
"""Parse frontend endpoint and build commands.
|
||||
|
||||
Returns (url, install_cmd, dev_cmd).
|
||||
Raises SystemExit(1) on invalid config.
|
||||
"""
|
||||
endpoints = parse_endpoint(endpoint, DEFAULT_VITE_PORT)
|
||||
endpoints = parse_endpoint(endpoint, default_port)
|
||||
|
||||
if "uds" in endpoints[0]:
|
||||
logger.warning("Unix sockets not supported with vite devserver")
|
||||
@@ -118,18 +163,53 @@ def setup_vite(endpoint: str) -> tuple[str, list[str], list[str]]:
|
||||
dev_cmd = find_dev_tool()
|
||||
if host != "localhost":
|
||||
dev_cmd.append("--host" if len(endpoints) > 1 else f"--host={host}")
|
||||
if port != 5173:
|
||||
dev_cmd.append(f"--port={port}")
|
||||
dev_cmd.append(f"--port={port}")
|
||||
|
||||
return f"http://{host}:{port}", install_cmd, dev_cmd
|
||||
|
||||
|
||||
def setup_fastapi(
|
||||
endpoint: str, module: str, default_port: int = DEFAULT_BACKEND_PORT
|
||||
endpoint: str,
|
||||
module: str,
|
||||
default_port: int = 8000,
|
||||
) -> tuple[str, list[str]]:
|
||||
"""Parse backend endpoint and build fastapi dev command.
|
||||
"""Parse backend endpoint and build uvicorn command.
|
||||
|
||||
Returns (url, cmd).
|
||||
Returns (url, uvicorn_cmd).
|
||||
Raises SystemExit(1) on invalid config.
|
||||
"""
|
||||
endpoints = parse_endpoint(endpoint, default_port)
|
||||
|
||||
if "uds" in endpoints[0]:
|
||||
logger.warning("Unix sockets not supported with vite devserver")
|
||||
raise SystemExit(1)
|
||||
|
||||
host = endpoints[0]["host"]
|
||||
port = endpoints[0]["port"]
|
||||
reload_dir = module.split(".", maxsplit=1)[0] # Don't reload on frontend changes
|
||||
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"uvicorn",
|
||||
module,
|
||||
f"--host={host}",
|
||||
f"--port={port}",
|
||||
"--reload",
|
||||
f"--reload-dir={reload_dir}",
|
||||
"--forwarded-allow-ips=*",
|
||||
]
|
||||
return f"http://{host}:{port}", cmd
|
||||
|
||||
|
||||
def setup_cli(
|
||||
cli: str,
|
||||
endpoint: str,
|
||||
default_port: int = 8000,
|
||||
) -> tuple[str, list[str]]:
|
||||
"""Parse backend endpoint and build CLI command.
|
||||
|
||||
Returns (url, cli_cmd).
|
||||
Raises SystemExit(1) on invalid config.
|
||||
"""
|
||||
endpoints = parse_endpoint(endpoint, default_port)
|
||||
@@ -141,14 +221,5 @@ def setup_fastapi(
|
||||
host = endpoints[0]["host"]
|
||||
port = endpoints[0]["port"]
|
||||
|
||||
cmd = [
|
||||
"fastapi",
|
||||
"dev",
|
||||
"--entrypoint",
|
||||
module,
|
||||
"--host",
|
||||
host,
|
||||
"--port",
|
||||
str(port),
|
||||
]
|
||||
cmd = [cli, f"--listen={host}:{port}"]
|
||||
return f"http://{host}:{port}", cmd
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
from http.cookies import SimpleCookie
|
||||
from pathlib import Path
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from sanic import Sanic
|
||||
|
||||
from cista import auth, config
|
||||
from cista.app import use_session
|
||||
from cista.auth import bp as auth_bp
|
||||
|
||||
|
||||
def _set_cookie_headers(response) -> list[str]:
|
||||
return list(response.headers.get_list("set-cookie"))
|
||||
|
||||
|
||||
def _cookie_header(response, name: str = "cista") -> dict[str, str]:
|
||||
for header in _set_cookie_headers(response):
|
||||
cookie = SimpleCookie()
|
||||
cookie.load(header)
|
||||
morsel = cookie.get(name)
|
||||
if morsel is not None and morsel.value:
|
||||
return {"Cookie": f"{name}={morsel.value}"}
|
||||
raise AssertionError(f"response did not set cookie {name!r}")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def setup_auth_config(tmp_path: Path):
|
||||
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},
|
||||
)
|
||||
return tmp_path
|
||||
|
||||
|
||||
@pytest_asyncio.fixture()
|
||||
async def client(setup_auth_config: Path):
|
||||
app = Sanic(f"auth-builtins-test-{uuid4().hex}", strict_slashes=True)
|
||||
|
||||
@app.on_request
|
||||
async def load_auth_context(request):
|
||||
await use_session(request)
|
||||
|
||||
app.blueprint(auth_bp)
|
||||
yield app.asgi_client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_restricted_page_renders_login_form_when_logged_out(client):
|
||||
_, res = await client.get("/auth/restricted/")
|
||||
|
||||
assert res.status_code == 200
|
||||
assert "Authentication Required" in res.text
|
||||
assert "Username:" in res.text
|
||||
assert "Password:" in res.text
|
||||
assert "/auth/login" in res.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_restricted_page_with_invalid_session_clears_cookie(client):
|
||||
_, res = await client.get(
|
||||
"/auth/restricted/",
|
||||
headers={"Cookie": "cista=missing-session"},
|
||||
)
|
||||
|
||||
assert res.status_code == 200
|
||||
assert "Authentication Required" in res.text
|
||||
assert any("cista=" in header.lower() for header in _set_cookie_headers(res))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_json_login_sets_session_cookie_and_allows_session_authenticated_api_access(
|
||||
client,
|
||||
):
|
||||
_, res = await client.post(
|
||||
"/auth/login",
|
||||
json={"username": "alice", "password": "secret"},
|
||||
)
|
||||
|
||||
assert res.status_code == 200
|
||||
assert res.json == {"data": {"username": "alice", "privileged": False}}
|
||||
|
||||
session_cookie = _cookie_header(res)
|
||||
|
||||
_, tokens_res = await client.get("/auth/tokens", headers=session_cookie)
|
||||
assert tokens_res.status_code == 200
|
||||
assert tokens_res.json == {"tokens": []}
|
||||
|
||||
_, restricted_res = await client.get("/auth/restricted/", headers=session_cookie)
|
||||
assert restricted_res.status_code == 200
|
||||
assert "auth-success" in restricted_res.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_json_login_rejects_missing_fields(client):
|
||||
_, res = await client.post(
|
||||
"/auth/login",
|
||||
json={"username": "alice"},
|
||||
)
|
||||
|
||||
assert res.status_code == 400
|
||||
assert "Missing username or password" in res.json["message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_json_login_rejects_invalid_password(client):
|
||||
_, res = await client.post(
|
||||
"/auth/login",
|
||||
json={"username": "alice", "password": "wrong"},
|
||||
)
|
||||
|
||||
assert res.status_code == 403
|
||||
assert "Invalid password" in res.json["message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_html_login_redirects_and_sets_flash_and_session_cookies(client):
|
||||
_, res = await client.post(
|
||||
"/auth/login",
|
||||
data={"username": "alice", "password": "secret"},
|
||||
headers={"Accept": "text/html"},
|
||||
follow_redirects=False,
|
||||
)
|
||||
|
||||
assert res.status_code == 302
|
||||
assert res.headers["location"] == "/"
|
||||
headers = _set_cookie_headers(res)
|
||||
assert any("cista=" in header.lower() for header in headers)
|
||||
assert any("message=" in header.lower() for header in headers)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logout_json_revokes_the_existing_session(client):
|
||||
_, login_res = await client.post(
|
||||
"/auth/login",
|
||||
json={"username": "alice", "password": "secret"},
|
||||
)
|
||||
session_cookie = _cookie_header(login_res)
|
||||
|
||||
_, logout_res = await client.post("/auth/api/logout", headers=session_cookie)
|
||||
|
||||
assert logout_res.status_code == 200
|
||||
assert logout_res.json == {"message": "Logged out"}
|
||||
assert any("cista=" in header.lower() for header in _set_cookie_headers(logout_res))
|
||||
|
||||
_, retry_res = await client.get("/auth/tokens", headers=session_cookie)
|
||||
assert retry_res.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logout_without_session_reports_not_logged_in(client):
|
||||
_, res = await client.post("/auth/api/logout")
|
||||
|
||||
assert res.status_code == 200
|
||||
assert res.json == {"message": "Not logged in"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_change_updates_credentials_and_reissues_session(client):
|
||||
_, change_res = await client.post(
|
||||
"/auth/password-change",
|
||||
json={
|
||||
"username": "alice",
|
||||
"password": "secret",
|
||||
"passwordChange": "fresh-secret",
|
||||
},
|
||||
)
|
||||
|
||||
assert change_res.status_code == 200
|
||||
assert change_res.json == {"message": "Password updated"}
|
||||
|
||||
session_cookie = _cookie_header(change_res)
|
||||
_, tokens_res = await client.get("/auth/tokens", headers=session_cookie)
|
||||
assert tokens_res.status_code == 200
|
||||
|
||||
_, old_login_res = await client.post(
|
||||
"/auth/login",
|
||||
json={"username": "alice", "password": "secret"},
|
||||
)
|
||||
assert old_login_res.status_code == 403
|
||||
|
||||
_, new_login_res = await client.post(
|
||||
"/auth/login",
|
||||
json={"username": "alice", "password": "fresh-secret"},
|
||||
)
|
||||
assert new_login_res.status_code == 200
|
||||
assert new_login_res.json == {"data": {"username": "alice", "privileged": False}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_change_rejects_wrong_current_password(client):
|
||||
_, res = await client.post(
|
||||
"/auth/password-change",
|
||||
json={
|
||||
"username": "alice",
|
||||
"password": "wrong",
|
||||
"passwordChange": "fresh-secret",
|
||||
},
|
||||
)
|
||||
|
||||
assert res.status_code == 403
|
||||
assert "Invalid password" in res.json["message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_change_rejects_missing_fields(client):
|
||||
_, res = await client.post(
|
||||
"/auth/password-change",
|
||||
json={"username": "alice", "password": "secret"},
|
||||
)
|
||||
|
||||
assert res.status_code == 400
|
||||
assert "Missing username, passwordChange or password" in res.json["message"]
|
||||
@@ -3,11 +3,11 @@ import hashlib
|
||||
import hmac
|
||||
import struct
|
||||
from pathlib import Path
|
||||
from time import time
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from Crypto.Hash import MD4
|
||||
from sanic import Sanic
|
||||
|
||||
from cista import auth, config, session, watching
|
||||
@@ -29,8 +29,6 @@ def _ntlm_type3(
|
||||
username: str, password: str, domain: str, challenge: bytes
|
||||
) -> dict[str, str]:
|
||||
"""Build an NTLMv2 Type 3 message for testing."""
|
||||
from Crypto.Hash import MD4
|
||||
|
||||
# NT hash
|
||||
nt_hash = MD4.new(password.encode("utf-16le")).digest()
|
||||
# NTLMv2 hash
|
||||
@@ -94,10 +92,7 @@ def _ntlm_type3(
|
||||
|
||||
def _session_cookie_header(username: str) -> dict[str, str]:
|
||||
token = "test-" + username
|
||||
session._sessions[token] = {
|
||||
"exp": int(time()) + session.max_age,
|
||||
"username": username,
|
||||
}
|
||||
session.put(token, username)
|
||||
return {"Cookie": f"cista={token}"}
|
||||
|
||||
|
||||
|
||||
@@ -160,7 +160,7 @@ async def test_mkcol_windows_drive_path_stays_within_root(client, setup_storage:
|
||||
# Either created inside the storage root (201) or sanitised away (400/404).
|
||||
# The important assertion: nothing was created outside the storage root.
|
||||
assert not (Path("/c:") / "secret").exists()
|
||||
assert not (Path("c:/secret")).exists()
|
||||
assert not (Path("c:/secret")).exists() # noqa: ASYNC240
|
||||
if res.status_code == 201:
|
||||
# Created safely inside tmp storage
|
||||
assert (setup_storage / "c:" / "secret").is_dir()
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
from pathlib import Path, PurePath
|
||||
from uuid import uuid4
|
||||
|
||||
import msgspec
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from sanic import Sanic
|
||||
@@ -12,10 +12,6 @@ from cista.auth import bp as auth_bp
|
||||
|
||||
|
||||
def _persist_config():
|
||||
from pathlib import PurePath
|
||||
|
||||
import msgspec
|
||||
|
||||
def enc_hook(obj):
|
||||
if isinstance(obj, PurePath):
|
||||
return obj.as_posix()
|
||||
@@ -27,8 +23,7 @@ def _persist_config():
|
||||
|
||||
@pytest.fixture
|
||||
def setup_storage(tmp_path: Path):
|
||||
os.environ["CISTA_HOME"] = str(tmp_path)
|
||||
config.init_confdir()
|
||||
config.init_confdir(tmp_path)
|
||||
user = config.User()
|
||||
auth.set_password(user, "secret")
|
||||
admin = config.User(privileged=True)
|
||||
|
||||
Reference in New Issue
Block a user