# PYTHON VERSION 3.8
#------IMPORT PYTHON MODULES---------------
import time
import random 
import numpy as np
import matplotlib
matplotlib.use('Qt5Agg') # Ensure you have PyQt5 installed
from matplotlib import pyplot as plt
import matplotlib.animation as animation
import warnings
import matplotlib.cbook
from matplotlib import MatplotlibDeprecationWarning
warnings.filterwarnings("ignore", category=MatplotlibDeprecationWarning)
#------------------------------------------

start_time = time.time()
current_time = time.strftime("%y-%m-%d-%H.%M.%S", time.localtime()) 

particles_list = []
# Particle Radius
r_ball = 2
# BOX DIMENSIONS
box_size = 200
# GRID SIZE (n x n)
n_grid = 20
total_particles = n_grid * n_grid
# ANIMATION FRAMES
# iterations = int(input("Enter number of iterations: "))
iterations = 200

# CLASS INITIALIZATION---------------------------
class Particle:
    def __init__(self, x, y, dx, dy, radii, color, shape):
        self.dx = dx
        self.dy = dy
        self.shape = shape
        self.color = color
        self.radii = radii
        self.x = x
        self.y = y
    
    def set_x(self, x_new):
        self.x = x_new

    def set_y(self, y_new):
        self.y = y_new

    def get_distance(self, p2):
        rx = self.x - p2.x
        ry = self.y - p2.y
        return np.sqrt(rx*rx + ry*ry)

    def handle_collision(self, b): 
        #----ELASTIC COLLISION LOGIC------------
        a = self
        v1 = np.array([a.dx, a.dy])
        v2 = np.array([b.dx, b.dy])
        x1 = np.array([a.x, a.y])
        x2 = np.array([b.x, b.y])
        
        # Conservation of momentum and kinetic energy formula
        v1_new = v1 - (v1-v2).dot(x1-x2)/(np.linalg.norm(x1-x2))**2 * (x1-x2)
        v2_new = v2 - (v2-v1).dot(x2-x1)/(np.linalg.norm(x2-x1))**2 * (x2-x1)
        return [v1_new, v2_new]

#------------INITIAL ENERGY & VELOCITY--------------
# TEMPERATURE CONTROL
total_energy_target = total_particles * 0.7
initial_energy = np.random.uniform(0.0, 1.5, total_particles) 
# Normalize energy
initial_energy = initial_energy * total_energy_target / initial_energy.sum()

v_initial = np.sqrt(2 * initial_energy)
print(f'Initial system energy = {np.array([p**2 for p in v_initial]).sum()*0.5}')

for i in range(total_particles):
    angle = 2 * np.pi * random.random()
    dx = v_initial[i] * np.cos(angle)
    dy = v_initial[i] * np.sin(angle)
    particles_list.append(Particle(0, 0, dx, dy, r_ball, "red", "circle"))

#--------------INITIAL PARTICLE POSITIONS (GRID) ----------------
counter = 0
for i in range(n_grid):
    for j in range(n_grid):
        x_start = (box_size/2 - box_size/n_grid) / n_grid * 2 * i + box_size/n_grid
        y_start = (box_size/2 - box_size/n_grid) / n_grid * 2 * j + box_size/n_grid
        particles_list[counter].set_x(x_start)
        particles_list[counter].set_y(y_start)
        counter += 1    

# Storage for simulation data
current_v_mags = np.zeros(total_particles)
history_x = np.array([])
history_y = np.array([])
history_v = np.array([])
frame_count = 0

