"""
Authoritative Endgame Atlas Engine for FreeExile.
Manages:
1. Transition from Basic Campaign Quests to Endgame Phase.
2. Infinite Procedural World Map (Atlas) Node Graph generation.
3. Deterministic Frontier Expansion upon clearing maps.
4. Procedural Map Content Encounter rolling (1 to 5 random contents per node).
5. Integration with Hideout 6-Portal Map Device.
6. Server-Authoritative traversal rules and Atlas Passive progression.
"""

from __future__ import annotations
import math
import random
import time
from typing import Dict, List, Optional, Tuple, Set, Any

from server.world.endgame_types import (
    EndGameStatus,
    MapBiomeType,
    ContentType,
    MapNodeStatus,
    MapContent,
    MapNode,
    AtlasWorldState,
)
from server.world.endgame_content_catalog import (
    ALL_CONTENT_TYPES,
    CONTENT_DEFINITIONS,
    BIOME_METADATA,
    AFFIX_POOL_BY_TIER,
    calculate_content_count_for_tier,
)
from server.world.hideout_engine import HideoutEngine, AstralMap
from server.world.quest_engine import QuestEngine, QuestStatus


CAMPAIGN_COMPLETION_MIN_LEVEL = 10
CAMPAIGN_COMPLETION_QUEST_ID = "quest_norm_05_khai_mo_van_gioi_cac"


