"""
============================================================
  WAR CHESS -- Hex Grid Strategy Game  (Pygame)
  No Chinese fonts required. Safe cross-platform build.
============================================================
"""

import pygame, math, random, sys
from enum import Enum
from dataclasses import dataclass, field
from typing import List, Tuple, Optional, Dict

# ── Screen ────────────────────────────────────────────────
W, H = 1280, 800
HEX = 38
COLS, ROWS = 16, 12
SIDE = 240
TOP = 44
BOT = 96

# ── Colours ────────────────────────────────────────────────
C = {
    "bg":       (18, 22, 32),
    "panel":    (26, 33, 48),
    "side":     (30, 38, 56),
    "grid":     (38, 50, 70),
    "text":     (220, 225, 235),
    "dim":      (130, 140, 155),
    "grey":     (95, 105, 120),
    "red":      (215, 55, 55),
    "red_d":    (130, 30, 30),
    "blue":     (50, 115, 215),
    "blue_d":   (28, 70, 135),
    "green":    (55, 195, 95),
    "gold":     (215, 175, 45),
    "purple":   (155, 75, 215),
    "orange":   (235, 135, 35),
    "plain":    (52, 66, 46),
    "forest":   (38, 82, 48),
    "mount":    (82, 70, 55),
    "water":    (42, 92, 145),
    "city":     (110, 95, 70),
    "hl_sel":   (100, 150, 200),
    "hl_mov":   (60, 175, 110),
    "hl_atk":   (200, 70, 70),
    "hl_hov":   (70, 95, 125),
    "hp_g":     (50, 200, 75),
    "hp_y":     (220, 200, 40),
    "hp_r":     (220, 50, 50),
}

# ── Enums ─────────────────────────────────────────────────
class Terrain(Enum):
    PLAIN = "Plain"
    FOREST = "Forest"
    MOUNTAIN = "Mount"
    WATER = "Water"
    CITY = "City"

class UType(Enum):
    INF = "Infantry"
    TANK = "Tank"
    ART = "Artillery"
    HELI = "Heli"
    CMD = "Commander"

class Side(Enum):
    RED = "RED"
    BLU = "BLUE"

# ── Unit data ─────────────────────────────────────────────
@dataclass
class Stats:
    type: UType
    hp: int
    atk: int
    dfn: int
    mov: int
    rng: int
    ammo: int
    fuel: int
    vis: int
    sym: str  # single char symbol

UDATA: Dict[UType, Stats] = {
    UType.INF:  Stats(UType.INF,  100, 25, 15, 3, 1, 10, 99, 3, "I"),
    UType.TANK: Stats(UType.TANK, 180, 55, 40, 4, 2,  8, 60, 4, "T"),
    UType.ART:  Stats(UType.ART,  120, 70, 10, 2, 4,  6, 40, 3, "A"),
    UType.HELI: Stats(UType.HELI,  90, 40, 20, 6, 2, 12, 50, 5, "H"),
    UType.CMD:  Stats(UType.CMD,  150, 30, 30, 3, 1, 99, 99, 5, "C"),
}

# ── Unit ──────────────────────────────────────────────────
@dataclass
class Unit:
    id: int
    side: Side
    type: UType
    col: int
    row: int
    stats: Stats
    hp: int = 100
    ammo: int = 10
    fuel: int = 60
    moved: bool = False
    attacked: bool = False

    @property
    def atk(self): return self.stats.atk
    @property
    def dfn(self): return self.stats.dfn

    def color(self):
        base = C["red"] if self.side == Side.RED else C["blue"]
        if self.moved and self.attacked:
            return C["red_d"] if self.side == Side.RED else C["blue_d"]
        return base

# ── Tile ──────────────────────────────────────────────────
@dataclass
class Tile:
    col: int
    row: int
    terrain: Terrain = Terrain.PLAIN
    unit: Optional[Unit] = None

# ── Hex math (pointy-top, axial q/r) ──────────────────
def hex2px(q, r, size=HEX):
    x = size * math.sqrt(3) * (q + r/2)
    y = size * 1.5 * r
    return x + (W-SIDE)//2, y + TOP + 40

def px2hex(px, py, size=HEX):
    ox = (W-SIDE)//2
    oy = TOP + 40
    x, y = (px-ox)/size, (py-oy)/size
    q = (math.sqrt(3)/3*x - 1/3*y)
    r = (2/3*y)
    return hex_round(q, r)

