#!/usr/bin/env python3
"""City-subcenter 460 device load test using only the Python standard library."""

from __future__ import annotations

import argparse
import concurrent.futures
import csv
import datetime as dt
import json
import math
import os
import socket
import statistics
import struct
import sys
import threading
import time
import urllib.error
import urllib.parse
import urllib.request
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Callable


@dataclass
class Result:
    sequence: int
    target: str
    success: bool
    latency_ms: float
    detail: str


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="City-subcenter device MQTT, heartbeat, and order load test."
    )
    parser.add_argument("--mode", choices=("mqtt", "heartbeat", "order"), required=True)
    parser.add_argument("--execute", action="store_true", help="Actually send requests.")
    parser.add_argument(
        "--confirm-order-load",
        action="store_true",
        help="Required for order mode because it creates real test orders.",
    )
    parser.add_argument("--base-url", default="http://172.26.154.109")
    parser.add_argument("--host-header", default="")
    parser.add_argument("--mqtt-host", default="172.26.154.109")
    parser.add_argument("--mqtt-port", type=int, default=1883)
    parser.add_argument("--mqtt-username", default=os.getenv("MQTT_USERNAME", ""))
    parser.add_argument("--mqtt-password", default=os.getenv("MQTT_PASSWORD", ""))
    parser.add_argument("--devices", type=Path)
    parser.add_argument("--device-prefix", default="CSFZX460")
    parser.add_argument("--device-start", type=int, default=1)
    parser.add_argument("--device-count", type=int, default=60)
    parser.add_argument("--device-width", type=int, default=3)
    parser.add_argument("--order-cases", type=Path)
    parser.add_argument("--workers", type=int, default=20)
    parser.add_argument("--rounds", type=int, default=1)
    parser.add_argument(
        "--duration",
        type=float,
        default=0,
        help="Spread requests evenly across seconds; 0 means burst.",
    )
    parser.add_argument("--timeout", type=float, default=10)
    parser.add_argument("--min-success-rate", type=float, default=99.5)
    parser.add_argument("--max-p99-ms", type=float, default=3000)
    parser.add_argument("--report-dir", type=Path)
    return parser.parse_args()


def validate_args(args: argparse.Namespace) -> None:
    if args.workers < 1 or args.rounds < 1 or args.device_count < 1:
        raise SystemExit("workers, rounds, and device-count must be positive.")
    if not 0 <= args.min_success_rate <= 100:
        raise SystemExit("min-success-rate must be between 0 and 100.")
    if args.max_p99_ms <= 0:
        raise SystemExit("max-p99-ms must be positive.")
    if args.mode == "order":
        if not args.order_cases:
            raise SystemExit("order mode requires --order-cases.")
        if args.rounds != 1:
            raise SystemExit("order mode requires --rounds 1 to avoid duplicate orders.")
        if args.execute and not args.confirm_order_load:
            raise SystemExit("order mode requires --execute and --confirm-order-load.")


def read_devices(args: argparse.Namespace) -> list[str]:
    if args.devices:
        devices = [
            line.strip()
            for line in args.devices.read_text(encoding="utf-8").splitlines()
            if line.strip() and not line.lstrip().startswith("#")
        ]
    else:
        devices = [
            f"{args.device_prefix}{number:0{args.device_width}d}"
            for number in range(
                args.device_start, args.device_start + args.device_count
            )
        ]
    if not devices:
        raise SystemExit("No device codes were loaded.")
    return devices


def read_order_cases(path: Path) -> list[dict[str, str]]:
    with path.open("r", encoding="utf-8-sig", newline="") as handle:
        rows = list(csv.DictReader(handle))
    required = {"staff_uuid", "equipment_code", "dishes_uuid", "weight"}
    if not rows or not required.issubset(rows[0]):
        raise SystemExit(
            "order case CSV must contain: staff_uuid,equipment_code,dishes_uuid,weight"
        )
    return rows


def percentile(values: list[float], percent: float) -> float:
    if not values:
        return 0.0
    ordered = sorted(values)
    rank = max(0, math.ceil(percent / 100 * len(ordered)) - 1)
    return ordered[rank]


def mqtt_string(value: str) -> bytes:
    encoded = value.encode("utf-8")
    return struct.pack("!H", len(encoded)) + encoded


def remaining_length(value: int) -> bytes:
    result = bytearray()
    while True:
        digit = value % 128
        value //= 128
        if value:
            digit |= 0x80
        result.append(digit)
        if not value:
            return bytes(result)


