# SPDX-FileCopyrightText: 2026 Tim Cocks For Adafruit Industries
#
# SPDX-License-Identifier: MIT
"""
USB Audio Effects Pedal for the Adafruit MacroPad.

Audio received from the host (USBSpeaker) is passed through up to four
user-configurable effect "slots" arranged as a serial chain, recombined at an
audiomixer output stage, and sent back to the host (USBMicrophone).

    host -> spk -> slot0 -> slot1 -> slot2 -> slot3 -> mixer -> mic -> host

NOTE on topology: a single audio source can only feed one consumer, so the four
slots are chained in series (like a real pedalboard) rather than mixed in
parallel.  The audiomixer is the final recombination / master-level stage.

Controls
--------
Encoder turn  : move selection / change the value being edited
Encoder push  : open / confirm / toggle edit
Keys (home)   : each row is a slot.
                  col 1 -> open slot,  col 2 -> bypass,  col 3 -> clear
Keys (editing): number pad for typing values
                  1 2 3 / 4 5 6 / 7 8 9 / . 0 Enter   (Enter = K12)
                K12 also = "back" when not typing a value.
"""

import board
import displayio
import terminalio
import digitalio
import rotaryio
import keypad
import neopixel

import usb_audio
import audiomixer
import audiofilters
import audiodelays
import audiofreeverb
import synthio

from adafruit_display_text import label
from displayio_listselect import ListSelect
from adafruit_displayio_layout.layouts.page_layout import PageLayout

# --------------------------------------------------------------------------- #
# Audio configuration
# --------------------------------------------------------------------------- #
SAMPLE_RATE = 16000
CHANNELS = 1
BITS = 16
BUFFER = 512

AUDIO_KW = {
    "sample_rate": SAMPLE_RATE,
    "channel_count": CHANNELS,
    "bits_per_sample": BITS,
    "samples_signed": True,
    "buffer_size": BUFFER,
}

NUM_SLOTS = 4

# Enum value shortcuts
_DM = audiofilters.DistortionMode
_FM = synthio.FilterMode


# --------------------------------------------------------------------------- #
# Parameter / effect descriptions
# --------------------------------------------------------------------------- #
class Param:
    """Describes one editable parameter of an effect."""

    def __init__(self, key, label_, kind, default, lo=0.0, hi=1.0, step=0.05,
                 options=None, target="effect", attr=None):
        self.key = key
        self.label = label_
        self.kind = kind          # "float" | "int" | "enum" | "bool"
        self.default = default
        self.lo = lo
        self.hi = hi
        self.step = step
        self.options = options    # list of (label, value) for enums
        self.target = target      # "effect" or "biquad"
        self.attr = attr if attr is not None else key

    @property
    def numeric(self):
        return self.kind in ("float", "int")


class EffectSpec:
    """Describes one effect type: how to build it and its parameters."""

    def __init__(self, label_, factory, params):
        self.label = label_
        self.factory = factory    # () -> (effect, aux)
        self.params = params


def _build_filter():
    bq = synthio.Biquad(mode=_FM.LOW_PASS, frequency=2000, Q=0.7)
    eff = audiofilters.Filter(filter=bq, **AUDIO_KW)
    return eff, bq


