"""管理后台促销 / 优惠券 API"""
from decimal import Decimal
from datetime import datetime
from typing import Optional, List

from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, Field
from sqlalchemy import select, func
from sqlalchemy.ext.asyncio import AsyncSession

from app.api.deps import get_db, get_admin_user
from app.core.models.discount import Discount
from app.core.models.discount_tier import DiscountTier
from app.core.models.discount_usage_log import DiscountUsageLog
from app.core.models.user import User

router = APIRouter(prefix="/coupons", tags=["促销管理"])

RULE_TYPES = {
    "order_amount_off", "order_percent_off",
    "product_amount_off", "product_buyxgety",
    "category_percent_off", "free_shipping",
    "points_multiplier",
    # 向后兼容
    "percentage", "fixed", "buy_x_get_y",
}


# ── Schemas ───────────────────────────────────────────────────

class TierIn(BaseModel):
    sort_order: int = 0
    min_amount: Optional[float] = None
    min_qty: Optional[int] = None
    discount_value: float


class TierOut(TierIn):
    id: int
    model_config = {"from_attributes": True}


class PromotionOut(BaseModel):
    id: int
    name: str
    code: Optional[str]
    type: str
    value: float
    min_order_amount: Optional[float]
    priority: int
    stop: bool
    apply_once: bool
    is_stackable: bool
    usage_limit: Optional[int]
    usage_limit_per_customer: Optional[int]
    used_count: int
    condition_min_qty: Optional[int]
    condition_order_count_min: Optional[int]
    condition_order_count_max: Optional[int]
    condition_customer_level_id: Optional[int]
    condition_product_ids: Optional[list]
    condition_category_ids: Optional[list]
    discount_product_ids: Optional[list]
    discount_category_ids: Optional[list]
    discount_qty: Optional[int]
    discount_qualifier: Optional[str]
    points_multiplier: Optional[float]
    start_at: Optional[datetime]
    end_at: Optional[datetime]
    is_active: bool
    tiers: List[TierOut] = []
    created_at: Optional[datetime]

    model_config = {"from_attributes": True}


class PromotionCreate(BaseModel):
    name: str = Field(..., max_length=100)
    code: Optional[str] = Field(None, max_length=50, description="为空则自动应用")
    type: str = Field(..., description="rule_type，见文档")
    value: float = Field(0, ge=0)
    min_order_amount: Optional[float] = None
    priority: int = 0
    stop: bool = False
    apply_once: bool = True
    is_stackable: bool = True
    usage_limit: Optional[int] = None
    usage_limit_per_customer: Optional[int] = None
    condition_min_qty: Optional[int] = None
    condition_order_count_min: Optional[int] = None
    condition_order_count_max: Optional[int] = None
    condition_customer_level_id: Optional[int] = None
    condition_product_ids: Optional[list] = None
    condition_category_ids: Optional[list] = None
    discount_product_ids: Optional[list] = None
    discount_category_ids: Optional[list] = None
    discount_qty: Optional[int] = None
    discount_qualifier: Optional[str] = None
    points_multiplier: Optional[float] = None
    start_at: Optional[datetime] = None
    end_at: Optional[datetime] = None
    is_active: bool = True
    tiers: List[TierIn] = []


class PromotionUpdate(BaseModel):
    name: Optional[str] = None
    type: Optional[str] = None
    value: Optional[float] = None
    min_order_amount: Optional[float] = None
    priority: Optional[int] = None
    stop: Optional[bool] = None
    apply_once: Optional[bool] = None
    is_stackable: Optional[bool] = None
    usage_limit: Optional[int] = None
    usage_limit_per_customer: Optional[int] = None
    condition_min_qty: Optional[int] = None
    condition_order_count_min: Optional[int] = None
    condition_order_count_max: Optional[int] = None
    condition_customer_level_id: Optional[int] = None
    condition_product_ids: Optional[list] = None
    condition_category_ids: Optional[list] = None
    discount_product_ids: Optional[list] = None
    discount_category_ids: Optional[list] = None
    discount_qty: Optional[int] = None
    discount_qualifier: Optional[str] = None
    points_multiplier: Optional[float] = None
    start_at: Optional[datetime] = None
    end_at: Optional[datetime] = None
    is_active: Optional[bool] = None
    tiers: Optional[List[TierIn]] = None


