"""Canonicalize saved manual taxonomy assignments against active master taxonomy nodes.

Default mode is dry-run. This repairs stale manual taxonomy slugs after taxonomy
renames/merges by resolving aliases to the current canonical path, refreshing the
stored path metadata, and merging duplicate rules created by the rewrite.
"""

from __future__ import annotations

import argparse
import json
from collections import defaultdict
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, Optional, Sequence

from sqlalchemy import select
from sqlalchemy.orm import Session

from channel.url_canonical import normalize_plp_path
from db.collection_landing_pages import _relative_shopping_l1_value, _relative_shopping_l2_value
from db.master_taxonomy_merge import resolve_canonical_node, resolve_canonical_path_slug
from db.models import ManualTaxonomyAssignment, ManualTaxonomyAssignmentCollection, MasterTaxonomyNode
from db.session import get_session


def cleanup_manual_taxonomy_assignments(
    session: Session,
    *,
    dry_run: bool = True,
    path_prefixes: Optional[Iterable[str]] = None,
    note: Optional[str] = None,
) -> Dict[str, Any]:
    cleanup_note = note or f"manual taxonomy cleanup at {datetime.now(timezone.utc).isoformat()}"
    prefixes = [_norm_slug(item) for item in (path_prefixes or []) if _norm_slug(item)]

    assignments = list(
        session.scalars(
            select(ManualTaxonomyAssignment)
            .where(ManualTaxonomyAssignment.assignment_status == "active")
            .order_by(ManualTaxonomyAssignment.source_sku, ManualTaxonomyAssignment.id)
        ).all()
    )
    result: Dict[str, Any] = {
        "status": "ok",
        "dry_run": dry_run,
        "path_prefixes": prefixes,
        "assignment_count": len(assignments),
        "candidate_count": 0,
        "would_update_assignments": 0,
        "updated_assignments": 0,
        "would_update_payloads": 0,
        "updated_payloads": 0,
        "would_merge_duplicate_rules": 0,
        "merged_duplicate_rules": 0,
        "would_delete_duplicate_children": 0,
        "deleted_duplicate_children": 0,
        "unresolved_count": 0,
        "skipped_out_of_scope": 0,
        "actions": [],
        "unresolved": [],
    }
    if not assignments:
        return result

    assignment_ids = [int(row.id) for row in assignments]
    children = list(
        session.scalars(
            select(ManualTaxonomyAssignmentCollection)
            .where(ManualTaxonomyAssignmentCollection.assignment_id.in_(assignment_ids))
            .where(ManualTaxonomyAssignmentCollection.is_active.is_(True))
            .order_by(
                ManualTaxonomyAssignmentCollection.assignment_id,
                ManualTaxonomyAssignmentCollection.sort_order,
                ManualTaxonomyAssignmentCollection.id,
            )
        ).all()
    )
    children_by_assignment: dict[int, list[ManualTaxonomyAssignmentCollection]] = defaultdict(list)
    for child in children:
        children_by_assignment[int(child.assignment_id)].append(child)

    planned_groups: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list)
    for assignment in assignments:
        source_sku = str(assignment.source_sku or "").strip().upper()
        raw_slug = _norm_slug(assignment.canonical_taxonomy_path_slug)
        if not raw_slug:
            result["unresolved_count"] += 1
            result["unresolved"].append({"assignment_id": int(assignment.id), "source_sku": source_sku, "reason": "blank_slug"})
            continue
        if prefixes and not any(raw_slug == prefix or raw_slug.startswith(f"{prefix}/") for prefix in prefixes):
            result["skipped_out_of_scope"] += 1
            continue
        canonical_slug = _norm_slug(resolve_canonical_path_slug(session, raw_slug) or raw_slug)
        node = resolve_canonical_node(session, path_slug=canonical_slug)
        if node is None or not node.is_active:
            result["unresolved_count"] += 1
            result["unresolved"].append(
                {
                    "assignment_id": int(assignment.id),
                    "source_sku": source_sku,
                    "path_slug": raw_slug,
                    "resolved_path_slug": canonical_slug,
                    "reason": "missing_active_taxonomy_node",
                }
            )
            continue
        metadata = _assignment_metadata(session, node)
        planned_groups[(source_sku, metadata["canonical_taxonomy_path_slug"])].append(
            {"assignment": assignment, "metadata": metadata}
        )

    result["candidate_count"] = sum(len(items) for items in planned_groups.values())

    for (_source_sku, _target_slug), group in sorted(planned_groups.items()):
        keep_entry = _choose_keeper(group)
        keep_assignment: ManualTaxonomyAssignment = keep_entry["assignment"]
        metadata: Dict[str, Optional[str]] = keep_entry["metadata"]
        duplicate_entries = [item for item in group if int(item["assignment"].id) != int(keep_assignment.id)]

        action = {
            "keep_assignment_id": int(keep_assignment.id),
            "source_sku": str(keep_assignment.source_sku or "").strip().upper(),
            "target_path_slug": metadata["canonical_taxonomy_path_slug"],
            "merged_assignment_ids": [int(item["assignment"].id) for item in duplicate_entries],
            "old_path_slugs": sorted(
                {
                    _norm_slug(item["assignment"].canonical_taxonomy_path_slug)
                    for item in group
                    if _norm_slug(item["assignment"].canonical_taxonomy_path_slug)
                }
            ),
        }

        assignment_changed = _assignment_needs_update(keep_assignment, metadata)
        merged_children, child_stats = _merge_assignment_children(
            keep_assignment,
            [item["assignment"] for item in duplicate_entries],
            children_by_assignment,
            note=cleanup_note,
            dry_run=dry_run,
        )
        payload = _rebuild_assignment_payload(keep_assignment, metadata, merged_children)
        current_payload = keep_assignment.raw_payload if isinstance(keep_assignment.raw_payload, dict) else {}
        payload_changed = payload != current_payload

        if assignment_changed:
            if dry_run:
                result["would_update_assignments"] += 1
            else:
                _apply_assignment_metadata(keep_assignment, metadata, cleanup_note)
                result["updated_assignments"] += 1

        if payload_changed:
            if dry_run:
                result["would_update_payloads"] += 1
            else:
                keep_assignment.raw_payload = payload
                result["updated_payloads"] += 1

        if not dry_run and (duplicate_entries or int(child_stats.get("deleted_duplicate_children") or 0) > 0):
            session.flush()

        if duplicate_entries:
            if dry_run:
                result["would_merge_duplicate_rules"] += len(duplicate_entries)
            else:
                for item in duplicate_entries:
                    session.delete(item["assignment"])
                    result["merged_duplicate_rules"] += 1

        result["would_delete_duplicate_children"] += int(child_stats.get("would_delete_duplicate_children") or 0)
        result["deleted_duplicate_children"] += int(child_stats.get("deleted_duplicate_children") or 0)
        action.update(
            {
                "assignment_changed": assignment_changed,
                "payload_changed": payload_changed,
                "duplicate_child_rows": int(child_stats.get("duplicate_child_rows") or 0),
            }
        )
        result["actions"].append(action)

    if not dry_run:
        session.flush()
    return result


