#!/usr/bin/python3
import mido
import time
import sys
import struct
import socket
import select
from blessed import Terminal
import threading

NOTE_OFFSET = 13
MIN_OFF_TIME = 0.06
TEMPO_OVERRIDE = 1

STRIKE_MAX = 40
STRIKE_MIN = 5

MAIN_VOLUME = int(64 * 0.5)  # normal = 64

# roll file header
roll_data = bytearray()
last_vel = -1

time_length = 0


def send_note(midi_note, midi_vel):  # sends note to serial port
    global roll_data
    global last_vel
    if midi_vel >= 0 and midi_vel <= 128:  # vel 0-127 is normal midi velocity; 128 is note off
        if midi_note >= 21 and midi_note <= 108:  # these midi notes are on the piano keyboard
            raw_note = midi_note + 108  # offset into note byte_command range
            if midi_vel != last_vel:
                last_vel = midi_vel
                data = bytearray([midi_vel, raw_note])
                roll_data += data
            else:
                data = bytearray([raw_note])
                roll_data += data


def render_midi(filename):
    if filename is None:
        return
    global roll_data
    global time_length
    # roll data header
    roll_data = bytearray([128, 128, 128, 128, 254, 219, 0, 220, 127, 221, STRIKE_MIN, 222, STRIKE_MAX, 223, STRIKE_MIN, 224, STRIKE_MAX])
    mid = mido.MidiFile(filename)
    events = []
    tempo = 500000
    ticks_per_beat = mid.ticks_per_beat
    t_sec = 0.0
    for msg in mido.merge_tracks(mid.tracks):
        dt_sec = mido.tick2second(msg.time, ticks_per_beat, tempo)
        t_sec += dt_sec
        if msg.type == "set_tempo":
            tempo = msg.tempo / TEMPO_OVERRIDE
        if hasattr(msg, "channel") and msg.channel == 9:
            continue
        if msg.type == "note_on" and msg.velocity > 0:
            # if msg.note in range(NOTE_OFFSET,88+NOTE_OFFSET):
            events.append((t_sec, "on", msg.note, msg.velocity))
        elif msg.type == "note_off" or (msg.type == "note_on" and msg.velocity == 0):
            # if msg.note in range(NOTE_OFFSET,88+NOTE_OFFSET):
            events.append((t_sec, "off", msg.note, 0))
    notes_active = [False] * 128
    notes_off_since = [0] * 128
    last_time = events[0][0]
    # add extra delay to back to back notes
    events_org = events
    for i in events_org:
        if i[1] == "on":
            if notes_active[i[2]] == True or i[0] - notes_off_since[i[2]] < MIN_OFF_TIME:
                events.append((i[0] - MIN_OFF_TIME, "off", i[2], 0))  # go back in time and turn the note off
            notes_active[i[2]] = True
        elif i[1] == "off":
            notes_active[i[2]] = False
            notes_off_since[i[2]] = i[0]
    events.sort()
    for i in events:
        dt = i[0] - last_time
        last_time = i[0]
        if dt > 0:
            time_length += dt
            delay_bytes = struct.pack(">H", int(dt * 1000))
            data = bytearray([217, delay_bytes[0], 218, delay_bytes[1]])
            roll_data += data
            pass
        if i[1] == "on":
            midi_vel = i[3]
            midi_note = i[2]
            send_note(midi_note, midi_vel)
        elif i[1] == "off":
            midi_note = i[2]
            send_note(midi_note, 128)
    roll_data += bytearray([255, 255])  # halt (twice just in case)
    return roll_data

send_header_on_start = True
exit_on_full_load = False
exit_at_song_end = True

filename = ""
if len(sys.argv) >= 2:
    filename = sys.argv[1]
    if filename.lower() == "null":
        filename = None
        send_header_on_start = False
        exit_at_song_end = False
        print("monitor mode")
    if len(sys.argv) >= 3:
        for i in range(2,len(sys.argv)):
            arg = sys.argv[i]
            if arg == "-n":
                send_header_on_start = False
            if arg == "-x":
                exit_on_full_load = True
                print("exit when entire fire is loaded into pico")
else:
    filename = None
    send_header_on_start = False
    exit_at_song_end = False

print("parsing midi file...")
roll_data = render_midi(filename)
CHUNK_SIZE = 256
if not filename is None:
    roll_chunks = [roll_data[i : i + CHUNK_SIZE] for i in range(0, len(roll_data), CHUNK_SIZE)]
