"""
⚔️ 战争兵棋推演 · War Chess
基于 pygame 的六边形格战争策略游戏
功能：六边形地图、多兵种、地形系统、回合制战斗、移动/攻击/AI
"""

import pygame
import random
import math
import sys
from enum import IntEnum
from dataclasses import dataclass, field
from typing import List, Tuple, Optional, Set, Dict
from collections import deque

# ══════════════════════════════════════════════
# 常量与配置
# ══════════════════════════════════════════════

# 屏幕
SCREEN_W = 1280
SCREEN_H = 800
FPS = 60

# 六边形网格（轴向坐标 axial coordinates）
TILE_SIZE = 48          # 六边形外接圆半径
GRID_W = 18             # 列数 (q)
GRID_H = 14             # 行数 (r)
GRID_ORIGIN_X = 80      # 网格左上角偏移
GRID_ORIGIN_Y = 120

# 颜色
COLORS = {
    "bg":          (30, 32, 40),
    "panel":       (45, 48, 58),
    "panel_border":(70, 75, 90),
    "text":        (230, 230, 235),
    "text_dim":    (160, 165, 175),
    "accent":      (80, 180, 255),
    "red":         (220, 70, 60),
    "red_dim":     (160, 50, 45),
    "blue":        (60, 130, 220),
    "blue_dim":    (45, 95, 160),
    "green":       (60, 180, 100),
    "yellow":      (230, 200, 60),
    "orange":      (230, 140, 50),
    "white":       (245, 245, 245),
    "black":       (20, 20, 25),
    "highlight":   (255, 255, 100),
    "attack":      (255, 80, 80),
    "move":        (80, 220, 120),
    "selected":    (100, 200, 255),
}

# 地形颜色映射
TERRAIN_COLORS = {
    0: (120, 180, 80),    # PLAIN   草绿
    1: (80, 140, 60),     # FOREST  深绿
    2: (160, 140, 100),   # MOUNTAIN 棕褐
    3: (60, 130, 200),    # WATER   蓝
    4: (200, 190, 150),   # DESERT  沙黄
    5: (100, 100, 110),   # ROAD    灰
    6: (90, 90, 100),     # CITY    深灰
}

# ══════════════════════════════════════════════
# 枚举
# ══════════════════════════════════════════════

class TerrainType(IntEnum):
    PLAIN    = 0
    FOREST   = 1
    MOUNTAIN = 2
    WATER    = 3
    DESERT   = 4
    ROAD     = 5
    CITY     = 6

class UnitType(IntEnum):
    INFANTRY  = 0   # 步兵
    TANK      = 1   # 坦克
    ARTILLERY = 2   # 炮兵
    RECON     = 3   # 侦察兵
    COMMANDER = 4   # 指挥官

class Faction(IntEnum):
    RED  = 0
    BLUE = 1

class Phase(IntEnum):
    MOVE   = 0
    ATTACK = 1
    END    = 2

# ══════════════════════════════════════════════
# 六边形坐标（轴向坐标 axial）
# ══════════════════════════════════════════════

@dataclass(frozen=True)
class HexCoord:
    """六边形轴向坐标 (q, r)"""
    q: int
    r: int

    def __add__(self, other):
        return HexCoord(self.q + other.q, self.r + other.r)

    def __sub__(self, other):
        return HexCoord(self.q - other.q, self.r - other.r)

    def neighbors(self) -> List['HexCoord']:
        """六方向邻居"""
        dirs = [(1,0), (1,-1), (0,-1), (-1,0), (-1,1), (0,1)]
        return [HexCoord(self.q + dq, self.r + dr) for dq, dr in dirs]

    def to_cube(self) -> Tuple[int, int, int]:
        """轴向 → 立方体坐标"""
        x = self.q
        z = self.r
        y = -x - z
        return (x, y, z)

# ── 六边形工具函数 ──

def hex_distance(a: HexCoord, b: HexCoord) -> int:
    """两六边形之间的格数距离"""
    ax, ay, az = a.to_cube()
    bx, by, bz = b.to_cube()
    return (abs(ax - bx) + abs(ay - by) + abs(az - bz)) // 2

def hex_to_pixel(h: HexCoord) -> Tuple[int, int]:
    """六边形坐标 → 像素坐标（尖顶方向）"""
    x = TILE_SIZE * math.sqrt(3) * (h.q + h.r / 2)
    y = TILE_SIZE * 3 / 2 * h.r
    return (int(x + GRID_ORIGIN_X + SCREEN_W * 0.22),
            int(y + GRID_ORIGIN_Y))

def pixel_to_hex(px: int, py: int) -> HexCoord:
    """像素坐标 → 最近的六边形坐标"""
    px -= GRID_ORIGIN_X + SCREEN_W * 0.22
    py -= GRID_ORIGIN_Y
    q = (math.sqrt(3)/3 * px - 1/3 * py) / TILE_SIZE
    r = (2/3 * py) / TILE_SIZE
    return hex_round(q, r)

def hex_round(qf: float, rf: float) -> HexCoord:
    """浮点轴向坐标 → 最近的整数坐标"""
    xf = qf
    zf = rf
    yf = -xf - zf
    rx = round(xf)
    ry = round(yf)
    rz = round(zf)
    dx = abs(rx - xf)
    dy = abs(ry - yf)
    dz = abs(rz - zf)
    if dx > dy and dx > dz:
        rx = -ry - rz
    elif dy > dz:
        ry = -rx - rz
    else:
        rz = -rx - ry
    return HexCoord(rx, rz)

# ══════════════════════════════════════════════
# 地形系统
# ══════════════════════════════════════════════

# 地形属性表
TERRAIN_DATA = {
    TerrainType.PLAIN:    {"name": "平原", "move_cost": 1, "defense": 0,  "attack_mod": 0,   "vision": 0},
    TerrainType.FOREST:   {"name": "森林", "move_cost": 2, "defense": 25, "attack_mod": -10, "vision": -1},
    TerrainType.MOUNTAIN: {"name": "山地", "move_cost": 3, "defense": 40, "attack_mod": -20, "vision": 2},
    TerrainType.WATER:    {"name": "水域", "move_cost": 99,"defense": 0,  "attack_mod": 0,   "vision": 0},
    TerrainType.DESERT:   {"name": "沙漠", "move_cost": 2, "defense": 5,  "attack_mod": 0,   "vision": 1},
    TerrainType.ROAD:     {"name": "道路", "move_cost": 1, "defense": 0,  "attack_mod": 0,   "vision": 0},
    TerrainType.CITY:     {"name": "城市", "move_cost": 1, "defense": 35, "attack_mod": 10,  "vision": 1},
}

def get_move_cost(terrain: TerrainType) -> int:
    return TERRAIN_DATA[terrain]["move_cost"]

def get_defense_bonus(terrain: TerrainType) -> int:
    return TERRAIN_DATA[terrain]["defense"]

def get_attack_modifier(terrain: TerrainType) -> int:
    return TERRAIN_DATA[terrain]["attack_mod"]

