"""Rebuild generated shell parents from saved base-SKU manual variation rules.

Examples:
  python -m app.jobs.rebuild_shell_parents_from_variation_rules --dry-run
  python -m app.jobs.rebuild_shell_parents_from_variation_rules --rules-json manual_variations.json --dry-run
  python -m app.jobs.rebuild_shell_parents_from_variation_rules --rules-json manual_variations.json --apply
  python -m app.jobs.rebuild_shell_parents_from_variation_rules --source-parent-skus B-PARENT,W-PARENT --apply
"""

from __future__ import annotations

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

from db.session import get_session

logger = logging.getLogger(__name__)


def run_rebuild_shell_parents(
    session,
    *,
    templates: Sequence[Dict[str, Any]],
    source_parent_skus: Optional[Sequence[str]] = None,
    assign_channels: Optional[Sequence[str]] = ("magento",),
    mapping_connection_ids: Optional[Dict[str, Optional[int]]] = None,
    create_identity_mappings: bool = True,
    apply: bool = False,
) -> Dict[str, Any]:
    from db.manual_variation_assignments import MANUAL_VARIATION_SOURCE, expand_manual_variation_templates
    from db.variation_builder import commit_variation_plan

    filtered_templates = _filter_templates(templates, source_parent_skus=source_parent_skus)
    expanded = expand_manual_variation_templates(session, filtered_templates)
    summary: Dict[str, Any] = {
        "status": "ok",
        "dry_run": not apply,
        "template_count": len(filtered_templates),
        "expanded_template_count": expanded.get("template_count", 0),
        "relation_count": expanded.get("relation_count", 0),
        "parent_count": expanded.get("parent_count", 0),
        "parent_skus": expanded.get("parent_skus", [])[:200],
        "assign_channels": list(assign_channels or []),
        "create_identity_mappings": bool(create_identity_mappings),
    }
    if not apply:
        return summary

    result = commit_variation_plan(
        session,
        groups=expanded.get("groups") or [],
        source_label=MANUAL_VARIATION_SOURCE,
        replace_existing=True,
        assign_channels=assign_channels or [],
        mapping_connection_ids=mapping_connection_ids or {},
        create_identity_mappings=bool(create_identity_mappings),
    )
    session.commit()
    summary.update(result)
    return summary


def load_templates_from_payload(payload: Dict[str, Any]) -> List[Dict[str, Any]]:
    templates = payload.get("templates")
    if not isinstance(templates, list):
        raise ValueError("rules payload must contain a templates list")
    return [item for item in templates if isinstance(item, dict)]


def load_templates_from_file(path: str) -> List[Dict[str, Any]]:
    with open(path, "r", encoding="utf-8") as fh:
        payload = json.load(fh)
    if not isinstance(payload, dict):
        raise ValueError("rules JSON must be an object with a templates array")
    return load_templates_from_payload(payload)


def _filter_templates(
    templates: Sequence[Dict[str, Any]],
    *,
    source_parent_skus: Optional[Sequence[str]],
) -> List[Dict[str, Any]]:
    normalized = [item for item in templates if isinstance(item, dict)]
    wanted = {
        str(sku or "").strip().upper()
        for sku in (source_parent_skus or [])
        if str(sku or "").strip()
    }
    if not wanted:
        return normalized
    return [
        item
        for item in normalized
        if str(item.get("source_parent_sku") or item.get("parent_sku") or "").strip().upper() in wanted
    ]


def _parse_csv(raw: Optional[str]) -> List[str]:
    if not raw:
        return []
    return [part.strip() for part in str(raw).split(",") if part.strip()]


def _parse_mapping_connection_ids(raw: Optional[str]) -> Dict[str, Optional[int]]:
    if not raw:
        return {}
    payload = json.loads(raw)
    if not isinstance(payload, dict):
        raise ValueError("--mapping-connection-ids-json must be a JSON object")
    out: Dict[str, Optional[int]] = {}
    for key, value in payload.items():
        name = str(key or "").strip()
        if not name:
            continue
        out[name] = int(value) if value is not None else None
    return out


def main() -> int:
    parser = argparse.ArgumentParser(description="Rebuild shell parents from saved manual variation rules")
    parser.add_argument(
        "--rules-json",
        default=None,
        help="Optional JSON export produced by remove_shell_parents --export-variation-rules",
    )
    parser.add_argument(
        "--source-parent-skus",
        default=None,
        help="Optional comma-separated base parent SKU filter (for example B-PARENT,W-PARENT)",
    )
    parser.add_argument(
        "--assign-channels",
        default="magento",
        help="Comma-separated channels to assign rebuilt parents to",
    )
    parser.add_argument(
        "--mapping-connection-ids-json",
        default=None,
        help='Optional JSON object like {"magento": 1, "shopify": 2}',
    )
    parser.add_argument("--no-identity-mappings", action="store_true")
    parser.add_argument("--apply", action="store_true", help="Persist rebuilt parent shells")
    parser.add_argument("--dry-run", action="store_true", help="Explicit dry-run")
    args = parser.parse_args()

    try:
        with get_session() as session:
            if args.rules_json:
                templates = load_templates_from_file(args.rules_json)
            else:
                from db.manual_variation_assignments import list_manual_variation_assignments

                templates = list_manual_variation_assignments(session).get("templates") or []
            result = run_rebuild_shell_parents(
                session,
                templates=templates,
                source_parent_skus=_parse_csv(args.source_parent_skus) or None,
                assign_channels=_parse_csv(args.assign_channels),
                mapping_connection_ids=_parse_mapping_connection_ids(args.mapping_connection_ids_json),
                create_identity_mappings=not args.no_identity_mappings,
                apply=bool(args.apply and not args.dry_run),
            )
    except Exception as exc:
        print(json.dumps({"status": "failed", "error": str(exc)}, indent=2))
        return 1

    print(json.dumps(result, indent=2, default=str))
    return 0 if result.get("status") == "ok" else 1


if __name__ == "__main__":
    logging.basicConfig(level=logging.INFO)
    sys.exit(main())
