# backend/app/plugins/advanced_stats/dimensions.py
"""维度白名单。

安全边界：dimension / granularity / sort 三个参数永远不进字符串拼接，
只走本模块的字典查表。查不到 → KeyError → router 转 400。
筛选值一律走 SQLAlchemy 参数绑定。
"""
from __future__ import annotations

from dataclasses import dataclass
from typing import Any, Callable

from sqlalchemy import case, exists, func, select, text

from app.core.models.customer import Customer
from app.core.models.member import MemberPointsLedger
from app.core.models.order import Order, OrderItem
from app.core.models.payment import Payment
from app.core.models.product import Product


def shift(col, tz_offset_minutes: int):
    """把时间列平移到用户时区，仅用于 GROUP BY 切天。

    tz_offset 已由 timerange.validate_tz_offset 限制在 [-840, 840]，
    此处仍走 bind 参数而非字符串格式化——不给注入留任何缝。
    同一条查询内 tz_off 取值恒定，重名 bind 不会冲突。
    """
    return func.date_add(
        col, text("INTERVAL :tz_off MINUTE").bindparams(tz_off=tz_offset_minutes)
    )


@dataclass(frozen=True)
class DimCtx:
    """维度表达式的上下文：时区、日期基准列、时间粒度"""
    tz: int
    date_col: Any          # Order.created_at 或 Order.paid_at
    granularity: str       # 仅 summary 维度用得上


@dataclass(frozen=True)
class Dimension:
    key: str
    label_key: str                       # i18n key，前端负责翻译
    base: str                            # 'order' | 'order_item'
    expr: Callable[[DimCtx], Any]        # 返回 GROUP BY 表达式
    joins: tuple[str, ...] = ()          # 需要的 JOIN，见 query.JOINS
    default_statuses: tuple[str, ...] | None = None   # 覆盖默认状态过滤
    label_from: str | None = None        # 额外 SELECT 的展示名列（如 product.name）


# ── 时间粒度（仅 summary 维度）───────────────────────────────

GRANULARITIES: dict[str, Callable[[DimCtx], Any]] = {
    "year":    lambda c: func.year(shift(c.date_col, c.tz)),
    "quarter": lambda c: func.concat(
        func.year(shift(c.date_col, c.tz)), "-Q", func.quarter(shift(c.date_col, c.tz))
    ),
    "month":   lambda c: func.date_format(shift(c.date_col, c.tz), "%Y-%m"),
    "week":    lambda c: func.yearweek(shift(c.date_col, c.tz), 3),   # 3 = ISO，周一为首日
    "day":     lambda c: func.date(shift(c.date_col, c.tz)),
    "order":   lambda c: Order.id,
}


def get_granularity(key: str) -> Callable[[DimCtx], Any]:
    return GRANULARITIES[key]


# ── 地址 JSON 提取 ───────────────────────────────────────────
# shipping_address 的 key 与 customer_addresses 表一致：
# country / province / city / district / zip_code
#
# ponytail: JSON 直接提取，不加生成列。报表必带日期范围，WHERE 已用
# ix_orders_tenant_status_date 把行数砍下来，JSON 只在已过滤结果集上算。
# 单次查询命中 10w+ 行时再换生成列 + 索引，只需改这里这几行。

def _addr(field: str):
    return func.trim(Order.shipping_address[field].as_string())


def points_used_flag():
    """把订单切成「用了积分」和「没用积分」两组。

    这是个真实的分析切面——两组的客单价、件数、退款率对比才是运营想看的东西，
    而不是又来一份按时间的汇总（那个 summary 维度已经有了）。

    # ponytail: 相关子查询 EXISTS，报表已被日期范围限流，够用。真慢了就在
    # orders 上加一个下单时写入的 used_points 布尔列，改这一个函数即可。
    """
    used = exists(
        select(1).where(
            MemberPointsLedger.order_id == Order.id,
            MemberPointsLedger.change_amount < 0,
        )
    )
    return case((used, "used"), else_="not_used")


# ── 维度注册表 ───────────────────────────────────────────────

