from __future__ import annotations

from collections import defaultdict
from typing import Any, Dict, Iterable, List, Optional, Sequence

import pandas as pd
from sqlalchemy import func, or_, select
from sqlalchemy.orm import Session

from channel.url_canonical import normalize_plp_path
from db.models import CollectionLandingPage, MasterCollectionRegistry, MasterProduct, MasterProductCollectionMembership

MEMBERSHIP_MODE_NATIVE = "native"
MEMBERSHIP_MODE_SHARED = "shared"
VALID_MEMBERSHIP_MODES = {MEMBERSHIP_MODE_NATIVE, MEMBERSHIP_MODE_SHARED}


def normalize_membership_mode(value: Optional[str]) -> str:
    mode = str(value or MEMBERSHIP_MODE_NATIVE).strip().lower()
    if mode not in VALID_MEMBERSHIP_MODES:
        raise ValueError(f"membership_mode must be one of: {', '.join(sorted(VALID_MEMBERSHIP_MODES))}")
    return mode


def split_master_sku_values(value: Any) -> List[str]:
    if isinstance(value, (list, tuple, set)):
        items = value
    elif value is None:
        items = []
    else:
        text = str(value).replace(";", "\n").replace(",", "\n")
        items = text.splitlines()
    normalized: List[str] = []
    seen = set()
    for item in items:
        sku = str(item or "").strip().upper()
        if sku and sku not in seen:
            seen.add(sku)
            normalized.append(sku)
    return normalized


def _split_path_values(value: Any) -> List[str]:
    if isinstance(value, (list, tuple, set)):
        items = value
    elif value is None:
        items = []
    else:
        text = str(value).replace(";", "\n").replace(",", "\n")
        items = text.splitlines()
    normalized: List[str] = []
    seen = set()
    for item in items:
        slug = normalize_plp_path(str(item or "").strip())
        if slug and slug not in seen:
            seen.add(slug)
            normalized.append(slug)
    return normalized


def _resolve_membership_targets(session: Session, path_slugs: Sequence[str]) -> Dict[str, Dict[str, Any]]:
    normalized = _split_path_values(path_slugs)
    if not normalized:
        return {}
    targets: Dict[str, Dict[str, Any]] = {}
    landing_rows = session.execute(
        select(
            CollectionLandingPage.path_slug,
            CollectionLandingPage.collection,
            CollectionLandingPage.collection_registry_id,
        ).where(
            CollectionLandingPage.path_slug.in_(normalized),
            CollectionLandingPage.is_active.is_(True),
        )
    ).all()
    for row in landing_rows:
        targets[str(row.path_slug)] = {
            "path_slug": str(row.path_slug),
            "collection_name": row.collection,
            "collection_registry_id": row.collection_registry_id,
        }

    unresolved = [slug for slug in normalized if slug not in targets]
    if unresolved:
        registry_rows = session.execute(
            select(
                MasterCollectionRegistry.path_slug,
                MasterCollectionRegistry.name,
                MasterCollectionRegistry.id,
            ).where(
                MasterCollectionRegistry.path_slug.in_(unresolved),
                MasterCollectionRegistry.is_active.is_(True),
            )
        ).all()
        for row in registry_rows:
            targets[str(row.path_slug)] = {
                "path_slug": str(row.path_slug),
                "collection_name": row.name,
                "collection_registry_id": row.id,
            }

    missing = [slug for slug in normalized if slug not in targets]
    if missing:
        raise ValueError(f"Unknown collection path slug(s): {', '.join(missing)}")

    return {slug: targets[slug] for slug in normalized}


def _require_product(session: Session, master_sku: str) -> MasterProduct:
    sku = str(master_sku or "").strip()
    if not sku:
        raise ValueError("master_sku is required")
    product = session.scalar(
        select(MasterProduct).where(
            MasterProduct.sku == sku,
        )
    )
    if not product:
        raise ValueError(f"Master SKU not found: {sku}")
    return product


def _serialize_memberships(rows: Iterable[MasterProductCollectionMembership]) -> List[Dict[str, Any]]:
    items = []
    for row in rows:
        items.append(
            {
                "path_slug": row.path_slug,
                "collection_registry_id": row.collection_registry_id,
                "collection_name": row.collection_name,
                "source_label": row.source_label,
                "notes": row.notes,
                "is_active": bool(row.is_active),
            }
        )
    items.sort(key=lambda item: (item["path_slug"] or "", item["collection_name"] or ""))
    return items


