"""Analyze a one-dimensional impedance step-response CSV.

Compatible with impedance_response.csv generated by impedance_1d.py.
Only the Python standard library is required.
"""

import argparse
import csv
import json
import math
from pathlib import Path


def mean(values: list[float]) -> float:
    if not values:
        raise ValueError("metric window contains no samples")
    return sum(values) / len(values)


def rms(values: list[float]) -> float:
    if not values:
        raise ValueError("metric window contains no samples")
    return math.sqrt(sum(value * value for value in values) / len(values))


def settling_time(
    rows: list[dict[str, float]],
    start_s: float,
    target_position_m: float,
    tolerance_m: float,
) -> float:
    """Return time after start when all remaining samples stay in the band."""
    if tolerance_m <= 0.0:
        raise ValueError("settling tolerance must be positive")
    last_violation = -1
    for index, row in enumerate(rows):
        if abs(row["position_m"] - target_position_m) > tolerance_m:
            last_violation = index
    settled_index = last_violation + 1
    if settled_index >= len(rows):
        return math.inf
    return max(0.0, rows[settled_index]["time_s"] - start_s)


def validate_rows(rows: list[dict[str, float]]) -> None:
    required = {"time_s", "force_n", "position_m"}
    if len(rows) < 3:
        raise ValueError("at least three samples are required")
    for index, row in enumerate(rows):
        if not isinstance(row, dict) or not required.issubset(row):
            raise ValueError(f"sample {index}: missing required fields")
        for name in required:
            value = row[name]
            if isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value):
                raise ValueError(f"sample {index}: {name} must be a finite number")
    if any(b["time_s"] <= a["time_s"] for a, b in zip(rows, rows[1:])):
        raise ValueError("time_s must be strictly increasing")


def load_rows(path: Path) -> list[dict[str, float]]:
    required = {"time_s", "force_n", "position_m"}
    with path.open(newline="", encoding="utf-8") as source:
        reader = csv.DictReader(source)
        fields = reader.fieldnames or []
        missing = required.difference(fields)
        if missing:
            raise ValueError(f"missing CSV columns: {sorted(missing)}")
        if len(set(fields)) != len(fields):
            raise ValueError("duplicate CSV column names")
        rows = []
        for row in reader:
            if None in row:
                raise ValueError(f"CSV line {reader.line_num}: extra cells without headers")
            try:
                rows.append({name: float(row[name]) for name in required})
            except (TypeError, ValueError) as error:
                raise ValueError(f"CSV line {reader.line_num}: missing or nonnumeric value") from error
    validate_rows(rows)
    return rows


def analyze(
    rows: list[dict[str, float]], step_on_s: float, step_off_s: float
) -> dict[str, float]:
    validate_rows(rows)  # Direct API calls must obey the same contract as CSV input.
    if not rows[0]["time_s"] < step_on_s < step_off_s < rows[-1]["time_s"]:
        raise ValueError("step times must lie strictly inside the log interval")

    step_duration = step_off_s - step_on_s
    baseline = [row for row in rows if row["time_s"] < step_on_s]
    steady = [
        row
        for row in rows
        if step_off_s - 0.2 * step_duration <= row["time_s"] < step_off_s
    ]
    released = [row for row in rows if row["time_s"] >= step_off_s]

    baseline_position = mean([row["position_m"] for row in baseline])
    steady_position = mean([row["position_m"] for row in steady])
    steady_force = mean([row["force_n"] for row in steady])
    displacement = steady_position - baseline_position
    if abs(displacement) < 1e-12:
        raise ValueError("steady displacement is too small to estimate stiffness")

    direction = 1.0 if displacement >= 0.0 else -1.0
    peak_magnitude = max(
        direction * (row["position_m"] - baseline_position)
        for row in rows
        if step_on_s <= row["time_s"] < step_off_s
    )
    steady_magnitude = abs(displacement)
    loaded = [
        row for row in rows if step_on_s <= row["time_s"] < step_off_s
    ]
    settling_band_m = 0.02 * steady_magnitude

    return {
        "baseline_position_m": baseline_position,
        "steady_displacement_m": displacement,
        "steady_force_n": steady_force,
        "estimated_stiffness_n_m": steady_force / displacement,
        "peak_displacement_m": direction * peak_magnitude,
        "overshoot_percent": max(
            0.0, (peak_magnitude - steady_magnitude) / steady_magnitude * 100.0
        ),
        "settling_time_2pct_s": settling_time(
            loaded, step_on_s, steady_position, settling_band_m
        ),
        "release_settling_time_2pct_s": settling_time(
            released, step_off_s, baseline_position, settling_band_m
        ),
        "released_position_rms_m": rms(
            [row["position_m"] - baseline_position for row in released]
        ),
        "sample_period_mean_s": (
            rows[-1]["time_s"] - rows[0]["time_s"]
        ) / (len(rows) - 1),
    }


