import unittest
import ciw
import ciw_des_in_sd
import ciw_sd_in_des
import numpy as np

class Test_SD_Component(unittest.TestCase):
    def test_init_method(self):
        SD = ciw_des_in_sd.SD_Component(
            S=200,
            I=2,
            M=0,
            D=0,
            visits_per_day=1/2,
            proportion_shop=1/3,
            infection_rate=1,
            recovery_rate=1/14,
            proportion_immune=1/4,
            rate_lose_immunity=1/75,
            fatality_rate=1/21,
        )
        self.assertEqual(SD.S, [200])
        self.assertEqual(SD.I, [2])
        self.assertEqual(SD.M, [0])
        self.assertEqual(SD.D, [0])
        self.assertEqual(SD.visits_per_day, 1/2)
        self.assertEqual(SD.proportion_shop, 1/3)
        self.assertEqual(SD.infection_rate, 1)
        self.assertEqual(SD.recovery_rate, 1/14)
        self.assertEqual(SD.proportion_immune, 1/4)
        self.assertEqual(SD.rate_lose_immunity, 1/75)
        self.assertEqual(SD.fatality_rate, 1/21)
        self.assertEqual(SD.time, np.array([0]))

    def test_differential_equations(self):
        SD = ciw_des_in_sd.SD_Component(
            S=200,
            I=2,
            M=0,
            D=0,
            visits_per_day=1/2,
            proportion_shop=1/3,
            infection_rate=1,
            recovery_rate=1/14,
            proportion_immune=1/4,
            rate_lose_immunity=1/75,
            fatality_rate=1/21,
        )

        # Everyone susceptible
        dS, dI, dM, dD = SD.differential_equations((100, 0, 0, 0), [0, 0.2, 0.4, 0.6, 0.8, 1.0], L=1)
        self.assertEqual(dS, 0)
        self.assertEqual(dI, 0)
        self.assertEqual(dM, 0)
        self.assertEqual(dD, 0)

        # Everyone infected
        dS, dI, dM, dD = SD.differential_equations((0, 100, 0, 0), [0, 0.2, 0.4, 0.6, 0.8, 1.0], L=1)
        self.assertEqual(dS, (1 - SD.proportion_immune) * SD.recovery_rate * 100)
        self.assertEqual(dI, -(SD.recovery_rate + SD.fatality_rate) * 100)
        self.assertEqual(dM, SD.proportion_immune * SD.recovery_rate * 100)
        self.assertEqual(dD, SD.fatality_rate * 100)

        # Everyone temporarily immune
        dS, dI, dM, dD = SD.differential_equations((0, 0, 100, 0), [0, 0.2, 0.4, 0.6, 0.8, 1.0], L=1)
        self.assertEqual(dS, SD.rate_lose_immunity * 100)
        self.assertEqual(dI, 0)
        self.assertEqual(dM, -SD.rate_lose_immunity * 100)
        self.assertEqual(dD, 0)

        # Some people everywhere (L=1, no conacts)
        dS, dI, dM, dD = SD.differential_equations((10, 10, 10, 10), [0, 0.2, 0.4, 0.6, 0.8, 1.0], L=1)
        self.assertEqual(dS, ((1 - SD.proportion_immune) * SD.recovery_rate * 10) + (SD.rate_lose_immunity * 10))
        self.assertEqual(dI, -(SD.recovery_rate + SD.fatality_rate) * 10)
        self.assertEqual(dM, ((SD.proportion_immune * SD.recovery_rate) * 10) - (SD.rate_lose_immunity * 10))
        self.assertEqual(dD, SD.fatality_rate * 10)

        # Some people everywhere (L=2, contacts)
        dS, dI, dM, dD = SD.differential_equations((10, 10, 10, 10), [0, 0.2, 0.4, 0.6, 0.8, 1.0], L=2)
        self.assertEqual(dS, (-(100 * SD.proportion_shop * SD.visits_per_day)/(20 + SD.proportion_shop*10)) + ((1 - SD.proportion_immune) * SD.recovery_rate * 10) + (SD.rate_lose_immunity * 10))
        self.assertEqual(dI, ((100 * SD.proportion_shop * SD.visits_per_day)/(20 + SD.proportion_shop*10)) -((SD.recovery_rate + SD.fatality_rate) * 10))
        self.assertEqual(dM, ((SD.proportion_immune * SD.recovery_rate) * 10) - (SD.rate_lose_immunity * 10))
        self.assertEqual(dD, SD.fatality_rate * 10)

    def test_solve(self):
        SD = ciw_des_in_sd.SD_Component(
            S=200,
            I=2,
            M=0,
            D=0,
            visits_per_day=1/2,
            proportion_shop=1/3,
            infection_rate=1,
            recovery_rate=1/14,
            proportion_immune=1/4,
            rate_lose_immunity=1/75,
            fatality_rate=1/21,
        )

        # Give the SD component $T^{\star}$
        SD.time_domain = np.linspace(0, 100, 201)

        # Ensure that the correct $T_i$ is added at each event
        SD.solve(3.2, L=3)
        expected_time_1 = [0.0, 0.0, 0.0, 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 3.2]
        self.assertEqual(list(SD.time), expected_time_1)
        self.assertEqual(len(SD.S), len(expected_time_1))
        self.assertEqual(len(SD.M), len(expected_time_1))
        self.assertEqual(len(SD.I), len(expected_time_1))
        self.assertEqual(len(SD.D), len(expected_time_1))

        SD.solve(7.4, L=3)
        expected_time_2 = [3.2, 3.5, 4.0, 4.5, 5.0, 5.5, 6.0, 6.5, 7.0, 7.4]
        self.assertEqual(list(SD.time), expected_time_1 + expected_time_2)
        self.assertEqual(len(SD.S), len(expected_time_1 + expected_time_2))
        self.assertEqual(len(SD.M), len(expected_time_1 + expected_time_2))
        self.assertEqual(len(SD.I), len(expected_time_1 + expected_time_2))
        self.assertEqual(len(SD.D), len(expected_time_1 + expected_time_2))

        SD.solve(10.1, L=3)
        expected_time_3 = [7.4, 7.5, 8.0, 8.5, 9.0, 9.5, 10.0, 10.1]
        self.assertEqual(list(SD.time), expected_time_1 + expected_time_2 + expected_time_3)
        self.assertEqual(len(SD.S), len(expected_time_1 + expected_time_2 + expected_time_3))
        self.assertEqual(len(SD.M), len(expected_time_1 + expected_time_2 + expected_time_3))
        self.assertEqual(len(SD.I), len(expected_time_1 + expected_time_2 + expected_time_3))
        self.assertEqual(len(SD.D), len(expected_time_1 + expected_time_2 + expected_time_3))