def get_shared_sku_membership_detail(session: Session, master_sku: str) -> Dict[str, Any]:
    product = _require_product(session, master_sku)
    membership_rows = session.scalars(
        select(MasterProductCollectionMembership).where(
            MasterProductCollectionMembership.master_sku == product.sku
        )
    ).all()
    active_rows = [row for row in membership_rows if row.is_active]
    inactive_rows = [row for row in membership_rows if not row.is_active]
    return {
        "master_sku": product.sku,
        "name": product.name,
        "category_l1": product.category_l1,
        "master_collection": product.collection,
        "membership_mode": normalize_membership_mode(product.membership_mode),
        "base_sku": product.base_sku,
        "variant_group_code": product.variant_group_code,
        "active_memberships": _serialize_memberships(active_rows),
        "inactive_memberships": _serialize_memberships(inactive_rows),
        "active_collection_path_slugs": [row.path_slug for row in sorted(active_rows, key=lambda item: item.path_slug or "")],
    }


def get_shared_sku_membership_details(session: Session, master_skus: Sequence[str]) -> List[Dict[str, Any]]:
    return [get_shared_sku_membership_detail(session, sku) for sku in split_master_sku_values(master_skus)]


def list_shared_sku_memberships(
    session: Session,
    *,
    search: Optional[str] = None,
    master_skus: Optional[Sequence[str]] = None,
    membership_mode: Optional[str] = MEMBERSHIP_MODE_SHARED,
    limit: int = 250,
) -> List[Dict[str, Any]]:
    stmt = select(MasterProduct)
    wanted_skus = split_master_sku_values(master_skus or [])
    if wanted_skus:
        stmt = stmt.where(func.upper(func.coalesce(MasterProduct.sku, "")).in_(wanted_skus))
    if membership_mode:
        stmt = stmt.where(func.lower(func.coalesce(MasterProduct.membership_mode, MEMBERSHIP_MODE_NATIVE)) == normalize_membership_mode(membership_mode))
    term = str(search or "").strip()
    if term and not wanted_skus:
        exact_skus = split_master_sku_values(term)
        if len(exact_skus) > 1:
            stmt = stmt.where(func.upper(func.coalesce(MasterProduct.sku, "")).in_(exact_skus))
        else:
            like = f"%{term.lower()}%"
            stmt = stmt.where(func.lower(func.coalesce(MasterProduct.sku, "")).like(like))
    products = session.scalars(stmt.order_by(MasterProduct.sku.asc()).limit(limit)).all()
    sku_list = [row.sku for row in products]
    memberships_by_sku: Dict[str, List[MasterProductCollectionMembership]] = defaultdict(list)
    if sku_list:
        membership_rows = session.scalars(
            select(MasterProductCollectionMembership).where(
                MasterProductCollectionMembership.master_sku.in_(sku_list),
                MasterProductCollectionMembership.is_active.is_(True),
            )
        ).all()
        for row in membership_rows:
            memberships_by_sku[str(row.master_sku)].append(row)

    rows: List[Dict[str, Any]] = []
    for product in products:
        memberships = sorted(memberships_by_sku.get(product.sku, []), key=lambda item: item.path_slug or "")
        rows.append(
            {
                "master_sku": product.sku,
                "name": product.name,
                "master_collection": product.collection,
                "membership_mode": normalize_membership_mode(product.membership_mode),
                "shared_collection_count": len(memberships),
                "collection_path_slugs": [row.path_slug for row in memberships],
                "collection_names": [row.collection_name or row.path_slug for row in memberships],
                "active_memberships": _serialize_memberships(memberships),
            }
        )
    return rows


def upsert_shared_sku_membership(
    session: Session,
    *,
    master_sku: str,
    membership_mode: Optional[str],
    collection_path_slugs: Optional[Sequence[str]] = None,
    notes: Optional[str] = None,
    source_label: Optional[str] = None,
) -> Dict[str, Any]:
    product = _require_product(session, master_sku)
    mode = normalize_membership_mode(membership_mode)
    targets = _resolve_membership_targets(session, collection_path_slugs or []) if mode == MEMBERSHIP_MODE_SHARED else {}
    if mode == MEMBERSHIP_MODE_SHARED and not targets:
        raise ValueError("Select at least one collection path for a shared SKU.")

    existing_rows = session.scalars(
        select(MasterProductCollectionMembership).where(
            MasterProductCollectionMembership.master_sku == product.sku
        )
    ).all()
    rows_by_slug = {str(row.path_slug): row for row in existing_rows}

    product.membership_mode = mode
    selected_slugs = set(targets)

    for row in existing_rows:
        if mode != MEMBERSHIP_MODE_SHARED or row.path_slug not in selected_slugs:
            row.is_active = False
            if notes is not None:
                row.notes = notes
            if source_label:
                row.source_label = source_label

    for slug, target in targets.items():
        row = rows_by_slug.get(slug)
        if row is None:
            row = MasterProductCollectionMembership(
                master_sku=product.sku,
                path_slug=slug,
            )
            session.add(row)
        row.collection_registry_id = target.get("collection_registry_id")
        row.collection_name = target.get("collection_name")
        row.source_label = source_label or row.source_label or "dashboard"
        row.notes = notes
        row.is_active = True

    session.flush()
    return get_shared_sku_membership_detail(session, product.sku)


