import taichi as ti
import numpy as np

ti.init(arch=ti.cuda, default_fp=ti.f32)


"""
Finite Difference Grid Setup
"""
N_NODES = 1000
L = 100
T = 2
v = L/T
dx = L/(N_NODES-1)
courant_number = 1
dt = courant_number * dx / v
s = v*dt/dx
s2 = s*s

"""
Damping Params
"""
USE_DAMPING=False
gamma = -0.0001
beta = gamma * dt / (dx**2)

now = ti.field(float,shape=())
"""
Finite Difference Coeff Matrix
Update coeffs derived from standard expression for finite differences via explicit scheme
apply to u_(t,i) ; u_(t,i-1) ; u_(t,i+1) ; u_(t-1,i)
"""
coeffs = ti.Matrix([2*(1-s2), s2, s2, -1])
coeffs_damping = ti.Matrix([2*(1-beta-s2), s2 + beta, s2 + beta, 2*beta -1, -beta, -beta])

"""
Initial Conditions Consts
"""
impulse_types = {
    "null": 0,
    "gaussian": 1,
    "splitting_gaussian": 2,
    "impulse": 3,
    "pulse_width": 4,
    "tri": 5,
    "fundamental": 6,
    "sine": 7,
    "harmonic": 8,
}
IMPULSE_MODE = impulse_types["sine"]

"""
WAGGLE
"""
WAGGLE = False
waggle_amp = 0.1
waggle_freq = 9.239*1/T

"""
Render/Window Config
"""
Y_VIEW_SCALE = 0.4
Y_VIEW_OFFSET = 0.5
UPDATE_BATCH_SIZE = int(N_NODES / 1000) # Updates / per frame

"""
Node Dataclass
"""
@ti.dataclass
class Node:
    i: int
    u: float
    u_next: float
    u_last: float
    x: float
    x_normalized: float

    @ti.func
    def compute_next_u(self, neighbor_l, neighbor_r, damping: bool):#neighbor_l_u: float, neighbor_r_u: float):
        if not damping:
            grid_vect = ti.Matrix.cols([[ self.u,  neighbor_l.u, neighbor_r.u, self.u_last]]) 
            self.u_next = (coeffs @ grid_vect)[0]
        else:
            grid_vect = ti.Matrix.cols([[ self.u,  neighbor_l.u, neighbor_r.u, self.u_last, neighbor_l.u_last, neighbor_r.u_last]]) 
            self.u_next = (coeffs_damping @ grid_vect)[0]

    
    @ti.func
    def shift_fields(self):
       self.u_last = self.u
       self.u = self.u_next 

    @ti.func 
    def apply_waggle(self, amp, freq, t):
        self.u_next = amp*ti.sin(freq*2*np.pi * t)

    @ti.func
    def apply_initial(self, mode: int):
        mean = 0.5
        sigma = 0.05
        if mode == impulse_types["null"]: # Moving Guassian
            self.u      = 0
            self.u_last = 0
        if mode == impulse_types["gaussian"]: # Moving Guassian
            self.u      = 0.1/ti.sqrt(2*np.pi*(sigma**2)) * ti.exp((-(self.x_normalized - mean)**2)/(2*(sigma**2)))
            self.u_last = 0.1/ti.sqrt(2*np.pi*(sigma**2)) * ti.exp((-(self.x_normalized - mean -v*dt/L)**2)/(2*(sigma**2)))
        if mode == impulse_types["splitting_gaussian"]: # Splitting Guassian
            self.u      = 0.1/ti.sqrt(2*np.pi*(sigma**2)) * ti.exp((-(self.x_normalized - mean)**2)/(2*(sigma**2)))
            self.u_last = 0.1/ti.sqrt(2*np.pi*(sigma**2)) * ti.exp((-(self.x_normalized - mean)**2)/(2*(sigma**2)))
        if mode == impulse_types["impulse"]:
            if self.i == int(N_NODES/2):
                self.u = 0.5
                self.u_last = 0.5
            else:
                self.u = 0
                self.u_last = 0
        if mode == impulse_types["pulse_width"]:
            if self.i > int(N_NODES/2)-10 and self.i < int(N_NODES/2)+10:
                self.u = 0.5
                self.u_last = 0.5
            else:
                self.u = 0
                self.u_last = 0
        if mode == impulse_types["tri"]:
            self.u = 0.5-ti.abs(0.5-self.x_normalized)
            self.u_last = 0.5-ti.abs(0.5-self.x_normalized)
        if mode == impulse_types["fundamental"]:
            self.u = ti.sin(self.x_normalized*np.pi)
            self.u_last = ti.sin(self.x_normalized*np.pi)
        if mode == impulse_types["sine"]:
            self.u = ti.sin(self.x_normalized*2*np.pi)
            self.u_last = ti.sin(self.x_normalized*2*np.pi)
        if mode == impulse_types["harmonic"]:
            self.u = ti.sin(self.x_normalized*3*np.pi)
            self.u_last = ti.sin(self.x_normalized*3*np.pi)

"""
Geometry and Node Fields
"""
line_points = ti.Vector.field(2,dtype=float, shape=(2*(N_NODES-1),))
circ_points = ti.Vector.field(2,dtype=float, shape=(N_NODES,))
nodes = Node.field(shape=(N_NODES))

"""
Kernels
"""
@ti.kernel
def init_nodes():
    for i in nodes:
        nodes[i].i = i
        nodes[i].x = L*i/(N_NODES-1)
        nodes[i].x_normalized = i / (N_NODES-1)
        nodes[i].u_next = 0
        nodes[i].apply_initial(mode=IMPULSE_MODE)
        if i == 0 or i == N_NODES-1:
            nodes[i].u = 0
            nodes[i].u_last = 0


@ti.kernel
def update_geo():
    for i in nodes:
        p_x = i / (N_NODES-1)
        p_y = nodes[i].u * Y_VIEW_SCALE + Y_VIEW_OFFSET

        circ_points[i].x  = p_x
        circ_points[i].y  = p_y 

        if (i != N_NODES -1):
            line_points[2*i].x = p_x
            line_points[2*i+1].x =  (i+1) / (N_NODES-1)
            line_points[2*i].y = p_y
            line_points[2*i+1].y = nodes[i+1].u * Y_VIEW_SCALE + Y_VIEW_OFFSET

@ti.kernel
def compute_next():
    for i in nodes:
        # Skip boundaries
        if i != 0 and i != N_NODES-1:
            nodes[i].compute_next_u(nodes[i-1], nodes[i+1], damping=USE_DAMPING)
        elif i ==0 and WAGGLE:
            nodes[i].apply_waggle(waggle_amp, waggle_freq, now[None])
    now[None] += dt

@ti.kernel
def shift_nodes():
    for i in nodes:
        if i != 0 and i != nodes.shape[0]:
            nodes[i].shift_fields()
        elif i==0 and WAGGLE:
            nodes[i].shift_fields()

"""
MAIN
"""
import time
if __name__ == '__main__':

    window = ti.ui.Window('Window Title', res = (1000,1000), pos = (150, 150))
    canvas = window.get_canvas()

    init_nodes()
    update_geo()


    while window.running:
        # time.sleep(0.030)
        canvas.set_background_color(color=(0,0,0))
        canvas.lines(line_points, 0.002, indices=None, color=(1,1,1))

        for i in range(UPDATE_BATCH_SIZE):
            compute_next()
            shift_nodes()
        update_geo()
        window.show()