else:
    roll_chunks = []
    roll_data = []

RX_IP = "0.0.0.0"
RX_PORT = 4243
TX_IP = "192.168.1.88"
TX_PORT = 4242

term = Terminal()

sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
sock.bind((RX_IP, RX_PORT))
sock.setblocking(False)

last_packet = ""
last_command = "none"

refill_rq_id = 0
roll_playing = 0
refill_enabled = 0
roll_fill = 0
roll_head = 0
roll_tail = 0
same_status_count = 0
last_stat_txt = ""
vis_line_txt = ""

ov_speed = 0
ov_vol = 0
ov_power = 0

local_ov_speed = None
local_ov_vol = None
local_ov_power = None


def load_chunk():
    packet = bytearray()
    try:
        packet.append(0xFF)  # roll data packet id
        packet += struct.pack("<H", refill_rq_id)
        for i in roll_chunks[refill_rq_id]:
            packet.append(i)
    except:
        packet = bytearray()
        # turn off refill
        packet.append(0x80)
        packet.append(0x04)
    sock.sendto(packet, (TX_IP, TX_PORT))


if send_header_on_start:
    # autostart song playing
    packet = bytearray()
    packet.append(0x80)  # command mode
    packet.append(0xFF)  # halt/reset
    packet.append(226)  # set speed
    packet.append(63)  # speed = normal
    packet.append(0x03)  # begin load
    packet.append(0x01)  # play
    sock.sendto(packet, (TX_IP, TX_PORT))

keepalive_send_time = time.time()

