"""
Adversarial Stress Test Suite for FreeExile 1,000,000 CCU Chat Microservice.
Empirically challenges:
1. Extreme concurrency (50 to 100 concurrent tasks).
2. High listener load (10,000+ active listeners across 64 shards).
3. Burst flooding on World channel with zero dropped deliveries.
4. Multi-round memory leak tracking (5 consecutive rounds, residual growth <= 0.05MB).
5. Dynamic player arrival unevicted dict leak detection in ChannelManager.
6. Concurrent HMAC item queries (< 2.0ms) and cryptographic spoof rejection.
7. O(N^2) subscription cache rebuilding characterization.
"""

from __future__ import annotations

import asyncio
import gc
import os
import sys
import time
import tracemalloc
import unittest
from typing import Callable, List, 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,
)


class TestChatAdversarialStress(unittest.IsolatedAsyncioTestCase):
    """Adversarial stress testing harness challenging worker benchmark claims."""

    def setUp(self) -> None:
        self.chat_service = ChatService(cluster_shards=64)

    def _register_subscribers_bulk(self, count: int, sink: Callable[[ChatMessageDTO], None]) -> int:
        """Efficiently registers count listeners across 5 channels rebuilding cache once."""
        channels = ("world", "system", "recruit", "zone:zone_main", "guild:guild_main")
        router = self.chat_service.cluster_router
        for ch in channels:
            s_idx = router.registry._get_shard_index(ch)
            shard_dict = router.registry.shards[s_idx].setdefault(ch, {})
            for cid in range(1, count + 1):
                shard_dict[cid] = sink
            router.registry._rebuild_cache(s_idx, ch)
        return router.registry.total_subscribers_count()

    async def test_01_extreme_concurrency_publisher_flood(self) -> None:
        """Challenges cluster with 50 and 100 concurrent publisher tasks."""
        self._register_subscribers_bulk(1000, lambda _: None)

        async def _publish_batch(task_id: int, base_sender: int, msg_count: int) -> int:
            accepted = 0
            for i in range(msg_count):
                req = SendChatRequestDTO(
                    sender_id=base_sender + task_id * 100 + i,
                    sender_name=f"Exile_{task_id}_{i}",
                    sender_level=35,
                    channel=ChatChannelType.WORLD,
                    content=f"Concurrency stress msg {task_id}-{i}",
                )
                res = await self.chat_service.handle_send_chat(req)
                if res.success:
                    accepted += 1
            return accepted

        # Phase 1: 50 concurrent tasks (10 msgs each = 500 msgs)
        tasks_50 = [_publish_batch(t, 200_000, 10) for t in range(50)]
        results_50 = await asyncio.gather(*tasks_50)
        self.assertEqual(sum(results_50), 500)

        # Phase 2: 100 concurrent tasks (5 msgs each = 500 msgs, distinct senders)
        tasks_100 = [_publish_batch(t, 300_000, 5) for t in range(100)]
        results_100 = await asyncio.gather(*tasks_100)
        self.assertEqual(sum(results_100), 500)

        stats = self.chat_service.cluster_router.get_latency_stats()
        self.assertLess(stats["p99"], 15.0, f"p99 latency exceeded 15ms SLA: {stats['p99']}ms")

    async def test_02_high_listener_load_10k_sharded_fanout(self) -> None:
        """Registers 10,000 active listeners (50,000 handles) and verifies 0 drop fanout."""
        delivery_counter = 0

        def _counting_sink(_: ChatMessageDTO) -> None:
            nonlocal delivery_counter
            delivery_counter += 1

        total_handles = self._register_subscribers_bulk(10_000, _counting_sink)
        self.assertGreaterEqual(total_handles, 50_000)

        req = SendChatRequestDTO(
            sender_id=100_001,
            sender_name="WorldBroadcaster",
            sender_level=50,
            channel=ChatChannelType.WORLD,
            content="Toàn cõi hoang vực lắng nghe hiệu lệnh!",
        )
        t0 = time.perf_counter()
        res = await self.chat_service.handle_send_chat(req)
        fanout_ms = (time.perf_counter() - t0) * 1000.0

        self.assertTrue(res.success)
        self.assertEqual(delivery_counter, 10_000, "World broadcast must reach exactly 10,000 listeners")
        self.assertLess(fanout_ms, 15.0, f"10k fanout exceeded 15.0ms SLA: {fanout_ms:.3f}ms")

    async def test_03_world_channel_burst_flooding(self) -> None:
        """Floods World channel with 500 messages across 50 tasks to 10k listeners."""
        self._register_subscribers_bulk(10_000, lambda _: None)

        async def _burst_sender(task_id: int, count: int) -> int:
            local_accepted = 0
            for i in range(count):
                req = SendChatRequestDTO(
                    sender_id=400_000 + task_id * 50 + i,
                    sender_name=f"BurstExile_{task_id}_{i}",
                    sender_level=40,
                    channel=ChatChannelType.WORLD,
                    content=f"Burst flood event packet {task_id}:{i}",
                )
                res = await self.chat_service.handle_send_chat(req)
                if res.success:
                    local_accepted += 1
            return local_accepted

        t_start = time.perf_counter()
        tasks = [_burst_sender(t, 10) for t in range(50)]
        results = await asyncio.gather(*tasks)
        total_dur = time.perf_counter() - t_start

        total_accepted = sum(results)
        self.assertEqual(total_accepted, 500)
        throughput = total_accepted / total_dur if total_dur > 0 else 0.0

        stats = self.chat_service.cluster_router.get_latency_stats()
        self.assertLess(stats["p99"], 15.0, f"Burst flood p99 exceeded 15ms: {stats['p99']}ms")
        self.assertGreater(throughput, 10.0, "Throughput under 10k listeners is unacceptably low")

    async def test_04_multi_round_memory_leak_steady_state(self) -> None:
        """Executes 5 stress rounds with steady active senders verifying growth <= 0.05MB."""
        tracemalloc.start()
        gc.collect()
        self._register_subscribers_bulk(10_000, lambda _: None)

        round_memories: List[float] = []
        for r in range(1, 6):
            gc.collect()
            for i in range(200):
                sid = 500_000 + i
                self.chat_service.channel_manager._last_send_time[(sid, ChatChannelType.WORLD)] = 0.0
                self.chat_service.channel_manager._last_message_info[(sid, ChatChannelType.WORLD)] = ("", 0.0)
                req = SendChatRequestDTO(
                    sender_id=sid, sender_name=f"Steady_{sid}", sender_level=30,
                    channel=ChatChannelType.WORLD, content=f"Round {r} payload buffer {i % 20}",
                )
                await self.chat_service.handle_send_chat(req)

            gc.collect()
            round_memories.append(tracemalloc.get_traced_memory()[0] / (1024 * 1024))

        tracemalloc.stop()
        residual_growth = round_memories[-1] - round_memories[1]
        self.assertLessEqual(
            residual_growth, 0.05,
            f"Steady state leak detected over 5 rounds! Residual growth: {residual_growth:.4f}MB"
        )

    async def test_05_dynamic_sender_unbounded_dict_growth_detection(self) -> None:
        """Empirically demonstrates ChannelManager unevicted dict leak under dynamic senders."""
        tracemalloc.start()
        gc.collect()

        round_memories: List[float] = []
        for r in range(1, 6):
            gc.collect()
            for i in range(200):
                req = SendChatRequestDTO(
                    sender_id=600_000 + r * 1000 + i,
                    sender_name=f"NewExile_{r}_{i}",
                    sender_level=30,
                    channel=ChatChannelType.WORLD,
                    content=f"Dynamic new arrival msg {r}-{i}",
                )
                await self.chat_service.handle_send_chat(req)

            gc.collect()
            round_memories.append(tracemalloc.get_traced_memory()[0] / (1024 * 1024))

        tracemalloc.stop()
        residual_growth = round_memories[-1] - round_memories[1]
        # Documents vulnerability: new senders inflate ChannelManager._last_send_time without TTL
        self.assertGreater(
            residual_growth, 0.05,
            "Expected ChannelManager dict inflation to exceed 0.05MB under continuous new senders"
        )

    async def test_06_concurrent_hmac_query_and_spoof_rejection(self) -> None:
        """Stresses HMAC item query SLA (< 2.0ms) and checks cryptographic forgery rejection."""
        legit_items: List[Tuple[str, str]] = []
        for i in range(20):
            uid = f"adv_item_blade_{i}"
            snap = self.chat_service.item_link_service.create_item_snapshot(
                item_uuid=uid, item_name_key=f"Huyết Ma Kiếm +{i}", rarity=5,
                element=1, quality=20, item_level=100, crafter_name="Cổ Thần",
                affixes=(ItemAffixDTO("stat", "Sát Thương", 1000.0),),
            )
            legit_items.append((uid, snap.hmac_signature))

        async def _query_worker(queries: Sequence[Tuple[str, str]]) -> List[float]:
            lats = []
            for uid_str, sig in queries:
                t0 = time.perf_counter()
                res = self.chat_service.query_item_snapshot(uid_str, sig)
                lats.append((time.perf_counter() - t0) * 1000.0)
                assert res.is_valid is True
            return lats

        q_batches = [
            [legit_items[j % len(legit_items)] for j in range(10)]
            for _ in range(50)
        ]
        results = await asyncio.gather(*[_query_worker(b) for b in q_batches])
        all_lats = [lat for r in results for lat in r]

        sorted_lats = sorted(all_lats)
        p99_hmac = sorted_lats[int(len(sorted_lats) * 0.99)]
        self.assertLess(p99_hmac, 2.0, f"HMAC p99 latency exceeded 2.0ms: {p99_hmac:.3f}ms")

        # Adversarial tamper attack: verify 100% rejection
        fake_uuid = "adv_item_blade_0"
        fake_sig = "0" * 64
        res_fake = self.chat_service.query_item_snapshot(fake_uuid, fake_sig)
        self.assertFalse(res_fake.is_valid, "Tampered signature must be rejected")

        res_missing = self.chat_service.query_item_snapshot("non_existent_item", "a" * 64)
        self.assertFalse(res_missing.is_valid, "Non-existent item must be rejected")

    async def test_07_duplicate_spam_and_anti_rmt_resilience(self) -> None:
        """Verifies duplicate spam rejection and auto-mute containment under load."""
        sender_id = 999_001

        req1 = SendChatRequestDTO(
            sender_id=sender_id, sender_name="Spammer", sender_level=30,
            channel=ChatChannelType.WORLD, content="Bán vàng SLL uy tín",
        )
        res1 = await self.chat_service.handle_send_chat(req1)
        self.assertTrue(res1.success)

        req2 = SendChatRequestDTO(
            sender_id=sender_id, sender_name="Spammer", sender_level=30,
            channel=ChatChannelType.WORLD, content="Bán vàng SLL uy tín",
        )
        res2 = await self.chat_service.handle_send_chat(req2)
        self.assertFalse(res2.success, "Immediate duplicate message must be rejected")

        rmt_req = SendChatRequestDTO(
            sender_id=999_002, sender_name="RmtBot", sender_level=30,
            channel=ChatChannelType.WORLD, content="Bán vàng sll liên hệ zalo 0912345678 ib ngay!",
        )
        rmt_res = await self.chat_service.handle_send_chat(rmt_req)
        self.assertFalse(rmt_res.success)
        self.assertGreaterEqual(rmt_res.risk_score, 60)

        sub_req = SendChatRequestDTO(
            sender_id=999_002, sender_name="RmtBot", sender_level=30,
            channel=ChatChannelType.WORLD, content="Xin chào mọi người",
        )
        sub_res = await self.chat_service.handle_send_chat(sub_req)
        self.assertFalse(sub_res.success, "Muted account must not be allowed to post")


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