"""Repair stored channel_sku_mapping rows to match current prefix rules (incl. RTA mirrors)."""

from __future__ import annotations

from typing import Any, Dict, Iterable, List, Optional, Set

from sqlalchemy import select
from sqlalchemy.orm import Session

from db.channel_assignments import active_skus_for_channel
from db.channel_exports import resolve_pipeline_channel_code
from db.channel_sku_mapping import ACTIVE_STATUS, reconcile_connection_id_for_channel
from db.channel_sku_prefix_mapping import RTA_ASSEMBLY_PREFIX, resolve_channel_sku_from_prefixes
from db.models import ChannelSkuMapping, MasterProduct


def plan_channel_sku_mapping_repairs(
    session: Session,
    channel_code: str,
    *,
    connection_id: Optional[int] = None,
    only_assigned: bool = False,
    skus: Optional[Iterable[str]] = None,
    rta_only: bool = False,
) -> Dict[str, Any]:
    """Find mapping rows whose channel_sku differs from prefix-rule resolution."""
    channel = resolve_pipeline_channel_code(channel_code)
    wanted = _master_skus(session, channel, only_assigned=only_assigned, skus=skus)
    mappings = _active_mappings(session, channel, wanted, connection_id=connection_id)

    repairs: List[Dict[str, Any]] = []
    conflicts: List[Dict[str, Any]] = []
    skipped = 0

    for master_sku in sorted(mappings):
        row = mappings[master_sku]
        if rta_only and not str(master_sku).upper().startswith(RTA_ASSEMBLY_PREFIX):
            continue
        current = str(row.channel_sku).strip()
        target = resolve_channel_sku_from_prefixes(
            session,
            master_sku,
            channel,
            connection_id=connection_id,
        )
        if target == current:
            skipped += 1
            continue

        occupant = _mapping_for_channel_sku(session, channel, connection_id, target)
        if occupant is not None and occupant.master_sku != master_sku:
            conflicts.append(
                {
                    "master_sku": master_sku,
                    "current_channel_sku": current,
                    "target_channel_sku": target,
                    "conflict_master_sku": occupant.master_sku,
                    "conflict_mapping_id": occupant.id,
                    "reason": "target_channel_sku_already_mapped",
                }
            )
            continue

        repairs.append(
            {
                "mapping_id": row.id,
                "master_sku": master_sku,
                "current_channel_sku": current,
                "target_channel_sku": target,
                "remote_id": row.remote_id,
                "match_source": row.match_source,
                "rta_mirror": str(master_sku).upper().startswith(RTA_ASSEMBLY_PREFIX),
            }
        )

    return {
        "channel_code": channel,
        "connection_id": connection_id,
        "only_assigned": only_assigned,
        "rta_only": rta_only,
        "mapping_count": len(mappings),
        "repair_count": len(repairs),
        "conflict_count": len(conflicts),
        "already_correct_count": skipped,
        "repairs": repairs,
        "conflicts": conflicts[:500],
    }


def apply_channel_sku_mapping_repairs(
    session: Session,
    repairs: List[Dict[str, Any]],
    *,
    dry_run: bool = True,
) -> Dict[str, Any]:
    """Update channel_sku on mapping rows planned by plan_channel_sku_mapping_repairs."""
    applied: List[Dict[str, Any]] = []
    if dry_run:
        return {"dry_run": True, "applied_count": 0, "preview": repairs[:500]}

    for item in repairs:
        row = session.get(ChannelSkuMapping, item["mapping_id"])
        if row is None:
            continue
        row.channel_sku = item["target_channel_sku"]
        row.match_source = "prefix_rule_repair"
        prior_notes = str(row.notes or "").strip()
        note = f"channel_sku {item['current_channel_sku']} → {item['target_channel_sku']} (prefix rules)"
        row.notes = f"{prior_notes}; {note}" if prior_notes else note
        applied.append(
            {
                "mapping_id": row.id,
                "master_sku": row.master_sku,
                "current_channel_sku": item["current_channel_sku"],
                "target_channel_sku": item["target_channel_sku"],
            }
        )
    session.flush()
    return {"dry_run": False, "applied_count": len(applied), "applied": applied[:500]}


