183 lines
5.6 KiB
Python
183 lines
5.6 KiB
Python
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,
|
|
}
|