"""
库存服务 — 校验 / 扣减 / 释放

职责：
1. validate_stock：下单前校验库存是否足够
2. deduct_stock：下单成功后扣减库存（含 variant 汇总处理）
3. restore_stock：取消/超时释放库存

所有写操作必须在外层事务中执行，使用 with_for_update() 防止并发超卖。
库存低于 low_stock_threshold 时发送 stock_low 信号。
库存变动时根据 status_type 自动切换库存状态。
"""
from decimal import Decimal
from typing import Optional

from fastapi import HTTPException
from sqlalchemy import and_, or_, select
from sqlalchemy.ext.asyncio import AsyncSession

from app.core.models.product import Product, ProductVariant
from app.core.models.stock_status import StockStatus
from app.core.models.order import Order
from app.core.models.tenant_settings import TenantSettings
from app.core.services.pricing import PricingCartItem
from app.core.signals import stock_low


async def _takeover_state(db: AsyncSession, tenant_id: int):
    """返回该租户的进销存生命周期状态；None 表示未接管，走本模块原有逻辑。"""
    from app.plugins.inventory.lifecycle import takeover_active
    from app.plugins.inventory.services import get_state
    state = await get_state(db, tenant_id)
    return state if takeover_active(state) else None


async def _guard_writable(db: AsyncSession, tenant_id: int) -> None:
    """suspended（主动停用）租户拒绝一切库存写入，且不回落旧逻辑。"""
    from app.plugins.inventory.lifecycle import assert_writable
    from app.plugins.inventory.services import get_state
    assert_writable(await get_state(db, tenant_id))


async def is_inventory_takeover(db: AsyncSession, tenant_id: int) -> bool:
    """该租户是否已被进销存接管（供订单路由判定「下单即锁库」，不再依赖遗留 deduct_statuses）。"""
    return await _takeover_state(db, tenant_id) is not None


async def _guarded_takeover(db: AsyncSession, tenant_id: int):
    """一次查询完成「拒写守卫 + 接管判定」，避免热路径重复查 inventory_settings。

    返回已接管的状态（active/arrears），未接管返回 None，suspended 直接抛。
    """
    from app.plugins.inventory.lifecycle import assert_writable, takeover_active
    from app.plugins.inventory.services import get_state
    state = await get_state(db, tenant_id)
    assert_writable(state)               # suspended → 抛
    return state if takeover_active(state) else None


async def _get_type_default_status(
    db: AsyncSession, tenant_id: int, status_type: str,
) -> Optional[StockStatus]:
    result = await db.execute(
        select(StockStatus).where(
            StockStatus.tenant_id == tenant_id,
            StockStatus.status_type == status_type,
            StockStatus.is_type_default == 1,
        )
    )
    return result.scalar_one_or_none()


async def _auto_switch_stock_status(
    db: AsyncSession, product: Product, new_stock_qty: Decimal,
) -> None:
    """根据库存变动自动切换状态。

    规则：
    - 库存从正数→0/负数：如果当前状态是有货类，切到缺货类默认
    - 库存从0/负数→正数：如果当前状态是缺货类，切到有货类默认
    - 同类型内不自动切换
    """
    if not product.stock_status_id:
        return

    current_ss = await db.execute(
        select(StockStatus).where(StockStatus.id == product.stock_status_id)
    )
    current = current_ss.scalar_one_or_none()
    if not current:
        return

    if new_stock_qty <= 0 and current.status_type == "in_stock":
        target = await _get_type_default_status(db, product.tenant_id, "out_of_stock")
        if target:
            product.stock_status_id = target.id
    elif new_stock_qty > 0 and current.status_type == "out_of_stock":
        target = await _get_type_default_status(db, product.tenant_id, "in_stock")
        if target:
            product.stock_status_id = target.id


