"""
Server-Authoritative Combat Engine & Five Elements System for FreeExile.
Calculates elemental interactions (Ngũ Hành Tương Khắc), resistance mitigation,
critical strikes, and enforces Huyễn Ảnh Bộ (Phantom Evasion) 0.25s i-frame invulnerability.
"""

from typing import Any, Callable, Dict, List, Optional, Tuple, TYPE_CHECKING
from dataclasses import dataclass, field

if TYPE_CHECKING:
    from server.world.martial_matrix import FiveElements
else:
    try:
        from server.world.martial_matrix import FiveElements
    except ImportError:
        from world.martial_matrix import FiveElements


@dataclass
class CombatActor:
    actor_id: int
    name: str
    element: FiveElements
    current_hp: float = 1000.0
    max_hp: float = 1000.0
    base_attack: float = 50.0
    crit_chance: float = 0.05
    crit_multiplier: float = 1.5
    resistances: Dict[FiveElements, float] = field(default_factory=dict)
    last_evasion_timestamp_ms: int = -1000
    evasion_iframe_duration_ms: int = 250  # 0.25s i-frame window
    is_player: bool = False
    level: int = 1
    player_id: Optional[str] = None
    # Sprint 2: Weapon Sets, Energy, Animation Lock & Cooldowns
    weapon_set_1: Any = "SWORD"
    weapon_set_2: Any = "PROJECTILE_WEAPON"
    active_weapon_set: int = 1
    current_energy: float = 100.0
    max_energy: float = 100.0
    animation_locked_until_ms: int = 0
    cooldowns: Dict[int, int] = field(default_factory=dict)


@dataclass
class DamageEventResult:
    attacker_id: int
    defender_id: int
    raw_damage: float
    final_damage: float
    is_critical: bool
    is_evaded: bool
    element: FiveElements
    is_fatal: bool = False


@dataclass
class SkillCastResult:
    success: bool
    skill_id: int
    attacker_id: int
    target_id: Optional[int]
    damage_result: Optional[DamageEventResult] = None
    projectiles: int = 1
    energy_consumed: float = 0.0
    swapped_weapon: bool = False
    new_weapon_set: int = 1
    error_reason: str = ""


try:
    from server.world.combat_skill_helpers import (
        extract_weapon_category,
        weapon_matches,
        parse_element,
        resolve_skill_data,
        resolve_sigil_data,
    )
except ImportError:
    from world.combat_skill_helpers import (
        extract_weapon_category,
        weapon_matches,
        parse_element,
        resolve_skill_data,
        resolve_sigil_data,
    )


