import pygame
import sys
import math
import random

# 初始化
pygame.init()

# 窗口设置
WIDTH, HEIGHT = 800, 600
screen = pygame.display.set_mode((WIDTH, HEIGHT))
pygame.display.set_caption("简易塔防")

# 颜色
WHITE = (255, 255, 255)
BLACK = (0, 0, 0)
RED   = (255, 0, 0)
GREEN = (0, 255, 0)
BLUE  = (0, 0, 255)
YELLOW= (255, 255, 0)
GRAY  = (128, 128, 128)
BROWN = (139, 69, 19)
DARK_GREEN = (0, 100, 0)

clock = pygame.time.Clock()
FPS = 60

# ---------- 游戏数据 ----------
# 路径点（敌人移动路线）
path_points = [(50, 300), (250, 300), (250, 100), (550, 100), (550, 500), (750, 500)]
tower_cost = 50           # 建塔花费
starting_gold = 200       # 初始金币
starting_lives = 20       # 初始生命
enemy_spawn_delay = 60    # 生成间隔帧数
enemies_per_wave = 5      # 每波敌人数量

# 炮塔属性
TOWER_RANGE = 120
TOWER_DAMAGE = 1
TOWER_FIRE_RATE = 20     # 攻击冷却帧数

# ---------- 类定义 ----------
class Enemy:
    def __init__(self, path):
        self.path = path
        self.pos = pygame.Vector2(path[0])
        self.target_index = 1
        self.speed = 2
        self.health = 3
        self.max_health = 3
        self.alive = True
        self.reached_end = False
        self.radius = 10

    def move(self):
        if self.target_index >= len(self.path):
            self.reached_end = True
            self.alive = False
            return

        target = pygame.Vector2(self.path[self.target_index])
        direction = target - self.pos
        dist = direction.length()

        if dist < self.speed:
            self.pos = target
            self.target_index += 1
        else:
            direction.scale_to_length(self.speed)
            self.pos += direction

    def take_damage(self, amount):
        self.health -= amount
        if self.health <= 0:
            self.alive = False

    def draw(self, surface):
        # 血条
        bar_width = 20
        bar_height = 4
        health_ratio = self.health / self.max_health
        bar_x = self.pos.x - bar_width//2
        bar_y = self.pos.y - 15
        pygame.draw.rect(surface, RED, (bar_x, bar_y, bar_width, bar_height))
        pygame.draw.rect(surface, GREEN, (bar_x, bar_y, bar_width * health_ratio, bar_height))
        # 敌人圆形
        pygame.draw.circle(surface, RED, (int(self.pos.x), int(self.pos.y)), self.radius)

class Tower:
    def __init__(self, pos):
        self.pos = pygame.Vector2(pos)
        self.range = TOWER_RANGE
        self.damage = TOWER_DAMAGE
        self.fire_rate = TOWER_FIRE_RATE
        self.cooldown = 0
        self.radius = 15

    def update(self, enemies):
        if self.cooldown > 0:
            self.cooldown -= 1
            return None

        # 寻找范围内最近的敌人
        target = None
        min_dist = self.range
        for enemy in enemies:
            if enemy.alive:
                dist = self.pos.distance_to(enemy.pos)
                if dist <= min_dist:
                    min_dist = dist
                    target = enemy

        if target:
            self.cooldown = self.fire_rate
            # 造成伤害
            target.take_damage(self.damage)
            # 返回攻击特效数据（起点、终点）
            return (self.pos, target.pos)
        return None

    def draw(self, surface):
        # 绘制炮塔
        pygame.draw.rect(surface, DARK_GREEN,
                         (self.pos.x - self.radius, self.pos.y - self.radius,
                          2*self.radius, 2*self.radius))
        # 绘制范围圈（半透明）
        range_surf = pygame.Surface((self.range*2, self.range*2), pygame.SRCALPHA)
        pygame.draw.circle(range_surf, (0, 255, 0, 30), (self.range, self.range), self.range)
        surface.blit(range_surf, (self.pos.x - self.range, self.pos.y - self.range))

