# backend/app/plugins/advanced_stats/router.py
from __future__ import annotations

from datetime import datetime
from decimal import Decimal
from typing import Any

from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession

from app.api.deps import (
    PermissionContext,
    get_db,
    get_permission_context,
    require_permission,
)
from app.core.models.brand import Brand
from app.core.models.category import Category
from app.core.models.discount import Discount
from app.core.models.member import MemberLevel
from app.core.models.order import Order
from app.core.models.payment import Payment
from app.core.models.product import Product
from app.core.models.shipping import ShippingMethod
from app.core.models.user import User
from app.core.services.plugin_helper import require_plugin
from app.plugins.advanced_stats import query as Q
from app.plugins.advanced_stats.derive import POINTS_PER_UNIT, derive_coupon_total
from app.plugins.advanced_stats.dimensions import (
    DIMENSIONS,
    ITEM_METRICS,
    ORDER_METRICS,
    WALLET_METRICS,
    get_dimension,
)
from app.plugins.advanced_stats.models import AdvReportPreset
from app.plugins.advanced_stats.schemas import PresetIn, PresetOut, ReportQuery, ReportResponse, ReportRow


async def _require_advanced_stats(db: AsyncSession = Depends(get_db)) -> None:
    """插件停用时返回 503，前端据此隐藏菜单"""
    await require_plugin("advanced_stats", db)


router = APIRouter(tags=["advanced_stats"], dependencies=[Depends(_require_advanced_stats)])


# ── 标签解析：把 id 换成人看得懂的名字 ──────────────────────

async def _label_map(db: AsyncSession, dimension: str, keys: list[Any], tenant_id: int) -> dict:
    """按维度批量取展示名。一条查询取回全部，不在循环里查库。"""
    ids = [k for k in keys if k is not None and str(k).isdigit()]
    if not ids:
        return {}
    model_col = {
        "member_level":    (MemberLevel, MemberLevel.name),
        "shipping_method": (ShippingMethod, ShippingMethod.name),
        "product":         (Product, Product.name),
        "category":        (Category, Category.name),
        "brand":           (Brand, Brand.name),
        "coupon":          (Discount, Discount.name),
    }.get(dimension)
    if not model_col:
        return {}
    model, name_col = model_col
    r = await db.execute(
        select(model.id, name_col).where(model.id.in_([int(i) for i in ids]))
    )
    return {str(row[0]): row[1] for row in r.all()}


WEEKDAY_I18N = {
    1: "advStats.weekday.sun", 2: "advStats.weekday.mon", 3: "advStats.weekday.tue",
    4: "advStats.weekday.wed", 5: "advStats.weekday.thu", 6: "advStats.weekday.fri",
    7: "advStats.weekday.sat",
}

POINTS_FLAG_I18N = {
    "used": "advStats.pointsFlag.used",
    "not_used": "advStats.pointsFlag.notUsed",
}

WALLET_TYPE_I18N = {
    "topup":        "advStats.walletType.topup",
    "order_pay":    "advStats.walletType.orderPay",
    "order_refund": "advStats.walletType.orderRefund",
    "admin_adjust": "advStats.walletType.adminAdjust",
}


def _render_label(dimension: str, dim_key: Any, names: dict) -> tuple[str, str | None]:
    """→ (展示文本, i18n key 或 None)

    i18n key 单独返回，不跟展示文本混在一个字段里靠前缀嗅探——
    支付网关叫 "advStats.foo" 这种边界情况就不会被误判成待翻译的 key。
    """
    if dim_key is None or dim_key == "":
        return "", "advStats.label.empty"          # 前端翻译成「未填写」
    k = str(dim_key)
    if dimension == "day_of_week":
        key = WEEKDAY_I18N.get(int(dim_key))
        return k, key
    if dimension == "hour":
        return f"{int(dim_key):02d}:00", None
    if dimension == "points":
        return k, POINTS_FLAG_I18N.get(k)
    if dimension == "wallet":
        return k, WALLET_TYPE_I18N.get(k)
    return names.get(k, k), None


def _to_num(v: Any) -> Any:
    return v if v is None else (float(v) if isinstance(v, Decimal) else v)


@router.get("/adv-stats/report", response_model=ReportResponse,
            summary="多维度统计报表")
