from __future__ import annotations

from dataclasses import dataclass
from datetime import datetime, timezone
from difflib import SequenceMatcher
import re
from typing import Iterable

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

from db.models_v2 import (
    V2Category,
    V2DuplicateReviewQueue,
    V2Product,
    V2ProductCategory,
    V2ProductCollection,
    V2ShoppingAttributeOptionMap,
    V2ShoppingCascadeJob,
)


NEAR_DUPLICATE_THRESHOLD = 85.0


def _utcnow() -> datetime:
    return datetime.now(timezone.utc)


def normalize_taxonomy_name(value: str) -> str:
    text = str(value or "").strip().lower()
    text = re.sub(r"\s+", " ", text)
    return text


def similarity_score(left: str, right: str) -> float:
    left_norm = normalize_taxonomy_name(left)
    right_norm = normalize_taxonomy_name(right)
    if not left_norm or not right_norm:
        return 0.0
    return round(SequenceMatcher(None, left_norm, right_norm).ratio() * 100, 2)


@dataclass
class TaxonomyDuplicateCandidate:
    category_id: int
    name: str
    normalized_name: str
    similarity_score: float
    parent_id: int | None
    node_type: str
    level: int | None


@dataclass
class TaxonomyCreateOutcome:
    status: str
    category: V2Category | None
    candidates: list[TaxonomyDuplicateCandidate]


@dataclass
class TaxonomyCascadeResult:
    cascade_job: V2ShoppingCascadeJob
    source_category_id: int
    target_category_id: int | None
    operation: str
    affected_product_category_links: int = 0
    affected_product_collection_links: int = 0
    affected_product_rows: int = 0
    affected_shopping_map_rows: int = 0
    reparented_child_categories: int = 0


def find_sibling_categories(
    session: Session,
    *,
    parent_id: int | None,
    node_type: str,
    exclude_category_id: int | None = None,
) -> list[V2Category]:
    stmt = (
        select(V2Category)
        .where(
            V2Category.merged_into_id.is_(None),
            V2Category.node_type == node_type,
        )
        .order_by(V2Category.name.asc())
    )
    if exclude_category_id is not None:
        stmt = stmt.where(V2Category.id != exclude_category_id)
    if parent_id is None:
        stmt = stmt.where(V2Category.parent_id.is_(None))
    else:
        stmt = stmt.where(V2Category.parent_id == parent_id)
    return list(session.execute(stmt).scalars().all())


def evaluate_taxonomy_creation(
    session: Session,
    *,
    name: str,
    node_type: str,
    parent_id: int | None,
    level: int | None,
    confirm_new: bool,
    created_by: str,
    threshold: float = NEAR_DUPLICATE_THRESHOLD,
) -> TaxonomyCreateOutcome:
    normalized = normalize_taxonomy_name(name)
    siblings = find_sibling_categories(session, parent_id=parent_id, node_type=node_type)

    exact_matches: list[TaxonomyDuplicateCandidate] = []
    near_matches: list[TaxonomyDuplicateCandidate] = []
    for sibling in siblings:
        sibling_norm = sibling.normalized_name or normalize_taxonomy_name(sibling.name)
        score = 100.0 if sibling_norm == normalized else similarity_score(name, sibling.name)
        candidate = TaxonomyDuplicateCandidate(
            category_id=int(sibling.id),
            name=sibling.name,
            normalized_name=sibling_norm,
            similarity_score=score,
            parent_id=sibling.parent_id,
            node_type=sibling.node_type,
            level=sibling.level,
        )
        if sibling_norm == normalized:
            exact_matches.append(candidate)
        elif score >= threshold:
            near_matches.append(candidate)

    if exact_matches:
        return TaxonomyCreateOutcome(
            status="existing_duplicate",
            category=None,
            candidates=sorted(exact_matches, key=lambda item: (-item.similarity_score, item.name)),
        )

    if near_matches and not confirm_new:
        return TaxonomyCreateOutcome(
            status="duplicate_requires_review",
            category=None,
            candidates=sorted(near_matches, key=lambda item: (-item.similarity_score, item.name)),
        )

    category = V2Category(
        name=name.strip(),
        parent_id=parent_id,
        level=level,
        node_type=node_type,
        created_by=created_by,
    )
    session.add(category)
    session.flush()
    session.refresh(category)
    return TaxonomyCreateOutcome(status="created", category=category, candidates=[])


