#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""热场配乐合成：120 BPM / 30 bar / 60.000s，原创合成，无版权问题。"""
import numpy as np, wave, os

SR   = 44100
BPM  = 120
SPB  = 60.0 / BPM          # 0.5s
BAR  = 4 * SPB             # 2.0s
BARS = 30
DUR  = BARS * BAR          # 60.0s
TAIL = 1.4
N    = int(SR * (DUR + TAIL))

BUF = [np.zeros(N), np.zeros(N)]
rng = np.random.default_rng(7)

def T(bar, beat=0.0):
    return (bar * 4 + beat) * SPB

def add(pos, sig, pan=0.0, gain=1.0):
    i = int(pos * SR)
    if i >= N: return
    j = min(N, i + len(sig))
    s = sig[:j - i] * gain
    th = (pan + 1) * np.pi / 4
    BUF[0][i:j] += s * np.cos(th) * 1.41421
    BUF[1][i:j] += s * np.sin(th) * 1.41421

def tt(d):
    return np.arange(int(SR * d)) / SR

def dec(d, tau, a=0.002):
    t = tt(d)
    e = np.exp(-t / tau)
    if a > 0:
        k = int(SR * a)
        if k > 1: e[:k] *= np.linspace(0, 1, k)
    return e

def lp(x, cut):
    """指数 FIR 近似低通，cut 越小越暗"""
    k = max(2, int(0.5 + 1.0 / max(cut, 1e-4)))
    k = min(k, 900)
    ker = np.exp(-np.arange(k) / (k / 3.0))
    ker /= ker.sum()
    return np.convolve(x, ker, mode="same")

def hp(x):
    y = np.diff(x, prepend=x[0])
    return y

def saw(f, d, det=0.0):
    t = tt(d)
    if det:
        s = np.zeros(len(t))
        for m in (1 - det, 1.0, 1 + det):
            s += 2 * ((t * f * m) % 1.0) - 1
        return s / 3
    return 2 * ((t * f) % 1.0) - 1

def sq(f, d):
    return np.sign(np.sin(2 * np.pi * f * tt(d)))

# ---------------- 音色 ----------------
def kick(d=0.52, punch=1.0):
    t = tt(d)
    f = 48 + 118 * np.exp(-t / 0.028)
    ph = 2 * np.pi * np.cumsum(f) / SR
    body = np.sin(ph) * np.exp(-t / 0.105)
    clk = rng.normal(0, 1, len(t)) * np.exp(-t / 0.0022) * 0.5
    s = body * 1.0 + clk * punch
    return np.tanh(s * 1.5) * 0.92

def snare(d=0.20):
    t = tt(d)
    n = rng.normal(0, 1, len(t))
    n = hp(lp(n, 0.30)) * np.exp(-t / 0.055)
    tone = (np.sin(2 * np.pi * 195 * t) + np.sin(2 * np.pi * 278 * t)) * np.exp(-t / 0.035)
    return (n * 1.0 + tone * 0.45) * 0.55

def clap(d=0.26):
    out = np.zeros(int(SR * d))
    for k, g in enumerate((0.7, 1.0, 0.85, 0.5)):
        o = int(SR * 0.0085 * k)
        t = tt(d - 0.0085 * k)
        n = hp(rng.normal(0, 1, len(t))) * np.exp(-t / (0.03 + 0.012 * k)) * g
        out[o:o + len(n)] += n
    return lp(out, 0.55) * 0.5

def hat(d=0.055, op=False):
    dd = 0.20 if op else d
    t = tt(dd)
    n = hp(hp(rng.normal(0, 1, len(t))))
    return n * np.exp(-t / (0.055 if op else 0.014)) * 0.30

def ride(d=0.32):
    t = tt(d)
    n = hp(rng.normal(0, 1, len(t)))
    tone = sum(np.sin(2 * np.pi * f * t) for f in (3200, 4700, 6100)) / 3
    return (n * 0.6 + tone * 0.4) * np.exp(-t / 0.10) * 0.16

def bass(f, d, bright=0.5):
    s = saw(f, d, 0.004) * 0.6 + np.sin(2 * np.pi * f * tt(d)) * 0.8
    s = lp(s, 0.06 + 0.22 * bright)
    e = dec(d, d * 0.55, 0.004)
    return s * e * 0.6

