"""Supplement unchanged upstream tests with neuron- and segment-aware diagnostics.

Run inside the pinned research environment, after run_regression.py. This reads
the preserved datastores and never changes the upstream tests or their results.
"""
import argparse
import json
import pathlib
import sys

import numpy as np


def numeric(signal, unit):
    return np.asarray(signal.rescale(unit).magnitude, dtype=float).reshape(-1)


def compare(a, b):
    result = {'run_count': int(a.size), 'reference_count': int(b.size),
              'same_shape': a.shape == b.shape}
    if a.shape != b.shape:
        return dict(result, exact_equal=False)
    equal = (a == b) | (np.isnan(a) & np.isnan(b))
    result.update(exact_equal=bool(equal.all()), different_samples=int((~equal).sum()))
    if not equal.all():
        i = int(np.flatnonzero(~equal)[0])
        result['first_difference'] = {'index': i, 'run': float(a[i]), 'reference': float(b[i])}
    if a.size and np.isfinite(a).all() and np.isfinite(b).all():
        result['max_absolute_difference'] = float(np.max(np.abs(a-b)))
    return result


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--run', type=pathlib.Path, required=True)
    args = parser.parse_args()
    run = args.run.resolve()
    framework = run / 'framework'
    sys.path.insert(0, str(framework))
    from tests.full_model.test_models import TestModel
    helper = TestModel()
    name = 'LSV1M_tiny_stepcurrentmodule'
    actual = helper.load_datastore(str(framework / 'tests/full_model/models' / name / 'LSV1M_pytest_____'))
    reference = helper.load_datastore(str(framework / 'tests/full_model/reference_data' / name))
    populations = ['X_ON', 'X_OFF', 'V1_Exc_L4', 'V1_Inh_L4', 'V1_Exc_L2/3', 'V1_Inh_L2/3']
    report = {'scope': 'Supplemental comparison; upstream JUnit remains the acceptance record',
              'units': {'spikes': 'ms', 'voltages': 'mV'}, 'populations': []}
    for population in populations:
        run_segments = helper.get_segments(actual, population)
        ref_segments = helper.get_segments(reference, population)
        run_by_id = {str(s.identifier): s for s in run_segments}
        ref_by_id = {str(s.identifier): s for s in ref_segments}
        row = {'population': population, 'run_segments': len(run_segments),
               'reference_segments': len(ref_segments), 'segments': []}
        if len(run_by_id) != len(run_segments) or len(ref_by_id) != len(ref_segments):
            raise ValueError('Duplicate segment identifier; manual semantic matching required')
        for identifier in sorted(set(run_by_id) | set(ref_by_id)):
            pair = {'segment_identifier': identifier, 'spikes': [], 'voltages': []}
            row['segments'].append(pair)
            if identifier not in run_by_id or identifier not in ref_by_id:
                pair['error'] = 'Missing segment in one datastore'
                continue
            a, b = run_by_id[identifier], ref_by_id[identifier]
            pair['run_stimulus'] = str(a.annotations.get('stimulus'))
            pair['reference_stimulus'] = str(b.annotations.get('stimulus'))
            pair['same_stimulus_annotation'] = pair['run_stimulus'] == pair['reference_stimulus']
            for kind, id_method, signal_method, unit, limit in [
                ('spikes', 'get_stored_spike_train_ids', 'get_spiketrain', 'ms', None),
                ('voltages', 'get_stored_vm_ids', 'get_vm', 'mV', 25),
            ]:
                if kind == 'voltages' and population in populations[:2]:
                    continue
                a_ids = sorted(getattr(a, id_method)())[:limit]
                b_ids = sorted(getattr(b, id_method)())[:limit]
                pair[kind + '_run_ids'] = [int(x) for x in a_ids]
                pair[kind + '_reference_ids'] = [int(x) for x in b_ids]
                for neuron_id in sorted(set(a_ids) | set(b_ids)):
                    result = {'neuron_id': int(neuron_id)}
                    pair[kind].append(result)
                    if neuron_id not in a_ids or neuron_id not in b_ids:
                        result['error'] = 'Missing recorded neuron in one datastore'
                        continue
                    a_signal = getattr(a, signal_method)(neuron_id)
                    b_signal = getattr(b, signal_method)(neuron_id)
                    result.update(compare(numeric(a_signal, unit), numeric(b_signal, unit)))
                    result.update(run_start_ms=float(a_signal.t_start.rescale('ms')),
                                  reference_start_ms=float(b_signal.t_start.rescale('ms')),
                                  run_stop_ms=float(a_signal.t_stop.rescale('ms')),
                                  reference_stop_ms=float(b_signal.t_stop.rescale('ms')))
        report['populations'].append(row)
    (run / 'diagnostics.json').write_text(json.dumps(report, indent=2, allow_nan=False) + '\n')
    # A reproducible, explicitly selected example; all comparisons stay in JSON.
    segment_index = 0
    population = 'V1_Exc_L4'
    example = {'population': population, 'ordered_segment_index': segment_index,
               'selection_rule': 'First upstream-ordered segment; lowest recorded voltage neuron ID'}
    for label, datastore in [('run', actual), ('reference', reference)]:
        segment = helper.get_segments(datastore, population)[segment_index]
        neuron_id = sorted(segment.get_stored_vm_ids())[0]
        signal = segment.get_vm(neuron_id)
        example[label] = {'segment_identifier': str(segment.identifier),
                          'stimulus': str(segment.annotations.get('stimulus')),
                          'voltage_neuron_id': int(neuron_id),
                          'time_ms': numeric(signal.times, 'ms').tolist(),
                          'voltage_mV': numeric(signal, 'mV').tolist(),
                          'spike_trains': [{'neuron_id': int(i), 'times_ms': numeric(segment.get_spiketrain(i), 'ms').tolist()}
                                           for i in sorted(segment.get_stored_spike_train_ids())]}
    (run / 'figure-data.json').write_text(json.dumps(example, allow_nan=False) + '\n')
    print(json.dumps({'diagnostic_populations': len(report['populations']), 'run': str(run)}))


if __name__ == '__main__':
    main()
