#!/usr/bin/env python3
"""
Visual Buddy Bridge
===================

Listens to everything Traktor already broadcasts and re-publishes it as a
simple JSON WebSocket feed that the Visual Buddy renderer consumes.

  Traktor Broadcasting (Icecast, port 8000)  ->  track metadata
  Traktor MIDI Clock  (virtual MIDI port)    ->  bpm / beat / bar
  Traktor MIDI notes  (virtual MIDI port)    ->  hotcue triggers
                                             ->  ws://0.0.0.0:8765

Nothing is injected into Traktor. No plugins, no memory reading.

Usage
-----
    pip install -r requirements.txt
    python vbuddy_bridge.py --list-midi
    python vbuddy_bridge.py --midi "IAC Driver Bus 1"

Optional OSC forwarding to Resolume / VDMX:
    python vbuddy_bridge.py --midi "IAC Driver Bus 1" --osc 127.0.0.1:7000
"""

from __future__ import annotations

import argparse
import asyncio
import json
import signal
import threading
import time
from dataclasses import dataclass, field
from typing import Any, Optional

try:
    import websockets
except ImportError:  # pragma: no cover
    raise SystemExit("Missing dependency: pip install websockets")

try:
    import mido
except ImportError:  # pragma: no cover
    mido = None

try:
    from traktor_nowplaying import Listener as TraktorListener
except ImportError:  # pragma: no cover
    TraktorListener = None

try:
    from pythonosc.udp_client import SimpleUDPClient
except ImportError:  # pragma: no cover
    SimpleUDPClient = None


PPQN = 24  # MIDI clock pulses per quarter note
BEATS_PER_BAR = 4


# --------------------------------------------------------------------------
# Shared state / fan-out
# --------------------------------------------------------------------------


@dataclass
class Hub:
    """Fans messages out to every connected WebSocket client."""

    loop: asyncio.AbstractEventLoop
    clients: set = field(default_factory=set)
    last_track: Optional[dict] = None
    osc: Any = None

    def register(self, ws) -> None:
        self.clients.add(ws)

    def unregister(self, ws) -> None:
        self.clients.discard(ws)

    async def _send(self, payload: dict) -> None:
        if not self.clients:
            return
        raw = json.dumps(payload)
        dead = []
        for ws in list(self.clients):
            try:
                await ws.send(raw)
            except Exception:
                dead.append(ws)
        for ws in dead:
            self.clients.discard(ws)

    def publish(self, payload: dict) -> None:
        """Thread-safe: callable from the MIDI or metadata threads."""
        payload.setdefault("ts", time.time())
        if payload.get("type") == "track":
            self.last_track = payload
        self._publish_osc(payload)
        try:
            asyncio.run_coroutine_threadsafe(self._send(payload), self.loop)
        except RuntimeError:
            pass

    def _publish_osc(self, payload: dict) -> None:
        if not self.osc:
            return
        kind = payload.get("type")
        try:
            if kind == "clock":
                self.osc.send_message("/vbuddy/bpm", float(payload.get("bpm") or 0.0))
                self.osc.send_message("/vbuddy/beat", int(payload.get("beat") or 0))
                self.osc.send_message("/vbuddy/playing", int(bool(payload.get("playing"))))
            elif kind == "track":
                self.osc.send_message(
                    "/vbuddy/track",
                    [str(payload.get("artist", "")), str(payload.get("title", ""))],
                )
            elif kind == "cue":
                self.osc.send_message("/vbuddy/cue", int(payload.get("index") or 0))
        except Exception as exc:  # pragma: no cover
            print(f"[osc ] send failed: {exc}")


# --------------------------------------------------------------------------
# MIDI clock + hotcues
# --------------------------------------------------------------------------