def pluck(f, d, bright=1.0):
    s = saw(f, d, 0.010) * 0.7 + sq(f, d) * 0.3
    s = lp(s, 0.10 + 0.40 * bright)
    return s * dec(d, d * 0.30, 0.003) * 0.34

def stab(freqs, d, bright=1.0):
    s = np.zeros(int(SR * d))
    for f in freqs:
        s += saw(f, d, 0.013)
    s /= len(freqs)
    s = lp(s, 0.14 + 0.42 * bright)
    t = tt(d)
    e = np.exp(-t / (d * 0.42))
    k = int(SR * 0.006); e[:k] *= np.linspace(0, 1, k)
    return s * e * 0.42

def pad(freqs, d, bright=0.35):
    s = np.zeros(int(SR * d))
    for f in freqs:
        s += saw(f, d, 0.006) + saw(f * 2, d, 0.004) * 0.4
    s /= len(freqs)
    s = lp(s, 0.02 + 0.10 * bright)
    t = tt(d)
    e = np.ones(len(t))
    a = int(SR * 0.25); r = int(SR * 0.45)
    e[:a] = np.linspace(0, 1, a)
    e[-r:] = np.linspace(1, 0, r)
    return s * e * 0.30

def riser(d):
    t = tt(d); p = t / d
    n = rng.normal(0, 1, len(t))
    sweep = np.zeros(len(t))
    step = 2048
    for i in range(0, len(t), step):
        c = 0.02 + 0.55 * (i / len(t)) ** 1.7
        seg = n[i:i + step]
        sweep[i:i + step] = hp(lp(seg, max(c, 0.02)))
    f = 180 * 2 ** (p * 3.2)
    ph = 2 * np.pi * np.cumsum(f) / SR
    tone = np.sin(ph) * 0.35
    env = p ** 2.2
    return (sweep * 0.9 + tone) * env * 0.42

def impact(d=1.6):
    t = tt(d)
    f = 62 * np.exp(-t / 0.22) + 28
    ph = 2 * np.pi * np.cumsum(f) / SR
    boom = np.sin(ph) * np.exp(-t / 0.34)
    cr = hp(rng.normal(0, 1, len(t))) * np.exp(-t / 0.42) * 0.30
    return np.tanh((boom * 1.1 + cr) * 1.3) * 0.72

def crash(d=2.2):
    t = tt(d)
    n = hp(rng.normal(0, 1, len(t)))
    return n * np.exp(-t / 0.70) * 0.30

# ---------------- 和声 ----------------
# A minor: Am - F - C - G
ROOT = [110.00, 87.31, 130.81, 98.00]
TRI  = [[220.00, 261.63, 329.63],
        [174.61, 220.00, 261.63],
        [261.63, 329.63, 392.00],
        [196.00, 246.94, 293.66]]

def chord(bar):
    return bar % 4

# ---------------- 编排 ----------------
FULL   = set(range(9, 20)) | set(range(26, 28))
KICKB  = set(range(2, 20)) | set(range(23, 25)) | set(range(26, 30))
BREAK  = set(range(20, 23))

