from __future__ import annotations

import json
import os
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any

from .config import get_mcp_home


def _utc_now() -> datetime:
    return datetime.now(timezone.utc)


def _parse_utc(value: Any) -> datetime | None:
    if not isinstance(value, str) or not value.strip():
        return None
    normalized = value.strip().replace("Z", "+00:00")
    try:
        parsed = datetime.fromisoformat(normalized)
    except ValueError:
        return None
    if parsed.tzinfo is None:
        return parsed.replace(tzinfo=timezone.utc)
    return parsed.astimezone(timezone.utc)


def _json_safe_key(value: str) -> str:
    keep = []
    for char in value:
        if char.isalnum() or char in {"-", "_"}:
            keep.append(char)
        else:
            keep.append("_")
    result = "".join(keep).strip("_")
    return result or "entry"


def _store_dir(env_var: str, default_name: str) -> Path:
    custom = os.getenv(env_var)
    if custom:
        return Path(custom).expanduser()
    return get_mcp_home() / default_name


@dataclass(slots=True)
class _JsonEntryStore:
    base_dir: Path
    ttl: timedelta

    def __post_init__(self) -> None:
        self.base_dir.mkdir(parents=True, exist_ok=True)
        self.prune()

    def put(self, entry_id: str, payload: dict[str, Any]) -> None:
        data = dict(payload)
        data["id"] = entry_id
        data["updated_at"] = _utc_now().isoformat()
        self._path(entry_id).write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")

    def get(self, entry_id: str) -> dict[str, Any] | None:
        path = self._path(entry_id)
        if not path.exists():
            return None
        try:
            payload = json.loads(path.read_text(encoding="utf-8"))
        except (OSError, json.JSONDecodeError):
            path.unlink(missing_ok=True)
            return None
        created_at = _parse_utc(payload.get("created_at")) or _parse_utc(payload.get("updated_at"))
        if created_at is None or _utc_now() - created_at > self.ttl:
            path.unlink(missing_ok=True)
            return None
        return payload

    def prune(self) -> None:
        for path in self.base_dir.glob("*.json"):
            try:
                payload = json.loads(path.read_text(encoding="utf-8"))
            except (OSError, json.JSONDecodeError):
                path.unlink(missing_ok=True)
                continue
            created_at = _parse_utc(payload.get("created_at")) or _parse_utc(payload.get("updated_at"))
            if created_at is None or _utc_now() - created_at > self.ttl:
                path.unlink(missing_ok=True)

    def list(self) -> list[dict[str, Any]]:
        entries: list[dict[str, Any]] = []
        self.prune()
        for path in self.base_dir.glob("*.json"):
            try:
                payload = json.loads(path.read_text(encoding="utf-8"))
            except (OSError, json.JSONDecodeError):
                continue
            created_at = _parse_utc(payload.get("created_at")) or _parse_utc(payload.get("updated_at"))
            if created_at is None:
                continue
            entries.append(payload)
        entries.sort(key=lambda item: item.get("created_at") or item.get("updated_at") or "", reverse=True)
        return entries

    def _path(self, entry_id: str) -> Path:
        return self.base_dir / f"{_json_safe_key(entry_id)}.json"


class ImportVerificationStore(_JsonEntryStore):
    def __init__(self, base_dir: Path | None = None, *, ttl_seconds: int = 3600) -> None:
        super().__init__(
            base_dir=base_dir or _store_dir("QINGFLOW_MCP_IMPORT_VERIFY_HOME", "import-verifications"),
            ttl=timedelta(seconds=ttl_seconds),
        )


class ImportJobStore(_JsonEntryStore):
    def __init__(self, base_dir: Path | None = None, *, ttl_seconds: int = 24 * 3600) -> None:
        super().__init__(
            base_dir=base_dir or _store_dir("QINGFLOW_MCP_IMPORT_JOB_HOME", "import-jobs"),
            ttl=timedelta(seconds=ttl_seconds),
        )
