"""
Empirical Adversarial Challenge Suite for Milestone M3 Chat Load Benchmark.
Tests:
1. CLI parameter edge cases and boundaries (negative, zero, invalid).
2. Concurrent tampered HMAC queries (bit-flip, bad hex, replay, empty, nonexistent).
3. Extreme shard imbalance (90% subscribers skewed into 1 shard).
4. Async task cancellation and clean shutdown without leaked tasks or tracing.
"""

from __future__ import annotations

import asyncio
import gc
import time
import tracemalloc
import unittest
from typing import List, Tuple

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


class TestChatLoadBenchmarkAdversarial(unittest.IsolatedAsyncioTestCase):
    """Empirical adversarial stress harness for chat load benchmark."""

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

    async def asyncSetUp(self) -> None:
        if tracemalloc.is_tracing():
            tracemalloc.stop()
        self.chat_service = ChatService(cluster_shards=64)

    async def asyncTearDown(self) -> None:
        if tracemalloc.is_tracing():
            tracemalloc.stop()

    # =========================================================================
    # VECTOR 1: CLI Parameter Boundary Testing
    # =========================================================================

    def test_cli_negative_and_zero_parameters_parsing(self) -> None:
        """Tests parse_args handling of negative and zero CLI values."""
        cfg = parse_args([
            "--simulated-ccu", "-50000",
            "--active-sample-subscribers", "0",
            "--message-count", "0",
            "--concurrency", "5",
            "--leak-check-rounds", "1",
            "--shards", "16",
        ])
        self.assertEqual(cfg.simulated_ccu, -50000)
        self.assertEqual(cfg.active_sample_subscribers, 0)
        self.assertEqual(cfg.message_count, 0)
        self.assertEqual(cfg.concurrency, 5)

    async def test_execution_with_zero_subscribers(self) -> None:
        """Validates benchmark execution when subscriber count is 0."""
        cfg = BenchmarkConfig(
            active_sample_subscribers=0,
            message_count=20,
            concurrency=2,
            leak_check_rounds=1,
            cluster_shards=16,
            item_query_count=10,
        )
        summary = await run_benchmark(cfg)
        self.assertEqual(len(summary.round_metrics), 1)
        m = summary.round_metrics[0]
        self.assertEqual(m.deliveries_count, 0)
        self.assertEqual(m.messages_accepted, 20)
        self.assertTrue(summary.sla_latency_passed)
        self.assertTrue(summary.overall_passed)

    async def test_execution_with_zero_messages(self) -> None:
        """Validates benchmark execution when message count is 0."""
        cfg = BenchmarkConfig(
            active_sample_subscribers=10,
            message_count=0,
            concurrency=2,
            leak_check_rounds=1,
            cluster_shards=16,
            item_query_count=10,
        )
        summary = await run_benchmark(cfg)
        m = summary.round_metrics[0]
        self.assertEqual(m.messages_accepted, 0)
        self.assertEqual(m.deliveries_count, 0)
        self.assertEqual(m.throughput_msg_per_sec, 0.0)

    async def test_execution_zero_concurrency_error(self) -> None:
        """Empirically confirms zero concurrency causes ZeroDivisionError without CLI guard."""
        cfg = BenchmarkConfig(active_sample_subscribers=10, message_count=10, concurrency=0)
        with self.assertRaises(ZeroDivisionError):
            await run_benchmark(cfg)

    async def test_execution_zero_shards_error(self) -> None:
        """Empirically confirms zero shards causes ZeroDivisionError in router."""
        cfg = BenchmarkConfig(active_sample_subscribers=10, cluster_shards=0)
        with self.assertRaises(ZeroDivisionError):
            await run_benchmark(cfg)

    # =========================================================================
    # VECTOR 2: Concurrent Tampered HMAC Queries
    # =========================================================================

    async def test_concurrent_tampered_hmac_queries(self) -> None:
        """
        Adversarial challenge: Concurrently query valid vs tampered HMAC signatures.
        Includes: bit-flipped, bad hex, empty signature, cross-item replay, nonexistent UUID.
        """
        # Provision authentic items
        authentic = provision_test_items(self.chat_service, count=10)
        valid_queries = list(authentic)

        # Generate adversarial tampered queries
        tampered_queries: List[Tuple[str, str]] = []
        for uid, sig in authentic:
            # 1. Bit-flipped hex signature
            flipped = sig[:-1] + ("0" if sig[-1] != "0" else "1")
            tampered_queries.append((uid, flipped))
            # 2. Fake signature (all zeros)
            tampered_queries.append((uid, "0" * 64))
            # 3. Truncated / empty signature
            tampered_queries.append((uid, ""))
            tampered_queries.append((uid, sig[:16]))
            # 4. Nonexistent item UUID with valid signature
            tampered_queries.append((f"nonexistent_{uid}", sig))

        # 5. Cross-item replay attack (valid sig of item 0 paired with item 1)
        tampered_queries.append((authentic[1][0], authentic[0][1]))

        # Mix valid (50) and tampered (50) into a single batch
        mixed_queries: List[Tuple[str, str, bool]] = []  # (uid, sig, expected_valid)
        for i in range(50):
            v = valid_queries[i % len(valid_queries)]
            mixed_queries.append((v[0], v[1], True))
            t = tampered_queries[i % len(tampered_queries)]
            mixed_queries.append((t[0], t[1], False))

        # Launch 10 concurrent workers executing the mixed queries
        async def query_tester(batch: List[Tuple[str, str, bool]]) -> Tuple[int, int, List[float]]:
            correct_rejects = 0
            correct_accepts = 0
            lats: List[float] = []
            for uid, sig, expected in batch:
                t0 = time.perf_counter()
                res = self.chat_service.query_item_snapshot(uid, sig)
                elapsed_ms = (time.perf_counter() - t0) * 1000.0
                lats.append(elapsed_ms)
                if expected:
                    if res.is_valid and res.item_snapshot is not None:
                        correct_accepts += 1
                else:
                    if not res.is_valid and res.error_message:
                        correct_rejects += 1
            return correct_accepts, correct_rejects, lats

        batch_size = max(1, len(mixed_queries) // 10)
        tasks = [
            query_tester(mixed_queries[i : i + batch_size])
            for i in range(0, len(mixed_queries), batch_size)
        ]
        results = await asyncio.gather(*tasks)

        total_accepts = sum(r[0] for r in results)
        total_rejects = sum(r[1] for r in results)
        all_lats = [lat for r in results for lat in r[2]]

        # Assert 100% precision: every valid query accepted, every tampered rejected
        self.assertEqual(total_accepts, 50, "Not all valid HMAC queries were accepted!")
        self.assertEqual(total_rejects, 50, "Some tampered HMAC queries were NOT rejected!")

        # Latency check: p99 must still be < 2.0ms even with rejection handling
        sorted_lats = sorted(all_lats)
        p99 = sorted_lats[int(len(sorted_lats) * 0.99)]
        self.assertLess(p99, 2.0, f"HMAC rejection p99 latency {p99:.3f}ms exceeded 2.0ms SLA!")

    # =========================================================================
    # VECTOR 3: Shard Imbalance (90% Skewed Subscriber Distribution)
    # =========================================================================

    async def test_shard_imbalance_90_percent_skew(self) -> None:
        """
        Adversarial challenge: 90% of subscribers concentrated in 1 single shard.
        Simulates extreme hotspotting (e.g. world channel burst or single mega-zone).
        """
        service = ChatService(cluster_shards=64)
        hot_channel = "world"
        cold_channels = [f"zone:zone_{i}" for i in range(1, 64)]

        received_hot = 0
        received_cold = 0

        def hot_sink(_: ChatMessageDTO) -> None:
            nonlocal received_hot
            received_hot += 1

        def cold_sink(_: ChatMessageDTO) -> None:
            nonlocal received_cold
            received_cold += 1

        # Register 900 subscribers on hot channel (90%)
        for cid in range(1, 901):
            service.subscribe_client(hot_channel, cid, hot_sink)

        # Register 100 subscribers spread evenly across 63 cold channels (10%)
        for cid in range(901, 1001):
            target_ch = cold_channels[cid % len(cold_channels)]
            service.subscribe_client(target_ch, cid, cold_sink)

        total_subs = service.cluster_router.registry.total_subscribers_count()
        self.assertEqual(total_subs, 1_000)

        # Verify that hot_channel is confined to 1 shard
        hot_shard_idx = hash(hot_channel) % 64
        hot_shard = service.cluster_router.registry.shards[hot_shard_idx]
        self.assertIn(hot_channel, hot_shard)
        self.assertEqual(len(hot_shard[hot_channel]), 900)

        # Execute 20 concurrent broadcast bursts into hot channel
        reqs = [
            SendChatRequestDTO(
                sender_id=500_000 + i,
                sender_name=f"HotTester_{i}",
                sender_level=35,
                channel=ChatChannelType.WORLD,
                content=f"Hot broadcast burst #{i}",
            )
            for i in range(20)
        ]

        async def send_req(r: SendChatRequestDTO) -> bool:
            service.channel_manager._last_send_time[(r.sender_id, r.channel)] = 0.0
            service.channel_manager._last_message_info[(r.sender_id, r.channel)] = ("", 0.0)
            res = await service.handle_send_chat(r)
            return res.success

        t0 = time.perf_counter()
        results = await asyncio.gather(*[send_req(r) for r in reqs])
        dur = time.perf_counter() - t0

        self.assertTrue(all(results))
        self.assertEqual(received_hot, 20 * 900)  # 18,000 deliveries

        # Latency statistics under hot shard pressure
        stats = service.cluster_router.get_latency_stats()
        self.assertLess(
            stats["p99"], 15.0,
            f"Hot-shard skewed p99 latency {stats['p99']}ms breached 15.0ms SLA! (Duration: {dur:.3f}s)",
        )

    # =========================================================================
    # VECTOR 4: Async Task Cancellation & Clean Shutdown
    # =========================================================================

    async def test_async_task_cancellation_during_benchmark(self) -> None:
        """
        Adversarial challenge: Cancel benchmark mid-execution via asyncio.CancelledError.
        Verifies clean propagation, no task leakage, and proper cleanup.
        """
        cfg = BenchmarkConfig(
            active_sample_subscribers=200,
            message_count=1_000,
            concurrency=4,
            leak_check_rounds=5,
            cluster_shards=16,
        )

        # Launch benchmark in background task
        task = asyncio.create_task(run_benchmark(cfg))
        # Allow benchmark to spin up and begin round 1
        await asyncio.sleep(0.08)

        # Force cancellation mid-flight
        task.cancel()

        with self.assertRaises(asyncio.CancelledError):
            await task

        # Verify task is cancelled cleanly
        self.assertTrue(task.cancelled())

        # Give event loop a cycle to finalize cancelled child tasks
        await asyncio.sleep(0.01)

        # Force gc to prove no circular memory leaks holding tasks
        gc.collect()

    async def test_tracemalloc_shutdown_resilience(self) -> None:
        """
        Verifies that tracemalloc tracing state is managed cleanly across runs.
        """
        if tracemalloc.is_tracing():
            tracemalloc.stop()

        cfg = BenchmarkConfig(
            active_sample_subscribers=50,
            message_count=20,
            concurrency=2,
            leak_check_rounds=1,
            cluster_shards=16,
        )
        summary = await run_benchmark(cfg)
        self.assertTrue(summary.overall_passed)
        # Verify tracemalloc was stopped at end of run_benchmark
        self.assertFalse(tracemalloc.is_tracing(), "tracemalloc was not cleanly stopped!")


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