def get_terrain_name(terrain: TerrainType) -> str:
    return TERRAIN_DATA[terrain]["name"]

# ══════════════════════════════════════════════
# 兵种系统
# ══════════════════════════════════════════════

@dataclass
class UnitStats:
    """兵种基础属性"""
    name: str
    max_hp: int
    attack: int          # 基础攻击力
    defense: int         # 基础防御力
    movement: int        # 移动力
    attack_range: int    # 攻击范围（格）
    vision: int          # 视野
    cost: int            # 生产/分值
    can_counter: bool = True   # 能否反击
    can_move_and_attack: bool = False  # 能否移动后攻击

UNIT_DATA = {
    UnitType.INFANTRY:  UnitStats("步兵",   100, 25, 10, 3, 1, 2, 100, True,  False),
    UnitType.TANK:      UnitStats("坦克",   180, 45, 25, 4, 1, 3, 300, True,  False),
    UnitType.ARTILLERY: UnitStats("炮兵",   120, 55, 5,  2, 3, 2, 250, False, False),
    UnitType.RECON:     UnitStats("侦察兵", 70, 15, 8,  5, 1, 4, 150, True,  True),
    UnitType.COMMANDER: UnitStats("指挥官", 150, 30, 15, 3, 1, 3, 200, True,  False),
}

# ══════════════════════════════════════════════
# 单位类
# ══════════════════════════════════════════════

class Unit:
    """战场单位"""

    def __init__(self, unit_type: UnitType, faction: Faction, pos: HexCoord, name: str = ""):
        self.type = unit_type
        self.faction = faction
        self.pos = pos
        self.stats = UNIT_DATA[unit_type]
        self.max_hp = self.stats.max_hp
        self.hp = self.max_hp
        self.movement_left = self.stats.movement
        self.has_attacked = False
        self.has_moved = False
        self.is_alive_flag = True
        self.name = name or self.stats.name
        self.experience = 0       # 经验值
        self.level = 1
        self.attack_count = 0      # 本回合攻击次数
        self.kills = 0
        self.id = id(self)         # 唯一标识

    @property
    def is_alive(self) -> bool:
        return self.is_alive_flag and self.hp > 0

    @property
    def can_move(self) -> bool:
        return self.is_alive and self.movement_left > 0 and not self.has_moved

    @property
    def can_attack(self) -> bool:
        if not self.is_alive or self.has_attacked:
            return False
        if self.stats.can_move_and_attack:
            return True
        return not self.has_moved  # 大多数兵种移动后不能攻击

    def take_damage(self, dmg: int) -> int:
        """承受伤害，返回实际伤害"""
        actual = min(dmg, self.hp)
        self.hp -= actual
        if self.hp <= 0:
            self.is_alive_flag = False
        return actual

    def heal(self, amount: int):
        self.hp = min(self.max_hp, self.hp + amount)

    def reset_turn(self):
        """回合开始时重置"""
        self.movement_left = self.stats.movement
        self.has_attacked = False
        self.has_moved = False
        self.attack_count = 0

    def get_attack_power(self) -> int:
        """当前攻击力（含等级加成）"""
        return self.stats.attack + (self.level - 1) * 5

    def get_defense_power(self) -> int:
        """当前防御力（含等级加成）"""
        return self.stats.defense + (self.level - 1) * 3

    def gain_exp(self, amount: int):
        self.experience += amount
        # 每100经验升一级
        while self.experience >= self.level * 100:
            self.experience -= self.level * 100
            self.level += 1
            self.max_hp += 10
            self.hp += 10

    def __repr__(self):
        return f"[{self.name} HP:{self.hp}/{self.max_hp} pos:{self.pos}]"


# ══════════════════════════════════════════════
# 地图系统
# ════════════════════════════════════════════

@dataclass
class Tile:
    """单个六边形地块"""
    pos: HexCoord
    terrain: TerrainType
    occupied_by: Optional[Unit] = None

    @property
    def defense_bonus(self) -> int:
        return get_defense_bonus(self.terrain)

    @property
    def move_cost(self) -> int:
        return get_move_cost(self.terrain)

    @property
    def is_water(self) -> bool:
        return self.terrain == TerrainType.WATER


