"""Pull all Shopify products/variants into shopify_product_source_snapshot.

CLI:
    python -m app.jobs.shopify_pull [--connection-id N] [--dry-run]
"""

from __future__ import annotations

import argparse
import hashlib
import json
import logging
import sys
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional

import pandas as pd

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)


PRODUCTS_PULL_QUERY = """
query shopifyProductsPull($first: Int!, $after: String) {
  products(first: $first, after: $after) {
    pageInfo { hasNextPage endCursor }
    nodes {
      id
      title
      handle
      descriptionHtml
      vendor
      productType
      status
      tags
      createdAt
      updatedAt
      seo { title description }
      options { name values }
      metafields(first: 50) {
        nodes { namespace key type value }
      }
      variants(first: 100) {
        nodes {
          id
          sku
          title
          price
          compareAtPrice
          barcode
          inventoryQuantity
          inventoryPolicy
          taxable
          selectedOptions { name value }
          inventoryItem {
            id
            measurement { weight { value unit } }
          }
          metafields(first: 50) {
            nodes { namespace key type value }
          }
        }
      }
    }
  }
}
"""
# Legacy static query kept for reference/tests; runtime uses shopify.pull_query.build_products_pull_query.


def run_shopify_pull(
    *,
    connection_id: Optional[int] = None,
    shop_code: Optional[str] = None,
    dry_run: bool = False,
    session=None,
) -> Dict[str, Any]:
    from db.compat_connections import SHOPIFY_COMPAT_ID_OFFSET, decode_compat_connection_id
    from db.models import ShopifyConnection
    from db.repositories import SqlAlchemyIngestRepository
    from db.session import get_session
    from db.source_imports import import_shopify_source_dataframe, normalize_column_key
    from settings import load_shopify_config
    from shopify.connections import build_client, get_active_connection

    own_session = session is None
    ctx = get_session() if own_session else _noop_ctx(session)
    with ctx as sess:
        connection: Optional[ShopifyConnection] = None
        native_id: Optional[int] = None
        if connection_id:
            native_id = connection_id
            if connection_id >= SHOPIFY_COMPAT_ID_OFFSET:
                channel_type, native_id = decode_compat_connection_id(connection_id)
                if channel_type != "shopify":
                    return {
                        "status": "failed",
                        "error": f"Connection {connection_id} is not a Shopify connection",
                    }
            connection = sess.get(ShopifyConnection, native_id)
            if connection is None:
                return {
                    "status": "failed",
                    "error": f"Shopify connection {connection_id} not found (native id {native_id})",
                }
        elif shop_code:
            connection = get_active_connection(sess, shop_code=shop_code)
            native_id = connection.id if connection else None

        cfg = load_shopify_config()
        client = build_client(connection)
        effective_shop_code = (connection.shop_code if connection else None) or shop_code or cfg.shop_code
        if not client.shop_domain:
            return {"status": "failed", "error": "Shopify shop_domain is required"}

        records = _fetch_product_records(client)
        if not records:
            return {"status": "failed", "error": "No Shopify products returned"}

        if dry_run:
            return {
                "status": "success",
                "dry_run": True,
                "shop_code": effective_shop_code,
                "connection_id": native_id,
                "total_products": len({r.get("shopify_product_id") for r in records}),
                "distinct_skus": len({r["sku"] for r in records if r.get("sku")}),
                "variant_rows": len(records),
            }

        df = pd.DataFrame(records)
        payload_hash = hashlib.sha256(
            json.dumps(records, sort_keys=True, default=str).encode("utf-8")
        ).hexdigest()
        repo = SqlAlchemyIngestRepository(sess)
        ingest_id, stats = import_shopify_source_dataframe(
            repo,
            df,
            file_name=f"shopify_api_pull_{native_id or effective_shop_code}_{datetime.now(timezone.utc).isoformat()}.json",
            file_hash=payload_hash,
            connection_id=native_id,
        )
        if own_session:
            sess.commit()
        return {
            "status": "success",
            "dry_run": False,
            "shop_code": effective_shop_code,
            "connection_id": native_id,
            "ingest_id": ingest_id,
            **stats,
        }


