from __future__ import annotations

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

from sqlalchemy import func, or_, select
from sqlalchemy.orm import Session

from db.models import (
    MagentoProductSourceSnapshot,
    MasterProduct,
    MasterProductAttributeValue,
    MasterProductPrice,
    PlytixProductSourceSnapshot,
    ShopifyProductSourceSnapshot,
)
from db.source_imports import normalize_column_key
from db.source_snapshot_enrichment import extract_image_urls
from db.source_snapshot_export import flatten_payload_for_export
from db.product_matrix_display import (
    load_magento_matrix_label_indexes,
    polish_magento_matrix_fields,
    polish_plytix_matrix_fields,
    polish_shopify_matrix_fields,
    sort_matrix_field_keys,
)


SNAPSHOT_CHANNELS = {"master", "plytix", "magento", "shopify"}
DEFAULT_PAGE_SIZE = 50
MAX_PAGE_SIZE = 200

PREVIEW_COLUMNS = (
    "sku",
    "label",
    "family",
    "title",
    "name",
    "thumbnail",
    "thumbnail_url",
    "featured_image_url",
    "primary_image_url",
    "image_urls",
    "base_image",
    "brand",
    "manufacturer",
    "product_type",
    "status",
    "price",
    "category_l1",
    "category_l2",
    "category_l3",
    "collection",
    "vendor",
    "handle",
)


