# -*- coding: utf-8 -*-
"""
Created on Fri Sep 16 11:34:42 2022

@author: rvans
"""
import numpy as np
import matplotlib.pyplot as plt
from mpl_toolkits.mplot3d import Axes3D
from math import pi, sin, cos, tan, sqrt
from scipy.integrate import cumtrapz
from scipy.spatial.transform import Rotation as R

g = 9.81

axtab = []
aytab = []
aztab = []

rot_xtab = []
rot_ytab = []
rot_ztab = []

ttab = []
dttab = []

xrottab = [0]
yrottab = [0]
zrottab = [0]

vxtab = [0]
vytab = [0]
vztab = [0]

pxtab = [0]
pytab = [0]
pztab = [0]

corr_axtab = [0]
corr_aytab = [0]
corr_aztab = [0]

# rotxtab = [0]
# rotytab = [0]
# rotztab = [0]


# axrealtab = [0]
# ayrealtab = [0]
# azrealtab = [0]

data = open('E://Kurkentrekker.txt', 'r')
for line in data.readlines():
    line = str(line).split(',')
    axtab.append((float(line[0])) * g)
    aytab.append((float(line[1])) * g)
    aztab.append((float(line[2]) - 1) * g)
    rot_xtab.append((float(line[3])) * pi/180)
    rot_ytab.append((float(line[4])) * pi/180)
    rot_ztab.append((float(line[5])) * pi/180)
    ttab.append(float(line[6]))
    
for i in range(0, len(ttab) - 1, 1):
    dt = abs(ttab[i+1] - ttab[i])/1000.0
    dttab.append(dt)
#print(f"Length ttab = {len(ttab)}, length dttab = {len(dttab)}, rot_xtab = {len(rot_xtab)}, rot_ytab = {len(rot_ytab)}, rot_ztab = {len(rot_ztab)}")

#filter
signal = np.array(axtab, dtype=float)
fourier = np.fft.rfft(signal)
n = signal.size
sample_rate = 100
freq = np.fft.rfftfreq(n, d=1./sample_rate)

fft_x = np.fft.rfft(axtab) 
fft_y = np.fft.rfft(aytab) 
fft_z = np.fft.rfft(aztab)
fft_rotx = np.fft.rfft(rot_xtab)
fft_roty = np.fft.rfft(rot_ytab)
fft_rotz = np.fft.rfft(rot_ztab)

# plt.figure()
# plt.plot(freq, abs(fft_x), label="raw ax")
# plt.plot(freq, abs(fft_y), label="raw ay")
# plt.plot(freq, abs(fft_z), label="raw az")
# plt.legend()
# plt.show()

# plt.figure()
# plt.plot(freq, abs(fft_rotx), label="raw rot X")
# plt.plot(freq, abs(fft_roty), label="raw rot Y")
# plt.plot(freq, abs(fft_rotz), label="raw rot Z")
# plt.legend()
# plt.show()


atten_x_fft = np.where(freq < 10, fft_x * 0.01, fft_x)
atten_y_fft = np.where(freq < 10, fft_y * 0.01, fft_y)
atten_z_fft = np.where((freq < 20) & (freq > 1), fft_z * 0.01, fft_z)

atten_rotx_fft = np.where(freq < 4, fft_rotx * 0.01, fft_rotx) 
atten_roty_fft = np.where(freq < 4, fft_roty * 0.01, fft_roty) 
atten_rotz_fft = np.where(freq < 4, fft_rotz * 0.01, fft_rotz) 

filter_ax = np.fft.irfft(atten_x_fft)
filter_ay = np.fft.irfft(atten_y_fft)
filter_az = np.fft.irfft(atten_z_fft)

filter_rotx = np.fft.irfft(atten_rotx_fft)
filter_roty = np.fft.irfft(atten_roty_fft)
filter_rotz = np.fft.irfft(atten_rotz_fft)

axtab = filter_ax
aytab = filter_ay
aztab = filter_az

rot_xtab = filter_rotx
rot_ytab = filter_roty
rot_ztab = filter_rotz

def R_x(x):
    return np.matrix([[1, 0, 0],
                      [0, cos(x), -sin(x)], 
                      [0, sin(x), cos(x)]])
def R_y(y):
    return np.array([[cos(y), 0, sin(y)],
                     [0, 1, 0],
                     [-sin(y), 0, cos(y)]])

def R_z(z):
    return np.array([[cos(z), -sin(z), 0],
                     [sin(z), cos(z), 0],
                     [0, 0, 1]])

for i in range(1, len(axtab), 1):
    xrot = xrottab[i - 1] + rot_xtab[i - 1] * dttab[i - 1]
    yrot = yrottab[i - 1] + rot_ytab[i - 1] * dttab[i - 1]
    zrot = zrottab[i - 1] + rot_ztab[i - 1] * dttab[i - 1]

    xrottab.append(xrot)
    yrottab.append(yrot)
    zrottab.append(zrot)
    
    a_vec = np.matrix([[axtab[i - 1]], 
                      [aytab[i - 1]], 
                      [aztab[i - 1]]])
    
    corr_a = np.dot(R_z(zrottab[i - 1]), np.dot(R_y(yrottab[i - 1]), np.dot(R_x(xrottab[i - 1]), a_vec)))
    
    corr_axtab.append(corr_a[0])
    corr_aytab.append(corr_a[1])
    corr_aztab.append(corr_a[2])
    
    vx = vxtab[i - 1] + corr_axtab[i - 1] * dttab[i - 1]
    vy = vytab[i - 1] + corr_aytab[i - 1] * dttab[i - 1]
    vz = vztab[i - 1] + corr_aztab[i - 1] * dttab[i - 1]

    vxtab.append(vx)
    vytab.append(vy)
    vztab.append(vz)
    
    px = pxtab[i - 1] + vxtab[i - 1] * dttab[i - 1] + ((corr_axtab[i - 1])*dttab[i - 1])**2/2
    py = pytab[i - 1] + vytab[i - 1] * dttab[i - 1] + ((corr_aytab[i - 1])*dttab[i - 1])**2/2
    pz = pztab[i - 1] + vztab[i - 1] * dttab[i - 1] + ((corr_aztab[i - 1])*dttab[i - 1])**2/2

    pxtab.append(px)
    pytab.append(py)
    pztab.append(pz)
      


plt.subplot(221)
plt.plot(ttab[0:-1:1], vxtab)
plt.title("Velocity in x-direction")
plt.xlabel("Time $ms$")
plt.ylabel("Velocity $m/s$")

plt.subplot(222)
plt.plot(ttab[0:-1:1], vytab)
plt.title("Velocity in y-direction")
plt.xlabel("Time $ms$")
plt.ylabel("Velocity $m/s$")

plt.subplot(223)
plt.plot(ttab[0:-1:1], pxtab)
plt.title("Position in x-direction")
plt.xlabel("Time $ms$")
plt.ylabel("Position $m$")

plt.subplot(224)
plt.plot(ttab[0:-1:1], pytab)
plt.title("Position in y-direction")
plt.xlabel("Time $ms$")
plt.ylabel("Position $m$")
plt.show()

