from __future__ import annotations

from dataclasses import asdict, dataclass
from typing import Any, Callable, Dict, List, Optional, Set, Tuple

from sqlalchemy import delete, select
from sqlalchemy.orm import Session

from db.channel_exports import resolve_pipeline_channel_code
from db.channel_sku_mapping import channel_sku_for_master, master_sku_for_channel_sku
from db.channel_assignments import remove_skus
from db.models import (
    ChannelPublishState,
    ChannelSkuMapping,
    MagentoCatalogState,
    MasterProduct,
    MasterProductRelation,
)
from db.product_channel_taxonomy import resolve_filtered_master_skus


ACTIVE_STATUS = "active"


@dataclass(frozen=True)
class RemovalCandidate:
    master_sku: Optional[str]
    channel_sku: str
    remote_id: str
    reason: str
    product_type: Optional[str] = None
    parent_sku: Optional[str] = None

    def as_dict(self) -> Dict[str, Any]:
        return asdict(self)


def _active_master_skus(session: Session) -> Set[str]:
    return {
        str(sku).strip()
        for (sku,) in session.execute(
            select(MasterProduct.sku).where(MasterProduct.is_active.is_(True))
        ).all()
        if str(sku).strip()
    }


def _relation_hints(session: Session, skus: Set[str]) -> Tuple[Set[str], Dict[str, str]]:
    """Return parent SKUs and child -> parent map from master_product_relation."""
    if not skus:
        return set(), {}
    parents: Set[str] = set()
    child_parent: Dict[str, str] = {}
    for parent_sku, child_sku in session.execute(
        select(MasterProductRelation.parent_sku, MasterProductRelation.child_sku).where(
            MasterProductRelation.parent_sku.in_(sorted(skus))
            | MasterProductRelation.child_sku.in_(sorted(skus))
        )
    ).all():
        parent = str(parent_sku or "").strip()
        child = str(child_sku or "").strip()
        if parent:
            parents.add(parent)
        if child and parent:
            child_parent[child] = parent
    return parents, child_parent


def _product_type_for_sku(
    sku: str,
    *,
    parent_skus: Set[str],
    child_parent: Dict[str, str],
) -> Tuple[Optional[str], Optional[str]]:
    if sku in parent_skus:
        return "configurable", None
    if sku in child_parent:
        return "simple", child_parent[sku]
    return "simple", None


def _sku_matches(sku: str, query: str, mode: str) -> bool:
    sku_norm = str(sku or "").strip().lower()
    query_norm = str(query or "").strip().lower()
    if not sku_norm or not query_norm:
        return False
    if mode == "exact":
        return sku_norm == query_norm
    if mode == "contains":
        return query_norm in sku_norm
    return sku_norm.startswith(query_norm)


def _magento_remote_entries(
    session: Session,
    connection_id: int,
) -> List[Tuple[str, str, Optional[str]]]:
    """channel_sku, remote_id, mapped master_sku from pulled Magento catalog."""
    rows = session.execute(
        select(
            MagentoCatalogState.sku,
            MagentoCatalogState.magento_product_id,
        ).where(MagentoCatalogState.connection_id == connection_id)
    ).all()
    entries: List[Tuple[str, str, Optional[str]]] = []
    for channel_sku, product_id in rows:
        sku_text = str(channel_sku or "").strip()
        if not sku_text or product_id is None:
            continue
        master = master_sku_for_channel_sku(
            session,
            sku_text,
            "magento",
            connection_id=connection_id,
        )
        entries.append((sku_text, str(product_id), master))
    return entries


