← Files Comic SolARCHIVED FILE

skills/comic-sol/scripts/project_io.py

22.9 KB · Sep 30, 2026 · 23:14 UTC

↓ Download file

"""Shared trust-boundary helpers for Comic Sol project input and paths."""

from __future__ import annotations

import errno
import json
import os
import re
import sys
import tempfile
import time
from contextlib import contextmanager
from pathlib import Path, PurePosixPath
from typing import BinaryIO, Iterator


MAX_SOURCE_BYTES = 200 * 1024
SOURCE_SUFFIXES = {".txt", ".md"}
_DRIVE = re.compile(r"^[A-Za-z]:")
_LOCK_RETRY_SECONDS = 0.05
# Windows byte-range locks are mandatory, so the locked byte must sit past any
# region readers touch. The PID metadata occupies the first bytes of the file.
_LOCK_BYTE_OFFSET = 4096
_O_NOFOLLOW = getattr(os, "O_NOFOLLOW", 0)
_HAS_NOFOLLOW = _O_NOFOLLOW != 0
_REPARSE_POINT = 0x400


class ProjectLock:
    """Cross-process advisory lock retained at ``.comic-sol.lock``."""

    def __init__(self, project_dir: Path, timeout: float = 10.0):
        self.project_dir = Path(project_dir)
        self.timeout = timeout
        self._handle: BinaryIO | None = None

    def __enter__(self) -> "ProjectLock":
        deadline = time.monotonic() + self.timeout
        path = self.project_dir / ".comic-sol.lock"
        try:
            descriptor = os.open(path, os.O_RDWR | os.O_CREAT | os.O_EXCL, 0o600)
        except FileExistsError:
            handle = self._open_retained(path)
        else:
            handle = os.fdopen(descriptor, "r+b")
            try:
                handle.write(b"\0")
                handle.flush()
            except BaseException:
                handle.close()
                raise
        self._handle = handle
        acquired = False
        try:
            while True:
                handle.seek(0, os.SEEK_END)
                if handle.tell() == 0:
                    # No metadata means writer crashed between create+write
                    # or truncate+write. Try to acquire flock directly; if
                    # we succeed, no real holder exists (stale lock).
                    try:
                        self._lock(handle)
                        acquired = True
                        break
                    except OSError:
                        pass  # real contention, keep waiting
                    if time.monotonic() >= deadline:
                        raise TimeoutError(
                            "project is locked by another process"
                        )
                else:
                    try:
                        self._lock(handle)
                        acquired = True
                        break
                    except OSError as error:
                        if not self._retryable(error):
                            raise
                        if time.monotonic() >= deadline:
                            raise TimeoutError(
                                "project is locked by another process"
                            ) from error
                remaining = max(0.0, deadline - time.monotonic())
                time.sleep(min(_LOCK_RETRY_SECONDS, remaining))
            handle.seek(0)
            handle.truncate()
            handle.write(f"{os.getpid()}\n".encode("ascii"))
            handle.flush()
            return self
        except BaseException:
            try:
                if acquired:
                    try:
                        self._unlock(handle)
                    except BaseException:
                        pass
            finally:
                handle.close()
                self._handle = None
            raise

    @staticmethod
    def _open_retained(path: Path) -> BinaryIO:
        """Reopen an existing lock file without ever following a symlink.

        A symlinked lock path would otherwise be truncated and overwritten with
        PID metadata when the lock is acquired.
        """
        if not _HAS_NOFOLLOW and path.is_symlink():
            raise ValueError("lock path must not be a symlink")
        try:
            descriptor = os.open(path, os.O_RDWR | _O_NOFOLLOW)
        except OSError as error:
            if error.errno in (errno.ELOOP, errno.EMLINK):
                raise ValueError("lock path must not be a symlink") from error
            raise
        return os.fdopen(descriptor, "r+b")

    @staticmethod
    def _lock(handle: BinaryIO) -> None:
        if os.name == "nt":
            import msvcrt

            handle.seek(_LOCK_BYTE_OFFSET)
            msvcrt.locking(handle.fileno(), msvcrt.LK_NBLCK, 1)
        else:
            import fcntl

            fcntl.flock(handle.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)

    @staticmethod
    def _retryable(error: OSError) -> bool:
        return error.errno in {errno.EACCES, errno.EAGAIN, errno.EDEADLK} or getattr(
            error, "winerror", None
        ) in {33, 36}

    @staticmethod
    def _unlock(handle: BinaryIO) -> None:
        if os.name == "nt":
            import msvcrt

            handle.seek(_LOCK_BYTE_OFFSET)
            msvcrt.locking(handle.fileno(), msvcrt.LK_UNLCK, 1)
        else:
            import fcntl

            fcntl.flock(handle.fileno(), fcntl.LOCK_UN)

    def __exit__(self, exc_type, exc, traceback) -> None:
        handle = self._handle
        self._handle = None
        if handle is None:
            return
        try:
            self._unlock(handle)
        finally:
            handle.close()


