"""数据看板路由"""
from datetime import datetime, timedelta, date
from decimal import Decimal
from typing import Optional
from fastapi import APIRouter, Depends, Query
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, func, and_, text, case

from app.api.deps import get_db, get_admin_user
from app.core.models.user import User
from app.core.models.order import Order, OrderItem
from app.core.models.customer import Customer
from app.core.models.product import Product, ProductVariant
from app.core.models.tenant_settings import TenantSettings
from app.core.models.discount import Discount
from app.core.models.review import ProductReview
from app.core.models.refund import RefundRequest
from app.schemas.dashboard import (
    DashboardStats, RevenueTrendItem, OrderStatusItem,
    TopProductItem, TopCustomerItem, CustomerGrowthItem,
    CategorySalesItem, InventoryAlertItem, CouponStatsItem,
    ReviewStatsResponse, ProductSalesItem, ProductSalesPage,
)

router = APIRouter(prefix="/dashboard", tags=["数据看板"])

PAID_STATUSES = ["paid", "shipped", "completed"]

def _effective_expiry_expr(product_date, variant_date):
    return case(
        (product_date.is_(None), variant_date),
        (variant_date.is_(None), product_date),
        else_=func.least(product_date, variant_date),
    )


async def _expiry_warning_days(db: AsyncSession, tenant_id: int) -> int:
    r = await db.execute(select(TenantSettings.extra).where(TenantSettings.tenant_id == tenant_id))
    extra = r.scalar_one_or_none() or {}
    try:
        return max(0, int(extra.get("expiry_warning_days", 30)))
    except (TypeError, ValueError):
        return 30

def _date_range(range_str: str) -> tuple[date, date]:
    """将 range 字符串转换为 (start_date, end_date)"""
    today = date.today()
    days = {"7d": 7, "30d": 30, "90d": 90}.get(range_str, 7)
    return today - timedelta(days=days - 1), today


def _pct_change(current: Decimal | int, previous: Decimal | int) -> float:
    if not previous:
        return 100.0 if current else 0.0
    return round(float((current - previous) / previous * 100), 1)