async def get_report(
    q: ReportQuery = Depends(),
    db: AsyncSession = Depends(get_db),
    user: User = Depends(require_permission("adv_stats.view")),
):
    tid = user.tenant_id
    dim = get_dimension(q.dimension)
    d_start, d_end, _ = Q.resolve_effective_range(q)
    notes: list[str] = []

    # 券维度：主查询就是券聚合，不走通用路径
    if q.dimension == "coupon":
        stmt = Q.build_coupon_aggregate(q, tid)
        metric_keys = ["orders", "customers", "total", "coupon_amount", "aov", "roi"]
        notes.append("advStats.note.couponScope")   # 不含管理员建单与 POS 单
    elif q.dimension == "wallet":
        # 余额维度查钱包流水（含不挂订单的充值），不走订单聚合
        stmt = Q.build_wallet_dimension_aggregate(q, tid)
        metric_keys = list(WALLET_METRICS)
        notes.append("advStats.note.walletScope")   # 时间按流水时间，不受订单状态筛选影响
    elif dim.base == "order_item":
        stmt = Q.build_item_aggregate(q, tid)
        metric_keys = list(ITEM_METRICS)
    else:
        stmt = Q.build_order_aggregate(q, tid)
        metric_keys = list(ORDER_METRICS)

    rows_r = await db.execute(stmt)
    raw = rows_r.mappings().all()

    main: dict[str, dict[str, Any]] = {}
    for r in raw:
        key = "" if r["dim_key"] is None else str(r["dim_key"])
        main[key] = {k: v for k, v in r.items() if k != "dim_key"}

    # 订单级维度才需要副表指标；行级维度、券维度、余额维度都有自己的口径
    if dim.base == "order" and q.dimension not in ("coupon", "wallet"):
        side: dict[str, dict[str, Any]] = {}
        for name, builder in (
            ("refunds", Q.build_refund_aggregate),
            ("points", Q.build_points_aggregate),
            ("wallet_paid", Q.build_wallet_aggregate),
        ):
            sr = await db.execute(builder(q, tid))
            for row in sr.mappings().all():
                k = "" if row["dim_key"] is None else str(row["dim_key"])
                for col, val in row.items():
                    if col == "dim_key":
                        continue
                    side.setdefault(col, {})[k] = val
        # 商品件数来自行级聚合，同样是副表
        ir = await db.execute(Q.build_item_aggregate(q, tid))
        for row in ir.mappings().all():
            k = "" if row["dim_key"] is None else str(row["dim_key"])
            side.setdefault("items_qty", {})[k] = row["quantity"]

        main = Q.merge_side_metrics(
            main, side,
            zero_keys=("refunds", "points_earned", "points_used", "wallet_paid", "items_qty"),
        )
        # 派生列
        for row in main.values():
            row["aov"] = Q.compute_aov(row.get("total") or 0, row.get("orders") or 0)
            pts_used = int(row.get("points_used") or 0)
            row["points_amount"] = Decimal(pts_used) / POINTS_PER_UNIT
            # 订单级维度的券抵扣走倒推（拿不到单券粒度）；coupon 维度不到这里，
            # 它在上面的分支里直接 SUM(discount_usage_logs.discount_amount)。
            # 倒推逻辑只有 derive 一份，别在这儿抄第二遍。
            row["coupon_amount"] = derive_coupon_total(
                discount_total=Decimal(row.get("discount") or 0), points_used=pts_used
            )

    if q.dimension == "coupon":
        for row in main.values():
            row["aov"] = Q.compute_aov(row.get("total") or 0, row.get("orders") or 0)
            amt = Decimal(row.get("coupon_amount") or 0)
            row["roi"] = float(Decimal(row.get("total") or 0) / amt) if amt else None

    # 全量合计（不是当前页合计）
    totals: dict[str, Any] = {}
    for key in metric_keys:
        if key in ("aov", "roi", "avg_unit_price"):
            continue
        vals = [r.get(key) or 0 for r in main.values()]
        totals[key] = _to_num(sum(Decimal(str(v)) for v in vals) if vals else 0)
    if q.dimension == "wallet":
        totals["burn_rate"] = Q.compute_burn_rate(main)
    if "total" in totals and "orders" in totals:
        totals["aov"] = float(Q.compute_aov(
            Decimal(str(totals["total"])), int(totals["orders"] or 0)
        ))

    # 排序 + 分页（在内存里做：分组数天然有限，不是订单行数量级）
    items = list(main.items())
    reverse = q.order == "desc"
    if q.sort == "dim_key":
        items.sort(key=lambda kv: kv[0], reverse=reverse)
    else:
        items.sort(key=lambda kv: (kv[1].get(q.sort) or 0), reverse=reverse)
    total_rows = len(items)
    start = (q.page - 1) * q.page_size
    page_items = items[start:start + q.page_size]

    names = await _label_map(db, q.dimension, [k for k, _ in page_items], tid)
    rows = []
    for k, v in page_items:
        label, label_key = _render_label(q.dimension, k, names)
        rows.append(ReportRow(
            dim_key=k,
            dim_label=label,
            dim_label_key=label_key,
            metrics={mk: _to_num(v.get(mk)) for mk in metric_keys},
        ))

    return ReportResponse(
        dimension=q.dimension,
        granularity=q.granularity,
        date_start=d_start,
        date_end=d_end,
        tz_offset=q.tz_offset,
        statuses=Q.effective_statuses(q),
        metric_keys=metric_keys,
        rows=rows,
        totals=totals,
        total_rows=total_rows,
        page=q.page,
        page_size=q.page_size,
        notes=notes,
    )