def mqtt_connect(
    sequence: int, device_code: str, args: argparse.Namespace
) -> Result:
    started = time.perf_counter()
    try:
        flags = 0x02
        payload = mqtt_string(f"load-{device_code}-{sequence}")
        if args.mqtt_username:
            flags |= 0x80
            payload += mqtt_string(args.mqtt_username)
        if args.mqtt_password:
            flags |= 0x40
            payload += mqtt_string(args.mqtt_password)
        variable_header = (
            mqtt_string("MQTT") + b"\x04" + bytes([flags]) + struct.pack("!H", 30)
        )
        body = variable_header + payload
        packet = b"\x10" + remaining_length(len(body)) + body
        with socket.create_connection(
            (args.mqtt_host, args.mqtt_port), timeout=args.timeout
        ) as mqtt_socket:
            mqtt_socket.sendall(packet)
            response = mqtt_socket.recv(4)
        success = (
            len(response) >= 4
            and response[0] == 0x20
            and response[1] == 0x02
            and response[3] == 0
        )
        detail = "CONNACK=0" if success else f"CONNACK={response.hex()}"
    except Exception as exc:
        success = False
        detail = f"{type(exc).__name__}: {exc}"
    return Result(
        sequence,
        device_code,
        success,
        (time.perf_counter() - started) * 1000,
        detail,
    )


def post_form(
    base_url: str,
    path: str,
    fields: dict[str, str],
    timeout: float,
    host_header: str = "",
) -> tuple[bool, str]:
    request = urllib.request.Request(
        f"{base_url.rstrip('/')}{path}",
        data=urllib.parse.urlencode(fields).encode("utf-8"),
        method="POST",
    )
    request.add_header("Content-Type", "application/x-www-form-urlencoded")
    if host_header:
        request.add_header("Host", host_header)
    opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
    try:
        with opener.open(request, timeout=timeout) as response:
            status = response.status
            body = response.read().decode("utf-8", errors="replace")
        payload = json.loads(body)
        api_code = payload.get("code", payload.get("status"))
        success = 200 <= status < 300 and api_code in (0, 200, "0", "200")
        return success, f"HTTP={status} API={api_code}"
    except urllib.error.HTTPError as exc:
        return False, f"HTTPError={exc.code}"
    except Exception as exc:
        return False, f"{type(exc).__name__}: {exc}"


def heartbeat(
    sequence: int, device_code: str, args: argparse.Namespace
) -> Result:
    started = time.perf_counter()
    success, detail = post_form(
        args.base_url,
        "/p/api/heartbeat",
        {"code": device_code},
        args.timeout,
        args.host_header,
    )
    return Result(
        sequence,
        device_code,
        success,
        (time.perf_counter() - started) * 1000,
        detail,
    )


def order_flow(
    sequence: int, case: dict[str, str], args: argparse.Namespace
) -> Result:
    started = time.perf_counter()
    plate_code = f"LT{dt.datetime.now():%m%d%H%M%S}{sequence:06d}"
    equipment_code = case["equipment_code"].strip()
    dishes_uuid = case["dishes_uuid"].strip()
    bind_success, bind_detail = post_form(
        args.base_url,
        "/p/api/bindPlateInfo",
        {
            "staff_uuid": case["staff_uuid"].strip(),
            "plate_code": plate_code,
            "equipment_code": equipment_code,
            "dishes_uuids": json.dumps([dishes_uuid], ensure_ascii=False),
        },
        args.timeout,
        args.host_header,
    )
    if not bind_success:
        return Result(
            sequence,
            equipment_code,
            False,
            (time.perf_counter() - started) * 1000,
            f"bind failed: {bind_detail}",
        )

    dish_info = json.dumps(
        {
            "create_date": dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
            "details": [
                {
                    "uuid": dishes_uuid,
                    "weight": float(case["weight"]),
                }
            ],
        },
        ensure_ascii=False,
    )
    meal_success, meal_detail = post_form(
        args.base_url,
        "/p/api/zhctPushMeal",
        {
            "dish_info": dish_info,
            "code": plate_code,
            "equipment_code": equipment_code,
        },
        args.timeout,
        args.host_header,
    )
    return Result(
        sequence,
        equipment_code,
        meal_success,
        (time.perf_counter() - started) * 1000,
        f"bind={bind_detail}; meal={meal_detail}",
    )