class EndGameAtlasEngine:
    def __init__(
        self,
        quest_engine: QuestEngine,
        hideout_engine: HideoutEngine
    ) -> None:
        self.quest_engine = quest_engine
        self.hideout_engine = hideout_engine
        # player_id -> AtlasWorldState
        self._player_atlases: Dict[str, AtlasWorldState] = {}

    def check_endgame_eligibility(self, player_id: str) -> bool:
        """Determines if the player has completed basic campaign quests to unlock endgame."""
        # 1. Player must exist in quest system and have sufficient level
        player_lvl = self.quest_engine._player_levels.get(player_id, 1)
        if player_lvl < CAMPAIGN_COMPLETION_MIN_LEVEL:
            return False

        # 2. Check canonical final milestone quest
        q_progress = self.quest_engine.get_player_quest(player_id, CAMPAIGN_COMPLETION_QUEST_ID)
        if q_progress is None:
            return False

        return q_progress.status in (QuestStatus.COMPLETED, QuestStatus.CLAIMED)

    def get_endgame_status(self, player_id: str) -> EndGameStatus:
        """Returns the current Endgame Phase status for a player."""
        state = self._player_atlases.get(player_id)
        if state and state.is_endgame_unlocked:
            return EndGameStatus.UNLOCKED

        if self.check_endgame_eligibility(player_id):
            return EndGameStatus.ELIGIBLE

        return EndGameStatus.LOCKED

    def get_atlas_state(self, player_id: str) -> Optional[AtlasWorldState]:
        return self._player_atlases.get(player_id)

    def unlock_endgame(self, player_id: str, seed: Optional[int] = None) -> AtlasWorldState:
        """Unlocks the Endgame Infinite World Map and generates the initial node cluster."""
        if seed is None:
            seed = int(time.time() * 1000) % 1_000_000

        atlas_state = AtlasWorldState(
            player_id=player_id,
            is_endgame_unlocked=True,
            world_seed=seed,
            current_origin_node_id="node_d0_b0",
            unlocked_depth=0,
            nodes={},
            total_maps_cleared=0,
            atlas_passive_points=0,
            created_timestamp_ms=int(time.time() * 1000),
        )

        # Generate initial cluster: Origin (Depth 0) and Branching children (Depth 1 & 2)
        self._generate_initial_cluster(atlas_state)
        self._player_atlases[player_id] = atlas_state
        return atlas_state

    def activate_map_node(self, player_id: str, node_id: str) -> Dict[str, Any]:
        """Activates a map node, opening exactly 6 portals on the player's Hideout Map Device."""
        state = self._player_atlases.get(player_id)
        if not state or not state.is_endgame_unlocked:
            return {"success": False, "error": "ENDGAME_NOT_UNLOCKED"}

        node = state.nodes.get(node_id)
        if not node:
            return {"success": False, "error": "NODE_NOT_FOUND"}

        if node.status not in (MapNodeStatus.ACCESSIBLE, MapNodeStatus.CLEARED):
            return {"success": False, "error": f"NODE_NOT_ACCESSIBLE_{node.status.value}"}

        # Convert MapNode to AstralMap format for HideoutEngine
        astral_map = AstralMap(
            map_id=node.node_id,
            name=node.name,
            tier=min(16, node.tier),
            zone_template_id=node.biome.value.lower(),
            item_quantity_bonus_pct=node.item_quantity_bonus_pct,
            item_rarity_bonus_pct=node.item_rarity_bonus_pct,
            monster_pack_size_pct=node.pack_size_bonus_pct,
            boss_name=node.boss_name,
            affixes=node.affixes,
        )

        success, msg, portals = self.hideout_engine.activate_map_device(player_id, astral_map)
        if not success:
            return {"success": False, "error": msg}

        hideout = self.hideout_engine.get_or_create_hideout(player_id)

        # Update node status in Atlas
        if node.status != MapNodeStatus.CLEARED:
            updated_node = self._copy_node_with_status(node, MapNodeStatus.IN_PROGRESS)
            state.nodes[node_id] = updated_node

        state.active_node_id = node_id
        return {
            "success": True,
            "node_id": node_id,
            "portals_opened": len(portals),
            "instance_id": hideout.map_device.active_instance_id,
            "contents": [c.to_dict() for c in node.contents],
        }

    def complete_map_node(self, player_id: str, node_id: str) -> Dict[str, Any]:
        """Marks a map node as cleared, awards points, and expands the infinite frontier."""
        state = self._player_atlases.get(player_id)
        if not state or not state.is_endgame_unlocked:
            return {"success": False, "error": "ENDGAME_NOT_UNLOCKED"}

        node = state.nodes.get(node_id)
        if not node:
            return {"success": False, "error": "NODE_NOT_FOUND"}

        was_already_cleared = (node.status == MapNodeStatus.CLEARED)
        if not was_already_cleared:
            state.nodes[node_id] = self._copy_node_with_status(
                node, MapNodeStatus.CLEARED, completed_ts=int(time.time() * 1000)
            )
            state.total_maps_cleared += 1
            state.atlas_passive_points += 1
            state.unlocked_depth = max(state.unlocked_depth, node.depth + 1)

            # Unlock direct child nodes
            for child_id in node.connected_node_ids:
                child = state.nodes.get(child_id)
                if child and child.status in (MapNodeStatus.UNREACHABLE, MapNodeStatus.DISCOVERED):
                    state.nodes[child_id] = self._copy_node_with_status(child, MapNodeStatus.ACCESSIBLE)

            # Infinite Frontier Expansion: generate deeper nodes if reaching boundary
            self._expand_infinite_frontier(state, node)

        state.active_node_id = None
        return {
            "success": True,
            "node_id": node_id,
            "atlas_points_awarded": 0 if was_already_cleared else 1,
            "unlocked_depth": state.unlocked_depth,
            "total_cleared": state.total_maps_cleared,
        }

    # -------------------------------------------------------------
    # INTERNAL PROCEDURAL GENERATION METHODS
    # -------------------------------------------------------------
    def _generate_initial_cluster(self, state: AtlasWorldState) -> None:
        """Synthesizes the initial cluster: Depth 0 (Origin), Depth 1, and Depth 2."""
        rng = random.Random(state.world_seed)

        # Depth 0: Origin Node
        origin_id = state.current_origin_node_id
        depth_1_ids = [f"node_d1_b{b}" for b in range(3)]

        origin_node = self._create_procedural_node(
            node_id=origin_id,
            depth=0,
            branch_index=0,
            grid_x=0,
            grid_y=0,
            tier=1,
            connected_ids=tuple(depth_1_ids),
            status=MapNodeStatus.ACCESSIBLE,
            rng=rng,
        )
        state.nodes[origin_id] = origin_node

        # Depth 1: 3 Branching nodes
        depth_2_all_ids: List[str] = []
        for b_idx, d1_id in enumerate(depth_1_ids):
            child_ids = (f"node_d2_b{b_idx * 2}", f"node_d2_b{b_idx * 2 + 1}")
            depth_2_all_ids.extend(child_ids)

            d1_node = self._create_procedural_node(
                node_id=d1_id,
                depth=1,
                branch_index=b_idx,
                grid_x=b_idx - 1,
                grid_y=1,
                tier=random.Random(state.world_seed + b_idx).randint(1, 2),
                connected_ids=child_ids,
                status=MapNodeStatus.DISCOVERED,
                rng=random.Random(state.world_seed + b_idx * 17),
            )
            state.nodes[d1_id] = d1_node

        # Depth 2: 6 Leaves (Unreachable initially until depth 1 cleared)
        for d2_idx, d2_id in enumerate(depth_2_all_ids):
            d2_node = self._create_procedural_node(
                node_id=d2_id,
                depth=2,
                branch_index=d2_idx,
                grid_x=(d2_idx - 2),
                grid_y=2,
                tier=random.Random(state.world_seed + d2_idx * 29).randint(2, 3),
                connected_ids=(),
                status=MapNodeStatus.UNREACHABLE,
                rng=random.Random(state.world_seed + d2_idx * 31),
            )
            state.nodes[d2_id] = d2_node

    def _expand_infinite_frontier(self, state: AtlasWorldState, cleared_node: MapNode) -> None:
        """Procedurally expands nodes at cleared_node.depth + 1 and beyond."""
        next_depth = cleared_node.depth + 1
        for child_id in cleared_node.connected_node_ids:
            child = state.nodes.get(child_id)
            if child and not child.connected_node_ids:
                # Generate 2 deeper grandchildren for this child
                deeper_depth = next_depth + 1
                grandchild_ids = [
                    f"node_d{deeper_depth}_b{child.branch_index * 2 + i}" for i in range(2)
                ]

                # Update child with connections to grandchildren
                state.nodes[child_id] = MapNode(
                    node_id=child.node_id,
                    depth=child.depth,
                    branch_index=child.branch_index,
                    grid_x=child.grid_x,
                    grid_y=child.grid_y,
                    tier=child.tier,
                    biome=child.biome,
                    name=child.name,
                    status=child.status,
                    connected_node_ids=tuple(grandchild_ids),
                    monster_level=child.monster_level,
                    item_quantity_bonus_pct=child.item_quantity_bonus_pct,
                    item_rarity_bonus_pct=child.item_rarity_bonus_pct,
                    pack_size_bonus_pct=child.pack_size_bonus_pct,
                    affixes=child.affixes,
                    boss_name=child.boss_name,
                    contents=child.contents,
                    completed_timestamp_ms=child.completed_timestamp_ms,
                )

                # Synthesize grandchildren nodes
                for gc_idx, gc_id in enumerate(grandchild_ids):
                    if gc_id not in state.nodes:
                        gc_seed = state.world_seed + deeper_depth * 100 + gc_idx
                        gc_rng = random.Random(gc_seed)
                        calculated_tier = min(16, max(1, deeper_depth + gc_rng.randint(0, 1)))

                        gc_node = self._create_procedural_node(
                            node_id=gc_id,
                            depth=deeper_depth,
                            branch_index=gc_idx,
                            grid_x=child.grid_x + (gc_idx * 2 - 1),
                            grid_y=deeper_depth,
                            tier=calculated_tier,
                            connected_ids=(),
                            status=MapNodeStatus.UNREACHABLE,
                            rng=gc_rng,
                        )
                        state.nodes[gc_id] = gc_node

    def _create_procedural_node(
        self,
        node_id: str,
        depth: int,
        branch_index: int,
        grid_x: int,
        grid_y: int,
        tier: int,
        connected_ids: Tuple[str, ...],
        status: MapNodeStatus,
        rng: random.Random,
    ) -> MapNode:
        """Synthesizes a single deterministic MapNode with biome, affixes, boss, and contents."""
        biome_list = list(MapBiomeType)
        biome = rng.choice(biome_list)
        biome_meta = BIOME_METADATA[biome]

        # Calculate map bonuses
        qty = 20 + tier * 4 + rng.randint(0, 10)
        rarity = 15 + tier * 3 + rng.randint(0, 8)
        pack_size = 10 + tier * 2 + rng.randint(0, 5)

        # Select Affixes by tier
        tier_category = "low" if tier <= 5 else ("mid" if tier <= 10 else ("high" if tier <= 15 else "uber"))
        affix_pool = AFFIX_POOL_BY_TIER[tier_category]
        affix_count = min(len(affix_pool), max(1, tier // 4 + 1))
        chosen_affixes = tuple(rng.sample(affix_pool, affix_count))

        # Procedural Map Name
        prefixes = ("Hoang Vu", "Khô Cốt", "Dị Biến", "Hắc Ngục", "Huyết Khí", "Cuồng Nộ")
        name = f"{rng.choice(prefixes)} · {biome_meta['name_vi']} (T{tier})"

        # Generate Map Contents (1 to 5 random encounters)
        contents = self._roll_node_contents(tier, rng)

        return MapNode(
            node_id=node_id,
            depth=depth,
            branch_index=branch_index,
            grid_x=grid_x,
            grid_y=grid_y,
            tier=tier,
            biome=biome,
            name=name,
            status=status,
            connected_node_ids=connected_ids,
            monster_level=67 + tier,
            item_quantity_bonus_pct=qty,
            item_rarity_bonus_pct=rarity,
            pack_size_bonus_pct=pack_size,
            affixes=chosen_affixes,
            boss_name=biome_meta["native_boss"],
            contents=contents,
        )

    def _roll_node_contents(self, tier: int, rng: random.Random) -> Tuple[MapContent, ...]:
        """Rolls a random count (1 to 5) and distinct types of savage map contents."""
        min_c, max_c = calculate_content_count_for_tier(tier)
        count = rng.randint(min_c, max_c)

        # Weighted sampling without replacement
        content_pool = list(ALL_CONTENT_TYPES)
        weights = [CONTENT_DEFINITIONS[ct]["weight"] for ct in content_pool]

        # Sample distinct content types
        selected_types: List[ContentType] = []
        for _ in range(min(count, len(content_pool))):
            total_w = sum(weights[i] for i, ct in enumerate(content_pool) if ct not in selected_types)
            if total_w <= 0:
                break
            pick_val = rng.uniform(0, total_w)
            curr = 0.0
            for i, ct in enumerate(content_pool):
                if ct in selected_types:
                    continue
                curr += weights[i]
                if curr >= pick_val:
                    selected_types.append(ct)
                    break

        # Construct MapContent objects with layout coordinates
        contents: List[MapContent] = []
        for idx, ct in enumerate(selected_types):
            meta = CONTENT_DEFINITIONS[ct]
            coord = (rng.uniform(-30.0, 30.0), rng.uniform(-30.0, 30.0))
            contents.append(
                MapContent(
                    content_id=f"mc_{ct.value.lower()}_{idx}_{rng.randint(100, 999)}",
                    content_type=ct,
                    name=meta["name"],
                    description=meta["description"],
                    difficulty_rating=min(5, max(1, meta["base_difficulty"] + (1 if tier > 10 else 0))),
                    reward_type=meta["reward_type"],
                    spawn_coord=coord,
                    is_completed=False,
                )
            )

        return tuple(contents)

    def _copy_node_with_status(
        self,
        node: MapNode,
        new_status: MapNodeStatus,
        completed_ts: Optional[int] = None
    ) -> MapNode:
        """Helper to create an updated copy of MapNode with modified status."""
        return MapNode(
            node_id=node.node_id,
            depth=node.depth,
            branch_index=node.branch_index,
            grid_x=node.grid_x,
            grid_y=node.grid_y,
            tier=node.tier,
            biome=node.biome,
            name=node.name,
            status=new_status,
            connected_node_ids=node.connected_node_ids,
            monster_level=node.monster_level,
            item_quantity_bonus_pct=node.item_quantity_bonus_pct,
            item_rarity_bonus_pct=node.item_rarity_bonus_pct,
            pack_size_bonus_pct=node.pack_size_bonus_pct,
            affixes=node.affixes,
            boss_name=node.boss_name,
            contents=node.contents,
            completed_timestamp_ms=completed_ts or node.completed_timestamp_ms,
        )