class GameMap:
    """六边形战场地图"""

    def __init__(self, width: int, height: int, seed: Optional[int] = None):
        self.width = width
        self.height = height
        self.tiles: Dict[HexCoord, Tile] = {}
        self.units: List[Unit] = []
        self._rng = random.Random(seed)

        self._generate_map()

    def _generate_map(self):
        """生成地图（带地形分布）"""
        for r in range(self.height):
            for q in range(self.width):
                pos = HexCoord(q, r)
                terrain = self._roll_terrain(q, r)
                self.tiles[pos] = Tile(pos=pos, terrain=terrain)

        # 后处理：确保连通性，修整水域
        self._smooth_terrain()

    def _roll_terrain(self, q: int, r: int) -> TerrainType:
        """按位置和随机性决定地形"""
        rn = self._rng.random()

        # 边缘更可能是水域
        edge_dist = min(q, r, self.width - 1 - q, self.height - 1 - r)
        if edge_dist <= 1 and rn < 0.35:
            return TerrainType.WATER

        # 中心区域更可能是平原/道路
        center_q, center_r = self.width / 2, self.height / 2
        dist_from_center = math.sqrt((q - center_q)**2 + (r - center_r)**2)

        if dist_from_center < 2 and rn < 0.2:
            return TerrainType.CITY

        if rn < 0.15:
            return TerrainType.FOREST
        elif rn < 0.22:
            return TerrainType.MOUNTAIN
        elif rn < 0.25:
            return TerrainType.DESERT
        elif rn < 0.28:
            return TerrainType.ROAD
        else:
            return TerrainType.PLAIN

    def _smooth_terrain(self):
        """平滑地形，避免零散水域"""
        for _ in range(2):
            new_terrain = {}
            for pos, tile in self.tiles.items():
                neighbors = [self.tiles[n].terrain for n in pos.neighbors() if n in self.tiles]
                if tile.terrain == TerrainType.WATER and neighbors.count(TerrainType.WATER) <= 1:
                    # 孤立水域变平原
                    new_terrain[pos] = TerrainType.PLAIN
                else:
                    new_terrain[pos] = tile.terrain
            for pos, t in new_terrain.items():
                self.tiles[pos].terrain = t

    def in_bounds(self, pos: HexCoord) -> bool:
        return 0 <= pos.q < self.width and 0 <= pos.r < self.height

    def get_tile(self, pos: HexCoord) -> Optional[Tile]:
        return self.tiles.get(pos)

    def is_passable(self, pos: HexCoord, unit: Optional[Unit] = None) -> bool:
        """检查地块是否可通过"""
        tile = self.get_tile(pos)
        if tile is None:
            return False
        if tile.terrain == TerrainType.WATER:
            return False  # 水域不可通过
        if tile.occupied_by is not None:
            return False
        return True

    def add_unit(self, unit: Unit):
        tile = self.get_tile(unit.pos)
        if tile:
            tile.occupied_by = unit
            unit.pos = tile.pos
        if unit not in self.units:
            self.units.append(unit)

    def remove_unit(self, unit: Unit):
        tile = self.get_tile(unit.pos)
        if tile and tile.occupied_by == unit:
            tile.occupied_by = None
        if unit in self.units:
            self.units.remove(unit)

    def get_unit_at(self, pos: HexCoord) -> Optional[Unit]:
        tile = self.get_tile(pos)
        return tile.occupied_by if tile else None

    def move_unit(self, unit: Unit, target: HexCoord) -> bool:
        """移动单位到目标格"""
        if not self.is_passable(target, unit):
            return False
        old_tile = self.get_tile(unit.pos)
        if old_tile:
            old_tile.occupied_by = None
        unit.pos = target
        new_tile = self.get_tile(target)
        if new_tile:
            new_tile.occupied_by = unit
        unit.has_moved = True
        unit.movement_left -= self.get_tile(target).move_cost
        return True

    def get_reachable(self, unit: Unit) -> Set[HexCoord]:
        """BFS 搜索可到达的所有格子"""
        if unit.movement_left <= 0:
            return {unit.pos}

        visited = {unit.pos: 0}
        queue = deque([(unit.pos, 0)])
        result = {unit.pos}  # 包含起始格

        while queue:
            current, cost = queue.popleft()
            for nxt in current.neighbors():
                tile = self.get_tile(nxt)
                if tile is None or tile.terrain == TerrainType.WATER:
                    continue
                new_cost = cost + tile.move_cost
                if new_cost <= unit.movement_left and nxt not in visited:
                    visited[nxt] = new_cost
                    result.add(nxt)
                    queue.append((nxt, new_cost))

        return result

    def get_attack_targets(self, unit: Unit) -> List[Unit]:
        """获取可攻击的敌方单位列表"""
        targets = []
        rng = unit.stats.attack_range
        for other in self.units:
            if other.faction != unit.faction and other.is_alive:
                if hex_distance(unit.pos, other.pos) <= rng:
                    targets.append(other)
        return targets

    def can_attack(self, attacker: Unit, target: Unit) -> bool:
        """检查能否攻击"""
        if not attacker.can_attack:
            return False
        if target.faction == attacker.faction:
            return False
        dist = hex_distance(attacker.pos, target.pos)
        return dist <= attacker.stats.attack_range

    def resolve_combat(self, attacker: Unit, defender: Unit) -> dict:
        """执行战斗，返回战斗日志"""
        if not self.can_attack(attacker, defender):
            return {"success": False, "reason": "无法攻击"}

        log = {"attacker": attacker, "defender": defender, "events": []}

        # 攻击者命中判定
        atk_tile = self.get_tile(attacker.pos)
        def_tile = self.get_tile(defender.pos)

        atk_terrain_mod = get_attack_modifier(atk_tile.terrain) if atk_tile else 0
        def_terrain_bonus = get_defense_bonus(def_tile.terrain) if def_tile else 0

        # 基础命中率
        hit_chance = 70 + attacker.stats.attack // 5 + atk_terrain_mod - def_terrain_bonus // 3
        hit_chance = max(20, min(95, hit_chance))

        roll = self._rng.randint(1, 100)
        if roll > hit_chance:
            log["events"].append(f"未命中 (掷骰{roll} > {hit_chance})")
            attacker.has_attacked = True
            attacker.attack_count += 1
            return log

        # 计算伤害
        damage = calculate_damage(
            attacker.type, defender.type,
            atk_tile.terrain if atk_tile else TerrainType.PLAIN,
            def_tile.terrain if def_tile else TerrainType.PLAIN,
            attacker.get_attack_power(),
            defender.get_defense_power(),
            self._rng
        )

        actual_dmg = defender.take_damage(damage)
        log["events"].append(f"命中! 造成 {actual_dmg} 点伤害")

        attacker.gain_exp(actual_dmg)
        attacker.has_attacked = True
        attacker.attack_count += 1
        attacker.movement_left = 0  # 攻击后不能再移动

        if not defender.is_alive:
            attacker.kills += 1
            attacker.gain_exp(50)
            log["events"].append(f"{defender.name} 被击毁!")
            self.remove_unit(defender)
        else:
            # 反击
            if defender.stats.can_counter and attacker.is_alive:
                counter_hit = 50 + defender.stats.attack // 6 - atk_terrain_mod // 3
                counter_hit = max(15, min(85, counter_hit))
                if self._rng.randint(1, 100) <= counter_hit:
                    counter_dmg = calculate_damage(
                        defender.type, attacker.type,
                        def_tile.terrain if def_tile else TerrainType.PLAIN,
                        atk_tile.terrain if atk_tile else TerrainType.PLAIN,
                        defender.get_attack_power(),
                        attacker.get_defense_power(),
                        self._rng
                    )
                    counter_actual = attacker.take_damage(counter_dmg)
                    log["events"].append(f"反击造成 {counter_actual} 点伤害")
                    if not attacker.is_alive:
                        log["events"].append(f"{attacker.name} 被反击击毁!")
                        self.remove_unit(attacker)

        log["success"] = True
        return log

    def get_units_in_range(self, pos: HexCoord, range_val: int, faction: Optional[Faction] = None) -> List[Unit]:
        """获取范围内的单位"""
        result = []
        for u in self.units:
            if u.is_alive and hex_distance(pos, u.pos) <= range_val:
                if faction is None or u.faction == faction:
                    result.append(u)
        return result


# ══════════════════════════════════════════════
# 战斗计算
# ════════════════════════════════════════════

def calculate_damage(
    atk_type: UnitType, def_type: UnitType,
    atk_terrain: TerrainType, def_terrain: TerrainType,
    atk_power: int, def_power: int,
    rng: Optional[random.Random] = None
) -> int:
    """计算攻击伤害"""
    if rng is None:
        rng = random.Random()

    # 兵种克制
    type_mod = 1.0
    if atk_type == UnitType.TANK and def_type == UnitType.INFANTRY:
        type_mod = 1.3
    elif atk_type == UnitType.INFANTRY and def_type == UnitType.TANK:
        type_mod = 0.7
    elif atk_type == UnitType.ARTILLERY and def_type == UnitType.TANK:
        type_mod = 1.5
    elif atk_type == UnitType.ARTILLERY and def_type == UnitType.INFANTRY:
        type_mod = 0.8
    elif atk_type == UnitType.RECON and def_type == UnitType.ARTILLERY:
        type_mod = 1.3
    elif atk_type == UnitType.COMMANDER:
        type_mod = 1.1  # 指挥官有光环加成

    # 地形修正
    atk_mod = 1.0 + get_attack_modifier(atk_terrain) / 100
    def_mod = 1.0 - get_defense_bonus(def_terrain) / 100

    # 基础公式
    base = atk_power * type_mod * atk_mod - def_power * def_mod * 0.5
    base = max(5, base)  # 最低伤害

    # 随机波动 ±20%
    variance = rng.uniform(0.8, 1.2)
    damage = int(base * variance)

    return max(1, damage)


