from __future__ import annotations

import json
import os
from copy import deepcopy
from dataclasses import dataclass
from pathlib import Path
from typing import Any

from ..config import get_mcp_home
from .design_session import DesignStage
from .run_store import resolve_storage_path, utc_now


def get_design_sessions_path() -> Path:
    custom_home = os.getenv("QINGFLOW_MCP_DESIGN_HOME")
    if custom_home:
        return Path(custom_home).expanduser()
    return get_mcp_home() / "design-sessions"


@dataclass(slots=True)
class DesignSessionStore:
    path: Path
    data: dict[str, Any]

    @classmethod
    def open(cls, *, session_id: str, metadata: dict[str, Any] | None = None) -> "DesignSessionStore":
        base_dir = get_design_sessions_path()
        base_dir.mkdir(parents=True, exist_ok=True)
        path = resolve_storage_path(base_dir, key=session_id, id_field="session_id")
        if path.exists():
            data = json.loads(path.read_text(encoding="utf-8"))
            stored_session_id = data.get("session_id")
            if stored_session_id != session_id:
                raise ValueError(f"existing design session at '{path}' belongs to '{stored_session_id}', not '{session_id}'")
        else:
            data = {
                "session_id": session_id,
                "status": "active",
                "current_stage": DesignStage.discover.value,
                "metadata": metadata or {},
                "stage_payloads": {
                    DesignStage.discover.value: {},
                    DesignStage.design.value: {},
                    DesignStage.experience.value: {},
                },
                "stage_results": {},
                "merged_design_spec": {},
                "normalized_solution_spec": None,
                "execution_plan": None,
                "created_at": utc_now(),
                "updated_at": utc_now(),
            }
            path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")
        return cls(path=path, data=data)

    def set_metadata(self, metadata: dict[str, Any]) -> None:
        if metadata:
            self.data["metadata"] = deepcopy(metadata)
            self._flush()

    def set_stage_payload(self, stage: str, payload: dict[str, Any]) -> None:
        self.data.setdefault("stage_payloads", {})
        self.data["stage_payloads"][stage] = deepcopy(payload)
        self._flush()

    def get_stage_payload(self, stage: str) -> dict[str, Any]:
        return deepcopy(self.data.get("stage_payloads", {}).get(stage, {}))

    def update_progress(self, *, status: str, current_stage: str, stage_results: dict[str, Any], merged_design_spec: dict[str, Any]) -> None:
        self.data["status"] = status
        self.data["current_stage"] = current_stage
        self.data["stage_results"] = deepcopy(stage_results)
        self.data["merged_design_spec"] = deepcopy(merged_design_spec)
        self._flush()

    def mark_finalized(self, *, normalized_solution_spec: dict[str, Any], execution_plan: dict[str, Any]) -> None:
        self.data["status"] = "finalized"
        self.data["current_stage"] = DesignStage.finalize.value
        self.data["normalized_solution_spec"] = deepcopy(normalized_solution_spec)
        self.data["execution_plan"] = deepcopy(execution_plan)
        self._flush()

    def summary(self) -> dict[str, Any]:
        return {
            "session_id": self.data["session_id"],
            "status": self.data["status"],
            "current_stage": self.data["current_stage"],
            "metadata": deepcopy(self.data.get("metadata", {})),
            "stage_results": deepcopy(self.data.get("stage_results", {})),
            "merged_design_spec": deepcopy(self.data.get("merged_design_spec", {})),
            "normalized_solution_spec": deepcopy(self.data.get("normalized_solution_spec")),
            "execution_plan": deepcopy(self.data.get("execution_plan")),
            "session_path": str(self.path),
        }

    def _flush(self) -> None:
        self.data["updated_at"] = utc_now()
        self.path.write_text(json.dumps(self.data, ensure_ascii=False, indent=2), encoding="utf-8")