class MidiWorker(threading.Thread):
    """Turns raw MIDI clock into BPM / beat / bar, and notes into cue events."""

    daemon = True

    def __init__(self, hub: Hub, port_name: str, cue_base: int, cue_channel: Optional[int]):
        super().__init__(name="midi")
        self.hub = hub
        self.port_name = port_name
        self.cue_base = cue_base
        self.cue_channel = cue_channel
        self._stop = threading.Event()
        self.pulses = 0
        self.beat = 0
        self.bar = 0
        self.playing = False
        self._tick_times: list[float] = []

    def stop(self) -> None:
        self._stop.set()

    def _bpm(self) -> float:
        """Rolling average over the last ~2 beats of clock ticks."""
        if len(self._tick_times) < 8:
            return 0.0
        span = self._tick_times[-1] - self._tick_times[0]
        if span <= 0:
            return 0.0
        ticks = len(self._tick_times) - 1
        seconds_per_tick = span / ticks
        return round(60.0 / (seconds_per_tick * PPQN), 2)

    def run(self) -> None:
        if mido is None:
            print("[midi] mido not installed - clock sync disabled")
            return
        try:
            port = mido.open_input(self.port_name)
        except Exception as exc:
            print(f"[midi] could not open '{self.port_name}': {exc}")
            print("[midi] run with --list-midi to see available ports")
            return

        print(f"[midi] listening on '{self.port_name}'")
        with port:
            for msg in port:
                if self._stop.is_set():
                    break

                if msg.type == "clock":
                    now = time.perf_counter()
                    self._tick_times.append(now)
                    if len(self._tick_times) > PPQN * 2:
                        self._tick_times.pop(0)

                    self.pulses += 1
                    if self.pulses % PPQN == 0:
                        self.beat = (self.beat + 1) % BEATS_PER_BAR
                        if self.beat == 0:
                            self.bar += 1
                        self.hub.publish(
                            {
                                "type": "clock",
                                "bpm": self._bpm(),
                                "beat": self.beat,
                                "bar": self.bar,
                                "playing": self.playing,
                            }
                        )

                elif msg.type in ("start", "continue"):
                    self.playing = True
                    self.pulses = 0
                    self.beat = 0
                    self.hub.publish({"type": "transport", "playing": True})

                elif msg.type == "stop":
                    self.playing = False
                    self.hub.publish({"type": "transport", "playing": False})

                elif msg.type == "note_on" and msg.velocity > 0:
                    if self.cue_channel is not None and msg.channel != self.cue_channel:
                        continue
                    index = msg.note - self.cue_base
                    if 0 <= index < 16:
                        deck = "A" if index < 8 else "B"
                        self.hub.publish(
                            {
                                "type": "cue",
                                "deck": deck,
                                "index": (index % 8) + 1,
                                "note": msg.note,
                                "channel": msg.channel,
                            }
                        )


# --------------------------------------------------------------------------
# Traktor broadcast metadata
# --------------------------------------------------------------------------


class MetadataWorker(threading.Thread):
    """Catches Traktor's Icecast broadcast and extracts artist / title."""

    daemon = True

    def __init__(self, hub: Hub, port: int):
        super().__init__(name="metadata")
        self.hub = hub
        self.port = port

    def run(self) -> None:
        if TraktorListener is None:
            print("[meta] traktor_nowplaying not installed - metadata disabled")
            print("[meta] pip install traktor-nowplaying")
            return

        def on_track(data: dict) -> None:
            artist = (data.get("artist") or "").strip()
            title = (data.get("title") or "").strip()
            if not artist and not title:
                return
            print(f"[meta] {artist} - {title}")
            self.hub.publish(
                {
                    "type": "track",
                    "deck": "A",
                    "artist": artist,
                    "title": title,
                    "filename": f"{artist} - {title}".strip(" -"),
                }
            )

        print(f"[meta] icecast sink on port {self.port} (Traktor > Preferences > Broadcasting)")
        try:
            listener = TraktorListener(port=self.port, quiet=True, custom_callback=on_track)
            listener.start()
        except TypeError:
            listener = TraktorListener(port=self.port, quiet=True)
            listener.start(callback=on_track)
        except Exception as exc:  # pragma: no cover
            print(f"[meta] listener stopped: {exc}")


