first commit

This commit is contained in:
eeymoo
2026-09-19 12:54:45 +08:00
commit 6fc5b64077
126 changed files with 8601 additions and 0 deletions
+113
View File
@@ -0,0 +1,113 @@
"""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)