import math
import random
import time
import tkinter as tk
from tkinter import messagebox
from tkinter import ttk

RAINBOW = ("#FF0000", "#FF8C00", "#FFFF00", "#00D000", "#1E90FF", "#4B0082", "#B000FF")
BG = "#080C18"
CAM_Z = 4.0
LEVELS = 14
NEIGHBORS = (
    (1, -1, -1), (1, -1, 0), (1, -1, 1),
    (1, 0, -1), (1, 0, 0), (1, 0, 1),
    (1, 1, -1), (1, 1, 0), (1, 1, 1),
    (0, 1, -1), (0, 1, 0), (0, 1, 1),
    (0, 0, 1),
)


def hex_rgb(color):
    return int(color[1:3], 16), int(color[3:5], 16), int(color[5:7], 16)


def mix(color_a, color_b, t):
    ra, ga, ba = hex_rgb(color_a)
    rb, gb, bb = hex_rgb(color_b)
    red = int(ra + (rb - ra) * t)
    green = int(ga + (gb - ga) * t)
    blue = int(ba + (bb - ba) * t)
    return "#%02x%02x%02x" % (red, green, blue)


def build_shades():
    body = []
    rim = []
    glow = []
    for base in RAINBOW:
        for k in range(LEVELS):
            far = 1.0 - k / (LEVELS - 1)
            body.append(mix(base, BG, 0.60 * far))
            rim.append(mix(base, "#000000", 0.25 + 0.35 * far))
            glow.append(mix(base, "#FFFFFF", 0.55 - 0.25 * far))
    return body, rim, glow


