"""Course DSP coefficients and numerical checks. Requires numpy and scipy.
Run: python design_filters.py ; outputs beside this file.
These checks validate mathematics, not MCU timing or the CMSIS binary.
"""
from pathlib import Path
import json
import numpy as np
from scipy import signal

FS, N = 8000, 256
b = signal.firwin(31, 1000, fs=FS, window="hamming")
sos = signal.butter(2, 1000, fs=FS, output="sos")
normalized = sos / sos[:, 3:4]
cmsis = np.column_stack((normalized[:, :3], -normalized[:, 4:]))
out = Path(__file__).resolve().parent

def c_array(values):
    return ",\n    ".join(f"{v:.10e}f" for v in np.asarray(values).ravel())

header = ("/* Generated for Fs=8000 Hz. Cortex-M scalar CMSIS-DSP f32. */\n"
          "#ifndef COURSE_FILTER_COEFFS_H\n#define COURSE_FILTER_COEFFS_H\n"
          "#define FIR_TAPS 31U\n#define IIR_STAGES 1U\n"
          "static const float fir_coeffs[FIR_TAPS] = {\n    " + c_array(b[::-1]) + "\n};\n"
          "static const float iir_coeffs[5U * IIR_STAGES] = {\n    " + c_array(cmsis) + "\n};\n#endif\n")
(out / "filter_coeffs.h").write_text(header, encoding="utf-8")

n = np.arange(N)
x = 2048 + 1000 * np.sin(2*np.pi*1000*n/FS)
w = 0.5 - 0.5*np.cos(2*np.pi*n/N)
bins = np.fft.rfft((x-x.mean())*w)
amp = 2*np.abs(bins)/w.sum()
amp[[0,-1]] *= 0.5
assert np.argmax(amp[1:-1])+1 == 32
assert abs(amp[32]-1000) < 1e-7

# Check the exact packing and scaling used in the tutorial.
packed = np.empty(N)
packed[0], packed[1] = bins[0].real, bins[-1].real
packed[2::2], packed[3::2] = bins[1:-1].real, bins[1:-1].imag
decoded = np.r_[abs(packed[0]), 2*np.hypot(packed[2::2],packed[3::2]), abs(packed[1])]/w.sum()
assert np.allclose(amp, decoded)

impulse = np.r_[1., np.zeros(511)]
assert np.allclose(signal.lfilter(b, [1.], impulse)[:31], b)
assert np.allclose(signal.lfilter(b, [1.], impulse)[31:], 0)
assert np.all(np.abs(np.roots(sos[0,3:])) < 1)

t = np.arange(4096)/FS
mixed = np.sin(2*np.pi*500*t)+0.5*np.sin(2*np.pi*2000*t)
f_state, i_state = np.zeros(30), np.zeros((1,2))
fir_blocks, iir_blocks = [], []
for block in mixed.reshape(-1, N):
    fy, f_state = signal.lfilter(b, [1.], block, zi=f_state)
    iy, i_state = signal.sosfilt(sos, block, zi=i_state)
    fir_blocks.append(fy); iir_blocks.append(iy)
fir_y, iir_y = np.concatenate(fir_blocks), np.concatenate(iir_blocks)
assert np.allclose(fir_y, signal.lfilter(b, [1.], mixed))
assert np.allclose(iir_y, signal.sosfilt(sos, mixed))

# Mirror DF1 float32 feedback convention and compare to SciPy reference.
c = cmsis.ravel().astype(np.float32)
x1=x2=y1=y2=np.float32(0)
manual=[]
for v in mixed.astype(np.float32):
    y=c[0]*v+c[1]*x1+c[2]*x2+c[3]*y1+c[4]*y2
    x2,x1,y2,y1=x1,v,y1,y
    manual.append(y)
assert np.max(np.abs(np.asarray(manual)-iir_y)) < 2e-6

freqs = np.array([500,1000,2000])
_,fh = signal.freqz(b, worN=freqs, fs=FS)
_,ih = signal.sosfreqz(sos, worN=freqs, fs=FS)
report={"sample_rate_hz":FS,"fft_points":N,"fft_peak_hz":1000,
        "fft_peak_amplitude":float(amp[32]),"frequencies_hz":freqs.tolist(),
        "fir_gain_db":(20*np.log10(abs(fh))).tolist(),
        "iir_gain_db":(20*np.log10(abs(ih))).tolist(),
        "iir_pole_magnitudes":abs(np.roots(sos[0,3:])).tolist(),
        "checks":"FFT packing/amplitude, impulse, poles, block continuity, f32 DF1 passed",
        "hardware_tested":False}
(out/"validation.json").write_text(json.dumps(report,indent=2)+"\n",encoding="utf-8")
np.savetxt(out/"filter_demo.csv",np.column_stack((t,mixed,fir_y,iir_y)),delimiter=",",
           header="time_s,input,fir_output,iir_output",comments="",fmt="%.9g")
print(json.dumps(report, indent=2))
