# -*- coding: utf-8 -*-
"""
A time, frequency and spectrogram real-time visualizer.

Based on the work of: https://github.com/flothesof/pyqtgraph-spectrographer
"""

from collections import deque, defaultdict
import numpy as np
import pyqtgraph as pg
from pyqtgraph.Qt import QtCore, QtWidgets, QtGui

import sys
import time

import jack

import multiprocessing as mp

global paused
paused = mp.Value('b',False)

global LINECOLORS
LINECOLORS = ['r', 'g', 'b', 'c', 'm', 'y', 'w']

def generatePgColormap(cm_name):
  cmap = pg.colormap.get(cm_name)
  lut = cmap.getLookupTable(nPts=256, mode=pg.ColorMap.BYTE)
  return lut

qt_keys = (
  (getattr(QtCore.Qt, attr), attr[4:])
  for attr in dir(QtCore.Qt)
  if attr.startswith("Key_")
)
global keys_mapping
keys_mapping = defaultdict(lambda: "unknown", qt_keys)

class KeyPressWindow(pg.GraphicsLayoutWidget):
  sigKeyPress = QtCore.pyqtSignal(object)
  
  def keyPressEvent(self, ev):
    self.scene().keyPressEvent(ev)
    self.sigKeyPress.emit(ev)