async def validate_stock(
    db: AsyncSession,
    tenant_id: int,
    items: list[PricingCartItem],
) -> None:
    """
    校验所有商品的库存是否满足购买数量。

    规则：
    - 有 variant_id：查 ProductVariant.stock_qty
    - 无 variant_id：查 Product.stock_qty
    - allow_oversell == 0 时，stock_qty < qty 抛出 400
    - allow_oversell == 1 时不校验（允许负库存）
    - 商品必须 status == "active"，variant 必须 is_active == 1
    """
    for it in items:
        if it.qty <= 0:
            raise HTTPException(status_code=400, detail="购买数量必须大于 0")

    # 接管后仍须复用既有商品/SKU/可购买状态校验；只有库存数量改由账本判断。
    inventory_takeover = await _takeover_state(db, tenant_id) is not None

    # 全局超卖开关：开启则跳过所有库存校验
    ts = (await db.execute(select(TenantSettings).where(TenantSettings.tenant_id == tenant_id))).scalar_one_or_none()
    global_oversell = bool(ts and ts.allow_oversell_global)

    product_ids = [it.product_id for it in items]
    result = await db.execute(
        select(Product).where(
            Product.id.in_(product_ids),
            Product.tenant_id == tenant_id,
            Product.status == "active",
        )
    )
    products = {p.id: p for p in result.scalars().all()}

    variant_map: dict[int, ProductVariant] = {}
    variant_pairs = [
        and_(ProductVariant.id == it.variant_id, ProductVariant.product_id == it.product_id)
        for it in items
        if it.variant_id is not None
    ]
    if variant_pairs:
        vr = await db.execute(
            select(ProductVariant).where(
                or_(*variant_pairs),
                ProductVariant.tenant_id == tenant_id,
                ProductVariant.is_active == 1,
            )
        )
        variant_map = {v.id: v for v in vr.scalars().all()}

    ss_ids = {p.stock_status_id for p in products.values() if p.stock_status_id}
    ss_map: dict[int, StockStatus] = {}
    if ss_ids:
        ss_r = await db.execute(select(StockStatus).where(StockStatus.id.in_(ss_ids)))
        ss_map = {s.id: s for s in ss_r.scalars().all()}

    errors: list[dict] = []
    for it in items:
        product = products.get(it.product_id)
        if product is None:
            errors.append({"product_id": it.product_id, "variant_id": it.variant_id, "unit_name": it.unit_name, "msg": f"商品 {it.product_id} 不存在或已下架"})
            continue

        if product.stock_status_id:
            ss = ss_map.get(product.stock_status_id)
            if ss and not ss.allow_purchase:
                label = ss.badge_text or ss.name
                errors.append({"product_id": it.product_id, "variant_id": it.variant_id, "unit_name": it.unit_name, "msg": f"《{product.name}》当前状态为「{label}」，不可购买"})
                continue

        variant: Optional[ProductVariant] = None
        if it.variant_id:
            variant = variant_map.get(it.variant_id)
            if variant is None:
                errors.append({"product_id": it.product_id, "variant_id": it.variant_id, "unit_name": it.unit_name, "msg": f"SKU {it.variant_id} 不存在或已禁用"})
                continue

        # 接管租户数量改由账本校验；全局/商品级超卖均跳过旧投影数量校验。
        if inventory_takeover or global_oversell or product.allow_oversell:
            continue

        # 取实际校验的库存
        qty_available: Decimal
        if variant:
            qty_available = variant.stock_qty
        else:
            qty_available = product.stock_qty

        deduct_amount = it.stock_qty_override if it.stock_qty_override is not None else it.qty
        if qty_available < deduct_amount:
            variant_name = variant.attributes.get("color", "") if variant else ""
            errors.append({
                "product_id": it.product_id,
                "variant_id": it.variant_id,
                "unit_name": it.unit_name,
                "available": str(qty_available),
                "required": str(deduct_amount),
                "msg": f"《{product.name}》{variant_name} 库存不足（剩余 {qty_available}，需 {deduct_amount}）",
            })

    if errors:
        raise HTTPException(status_code=400, detail=errors)

    if inventory_takeover:
        from app.plugins.inventory.services import validate_via_ledger
        await validate_via_ledger(
            db, tenant_id, items,
            oversell_product_ids={p.id for p in products.values() if p.allow_oversell},
        )


