from __future__ import annotations

import json
import os
from copy import deepcopy
from dataclasses import dataclass
from datetime import datetime, timezone
from hashlib import sha256
from pathlib import Path
from typing import Any

from ..config import get_mcp_home


def get_bootstrap_runs_path() -> Path:
    custom_home = os.getenv("QINGFLOW_MCP_BOOTSTRAP_HOME")
    if custom_home:
        return Path(custom_home).expanduser()
    return get_mcp_home() / "bootstrap-runs"


def fingerprint_payload(payload: dict[str, Any]) -> str:
    text = json.dumps(payload, sort_keys=True, ensure_ascii=False, separators=(",", ":"))
    return sha256(text.encode("utf-8")).hexdigest()


def utc_now() -> str:
    return datetime.now(timezone.utc).isoformat()


def storage_file_stem(value: str) -> str:
    safe = sanitize_key(value).strip("_") or "item"
    safe = safe[:80].rstrip("_") or "item"
    digest = sha256(value.encode("utf-8")).hexdigest()[:12]
    return f"{safe}--{digest}"


def resolve_storage_path(base_dir: Path, *, key: str, id_field: str) -> Path:
    preferred = base_dir / f"{storage_file_stem(key)}.json"
    if preferred.exists():
        return preferred
    legacy = base_dir / f"{sanitize_key(key)}.json"
    if not legacy.exists():
        return preferred
    try:
        payload = json.loads(legacy.read_text(encoding="utf-8"))
    except Exception:
        return preferred
    if payload.get(id_field) == key:
        return legacy
    return preferred


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

    @classmethod
    def open(
        cls,
        *,
        idempotency_key: str,
        normalized_solution_spec: dict[str, Any],
        request_fingerprint: str,
        run_label: str | None,
        initial_artifacts: dict[str, Any] | None = None,
    ) -> "RunArtifactStore":
        base_dir = get_bootstrap_runs_path()
        base_dir.mkdir(parents=True, exist_ok=True)
        path = resolve_storage_path(base_dir, key=idempotency_key, id_field="idempotency_key")
        if path.exists():
            data = json.loads(path.read_text(encoding="utf-8"))
            stored_key = data.get("idempotency_key")
            if stored_key != idempotency_key:
                raise ValueError(f"existing run artifact at '{path}' belongs to '{stored_key}', not '{idempotency_key}'")
            existing_fingerprint = data.get("request_fingerprint")
            if isinstance(initial_artifacts, dict) and (not existing_fingerprint or existing_fingerprint == request_fingerprint):
                merged_artifacts = _merge_nested_dicts(data.get("artifacts", {}), initial_artifacts)
                if merged_artifacts != data.get("artifacts", {}):
                    data["artifacts"] = merged_artifacts
                    data["updated_at"] = utc_now()
                    path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")
        else:
            data = {
                "idempotency_key": idempotency_key,
                "request_fingerprint": request_fingerprint,
                "normalized_solution_spec": normalized_solution_spec,
                "run_label": run_label,
                "artifacts": deepcopy(initial_artifacts)
                if isinstance(initial_artifacts, dict)
                else {
                    "package": {},
                    "roles": {},
                    "apps": {},
                    "views": {},
                    "charts": {},
                    "records": {},
                    "portal": {},
                    "navigation": {},
                    "field_maps": {},
                },
                "steps": {},
                "errors": [],
                "status": "pending",
                "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 ensure_apply_fingerprint(self, request_fingerprint: str) -> None:
        existing = self.data.get("request_fingerprint")
        if existing and existing != request_fingerprint:
            raise ValueError("idempotency_key already exists with a different request fingerprint")

    def record_step_started(self, step_name: str, request_fingerprint: str, debug_context: dict[str, Any] | None = None) -> None:
        step = self.data["steps"].get(step_name, {})
        step.update(
            {
                "status": "running",
                "request_fingerprint": request_fingerprint,
                "started_at": step.get("started_at") or utc_now(),
                "finished_at": None,
                "error": None,
                "debug_context": deepcopy(debug_context) if debug_context is not None else step.get("debug_context"),
            }
        )
        self.data["steps"][step_name] = step
        self.data["status"] = "running"
        self._flush()

    def record_step_completed(
        self,
        step_name: str,
        artifact_keys: dict[str, Any] | None = None,
        result: dict[str, Any] | None = None,
        debug_context: dict[str, Any] | None = None,
    ) -> None:
        step = self.data["steps"].setdefault(step_name, {})
        step.update(
            {
                "status": "completed",
                "artifact_keys": artifact_keys or step.get("artifact_keys") or {},
                "result": result or step.get("result") or {},
                "finished_at": utc_now(),
                "error": None,
                "debug_context": deepcopy(debug_context) if debug_context is not None else step.get("debug_context"),
            }
        )
        self.data["errors"] = [entry for entry in self.data.get("errors", []) if entry.get("step_name") != step_name]
        self._flush()

    def record_step_failed(self, step_name: str, error: Any, debug_context: dict[str, Any] | None = None) -> None:
        error_text, error_payload = _normalize_error_payload(error)
        step = self.data["steps"].setdefault(step_name, {})
        step.update(
            {
                "status": "failed",
                "finished_at": utc_now(),
                "error": error_text,
                "debug_context": deepcopy(debug_context) if debug_context is not None else step.get("debug_context"),
            }
        )
        if error_payload is not None:
            step["error_payload"] = deepcopy(error_payload)
        entry = {
            "step_name": step_name,
            "error": error_text,
            "at": utc_now(),
            "debug_context": deepcopy(debug_context) if debug_context is not None else step.get("debug_context"),
        }
        if error_payload is not None:
            entry["error_payload"] = deepcopy(error_payload)
            if isinstance(error_payload.get("category"), str):
                entry["category"] = error_payload["category"]
            if error_payload.get("details") is not None:
                entry["detail"] = deepcopy(error_payload["details"])
            elif error_payload.get("message") is not None:
                entry["detail"] = deepcopy(error_payload["message"])
        self.data.setdefault("errors", []).append(entry)
        self.data["status"] = "failed"
        self._flush()

    def set_artifact(self, section: str, key: str, value: Any) -> None:
        self.data["artifacts"].setdefault(section, {})
        self.data["artifacts"][section][key] = deepcopy(value)
        self._flush()

    def get_artifact(self, section: str, key: str, default: Any = None) -> Any:
        return self.data.get("artifacts", {}).get(section, {}).get(key, default)

    def get_step_status(self, step_name: str) -> str | None:
        step = self.data.get("steps", {}).get(step_name)
        return step.get("status") if step else None

    def should_run(self, step_name: str, *, force: bool = False) -> bool:
        if force:
            return True
        return self.get_step_status(step_name) != "completed"

    def mark_finished(self, *, status: str) -> None:
        self.data["status"] = status
        self._flush()

    def summary(self) -> dict[str, Any]:
        return {
            "artifacts": deepcopy(self.data.get("artifacts", {})),
            "step_results": deepcopy(self.data.get("steps", {})),
            "errors": deepcopy(self.data.get("errors", [])),
            "status": self.data.get("status", "pending"),
            "run_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")


def sanitize_key(value: str) -> str:
    return "".join(ch if ch.isalnum() or ch in ("-", "_") else "_" for ch in value)


def _merge_nested_dicts(base: dict[str, Any], incoming: dict[str, Any]) -> dict[str, Any]:
    merged = deepcopy(base)
    for key, value in incoming.items():
        if isinstance(value, dict) and isinstance(merged.get(key), dict):
            merged[key] = _merge_nested_dicts(merged[key], value)
        else:
            merged[key] = deepcopy(value)
    return merged


def _normalize_error_payload(error: Any) -> tuple[str, dict[str, Any] | None]:
    if isinstance(error, dict):
        return json.dumps(error, ensure_ascii=False), deepcopy(error)
    text = str(error)
    try:
        parsed = json.loads(text)
    except Exception:
        return text, None
    if isinstance(parsed, dict):
        return text, parsed
    return text, None
