from __future__ import annotations

from collections import defaultdict
from datetime import datetime, timezone
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, filter_preferred_product_image_rows
from db.models import (
    ManualTaxonomyAssignment,
    ManualTaxonomyAssignmentCollection,
    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 sync_shared_sku_images_from_manual_taxonomy(
    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,
    replace_existing: bool = True,
    source_label: str = "shared_sku_base_image_sync",
) -> 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,
            "processed": 0,
            "updated": 0,
            "skipped": 0,
            "errors": [],
            "sample": [],
            "message": "No shared SKUs matched the requested sync scope.",
        }

    products = session.scalars(
        select(MasterProduct)
        .where(MasterProduct.sku.in_(target_skus))
        .order_by(MasterProduct.sku.asc())
    ).all()
    product_by_sku = {str(row.sku or "").strip().upper(): row for row in products}

    source_items = sorted(
        {
            str(getattr(row, "source_item", "") or "").strip().upper()
            for row in products
            if str(getattr(row, "source_item", "") or "").strip()
        }
    )
    membership_slugs = sorted(
        {
            normalize_plp_path(slug)
            for slugs in scope["memberships_by_sku"].values()
            for slug in (slugs or [])
            if normalize_plp_path(slug)
        }
    )

    assignment_rows = session.execute(
        select(
            ManualTaxonomyAssignment.source_sku,
            ManualTaxonomyAssignmentCollection.collection_path_slug,
            ManualTaxonomyAssignmentCollection.master_sku,
        )
        .join(
            ManualTaxonomyAssignmentCollection,
            ManualTaxonomyAssignmentCollection.assignment_id == ManualTaxonomyAssignment.id,
        )
        .where(ManualTaxonomyAssignment.assignment_status == "active")
        .where(ManualTaxonomyAssignmentCollection.is_active.is_(True))
        .where(ManualTaxonomyAssignment.source_sku.in_(source_items))
        .where(ManualTaxonomyAssignmentCollection.collection_path_slug.in_(membership_slugs))
    ).all() if source_items and membership_slugs else []
    assignment_map: Dict[tuple[str, str], str] = {}
    for source_sku, path_slug, base_sku in assignment_rows:
        source_key = str(source_sku or "").strip().upper()
        path_key = normalize_plp_path(path_slug or "")
        base_key = str(base_sku or "").strip().upper()
        if source_key and path_key and base_key:
            assignment_map[(source_key, path_key)] = base_key

    collection_code_rows = session.execute(
        select(MasterCollectionRegistry.path_slug, MasterCollectionRegistry.code)
        .where(MasterCollectionRegistry.path_slug.in_(membership_slugs))
        .where(MasterCollectionRegistry.is_active.is_(True))
    ).all() if membership_slugs else []
    collection_code_by_path = {
        normalize_plp_path(path_slug or ""): str(code or "").strip().upper()
        for path_slug, code in collection_code_rows
        if normalize_plp_path(path_slug or "") and str(code or "").strip()
    }

    current_images = session.scalars(
        select(MasterProductImage)
        .where(MasterProductImage.sku.in_(target_skus))
        .order_by(MasterProductImage.sku.asc(), MasterProductImage.sort_order.asc(), MasterProductImage.id.asc())
    ).all()
    current_by_sku: Dict[str, List[MasterProductImage]] = defaultdict(list)
    for row in current_images:
        current_by_sku[str(row.sku or "").strip().upper()].append(row)

    now = datetime.now(timezone.utc)
    results: List[Dict[str, Any]] = []
    updated = 0
    skipped = 0
    errors: List[Dict[str, Any]] = []
    resolved_base_skus_for_scope: set[str] = set()
    base_source_rows_by_sku: Dict[str, List[MasterProductImage]] = {}

    def _load_base_source_rows(base_sku: str) -> List[MasterProductImage]:
        key = str(base_sku or "").strip().upper()
        if not key:
            return []
        cached = base_source_rows_by_sku.get(key)
        if cached is not None:
            return cached
        rows = session.scalars(
            select(MasterProductImage)
            .where(MasterProductImage.sku == key)
            .order_by(MasterProductImage.sort_order.asc(), MasterProductImage.id.asc())
        ).all()
        filtered_rows = filter_preferred_product_image_rows(rows)
        base_source_rows_by_sku[key] = filtered_rows
        return filtered_rows

    for sku in target_skus:
        product = product_by_sku.get(sku)
        if product is None:
            errors.append({"sku": sku, "reason": "shared_product_missing"})
            continue
        source_item = str(getattr(product, "source_item", "") or "").strip().upper()
        memberships = [normalize_plp_path(slug) for slug in (scope["memberships_by_sku"].get(sku) or []) if normalize_plp_path(slug)]
        candidate_map: Dict[str, List[str]] = defaultdict(list)
        for slug in memberships:
            base_sku = assignment_map.get((source_item, slug))
            if not base_sku:
                collection_code = collection_code_by_path.get(slug) or ""
                if collection_code:
                    base_sku = f"{collection_code}-{source_item}"
            if base_sku:
                candidate_map[base_sku].append(slug)
        resolved_base_skus = sorted(candidate_map.keys())
        existing_rows = current_by_sku.get(sku, [])
        result_row: Dict[str, Any] = {
            "sku": sku,
            "source_item": source_item,
            "memberships": memberships,
            "existing_image_count": len(existing_rows),
            "resolved_base_skus": resolved_base_skus,
        }

        if not source_item:
            result_row["status"] = "skipped"
            result_row["reason"] = "missing_source_item"
            skipped += 1
            results.append(result_row)
            continue
        if not resolved_base_skus:
            result_row["status"] = "skipped"
            result_row["reason"] = "no_manual_taxonomy_base_match"
            skipped += 1
            results.append(result_row)
            continue
        if len(resolved_base_skus) > 1:
            result_row["status"] = "skipped"
            result_row["reason"] = "ambiguous_base_sku"
            result_row["base_membership_map"] = dict(candidate_map)
            skipped += 1
            results.append(result_row)
            continue

        base_sku = resolved_base_skus[0]
        resolved_base_skus_for_scope.add(base_sku)
        source_rows = _load_base_source_rows(base_sku)
        result_row["base_sku"] = base_sku
        result_row["base_image_count"] = len(source_rows)
        if not source_rows:
            result_row["status"] = "skipped"
            result_row["reason"] = "base_sku_has_no_images"
            skipped += 1
            results.append(result_row)
            continue

        result_row["status"] = "would_update" if dry_run else "updated"
        result_row["replaced_image_count"] = len(source_rows)
        results.append(result_row)

        if dry_run:
            continue

        if replace_existing:
            session.execute(
                delete(MasterProductImage).where(MasterProductImage.sku == sku)
            )

        for source_row in source_rows:
            session.add(
                MasterProductImage(
                    product_id=product.id,
                    sku=sku,
                    file_name=str(source_row.file_name or ""),
                    image_url=str(source_row.image_url or ""),
                    image_role=source_row.image_role,
                    view_suffix=source_row.view_suffix,
                    sort_order=int(source_row.sort_order or 0),
                    source_label=f"{source_label}:{base_sku}",
                    updated_at=now,
                )
            )
        product.base_sku = base_sku
        updated += 1

    return {
        "status": "ok",
        "dry_run": bool(dry_run),
        "scope": scope,
        "processed": len(target_skus),
        "updated": updated,
        "skipped": skipped,
        "errors": errors,
        "resolved_base_sku_count": len(resolved_base_skus_for_scope),
        "sample": results[:100],
        "results": results,
    }


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
