"""Duplex stream — DDJ multi-channel aware."""
from __future__ import annotations

import threading
from typing import Callable

import numpy as np
import sounddevice as sd

from kandi_devices import LIVE_SR, is_ddj


def pick_stereo(raw: np.ndarray, hint: int) -> tuple[np.ndarray, int]:
    x = np.asarray(raw, dtype=np.float32)
    if x.ndim == 1:
        return np.clip(x[:, None], -1.0, 1.0), 0
    if x.shape[1] <= 2:
        return np.clip(x[:, :2], -1.0, 1.0), 0
    best, best_e = 0, -1.0
    for off in range(0, min(x.shape[1] - 1, 6), 2):
        e = float(np.sqrt(np.mean(x[:, off : off + 2] ** 2) + 1e-18))
        if e > best_e:
            best_e, best = e, off
    if hint >= 0 and hint + 1 < x.shape[1]:
        eh = float(np.sqrt(np.mean(x[:, hint : hint + 2] ** 2) + 1e-18))
        if eh >= best_e * 0.85:
            best = hint
    return np.clip(x[:, best : best + 2], -1.0, 1.0), best


class Duplex:
    def __init__(
        self,
        in_dev: int,
        out_dev: int,
        in_name: str,
        process: Callable[[np.ndarray, float], np.ndarray],
        meter: Callable[[np.ndarray, np.ndarray], None],
    ) -> None:
        self.process = process
        self.meter = meter
        self.sr = LIVE_SR
        self.in_ch = min(4, int(sd.query_devices(in_dev).get("max_input_channels") or 2)) if is_ddj(in_name) else 2
        self.pair = 0
        self.err = ""
        self._lock = threading.Lock()
        block = 512
        tries = (
            dict(blocksize=block, dtype="float32"),
            dict(blocksize=block, dtype="float32", latency="high"),
            dict(blocksize=1024, dtype="float32"),
        )

        def cb(indata, outdata, _frames, _time, _status) -> None:
            try:
                raw = np.asarray(indata, dtype=np.float32)
                stereo, self.pair = pick_stereo(raw, self.pair)
                try:
                    y = self.process(stereo, float(self.sr))
                except Exception:
                    y = stereo
                y = np.nan_to_num(np.asarray(y, dtype=np.float32), copy=False)
                y = np.clip(y, -1.0, 1.0)
                if y.ndim == 1:
                    y = np.column_stack([y, y])
                n = min(len(y), len(outdata))
                outdata[:n, :2] = y[:n, :2]
                if n < len(outdata):
                    outdata[n:, :] = 0.0
                in_mono = np.max(np.abs(stereo), axis=1)
                out_mono = np.max(np.abs(y[:n, :2]), axis=1) if n else in_mono
                with self._lock:
                    self.meter(in_mono, out_mono)
            except Exception:
                outdata.fill(0.0)

        rates = [LIVE_SR, 44100]
        for in_native, out_native in (
            (int(sd.query_devices(in_dev).get("default_samplerate") or LIVE_SR), int(sd.query_devices(out_dev).get("default_samplerate") or LIVE_SR)),
        ):
            for r in (in_native, out_native):
                if r not in rates:
                    rates.insert(0, r)
        last_err = ""
        self._stream = None
        for kw in tries:
            for sr in rates:
                for ch in (self.in_ch, 2) if self.in_ch > 2 else (2,):
                    try:
                        self._stream = sd.Stream(
                            device=(in_dev, out_dev),
                            samplerate=sr,
                            channels=(ch, 2),
                            callback=cb,
                            **kw,
                        )
                        self.sr = sr
                        self.in_ch = ch
                        break
                    except OSError as e:
                        last_err = str(e)
                if self._stream:
                    break
            if self._stream:
                break
        if self._stream is None:
            raise RuntimeError(last_err or "Audio open failed")
        self._stream.start()

    def close(self) -> None:
        if self._stream is None:
            return
        self._stream.stop()
        self._stream.close()
        self._stream = None
