"""AutoPOE2 - Atlas Graph Pathfinder & Route Optimizer (Doc 56 / Doc 65 / Doc 69).
=================================================================================
Thuật toán tìm đường trên đồ thị Atlas G = (V, E) phục vụ tiến trình mở khóa toàn bộ nodes:
- Tìm đường ngắn nhất Dijkstra / A* tới các Node mục tiêu (Citadels, Quest Nodes, Unique Maps).
- Tìm kiếm mở rộng biên (Frontier Search) hướng dẫn bot mở khóa full 115+ nodes Atlas.
- Hàm chi phí đa mục tiêu (Multi-Objective Cost Function) kết hợp Tier, Bonus, Nguy hiểm và Encounter.

Tuân thủ nghiêm ngặt:
- Rule 1: Cold Path Tier 2 (Python 3.11).
- Rule 4: Module hóa, trần file < 500 dòng.
- Rule 8: Bất biến kiến trúc INV-ATLAS-FARM-NODE & INV-ATLAS-QUEST-PROG.
"""

from __future__ import annotations

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

from src.assistant_tool.atlas.graph import AtlasGraph, AtlasGraphNode, NodeProgressionState
from src.assistant_tool.atlas.encounter_matcher import EncounterMatcher


@dataclass
class AtlasPathResult:
    """Kết quả tính toán đường đi trên đồ thị Atlas."""
    path_node_ids: List[str]
    total_cost: float
    total_tiers_traversed: int
    uncompleted_nodes_count: int
    target_node_id: str
    target_node_name: str = ""

    @property
    def length(self) -> int:
        return len(self.path_node_ids)

    @property
    def next_hop_node_id(self) -> Optional[str]:
        """Node kế tiếp cần kích hoạt để di chuyển theo lộ trình."""
        if len(self.path_node_ids) >= 2:
            return self.path_node_ids[1]
        elif len(self.path_node_ids) == 1:
            return self.path_node_ids[0]
        return None


class AtlasPathfinder:
    """Bộ định tuyến và tìm đường tối ưu trên cây bản đồ Atlas."""

    def __init__(
        self,
        tier_weight: float = 0.3,
        uncompleted_bonus_incentive: float = 0.8,
        danger_penalty: float = 1.2,
    ):
        self.tier_weight = tier_weight
        self.uncompleted_bonus_incentive = uncompleted_bonus_incentive
        self.danger_penalty = danger_penalty

    def calculate_edge_cost(self, u: AtlasGraphNode, v: AtlasGraphNode) -> float:
        """
        Hàm chi phí cạnh (Edge Cost) có trọng số:
        Cost(u, v) = Base + TierDelta * w_tier - EncounterScore * w_enc + UncompletedDelta
        """
        base_cost = 1.0

        # Nếu v đã hoàn thành với bonus, đường đi thông thoáng nhưng không mang lại điểm mới
        if v.is_bonus_completed:
            state_cost = 0.5
        elif v.is_frontier:
            # Ưu tiên cực lớn cho các node biên chưa hoàn thành để mở rộng bản đồ
            state_cost = 0.2
        elif v.progression_state == NodeProgressionState.COMPLETED:
            state_cost = 0.6  # Cần chạy lại để lấy bonus
        else:
            state_cost = 2.0  # LOCKED

        # Độ lệch Tier giữa hai node
        tier_cost = abs(v.tier - u.tier) * self.tier_weight

        # Giảm chi phí nếu v có các cơ chế giá trị cao (Encounter Incentive)
        enc_score = EncounterMatcher.calculate_score(v.encounters, v.tower_coverage_count)
        enc_discount = min(0.6, (enc_score - 1.0) * 0.15)

        total_cost = max(0.1, base_cost + state_cost + tier_cost - enc_discount)
        return round(total_cost, 3)

    def find_shortest_path(
        self,
        graph: AtlasGraph,
        start_id: str,
        target_id: str,
    ) -> Optional[AtlasPathResult]:
        """Thuật toán Dijkstra tìm lộ trình tối ưu từ start_id tới target_id."""
        start_node = graph.get_node(start_id)
        target_node = graph.get_node(target_id)
        if not start_node or not target_node:
            return None

        if start_id == target_id:
            return AtlasPathResult(
                path_node_ids=[start_id],
                total_cost=0.0,
                total_tiers_traversed=0,
                uncompleted_nodes_count=0 if start_node.is_bonus_completed else 1,
                target_node_id=target_id,
                target_node_name=target_node.name,
            )

        # Min-heap: (accumulated_cost, current_node_id, path)
        pq: List[Tuple[float, str, List[str]]] = [(0.0, start_id, [start_id])]
        visited: Dict[str, float] = {start_id: 0.0}

        while pq:
            cost, curr_id, path = heapq.heappop(pq)

            if curr_id == target_id:
                # Xây dựng kết quả lộ trình
                uncompleted = sum(
                    1 for nid in path
                    if graph.get_node(nid) and not graph.get_node(nid).is_bonus_completed
                )
                tiers = sum(
                    abs(graph.get_node(path[i]).tier - graph.get_node(path[i - 1]).tier)
                    for i in range(1, len(path))
                    if graph.get_node(path[i]) and graph.get_node(path[i - 1])
                )
                return AtlasPathResult(
                    path_node_ids=path,
                    total_cost=round(cost, 3),
                    total_tiers_traversed=tiers,
                    uncompleted_nodes_count=uncompleted,
                    target_node_id=target_id,
                    target_node_name=target_node.name,
                )

            if cost > visited.get(curr_id, float("inf")):
                continue

            curr_node = graph.get_node(curr_id)
            if not curr_node:
                continue

            for nbr in graph.get_neighbors(curr_id):
                edge_cost = self.calculate_edge_cost(curr_node, nbr)
                new_cost = cost + edge_cost

                if new_cost < visited.get(nbr.node_id, float("inf")):
                    visited[nbr.node_id] = new_cost
                    heapq.heappush(pq, (new_cost, nbr.node_id, path + [nbr.node_id]))

        return None

    def find_nearest_frontier(
        self,
        graph: AtlasGraph,
        current_id: str,
        preferred_tier: Optional[int] = None,
    ) -> Optional[AtlasPathResult]:
        """
        Tìm kiếm node biên (Frontier Node) gần nhất chưa hoàn thành để mở rộng lãnh thổ Atlas.
        """
        frontiers = graph.get_frontier_nodes()
        if not frontiers:
            return None

        # Sắp xếp các node biên theo khoảng cách tier ưu tiên
        if preferred_tier is not None:
            frontiers.sort(key=lambda n: abs(n.tier - preferred_tier))

        best_path: Optional[AtlasPathResult] = None
        min_cost = float("inf")

        for f_node in frontiers:
            path_res = self.find_shortest_path(graph, current_id, f_node.node_id)
            if path_res and path_res.total_cost < min_cost:
                min_cost = path_res.total_cost
                best_path = path_res

        return best_path