class GasApp:
    def __init__(self, root):
        self.root = root
        root.title("Движение молекул идеального газа")
        root.geometry("1180x800")
        root.minsize(900, 620)

        self.body, self.rim, self.glow = build_shades()
        self.count = 160
        self.speed = 1.2
        self.diameter = 5.0
        self.rw = 0.012
        self.running = True
        self.yaw = 0.62
        self.pitch = 0.34
        self.cyw = math.cos(self.yaw)
        self.syw = math.sin(self.yaw)
        self.cpt = math.cos(self.pitch)
        self.spt = math.sin(self.pitch)

        self.mx = []
        self.my = []
        self.mz = []
        self.vx = []
        self.vy = []
        self.vz = []
        self.mcol = []

        self.edges = self.make_edges()
        self.last = time.perf_counter()
        self.fps_time = self.last
        self.frames = 0
        self.fps = 0.0
        self.drag = None

        self.collisions = tk.BooleanVar(value=True)
        self.autorotate = tk.BooleanVar(value=True)

        self.build_ui()
        self.reset()
        self.loop()

    def make_edges(self):
        edges = []
        for axis in range(3):
            b = (axis + 1) % 3
            c = (axis + 2) % 3
            for sb in (-1, 1):
                for sc in (-1, 1):
                    p1 = [0.0, 0.0, 0.0]
                    p2 = [0.0, 0.0, 0.0]
                    p1[axis] = -1.0
                    p2[axis] = 1.0
                    p1[b] = p2[b] = float(sb)
                    p1[c] = p2[c] = float(sc)
                    f1 = b * 2 + (1 if sb > 0 else 0)
                    f2 = c * 2 + (1 if sc > 0 else 0)
                    edges.append((tuple(p1), tuple(p2), f1, f2))
        return edges

    def build_ui(self):
        bar = ttk.Frame(self.root, padding=(10, 8, 10, 4))
        bar.pack(side="top", fill="x")

        ttk.Label(bar, text="Молекул:").grid(row=0, column=0, sticky="w")
        self.count_scale = ttk.Scale(bar, from_=1, to=500, length=170)
        self.count_scale.set(self.count)
        self.count_scale.grid(row=0, column=1, padx=(6, 4))
        self.count_lbl = ttk.Label(bar, text=str(self.count), width=4, anchor="w")
        self.count_lbl.grid(row=0, column=2, sticky="w")

        ttk.Label(bar, text="Начальная скорость:").grid(row=0, column=3, sticky="w", padx=(20, 0))
        self.speed_scale = ttk.Scale(bar, from_=0.1, to=4.0, length=170)
        self.speed_scale.set(self.speed)
        self.speed_scale.grid(row=0, column=4, padx=(6, 4))
        self.speed_lbl = ttk.Label(bar, text="%.2f" % self.speed, width=5, anchor="w")
        self.speed_lbl.grid(row=0, column=5, sticky="w")

        ttk.Label(bar, text="Диаметр шарика, px:").grid(row=0, column=6, sticky="w", padx=(20, 0))
        self.diam_scale = ttk.Scale(bar, from_=2, to=24, length=170)
        self.diam_scale.set(self.diameter)
        self.diam_scale.grid(row=0, column=7, padx=(6, 4))
        self.diam_lbl = ttk.Label(bar, text="%.0f" % self.diameter, width=4, anchor="w")
        self.diam_lbl.grid(row=0, column=8, sticky="w")

        ttk.Checkbutton(bar, text="Столкновения молекул", variable=self.collisions).grid(
            row=1, column=0, columnspan=2, sticky="w", pady=(10, 0))
        ttk.Checkbutton(bar, text="Вращение камеры", variable=self.autorotate).grid(
            row=1, column=3, columnspan=2, sticky="w", pady=(10, 0))
        self.run_btn = ttk.Button(bar, text="Пауза", command=self.toggle_run)
        self.run_btn.grid(row=1, column=6, sticky="w", pady=(10, 0))
        ttk.Button(bar, text="Сброс", command=self.reset).grid(row=1, column=7, sticky="w", pady=(10, 0))
        ttk.Button(bar, text="Случайный ракурс", command=self.random_view).grid(
            row=1, column=8, sticky="w", pady=(10, 0))

        self.status = ttk.Label(self.root, anchor="w", padding=(12, 2, 12, 8))
        self.status.pack(side="bottom", fill="x")
        self.status.configure(text="Мышью по полю — вращение камеры")

        self.canvas = tk.Canvas(self.root, bg=BG, highlightthickness=0, bd=0)
        self.canvas.pack(side="top", fill="both", expand=True, padx=10, pady=(4, 4))
        self.canvas.bind("<ButtonPress-1>", self.on_press)
        self.canvas.bind("<B1-Motion>", self.on_drag)
        self.canvas.bind("<ButtonRelease-1>", self.on_release)

        self.count_scale.configure(command=self.on_count)
        self.speed_scale.configure(command=self.on_speed)
        self.diam_scale.configure(command=self.on_diam)

    def on_press(self, event):
        self.drag = (event.x, event.y)

    def on_drag(self, event):
        if self.drag is None:
            return
        dx = event.x - self.drag[0]
        dy = event.y - self.drag[1]
        self.drag = (event.x, event.y)
        self.yaw += dx * 0.01
        self.pitch += dy * 0.01
        if self.pitch > 1.45:
            self.pitch = 1.45
        elif self.pitch < -1.45:
            self.pitch = -1.45

    def on_release(self, event):
        self.drag = None

    def random_view(self):
        self.yaw = random.uniform(-math.pi, math.pi)
        self.pitch = random.uniform(-0.9, 0.9)

    def on_count(self, value):
        try:
            target = int(float(value))
        except (TypeError, ValueError):
            return
        if target < 1:
            target = 1
        self.count = target
        self.count_lbl.configure(text=str(target))
        self.set_count(target)

    def on_speed(self, value):
        try:
            speed = float(value)
        except (TypeError, ValueError):
            return
        if speed < 0.01:
            speed = 0.01
        self.speed = speed
        self.speed_lbl.configure(text="%.2f" % speed)
        self.set_speeds(speed)

    def on_diam(self, value):
        try:
            diameter = float(value)
        except (TypeError, ValueError):
            return
        if diameter < 1.0:
            diameter = 1.0
        self.diameter = diameter
        self.diam_lbl.configure(text="%.0f" % diameter)

    def set_count(self, target):
        while len(self.mx) < target:
            self.add_molecule()
        if len(self.mx) > target:
            del self.mx[target:]
            del self.my[target:]
            del self.mz[target:]
            del self.vx[target:]
            del self.vy[target:]
            del self.vz[target:]
            del self.mcol[target:]

    def set_speeds(self, speed):
        for i in range(len(self.mx)):
            magnitude = math.sqrt(self.vx[i] ** 2 + self.vy[i] ** 2 + self.vz[i] ** 2)
            if magnitude < 1e-9:
                u = random.uniform(-1.0, 1.0)
                phi = random.uniform(0.0, 2.0 * math.pi)
                s = math.sqrt(max(0.0, 1.0 - u * u))
                self.vx[i] = s * math.cos(phi) * speed
                self.vy[i] = s * math.sin(phi) * speed
                self.vz[i] = u * speed
            else:
                k = speed / magnitude
                self.vx[i] *= k
                self.vy[i] *= k
                self.vz[i] *= k

    def add_molecule(self):
        limit = 1.0 - self.rw
        if limit < 0.05:
            limit = 0.05
        self.mx.append(random.uniform(-limit, limit))
        self.my.append(random.uniform(-limit, limit))
        self.mz.append(random.uniform(-limit, limit))
        u = random.uniform(-1.0, 1.0)
        phi = random.uniform(0.0, 2.0 * math.pi)
        s = math.sqrt(max(0.0, 1.0 - u * u))
        speed = self.speed
        self.vx.append(s * math.cos(phi) * speed)
        self.vy.append(s * math.sin(phi) * speed)
        self.vz.append(u * speed)
        self.mcol.append(random.randrange(len(RAINBOW)))

    def reset(self):
        self.mx.clear()
        self.my.clear()
        self.mz.clear()
        self.vx.clear()
        self.vy.clear()
        self.vz.clear()
        self.mcol.clear()
        self.set_count(self.count)

    def toggle_run(self):
        self.running = not self.running
        self.run_btn.configure(text="Пауза" if self.running else "Старт")

    @staticmethod
    def bounce(pos, vel, dt, limit):
        p = pos + vel * dt
        if p > limit:
            p = 2.0 * limit - p
            vel = -vel
        elif p < -limit:
            p = -2.0 * limit - p
            vel = -vel
        if p > limit:
            p = limit
        elif p < -limit:
            p = -limit
        return p, vel

    def step(self, dt):
        limit = 1.0 - self.rw
        if limit < 0.05:
            limit = 0.05
        mx = self.mx
        my = self.my
        mz = self.mz
        vx = self.vx
        vy = self.vy
        vz = self.vz
        for i in range(len(mx)):
            p, v = self.bounce(mx[i], vx[i], dt, limit)
            mx[i] = p
            vx[i] = v
            p, v = self.bounce(my[i], vy[i], dt, limit)
            my[i] = p
            vy[i] = v
            p, v = self.bounce(mz[i], vz[i], dt, limit)
            mz[i] = p
            vz[i] = v

    def resolve(self):
        rw = self.rw
        if rw <= 0.0:
            return
        cell = 2.0 * rw
        inv = 1.0 / cell
        mx = self.mx
        my = self.my
        mz = self.mz
        grid = {}
        for i in range(len(mx)):
            key = (int((mx[i] + 1.0) * inv), int((my[i] + 1.0) * inv), int((mz[i] + 1.0) * inv))
            if key in grid:
                grid[key].append(i)
            else:
                grid[key] = [i]
        rad = 2.0 * rw
        rad2 = rad * rad
        for key, ids in grid.items():
            kx, ky, kz = key
            count = len(ids)
            for a in range(count):
                i = ids[a]
                for b in range(a + 1, count):
                    self.pair(i, ids[b], rad, rad2)
                for ox, oy, oz in NEIGHBORS:
                    other = grid.get((kx + ox, ky + oy, kz + oz))
                    if other:
                        for j in other:
                            self.pair(i, j, rad, rad2)

    def pair(self, i, j, rad, rad2):
        dx = self.mx[j] - self.mx[i]
        dy = self.my[j] - self.my[i]
        dz = self.mz[j] - self.mz[i]
        d2 = dx * dx + dy * dy + dz * dz
        if d2 >= rad2 or d2 <= 1e-16:
            return
        d = math.sqrt(d2)
        nx = dx / d
        ny = dy / d
        nz = dz / d
        rel = (self.vx[j] - self.vx[i]) * nx + (self.vy[j] - self.vy[i]) * ny + (self.vz[j] - self.vz[i]) * nz
        if rel < 0.0:
            self.vx[i] += rel * nx
            self.vy[i] += rel * ny
            self.vz[i] += rel * nz
            self.vx[j] -= rel * nx
            self.vy[j] -= rel * ny
            self.vz[j] -= rel * nz
        push = (rad - d) * 0.5
        self.mx[i] -= nx * push
        self.my[i] -= ny * push
        self.mz[i] -= nz * push
        self.mx[j] += nx * push
        self.my[j] += ny * push
        self.mz[j] += nz * push

    def render(self):
        canvas = self.canvas
        w = canvas.winfo_width()
        h = canvas.winfo_height()
        if w < 60 or h < 60:
            return
        cyw = math.cos(self.yaw)
        syw = math.sin(self.yaw)
        cpt = math.cos(self.pitch)
        spt = math.sin(self.pitch)
        self.cyw = cyw
        self.syw = syw
        self.cpt = cpt
        self.spt = spt
        base = min(w, h) * 0.25
        focal = base * CAM_Z
        rw = (self.diameter * 0.5) / base
        if rw > 0.2:
            rw = 0.2
        self.rw = rw
        cx = w * 0.5
        cy = h * 0.5
        canvas.delete("all")

        visible = [False] * 6
        for axis in range(3):
            for slot, sign in enumerate((-1.0, 1.0)):
                nx = sign if axis == 0 else 0.0
                ny = sign if axis == 1 else 0.0
                nz = sign if axis == 2 else 0.0
                x1 = nx * cyw + nz * syw
                z1 = -nx * syw + nz * cyw
                y2 = ny * cpt - z1 * spt
                z2 = ny * spt + z1 * cpt
                visible[axis * 2 + slot] = (z2 * CAM_Z - 1.0) > 0.0

        def screen(x, y, z):
            x1 = x * cyw + z * syw
            z1 = -x * syw + z * cyw
            y2 = y * cpt - z1 * spt
            z2 = y * spt + z1 * cpt
            k = focal / (CAM_Z - z2)
            return cx + x1 * k, cy - y2 * k

        for p1, p2, f1, f2 in self.edges:
            x1, y1 = screen(p1[0], p1[1], p1[2])
            x2, y2 = screen(p2[0], p2[1], p2[2])
            if visible[f1] or visible[f2]:
                canvas.create_line(x1, y1, x2, y2, fill="#3C5C86")
            else:
                canvas.create_line(x1, y1, x2, y2, fill="#1C2740", dash=(5, 5))

        mx = self.mx
        my = self.my
        mz = self.mz
        mcol = self.mcol
        dn = CAM_Z - 1.75
        df = CAM_Z + 1.75
        span = df - dn
        last = LEVELS - 1
        items = []
        for i in range(len(mx)):
            x = mx[i]
            y = my[i]
            z = mz[i]
            x1 = x * cyw + z * syw
            z1 = -x * syw + z * cyw
            y2 = y * cpt - z1 * spt
            z2 = y * spt + z1 * cpt
            depth = CAM_Z - z2
            k = focal / depth
            radius = rw * k
            if radius < 0.6:
                radius = 0.6
            b = (df - depth) / span
            level = int(b * last + 0.5)
            if level < 0:
                level = 0
            elif level > last:
                level = last
            items.append((depth, cx + x1 * k, cy - y2 * k, radius, mcol[i] * LEVELS + level))
        items.sort(key=lambda item: -item[0])

        for depth, px, py, radius, idx in items:
            if radius < 2.2:
                canvas.create_oval(px - radius, py - radius, px + radius, py + radius,
                                   fill=self.body[idx], outline="")
            else:
                canvas.create_oval(px - radius, py - radius, px + radius, py + radius,
                                   fill=self.body[idx], outline=self.rim[idx])
                if radius >= 4.5:
                    hr = radius * 0.45
                    hx = px - radius * 0.33
                    hy = py - radius * 0.33
                    canvas.create_oval(hx - hr, hy - hr, hx + hr, hy + hr,
                                       fill=self.glow[idx], outline="")

    def update_status(self):
        n = len(self.mx)
        total = 0.0
        for i in range(n):
            total += math.sqrt(self.vx[i] ** 2 + self.vy[i] ** 2 + self.vz[i] ** 2)
        average = total / n if n else 0.0
        self.status.configure(
            text="Молекул: %d    Кадров/с: %.0f    Средняя скорость: %.2f ед/с    "
                 "Перетаскивайте мышью по полю, чтобы повернуть камеру" % (n, self.fps, average))

    def loop(self):
        try:
            now = time.perf_counter()
            dt = now - self.last
            self.last = now
            if dt > 0.05:
                dt = 0.05
            elif dt < 0.0:
                dt = 0.0
            if self.autorotate.get():
                self.yaw += 0.25 * dt
            if self.running:
                self.step(dt)
                if self.collisions.get():
                    self.resolve()
            self.render()
            self.frames += 1
            if now - self.fps_time >= 0.5:
                self.fps = self.frames / (now - self.fps_time)
                self.frames = 0
                self.fps_time = now
                self.update_status()
            self.root.after(16, self.loop)
        except Exception as exc:
            messagebox.showerror("Ошибка", "Произошла ошибка:\n%s" % exc)


def main():
    try:
        root = tk.Tk()
        GasApp(root)
        root.mainloop()
    except Exception as exc:
        messagebox.showerror("Ошибка", "Не удалось запустить программу:\n%s" % exc)


if __name__ == "__main__":
    main()