DIMENSIONS: dict[str, Dimension] = {
    "summary": Dimension(
        key="summary", label_key="advStats.dim.summary", base="order",
        expr=lambda c: get_granularity(c.granularity)(c),
    ),
    "day_of_week": Dimension(
        key="day_of_week", label_key="advStats.dim.dayOfWeek", base="order",
        expr=lambda c: func.dayofweek(shift(c.date_col, c.tz)),   # 1=周日 … 7=周六
    ),
    "hour": Dimension(
        key="hour", label_key="advStats.dim.hour", base="order",
        expr=lambda c: func.hour(shift(c.date_col, c.tz)),
    ),
    "member_level": Dimension(
        key="member_level", label_key="advStats.dim.memberLevel", base="order",
        expr=lambda c: Customer.member_level_id,
        joins=("customer",), label_from="member_level_name",
    ),
    "payment_method": Dimension(
        key="payment_method", label_key="advStats.dim.paymentMethod", base="order",
        expr=lambda c: Payment.gateway, joins=("payment",),
    ),
    "shipping_method": Dimension(
        key="shipping_method", label_key="advStats.dim.shippingMethod", base="order",
        expr=lambda c: Order.shipping_method_id, label_from="shipping_method_name",
    ),
    "currency": Dimension(
        key="currency", label_key="advStats.dim.currency", base="order",
        expr=lambda c: func.coalesce(Order.display_currency, Order.currency),
    ),
    "tax": Dimension(
        key="tax", label_key="advStats.dim.tax", base="order_item",
        expr=lambda c: OrderItem.tax_rate,
    ),
    "unpaid": Dimension(
        key="unpaid", label_key="advStats.dim.unpaid", base="order",
        expr=lambda c: Order.status,
        default_statuses=("pending", "cancelled"),
    ),
    # ── 地理五级 ──
    "country": Dimension(
        key="country", label_key="advStats.dim.country", base="order",
        expr=lambda c: _addr("country"),
    ),
    "province": Dimension(
        key="province", label_key="advStats.dim.province", base="order",
        expr=lambda c: _addr("province"),
    ),
    "city": Dimension(
        key="city", label_key="advStats.dim.city", base="order",
        expr=lambda c: _addr("city"),
    ),
    "district": Dimension(
        key="district", label_key="advStats.dim.district", base="order",
        expr=lambda c: _addr("district"),
    ),
    "postcode": Dimension(
        key="postcode", label_key="advStats.dim.postcode", base="order",
        expr=lambda c: _addr("zip_code"),
    ),
    # ── 商品维度（行级）──
    "product": Dimension(
        key="product", label_key="advStats.dim.product", base="order_item",
        expr=lambda c: OrderItem.product_id,
        joins=("product",), label_from="product_name",
    ),
    "category": Dimension(
        key="category", label_key="advStats.dim.category", base="order_item",
        expr=lambda c: Product.category_id,
        joins=("product",), label_from="category_name",
    ),
    "brand": Dimension(
        key="brand", label_key="advStats.dim.brand", base="order_item",
        expr=lambda c: Product.brand_id,
        joins=("product",), label_from="brand_name",
    ),
    # ── 营销/资金维度（副表，见 query.py 的 side query）──
    "coupon": Dimension(
        key="coupon", label_key="advStats.dim.coupon", base="order",
        expr=lambda c: text("adv_coupon_dim"),   # 由 query.py 的 coupon 专用分支处理
    ),
    "points": Dimension(
        key="points", label_key="advStats.dim.points", base="order",
        expr=lambda c: points_used_flag(),
    ),
    "wallet": Dimension(
        key="wallet", label_key="advStats.dim.wallet", base="order",
        expr=lambda c: text("adv_wallet_dim"),   # 由 query.py 的 wallet 专用分支处理
    ),
}


def get_dimension(key: str) -> Dimension:
    return DIMENSIONS[key]


# ── 指标集 ───────────────────────────────────────────────────
# ORDER_METRICS：订单级，base='order' 的维度输出
# ITEM_METRICS：行级，base='order_item' 的维度输出
#   行级不含 shipping / aov —— 运费和客单价属于整单，分不到某个分类上

ORDER_METRICS: tuple[str, ...] = (
    "orders", "customers", "items_qty",
    "subtotal", "discount", "coupon_amount", "points_amount",
    "shipping", "tax", "total", "aov",
    "refunds", "points_earned", "points_used", "wallet_paid",
)

ITEM_METRICS: tuple[str, ...] = (
    "quantity", "orders", "revenue", "tax", "avg_unit_price", "refund_qty",
)

# 余额维度专用：它查的是钱包流水不是订单，指标自成一套
WALLET_METRICS: tuple[str, ...] = ("wallet_txns", "wallet_customers", "wallet_amount")

SORTABLE: tuple[str, ...] = ("dim_key",) + ORDER_METRICS + ITEM_METRICS + WALLET_METRICS


def get_sort_key(key: str) -> str:
    if key not in SORTABLE:
        raise KeyError(f"不可排序的字段: {key!r}")
    return key