# SPDX-License-Identifier: MIT
# SPDX-FileCopyrightText: Copyright 2025 Sam Blenny
#
# Related Docs:
# - https://docs.circuitpython.org/projects/tlv320/en/latest/api.html
# - https://learn.adafruit.com/adafruit-tlv320dac3100-i2s-dac/overview
# - https://docs.circuitpython.org/en/latest/shared-bindings/audiobusio/
# - https://docs.circuitpython.org/en/latest/shared-bindings/audiomixer/
# - https://midi.org/specs
# - https://github.com/todbot/circuitpython-synthio-tricks
#
from audiobusio import I2SOut
from board import (
    I2C, I2S_BCLK, I2S_DIN, I2S_MCLK, I2S_WS, PERIPH_RESET
)
from digitalio import DigitalInOut, Direction, Pull
import gc
from micropython import const
import synthio
import sys
from time import sleep
from usb.core import USBError, USBTimeoutError
import usb_host

from adafruit_tlv320 import TLV320DAC3100

from sb_usb_midi import find_usb_device, MIDIInputDevice


# DAC and Synthesis parameters
SAMPLE_RATE = const(11025)
CHAN_COUNT  = const(2)
BUFFER_SIZE = const(1024)
#==============================================================
# CAUTION! When this is set to True, the headphone jack will
# send a line-level output suitable use with a mixer or powered
# speakers, but that will be _way_ too loud for earbuds. For
# finer control of line level volume, adjust LL_DAC_VOLUME.
LINE_LEVEL = const(True)
LL_HEADPHONE_VOLUME = -6.0
#==============================================================

# Change this to True if you want more MIDI output on the serial console
DEBUG = True

# Adjust this to balance the probability of dropped/stuck notes against the
# amount of time you're willing to let event loop spend on blocking IO to wait
# for a USB host read() call to finish. Note that background tasks like synthio
# should work fine while read() is blocked. In my testing, short timeouts like
# 3 or 33 ms would result in some dropped or stuck notes. But, I didn't get
# any of that with longer timeouts like 50, 100, or 300 ms.
#
READ_TIMEOUT_MS = const(50)

# Constants for parsing MIDI messages
CIN_NOTE_OFF = const(0x08)  # Note Off
CIN_NOTE_ON  = const(0x09)  # Note On
CIN_CC       = const(0x0b)  # Control Change
CIN_MPE      = const(0x0a)  # Polyphonic Expression (individual key pressure)
CIN_CP       = const(0x0d)  # Channel Pressure (grouped key pressure)
CIN_PB       = const(0x0e)  # Pitch Bend


def init_dac_audio_synth(i2c):
    # Configure Fruit Jam rev D TLV320 I2S DAC and make a Synthesizer.
    # - i2c: a reference to board.I2C()
    # - returns tuple: (dac: TLV320DAC3100, audio: I2SOut, synth: Synthesizer)

    # 1. Reset DAC (reset is active low)
    rst = DigitalInOut(PERIPH_RESET)
    rst.direction = Direction.OUTPUT
    rst.value = False
    sleep(0.1)
    rst.value = True
    sleep(0.05)
    # 2. Configure sample rate, bit depth, and output port
    dac = TLV320DAC3100(i2c)
    dac.configure_clocks(sample_rate=SAMPLE_RATE, bit_depth=16)
    dac.speaker_output = False
    dac.headphone_output = True
    # 3. Adjust volume for for line-level if needed, otherwise use default
    #    volume set by `dac.headphone_output = True`
    if LINE_LEVEL:
        # This gives a line output level suitable for plugging into a mixer or
        # the AUX input of a powered speaker (THIS IS TOO LOUD FOR HEADPHONES!)
        dac.headphone_volume = LL_HEADPHONE_VOLUME
    if DEBUG:
        print(f"dac.dac_volume = {dac.dac_volume:.1f}")
        print(f"dac.headphone_volume = {dac.headphone_volume:.1f}")
    # 4. Configure I2S for Fruit Jam rev D (rev B swapped WS and MCLK)
    audio = I2SOut(bit_clock=I2S_BCLK, word_select=I2S_WS, data=I2S_DIN)
    # 5. Configure synthio patch to generate audio
    vca = synthio.Envelope(
        attack_time=0.002, decay_time=0.01, sustain_level=0.4,
        release_time=0, attack_level=0.6
    )
    synth = synthio.Synthesizer(
        sample_rate=SAMPLE_RATE, channel_count=CHAN_COUNT, envelope=vca
    )
    audio.play(synth)
    return (dac, audio, synth)