def _shopify_remote_entries(
    session: Session,
    *,
    connection_id: Optional[int] = None,
) -> List[Tuple[str, str, Optional[str]]]:
    """channel_sku, remote_id, master_sku from publish state + mappings (no live API)."""
    by_channel_sku: Dict[str, Tuple[str, Optional[str]]] = {}

    for sku, remote_id in session.execute(
        select(ChannelPublishState.sku, ChannelPublishState.remote_id).where(
            ChannelPublishState.channel_code == "shopify",
            ChannelPublishState.remote_id.is_not(None),
        )
    ).all():
        master = str(sku or "").strip()
        remote = str(remote_id or "").strip()
        if not master or not remote:
            continue
        channel_sku = channel_sku_for_master(
            session,
            master,
            "shopify",
            connection_id=connection_id,
        )
        by_channel_sku[channel_sku] = (remote, master)

    stmt = select(
        ChannelSkuMapping.channel_sku,
        ChannelSkuMapping.remote_id,
        ChannelSkuMapping.master_sku,
    ).where(
        ChannelSkuMapping.channel_code == "shopify",
        ChannelSkuMapping.mapping_status == ACTIVE_STATUS,
        ChannelSkuMapping.remote_id.is_not(None),
    )
    if connection_id is not None:
        stmt = stmt.where(
            (ChannelSkuMapping.connection_id == connection_id)
            | (ChannelSkuMapping.connection_id.is_(None))
        )
    for channel_sku, remote_id, master_sku in session.execute(stmt).all():
        sku_text = str(channel_sku or "").strip()
        remote = str(remote_id or "").strip()
        master = str(master_sku or "").strip() or None
        if sku_text and remote:
            by_channel_sku.setdefault(sku_text, (remote, master))

    return [(sku, remote, master) for sku, (remote, master) in sorted(by_channel_sku.items())]


def _resolve_remote_entries(
    session: Session,
    channel_code: str,
    *,
    connection_id: Optional[int] = None,
) -> List[Tuple[str, str, Optional[str]]]:
    channel = resolve_pipeline_channel_code(channel_code)
    if channel == "magento":
        if not connection_id:
            return []
        return _magento_remote_entries(session, connection_id)
    if channel == "shopify":
        return _shopify_remote_entries(session, connection_id=connection_id)
    return []


def plan_product_removals(
    session: Session,
    channel_code: str,
    *,
    connection_id: Optional[int] = None,
    selection: str = "not_in_master",
    category_l1: Optional[str] = None,
    category_l2: Optional[str] = None,
    category_l3: Optional[str] = None,
    collection: Optional[str] = None,
    product_family: Optional[str] = None,
    search: Optional[str] = None,
    master_filters: Optional[Any] = None,
    only_assigned: bool = False,
    sku_filter: Optional[str] = None,
    sku_match_mode: str = "starts_with",
    limit: Optional[int] = None,
) -> Dict[str, Any]:
    """
    Plan remote product deletions using DB catalog state only.

    selection:
      - not_in_master: remote SKU has no active master catalog row
      - master_filter: active master SKUs matching filters that exist on the remote channel
      - sku_filter: remote channel SKU matches starts_with / contains / exact text
    """
    channel = resolve_pipeline_channel_code(channel_code)
    selection_norm = (selection or "not_in_master").strip().lower()
    if selection_norm not in {"not_in_master", "master_filter", "sku_filter"}:
        raise ValueError("selection must be 'not_in_master', 'master_filter', or 'sku_filter'")

    remote_entries = _resolve_remote_entries(session, channel, connection_id=connection_id)
    master_skus = _active_master_skus(session)
    candidates: List[RemovalCandidate] = []
    seen_remote: Set[str] = set()

    if selection_norm == "not_in_master":
        for channel_sku, remote_id, mapped_master in remote_entries:
            if remote_id in seen_remote:
                continue
            master = mapped_master or master_sku_for_channel_sku(
                session, channel_sku, channel, connection_id=connection_id
            )
            if master and master in master_skus:
                continue
            if channel_sku in master_skus:
                continue
            seen_remote.add(remote_id)
            lookup_sku = master or channel_sku
            parent_skus, child_parent = _relation_hints(session, {lookup_sku})
            product_type, parent_sku = _product_type_for_sku(
                lookup_sku, parent_skus=parent_skus, child_parent=child_parent
            )
            candidates.append(
                RemovalCandidate(
                    master_sku=master,
                    channel_sku=channel_sku,
                    remote_id=remote_id,
                    reason="not_in_master",
                    product_type=product_type,
                    parent_sku=parent_sku,
                )
            )
    elif selection_norm == "master_filter":
        filtered_masters = resolve_filtered_master_skus(
            session,
            channel_code=channel if only_assigned else None,
            only_assigned=only_assigned,
            category_l1=category_l1,
            category_l2=category_l2,
            category_l3=category_l3,
            collection=collection,
            product_family=product_family,
            search=search,
            master_filters=master_filters,
            limit=limit,
        )
        remote_by_channel_sku = {sku: (remote, master) for sku, remote, master in remote_entries}
        remote_by_master: Dict[str, Tuple[str, str]] = {}
        for channel_sku, remote_id, mapped_master in remote_entries:
            key = mapped_master or channel_sku
            if key:
                remote_by_master.setdefault(key, (channel_sku, remote_id))

        for master_sku in filtered_masters:
            channel_sku = channel_sku_for_master(
                session,
                master_sku,
                channel,
                connection_id=connection_id,
            )
            remote_id = None
            resolved_channel_sku = channel_sku
            if channel_sku in remote_by_channel_sku:
                remote_id, _ = remote_by_channel_sku[channel_sku]
            elif master_sku in remote_by_master:
                resolved_channel_sku, remote_id = remote_by_master[master_sku]
            if not remote_id or remote_id in seen_remote:
                continue
            seen_remote.add(remote_id)
            parent_skus, child_parent = _relation_hints(session, {master_sku})
            product_type, parent_sku = _product_type_for_sku(
                master_sku, parent_skus=parent_skus, child_parent=child_parent
            )
            candidates.append(
                RemovalCandidate(
                    master_sku=master_sku,
                    channel_sku=resolved_channel_sku,
                    remote_id=remote_id,
                    reason="master_filter",
                    product_type=product_type,
                    parent_sku=parent_sku,
                )
            )
    else:
        query = str(sku_filter or "").strip()
        mode = str(sku_match_mode or "starts_with").strip().lower()
        if mode not in {"starts_with", "contains", "exact"}:
            raise ValueError("sku_match_mode must be 'starts_with', 'contains', or 'exact'")
        if not query:
            raise ValueError("sku_filter is required when selection='sku_filter'")
        for channel_sku, remote_id, mapped_master in remote_entries:
            if remote_id in seen_remote:
                continue
            if not _sku_matches(channel_sku, query, mode):
                continue
            seen_remote.add(remote_id)
            lookup_sku = mapped_master or channel_sku
            parent_skus, child_parent = _relation_hints(session, {lookup_sku})
            product_type, parent_sku = _product_type_for_sku(
                lookup_sku, parent_skus=parent_skus, child_parent=child_parent
            )
            candidates.append(
                RemovalCandidate(
                    master_sku=mapped_master,
                    channel_sku=channel_sku,
                    remote_id=remote_id,
                    reason=f"sku_filter:{mode}:{query}",
                    product_type=product_type,
                    parent_sku=parent_sku,
                )
            )

    if selection_norm in {"not_in_master", "sku_filter"} and limit:
        candidates = candidates[: int(limit)]

    return {
        "channel_code": channel,
        "connection_id": connection_id,
        "selection": selection_norm,
        "remote_catalog_count": len(remote_entries),
        "active_master_count": len(master_skus),
        "candidate_count": len(candidates),
        "candidates": [item.as_dict() for item in candidates],
        "samples": [item.as_dict() for item in candidates[:25]],
    }