def _assignment_metadata(session: Session, node: MasterTaxonomyNode) -> Dict[str, Optional[str]]:
    labels = _taxonomy_lineage_labels(session, node)
    path_slug = _norm_slug(node.path_slug)
    return {
        "canonical_taxonomy_path_slug": path_slug,
        "canonical_taxonomy_path_label": " / ".join(labels) if labels else None,
        "shopping_l1": _relative_shopping_l1_value("/".join(path_slug.split("/")[:2]) if len(path_slug.split("/")) >= 2 else ""),
        "shopping_l2": _relative_shopping_l2_value(path_slug),
        "l2_category": labels[1] if len(labels) >= 2 else None,
        "l3_category": labels[2] if len(labels) >= 3 else None,
    }


def _taxonomy_lineage_labels(session: Session, node: MasterTaxonomyNode) -> list[str]:
    labels: list[str] = []
    current: Optional[MasterTaxonomyNode] = node
    seen: set[int] = set()
    while current is not None and int(current.id) not in seen:
        seen.add(int(current.id))
        labels.append(str(current.name or "").strip() or _slug_label(current.path_slug))
        current = session.get(MasterTaxonomyNode, int(current.parent_id)) if current.parent_id else None
    labels.reverse()
    return [label for label in labels if label]


def _slug_label(path_slug: Optional[str]) -> str:
    slug = _norm_slug(path_slug)
    if not slug:
        return ""
    return slug.split("/")[-1].replace("-", " ").title()


