"""
Unit Tests for 1M CCU Chat Load Benchmark Suite (tools/stress/chat_load_benchmark.py).
Verifies CLI Parsing, Percentile Math, 8-Channel Generation, Concurrent Publishers,
HMAC Queries, Memory Leak Slope Calculation, and End-to-End Benchmark Execution.
"""

from __future__ import annotations

import unittest
from typing import List

from server.chat.chat_service import ChatService
from server.chat.chat_types import ChatChannelType, ChatMessageDTO
from tools.stress.chat_load_benchmark import (
    BenchmarkConfig,
    BenchmarkSummary,
    RoundMetrics,
    _dispatch_concurrent_work,
    _hmac_query_worker,
    _publisher_worker,
    calculate_percentiles,
    execute_round,
    generate_round_requests,
    parse_args,
    provision_subscribers,
    provision_test_items,
    run_benchmark,
)


class TestChatLoadBenchmark(unittest.IsolatedAsyncioTestCase):
    """Unit test suite covering the 1M CCU stress benchmark tool."""

    @classmethod
    def tearDownClass(cls) -> None:
        import tracemalloc
        if tracemalloc.is_tracing():
            tracemalloc.stop()

    async def asyncSetUp(self) -> None:
        self.chat_service = ChatService(cluster_shards=16)

    def test_01_parse_args_defaults_and_custom(self) -> None:
        """Verifies CLI argument parser defaults and overrides."""
        default_cfg = parse_args([])
        self.assertEqual(default_cfg.simulated_ccu, 1_000_000)
        self.assertEqual(default_cfg.active_sample_subscribers, 5_000)
        self.assertEqual(default_cfg.message_count, 1_000)
        self.assertEqual(default_cfg.concurrency, 10)
        self.assertEqual(default_cfg.leak_check_rounds, 3)
        self.assertEqual(default_cfg.cluster_shards, 64)

        custom_cfg = parse_args([
            "--simulated-ccu", "500000",
            "--active-sample-subscribers", "2000",
            "--message-count", "400",
            "--concurrency", "8",
            "--leak-check-rounds", "2",
            "--shards", "32"
        ])
        self.assertEqual(custom_cfg.simulated_ccu, 500_000)
        self.assertEqual(custom_cfg.active_sample_subscribers, 2_000)
        self.assertEqual(custom_cfg.message_count, 400)
        self.assertEqual(custom_cfg.concurrency, 8)
        self.assertEqual(custom_cfg.leak_check_rounds, 2)
        self.assertEqual(custom_cfg.cluster_shards, 32)

    def test_02_calculate_percentiles_math(self) -> None:
        """Validates statistical percentile calculation including p50, p95, p99."""
        empty = calculate_percentiles([])
        self.assertEqual(empty["p50"], 0.0)
        self.assertEqual(empty["p99"], 0.0)

        single = calculate_percentiles([15.5])
        self.assertEqual(single["p50"], 15.5)
        self.assertEqual(single["p99"], 15.5)
        self.assertEqual(single["avg"], 15.5)

        data = [float(x) for x in range(1, 101)]
        stats = calculate_percentiles(data)
        self.assertEqual(stats["p50"], 51.0)
        self.assertEqual(stats["p95"], 96.0)
        self.assertEqual(stats["p99"], 100.0)
        self.assertEqual(stats["avg"], 50.5)

    def test_03_provision_test_items(self) -> None:
        """Verifies cryptographic item snapshots are generated and valid."""
        snapshots = provision_test_items(self.chat_service, count=5)
        self.assertEqual(len(snapshots), 5)
        for uuid_str, sig in snapshots:
            self.assertTrue(uuid_str.startswith("bench_item_blade_"))
            self.assertTrue(bool(sig))
            res = self.chat_service.query_item_snapshot(uuid_str, sig)
            self.assertTrue(res.is_valid)
            self.assertIsNotNone(res.item_snapshot)

    def test_04_provision_subscribers(self) -> None:
        """Verifies subscriber handles are registered across shards for all channels."""
        delivered: List[ChatMessageDTO] = []
        count = provision_subscribers(self.chat_service, count=50, sink=delivered.append)
        self.assertGreater(count, 50)
        reg = self.chat_service.cluster_router.registry
        self.assertGreater(len(reg.get_callbacks("world")), 0)
        self.assertGreater(len(reg.get_callbacks("system")), 0)
        self.assertGreater(len(reg.get_callbacks("recruit")), 0)
        self.assertGreater(len(reg.get_callbacks("zone:zone_main")), 0)
        self.assertGreater(len(reg.get_callbacks("guild:guild_main")), 0)
        self.assertGreater(len(reg.get_callbacks("party:party_main")), 0)

    def test_05_generate_round_requests_8_channels(self) -> None:
        """Verifies round request generator covers all 8 channel types."""
        reqs = generate_round_requests(round_idx=1, total_count=40)
        self.assertEqual(len(reqs), 40)
        channels_found = {r.channel for r in reqs}
        self.assertEqual(len(channels_found), 8)
        self.assertIn(ChatChannelType.WORLD, channels_found)
        self.assertIn(ChatChannelType.ZONE, channels_found)
        self.assertIn(ChatChannelType.GUILD, channels_found)
        self.assertIn(ChatChannelType.PARTY, channels_found)
        self.assertIn(ChatChannelType.WHISPER, channels_found)
        self.assertIn(ChatChannelType.SYSTEM, channels_found)
        self.assertIn(ChatChannelType.RECRUIT, channels_found)
        self.assertIn(ChatChannelType.FEEDBACK, channels_found)

        # Check whisper target id
        whispers = [r for r in reqs if r.channel == ChatChannelType.WHISPER]
        self.assertTrue(all(w.target_id > 0 for w in whispers))

    async def test_06_concurrent_publisher_and_hmac_workers(self) -> None:
        """Tests concurrent worker execution with genuine coroutine interleaving."""
        provision_subscribers(self.chat_service, count=10, sink=lambda _: None)
        snapshots = provision_test_items(self.chat_service, count=5)
        reqs = generate_round_requests(round_idx=1, total_count=20)
        trace: List[str] = []
        orig_send = self.chat_service.handle_send_chat
        orig_query = self.chat_service.item_link_service.query_item_snapshot

        async def tracked_send(req: Any) -> Any:
            trace.append("PUB")
            return await orig_send(req)

        def tracked_query(uuid_str: str, sig: str) -> Any:
            trace.append("QUERY")
            return orig_query(uuid_str, sig)

        self.chat_service.handle_send_chat = tracked_send  # type: ignore[assignment]
        self.chat_service.item_link_service.query_item_snapshot = tracked_query  # type: ignore[assignment]
        try:
            accepted, deliv, hmac_ok, lats = await _dispatch_concurrent_work(
                chat_service=self.chat_service, requests=reqs, snapshots=snapshots,
                query_count=10, concurrency=2,
            )
        finally:
            self.chat_service.handle_send_chat = orig_send
            self.chat_service.item_link_service.query_item_snapshot = orig_query

        self.assertEqual(accepted, 20)
        self.assertGreater(deliv, 0)
        self.assertEqual(hmac_ok, 10)
        self.assertEqual(len(lats), 10)
        self.assertTrue(all(lat >= 0.0 for lat in lats))
        pub_idx = [i for i, kind in enumerate(trace) if kind == "PUB"]
        q_idx = [i for i, kind in enumerate(trace) if kind == "QUERY"]
        self.assertTrue(len(pub_idx) > 0 and len(q_idx) > 0)
        self.assertLess(min(q_idx), max(pub_idx), "Workers must interleave concurrently")

    async def test_07_execute_round_telemetry(self) -> None:
        """Verifies round execution captures valid metrics and latency telemetry."""
        snapshots = provision_test_items(self.chat_service, count=4)
        provision_subscribers(self.chat_service, count=10, sink=lambda _: None)
        config = BenchmarkConfig(message_count=16, concurrency=2, item_query_count=8, cluster_shards=16)
        m = await execute_round(1, config, self.chat_service, snapshots)
        self.assertEqual(m.round_index, 1)
        self.assertEqual(m.messages_accepted, 16)
        self.assertGreater(m.deliveries_count, 0)
        self.assertGreater(m.duration_sec, 0.0)
        self.assertGreater(m.throughput_msg_per_sec, 0.0)
        self.assertEqual(m.hmac_queries_passed, 8)
        self.assertLess(m.p99_ms, 15.0)
        self.assertLess(m.hmac_p99_ms, 2.0)

    async def test_08_end_to_end_benchmark_run(self) -> None:
        """Executes full multi-round benchmark on scaled parameters and verifies summary."""
        config = BenchmarkConfig(
            simulated_ccu=10_000,
            active_sample_subscribers=50,
            message_count=32,
            concurrency=4,
            leak_check_rounds=2,
            cluster_shards=16,
            item_query_count=16,
        )
        summary = await run_benchmark(config)
        self.assertIsInstance(summary, BenchmarkSummary)
        self.assertEqual(len(summary.round_metrics), 2)
        self.assertTrue(summary.sla_latency_passed)
        self.assertTrue(summary.sla_hmac_passed)
        self.assertTrue(summary.sla_leak_passed)
        self.assertTrue(summary.overall_passed)
        self.assertLessEqual(summary.residual_growth_mb, 0.05)


if __name__ == "__main__":
    unittest.main()
