"""
FreeExile Character Stat Aggregator Service.
Combines Base Attributes, Equipped Item Affixes, and Meridian Passives.
Evaluates Path of Exile modifier math (Flat, Inc/Red, More/Less), tag filters,
and conditional triggers. Constructs calculation ASTs and persists them to SQLite.
Strict typing and modular methods (<= 50 lines each) following Elite Standards 2026.
"""

from __future__ import annotations
import uuid
import time
from typing import Dict, List, Optional, Any, Tuple
from server.stats.stat_types import (
    ModifierType,
    StatModifier,
    EvaluationContext,
    AggregatedCharacterStats,
    ConstantNode,
    SumNode,
    ScaleFactorNode,
    ProductNode,
    ModifierContributionNode,
)
from server.stats.formula_persistence import FormulaPersistenceService


class CharacterStatAggregator:
    """Authoritative service aggregating character stats under Path of Exile mathematics."""

    def __init__(
        self,
        persistence_service: Optional[FormulaPersistenceService] = None,
        inventory_service: Optional[Any] = None,
        meridian_service: Optional[Any] = None,
    ) -> None:
        self.persistence = persistence_service or FormulaPersistenceService(":memory:")
        self.inventory_service = inventory_service
        self.meridian_service = meridian_service

    def calculate_stats(
        self,
        character_loadout: Optional[Any] = None,
        inventory: Optional[Any] = None,
        meridian_bonus: Optional[Any] = None,
        context: Optional[EvaluationContext] = None,
        player_id: Optional[str] = None,
        custom_modifiers: Optional[List[StatModifier]] = None,
    ) -> AggregatedCharacterStats:
        """Executes full stat aggregation pipeline from all sources, AST build, and DB persist."""
        ctx = context or EvaluationContext()
        pid = player_id or "anon_player"
        mods: List[StatModifier] = []

        self._collect_attribute_modifiers(character_loadout, mods)
        self._collect_inventory_modifiers(inventory, mods)
        self._collect_meridian_modifiers(meridian_bonus, pid, mods)
        if custom_modifiers:
            mods.extend(custom_modifiers)

        stats_dict, ast_tree = self._evaluate_all_stats(mods, ctx)
        final_stats = self._build_aggregated_stats(stats_dict)

        calc_id = f"calc_{pid}_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
        self.persistence.save_calculation(
            player_id=pid,
            calculation_id=calc_id,
            formula_ast=ast_tree,
            final_stats=final_stats.to_dict(),
            context_tags=list(ctx.active_tags),
        )
        return final_stats

    def aggregate_player_stats(
        self,
        player_id: str,
        account_id: Optional[str] = None,
        character_id: Optional[str] = None,
        active_tags: Optional[Any] = None,
        conditions: Optional[Any] = None,
        trigger_reason: str = "MAP_JOIN",
    ) -> AggregatedCharacterStats:
        """High-level facade querying upstream services by player_id."""
        inv = None
        if self.inventory_service and account_id and character_id:
            try:
                inv = self.inventory_service.get_character_inventory(account_id, character_id)
            except Exception:
                pass
        meridian = None
        if self.meridian_service and player_id:
            try:
                meridian = self.meridian_service.compute_total_stats(player_id)
            except Exception:
                pass

        tags = frozenset(active_tags) if active_tags else frozenset()
        conds = frozenset(conditions.keys() if isinstance(conditions, dict) else (conditions or ()))
        ctx = EvaluationContext(active_tags=tags, conditions=conds)
        return self.calculate_stats(inventory=inv, meridian_bonus=meridian, context=ctx, player_id=player_id)

    def _collect_attribute_modifiers(self, loadout: Optional[Any], mods: List[StatModifier]) -> None:
        """Extracts base attribute scaling (STR, DEX, INT)."""
        str_val = getattr(loadout, "cuong_the", 50) if loadout else 50
        dex_val = getattr(loadout, "than_phap", 50) if loadout else 50
        int_val = getattr(loadout, "than_niem", 50) if loadout else 50

        # STR (+1 HP/pt, +0.2% phys/melee dmg/pt)
        mods.append(StatModifier("max_hp", ModifierType.FLAT, float(str_val), source="attr:STR"))
        mods.append(StatModifier(
            "attack_damage", ModifierType.INCREASED, str_val * 0.2,
            tags=frozenset({"physical", "melee", "attack"}), source="attr:STR"
        ))
        # DEX (+0.15% atk speed/pt, +0.1% crit/pt)
        mods.append(StatModifier("attack_speed", ModifierType.INCREASED, dex_val * 0.15, source="attr:DEX"))
        mods.append(StatModifier("crit_chance", ModifierType.FLAT, (dex_val * 0.1) / 100.0, source="attr:DEX"))
        # INT (+0.2% spell/elemental/pt)
        mods.append(StatModifier(
            "attack_damage", ModifierType.INCREASED, int_val * 0.2,
            tags=frozenset({"spell", "elemental"}), source="attr:INT"
        ))

    def _collect_inventory_modifiers(self, inventory: Optional[Any], mods: List[StatModifier]) -> None:
        """Ingests equipment affixes and weapon grip mechanics."""
        if not inventory:
            return
        items_dict = getattr(inventory, "equipment", inventory)
        items = list(items_dict.values()) if isinstance(items_dict, dict) else list(items_dict)

        has_main_2h, has_main_1h, has_off_1h = self._inspect_weapon_grips(items)
        if has_main_2h:
            mods.append(StatModifier("attack_damage", ModifierType.MORE, 50.0, source="grip:TWO_HANDED"))
        elif has_main_1h and has_off_1h:
            mods.append(StatModifier("attack_speed", ModifierType.MORE, 10.0, source="grip:DUAL_WIELD"))
            mods.append(StatModifier("block_chance", ModifierType.FLAT, 15.0, source="grip:DUAL_WIELD"))

        for it in items:
            self._parse_item_affixes(it, mods)

    def _inspect_weapon_grips(self, items: List[Any]) -> Tuple[bool, bool, bool]:
        """Detects 2H or Dual-Wield configurations."""
        main_2h = False
        main_1h = False
        off_1h = False
        for it in items:
            name = str(getattr(it, "name", "")).lower()
            item_id = str(getattr(it, "item_id", "")).lower()
            slot = str(getattr(it, "slot", "")).upper()
            meta = getattr(it, "metadata", {}) or {}
            is_2h = (
                "2h" in name
                or "2h" in item_id
                or meta.get("grip") == "TWO_HANDED"
                or meta.get("is_two_handed") is True
            )
            item_type_str = str(getattr(it, "item_type", "")).lower()
            is_wpn = (
                "weapon" in item_type_str
                or "sword" in name
                or "blade" in name
                or "kiếm" in name
                or "đao" in name
                or is_2h
            )
            if is_wpn:
                if is_2h:
                    main_2h = True
                elif "off_hand" in slot.lower() or "offhand" in slot.lower():
                    off_1h = True
                else:
                    main_1h = True
        return main_2h, main_1h, off_1h

    def _parse_item_affixes(self, item: Any, mods: List[Any]) -> None:
        """Parses legacy Affix and 15-tier AffixMod from an equipped item."""
        affixes = getattr(item, "affixes", [])
        src = f"item:{getattr(item, 'name', 'gear')}"
        for aff in affixes:
            # 1. Legacy Affix (stat_key, current_val)
            s_key = getattr(aff, "stat_key", None)
            if s_key:
                val = float(getattr(aff, "current_val", 0))
                self._map_legacy_affix(s_key, val, src, mods)
                continue
            # 2. Canonical AffixMod (mod_id, value)
            m_id = getattr(aff, "mod_id", None)
            if m_id:
                val = float(getattr(aff, "value", 0))
                self._map_affix_mod(m_id, val, src, mods)

    def _map_legacy_affix(self, key: str, val: float, src: str, mods: List[StatModifier]) -> None:
        """Maps legacy stone affixes to StatModifiers."""
        if key == "phys_dmg":
            mods.append(StatModifier("attack_damage", ModifierType.FLAT, val, source=src))
        elif key == "fire_dmg":
            mods.append(StatModifier("attack_damage", ModifierType.FLAT, val, frozenset({"fire"}), source=src))
        elif key == "max_hp":
            mods.append(StatModifier("max_hp", ModifierType.FLAT, val, source=src))
        elif key == "atk_speed":
            mods.append(StatModifier("attack_speed", ModifierType.INCREASED, val, source=src))
        elif key == "crit_rate":
            mods.append(StatModifier("crit_chance", ModifierType.FLAT, val / 100.0 if val > 1.0 else val, source=src))
        elif key in ("all_res", "fire_res", "cold_res", "chaos_res"):
            elems = ("hoa", "thuy", "kim", "moc", "tho") if key == "all_res" else (
                ("hoa",) if "fire" in key else (("thuy",) if "cold" in key else ("moc",))
            )
            for e in elems:
                mods.append(StatModifier(f"res_{e}", ModifierType.FLAT, val, source=src))

    def _map_affix_mod(self, mod_id: str, val: float, src: str, mods: List[StatModifier]) -> None:
        """Maps 15-tier AffixMod instances to StatModifiers."""
        if "pref_flat_phys" in mod_id or "pref_phys_pct" in mod_id:
            m_type = ModifierType.INCREASED if "pct" in mod_id else ModifierType.FLAT
            mods.append(StatModifier("attack_damage", m_type, val, source=src))
        elif "pref_fire_flat" in mod_id:
            mods.append(StatModifier("attack_damage", ModifierType.FLAT, val, frozenset({"fire"}), source=src))
        elif "pref_life" in mod_id:
            mods.append(StatModifier("max_hp", ModifierType.FLAT, val, source=src))
        elif "suff_move_spd" in mod_id:
            mods.append(StatModifier("move_speed", ModifierType.INCREASED, val, source=src))
        elif "suff_atk_spd" in mod_id:
            mods.append(StatModifier("attack_speed", ModifierType.INCREASED, val, source=src))
        elif "suff_crit_multi" in mod_id:
            mods.append(StatModifier("crit_multiplier", ModifierType.FLAT, val / 100.0 if val > 1.0 else val, source=src))
        elif "res" in mod_id:
            elem = "hoa" if "fire" in mod_id else ("thuy" if "cold" in mod_id else "moc")
            mods.append(StatModifier(f"res_{elem}", ModifierType.FLAT, val, source=src))

    def _collect_meridian_modifiers(self, meridian: Optional[Any], player_id: str, mods: List[StatModifier]) -> None:
        """Ingests Meridian passive constellation stats."""
        b = meridian
        if b is None and self.meridian_service and player_id != "anon_player":
            try:
                b = self.meridian_service.compute_total_stats(player_id)
            except Exception:
                b = None
        if not b:
            return
        src = "meridian"
        if getattr(b, "hp", 0) > 0:
            mods.append(StatModifier("max_hp", ModifierType.FLAT, float(b.hp), source=src))
        if getattr(b, "dps", 0) > 0:
            mods.append(StatModifier("attack_damage", ModifierType.FLAT, float(b.dps), source=src))
        if getattr(b, "dps_mult", 0.0) > 0.0:
            d_pct = b.dps_mult * 100.0 if b.dps_mult <= 2.0 else b.dps_mult
            mods.append(StatModifier("attack_damage", ModifierType.INCREASED, d_pct, source=src))
        if getattr(b, "crit_rate", 0.0) > 0.0:
            mods.append(StatModifier("crit_chance", ModifierType.FLAT, float(b.crit_rate), source=src))
        if getattr(b, "crit_dmg", 0.0) > 0.0:
            mods.append(StatModifier("crit_multiplier", ModifierType.FLAT, float(b.crit_dmg), source=src))
        if getattr(b, "resist", 0.0) > 0.0:
            for e in ("hoa", "thuy", "kim", "moc", "tho"):
                mods.append(StatModifier(f"res_{e}", ModifierType.FLAT, float(b.resist), source=src))

    def _evaluate_all_stats(
        self, mods: List[StatModifier], ctx: EvaluationContext
    ) -> Tuple[Dict[str, float], Dict[str, Any]]:
        """Evaluates mathematical resolution for each stat."""
        base_vals = {
            "attack_damage": 50.0, "max_hp": 1000.0, "crit_chance": 0.05,
            "crit_multiplier": 1.50, "move_speed": 6.0, "attack_speed": 1.0,
            "res_hoa": 0.0, "res_thuy": 0.0, "res_kim": 0.0, "res_moc": 0.0, "res_tho": 0.0,
        }
        all_stat_keys = set(base_vals.keys()) | {m.stat_key for m in mods}
        res_stats, ast_tree = {}, {}
        for s_key in all_stat_keys:
            final_val, stat_ast = self._compute_stat(s_key, base_vals.get(s_key, 0.0), mods, ctx)
            res_stats[s_key] = final_val
            ast_tree[s_key] = stat_ast
        return res_stats, ast_tree

    def _compute_stat(
        self, s_key: str, base_val: float, mods: List[StatModifier], ctx: EvaluationContext
    ) -> Tuple[float, Dict[str, Any]]:
        """Resolves PoE formula: (Base + Flat) * (1 + (Inc - Red)/100) * More/Less."""
        flat_total, inc_total, red_total, more_factor = 0.0, 0.0, 0.0, 1.0
        flat_nodes, inc_nodes, more_nodes = [], [], []

        for m in mods:
            if m.stat_key != s_key:
                continue
            active = ctx.matches_tags(m.tags) and ctx.is_condition_met(m.condition)
            c_node = ModifierContributionNode(
                source=m.source, modifier_type=m.mod_type.value, value=m.value,
                tags=list(m.tags), condition=m.condition, condition_met=active
            ).to_dict()
            if not active:
                continue
            if m.mod_type == ModifierType.FLAT:
                flat_total += m.value
                flat_nodes.append(c_node)
            elif m.mod_type == ModifierType.INCREASED:
                inc_total += m.value
                inc_nodes.append(c_node)
            elif m.mod_type == ModifierType.REDUCED:
                red_total += m.value
                inc_nodes.append(c_node)
            elif m.mod_type == ModifierType.MORE:
                more_factor *= (1.0 + m.value / 100.0)
                more_nodes.append(c_node)
            elif m.mod_type == ModifierType.LESS:
                more_factor *= (1.0 - m.value / 100.0)
                more_nodes.append(c_node)

        scale_factor = max(0.0, 1.0 + (inc_total - red_total) / 100.0)
        final_val = (base_val + flat_total) * scale_factor * more_factor
        precision = 4 if "crit" in s_key else 2
        final_rounded = round(final_val, precision)

        stat_ast = {
            "stat_name": s_key, "final_value": final_rounded,
            "base": ConstantNode("base", base_val, "base_rules").to_dict(),
            "flat": SumNode("flat_sum", flat_total, flat_nodes).to_dict(),
            "scale": ScaleFactorNode("inc_red_factor", inc_total - red_total, scale_factor, inc_nodes).to_dict(),
            "more": ProductNode("more_less_factor", more_factor, more_nodes).to_dict(),
        }
        return final_rounded, stat_ast

    def _build_aggregated_stats(self, s_dict: Dict[str, float]) -> AggregatedCharacterStats:
        """Assembles AggregatedCharacterStats from evaluated dictionary."""
        resistances = {e: s_dict.get(f"res_{e}", 0.0) for e in ("hoa", "thuy", "kim", "moc", "tho")}
        return AggregatedCharacterStats(
            attack_damage=s_dict["attack_damage"],
            max_hp=s_dict["max_hp"],
            crit_chance=min(1.0, max(0.0, s_dict["crit_chance"])),
            crit_multiplier=max(1.0, s_dict["crit_multiplier"]),
            move_speed=max(1.0, s_dict["move_speed"]),
            resistances=resistances,
            all_stats=dict(s_dict),
        )