def evaluate_taxonomy_update(
    session: Session,
    *,
    category_id: int,
    name: str,
    node_type: str,
    parent_id: int | None,
    level: int | None,
    threshold: float = NEAR_DUPLICATE_THRESHOLD,
) -> TaxonomyCreateOutcome:
    normalized = normalize_taxonomy_name(name)
    siblings = find_sibling_categories(
        session,
        parent_id=parent_id,
        node_type=node_type,
        exclude_category_id=category_id,
    )

    exact_matches: list[TaxonomyDuplicateCandidate] = []
    near_matches: list[TaxonomyDuplicateCandidate] = []
    for sibling in siblings:
        sibling_norm = sibling.normalized_name or normalize_taxonomy_name(sibling.name)
        score = 100.0 if sibling_norm == normalized else similarity_score(name, sibling.name)
        candidate = TaxonomyDuplicateCandidate(
            category_id=int(sibling.id),
            name=sibling.name,
            normalized_name=sibling_norm,
            similarity_score=score,
            parent_id=sibling.parent_id,
            node_type=sibling.node_type,
            level=sibling.level,
        )
        if sibling_norm == normalized:
            exact_matches.append(candidate)
        elif score >= threshold:
            near_matches.append(candidate)

    if exact_matches:
        return TaxonomyCreateOutcome(
            status="existing_duplicate",
            category=None,
            candidates=sorted(exact_matches, key=lambda item: (-item.similarity_score, item.name)),
        )
    if near_matches:
        return TaxonomyCreateOutcome(
            status="duplicate_requires_review",
            category=None,
            candidates=sorted(near_matches, key=lambda item: (-item.similarity_score, item.name)),
        )

    category = session.get(V2Category, category_id)
    if category is None:
        raise ValueError("Category not found")
    category.name = name.strip()
    category.parent_id = parent_id
    category.level = level
    session.flush()
    session.refresh(category)
    return TaxonomyCreateOutcome(status="updated", category=category, candidates=[])


def _iter_duplicate_candidates(
    categories: Iterable[V2Category],
    threshold: float,
) -> list[tuple[V2Category, V2Category, float]]:
    items = list(categories)
    results: list[tuple[V2Category, V2Category, float]] = []
    for index, left in enumerate(items):
        for right in items[index + 1 :]:
            score = similarity_score(left.name, right.name)
            if score >= threshold:
                results.append((left, right, score))
    return results


def scan_duplicate_review_queue(
    session: Session,
    *,
    threshold: float = NEAR_DUPLICATE_THRESHOLD,
) -> int:
    siblings = (
        session.execute(
            select(V2Category)
            .where(V2Category.merged_into_id.is_(None))
            .order_by(V2Category.dedupe_scope.asc(), V2Category.node_type.asc(), V2Category.name.asc())
        )
        .scalars()
        .all()
    )

    grouped: dict[tuple[int, str], list[V2Category]] = {}
    for category in siblings:
        key = (int(category.dedupe_scope), category.node_type)
        grouped.setdefault(key, []).append(category)

    existing_pairs = {
        (min(int(row.category_id_a), int(row.category_id_b)), max(int(row.category_id_a), int(row.category_id_b)))
        for row in session.execute(select(V2DuplicateReviewQueue)).scalars().all()
    }

    created = 0
    for categories in grouped.values():
        if len(categories) < 2:
            continue
        for left, right, score in _iter_duplicate_candidates(categories, threshold):
            pair = (min(int(left.id), int(right.id)), max(int(left.id), int(right.id)))
            if pair in existing_pairs:
                continue
            session.add(
                V2DuplicateReviewQueue(
                    category_id_a=pair[0],
                    category_id_b=pair[1],
                    similarity_score=score,
                )
            )
            existing_pairs.add(pair)
            created += 1

    session.flush()
    return created


def mark_duplicate_review_not_duplicate(
    session: Session,
    *,
    review_id: int,
    reviewed_by: str,
) -> V2DuplicateReviewQueue | None:
    review = session.get(V2DuplicateReviewQueue, review_id)
    if review is None:
        return None
    review.status = "not_duplicate"
    review.reviewed_by = reviewed_by
    return review


def get_cascade_job(session: Session, job_id: str) -> V2ShoppingCascadeJob | None:
    return session.get(V2ShoppingCascadeJob, job_id)