with term.fullscreen(), term.cbreak(), term.hidden_cursor():
    running = True

    print(term.home + term.clear, end="")

    while running:
        now = time.time()
        if now - keepalive_send_time >= 5:
            keepalive_send_time = now
            packet = bytearray()
            packet.append(0x80)  # command mode
            packet.append(0x00)
            sock.sendto(packet, (TX_IP, TX_PORT))


        # Wait for either keyboard or UDP activity
        readable, _, _ = select.select([sock], [], [], 0.05)

        # Receive UDP packets
        for r in readable:
            if r == sock:
                data, addr = sock.recvfrom(4096)

                stat_txt = ""

                # try:
                last_packet = data.hex(" ")
                if data[0] == 0x80:
                    roll_playing = data[1]
                    refill_enabled = data[2]
                    refill_rq_id = struct.unpack("<H", data[3:5])[0]
                    roll_fill = struct.unpack("<H", data[5:7])[0]
                    roll_head = struct.unpack("<H", data[7:9])[0]
                    roll_tail = struct.unpack("<H", data[9:11])[0]
                    ov_speed = data[11]
                    ov_vol = data[12]
                    ov_power = data[13]
                    if local_ov_power is None:
                        local_ov_power = ov_power
                    if local_ov_speed is None:
                        local_ov_speed = ov_speed
                    if local_ov_vol is None:
                        local_ov_vol = ov_vol
                    stat_txt = data.hex(" ")
                    if last_stat_txt == stat_txt:
                        same_status_count += 1
                    else:
                        same_status_count = 0
                    last_stat_txt = stat_txt

                    if exit_on_full_load:
                        if refill_rq_id == len(roll_chunks):
                            quit()


                    if refill_enabled and roll_fill <= (65535 - CHUNK_SIZE * 2):
                        load_chunk()
                elif data[0] == 88:
                    vis_line_txt = ""
                    for i in data:
                        if i == 128:
                            vis_line_txt += "."
                        else:
                            vis_line_txt += chr(96 - int(i / 4))
                    vis_line_txt = vis_line_txt[1:]  # remove first char
                # except:
                # pass

        # Keyboard handling
        key = term.inkey(timeout=0.0)

        # Draw screen
        print(term.home, end="")

        if key:
            if key.lower() == "q":
                running = False

            elif key.lower() == "y":
                last_command = "status_en"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(0x00)
                sock.sendto(packet, (TX_IP, TX_PORT))

            elif key.lower() == "e":
                last_command = "autoload"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(0x03)
                sock.sendto(packet, (TX_IP, TX_PORT))

            elif key.lower() == "h":
                last_command = "Halt / Reset"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(0xFF)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "p":
                last_command = "pause"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(0x02)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "r":
                last_command = "play"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(0x01)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "u":
                local_ov_speed += 1
                if local_ov_speed >= 255:
                    local_ov_speed = 255
                last_command = f"speed +1 ({local_ov_speed},{local_ov_speed/0.63:03.0f}%)"
            elif key.lower() == "j":
                local_ov_speed -= 1
                if local_ov_speed <= 0:
                    local_ov_speed = 0
                last_command = f"speed -1 ({local_ov_speed},{local_ov_speed/0.63:03.0f}%)"
            elif key.lower() == "n":
                local_ov_speed = 63
                last_command = f"speed normalF ({local_ov_speed},{local_ov_speed/0.63:03.0f}%)"
            elif key.lower() == "i":
                local_ov_vol += 1
                if local_ov_vol >= 255:
                    local_ov_vol = 255
                last_command = f"vol +1 ({local_ov_vol},{local_ov_vol/0.64:03.0f}%)"
            elif key.lower() == "k":
                local_ov_vol -= 1
                if local_ov_vol <= 0:
                    local_ov_vol = 0
                last_command = f"vol -1 ({local_ov_vol},{local_ov_vol/0.64:03.0f}%)"
            elif key.lower() == "j":
                local_ov_speed -= 1
                if local_ov_speed <= 0:
                    local_ov_speed = 0
                last_command = f"speed -1 ({local_ov_speed},{local_ov_speed/0.63:03.0f}%)"
            elif key.lower() == "o":
                local_ov_power += 1
                if local_ov_power >= 255:
                    local_ov_power = 255
                last_command = f"power +1 ({local_ov_power},{local_ov_power/2.55:03.0f}%)"
            elif key.lower() == "l":
                local_ov_power -= 1
                if local_ov_power <= 0:
                    local_ov_power = 0
                last_command = f"power -1 ({local_ov_power},{local_ov_power/2.55:03.0f}%)"
            elif key.lower() == " ":
                last_command = "send overrides"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(226)
                packet.append(local_ov_speed)
                packet.append(225)
                packet.append(local_ov_vol)
                packet.append(227)
                packet.append(local_ov_power)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "1":
                last_command = "volume 25%"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(225)
                packet.append(16)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "2":
                last_command = "volume 50%"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(225)
                packet.append(32)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "3":
                last_command = "volume 75%"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(225)
                packet.append(48)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "4":
                last_command = "volume 100%"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(225)
                packet.append(64)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "5":
                last_command = "volume 125%"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(225)
                packet.append(80)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "6":
                last_command = "volume 150%"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(225)
                packet.append(96)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "7":
                last_command = "volume 172%"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(225)
                packet.append(112)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "8":
                last_command = "volume 200%"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(225)
                packet.append(128)
                sock.sendto(packet, (TX_IP, TX_PORT))
            elif key.lower() == "\\":
                last_command = "step event"
                packet = bytearray()
                packet.append(0x80)  # command mode
                packet.append(5)
                sock.sendto(packet, (TX_IP, TX_PORT))

        print("file:", filename)
        print(f"PLAYER PIANO SENDER - {time_length:0.2f} seconds, {len(roll_data)} bytes, {len(roll_chunks)} chunks")
        same_stat_txt = "    "
        if same_status_count >= 1:
            same_stat_txt = f"{same_status_count:03}"
        print(
            f'{"playing" if roll_playing==1 else "paused " if roll_playing==0 else "done   "} {"auto_load" if refill_enabled else "load_stop"} rq:{refill_rq_id:05} {roll_fill:05}:{roll_head:05}:{roll_tail:05} speed={ov_speed/0.63:03.0f}% vol={ov_vol/0.64:03.0f}% power={ov_power/2.55:03.0f}% {same_stat_txt}'
        )
        print()
        print(f"last command: {last_command.ljust(64)}")
        print()
        print(vis_line_txt[0:72])
        print(vis_line_txt[72:])
        print()
        print("  y = send test packet")
        print("  p = pause        r = play")
        print("  h = halt/reset   e = autoload on")
        print("  u = +speed       j = -speed")
        print("  i = +volume      k = -volume")
        print("  o = +power       l = -power")
        print("  space = send overrides")
        print("  q = quit")

        if roll_playing == 2 and same_status_count > 2 and exit_at_song_end:
            quit()
