import json
from pathlib import Path

from sqlalchemy.orm import Session

from models.db_models import AppMetaModel, PresetModel
from presets.template_utils import extract_variables, render_template

DEFAULT_PRESETS_PATH = Path(__file__).parent / "default_presets.json"
PRESETS_SEEDED_KEY = "presets_seeded"


class PresetStore:
    def __init__(self, db: Session):
        self.db = db

    def _get_meta(self, key: str) -> str | None:
        row = self.db.get(AppMetaModel, key)
        return row.value if row else None

    def _set_meta(self, key: str, value: str) -> None:
        row = self.db.get(AppMetaModel, key)
        if row:
            row.value = value
        else:
            self.db.add(AppMetaModel(key=key, value=value))
        self.db.commit()

    def seed_defaults(self) -> None:
        """仅在首次安装时写入默认预设，之后不再自动恢复。"""
        if self._get_meta(PRESETS_SEEDED_KEY) == "1":
            return

        if self.db.query(PresetModel).count() > 0:
            self._set_meta(PRESETS_SEEDED_KEY, "1")
            return

        data = json.loads(DEFAULT_PRESETS_PATH.read_text(encoding="utf-8"))
        for item in data:
            self.db.add(PresetModel(**item))
        self.db.commit()
        self._set_meta(PRESETS_SEEDED_KEY, "1")

    def list_presets(self, category: str | None = None) -> list[PresetModel]:
        query = self.db.query(PresetModel).order_by(
            PresetModel.sort_order, PresetModel.created_at
        )
        if category:
            query = query.filter(PresetModel.category == category)
        return query.all()

    def get_preset(self, preset_id: str) -> PresetModel | None:
        return self.db.get(PresetModel, preset_id)

    def create_preset(self, data: dict) -> PresetModel:
        preset = PresetModel(**data)
        self.db.add(preset)
        self.db.commit()
        self.db.refresh(preset)
        return preset

    def update_preset(self, preset_id: str, data: dict) -> PresetModel | None:
        preset = self.get_preset(preset_id)
        if not preset:
            return None
        for key, value in data.items():
            if value is not None:
                setattr(preset, key, value)
        self.db.commit()
        self.db.refresh(preset)
        return preset

    def delete_preset(self, preset_id: str) -> bool:
        preset = self.get_preset(preset_id)
        if not preset:
            return False
        self.db.delete(preset)
        self.db.commit()
        return True

    def render_prompt(
        self, preset_id: str, variables: dict[str, str] | None = None
    ) -> tuple[PresetModel, str]:
        preset = self.get_preset(preset_id)
        if not preset:
            raise LookupError(f"预设不存在: {preset_id}")
        prompt = render_template(preset.prompt_template, variables)
        return preset, prompt

    @staticmethod
    def get_variables(template: str) -> list[str]:
        return extract_variables(template)