class SimulateIn(BaseModel):
    items: list = Field(..., description='[{"product_id": 1, "qty": 2, "unit_price": 99.0, "category_id": 5}]')
    coupon_code: Optional[str] = None


class SimulateOut(BaseModel):
    subtotal: float
    promotions_applied: list
    total_discount: float
    free_shipping: bool
    points_multiplier: float


class UsageLogOut(BaseModel):
    id: int
    discount_id: int
    customer_id: int
    order_id: Optional[int]
    created_at: datetime
    model_config = {"from_attributes": True}


# ── 辅助 ──────────────────────────────────────────────────────

async def _get_promo_or_404(db: AsyncSession, promo_id: int, tenant_id: int) -> Discount:
    result = await db.execute(
        select(Discount).where(Discount.id == promo_id, Discount.tenant_id == tenant_id)
    )
    promo = result.scalar_one_or_none()
    if not promo:
        raise HTTPException(status_code=404, detail="促销不存在")
    return promo


async def _load_tiers(db: AsyncSession, promo_id: int) -> list[DiscountTier]:
    result = await db.execute(
        select(DiscountTier).where(DiscountTier.discount_id == promo_id).order_by(DiscountTier.sort_order)
    )
    return result.scalars().all()


def _build_out(promo: Discount, tiers: list[DiscountTier]) -> PromotionOut:
    return PromotionOut(
        id=promo.id,
        name=promo.name,
        code=promo.code,
        type=promo.type,
        value=float(promo.value),
        min_order_amount=float(promo.min_order_amount) if promo.min_order_amount else None,
        priority=promo.priority,
        stop=bool(promo.stop),
        apply_once=bool(promo.apply_once),
        is_stackable=bool(promo.is_stackable),
        usage_limit=promo.usage_limit,
        usage_limit_per_customer=promo.usage_limit_per_customer,
        used_count=promo.used_count,
        condition_min_qty=promo.condition_min_qty,
        condition_order_count_min=promo.condition_order_count_min,
        condition_order_count_max=promo.condition_order_count_max,
        condition_customer_level_id=promo.condition_customer_level_id,
        condition_product_ids=promo.condition_product_ids,
        condition_category_ids=promo.condition_category_ids,
        discount_product_ids=promo.discount_product_ids,
        discount_category_ids=promo.discount_category_ids,
        discount_qty=promo.discount_qty,
        discount_qualifier=promo.discount_qualifier,
        points_multiplier=float(promo.points_multiplier) if promo.points_multiplier else None,
        start_at=promo.start_at,
        end_at=promo.end_at,
        is_active=bool(promo.is_active),
        tiers=[TierOut(id=t.id, sort_order=t.sort_order, min_amount=float(t.min_amount) if t.min_amount else None,
                       min_qty=t.min_qty, discount_value=float(t.discount_value)) for t in tiers],
        created_at=promo.created_at,
    )


# ── 路由 ──────────────────────────────────────────────────────

@router.get("", response_model=List[PromotionOut])
async def list_promotions(
    db: AsyncSession = Depends(get_db),
    admin_user: User = Depends(get_admin_user),
):
    result = await db.execute(
        select(Discount)
        .where(Discount.tenant_id == admin_user.tenant_id)
        .order_by(Discount.priority.desc(), Discount.id.desc())
    )
    promos = result.scalars().all()
    out = []
    for p in promos:
        tiers = await _load_tiers(db, p.id)
        out.append(_build_out(p, tiers))
    return out


