import taichi as ti
import numpy as np

from utils import Sinusoid

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


"""
Finite Difference Grid Setup
"""
N_NODES = 500
dx = 1
D = 1
L = (N_NODES-1)*dx
# courant condition: 2D*Delta T/Delta x^2 <= 1
dt = 0.5*(dx**2)/D 
s = D*dt/(dx**2)

UPDATE_BATCH_SIZE = 10# Timesteps / per frame

now = ti.field(float,shape=())
"""
Finite Difference Coeff Matrix
Update coeffs derived from standard expression for finite differences via explicit scheme
see colab
"""
coeffs = ti.Matrix([1-2*s, s, s])

"""
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["null"]

"""
WAGGLE
"""
WAGGLE = True


annual_to_daily_ratio = 1/50
daily = 0.0005*dt
daily_sine = Sinusoid(
    amp=0.1, 
    freq=daily,
    bias=0
)
annual_sine = Sinusoid(
    amp=0.8, 
    freq=daily*annual_to_daily_ratio,
    bias=-0.2
)

"""
Render/Window Config
"""
Y_VIEW_SCALE = 0.4
Y_VIEW_OFFSET = 0.5

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

    @ti.func
    def compute_next_u(self, neighbor_l, neighbor_r):
        grid_vect = ti.Matrix.cols([[ self.u,  neighbor_l.u, neighbor_r.u]]) 
        self.u_next = (coeffs @ grid_vect)[0]

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

    @ti.func 
    def apply_waggle(self, t: float):#amp, freq, bias, t):
        self.u_next = daily_sine.z(t) + annual_sine.z(t)#amp*ti.sin(freq*2*np.pi * t) + bias

    @ti.func
    def apply_initial(self, mode: int):
        mean = 0.5
        sigma = 0.05
        pulse_width = 0.15
        if mode == impulse_types["null"]: # Moving Guassian
            self.u      = 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)))
        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)))
        if mode == impulse_types["impulse"]:
            if self.i == int(N_NODES/2):
                self.u = 0.5
            else:
                self.u = 0
        if mode == impulse_types["pulse_width"]:
            if self.i > int(N_NODES/2)-int(pulse_width*N_NODES) and self.i < int(N_NODES/2)+int(pulse_width*N_NODES):
                self.u = 0.5
            else:
                self.u = 0
        if mode == impulse_types["tri"]:
            self.u = 0.5-ti.abs(0.5-self.x_normalized)
        if mode == impulse_types["fundamental"]:
            self.u = ti.sin(self.x_normalized*np.pi)
        if mode == impulse_types["sine"]:
            self.u = ti.sin(self.x_normalized*2*np.pi)
        if mode == impulse_types["harmonic"]:
            self.u = 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


@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])
        elif i ==0 and WAGGLE:
            nodes[i].apply_waggle(now[None])#waggle_amp, waggle_freq, waggle_bias, 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.015)
        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()