class CombatEngine:
    # Ngũ Hành Tương Khắc (Overcoming relationship: Attacker -> Defender)
    # Kim khắc Mộc, Mộc khắc Thổ, Thổ khắc Thủy, Thủy khắc Hỏa, Hỏa khắc Kim
    ELEMENT_OVERCOMING = {
        FiveElements.KIM: FiveElements.MOC,
        FiveElements.MOC: FiveElements.THO,
        FiveElements.THO: FiveElements.THUY,
        FiveElements.THUY: FiveElements.HOA,
        FiveElements.HOA: FiveElements.KIM,
    }

    def __init__(self) -> None:
        self.actors: Dict[int, CombatActor] = {}
        self.on_fatal_damage: Optional[Callable[[DamageEventResult, CombatActor, CombatActor], None]] = None

    def set_fatal_damage_hook(
        self, callback: Optional[Callable[[DamageEventResult, CombatActor, CombatActor], None]]
    ) -> None:
        """Sets an optional observer callback triggered on fatal damage."""
        self.on_fatal_damage = callback

    def attach_progression_service(self, service: Any) -> None:
        """Attaches a LevelProgressionService instance via an automated fatal damage hook."""
        def _hook(res: DamageEventResult, atk: CombatActor, dfn: CombatActor) -> None:
            if not getattr(dfn, "is_player", False) and getattr(atk, "is_player", False):
                pid = getattr(atk, "player_id", None) or str(atk.actor_id)
                service.award_monster_exp(player_id=pid, monster_level=getattr(dfn, "level", 1))
            elif getattr(dfn, "is_player", False):
                pid = getattr(dfn, "player_id", None) or str(dfn.actor_id)
                service.apply_death_penalty(player_id=pid)

        self.set_fatal_damage_hook(_hook)

    def register_actor(self, actor: CombatActor) -> None:
        self.actors[actor.actor_id] = actor

    def trigger_phantom_evasion(self, actor_id: int, timestamp_ms: int) -> bool:
        """Triggers Huyễn Ảnh Bộ, canceling animation lock and granting 250ms i-frame."""
        actor = self.actors.get(actor_id)
        if not actor:
            return False
        actor.last_evasion_timestamp_ms = timestamp_ms
        actor.animation_locked_until_ms = 0
        return True

    def is_animation_locked(self, actor_id: int, timestamp_ms: int) -> bool:
        """Checks if the actor is currently locked in a skill animation."""
        actor = self.actors.get(actor_id)
        if not actor:
            return False
        return timestamp_ms < actor.animation_locked_until_ms

    def swap_weapon_set(self, actor_id: int) -> int:
        """Manually swaps active weapon set between Set 1 and Set 2."""
        actor = self.actors.get(actor_id)
        if not actor:
            return 1
        actor.active_weapon_set = 2 if actor.active_weapon_set == 1 else 1
        return actor.active_weapon_set

    def _calculate_mitigated_damage(
        self,
        attacker: CombatActor,
        defender: CombatActor,
        raw_damage: float,
        damage_element: FiveElements,
        force_crit: Optional[bool],
    ) -> tuple[float, bool]:
        """Calculates elemental overcoming, resistance mitigation, and critical multiplier."""
        elem_mult = 1.25 if self.ELEMENT_OVERCOMING.get(damage_element) == defender.element else 1.0
        res = defender.resistances.get(damage_element, 0.0)
        effective_res = min(0.75, max(-0.50, res))
        mitigated = raw_damage * elem_mult * (1.0 - effective_res)
        is_crit = force_crit if force_crit is not None else False
        if is_crit:
            mitigated *= attacker.crit_multiplier
        return max(1.0, mitigated), is_crit

    def _check_special_damage_cases(
        self,
        attacker: Optional[CombatActor],
        defender: Optional[CombatActor],
        attacker_id: int,
        defender_id: int,
        raw_damage: float,
        damage_element: FiveElements,
        current_timestamp_ms: int,
    ) -> Optional[DamageEventResult]:
        """Handles missing actors or evasion i-frames without full combat resolution."""
        if not attacker or not defender:
            return DamageEventResult(
                attacker_id, defender_id, raw_damage, 0.0, False, False, damage_element, False
            )
        if defender.last_evasion_timestamp_ms >= 0:
            elapsed = current_timestamp_ms - defender.last_evasion_timestamp_ms
            if 0 <= elapsed <= defender.evasion_iframe_duration_ms:
                return DamageEventResult(
                    attacker_id, defender_id, raw_damage, 0.0, False, True, damage_element, False
                )
        return None

    def calculate_damage(
        self,
        attacker_id: int,
        defender_id: int,
        raw_damage: float,
        damage_element: FiveElements,
        current_timestamp_ms: int,
        force_crit: Optional[bool] = None,
    ) -> DamageEventResult:
        attacker = self.actors.get(attacker_id)
        defender = self.actors.get(defender_id)

        special = self._check_special_damage_cases(
            attacker, defender, attacker_id, defender_id, raw_damage, damage_element, current_timestamp_ms
        )
        if special is not None:
            return special
        assert attacker is not None and defender is not None

        final_damage, is_crit = self._calculate_mitigated_damage(
            attacker, defender, raw_damage, damage_element, force_crit
        )
        was_alive = (defender.current_hp > 0.0)
        new_hp = max(0.0, defender.current_hp - final_damage)
        defender.current_hp = new_hp
        is_fatal = was_alive and (new_hp <= 0.0)

        result = DamageEventResult(
            attacker_id=attacker_id,
            defender_id=defender_id,
            raw_damage=raw_damage,
            final_damage=round(final_damage, 2),
            is_critical=is_crit,
            is_evaded=False,
            element=damage_element,
            is_fatal=is_fatal,
        )

        if is_fatal and self.on_fatal_damage is not None:
            self.on_fatal_damage(result, attacker, defender)

        return result

    def _resolve_weapon_for_cast(
        self, attacker: CombatActor, req_weapon: Any
    ) -> Tuple[bool, bool]:
        """Validates weapon compatibility and auto-swaps if alternate weapon fits."""
        if req_weapon is None or str(req_weapon) == "ANY" or getattr(req_weapon, "value", "") == "ANY":
            return False, True
        cur_weapon = attacker.weapon_set_1 if attacker.active_weapon_set == 1 else attacker.weapon_set_2
        if weapon_matches(cur_weapon, req_weapon):
            return False, True
        alt_weapon = attacker.weapon_set_2 if attacker.active_weapon_set == 1 else attacker.weapon_set_1
        if weapon_matches(alt_weapon, req_weapon):
            attacker.active_weapon_set = 2 if attacker.active_weapon_set == 1 else 1
            return True, True
        return False, False

    def _calculate_sigil_modifiers(
        self,
        skill_id: int,
        skill_data: Dict[str, Any],
        linked_sigil_ids: List[int],
        skill_db: Any,
    ) -> Tuple[float, int, float]:
        """Calculates combined damage, projectiles, and energy cost from support sigils."""
        eff_damage = float(skill_data.get("base_damage", 0.0))
        eff_proj = int(skill_data.get("base_projectile_count", 1))
        eff_cost = float(skill_data.get("energy_cost", 0.0))
        for sigil_id in linked_sigil_ids:
            if skill_db and hasattr(skill_db, "validate_gem_link"):
                valid, _ = skill_db.validate_gem_link(skill_id, sigil_id)
                if not valid:
                    continue
            sigil = resolve_sigil_data(sigil_id, skill_db)
            if sigil:
                eff_damage *= float(sigil.get("damage_multiplier", 1.0))
                eff_proj += int(sigil.get("projectile_bonus", 0))
                eff_cost *= float(sigil.get("energy_cost_multiplier", 1.0))
        return max(0.0, eff_damage), max(1, eff_proj), max(0.0, eff_cost)

    def execute_skill_cast(
        self,
        attacker_id: int,
        skill_id: int,
        target_id: Optional[int],
        timestamp_ms: int,
        linked_sigil_ids: Optional[List[int]] = None,
        skill_db: Any = None,
    ) -> SkillCastResult:
        """Executes a server-authoritative skill cast with PoE2 mechanics."""
        attacker = self.actors.get(attacker_id)
        if not attacker:
            return SkillCastResult(False, skill_id, attacker_id, target_id, error_reason="ATTACKER_NOT_FOUND")

        if self.is_animation_locked(attacker_id, timestamp_ms):
            return SkillCastResult(False, skill_id, attacker_id, target_id, error_reason="ANIMATION_LOCKED")

        skill_data = resolve_skill_data(skill_id, skill_db)
        if not skill_data:
            return SkillCastResult(False, skill_id, attacker_id, target_id, error_reason="SKILL_NOT_FOUND")

        cooldown_until = attacker.cooldowns.get(skill_id, 0)
        if timestamp_ms < cooldown_until:
            return SkillCastResult(False, skill_id, attacker_id, target_id, error_reason="COOLDOWN_ACTIVE")

        req_weapon = skill_data.get("weapon_requirement")
        swapped, swap_ok = self._resolve_weapon_for_cast(attacker, req_weapon)
        if not swap_ok:
            return SkillCastResult(
                False, skill_id, attacker_id, target_id,
                swapped_weapon=False, new_weapon_set=attacker.active_weapon_set,
                error_reason="WEAPON_INCOMPATIBLE",
            )

        eff_damage, eff_proj, eff_cost = self._calculate_sigil_modifiers(
            skill_id, skill_data, linked_sigil_ids or [], skill_db
        )

        if attacker.current_energy < eff_cost:
            return SkillCastResult(
                False, skill_id, attacker_id, target_id,
                swapped_weapon=swapped, new_weapon_set=attacker.active_weapon_set,
                error_reason="INSUFFICIENT_ENERGY",
            )

        attacker.current_energy = max(0.0, attacker.current_energy - eff_cost)
        attacker.cooldowns[skill_id] = timestamp_ms + int(skill_data.get("cooldown_ms", 0))
        attacker.animation_locked_until_ms = timestamp_ms + int(skill_data.get("animation_lock_ms", 0))

        damage_res = None
        if target_id is not None and eff_damage > 0:
            elem = parse_element(skill_data.get("element", FiveElements.KIM))
            raw_dmg = eff_damage * (attacker.base_attack / 50.0)
            damage_res = self.calculate_damage(
                attacker_id=attacker_id,
                defender_id=target_id,
                raw_damage=raw_dmg,
                damage_element=elem,
                current_timestamp_ms=timestamp_ms,
            )

        return SkillCastResult(
            success=True,
            skill_id=skill_id,
            attacker_id=attacker_id,
            target_id=target_id,
            damage_result=damage_res,
            projectiles=eff_proj,
            energy_consumed=round(eff_cost, 2),
            swapped_weapon=swapped,
            new_weapon_set=attacker.active_weapon_set,
        )
