Files
2026-09-19 12:54:45 +08:00

114 lines
3.6 KiB
Python

"""Prompt template storage, versioning, and rendering service."""
from typing import Any, Dict, List, Optional
from jinja2 import BaseLoader, Environment
from sqlalchemy.orm import Session
from app.models import PromptTemplate, PromptTemplateVersion
class PromptService:
def __init__(self, db: Session):
self.db = db
self.jinja = Environment(loader=BaseLoader())
def get_or_create_template(self, strategy_id: str, level: str) -> PromptTemplate:
template = (
self.db.query(PromptTemplate)
.filter_by(strategy_id=strategy_id, level=level)
.first()
)
if not template:
template = PromptTemplate(strategy_id=strategy_id, level=level)
self.db.add(template)
self.db.commit()
self.db.refresh(template)
return template
def create_version(
self,
strategy_id: str,
level: str,
body: str,
variables_schema: Optional[Dict[str, Any]] = None,
) -> PromptTemplateVersion:
template = self.get_or_create_template(strategy_id, level)
next_version = (
self.db.query(PromptTemplateVersion)
.filter_by(template_id=template.id)
.count()
+ 1
)
version = PromptTemplateVersion(
template_id=template.id,
version_number=next_version,
body=body,
variables_schema=variables_schema or self._infer_schema(body),
)
self.db.add(version)
self.db.commit()
self.db.refresh(version)
return version
def get_version(self, version_id: str) -> Optional[PromptTemplateVersion]:
from uuid import UUID
try:
return self.db.query(PromptTemplateVersion).filter_by(id=UUID(version_id)).first()
except ValueError:
return None
def list_versions(self, strategy_id: str, level: str) -> List[PromptTemplateVersion]:
template = (
self.db.query(PromptTemplate)
.filter_by(strategy_id=strategy_id, level=level)
.first()
)
if not template:
return []
return (
self.db.query(PromptTemplateVersion)
.filter_by(template_id=template.id)
.order_by(PromptTemplateVersion.version_number)
.all()
)
def render(
self,
strategy_id: str,
version: Optional[int],
context: Dict[str, Any],
) -> str:
template = (
self.db.query(PromptTemplate)
.filter_by(strategy_id=strategy_id)
.first()
)
if not template:
raise ValueError(f"Prompt template not found: {strategy_id}")
query = self.db.query(PromptTemplateVersion).filter_by(template_id=template.id)
if version:
version_obj = query.filter_by(version_number=version).first()
else:
version_obj = query.order_by(PromptTemplateVersion.version_number.desc()).first()
if not version_obj:
raise ValueError(f"Prompt version not found: {strategy_id} v{version}")
jinja_template = self.jinja.from_string(version_obj.body)
return jinja_template.render(**context)
def _infer_schema(self, body: str) -> Dict[str, Any]:
"""Infer required variables from Jinja2 template."""
from jinja2.meta import find_undeclared_variables
ast = self.jinja.parse(body)
variables = find_undeclared_variables(ast)
return {var: {"type": "string"} for var in variables}
def get_prompt_service(db: Session) -> PromptService:
return PromptService(db)