def _cleanup_local_state(
    session: Session,
    channel_code: str,
    *,
    connection_id: Optional[int],
    candidate: RemovalCandidate,
) -> None:
    channel = resolve_pipeline_channel_code(channel_code)
    channel_sku = candidate.channel_sku
    master_sku = candidate.master_sku or channel_sku

    if channel == "magento" and connection_id:
        session.execute(
            delete(MagentoCatalogState).where(
                MagentoCatalogState.connection_id == connection_id,
                MagentoCatalogState.sku == channel_sku,
            )
        )

    session.execute(
        delete(ChannelPublishState).where(
            ChannelPublishState.channel_code == channel,
            ChannelPublishState.sku == master_sku,
        )
    )

    if master_sku:
        remove_skus(session, skus=[master_sku], channel_code=channel, reason="remote_product_removed")


ProgressCallback = Callable[[Dict[str, Any]], None]


def _report_removal_progress(
    on_progress: Optional[ProgressCallback],
    *,
    total: int,
    completed: int,
    deleted: int,
    failed: int,
) -> None:
    if on_progress:
        on_progress(
            {
                "total": total,
                "completed": completed,
                "pushed": deleted,
                "failed": failed,
            }
        )


def plan_options_from_payload(payload: Dict[str, Any]) -> Dict[str, Any]:
    """Normalize API payload into plan_product_removals keyword args."""
    return {
        "selection": str(payload.get("selection") or "not_in_master"),
        "category_l1": payload.get("category_l1"),
        "category_l2": payload.get("category_l2"),
        "category_l3": payload.get("category_l3"),
        "collection": payload.get("collection"),
        "product_family": payload.get("product_family"),
        "search": payload.get("search"),
        "master_filters": payload.get("master_filters"),
        "only_assigned": bool(payload.get("only_assigned", False)),
        "sku_filter": payload.get("sku_filter"),
        "sku_match_mode": payload.get("sku_match_mode") or "starts_with",
        "limit": payload.get("limit"),
        "magento_action": payload.get("magento_action"),
        "cleanup_local": bool(payload.get("cleanup_local", True)),
    }


