from __future__ import annotations

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

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

from channel.url_canonical import normalize_plp_path
from db.master_product_images import COLLECTION_IMAGE_ROLES
from db.models import (
    MasterCollectionRegistry,
    MasterProduct,
    MasterProductCollectionMembership,
    MasterProductImage,
)
from db.shared_sku_membership import MEMBERSHIP_MODE_SHARED, split_master_sku_values
from db.tribeca_sku_parse import collection_asset_key


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


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


def _dedupe_connections(connections: Iterable[Mapping[str, Any]]) -> List[Dict[str, Any]]:
    rows: List[Dict[str, Any]] = []
    seen = set()
    for raw in connections or []:
        compat_id = raw.get("id")
        native_id = raw.get("native_id")
        channel_type = str(raw.get("channel_type") or raw.get("channel_code") or "").strip().lower()
        if compat_id is None or native_id is None or channel_type not in {"magento", "shopify"}:
            continue
        key = (channel_type, int(compat_id), int(native_id))
        if key in seen:
            continue
        seen.add(key)
        rows.append(
            {
                "id": int(compat_id),
                "native_id": int(native_id),
                "channel_type": channel_type,
                "channel_code": str(raw.get("channel_code") or channel_type),
                "store_code": raw.get("store_code"),
            }
        )
    return rows


def resolve_shared_sku_image_cleanup_scope(
    session: Session,
    *,
    master_skus: Optional[Sequence[str]] = None,
    collection_path_slugs: Optional[Sequence[str]] = None,
    collection_codes: Optional[Sequence[str]] = None,
    all_shared: bool = False,
) -> Dict[str, Any]:
    wanted_skus = split_master_sku_values(master_skus or [])
    wanted_codes = _split_collection_codes(collection_codes or [])
    wanted_slugs = _split_path_slugs(collection_path_slugs or [])

    resolved_code_rows: List[Dict[str, str]] = []
    missing_codes: List[str] = []
    if wanted_codes:
        rows = session.execute(
            select(MasterCollectionRegistry.code, MasterCollectionRegistry.path_slug)
            .where(func.upper(MasterCollectionRegistry.code).in_(wanted_codes))
            .where(MasterCollectionRegistry.is_active.is_(True))
        ).all()
        code_map = {
            str(row.code).strip().upper(): normalize_plp_path(str(row.path_slug or ""))
            for row in rows
            if normalize_plp_path(str(row.path_slug or ""))
        }
        missing_codes = [code for code in wanted_codes if code not in code_map]
        wanted_slugs = sorted(set(wanted_slugs).union(code_map.values()))
        resolved_code_rows = [
            {"collection_code": code, "path_slug": code_map[code]}
            for code in wanted_codes
            if code in code_map
        ]

    if not all_shared and not wanted_skus and not wanted_slugs:
        raise ValueError(
            "Provide master_skus, collection_path_slugs, collection_codes, or all_shared=true."
        )
    if missing_codes:
        raise ValueError(f"Unknown active collection code(s): {', '.join(missing_codes)}")

    stmt = (
        select(MasterProduct.sku)
        .where(MasterProduct.is_active.is_(True))
        .where(
            func.lower(func.coalesce(MasterProduct.membership_mode, "native"))
            == MEMBERSHIP_MODE_SHARED
        )
    )
    if wanted_skus:
        stmt = stmt.where(func.upper(MasterProduct.sku).in_(wanted_skus))
    if wanted_slugs:
        stmt = stmt.join(
            MasterProductCollectionMembership,
            MasterProductCollectionMembership.master_sku == MasterProduct.sku,
        )
        stmt = stmt.where(MasterProductCollectionMembership.is_active.is_(True))
        stmt = stmt.where(MasterProductCollectionMembership.path_slug.in_(wanted_slugs))

    target_skus = sorted({str(sku).strip().upper() for sku in session.scalars(stmt).all() if str(sku or "").strip()})
    memberships_by_sku: Dict[str, List[str]] = defaultdict(list)
    if target_skus:
        membership_rows = session.execute(
            select(
                MasterProductCollectionMembership.master_sku,
                MasterProductCollectionMembership.path_slug,
            )
            .where(MasterProductCollectionMembership.master_sku.in_(target_skus))
            .where(MasterProductCollectionMembership.is_active.is_(True))
        ).all()
        for row in membership_rows:
            sku = str(row.master_sku).strip().upper()
            slug = normalize_plp_path(str(row.path_slug or ""))
            if slug:
                memberships_by_sku[sku].append(slug)

    for sku, slugs in memberships_by_sku.items():
        memberships_by_sku[sku] = sorted(set(slugs))

    return {
        "all_shared": bool(all_shared),
        "requested_master_skus": wanted_skus,
        "requested_collection_codes": wanted_codes,
        "requested_collection_path_slugs": _split_path_slugs(collection_path_slugs or []),
        "resolved_collection_paths": wanted_slugs,
        "resolved_collection_code_paths": resolved_code_rows,
        "target_skus": target_skus,
        "target_sku_count": len(target_skus),
        "memberships_by_sku": dict(memberships_by_sku),
    }