@router.get("/adv-stats/report/options", summary="筛选器下拉选项")
async def get_options(
    db: AsyncSession = Depends(get_db),
    user: User = Depends(require_permission("adv_stats.view")),
):
    """一次返回所有下拉数据，前端进页面只发一个请求。"""
    tid = user.tenant_id

    levels = (await db.execute(
        select(MemberLevel.id, MemberLevel.name)
        .where(MemberLevel.tenant_id == tid).order_by(MemberLevel.rank)
    )).all()
    methods = (await db.execute(
        select(ShippingMethod.id, ShippingMethod.name)
        .where(ShippingMethod.tenant_id == tid)
    )).all()
    brands = (await db.execute(
        select(Brand.id, Brand.name).where(Brand.tenant_id == tid).order_by(Brand.name)
    )).all()
    categories = (await db.execute(
        select(Category.id, Category.name).where(Category.tenant_id == tid).order_by(Category.name)
    )).all()
    discounts = (await db.execute(
        select(Discount.id, Discount.name, Discount.code)
        .where(Discount.tenant_id == tid).order_by(Discount.id.desc())
    )).all()

    # 实际用过的支付网关、币种、订单状态，从历史订单里取，避免列一堆没用过的，
    # 也避免项目新增状态后这里漏掉
    gateways = (await db.execute(
        select(Payment.gateway).where(Payment.tenant_id == tid).distinct()
    )).scalars().all()
    currencies = (await db.execute(
        select(func.coalesce(Order.display_currency, Order.currency))
        .where(Order.tenant_id == tid).distinct()
    )).scalars().all()
    used_statuses = (await db.execute(
        select(Order.status).where(Order.tenant_id == tid).distinct()
    )).scalars().all()
    # 并上默认口径与未支付维度用到的状态：新租户一单没有时下拉也不能是空的
    statuses = sorted(
        {s for s in used_statuses if s} | set(Q.PAID_STATUSES) | {"pending", "cancelled"}
    )

    return {
        "dimensions": [
            {"key": k, "label_key": d.label_key, "base": d.base}
            for k, d in DIMENSIONS.items()
        ],
        "member_levels":    [{"id": i, "name": n} for i, n in levels],
        "shipping_methods": [{"id": i, "name": n} for i, n in methods],
        "brands":           [{"id": i, "name": n} for i, n in brands],
        "categories":       [{"id": i, "name": n} for i, n in categories],
        "discounts":        [{"id": i, "name": n, "code": c} for i, n, c in discounts],
        "payment_gateways": [g for g in gateways if g],
        "currencies":       [c for c in currencies if c],
        "statuses":         statuses,
    }


@router.get("/adv-stats/presets", response_model=list[PresetOut], summary="预设报表列表")
async def list_presets(
    db: AsyncSession = Depends(get_db),
    user: User = Depends(require_permission("adv_stats.view")),
):
    r = await db.execute(
        select(AdvReportPreset)
        .where(AdvReportPreset.tenant_id == user.tenant_id,
               AdvReportPreset.user_id == user.id)
        .order_by(AdvReportPreset.id.desc())
    )
    return r.scalars().all()


@router.post("/adv-stats/presets", response_model=PresetOut, status_code=201,
             summary="保存预设报表")
async def create_preset(
    body: PresetIn,
    db: AsyncSession = Depends(get_db),
    user: User = Depends(require_permission("adv_stats.view")),
):
    # 校验存进去的确实是一份合法查询，避免加载时才炸
    try:
        ReportQuery(**body.query)
    except Exception as e:
        raise HTTPException(status_code=422, detail=f"预设内容不是合法的查询参数: {e}")

    # 不做前置查重：查完再插中间有窗口，两个并发请求会双双通过检查然后撞唯一键，
    # IntegrityError 冒成 500。让 DB 的唯一约束做唯一裁判，捕获后转 409。
    p = AdvReportPreset(
        tenant_id=user.tenant_id, user_id=user.id,
        title=body.title, query=body.query,
    )
    db.add(p)
    try:
        await db.commit()
    except IntegrityError:
        await db.rollback()
        raise HTTPException(status_code=409, detail="同名预设已存在")
    await db.refresh(p)
    return p