def validate_source_bytes(source: bytes, suffix: str | None = None) -> str:
    if not isinstance(source, bytes):
        raise TypeError("source must be bytes")
    if len(source) > MAX_SOURCE_BYTES:
        raise ValueError("source must be at most 200 KiB as UTF-8 bytes")
    if suffix is not None and suffix.lower() not in SOURCE_SUFFIXES:
        raise ValueError("source file must use .txt or .md")
    try:
        return source.decode("utf-8")
    except UnicodeDecodeError as error:
        raise ValueError("source must be valid UTF-8") from error


def contained_project_path(
    project_dir: Path,
    relative: str | Path,
    *,
    must_exist: bool = False,
) -> Path:
    text = os.fspath(relative).replace("\\", "/")
    if not text or text.startswith("/") or _DRIVE.match(text) or ".." in text.split("/"):
        raise ValueError("path must be a relative project path")
    root = Path(project_dir).resolve(strict=True)
    unresolved = root.joinpath(*PurePosixPath(text).parts)
    current = unresolved
    while current != root:
        if current.is_symlink():
            raise ValueError("project path must not contain symlinks")
        if os.name == "nt":
            try:
                attributes = getattr(
                    current.stat(follow_symlinks=False), "st_file_attributes", 0
                )
            except FileNotFoundError:
                attributes = 0
            if attributes & _REPARSE_POINT:
                raise ValueError(
                    "project path must not contain symlinks or reparse points"
                )
        current = current.parent
    candidate = unresolved.resolve(strict=must_exist)
    if candidate != root and root not in candidate.parents:
        raise ValueError("path escapes the project directory")
    return candidate


def _relative_parts(relative: str | Path) -> tuple[str, ...]:
    text = os.fspath(relative).replace("\\", "/")
    if not text or text.startswith("/") or _DRIVE.match(text):
        raise ValueError("path must be a relative project path")
    parts = PurePosixPath(text).parts
    if not parts or any(part in {"", ".", ".."} for part in parts):
        raise ValueError("path must be a relative project path")
    return parts


def _open_parent_fd(project_dir: Path, parts: tuple[str, ...], *, create: bool) -> tuple[int, str]:
    root = Path(project_dir).resolve(strict=True)
    if os.name == "nt" or not _HAS_NOFOLLOW:
        raise NotImplementedError
    flags = os.O_RDONLY | getattr(os, "O_DIRECTORY", 0) | _O_NOFOLLOW
    current = os.open(root, flags)
    try:
        for part in parts[:-1]:
            try:
                child = os.open(part, flags, dir_fd=current)
            except FileNotFoundError:
                if not create:
                    raise
                os.mkdir(part, 0o755, dir_fd=current)
                child = os.open(part, flags, dir_fd=current)
            os.close(current)
            current = child
        return current, parts[-1]
    except BaseException:
        os.close(current)
        raise


@contextmanager
def open_contained(project_dir: Path, relative: str | Path, *, flags: int = os.O_RDONLY, mode: int = 0) -> Iterator[BinaryIO]:
    """Open project file without following symlinked path components on POSIX."""
    parts = _relative_parts(relative)
    if os.name == "nt" or not _HAS_NOFOLLOW:
        path = contained_project_path(project_dir, relative, must_exist=not (flags & os.O_CREAT))
        stream = open_path_nofollow(path, flags=flags, mode=mode)
        try:
            yield stream
        finally:
            stream.close()
        return
    else:
        parent_fd, name = _open_parent_fd(project_dir, parts, create=bool(flags & os.O_CREAT))
        try:
            descriptor = os.open(name, flags | _O_NOFOLLOW, mode, dir_fd=parent_fd)
        finally:
            os.close(parent_fd)
    stream = os.fdopen(descriptor, "r+b" if flags & os.O_RDWR else "wb" if flags & os.O_WRONLY else "rb")
    try:
        yield stream
    finally:
        stream.close()


