import serial, sys, time
import matplotlib.pyplot as plt

import numpy as np

num_samples = 1000

sampling_interval = 0.016 # in seconds
f_sampling = 1/sampling_interval
aliased_60Hz_noise_freq = f_sampling - 60 # assuming we're NOT sampling above 120 Hz

freqs = np.fft.fftfreq(num_samples, d=sampling_interval)
ledOn = True

if (len(sys.argv) != 3):
   print "command line: rx.py serial_port port_speed"
   sys.exit()
port = sys.argv[1]
speed = int(sys.argv[2])

ser = serial.Serial(port,speed)
ser.setDTR()
ser.flushInput()

sensor1_samples = []
sensor2_samples = []
sensorS_samples = []

# set up a MatPlotLib figure with three parts (corresponding to three input streams from SlaveSampler)
fig = plt.figure()
ax1 = fig.add_subplot(411)
ax2 = fig.add_subplot(412)
ax3 = fig.add_subplot(413)
ax4 = fig.add_subplot(414)


ax1.set_autoscaley_on(False)
ax1.set_ylabel("EEG \n raw")
ax1.set_ylim([190,220])

ax2.set_autoscaley_on(False)
ax2.set_ylabel("EEG \n filtered")
ax2.set_ylim([190,220])

ax3.set_autoscaley_on(False)
ax3.set_ylabel("FFT")
ax3.set_ylim([0,100])

ax4.set_autoscaley_on(False)
ax4.set_ylabel("FFT \n filtered")
ax4.set_ylim([0,100])

ax3.set_xlabel("Sample index")

line1, = ax1.plot(range(1,num_samples+1), [0 for i in range(num_samples)], color='b', lw=2)
line2, = ax2.plot(range(1,num_samples+1), [0 for i in range(num_samples)], color='r', lw=2)
line3, = ax3.plot(freqs, [0 for i in range(num_samples)], color='g', lw=2)
line4, = ax4.plot(freqs, [0 for i in range(num_samples)], color='g', lw=2)


while(1):
    # toggle LED
    if ledOn:
        ser.write("L0\r\n")
        ledOn = False
    else:
        ser.write("L1\r\n")
        ledOn = True
        
    sensor1_samples = []
    sensor2_samples = []
    sensorS_samples = []
    
    millis_first = 0
    millis_second = 0
    
    for j in range(num_samples):
        if j >= 1:
            millis_first = millis_second
            millis_second = int(round(time.time() * 1000))
            # print "inter-sample interval = " + str(millis_second - millis_first) + " msec\n"
        else:
            millis_second =  int(round(time.time() * 1000))   
        
        millis1 = int(round(time.time() * 1000))
        ser.write("R1\r\n")
        x1 = ser.read()
        x2 = ser.read()
        y = ord(x1)*256 + ord(x2)
        sensor1_samples.append(y)
        
        # eat up the carriage return and newline characters
        ser.read()
        ser.read()
    
    # FFT 
    f1 = np.fft.fft(sensor1_samples)
    
    '''Filtering out aliased 60 Hz noise'''
    freq_interval = f_sampling / num_samples
    pos = int(aliased_60Hz_noise_freq / freq_interval)
    range_freq = 2
    range_sample_num = freq_interval
    pos_upper = int((aliased_60Hz_noise_freq + range_freq) / freq_interval)
    pos_lower = int((aliased_60Hz_noise_freq - range_freq) / freq_interval)
    
    f1_filtered = f1.copy() # deep copy
    
    f1_filtered[pos_lower:pos_upper] = 0
    f1_filtered[len(f1) - pos_upper:len(f1) - pos_lower] = 0
    
    sensor1_samples_filtered = np.fft.ifft(f1_filtered)
        
    # plot the data
    line1.set_ydata(sensor1_samples)
    line2.set_ydata(sensor1_samples_filtered)
    line3.set_ydata(f1)
    line4.set_ydata(f1_filtered)
    
    plt.draw()
    
    # print sensor1_samples
    # print sensor2_samples
    # print sensorS_samples
    #print "\n"
    
    fig.savefig('temp.png')
    