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

1from __future__ import annotations 

2 

3from typing import Any 

4 

5from jinja2 import Template 

6 

7from .types import Prompt 

8 

9 

10class PromptRegistry: 

11 """Prompt 模板库:注册/取用/版本管理/渲染(集中管理待迁移 prompts/ 目录)。""" 

12 

13 def __init__(self) -> None: 

14 self._templates: dict[tuple[str, str], str] = {} 

15 

16 def register(self, name: str, version: str, template: str) -> None: 

17 """注册(或覆盖)一个版本的模板。""" 

18 self._templates[(name, version)] = template 

19 

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 

39 

40 def list_versions(self, name: str) -> list[str]: 

41 """返回某 name 的已注册版本(按注册顺序)。""" 

42 return [v for (n, v) in self._templates if n == name] 

43 

44 def render(self, template: str, variables: dict[str, Any]) -> str: 

45 """用 jinja2 渲染模板。""" 

46 from jinja2 import Template 

47 

48 return Template(template).render(**variables)