def roll_dice(num: int, sides: int, rng: Optional[random.Random] = None) -> int:
    """掷骰子"""
    if rng is None:
        rng = random.Random()
    return sum(rng.randint(1, sides) for _ in range(num))


# ══════════════════════════════════════════════
# 游戏状态管理
# ════════════════════════════════════════════

class GameState:
    """全局游戏状态"""

    def __init__(self, game_map: GameMap):
        self.map = game_map
        self.turn = 1
        self.current_faction = Faction.RED
        self.phase: str = "move"  # move / attack / end
        self.selected_unit: Optional[Unit] = None
        self.reachable: Set[HexCoord] = set()
        self.attack_targets: List[Unit] = []
        self.battle_log: List[str] = []
        self.score = {Faction.RED: 0, Faction.BLUE: 0}
        self.game_over = False
        self.winner: Optional[Faction] = None
        self.action_points = {Faction.RED: 10, Faction.BLUE: 10}
        self.max_ap = 10

    def get_units(self, faction: Faction) -> List[Unit]:
        return [u for u in self.map.units if u.faction == faction and u.is_alive]

    def end_turn(self):
        """结束当前回合"""
        # 切换阵营
        self.current_faction = Faction.BLUE if self.current_faction == Faction.RED else Faction.RED
        self.turn += 1
        self.phase = "move"
        self.selected_unit = None
        self.reachable = set()
        self.attack_targets = []

        # 重置新回合阵营的单位
        for u in self.get_units(self.current_faction):
            u.reset_turn()

        # 恢复行动力
        self.action_points[self.current_faction] = self.max_ap

        self.log(f"── 第 {self.turn} 回合 · {'红方' if self.current_faction == Faction.RED else '蓝方'}行动 ──")

    def check_victory(self) -> Optional[Faction]:
        """检查胜负"""
        red_alive = len(self.get_units(Faction.RED))
        blue_alive = len(self.get_units(Faction.BLUE))

        if red_alive == 0 and blue_alive == 0:
            self.game_over = True
            self.winner = None  # 平局
        elif red_alive == 0:
            self.game_over = True
            self.winner = Faction.BLUE
        elif blue_alive == 0:
            self.game_over = True
            self.winner = Faction.RED

        return self.winner

    def log(self, msg: str):
        self.battle_log.append(msg)
        if len(self.battle_log) > 100:
            self.battle_log = self.battle_log[-100:]

    def select_unit(self, unit: Optional[Unit]):
        self.selected_unit = unit
        if unit and unit.can_move:
            self.reachable = self.map.get_reachable(unit)
        else:
            self.reachable = set()
        if unit:
            self.attack_targets = self.map.get_attack_targets(unit)
        else:
            self.attack_targets = []

    def try_move(self, target: HexCoord) -> bool:
        """尝试移动选中单位"""
        if self.selected_unit is None:
            return False
        unit = self.selected_unit
        if target not in self.reachable:
            return False
        if self.map.move_unit(unit, target):
            self.log(f"{unit.name} 移动至 {target}")
            # 移动后更新可达和攻击目标
            self.reachable = self.map.get_reachable(unit)
            self.attack_targets = self.map.get_attack_targets(unit)
            return True
        return False

    def try_attack(self, target_unit: Unit) -> Optional[dict]:
        """尝试攻击"""
        if self.selected_unit is None:
            return None
        attacker = self.selected_unit
        if not self.map.can_attack(attacker, target_unit):
            return None
        result = self.map.resolve_combat(attacker, target_unit)
        for evt in result.get("events", []):
            self.log(f"  {evt}")
        # 更新攻击目标
        self.attack_targets = self.map.get_attack_targets(attacker)
        # 检查胜负
        self.check_victory()
        return result


# ══════════════════════════════════════════════
# AI 系统
# ════════════════════════════════════════════

class AIController:
    """简单的 AI 控制器"""

    def __init__(self, game_state: GameState, faction: Faction):
        self.state = game_state
        self.faction = faction
        self.rng = random.Random(42)

    def take_turn(self):
        """执行 AI 回合"""
        units = self.state.get_units(self.faction)
        self.state.log(f"🤖 { '蓝方' if self.faction == Faction.BLUE else '红方'} AI 行动中...")

        for unit in units:
            if not unit.is_alive:
                continue
            self._process_unit(unit)

        self.state.log("🤖 AI 回合结束")
        self.state.end_turn()

    def _process_unit(self, unit: Unit):
        """处理单个 AI 单位"""
        # 1. 尝试攻击最近的敌人
        enemies = [u for u in self.state.map.units
                   if u.faction != self.faction and u.is_alive]
        if not enemies:
            return

        # 按距离排序
        enemies.sort(key=lambda e: hex_distance(unit.pos, e.pos))

        # 尝试攻击范围内的敌人
        for enemy in enemies:
            if self.state.map.can_attack(unit, enemy):
                result = self.state.try_attack(enemy)
                if result and result.get("success"):
                    self.state.log(f"  💥 {unit.name} 攻击 {enemy.name}")
                break

        # 2. 如果还能移动，向最近敌人靠近
        if unit.can_move and enemies:
            nearest = min(enemies, key=lambda e: hex_distance(unit.pos, e.pos))
            path = self._find_path_towards(unit, nearest.pos)
            if path:
                for step in path[:unit.movement_left]:
                    if step in self.state.map.get_reachable(unit):
                        self.state.map.move_unit(unit, step)
                    else:
                        break

    def _find_path_towards(self, unit: Unit, target: HexCoord) -> List[HexCoord]:
        """简单贪心寻路：每步向目标靠近"""
        current = unit.pos
        path = []
        visited = {current}
        max_steps = unit.movement_left

        for _ in range(max_steps):
            neighbors = [n for n in current.neighbors()
                         if self.state.map.in_bounds(n)
                         and self.state.map.is_passable(n, unit)
                         and n not in visited]
            if not neighbors:
                break
            # 选距离目标最近的
            neighbors.sort(key=lambda n: hex_distance(n, target))
            best = neighbors[0]
            if hex_distance(best, target) >= hex_distance(current, target):
                break  # 无法更接近
            path.append(best)
            visited.add(best)
            current = best

        return path


# ══════════════════════════════════════════════
# 渲染系统（pygame）
# ════════════════════════════════════════════

