import tkinter as tk
from tkinter import messagebox
import torch
import chess
import chess.svg
import hashlib
import numpy as np
from PIL import Image, ImageTk
import io

# ------------------------------------------------------------
# 1. Feature extraction (same as before)
# ------------------------------------------------------------
def board_to_tensor(board: chess.Board) -> torch.Tensor:
    piece_map = board.piece_map()
    tensor = torch.zeros(64, 12, dtype=torch.float32)
    piece_to_idx = {
        chess.PAWN: 0, chess.KNIGHT: 1, chess.BISHOP: 2,
        chess.ROOK: 3, chess.QUEEN: 4, chess.KING: 5
    }
    for sq, piece in piece_map.items():
        idx = piece_to_idx[piece.piece_type]
        if piece.color == chess.BLACK:
            idx += 6
        tensor[sq, idx] = 1.0
    return tensor.flatten()

# ------------------------------------------------------------
# 2. Telepathic memory generator (interactive version)
# ------------------------------------------------------------
class TelepathicChessMemory:
    def __init__(self, board_features: torch.Tensor, hidden_dim=32,
                 gamma=0.2, epsilon=1e-4, max_iter=100):
        self.board = board_features
        self.hidden_dim = hidden_dim
        self.gamma = gamma
        self.epsilon = epsilon
        self.max_iter = max_iter

        self.W_board = torch.randn(len(board_features), hidden_dim) * 0.1
        self.W_mem   = torch.randn(hidden_dim, hidden_dim) * 0.1

        self.xi = self._compute_xi()
        self.memory = torch.zeros(hidden_dim)

    def _compute_xi(self) -> torch.Tensor:
        board_bytes = self.board.numpy().tobytes()
        h = hashlib.md5(board_bytes).hexdigest()
        n = self.hidden_dim
        bytes_needed = n * 4
        hex_bytes = h[:bytes_needed*2]
        if len(hex_bytes) < bytes_needed*2:
            hex_bytes = hex_bytes.ljust(bytes_needed*2, '0')
        arr = np.frombuffer(bytes.fromhex(hex_bytes), dtype=np.float32)
        if len(arr) < n:
            arr = np.pad(arr, (0, n - len(arr)))
        else:
            arr = arr[:n]
        xi = torch.tensor(arr)
        return xi / (xi.norm() + 1e-8)

    def telepath(self) -> torch.Tensor:
        board_proj = self.board @ self.W_board
        mem_proj   = self.memory @ self.W_mem
        target = torch.sigmoid(board_proj + mem_proj)
        grad = target - torch.sigmoid(self.memory)
        return grad

    def addp(self, grad: torch.Tensor):
        self.memory += self.gamma * grad
        self.memory = torch.clamp(self.memory, -1.0, 1.0)

    def collapse_if_divergent(self) -> bool:
        if torch.norm(self.memory - self.xi) > 0.8:
            self.memory = self.xi.clone()
            return True
        return False

    def run(self) -> torch.Tensor:
        for _ in range(self.max_iter):
            prev = self.memory.clone()
            grad = self.telepath()
            self.addp(grad)
            self.collapse_if_divergent()
            if torch.norm(self.memory - prev) < self.epsilon:
                break
        return self.memory