def hex_round(q, r):
    s = -q-r
    rq, rr, rs = round(q), round(r), round(s)
    dq, dr, ds = abs(rq-q), abs(rr-r), abs(rs-s)
    if dq > dr and dq > ds: rq = -rr-rs
    elif dr > ds: rr = -rq-rs
    else: rs = -rq-rr
    return int(rq), int(rr)

def hex_neigh(c, r):
    d = [(1,0),(1,-1),(0,-1),(-1,0),(-1,1),(0,1)]
    return [(c+i, r+j) for i,j in d]

def hex_dist(c1, r1, c2, r2):
    return (abs(c1-c2)+abs(c1+r1-c2-r2)+abs(r1-r2))//2

def hex_corner(cx, cy, s, i):
    a = math.radians(60*i - 30)
    return cx+s*math.cos(a), cy+s*math.sin(a)

# ── Map gen ───────────────────────────────────────────────
def gen_map():
    tiles = {}
    for r in range(ROWS):
        for c in range(COLS):
            p = random.random()
            if p < .10: t = Terrain.FOREST
            elif p < .17: t = Terrain.MOUNTAIN
            elif p < .24: t = Terrain.WATER
            elif p < .31: t = Terrain.CITY
            else: t = Terrain.PLAIN
            if c < 2 or c >= COLS-2 or r < 1 or r >= ROWS-1:
                if t == Terrain.WATER: t = Terrain.PLAIN
            tiles[(c,r)] = Tile(c, r, t)
    return tiles

def place_units(tiles):
    units, uid = [], 0
    red_cfg = [
        (UType.CMD,1,5),(UType.INF,0,3),(UType.INF,0,7),
        (UType.TANK,2,4),(UType.TANK,2,6),(UType.ART,1,2),
        (UType.ART,1,8),(UType.HELI,3,5),
    ]
    blue_cfg = [
        (UType.CMD,14,5),(UType.INF,15,3),(UType.INF,15,7),
        (UType.TANK,13,4),(UType.TANK,13,6),(UType.ART,14,2),
        (UType.ART,14,8),(UType.HELI,12,5),
    ]
    for fac, cfg in [(Side.RED, red_cfg), (Side.BLU, blue_cfg)]:
        for ut, c, r in cfg:
            s = UDATA[ut]
            while tiles[(c,r)].terrain == Terrain.WATER:
                c += 1 if fac == Side.RED else -1
            u = Unit(uid, fac, ut, c, r, s, s.hp, s.ammo, s.fuel)
            tiles[(c,r)].unit = u
            units.append(u); uid += 1
    return units

# ── Combat ────────────────────────────────────────────────
def calc_dmg(atk: Unit, dfn: Unit, tiles: Dict):
    atk_t = tiles[(atk.col, atk.row)].terrain
    dfn_t = tiles[(dfn.col, dfn.row)].terrain
    dmg = atk.atk*2 - dfn.dfn
    # terrain def
    td = 0
    if dfn_t == Terrain.FOREST: td = 8
    elif dfn_t == Terrain.MOUNTAIN: td = 15
    elif dfn_t == Terrain.CITY: td = 10
    elif dfn_t == Terrain.WATER: td = -5
    # terrain atk
    ta = 0
    if atk_t == Terrain.MOUNTAIN: ta = 5
    elif atk_t == Terrain.FOREST: ta = -3
    # type bonus
    tb = 0
    pairs = {
        (UType.TANK, UType.INF): 15,
        (UType.INF, UType.ART): 20,
        (UType.ART, UType.TANK): 10,
        (UType.HELI, UType.TANK): 15,
        (UType.ART, UType.HELI): -10,
    }
    tb = pairs.get((atk.type, dfn.type), 0)
    var = random.uniform(.85, 1.15)
    dmg = int((dmg + td + ta + tb) * var)
    dmg = max(5, min(dmg, 95))
    tag = ""
    if tb >= 10: tag = "Counter!"
    elif tb <= -5: tag = "Disadv"
    if td >= 10: tag += " Cover" if tag else "Cover"
    if dmg >= 60: tag = "CRIT! " + tag if tag else "CRIT!"
    elif dmg >= 40: tag = "Heavy " + tag if tag else "Heavy"
    return dmg, tag.strip()