@router.post("", response_model=PromotionOut, status_code=status.HTTP_201_CREATED)
async def create_promotion(
    body: PromotionCreate,
    db: AsyncSession = Depends(get_db),
    admin_user: User = Depends(get_admin_user),
):
    if body.type not in RULE_TYPES:
        raise HTTPException(status_code=400, detail=f"不支持的 rule_type: {body.type}")

    if body.code:
        existing = await db.execute(
            select(Discount).where(
                Discount.tenant_id == admin_user.tenant_id,
                Discount.code == body.code.upper(),
            )
        )
        if existing.scalar_one_or_none():
            raise HTTPException(status_code=400, detail="该券码已存在")

    promo = Discount(
        tenant_id=admin_user.tenant_id,
        name=body.name,
        code=body.code.upper() if body.code else None,
        type=body.type,
        value=Decimal(str(body.value)),
        min_order_amount=Decimal(str(body.min_order_amount)) if body.min_order_amount else None,
        priority=body.priority,
        stop=1 if body.stop else 0,
        apply_once=1 if body.apply_once else 0,
        is_stackable=1 if body.is_stackable else 0,
        usage_limit=body.usage_limit,
        usage_limit_per_customer=body.usage_limit_per_customer,
        condition_min_qty=body.condition_min_qty,
        condition_order_count_min=body.condition_order_count_min,
        condition_order_count_max=body.condition_order_count_max,
        condition_customer_level_id=body.condition_customer_level_id,
        condition_product_ids=body.condition_product_ids,
        condition_category_ids=body.condition_category_ids,
        discount_product_ids=body.discount_product_ids,
        discount_category_ids=body.discount_category_ids,
        discount_qty=body.discount_qty,
        discount_qualifier=body.discount_qualifier,
        points_multiplier=Decimal(str(body.points_multiplier)) if body.points_multiplier else None,
        start_at=body.start_at,
        end_at=body.end_at,
        is_active=1 if body.is_active else 0,
    )
    db.add(promo)
    await db.flush()

    tiers = []
    for t in body.tiers:
        tier = DiscountTier(
            tenant_id=admin_user.tenant_id,
            discount_id=promo.id,
            sort_order=t.sort_order,
            min_amount=Decimal(str(t.min_amount)) if t.min_amount else None,
            min_qty=t.min_qty,
            discount_value=Decimal(str(t.discount_value)),
        )
        db.add(tier)
        tiers.append(tier)

    await db.commit()
    await db.refresh(promo)
    return _build_out(promo, tiers)


@router.put("/{promo_id}", response_model=PromotionOut)
async def update_promotion(
    promo_id: int,
    body: PromotionUpdate,
    db: AsyncSession = Depends(get_db),
    admin_user: User = Depends(get_admin_user),
):
    promo = await _get_promo_or_404(db, promo_id, admin_user.tenant_id)

    scalar_fields = {
        "name", "type", "priority", "usage_limit", "usage_limit_per_customer",
        "condition_min_qty", "condition_order_count_min", "condition_order_count_max",
        "condition_customer_level_id", "condition_product_ids", "condition_category_ids",
        "discount_product_ids", "discount_category_ids", "discount_qty", "discount_qualifier",
        "start_at", "end_at",
    }
    decimal_fields = {"value", "min_order_amount", "points_multiplier"}
    bool_fields = {"stop", "apply_once", "is_stackable", "is_active"}

    for k, v in body.model_dump(exclude_unset=True, exclude={"tiers"}).items():
        if v is None:
            continue
        if k in decimal_fields:
            setattr(promo, k, Decimal(str(v)))
        elif k in bool_fields:
            setattr(promo, k, 1 if v else 0)
        elif k in scalar_fields:
            setattr(promo, k, v)

    if body.tiers is not None:
        await db.execute(
            DiscountTier.__table__.delete().where(DiscountTier.discount_id == promo_id)
        )
        tiers = []
        for t in body.tiers:
            tier = DiscountTier(
                tenant_id=admin_user.tenant_id,
                discount_id=promo_id,
                sort_order=t.sort_order,
                min_amount=Decimal(str(t.min_amount)) if t.min_amount else None,
                min_qty=t.min_qty,
                discount_value=Decimal(str(t.discount_value)),
            )
            db.add(tier)
            tiers.append(tier)
    else:
        tiers = await _load_tiers(db, promo_id)

    await db.commit()
    await db.refresh(promo)
    return _build_out(promo, tiers)