def main():
    gc.collect()

    # Set up the audio stuff for a basic synthesizer
    i2c = I2C()
    (dac, audio, synth) = init_dac_audio_synth(i2c)

    # Cache function references (MicroPython performance boost trick)
    wr = sys.stdout.write
    panic = synth.release_all
    press = synth.press
    release = synth.release

    # Dictionary to keep track of which notes are active.
    # This is useful for debug printing to watch for stuck notes
    notes = {}

    # Note On helper function with closure for notes dictionary
    def note_off(num):
        release(num)
        if num in notes:
            notes.pop(num)

    # Note Off helper function with closure for notes dictionary
    def note_on(num):
        press(num)
        notes[num] = True

    # Main loop: scan USB host bus for MIDI device, connect, start event loop.
    # This grabs the first MIDI device it finds.
    while True:
        wr("USB Host: scanning bus...\n")
        gc.collect()
        device_cache = {}
        try:
            # This loop will end as soon as it finds a ScanResult object (r)
            r = None
            while r is None:
                sleep(0.4)
                r = find_usb_device(device_cache)
            # Use ScanResult object (r) to check if USB device descriptor info
            # matches the class/subclass/protocol pattern for a MIDI device. If
            # the device doesn't match, MIDIInputDevice will raise an exception
            # and trigger another iteration through the outer while True loop.
            dev = MIDIInputDevice(r, read_timeout=READ_TIMEOUT_MS)
            wr(" found MIDI device vid:pid %04X:%04X\n" % (r.vid, r.pid))
            # Collect garbage to hopefully limit heap fragmentation.
            r = None
            device_cache = {}
            gc.collect()
            # MIDI Event Input Loop: Poll for input until USB error.
            cin = chan = num = val = 0x00
            for data in dev.input_event_generator():
                # Begin handling midi packet which should be None or a 4-byte
                # memoryview.
                if data is None:
                    continue

                # data[0] has CN (Cable Number) and CIN (Code Index Number). By
                # discarding CN with `& 0x0f`, we ignore the virtual MIDI port
                # that the messages arrive from. Ignoring CN would be bad for a
                # fancy DAW or synth setup where you needed to route MIDI
                # among multiple devices. But, that doesn't matter here. We do
                # need CIN to distinguish between note on, note off, Control
                # Change (CC), and so on. For the channel, adding 1 gives us
                # human-friendly channel numbers in the range of 1 to 16.
                #
                # NOTE: As far as I can tell from reading the USB MIDI specs,
                # each MIDI packet will always be 32-bits long and CIN will
                # always be set. That means there's no need to worry about
                # parsing "running status" as would be the case when using UART
                # MIDI with DIN-5 or TRS cables.
                #
                cin = data[0] & 0x0f
                chan = (data[1] & 0xf) + 1
                num = data[2]
                val = data[3]

                # This decodes MIDI events by comparing constants against bytes
                # from a memoryview. Using a class to do this parsing would add
                # many extra heap allocations and dictionary lookups. That
                # stuff is slow, and we want to go _fast_. For details about
                # the MIDI 1.0 standard, see https://midi.org/specs
                #
                if cin == CIN_NOTE_OFF and (21 <= num <= 108):
                    # Note off
                    note_off(num)
                elif cin == CIN_NOTE_ON and (21 <= num <= 108):
                    if val == 0:
                        # Some devices send zero velocity instead of note off
                        note_off(num)
                    else:
                        # Note on
                        note_on(num)
                elif cin == CIN_CC and num == 123 and val == 0:
                    # CC 123 means stop all notes ("panic")
                    panic()
                    wr('PANIC %d %d %d\n' % (chan, num, val))
                if DEBUG:
                    if cin == CIN_NOTE_OFF:
                        # Note On
                        # Debug print this message + list of active notes
                        n = ' '.join(sorted([str(n) for n in notes.keys()]))
                        wr('Off %d %d %3d  notes: %s\n' % (chan, num, val, n))
                    elif cin == CIN_NOTE_ON:
                        # Note Off
                        # Debug print this message + list of active notes
                        n = ' '.join(sorted([str(n) for n in notes.keys()]))
                        wr('On  %d %d %3d  notes: %s\n' % (chan, num, val, n))
                    elif cin == CIN_CC:
                        # CC control change
                        wr('CC  %d %d %d\n' % (chan, num, val))
                    elif cin == CIN_MPE:
                        # MPE polyphonic key pressure (aftertouch)
                        wr('MPE %d %d %d\n' % (chan, num, val))
                    elif cin == CIN_CP:
                        # CP channel key pressure (aftertouch)
                        wr('CP  %d %d %d\n' % (chan, num, val))
                    elif cin == CIN_PB:
                        # PB pitch bend
                        wr('PB  %d %d %d\n' % (chan, num, val))
                    # Ignore the rest: SysEx, System Realtime, or whatever
        except USBError as e:
            # This sometimes happens when devices are unplugged. Not always.
            print("USBError: '%s' (device unplugged?)" % e)
        except ValueError as e:
            # This can happen if an initialization handshake glitches
            print(e)


main()
