"""
High-Concurrency Stress Testing & Benchmark Suite for FreeExile Chat Architecture.
Simulates partitioned distributed loads representing 1,000,000 CCU across all 8 channels:
World, Zone, Guild, Party, Whisper, System, Recruit, and Feedback.
Measures p50, p95, p99 fan-out latency (< 15ms), HMAC query SLA (< 2ms),
and multi-round heap regression (residual growth <= 0.05MB).
"""

from __future__ import annotations

import argparse
import asyncio
from dataclasses import dataclass
import gc
import os
import sys
import time
import tracemalloc
from typing import Callable, Dict, List, Optional, Sequence, Tuple

sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))

from server.chat.chat_service import ChatService
from server.chat.chat_types import (
    ChatChannelType,
    ChatMessageDTO,
    ItemAffixDTO,
    SendChatRequestDTO,
)


@dataclass(slots=True, frozen=True)
class BenchmarkConfig:
    simulated_ccu: int = 1_000_000
    active_sample_subscribers: int = 5_000
    message_count: int = 1_000
    concurrency: int = 10
    leak_check_rounds: int = 3
    cluster_shards: int = 64
    item_query_count: int = 200


@dataclass(slots=True, frozen=True)
class RoundMetrics:
    round_index: int
    messages_accepted: int
    deliveries_count: int
    duration_sec: float
    throughput_msg_per_sec: float
    p50_ms: float
    p95_ms: float
    p99_ms: float
    hmac_queries_passed: int
    hmac_avg_ms: float
    hmac_p99_ms: float
    start_heap_mb: float
    peak_heap_mb: float
    end_heap_mb: float
    net_growth_mb: float


@dataclass(slots=True, frozen=True)
class BenchmarkSummary:
    config: BenchmarkConfig
    round_metrics: tuple[RoundMetrics, ...]
    sla_latency_passed: bool
    sla_hmac_passed: bool
    sla_leak_passed: bool
    residual_growth_mb: float
    overall_passed: bool


def calculate_percentiles(latencies: Sequence[float]) -> Dict[str, float]:
    """Calculates p50, p95, p99, and average latency in ms."""
    if not latencies:
        return {"p50": 0.0, "p95": 0.0, "p99": 0.0, "avg": 0.0}
    vals = sorted(latencies)
    n = len(vals)
    return {
        "p50": round(vals[int(n * 0.50)], 3),
        "p95": round(vals[min(int(n * 0.95), n - 1)], 3),
        "p99": round(vals[min(int(n * 0.99), n - 1)], 3),
        "avg": round(sum(vals) / n, 3),
    }


def provision_test_items(chat_service: ChatService, count: int = 10) -> tuple[tuple[str, str], ...]:
    """Pre-generates signed items for concurrent HMAC inspection queries."""
    items: List[Tuple[str, str]] = []
    for i in range(count):
        uid = f"bench_item_blade_{i}"
        snap = chat_service.item_link_service.create_item_snapshot(
            item_uuid=uid, item_name_key=f"Huyền Vũ Kiếm Bậc {i + 1}", rarity=5,
            element=(i % 5) + 1, quality=20, item_level=90 + i, crafter_name=f"Đại Sư #{i + 1}",
            affixes=(ItemAffixDTO("aff_stat", "Huyết Lực", 500.0 * (i + 1)),),
        )
        items.append((uid, snap.hmac_signature))
    return tuple(items)


def provision_subscribers(chat_service: ChatService, count: int, sink: Callable[[ChatMessageDTO], None]) -> int:
    """Registers listeners across World, System, Recruit, Zones, Guilds, and Parties."""
    for cid in range(1, count + 1):
        for ch in ("world", "system", "recruit", "zone:zone_main", "guild:guild_main"):
            chat_service.subscribe_client(ch, cid, sink)
        if cid <= 1000:
            chat_service.subscribe_client("party:party_main", cid, sink)
        if cid == 1:
            chat_service.subscribe_client("whisper:100001:200001", cid, sink)
            chat_service.subscribe_client("feedback:100001", cid, sink)
    return chat_service.cluster_router.registry.total_subscribers_count()


def generate_round_requests(round_idx: int, total_count: int, pool_size: int = 5000) -> List[SendChatRequestDTO]:
    """Generates distributed chat requests cycling through all 8 channels across simulated users."""
    channels = (
        ChatChannelType.WORLD, ChatChannelType.ZONE, ChatChannelType.GUILD,
        ChatChannelType.PARTY, ChatChannelType.WHISPER, ChatChannelType.SYSTEM,
        ChatChannelType.RECRUIT, ChatChannelType.FEEDBACK,
    )
    reqs: List[SendChatRequestDTO] = []
    for i in range(total_count):
        ch = channels[i % len(channels)]
        # Distribute senders across the active subscriber pool without collisions within a round
        safe_pool = max(total_count, pool_size, 1000)
        sid = 0 if ch == ChatChannelType.SYSTEM else 100_001 + (i % safe_pool)
        sname = "Hệ Thống" if ch == ChatChannelType.SYSTEM else f"Exile_{sid}"
        tid = 200_001 if ch == ChatChannelType.WHISPER else 0
        reqs.append(SendChatRequestDTO(
            sender_id=sid, sender_name=sname, sender_level=35, channel=ch,
            content=f"Benchmark msg r{round_idx}_seq{i} on {ch.name} by {sid}",
            zone_id="zone_main", guild_id="guild_main", party_id="party_main",
            target_id=tid,
        ))
    return reqs


