from __future__ import annotations

from typing import Any, Dict, Optional

from sqlalchemy.orm import Session

from db.models import ChannelCatalogPolicy, ChannelVariationPolicy
from db.models import ChannelConnection
from db.source_imports import normalize_column_key


DEFAULT_PRODUCT_STRUCTURE = {
    "shopify": "flat",
    "magento": "parent_children",
    "plytix": "flat",
}

ProductStructure = str  # "flat" | "parent_children"


def resolve_product_structure(
    channel_code: str,
    *,
    capabilities: Optional[Dict[str, Any]] = None,
) -> ProductStructure:
    channel = normalize_column_key(channel_code)
    if capabilities:
        raw = capabilities.get("product_structure") or capabilities.get("productStructure")
        if raw:
            normalized = str(raw).strip().lower()
            if normalized in {"flat", "parent_children", "parent-children", "parentchildren"}:
                return "parent_children" if "parent" in normalized else "flat"
    return DEFAULT_PRODUCT_STRUCTURE.get(channel, "flat")


def product_structure_for_connection(session: Session, connection: ChannelConnection) -> ProductStructure:
    return resolve_product_structure(
        connection.channel_type,
        capabilities=connection.capabilities or {},
    )


def resolve_product_structure_for_connection(
    session: Session,
    channel_code: str,
    *,
    connection_id: Optional[int] = None,
    capabilities: Optional[Dict[str, Any]] = None,
) -> ProductStructure:
    channel = normalize_column_key(channel_code)
    base = resolve_product_structure(channel, capabilities=capabilities)
    native_connection_id = connection_id
    if native_connection_id is not None and channel in {"shopify", "magento"}:
        from db.compat_connections import resolve_native_connection_id

        native_connection_id = resolve_native_connection_id(channel, native_connection_id)

    variation_stmt = (
        session.query(ChannelVariationPolicy)
        .filter(ChannelVariationPolicy.channel_code == channel)
        .filter(ChannelVariationPolicy.is_active.is_(True))
    )
    if native_connection_id is None:
        variation_stmt = variation_stmt.filter(ChannelVariationPolicy.connection_id.is_(None))
    else:
        variation_stmt = variation_stmt.filter(
            (ChannelVariationPolicy.connection_id == native_connection_id)
            | (ChannelVariationPolicy.connection_id.is_(None))
        )
    variation = variation_stmt.order_by(ChannelVariationPolicy.connection_id.is_(None), ChannelVariationPolicy.id.desc()).first()
    mode = normalize_column_key(getattr(variation, "variation_mode", "") or "")
    if mode in {"simple_only", "simple"}:
        return "flat"
    if mode in {"single_axis", "multi_axis", "variants", "variant"}:
        return "parent_children"

    policy_stmt = (
        session.query(ChannelCatalogPolicy)
        .filter(ChannelCatalogPolicy.channel_code == channel)
        .filter(ChannelCatalogPolicy.is_active.is_(True))
    )
    if native_connection_id is None:
        policy_stmt = policy_stmt.filter(ChannelCatalogPolicy.connection_id.is_(None))
    else:
        policy_stmt = policy_stmt.filter(
            (ChannelCatalogPolicy.connection_id == native_connection_id)
            | (ChannelCatalogPolicy.connection_id.is_(None))
        )
    policy = policy_stmt.order_by(ChannelCatalogPolicy.connection_id.is_(None), ChannelCatalogPolicy.id.desc()).first()
    mode = normalize_column_key(getattr(policy, "variation_mode", "") or "")
    if mode in {"simple_only", "simple"}:
        return "flat"
    if mode in {"single_axis", "multi_axis", "variants", "variant"}:
        return "parent_children"
    return base


def default_capabilities_for_channel(channel_code: str) -> Dict[str, Any]:
    channel = normalize_column_key(channel_code)
    if channel not in DEFAULT_PRODUCT_STRUCTURE:
        raise ValueError(f"Unsupported channel_code: {channel_code!r}")
    return {"product_structure": DEFAULT_PRODUCT_STRUCTURE[channel]}