def process_graph(queue,ports_num,event,paused,plot_secs,CHUNKSIZE,SAMPLE_RATE,lock):
  try:
    global LINECOLORS
    print("PyQtGraph: Setting up plot.")
    paused.value = False
    
    # PyQtGraph setup
    app = QtWidgets.QApplication(sys.argv)
    
    TIME_VECTOR = np.arange(CHUNKSIZE) / SAMPLE_RATE
    TIME_FRAMES = int(plot_secs*SAMPLE_RATE)
    FULL_TIME_VECTOR = np.arange(TIME_FRAMES) / SAMPLE_RATE
    
    N_FFT = CHUNKSIZE
    FREQ_VECTOR = np.fft.rfftfreq(N_FFT, d=TIME_VECTOR[1] - TIME_VECTOR[0])
    SPECTOGRAM_FRAMES = int(TIME_FRAMES / N_FFT)
    TIMEOUT = int(TIME_VECTOR.max()*1000)
    EPS = 1e-8
    
    win = KeyPressWindow(title="rtfftvisualizer")
    win.resize(1200, 600)
    win.setWindowTitle('rtfftvisualizer')
    win.show()
    
    global jack_data
    jack_data = np.zeros((CHUNKSIZE,ports_num))
    
    def get_jack_data():
      global jack_data
      
      if event.is_set():
        print("PyQtGraph: JACKClient is shutting down.")
        app.quit()
      
      # Read microphone data from jack
      with lock:
        jack_data = jack_audio.copy()
    
    timer_jack = QtCore.QTimer()
    timer_jack.timeout.connect(get_jack_data)
    timer_jack.start(TIMEOUT)
    
    global waveform_mag, waveform_mag_step, waveform_mag_max
    waveform_mag = 0.5
    waveform_mag_step = 0.05
    waveform_mag_max = 1.0
    
    global waveform_range, waveform_range_step, waveform_range_max
    waveform_range = plot_secs
    waveform_range_step = 0.5
    waveform_range_max = plot_secs
    
    global time_data
    time_data = np.zeros((TIME_FRAMES,ports_num))
    
    time_plot = win.addPlot(title="Time", row=0, col=1, colspan=2)
    time_plot.hideButtons()
    time_plot.showGrid(x=True, y=True)
    time_plot.enableAutoRange('xy', False)
    time_plot.setXRange(0, plot_secs)
    time_plot.setYRange(-waveform_mag, waveform_mag)
    time_plot.setLabel('left', "Amplitude", units='')
    time_plot.setLabel('bottom', "Time", units='s')
    time_curves = []
    for i in range(ports_num):
      time_curves.append(time_plot.plot(pen=LINECOLORS[i]))
    
    waveform_plot = win.addPlot(title="Segment: Time", row=0, col=3)
    waveform_plot.hideButtons()
    waveform_plot.showGrid(x=True, y=True)
    waveform_plot.enableAutoRange('xy', False)
    waveform_plot.setXRange(TIME_VECTOR.min(), TIME_VECTOR.max())
    waveform_plot.setYRange(-waveform_mag, waveform_mag)
    waveform_plot.setLabel('left', "Amplitude", units='')
    waveform_plot.setLabel('bottom', "Time", units='s')
    waveform_curves = []
    for i in range(ports_num):
      waveform_curves.append(waveform_plot.plot(pen=LINECOLORS[i]))
    
    win.nextRow()
    
    global spectrogram_data
    spectrogram_data = deque(maxlen=SPECTOGRAM_FRAMES)
    
    image_data = np.random.rand(20, 20)
    spectrogram_plot = win.addPlot(title='Spectrogram', row=1, col=1, colspan=2)
    spectrogram_plot.hideButtons()
    spectrogram_plot.setLabel('left', "Frequency", units='Hz')
    spectrogram_plot.setLabel('bottom', "Time", units='s')
    spectrogram_plot.setXRange(0, plot_secs)
    
    global spectrogram_image
    spectrogram_image = pg.ImageItem()
    spectrogram_plot.addItem(spectrogram_image)
    spectrogram_image.setImage(image_data)
    lut = generatePgColormap('viridis')
    spectrogram_image.setLookupTable(lut)
    
    # set scale: x in seconds, y in Hz
    spectrogram_image.setTransform(QtGui.QTransform.fromScale(CHUNKSIZE / SAMPLE_RATE, FREQ_VECTOR.max() * 2. / N_FFT))
    
    fft_plot = win.addPlot(title='Segment: Frequency', row=1, col=3)
    fft_plot.hideButtons()
    fft_plot.showGrid(x=True, y=True)
    fft_plot.enableAutoRange('xy', False)
    fft_plot.setXRange(FREQ_VECTOR.min(), FREQ_VECTOR.max())
    fft_plot.setYRange(-100, 50)
    fft_plot.setLabel('left', "Amplitude", units='')
    fft_plot.setLabel('bottom', "Frequency", units='Hz')
    fft_curves = []
    for i in range(ports_num):
      fft_curves.append(fft_plot.plot(pen=LINECOLORS[i]))
    
    global vLine_pos, vLine_pos_past
    vLine_pos = -10.0
    vLine_pos_past = -10.0
    
    time_vLine = pg.InfiniteLine(angle=90, movable=False, pen='w')
    time_plot.addItem(time_vLine)
    time_vLine.setPos(-10.0)
    spec_vLine = pg.InfiniteLine(angle=90, movable=False, pen='w')
    spectrogram_plot.addItem(spec_vLine)
    spec_vLine.setPos(-10.0)
    
    global time_data_write_i
    time_data_write_i = 0
    def update_waveform_time_fft():
      global spectrogram_data, time_data, vLine_pos, vLine_pos_past, jack_data, time_data_write_i
      
      if not paused.value:
        # updating waveform
        for i in range(ports_num):
          waveform_curves[i].setData(TIME_VECTOR,jack_data[:,i])
        
        # updating time
        if time_data_write_i < TIME_FRAMES-CHUNKSIZE:
          time_data[time_data_write_i:time_data_write_i+CHUNKSIZE] = jack_data
          time_data_write_i += CHUNKSIZE
        else:
          time_data = np.roll(time_data, -CHUNKSIZE, axis=0)
          time_data[-CHUNKSIZE:,:] = jack_data
        for i in range(ports_num):
          time_curves[i].setData(FULL_TIME_VECTOR,time_data[:,i])
        
        # updating fft
        for i in range(ports_num):
          X = np.abs(np.fft.rfft(np.hanning(jack_data[:,i].size) * jack_data[:,i], n=N_FFT))
          magn = 20 * np.log10(X + EPS)
          fft_curves[i].setData(x=FREQ_VECTOR, y=magn)
          if i == 0:
            spectrogram_data.append(magn)
        
      else:
        if vLine_pos != -10.0:
          if vLine_pos != vLine_pos_past:
            #getting frame number from vLine_pos
            vLine_pos_i = int(vLine_pos*SAMPLE_RATE)
            #print(f"{vLine_pos} : {vLine_pos_i}")
            
            for i in range(ports_num):
              #getting window from mouse pointer
              this_time_curve_t, this_time_curve_data = time_curves[i].getData()
              this_data = this_time_curve_data[vLine_pos_i:vLine_pos_i+CHUNKSIZE]
              
              # updating waveform
              waveform_curves[i].setData(TIME_VECTOR,this_data)
              
              # updating fft
              X = np.abs(np.fft.rfft(np.hanning(this_data.size) * this_data, n=N_FFT))
              magn = 20 * np.log10(X + EPS)
              fft_curves[i].setData(x=FREQ_VECTOR, y=magn)
            
            vLine_pos_past = vLine_pos
    
    
    proxy = QtWidgets.QGraphicsProxyWidget()
    button = QtWidgets.QPushButton('Help')
    proxy.setWidget(button)
    win.addItem(proxy, row=0, col=0)
    
    label_help = win.addLabel("<span style='font-size: 10pt'>(spacebar) pause</span><br><span style='font-size: 10pt'>(←, →) zoom time</span><br><span style='font-size: 10pt'>(↑, ↓) zoom amplitude</span><br><span style='font-size: 10pt'>when paused, cursor selects segment</span>", row=0, col=0)
    label_help.hide()
    
    def toggle_label():
      if label_help.isVisible():
        label_help.hide()
      else:
        label_help.show()
      button.clearFocus()
    
    button.clicked.connect(toggle_label)
    
    timer_fft = QtCore.QTimer()
    timer_fft.timeout.connect(update_waveform_time_fft)
    timer_fft.start(TIMEOUT)
    
    
    def update_spectrogram():
      global spectrogram_data, spectrogram_image, time_data
      arr = np.c_[spectrogram_data]
      if arr.size > 0:
        if arr.ndim == 1:
            arr = arr[:, np.newaxis]
        max = arr.max()
        min = max / 10
        spectrogram_image.setImage(arr, levels=(min, max))
    
    
    timer_spectrogram = QtCore.QTimer()
    timer_spectrogram.timeout.connect(update_spectrogram)
    timer_spectrogram.start(TIMEOUT)
    
    def toggle_pause():
      global waveform_range
      
      paused.value = not paused.value
      if not paused.value:
          spec_vLine.setPos(-10.0)
          time_vLine.setPos(-10.0)
          time_plot.setXRange(0.0, plot_secs)
          waveform_range = plot_secs
    
    def increase_mag():
      global waveform_mag, waveform_mag_step, waveform_mag_max
      if waveform_mag < waveform_mag_max:
        waveform_mag += waveform_mag_step
        waveform_plot.setYRange(-waveform_mag, waveform_mag)
        time_plot.setYRange(-waveform_mag, waveform_mag)
    
    def decrease_mag():
      global waveform_mag, waveform_mag_step, waveform_mag_max
      if waveform_mag > waveform_mag_step*2:
        waveform_mag -= waveform_mag_step
        waveform_plot.setYRange(-waveform_mag, waveform_mag)
        time_plot.setYRange(-waveform_mag, waveform_mag)
    
    def apply_range(middle_value):
      global waveform_range, waveform_range_max
      range_min = middle_value-(waveform_range/2)
      if range_min < 0.0:
        range_min = 0.0
        
      range_max = middle_value+(waveform_range/2)
      if range_max > waveform_range_max:
        range_max = waveform_range_max
      
      time_plot.setXRange(range_min, range_max)
    
    def decrease_range():
      global waveform_range, waveform_range_step, waveform_range_max, vLine_pos
      if waveform_range > waveform_mag_step*2:
        waveform_range -= waveform_mag_step
        if paused.value and vLine_pos != -10.0:
          apply_range(vLine_pos)
        else:
          apply_range(waveform_range_max/2)
    
    def increase_range():
      global waveform_range, waveform_range_step, waveform_range_max, vLine_pos
      if waveform_range < waveform_range_max:
        waveform_range += waveform_mag_step
        if paused.value and vLine_pos != -10.0:
          apply_range(vLine_pos)
        else:
          apply_range(waveform_range_max/2)
    
    
    def keyPressEvent(kbevent):
      global keys_mapping
      
      key = keys_mapping[kbevent.key()]
      #print(key)
      
      if key == 'Space':
        toggle_pause()
      elif key == 'Up':
        increase_mag()
      elif key == 'Down':
        decrease_mag()
      elif key == 'Left':
        decrease_range()
      elif key == 'Right':
        increase_range()
    
    win.sigKeyPress.connect(keyPressEvent)
    
    def mouseMoved(evt):
      global vLine_pos
      if spectrogram_plot.sceneBoundingRect().contains(evt[0]):
        mousePoint = spectrogram_plot.vb.mapSceneToView(evt[0])
        if paused.value:
          vLine_pos = mousePoint.x()
          if vLine_pos >= 0.0:
            time_vLine.setPos(vLine_pos)
            spec_vLine.setPos(vLine_pos)
          else:
            vLine_pos = -10.0
      elif time_plot.sceneBoundingRect().contains(evt[0]):
        mousePoint = time_plot.vb.mapSceneToView(evt[0])
        if paused.value:
          vLine_pos = mousePoint.x()
          if vLine_pos >= 0.0:
            time_vLine.setPos(vLine_pos)
            spec_vLine.setPos(vLine_pos)
          else:
            vLine_pos = -10.0
      else:
        vLine_pos = -10.0
    
    mouse_proxy = pg.SignalProxy(spectrogram_plot.scene().sigMouseMoved, rateLimit=60, slot=mouseMoved) #using spectrogram_plot as proxy, but it should get all of the window
    
    print("PyQtGraph: Starting to graph.")
    app.exec_()
  except KeyboardInterrupt:
    print("PyQtGraph: Cleaning up before exit.")
    app.quit()


