"""
Empirical Stress Test Suite: Formula Persistence & Server Engine Loop Integration.
Tests SQLite concurrency, multi-threaded stress, JSON AST round-trip fidelity,
data integrity, and ServerEngineLoop registration with high-tier gear and passives.
Conforms to FreeExile Elite Standards 2026 (<= 350 lines, method <= 50 lines).
"""

from __future__ import annotations
import os
import sys
import json
import time
import tempfile
import unittest
import concurrent.futures
from typing import Dict, List, Any

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

from server.stats.stat_types import (
    ModifierType,
    StatModifier,
    EvaluationContext,
    AggregatedCharacterStats,
)
from server.stats.formula_persistence import FormulaPersistenceService
from server.stats.stat_aggregator import CharacterStatAggregator
from server.world.server_engine_loop import ServerEngineLoop, FiveElements
from server.world.meridian_types import MeridianStatBonus
from server.world.primal_stones_crafting import Affix, AffixType
from server.inventory.inventory_types import ItemType


class MockGearItem:
    """Mock item for high-tier equipment stress testing."""
    def __init__(
        self,
        name: str,
        item_id: str,
        affixes: list,
        is_2h: bool = False,
        slot: str = "MAIN_HAND",
        item_type: ItemType = ItemType.WEAPON,
    ) -> None:
        self.name = name
        self.item_id = item_id
        self.affixes = affixes
        self.item_type = item_type
        self.slot = slot
        self.metadata = {"is_two_handed": is_2h, "grip": "TWO_HANDED" if is_2h else "ONE_HANDED"}