def cleanup_shared_sku_collection_images(
    session: Session,
    *,
    master_skus: Optional[Sequence[str]] = None,
    collection_path_slugs: Optional[Sequence[str]] = None,
    collection_codes: Optional[Sequence[str]] = None,
    all_shared: bool = False,
    dry_run: bool = True,
) -> Dict[str, Any]:
    scope = resolve_shared_sku_image_cleanup_scope(
        session,
        master_skus=master_skus,
        collection_path_slugs=collection_path_slugs,
        collection_codes=collection_codes,
        all_shared=all_shared,
    )
    target_skus = list(scope["target_skus"])
    if not target_skus:
        return {
            "status": "ok",
            "dry_run": bool(dry_run),
            "scope": scope,
            "image_rows_scanned": 0,
            "removable_row_count": 0,
            "deleted_row_count": 0,
            "kept_row_count": 0,
            "sku_breakdown": [],
            "sample_removed": [],
            "message": "No shared SKUs matched the requested cleanup scope.",
        }

    rows = session.scalars(
        select(MasterProductImage)
        .where(MasterProductImage.sku.in_(target_skus))
        .order_by(MasterProductImage.sku.asc(), MasterProductImage.sort_order.asc(), MasterProductImage.id.asc())
    ).all()

    removable_ids: List[int] = []
    removable_samples: List[Dict[str, Any]] = []
    sku_stats: Dict[str, Dict[str, Any]] = defaultdict(
        lambda: {"sku": "", "rows_scanned": 0, "removable_rows": 0, "kept_rows": 0}
    )
    for row in rows:
        sku = str(row.sku or "").strip().upper()
        stats = sku_stats[sku]
        stats["sku"] = sku
        stats["rows_scanned"] += 1
        role = str(row.image_role or "").strip().lower()
        removable = role in COLLECTION_IMAGE_ROLES or bool(collection_asset_key(row.file_name))
        if removable:
            removable_ids.append(int(row.id))
            stats["removable_rows"] += 1
            if len(removable_samples) < 25:
                removable_samples.append(
                    {
                        "id": row.id,
                        "sku": sku,
                        "file_name": row.file_name,
                        "image_role": row.image_role,
                        "image_url": row.image_url,
                    }
                )
        else:
            stats["kept_rows"] += 1

    deleted_count = 0
    if removable_ids and not dry_run:
        result = session.execute(delete(MasterProductImage).where(MasterProductImage.id.in_(removable_ids)))
        deleted_count = int(result.rowcount or len(removable_ids))

    breakdown = [sku_stats.get(sku, {"sku": sku, "rows_scanned": 0, "removable_rows": 0, "kept_rows": 0}) for sku in target_skus]
    return {
        "status": "ok",
        "dry_run": bool(dry_run),
        "scope": scope,
        "image_rows_scanned": len(rows),
        "removable_row_count": len(removable_ids),
        "deleted_row_count": deleted_count,
        "kept_row_count": len(rows) - len(removable_ids),
        "sku_breakdown": breakdown,
        "sample_removed": removable_samples,
    }