ports_num = 2
plot_secs = 3.0

if len(sys.argv) > 1:
  ports_num = int(sys.argv[1])

if ports_num > len(LINECOLORS):
  print("input port number is too large, defaulting to maximum "+str(len(LINECOLORS)))
  ports_num = len(LINECOLORS)

if len(sys.argv) > 2:
  plot_secs = float(sys.argv[2])

jackclient = jack.Client('rtfftvisualizer', use_exact_name=True)

# Multiprocessing variables
event = mp.Event()
lock = mp.Lock()
queue = mp.Manager().Queue()

print("input port number      : "+str(ports_num))
print("sample rate            : "+str(jackclient.samplerate))
print("jack win. length       : "+str(jackclient.blocksize))
print("seconds to plot        : "+str(plot_secs))

jack_audio_shared_memory = mp.shared_memory.SharedMemory(create=True, size=jackclient.blocksize*ports_num*64)
jack_audio = np.ndarray((jackclient.blocksize,ports_num), np.float32, buffer=jack_audio_shared_memory.buf)

@jackclient.set_process_callback
def process(frames):
  global jack_audio
  global paused
  
  assert frames == jackclient.blocksize
  
  if not paused.value:
    #start_time = time.time()
    with lock:
      for i in range(ports_num):
        jack_audio[:,i] = inputs[i].get_array()
    #exec_time = time.time() - start_time
    #print('copy input to audio : %f secs' % (exec_time))