async def _publisher_worker(chat_service: ChatService, batch: Sequence[SendChatRequestDTO]) -> Tuple[int, int]:
    """Asynchronous worker executing a subset of chat publish requests with cooperative yields."""
    accepted, delivered = 0, 0
    for req in batch:
        ch_key = chat_service.channel_manager.get_channel_key(
            channel=req.channel, zone_id=req.zone_id, guild_id=req.guild_id,
            party_id=req.party_id, player_a=req.sender_id, player_b=req.target_id,
        )
        deliv_count = len(chat_service.cluster_router.registry.get_callbacks(ch_key))
        res = await chat_service.handle_send_chat(req)
        if res.success:
            accepted += 1
            delivered += deliv_count
        await asyncio.sleep(0)
    return accepted, delivered


async def _hmac_query_worker(chat_service: ChatService, queries: Sequence[Tuple[str, str]]) -> Tuple[int, List[float]]:
    """Worker executing concurrent HMAC item snapshot tooltip lookups with cooperative yields."""
    passed, latencies = 0, []
    for uuid_str, sig in queries:
        t0 = time.perf_counter()
        res = chat_service.query_item_snapshot(uuid_str, sig)
        latencies.append((time.perf_counter() - t0) * 1000.0)
        if res.is_valid:
            passed += 1
        await asyncio.sleep(0)
    return passed, latencies