def _choose_keeper(group: Sequence[dict[str, Any]]) -> dict[str, Any]:
    def score(item: dict[str, Any]) -> tuple[int, int]:
        assignment: ManualTaxonomyAssignment = item["assignment"]
        current_slug = _norm_slug(assignment.canonical_taxonomy_path_slug)
        target_slug = _norm_slug(item["metadata"]["canonical_taxonomy_path_slug"])
        return (1 if current_slug == target_slug else 0, -int(assignment.id))

    return max(group, key=score)


def _assignment_needs_update(
    assignment: ManualTaxonomyAssignment,
    metadata: Dict[str, Optional[str]],
) -> bool:
    return any(
        [
            _norm_slug(assignment.canonical_taxonomy_path_slug) != _norm_slug(metadata["canonical_taxonomy_path_slug"]),
            _clean(assignment.canonical_taxonomy_path_label) != _clean(metadata["canonical_taxonomy_path_label"]),
            _clean(assignment.shopping_l1) != _clean(metadata["shopping_l1"]),
            _clean(assignment.shopping_l2) != _clean(metadata["shopping_l2"]),
            _clean(assignment.l2_category) != _clean(metadata["l2_category"]),
            _clean(assignment.l3_category) != _clean(metadata["l3_category"]),
        ]
    )


def _apply_assignment_metadata(
    assignment: ManualTaxonomyAssignment,
    metadata: Dict[str, Optional[str]],
    note: str,
) -> None:
    assignment.canonical_taxonomy_path_slug = str(metadata["canonical_taxonomy_path_slug"] or "").strip()
    assignment.canonical_taxonomy_path_label = _none_if_blank(metadata["canonical_taxonomy_path_label"])
    assignment.shopping_l1 = _none_if_blank(metadata["shopping_l1"])
    assignment.shopping_l2 = _none_if_blank(metadata["shopping_l2"])
    assignment.l2_category = _none_if_blank(metadata["l2_category"])
    assignment.l3_category = _none_if_blank(metadata["l3_category"])
    assignment.notes = _append_note(assignment.notes, note)


def _merge_assignment_children(
    keep_assignment: ManualTaxonomyAssignment,
    duplicate_assignments: Sequence[ManualTaxonomyAssignment],
    children_by_assignment: dict[int, list[ManualTaxonomyAssignmentCollection]],
    *,
    note: str,
    dry_run: bool,
) -> tuple[list[ManualTaxonomyAssignmentCollection], Dict[str, int]]:
    kept_children = list(children_by_assignment.get(int(keep_assignment.id), []))
    survivors: dict[str, ManualTaxonomyAssignmentCollection] = {
        _child_key(child): child for child in kept_children if _child_key(child)
    }
    merged_children = list(kept_children)
    stats = {"duplicate_child_rows": 0, "would_delete_duplicate_children": 0, "deleted_duplicate_children": 0}

    for assignment in duplicate_assignments:
        for child in children_by_assignment.get(int(assignment.id), []):
            key = _child_key(child)
            existing = survivors.get(key) if key else None
            if existing is not None:
                stats["duplicate_child_rows"] += 1
                if dry_run:
                    stats["would_delete_duplicate_children"] += 1
                else:
                    session = Session.object_session(existing)
                    _merge_notes(existing, child, note)
                    if session is not None:
                        session.delete(child)
                    stats["deleted_duplicate_children"] += 1
                continue
            survivors[key] = child
            merged_children.append(child)
            if not dry_run:
                child.assignment_id = int(keep_assignment.id)
                child.notes = _append_note(child.notes, note)
    return merged_children, stats


