Coverage for src\genesis\inference\prompt_registry.py: 100%
27 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 14:20 +0800
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 14:20 +0800
1from __future__ import annotations
3from typing import Any
5from jinja2 import Template
7from .types import Prompt
10class PromptRegistry:
11 """Prompt 模板库:注册/取用/版本管理/渲染(集中管理待迁移 prompts/ 目录)。"""
13 def __init__(self) -> None:
14 self._templates: dict[tuple[str, str], str] = {}
16 def register(self, name: str, version: str, template: str) -> None:
17 """注册(或覆盖)一个版本的模板。"""
18 self._templates[(name, version)] = template
20 def get(
21 self,
22 name: str,
23 version: str | None = None,
24 variables: dict[str, Any] | None = None,
25 ) -> str:
26 """取模板;version=None 返回该 name 最新注册版本;variables 非空时渲染。"""
27 if version is None:
28 versions = self.list_versions(name)
29 if not versions:
30 raise KeyError(f"prompt not found: {name}")
31 version = versions[-1]
32 key = (name, version)
33 if key not in self._templates:
34 raise KeyError(f"prompt version not found: {name}@{version}")
35 template = self._templates[key]
36 if variables:
37 return self.render(template, variables)
38 return template
40 def list_versions(self, name: str) -> list[str]:
41 """返回某 name 的已注册版本(按注册顺序)。"""
42 return [v for (n, v) in self._templates if n == name]
44 def render(self, template: str, variables: dict[str, Any]) -> str:
45 """用 jinja2 渲染模板。"""
46 from jinja2 import Template
48 return Template(template).render(**variables)