def plan_and_build_candidates(
    session: Session,
    channel_code: str,
    *,
    connection_id: Optional[int],
    payload: Dict[str, Any],
) -> Tuple[Dict[str, Any], List[RemovalCandidate]]:
    opts = plan_options_from_payload(payload)
    plan = plan_product_removals(
        session,
        channel_code,
        connection_id=connection_id,
        selection=opts["selection"],
        category_l1=opts["category_l1"],
        category_l2=opts["category_l2"],
        category_l3=opts["category_l3"],
        collection=opts["collection"],
        product_family=opts["product_family"],
        search=opts["search"],
        master_filters=opts["master_filters"],
        only_assigned=opts["only_assigned"],
        sku_filter=opts["sku_filter"],
        sku_match_mode=opts["sku_match_mode"],
        limit=opts["limit"],
    )
    candidates = [RemovalCandidate(**item) for item in plan.get("candidates") or []]
    return plan, candidates


def run_remove_products_job(
    session: Session,
    *,
    channel_code: str,
    connection_id: Optional[int],
    options: Dict[str, Any],
    dry_run: bool,
    on_progress: Optional[ProgressCallback] = None,
) -> Dict[str, Any]:
    """Execute remove-products work inside a channel_job worker."""
    payload = {**options, "connection_id": connection_id}
    plan, candidates = plan_and_build_candidates(
        session,
        channel_code,
        connection_id=connection_id,
        payload=payload,
    )
    if dry_run:
        return {
            "status": "completed",
            "backend": "channel_product_removal",
            "dry_run": True,
            "total_count": len(candidates),
            "would_delete": len(candidates),
            "plan": plan,
            "samples": [c.as_dict() for c in candidates[:25]],
        }

    channel = resolve_pipeline_channel_code(channel_code)
    if channel == "magento":
        if not connection_id:
            raise ValueError("connection_id is required for Magento removals")
        if not options.get("magento_action"):
            raise ValueError("magento_action is required ('hard' or 'disable')")
        execution = execute_magento_removals(
            session,
            connection_id=int(connection_id),
            candidates=candidates,
            dry_run=False,
            cleanup_local=bool(options.get("cleanup_local", True)),
            magento_action=options.get("magento_action"),
            on_progress=on_progress,
            commit_each=bool(on_progress),
        )
    elif channel == "shopify":
        execution = execute_shopify_removals(
            session,
            connection_id=int(connection_id) if connection_id is not None else None,
            candidates=candidates,
            dry_run=False,
            cleanup_local=bool(options.get("cleanup_local", True)),
            on_progress=on_progress,
            commit_each=bool(on_progress),
        )
    else:
        raise ValueError(f"Unsupported channel for remove_products: {channel_code}")

    deleted = execution.get("deleted") or 0
    failed = execution.get("failed") or 0
    total = len(candidates)
    action = execution.get("magento_action")
    message = (
        f"Removed {deleted}/{total} product(s)"
        + (f" ({action})" if action else "")
        + (f", {failed} failed" if failed else "")
        + "."
    )
    return {
        "status": "completed" if failed == 0 else "completed",
        "backend": "channel_product_removal",
        "total_count": total,
        "success_count": deleted,
        "error_count": failed,
        "deleted": deleted,
        "failed": failed,
        "magento_action": action,
        "plan": plan,
        "message": message,
        **{k: v for k, v in execution.items() if k not in {"deleted", "failed"}},
    }