# ── Pathfinding ───────────────────────────────────────────
def reachable(unit: Unit, tiles: Dict):
    out, q = {}, [(unit.col, unit.row, 0)]
    seen = {(unit.col, unit.row): 0}
    while q:
        c, r, cost = q.pop(0)
        if cost > unit.stats.mov: continue
        out[(c,r)] = cost
        for nc, nr in hex_neigh(c, r):
            if (nc, nr) not in tiles: continue
            t = tiles[(nc, nr)]
            if t.terrain == Terrain.WATER and unit.type != UType.HELI: continue
            if t.unit and (nc, nr) != (unit.col, unit.row): continue
            mc = 2 if t.terrain == Terrain.MOUNTAIN and unit.type != UType.HELI else 1
            nc2 = cost + mc
            if nc2 <= unit.stats.mov and (nc, nr) not in seen:
                seen[(nc, nr)] = nc2
                q.append((nc, nr, nc2))
    return out

def attackable(unit: Unit, tiles: Dict):
    out = []
    for dc in range(-unit.stats.rng, unit.stats.rng+1):
        for dr in range(-unit.stats.rng, unit.stats.rng+1):
            nc, nr = unit.col+dc, unit.row+dr
            if 1 <= hex_dist(unit.col,unit.row,nc,nr) <= unit.stats.rng and (nc,nr) in tiles:
                out.append((nc,nr))
    return out