@router.delete("/adv-stats/presets/{preset_id}", status_code=204, summary="删除预设报表")
async def delete_preset(
    preset_id: int,
    db: AsyncSession = Depends(get_db),
    user: User = Depends(require_permission("adv_stats.view")),
):
    r = await db.execute(
        select(AdvReportPreset).where(
            AdvReportPreset.id == preset_id,
            AdvReportPreset.tenant_id == user.tenant_id,
            AdvReportPreset.user_id == user.id,
        )
    )
    p = r.scalar_one_or_none()
    if not p:
        raise HTTPException(status_code=404, detail="预设不存在")
    await db.delete(p)
    await db.commit()


@router.get("/adv-stats/report/details", summary="明细下钻（含客户信息）")
async def get_details(
    q: ReportQuery = Depends(),
    db: AsyncSession = Depends(get_db),
    user: User = Depends(require_permission("adv_stats.detail")),
):
    """明细层输出客户 PII，权限在依赖里硬拦——前端隐藏不算数。"""
    from app.plugins.advanced_stats import details as D

    if q.details not in ("order", "item"):
        raise HTTPException(status_code=400, detail="details 必须是 order 或 item")

    stmt = D.build_order_details(q, user.tenant_id) if q.details == "order" \
        else D.build_item_details(q, user.tenant_id)

    count_r = await db.execute(select(func.count()).select_from(stmt.subquery()))
    total = count_r.scalar() or 0

    stmt = stmt.limit(q.page_size).offset((q.page - 1) * q.page_size)
    rows = (await db.execute(stmt)).mappings().all()

    return {
        "rows": [dict(r) for r in rows],
        "total": total,
        "page": q.page,
        "page_size": q.page_size,
    }


@router.get("/adv-stats/report/export", summary="导出报表")
async def export_report(
    q: ReportQuery = Depends(),
    fmt: str = Query("xlsx", pattern="^(xlsx|csv)$"),
    columns: str | None = Query(None, description="逗号分隔的指标列，缺省导出全部"),
    db: AsyncSession = Depends(get_db),
    user: User = Depends(require_permission("adv_stats.export")),
    ctx: PermissionContext = Depends(get_permission_context),
):
    """导出当前筛选条件下的数据，**最多 1000 行**（不是全量）。

    1000 是 ReportQuery.page_size 的上限。要真正无上限导出，CSV 得改流式追加、
    xlsx 得用 openpyxl write_only 模式，否则大结果集会把内存打爆。
    # ponytail: 本计划规模下 1000 个分组足够（按日看一年才 365 行）。命中上限时
    # 前端提示用户缩小时间范围或换粗粒度，比上流式导出少写几十行。
    """
    # 导出权限校验：必须同时有 view（前端隐藏不算数）。get_report 是 in-process
    # 调用，Depends() 默认在直接调用时跳，所以 get_report 自己那个 view 校验
    # 这里拿不到——必须在这里手动再校一遍。
    if not ctx.can("adv_stats.view"):
        raise HTTPException(status_code=403, detail="Permission denied: adv_stats.view")

    # 明细导出额外要求 detail 权限：导出 + 明细 = 全量客户名单
    if q.details != "none" and not ctx.can("adv_stats.detail"):
        raise HTTPException(status_code=403, detail="Permission denied: adv_stats.detail")

    from fastapi.responses import Response

    from app.plugins.advanced_stats.export import rows_to_csv, rows_to_xlsx

    # 导出取全量：把 page_size 放到上限，page 归 1
    q_all = q.model_copy(update={"page": 1, "page_size": 1000})
    resp = await get_report(q=q_all, db=db, user=user)

    wanted = [c for c in (columns.split(",") if columns else resp.metric_keys)
              if c in resp.metric_keys]          # 过滤非法列名，不信任前端传值
    headers = ["dim_label"] + wanted
    rows = [[r.dim_label] + [r.metrics.get(k) for k in wanted] for r in resp.rows]
    rows.append(["TOTAL"] + [resp.totals.get(k) for k in wanted])

    if fmt == "csv":
        data = rows_to_csv(headers, rows)
        media = "text/csv"
        ext = "csv"
    else:
        data = rows_to_xlsx(headers, rows, sheet_title=q.dimension)
        media = "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"
        ext = "xlsx"

    # 带查询时刻：all_time 报表的 date_end 恒为今天，只用它会导出一堆同名文件
    stamp = datetime.now().strftime("%Y%m%d-%H%M%S")
    filename = f"report-{q.dimension}-{stamp}.{ext}"
    return Response(
        content=data, media_type=media,
        headers={"Content-Disposition": f'attachment; filename="{filename}"'},
    )