def repair_channel_sku_mappings(
    session: Session,
    *,
    channels: Optional[Iterable[str]] = None,
    connection_id: Optional[int] = None,
    only_assigned: bool = False,
    skus: Optional[Iterable[str]] = None,
    rta_only: bool = False,
    dry_run: bool = True,
    rename_remote: bool = False,
    rename_limit: Optional[int] = None,
    magento_connection_id: Optional[int] = None,
    shopify_connection_id: Optional[int] = None,
) -> Dict[str, Any]:
    """Sync mapping rows to prefix targets; optionally rename live remote SKUs first."""
    selected = [resolve_pipeline_channel_code(ch) for ch in (channels or ("shopify",))]
    summary: Dict[str, Any] = {
        "status": "ok",
        "dry_run": dry_run,
        "rename_remote": rename_remote,
        "channels": {},
    }

    remote_rename_result = None
    if rename_remote and not dry_run:
        from db.channel_sku_rename import apply_channel_sku_renames

        remote_rename_result = apply_channel_sku_renames(
            session,
            channels=selected,
            only_assigned=only_assigned,
            linked_only=True,
            dry_run=False,
            skus=list(skus) if skus else None,
            limit=rename_limit,
            magento_connection_id=magento_connection_id,
            shopify_connection_id=shopify_connection_id,
        )
        summary["remote_rename"] = remote_rename_result

    for channel in selected:
        scoped_connection_id = connection_id or reconcile_connection_id_for_channel(
            channel,
            compat_connection_id=shopify_connection_id if channel == "shopify" else None,
            native_connection_id=magento_connection_id if channel == "magento" else None,
        )
        plan = plan_channel_sku_mapping_repairs(
            session,
            channel,
            connection_id=scoped_connection_id,
            only_assigned=only_assigned,
            skus=skus,
            rta_only=rta_only,
        )
        apply_result = apply_channel_sku_mapping_repairs(
            session,
            plan["repairs"],
            dry_run=dry_run,
        )
        summary["channels"][channel] = {
            "connection_id": scoped_connection_id,
            **plan,
            "apply": apply_result,
        }
        if plan["conflict_count"]:
            summary["status"] = "partial"

    if rename_remote and dry_run:
        from db.channel_sku_rename import apply_channel_sku_renames

        summary["remote_rename_preview"] = apply_channel_sku_renames(
            session,
            channels=selected,
            only_assigned=only_assigned,
            linked_only=True,
            dry_run=True,
            skus=list(skus) if skus else None,
            limit=rename_limit,
            magento_connection_id=magento_connection_id,
            shopify_connection_id=shopify_connection_id,
        )

    return summary


def _master_skus(
    session: Session,
    channel_code: str,
    *,
    only_assigned: bool,
    skus: Optional[Iterable[str]] = None,
) -> Set[str]:
    if skus:
        return {str(s).strip() for s in skus if str(s).strip()}
    stmt = select(MasterProduct.sku).where(MasterProduct.is_active.is_(True))
    if only_assigned:
        assigned = active_skus_for_channel(session, channel_code)
        if not assigned:
            return set()
        stmt = stmt.where(MasterProduct.sku.in_(sorted(assigned)))
    return {str(sku).strip() for (sku,) in session.execute(stmt).all() if str(sku).strip()}


def _active_mappings(
    session: Session,
    channel_code: str,
    master_skus: Set[str],
    *,
    connection_id: Optional[int],
) -> Dict[str, ChannelSkuMapping]:
    if not master_skus:
        return {}
    stmt = (
        select(ChannelSkuMapping)
        .where(ChannelSkuMapping.channel_code == channel_code)
        .where(ChannelSkuMapping.mapping_status == ACTIVE_STATUS)
        .where(ChannelSkuMapping.master_sku.in_(sorted(master_skus)))
    )
    if connection_id is not None:
        stmt = stmt.where(ChannelSkuMapping.connection_id == connection_id)
    else:
        stmt = stmt.where(ChannelSkuMapping.connection_id.is_(None))
    return {row.master_sku: row for row in session.scalars(stmt).all()}


def _mapping_for_channel_sku(
    session: Session,
    channel_code: str,
    connection_id: Optional[int],
    channel_sku: str,
) -> Optional[ChannelSkuMapping]:
    stmt = (
        select(ChannelSkuMapping)
        .where(ChannelSkuMapping.channel_code == channel_code)
        .where(ChannelSkuMapping.channel_sku == channel_sku)
        .where(ChannelSkuMapping.mapping_status == ACTIVE_STATUS)
    )
    if connection_id is not None:
        stmt = stmt.where(ChannelSkuMapping.connection_id == connection_id)
    else:
        stmt = stmt.where(ChannelSkuMapping.connection_id.is_(None))
    return session.scalars(stmt).first()
