"""
Examples:
    python -m app.jobs.audit_manual_taxonomy_rule_mismatches
    python -m app.jobs.audit_manual_taxonomy_rule_mismatches --rule-batch-size 50 --sample-limit 10
    python -m app.jobs.audit_manual_taxonomy_rule_mismatches --source-sku BF3
"""

from __future__ import annotations

import argparse
import json
from typing import Any, Dict, List, Optional, Sequence

from sqlalchemy import select

from app.jobs.audit_manual_taxonomy_suffix_mismatches import audit_manual_taxonomy_suffix_mismatches
from db.models import ManualTaxonomyAssignment
from db.session import get_session


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


def audit_manual_taxonomy_rule_mismatches(
    *,
    source_sku: Optional[str] = None,
    rule_batch_size: int = 50,
    sample_limit: int = 25,
    only_mismatches: bool = True,
) -> Dict[str, Any]:
    normalized_source = _clean(source_sku).upper()
    rule_batch_size = max(int(rule_batch_size or 50), 1)

    with get_session() as session:
        stmt = (
            select(ManualTaxonomyAssignment.source_sku)
            .where(ManualTaxonomyAssignment.assignment_status == "active")
            .order_by(ManualTaxonomyAssignment.source_sku)
        )
        if normalized_source:
            stmt = stmt.where(ManualTaxonomyAssignment.source_sku == normalized_source)
        source_skus = [
            _clean(value).upper()
            for value in session.scalars(stmt).all()
            if _clean(value)
        ]

    total_rules = len(source_skus)
    rule_batches = list(range(0, total_rules, rule_batch_size))
    reports: List[Dict[str, Any]] = []
    mismatch_rule_count = 0
    mismatch_row_count = 0

    for start in rule_batches:
        batch = source_skus[start : start + rule_batch_size]
        for current_source_sku in batch:
            result = audit_manual_taxonomy_suffix_mismatches(
                source_sku=current_source_sku,
                batch_size=500,
                sample_limit=sample_limit,
            )
            if int(result.get("mismatch_count") or 0) > 0:
                mismatch_rule_count += 1
                mismatch_row_count += int(result.get("mismatch_count") or 0)
            if only_mismatches and int(result.get("mismatch_count") or 0) <= 0:
                continue
            reports.append(result)

    return {
        "status": "ok",
        "source_sku": normalized_source or None,
        "rule_batch_size": rule_batch_size,
        "active_rule_count": total_rules,
        "reported_rule_count": len(reports),
        "mismatch_rule_count": mismatch_rule_count,
        "mismatch_row_count": mismatch_row_count,
        "reports": reports,
    }


def main(argv: Optional[Sequence[str]] = None) -> int:
    parser = argparse.ArgumentParser(description="Audit manual taxonomy mismatches grouped by source_sku rule")
    parser.add_argument("--source-sku", default=None, help="Optional single manual source_sku like BF3")
    parser.add_argument("--rule-batch-size", type=int, default=50, help="How many manual rules to inspect per batch")
    parser.add_argument("--sample-limit", type=int, default=25, help="Mismatch sample size per source_sku rule")
    parser.add_argument(
        "--include-clean",
        action="store_true",
        help="Include rules with zero mismatches in the output",
    )
    args = parser.parse_args(list(argv) if argv is not None else None)

    result = audit_manual_taxonomy_rule_mismatches(
        source_sku=args.source_sku,
        rule_batch_size=args.rule_batch_size,
        sample_limit=args.sample_limit,
        only_mismatches=not args.include_clean,
    )
    print(json.dumps(result, indent=2, default=str))
    return 0


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