"""Analytic synthetic edge cases. These are not neural simulation results."""
import math
import unittest
from metrics import spike_metrics, analog_metrics, printed_precision_comparison


def row(i, times):
    return {'id': i, 'times_ms': times}


class MetricTests(unittest.TestCase):
    def test_silent_neuron_remains_in_rate(self):
        r = spike_metrics([row(1, []), row(2, [1, 11, 21, 31, 41, 51, 61])], 0, 100)
        self.assertEqual(r['firing_rate_hz'], {'mean': 35., 'sd': 35., 'n': 2})

    def test_six_excluded_seven_included(self):
        r = spike_metrics([row(1, list(range(6))), row(2, list(range(7)))], 0, 100)
        self.assertEqual(r['eligible_ids'], [2])
        self.assertEqual(r['isi_cv']['mean'], 0.)
        self.assertIsNone(r['correlation_10ms']['mean'])

    def test_known_cv_population_sd(self):
        # Intervals 1,2,3,4,5,6 have mean 3.5 and population variance 35/12.
        r = spike_metrics([row(1, [0, 1, 3, 6, 10, 15, 21])], 0, 100)
        self.assertAlmostEqual(r['isi_cv']['mean'], math.sqrt(35/12)/3.5)

    def test_pair_excludes_self_and_duplicate_direction(self):
        r = spike_metrics([row(1, list(range(1, 8))), row(2, list(range(21, 28)))], 0, 40)
        self.assertEqual(r['correlation_10ms']['n'], 1)
        self.assertAlmostEqual(r['correlation_10ms']['mean'], -1/3)

    def test_undefined_correlation_is_counted_and_zeroed(self):
        r = spike_metrics([row(1, list(range(1, 80, 10))), row(2, list(range(2, 81, 10)))], 0, 80)
        self.assertEqual(r['undefined_pair_values_converted_to_zero'], 1)
        self.assertEqual(r['correlation_10ms']['mean'], 0.)

    def test_window_endpoint_and_bin_count(self):
        r = spike_metrics([row(1, [0, 40320])], 0, 40320)
        self.assertEqual(r['bin_count'], 4032)
        self.assertAlmostEqual(r['firing_rate_hz']['mean'], 2/40.32)

    def test_invalid_spikes_rejected(self):
        for times in ([2, 1], [1, 1], [-1], [101], [float('nan')]):
            with self.subTest(times=times), self.assertRaises(ValueError):
                spike_metrics([row(1, times)], 0, 100)

    def test_invalid_metadata_rejected(self):
        for start, stop, width in ((0, 0, 10), (0, 101, 10), (0, 100, 0)):
            with self.subTest(stop=stop), self.assertRaises(ValueError):
                spike_metrics([], start, stop, width)
        with self.assertRaises(ValueError):
            spike_metrics([row(1, []), row(1, [])], 0, 100)

    def test_analog_mean_is_equal_weight_per_neuron(self):
        r = analog_metrics([{'id': 1, 'values': [-70, -60]}, {'id': 2, 'values': [-75, -75]}], 'mV')
        self.assertEqual(r, {'mean': -70., 'sd': 5., 'n': 2, 'unit': 'mV'})

    def test_conductance_units_and_invalid_values(self):
        self.assertEqual(analog_metrics([{'id': 1, 'values': [.001, .003]}], 'uS')['mean'], 2.)
        for values in ([], [-.001], [float('inf')]):
            with self.subTest(values=values), self.assertRaises(ValueError):
                analog_metrics([{'id': 1, 'values': values}], 'uS')

    def test_rounding_does_not_certify_a_pass(self):
        r = printed_precision_comparison(1.41, 1.4)
        self.assertTrue(r['same_two_significant_digit_display'])
        self.assertIsNone(r['benchmark_pass'])
        self.assertFalse(printed_precision_comparison(1.46, 1.4)['same_two_significant_digit_display'])


if __name__ == '__main__':
    unittest.main(verbosity=2)