def list_cascade_jobs(session: Session, *, limit: int = 100) -> list[V2ShoppingCascadeJob]:
    return list(
        session.execute(
            select(V2ShoppingCascadeJob)
            .order_by(V2ShoppingCascadeJob.created_at.desc())
            .limit(limit)
        ).scalars().all()
    )


def can_delete_taxonomy_node(session: Session, *, category_id: int) -> tuple[bool, str | None]:
    category = session.get(V2Category, category_id)
    if category is None:
        return False, "Category not found"
    child_count = len(
        session.execute(select(V2Category.id).where(V2Category.parent_id == category_id)).all()
    )
    if child_count:
        return False, "Category still has child categories"
    merge_target_count = len(
        session.execute(select(V2Category.id).where(V2Category.merged_into_id == category_id)).all()
    )
    if merge_target_count:
        return False, "Category is the target of a completed merge and cannot be deleted"
    cascade_job_count = len(
        session.execute(
            select(V2ShoppingCascadeJob.id).where(
                (V2ShoppingCascadeJob.source_category_id == category_id)
                | (V2ShoppingCascadeJob.target_category_id == category_id)
            )
        ).all()
    )
    if cascade_job_count:
        return False, "Category is referenced by a cascade job history record and cannot be deleted"
    product_category_count = len(
        session.execute(select(V2ProductCategory.product_id).where(V2ProductCategory.category_id == category_id)).all()
    )
    if product_category_count:
        return False, "Category is still assigned through product categories"
    product_collection_count = len(
        session.execute(select(V2ProductCollection.product_id).where(V2ProductCollection.collection_id == category_id)).all()
    )
    if product_collection_count:
        return False, "Category is still assigned through product collections"
    shopping_map_count = len(
        session.execute(select(V2ShoppingAttributeOptionMap.id).where(V2ShoppingAttributeOptionMap.category_id == category_id)).all()
    )
    if shopping_map_count:
        return False, "Category is still referenced by Start Shopping mappings"
    product_fk_count = len(
        session.execute(
            select(V2Product.id).where(
                (V2Product.primary_category_id == category_id)
                | (V2Product.shopping_l1_category_id == category_id)
                | (V2Product.shopping_l2_category_id == category_id)
            )
        ).all()
    )
    if product_fk_count:
        return False, "Category is still referenced directly by products"
    return True, None


def delete_taxonomy_node(session: Session, *, category_id: int) -> bool:
    category = session.get(V2Category, category_id)
    if category is None:
        return False
    ok, reason = can_delete_taxonomy_node(session, category_id=category_id)
    if not ok:
        raise ValueError(reason or "Category cannot be deleted")
    session.delete(category)
    session.flush()
    return True