def paginate_product_snapshots(
    session: Session,
    channel: str,
    *,
    page: int = 1,
    page_size: int = DEFAULT_PAGE_SIZE,
    q: Optional[str] = None,
    connection_id: Optional[int] = None,
) -> Dict[str, Any]:
    channel = (channel or "master").strip().lower()
    if channel not in SNAPSHOT_CHANNELS:
        raise ValueError(f"Unsupported channel: {channel}")

    page = max(1, int(page))
    page_size = min(MAX_PAGE_SIZE, max(1, int(page_size)))
    offset = (page - 1) * page_size

    if channel == "master":
        total, rows = _paginate_master(session, offset=offset, limit=page_size, q=q)
    else:
        total, rows = _paginate_source_snapshot(
            session,
            channel,
            offset=offset,
            limit=page_size,
            q=q,
            connection_id=connection_id,
        )

    columns = _discover_columns(rows)
    return {
        "channel": channel,
        "connection_id": connection_id,
        "page": page,
        "page_size": page_size,
        "total": total,
        "total_pages": max(1, (total + page_size - 1) // page_size) if total else 0,
        "columns": columns,
        "rows": rows,
    }


def paginate_product_matrix(
    session: Session,
    *,
    page: int = 1,
    page_size: int = DEFAULT_PAGE_SIZE,
    q: Optional[str] = None,
    magento_connection_id: Optional[int] = None,
    shopify_connection_id: Optional[int] = None,
    max_fields: int = 48,
    master_filters: Optional[Any] = None,
) -> Dict[str, Any]:
    page = max(1, int(page))
    page_size = min(MAX_PAGE_SIZE, max(1, int(page_size)))
    max_fields = min(120, max(8, int(max_fields)))
    offset = (page - 1) * page_size

    from db.master_product_filters import apply_master_attribute_filters, parse_master_filters

    base = select(MasterProduct.sku).where(MasterProduct.is_active.is_(True))
    if q:
        like = f"%{str(q).strip()}%"
        base = base.where(MasterProduct.sku.ilike(like))
    base = apply_master_attribute_filters(base, filters=parse_master_filters(master_filters))
    total = session.scalar(select(func.count()).select_from(base.subquery())) or 0
    skus = [
        sku
        for (sku,) in session.execute(
            base.order_by(MasterProduct.sku).offset(offset).limit(page_size)
        ).all()
        if sku
    ]

    master_map = _master_fields_by_sku(session, skus)
    plytix_map = _source_fields_by_sku(session, "plytix", skus)
    magento_map = _source_fields_by_sku(session, "magento", skus, connection_id=magento_connection_id)
    shopify_map = _source_fields_by_sku(session, "shopify", skus, connection_id=shopify_connection_id)
    magento_option_labels, magento_attr_labels = load_magento_matrix_label_indexes(
        session,
        magento_connection_id,
    )

    rows = []
    for sku in skus:
        rows.append(
            {
                "sku": sku,
                "master": _compact_fields(master_map.get(sku), max_fields=max_fields),
                "plytix": _compact_fields(
                    polish_plytix_matrix_fields(plytix_map.get(sku)),
                    max_fields=max_fields,
                ),
                "magento": _compact_fields(
                    polish_magento_matrix_fields(
                        magento_map.get(sku),
                        option_labels_by_code=magento_option_labels,
                        attribute_labels_by_code=magento_attr_labels,
                    ),
                    max_fields=max_fields,
                ),
                "shopify": _compact_fields(
                    polish_shopify_matrix_fields(shopify_map.get(sku)),
                    max_fields=max_fields,
                ),
            }
        )

    return {
        "page": page,
        "page_size": page_size,
        "total": int(total),
        "total_pages": max(1, (int(total) + page_size - 1) // page_size) if total else 0,
        "magento_connection_id": magento_connection_id,
        "shopify_connection_id": shopify_connection_id,
        "rows": rows,
    }


def export_snapshot_csv_rows(
    session: Session,
    channel: str,
    *,
    connection_id: Optional[int] = None,
    q: Optional[str] = None,
    limit: Optional[int] = None,
) -> List[Dict[str, Any]]:
    channel = (channel or "master").strip().lower()
    if channel == "master":
        _, rows = _paginate_master(session, offset=0, limit=limit or 50000, q=q)
        return [{**row["fields"], "sku": row["sku"]} for row in rows]
    from db.source_snapshot_export import export_source_snapshot_csv_rows

    return export_source_snapshot_csv_rows(
        session,
        channel,
        connection_id=connection_id,
        limit=limit,
    )


def _paginate_master(
    session: Session,
    *,
    offset: int,
    limit: int,
    q: Optional[str],
) -> Tuple[int, List[Dict[str, Any]]]:
    base = select(MasterProduct).where(MasterProduct.is_active.is_(True))
    if q:
        like = f"%{str(q).strip()}%"
        base = base.where(or_(MasterProduct.sku.ilike(like), MasterProduct.name.ilike(like)))
    total = session.scalar(select(func.count()).select_from(base.subquery())) or 0
    products = session.scalars(base.order_by(MasterProduct.sku).offset(offset).limit(limit)).all()
    skus = [p.sku for p in products if p.sku]
    fields_map = _master_fields_by_sku(session, skus)
    rows = [{"sku": sku, "fields": fields_map.get(sku, {"sku": sku})} for sku in skus]
    return int(total), rows


def _paginate_source_snapshot(
    session: Session,
    channel: str,
    *,
    offset: int,
    limit: int,
    q: Optional[str],
    connection_id: Optional[int],
) -> Tuple[int, List[Dict[str, Any]]]:
    model, channel_key = _snapshot_model(channel)
    base = select(model.sku).where(model.valid_to.is_(None))
    if connection_id is not None and channel in {"magento", "shopify"}:
        base = base.where(model.connection_id == connection_id)
    if q:
        like = f"%{str(q).strip()}%"
        base = base.where(model.sku.ilike(like))
    total = session.scalar(select(func.count()).select_from(base.subquery())) or 0
    skus = [
        sku
        for (sku,) in session.execute(base.order_by(model.sku).offset(offset).limit(limit)).all()
        if sku
    ]
    fields_map = _source_fields_by_sku(session, channel, skus, connection_id=connection_id)
    rows = [{"sku": sku, "fields": fields_map.get(sku, {"sku": sku})} for sku in skus]
    return int(total), rows


def _master_fields_by_sku(session: Session, skus: Sequence[str]) -> Dict[str, Dict[str, Any]]:
    if not skus:
        return {}
    products = session.scalars(select(MasterProduct).where(MasterProduct.sku.in_(list(skus)))).all()
    product_ids = [p.id for p in products]
    attrs_by_sku: Dict[str, Dict[str, str]] = {sku: {} for sku in skus}
    if product_ids:
        for sku, code, value in session.execute(
            select(
                MasterProductAttributeValue.sku,
                MasterProductAttributeValue.attribute_code,
                MasterProductAttributeValue.value,
            ).where(MasterProductAttributeValue.product_id.in_(product_ids))
        ).all():
            if value is None:
                continue
            attrs_by_sku.setdefault(str(sku), {})[normalize_column_key(str(code))] = str(value)

    prices_by_sku: Dict[str, str] = {}
    for sku, price_type, amount in session.execute(
        select(MasterProductPrice.sku, MasterProductPrice.price_type, MasterProductPrice.amount).where(
            MasterProductPrice.sku.in_(list(skus))
        )
    ).all():
        if amount is None:
            continue
        prices_by_sku[f"{sku}:{price_type}"] = str(amount)

    out: Dict[str, Dict[str, Any]] = {}
    for product in products:
        sku = product.sku
        fields: Dict[str, Any] = {"sku": sku}
        if product.name:
            fields["title"] = product.name
        for col in (
            "brand",
            "manufacturer",
            "product_family",
            "category_l1",
            "category_l2",
            "category_l3",
            "collection",
            "assembly_type",
            "item_size",
            "item_style",
            "status",
        ):
            value = getattr(product, col, None)
            if value not in (None, ""):
                fields[col] = str(value)
        for code, value in attrs_by_sku.get(sku, {}).items():
            fields.setdefault(code, value)
        for key, value in prices_by_sku.items():
            if key.startswith(f"{sku}:"):
                fields[f"price_{key.split(':', 1)[1]}"] = value
        out[sku] = fields
    return out


def _source_fields_by_sku(
    session: Session,
    channel: str,
    skus: Sequence[str],
    *,
    connection_id: Optional[int] = None,
) -> Dict[str, Dict[str, Any]]:
    if not skus:
        return {}
    model, channel_key = _snapshot_model(channel)
    stmt = select(model.sku, model.payload).where(model.valid_to.is_(None)).where(model.sku.in_(list(skus)))
    if connection_id is not None and channel in {"magento", "shopify"}:
        stmt = stmt.where(model.connection_id == connection_id)
    out: Dict[str, Dict[str, Any]] = {}
    for sku, payload in session.execute(stmt).all():
        if not sku or not isinstance(payload, dict):
            continue
        flat = flatten_payload_for_export(payload, channel=channel_key)
        flat["sku"] = str(sku)
        urls = extract_image_urls(flat, channel=channel_key)
        if urls and not flat.get("primary_image_url"):
            flat["primary_image_url"] = urls[0]
        if urls and not flat.get("thumbnail_url"):
            flat["thumbnail_url"] = urls[0]
        out[str(sku)] = flat
    return out


def _snapshot_model(channel: str):
    if channel == "plytix":
        return PlytixProductSourceSnapshot, "plytix"
    if channel == "magento":
        return MagentoProductSourceSnapshot, "magento"
    if channel == "shopify":
        return ShopifyProductSourceSnapshot, "shopify"
    raise ValueError(f"Unsupported source channel: {channel}")


def _discover_columns(rows: List[Dict[str, Any]]) -> List[str]:
    columns: List[str] = ["sku"]
    seen: set[str] = {"sku"}
    for preview in PREVIEW_COLUMNS:
        if preview != "sku" and preview not in seen:
            columns.append(preview)
            seen.add(preview)
    for row in rows:
        for key in row.get("fields", {}):
            if key not in seen:
                columns.append(key)
                seen.add(key)
    return columns[:40]


def _compact_fields(fields: Optional[Dict[str, Any]], *, max_fields: int) -> Optional[Dict[str, Any]]:
    if not fields:
        return None
    ordered_keys: List[str] = []
    for key in PREVIEW_COLUMNS:
        if key in fields and fields[key] not in (None, ""):
            ordered_keys.append(key)
    for key in sort_matrix_field_keys(fields):
        if key not in ordered_keys and fields[key] not in (None, ""):
            ordered_keys.append(key)
    compact: Dict[str, Any] = {}
    for key in ordered_keys[:max_fields]:
        value = fields[key]
        if isinstance(value, (dict, list)):
            text = json.dumps(value, sort_keys=True, default=str)
        else:
            text = str(value)
        if len(text) > 240:
            text = text[:237] + "..."
        compact[key] = text
    return compact or None