@router.get("/stats", response_model=DashboardStats, summary="核心统计指标（含环比）")
async def get_stats(
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    today = date.today()
    yesterday = today - timedelta(days=1)
    month_start = today.replace(day=1)
    last_month_end = month_start - timedelta(days=1)
    last_month_start = last_month_end.replace(day=1)

    async def _revenue(start, end):
        r = await db.execute(
            select(func.coalesce(func.sum(Order.grand_total), 0)).where(
                Order.tenant_id == tid,
                Order.status.in_(PAID_STATUSES),
                func.date(Order.created_at) >= start,
                func.date(Order.created_at) <= end,
            )
        )
        return r.scalar() or Decimal("0")

    async def _order_count(d):
        r = await db.execute(
            select(func.count()).where(
                Order.tenant_id == tid,
                func.date(Order.created_at) == d,
            )
        )
        return r.scalar() or 0

    revenue_today = await _revenue(today, today)
    revenue_yesterday = await _revenue(yesterday, yesterday)
    revenue_month = await _revenue(month_start, today)
    revenue_last_month = await _revenue(last_month_start, last_month_end)

    orders_today = await _order_count(today)
    orders_yesterday = await _order_count(yesterday)

    r = await db.execute(
        select(func.count()).where(Order.tenant_id == tid, Order.status == "pending")
    )
    orders_pending = r.scalar() or 0

    r = await db.execute(select(func.count()).where(Customer.tenant_id == tid))
    customers_total = r.scalar() or 0

    r = await db.execute(
        select(func.count()).where(
            Customer.tenant_id == tid,
            func.date(Customer.created_at) == today,
        )
    )
    customers_new_today = r.scalar() or 0

    r = await db.execute(
        select(func.count()).where(Product.tenant_id == tid, Product.status == "active")
    )
    products_active = r.scalar() or 0

    warning_days = await _expiry_warning_days(db, tid)
    cutoff = today + timedelta(days=warning_days)
    variant_expiry = (
        select(
            ProductVariant.product_id.label("product_id"),
            func.min(ProductVariant.expiry_date).label("variant_expiry_date"),
        )
        .where(ProductVariant.tenant_id == tid)
        .group_by(ProductVariant.product_id)
        .subquery()
    )
    effective_expiry = _effective_expiry_expr(Product.expiry_date, variant_expiry.c.variant_expiry_date)
    expiry_base = (
        select(Product.id.label("product_id"), effective_expiry.label("expiry_date"))
        .outerjoin(variant_expiry, variant_expiry.c.product_id == Product.id)
        .where(Product.tenant_id == tid, Product.status == "active")
        .subquery()
    )
    r = await db.execute(select(func.count()).select_from(expiry_base).where(expiry_base.c.expiry_date < today))
    products_expired = r.scalar() or 0
    r = await db.execute(
        select(func.count()).select_from(expiry_base).where(
            expiry_base.c.expiry_date >= today,
            expiry_base.c.expiry_date <= cutoff,
        )
    )
    products_expiring = r.scalar() or 0
    # ── 退款统计 ──────────────────────────────────────────────────
    r = await db.execute(
        select(func.count()).where(
            RefundRequest.tenant_id == tid,
            RefundRequest.status == "pending",
        )
    )
    refunds_pending = r.scalar() or 0

    r = await db.execute(
        select(func.coalesce(func.sum(RefundRequest.refund_amount), 0)).where(
            RefundRequest.tenant_id == tid,
            RefundRequest.status == "completed",
            func.date(RefundRequest.completed_at) >= month_start,
            func.date(RefundRequest.completed_at) <= today,
        )
    )
    refunds_month_amount = r.scalar() or Decimal("0")

    # ── 联系表单未读统计 ──────────────────────────────────────────
    contact_form_unread = 0
    try:
        from app.core.models import ContactFormSubmission
        r = await db.execute(
            select(func.count()).where(
                ContactFormSubmission.tenant_id == tid,
                ContactFormSubmission.status == "unread",
            )
        )
        contact_form_unread = r.scalar() or 0
    except Exception:
        pass

    return DashboardStats(
        revenue_today=revenue_today,
        revenue_month=revenue_month,
        orders_today=orders_today,
        orders_pending=orders_pending,
        customers_total=customers_total,
        products_active=products_active,
        products_expiring=products_expiring,
        products_expired=products_expired,
        revenue_today_change=_pct_change(revenue_today, revenue_yesterday),
        revenue_month_change=_pct_change(revenue_month, revenue_last_month),
        orders_today_change=_pct_change(orders_today, orders_yesterday),
        customers_new_today=customers_new_today,
        refunds_pending=refunds_pending,
        refunds_month_amount=refunds_month_amount,
        contact_form_unread=contact_form_unread,
    )


@router.get("/trend", response_model=list[RevenueTrendItem], summary="营收趋势（支持 7d/30d/90d）")
async def get_trend(
    range: str = Query("7d", pattern="^(7d|30d|90d)$"),
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    start_date, end_date = _date_range(range)
    result = []
    d = start_date
    while d <= end_date:
        r_rev = await db.execute(
            select(func.coalesce(func.sum(Order.grand_total), 0)).where(
                Order.tenant_id == tid,
                Order.status.in_(PAID_STATUSES),
                func.date(Order.created_at) == d,
            )
        )
        r_ord = await db.execute(
            select(func.count()).where(
                Order.tenant_id == tid,
                func.date(Order.created_at) == d,
            )
        )
        revenue = r_rev.scalar() or Decimal("0")
        orders = r_ord.scalar() or 0
        r_ref = await db.execute(
            select(func.coalesce(func.sum(RefundRequest.refund_amount), 0)).where(
                RefundRequest.tenant_id == tid,
                RefundRequest.status == "completed",
                func.date(RefundRequest.completed_at) == d,
            )
        )
        refund_amount = r_ref.scalar() or Decimal("0")
        label = d.strftime("%m/%d") if range == "7d" else d.strftime("%m/%d")
        result.append(RevenueTrendItem(date=label, revenue=revenue, orders=orders, refund_amount=refund_amount))
        d += timedelta(days=1)
    return result


@router.get("/order-status", response_model=list[OrderStatusItem], summary="订单状态分布")
async def get_order_status(
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    status_map = [
        ("pending",   "待付款", "#e6a23c"),
        ("paid",      "已付款", "#409eff"),
        ("shipped",   "已发货", "#9c27b0"),
        ("completed", "已完成", "#67c23a"),
        ("cancelled", "已取消", "#f56c6c"),
        ("refunding", "退款中", "#ff9800"),
        ("refunded",  "已退款", "#795548"),
    ]
    result = []
    for status, label, color in status_map:
        r = await db.execute(
            select(func.count()).where(Order.tenant_id == tid, Order.status == status)
        )
        count = r.scalar() or 0
        result.append(OrderStatusItem(name=label, value=count, color=color))
    return result


@router.get("/top-products", response_model=list[TopProductItem], summary="销售额 Top N 商品")
async def get_top_products(
    range: str = Query("30d", pattern="^(7d|30d|90d)$"),
    limit: int = Query(10, ge=1, le=200),
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    start_date, end_date = _date_range(range)

    rows = await db.execute(
        select(
            OrderItem.product_id,
            func.sum(OrderItem.total_price).label("revenue"),
            func.count(func.distinct(OrderItem.order_id)).label("orders_count"),
            func.sum(OrderItem.quantity).label("quantity"),
        )
        .join(Order, Order.id == OrderItem.order_id)
        .where(
            OrderItem.tenant_id == tid,
            Order.status.in_(PAID_STATUSES),
            func.date(Order.created_at) >= start_date,
            func.date(Order.created_at) <= end_date,
            OrderItem.product_id.isnot(None),
        )
        .group_by(OrderItem.product_id)
        .order_by(func.sum(OrderItem.total_price).desc())
        .limit(limit)
    )
    rows = rows.all()

    result = []
    for row in rows:
        p = await db.get(Product, row.product_id)
        if not p:
            continue
        result.append(TopProductItem(
            id=p.id,
            name=p.name,
            sku=p.sku,
            revenue=row.revenue or Decimal("0"),
            orders_count=row.orders_count or 0,
            quantity=row.quantity or 0,
        ))
    return result


@router.get("/top-customers", response_model=list[TopCustomerItem], summary="消费额 Top N 客户")
async def get_top_customers(
    range: str = Query("30d", pattern="^(7d|30d|90d)$"),
    limit: int = Query(10, ge=1, le=50),
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    start_date, end_date = _date_range(range)

    rows = await db.execute(
        select(
            Order.customer_id,
            func.sum(Order.grand_total).label("total_spent"),
            func.count(Order.id).label("orders_count"),
        )
        .where(
            Order.tenant_id == tid,
            Order.status.in_(PAID_STATUSES),
            func.date(Order.created_at) >= start_date,
            func.date(Order.created_at) <= end_date,
        )
        .group_by(Order.customer_id)
        .order_by(func.sum(Order.grand_total).desc())
        .limit(limit)
    )
    rows = rows.all()

    result = []
    for row in rows:
        c = await db.get(Customer, row.customer_id)
        if not c:
            continue
        result.append(TopCustomerItem(
            id=c.id,
            name=c.name,
            email=c.email,
            orders_count=row.orders_count or 0,
            total_spent=row.total_spent or Decimal("0"),
        ))
    return result


@router.get("/customer-growth", response_model=list[CustomerGrowthItem], summary="新客增长趋势")
async def get_customer_growth(
    range: str = Query("30d", pattern="^(7d|30d|90d)$"),
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    start_date, end_date = _date_range(range)

    # 获取截止开始日期前的累计客户数
    r = await db.execute(
        select(func.count()).where(
            Customer.tenant_id == tid,
            func.date(Customer.created_at) < start_date,
        )
    )
    base_count = r.scalar() or 0

    result = []
    cumulative = base_count
    d = start_date
    while d <= end_date:
        r = await db.execute(
            select(func.count()).where(
                Customer.tenant_id == tid,
                func.date(Customer.created_at) == d,
            )
        )
        new_count = r.scalar() or 0
        cumulative += new_count
        result.append(CustomerGrowthItem(
            date=d.strftime("%m/%d"),
            new_count=new_count,
            cumulative=cumulative,
        ))
        d += timedelta(days=1)
    return result


@router.get("/category-sales", response_model=list[CategorySalesItem], summary="品类销售占比")
async def get_category_sales(
    range: str = Query("30d", pattern="^(7d|30d|90d)$"),
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    start_date, end_date = _date_range(range)

    # 从 product_snapshot 中取分类名，或 join products 表
    rows = await db.execute(
        select(
            Product.category_id,
            func.sum(OrderItem.total_price).label("revenue"),
            func.count(func.distinct(OrderItem.order_id)).label("orders_count"),
        )
        .join(Order, Order.id == OrderItem.order_id)
        .join(Product, Product.id == OrderItem.product_id)
        .where(
            OrderItem.tenant_id == tid,
            Order.status.in_(PAID_STATUSES),
            func.date(Order.created_at) >= start_date,
            func.date(Order.created_at) <= end_date,
            OrderItem.product_id.isnot(None),
        )
        .group_by(Product.category_id)
        .order_by(func.sum(OrderItem.total_price).desc())
    )
    rows = rows.all()

    if not rows:
        return []

    # 获取分类名称（从 categories 表）
    from app.core.models.category import Category
    total_revenue = sum(r.revenue or Decimal("0") for r in rows)

    result = []
    for row in rows:
        cat_name = "未分类"
        if row.category_id:
            cat = await db.get(Category, row.category_id)
            if cat:
                cat_name = cat.name
        rev = row.revenue or Decimal("0")
        pct = round(float(rev / total_revenue * 100), 1) if total_revenue else 0.0
        result.append(CategorySalesItem(
            category=cat_name,
            revenue=rev,
            orders_count=row.orders_count or 0,
            percentage=pct,
        ))
    return result


@router.get("/inventory-alerts", response_model=list[InventoryAlertItem], summary="库存预警商品")
async def get_inventory_alerts(
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    rows = await db.execute(
        select(Product)
        .where(
            Product.tenant_id == tid,
            Product.status == "active",
            Product.stock_qty <= Product.low_stock_threshold,
        )
        .order_by(Product.stock_qty.asc())
        .limit(50)
    )
    products = rows.scalars().all()
    return [
        InventoryAlertItem(
            id=p.id, name=p.name, sku=p.sku,
            stock_qty=p.stock_qty,
            low_stock_threshold=p.low_stock_threshold,
            status="out_of_stock" if p.stock_qty == 0 else "low_stock",
        )
        for p in products
    ]


@router.get("/coupon-stats", response_model=list[CouponStatsItem], summary="优惠券使用统计")
async def get_coupon_stats(
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    rows = await db.execute(
        select(Discount)
        .where(Discount.tenant_id == tid)
        .order_by(Discount.used_count.desc())
        .limit(20)
    )
    discounts = rows.scalars().all()
    result = []
    for d in discounts:
        limit = d.usage_limit or 0
        used = d.used_count or 0
        rate = round(used / limit * 100, 1) if limit > 0 else 0.0
        result.append(CouponStatsItem(
            id=d.id, name=d.name, code=d.code, type=d.type,
            value=d.value, usage_limit=d.usage_limit,
            used_count=used, usage_rate=rate,
        ))
    return result


@router.get("/review-stats", response_model=ReviewStatsResponse, summary="评价数据统计")
async def get_review_stats(
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id

    r = await db.execute(
        select(
            func.avg(ProductReview.rating).label("avg_rating"),
            func.count().label("total"),
        ).where(
            ProductReview.tenant_id == tid,
            ProductReview.status == "approved",
        )
    )
    row = r.one()
    avg_rating = round(float(row.avg_rating or 0), 2)
    total = row.total or 0

    r_pending = await db.execute(
        select(func.count()).where(
            ProductReview.tenant_id == tid,
            ProductReview.status == "pending",
        )
    )
    pending = r_pending.scalar() or 0

    distribution = {}
    for star in range(1, 6):
        r_star = await db.execute(
            select(func.count()).where(
                ProductReview.tenant_id == tid,
                ProductReview.status == "approved",
                ProductReview.rating == star,
            )
        )
        distribution[str(star)] = r_star.scalar() or 0

    return ReviewStatsResponse(
        avg_rating=avg_rating,
        total=total,
        approved=total,
        pending=pending,
        distribution=distribution,
    )


@router.get("/product-sales", response_model=ProductSalesPage, summary="商品销售明细（搜索+分页）")
async def get_product_sales(
    range: str = Query("30d", pattern="^(7d|30d|90d)$"),
    keyword: str = Query("", description="搜索商品名称或SKU"),
    sort_by: str = Query("revenue", pattern="^(revenue|quantity|orders_count|avg_unit_price)$"),
    page: int = Query(1, ge=1),
    page_size: int = Query(20, ge=5, le=100),
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    tid = current_user.tenant_id
    start_date, end_date = _date_range(range)

    # 聚合所有在此时间段内有销售记录的商品
    sort_col = {
        "revenue": func.sum(OrderItem.total_price),
        "quantity": func.sum(OrderItem.quantity),
        "orders_count": func.count(func.distinct(OrderItem.order_id)),
        "avg_unit_price": func.avg(OrderItem.unit_price),
    }.get(sort_by, func.sum(OrderItem.total_price))

    base_q = (
        select(
            OrderItem.product_id,
            func.sum(OrderItem.total_price).label("revenue"),
            func.sum(OrderItem.quantity).label("quantity"),
            func.count(func.distinct(OrderItem.order_id)).label("orders_count"),
            func.avg(OrderItem.unit_price).label("avg_unit_price"),
        )
        .join(Order, Order.id == OrderItem.order_id)
        .join(Product, Product.id == OrderItem.product_id)
        .where(
            OrderItem.tenant_id == tid,
            Order.status.in_(PAID_STATUSES),
            func.date(Order.created_at) >= start_date,
            func.date(Order.created_at) <= end_date,
            OrderItem.product_id.isnot(None),
        )
        .group_by(OrderItem.product_id)
    )

    if keyword:
        kw = f"%{keyword}%"
        base_q = base_q.where(
            (Product.name.like(kw)) | (Product.sku.like(kw))
        )

    # 计算总数
    count_q = select(func.count()).select_from(base_q.subquery())
    total_r = await db.execute(count_q)
    total = total_r.scalar() or 0

    # 分页排序
    rows_q = base_q.order_by(sort_col.desc()).offset((page - 1) * page_size).limit(page_size)
    rows = (await db.execute(rows_q)).all()

    # 获取商品详情（分类名）
    from app.core.models.category import Category
    items = []
    for row in rows:
        p = await db.get(Product, row.product_id)
        if not p:
            continue
        cat_name = "未分类"
        if p.category_id:
            cat = await db.get(Category, p.category_id)
            if cat:
                cat_name = cat.name
        avg_price = row.avg_unit_price or Decimal("0")
        items.append(ProductSalesItem(
            id=p.id,
            name=p.name,
            sku=p.sku,
            category=cat_name,
            base_price=p.base_price,
            revenue=row.revenue or Decimal("0"),
            quantity=row.quantity or 0,
            orders_count=row.orders_count or 0,
            avg_unit_price=Decimal(str(round(float(avg_price), 2))),
            stock_qty=p.stock_qty,
        ))

    return ProductSalesPage(items=items, total=total, page=page, page_size=page_size)
