from __future__ import annotations from typing import Any from jinja2 import Template from .types import Prompt class PromptRegistry: """Prompt 模板库:注册/取用/版本管理/渲染(集中管理待迁移 prompts/ 目录)。""" def __init__(self) -> None: self._templates: dict[tuple[str, str], str] = {} def register(self, name: str, version: str, template: str) -> None: """注册(或覆盖)一个版本的模板。""" self._templates[(name, version)] = template def get( self, name: str, version: str | None = None, variables: dict[str, Any] | None = None, ) -> str: """取模板;version=None 返回该 name 最新注册版本;variables 非空时渲染。""" if version is None: versions = self.list_versions(name) if not versions: raise KeyError(f"prompt not found: {name}") version = versions[-1] key = (name, version) if key not in self._templates: raise KeyError(f"prompt version not found: {name}@{version}") template = self._templates[key] if variables: return self.render(template, variables) return template def list_versions(self, name: str) -> list[str]: """返回某 name 的已注册版本(按注册顺序)。""" return [v for (n, v) in self._templates if n == name] def render(self, template: str, variables: dict[str, Any]) -> str: """用 jinja2 渲染模板。""" from jinja2 import Template return Template(template).render(**variables)