"""
Sprint 4 TDD Unit & Integration Test Suite for FreeExile:
PHÁO ĐÀI BẢO MẬT & TRÍ TUỆ NHÂN TẠO CHỐNG BOT (TUẦN 7)
- 1. Sinh trắc học cảm ứng & AI / Heuristic Bot Detection (> 99.2% bot detection)
- 2. Apple App Attest & Client Integrity Verification (Secure Enclave, Anti-Jailbreak)
- 3. Chống tấn công Replay, Timestamp Gating & Nonce Cache (< 1ms rejection)
- 4. Dynamic Polymorphic XOR Key RAM Encryption (Anti-Cheat Engine / GameGem)
- 5. Penetration Test Suite (Speedhack, Teleport, Packet Injection, Dupe Replay, Wallhack)
"""

from __future__ import annotations
import math
import os
import sys
import time
import unittest
from typing import List

PROJECT_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "../.."))
SERVER_DIR = os.path.join(PROJECT_ROOT, "server")
if SERVER_DIR not in sys.path:
    sys.path.insert(0, SERVER_DIR)
if PROJECT_ROOT not in sys.path:
    sys.path.insert(0, PROJECT_ROOT)

from server.security.touch_biometrics import (
    TouchBiometricsValidator,
    TouchPoint,
    BiometricEvaluationResult,
)
from server.security.app_attest import (
    AppleAppAttestValidator,
    AttestationToken,
    ClientEnvironmentInfo,
)
from server.security.packet_cipher import (
    TimestampGatedPacketGuard,
    PacketCipherEngine,
    AntiReplayWindow,
)
from server.security.polymorphic_memory import (
    PolymorphicRAMVariable,
    ProtectedPlayerStats,
    MemoryTamperViolationError,
)
from server.security.penetration_test_harness import (
    PenetrationTestHarness,
    PenetrationSuiteSummary,
)


class TestSprint4TouchBiometricsML(unittest.TestCase):
    """Test suite for Touch Biometrics AI & Micro-tremor Bot Detection."""

    def setUp(self) -> None:
        self.validator = TouchBiometricsValidator()

    def _generate_synthetic_human_gesture(self, seed: int) -> List[TouchPoint]:
        """Generates an authentic human curved swipe with natural tremors."""
        points: List[TouchPoint] = []
        for i in range(12):
            jitter = math.sin(i * 0.7 + seed) * 2.8
            tremor = math.sin(i * 1.5 + seed) * 0.4
            points.append(
                TouchPoint(
                    x=float(i * 12 + jitter + tremor),
                    y=float(i * 18 + math.cos(i * 0.5) * 3.5),
                    timestamp_ms=float(i * 40 + (math.sin(i + seed) * 3.5)),
                    major_radius=18.0 + math.sin(i * 0.8) * 3.0,
                    force=0.85 + math.sin(i * 0.4) * 0.15,
                )
            )
        return points

    def test_authentic_human_touch_accepted_with_high_trust(self) -> None:
        """Genuine human swipes are accepted with trust score >= 0.70."""
        for seed in range(20):
            sample = self._generate_synthetic_human_gesture(seed)
            res = self.validator.evaluate_detailed(sample)
            self.assertTrue(res.is_human, f"Failed on seed {seed}: {res.anomaly_reasons}")
            self.assertGreaterEqual(res.trust_score, 0.70)
            self.assertLess(res.bot_probability, 0.50)

    def test_bot_detection_rate_exceeds_99_point_2_percent(self) -> None:
        # Benchmark: Evaluates 1,200 synthetic bot trajectories across 6 bot archetypes
        archetypes = [
            lambda j, i: TouchPoint(float(j * 15), float(j * 25), float(j * 30), 15.0 + (j % 2), 1.0),
            lambda j, i: TouchPoint(float(100 + i + j), float(200 + j), float(j * 50), 0.0, 1.0),
            lambda j, i: TouchPoint(150.0, 300.0, float(j * 40), 16.0, 1.0),
            lambda j, i: TouchPoint(float(j * 10 + (j % 3)), float(j * 12), float(j * 33.333), 16.0, 1.0),
            lambda j, i: TouchPoint(float(j * 20), float(j * 20), float(j * 20), 18.0 + (j % 3), 1.0),
            lambda j, i: TouchPoint(100.0 + 30.0 * math.cos(j), 100.0 + 30.0 * math.sin(j), 1000.0, 15.0, 1.0),
        ]
        bot_samples: List[List[TouchPoint]] = [
            [fn(j, i) for j in range(8)] for fn in archetypes for i in range(200)
        ]

        total_bots = len(bot_samples)
        detected_bots = sum(1 for sample in bot_samples if not self.validator.evaluate_detailed(sample).is_human)
        detection_rate = (detected_bots / total_bots) * 100.0

        self.assertGreaterEqual(
            detection_rate, 99.2,
            f"Bot detection rate {detection_rate:.2f}% is below required 99.2% threshold!"
        )