def load_acceptance(path: Path) -> dict[str, float]:
    required = {
        "expected_stiffness_n_m",
        "stiffness_relative_tolerance",
        "max_overshoot_percent",
        "max_released_position_rms_m",
        "max_settling_time_2pct_s",
        "max_release_settling_time_2pct_s",
        "expected_sample_period_s",
        "sample_period_relative_tolerance",
    }
    with path.open(encoding="utf-8") as source:
        raw = json.load(source)
    if not isinstance(raw, dict):
        raise ValueError("acceptance JSON root must be an object")
    missing = required.difference(raw)
    if missing:
        raise ValueError(f"missing acceptance fields: {sorted(missing)}")

    acceptance: dict[str, float] = {}
    for name in required:
        value = raw[name]
        if (
            isinstance(value, bool)
            or not isinstance(value, (int, float))
            or not math.isfinite(value)
        ):
            raise ValueError(f"acceptance field {name} must be finite")
        acceptance[name] = float(value)
    if acceptance["expected_stiffness_n_m"] <= 0.0:
        raise ValueError("expected_stiffness_n_m must be positive")
    if acceptance["expected_sample_period_s"] <= 0.0:
        raise ValueError("expected_sample_period_s must be positive")
    for name in required - {
        "expected_stiffness_n_m",
        "expected_sample_period_s",
    }:
        if acceptance[name] < 0.0:
            raise ValueError(f"acceptance field {name} cannot be negative")
    return acceptance


def acceptance_failures(
    metrics: dict[str, float], acceptance: dict[str, float]
) -> list[str]:
    stiffness_error = abs(
        metrics["estimated_stiffness_n_m"]
        - acceptance["expected_stiffness_n_m"]
    ) / acceptance["expected_stiffness_n_m"]
    period_error = abs(
        metrics["sample_period_mean_s"]
        - acceptance["expected_sample_period_s"]
    ) / acceptance["expected_sample_period_s"]

    checks = (
        (
            "stiffness_relative_error",
            stiffness_error,
            acceptance["stiffness_relative_tolerance"],
        ),
        (
            "overshoot_percent",
            metrics["overshoot_percent"],
            acceptance["max_overshoot_percent"],
        ),
        (
            "released_position_rms_m",
            metrics["released_position_rms_m"],
            acceptance["max_released_position_rms_m"],
        ),
        (
            "settling_time_2pct_s",
            metrics["settling_time_2pct_s"],
            acceptance["max_settling_time_2pct_s"],
        ),
        (
            "release_settling_time_2pct_s",
            metrics["release_settling_time_2pct_s"],
            acceptance["max_release_settling_time_2pct_s"],
        ),
        (
            "sample_period_relative_error",
            period_error,
            acceptance["sample_period_relative_tolerance"],
        ),
    )
    return [
        f"{name}: actual={actual:.9g}, limit={limit:.9g}"
        for name, actual, limit in checks
        if not math.isfinite(actual) or actual > limit
    ]


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("csv_path", type=Path)
    parser.add_argument("--step-on", type=float, default=0.5)
    parser.add_argument("--step-off", type=float, default=1.5)
    parser.add_argument("--acceptance", type=Path)
    args = parser.parse_args()

    metrics = analyze(load_rows(args.csv_path), args.step_on, args.step_off)
    print(json.dumps(metrics, indent=2, ensure_ascii=False))
    if args.acceptance is not None:
        failures = acceptance_failures(metrics, load_acceptance(args.acceptance))
        if failures:
            print("FAIL:")
            for failure in failures:
                print(f"- {failure}")
            raise SystemExit(1)
        print("PASS: all acceptance checks satisfied")


if __name__ == "__main__":
    main()