# Effect registry.  Order here is the order shown in the "add effect" menu.
EFFECTS = {
    "Distortion": EffectSpec(
        "Distortion",
        lambda: (audiofilters.Distortion(**AUDIO_KW), None),
        [
            Param("drive", "Drive", "float", 0.5, 0.0, 1.0, 0.05),
            Param("pre_gain", "PreGain", "float", 0.0, -60.0, 60.0, 1.0),
            Param("post_gain", "PostGain", "float", 0.0, -80.0, 24.0, 1.0),
            Param("mode", "Mode", "enum", _DM.CLIP,
                  options=[("Clip", _DM.CLIP), ("LoFi", _DM.LOFI),
                           ("Drive", _DM.OVERDRIVE), ("Wave", _DM.WAVESHAPE)]),
            Param("soft_clip", "SoftClip", "bool", False),
            Param("mix", "Mix", "float", 1.0, 0.0, 1.0, 0.05),
        ],
    ),
    "Filter": EffectSpec(
        "Filter",
        _build_filter,
        [
            Param("type", "Type", "enum", _FM.LOW_PASS, target="biquad", attr="mode",
                  options=[("LowPass", _FM.LOW_PASS), ("HighPass", _FM.HIGH_PASS),
                           ("BandPass", _FM.BAND_PASS)]),
            Param("freq", "Freq", "float", 2000.0, 20.0, 7000.0, 50.0,
                  target="biquad", attr="frequency"),
            Param("Q", "Q", "float", 0.7, 0.1, 10.0, 0.1,
                  target="biquad", attr="Q"),
            Param("mix", "Mix", "float", 1.0, 0.0, 1.0, 0.05),
        ],
    ),
    "Phaser": EffectSpec(
        "Phaser",
        lambda: (audiofilters.Phaser(**AUDIO_KW), None),
        [
            Param("frequency", "Freq", "float", 1000.0, 50.0, 4000.0, 50.0),
            Param("feedback", "Feedbk", "float", 0.7, 0.0, 1.0, 0.05),
            Param("stages", "Stages", "int", 6, 1, 8, 1),
            Param("mix", "Mix", "float", 1.0, 0.0, 1.0, 0.05),
        ],
    ),
    "Chorus": EffectSpec(
        "Chorus",
        lambda: (audiodelays.Chorus(max_delay_ms=100, **AUDIO_KW), None),
        [
            Param("delay_ms", "Delay", "float", 50.0, 1.0, 99.0, 1.0),
            Param("voices", "Voices", "int", 2, 1, 5, 1),
            Param("mix", "Mix", "float", 0.5, 0.0, 1.0, 0.05),
        ],
    ),
    "Echo": EffectSpec(
        "Echo",
        lambda: (audiodelays.Echo(max_delay_ms=500, **AUDIO_KW), None),
        [
            Param("delay_ms", "Delay", "float", 250.0, 1.0, 500.0, 10.0),
            Param("decay", "Decay", "float", 0.5, 0.0, 1.0, 0.05),
            Param("mix", "Mix", "float", 0.5, 0.0, 1.0, 0.05),
            Param("freq_shift", "FrqShift", "bool", False),
        ],
    ),
    "MultiTap": EffectSpec(
        "MultiTap",
        lambda: (audiodelays.MultiTapDelay(max_delay_ms=500, **AUDIO_KW), None),
        [
            Param("delay_ms", "Delay", "float", 250.0, 1.0, 500.0, 10.0),
            Param("decay", "Decay", "float", 0.5, 0.0, 1.0, 0.05),
            Param("mix", "Mix", "float", 0.5, 0.0, 1.0, 0.05),
        ],
    ),
    "Pitch": EffectSpec(
        "Pitch",
        lambda: (audiodelays.PitchShift(**AUDIO_KW), None),
        [
            Param("semitones", "Semis", "float", 0.0, -12.0, 12.0, 1.0),
            Param("mix", "Mix", "float", 1.0, 0.0, 1.0, 0.05),
        ],
    ),
    "Reverb": EffectSpec(
        "Reverb",
        lambda: (audiofreeverb.Freeverb(**AUDIO_KW), None),
        [
            Param("roomsize", "Room", "float", 0.5, 0.0, 1.0, 0.05),
            Param("damp", "Damp", "float", 0.5, 0.0, 1.0, 0.05),
            Param("mix", "Mix", "float", 0.5, 0.0, 1.0, 0.05),
        ],
    ),
}
EFFECT_NAMES = list(EFFECTS.keys())


class Slot:
    """One position in the effect chain."""

    def __init__(self):
        self.effect_name = None
        self.effect = None
        self.biquad = None
        self.values = {}
        self.bypassed = False


# --------------------------------------------------------------------------- #
# Number-pad key map (key_number -> character).  K12 (index 11) = Enter.
# --------------------------------------------------------------------------- #
NUMPAD = {0: "1", 1: "2", 2: "3",
          3: "4", 4: "5", 5: "6",
          6: "7", 7: "8", 8: "9",
          9: ".", 10: "0"}
ENTER_KEY = 11

