"""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)