"""Sum processed stems with genre balance trims; peak-normalize to -6 dBFS."""
import numpy as np
import pyloudnorm as pyln

HEADROOM_DBFS = -6.0
TRIM_CAP_DB = 6.0


def _stem_level_db(audio, sr):
    """Integrated LUFS, RMS dB fallback, or None if silent/too short."""
    if len(audio) < int(0.4 * sr):
        rms = float(np.sqrt(np.mean(audio ** 2)))
        return None if rms < 1e-9 else 20 * np.log10(rms)
    try:
        lufs = pyln.Meter(sr).integrated_loudness(audio)
        if np.isfinite(lufs):
            return float(lufs)
    except Exception:
        pass
    rms = float(np.sqrt(np.mean(audio ** 2)))
    return None if rms < 1e-9 else 20 * np.log10(rms)


def _balance_trims_db(stems, sr, profile):
    """Profile balance targets minus measured relative stem loudness, capped."""
    targets = profile.get("stem_trims_db", {})
    n = min(s.shape[0] for s in stems.values())
    levels = {name: _stem_level_db(audio[:n], sr) for name, audio in stems.items()}
    valid = [v for v in levels.values() if v is not None]
    if not valid:
        return {name: 0.0 for name in stems}
    ref = float(np.mean(valid))
    measured = {name: (lv - ref if lv is not None else 0.0) for name, lv in levels.items()}
    return {
        name: float(np.clip(float(targets.get(name, 0.0)) - measured[name],
                             -TRIM_CAP_DB, TRIM_CAP_DB))
        for name in stems
    }


def remix(stems, sr, profile):
    n = min(s.shape[0] for s in stems.values())
    trims = _balance_trims_db(stems, sr, profile)
    mix = np.zeros((n, 2))
    for name, audio in stems.items():
        g = 10 ** (trims.get(name, 0.0) / 20.0)
        mix += audio[:n] * g
    peak = np.max(np.abs(mix)) + 1e-12
    return mix * (10 ** (HEADROOM_DBFS / 20.0) / peak)