# ---------- 游戏状态 ----------
class Game:
    def __init__(self):
        self.gold = starting_gold
        self.lives = starting_lives
        self.towers = []
        self.enemies = []
        self.wave = 1
        self.enemies_spawned = 0
        self.spawn_timer = 0
        self.wave_complete = False
        self.game_over = False
        self.attack_effects = []  # 攻击动画线条
        self.font = pygame.font.SysFont("SimHei", 24)
        self.selected_tower = None

    def spawn_enemy(self):
        if self.enemies_spawned < enemies_per_wave:
            self.enemies.append(Enemy(path_points))
            self.enemies_spawned += 1

    def update(self):
        if self.game_over:
            return

        # 生成敌人逻辑
        if not self.wave_complete:
            if self.spawn_timer <= 0:
                self.spawn_enemy()
                self.spawn_timer = enemy_spawn_delay
            else:
                self.spawn_timer -= 1

        # 更新敌人
        for enemy in self.enemies:
            if enemy.alive:
                enemy.move()
                if enemy.reached_end:
                    self.lives -= 1
                    if self.lives <= 0:
                        self.game_over = True

        # 移除死亡/到达终点的敌人
        self.enemies = [e for e in self.enemies if e.alive and not e.reached_end]

        # 更新炮塔并收集攻击特效
        self.attack_effects.clear()
        for tower in self.towers:
            effect = tower.update(self.enemies)
            if effect:
                self.attack_effects.append(effect)

        # 处理金币（击杀奖励）
        for enemy in self.enemies:
            if not enemy.alive and enemy.health <= 0:
                self.gold += 20  # 击杀奖励
                self.enemies.remove(enemy)  # 需重新清理

        # 检查波次是否结束
        if self.enemies_spawned >= enemies_per_wave and len(self.enemies) == 0:
            self.wave_complete = True
            # 波次完成奖励
            self.gold += 50
            # 下一波
            self.wave += 1
            self.enemies_spawned = 0
            self.spawn_timer = 0
            self.wave_complete = False

    def add_tower(self, pos):
        # 检查是否与已有炮塔或路径重叠（简单判断：离路径点距离）
        min_dist_to_path = min([pygame.Vector2(pos).distance_to(p) for p in path_points])
        if min_dist_to_path < 40:  # 不能建在路径上
            return False
        for t in self.towers:
            if t.pos.distance_to(pos) < 30:
                return False
        if self.gold >= tower_cost:
            self.towers.append(Tower(pos))
            self.gold -= tower_cost
            return True
        return False

    def draw(self, surface):
        surface.fill(WHITE)
        # 绘制路径
        if len(path_points) > 1:
            pygame.draw.lines(surface, BROWN, False, path_points, 8)

        # 绘制炮塔
        for tower in self.towers:
            tower.draw(surface)

        # 绘制敌人
        for enemy in self.enemies:
            enemy.draw(surface)

        # 绘制攻击特效
        for start, end in self.attack_effects:
            pygame.draw.line(surface, YELLOW, (int(start.x), int(start.y)),
                             (int(end.x), int(end.y)), 3)

        # UI信息
        gold_text = self.font.render(f"金币: {self.gold}", True, BLACK)
        lives_text = self.font.render(f"生命: {self.lives}", True, BLACK)
        wave_text = self.font.render(f"波次: {self.wave}", True, BLACK)
        surface.blit(gold_text, (10, 10))
        surface.blit(lives_text, (10, 40))
        surface.blit(wave_text, (10, 70))

        # 当前鼠标位置建塔提示
        mx, my = pygame.mouse.get_pos()
        if self.gold >= tower_cost:
            # 检查是否可建
            min_dist = min([pygame.Vector2(mx, my).distance_to(p) for p in path_points])
            can_place = min_dist >= 40
            for t in self.towers:
                if t.pos.distance_to((mx, my)) < 30:
                    can_place = False
                    break
            color = (0, 255, 0, 100) if can_place else (255, 0, 0, 100)
            preview_surf = pygame.Surface((TOWER_RANGE*2, TOWER_RANGE*2), pygame.SRCALPHA)
            pygame.draw.circle(preview_surf, color, (TOWER_RANGE, TOWER_RANGE), TOWER_RANGE)
            surface.blit(preview_surf, (mx - TOWER_RANGE, my - TOWER_RANGE))
            pygame.draw.rect(surface, DARK_GREEN if can_place else RED,
                             (mx-15, my-15, 30, 30))

        if self.game_over:
            over_text = self.font.render("游戏结束！按R重新开始", True, RED)
            surface.blit(over_text, (WIDTH//2 - 150, HEIGHT//2 - 20))

def main():
    game = Game()
    running = True

    while running:
        clock.tick(FPS)
        for event in pygame.event.get():
            if event.type == pygame.QUIT:
                running = False
            if event.type == pygame.KEYDOWN:
                if event.key == pygame.K_r and game.game_over:
                    game = Game()  # 重置
            if event.type == pygame.MOUSEBUTTONDOWN:
                if event.button == 1 and not game.game_over:
                    game.add_tower(pygame.mouse.get_pos())

        game.update()
        game.draw(screen)
        pygame.display.flip()

    pygame.quit()
    sys.exit()

if __name__ == "__main__":
    main()