class TestSprint4AppAttestAndIntegrity(unittest.TestCase):
    """Test suite for Apple App Attest & Client Integrity Verification."""

    def setUp(self) -> None:
        self.validator = AppleAppAttestValidator(expected_bundle_id="com.freeexile.game.ios")

    def test_app_attest_challenge_response_handshake(self) -> None:
        """Challenge nonce generation, TTL expiration, and single-use consumption."""
        device_id = "iphone_16_pro_max_01"
        nonce = self.validator.generate_challenge(device_id=device_id, ttl_seconds=120)
        self.assertEqual(len(nonce), 48)

        # Attestation with valid challenge passes
        token = self.validator.generate_mock_attestation(
            device_id=device_id, bundle_id="com.freeexile.game.ios", counter=1, challenge_nonce=nonce
        )
        ok, msg = self.validator.verify_attestation(token, required_challenge=nonce)
        self.assertTrue(ok)
        self.assertIn("verified", msg.lower())

        # Challenge consumption prevents replay
        replay_token = self.validator.generate_mock_attestation(
            device_id=device_id, bundle_id="com.freeexile.game.ios", counter=2, challenge_nonce=nonce
        )
        ok_replay, _ = self.validator.verify_attestation(replay_token, required_challenge=nonce)
        self.assertFalse(ok_replay)

    def test_assertion_counter_monotonicity(self) -> None:
        """Monotonic counter prevents hardware assertion replay attacks."""
        device_id = "device_test_counter"
        token = self.validator.generate_mock_attestation(device_id, "com.freeexile.game.ios", counter=10)
        self.validator.verify_attestation(token)

        self.assertTrue(self.validator.verify_assertion_counter(device_id, 11)[0])
        self.assertTrue(self.validator.verify_assertion_counter(device_id, 15)[0])
        # Replayed counter 15 rejected
        self.assertFalse(self.validator.verify_assertion_counter(device_id, 15)[0])
        # Decreasing counter 12 rejected
        self.assertFalse(self.validator.verify_assertion_counter(device_id, 12)[0])
        # Replaying earlier attestation token 10 after counter reached 15 rejected
        ok, msg = self.validator.verify_attestation(token)
        self.assertFalse(ok)
        self.assertIn("replay", msg.lower())
        # Non-positive counter rejected
        self.assertFalse(self.validator.verify_assertion_counter(device_id, 0)[0])
        self.assertFalse(self.validator.verify_assertion_counter(device_id, -1)[0])

    def test_client_environment_integrity_checks(self) -> None:
        """Flags jailbreak artifacts, attached debuggers, and modified bundle ID."""
        clean_env = ClientEnvironmentInfo(
            bundle_id="com.freeexile.game.ios",
            binary_hash="SHA256:OFFICIAL_APPLE_CERT_2026",
            is_jailbroken=False,
            debugger_attached=False,
        )
        clean_res = self.validator.verify_client_environment(clean_env)
        self.assertTrue(clean_res.is_trusted)
        self.assertEqual(len(clean_res.violation_flags), 0)

        hostile_env = ClientEnvironmentInfo(
            bundle_id="com.cracked.freeexile",
            binary_hash="SHA256:TAMPERED_MOD_IPA",
            is_jailbroken=True,
            debugger_attached=True,
            suspicious_dylib_loaded=True,
        )
        hostile_res = self.validator.verify_client_environment(hostile_env)
        self.assertFalse(hostile_res.is_trusted)
        self.assertIn("BUNDLE_ID_MISMATCH", hostile_res.violation_flags)
        self.assertIn("JAILBREAK_DETECTED", hostile_res.violation_flags)
        self.assertIn("DEBUGGER_ATTACHED", hostile_res.violation_flags)