def open_path_nofollow(path: Path, *, flags: int = os.O_RDONLY, mode: int = 0) -> BinaryIO:
    """Open absolute path while refusing symlink components on POSIX."""
    path = Path(path)
    if not path.is_absolute():
        raise ValueError("path must be absolute")
    absolute = path.absolute()
    parts = absolute.parts
    if not absolute.is_absolute() or len(parts) < 2:
        raise ValueError("path must be absolute")
    if sys.platform == "darwin" and parts[:2] == ("/", "var"):
        absolute = Path("/private", *parts[1:])
        parts = absolute.parts
    if os.name == "nt" or not _HAS_NOFOLLOW:
        current = Path(parts[0])
        for part in parts[1:]:
            current /= part
            if current.is_symlink():
                raise ValueError("path must not contain symlinks or reparse points")
            try:
                attributes = getattr(current.stat(follow_symlinks=False), "st_file_attributes", 0)
            except AttributeError:
                attributes = 0
            if attributes & _REPARSE_POINT:
                raise ValueError("path must not contain symlinks or reparse points")
        return os.fdopen(os.open(absolute, flags, mode), "rb")
    directory_flags = os.O_RDONLY | getattr(os, "O_DIRECTORY", 0) | _O_NOFOLLOW
    current = os.open(parts[0], directory_flags)
    try:
        for part in parts[1:-1]:
            child = os.open(part, directory_flags, dir_fd=current)
            os.close(current)
            current = child
        descriptor = os.open(parts[-1], flags | _O_NOFOLLOW, mode, dir_fd=current)
    finally:
        os.close(current)
    return os.fdopen(descriptor, "rb" if not flags & os.O_WRONLY else "wb")


def read_contained_bytes(project_dir: Path, relative: str | Path) -> bytes:
    try:
        with open_contained(project_dir, relative) as stream:
            return stream.read()
    except OSError as error:
        if error.errno in (errno.ELOOP, errno.EMLINK):
            raise ValueError("project path must not contain symlinks") from error
        raise


def remove_contained(project_dir: Path, relative: str | Path) -> None:
    parts = _relative_parts(relative)
    if os.name == "nt" or not _HAS_NOFOLLOW:
        path = contained_project_path(project_dir, relative)
        path.unlink(missing_ok=True)
        return
    parent_fd, name = _open_parent_fd(project_dir, parts, create=False)
    try:
        os.unlink(name, dir_fd=parent_fd)
    except FileNotFoundError:
        pass
    finally:
        os.close(parent_fd)


def replace_contained(project_dir: Path, source: str | Path, destination: str | Path) -> None:
    """Atomically replace destination from source with no-follow parent traversal."""
    source_parts = _relative_parts(source)
    destination_parts = _relative_parts(destination)
    if os.name == "nt" or not _HAS_NOFOLLOW:
        source_path = contained_project_path(project_dir, source, must_exist=True)
        destination_path = contained_project_path(project_dir, destination)
        destination_path.parent.mkdir(parents=True, exist_ok=True)
        os.replace(source_path, destination_path)
        return
    source_fd, source_name = _open_parent_fd(project_dir, source_parts, create=False)
    try:
        destination_fd, destination_name = _open_parent_fd(project_dir, destination_parts, create=True)
        try:
            os.replace(source_name, destination_name, src_dir_fd=source_fd, dst_dir_fd=destination_fd)
        finally:
            os.close(destination_fd)
    finally:
        os.close(source_fd)


def fsync_directory(path: Path) -> None:
    """Persist directory metadata; Windows has no stdlib directory fsync."""
    if os.name == "nt":
        return
    descriptor = os.open(path, os.O_RDONLY)
    try:
        os.fsync(descriptor)
    finally:
        os.close(descriptor)


