import pygame
import sys
import math

pygame.init()

BOARD_SIZE = 19
CELL_SIZE = 30
MARGIN_LEFT_RIGHT = 40
TOP_MARGIN = 100
BOTTOM_MARGIN = 60
BOARD_WIDTH = CELL_SIZE * (BOARD_SIZE - 1)
WIDTH = MARGIN_LEFT_RIGHT * 2 + BOARD_WIDTH
HEIGHT = TOP_MARGIN + BOARD_WIDTH + BOTTOM_MARGIN

screen = pygame.display.set_mode((WIDTH, HEIGHT))
pygame.display.set_caption("Go Game - Enhanced")

# 颜色
WOOD_LIGHT = (222, 184, 135)
WOOD_DARK = (188, 143, 100)
BOARD_BG = (210, 170, 110)
LINE_COLOR = (70, 40, 20)
BLACK_STONE = (30, 30, 30)
WHITE_STONE = (245, 245, 245)
BLACK_GLOW = (100, 100, 100)
WHITE_GLOW = (200, 200, 200)
GREEN = (50, 200, 50)
RED = (220, 50, 50)
GOLD = (255, 215, 0)

font = pygame.font.Font(None, 26)
small_font = pygame.font.Font(None, 18)

board = [[0] * BOARD_SIZE for _ in range(BOARD_SIZE)]
current_player = 1
game_over = False
winner = 0
animation = None  # 动画状态字典

STAR_POINTS = [(3,3), (3,9), (3,15), (9,3), (9,9), (9,15), (15,3), (15,9), (15,15)]
BLACK_REST = (80, 50)
WHITE_REST = (WIDTH - 80, 50)

def board_to_screen(row, col):
    return MARGIN_LEFT_RIGHT + col * CELL_SIZE, TOP_MARGIN + row * CELL_SIZE

def screen_to_board(x, y):
    col = round((x - MARGIN_LEFT_RIGHT) / CELL_SIZE)
    row = round((y - TOP_MARGIN) / CELL_SIZE)
    if 0 <= row < BOARD_SIZE and 0 <= col < BOARD_SIZE:
        px, py = board_to_screen(row, col)
        if abs(x - px) <= CELL_SIZE//2 and abs(y - py) <= CELL_SIZE//2:
            return row, col
    return None

def has_liberty(row, col, player):
    visited = set()
    stack = [(row, col)]
    while stack:
        r, c = stack.pop()
        if (r, c) in visited: continue
        visited.add((r, c))
        for dr, dc in [(-1,0),(1,0),(0,-1),(0,1)]:
            nr, nc = r+dr, c+dc
            if 0 <= nr < BOARD_SIZE and 0 <= nc < BOARD_SIZE:
                if board[nr][nc] == 0: return True
                elif board[nr][nc] == player and (nr,nc) not in visited:
                    stack.append((nr,nc))
    return False

def remove_group(row, col, player):
    group = []
    stack = [(row,col)]
    visited = set()
    while stack:
        r,c = stack.pop()
        if (r,c) in visited: continue
        visited.add((r,c))
        if board[r][c] == player:
            group.append((r,c))
            for dr,dc in [(-1,0),(1,0),(0,-1),(0,1)]:
                nr,nc = r+dr, c+dc
                if 0 <= nr < BOARD_SIZE and 0 <= nc < BOARD_SIZE and board[nr][nc] == player:
                    stack.append((nr,nc))
    for r,c in group:
        board[r][c] = 0
    return len(group)

def is_valid_move(row, col, player):
    if board[row][col] != 0: return False
    board[row][col] = player
    opponent = 3 - player
    captured = 0
    for dr,dc in [(-1,0),(1,0),(0,-1),(0,1)]:
        nr,nc = row+dr, col+dc
        if 0 <= nr < BOARD_SIZE and 0 <= nc < BOARD_SIZE:
            if board[nr][nc] == opponent and not has_liberty(nr,nc,opponent):
                captured += remove_group(nr,nc,opponent)
    if has_liberty(row, col, player):
        return True
    else:
        if captured == 0:
            board[row][col] = 0
            return False
        return True