@jackclient.set_xrun_callback
def xrun_callback(delay):
  print('JACKClient: XRUN occurred : -> '+str(delay))

@jackclient.set_shutdown_callback
def shutdown(status, reason):
  print('JACKClient: JACK shutdown!')
  print('JACKClient: status -> '+str(status))
  print('JACKClient: reason -> '+str(reason))
  event.set()

inputs = []
for i in range(ports_num):
  inputs.append(jackclient.inports.register('input_'+str(i+1)))

with jackclient:
  print('JACKClient: JACK started!')
  
  capture = jackclient.get_ports(is_physical=True, is_output=True)
  if not capture:
    raise RuntimeError('No physical capture ports')
  
  for src, dest in zip(capture, jackclient.inports):
    jackclient.connect(src, dest)
  
  graph_thread = mp.Process(target=process_graph, args=(queue,ports_num,event,paused,plot_secs,jackclient.blocksize,jackclient.samplerate,lock), daemon=True)
  graph_thread.start()
  
  try:
    graph_thread.join()
  except KeyboardInterrupt:
    print('JACKClient: stopped by user (CTRL+C).')
    graph_thread.terminate() # Forcefully stop if join was interrupted
    graph_thread.join()      # Ensure the process is fully reaped
    print("Graph process terminated.")
  finally:
    print("Terminating JACKClient.")
    jackclient.deactivate()
    jack_audio_shared_memory.close()
    jack_audio_shared_memory.unlink()
    print("JACKClient terminated.")
    print("Main process exited.")