@router.delete("/{promo_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_promotion(
    promo_id: int,
    db: AsyncSession = Depends(get_db),
    admin_user: User = Depends(get_admin_user),
):
    promo = await _get_promo_or_404(db, promo_id, admin_user.tenant_id)
    await db.delete(promo)
    await db.commit()


@router.get("/{promo_id}/usage", response_model=List[UsageLogOut])
async def get_usage_logs(
    promo_id: int,
    limit: int = 50,
    offset: int = 0,
    db: AsyncSession = Depends(get_db),
    admin_user: User = Depends(get_admin_user),
):
    """查看某促销的使用明细"""
    await _get_promo_or_404(db, promo_id, admin_user.tenant_id)
    result = await db.execute(
        select(DiscountUsageLog)
        .where(DiscountUsageLog.discount_id == promo_id)
        .order_by(DiscountUsageLog.created_at.desc())
        .limit(limit).offset(offset)
    )
    return result.scalars().all()


@router.post("/simulate", response_model=SimulateOut)
async def simulate_promotion(
    body: SimulateIn,
    db: AsyncSession = Depends(get_db),
    admin_user: User = Depends(get_admin_user),
):
    """
    促销模拟接口 — 传入虚拟购物车，返回促销计算结果，不产生任何副作用。

    items 格式：
      [{"product_id": 1, "qty": 2, "unit_price": 99.0, "category_id": 5}]
    """
    from app.core.services.promotion_engine import CartLine, apply_promotions
    from datetime import timezone

    cart_lines: list[CartLine] = []
    subtotal = Decimal("0")
    for item in body.items:
        unit_price = Decimal(str(item.get("unit_price", 0)))
        qty = int(item.get("qty", 1))
        line_total = unit_price * qty
        subtotal += line_total
        cart_lines.append(CartLine(
            product_id=int(item["product_id"]),
            variant_id=item.get("variant_id"),
            qty=qty,
            unit_price=unit_price,
            line_total=line_total,
            category_id=item.get("category_id"),
        ))

    now = datetime.now(timezone.utc)
    code_filter = [Discount.code.is_(None)]
    if body.coupon_code:
        code_filter.append(Discount.code == body.coupon_code.upper())

    from sqlalchemy import or_
    result = await db.execute(
        select(Discount).where(
            Discount.tenant_id == admin_user.tenant_id,
            Discount.is_active == 1,
            or_(Discount.start_at.is_(None), Discount.start_at <= now),
            or_(Discount.end_at.is_(None), Discount.end_at >= now),
            or_(*code_filter),
        ).order_by(Discount.priority.desc())
    )
    promos = result.scalars().all()

    tier_result = await db.execute(
        select(DiscountTier).where(DiscountTier.discount_id.in_([p.id for p in promos]))
    )
    tiers_by_promo: dict[int, list] = {}
    for t in tier_result.scalars().all():
        tiers_by_promo.setdefault(t.discount_id, []).append(t)

    promo_results = apply_promotions(
        promos_with_tiers=[(p, tiers_by_promo.get(p.id, [])) for p in promos],
        customer=None,
        lines=cart_lines,
        subtotal=subtotal,
        usage_counts={},
    )

    total_discount = sum(r.amount for r in promo_results)
    free_shipping = any(r.free_shipping for r in promo_results)
    pts_mult = max((r.points_multiplier for r in promo_results), default=Decimal("1"))

    return SimulateOut(
        subtotal=float(subtotal),
        promotions_applied=[
            {"discount_id": r.discount_id, "label": r.label,
             "amount": float(r.amount), "free_shipping": r.free_shipping,
             "points_multiplier": float(r.points_multiplier)}
            for r in promo_results if r.applied
        ],
        total_discount=float(total_discount),
        free_shipping=free_shipping,
        points_multiplier=float(pts_mult),
    )