def place_stone(row, col):
    global current_player, animation
    if game_over or board[row][col] != 0: return
    player = current_player
    if not is_valid_move(row, col, player): return
    target = board_to_screen(row, col)
    start = BLACK_REST if player == 1 else WHITE_REST
    animation = {
        'player': player,
        'start_pos': start,
        'target_pos': target,
        'time': 0,
        'duration': 20,          # 每半程帧数
        'phase': 'out',          # 'out' 去, 'back' 回
        'stone_placed': False,
        'board_pos': (row, col),
        'current_pos': start
    }

def update_animation():
    global animation, current_player
    if not animation: return
    anim = animation
    anim['time'] += 1
    progress = anim['time'] / anim['duration']
    if anim['phase'] == 'out':
        if progress >= 1.0:
            # 到达目标，落子
            anim['stone_placed'] = True
            row, col = anim['board_pos']
            board[row][col] = anim['player']
            anim['phase'] = 'back'
            anim['time'] = 0
            current_player = 3 - anim['player']
        else:
            t = progress * progress * (3 - 2 * progress)  # smoothstep
            x = anim['start_pos'][0] + (anim['target_pos'][0] - anim['start_pos'][0]) * t
            y = anim['start_pos'][1] + (anim['target_pos'][1] - anim['start_pos'][1]) * t
            anim['current_pos'] = (int(x), int(y))
    else:  # back
        if progress >= 1.0:
            animation = None
        else:
            t = progress * progress * (3 - 2 * progress)
            x = anim['target_pos'][0] + (anim['start_pos'][0] - anim['target_pos'][0]) * t
            y = anim['target_pos'][1] + (anim['start_pos'][1] - anim['target_pos'][1]) * t
            anim['current_pos'] = (int(x), int(y))