class TestSprint4AntiReplayAndTimestampGating(unittest.TestCase):
    """Test suite for Anti-Replay, Timestamp Gating & Sub-Millisecond Nonce Cache."""

    def setUp(self) -> None:
        self.shared_key = b"SHARED_SECRET_PACKET_GUARD_KEY_2026"
        self.guard = TimestampGatedPacketGuard(shared_key=self.shared_key, max_drift_ms=3000)

    def test_packet_acceptance_within_drift_window(self) -> None:
        """Valid packet within 3000ms drift window is accepted."""
        now = int(time.time() * 1000)
        pkt = self.guard.create_packet(payload=b"MOVE_DIR_1_0", seq_num=1, timestamp_ms=now)
        ok, reason, latency = self.guard.verify_and_process_packet(pkt, server_now_ms=now + 50)
        self.assertTrue(ok)
        self.assertEqual(reason, "PACKET_ACCEPTED")
        self.assertLess(latency, 1.0)

        cipher = PacketCipherEngine(self.shared_key)
        with self.assertRaises(ValueError):
            cipher.encrypt(b"MOVE_DIR_1_0", seq_num=-1)

    def test_timestamp_drift_exceeded_rejected(self) -> None:
        """Packets with timestamp drift > 3000ms are rejected."""
        now = int(time.time() * 1000)
        stale_pkt = self.guard.create_packet(payload=b"CAST_SKILL_01", seq_num=2, timestamp_ms=now - 5000)
        ok, reason, latency = self.guard.verify_and_process_packet(stale_pkt, server_now_ms=now)
        self.assertFalse(ok)
        self.assertIn("TIMESTAMP_DRIFT_EXCEEDED", reason)
        self.assertLess(latency, 1.0)

    def test_duplicate_nonce_replay_rejected_in_sub_millisecond(self) -> None:
        """Replay attack with duplicated nonce is rejected in < 1ms."""
        now = int(time.time() * 1000)
        pkt = self.guard.create_packet(payload=b"BUY_ITEM_01", seq_num=3, timestamp_ms=now)
        ok1, _, _ = self.guard.verify_and_process_packet(pkt, server_now_ms=now)
        self.assertTrue(ok1)

        # Same packet replayed
        ok2, reason, latency = self.guard.verify_and_process_packet(pkt, server_now_ms=now + 10)
        self.assertFalse(ok2)
        self.assertEqual(reason, "REPLAY_DUPLICATE_NONCE")
        self.assertLess(latency, 1.0)

    def test_rejection_latency_benchmark_under_one_millisecond(self) -> None:
        """Verifies that 100 replayed/tampered packets are all rejected in < 1ms each."""
        now = int(time.time() * 1000)
        replayed_pkt = self.guard.create_packet(payload=b"SPAM_REPLAY", seq_num=10, timestamp_ms=now)
        self.guard.verify_and_process_packet(replayed_pkt, server_now_ms=now)

        latencies: List[float] = []
        for _ in range(100):
            _, _, lat = self.guard.verify_and_process_packet(replayed_pkt, server_now_ms=now)
            latencies.append(lat)

        avg_lat = sum(latencies) / len(latencies)
        max_lat = max(latencies)
        self.assertLess(avg_lat, 0.5, f"Average rejection latency {avg_lat:.4f}ms exceeded 0.5ms")
        self.assertLess(max_lat, 1.0, f"Max rejection latency {max_lat:.4f}ms exceeded 1.0ms")