def build_tasks(
    args: argparse.Namespace,
) -> tuple[list[object], Callable[[int, object, argparse.Namespace], Result]]:
    if args.mode == "order":
        return read_order_cases(args.order_cases), order_flow
    devices = read_devices(args)
    tasks = [device for _ in range(args.rounds) for device in devices]
    return tasks, mqtt_connect if args.mode == "mqtt" else heartbeat


def run_tasks(
    tasks: list[object],
    worker: Callable[[int, object, argparse.Namespace], Result],
    args: argparse.Namespace,
) -> list[Result]:
    results: list[Result] = []
    started = time.monotonic()
    lock = threading.Lock()

    def scheduled(sequence: int, task: object) -> Result:
        if args.duration > 0 and len(tasks) > 1:
            target = started + args.duration * sequence / (len(tasks) - 1)
            delay = target - time.monotonic()
            if delay > 0:
                time.sleep(delay)
        result = worker(sequence, task, args)
        with lock:
            state = "PASS" if result.success else "FAIL"
            print(
                f"[{sequence + 1}/{len(tasks)}] {state} "
                f"{result.target} {result.latency_ms:.1f}ms {result.detail}",
                flush=True,
            )
        return result

    with concurrent.futures.ThreadPoolExecutor(max_workers=args.workers) as pool:
        futures = [
            pool.submit(scheduled, sequence, task)
            for sequence, task in enumerate(tasks)
        ]
        for future in concurrent.futures.as_completed(futures):
            results.append(future.result())
    return sorted(results, key=lambda item: item.sequence)


def save_report(
    args: argparse.Namespace, results: list[Result], elapsed: float
) -> dict[str, object]:
    timestamp = dt.datetime.now().strftime("%Y%m%d-%H%M%S")
    report_dir = args.report_dir or Path("reports") / f"device-load-{args.mode}-{timestamp}"
    report_dir.mkdir(parents=True, exist_ok=True)

    latencies = [result.latency_ms for result in results]
    success_count = sum(result.success for result in results)
    total = len(results)
    success_rate = 100 * success_count / total if total else 0.0
    p99_latency = percentile(latencies, 99)
    summary: dict[str, object] = {
        "mode": args.mode,
        "target": (
            f"{args.mqtt_host}:{args.mqtt_port}"
            if args.mode == "mqtt"
            else args.base_url
        ),
        "total": total,
        "success": success_count,
        "failure": total - success_count,
        "success_rate": round(success_rate, 3),
        "elapsed_seconds": round(elapsed, 3),
        "throughput_per_second": round(total / elapsed, 3) if elapsed else 0,
        "latency_ms": {
            "min": round(min(latencies), 3) if latencies else 0,
            "avg": round(statistics.fmean(latencies), 3) if latencies else 0,
            "p50": round(percentile(latencies, 50), 3),
            "p95": round(percentile(latencies, 95), 3),
            "p99": round(p99_latency, 3),
            "max": round(max(latencies), 3) if latencies else 0,
        },
        "thresholds": {
            "min_success_rate": args.min_success_rate,
            "max_p99_ms": args.max_p99_ms,
        },
        "passed": (
            success_rate >= args.min_success_rate
            and p99_latency <= args.max_p99_ms
        ),
    }

    (report_dir / "summary.json").write_text(
        json.dumps(summary, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
    )
    with (report_dir / "results.csv").open("w", encoding="utf-8", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=asdict(results[0]).keys())
        writer.writeheader()
        writer.writerows(asdict(result) for result in results)

    print(json.dumps(summary, ensure_ascii=False, indent=2))
    print(f"Report: {report_dir.resolve()}")
    return summary


def main() -> int:
    args = parse_args()
    validate_args(args)
    tasks, worker = build_tasks(args)

    print(
        json.dumps(
            {
                "mode": args.mode,
                "task_count": len(tasks),
                "workers": args.workers,
                "duration": args.duration,
                "execute": args.execute,
            },
            ensure_ascii=False,
        )
    )
    if not args.execute:
        print("DRY_RUN: add --execute to start the load test.")
        return 0

    started = time.perf_counter()
    results = run_tasks(tasks, worker, args)
    summary = save_report(args, results, time.perf_counter() - started)
    return 0 if summary["passed"] else 1


if __name__ == "__main__":
    sys.exit(main())