for bar in range(BARS):
    ci = chord(bar)
    tri = TRI[ci]; root = ROOT[ci]
    b0 = T(bar)

    # --- pad：全程铺底 ---
    pg = 0.55 if bar in BREAK else (0.85 if bar in FULL else 0.6)
    add(b0, pad([tri[0] / 2, tri[1] / 2, tri[2] / 2], BAR * 1.02,
                0.55 if bar in FULL else 0.3), 0.0, pg)

    # --- kick 四踩 ---
    if bar in KICKB:
        for b in range(4):
            g = 1.0 if bar in FULL else 0.72
            add(T(bar, b), kick(), 0.0, g)
        if bar in FULL and bar % 4 == 3:
            add(T(bar, 3.5), kick(0.36, 0.6), 0.0, 0.8)

    # --- clap / snare 2&4 ---
    if bar >= 3 and bar not in BREAK:
        for b in (1, 3):
            add(T(bar, b), clap(), 0.06, 0.9 if bar in FULL else 0.6)
        if bar in FULL and bar % 2 == 1:
            add(T(bar, 3.75), snare(0.14), -0.15, 0.45)

    # --- hats 16分 ---
    if bar >= 1:
        for s in range(8):
            beat = s * 0.5
            if bar in BREAK and s % 2 == 0:
                continue
            g = (0.9 if s % 2 else 0.45) * (1.0 if bar in FULL else 0.62)
            add(T(bar, beat + 0.25), hat(op=(s == 7 and bar % 4 == 3)),
                0.22 if s % 2 else -0.18, g)

    # --- ride 点缀 ---
    if bar in FULL and bar % 2 == 0:
        for b in (0.5, 2.5):
            add(T(bar, b), ride(), 0.3, 0.7)

    # --- bass ---
    if bar in KICKB:
        pat = [(0, 0.5), (1, 0.5), (1.75, 0.25), (2, 0.5), (3, 0.5), (3.5, 0.5)]
        for b, ln in pat:
            f = root if b < 3 else root * (1.5 if bar % 2 else 1.0)
            add(T(bar, b), bass(f, ln * SPB * 0.95,
                                0.75 if bar in FULL else 0.4),
                0.0, 1.0 if bar in FULL else 0.7)

    # --- arp 16分 ---
    if bar not in BREAK:
        seq = [0, 1, 2, 1, 2, 0, 1, 2]
        for s in range(8):
            f = tri[seq[s % len(seq)]] * (2 if (s in (4, 6) and bar in FULL) else 1)
            br = 1.0 if bar in FULL else (0.45 if bar < 5 else 0.7)
            g = 1.0 if bar in FULL else (0.5 if bar < 2 else 0.75)
            add(T(bar, s * 0.5), pluck(f, SPB * 0.48, br),
                -0.28 if s % 2 else 0.28, g)
    else:
        for s in (0, 3, 5):
            add(T(bar, s * 0.5 + 0.5), pluck(tri[s % 3] * 2, SPB * 0.9, 0.25),
                0.2, 0.45)

    # --- 主音 stab ---
    if bar in FULL:
        add(T(bar, 0), stab([f * 2 for f in tri], SPB * 1.6, 1.0), 0.0, 0.85)
        if bar % 2 == 1:
            add(T(bar, 2.5), stab([f * 2 for f in tri], SPB * 0.9, 0.8), 0.1, 0.6)

# --- riser / impact / crash ---
add(T(7), riser(BAR * 2.0), 0.0, 0.85)          # 进 drop
add(T(9), impact(), 0.0, 1.0)
add(T(9), crash(), 0.15, 0.9)
add(T(15, 2), riser(BAR * 0.5), 0.0, 0.55)
add(T(16), impact(1.0), 0.0, 0.7)
add(T(16), crash(1.6), -0.15, 0.7)
add(T(20), impact(1.8), 0.0, 0.95)              # 进 breakdown（数据红线）
add(T(20), crash(2.0), 0.0, 0.6)
add(T(24), riser(BAR * 2.0), 0.0, 0.9)          # 回升
add(T(26), impact(), 0.0, 1.0)                  # 终段 drop
add(T(26), crash(), 0.15, 0.95)
add(T(28), impact(1.9), 0.0, 0.9)               # 收尾
add(T(28), crash(2.4), 0.0, 0.8)
add(T(29, 2), crash(1.8), 0.0, 0.5)

# ---------------- 母带 ----------------
out = []
for ch in BUF:
    x = ch
    # 侧链：kick 位置轻微压低
    duck = np.ones(N)
    for bar in range(BARS):
        if bar not in KICKB: continue
        for b in range(4):
            i = int(T(bar, b) * SR)
            L = int(SR * 0.16)
            if i + L < N:
                duck[i:i + L] = np.minimum(duck[i:i + L],
                                           0.62 + 0.38 * np.linspace(0, 1, L) ** 0.6)
    x = x * duck
    x = np.tanh(x * 0.62) / np.tanh(0.62)     # 软限
    out.append(x)

# 尾部淡出
fo = int(SR * 1.0)
for x in out:
    x[-fo:] *= np.linspace(1, 0, fo)
    x[:int(SR * 0.02)] *= np.linspace(0, 1, int(SR * 0.02))

peak = max(np.abs(out[0]).max(), np.abs(out[1]).max())
g = 0.90 / peak
inter = np.empty(N * 2)
inter[0::2] = out[0] * g
inter[1::2] = out[1] * g
pcm = (inter * 32767).astype("<i2")

path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "bgm.wav")
w = wave.open(path, "wb")
w.setnchannels(2); w.setsampwidth(2); w.setframerate(SR)
w.writeframes(pcm.tobytes()); w.close()
print("wrote", path, "%.2fs" % ((N) / SR), "peak %.3f" % peak)
