114 lines
3.6 KiB
Python
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)
|