# --------------------------------------------------------------------------
# WebSocket server
# --------------------------------------------------------------------------


async def serve(hub: Hub, host: str, port: int) -> None:
    async def handler(ws):
        hub.register(ws)
        peer = getattr(ws, "remote_address", ("?",))[0]
        print(f"[ws  ] client connected: {peer} ({len(hub.clients)} total)")
        await ws.send(json.dumps({"type": "hello", "app": "vbuddy-bridge", "version": "0.1.0"}))
        if hub.last_track:
            await ws.send(json.dumps(hub.last_track))
        try:
            async for _ in ws:
                pass  # renderer is receive-only for now
        except Exception:
            pass
        finally:
            hub.unregister(ws)
            print(f"[ws  ] client gone ({len(hub.clients)} total)")

    async with websockets.serve(handler, host, port, ping_interval=20):
        print(f"[ws  ] serving ws://{host}:{port}")
        print("[ws  ] open the Visual Buddy renderer and connect")
        await asyncio.Future()


# --------------------------------------------------------------------------
# Entrypoint
# --------------------------------------------------------------------------


def list_midi_ports() -> None:
    if mido is None:
        print("mido is not installed: pip install mido python-rtmidi")
        return
    ports = mido.get_input_names()
    if not ports:
        print("No MIDI input ports found.")
        print("macOS:   enable the IAC Driver in Audio MIDI Setup.")
        print("Windows: install loopMIDI and create a port.")
        return
    print("Available MIDI input ports:")
    for name in ports:
        print(f"  - {name}")


def main() -> None:
    parser = argparse.ArgumentParser(description="Visual Buddy Traktor bridge")
    parser.add_argument("--host", default="0.0.0.0", help="WebSocket bind host")
    parser.add_argument("--ws-port", type=int, default=8765, help="WebSocket port")
    parser.add_argument("--meta-port", type=int, default=8000, help="Traktor broadcast port")
    parser.add_argument("--midi", default=None, help="MIDI input port name (see --list-midi)")
    parser.add_argument("--cue-base", type=int, default=36, help="MIDI note of hotcue 1 (deck A)")
    parser.add_argument("--cue-channel", type=int, default=None, help="Restrict cues to a channel (0-15)")
    parser.add_argument("--osc", default=None, help="Forward events to host:port, e.g. 127.0.0.1:7000")
    parser.add_argument("--no-meta", action="store_true", help="Disable the metadata listener")
    parser.add_argument("--list-midi", action="store_true", help="List MIDI inputs and exit")
    args = parser.parse_args()

    if args.list_midi:
        list_midi_ports()
        return

    loop = asyncio.new_event_loop()
    asyncio.set_event_loop(loop)
    hub = Hub(loop=loop)

    if args.osc:
        if SimpleUDPClient is None:
            print("[osc ] python-osc not installed - skipping OSC output")
        else:
            osc_host, _, osc_port = args.osc.partition(":")
            hub.osc = SimpleUDPClient(osc_host, int(osc_port or 7000))
            print(f"[osc ] forwarding to {osc_host}:{osc_port or 7000}")

    workers: list[threading.Thread] = []

    if args.midi:
        midi = MidiWorker(hub, args.midi, args.cue_base, args.cue_channel)
        midi.start()
        workers.append(midi)
    else:
        print("[midi] no --midi port given; beat sync disabled (use --list-midi)")

    if not args.no_meta:
        meta = MetadataWorker(hub, args.meta_port)
        meta.start()
        workers.append(meta)

    def shutdown(*_: Any) -> None:
        print("\n[bye ] shutting down")
        loop.stop()

    try:
        signal.signal(signal.SIGINT, shutdown)
        signal.signal(signal.SIGTERM, shutdown)
    except ValueError:
        pass

    try:
        loop.run_until_complete(serve(hub, args.host, args.ws_port))
    except (KeyboardInterrupt, RuntimeError):
        pass
    finally:
        loop.close()


if __name__ == "__main__":
    main()