def execute_taxonomy_merge(
    session: Session,
    *,
    source_category_id: int,
    target_category_id: int,
    reviewed_by: str,
    delete_downstream_option: bool = False,
) -> TaxonomyCascadeResult:
    if source_category_id == target_category_id:
        raise ValueError("source_category_id and target_category_id must differ")
    source = session.get(V2Category, source_category_id)
    target = session.get(V2Category, target_category_id)
    if source is None or target is None:
        raise ValueError("Source or target category not found")
    if source.merged_into_id is not None:
        raise ValueError("Source category is already merged")

    cascade_job = V2ShoppingCascadeJob(
        operation="merge",
        source_category_id=source_category_id,
        target_category_id=target_category_id,
        delete_downstream_option=delete_downstream_option,
        status="running",
        created_by=reviewed_by,
    )
    session.add(cascade_job)
    session.flush()

    reparented_children = 0
    for child in session.execute(select(V2Category).where(V2Category.parent_id == source_category_id)).scalars().all():
        child.parent_id = target_category_id
        reparented_children += 1

    affected_product_rows = 0
    affected_product_rows += session.execute(
        update(V2Product)
        .where(V2Product.primary_category_id == source_category_id)
        .values(primary_category_id=target_category_id, updated_at=_utcnow())
    ).rowcount or 0
    affected_product_rows += session.execute(
        update(V2Product)
        .where(V2Product.shopping_l1_category_id == source_category_id)
        .values(shopping_l1_category_id=target_category_id, updated_at=_utcnow())
    ).rowcount or 0
    affected_product_rows += session.execute(
        update(V2Product)
        .where(V2Product.shopping_l2_category_id == source_category_id)
        .values(shopping_l2_category_id=target_category_id, updated_at=_utcnow())
    ).rowcount or 0

    product_category_rows = session.execute(
        select(V2ProductCategory.product_id).where(V2ProductCategory.category_id == source_category_id)
    ).all()
    affected_product_category_links = len(product_category_rows)
    for (product_id,) in product_category_rows:
        existing = session.execute(
            select(V2ProductCategory).where(
                V2ProductCategory.product_id == product_id,
                V2ProductCategory.category_id == target_category_id,
            )
        ).scalar_one_or_none()
        source_row = session.execute(
            select(V2ProductCategory).where(
                V2ProductCategory.product_id == product_id,
                V2ProductCategory.category_id == source_category_id,
            )
        ).scalar_one_or_none()
        if source_row is None:
            continue
        if existing is not None:
            session.delete(source_row)
        else:
            source_row.category_id = target_category_id

    product_collection_rows = session.execute(
        select(V2ProductCollection.product_id, V2ProductCollection.is_primary).where(
            V2ProductCollection.collection_id == source_category_id
        )
    ).all()
    affected_product_collection_links = len(product_collection_rows)
    for product_id, is_primary in product_collection_rows:
        existing = session.execute(
            select(V2ProductCollection).where(
                V2ProductCollection.product_id == product_id,
                V2ProductCollection.collection_id == target_category_id,
            )
        ).scalar_one_or_none()
        source_row = session.execute(
            select(V2ProductCollection).where(
                V2ProductCollection.product_id == product_id,
                V2ProductCollection.collection_id == source_category_id,
            )
        ).scalar_one_or_none()
        if source_row is None:
            continue
        if existing is not None:
            existing.is_primary = bool(existing.is_primary or is_primary)
            session.delete(source_row)
        else:
            source_row.collection_id = target_category_id
            source_row.is_primary = bool(is_primary)

    shopping_map_rows = session.execute(
        select(V2ShoppingAttributeOptionMap).where(V2ShoppingAttributeOptionMap.category_id == source_category_id)
    ).scalars().all()
    affected_shopping_map_rows = len(shopping_map_rows)
    for row in shopping_map_rows:
        existing = session.execute(
            select(V2ShoppingAttributeOptionMap).where(
                V2ShoppingAttributeOptionMap.channel == row.channel,
                V2ShoppingAttributeOptionMap.shopping_attribute == row.shopping_attribute,
                V2ShoppingAttributeOptionMap.category_id == target_category_id,
            )
        ).scalar_one_or_none()
        if existing is not None:
            session.delete(row)
        else:
            row.category_id = target_category_id
            row.updated_at = _utcnow()

    session.execute(
        update(V2Category)
        .where(V2Category.merged_into_id == source_category_id)
        .values(merged_into_id=target_category_id, updated_at=_utcnow())
    )

    source.merged_into_id = target_category_id
    source.updated_at = _utcnow()
    target.updated_at = _utcnow()

    cascade_job.status = "completed"
    cascade_job.affected_product_count = affected_product_rows
    cascade_job.completed_at = _utcnow()
    session.flush()

    return TaxonomyCascadeResult(
        cascade_job=cascade_job,
        source_category_id=source_category_id,
        target_category_id=target_category_id,
        operation="merge",
        affected_product_category_links=affected_product_category_links,
        affected_product_collection_links=affected_product_collection_links,
        affected_product_rows=affected_product_rows,
        affected_shopping_map_rows=affected_shopping_map_rows,
        reparented_child_categories=reparented_children,
    )


def confirm_duplicate_review(
    session: Session,
    *,
    review_id: int,
    reviewed_by: str,
    survivor_category_id: int,
    delete_downstream_option: bool = False,
) -> tuple[V2DuplicateReviewQueue | None, V2ShoppingCascadeJob | None]:
    review = session.get(V2DuplicateReviewQueue, review_id)
    if review is None:
        return None, None

    source_category_id: int
    if survivor_category_id == int(review.category_id_a):
        source_category_id = int(review.category_id_b)
    elif survivor_category_id == int(review.category_id_b):
        source_category_id = int(review.category_id_a)
    else:
        raise ValueError("survivor_category_id must match one of the review pair category ids")

    review.status = "confirmed_duplicate"
    review.reviewed_by = reviewed_by
    result = execute_taxonomy_merge(
        session,
        source_category_id=source_category_id,
        target_category_id=survivor_category_id,
        reviewed_by=reviewed_by,
        delete_downstream_option=delete_downstream_option,
    )
    return review, result.cascade_job