async def deduct_stock(
    db: AsyncSession,
    tenant_id: int,
    items: list[PricingCartItem],
    order_id: int,
    force_oversell: bool = False,
) -> None:
    """
    扣减库存（在下单事务中调用）。

    规则：
    - 有 variant_id：锁 ProductVariant 行，扣 variant.stock_qty
      同时汇总更新 Product.stock_qty（减少对应数量）
    - 无 variant_id：锁 Product 行，扣 product.stock_qty
    - allow_oversell == 1 时允许负库存，不做额外限制
    - 扣减后若库存低于 low_stock_threshold，发送 stock_low 信号
    - 扣减后根据 stock_qty 自动切换库存状态
    """
    for it in items:
        if it.qty <= 0:
            raise HTTPException(status_code=400, detail="购买数量必须大于 0")

    # 进销存接管：suspended 拒写；active/arrears 委托账本扣减（一次查询完成守卫+判定）
    if await _guarded_takeover(db, tenant_id) is not None:
        # force_oversell=True 来自 POS 离线同步：走离线入账，永远允许负库存，
        # 批次缺口进未分配异常桶
        if force_oversell:
            from app.plugins.inventory.services import deduct_offline_via_ledger
            return await deduct_offline_via_ledger(db, tenant_id, items, order_id)
        # 在线订单：下单锁库（sellable 不动、reserved 增加），发货时才实扣
        from app.plugins.inventory.services import lock_via_ledger
        return await lock_via_ledger(db, tenant_id, items, order_id)

    product_ids = [it.product_id for it in items]
    result = await db.execute(
        select(Product).where(
            Product.id.in_(product_ids),
            Product.tenant_id == tenant_id,
        ).with_for_update()
    )
    products = {p.id: p for p in result.scalars().all()}

    variant_map: dict[int, ProductVariant] = {}
    variant_pairs = [
        and_(ProductVariant.id == it.variant_id, ProductVariant.product_id == it.product_id)
        for it in items
        if it.variant_id is not None
    ]
    if variant_pairs:
        vr = await db.execute(
            select(ProductVariant).where(
                or_(*variant_pairs),
                ProductVariant.tenant_id == tenant_id,
            ).with_for_update()
        )
        variant_map = {v.id: v for v in vr.scalars().all()}

    for it in items:
        product = products[it.product_id]
        variant: Optional[ProductVariant] = variant_map.get(it.variant_id) if it.variant_id else None

        deduct_amount = it.stock_qty_override if it.stock_qty_override is not None else it.qty
        if variant:
            variant.stock_qty -= deduct_amount
            product.stock_qty -= deduct_amount
            if variant.stock_qty < product.low_stock_threshold:
                stock_low.send(
                    sender=Order,
                    product_id=product.id,
                    variant_id=variant.id,
                    current_qty=variant.stock_qty,
                    tenant_id=tenant_id,
                )
        else:
            product.stock_qty -= deduct_amount
            if product.stock_qty < product.low_stock_threshold:
                stock_low.send(
                    sender=Order,
                    product_id=product.id,
                    variant_id=None,
                    current_qty=product.stock_qty,
                    tenant_id=tenant_id,
                )

        # 销量随库存扣减同步自增（与 restore_stock 对称）。
        # ponytail: 按购买件数 it.qty 计（整数化），weight/按重量商品的小数份量会向下取整
        product.sales_count = (product.sales_count or 0) + int(it.qty)

        await _auto_switch_stock_status(db, product, product.stock_qty)


def get_deduct_statuses(ts_extra: dict | None) -> list[str]:
    """从 TenantSettings.extra 读取库存扣减状态列表，默认 ['pending']。"""
    if ts_extra:
        val = ts_extra.get("inventory_deduct_statuses")
        if isinstance(val, list):
            return val
    return ["pending"]


