feat: PromptRegistry 注册/版本/渲染(jinja2)
This commit is contained in:
@@ -0,0 +1,48 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user