# Neopixel colors
C_EMPTY = (0, 0, 12)
C_ACTIVE = (0, 40, 0)
C_BYPASS = (40, 0, 0)
C_BYPASS_KEY = (30, 15, 0)
C_CLEAR_KEY = (15, 0, 0)
C_DIGIT = (18, 18, 18)
C_ENTER = (0, 40, 0)
C_BACK = (40, 0, 0)


# --------------------------------------------------------------------------- #
# The application
# --------------------------------------------------------------------------- #
class EffectsPedal:
    def __init__(self):
        # --- audio graph ---
        # self.mic = usb_audio.USBMicrophone()
        # self.spk = usb_audio.USBSpeaker()
        self.mic = usb_audio.usb_microphone
        self.spk = usb_audio.usb_speaker
        self.mixer = audiomixer.Mixer(
            voice_count=1, sample_rate=SAMPLE_RATE, channel_count=CHANNELS,
            bits_per_sample=BITS, samples_signed=True)
        self.master = 0.8
        self.mic.play(self.mixer)

        self.slots = [Slot() for _ in range(NUM_SLOTS)]

        # --- hardware ---
        key_pins = (board.KEY1, board.KEY2, board.KEY3, board.KEY4,
                    board.KEY5, board.KEY6, board.KEY7, board.KEY8,
                    board.KEY9, board.KEY10, board.KEY11, board.KEY12)
        self.keys = keypad.Keys(key_pins, value_when_pressed=False, pull=True)
        self.encoder = rotaryio.IncrementalEncoder(board.ROTA, board.ROTB)
        self.button = digitalio.DigitalInOut(board.BUTTON)
        self.button.switch_to_input(pull=digitalio.Pull.UP)
        self.pixels = neopixel.NeoPixel(board.NEOPIXEL, 12, brightness=0.3,
                                        auto_write=True)
        self._last_pos = self.encoder.position
        self._last_btn = True  # not pressed (pull-up)

        # --- UI state ---
        self.mode = "home"            # "home" | "edit"
        self.edit_state = "param_nav"  # "effect_choose"|"param_nav"|"param_edit"
        self.focus = 0                # slot being edited
        self.param_idx = 0            # param being edited within a slot
        self.entry_str = None         # numeric typing buffer
        self.master_edit = False

        self._build_ui()
        self.rebuild_chain()
        self.refresh_home()
        self.update_pixels()

    # ------------------------------------------------------------------ #
    # UI construction
    # ------------------------------------------------------------------ #
    def _build_ui(self):
        self.display = board.DISPLAY
        self.main_group = displayio.Group()
        self.pages = PageLayout(x=0, y=0)

        # Home page ----------------------------------------------------
        home = displayio.Group()
        self.home_title = label.Label(
            terminalio.FONT, text="USB FX Pedal", color=0xFFFFFF, x=2, y=5)
        self.home_title.hidden = True
        self.home_list = ListSelect(
            items=["----"], x=4, y=0, visible_items_count=4)
        home.append(self.home_title)
        home.append(self.home_list)

        # Edit page ----------------------------------------------------
        edit = displayio.Group()
        self.edit_title = label.Label(
            terminalio.FONT, text="", color=0xFFFFFF, x=2, y=5)
        self.edit_list = ListSelect(
            items=["----"], x=4, y=16, visible_items_count=3)
        self.edit_list._label.line_spacing = 1.0

        self.edit_hint = label.Label(
            terminalio.FONT, text="", color=0xFFFFFF, x=2, y=58)
        edit.append(self.edit_title)
        edit.append(self.edit_list)
        edit.append(self.edit_hint)

        self.pages.add_content(home, "home")
        self.pages.add_content(edit, "edit")
        self.main_group.append(self.pages)
        try:
            self.display.root_group = self.main_group
        except AttributeError:
            self.display.show(self.main_group)

    # ------------------------------------------------------------------ #
    # Value formatting / helpers
    # ------------------------------------------------------------------ #
    @staticmethod
    def fmt(p, v):
        if p.kind == "bool":
            return "On" if v else "Off"
        if p.kind == "enum":
            for lbl, val in p.options:
                if val == v:
                    return lbl
            return "?"
        if p.kind == "int":
            return str(int(v))
        return "%.2f" % v

    @staticmethod
    def clamp(p, v):
        if v < p.lo:
            return p.lo
        if v > p.hi:
            return p.hi
        return v

    def _cur_param(self):
        spec = EFFECTS[self.slots[self.focus].effect_name]
        return spec.params[self.param_idx]

    def _cur_param_numeric(self):
        return (self.mode == "edit" and self.edit_state == "param_edit"
                and self._cur_param().numeric)

    # ------------------------------------------------------------------ #
    # Audio graph management
    # ------------------------------------------------------------------ #
    def apply_all(self, slot):
        """Write every stored parameter value onto the live objects."""
        for p in EFFECTS[slot.effect_name].params:
            obj = slot.biquad if p.target == "biquad" else slot.effect
            if obj is not None:
                try:
                    setattr(obj, p.attr, slot.values[p.key])
                except (ValueError, AttributeError):
                    pass

    def set_param_value(self, slot, p, v):
        slot.values[p.key] = v
        # synthio.Biquad.mode is read-only, so a filter-type change means
        # rebuilding the biquad (keeping the current frequency / Q) and
        # reassigning it to the Filter effect.
        if slot.effect_name == "Filter" and p.target == "biquad" and p.attr == "mode":
            if slot.effect is not None:
                bq = synthio.Biquad(
                    mode=v,
                    frequency=slot.values["freq"],
                    Q=slot.values["Q"])
                slot.biquad = bq
                try:
                    slot.effect.filter = bq
                except (ValueError, AttributeError):
                    pass
            return
        obj = slot.biquad if p.target == "biquad" else slot.effect
        if obj is not None:
            try:
                setattr(obj, p.attr, v)
            except (ValueError, AttributeError):
                pass

    def rebuild_chain(self):
        """Reconnect the active, non-bypassed slots in series spk -> mixer."""
        try:
            if self.mixer.voice[0].playing:
                self.mixer.voice[0].stop()
        except (RuntimeError, ValueError):
            pass

        source = self.spk
        for slot in self.slots:
            if slot.effect_name is None or slot.bypassed:
                continue
            try:
                if slot.effect.playing:
                    slot.effect.stop()
            except (RuntimeError, ValueError):
                pass
            slot.effect.play(source)
            source = slot.effect

        self.mixer.voice[0].play(source)
        self.mixer.voice[0].level = self.master

    def assign_effect(self, name):
        slot = self.slots[self.focus]
        self.clear_slot(self.focus, rebuild=False)
        spec = EFFECTS[name]
        slot.effect_name = name
        slot.values = {p.key: p.default for p in spec.params}
        slot.bypassed = False
        if spec.factory is not None:
            slot.effect, slot.biquad = spec.factory()
            self.apply_all(slot)
        self.rebuild_chain()

    def clear_slot(self, idx, rebuild=True):
        slot = self.slots[idx]
        if slot.effect is not None:
            try:
                slot.effect.deinit()
            except (RuntimeError, ValueError):
                pass
        slot.effect = None
        slot.biquad = None
        slot.effect_name = None
        slot.values = {}
        slot.bypassed = False
        if rebuild:
            self.rebuild_chain()

    def toggle_bypass(self, idx):
        slot = self.slots[idx]
        if slot.effect_name is None:
            return
        slot.bypassed = not slot.bypassed
        self.rebuild_chain()

    # ------------------------------------------------------------------ #
    # Display refresh
    # ------------------------------------------------------------------ #
    def refresh_home(self):
        items = []
        for i, slot in enumerate(self.slots):
            name = slot.effect_name if slot.effect_name else "--"
            tag = " off" if (slot.bypassed and slot.effect_name) else ""
            items.append("S%d %s%s" % (i + 1, name, tag))
        items.append("Master %.2f" % self.master)
        sel = self.home_list.selected_index
        if sel >= len(items):
            sel = len(items) - 1
        self.home_list.items = items
        self.home_list.cursor_char = "*" if self.master_edit else ">"
        self.home_list.selected_index = sel

    def refresh_edit(self):
        slot = self.slots[self.focus]
        if self.edit_state == "effect_choose":
            self.edit_title.text = "S%d add FX" % (self.focus + 1)
            self.edit_hint.text = "turn pick  push add"
            sel = self.edit_list.selected_index
            self.edit_list.items = ["< Back"] + [
                EFFECTS[n].label for n in EFFECT_NAMES]
            self.edit_list.cursor_char = ">"
            self.edit_list.selected_index = min(sel, len(self.edit_list.items) - 1)
            return

        spec = EFFECTS[slot.effect_name]
        self.edit_title.text = "S%d %s" % (self.focus + 1, spec.label)
        items = []
        for i, p in enumerate(spec.params):
            if (self.edit_state == "param_edit" and i == self.param_idx
                    and self.entry_str is not None):
                val = self.entry_str + "_"
            else:
                val = self.fmt(p, slot.values[p.key])
            items.append("%s:%s" % (p.label, val))
        items.append("[Change FX]")
        items.append("[Remove]")

        sel = self.edit_list.selected_index
        if self.edit_state == "param_edit":
            sel = self.param_idx
            self.edit_list.cursor_char = "="
            self.edit_hint.text = "turn/keys  K12 ok"
        else:
            self.edit_list.cursor_char = ">"
            self.edit_hint.text = "push edit  K12 home"
        self.edit_list.items = items
        self.edit_list.selected_index = min(sel, len(items) - 1)

    # ------------------------------------------------------------------ #
    # Navigation
    # ------------------------------------------------------------------ #
    def open_slot(self, idx):
        self.focus = idx
        self.mode = "edit"
        self.master_edit = False
        slot = self.slots[idx]
        if slot.effect_name is None:
            self.edit_state = "effect_choose"
        else:
            self.edit_state = "param_nav"
        self.edit_list.selected_index = 0
        self.pages.show_page("edit")
        self.refresh_edit()
        self.update_pixels()

    def go_home(self):
        self.mode = "home"
        self.master_edit = False
        self.pages.show_page("home")
        self.refresh_home()
        self.update_pixels()

    # ------------------------------------------------------------------ #
    # Encoder
    # ------------------------------------------------------------------ #
    def on_encoder(self, direction):
        if self.mode == "home":
            if self.master_edit:
                self.master = self.clamp_master(self.master + 0.05 * direction)
                self.mixer.voice[0].level = self.master
                self.refresh_home()
            else:
                if direction > 0:
                    self.home_list.move_selection_down()
                else:
                    self.home_list.move_selection_up()
            return

        # edit mode
        if self.edit_state == "param_edit":
            self.entry_str = None
            self.adjust_param(self._cur_param(), direction)
            self.refresh_edit()
        else:
            if direction > 0:
                self.edit_list.move_selection_down()
            else:
                self.edit_list.move_selection_up()

    @staticmethod
    def clamp_master(v):
        return 0.0 if v < 0.0 else (1.0 if v > 1.0 else v)

    def adjust_param(self, p, direction):
        slot = self.slots[self.focus]
        v = slot.values[p.key]
        if p.kind == "bool":
            v = not v
        elif p.kind == "enum":
            idx = 0
            for i, (lbl, val) in enumerate(p.options):
                if val == v:
                    idx = i
            idx = (idx + direction) % len(p.options)
            v = p.options[idx][1]
        else:
            v = self.clamp(p, v + p.step * direction)
            if p.kind == "int":
                v = int(round(v))
        self.set_param_value(slot, p, v)

    # ------------------------------------------------------------------ #
    # Encoder push button
    # ------------------------------------------------------------------ #
    def on_button(self):
        if self.mode == "home":
            sel = self.home_list.selected_index
            if sel < NUM_SLOTS:
                self.open_slot(sel)
            else:
                self.master_edit = not self.master_edit
                self.refresh_home()
            return

        # edit mode
        if self.edit_state == "effect_choose":
            sel = self.edit_list.selected_index
            if sel == 0:  # "< Back"
                if self.slots[self.focus].effect_name is None:
                    self.go_home()
                else:
                    self.edit_state = "param_nav"
                    self.edit_list.selected_index = 0
                    self.refresh_edit()
            else:
                self.assign_effect(EFFECT_NAMES[sel - 1])
                self.edit_state = "param_nav"
                self.edit_list.selected_index = 0
                self.refresh_edit()
            self.update_pixels()
            return

        if self.edit_state == "param_nav":
            spec = EFFECTS[self.slots[self.focus].effect_name]
            sel = self.edit_list.selected_index
            if sel < len(spec.params):
                self.param_idx = sel
                self.edit_state = "param_edit"
                self.entry_str = None
                self.refresh_edit()
            elif sel == len(spec.params):       # [Change FX]
                self.edit_state = "effect_choose"
                self.edit_list.selected_index = 0
                self.refresh_edit()
            else:                                # [Remove]
                self.clear_slot(self.focus)
                self.go_home()
            self.update_pixels()
            return

        if self.edit_state == "param_edit":
            self.commit_entry()
            self.edit_state = "param_nav"
            self.refresh_edit()
            self.update_pixels()

    def commit_entry(self):
        if self.entry_str:
            p = self._cur_param()
            try:
                v = float(self.entry_str)
            except ValueError:
                v = None
            if v is not None:
                v = self.clamp(p, v)
                if p.kind == "int":
                    v = int(round(v))
                self.set_param_value(self.slots[self.focus], p, v)
        self.entry_str = None

    # ------------------------------------------------------------------ #
    # Keys
    # ------------------------------------------------------------------ #
    def on_key(self, k):
        if self.mode == "home":
            slot_idx = k // 3
            col = k % 3
            if slot_idx >= NUM_SLOTS:
                return
            if col == 0:
                self.open_slot(slot_idx)
            elif col == 1:
                self.toggle_bypass(slot_idx)
                self.refresh_home()
                self.update_pixels()
            else:
                self.clear_slot(slot_idx)
                self.refresh_home()
                self.update_pixels()
            return

        # edit mode -- number pad while editing a numeric parameter
        if self._cur_param_numeric():
            if k in NUMPAD:
                self.entry_str = (self.entry_str or "") + NUMPAD[k]
                self.refresh_edit()
                return
            if k == ENTER_KEY:
                self.commit_entry()
                self.edit_state = "param_nav"
                self.refresh_edit()
                self.update_pixels()
                return

        # K12 acts as "back" everywhere else in edit mode
        if k == ENTER_KEY:
            if self.edit_state == "param_edit":
                self.entry_str = None
                self.edit_state = "param_nav"
                self.refresh_edit()
                self.update_pixels()
            else:
                self.go_home()

    # ------------------------------------------------------------------ #
    # Neopixels
    # ------------------------------------------------------------------ #
    def update_pixels(self):
        self.pixels.fill(0)
        if self.mode == "home":
            for i, slot in enumerate(self.slots):
                base = i * 3
                if slot.effect_name is None:
                    self.pixels[base] = C_EMPTY
                elif slot.bypassed:
                    self.pixels[base] = C_BYPASS
                    self.pixels[base + 1] = C_BYPASS_KEY
                    self.pixels[base + 2] = C_CLEAR_KEY
                else:
                    self.pixels[base] = C_ACTIVE
                    self.pixels[base + 1] = C_BYPASS_KEY
                    self.pixels[base + 2] = C_CLEAR_KEY
        else:
            if self._cur_param_numeric():
                for kk in range(11):
                    self.pixels[kk] = C_DIGIT
                self.pixels[ENTER_KEY] = C_ENTER
            else:
                self.pixels[ENTER_KEY] = C_BACK

    # ------------------------------------------------------------------ #
    # Main loop
    # ------------------------------------------------------------------ #
    def run(self):
        while True:
            # encoder rotation
            pos = self.encoder.position
            if pos != self._last_pos:
                delta = pos - self._last_pos
                self._last_pos = pos
                step = 1 if delta > 0 else -1
                for _ in range(min(abs(delta), 10)):
                    self.on_encoder(step)

            # encoder push (active low)
            btn = self.button.value
            if (not btn) and self._last_btn:
                self.on_button()
            self._last_btn = btn

            # keys
            event = self.keys.events.get()
            if event and event.pressed:
                self.on_key(event.key_number)


EffectsPedal().run()