# ═════════════════════════════════════════════════════════
#  GAME CLASS
# ═════════════════════════════════════════════════════════
class WarGame:
    def __init__(self):
        pygame.init()
        self.screen = pygame.display.set_mode((W, H))
        pygame.display.set_caption("WAR CHESS -- Hex Strategy")
        self.clock = pygame.time.Clock()
        # fonts -- default only, no Chinese
        self.f1 = pygame.font.Font(None, 14)
        self.f2 = pygame.font.Font(None, 16)
        self.f3 = pygame.font.Font(None, 19)
        self.f4 = pygame.font.Font(None, 22)
        self.f5 = pygame.font.Font(None, 28)
        self.f6 = pygame.font.Font(None, 36)
        self.fb = pygame.font.Font(None, 20)  # bold-ish (just bigger)

        self.tiles = gen_map()
        self.units = place_units(self.tiles)
        self.turn = 1
        self.side = Side.RED
        self.sel: Optional[Unit] = None
        self.mode = "sel"  # sel / mov / atk
        self.reach: Dict = {}
        self.atk_tiles: List = []
        self.log: List[str] = []
        self.over = False
        self.winner = None
        self.mouse = (0, 0)
        self.hover = None
        self.ai = False
        self.ai_t = 0

        self.log.append(">> WAR CHESS -- RED moves first!")
        self.log.append(">> Goal: destroy enemy Commander!")

    # ── helpers ─────────────────────────────────────────
    def add_log(self, m):
        self.log.append(m)
        if len(self.log) > 60: self.log = self.log[-60:]

    def end_unit(self):
        if self.sel: self.sel.moved = self.sel.attacked = True
        self.sel = None
        self.mode = "sel"
        self.reach = {}
        self.atk_tiles = []

    def switch(self):
        self.side = Side.BLU if self.side == Side.RED else Side.RED
        if self.side == Side.RED: self.turn += 1
        for u in self.units:
            if u.side == self.side: u.moved = u.attacked = False
        self.sel = None; self.mode = "sel"; self.reach = {}; self.atk_tiles = []
        self.add_log(f"-- Turn {self.turn} : {'RED' if self.side==Side.RED else 'BLUE'} --")

    def check_win(self):
        red = any(u.side==Side.RED and u.type==UType.CMD and u.hp>0 for u in self.units)
        blu = any(u.side==Side.BLU and u.type==UType.CMD and u.hp>0 for u in self.units)
        if not red: self.over, self.winner = True, Side.BLU
        elif not blu: self.over, self.winner = True, Side.RED

    # ── Click ──────────────────────────────────────────
    def click(self, pos):
        if self.over: return
        hx, hy = px2hex(pos[0], pos[1])
        if (hx, hy) not in self.tiles: return
        tile = self.tiles[(hx, hy)]

        if self.mode == "sel":
            if tile.unit and tile.unit.side == self.side:
                self.sel = tile.unit
                self.mode = "mov"
                self.reach = reachable(tile.unit, self.tiles)
                self.atk_tiles = attackable(tile.unit, self.tiles) if not tile.unit.attacked else []
                self.add_log(f"Sel: {tile.unit.stats.sym} {tile.unit.type.value} ({hx},{hy})")

        elif self.mode == "mov":
            if not self.sel: return
            if (hx, hy) in self.reach and not tile.unit:
                oc, or_ = self.sel.col, self.sel.row
                self.tiles[(oc, or_)].unit = None
                self.sel.col, self.sel.row = hx, hy
                self.tiles[(hx, hy)].unit = self.sel
                self.sel.moved = True
                self.sel.fuel -= 1
                self.add_log(f"  {self.sel.type.value} -> ({hx},{hy})")
                if not self.sel.attacked and self.sel.ammo > 0:
                    self.mode = "atk"
                    self.atk_tiles = attackable(self.sel, self.tiles)
                else: self.end_unit()
            elif tile.unit and tile.unit.side == self.side:
                self.sel = tile.unit
                self.reach = reachable(tile.unit, self.tiles)
                self.atk_tiles = attackable(tile.unit, self.tiles) if not tile.unit.attacked else []

        elif self.mode == "atk":
            if not self.sel: return
            if (hx, hy) in self.atk_tiles:
                tgt = tile.unit
                if tgt and tgt.side != self.side:
                    self.do_atk(self.sel, tgt)
            self.end_unit()

    def do_atk(self, a: Unit, d: Unit):
        if a.ammo <= 0:
            self.add_log("  No ammo!"); return
        a.ammo -= 1; a.attacked = True
        dmg, tag = calc_dmg(a, d, self.tiles)
        d.hp -= dmg
        self.add_log(f"  {a.type.value} >> {d.type.value}  dmg={dmg} {tag}")
        # counter
        if d.hp > 0 and d.ammo > 0 and hex_dist(a.col,a.row,d.col,d.row) <= d.stats.rng:
            cd, ctag = calc_dmg(d, a, self.tiles)
            a.hp -= cd
            self.add_log(f"    << {d.type.value} counter dmg={cd} {ctag}")
        for u in (d, a):
            if u.hp <= 0:
                self.add_log(f"    XX {u.type.value} DESTROYED")
                self.tiles[(u.col, u.row)].unit = None
                self.units.remove(u)
        self.check_win()

    # ── AI ─────────────────────────────────────────────
    def ai_step(self):
        if self.over: return
        fac = self.side
        allies = [u for u in self.units if u.side==fac and not (u.moved and u.attacked)]
        random.shuffle(allies)
        enemies = [u for u in self.units if u.side != fac]
        if not enemies: return

        for u in allies:
            if u.moved and u.attacked: continue
            # attack first
            if not u.attacked and u.ammo > 0:
                atks = attackable(u, self.tiles)
                # prefer commander
                tgt = None
                for ac, ar in atks:
                    t = self.tiles[(ac,ar)].unit
                    if t and t.side != fac:
                        if t.type == UType.CMD: tgt = t; break
                if not tgt:
                    for ac, ar in atks:
                        t = self.tiles[(ac,ar)].unit
                        if t and t.side != fac: tgt = t; break
                if tgt: self.do_atk(u, tgt)
                if self.over: return
            # move toward enemy
            if not u.moved:
                reach = reachable(u, self.tiles)
                best, bd = None, 999
                for (rc, rr), _ in reach.items():
                    if self.tiles[(rc,rr)].unit: continue
                    for e in enemies:
                        d = hex_dist(rc, rr, e.col, e.row)
                        if d < bd: bd, best = d, (rc, rr)
                if best and best != (u.col, u.row):
                    self.tiles[(u.col,u.row)].unit = None
                    u.col, u.row = best
                    self.tiles[best].unit = u
                    u.moved = True; u.fuel -= 1
            # attack after move
            if not u.attacked and u.ammo > 0:
                for ac, ar in attackable(u, self.tiles):
                    t = self.tiles[(ac,ar)].unit
                    if t and t.side != fac: self.do_atk(u, t); break
                if self.over: return
        self.switch()

    # ── Drawing ────────────────────────────────────────
    def draw_hex(self, c, r, col, alpha=255, border=None, bw=1):
        cx, cy = hex2px(c, r)
        pts = [hex_corner(cx, cy, HEX-1, i) for i in range(6)]
        if alpha < 255:
            tmp = pygame.Surface((HEX*2, HEX*2), pygame.SRCALPHA)
            tr = tmp.get_rect(center=(cx, cy))
            pts2 = [hex_corner(HEX, HEX, HEX-1, i) for i in range(6)]
            pygame.draw.polygon(tmp, (*col, alpha), pts2)
            if border: pygame.draw.polygon(tmp, (*border, alpha), pts2, bw)
            self.screen.blit(tmp, tr)
        else:
            pygame.draw.polygon(self.screen, col, pts)
            if border: pygame.draw.polygon(self.screen, border, pts, bw)
            else: pygame.draw.polygon(self.screen, C["grid"], pts, 1)

    def draw_terrain(self, c, r, t: Terrain):
        cx, cy = hex2px(c, r)
        if t == Terrain.FOREST:
            for dx, dy in [(-7,-4),(4,1),(-1,7)]:
                pygame.draw.circle(self.screen, (28,90,38), (int(cx+dx),int(cy+dy)), 5)
                pygame.draw.line(self.screen, (75,55,25), (cx+dx,cy+dy+5), (cx+dx,cy+dy+10), 2)
        elif t == Terrain.MOUNTAIN:
            pts = [(cx-9,cy+7),(cx-2,cy-9),(cx+5,cy-2),(cx+9,cy+7)]
            pygame.draw.polygon(self.screen, (105,90,65), pts)
            pygame.draw.polygon(self.screen, (190,190,210), [(cx-2,cy-9),(cx+2,cy-5),(cx-3,cy-3)])
        elif t == Terrain.WATER:
            for i in range(2):
                pygame.draw.line(self.screen, (58,125,185), (cx-11,cy-2+i*6), (cx+11,cy-2+i*6), 2)
        elif t == Terrain.CITY:
            pygame.draw.rect(self.screen, (95,80,55), (cx-9,cy-7,18,14))
            pygame.draw.polygon(self.screen, (135,45,45), [(cx-9,cy-7),(cx,cy-14),(cx+9,cy-7)])
            pygame.draw.rect(self.screen, (55,45,35), (cx-2,cy+2,4,7))

    def draw_unit(self, u: Unit):
        cx, cy = hex2px(u.col, u.row)
        col = u.color()
        s = HEX-10
        pts = [hex_corner(cx, cy, s, i) for i in range(6)]
        pygame.draw.polygon(self.screen, col, pts)
        pygame.draw.polygon(self.screen, C["text"], pts, 2)
        # symbol
        t = self.fb.render(u.stats.sym, True, C["text"])
        self.screen.blit(t, t.get_rect(center=(cx, cy-3)))
        # hp bar
        bw, bh = 26, 4
        pct = max(0, u.hp/u.stats.hp)
        hc = C["hp_g"] if pct>.6 else C["hp_y"] if pct>.3 else C["hp_r"]
        pygame.draw.rect(self.screen, (15,15,15), (cx-bw//2, cy+8, bw, bh))
        pygame.draw.rect(self.screen, hc, (cx-bw//2, cy+8, int(bw*pct), bh))
        # ammo dot
        if u.ammo <= 2 and u.ammo > 0:
            pygame.draw.circle(self.screen, C["orange"], (cx+10, cy-10), 3)
        # done X
        if u.moved and u.attacked:
            pygame.draw.line(self.screen, C["grey"], (cx-7,cy-7),(cx+7,cy+7),2)
            pygame.draw.line(self.screen, C["grey"], (cx-7,cy+7),(cx+7,cy-7),2)

    def draw_top(self):
        pygame.draw.rect(self.screen, C["panel"], (0,0,W-SIDE,TOP))
        pygame.draw.line(self.screen, C["grid"], (0,TOP),(W-SIDE,TOP),2)
        t1 = self.f4.render("WAR CHESS -- Hex Strategy", True, C["gold"])
        self.screen.blit(t1, (12, 10))
        fc = C["red"] if self.side==Side.RED else C["blue"]
        bg = (60,20,20) if self.side==Side.RED else (20,30,60)
        r = pygame.Rect(340, 8, 110, 28)
        pygame.draw.rect(self.screen, bg, r, border_radius=4)
        pygame.draw.rect(self.screen, fc, r, 2, border_radius=4)
        t2 = self.f3.render(f"Turn {self.turn} : {self.side.value}", True, fc)
        self.screen.blit(t2, (350, 12))
        # counts
        rc = sum(1 for u in self.units if u.side==Side.RED)
        bc = sum(1 for u in self.units if u.side==Side.BLU)
        self.screen.blit(self.f2.render(f"RED:{rc}", True, C["red"]), (490,14))
        self.screen.blit(self.f2.render(f"BLU:{bc}", True, C["blue"]), (560,14))
        # buttons
        for txt, bx, col in [("End Turn", W-SIDE-170, (75,45,45)), ("AI Auto", W-SIDE-90, (45,45,75))]:
            r = pygame.Rect(bx, 8, 72, 28)
            pygame.draw.rect(self.screen, col, r, border_radius=4)
            pygame.draw.rect(self.screen, C["text"], r, 1, border_radius=4)
            self.screen.blit(self.f2.render(txt, True, C["text"]), (bx+6, 12))

    def draw_side(self):
        sx = W - SIDE
        pygame.draw.rect(self.screen, C["side"], (sx,0,SIDE,H))
        pygame.draw.line(self.screen, C["grid"], (sx,0),(sx,H),2)
        y = 12
        self.screen.blit(self.f4.render("BATTLE INFO", True, C["gold"]), (sx+16, y)); y+=30
        # turn
        fc = C["red"] if self.side==Side.RED else C["blue"]
        self.screen.blit(self.f3.render(f"Turn {self.turn}", True, C["text"]), (sx+16,y))
        self.screen.blit(self.f3.render(f"{self.side.value} ACTIVE", True, fc), (sx+110,y)); y+=26
        pygame.draw.line(self.screen, C["grid"], (sx+8,y),(sx+SIDE-8,y),1); y+=8

        if self.sel:
            u = self.sel
            hc = C["red"] if u.side==Side.RED else C["blue"]
            self.screen.blit(self.f3.render(f"{u.side.value} {u.type.value}", True, hc), (sx+16,y)); y+=24
            for s in [f"HP  {u.hp}/{u.stats.hp}", f"ATK {u.atk}  DEF {u.dfn}",
                      f"MOV {u.stats.mov}  RNG {u.stats.rng}",
                      f"Ammo {u.ammo}  Fuel {u.fuel}", f"Pos ({u.col},{u.row})"]:
                self.screen.blit(self.f1.render(s, True, C["dim"]), (sx+22,y)); y+=18
            y+=4
            pygame.draw.line(self.screen, C["grid"], (sx+8,y),(sx+SIDE-8,y),1); y+=6
            if self.mode=="mov":
                self.screen.blit(self.f1.render("Click GREEN to move", True, C["green"]), (sx+16,y)); y+=18
            elif self.mode=="atk":
                self.screen.blit(self.f1.render("Click RED to attack", True, C["hl_atk"]), (sx+16,y)); y+=18
        else:
            self.screen.blit(self.f1.render("Click your unit", True, C["grey"]), (sx+16,y)); y+=22

        pygame.draw.line(self.screen, C["grid"], (sx+8,y),(sx+SIDE-8,y),1); y+=8
        self.screen.blit(self.f3.render("UNITS", True, C["gold"]), (sx+16,y)); y+=24
        for ut in UType:
            s = UDATA[ut]
            pygame.draw.rect(self.screen, C["red"], (sx+18,y+2,12,12))
            pygame.draw.rect(self.screen, C["blue"], (sx+32,y+2,12,12))
            info = f"{s.sym} {ut.value:9s} A{s.atk} D{s.dfn} M{s.mov} R{s.rng}"
            self.screen.blit(self.f1.render(info, True, C["dim"]), (sx+50,y))
            y+=18

        y+=4
        pygame.draw.line(self.screen, C["grid"], (sx+8,y),(sx+SIDE-8,y),1); y+=8
        self.screen.blit(self.f3.render("TERRAIN", True, C["gold"]), (sx+16,y)); y+=22
        for name, col in [("Plain",C["plain"]),("Forest",C["forest"]),("Mount",C["mount"]),
                          ("Water",C["water"]),("City",C["city"])]:
            pygame.draw.rect(self.screen, col, (sx+18,y+2,10,10))
            self.screen.blit(self.f1.render(name, True, C["dim"]), (sx+34,y))
            y+=16

    def draw_bot(self):
        by = H - BOT
        pygame.draw.rect(self.screen, C["panel"], (0,by,W-SIDE,BOT))
        pygame.draw.line(self.screen, C["grid"], (0,by),(W-SIDE,by),2)
        self.screen.blit(self.f3.render("BATTLE LOG", True, C["gold"]), (12, by+6))
        for i, m in enumerate(self.log[-4:]):
            col = C["text"]
            if "DESTROY" in m or "win" in m.lower(): col = C["hl_atk"]
            elif ">>" in m or "<<" in m: col = C["orange"]
            elif "->" in m: col = C["dim"]
            self.screen.blit(self.f1.render(m, True, col), (16, by+28+i*16))
        hx = 480
        for i, s in enumerate(["LClick: Sel/Move/Atk","Space: End Turn","R: Restart","A: AI Toggle"]):
            self.screen.blit(self.f1.render(s, True, C["grey"]), (hx+i*170, by+8))

    def draw_over(self):
        ov = pygame.Surface((W,H), pygame.SRCALPHA); ov.fill((0,0,0,170))
        self.screen.blit(ov, (0,0))
        cx, cy = W//2, H//2
        wc = C["red"] if self.winner==Side.RED else C["blue"]
        self.screen.blit(self.f6.render("GAME OVER", True, C["gold"]), (cx-100, cy-60))
        self.screen.blit(self.f5.render(f"{self.winner.value} WINS!", True, wc), (cx-90, cy-16))
        self.screen.blit(self.f2.render("Press R to restart", True, C["dim"]), (cx-70, cy+30))

    def draw(self):
        self.screen.fill(C["bg"])
        # tiles
        for (c,r), tile in self.tiles.items():
            tc = {Terrain.PLAIN:C["plain"],Terrain.FOREST:(33,75,42),
                  Terrain.MOUNTAIN:C["mount"],Terrain.WATER:C["water"],
                  Terrain.CITY:(88,76,55)}[tile.terrain]
            col, alpha, border, bw = tc, 255, None, 1
            if self.sel and (c,r)==(self.sel.col,self.sel.row):
                border, bw = C["hl_sel"], 3
            elif self.mode=="mov" and (c,r) in self.reach:
                col, alpha, border, bw = C["hl_mov"], 200, C["green"], 2
            elif self.mode=="atk" and (c,r) in self.atk_tiles:
                t = self.tiles[(c,r)].unit
                if t and t.side != self.side:
                    col, alpha, border, bw = C["hl_atk"], 220, C["red"], 2
            elif self.hover==(c,r):
                col, alpha = C["hl_hov"], 230
            self.draw_hex(c, r, col, alpha, border, bw)
            self.draw_terrain(c, r, tile.terrain)
        # units
        for u in self.units: self.draw_unit(u)
        # ui
        self.draw_top(); self.draw_side(); self.draw_bot()
        if self.over: self.draw_over()
        pygame.display.flip()

    # ── Main loop ─────────────────────────────────────
    def run(self):
        while True:
            dt = self.clock.tick(60)
            for ev in pygame.event.get():
                if ev.type == pygame.QUIT: pygame.quit(); sys.exit()
                elif ev.type == pygame.MOUSEMOTION:
                    self.mouse = ev.pos
                    h = px2hex(*ev.pos)
                    self.hover = h if h in self.tiles else None
                elif ev.type == pygame.MOUSEBUTTONDOWN and ev.button == 1:
                    x,y = ev.pos
                    # End Turn btn
                    if 0<=y<=TOP and W-SIDE-170<=x<=W-SIDE-98: self.switch(); self.ai=False; continue
                    if 0<=y<=TOP and W-SIDE-90<=x<=W-SIDE-18: self.ai=not self.ai; continue
                    if not self.over and x < W-SIDE: self.click(ev.pos)
                elif ev.type == pygame.KEYDOWN:
                    if ev.key == pygame.K_SPACE: self.switch(); self.ai=False
                    elif ev.key == pygame.K_r: self.__init__()
                    elif ev.key == pygame.K_a: self.ai = not self.ai

            if self.ai and not self.over:
                self.ai_t += dt
                if self.ai_t > 700:
                    self.ai_t = 0
                    self.ai_step()

            self.draw()

# ── Run ───────────────────────────────────────────────────
if __name__ == "__main__":
    WarGame().run()
