#!/usr/bin/env python3
"""
IEPE accelerometer / microphone capture with a spectrum — USB-4431 and USB-4432.

The 4431 series supplies the 2.1 mA constant-current excitation and the AC
coupling, so an accelerometer connects straight to the BNC with no external
signal conditioner. This script captures a hardware-timed block and prints the
largest spectral peaks, which is the fastest way to sanity-check a vibration
setup before writing any analysis code.

    pip install nidaqmx numpy
    python python-iepe-vibration.py --device Dev2 --channels ai0:3 --rate 51200

Sensitivity comes from the accelerometer datasheet — 100 mV/g is the usual
value for a general-purpose industrial accelerometer.
"""

from __future__ import annotations

import argparse
import sys

try:
    import numpy as np
    import nidaqmx
    from nidaqmx.constants import AcquisitionType, Coupling
except ImportError:  # pragma: no cover
    sys.exit("Requires nidaqmx and numpy. Run: pip install nidaqmx numpy")


def spectrum(samples: np.ndarray, rate: float) -> tuple[np.ndarray, np.ndarray]:
    """Single-sided amplitude spectrum in g, Hann windowed and corrected."""
    n = len(samples)
    window = np.hanning(n)
    # Amplitude correction: coherent gain of the Hann window is 0.5.
    coherent_gain = window.sum() / n
    spectrum = np.fft.rfft(samples * window)
    amplitude = np.abs(spectrum) * (2.0 / (n * coherent_gain))
    amplitude[0] /= 2.0
    freqs = np.fft.rfftfreq(n, d=1.0 / rate)
    return freqs, amplitude


def main() -> int:
    parser = argparse.ArgumentParser(description="IEPE vibration capture")
    parser.add_argument("--device", default="Dev2")
    parser.add_argument("--channels", default="ai0:3", help="IEPE input channels")
    parser.add_argument("--rate", type=float, default=51_200.0, help="S/s per channel")
    parser.add_argument("--seconds", type=float, default=1.0)
    parser.add_argument("--sensitivity", type=float, default=100.0, help="mV per g")
    parser.add_argument("--peaks", type=int, default=5, help="peaks to print per channel")
    args = parser.parse_args()

    samples_per_channel = int(args.rate * args.seconds)
    scale = 1000.0 / args.sensitivity  # mV/g -> g per volt

    print(f"Device  : {args.device}")
    print(f"Channels: {args.channels}")
    print(f"Rate    : {args.rate:,.0f} S/s per channel")
    print(f"Block   : {samples_per_channel:,} samples ({args.seconds:g} s)")
    print(f"Scale   : {scale:.4g} g per volt ({args.sensitivity:g} mV/g)")
    print()

    with nidaqmx.Task() as task:
        task.ai_channels.add_ai_accel_chan(
            f"{args.device}/{args.channels}",
            min_val=-10.0,
            max_val=10.0,
            sensitivity=args.sensitivity / 1000.0,  # V per g
            current_excit_val=0.0021,               # 2.1 mA IEPE excitation
        )
        task.timing.cfg_samp_clk_timing(
            rate=args.rate,
            sample_mode=AcquisitionType.FINITE,
            samps_per_chan=samples_per_channel,
        )
        data = task.read(number_of_samples_per_channel=samples_per_channel)

    channels = data if data and isinstance(data[0], list) else [data]

    for index, column in enumerate(channels):
        signal = np.asarray(column, dtype=np.float64) * scale
        freqs, amplitude = spectrum(signal, args.rate)

        overall = float(np.sqrt(np.mean(signal**2)))
        peak = float(np.max(np.abs(signal)))

        print(f"ch{index}  RMS {overall:.4f} g   peak {peak:.4f} g")

        # Ignore the DC bin and anything below 5 Hz, which is AC-coupling noise.
        usable = freqs >= 5.0
        f_use = freqs[usable]
        a_use = amplitude[usable]
        top = np.argsort(a_use)[::-1][: args.peaks]

        for rank, idx in enumerate(top, start=1):
            print(f"      #{rank}  {f_use[idx]:9.1f} Hz   {a_use[idx]:.4f} g")
        print()

    print("Tip: a bearing defect shows up as a peak at the ball-pass frequency")
    print("     and its harmonics. Compare against the shaft speed to identify it.")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
