#!/usr/bin/env python3
from __future__ import annotations

import argparse
import csv
import json
from datetime import datetime, timedelta
from decimal import Decimal
from pathlib import Path
from zoneinfo import ZoneInfo

import pymysql


ROOT = Path(__file__).resolve().parents[5]
DEFAULT_CREDS = ROOT / "config/local/zhct-db-credentials.json"
DEFAULT_OUTPUT_ROOT = ROOT / "work_wecom_customer_service/runtime/data_queries/all-project-orders"
TZ = ZoneInfo("Asia/Shanghai")


def resolve_date(value: str) -> str:
    today = datetime.now(TZ).date()
    if value == "today":
        return today.isoformat()
    if value == "yesterday":
        return (today - timedelta(days=1)).isoformat()
    datetime.strptime(value, "%Y-%m-%d")
    return value


def decimal_text(value: object) -> str:
    if isinstance(value, Decimal):
        return str(value)
    return str(value or 0)


def write_csv(path: Path, rows: list[dict[str, object]]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    fields: list[str] = []
    for row in rows:
        for key in row:
            if key not in fields:
                fields.append(key)
    with path.open("w", encoding="utf-8-sig", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=fields)
        writer.writeheader()
        for row in rows:
            writer.writerow({key: decimal_text(value) if isinstance(value, Decimal) else value for key, value in row.items()})


def choose_identity_column(cursor: pymysql.cursors.DictCursor) -> str | None:
    cursor.execute("SHOW COLUMNS FROM ydy_meal_order")
    columns = {row["Field"] for row in cursor.fetchall()}
    for column in ("user_id", "staff_uuid", "staff_id", "uid"):
        if column in columns:
            return column
    return None


def query_connection(conn_cfg: dict[str, object], day: str, include_dau: bool) -> tuple[dict[str, object], list[dict[str, object]]]:
    host = str(conn_cfg.get("external_host") or conn_cfg.get("host") or "")
    port = int(conn_cfg.get("port") or 3306)
    database = str(conn_cfg.get("database") or "")
    summary = {
        "配置目录": conn_cfg.get("config_dir", ""),
        "项目名称": conn_cfg.get("project_name", ""),
        "查询host": host,
        "数据库名称": database,
    }
    meal_rows: list[dict[str, object]] = []
    try:
        connection = pymysql.connect(
            host=host,
            port=port,
            user=str(conn_cfg.get("username") or ""),
            password=str(conn_cfg.get("password") or ""),
            database=database,
            charset="utf8mb4",
            connect_timeout=8,
            read_timeout=45,
            write_timeout=10,
            cursorclass=pymysql.cursors.DictCursor,
        )
        with connection.cursor() as cursor:
            cursor.execute(
                """
                SELECT COUNT(*) AS order_num,
                       COALESCE(SUM(total_price), 0) AS total_price
                FROM ydy_meal_order
                WHERE meal_date = %s
                """,
                (day,),
            )
            stats = cursor.fetchone() or {}
            summary.update(
                {
                    "查询状态": "成功",
                    "订单数": stats.get("order_num") or 0,
                    "流水(元)": stats.get("total_price") or Decimal("0"),
                }
            )

            if include_dau:
                identity_column = choose_identity_column(cursor)
                summary["日活口径"] = identity_column or ""
                if identity_column:
                    cursor.execute(
                        f"""
                        SELECT COUNT(DISTINCT NULLIF(CAST({identity_column} AS CHAR), '')) AS dau
                        FROM ydy_meal_order
                        WHERE meal_date = %s
                        """,
                        (day,),
                    )
                    summary["日活人数"] = (cursor.fetchone() or {}).get("dau") or 0
                else:
                    summary["日活人数"] = ""

            cursor.execute(
                """
                SELECT meal_times,
                       CASE meal_times
                         WHEN 1 THEN '早餐'
                         WHEN 2 THEN '午餐'
                         WHEN 3 THEN '晚餐'
                         WHEN 4 THEN '夜宵'
                         ELSE CONCAT('其他(', meal_times, ')')
                       END AS meal_segment,
                       COUNT(*) AS order_num,
                       COALESCE(SUM(total_price), 0) AS total_price
                FROM ydy_meal_order
                WHERE meal_date = %s
                GROUP BY meal_times
                ORDER BY meal_times
                """,
                (day,),
            )
            for row in cursor.fetchall():
                meal_rows.append(
                    {
                        "配置目录": conn_cfg.get("config_dir", ""),
                        "项目名称": conn_cfg.get("project_name", ""),
                        "日期": day,
                        "餐段编码": row.get("meal_times"),
                        "餐段": row.get("meal_segment"),
                        "订单数": row.get("order_num") or 0,
                        "流水(元)": row.get("total_price") or Decimal("0"),
                    }
                )
        connection.close()
    except Exception as exc:  # noqa: BLE001 - reporting connection/query failure class is intentional.
        summary.update(
            {
                "查询状态": "失败",
                "订单数": "",
                "流水(元)": "",
                "失败原因": f"{type(exc).__name__}: {str(exc)[:240]}",
            }
        )
    return summary, meal_rows


def main() -> int:
    parser = argparse.ArgumentParser(description="Query all ZHCT smart-canteen project order counts for a meal_date.")
    parser.add_argument("--date", default="today", help="today, yesterday, or YYYY-MM-DD. Uses Asia/Shanghai for relative dates.")
    parser.add_argument("--creds", default=str(DEFAULT_CREDS), help="Local ignored credential JSON.")
    parser.add_argument("--output-root", default=str(DEFAULT_OUTPUT_ROOT), help="Directory for runtime CSV outputs.")
    parser.add_argument("--include-ai", action="store_true", help="Also query role=ai auxiliary connections. Not for routine order counts.")
    parser.add_argument("--include-dau", action="store_true", help="Include distinct user/person count when a supported identifier column exists.")
    args = parser.parse_args()

    day = resolve_date(args.date)
    creds_path = Path(args.creds)
    if not creds_path.exists():
        raise SystemExit(f"Missing local credential file: {creds_path}")

    data = json.loads(creds_path.read_text(encoding="utf-8"))
    connections = data.get("connections", [])
    if not args.include_ai:
        connections = [item for item in connections if item.get("role") == "main"]

    output_dir = Path(args.output_root) / f"{day}-all-project-order-count"
    summary_rows: list[dict[str, object]] = []
    meal_rows: list[dict[str, object]] = []
    for conn_cfg in connections:
        summary, meals = query_connection(conn_cfg, day, args.include_dau)
        summary_rows.append(summary)
        meal_rows.extend(meals)

    summary_path = output_dir / f"订单汇总_{day}.csv"
    meal_path = output_dir / f"餐段订单_{day}.csv"
    write_csv(summary_path, summary_rows)
    write_csv(meal_path, meal_rows)

    successful = [row for row in summary_rows if row.get("查询状态") == "成功"]
    failed = [row for row in summary_rows if row.get("查询状态") != "成功"]
    total_orders = sum(Decimal(str(row.get("订单数") or 0)) for row in successful)
    total_amount = sum(Decimal(str(row.get("流水(元)") or 0)) for row in successful)

    print(f"日期: {day}")
    print(f"成功查询: {len(successful)} / {len(summary_rows)}; 失败: {len(failed)}")
    print(f"合计订单数: {int(total_orders)}")
    print(f"合计流水(元): {total_amount:.2f}")
    print("")
    if args.include_dau:
        print("| 项目 | 订单数 | 流水(元) | 日活人数 | 状态 |")
        print("|---|---:|---:|---:|---|")
    else:
        print("| 项目 | 订单数 | 流水(元) | 状态 |")
        print("|---|---:|---:|---|")
    sorted_rows = sorted(
        summary_rows,
        key=lambda row: Decimal(str(row.get("订单数") or 0)) if row.get("查询状态") == "成功" else Decimal("-1"),
        reverse=True,
    )
    for row in sorted_rows:
        amount = Decimal(str(row.get("流水(元)") or 0)) if row.get("查询状态") == "成功" else Decimal("0")
        if args.include_dau:
            print(f"| {row.get('项目名称')} | {row.get('订单数')} | {amount:.2f} | {row.get('日活人数', '')} | {row.get('查询状态')} |")
        else:
            print(f"| {row.get('项目名称')} | {row.get('订单数')} | {amount:.2f} | {row.get('查询状态')} |")

    if failed:
        print("\n失败项目:")
        for row in failed:
            print(f"- {row.get('配置目录')} {row.get('项目名称')}: {row.get('失败原因', '')}")

    print(f"\n输出目录: {output_dir}")
    print(f"汇总CSV: {summary_path}")
    print(f"餐段CSV: {meal_path}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