def _fetch_product_records(client: Any) -> List[Dict[str, Any]]:
    from db.source_imports import normalize_column_key
    from shopify.pull_query import build_products_pull_query

    query = build_products_pull_query(client)
    records: List[Dict[str, Any]] = []
    after: Optional[str] = None
    while True:
        data = client.graphql(query, {"first": 50, "after": after})
        connection = data.get("products") or {}
        for product in connection.get("nodes") or []:
            if not isinstance(product, dict):
                continue
            records.extend(_product_to_variant_records(product, normalize_column_key))
        page_info = connection.get("pageInfo") or {}
        if not page_info.get("hasNextPage"):
            break
        after = page_info.get("endCursor")
        if not after:
            break
    return records


def _product_to_variant_records(product: Dict[str, Any], normalize_key) -> List[Dict[str, Any]]:
    from shopify.pull_query import extract_shopify_media, variant_weight

    product_id = product.get("id")
    media_fields = extract_shopify_media(product)
    base: Dict[str, Any] = {
        "shopify_product_id": product_id,
        "title": product.get("title"),
        "handle": product.get("handle"),
        "description_html": product.get("descriptionHtml"),
        "vendor": product.get("vendor"),
        "product_type": product.get("productType"),
        "status": product.get("status"),
        "tags": ", ".join(product.get("tags") or []) if isinstance(product.get("tags"), list) else product.get("tags"),
        "created_at": product.get("createdAt"),
        "updated_at": product.get("updatedAt"),
    }
    seo = product.get("seo") or {}
    if isinstance(seo, dict):
        if seo.get("title"):
            base["seo_title"] = seo.get("title")
        if seo.get("description"):
            base["seo_description"] = seo.get("description")

    for idx, option in enumerate(product.get("options") or []):
        if not isinstance(option, dict):
            continue
        name = str(option.get("name") or f"option_{idx + 1}").strip()
        base[f"option_{normalize_key(name)}"] = ", ".join(option.get("values") or [])

    for mf in (product.get("metafields") or {}).get("nodes") or []:
        if not isinstance(mf, dict):
            continue
        ns = str(mf.get("namespace") or "").strip()
        key = str(mf.get("key") or "").strip()
        if ns and key:
            base[f"metafield_{normalize_key(f'{ns}_{key}')}"] = mf.get("value")

    variants = ((product.get("variants") or {}).get("nodes") or [])
    if not variants:
        return []

    rows: List[Dict[str, Any]] = []
    for variant in variants:
        if not isinstance(variant, dict):
            continue
        sku = str(variant.get("sku") or "").strip()
        if not sku:
            inventory_item = variant.get("inventoryItem") or {}
            if isinstance(inventory_item, dict):
                sku = str(inventory_item.get("sku") or "").strip()
        if not sku:
            continue
        weight, weight_unit = variant_weight(variant)
        row = dict(base)
        row.update(media_fields)
        row.update(
            {
                "sku": sku,
                "shopify_variant_id": variant.get("id"),
                "variant_title": variant.get("title"),
                "price": variant.get("price"),
                "compare_at_price": variant.get("compareAtPrice"),
                "barcode": variant.get("barcode"),
                "inventory_quantity": variant.get("inventoryQuantity"),
                "inventory_policy": variant.get("inventoryPolicy"),
                "taxable": variant.get("taxable"),
                "weight": weight,
                "weight_unit": weight_unit,
            }
        )
        for opt in variant.get("selectedOptions") or []:
            if not isinstance(opt, dict):
                continue
            name = str(opt.get("name") or "").strip()
            if name:
                row[f"variant_option_{normalize_key(name)}"] = opt.get("value")
        for mf in (variant.get("metafields") or {}).get("nodes") or []:
            if not isinstance(mf, dict):
                continue
            ns = str(mf.get("namespace") or "").strip()
            key = str(mf.get("key") or "").strip()
            if ns and key:
                row[f"variant_metafield_{normalize_key(f'{ns}_{key}')}"] = mf.get("value")
        rows.append(row)
    return rows


class _noop_ctx:
    def __init__(self, obj):
        self._obj = obj

    def __enter__(self):
        return self._obj

    def __exit__(self, *args):
        pass


def main() -> int:
    parser = argparse.ArgumentParser(description="Pull Shopify products into source snapshots")
    parser.add_argument("--connection-id", type=int, default=None)
    parser.add_argument("--shop-code", default=None)
    parser.add_argument("--dry-run", action="store_true")
    args = parser.parse_args()

    result = run_shopify_pull(
        connection_id=args.connection_id,
        shop_code=args.shop_code,
        dry_run=args.dry_run,
    )
    if result.get("status") != "success":
        print(result.get("error", "Shopify pull failed"), file=sys.stderr)
        return 1
    print(json.dumps(result, indent=2, default=str))
    return 0


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