async def _dispatch_concurrent_work(
    chat_service: ChatService, requests: Sequence[SendChatRequestDTO],
    snapshots: Sequence[Tuple[str, str]], query_count: int, concurrency: int,
) -> Tuple[int, int, int, List[float]]:
    """Dispatches publisher and HMAC query workers concurrently via asyncio.gather."""
    c_size = max(1, len(requests) // concurrency)
    p_batches = [requests[i : i + c_size] for i in range(0, len(requests), c_size)]
    queries = [snapshots[i % len(snapshots)] for i in range(max(1, query_count))]
    q_size = max(1, len(queries) // concurrency)
    q_batches = [queries[i : i + q_size] for i in range(0, len(queries), q_size)]

    p_tasks = [_publisher_worker(chat_service, b) for b in p_batches]
    q_tasks = [_hmac_query_worker(chat_service, qb) for qb in q_batches]
    res_pub, res_queries = await asyncio.gather(asyncio.gather(*p_tasks), asyncio.gather(*q_tasks))

    accepted_total = sum(a for a, _ in res_pub)
    deliv_total = sum(d for _, d in res_pub)
    all_lats = [lat for _, lats in res_queries for lat in lats]
    return accepted_total, deliv_total, sum(p for p, _ in res_queries), all_lats


async def execute_round(
    round_idx: int, config: BenchmarkConfig, chat_service: ChatService,
    snapshots: Sequence[Tuple[str, str]],
) -> RoundMetrics:
    """Executes a single benchmark load round with memory and latency tracking."""
    chat_service.reset_rate_limits()
    gc.collect()
    mem_start = tracemalloc.get_traced_memory()[0]
    requests = generate_round_requests(round_idx, config.message_count, config.active_sample_subscribers)
    t0 = time.perf_counter()
    accepted, deliveries, hmac_ok, hmac_lats = await _dispatch_concurrent_work(
        chat_service, requests, snapshots, config.item_query_count, config.concurrency,
    )
    dur = time.perf_counter() - t0
    gc.collect()
    mem_end, mem_peak = tracemalloc.get_traced_memory()
    r_stats = chat_service.cluster_router.get_latency_stats()
    h_stats = calculate_percentiles(hmac_lats)
    return RoundMetrics(
        round_index=round_idx, messages_accepted=accepted, deliveries_count=deliveries,
        duration_sec=dur, throughput_msg_per_sec=accepted / dur if dur > 0 else 0.0,
        p50_ms=r_stats["p50"], p95_ms=r_stats["p95"], p99_ms=r_stats["p99"],
        hmac_queries_passed=hmac_ok, hmac_avg_ms=h_stats["avg"], hmac_p99_ms=h_stats["p99"],
        start_heap_mb=mem_start / (1024 * 1024), peak_heap_mb=mem_peak / (1024 * 1024),
        end_heap_mb=mem_end / (1024 * 1024), net_growth_mb=(mem_end - mem_start) / (1024 * 1024),
    )


def print_banner(cfg: BenchmarkConfig) -> None:
    """Prints benchmark configuration parameters."""
    print("=" * 70 + "\nFREEEXILE 2026: 1,000,000 CCU DISTRIBUTED CHAT STRESS BENCHMARK\n" + "=" * 70)
    print(f"Scale CCU: {cfg.simulated_ccu:,} | Subscribers: {cfg.active_sample_subscribers:,} | Shards: {cfg.cluster_shards}")
    print(f"Msgs/Round: {cfg.message_count:,} | Workers: {cfg.concurrency} | HMAC: {cfg.item_query_count:,} | Rounds: {cfg.leak_check_rounds}\n" + "-" * 70)


def print_round_metrics(m: RoundMetrics) -> None:
    """Formats telemetry output for a completed round."""
    print(f"[ROUND {m.round_index}] Duration: {m.duration_sec:.3f}s | Throughput: {m.throughput_msg_per_sec:,.1f} msg/s | Deliveries: {m.deliveries_count:,}")
    print(f"  Latency p50/p95/p99 : {m.p50_ms:.3f}ms / {m.p95_ms:.3f}ms / {m.p99_ms:.3f}ms")
    print(f"  HMAC Item Query p99 : {m.hmac_p99_ms:.3f}ms (avg {m.hmac_avg_ms:.3f}ms, valid: {m.hmac_queries_passed})")
    print(f"  Heap Memory Start/End: {m.start_heap_mb:.2f}MB -> {m.end_heap_mb:.2f}MB (Net: {m.net_growth_mb:.4f}MB)")


def print_final_summary(s: BenchmarkSummary) -> None:
    """Prints final multi-round verification report."""
    print("-" * 70 + "\nBENCHMARK VERIFICATION & SLA COMPLIANCE\n" + "-" * 70)
    print(f"Fan-Out Latency SLA (p99 < 15.0ms)     : [{'PASS' if s.sla_latency_passed else 'FAIL'}]")
    print(f"HMAC Item Query SLA (p99 < 2.0ms)      : [{'PASS' if s.sla_hmac_passed else 'FAIL'}]")
    print(f"Memory Leak Slope (Residual <= 0.05MB) : [{'PASS' if s.sla_leak_passed else 'FAIL'}] ({s.residual_growth_mb:.4f} MB)")
    print(f"Overall Benchmark Status               : [{'PASS' if s.overall_passed else 'FAIL'}]\n" + "=" * 70)


async def run_benchmark(config: BenchmarkConfig) -> BenchmarkSummary:
    """Coordinates full 1M CCU stress benchmark suite across multiple rounds."""
    try:
        sys.stdout.reconfigure(encoding="utf-8")
    except Exception:
        pass
    print_banner(config)
    if not tracemalloc.is_tracing():
        tracemalloc.start()

    chat_service = ChatService(cluster_shards=config.cluster_shards)
    snapshots = provision_test_items(chat_service, count=10)
    total_subs = provision_subscribers(chat_service, config.active_sample_subscribers, lambda _: None)
    print(f"[✓] Subscriber routing table established: {total_subs:,} handles across {config.cluster_shards} shards.\n")

    rounds: List[RoundMetrics] = []
    for r in range(1, config.leak_check_rounds + 1):
        metric = await execute_round(r, config, chat_service, snapshots)
        rounds.append(metric)
        print_round_metrics(metric)

    all_p99_ok = all(m.p99_ms < 15.0 for m in rounds)
    all_hmac_ok = all(m.hmac_p99_ms < 2.0 for m in rounds)
    res_growth = (rounds[-1].end_heap_mb - rounds[0].end_heap_mb) if len(rounds) > 1 else 0.0
    leak_ok = res_growth <= 0.05
    summary = BenchmarkSummary(
        config=config, round_metrics=tuple(rounds), sla_latency_passed=all_p99_ok,
        sla_hmac_passed=all_hmac_ok, sla_leak_passed=leak_ok,
        residual_growth_mb=res_growth, overall_passed=all_p99_ok and all_hmac_ok and leak_ok,
    )
    print_final_summary(summary)
    if tracemalloc.is_tracing():
        tracemalloc.stop()
    return summary


def parse_args(argv: Optional[Sequence[str]] = None) -> BenchmarkConfig:
    """Parses command-line arguments into BenchmarkConfig."""
    parser = argparse.ArgumentParser(description="1M CCU FreeExile Chat Stress Benchmark")
    parser.add_argument("--simulated-ccu", type=int, default=1_000_000)
    parser.add_argument("--active-sample-subscribers", type=int, default=5_000)
    parser.add_argument("--message-count", type=int, default=1_000)
    parser.add_argument("--concurrency", type=int, default=10)
    parser.add_argument("--leak-check-rounds", type=int, default=3)
    parser.add_argument("--shards", type=int, default=64)
    args = parser.parse_args(argv)
    return BenchmarkConfig(
        simulated_ccu=args.simulated_ccu,
        active_sample_subscribers=args.active_sample_subscribers,
        message_count=args.message_count,
        concurrency=args.concurrency,
        leak_check_rounds=args.leak_check_rounds,
        cluster_shards=args.shards,
    )


def main() -> None:
    """CLI Entry point."""
    cfg = parse_args()
    summary = asyncio.run(run_benchmark(cfg))
    sys.exit(0 if summary.overall_passed else 1)


if __name__ == "__main__":
    main()
