"""
Adversarial Stress Testing Suite: Spatial Grid Isolation & Portal Roundtrip.
Validates:
1. High entity density (500+ entities) in SpatialGrid for zone_player_hideout.
2. Boundary coordinate precision (63.99 vs 64.01) and boundary jitter without ghost leaks.
3. Cross-zone spatial grid isolation between hideout, sanctuary, and outer maps.
4. Rapid portal roundtrip endurance (1,000x roundtrips) and concurrent multi-player storms.
5. Adversarial portal attacks (out-of-zone, under-level, non-existent portals).
"""

from __future__ import annotations
import concurrent.futures
import time
from typing import List, Set, Tuple
import pytest

from server.world.spatial_grid import Entity, SpatialGrid
from server.world.zone_engine import ZoneEngine
from server.world.zone_types import ZoneType


@pytest.fixture
def engine() -> ZoneEngine:
    return ZoneEngine()


class TestSpatialGridPortalAdversarial:
    """Adversarial stress and isolation test harness for Milestone 1."""

    def test_spatial_grid_high_volume_hideout_500_entities(self, engine: ZoneEngine) -> None:
        """Stress-tests hideout spatial grid with 500 entities, AOI querying, and cleanup."""
        grid = engine.zone_spatial_grids["zone_player_hideout"]
        total_entities = 500

        # Ingest 500 entities across a 10x10 cell cluster
        for i in range(total_entities):
            # Coordinates distribute entities across cells from -5 to +5
            x = (i % 25) * 16.0 - 200.0
            y = (i // 25) * 16.0 - 200.0
            entity = Entity(entity_id=1000 + i, x=x, y=y)
            grid.add_entity(entity)

        assert len(grid.entities) == total_entities
        assert len(grid.entity_cell_map) == total_entities

        # Invariant: sum of entities across all cells must exactly equal 500
        total_in_cells = sum(len(s) for s in grid.cells.values())
        assert total_in_cells == total_entities

        # AOI query at center (0, 0) with radius 1 (3x3 cells: -64 to +128 in x and y)
        aoi = grid.get_entities_in_aoi(0.0, 0.0, radius_cells=1)
        assert len(aoi) > 0

        # Verify all returned entities are indeed within the queried cells
        center_cell = grid.get_cell_coords(0.0, 0.0)
        expected_cells = {
            (center_cell[0] + dx, center_cell[1] + dy)
            for dx in range(-1, 2)
            for dy in range(-1, 2)
        }
        for eid in aoi:
            assert grid.entity_cell_map[eid] in expected_cells

        # Move all 500 entities simultaneously by +256.0 in X
        for i in range(total_entities):
            eid = 1000 + i
            ent = grid.entities[eid]
            grid.update_entity_position(ent, ent.x + 256.0, ent.y)

        # Invariant maintained after bulk movement
        assert len(grid.entities) == total_entities
        assert sum(len(s) for s in grid.cells.values()) == total_entities
        for cell_set in grid.cells.values():
            assert len(cell_set) > 0, "Empty cells must be garbage collected"

        # Remove all 500 entities
        for i in range(total_entities):
            removed = grid.remove_entity(1000 + i)
            assert removed is not None

        assert len(grid.entities) == 0
        assert len(grid.entity_cell_map) == 0
        assert len(grid.cells) == 0, "All cell sets must be fully cleaned up"

    def test_spatial_grid_concurrent_worker_stress(self, engine: ZoneEngine) -> None:
        """Stresses spatial grid with concurrent thread operations to verify robustness."""
        grid = engine.zone_spatial_grids["zone_player_hideout"]
        worker_count = 8
        items_per_worker = 50

        def worker_task(worker_id: int) -> int:
            base_id = 5000 + worker_id * 100
            for k in range(items_per_worker):
                eid = base_id + k
                ent = Entity(entity_id=eid, x=float(k * 10), y=float(worker_id * 20))
                grid.add_entity(ent)
                grid.update_entity_position(ent, ent.x + 32.0, ent.y + 32.0)
                _ = grid.get_entities_in_aoi(ent.x, ent.y, radius_cells=1)
            return items_per_worker

        with concurrent.futures.ThreadPoolExecutor(max_workers=worker_count) as executor:
            futures = [executor.submit(worker_task, w) for w in range(worker_count)]
            results = [f.result() for f in concurrent.futures.as_completed(futures)]

        assert sum(results) == worker_count * items_per_worker
        assert len(grid.entities) == worker_count * items_per_worker

    def test_boundary_coordinate_precision_and_no_leak(self, engine: ZoneEngine) -> None:
        """Adversarial boundary coordinate check (63.99 vs 64.01, negative bounds)."""
        grid = engine.zone_spatial_grids["zone_player_hideout"]

        # Cell size is 64.0. Math.floor dictates cell index.
        assert grid.get_cell_coords(63.999, 0.0) == (0, 0)
        assert grid.get_cell_coords(64.000, 0.0) == (1, 0)
        assert grid.get_cell_coords(64.001, 0.0) == (1, 0)

        assert grid.get_cell_coords(-0.001, 0.0) == (-1, 0)
        assert grid.get_cell_coords(0.000, 0.0) == (0, 0)
        assert grid.get_cell_coords(0.001, 0.0) == (0, 0)

        assert grid.get_cell_coords(-64.001, 0.0) == (-2, 0)
        assert grid.get_cell_coords(-64.000, 0.0) == (-1, 0)
        assert grid.get_cell_coords(-63.999, 0.0) == (-1, 0)

        # Boundary jitter test: 500 oscillations across boundary x=63.99 <-> x=64.01
        jitter_ent = Entity(entity_id=9901, x=63.99, y=0.0)
        grid.add_entity(jitter_ent)

        for step in range(500):
            new_x = 64.01 if (step % 2 == 0) else 63.99
            grid.update_entity_position(jitter_ent, new_x, 0.0)
            expected_cell = (1, 0) if (step % 2 == 0) else (0, 0)
            other_cell = (0, 0) if (step % 2 == 0) else (1, 0)

            assert grid.entity_cell_map[9901] == expected_cell
            assert 9901 in grid.cells[expected_cell]
            # Ensure entity NEVER leaves a ghost remnant in the other cell
            if other_cell in grid.cells:
                assert 9901 not in grid.cells[other_cell]

        grid.remove_entity(9901)
        assert 9901 not in grid.entities

    def test_cross_zone_spatial_isolation_under_identical_coords(self, engine: ZoneEngine) -> None:
        """Verifies hideout and sanctuary spatial grids never cross-pollinate."""
        hideout_grid = engine.zone_spatial_grids["zone_player_hideout"]
        sanctuary_grid = engine.zone_spatial_grids["zone_boundless_sanctuary"]

        # Entity A at boundary in hideout, Entity B at adjacent boundary in sanctuary
        ent_a = Entity(entity_id=7001, x=63.99, y=63.99)
        ent_b = Entity(entity_id=7002, x=64.01, y=64.01)

        hideout_grid.add_entity(ent_a)
        sanctuary_grid.add_entity(ent_b)

        # Hideout AOI must contain ONLY 7001, never 7002
        aoi_hideout = hideout_grid.get_entities_in_aoi(64.0, 64.0, radius_cells=2)
        assert 7001 in aoi_hideout
        assert 7002 not in aoi_hideout

        # Sanctuary AOI must contain ONLY 7002, never 7001
        aoi_sanctuary = sanctuary_grid.get_entities_in_aoi(64.0, 64.0, radius_cells=2)
        assert 7002 in aoi_sanctuary
        assert 7001 not in aoi_sanctuary

        # Identical entity_id (7777) present in both distinct zones
        ent_h = Entity(entity_id=7777, x=10.0, y=10.0)
        ent_s = Entity(entity_id=7777, x=200.0, y=200.0)
        hideout_grid.add_entity(ent_h)
        sanctuary_grid.add_entity(ent_s)

        # Move in hideout, verify sanctuary copy is completely unaffected
        hideout_grid.update_entity_position(ent_h, 15.0, 15.0)
        assert hideout_grid.entities[7777].x == 15.0
        assert sanctuary_grid.entities[7777].x == 200.0

        # Cleanup
        hideout_grid.remove_entity(7001)
        hideout_grid.remove_entity(7777)
        sanctuary_grid.remove_entity(7002)
        sanctuary_grid.remove_entity(7777)

    def test_portal_roundtrip_rapid_transitions_1000x(self, engine: ZoneEngine) -> None:
        """Simulates 1,000 rapid back-and-forth roundtrips (2,000 traversals)."""
        pid = "player_stress_commuter"
        engine.spawn_player(pid, "zone_boundless_sanctuary")

        start_time = time.perf_counter()
        roundtrips = 1000

        for _ in range(roundtrips):
            # Sanctuary -> Hideout
            ok1, dest1, _ = engine.traverse_portal(pid, "portal_sanctuary_to_hideout", player_level=1)
            assert ok1 is True
            assert dest1 == "zone_player_hideout"
            loc1 = engine.get_player_location(pid)
            assert loc1.zone_id == "zone_player_hideout"
            assert loc1.x == 0.0
            assert loc1.y == 0.0

            # Hideout -> Sanctuary
            ok2, dest2, _ = engine.traverse_portal(pid, "portal_hideout_to_sanctuary", player_level=1)
            assert ok2 is True
            assert dest2 == "zone_boundless_sanctuary"
            loc2 = engine.get_player_location(pid)
            assert loc2.zone_id == "zone_boundless_sanctuary"
            assert loc2.x == 360.0
            assert loc2.y == 0.0

        duration = time.perf_counter() - start_time
        # 2,000 traversals should complete within 0.5s in Python
        assert duration < 1.0, f"Roundtrip performance degraded: {duration:.3f}s for 2000 traversals"

    def test_portal_concurrent_multiplayer_storm(self, engine: ZoneEngine) -> None:
        """Simulates 50 players concurrently traversing portals across threads."""
        player_count = 50
        cycles_per_player = 20

        # Pre-spawn players
        for i in range(player_count):
            engine.spawn_player(f"player_storm_{i}", "zone_boundless_sanctuary")

        def player_storm_task(player_idx: int) -> int:
            pid = f"player_storm_{player_idx}"
            success_count = 0
            for _ in range(cycles_per_player):
                ok1, dest1, _ = engine.traverse_portal(pid, "portal_sanctuary_to_hideout", player_level=1)
                if ok1 and dest1 == "zone_player_hideout":
                    success_count += 1
                ok2, dest2, _ = engine.traverse_portal(pid, "portal_hideout_to_sanctuary", player_level=1)
                if ok2 and dest2 == "zone_boundless_sanctuary":
                    success_count += 1
            return success_count

        with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
            futures = [executor.submit(player_storm_task, i) for i in range(player_count)]
            total_success = sum(f.result() for f in concurrent.futures.as_completed(futures))

        expected = player_count * cycles_per_player * 2
        assert total_success == expected

    def test_adversarial_portal_invalid_traversal_attacks(self, engine: ZoneEngine) -> None:
        """Tests invalid portal requests: wrong zone, wrong level, fake portal ID."""
        pid = "player_adversary_01"
        engine.spawn_player(pid, "zone_boundless_sanctuary")

        # 1. Player is in Sanctuary, tries to traverse Hideout-to-Sanctuary portal
        ok, dest, msg = engine.traverse_portal(pid, "portal_hideout_to_sanctuary", player_level=10)
        assert ok is False
        assert dest == ""
        assert "Không tìm thấy Lối Qua" in msg

        # 2. Player level 0 (under required level 1)
        ok_lvl, dest_lvl, msg_lvl = engine.traverse_portal(
            pid, "portal_sanctuary_to_ancient_sword", player_level=5  # requires min_level 8
        )
        assert ok_lvl is False
        assert dest_lvl == ""
        assert "Chưa đạt cấp độ yêu cầu" in msg_lvl

        # 3. Non-existent portal ID
        ok_fake, dest_fake, msg_fake = engine.traverse_portal(pid, "portal_non_existent_id", player_level=100)
        assert ok_fake is False
        assert dest_fake == ""
        assert "Không tìm thấy Lối Qua" in msg_fake

    def test_cross_zone_isolation_all_canonical_zones(self, engine: ZoneEngine) -> None:
        """Verifies spatial grid isolation across all registered canonical zones."""
        all_zone_ids = list(engine.zones.keys())
        assert len(all_zone_ids) >= 11
        assert "zone_player_hideout" in all_zone_ids
        assert "zone_boundless_sanctuary" in all_zone_ids

        # Ensure each zone has a dedicated, non-shared SpatialGrid instance
        grid_instances = [engine.zone_spatial_grids[zid] for zid in all_zone_ids]
        unique_instances = {id(g) for g in grid_instances}
        assert len(unique_instances) == len(all_zone_ids), "Each zone must own a unique SpatialGrid"

        # Ingest entity 8888 into all zones with distinct coords
        for idx, zid in enumerate(all_zone_ids):
            grid = engine.zone_spatial_grids[zid]
            grid.add_entity(Entity(entity_id=8888, x=float(idx * 100), y=0.0))

        # Assert no zone overrides another zone's entity position
        for idx, zid in enumerate(all_zone_ids):
            grid = engine.zone_spatial_grids[zid]
            assert grid.entities[8888].x == float(idx * 100)

        # Cleanup
        for zid in all_zone_ids:
            engine.zone_spatial_grids[zid].remove_entity(8888)
