"""Independent, single-window Figure 4 metric helpers; not a simulator or validator.

Definitions traced to LSV1M v1.0 SpontStatisticsOverview and Mozaik v0.4.0.
Inputs must already be the intended population, stimulus and recording window.
No spatial filtering, warm-up deletion, trial averaging or datastore loading here.
"""
import numpy as np


def summary(values):
    a = np.asarray(values, dtype=float)
    if a.ndim != 1 or not np.isfinite(a).all():
        raise ValueError('Expected a finite one-dimensional sample')
    return {'mean': float(a.mean()) if len(a) else None,
            'sd': float(a.std(ddof=0)) if len(a) else None, 'n': len(a)}


def spike_metrics(records, start_ms, stop_ms, bin_ms=10.0):
    """records = [{'id': unique ID, 'times_ms': ordered spike times}, ...].

    Includes silent cells in rates, >=7-spike cells in CV and correlation.
    Uses NumPy histogram semantics, including the rightmost endpoint.
    Rejects nonintegral bin counts: no silent change of effective bin width.
    """
    if not np.isfinite([start_ms, stop_ms, bin_ms]).all():
        raise ValueError('Nonfinite time metadata')
    duration = stop_ms - start_ms
    if duration <= 0 or bin_ms <= 0:
        raise ValueError('Invalid duration or bin size')
    bins = round(duration / bin_ms)
    if bins < 2 or not np.isclose(bins * bin_ms, duration, rtol=0, atol=1e-7):
        raise ValueError('Window must contain at least two complete bins')
    ids = [r['id'] for r in records]
    if len(ids) != len(set(ids)):
        raise ValueError('Duplicate neuron IDs')
    rates, cvs, histograms, eligible = [], [], [], []
    for r in records:
        t = np.asarray(r['times_ms'], dtype=float)
        if (t.ndim != 1 or not np.isfinite(t).all() or
                np.any(np.diff(t) <= 0) or np.any(t < start_ms) or np.any(t > stop_ms)):
            raise ValueError('Invalid, unordered, duplicate or out-of-window spikes')
        rates.append(len(t) / (duration / 1000))
        if len(t) >= 7:
            isi = np.diff(t)
            cvs.append(float(isi.std(ddof=0) / isi.mean()))
            histograms.append(np.histogram(t, bins=bins, range=(start_ms, stop_ms))[0])
            eligible.append(r['id'])
    raw_pairs = np.array([], dtype=float)
    if len(histograms) >= 2:
        # Subsetting before correlation gives the same selected pair values,
        # while avoiding correlations for excluded neurons.
        with np.errstate(divide='ignore', invalid='ignore'):
            corr = np.corrcoef(np.asarray(histograms))
        raw_pairs = corr[np.triu_indices(len(histograms), k=1)]
    undefined = int(np.count_nonzero(~np.isfinite(raw_pairs)))
    pairs = np.nan_to_num(raw_pairs)
    return {'firing_rate_hz': summary(rates), 'isi_cv': summary(cvs),
            'correlation_10ms' if bin_ms == 10 else 'correlation': summary(pairs),
            'eligible_ids': eligible, 'excluded_from_cv_and_correlation': len(ids)-len(eligible),
            'undefined_pair_values_converted_to_zero': undefined,
            'window_ms': [start_ms, stop_ms], 'bin_ms': bin_ms, 'bin_count': bins}


def analog_metrics(records, unit):
    """Equal-weight summary of each neuron's temporal mean, one trial only.

    records = [{'id': unique ID, 'values': one-dimensional sampled signal}, ...].
    Caller must independently verify common time window, sample interval and units.
    unit='mV' preserves signed voltage; 'uS' converts conductance to nS.
    """
    if unit not in ('mV', 'uS'):
        raise ValueError('Expected mV or uS')
    ids = [r['id'] for r in records]
    if len(ids) != len(set(ids)):
        raise ValueError('Duplicate neuron IDs')
    means = []
    for r in records:
        a = np.asarray(r['values'], dtype=float)
        if a.ndim != 1 or len(a) == 0 or not np.isfinite(a).all():
            raise ValueError('Invalid analog trace')
        if unit == 'uS' and np.any(a < 0):
            raise ValueError('Negative conductance')
        means.append(float(a.mean()) * (1000 if unit == 'uS' else 1))
    return dict(summary(means), unit='nS' if unit == 'uS' else 'mV')


def printed_precision_comparison(observed, reference):
    """Descriptive only: agreement at the source's two-significant-digit display.

    This is NOT an equivalence margin, statistical test or biological acceptance.
    """
    if not np.isfinite([observed, reference]).all():
        raise ValueError('Nonfinite comparison')
    return {'difference_from_rounded_reference': observed-reference,
            'same_two_significant_digit_display': format(observed, '.2g') == format(reference, '.2g'),
            'benchmark_pass': None}
