"""SnapshotService — 체크포인트 스냅샷 저장/목록/복원.

Phase 2.3 (architecture-refactor-final/02-final-roadmap.md §Phase 2.3).
기존 `steps.py:list_snapshots/create_snapshot/restore_snapshot`에서 이관.
"""
from __future__ import annotations

import json
import logging
import re
import shutil
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, List, Optional

from sqlalchemy import text as sql_text
from sqlalchemy.orm import Session as OrmSession

from app.core.config import settings
from app.core.errors import AppError
from app.core.step_catalog import get_all_downstream_recursive


logger = logging.getLogger(__name__)


_SNAP_TS_RE = re.compile(
    r"^manifest_(\d{8}_\d{6})(?:_[a-zA-Z0-9_-]+)?\.json$"
)

_VALID_STATUSES = {
    "completed", "failed", "partial", "stale", "pending", "running",
    "not_applicable", "skipped", "error",
}


class SnapshotService:
    def __init__(self, db: OrmSession, project_id: str, episode_id: str):
        self.db = db
        self.project_id = project_id
        self.episode_id = episode_id

    @property
    def _base(self) -> Path:
        return Path(settings.projects_dir) / self.project_id / "checkpoints" / "episodes" / self.episode_id

    # ── list ──────────────────────────────────────────────

    def list_versions(self, step_id: Optional[str] = None) -> Dict[str, Any]:
        base = self._base
        if not base.exists():
            return {"snapshots": []}

        step_dirs = [base / step_id] if step_id else sorted(base.iterdir())

        ts_map: Dict[str, List[Dict]] = {}
        for step_dir in step_dirs:
            if not step_dir.is_dir():
                continue
            sid = step_dir.name
            for f in sorted(step_dir.glob("manifest_*.json")):
                if "prerestore" in f.name:
                    continue
                m = _SNAP_TS_RE.match(f.name)
                if not m:
                    continue
                ts = m.group(1)
                ts_map.setdefault(ts, []).append({
                    "step_id": sid,
                    "file": f.name,
                    "size": f.stat().st_size,
                })

        snapshots = []
        for ts in sorted(ts_map.keys(), reverse=True):
            steps = ts_map[ts]
            snapshots.append({
                "version": ts,
                "timestamp": f"{ts[:4]}-{ts[4:6]}-{ts[6:8]} {ts[9:11]}:{ts[11:13]}:{ts[13:15]}",
                "step_count": len(steps),
                "steps": [s["step_id"] for s in steps],
            })

        return {"snapshots": snapshots}

    # ── create ────────────────────────────────────────────

    def save_snapshot(
        self,
        step_id: Optional[str] = None,
        label: Optional[str] = None,
    ) -> Dict[str, Any]:
        base = self._base
        if not base.exists():
            raise AppError(code="snapshot.no_checkpoints", message="체크포인트 디렉토리 없음", status_code=404)

        ts = datetime.now().strftime("%Y%m%d_%H%M%S")
        step_dirs = [base / step_id] if step_id else sorted(base.iterdir())
        saved: List[str] = []

        for step_dir in step_dirs:
            if not step_dir.is_dir():
                continue
            manifest = step_dir / "manifest.json"
            if not manifest.exists():
                continue
            suffix = f"_{label}" if label else ""
            archive = step_dir / f"manifest_{ts}{suffix}.json"
            shutil.copy2(str(manifest), str(archive))
            saved.append(step_dir.name)

        return {"version": ts, "label": label, "saved_steps": saved, "count": len(saved)}

    # ── restore ───────────────────────────────────────────

    def restore(self, version: str, step_id: Optional[str] = None) -> Dict[str, Any]:
        base = self._base
        if not base.exists():
            raise AppError(code="snapshot.not_found", message="체크포인트 디렉토리 없음", status_code=404)

        pattern = re.compile(
            rf"^manifest_{re.escape(version)}(?:_[a-zA-Z0-9_-]+)?\.json$"
        )
        step_dirs = [base / step_id] if step_id else sorted(base.iterdir())

        now_iso = datetime.now(timezone.utc).isoformat()
        restored: List[str] = []
        for step_dir in step_dirs:
            if not step_dir.is_dir():
                continue
            sid = step_dir.name
            candidates = [
                f for f in step_dir.iterdir()
                if pattern.match(f.name) and "prerestore" not in f.name
            ]
            if not candidates:
                continue
            archive = sorted(candidates)[-1]

            manifest = step_dir / "manifest.json"
            if manifest.exists():
                bak_ts = datetime.now().strftime("%Y%m%d_%H%M%S")
                shutil.copy2(str(manifest), str(step_dir / f"manifest_{bak_ts}_prerestore.json"))

            shutil.copy2(str(archive), str(manifest))

            # 아카이브의 실제 status 읽어서 step_run에 반영
            try:
                restored_data = json.loads(manifest.read_text(encoding="utf-8"))
                archived_status = restored_data.get("status") or "completed"
            except Exception:
                archived_status = "completed"
            if archived_status not in _VALID_STATUSES:
                archived_status = "completed"

            self.db.execute(sql_text(
                "UPDATE step_run SET status = :status, updated_at = :now "
                "WHERE project_id = :pid AND episode_id = :eid AND step_id = :sid"
            ), {
                "pid": self.project_id, "eid": self.episode_id, "sid": sid,
                "now": now_iso, "status": archived_status,
            })
            restored.append(sid)

        if not restored:
            raise AppError(code="snapshot.version_not_found", message=f"버전 {version}에 해당하는 스냅샷 없음", status_code=404)

        self.db.commit()

        if step_id:
            # 단일 step 복원 → 하위 의존 step들 stale
            downstream = get_all_downstream_recursive(step_id)
            if downstream:
                for ds_id in downstream:
                    self.db.execute(sql_text(
                        "UPDATE step_run SET status = 'stale', updated_at = :now "
                        "WHERE project_id = :pid AND episode_id = :eid AND step_id = :sid "
                        "AND status = 'completed'"
                    ), {
                        "pid": self.project_id, "eid": self.episode_id,
                        "sid": ds_id, "now": now_iso,
                    })
                self.db.commit()
                logger.info(
                    "Snapshot restore: marked %d downstream steps as stale for %s",
                    len(downstream), step_id,
                )
        else:
            # 전체 복원 → DB 파생 테이블 재구축 (Phase 2.1 orchestrator)
            try:
                from app.services.checkpoint_sync import orchestrate_full_sync
                orchestrate_full_sync(self.project_id, self.episode_id, self.db)
                logger.info(
                    "Snapshot restore: DB resync completed for %d steps",
                    len(restored),
                )
            except Exception as exc:
                logger.warning("Snapshot restore: DB resync failed — %s", exc)
                self.db.rollback()

        return {"version": version, "restored_steps": restored, "count": len(restored)}