def execute_magento_removals(
    session: Session,
    *,
    connection_id: int,
    candidates: List[RemovalCandidate],
    dry_run: bool = True,
    cleanup_local: bool = True,
    magento_action: Optional[str] = None,
    on_progress: Optional[ProgressCallback] = None,
    commit_each: bool = False,
) -> Dict[str, Any]:
    from db.magento_repositories import SqlAlchemyMagentoConnectionRepository
    from magento.deletion_service import execute_deletion, normalize_magento_removal_action
    from magento.magento_api import MagentoRestClient
    from magento.oauth_client import MagentoOAuthClient, build_magento_oauth_kwargs

    if dry_run:
        return {
            "channel_code": "magento",
            "connection_id": connection_id,
            "dry_run": True,
            "would_delete": len(candidates),
            "samples": [c.as_dict() for c in candidates[:25]],
        }

    action = normalize_magento_removal_action(magento_action)

    conn_repo = SqlAlchemyMagentoConnectionRepository(session)
    conn = conn_repo.get_for_sync(connection_id)
    if not conn:
        raise ValueError(f"Magento connection {connection_id} not found")

    oauth = MagentoOAuthClient(**build_magento_oauth_kwargs(conn))
    api = MagentoRestClient(oauth)

    deleted = 0
    failed = 0
    results: List[Dict[str, Any]] = []
    total = len(candidates)
    for index, candidate in enumerate(candidates):
        try:
            result = execute_deletion(
                api=api,
                sku=candidate.channel_sku,
                product_type=candidate.product_type,
                parent_sku=candidate.parent_sku,
                action=action,
                remote_id=str(candidate.remote_id or "") or None,
            )
            deleted += 1
            if cleanup_local:
                _cleanup_local_state(
                    session,
                    "magento",
                    connection_id=connection_id,
                    candidate=candidate,
                )
            results.append(
                {
                    "channel_sku": candidate.channel_sku,
                    "remote_id": candidate.remote_id,
                    "status": "ok",
                    "result": result,
                }
            )
        except Exception as exc:
            failed += 1
            results.append(
                {
                    "channel_sku": candidate.channel_sku,
                    "remote_id": candidate.remote_id,
                    "status": "failed",
                    "error": str(exc),
                }
            )
        completed = index + 1
        _report_removal_progress(
            on_progress,
            total=total,
            completed=completed,
            deleted=deleted,
            failed=failed,
        )
        if commit_each and cleanup_local:
            session.commit()

    return {
        "channel_code": "magento",
        "connection_id": connection_id,
        "dry_run": False,
        "magento_action": action,
        "deleted": deleted,
        "failed": failed,
        "results": results[:100],
    }


def execute_shopify_removals(
    session: Session,
    *,
    connection_id: Optional[int],
    candidates: List[RemovalCandidate],
    dry_run: bool = True,
    cleanup_local: bool = True,
    on_progress: Optional[ProgressCallback] = None,
    commit_each: bool = False,
) -> Dict[str, Any]:
    from db.models import ShopifyConnection
    from shopify.connections import build_client, get_active_connection
    from shopify.product_delete import delete_product_by_id

    if dry_run:
        return {
            "channel_code": "shopify",
            "connection_id": connection_id,
            "dry_run": True,
            "would_delete": len(candidates),
            "samples": [c.as_dict() for c in candidates[:25]],
        }

    connection = session.get(ShopifyConnection, connection_id) if connection_id else None
    if connection is None:
        connection = get_active_connection(session)
    if connection is None or str(connection.status or "").strip().lower() != "active":
        raise ValueError("No active Shopify connection")

    client = build_client(connection)
    deleted = 0
    failed = 0
    results: List[Dict[str, Any]] = []
    total = len(candidates)
    for index, candidate in enumerate(candidates):
        remote_id = str(candidate.remote_id or "").strip()
        if not remote_id:
            failed += 1
            results.append(
                {
                    "remote_id": remote_id,
                    "status": "failed",
                    "error": "missing product id",
                }
            )
        else:
            try:
                item = delete_product_by_id(client, remote_id)
            except Exception as exc:
                item = {"remote_id": remote_id, "status": "failed", "error": str(exc)}
            results.append(item)
            if item.get("status") == "ok":
                deleted += 1
                if cleanup_local:
                    _cleanup_local_state(
                        session,
                        "shopify",
                        connection_id=connection_id,
                        candidate=candidate,
                    )
            else:
                failed += 1
        completed = index + 1
        _report_removal_progress(
            on_progress,
            total=total,
            completed=completed,
            deleted=deleted,
            failed=failed,
        )
        if commit_each and cleanup_local:
            session.commit()

    return {
        "channel_code": "shopify",
        "connection_id": connection_id,
        "dry_run": False,
        "deleted": deleted,
        "failed": failed,
        "results": results[:100],
    }
