Files
hermes-agent/gateway/browser_control_artifacts.py

413 lines
16 KiB
Python

"""One-shot artifact transport for browser control (Gateway side).
Transport-neutral store core for bounded browser-control artifacts
(screenshots, PDFs, uploads). It knows nothing about aiohttp: the routes in
:mod:`gateway.platforms.api_server` authenticate callers and rate-limit, then
hand bytes here. The controller WebSocket is a command channel, not a file
pipe — frames carry only ``artifact_id`` strings and the bytes live on disk
under a controlled root for a short TTL.
Contract (tests/gateway/test_browser_control_artifacts.py):
- Server-minted ``[0-9a-f]{32}`` ids resolved strictly inside the root;
client filenames are metadata only, never paths.
- Exact size and MIME caps enforced before any disk write.
- SHA-256 recorded in the receipt and re-verified by ``load``/``validate``.
- ``load`` requires the exact scope key and consumes atomically;
``validate`` checks existence/TTL/scope without consuming.
- No overwrite: an id collision is retried with a fresh id.
- ``prune_expired`` removes expired entries; the API server sweeps on demand.
Thread-safety: the in-memory index is lock-guarded; files are written to a
temp name and atomically renamed so a concurrent ``load`` never sees a
partial artifact.
"""
from __future__ import annotations
import contextlib
import hashlib
import logging
import os
import re
import secrets
import threading
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable, Optional
logger = logging.getLogger(__name__)
DEFAULT_ARTIFACT_TTL_SECONDS = 300.0
DEFAULT_MAX_ARTIFACT_BYTES = 10 * 1024 * 1024
#: Exact allowlist — parameterized/unknown variants are rejected.
DEFAULT_ALLOWED_MIME_TYPES = frozenset({
"application/json", "application/pdf", "image/gif", "image/jpeg",
"image/png", "image/webp", "text/plain",
})
_ARTIFACT_ID_HEX = 32
_ARTIFACT_ID_RE = re.compile(r"^[0-9a-f]{32}$")
_TEMP_SUFFIX = ".tmp"
class ArtifactError(Exception):
"""Base class for artifact store contract failures."""
class ArtifactNotFound(ArtifactError):
"""The artifact id is unknown (or already consumed)."""
class ArtifactExpired(ArtifactError):
"""The artifact outlived its TTL."""
class ArtifactTooLarge(ArtifactError):
"""The upload exceeds the configured byte cap."""
class ArtifactMimeRejected(ArtifactError):
"""The content type is outside the exact allowlist."""
class ArtifactScopeMismatch(ArtifactError):
"""The artifact exists but belongs to a different scope."""
class ArtifactChecksumMismatch(ArtifactError):
"""The stored bytes do not match the recorded SHA-256."""
class ArtifactTraversal(ArtifactError):
"""A caller-supplied id is not a valid minted artifact id."""
@dataclass(frozen=True)
class ArtifactReceipt:
"""Provenance record returned to the caller of ``store``."""
artifact_id: str
sha256: str
size_bytes: int
content_type: str
filename: str
created_at: float
expires_at: float
ttl_seconds: float
scope_key: str
def to_dict(self, *, download_path: str = "") -> dict[str, Any]:
"""Serialize to the wire receipt (never contains file paths)."""
receipt = {
"artifact_id": self.artifact_id,
"sha256": self.sha256,
"size_bytes": self.size_bytes,
"content_type": self.content_type,
"filename": self.filename,
"created_at": self.created_at,
"expires_at": self.expires_at,
"ttl_seconds": self.ttl_seconds,
"one_shot": True,
}
if download_path:
receipt["download_path"] = download_path
return receipt
def artifact_scope_key(scope: Any) -> str:
"""Derive the stable scope key an artifact is bound to.
Only principal (mandatory) + transport family participate. ``session_id``
is deliberately EXCLUDED: HTTP artifact routes authenticate by API key and
can't resolve a session, while broker dispatch always carries one — hashing
it would make upload and dispatch never compose. Ids are unguessable and
downloads one-shot, so cross-session reuse within one principal is by
design. Capabilities/optional ids are excluded so a reconnect keeps its
artifacts.
"""
principal = ""
family = ""
try:
principal = str(getattr(scope, "principal_id", "") or "")
family = str(getattr(scope, "transport_family", "") or "")
except Exception:
pass
if not principal:
# Fail closed: only an authenticated principal may mint artifacts.
raise ArtifactError("artifact scope must carry a resolved principal")
material = f"{principal}\x00{family}".encode("utf-8")
return hashlib.sha256(material).hexdigest()
def _sha256(data: bytes) -> str:
return hashlib.sha256(data).hexdigest()
@dataclass
class _ArtifactEntry:
receipt: ArtifactReceipt
path: Path
class ArtifactStore:
"""Thread-safe, TTL-bounded, scope-bound one-shot artifact store."""
def __init__(
self, root: Path, *,
ttl_seconds: float = DEFAULT_ARTIFACT_TTL_SECONDS,
max_bytes: int = DEFAULT_MAX_ARTIFACT_BYTES,
allowed_mime_types: frozenset = DEFAULT_ALLOWED_MIME_TYPES,
clock: Optional[Callable[[], float]] = None,
) -> None:
self._root = Path(root)
self._root.mkdir(parents=True, exist_ok=True)
self._ttl_seconds = max(1.0, float(ttl_seconds))
self._max_bytes = max(1, int(max_bytes))
self._allowed_mime_types = frozenset(allowed_mime_types)
self._clock = clock if clock is not None else time.time
self._lock = threading.RLock()
self._entries: dict[str, _ArtifactEntry] = {}
# Receipts live only in memory, so files left by a previous process
# are unreachable orphans past their TTL by definition — sweep them.
self._sweep_orphan_files()
def _sweep_orphan_files(self) -> None:
"""Delete on-disk files with no live index entry.
Only names matching the minted 32-hex id shape or the ``*.tmp``
staging suffix are touched; anything else in the directory is left.
"""
try:
candidates = list(self._root.iterdir())
except OSError:
return
with self._lock:
live = set(self._entries)
for path in candidates:
name = path.name
is_temp = name.endswith(_TEMP_SUFFIX)
if not path.is_file() or not (is_temp or _ARTIFACT_ID_RE.fullmatch(name)):
continue
if not is_temp and name in live:
continue
with contextlib.suppress(OSError):
path.unlink(missing_ok=True)
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
@property
def root(self) -> Path:
"""Controlled artifact root (never exposed to callers by default)."""
return self._root
@property
def max_bytes(self) -> int:
return self._max_bytes
@property
def allowed_mime_types(self) -> frozenset:
return self._allowed_mime_types
def store(self, data: bytes, *, filename: str, content_type: str, scope: Any) -> ArtifactReceipt:
"""Validate and store one artifact, returning its provenance receipt.
Raises :class:`ArtifactTooLarge` / :class:`ArtifactMimeRejected`
before any disk write; :class:`ArtifactError` on an unresolved scope.
"""
size = len(data)
if size > self._max_bytes:
raise ArtifactTooLarge(f"artifact is {size} bytes; cap is {self._max_bytes}")
normalized_type = _normalize_content_type(content_type)
if normalized_type not in self._allowed_mime_types:
raise ArtifactMimeRejected(f"content type {content_type!r} is outside the exact allowlist")
scope_key = artifact_scope_key(scope)
now = self._clock()
# Mint a fresh id; retry on an astronomically unlikely collision.
while True:
artifact_id = secrets.token_hex(_ARTIFACT_ID_HEX // 2)
target = self._artifact_path(artifact_id)
with self._lock:
if artifact_id in self._entries or target.exists():
continue
receipt = ArtifactReceipt(
artifact_id=artifact_id, sha256=_sha256(data), size_bytes=size,
content_type=normalized_type, filename=_bounded_filename(filename),
created_at=now, expires_at=now + self._ttl_seconds,
ttl_seconds=self._ttl_seconds, scope_key=scope_key,
)
self._entries[artifact_id] = _ArtifactEntry(receipt=receipt, path=target)
break
# Temp + atomic rename so readers never observe a partial artifact.
temp = target.with_name(f"{target.name}{_TEMP_SUFFIX}")
try:
with open(temp, "wb") as handle:
handle.write(data)
handle.flush()
os.fsync(handle.fileno())
os.replace(temp, target)
except Exception:
with self._lock:
self._entries.pop(artifact_id, None)
with contextlib.suppress(Exception):
temp.unlink(missing_ok=True)
raise
return receipt
def validate(self, artifact_id: str, *, scope: Any) -> ArtifactReceipt:
"""Return the receipt when the artifact is live for ``scope``
(existence, TTL, scope) without consuming it; raises otherwise."""
return self._entry_for(artifact_id, scope=scope).receipt
def load(self, artifact_id: str, *, scope: Any) -> tuple[bytes, ArtifactReceipt]:
"""One-shot download: verify, read, checksum, then consume.
A second ``load`` raises :class:`ArtifactNotFound`. A checksum
mismatch raises :class:`ArtifactChecksumMismatch` without consuming.
"""
with self._lock:
entry = self._entry_for(artifact_id, scope=scope)
path = entry.path
if not path.exists():
self._entries.pop(artifact_id, None)
raise ArtifactNotFound(f"artifact {artifact_id!r} is gone")
try:
data = path.read_bytes()
except OSError as exc:
raise ArtifactError(f"artifact read failed: {exc}") from exc
if _sha256(data) != entry.receipt.sha256:
raise ArtifactChecksumMismatch(f"artifact {artifact_id!r} failed SHA-256 validation")
# Drop the index entry first so a concurrent load fails closed.
self._entries.pop(artifact_id, None)
try:
path.unlink(missing_ok=True)
except OSError:
logger.warning("artifact %s: file removal failed; TTL sweep will retry", artifact_id)
return data, entry.receipt
def prune_expired(self, now: Optional[float] = None) -> int:
"""Delete every artifact past its TTL (and stale temp files); return
the count removed. Idempotent."""
now = self._clock() if now is None else float(now)
with self._lock:
removed = self._prune_expired_locked(now)
for temp in self._root.glob(f"*{_TEMP_SUFFIX}"):
try:
if temp.stat().st_mtime <= now - self._ttl_seconds:
temp.unlink(missing_ok=True)
except OSError:
continue
return removed
def count(self) -> int:
"""Number of live (unconsumed, not-yet-pruned) artifacts."""
with self._lock:
return len(self._entries)
# ------------------------------------------------------------------
# Internals
# ------------------------------------------------------------------
def _entry_for(self, artifact_id: str, *, scope: Any) -> _ArtifactEntry:
path = self._artifact_path(artifact_id)
scope_key = artifact_scope_key(scope)
now = self._clock()
with self._lock:
entry = self._entries.get(artifact_id)
# Check the target's own expiry BEFORE sweeping so an expired
# artifact surfaces as ArtifactExpired, not ArtifactNotFound.
if entry is None:
self._prune_expired_locked(now)
entry = self._entries.get(artifact_id)
if entry is None:
raise ArtifactNotFound(f"unknown artifact {artifact_id!r}")
if entry.receipt.expires_at <= now:
self._entries.pop(artifact_id, None)
with contextlib.suppress(OSError):
path.unlink(missing_ok=True)
raise ArtifactExpired(f"artifact {artifact_id!r} expired")
if entry.receipt.scope_key != scope_key:
raise ArtifactScopeMismatch(f"artifact {artifact_id!r} is bound to a different scope")
return entry
def _prune_expired_locked(self, now: float) -> int:
removed = 0
for artifact_id, entry in list(self._entries.items()):
if entry.receipt.expires_at <= now:
self._entries.pop(artifact_id, None)
with contextlib.suppress(OSError):
entry.path.unlink(missing_ok=True)
removed += 1
return removed
def _artifact_path(self, artifact_id: str) -> Path:
"""Resolve a minted id strictly inside the controlled root."""
if not isinstance(artifact_id, str) or not _ARTIFACT_ID_RE.fullmatch(artifact_id):
raise ArtifactTraversal(f"invalid artifact id {artifact_id!r}")
candidate = (self._root / artifact_id).resolve()
try:
root_resolved = self._root.resolve()
except OSError:
root_resolved = self._root.absolute()
if candidate.parent != root_resolved or candidate.name != artifact_id:
raise ArtifactTraversal(f"artifact path escapes root for {artifact_id!r}")
return candidate
def _normalize_content_type(value: str) -> str:
"""Return the canonical MIME type, or ``""`` for malformed input."""
if not isinstance(value, str):
return ""
return value.strip().split(";", 1)[0].strip().lower()
def _bounded_filename(value: str, limit: int = 160) -> str:
"""Sanitize a display-only filename; never used as a filesystem path."""
if not isinstance(value, str):
return ""
cleaned = value.strip().replace("\\", "_").replace("/", "_")
cleaned = "".join(character for character in cleaned if ord(character) >= 32)
return cleaned[:limit]
# ----------------------------------------------------------------------
# Rate limiting (route-level, per principal)
# ----------------------------------------------------------------------
class ArtifactRateLimiter:
"""Sliding-window per-key limiter; the API server keys it by principal."""
def __init__(
self, *, window_seconds: float = 60.0, max_requests: int = 30,
clock: Optional[Callable[[], float]] = None,
) -> None:
self._window_seconds = max(1.0, float(window_seconds))
self._max_requests = max(1, int(max_requests))
self._clock = clock if clock is not None else time.time
self._lock = threading.Lock()
self._hits: dict[str, list[float]] = {}
def allow(self, key: str) -> bool:
"""Return True when ``key`` is under the window cap; else False."""
if not isinstance(key, str) or not key:
return False
now = self._clock()
window_start = now - self._window_seconds
with self._lock:
hits = [hit for hit in self._hits.get(key, []) if hit > window_start]
allowed = len(hits) < self._max_requests
if allowed:
hits.append(now)
self._hits[key] = hits
return allowed
def reset(self, key: str) -> None:
"""Drop the recorded hits for ``key`` (tests/diagnostics)."""
with self._lock:
self._hits.pop(key, None)