"""Inspect channel projection cache drift and optionally invalidate stale rows.

Examples:
    python -m app.jobs.inspect_channel_projection --channel shopify
    python -m app.jobs.inspect_channel_projection --channel magento --connection-id 1 --fix
    python -m app.jobs.inspect_channel_projection --channel shopify --sku ACH-B12 --sku ACH-B15 --fix
"""

from __future__ import annotations

import argparse
import json
import logging
import sys
from typing import List, Optional

from db.channel_projection_invalidation import inspect_projection_drift
from db.session import get_session


logger = logging.getLogger(__name__)


def main(argv: Optional[List[str]] = None) -> int:
    logging.basicConfig(level=logging.INFO)
    parser = argparse.ArgumentParser(description="Inspect/invalidate channel projection cache drift")
    parser.add_argument("--channel", required=True, choices=["magento", "shopify", "plytix"])
    parser.add_argument("--connection-id", type=int, default=None)
    parser.add_argument("--sku", action="append", dest="skus", help="Limit to SKU(s); repeatable")
    parser.add_argument(
        "--fix",
        action="store_true",
        help="Mark drifted projection rows stale (default is inspect-only)",
    )
    args = parser.parse_args(argv)

    with get_session() as session:
        report = inspect_projection_drift(
            session,
            channel_code=args.channel,
            connection_id=args.connection_id,
            skus=args.skus,
            apply=args.fix,
        )
        if args.fix:
            session.commit()
        else:
            session.rollback()

    print(json.dumps(report, indent=2, default=str))
    logger.info(
        "projection-inspect channel=%s drift=%s applied=%s invalidated=%s",
        report.get("channel_code"),
        report.get("drift_count"),
        report.get("applied"),
        report.get("invalidated"),
    )
    return 0


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