class Test_HybridSimulation(unittest.TestCase):
    def test_SD_setup(self):
        N = ciw.create_network(
            arrival_distributions=[ciw.dists.Exponential(40)],
            service_distributions=[ciw.dists.Exponential(48)],
            number_of_servers=[3])
        Q = ciw_des_in_sd.HybridSimulation(
            network=N,
            tracker=ciw.trackers.SystemPopulation(),
            S=200, I=2, M=0, D=0,
            visits_per_day=1/2,
            proportion_shop=1/3,
            infection_rate=1,
            recovery_rate=1/14,
            proportion_immune=1/4,
            rate_lose_immunity=1/75,
            fatality_rate=1/21
        )

        self.assertEqual(Q.SD_Component.S, [200])
        self.assertEqual(Q.SD_Component.I, [2])
        self.assertEqual(Q.SD_Component.M, [0])
        self.assertEqual(Q.SD_Component.D, [0])
        self.assertEqual(Q.SD_Component.visits_per_day, 1/2)
        self.assertEqual(Q.SD_Component.proportion_shop, 1/3)
        self.assertEqual(Q.SD_Component.infection_rate, 1)
        self.assertEqual(Q.SD_Component.recovery_rate, 1/14)
        self.assertEqual(Q.SD_Component.proportion_immune, 1/4)
        self.assertEqual(Q.SD_Component.rate_lose_immunity, 1/75)
        self.assertEqual(Q.SD_Component.fatality_rate, 1/21)
        self.assertEqual(Q.SD_Component.time, np.array([0]))

    def test_tstar(self):
        N = ciw.create_network(
            arrival_distributions=[ciw.dists.Deterministic(0.2)],
            service_distributions=[ciw.dists.Deterministic(3.34)],
            number_of_servers=[1])

        # max_simulation_time=10, and 11 time points
        Q = ciw_des_in_sd.HybridSimulation(
            network=N,
            tracker=ciw.trackers.SystemPopulation(),
            S=200, I=2, M=0, D=0,
            visits_per_day=1/2,
            proportion_shop=1/3,
            infection_rate=1,
            recovery_rate=1/14,
            proportion_immune=1/4,
            rate_lose_immunity=1/75,
            fatality_rate=1/21
        )
        Q.simulate_until_max_time(10, 11)
        self.assertEqual(list(Q.SD_Component.time_domain), [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0])
        # There should only be 'extra' time points when a customer is released (3.54 and 6.88),
        # and doubles at the endpoints (0.0, 3.54, 6.8, 10.0):
        self.assertEqual(list(Q.SD_Component.time), [0.0, 0.0, 0.0, 1.0, 2.0, 3.0, 3.54, 3.54, 4.0, 5.0, 6.0, 6.88, 6.88, 7.0, 8.0, 9.0, 10.0, 10.0])


        # max_simulation_time=15, and 6 time points
        Q = ciw_des_in_sd.HybridSimulation(
            network=N,
            tracker=ciw.trackers.SystemPopulation(),
            S=200, I=2, M=0, D=0,
            visits_per_day=1/2,
            proportion_shop=1/3,
            infection_rate=1,
            recovery_rate=1/14,
            proportion_immune=1/4,
            rate_lose_immunity=1/75,
            fatality_rate=1/21
        )
        Q.simulate_until_max_time(15, 6)
        self.assertEqual(list(Q.SD_Component.time_domain), [0.0, 3.0, 6.0, 9.0, 12.0, 15.0])
        # There should only be 'extra' time points when a customer is released (3.54, 6.88, 10.22, 13.56),
        # and doubles at the endpoints (0.0, 3.54, 6.8, 10.22, 13.56, 15.0):
        self.assertEqual([round(x, 3) for x in Q.SD_Component.time], [0.0, 0.0, 0.0, 3.0, 3.54, 3.54, 6.0, 6.88, 6.88, 9.0, 10.22, 10.22, 12.0, 13.56, 13.56, 15.0, 15.0])


        # Now test for arrivals
        ciw.seed(0)
        N = ciw.create_network(
            arrival_distributions=[ciw_des_in_sd.SolveSDArrivals()],
            service_distributions=[ciw.dists.Deterministic(float('Inf'))],
            number_of_servers=[1])

        # max_simulation_time=10, and 11 time points
        Q = ciw_des_in_sd.HybridSimulation(
            network=N,
            tracker=ciw.trackers.SystemPopulation(),
            S=200, I=2, M=0, D=0,
            visits_per_day=1/200,
            proportion_shop=1/3,
            infection_rate=1,
            recovery_rate=1/14,
            proportion_immune=1/4,
            rate_lose_immunity=1/75,
            fatality_rate=1/21
        )
        Q.simulate_until_max_time(10, 11)
        self.assertEqual(list(Q.SD_Component.time_domain), [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0])
        # There should only be 'extra' time points when a customer arrives, these should repeat, as well as 0.0 and 10.0
        self.assertEqual([round(x, 2) for x in Q.SD_Component.time], [0.0, 0.0, 0.0, 1.0, 1.85, 1.85, 2.0, 3.0, 3.27, 3.27, 3.81, 3.81, 4.0, 4.11, 4.11, 4.82, 4.82, 5.0, 5.34, 5.34, 6.0, 6.86, 6.86, 7.0, 7.22, 7.22, 7.87, 7.87, 8.0, 8.74, 8.74, 9.0, 10.0, 10.0])