class TestSprint4PolymorphicMemoryEncryption(unittest.TestCase):
    """Test suite for Dynamic Polymorphic XOR Key RAM Encryption."""

    def test_polymorphic_key_rotates_on_access(self) -> None:
        """Internal XOR mask key rotates dynamically on read/write without mutating value."""
        hp_var = PolymorphicRAMVariable[float](1250.75, "player_hp")
        initial_key = hp_var._mask_key

        # Reading value decrypts accurately and cycles mask key
        val = hp_var.get_value()
        self.assertAlmostEqual(val, 1250.75, places=2)
        second_key = hp_var._mask_key
        self.assertNotEqual(initial_key, second_key)

        hp_var.cycle_mask()
        third_key = hp_var._mask_key
        self.assertNotEqual(second_key, third_key)
        self.assertAlmostEqual(hp_var.get_value(), 1250.75, places=2)

        # Exact signed 64-bit integer roundtrip test
        signed_var = PolymorphicRAMVariable[int](-42, "signed_stat")
        self.assertEqual(signed_var.get_value(), -42)
        signed_var.set_value(-999999)
        self.assertEqual(signed_var.get_value(), -999999)

    def test_cheat_engine_memory_freeze_triggers_tamper_violation(self) -> None:
        """Overwriting or freezing masked memory bytes causes canary verification exception."""
        mana_var = PolymorphicRAMVariable[int](500, "player_mana")
        self.assertEqual(mana_var.get_value(), 500)

        # Attacker injects forged raw integer into RAM
        mana_var.inject_raw_memory_corruption(0xBADC0DE999)
        with self.assertRaises(MemoryTamperViolationError):
            mana_var.get_value()

    def test_protected_player_stats_lifecycle(self) -> None:
        """Validates damage, healing, mana spend, currency transactions, and frame key cycling."""
        stats = ProtectedPlayerStats(hp=1000.0, max_hp=1000.0, mana=500.0, max_mana=500.0, currency_hon_nguyen=100)
        self.assertTrue(stats.verify_all_integrity())

        # Take damage & heal
        stats.take_damage(350.0)
        self.assertAlmostEqual(stats.hp.get_value(), 650.0, places=2)
        stats.heal(150.0)
        self.assertAlmostEqual(stats.hp.get_value(), 800.0, places=2)

        # Spend mana
        self.assertTrue(stats.spend_mana(100.0))
        self.assertAlmostEqual(stats.mana.get_value(), 400.0, places=2)
        self.assertFalse(stats.spend_mana(999.0))

        # Currency operations
        stats.add_currency(50)
        self.assertEqual(stats.currency_hon_nguyen.get_value(), 150)
        self.assertTrue(stats.spend_currency(70))
        self.assertEqual(stats.currency_hon_nguyen.get_value(), 80)

        # Frame mask shuffle (120Hz simulation)
        stats.cycle_all_masks()
        self.assertTrue(stats.verify_all_integrity())


class TestSprint4PenetrationHarness(unittest.TestCase):
    """Test suite for Automated Penetration & Exploit Fuzzing Harness."""

    def setUp(self) -> None:
        self.harness = PenetrationTestHarness()

    def test_penetration_suite_achieves_100_percent_block_rate(self) -> None:
        """Executes full penetration matrix (Speedhack, Teleport, Tamper, Dupe, Wallhack, Memory)."""
        summary: PenetrationSuiteSummary = self.harness.run_comprehensive_penetration_suite()
        self.assertEqual(summary.total_attacks, 6)
        self.assertEqual(summary.blocked_attacks, 6)
        self.assertEqual(summary.block_rate_percent, 100.0)

        for verdict in summary.verdicts:
            self.assertTrue(verdict.is_blocked, f"Attack vector {verdict.vector_name} bypassed defense!")


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