"""
Collections service - persistent user-curated sets of assets.

Collections are stored as JSON files under `config.COLLECTIONS_DIR`.
"""

from __future__ import annotations

import builtins
import hashlib
import json
import os
import re
import tempfile
import threading
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from uuid import uuid4

from ...config import COLLECTIONS_DIR_PATH
from ...shared import ErrorCode, Result, classify_file, get_logger

logger = get_logger(__name__)

MIN_COLLECTION_ID_LEN = 6
MAX_COLLECTION_ID_LEN = 80
MAX_COLLECTION_NAME_LEN = 80

DEFAULT_MAX_COLLECTION_ITEMS = 50_000
MIN_COLLECTION_ITEMS = 100
HARD_MAX_COLLECTION_ITEMS = 500_000

_ID_RE = re.compile(rf"^[a-zA-Z0-9_-]{{{MIN_COLLECTION_ID_LEN},{MAX_COLLECTION_ID_LEN}}}$")
_USER_SEGMENT_RE = re.compile(r"[^A-Za-z0-9._-]+")
_DEFAULT_MAX_ITEMS = DEFAULT_MAX_COLLECTION_ITEMS
try:
    _MAX_COLLECTION_ITEMS = int(os.environ.get("MJR_COLLECTION_MAX_ITEMS", str(_DEFAULT_MAX_ITEMS)))
except Exception:
    _MAX_COLLECTION_ITEMS = _DEFAULT_MAX_ITEMS
_MAX_COLLECTION_ITEMS = max(MIN_COLLECTION_ITEMS, min(HARD_MAX_COLLECTION_ITEMS, int(_MAX_COLLECTION_ITEMS or _DEFAULT_MAX_ITEMS)))


def _now_iso() -> str:
    try:
        return time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
    except Exception:
        return ""


def _normalize_fp(fp: str) -> str:
    try:
        p = Path(fp)
        # Prefer strict resolution when the path exists (resolves symlinks/junctions).
        try:
            if p.exists():
                p = p.resolve(strict=True)
            else:
                p = p.resolve(strict=False)
        except Exception:
            p = p.resolve(strict=False)
        normalized = str(p)
        # Windows filesystems are case-insensitive by default; preserve case on POSIX.
        if os.name == "nt":
            return os.path.normcase(normalized)
        return normalized
    except Exception:
        fallback = str(fp or "").strip()
        if os.name == "nt":
            return os.path.normcase(fallback)
        return fallback


def _safe_id(value: str) -> str | None:
    s = str(value or "").strip()
    if not s:
        return None
    if not _ID_RE.match(s):
        return None
    return s


def _safe_name(value: str) -> str | None:
    s = str(value or "").strip()
    if not s:
        return None
    if len(s) > MAX_COLLECTION_NAME_LEN:
        s = s[:MAX_COLLECTION_NAME_LEN].strip()
    return s or None


def _current_user_id() -> str:
    try:
        from ...routes.core.security import _current_user_id as _get_current_user_id

        return str(_get_current_user_id() or "").strip()
    except Exception:
        return ""


def _safe_user_store_segment(user_id: str | None) -> str:
    raw = str(user_id or "").strip()
    if not raw:
        return ""
    cleaned = _USER_SEGMENT_RE.sub("_", raw).strip("._-")
    if not cleaned:
        cleaned = hashlib.sha256(raw.encode("utf-8", errors="ignore")).hexdigest()[:24]
    if len(cleaned) > 80:
        digest = hashlib.sha256(raw.encode("utf-8", errors="ignore")).hexdigest()[:16]
        cleaned = f"{cleaned[:48]}_{digest}"
    return cleaned


def collections_base_dir_for_user(
    user_id: str | None,
    *,
    base_dir: str | Path | None = None,
) -> Path:
    base = Path(base_dir) if base_dir is not None else Path(COLLECTIONS_DIR_PATH)
    segment = _safe_user_store_segment(user_id)
    if not segment:
        return base
    return base / "users" / segment


