276 lines
9.4 KiB
Python
276 lines
9.4 KiB
Python
"""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",
|
|
]
|