def bulk_upsert_shared_sku_memberships(
    session: Session,
    *,
    master_skus: Sequence[str],
    membership_mode: Optional[str],
    collection_path_slugs: Optional[Sequence[str]] = None,
    notes: Optional[str] = None,
    source_label: Optional[str] = None,
) -> Dict[str, Any]:
    wanted_skus = split_master_sku_values(master_skus)
    items = []
    errors = []
    for sku in wanted_skus:
        try:
            items.append(
                upsert_shared_sku_membership(
                    session,
                    master_sku=sku,
                    membership_mode=membership_mode,
                    collection_path_slugs=collection_path_slugs,
                    notes=notes,
                    source_label=source_label,
                )
            )
        except ValueError as exc:
            errors.append({"master_sku": sku, "error": str(exc)})
    return {
        "processed": len(wanted_skus),
        "updated": len(items),
        "error_count": len(errors),
        "items": items,
        "errors": errors,
    }


def import_shared_sku_membership_dataframe(
    session: Session,
    df: pd.DataFrame,
    *,
    default_source_label: str = "shared_sku_upload",
) -> Dict[str, Any]:
    if df is None or df.empty:
        raise ValueError("CSV is empty.")
    rename = {str(column): str(column).strip().lower() for column in df.columns}
    frame = df.rename(columns=rename).copy()

    sku_column = next((name for name in ("master_sku", "sku") if name in frame.columns), None)
    path_column = next((name for name in ("collection_path_slug", "path_slug", "collection_path_slugs") if name in frame.columns), None)
    mode_column = next((name for name in ("membership_mode", "mode") if name in frame.columns), None)
    notes_column = "notes" if "notes" in frame.columns else None
    source_column = "source_label" if "source_label" in frame.columns else None

    if not sku_column:
        raise ValueError("CSV must include sku or master_sku column.")
    if not path_column:
        raise ValueError("CSV must include collection_path_slug, path_slug, or collection_path_slugs column.")

    grouped: Dict[str, Dict[str, Any]] = {}
    for idx, row in frame.iterrows():
        sku = str(row.get(sku_column) or "").strip()
        if not sku:
            continue
        bucket = grouped.setdefault(
            sku,
            {
                "membership_mode": MEMBERSHIP_MODE_SHARED,
                "collection_path_slugs": [],
                "notes": None,
                "source_label": None,
            },
        )
        raw_mode = row.get(mode_column) if mode_column else None
        if str(raw_mode or "").strip():
            bucket["membership_mode"] = normalize_membership_mode(str(raw_mode))
        bucket["collection_path_slugs"].extend(_split_path_values(row.get(path_column)))
        if notes_column and str(row.get(notes_column) or "").strip():
            bucket["notes"] = str(row.get(notes_column)).strip()
        if source_column and str(row.get(source_column) or "").strip():
            bucket["source_label"] = str(row.get(source_column)).strip()

    updated = []
    errors = []
    for sku, payload in grouped.items():
        try:
            detail = upsert_shared_sku_membership(
                session,
                master_sku=sku,
                membership_mode=payload.get("membership_mode"),
                collection_path_slugs=payload.get("collection_path_slugs") or [],
                notes=payload.get("notes"),
                source_label=payload.get("source_label") or default_source_label,
            )
            updated.append(
                {
                    "master_sku": detail["master_sku"],
                    "membership_mode": detail["membership_mode"],
                    "active_collection_path_slugs": detail["active_collection_path_slugs"],
                }
            )
        except ValueError as exc:
            errors.append({"master_sku": sku, "error": str(exc)})

    return {
        "processed": len(grouped),
        "updated": len(updated),
        "error_count": len(errors),
        "rows": updated,
        "errors": errors[:100],
    }
