"""
Bộ Não Thích Ứng Trực Tuyến: Contextual Bandit (LinUCB with Online Learning)
Ra quyết định hành động tối ưu trong chiến đấu và tự học qua từng trải nghiệm thực chiến.
Thiết kế thuần Python 3.11 (Zero Dependency / No NumPy required) đảm bảo siêu nhẹ và không crash.
"""

import json
import logging
import math
import os
from typing import Dict, List, Optional, Tuple

logger = logging.getLogger("ContextualBandit")


def _create_identity(d: int) -> List[List[float]]:
    return [[1.0 if i == j else 0.0 for j in range(d)] for i in range(d)]


def _mat_add(A: List[List[float]], B: List[List[float]]) -> List[List[float]]:
    d = len(A)
    return [[A[i][j] + B[i][j] for j in range(d)] for i in range(d)]


def _outer_product(x: List[float], y: List[float]) -> List[List[float]]:
    d = len(x)
    return [[x[i] * y[j] for j in range(d)] for i in range(d)]


def _mat_vec_mul(A: List[List[float]], x: List[float]) -> List[float]:
    d = len(A)
    return [sum(A[i][j] * x[j] for j in range(d)) for i in range(d)]


def _vec_dot(x: List[float], y: List[float]) -> float:
    return sum(a * b for a, b in zip(x, y))


def _invert_matrix(A: List[List[float]]) -> List[List[float]]:
    """Đảo ma trận kích thước nhỏ bằng phương pháp Gauss-Jordan với Partial Pivoting."""
    n = len(A)
    aug = [row[:] + [1.0 if i == j else 0.0 for j in range(n)] for i, row in enumerate(A)]

    for col in range(n):
        pivot_row = col
        max_val = abs(aug[col][col])
        for r in range(col + 1, n):
            if abs(aug[r][col]) > max_val:
                max_val = abs(aug[r][col])
                pivot_row = r

        if max_val < 1e-12:
            aug[col][col] += 1e-5

        if pivot_row != col:
            aug[col], aug[pivot_row] = aug[pivot_row], aug[col]

        pivot = aug[col][col]
        for c in range(2 * n):
            aug[col][c] /= pivot

        for r in range(n):
            if r != col:
                factor = aug[r][col]
                if factor != 0.0:
                    for c in range(2 * n):
                        aug[r][c] -= factor * aug[col][c]

    return [row[n:] for row in aug]


class LinUCBTacticalAdvisor:
    """
    Bộ thuật toán Contextual Bandit theo trường phái LinUCB (Linear Upper Confidence Bound).
    """

    DEFAULT_ACTIONS = [
        "AGGRESSIVE_BURST",
        "KITE_STABILIZE",
        "RETREAT_SAFE",
        "EMERGENCY_PORTAL",
        "LOOT_WINDOW"
    ]

    def __init__(
        self,
        actions: Optional[List[str]] = None,
        feature_dim: int = 8,
        alpha: float = 0.8,
        model_path: str = "captures/bandit_model.json"
    ):
        self.actions = actions or self.DEFAULT_ACTIONS
        self.d = feature_dim
        self.alpha = alpha
        self.model_path = model_path

        self.A: Dict[str, List[List[float]]] = {a: _create_identity(self.d) for a in self.actions}
        self.b: Dict[str, List[float]] = {a: [0.0] * self.d for a in self.actions}
        self.decision_count = 0
        self._load_model()

    def extract_context(
        self,
        hp_pct: float,
        es_pct: float,
        monster_count: int,
        boss_poise_pct: float,
        flask_charges_pct: float,
        map_tier: int,
        recent_dps: float,
        damage_taken_rate: float
    ) -> List[float]:
        def clamp(v: float) -> float:
            return max(0.0, min(1.0, float(v)))

        return [
            clamp(hp_pct / 100.0),
            clamp(es_pct / 100.0),
            clamp(monster_count / 20.0),
            clamp(boss_poise_pct / 100.0),
            clamp(flask_charges_pct / 100.0),
            clamp(map_tier / 16.0),
            clamp(recent_dps / 50000.0),
            clamp(damage_taken_rate / 100.0),
        ]

    def recommend_action(self, context_vector: List[float]) -> Tuple[str, float, Dict[str, float]]:
        best_action = self.actions[0]
        max_score = -float("inf")
        scores: Dict[str, float] = {}

        for action in self.actions:
            A_inv = _invert_matrix(self.A[action])
            theta = _mat_vec_mul(A_inv, self.b[action])

            mean_reward = _vec_dot(theta, context_vector)
            A_inv_x = _mat_vec_mul(A_inv, context_vector)
            variance = max(0.0, _vec_dot(context_vector, A_inv_x))
            ucb_score = mean_reward + self.alpha * math.sqrt(variance)

            scores[action] = round(ucb_score, 4)
            if ucb_score > max_score:
                max_score = ucb_score
                best_action = action

        self.decision_count += 1
        return best_action, max_score, scores

    def update_feedback(self, action: str, context_vector: List[float], reward: float) -> None:
        if action not in self.A:
            return

        outer = _outer_product(context_vector, context_vector)
        self.A[action] = _mat_add(self.A[action], outer)

        for i in range(self.d):
            self.b[action][i] += reward * context_vector[i]

        if self.decision_count % 20 == 0:
            self._save_model()

    def compute_combat_reward(
        self,
        action_taken: str,
        hp_end_pct: float,
        damage_dealt: float,
        monsters_cleared: int,
        is_player_slain: bool
    ) -> float:
        if is_player_slain:
            return -15.0

        reward = 0.0
        if hp_end_pct > 70.0:
            reward += 1.5
        elif hp_end_pct < 30.0:
            reward -= 2.0

        reward += min(2.0, monsters_cleared * 0.5)
        reward += min(2.0, damage_dealt / 20000.0)

        if action_taken == "AGGRESSIVE_BURST" and hp_end_pct > 50.0:
            reward += 1.0
        elif action_taken in ["KITE_STABILIZE", "RETREAT_SAFE"] and hp_end_pct >= 40.0:
            reward += 0.8

        return round(reward, 3)

    def _save_model(self) -> None:
        os.makedirs(os.path.dirname(self.model_path) or ".", exist_ok=True)
        try:
            data = {
                "decision_count": self.decision_count,
                "actions": self.actions,
                "A": self.A,
                "b": self.b
            }
            with open(self.model_path, "w", encoding="utf-8") as f:
                json.dump(data, f)
        except Exception as e:
            logger.error(f"[ContextualBandit] Lỗi lưu model: {e}")

    def _load_model(self) -> None:
        if os.path.exists(self.model_path):
            try:
                with open(self.model_path, "r", encoding="utf-8") as f:
                    data = json.load(f)
                    self.decision_count = data.get("decision_count", 0)
                    for a in self.actions:
                        if a in data.get("A", {}):
                            self.A[a] = data["A"][a]
                        if a in data.get("b", {}):
                            self.b[a] = data["b"][a]
            except Exception as e:
                logger.error(f"[ContextualBandit] Lỗi đọc model: {e}")