def _collection_path_for_base(collection_id: str, base_dir: str | Path) -> Path | None:
    cid = _safe_id(collection_id)
    if not cid:
        return None
    try:
        base = Path(base_dir).resolve(strict=False)
        p = (base / f"{cid}.json").resolve(strict=False)
        if p.parent != base:
            return None
        return p
    except Exception:
        return None


def _collection_path(collection_id: str) -> Path | None:
    return _collection_path_for_base(collection_id, COLLECTIONS_DIR_PATH)


@dataclass(frozen=True)
class CollectionSummary:
    """Lightweight summary row for collection listings."""
    id: str
    name: str
    count: int
    updated_at: str


def _read_collection_summary_row(path: Path) -> CollectionSummary | None:
    try:
        data = json.loads(path.read_text(encoding="utf-8"))
    except Exception:
        return None
    if not isinstance(data, dict):
        return None
    cid = str(data.get("id") or path.stem)
    if not _safe_id(cid):
        return None
    name = str(data.get("name") or "").strip() or cid
    count = int(len(data.get("items") or []))
    updated_at = str(data.get("updated_at") or "")
    return CollectionSummary(cid, name, count, updated_at)


class CollectionsService:
    """Filesystem-backed collections service (JSON files under `COLLECTIONS_DIR`)."""

    def __init__(self, base_dir: str | Path | None = None, *, user_id: str | None = None):
        self._base_root = Path(base_dir) if base_dir is not None else Path(COLLECTIONS_DIR_PATH)
        self._forced_user_id = str(user_id or "").strip()
        self._lock = threading.Lock()
        try:
            self._base_root.mkdir(parents=True, exist_ok=True)
        except Exception:
            pass

    def _effective_user_id(self) -> str:
        if self._forced_user_id:
            return self._forced_user_id
        return _current_user_id()

    def _effective_base_dir(self) -> Path:
        base = collections_base_dir_for_user(self._effective_user_id(), base_dir=self._base_root)
        try:
            base.mkdir(parents=True, exist_ok=True)
        except Exception:
            pass
        return base

    def _collection_path(self, collection_id: str) -> Path | None:
        return _collection_path_for_base(collection_id, self._effective_base_dir())

    def list(self) -> Result[builtins.list[dict[str, Any]]]:
        """List collections with basic metadata and item counts."""
        base_res = self._resolve_collections_base_dir()
        if not base_res.ok:
            return Result.Err(base_res.code or ErrorCode.DB_ERROR, base_res.error or "Collections dir unavailable")
        base = base_res.data
        if not isinstance(base, Path):
            return Result.Err(ErrorCode.DB_ERROR, "Collections dir unavailable")

        items: list[CollectionSummary] = []
        with self._lock:
            try:
                for p in base.glob("*.json"):
                    summary = _read_collection_summary_row(p)
                    if summary is not None:
                        items.append(summary)
            except Exception as exc:
                return Result.Err(ErrorCode.DB_ERROR, f"Failed to list collections: {exc}")

        # Sort newest-ish first (fallback to id)
        items.sort(key=lambda x: (x.updated_at or "", x.id), reverse=True)
        return Result.Ok([i.__dict__ for i in items])

    def _resolve_collections_base_dir(self) -> Result[Path]:
        try:
            return Result.Ok(self._effective_base_dir().resolve(strict=False))
        except Exception as exc:
            return Result.Err(ErrorCode.DB_ERROR, f"Collections dir unavailable: {exc}")

    def create(self, name: str) -> Result[dict[str, Any]]:
        """Create a new empty collection."""
        cname = _safe_name(name)
        if not cname:
            return Result.Err(ErrorCode.INVALID_INPUT, "Missing collection name")

        cid = uuid4().hex[:12]
        path = self._collection_path(cid)
        if not path:
            return Result.Err(ErrorCode.DB_ERROR, "Failed to allocate collection id")

        payload: dict[str, Any] = {
            "id": cid,
            "name": cname,
            "created_at": _now_iso(),
            "updated_at": _now_iso(),
            "items": [],
        }

        with self._lock:
            try:
                path.parent.mkdir(parents=True, exist_ok=True)
                path.write_text(json.dumps(payload, ensure_ascii=False, separators=(",", ":")), encoding="utf-8")
            except Exception as exc:
                return Result.Err(ErrorCode.DB_ERROR, f"Failed to create collection: {exc}")

        return Result.Ok({"id": cid, "name": cname})

    def get(self, collection_id: str) -> Result[dict[str, Any]]:
        """Get a collection by id (including items)."""
        path = self._collection_path(collection_id)
        if not path:
            return Result.Err(ErrorCode.INVALID_INPUT, "Invalid collection id")
        loaded = self._load_collection_data(path)
        if not loaded.ok:
            return loaded
        return self._build_collection_get_result(loaded.data or {}, collection_id)

    @staticmethod
    def _load_collection_data(path: Path) -> Result[dict[str, Any]]:
        try:
            return Result.Ok(json.loads(path.read_text(encoding="utf-8")))
        except FileNotFoundError:
            return Result.Err(ErrorCode.NOT_FOUND, "Collection not found")
        except Exception as exc:
            return Result.Err(ErrorCode.PARSE_ERROR, f"Failed to read collection: {exc}")

    @staticmethod
    def _build_collection_get_result(data: dict[str, Any], collection_id: str) -> Result[dict[str, Any]]:
        cid = str(data.get("id") or collection_id)
        if not _safe_id(cid):
            return Result.Err(ErrorCode.PARSE_ERROR, "Collection file is corrupted (bad id)")
        return Result.Ok(
            {
                "id": cid,
                "name": str(data.get("name") or "").strip() or cid,
                "created_at": data.get("created_at") or "",
                "updated_at": data.get("updated_at") or "",
                "items": data.get("items") if isinstance(data.get("items"), list) else [],
            }
        )

    def delete(self, collection_id: str) -> Result[bool]:
        """Delete a collection file by id."""
        path = self._collection_path(collection_id)
        if not path:
            return Result.Err(ErrorCode.INVALID_INPUT, "Invalid collection id")
        with self._lock:
            try:
                path.unlink(missing_ok=True)  # py3.8+; safe on 3.11+
            except TypeError:
                try:
                    if path.exists():
                        path.unlink()
                except Exception as exc:
                    return Result.Err(ErrorCode.DB_ERROR, f"Failed to delete collection: {exc}")
            except Exception as exc:
                return Result.Err(ErrorCode.DB_ERROR, f"Failed to delete collection: {exc}")
        return Result.Ok(True)

    def _load_for_update(self, collection_id: str) -> tuple[Path | None, dict[str, Any] | None, Result[Any] | None]:
        path = self._collection_path(collection_id)
        if not path:
            return None, None, Result.Err(ErrorCode.INVALID_INPUT, "Invalid collection id")
        try:
            data = json.loads(path.read_text(encoding="utf-8"))
        except FileNotFoundError:
            return path, None, Result.Err(ErrorCode.NOT_FOUND, "Collection not found")
        except Exception as exc:
            return path, None, Result.Err(ErrorCode.PARSE_ERROR, f"Failed to read collection: {exc}")
        if not isinstance(data, dict):
            return path, None, Result.Err(ErrorCode.PARSE_ERROR, "Invalid collection format")
        if not _safe_id(str(data.get("id") or collection_id)):
            return path, None, Result.Err(ErrorCode.PARSE_ERROR, "Collection file is corrupted (bad id)")
        if not isinstance(data.get("items"), list):
            data["items"] = []
        return path, data, None

    @staticmethod
    def _normalize_add_asset_item(asset: dict[str, Any]) -> dict[str, Any] | None:
        if not isinstance(asset, dict):
            return None
        fp = str(asset.get("filepath") or "").strip()
        if not fp:
            return None
        root_id = CollectionsService._normalize_asset_root_id(asset)
        kind = CollectionsService._normalize_asset_kind(asset, fp)
        return {
            "filepath": fp,
            "filename": str(asset.get("filename") or "").strip(),
            "subfolder": str(asset.get("subfolder") or "").strip(),
            "type": str(asset.get("type") or "output").strip().lower(),
            "root_id": root_id,
            "kind": kind,
            "added_at": _now_iso(),
        }

    @staticmethod
    def _normalize_asset_root_id(asset: dict[str, Any]) -> str | None:
        root_id = str(
            asset.get("root_id")
            or asset.get("rootId")
            or asset.get("custom_root_id")
            or ""
        ).strip()
        return root_id or None

    @staticmethod
    def _normalize_asset_kind(asset: dict[str, Any], filepath: str) -> str:
        return str(asset.get("kind") or classify_file(filepath) or "unknown").strip().lower()

    def _clean_add_assets_input(self, assets: builtins.list[dict[str, Any]]) -> builtins.list[dict[str, Any]]:
        cleaned: list[dict[str, Any]] = []
        for asset in assets:
            item = self._normalize_add_asset_item(asset)
            if item is None:
                continue
            cleaned.append(item)
        return cleaned

    @staticmethod
    def _existing_items_with_indexes(raw_items: Any) -> tuple[builtins.list[dict[str, Any]], set[str], set[str]]:
        existing: list[dict[str, Any]] = (
            [it for it in raw_items if isinstance(it, dict)] if isinstance(raw_items, list) else []
        )
        seen: set[str] = set()
        existing_seen: set[str] = set()
        for item in existing:
            existing_fp = item.get("filepath")
            if not existing_fp:
                continue
            key = _normalize_fp(str(existing_fp))
            seen.add(key)
            existing_seen.add(key)
        return existing, seen, existing_seen

    def _append_collection_items(
        self,
        existing: builtins.list[dict[str, Any]],
        cleaned: builtins.list[dict[str, Any]],
        seen: set[str],
        existing_seen: set[str],
    ) -> dict[str, int]:
        added = 0
        skipped_existing = 0
        skipped_duplicate = 0
        skipped_limit = 0
        for item in cleaned:
            if len(existing) >= _MAX_COLLECTION_ITEMS:
                skipped_limit += 1
                continue
            key = _normalize_fp(item["filepath"])
            if key in seen:
                if key in existing_seen:
                    skipped_existing += 1
                else:
                    skipped_duplicate += 1
                continue
            existing.append(item)
            seen.add(key)
            added += 1
        return {
            "added": int(added),
            "skipped_existing": int(skipped_existing),
            "skipped_duplicate": int(skipped_duplicate),
            "skipped_limit": int(skipped_limit),
        }

    @staticmethod
    def _write_collection_payload(path: Path, data: dict[str, Any]) -> Result[Any] | None:
        """Atomically write collection JSON: write to a temp file then os.replace() into place.

        Using a temp file in the same directory guarantees the rename is atomic on the same
        filesystem (single-device rename), preventing partial-write corruption if the process
        is killed mid-write.
        """
        try:
            payload = json.dumps(data, ensure_ascii=False, separators=(",", ":"))
            dir_path = path.parent
            fd, tmp_path = tempfile.mkstemp(dir=dir_path, suffix=".tmp")
            try:
                with os.fdopen(fd, "w", encoding="utf-8") as fh:
                    fh.write(payload)
                os.replace(tmp_path, path)
            except Exception:
                # Clean up the temp file if the rename or write failed.
                try:
                    os.unlink(tmp_path)
                except OSError:
                    pass
                raise
        except Exception as exc:
            return Result.Err(ErrorCode.DB_ERROR, f"Failed to update collection: {exc}")
        return None

    def _add_assets_sync(self, collection_id: str, cleaned: builtins.list[dict[str, Any]]) -> Result[dict[str, Any]]:
        """Synchronous core of add_assets — must be called from a thread (not the event loop).

        Holds ``self._lock`` for the entire read-mutate-write sequence so that no two
        concurrent callers can read the same state and then overwrite each other (lost-update
        race, H-15/H-16).
        """
        with self._lock:
            path, data, err = self._load_for_update(collection_id)
            if err:
                return err
            if path is None or data is None:
                return Result.Err(ErrorCode.INVALID_INPUT, "Collection load failed: missing path or data")

            existing, seen, existing_seen = self._existing_items_with_indexes(data.get("items"))
            counts = self._append_collection_items(existing, cleaned, seen, existing_seen)

            data["items"] = existing
            data["updated_at"] = _now_iso()

            write_err = self._write_collection_payload(path, data)
            if write_err:
                return write_err

        return Result.Ok(
            {
                "id": str(data.get("id") or collection_id),
                "added": counts["added"],
                "skipped_existing": counts["skipped_existing"],
                "skipped_duplicate": counts["skipped_duplicate"],
                "skipped_limit": counts["skipped_limit"],
                "max_items": int(_MAX_COLLECTION_ITEMS),
                "count": len(existing),
            }
        )

    def add_assets(self, collection_id: str, assets: builtins.list[dict[str, Any]]) -> Result[dict[str, Any]]:
        """Add assets (by filepath) to a collection (deduplicated, bounded).

        Synchronous; callers in an async context should wrap with ``asyncio.to_thread``.
        ``self._lock`` inside ``_add_assets_sync`` ensures the read-mutate-write
        sequence is atomic with respect to other concurrent callers (H-14/H-15/H-16).
        """
        if not isinstance(assets, list) or not assets:
            return Result.Err(ErrorCode.INVALID_INPUT, "No assets provided")

        cleaned = self._clean_add_assets_input(assets)
        if not cleaned:
            return Result.Err(ErrorCode.INVALID_INPUT, "No valid assets provided")

        return self._add_assets_sync(collection_id, cleaned)

    def _remove_filepaths_sync(self, collection_id: str, targets: set[str]) -> Result[dict[str, Any]]:
        """Synchronous core of remove_filepaths — must be called from a thread (not the event loop).

        Holds ``self._lock`` for the entire read-mutate-write sequence (H-15/H-16).
        """
        with self._lock:
            path, data, err = self._load_for_update(collection_id)
            if err:
                return err
            if path is None or data is None:
                return Result.Err(ErrorCode.INVALID_INPUT, "Collection load failed: missing path or data")

            existing = self._collection_items_as_dicts(data.get("items"))
            kept, removed = self._partition_removed_items(existing, targets)
            data["items"] = kept
            data["updated_at"] = _now_iso()
            write_err = self._write_collection_payload(path, data)
            if write_err:
                return write_err

        return Result.Ok({"id": str(data.get("id") or collection_id), "removed": int(removed), "count": len(kept)})

    def remove_filepaths(self, collection_id: str, filepaths: builtins.list[str]) -> Result[dict[str, Any]]:
        """Remove items from a collection by filepath.

        Synchronous; callers in an async context should wrap with ``asyncio.to_thread``.
        ``self._lock`` inside ``_remove_filepaths_sync`` ensures the read-mutate-write
        sequence is atomic with respect to other concurrent callers (H-14/H-15/H-16).
        """
        targets_res = self._remove_targets(filepaths)
        if not targets_res.ok:
            return Result.Err(targets_res.code or ErrorCode.INVALID_INPUT, targets_res.error or "No valid filepaths provided")
        targets = targets_res.data or set()

        return self._remove_filepaths_sync(collection_id, targets)

    @staticmethod
    def _remove_targets(filepaths: builtins.list[str]) -> Result[set[str]]:
        if not isinstance(filepaths, list) or not filepaths:
            return Result.Err(ErrorCode.INVALID_INPUT, "No filepaths provided")
        targets = set(_normalize_fp(str(p)) for p in filepaths if str(p or "").strip())
        if not targets:
            return Result.Err(ErrorCode.INVALID_INPUT, "No valid filepaths provided")
        return Result.Ok(targets)

    @staticmethod
    def _collection_items_as_dicts(raw_items: Any) -> builtins.list[dict[str, Any]]:
        return [it for it in raw_items if isinstance(it, dict)] if isinstance(raw_items, list) else []

    @staticmethod
    def _partition_removed_items(existing: builtins.list[dict[str, Any]], targets: set[str]) -> tuple[builtins.list[Any], int]:
        kept: list[Any] = []
        removed = 0
        for item in existing:
            fp = item.get("filepath")
            if fp and _normalize_fp(str(fp)) in targets:
                removed += 1
                continue
            kept.append(item)
        return kept, removed