#------SIMULATION LOOP-------------------------
print('Simulating physical interactions...')
step = -1
while (step < iterations):
    step += 1
    idx = 0
    temp_x = np.array([])
    temp_y = np.array([])

    #--------- CHECK INTER-PARTICLE COLLISIONS --------------
    for i in range(0, len(particles_list)):
        for j in range(i + 1, len(particles_list)): 
            if particles_list[i].get_distance(particles_list[j]) < r_ball * 2:
                result = particles_list[i].handle_collision(particles_list[j])
                [particles_list[i].dx, particles_list[i].dy] = result[0]
                [particles_list[j].dx, particles_list[j].dy] = result[1]

    #---------- CHECK WALL COLLISIONS ----------
    for ball in particles_list:
        temp_x = np.append(temp_x, ball.x)
        temp_y = np.append(temp_y, ball.y)

        ball.set_y(ball.y + ball.dy)
        ball.set_x(ball.x + ball.dx)

        # Bounce off Top/Bottom
        if ball.y > box_size - r_ball or ball.y < r_ball:
            ball.dy *= -1
            ball.y = np.clip(ball.y, r_ball, box_size - r_ball)

        # Bounce off Sides
        if ball.x > box_size - r_ball or ball.x < r_ball:
            ball.dx *= -1
            ball.x = np.clip(ball.x, r_ball, box_size - r_ball)

        current_v_mags[idx] = np.sqrt(ball.dx**2 + ball.dy**2)
        idx += 1

    # Record data every step
    if step == 0:
        history_v = current_v_mags
        history_x = temp_x
        history_y = temp_y
    else:
        history_v = np.vstack([history_v, current_v_mags])
        history_x = np.vstack([history_x, temp_x])
        history_y = np.vstack([history_y, temp_y])
    
    frame_count += 1

    if step % 10 == 0:
        print(f'Step: {step} | System Energy: {np.array([p**2 for p in current_v_mags]).sum()*0.5:.5f}')

print('===========================================')
print(f"Simulation complete. Elapsed time: {time.time()-start_time:.2f} seconds")

#-------- THEORETICAL PDF CURVE --------------
v_range = np.linspace(0, 3, 100)
def maxwell_boltzmann_2d(v, energy_factor):
    # beta is inversely proportional to temperature
    beta = 1.0 / energy_factor 
    return beta * v * np.exp(-(beta * v**2) / 2)

def update_animation(frame_idx, x_data, y_data):
    # Clear and redraw the particle plot
    plt.subplot(211)
    plt.cla()
    plt.scatter(x_data[frame_idx], y_data[frame_idx], s=18, color='royalblue', alpha=0.7)
    plt.xlim(0, box_size)
    plt.ylim(0, box_size)
    plt.title('Particle Position Plot')

    # Clear and redraw the speed distribution
    plt.subplot(212)
    plt.cla()
    weights = np.ones_like(history_v[frame_idx]) / len(history_v[frame_idx])
    plt.hist(history_v[frame_idx], bins=np.linspace(0, 3, 18), edgecolor='black', density=True, alpha=0.6, label='Simulated')

    # --- UPDATED SECTION ---
    # Calculate the scale based on your target energy
    energy_factor = total_energy_target / total_particles
    
    # Pass the energy_factor to the theoretical curve function
    plt.plot(v_range, maxwell_boltzmann_2d(v_range, energy_factor), color='red', lw=2, label='Theoretical PDF')
    # -----------------------

    energy_val = np.array([p**2 for p in history_v[frame_idx]]).sum() * 0.5
    plt.text(1.2, 0.37, f'System Energy = {energy_val:.5f}')
    plt.text(1.2, 0.34, f'Iteration: {frame_idx}')

    plt.title('Speed Distribution')
    plt.xlabel('Velocity (v)')
    plt.ylabel('Probability Density')
    plt.ylim(0, 1.2)
    plt.xlim(0, 3)
    plt.legend(loc='upper right')

# --- INTERACTIVE MENU ---
while True:
    print("\n--- OPTIONS ---")
    print("1: Show Live Plot")
    print("2: Save to MP4 Video")
    print("0: Exit")
    user_choice = input('Select an option: ')

    if user_choice == '1':
        fig = plt.figure(figsize=(6, 10))
        anim = animation.FuncAnimation(fig, update_animation, frames=frame_count, fargs=(history_x, history_y), interval=20)
        plt.tight_layout()
        plt.show()
    elif user_choice == '2':
        fig = plt.figure(figsize=(6, 10))
        anim = animation.FuncAnimation(fig, update_animation, frames=frame_count, fargs=(history_x, history_y), interval=20)
        # Note: ffmpeg must be installed and path configured correctly
        # plt.rcParams['animation.ffmpeg_path'] = r'C:\path\to\ffmpeg.exe'
        writer = animation.FFMpegWriter(fps=15)
        output_name = f'simulation_result_{current_time}.mp4'
        print(f"Saving to {output_name}...")
        anim.save(output_name, writer=writer)
        print("Save complete.")
        break
    elif user_choice == '0':
        break

print(f"Final Energy Check: {np.array([p**2 for p in history_v[-1]]).sum()*0.5:.5f}")