def draw_stick_figure(surface, x, y, color, scale=1.0):
    """绘制更精细的火柴人"""
    head_r = int(8 * scale)
    body_len = int(22 * scale)
    arm_len = int(12 * scale)
    leg_len = int(18 * scale)
    # 头
    pygame.draw.circle(surface, color, (x, y), head_r)
    pygame.draw.circle(surface, (200,200,200) if color==WHITE_STONE else (80,80,80),
                       (x-2, y-2), max(2, head_r//3))
    # 身体
    pygame.draw.line(surface, color, (x, y+head_r), (x, y+head_r+body_len), 2)
    # 手臂
    shoulder_y = y + head_r + int(body_len*0.2)
    hand_y = y + head_r + int(body_len*0.7)
    pygame.draw.line(surface, color, (x, shoulder_y), (x - arm_len, hand_y), 2)
    pygame.draw.line(surface, color, (x, shoulder_y), (x + arm_len, hand_y), 2)
    # 腿
    hip_y = y + head_r + body_len
    foot_y = hip_y + leg_len
    pygame.draw.line(surface, color, (x, hip_y), (x - arm_len, foot_y), 2)
    pygame.draw.line(surface, color, (x, hip_y), (x + arm_len, foot_y), 2)

def draw_board():
    screen.fill(WOOD_LIGHT)
    # 绘制木纹条纹
    for i in range(0, WIDTH, 40):
        pygame.draw.line(screen, WOOD_DARK, (i, 0), (i+20, HEIGHT), 3)
    # 棋盘
    board_rect = pygame.Rect(MARGIN_LEFT_RIGHT-10, TOP_MARGIN-10, BOARD_WIDTH+20, BOARD_WIDTH+20)
    pygame.draw.rect(screen, BOARD_BG, board_rect)
    pygame.draw.rect(screen, LINE_COLOR, board_rect, 3)
    # 网格
    for i in range(BOARD_SIZE):
        sx, sy = board_to_screen(i, 0)
        ex, ey = board_to_screen(i, BOARD_SIZE-1)
        pygame.draw.line(screen, LINE_COLOR, (sx, sy), (ex, ey), 1)
        sx, sy = board_to_screen(0, i)
        ex, ey = board_to_screen(BOARD_SIZE-1, i)
        pygame.draw.line(screen, LINE_COLOR, (sx, sy), (ex, ey), 1)
    # 星位
    for r, c in STAR_POINTS:
        x, y = board_to_screen(r, c)
        pygame.draw.circle(screen, LINE_COLOR, (x, y), 5)
        pygame.draw.circle(screen, BOARD_BG, (x, y), 2)
    # 棋子
    for r in range(BOARD_SIZE):
        for c in range(BOARD_SIZE):
            if board[r][c] != 0:
                x, y = board_to_screen(r, c)
                draw_stone(x, y, board[r][c])
    # 动画中的棋子
    if animation and animation['stone_placed']:
        row, col = animation['board_pos']
        x, y = board_to_screen(row, col)
        draw_stone(x, y, animation['player'])

def draw_stone(x, y, player):
    """绘制带光泽的棋子"""
    radius = CELL_SIZE // 2 - 2
    if player == 1:
        pygame.draw.circle(screen, BLACK_STONE, (x, y), radius)
        # 高光
        pygame.draw.circle(screen, (100,100,100), (x-3, y-3), radius//3)
    else:
        pygame.draw.circle(screen, WHITE_STONE, (x, y), radius)
        pygame.draw.circle(screen, (150,150,150), (x-3, y-3), radius//3)
    pygame.draw.circle(screen, LINE_COLOR, (x, y), radius, 1)

def draw_ui():
    # 静止时画两个火柴人
    if animation is None or animation['player'] != 1:
        draw_stick_figure(screen, BLACK_REST[0], BLACK_REST[1], BLACK_STONE)
    if animation is None or animation['player'] != 2:
        draw_stick_figure(screen, WHITE_REST[0], WHITE_REST[1], WHITE_STONE)
    # 动画中的火柴人
    if animation:
        pos = animation['current_pos']
        color = BLACK_STONE if animation['player'] == 1 else WHITE_STONE
        draw_stick_figure(screen, pos[0], pos[1], color)
    # 底部状态栏
    pygame.draw.rect(screen, (60,40,30), (0, HEIGHT-BOTTOM_MARGIN, WIDTH, BOTTOM_MARGIN))
    if not game_over:
        if current_player == 1:
            text = "Black's Turn"
        else:
            text = "White's Turn"
        color = GREEN
    else:
        if winner == 1: text = "Black Wins!"
        elif winner == 2: text = "White Wins!"
        else: text = "Draw"
        color = RED
    turn_surf = font.render(text, True, color)
    screen.blit(turn_surf, (20, HEIGHT-BOTTOM_MARGIN+20))
    # 新游戏按钮
    button_rect = pygame.Rect(WIDTH-140, HEIGHT-BOTTOM_MARGIN+10, 120, 40)
    pygame.draw.rect(screen, GOLD, button_rect, border_radius=8)
    pygame.draw.rect(screen, (100,80,50), button_rect, 2, border_radius=8)
    btn_text = font.render("New Game", True, (60,40,20))
    screen.blit(btn_text, (button_rect.x+20, button_rect.y+8))
    return button_rect

def reset_game():
    global board, current_player, game_over, winner, animation
    board = [[0]*BOARD_SIZE for _ in range(BOARD_SIZE)]
    current_player = 1
    game_over = False
    winner = 0
    animation = None

def main():
    global current_player, game_over, winner, animation
    reset_game()
    button_rect = None
    clock = pygame.time.Clock()

    running = True
    while running:
        for event in pygame.event.get():
            if event.type == pygame.QUIT:
                running = False
            elif event.type == pygame.MOUSEBUTTONDOWN and event.button == 1:
                x, y = event.pos
                if button_rect and button_rect.collidepoint(x, y):
                    reset_game()
                elif animation is None and y >= TOP_MARGIN and y < TOP_MARGIN + BOARD_WIDTH:
                    pos = screen_to_board(x, y)
                    if pos:
                        row, col = pos
                        place_stone(row, col)

        if animation:
            update_animation()

        draw_board()
        button_rect = draw_ui()
        pygame.display.flip()
        clock.tick(60)

    pygame.quit()
    sys.exit()

if __name__ == "__main__":
    main()