zcbot/core/web_previews.py

230 lines
8.0 KiB
Python

"""Static multi-file web preview lifecycle and workspace path validation."""
from __future__ import annotations
import os
from datetime import datetime, timedelta, timezone
from pathlib import Path, PurePosixPath
from uuid import UUID
from sqlalchemy import select, update
from core.storage import session_scope
from core.storage.models import Task, WebPreview
DEFAULT_TTL_SECONDS = 7 * 24 * 3600
MIN_TTL_SECONDS = 300
MAX_TTL_SECONDS = 30 * 24 * 3600
MAX_ACTIVE_PER_USER = 10
MAX_FILES = 5_000
MAX_TOTAL_BYTES = 200 * 1024 * 1024
MAX_FILE_BYTES = 50 * 1024 * 1024
class WebPreviewError(ValueError):
pass
def _now() -> datetime:
return datetime.now(timezone.utc)
def _ttl_seconds() -> int:
raw = os.getenv("ZCBOT_WEB_PREVIEW_TTL_SECONDS", "").strip()
try:
value = int(raw) if raw else DEFAULT_TTL_SECONDS
except ValueError:
value = DEFAULT_TTL_SECONDS
return max(MIN_TTL_SECONDS, min(value, MAX_TTL_SECONDS))
def _safe_relative(raw: str, *, label: str, allow_dot: bool = False) -> str:
value = str(raw or "").strip().replace("\\", "/")
if allow_dot and value in {"", ".", "./"}:
return "."
if not value or "\x00" in value or value.startswith("/"):
raise WebPreviewError(f"{label} must be a non-empty relative path")
path = PurePosixPath(value)
if path.is_absolute() or any(part in {"", ".", ".."} for part in path.parts):
raise WebPreviewError(f"{label} contains an invalid path segment")
return path.as_posix()
def resolve_preview_root(working_dir: Path, root_path: str) -> tuple[Path, str]:
wd = Path(working_dir).resolve()
rel = _safe_relative(root_path, label="directory", allow_dot=True)
candidate = wd if rel == "." else wd.joinpath(*PurePosixPath(rel).parts)
resolved = candidate.resolve()
try:
resolved.relative_to(wd)
except ValueError as exc:
raise WebPreviewError("preview directory escapes the task working directory") from exc
if not resolved.is_dir():
raise WebPreviewError(f"preview directory not found: {rel}")
return resolved, rel
def inspect_preview_tree(root: Path, entry_path: str) -> tuple[str, int, int]:
entry = _safe_relative(entry_path or "index.html", label="entry")
entry_abs = root.joinpath(*PurePosixPath(entry).parts).resolve()
try:
entry_abs.relative_to(root.resolve())
except ValueError as exc:
raise WebPreviewError("entry escapes the preview directory") from exc
if not entry_abs.is_file() or entry_abs.suffix.lower() not in {".html", ".htm"}:
raise WebPreviewError(f"preview entry must be an existing HTML file: {entry}")
count = 0
total = 0
for current, dirs, files in os.walk(root, followlinks=False):
current_path = Path(current)
for name in tuple(dirs):
if (current_path / name).is_symlink():
raise WebPreviewError("preview directory cannot contain symbolic links")
for name in files:
path = current_path / name
if path.is_symlink():
raise WebPreviewError("preview directory cannot contain symbolic links")
try:
size = path.stat().st_size
except OSError as exc:
raise WebPreviewError(f"cannot inspect preview file: {name}") from exc
if size > MAX_FILE_BYTES:
raise WebPreviewError(f"preview file exceeds 50 MiB: {name}")
count += 1
total += size
if count > MAX_FILES:
raise WebPreviewError(f"preview contains more than {MAX_FILES} files")
if total > MAX_TOTAL_BYTES:
raise WebPreviewError("preview directory exceeds 200 MiB")
return entry, count, total
def preview_dict(row: WebPreview) -> dict:
return {
"type": "web_preview",
"preview_id": str(row.preview_id),
"task_id": str(row.task_id),
"name": row.name,
"root": row.root_path,
"entry": row.entry_path,
"spa_fallback": bool(row.spa_fallback),
"status": row.status,
"file_count": row.file_count,
"size_bytes": row.size_bytes,
"expires_at": row.expires_at.isoformat(),
}
def create_web_preview(
*,
user_id: UUID,
task_id: UUID,
working_dir: Path,
directory: str,
entry: str = "index.html",
name: str = "网页项目预览",
spa_fallback: bool = True,
) -> dict:
root, root_rel = resolve_preview_root(working_dir, directory)
entry_rel, file_count, size_bytes = inspect_preview_tree(root, entry)
display_name = str(name or "").strip()[:120] or "网页项目预览"
now = _now()
expires_at = now + timedelta(seconds=_ttl_seconds())
with session_scope() as session:
task_exists = session.execute(
select(Task.task_id).where(Task.task_id == task_id, Task.user_id == user_id)
).scalar_one_or_none()
if task_exists is None:
raise WebPreviewError("task not found")
session.execute(
update(WebPreview)
.where(
WebPreview.user_id == user_id,
WebPreview.status == "active",
WebPreview.expires_at <= now,
)
.values(status="expired")
)
existing = session.execute(
select(WebPreview)
.where(
WebPreview.user_id == user_id,
WebPreview.task_id == task_id,
WebPreview.root_path == root_rel,
WebPreview.entry_path == entry_rel,
WebPreview.status == "active",
)
.order_by(WebPreview.created_at.desc())
.limit(1)
).scalar_one_or_none()
if existing is not None:
existing.name = display_name
existing.spa_fallback = bool(spa_fallback)
existing.file_count = file_count
existing.size_bytes = size_bytes
existing.expires_at = expires_at
session.flush()
return preview_dict(existing)
active = session.execute(
select(WebPreview.preview_id).where(
WebPreview.user_id == user_id,
WebPreview.status == "active",
)
).scalars().all()
if len(active) >= MAX_ACTIVE_PER_USER:
raise WebPreviewError(
f"active web preview limit reached ({MAX_ACTIVE_PER_USER}); stop an old preview first"
)
row = WebPreview(
user_id=user_id,
task_id=task_id,
root_path=root_rel,
entry_path=entry_rel,
name=display_name,
spa_fallback=bool(spa_fallback),
file_count=file_count,
size_bytes=size_bytes,
expires_at=expires_at,
)
session.add(row)
session.flush()
return preview_dict(row)
def get_web_preview(preview_id: UUID, user_id: UUID | None = None) -> tuple[WebPreview, str]:
with session_scope() as session:
statement = (
select(WebPreview, Task.working_dir)
.join(Task, Task.task_id == WebPreview.task_id)
.where(WebPreview.preview_id == preview_id)
)
if user_id is not None:
statement = statement.where(WebPreview.user_id == user_id)
result = session.execute(statement).first()
if result is None:
raise WebPreviewError("preview not found")
row, working_dir = result
if row.status == "active" and row.expires_at <= _now():
row.status = "expired"
session.flush()
return row, working_dir
def stop_web_preview(preview_id: UUID, user_id: UUID) -> dict:
with session_scope() as session:
row = session.execute(
select(WebPreview).where(
WebPreview.preview_id == preview_id,
WebPreview.user_id == user_id,
)
).scalar_one_or_none()
if row is None:
raise WebPreviewError("preview not found")
if row.status == "active":
row.status = "stopped"
session.flush()
return preview_dict(row)