"""Magento attribute preflight for a specific connection.

Use before product pushes to verify that active Magento channel aliases exist in
the target connection's default attribute set registry.

CLI:
    python -m app.jobs.magento_attribute_preflight --connection-id 2
    python -m app.jobs.magento_attribute_preflight --connection-id 2 --mark-missing-create-new
    python -m app.jobs.magento_attribute_preflight --connection-id 2 --mark-missing-create-new --provision --baseline-after
"""

from __future__ import annotations

import argparse
import json
import re
import sys
from typing import Any, Dict, List

from sqlalchemy import select
from sqlalchemy.orm import Session

from db.models import ChannelAttributeAlias, MagentoAttributeRegistry


def run_magento_attribute_preflight(
    *,
    connection_id: int,
    mark_missing_create_new: bool = False,
    provision: bool = False,
    baseline_after: bool = False,
) -> Dict[str, Any]:
    from db.session import get_session

    with get_session() as session:
        result = analyze_magento_attribute_preflight(session, connection_id=connection_id)
        mark_result = None
        provision_result = None
        baseline_result = None

        if mark_missing_create_new:
            mark_result = mark_missing_aliases_create_new(session, connection_id=connection_id)
            session.commit()
            result = analyze_magento_attribute_preflight(session, connection_id=connection_id)

        if provision:
            from db.channel_attribute_provision import provision_pending_for_channels

            provision_result = provision_pending_for_channels(
                session,
                ["magento"],
                magento_connection_id=connection_id,
                dry_run=False,
            )
            session.commit()
            result = analyze_magento_attribute_preflight(session, connection_id=connection_id)

        if baseline_after:
            from app.jobs.magento_baseline_sync import run_baseline_sync

            baseline_result = run_baseline_sync(connection_id, force=True, session=session)
            result = analyze_magento_attribute_preflight(session, connection_id=connection_id)

        return {
            **result,
            "mark_missing_create_new": mark_result,
            "provision": provision_result,
            "baseline_after": baseline_result,
        }


def analyze_magento_attribute_preflight(session: Session, *, connection_id: int) -> Dict[str, Any]:
    aliases = _active_magento_aliases(session)
    existing_codes = _existing_magento_codes(session, connection_id)
    missing = [
        _alias_dict(alias)
        for alias in aliases
        if _normalize_code(alias.channel_attribute_code) not in existing_codes
    ]
    invalid = [
        _alias_dict(alias)
        for alias in aliases
        if not _valid_magento_attribute_code(alias.channel_attribute_code)
    ]
    missing_create_new = [row for row in missing if row["action"] == "create_new"]
    missing_use_existing = [row for row in missing if row["action"] != "create_new"]
    status = "blocked" if missing_use_existing or invalid else "ok"
    if missing_create_new and not missing_use_existing and not invalid:
        status = "needs_provision"
    return {
        "status": status,
        "connection_id": connection_id,
        "existing_magento_attribute_count": len(existing_codes),
        "active_magento_alias_count": len(aliases),
        "existing_alias_count": len(aliases) - len(missing),
        "missing_alias_count": len(missing),
        "missing_create_new_count": len(missing_create_new),
        "missing_use_existing_count": len(missing_use_existing),
        "invalid_alias_count": len(invalid),
        "missing_aliases": missing,
        "invalid_aliases": invalid,
    }


def mark_missing_aliases_create_new(session: Session, *, connection_id: int) -> Dict[str, Any]:
    aliases = _active_magento_aliases(session)
    existing_codes = _existing_magento_codes(session, connection_id)
    marked: List[Dict[str, Any]] = []
    skipped_invalid: List[Dict[str, Any]] = []
    for alias in aliases:
        code = _normalize_code(alias.channel_attribute_code)
        if not code or code in existing_codes:
            continue
        if not _valid_magento_attribute_code(code):
            skipped_invalid.append(_alias_dict(alias))
            continue
        alias.action = "create_new"
        alias.notes = _append_note(alias.notes, f"marked create_new for Magento connection {connection_id}")
        marked.append(_alias_dict(alias))
    session.flush()
    return {
        "marked_count": len(marked),
        "skipped_invalid_count": len(skipped_invalid),
        "marked": marked,
        "skipped_invalid": skipped_invalid,
    }


def _active_magento_aliases(session: Session) -> List[ChannelAttributeAlias]:
    return list(
        session.scalars(
            select(ChannelAttributeAlias)
            .where(ChannelAttributeAlias.channel_code == "magento")
            .where(ChannelAttributeAlias.is_active.is_(True))
            .order_by(ChannelAttributeAlias.canonical_code, ChannelAttributeAlias.channel_attribute_code)
        ).all()
    )


def _existing_magento_codes(session: Session, connection_id: int) -> set[str]:
    return {
        _normalize_code(code)
        for (code,) in session.execute(
            select(MagentoAttributeRegistry.attribute_code).where(
                MagentoAttributeRegistry.connection_id == connection_id
            )
        ).all()
        if _normalize_code(code)
    }


def _alias_dict(alias: ChannelAttributeAlias) -> Dict[str, Any]:
    return {
        "id": alias.id,
        "canonical_code": alias.canonical_code,
        "channel_attribute_code": _normalize_code(alias.channel_attribute_code),
        "mapping_scope": alias.mapping_scope,
        "action": alias.action,
        "data_type": alias.data_type,
    }


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


def _valid_magento_attribute_code(value: Any) -> bool:
    return bool(re.match(r"^[a-z0-9_]{1,64}$", _normalize_code(value)))


def _append_note(existing: Any, note: str) -> str:
    text = str(existing or "").strip()
    return f"{text};{note}" if text else note


def main() -> int:
    parser = argparse.ArgumentParser(description="Magento attribute preflight for one connection")
    parser.add_argument("--connection-id", type=int, required=True)
    parser.add_argument(
        "--mark-missing-create-new",
        action="store_true",
        help="Mark missing active Magento aliases as create_new",
    )
    parser.add_argument("--provision", action="store_true", help="Provision pending create_new Magento aliases")
    parser.add_argument("--baseline-after", action="store_true", help="Force baseline refresh after provisioning")
    args = parser.parse_args()

    result = run_magento_attribute_preflight(
        connection_id=args.connection_id,
        mark_missing_create_new=args.mark_missing_create_new,
        provision=args.provision,
        baseline_after=args.baseline_after,
    )
    print(json.dumps(result, indent=2, default=str))
    return 1 if result.get("status") == "blocked" else 0


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