def _child_key(child: ManualTaxonomyAssignmentCollection) -> str:
    return str(child.collection_code or "").strip().upper()


def _rebuild_assignment_payload(
    assignment: ManualTaxonomyAssignment,
    metadata: Dict[str, Optional[str]],
    children: Sequence[ManualTaxonomyAssignmentCollection],
) -> dict[str, Any]:
    payload = dict(assignment.raw_payload) if isinstance(assignment.raw_payload, dict) else {}
    payload["source_sku"] = str(assignment.source_sku or "").strip().upper()
    payload["canonical_taxonomy_path_slug"] = str(metadata["canonical_taxonomy_path_slug"] or "").strip()
    payload["canonical_taxonomy_path_label"] = _none_if_blank(metadata["canonical_taxonomy_path_label"])
    payload["shopping_l1"] = _none_if_blank(metadata["shopping_l1"])
    payload["shopping_l2"] = _none_if_blank(metadata["shopping_l2"])
    payload["l2_category"] = _none_if_blank(metadata["l2_category"])
    payload["l3_category"] = _none_if_blank(metadata["l3_category"])
    payload["collections"] = [
        {
            "master_sku": str(child.master_sku or "").strip().upper(),
            "sku_prefix": str(child.sku_prefix or "").strip().upper(),
            "alias_codes": child.alias_codes if isinstance(child.alias_codes, list) else [],
            "collection_code": str(child.collection_code or "").strip().upper(),
            "source_collection_code": str(child.source_collection_code or "").strip().upper() or None,
            "collection_name": str(child.collection_name or "").strip(),
            "collection_path_slug": str(child.collection_path_slug or "").strip() or None,
            "shopping_collection": str(child.shopping_collection or "").strip() or None,
        }
        for child in sorted(children, key=lambda item: (int(item.sort_order or 0), int(item.id or 0)))
        if child.is_active
    ]
    return payload


def _merge_notes(target_row: Any, source_row: Any, note: str) -> None:
    extra = f"merged duplicate child row id={getattr(source_row, 'id', '?')}"
    target_row.notes = _append_note(getattr(target_row, "notes", None), f"{note}; {extra}")


def _append_note(existing: Optional[str], addition: str) -> str:
    base = _clean(existing)
    extra = _clean(addition)
    if not base:
        return extra
    if not extra:
        return base
    return f"{base}\n{extra}"


def _none_if_blank(value: Optional[str]) -> Optional[str]:
    text = _clean(value)
    return text or None


def _clean(value: Any) -> str:
    return str(value or "").strip()


def _norm_slug(value: Any) -> str:
    return normalize_plp_path(str(value or ""))


def main() -> int:
    parser = argparse.ArgumentParser(description="Canonicalize saved manual taxonomy assignments from the database.")
    parser.add_argument("--apply", action="store_true", help="Apply changes. Default is dry-run.")
    parser.add_argument(
        "--path-prefix",
        action="append",
        default=None,
        help="Limit cleanup to canonical_taxonomy_path_slug values under this prefix. Repeatable.",
    )
    args = parser.parse_args()

    with get_session() as session:
        try:
            summary = cleanup_manual_taxonomy_assignments(
                session,
                dry_run=not args.apply,
                path_prefixes=args.path_prefix,
            )
            if args.apply:
                session.commit()
            else:
                session.rollback()
        except Exception:
            session.rollback()
            raise
    print(json.dumps(summary, indent=2, default=str))
    return 0 if summary.get("status") == "ok" else 1


if __name__ == "__main__":
    raise SystemExit(main())