async def apply_inventory_transition(
    db: AsyncSession,
    tenant_id: int,
    items: list[dict],
    old_status: str | None,
    new_status: str,
    deduct_statuses: list[str],
    order_id: int = 0,
    idem_suffix: str = "",
) -> None:
    """
    根据订单状态变更自动处理库存。

    - old_status=None 表示新建订单（None → new_status）
    - old 不在扣库存集合 & new 在 → 扣库存
    - old 在扣库存集合 & new 不在 → 恢复库存
    - 两者都在或都不在 → 不动

    items 格式: [{"product_id": int, "variant_id": int|None, "qty": Decimal}]
    """
    if not items:
        return
    # 已接管租户：不用订单状态集合，改由两段式锁库/出库状态机驱动
    if await _takeover_state(db, tenant_id) is not None:
        from app.plugins.inventory.services import apply_transition_via_ledger
        return await apply_transition_via_ledger(
            db, tenant_id, items, old_status, new_status, order_id, idem_suffix=idem_suffix)
    old_in = old_status in deduct_statuses if old_status else False
    new_in = new_status in deduct_statuses
    if old_in == new_in:
        return
    if new_in:
        cart_items = [
            PricingCartItem(
                product_id=it["product_id"],
                variant_id=it.get("variant_id"),
                qty=it["qty"],
            )
            for it in items
        ]
        await deduct_stock(db, tenant_id, cart_items, order_id)
    else:
        await restore_stock(db, tenant_id, items, order_id=order_id, restock=True)


async def restore_stock(
    db: AsyncSession,
    tenant_id: int,
    items: list[dict],
    *,
    order_id: int = 0,
    restock: bool = False,
    idem_suffix: str = "",
) -> None:
    """
    释放已占用的库存（取消订单 / 超时回滚时调用）。

    items 格式: [{"product_id": int, "variant_id": int|None, "qty": Decimal}, ...]
    """
    if not items:
        return

    for it in items:
        if Decimal(str(it.get("qty", 0))) <= 0:
            raise HTTPException(status_code=400, detail="库存恢复数量必须大于 0")

    # 进销存接管：未发货取消释放 reserved；已发货退款/退货回补 sellable。
    if await _guarded_takeover(db, tenant_id) is not None:
        if order_id <= 0:
            raise HTTPException(status_code=400, detail="账本库存恢复必须提供业务单据号")
        if restock:
            from app.plugins.inventory.services import restore_via_ledger
            return await restore_via_ledger(db, tenant_id, items, order_id, suffix=idem_suffix)
        from app.plugins.inventory.services import release_via_ledger
        return await release_via_ledger(db, tenant_id, items, order_id, suffix=idem_suffix)

    product_ids = list({it["product_id"] for it in items})
    result = await db.execute(
        select(Product).where(
            Product.id.in_(product_ids),
            Product.tenant_id == tenant_id,
        ).with_for_update()
    )
    products = {p.id: p for p in result.scalars().all()}

    variant_map: dict[int, ProductVariant] = {}
    variant_pairs = [
        and_(ProductVariant.id == it["variant_id"], ProductVariant.product_id == it["product_id"])
        for it in items
        if it.get("variant_id")
    ]
    if variant_pairs:
        vr = await db.execute(
            select(ProductVariant).where(
                or_(*variant_pairs),
                ProductVariant.tenant_id == tenant_id,
            ).with_for_update()
        )
        variant_map = {v.id: v for v in vr.scalars().all()}

    for it in items:
        product = products.get(it["product_id"])
        if product is None:
            continue
        variant: Optional[ProductVariant] = variant_map.get(it["variant_id"]) if it.get("variant_id") else None
        qty = Decimal(str(it["qty"]))

        if variant:
            variant.stock_qty += qty
            product.stock_qty += qty
        else:
            product.stock_qty += qty

        # 销量随库存恢复同步减回（取消/超时/退款），下限 0 防止意外负数
        product.sales_count = max(0, (product.sales_count or 0) - int(qty))

        await _auto_switch_stock_status(db, product, product.stock_qty)