def durable_atomic_write(path: Path, payload: bytes) -> None:
    """Atomically publish bytes and durably persist file and directory metadata."""
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    temporary: Path | None = None
    try:
        with tempfile.NamedTemporaryFile(
            "wb",
            dir=path.parent,
            prefix=f".{path.name}.",
            suffix=".tmp",
            delete=False,
        ) as handle:
            temporary = Path(handle.name)
            handle.write(payload)
            handle.flush()
            os.fsync(handle.fileno())
        os.replace(temporary, path)
        temporary = None
        fsync_directory(path.parent)
    finally:
        if temporary is not None:
            temporary.unlink(missing_ok=True)


def _find_transaction_dir(transaction_dir: Path) -> int:
    """Return the next available numeric transaction ID."""
    biggest = 0
    if transaction_dir.is_dir():
        for entry in transaction_dir.iterdir():
            try:
                value = int(entry.name)
                if value > biggest:
                    biggest = value
            except (ValueError, OSError):
                pass
    return biggest + 1


def _canonical_json_bytes(value: object) -> bytes:
    return json.dumps(value, ensure_ascii=False, separators=(",", ":"), sort_keys=True).encode("utf-8")


class ProjectTransaction:
    """Durable journal-backed all-or-nothing batch of file replacements.

    Acquires ``ProjectLock`` on enter, creates a numbered transaction
    directory under ``logs/transactions/<id>/``, writes a durable canonical
    journal before the first replace, and either commits or rolls back on exit.
    """

    JOURNAL_SCHEMA_VERSION = "1.0"

    def __init__(self, project_dir: Path, operation: str) -> None:
        self.project_dir = Path(project_dir)
        self.operation = operation
        self._lock: ProjectLock | None = None
        self._dir: Path | None = None
        self._journal: list[dict] = []
        self._phase: str | None = None
        self._id: int | None = None

    def __enter__(self) -> "ProjectTransaction":
        self._lock = ProjectLock(self.project_dir).__enter__()
        try:
            base = contained_project_path(self.project_dir, "logs/transactions")
            base.mkdir(parents=True, exist_ok=True)
            self._id = _find_transaction_dir(base)
            self._dir = base / str(self._id)
            self._dir.mkdir(parents=True)
            self._phase = "staging"
            return self
        except BaseException:
            self._lock.__exit__(*sys.exc_info())
            raise

    def stage_bytes(self, relative: str, payload: bytes) -> None:
        """Back up old destination (if any) and store staged payload under the
        transaction directory, recording an entry in the in-memory journal."""
        if self._dir is None:
            raise RuntimeError("transaction not started")
        path = Path(relative)
        if path.is_absolute():
            raise ValueError("stage_bytes requires a relative path")
        # Reject traversal and use one canonical path for validation and publication.
        resolved = contained_project_path(self.project_dir, relative)
        if resolved.resolve() == self.project_dir.resolve():
            raise ValueError(f"path '{relative}' resolves to the project root")
        relative = resolved.relative_to(self.project_dir.resolve()).as_posix()
        dest = resolved
        index = len(self._journal) + 1
        backup_name = f"backup-{index:03d}-{path.name}"
        staged_name = f"staged-{index:03d}-{path.name}"
        backup_path = self._dir / backup_name
        staged_path = self._dir / staged_name
        if dest.is_file():
            durable_atomic_write(
                backup_path,
                read_contained_bytes(self.project_dir, relative),
            )
        durable_atomic_write(staged_path, payload)
        entry = {
            "path": relative,
            "backup": (
                f"logs/transactions/{self._id}/{backup_name}"
                if dest.is_file() else None
            ),
            "staged": f"logs/transactions/{self._id}/{staged_name}",
        }
        self._journal.append(entry)

    def commit(self) -> None:
        """Durably write the canonical journal, then atomically replace each
        target. On any replace failure, restore backups in reverse order."""
        if self._dir is None:
            raise RuntimeError("transaction not started")
        if self._phase != "staging":
            raise RuntimeError("transaction already committed or rolling back")
        self._phase = "publishing"
        self._write_journal()
        published: list[tuple[Path, dict]] = []
        try:
            for entry in self._journal:
                dest = contained_project_path(self.project_dir, entry["path"])
                staged = contained_project_path(
                    self.project_dir, entry["staged"], must_exist=True
                )
                if os.name == "nt" or not _HAS_NOFOLLOW:
                    dest.parent.mkdir(parents=True, exist_ok=True)
                    os.replace(staged, dest)
                else:
                    replace_contained(self.project_dir, entry["staged"], entry["path"])
                published.append((dest, entry))
                fsync_directory(dest.parent)
            self._phase = "committed"
            self._write_journal()
            self._cleanup()
        except BaseException:
            for dest, entry in reversed(published):
                if entry.get("backup"):
                    backup = contained_project_path(self.project_dir, entry["backup"])
                    if backup.is_file():
                        if os.name == "nt" or not _HAS_NOFOLLOW:
                            os.replace(backup, dest)
                        else:
                            replace_contained(self.project_dir, entry["backup"], entry["path"])
                        fsync_directory(dest.parent)
                else:
                    try:
                        remove_contained(self.project_dir, entry["path"])
                    except OSError:
                        pass
            self._phase = "rolled_back"
            self._write_journal()
            raise

    def _write_journal(self) -> None:
        if self._dir is None:
            return
        journal = {
            "schema_version": self.JOURNAL_SCHEMA_VERSION,
            "operation": self.operation,
            "phase": self._phase,
            "targets": self._journal,
        }
        durable_atomic_write(self._dir / "journal.json", _canonical_json_bytes(journal))

    def _cleanup(self) -> None:
        if self._dir is None or not self._dir.is_dir():
            return
        for child in self._dir.iterdir():
            try:
                child.unlink()
            except OSError:
                pass
        parent = self._dir.parent
        try:
            self._dir.rmdir()
        except OSError:
            pass
        fsync_directory(parent)
        self._dir = None

    def __exit__(self, exc_type, exc, traceback) -> None:
        try:
            if exc_type is None and self._phase == "staging":
                self.commit()
            elif exc_type is not None and self._phase in ("staging", "publishing"):
                self._phase = "rolled_back"
                self._write_journal()
            if self._phase in ("committed", "rolled_back"):
                self._cleanup()
        finally:
            lock = self._lock
            self._lock = None
            if lock is not None:
                lock.__exit__(exc_type, exc, traceback)

    @staticmethod
    def recover(project_dir: Path) -> None:
        """Roll back incomplete journals while holding the project lock."""
        project_dir = Path(project_dir)
        base = contained_project_path(project_dir, "logs/transactions")
        if not base.is_dir():
            return
        with ProjectLock(project_dir):
            ids: list[int] = []
            for entry in base.iterdir():
                try:
                    ids.append(int(entry.name))
                except (ValueError, OSError):
                    continue
            for tid in sorted(ids):
                tx_dir = base / str(tid)
                journal_path = tx_dir / "journal.json"
                if not journal_path.is_file():
                    continue
                try:
                    journal = json.loads(journal_path.read_text("utf-8"))
                except (json.JSONDecodeError, OSError):
                    continue
                phase = journal.get("phase")
                targets = journal.get("targets")
                if not isinstance(targets, list):
                    continue
                if phase in ("staging", "publishing", "rolled_back"):
                    for entry in reversed(targets):
                        dest = contained_project_path(project_dir, entry["path"])
                        backup_path = entry.get("backup")
                        if backup_path:
                            backup = contained_project_path(project_dir, backup_path)
                            if backup.is_file():
                                if os.name == "nt" or not _HAS_NOFOLLOW:
                                    os.replace(backup, dest)
                                else:
                                    replace_contained(project_dir, backup_path, entry["path"])
                                fsync_directory(dest.parent)
                        else:
                            try:
                                remove_contained(project_dir, entry["path"])
                                fsync_directory(dest.parent)
                            except OSError:
                                pass
                for child in tx_dir.iterdir():
                    try:
                        child.unlink()
                    except OSError:
                        pass
                try:
                    tx_dir.rmdir()
                except OSError:
                    pass
                fsync_directory(base)

SHA-256: 65aa782b22b3ac1e21e6bf83872a75cb8b1a58af5cee424056d9085ee47934b8