class Test_Cancer_Patient(unittest.TestCase):
    def test_patient_init(self):
        for cust in range(30):
            P = ciw_sd_in_des.CancerPatient(cust)
            self.assertEqual(P.id_number, cust)
            self.assertEqual(P.customer_class, 0)
            self.assertEqual(P.priority_class, 0)
            self.assertEqual(P.degenerative_rate, 0.08)
            self.assertEqual(P.cure_rate, 0.4)
            self.assertEqual(len(P.H), 1)
            self.assertEqual(len(P.U), 1)
            self.assertTrue(P.H[0] < 91)
            self.assertTrue(P.H[0] > 39)
            self.assertEqual(P.H[0], 100 - P.U[0])
            self.assertEqual(P.U[0], 100 - P.H[0])

    def test_cancer_differential_equations(self):
        P = ciw_sd_in_des.CancerPatient(0)

        # No unhealthy cells, no treatment
        dH, dU = P.differential_equations(y=(100, 0), time_domain=[0, 1, 2, 3], in_treatment=False)
        self.assertEqual(dH, 0)
        self.assertEqual(dU, 0)
        # No unhealthy cells, treatment
        dH, dU = P.differential_equations(y=(100, 0), time_domain=[0, 1, 2, 3], in_treatment=True)
        self.assertEqual(dH, 0)
        self.assertEqual(dU, 0)
        # No healthy cells, no treatment
        dH, dU = P.differential_equations(y=(0, 100), time_domain=[0, 1, 2, 3], in_treatment=False)
        self.assertEqual(dH, 0)
        self.assertEqual(dU, 0)
        # No healthy cells, treatment
        dH, dU = P.differential_equations(y=(0, 100), time_domain=[0, 1, 2, 3], in_treatment=True)
        self.assertEqual(dH, 0)
        self.assertEqual(dU, 0)
        # Half healthy half unhealthy, no treatment
        dH, dU = P.differential_equations(y=(50, 50), time_domain=[0, 1, 2, 3], in_treatment=False)
        self.assertEqual(dH, -25*P.degenerative_rate)
        self.assertEqual(dU, 25*P.degenerative_rate)
        # Half healthy half unhealthy, treatment
        dH, dU = P.differential_equations(y=(50, 50), time_domain=[0, 1, 2, 3], in_treatment=True)
        self.assertEqual(dH, 25*(P.cure_rate - P.degenerative_rate))
        self.assertEqual(dU, 25*(P.degenerative_rate - P.cure_rate))

    def test_cancer_solve(self):
        N = ciw.create_network(
            arrival_distributions=[ciw.dists.Exponential(4)],
            service_distributions=[ciw.dists.Exponential(5)],
            number_of_servers=[3]
        )
        Q = ciw.Simulation(N)
        Q.time_domain = np.linspace(0, 100, 201)
        
        P = ciw_sd_in_des.CancerPatient(0)
        P.simulation = Q
        P.time = np.array([37.2])

        expected_time_1 = [37.2]
        self.assertEqual(list(P.time), expected_time_1)
        self.assertEqual(len(P.H), len(expected_time_1))
        self.assertEqual(len(P.U), len(expected_time_1))

        P.solve(41.1, in_treatment=False)
        expected_time_2 = [37.2, 37.5, 38, 38.5, 39, 39.5, 40, 40.5, 41, 41.1]
        self.assertEqual(list(P.time), expected_time_1 + expected_time_2)
        self.assertEqual(len(P.H), len(expected_time_1 + expected_time_2))
        self.assertEqual(len(P.U), len(expected_time_1 + expected_time_2))

        P.solve(43.3, in_treatment=True)
        expected_time_3 = [41.1, 41.5, 42, 42.5, 43, 43.3]
        self.assertEqual(list(P.time), expected_time_1 + expected_time_2 + expected_time_3)
        self.assertEqual(len(P.H), len(expected_time_1 + expected_time_2 + expected_time_3))
        self.assertEqual(len(P.U), len(expected_time_1 + expected_time_2 + expected_time_3))


class Test_Cancer_Simulation(unittest.TestCase):
    def test_sd_triggering_correctly(self):
        N = ciw.create_network(
            arrival_distributions=[ciw.dists.Exponential(4)],
            service_distributions=[ciw_sd_in_des.Treatment()],
            number_of_servers=[1]
        )
        ciw.seed(0)
        Q = ciw_sd_in_des.HybridSimulation(network=N)
        Q.simulate_until_max_time(500, 5000, progress_bar=False)

        for ind in Q.nodes[-1].all_individuals:
            arrival = ind.data_records[0].arrival_date
            service_start = ind.data_records[0].service_start_date
            service_end = ind.data_records[0].service_end_date
            self.assertIn(arrival, ind.time)
            self.assertIn(service_start, ind.time)
            self.assertIn(service_end, ind.time)
            if ind.status == 'Treated':
                self.assertTrue(ind.H[-1] > 98)
                self.assertTrue(ind.H[-1] < 100)
            if ind.status == 'Untreatable':
                self.assertTrue(ind.H[-1] < 10)
                self.assertTrue(ind.H[-1] > 0)