class TestFormulaPersistenceStress(unittest.TestCase):
    """Adversarial and empirical stress test harness."""

    def test_disk_and_memory_initialization_lifecycle(self) -> None:
        """Verifies clean creation, WAL pragma setup, and closing on disk vs :memory:."""
        mem_svc = FormulaPersistenceService(":memory:")
        self.assertTrue(mem_svc._is_memory)
        self.assertIsNotNone(mem_svc._memory_conn)
        row_id = mem_svc.save_calculation("p1", "c1", {"ast": 1}, {"atk": 50.0}, ["melee"])
        self.assertGreater(row_id, 0)
        self.assertIsNotNone(mem_svc.get_calculation("c1"))
        mem_svc.close()
        self.assertIsNone(mem_svc._memory_conn)

        with tempfile.TemporaryDirectory() as tmp_dir:
            disk_path = os.path.join(tmp_dir, "stat_audit.db")
            disk_svc = FormulaPersistenceService(disk_path)
            self.assertFalse(disk_svc._is_memory)
            self.assertTrue(os.path.exists(disk_path))
            r_id = disk_svc.save_calculation("p2", "c2", {"ast": 2}, {"atk": 100.0}, ["fire"])
            self.assertGreater(r_id, 0)
            rec = disk_svc.get_calculation("c2")
            self.assertIsNotNone(rec)
            self.assertEqual(rec["player_id"], "p2")
            disk_svc.close()

    def test_concurrent_writes_and_reads_threadpool_stress(self) -> None:
        """Stresses disk-backed SQLite under concurrent multi-threaded writers and readers."""
        with tempfile.TemporaryDirectory() as tmp_dir:
            db_path = os.path.join(tmp_dir, "concurrent_stress.db")
            svc = FormulaPersistenceService(db_path)
            total_ops = 80
            player_ids = [f"player_{i % 8}" for i in range(total_ops)]

            def _worker_write_and_read(idx: int) -> Dict[str, Any]:
                cid = f"calc_conc_{idx}_{time.time_ns()}"
                pid = player_ids[idx]
                ast = {
                    "stat_name": "attack_damage",
                    "step_data": {"idx": idx, "formula": "50 + 100 * 1.5"},
                    "tags": ["melee", "savage", "việt_hóa_kiếm"],
                }
                stats = {"attack_damage": 225.0 + idx, "max_hp": 1500.0}
                tags = ["melee", f"tag_{idx}"]
                svc.save_calculation(pid, cid, ast, stats, tags)
                read_back = svc.get_calculation(cid)
                if read_back is None:
                    raise ValueError(f"Failed to read back calculation {cid}")
                return read_back

            with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
                futures = [executor.submit(_worker_write_and_read, i) for i in range(total_ops)]
                results = [f.result() for f in concurrent.futures.as_completed(futures)]

            self.assertEqual(len(results), total_ops)
            for res in results:
                self.assertIn("attack_damage", res["final_stats"])
                self.assertIn("step_data", res["formula_ast"])
            svc.close()

    def test_json_ast_fidelity_deep_structure_and_unicode(self) -> None:
        """Verifies 100% data fidelity of deep AST structures and Unicode diacritics."""
        mem_svc = FormulaPersistenceService(":memory:")
        cid = "fidelity_check_calc_001"
        pid = "player_hoang_gia_999"
        complex_ast = {
            "attack_damage": {
                "stat_name": "attack_damage",
                "final_value": 384.75,
                "base": {"ast_type": "Constant", "name": "base", "value": 50.0, "source": "Đặc Tính Khởi Nguyên"},
                "flat": {
                    "ast_type": "Sum",
                    "name": "flat_sum",
                    "total": 135.0,
                    "children": [
                        {"source": "Vũ Khí: Long Uyên Kiếm", "value": 85.0, "type": "FLAT"},
                        {"source": "Kinh Mạch: Đốc Mạch Chi Huyết", "value": 50.0, "type": "FLAT"},
                    ],
                },
                "scale": {
                    "ast_type": "ScaleFactor",
                    "percentage_sum": 65.0,
                    "multiplier": 1.65,
                    "children": [{"source": "Tà Ấn: Huyết Sát", "value": 65.0}],
                },
                "more": {
                    "ast_type": "Product",
                    "multiplier": 1.25,
                    "children": [{"source": "Thế Đao: Song Binh", "value": 25.0}],
                },
            }
        }
        final_stats = {
            "attack_damage": 384.75,
            "max_hp": 2450.0,
            "crit_chance": 0.3542,
            "crit_multiplier": 2.25,
            "move_speed": 8.4,
            "res_hoa": 45.0,
            "res_thuy": 30.0,
        }
        context_tags = ["cổ_võ", "huyết_sát", "song_binh", "physical", "melee"]

        mem_svc.save_calculation(pid, cid, complex_ast, final_stats, context_tags)
        retrieved = mem_svc.get_calculation(cid)
        self.assertIsNotNone(retrieved)

        # 100% AST fidelity check
        self.assertEqual(retrieved["formula_ast"], complex_ast)
        self.assertEqual(retrieved["final_stats"], final_stats)
        self.assertEqual(retrieved["context_tags"], context_tags)
        self.assertEqual(
            retrieved["formula_ast"]["attack_damage"]["base"]["source"],
            "Đặc Tính Khởi Nguyên"
        )
        self.assertEqual(
            retrieved["formula_ast"]["attack_damage"]["flat"]["children"][0]["source"],
            "Vũ Khí: Long Uyên Kiếm"
        )
        mem_svc.close()

    def test_rapid_succession_and_ordering_limits(self) -> None:
        """Tests rapid sequential inserts and get_player_calculations limit and order."""
        mem_svc = FormulaPersistenceService(":memory:")
        pid = "batch_player_01"
        for i in range(25):
            cid = f"calc_batch_{i:02d}"
            t = 1000.0 + i
            mem_svc.save_calculation(
                pid, cid, {"seq": i}, {"attack_damage": float(50 + i)}, ["tag"], timestamp=t
            )

        history_5 = mem_svc.get_player_calculations(pid, limit=5)
        self.assertEqual(len(history_5), 5)
        # Should be ordered by timestamp DESC
        self.assertEqual(history_5[0]["calculation_id"], "calc_batch_24")
        self.assertEqual(history_5[4]["calculation_id"], "calc_batch_20")

        history_all = mem_svc.get_player_calculations(pid, limit=100)
        self.assertEqual(len(history_all), 25)
        self.assertEqual(history_all[-1]["calculation_id"], "calc_batch_00")

        # Non-existent query checks
        self.assertIsNone(mem_svc.get_calculation("non_existent_id"))
        self.assertEqual(mem_svc.get_player_calculations("unknown_user"), [])
        mem_svc.close()

    def test_insert_or_replace_behavior(self) -> None:
        """Verifies updating an existing calculation_id safely overwrites record."""
        mem_svc = FormulaPersistenceService(":memory:")
        cid = "same_calculation_key"
        mem_svc.save_calculation("p1", cid, {"version": 1}, {"attack_damage": 60.0})
        rec1 = mem_svc.get_calculation(cid)
        self.assertEqual(rec1["final_stats"]["attack_damage"], 60.0)

        # Overwrite with version 2
        mem_svc.save_calculation("p1", cid, {"version": 2}, {"attack_damage": 95.0})
        rec2 = mem_svc.get_calculation(cid)
        self.assertEqual(rec2["final_stats"]["attack_damage"], 95.0)
        self.assertEqual(rec2["formula_ast"]["version"], 2)

        # Ensure no duplicates in player history
        history = mem_svc.get_player_calculations("p1")
        self.assertEqual(len(history), 1)
        mem_svc.close()

    def test_server_engine_registration_with_high_gear_stats(self) -> None:
        """Verifies CombatActor.base_attack > 50.0 and move_speed properly initialized."""
        affix_phys_flat = Affix("Huyết Long", AffixType.PREFIX, "phys_dmg", 120, 120, 120)
        affix_hp_flat = Affix("Kim Cương Thể", AffixType.PREFIX, "max_hp", 350, 350, 350)
        weapon_2h = MockGearItem("Huyết Ma Cự Đao 2H", "wpn_2h_001", [affix_phys_flat], is_2h=True)
        chest_armor = MockGearItem("Chiến Giáp Hắc Thiết", "arm_chest_001", [affix_hp_flat], is_2h=False, slot="CHEST")

        meridian_data = MeridianStatBonus(hp=200, dps=45.0, dps_mult=0.25, crit_rate=0.15, crit_dmg=0.50, resist=25.0)
        context = EvaluationContext(active_tags=frozenset({"physical", "melee", "attack"}))

        persistence = FormulaPersistenceService(":memory:")
        aggregator = CharacterStatAggregator(persistence_service=persistence)
        computed_stats = aggregator.calculate_stats(
            inventory=[weapon_2h, chest_armor],
            meridian_bonus=meridian_data,
            context=context,
            player_id="test_elite_hero",
        )

        # Calculations:
        # Base attack = 50.0. STR 50 gives +50*0.2 = +10% inc. Meridian gives +45 flat, +25% inc.
        # Weapon gives +120 flat, 2H gives +50% MORE.
        # Flat = 120 + 45 = 165. Total flat + base = 50 + 165 = 215.
        # Inc = 10% (STR) + 25% (Meridian) = 35% -> 1.35x.
        # More = 1.50x.
        # Expected attack = 215 * 1.35 * 1.50 = 435.375 -> rounded 435.38.
        self.assertGreater(computed_stats.attack_damage, 50.0)
        self.assertAlmostEqual(computed_stats.attack_damage, 435.38, places=1)
        self.assertGreater(computed_stats.max_hp, 1000.0)

        # Now wire into ServerEngineLoop
        engine = ServerEngineLoop(cell_size=64.0, stat_aggregator=aggregator)
        actor = engine.register_player(
            entity_id=101,
            initial_x=12.5,
            initial_y=-8.0,
            element=FiveElements.HOA,
            player_id="test_elite_hero",
            aggregated_stats=computed_stats,
        )

        # Verification of CombatActor initialized fields
        self.assertGreater(actor.base_attack, 50.0)
        self.assertEqual(actor.base_attack, computed_stats.attack_damage)
        self.assertEqual(actor.max_hp, computed_stats.max_hp)
        self.assertEqual(actor.current_hp, computed_stats.max_hp)
        self.assertEqual(actor.crit_chance, computed_stats.crit_chance)
        self.assertEqual(actor.crit_multiplier, computed_stats.crit_multiplier)
        self.assertEqual(actor.resistances[FiveElements.HOA], 25.0)
        self.assertEqual(actor.resistances[FiveElements.THUY], 25.0)

        # Verification of PlayerCharacter.move_speed synchronized with movement_authority
        registered_char = engine.movement_authority.players[101]
        self.assertEqual(registered_char.move_speed, computed_stats.move_speed)

    def test_server_engine_default_fallback_without_stats(self) -> None:
        """Verifies register_player without stats strictly preserves default 50.0 attack, 1000.0 HP, 6.0 move_speed."""
        engine = ServerEngineLoop(cell_size=64.0)
        actor = engine.register_player(entity_id=202, initial_x=0.0, initial_y=0.0)

        self.assertEqual(actor.base_attack, 50.0)
        self.assertEqual(actor.max_hp, 1000.0)
        self.assertEqual(actor.current_hp, 1000.0)
        self.assertEqual(actor.crit_chance, 0.05)
        self.assertEqual(actor.crit_multiplier, 1.50)
        self.assertEqual(engine.movement_authority.players[202].move_speed, 6.0)

    def test_multi_player_concurrent_registration_isolation(self) -> None:
        """Stress-tests multi-player registrations ensuring no cross-player stat contamination."""
        engine = ServerEngineLoop(cell_size=64.0)
        actors = []
        for i in range(20):
            stats = AggregatedCharacterStats(
                attack_damage=60.0 + i * 10.0,
                max_hp=1000.0 + i * 50.0,
                crit_chance=0.05 + i * 0.01,
                crit_multiplier=1.5 + i * 0.05,
                move_speed=6.0 + i * 0.2,
                resistances={"hoa": float(i)},
            )
            act = engine.register_player(
                entity_id=1000 + i,
                initial_x=float(i),
                initial_y=float(i),
                player_id=f"hero_{i}",
                aggregated_stats=stats,
            )
            actors.append(act)

        for i, act in enumerate(actors):
            expected_atk = 60.0 + i * 10.0
            expected_spd = 6.0 + i * 0.2
            self.assertEqual(act.base_attack, expected_atk)
            self.assertEqual(engine.movement_authority.players[1000 + i].move_speed, expected_spd)


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