class Renderer:
    """游戏渲染器"""

    def __init__(self, screen: pygame.Surface, game_state: GameState):
        self.screen = screen
        self.state = game_state
        self.font_small = pygame.font.Font(None, 16)
        self.font_normal = pygame.font.Font(None, 20)
        self.font_large = pygame.font.Font(None, 28)
        self.font_title = pygame.font.Font(None, 36)
        self.font_tiny = pygame.font.Font(None, 13)

        # 预计算六边形顶点
        self.hex_vertices = self._make_hex_vertices()

        # 动画
        self.anim_time = 0

    def _make_hex_vertices(self) -> List[Tuple[int, int]]:
        """预计算六边形六个顶点（相对中心）"""
        verts = []
        for i in range(6):
            angle = math.pi / 180 * (60 * i - 30)  # 尖顶
            x = TILE_SIZE * math.cos(angle)
            y = TILE_SIZE * math.sin(angle)
            verts.append((int(x), int(y)))
        return verts

    def get_hex_corners(self, center: Tuple[int, int]) -> List[Tuple[int, int]]:
        """获取六边形在世界坐标中的六个角"""
        cx, cy = center
        return [(cx + vx, cy + vy) for vx, vy in self.hex_vertices]

    def draw_hex(self, surface: pygame.Surface, center: Tuple[int, int],
                 fill_color: Tuple[int, int, int], border_color: Tuple[int, int, int],
                 border_width: int = 1, alpha: int = 255):
        """绘制单个六边形"""
        corners = self.get_hex_corners(center)
        if alpha < 255:
            # 使用临时 surface 实现半透明
            temp = pygame.Surface((TILE_SIZE * 2, TILE_SIZE * 2), pygame.SRCALPHA)
            temp_center = (TILE_SIZE, TILE_SIZE)
            temp_corners = [(c[0] - center[0] + TILE_SIZE, c[1] - center[1] + TILE_SIZE) for c in corners]
            pygame.draw.polygon(temp, (*fill_color, alpha), temp_corners)
            pygame.draw.polygon(temp, (*border_color, alpha), temp_corners, border_width)
            surface.blit(temp, (center[0] - TILE_SIZE, center[1] - TILE_SIZE))
        else:
            pygame.draw.polygon(surface, fill_color, corners)
            pygame.draw.polygon(surface, border_color, corners, border_width)

    def draw_tile(self, surface: pygame.Surface, tile: Tile):
        """绘制地块"""
        center = hex_to_pixel(tile.pos)
        base_color = TERRAIN_COLORS[tile.terrain.value]

        # 轻微的颜色变化
        variant = (tile.pos.q * 7 + tile.pos.r * 13) % 10
        color = tuple(max(0, min(255, c + variant - 5)) for c in base_color)

        # 选中/高亮
        is_selected = (self.state.selected_unit
                       and self.state.selected_unit.pos == tile.pos)
        is_reachable = tile.pos in self.state.reachable
        is_attack_target = any(t.pos == tile.pos for t in self.state.attack_targets)

        border_color = COLORS["panel_border"]
        border_width = 1

        if is_selected:
            border_color = COLORS["selected"]
            border_width = 3
        elif is_attack_target:
            border_color = COLORS["attack"]
            border_width = 2
        elif is_reachable:
            border_color = COLORS["move"]
            border_width = 2

        self.draw_hex(surface, center, color, border_color, border_width)

        # 地形图标
        terrain_name = get_terrain_name(tile.terrain)
        if tile.terrain == TerrainType.MOUNTAIN:
            self._draw_mountain(surface, center)
        elif tile.terrain == TerrainType.FOREST:
            self._draw_tree(surface, center)
        elif tile.terrain == TerrainType.WATER:
            self._draw_wave(surface, center)
        elif tile.terrain == TerrainType.CITY:
            self._draw_city(surface, center)
        elif tile.terrain == TerrainType.ROAD:
            self._draw_road(surface, center)
        elif tile.terrain == TerrainType.DESERT:
            self._draw_sand(surface, center)

        # 坐标文字（小字）
        coord_text = f"{tile.pos.q},{tile.pos.r}"
        txt = self.font_tiny.render(coord_text, True, (255, 255, 255, 128))
        txt.set_alpha(80)
        surface.blit(txt, (center[0] - 12, center[1] + TILE_SIZE // 2 - 8))

    def _draw_mountain(self, surface, center):
        pts = [(center[0], center[1] - 12), (center[0] - 10, center[1] + 5),
               (center[0] + 10, center[1] + 5)]
        pygame.draw.polygon(surface, (140, 120, 100), pts)
        snow_pts = [(center[0], center[1] - 12), (center[0] - 4, center[1] - 5),
                    (center[0] + 4, center[1] - 5)]
        pygame.draw.polygon(surface, (240, 240, 240), snow_pts)

    def _draw_tree(self, surface, center):
        pygame.draw.circle(surface, (40, 100, 40), (center[0], center[1] - 3), 8)
        pygame.draw.circle(surface, (50, 120, 50), (center[0] - 4, center[1] + 2), 6)
        pygame.draw.circle(surface, (45, 110, 45), (center[0] + 4, center[1] + 2), 6)

    def _draw_wave(self, surface, center):
        for i in range(3):
            offset = math.sin(self.anim_time * 0.05 + i) * 2
            y = center[1] - 4 + i * 4
            pygame.draw.line(surface, (180, 220, 255), (center[0] - 10, int(y + offset)),
                            (center[0] + 10, int(y + offset)), 2)

    def _draw_city(self, surface, center):
        pygame.draw.rect(surface, (120, 120, 130), (center[0] - 8, center[1] - 6, 16, 14))
        pygame.draw.polygon(surface, (100, 100, 110),
                           [(center[0] - 8, center[1] - 6), (center[0], center[1] - 12),
                            (center[0] + 8, center[1] - 6)])
        # 窗户
        for wx in range(-5, 6, 5):
            pygame.draw.rect(surface, (255, 255, 150), (center[0] + wx, center[1] - 2, 2, 3))

    def _draw_road(self, surface, center):
        pygame.draw.line(surface, (180, 180, 180), (center[0] - 12, center[1]),
                        (center[0] + 12, center[1]), 3)
        pygame.draw.line(surface, (180, 180, 180), (center[0], center[1] - 10),
                        (center[0], center[1] + 10), 3)

    def _draw_sand(self, surface, center):
        for i in range(4):
            angle = self.anim_time * 0.02 + i * 1.5
            r = 6 + i * 2
            x = center[0] + int(r * math.cos(angle))
            y = center[1] + int(r * math.sin(angle))
            pygame.draw.circle(surface, (210, 200, 160), (x, y), 2)

    def draw_unit(self, surface: pygame.Surface, unit: Unit):
        """绘制单位"""
        center = hex_to_pixel(unit.pos)

        # 阵营颜色
        if unit.faction == Faction.RED:
            main_color = COLORS["red"]
            dim_color = COLORS["red_dim"]
        else:
            main_color = COLORS["blue"]
            dim_color = COLORS["blue_dim"]

        # 是否被选中
        is_selected = self.state.selected_unit == unit
        pulse = math.sin(self.anim_time * 0.08) * 3 if is_selected else 0

        # 单位底盘（圆形）
        radius = TILE_SIZE // 2 - 6
        pygame.draw.circle(surface, dim_color, center, radius + 1)
        pygame.draw.circle(surface, main_color, center, radius - 1)

        # 兵种图标
        self._draw_unit_icon(surface, unit, center)

        # HP 条
        hp_ratio = unit.hp / unit.max_hp
        bar_w = radius * 2
        bar_h = 4
        bar_x = center[0] - radius
        bar_y = center[1] + radius + 2
        pygame.draw.rect(surface, (60, 60, 60), (bar_x, bar_y, bar_w, bar_h))
        hp_color = COLORS["green"] if hp_ratio > 0.5 else COLORS["yellow"] if hp_ratio > 0.25 else COLORS["red"]
        pygame.draw.rect(surface, hp_color, (bar_x, bar_y, int(bar_w * hp_ratio), bar_h))

        # 等级标记
        if unit.level > 1:
            lvl_text = self.font_tiny.render(f"Lv{unit.level}", True, COLORS["white"])
            surface.blit(lvl_text, (center[0] + radius - 12, center[1] - radius - 2))

        # 可移动/可攻击指示
        if unit.faction == self.state.current_faction:
            if unit.can_move:
                pygame.draw.circle(surface, COLORS["move"], (center[0] + radius - 2, center[1] - radius + 2), 3)
            if unit.can_attack:
                pygame.draw.circle(surface, COLORS["attack"], (center[0] + radius - 2, center[1] - radius + 8), 3)

        # 选中光环
        if is_selected:
            ring_r = radius + 4 + int(pulse)
            pygame.draw.circle(surface, COLORS["highlight"], center, ring_r, 2)

    def _draw_unit_icon(self, surface, unit, center):
        """绘制兵种图标"""
        cx, cy = center
        if unit.type == UnitType.INFANTRY:
            # 步兵：小三角形
            pts = [(cx, cy - 8), (cx - 6, cy + 5), (cx + 6, cy + 5)]
            pygame.draw.polygon(surface, COLORS["white"], pts)
        elif unit.type == UnitType.TANK:
            # 坦克：矩形+炮管
            pygame.draw.rect(surface, COLORS["white"], (cx - 7, cy - 5, 14, 10))
            pygame.draw.line(surface, COLORS["white"], (cx + 7, cy), (cx + 14, cy), 3)
            pygame.draw.circle(surface, COLORS["white"], (cx - 3, cy), 3)
        elif unit.type == UnitType.ARTILLERY:
            # 炮兵：大炮形状
            pygame.draw.circle(surface, COLORS["white"], (cx, cy), 5)
            pygame.draw.line(surface, COLORS["white"], (cx, cy), (cx + 10, cy - 8), 3)
        elif unit.type == UnitType.RECON:
            # 侦察：菱形
            pts = [(cx, cy - 9), (cx + 7, cy), (cx, cy + 9), (cx - 7, cy)]
            pygame.draw.polygon(surface, COLORS["white"], pts)
        elif unit.type == UnitType.COMMANDER:
            # 指挥官：星形
            for i in range(5):
                angle = math.pi / 2 + i * 2 * math.pi / 5
                inner = angle + math.pi / 5
                outer_r = 9
                inner_r = 4
                x1, y1 = cx + outer_r * math.cos(angle), cy - outer_r * math.sin(angle)
                x2, y2 = cx + inner_r * math.cos(inner), cy - inner_r * math.sin(inner)
                if i == 0:
                    pts = [(x1, y1)]
                pts.append((x2, y2))
                angle2 = angle + 2 * math.pi / 5
                x3, y3 = cx + outer_r * math.cos(angle2), cy - outer_r * math.sin(angle2)
                pts.append((x3, y3))
            pygame.draw.polygon(surface, COLORS["yellow"], pts)

    def draw_panel(self, surface: pygame.Surface):
        """绘制右侧信息面板"""
        panel_x = SCREEN_W - 300
        panel_w = 300
        panel_rect = pygame.Rect(panel_x, 0, panel_w, SCREEN_H)
        pygame.draw.rect(surface, COLORS["panel"], panel_rect)
        pygame.draw.line(surface, COLORS["panel_border"], (panel_x, 0), (panel_x, SCREEN_H), 2)

        y = 10

        # 标题
        title = self.font_title.render("⚔️ 兵棋推演", True, COLORS["text"])
        surface.blit(title, (panel_x + 20, y))
        y += 45

        # 回合信息
        faction_name = "红方" if self.state.current_faction == Faction.RED else "蓝方"
        faction_color = COLORS["red"] if self.state.current_faction == Faction.RED else COLORS["blue"]
        turn_text = self.font_large.render(f"第 {self.state.turn} 回合", True, COLORS["text"])
        surface.blit(turn_text, (panel_x + 20, y))
        y += 30
        faction_text = self.font_normal.render(f"当前: {faction_name}", True, faction_color)
        surface.blit(faction_text, (panel_x + 20, y))
        y += 25

        # 阶段
        phase_text = self.font_normal.render(f"阶段: {self.state.phase}", True, COLORS["accent"])
        surface.blit(phase_text, (panel_x + 20, y))
        y += 25

        # 行动力
        ap = self.state.action_points[self.state.current_faction]
        ap_text = self.font_normal.render(f"行动力: {ap}/{self.state.max_ap}", True, COLORS["text"])
        surface.blit(ap_text, (panel_x + 20, y))
        y += 30

        # 分隔线
        pygame.draw.line(surface, COLORS["panel_border"], (panel_x + 10, y), (panel_x + panel_w - 10, y), 1)
        y += 10

        # 选中单位信息
        if self.state.selected_unit:
            u = self.state.selected_unit
            u_color = COLORS["red"] if u.faction == Faction.RED else COLORS["blue"]
            name_text = self.font_normal.render(f"{u.name} (Lv{u.level})", True, u_color)
            surface.blit(name_text, (panel_x + 20, y))
            y += 22

            hp_text = self.font_small.render(f"HP: {u.hp}/{u.max_hp}", True, COLORS["text"])
            surface.blit(hp_text, (panel_x + 20, y))
            y += 18

            atk_text = self.font_small.render(f"攻击: {u.get_attack_power()}", True, COLORS["text"])
            surface.blit(atk_text, (panel_x + 20, y))
            y += 18

            def_text = self.font_small.render(f"防御: {u.get_defense_power()}", True, COLORS["text"])
            surface.blit(def_text, (panel_x + 20, y))
            y += 18

            mov_text = self.font_small.render(f"移动力: {u.movement_left}/{u.stats.movement}", True, COLORS["text"])
            surface.blit(mov_text, (panel_x + 20, y))
            y += 18

            rng_text = self.font_small.render(f"射程: {u.stats.attack_range}", True, COLORS["text"])
            surface.blit(rng_text, (panel_x + 20, y))
            y += 18

            exp_text = self.font_small.render(f"经验: {u.experience}", True, COLORS["text_dim"])
            surface.blit(exp_text, (panel_x + 20, y))
            y += 18

            kills_text = self.font_small.render(f"击杀: {u.kills}", True, COLORS["text_dim"])
            surface.blit(kills_text, (panel_x + 20, y))
            y += 20
        else:
            hint = self.font_small.render("点击单位选中", True, COLORS["text_dim"])
            surface.blit(hint, (panel_x + 20, y))
            y += 20

        # 分隔线
        pygame.draw.line(surface, COLORS["panel_border"], (panel_x + 10, y), (panel_x + panel_w - 10, y), 1)
        y += 10

        # 单位统计
        red_units = self.state.get_units(Faction.RED)
        blue_units = self.state.get_units(Faction.BLUE)
        red_alive = sum(1 for u in red_units if u.is_alive)
        blue_alive = sum(1 for u in blue_units if u.is_alive)

        red_text = self.font_normal.render(f"🔴 红方单位: {red_alive}", True, COLORS["red"])
        surface.blit(red_text, (panel_x + 20, y))
        y += 22

        blue_text = self.font_normal.render(f"🔵 蓝方单位: {blue_alive}", True, COLORS["blue"])
        surface.blit(blue_text, (panel_x + 20, y))
        y += 25

        # 分隔线
        pygame.draw.line(surface, COLORS["panel_border"], (panel_x + 10, y), (panel_x + panel_w - 10, y), 1)
        y += 10

        # 战斗日志
        log_title = self.font_normal.render("📜 战斗日志", True, COLORS["text"])
        surface.blit(log_title, (panel_x + 20, y))
        y += 22

        # 只显示最近10条
        recent_logs = self.state.battle_log[-12:]
        for log in recent_logs:
            log_color = COLORS["text_dim"]
            if "击毁" in log or "击杀" in log:
                log_color = COLORS["yellow"]
            elif "未命中" in log:
                log_color = COLORS["text_dim"]
            elif "命中" in log or "攻击" in log:
                log_color = COLORS["orange"]
            log_surf = self.font_tiny.render(log[:38], True, log_color)
            surface.blit(log_surf, (panel_x + 10, y))
            y += 15

        # 底部操作提示
        y = SCREEN_H - 80
        pygame.draw.line(surface, COLORS["panel_border"], (panel_x + 10, y), (panel_x + panel_w - 10, y), 1)
        y += 8

        hints = [
            "左键: 选择/移动/攻击",
            "右键: 取消选择",
            "空格: 结束回合",
            "A: AI 自动行动",
        ]
        for h in hints:
            hs = self.font_tiny.render(h, True, COLORS["text_dim"])
            surface.blit(hs, (panel_x + 20, y))
            y += 16

        # 结束回合按钮
        btn_rect = pygame.Rect(panel_x + 60, SCREEN_H - 45, 180, 32)
        btn_color = COLORS["accent"] if not self.state.game_over else COLORS["panel_border"]
        pygame.draw.rect(surface, btn_color, btn_rect, border_radius=6)
        btn_text = "结束回合" if not self.state.game_over else "游戏结束"
        btn_surf = self.font_normal.render(btn_text, True, COLORS["white"])
        surface.blit(btn_surf, (btn_rect.x + 55, btn_rect.y + 6))

        # 存储按钮区域供事件处理使用
        self.end_turn_btn = btn_rect

    def draw_top_bar(self, surface: pygame.Surface):
        """绘制顶部信息栏"""
        bar_rect = pygame.Rect(0, 0, SCREEN_W - 300, 50)
        pygame.draw.rect(surface, COLORS["panel"], bar_rect)
        pygame.draw.line(surface, COLORS["panel_border"], (0, 50), (SCREEN_W - 300, 50), 2)

        # 分数
        red_score = self.state.score[Faction.RED]
        blue_score = self.state.score[Faction.BLUE]
        score_text = self.font_normal.render(f"🔴 {red_score}  :  {blue_score} 🔵", True, COLORS["text"])
        surface.blit(score_text, (20, 15))

        # 当前地形信息
        if self.state.selected_unit:
            u = self.state.selected_unit
            tile = self.state.map.get_tile(u.pos)
            if tile:
                terrain_info = f"当前地形: {get_terrain_name(tile.terrain)} (防御+{tile.defense_bonus}%)"
                terrain_surf = self.font_small.render(terrain_info, True, COLORS["text_dim"])
                surface.blit(terrain_surf, (200, 18))

    def draw_game_over(self, surface: pygame.Surface):
        """绘制游戏结束画面"""
        if not self.state.game_over:
            return

        overlay = pygame.Surface((SCREEN_W, SCREEN_H), pygame.SRCALPHA)
        overlay.fill((0, 0, 0, 180))
        surface.blit(overlay, (0, 0))

        cx, cy = SCREEN_W // 2, SCREEN_H // 2

        if self.state.winner is None:
            msg = "平局！"
            color = COLORS["yellow"]
        else:
            msg = f"{'🔴 红方' if self.state.winner == Faction.RED else '🔵 蓝方'} 胜利！"
            color = COLORS["red"] if self.state.winner == Faction.RED else COLORS["blue"]

        text = self.font_title.render(msg, True, color)
        surface.blit(text, (cx - text.get_width() // 2, cy - 40))

        hint = self.font_normal.render("按 R 重新开始", True, COLORS["text"])
        surface.blit(hint, (cx - hint.get_width() // 2, cy + 20))

    def draw_attack_preview(self, surface: pygame.Surface, mouse_pos: Tuple[int, int]):
        """绘制攻击预览线"""
        if not self.state.selected_unit:
            return
        unit = self.state.selected_unit
        if not unit.can_attack:
            return

        center = hex_to_pixel(unit.pos)
        pygame.draw.line(surface, COLORS["attack"], center, mouse_pos, 2)

        # 在鼠标位置显示命中率预估
        target_hex = pixel_to_hex(mouse_pos[0], mouse_pos[1])
        target_unit = self.state.map.get_unit_at(target_hex)
        if target_unit and target_unit.faction != unit.faction:
            tile = self.state.map.get_tile(unit.pos)
            atk_mod = get_attack_modifier(tile.terrain) if tile else 0
            def_tile = self.state.map.get_tile(target_unit.pos)
            def_bonus = get_defense_bonus(def_tile.terrain) if def_tile else 0
            hit = max(20, min(95, 70 + unit.stats.attack // 5 + atk_mod - def_bonus // 3))
            hit_text = self.font_small.render(f"命中率: {hit}%", True, COLORS["attack"])
            surface.blit(hit_text, (mouse_pos[0] + 15, mouse_pos[1] - 20))

    def render(self, surface: pygame.Surface):
        """主渲染函数"""
        self.anim_time += 1

        # 背景
        surface.fill(COLORS["bg"])

        # 绘制所有地块
        for tile in self.state.map.tiles.values():
            self.draw_tile(surface, tile)

        # 绘制可达范围高亮（半透明覆盖）
        for pos in self.state.reachable:
            center = hex_to_pixel(pos)
            self.draw_hex(surface, center, (80, 220, 120), (80, 220, 120), 0, 50)

        # 绘制攻击目标高亮
        for target in self.state.attack_targets:
            center = hex_to_pixel(target.pos)
            self.draw_hex(surface, center, (255, 80, 80), (255, 80, 80), 0, 40)

        # 绘制所有单位
        for unit in self.state.map.units:
            if unit.is_alive:
                self.draw_unit(surface, unit)

        # 攻击预览
        mouse_pos = pygame.mouse.get_pos()
        self.draw_attack_preview(surface, mouse_pos)

        # UI
        self.draw_top_bar(surface)
        self.draw_panel(surface)

        # 游戏结束
        self.draw_game_over(surface)


# ══════════════════════════════════════════════
# 输入处理
# ════════════════════════════════════════════

class InputHandler:
    """处理鼠标和键盘输入"""

    def __init__(self, game_state: GameState, renderer: Renderer, screen: pygame.Surface):
        self.state = game_state
        self.renderer = renderer
        self.screen = screen
        self.ai = AIController(game_state, Faction.BLUE)

    def handle_event(self, event: pygame.event.Event):
        if self.state.game_over:
            if event.type == pygame.KEYDOWN and event.key == pygame.K_r:
                self._restart_game()
            return

        if event.type == pygame.MOUSEBUTTONDOWN:
            self._handle_click(event)
        elif event.type == pygame.KEYDOWN:
            self._handle_keydown(event)

    def _handle_click(self, event):
        mx, my = event.pos

        # 检查是否点击了结束回合按钮
        if hasattr(self.renderer, 'end_turn_btn'):
            btn = self.renderer.end_turn_btn
            if btn.collidepoint(mx, my):
                self.state.end_turn()
                return

        # 只处理左键
        if event.button != 1:
            if event.button == 3:  # 右键取消
                self.state.select_unit(None)
            return

        # 转换为六边形坐标
        hex_pos = pixel_to_hex(mx, my)
        tile = self.state.map.get_tile(hex_pos)
        if tile is None:
            return

        clicked_unit = tile.occupied_by

        if self.state.selected_unit is None:
            # 没有选中单位 → 尝试选中
            if clicked_unit and clicked_unit.faction == self.state.current_faction:
                self.state.select_unit(clicked_unit)
                self.state.log(f"选中 {clicked_unit.name} @ {clicked_unit.pos}")
        else:
            selected = self.state.selected_unit

            if clicked_unit and clicked_unit.faction == self.state.current_faction:
                # 点击己方单位 → 切换选中
                self.state.select_unit(clicked_unit)
            elif clicked_unit and clicked_unit.faction != self.state.current_faction:
                # 点击敌方单位 → 尝试攻击
                result = self.state.try_attack(clicked_unit)
                if result and result.get("success"):
                    self.state.log(f"✅ 攻击完成")
                else:
                    self.state.log(f"❌ 无法攻击 (距离/状态不符)")
            elif tile.pos in self.state.reachable:
                # 移动到空地
                self.state.try_move(tile.pos)
            else:
                # 点击无关区域 → 取消选择
                self.state.select_unit(None)

    def _handle_keydown(self, event):
        if event.key == pygame.K_SPACE:
            self.state.end_turn()
            self.state.log("⏭ 回合结束")
        elif event.key == pygame.K_a:
            # AI 自动行动（用于测试）
            self.ai.take_turn()
        elif event.key == pygame.K_ESCAPE:
            self.state.select_unit(None)
        elif event.key == pygame.K_r:
            self._restart_game()

    def _restart_game(self):
        """重新开始游戏"""
        from war_game import create_game  # 避免循环导入
        new_state = create_game()
        self.state.__dict__.update(new_state.__dict__)


# ══════════════════════════════════════════════
# 游戏初始化与预设
# ════════════════════════════════════════════

def create_game(seed: Optional[int] = 42) -> GameState:
    """创建一局新游戏"""
    game_map = GameMap(GRID_W, GRID_H, seed=seed)
    state = GameState(game_map)

    # 红方初始部队（左下角）
    red_setups = [
        (UnitType.COMMANDER, HexCoord(2, 11)),
        (UnitType.INFANTRY,  HexCoord(1, 12)),
        (UnitType.INFANTRY,  HexCoord(2, 12)),
        (UnitType.INFANTRY,  HexCoord(3, 12)),
        (UnitType.TANK,      HexCoord(1, 11)),
        (UnitType.TANK,      HexCoord(3, 11)),
        (UnitType.ARTILLERY, HexCoord(2, 10)),
        (UnitType.RECON,     HexCoord(0, 12)),
    ]

    # 蓝方初始部队（右上角）
    blue_setups = [
        (UnitType.COMMANDER, HexCoord(15, 2)),
        (UnitType.INFANTRY,  HexCoord(14, 1)),
        (UnitType.INFANTRY,  HexCoord(15, 1)),
        (UnitType.INFANTRY,  HexCoord(16, 1)),
        (UnitType.TANK,      HexCoord(14, 2)),
        (UnitType.TANK,      HexCoord(16, 2)),
        (UnitType.ARTILLERY, HexCoord(15, 3)),
        (UnitType.RECON,     HexCoord(17, 1)),
    ]

    # 放置红方
    for utype, pos in red_setups:
        # 确保位置是平原或道路
        tile = game_map.get_tile(pos)
        if tile and tile.terrain == TerrainType.WATER:
            # 找附近非水域
            for n in pos.neighbors():
                nt = game_map.get_tile(n)
                if nt and nt.terrain != TerrainType.WATER:
                    pos = n
                    break
        unit = Unit(utype, Faction.RED, pos, f"红方{UNIT_DATA[utype].name}")
        game_map.add_unit(unit)

    # 放置蓝方
    for utype, pos in blue_setups:
        tile = game_map.get_tile(pos)
        if tile and tile.terrain == TerrainType.WATER:
            for n in pos.neighbors():
                nt = game_map.get_tile(n)
                if nt and nt.terrain != TerrainType.WATER:
                    pos = n
                    break
        unit = Unit(utype, Faction.BLUE, pos, f"蓝方{UNIT_DATA[utype].name}")
        game_map.add_unit(unit)

    state.log("⚔️ 战争兵棋推演 开始！")
    state.log(f"红方 vs 蓝方 · 地图 {GRID_W}x{GRID_H}")
    state.log("── 第 1 回合 · 红方行动 ──")

    return state


# ══════════════════════════════════════════════
# 主程序入口
# ════════════════════════════════════════════

def main():
    pygame.init()
    pygame.display.set_caption("⚔️ 战争兵棋推演 · War Chess")
    screen = pygame.display.set_mode((SCREEN_W, SCREEN_H))
    clock = pygame.time.Clock()

    # 创建游戏
    game_state = create_game(seed=None)  # None = 随机地图
    renderer = Renderer(screen, game_state)
    input_handler = InputHandler(game_state, renderer, screen)

    running = True
    while running:
        for event in pygame.event.get():
            if event.type == pygame.QUIT:
                running = False
            else:
                input_handler.handle_event(event)

        # 渲染
        renderer.render(screen)
        pygame.display.flip()
        clock.tick(FPS)

    pygame.quit()
    sys.exit(0)


if __name__ == "__main__":
    main()