# ------------------------------------------------------------
# 3. GUI using tkinter
# ------------------------------------------------------------
class TelepathicChessGUI:
    def __init__(self, master):
        self.master = master
        master.title("♜ Telepathic Chess Memory ♞")
        master.geometry("900x600")

        self.board = chess.Board()
        self.selected_square = None
        self.prev_memory = None

        # Canvas for chess board (500x500)
        self.canvas = tk.Canvas(master, width=500, height=500, bg="white")
        self.canvas.pack(side=tk.LEFT, padx=10, pady=10)
        self.canvas.bind("<Button-1>", self.on_click)

        # Right panel for info
        self.info_frame = tk.Frame(master)
        self.info_frame.pack(side=tk.RIGHT, fill=tk.BOTH, expand=True, padx=10, pady=10)

        self.metric_label = tk.Label(self.info_frame, text="Memory Metric", font=("Arial", 14, "bold"))
        self.metric_label.pack(pady=5)

        self.change_label = tk.Label(self.info_frame, text="Change: --", font=("Arial", 12))
        self.change_label.pack(anchor="w")

        self.xi_label = tk.Label(self.info_frame, text="Distance to Ξ: --", font=("Arial", 12))
        self.xi_label.pack(anchor="w")

        self.memory_values_label = tk.Label(self.info_frame, text="Memory (first 5): --", font=("Arial", 10), wraplength=300)
        self.memory_values_label.pack(anchor="w", pady=10)

        self.reset_button = tk.Button(self.info_frame, text="Reset Board", command=self.reset_board)
        self.reset_button.pack(pady=5)

        self.quit_button = tk.Button(self.info_frame, text="Quit", command=master.quit)
        self.quit_button.pack(pady=5)

        # Initial memory
        self.update_memory_for_current_board()
        self.draw_board()

    def get_square_from_click(self, x, y):
        """Convert canvas coordinates to chess square (0..63)."""
        size = 500 // 8
        col = x // size
        row = 7 - (y // size)  # because rank 8 is y=0
        if 0 <= col < 8 and 0 <= row < 8:
            return row * 8 + col
        return None

    def on_click(self, event):
        sq = self.get_square_from_click(event.x, event.y)
        if sq is None:
            return

        if self.selected_square is None:
            # First click: select piece if it belongs to current player
            piece = self.board.piece_at(sq)
            if piece and piece.color == self.board.turn:
                self.selected_square = sq
                self.draw_board(highlight=sq)
            else:
                pass
        else:
            # Second click: try to move
            move = chess.Move(self.selected_square, sq)
            # Check promotion (simplified: always promote to queen)
            if self.board.piece_at(self.selected_square) and self.board.piece_at(self.selected_square).piece_type == chess.PAWN:
                if chess.square_rank(sq) in (0, 7):
                    move = chess.Move(self.selected_square, sq, promotion=chess.QUEEN)

            if move in self.board.legal_moves:
                self.board.push(move)
                self.update_memory_for_current_board()
                self.draw_board()
                # Check game over
                if self.board.is_game_over():
                    messagebox.showinfo("Game Over", "Game finished!\n" + self.board.result())
                    self.reset_board()
            else:
                # Illegal move: clear selection
                pass
            self.selected_square = None
            self.draw_board()

    def draw_board(self, highlight=None):
        """Draw chess board with pieces (using Unicode on canvas text)."""
        self.canvas.delete("all")
        size = 500 // 8
        colors = ["#F0D9B5", "#B58863"]  # light, dark
        for row in range(8):
            for col in range(8):
                x1 = col * size
                y1 = (7 - row) * size
                x2 = x1 + size
                y2 = y1 + size
                color = colors[(row + col) % 2]
                self.canvas.create_rectangle(x1, y1, x2, y2, fill=color, outline="black")
                sq = row * 8 + col
                piece = self.board.piece_at(sq)
                if piece:
                    # Unicode symbols
                    symbols = {
                        "K": "♔", "Q": "♕", "R": "♖", "B": "♗", "N": "♘", "P": "♙",
                        "k": "♚", "q": "♛", "r": "♜", "b": "♝", "n": "♞", "p": "♟"
                    }
                    symbol = symbols[piece.symbol()]
                    self.canvas.create_text(x1 + size/2, y1 + size/2, text=symbol, font=("Arial", size-10), fill="black")
        # Highlight selected square
        if highlight is not None:
            row = highlight // 8
            col = highlight % 8
            x1 = col * size
            y1 = (7 - row) * size
            self.canvas.create_rectangle(x1, y1, x1+size, y1+size, outline="yellow", width=4)

    def update_memory_for_current_board(self):
        """Compute stationary memory for current board and update GUI labels."""
        features = board_to_tensor(self.board)
        tmem = TelepathicChessMemory(features, hidden_dim=32, gamma=0.2)
        new_memory = tmem.run()

        if self.prev_memory is None:
            self.prev_memory = new_memory
            change = 0.0
        else:
            change = torch.norm(new_memory - self.prev_memory).item()

        divergence_xi = torch.norm(new_memory - tmem.xi).item()
        mem_vals = new_memory[:5].tolist()
        mem_str = ", ".join(f"{v:.3f}" for v in mem_vals)

        self.change_label.config(text=f"Change from previous: {change:.4f}")
        self.xi_label.config(text=f"Distance to Ξ: {divergence_xi:.4f}")
        self.memory_values_label.config(text=f"Memory (first 5): [{mem_str}]")

        self.prev_memory = new_memory

    def reset_board(self):
        self.board = chess.Board()
        self.selected_square = None
        self.prev_memory = None
        self.update_memory_for_current_board()
        self.draw_board()

# ------------------------------------------------------------
# 4. Run the GUI
# ------------------------------------------------------------
if __name__ == "__main__":
    root = tk.Tk()
    app = TelepathicChessGUI(root)
    root.mainloop()