"""Fill missing filter attributes from dominant values within the same collection.

Collection-shareable attrs (door_style, color, finish, cabinet_construction, …)
are uniform within a collection-style group in the SKU Attribute Workbook.
Leaf SKUs usually already have them; gaps are mostly variation PARENT shells.

For collections that are entirely empty (no peer donors), use
``python -m app.jobs.backfill_collection_filter_attributes`` instead.

Examples:
    python -m app.jobs.impute_collection_filter_attributes --dry-run
    python -m app.jobs.impute_collection_filter_attributes --attribute door_style --attribute color
    python -m app.jobs.impute_collection_filter_attributes --apply
"""

from __future__ import annotations

import argparse
import json
import sys
from collections import Counter, defaultdict
from typing import Any, Dict, List, Optional, Sequence

from sqlalchemy import select
from sqlalchemy.orm import Session

from db.channel_exports import compose_canonical_fields
from db.collection_landing_pages import canonical_collection
from db.collection_path_guard import is_polluted_collection_name
from db.models import MasterProduct, MasterProductAttributeValue
from db.session import get_session
from db.source_imports import normalize_column_key


DEFAULT_IMPUTE_ATTRIBUTES = (
    "door_style",
    "color",
    "finish",
    "cabinet_construction",
    "cabinet_door_overlay",
)


def impute_collection_filter_attributes(
    session: Session,
    *,
    attributes: Optional[Sequence[str]] = None,
    min_coverage_pct: float = 70.0,
    min_observed: int = 2,
    source_label: str = "collection_filter_impute",
    dry_run: bool = True,
) -> Dict[str, Any]:
    wanted = _codes(attributes or DEFAULT_IMPUTE_ATTRIBUTES)
    products = list(
        session.scalars(
            select(MasterProduct)
            .where(MasterProduct.is_active.is_(True))
            .order_by(MasterProduct.collection, MasterProduct.sku)
        ).all()
    )
    attrs_by_product = _attrs_by_product(session, [product.id for product in products])
    field_by_sku: Dict[str, Dict[str, Any]] = {}
    products_by_collection: Dict[str, List[MasterProduct]] = defaultdict(list)
    for product in products:
        fields, _ = compose_canonical_fields(_core_fields(product), attrs_by_product.get(product.id, {}), {})
        field_by_sku[product.sku] = fields
        collection = canonical_collection(fields.get("collection") or product.collection)
        if collection:
            products_by_collection[collection].append(product)

    updates: List[Dict[str, Any]] = []
    suggestions: List[Dict[str, Any]] = []
    actionable_collection_count = 0
    polluted_skipped_count = 0
    for collection, members in sorted(products_by_collection.items()):
        category_l1 = _dominant_category_l1(members)
        if is_polluted_collection_name(session, collection, category_l1=category_l1):
            polluted_skipped_count += 1
            continue
        actionable_collection_count += 1
        for code in wanted:
            values = [
                str(field_by_sku.get(product.sku, {}).get(code) or "").strip()
                for product in members
            ]
            non_empty = [value for value in values if value]
            if len(non_empty) < min_observed:
                continue
            value, count = Counter(non_empty).most_common(1)[0]
            coverage_pct = round(100.0 * count / len(non_empty), 1) if non_empty else 0.0
            if coverage_pct < min_coverage_pct:
                continue
            missing = [
                product
                for product in members
                if not str(field_by_sku.get(product.sku, {}).get(code) or "").strip()
            ]
            if not missing:
                continue
            suggestion = {
                "collection": collection,
                "attribute": code,
                "value": value,
                "observed_count": len(non_empty),
                "dominant_count": count,
                "dominant_pct": coverage_pct,
                "missing_count": len(missing),
                "sample_missing_skus": [product.sku for product in missing[:25]],
            }
            suggestions.append(suggestion)
            for product in missing:
                updates.append(
                    {
                        "product": product,
                        "sku": product.sku,
                        "attribute": code,
                        "value": value,
                        "collection": collection,
                    }
                )

    if not dry_run:
        for update in updates:
            _upsert_attr(
                session,
                update["product"],
                update["attribute"],
                update["value"],
                source_label,
            )

    return {
        "status": "ok",
        "dry_run": dry_run,
        "attribute_count": len(wanted),
        "collection_count": actionable_collection_count,
        "raw_collection_count": len(products_by_collection),
        "polluted_skipped_count": polluted_skipped_count,
        "suggestion_count": len(suggestions),
        "would_update": len(updates),
        "updated": 0 if dry_run else len(updates),
        "suggestions": suggestions[:100],
    }


def _attrs_by_product(session: Session, product_ids: Sequence[int]) -> Dict[int, Dict[str, Any]]:
    if not product_ids:
        return {}
    rows: Dict[int, Dict[str, Any]] = defaultdict(dict)
    for row in session.scalars(
        select(MasterProductAttributeValue).where(MasterProductAttributeValue.product_id.in_(list(product_ids)))
    ).all():
        if row.value is not None and str(row.value).strip():
            rows[int(row.product_id)][normalize_column_key(row.attribute_code)] = row.value
    return rows


def _upsert_attr(
    session: Session,
    product: MasterProduct,
    attribute_code: str,
    value: str,
    source_label: str,
) -> None:
    row = session.scalar(
        select(MasterProductAttributeValue)
        .where(MasterProductAttributeValue.product_id == product.id)
        .where(MasterProductAttributeValue.attribute_code == attribute_code)
    )
    if row is None:
        row = MasterProductAttributeValue(
            product_id=product.id,
            sku=product.sku,
            attribute_code=attribute_code,
        )
        session.add(row)
    row.value = value
    row.source_label = source_label


def _core_fields(product: MasterProduct) -> Dict[str, Any]:
    return {
        "title": product.name,
        "brand": product.brand,
        "manufacturer": product.manufacturer,
        "product_family": product.product_family,
        "category_l1": product.category_l1,
        "category_l2": product.category_l2,
        "category_l3": product.category_l3,
        "collection": product.collection,
        "assembly_type": product.assembly_type,
        "item_size": product.item_size,
        "status": product.status,
        "stocked": product.stocked,
    }


def _codes(values: Sequence[str]) -> List[str]:
    out: List[str] = []
    for value in values:
        code = normalize_column_key(value)
        if code and code not in out:
            out.append(code)
    return out


def _dominant_category_l1(products: Sequence[MasterProduct]) -> Optional[str]:
    counts = Counter(str(product.category_l1 or "").strip() for product in products if str(product.category_l1 or "").strip())
    if not counts:
        return None
    return counts.most_common(1)[0][0]


def main() -> int:
    parser = argparse.ArgumentParser(description="Impute missing filter attributes from collection-dominant values")
    parser.add_argument("--attribute", action="append", default=None, help="Attribute code to impute; repeatable")
    parser.add_argument("--min-coverage-pct", type=float, default=70.0)
    parser.add_argument("--min-observed", type=int, default=2)
    parser.add_argument("--source-label", default="collection_filter_impute")
    parser.add_argument("--apply", action="store_true", help="Write suggested values")
    parser.add_argument("--dry-run", action="store_true", help="Preview only; default unless --apply is set")
    args = parser.parse_args()

    dry_run = not args.apply
    with get_session() as session:
        result = impute_collection_filter_attributes(
            session,
            attributes=args.attribute,
            min_coverage_pct=args.min_coverage_pct,
            min_observed=args.min_observed,
            source_label=args.source_label,
            dry_run=dry_run,
        )
        if not dry_run:
            session.commit()
    print(json.dumps(result, indent=2, default=str))
    return 0


if __name__ == "__main__":
    sys.exit(main())