def enqueue_shared_sku_image_cleanup_outbound(
    session: Session,
    *,
    skus: Sequence[str],
    magento_connections: Sequence[Mapping[str, Any]] = (),
    shopify_connections: Sequence[Mapping[str, Any]] = (),
    dry_run: bool = True,
    prune_magento_orphans: bool = True,
    push_magento_images: bool = True,
    push_shopify_images: bool = True,
    batch_size: int = 250,
    notes: str = "shared_sku_image_cleanup",
) -> Dict[str, Any]:
    from app.jobs.channel_jobs import (
        enqueue_channel_job,
        enqueue_magento_media_cleanup_job,
        job_to_dict,
    )

    target_skus = split_master_sku_values(skus)
    magento_rows = _dedupe_connections(magento_connections)
    shopify_rows = _dedupe_connections(shopify_connections)
    result: Dict[str, Any] = {
        "status": "ok",
        "dry_run": bool(dry_run),
        "target_sku_count": len(target_skus),
        "skus": target_skus,
        "prune_magento_orphans": bool(prune_magento_orphans),
        "push_magento_images": bool(push_magento_images),
        "push_shopify_images": bool(push_shopify_images),
        "would_queue": 0,
        "queued": 0,
        "jobs": [],
        "skipped_channels": [],
    }
    if not target_skus:
        result["message"] = "No SKUs available for outbound image cleanup."
        return result

    if not magento_rows and (prune_magento_orphans or push_magento_images):
        result["skipped_channels"].append("magento")
    if not shopify_rows and push_shopify_images:
        result["skipped_channels"].append("shopify")

    for connection in magento_rows:
        if prune_magento_orphans:
            entry = {
                "channel": "magento",
                "connection_id": connection["native_id"],
                "channel_connection_id": connection["id"],
                "job_type": "media_cleanup",
                "sku_count": len(target_skus),
            }
            if dry_run:
                result["would_queue"] += 1
                result["jobs"].append({**entry, "status": "would_queue"})
            else:
                job = enqueue_magento_media_cleanup_job(
                    session,
                    channel_connection_id=connection["id"],
                    native_connection_id=connection["native_id"],
                    channel_code=connection["channel_code"],
                    skus=target_skus,
                    all_assigned=False,
                    batch_size=max(1, int(batch_size or 250)),
                    purge_unmapped=True,
                    dry_run=False,
                    notes=f"{notes}:magento:prune",
                )
                result["queued"] += 1
                result["jobs"].append(
                    {
                        **entry,
                        "status": "queued",
                        **job_to_dict(job),
                    }
                )
        if push_magento_images:
            entry = {
                "channel": "magento",
                "connection_id": connection["native_id"],
                "channel_connection_id": connection["id"],
                "job_type": "push_images",
                "sku_count": len(target_skus),
            }
            if dry_run:
                result["would_queue"] += 1
                result["jobs"].append({**entry, "status": "would_queue"})
            else:
                job = enqueue_channel_job(
                    session,
                    channel_connection_id=connection["id"],
                    channel_type="magento",
                    native_connection_id=connection["native_id"],
                    channel_code=connection["channel_code"],
                    job_type="push_images",
                    dry_run=False,
                    mode="images_only",
                    notes=f"{notes}:magento:push_images",
                    options={"limit_skus": target_skus},
                )
                result["queued"] += 1
                result["jobs"].append(
                    {
                        **entry,
                        "status": "queued",
                        **job_to_dict(job),
                    }
                )

    for connection in shopify_rows:
        if not push_shopify_images:
            continue
        entry = {
            "channel": "shopify",
            "connection_id": connection["native_id"],
            "channel_connection_id": connection["id"],
            "job_type": "push_images",
            "sku_count": len(target_skus),
        }
        if dry_run:
            result["would_queue"] += 1
            result["jobs"].append({**entry, "status": "would_queue"})
            continue
        job = enqueue_channel_job(
            session,
            channel_connection_id=connection["id"],
            channel_type="shopify",
            native_connection_id=connection["native_id"],
            channel_code=connection["channel_code"],
            job_type="push_images",
            dry_run=False,
            mode="images_only",
            notes=f"{notes}:shopify:push_images",
            options={
                "skus": target_skus,
                "shop_code": connection.get("store_code"),
                "connection_id": connection["id"],
            },
        )
        result["queued"] += 1
        result["jobs"].append(
            {
                **entry,
                "status": "queued",
                **job_to_dict(job),
            }
        )

    return result
