"""Versioned prompt loading with strict, traversal-safe front matter parsing.""" from __future__ import annotations from collections.abc import Iterator, Mapping from dataclasses import dataclass from pathlib import Path import re from types import MappingProxyType from typing import Any import yaml from yaml.composer import ComposerError from yaml.constructor import ConstructorError from yaml.events import AliasEvent from yaml.nodes import MappingNode _PROMPT_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$") _VERSION = re.compile( r"^[0-9]+\.[0-9]+\.[0-9]+(?:-[0-9A-Za-z.-]+)?(?:\+[0-9A-Za-z.-]+)?$" ) _OUTPUT_MODEL = re.compile(r"^[A-Za-z_][A-Za-z0-9_.]*$") _MAX_PROMPT_BYTES = 1_000_000 _MAX_FRONT_MATTER_BYTES = 65_536 class PromptRepositoryError(ValueError): """Base error for invalid or unsafe prompt repositories.""" class PromptFormatError(PromptRepositoryError): """Raised when one prompt does not satisfy the file contract.""" class DuplicatePromptIdError(PromptRepositoryError): """Raised when two files claim the same logical prompt identifier.""" class _StrictSafeLoader(yaml.SafeLoader): """SafeLoader variant that also rejects aliases and duplicate mapping keys.""" def compose_node(self, parent: Any, index: Any) -> Any: if self.check_event(AliasEvent): event = self.peek_event() raise ComposerError( None, None, "YAML aliases are not allowed in prompt front matter", event.start_mark, ) return super().compose_node(parent, index) def _construct_unique_mapping( loader: _StrictSafeLoader, node: MappingNode, deep: bool = False ) -> dict[str, Any]: if not isinstance(node, MappingNode): raise ConstructorError( None, None, "front matter must be a mapping", node.start_mark ) mapping: dict[str, Any] = {} for key_node, value_node in node.value: key = loader.construct_object(key_node, deep=deep) if not isinstance(key, str): raise ConstructorError( "while constructing prompt front matter", node.start_mark, "front matter keys must be strings", key_node.start_mark, ) if key in mapping: raise ConstructorError( "while constructing prompt front matter", node.start_mark, f"duplicate front matter key: {key!r}", key_node.start_mark, ) mapping[key] = loader.construct_object(value_node, deep=deep) return mapping _StrictSafeLoader.add_constructor( yaml.resolver.BaseResolver.DEFAULT_MAPPING_TAG, _construct_unique_mapping ) @dataclass(frozen=True, slots=True) class PromptTemplate: """One immutable, versioned prompt document.""" id: str version: str body: str output_model: str | None metadata: Mapping[str, str] source_path: Path @property def prompt_id(self) -> str: return self.id @property def content(self) -> str: return self.body class PromptRepository: """Eagerly validate and index the direct ``*.md`` children of a directory.""" def __init__(self, root: str | Path | None = None) -> None: # Prompt templates are package data so the default repository also works # after installation from a wheel. The top-level ``prompts/`` directory # is a development mirror whose byte-for-byte parity is covered by tests. default_root = Path(__file__).resolve().with_name("prompt_templates") requested_root = default_root if root is None else Path(root) if not requested_root.exists(): raise FileNotFoundError( f"prompt directory does not exist: {requested_root}" ) if not requested_root.is_dir(): raise NotADirectoryError( f"prompt repository is not a directory: {requested_root}" ) self.root = requested_root.resolve() self._prompts = MappingProxyType(self._read_all()) def _read_all(self) -> dict[str, PromptTemplate]: prompts: dict[str, PromptTemplate] = {} normalised_ids: dict[str, Path] = {} candidates = sorted( (item for item in self.root.iterdir() if item.suffix.casefold() == ".md"), key=lambda item: item.name.casefold(), ) for candidate in candidates: if candidate.is_symlink(): raise PromptRepositoryError( "symbolic links are not allowed in prompt repositories: " f"{candidate.name}" ) resolved = candidate.resolve(strict=True) if not resolved.is_relative_to(self.root) or not resolved.is_file(): raise PromptRepositoryError( f"prompt path escapes the repository: {candidate.name}" ) prompt = _parse_prompt_file(resolved) normalised_id = prompt.id.casefold() if normalised_id in normalised_ids: first = normalised_ids[normalised_id].name raise DuplicatePromptIdError( f"duplicate prompt id {prompt.id!r} in {first!r} and " f"{candidate.name!r}" ) prompts[prompt.id] = prompt normalised_ids[normalised_id] = candidate return prompts def get(self, prompt_id: str) -> PromptTemplate: _validate_prompt_id(prompt_id, context="requested prompt id") try: return self._prompts[prompt_id] except KeyError as exc: raise KeyError(f"unknown prompt id: {prompt_id!r}") from exc def load(self, prompt_id: str) -> PromptTemplate: """Alias for ``get`` retained for stage-runner readability.""" return self.get(prompt_id) def list_ids(self) -> tuple[str, ...]: return tuple(sorted(self._prompts, key=str.casefold)) def all(self) -> tuple[PromptTemplate, ...]: return tuple(self._prompts[prompt_id] for prompt_id in self.list_ids()) def __contains__(self, prompt_id: object) -> bool: return isinstance(prompt_id, str) and prompt_id in self._prompts def __iter__(self) -> Iterator[str]: return iter(self.list_ids()) def __len__(self) -> int: return len(self._prompts) def _parse_prompt_file(path: Path) -> PromptTemplate: size = path.stat().st_size if size > _MAX_PROMPT_BYTES: raise PromptFormatError(f"prompt file is too large: {path.name}") try: text = path.read_text(encoding="utf-8-sig") except UnicodeDecodeError as exc: raise PromptFormatError(f"prompt must be UTF-8: {path.name}") from exc lines = text.splitlines(keepends=True) if not lines or lines[0].strip() != "---": raise PromptFormatError(f"prompt is missing opening front matter: {path.name}") closing_index: int | None = None front_matter_size = 0 for index, line in enumerate(lines[1:], start=1): if line.strip() == "---": closing_index = index break front_matter_size += len(line.encode("utf-8")) if front_matter_size > _MAX_FRONT_MATTER_BYTES: raise PromptFormatError(f"prompt front matter is too large: {path.name}") if closing_index is None: raise PromptFormatError(f"prompt is missing closing front matter: {path.name}") front_matter = "".join(lines[1:closing_index]) try: loaded = yaml.load(front_matter, Loader=_StrictSafeLoader) except yaml.YAMLError as exc: raise PromptFormatError( f"invalid prompt front matter: {path.name}: {exc}" ) from exc if not isinstance(loaded, dict): raise PromptFormatError(f"prompt front matter must be a mapping: {path.name}") metadata: dict[str, str] = {} for key, value in loaded.items(): if not isinstance(key, str) or not isinstance(value, str): raise PromptFormatError( "front matter values must be strings in " f"{path.name}: {key!r}" ) metadata[key] = value.strip() prompt_id = metadata.get("id") version = metadata.get("version") if prompt_id is None or version is None: raise PromptFormatError( f"prompt front matter requires id and version: {path.name}" ) _validate_prompt_id(prompt_id, context=f"prompt id in {path.name}") if not _VERSION.fullmatch(version): raise PromptFormatError(f"invalid prompt version in {path.name}: {version!r}") output_model = metadata.get("output_model") if output_model is not None and not _OUTPUT_MODEL.fullmatch(output_model): raise PromptFormatError( f"invalid output_model in {path.name}: {output_model!r}" ) body = "".join(lines[closing_index + 1 :]).strip() if not body: raise PromptFormatError(f"prompt body must not be empty: {path.name}") return PromptTemplate( id=prompt_id, version=version, body=body, output_model=output_model, metadata=MappingProxyType(metadata), source_path=path, ) def _validate_prompt_id(prompt_id: str, *, context: str) -> None: if not isinstance(prompt_id, str) or not _PROMPT_ID.fullmatch(prompt_id): raise PromptRepositoryError(f"invalid {context}: {prompt_id!r}") __all__ = [ "DuplicatePromptIdError", "PromptFormatError", "PromptRepository", "PromptRepositoryError", "PromptTemplate", ]