from __future__ import annotations from dataclasses import dataclass from datetime import UTC, datetime import json from functools import lru_cache from pathlib import Path from typing import Any from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.system_setting import SystemSetting AI_PROMPTS_CATEGORY = "ai_prompts" DEFAULT_PROMPTS_PATH = Path(__file__).with_name("default_prompts.json") @dataclass(frozen=True) class AIPromptDefinition: key: str label: str group: str version: str system_prompt: str prompt: str @dataclass(frozen=True) class EffectiveAIPrompt: key: str label: str group: str version: str default_system_prompt: str default_prompt: str system_prompt: str prompt: str is_custom: bool updated_at: str | None = None @lru_cache(maxsize=1) def list_prompt_definitions() -> tuple[AIPromptDefinition, ...]: raw_items = json.loads(DEFAULT_PROMPTS_PATH.read_text(encoding="utf-8")) return tuple( AIPromptDefinition( key=str(item["key"]), label=str(item["label"]), group=str(item["group"]), version=str(item["version"]), system_prompt=str(item.get("system_prompt") or ""), prompt=str(item.get("prompt") or ""), ) for item in raw_items ) def get_prompt_definition(task_key: str) -> AIPromptDefinition: for definition in list_prompt_definitions(): if definition.key == task_key: return definition raise KeyError(task_key) async def _get_prompt_setting(db: AsyncSession) -> SystemSetting | None: result = await db.execute( select(SystemSetting).where(SystemSetting.category == AI_PROMPTS_CATEGORY) ) return result.scalar_one_or_none() def _normalize_overrides(payload: dict[str, Any] | None) -> dict[str, dict[str, Any]]: raw = (payload or {}).get("overrides") if not isinstance(raw, dict): return {} return { str(key): dict(value) for key, value in raw.items() if isinstance(value, dict) } async def get_prompt_overrides(db: AsyncSession) -> dict[str, dict[str, Any]]: if not hasattr(db, "execute"): return {} setting = await _get_prompt_setting(db) return _normalize_overrides(setting.payload if setting else None) def _effective_prompt( definition: AIPromptDefinition, override: dict[str, Any] | None, ) -> EffectiveAIPrompt: override = override or {} custom_system = override.get("system_prompt") custom_prompt = override.get("prompt") has_custom_system = isinstance(custom_system, str) has_custom_prompt = isinstance(custom_prompt, str) return EffectiveAIPrompt( key=definition.key, label=definition.label, group=definition.group, version=definition.version, default_system_prompt=definition.system_prompt, default_prompt=definition.prompt, system_prompt=custom_system if has_custom_system else definition.system_prompt, prompt=custom_prompt if has_custom_prompt else definition.prompt, is_custom=has_custom_system or has_custom_prompt, updated_at=str(override.get("updated_at") or "") or None, ) async def list_effective_prompts(db: AsyncSession) -> list[EffectiveAIPrompt]: overrides = await get_prompt_overrides(db) return [ _effective_prompt(definition, overrides.get(definition.key)) for definition in list_prompt_definitions() ] async def get_effective_prompt(db: AsyncSession | None, task_key: str) -> EffectiveAIPrompt: definition = get_prompt_definition(task_key) if db is None: return _effective_prompt(definition, None) overrides = await get_prompt_overrides(db) return _effective_prompt(definition, overrides.get(task_key)) async def save_prompt_override( db: AsyncSession, task_key: str, *, system_prompt: str, prompt: str, ) -> EffectiveAIPrompt: definition = get_prompt_definition(task_key) setting = await _get_prompt_setting(db) payload = dict(setting.payload or {}) if setting else {} overrides = _normalize_overrides(payload) overrides[definition.key] = { "system_prompt": system_prompt, "prompt": prompt, "updated_at": datetime.now(UTC).isoformat().replace("+00:00", "Z"), } payload["overrides"] = overrides if setting is None: setting = SystemSetting(category=AI_PROMPTS_CATEGORY, payload=payload) db.add(setting) else: setting.payload = payload await db.commit() return _effective_prompt(definition, overrides[definition.key]) async def reset_prompt_override(db: AsyncSession, task_key: str) -> EffectiveAIPrompt: definition = get_prompt_definition(task_key) setting = await _get_prompt_setting(db) if setting is None: return _effective_prompt(definition, None) payload = dict(setting.payload or {}) overrides = _normalize_overrides(payload) overrides.pop(definition.key, None) payload["overrides"] = overrides setting.payload = payload await db.commit() return _effective_prompt(definition, None) def serialize_effective_prompt(prompt: EffectiveAIPrompt) -> dict[str, Any]: return { "key": prompt.key, "label": prompt.label, "group": prompt.group, "version": prompt.version, "default_system_prompt": prompt.default_system_prompt, "default_prompt": prompt.default_prompt, "system_prompt": prompt.system_prompt, "prompt": prompt.prompt, "is_custom": prompt.is_custom, "updated_at": prompt.updated_at, }