import unittest
from fractions import Fraction as F
import numpy as np
from memory import *
from frontier_proportion import frontier


class ScientificChecks(unittest.TestCase):
    def test_generator_and_object_conservation(self):
        r=SupportReactor(8);states,G=r.generator()
        self.assertLess(np.max(np.abs(G.sum(axis=1))),1e-12)
        for z in states:
            for dest,a,channel in r.transitions(z):
                self.assertGreater(a,0);self.assertGreaterEqual(sum(dest[:2]),1)
                self.assertEqual(dest[0]+2*dest[1]+dest[2]+r.food(dest),8)
        transitions=list(r.transitions((2,0,0)))
        self.assertEqual(next(a for _,a,i in transitions if i==1),F(1,50))
        self.assertEqual(len(kernel.states_of(32)),3384)
        self.assertEqual(sum(p==0 and n+c<=31 for n,c,p in kernel.states_of(32)),287)

    def test_joint_payoff_not_independent_daughters(self):
        a=ComplementaryAllocation()
        self.assertEqual(a.support_return(1),0)
        self.assertEqual(a.support_return(2),F(1,2))
        self.assertNotEqual(a.support_return(2),(1-F(1,4))**2)
        states,v=SupportReactor(6).deadline(.3,1)
        self.assertLess(np.max(np.abs(v[:,:3].sum(axis=1)-1)),1e-12)
        for i,z in enumerate(states):self.assertGreaterEqual(v[i,0],-1e-14)

    def test_exact_envelope_encloses_full_backward_equation(self):
        groups=tuple(kernel.nominal_groups(2,3))
        cert=IntegerEnvelope(6,groups,1,1).run()
        for k in [F(2),F(5,2),F(3)]:
            r=SupportReactor(6,SupportRates(dissociation=k,release=k));states,v=r.deadline(1,1)
            value=min(v[i,0] for i,(n,c,p) in enumerate(states) if p==0 and n+c<=5)
            self.assertLessEqual(cert['lower']/kernel.SCALE,value)
            self.assertGreaterEqual(cert['upper']/kernel.SCALE,value)
        with self.assertRaises(ValueError):IntegerEnvelope(6,groups+groups,1,1).run()

    def test_protocol_charges_both_daughters_even_after_failure(self):
        rows=SerialProtocol(SupportReactor(8),.1,4).lineage((0,1,0),5,np.random.default_rng(4))
        self.assertTrue(any(not r['joint_success'] for r in rows))
        for row in rows:self.assertEqual(row['refill_a']+row['refill_b'],8+row['terminal'][2])
        self.assertEqual(rows[-1]['total_supplied'],8+5*8+rows[-1]['total_harvested'])
        for a,b in zip(rows,rows[1:]):self.assertEqual(b['start'],a['daughter_a'][:2]+[0])

    def test_corridor_endpoint_reduction_including_bias(self):
        for l in range(1,5):
            for u in range(l,9):
                for g in range(l,u+1):
                    c=Corridor(l,u,g)
                    for theta in [F(1,2),F(2,5)]:
                        direct=min(psi_theta(n,l,u,theta) for n in range(l+g,u+g+1))
                        self.assertEqual(c.uniform_return(theta),direct)
        with self.assertRaises(ValueError):Corridor(4,8,3)

    def test_gate_rates_and_fresh_exact_class_minima(self):
        gate=CooperativeGate(Corridor(65,195,121),Corridor(8,64,30))
        self.assertEqual((gate.core,gate.molecularity),(414,229))
        self.assertGreater(gate.joint_lower(),F(9997,10000))
        self.assertGreater(gate.joint_lower(F(49,100)),F(9996,10000))
        for x in [65,100,195]:
            for y in [8,40,64]:
                a,d=gate.pair_rates(x,y);self.assertGreaterEqual(a,20);self.assertLessEqual(d,F(1,100000))
                z=(x,y,gate.core-x-y,1,0,0,0);forward=list(gate.transitions(z))
                self.assertEqual(len(forward),1);self.assertEqual(forward[0][1],a)
                reverse=list(gate.transitions(forward[0][0]));self.assertEqual(len(reverse),1)
                self.assertEqual(reverse[0][:2],(z,d))
                mirror=(y,x,gate.core-x-y,1,0,0,0)
                self.assertEqual(list(gate.transitions(mirror))[0][2],'mirror_forward')
        for target,floor,expected in [(F(999,1000),1,230),(F(1999,2000),1,253),(F(1999,2000),8,375)]:
            result=frontier(target,ly_min=floor,cap=800)
            self.assertEqual(result['N'],expected);self.assertGreater(F(result['joint']),target)
        with self.assertRaises(RuntimeError):frontier(F(999,1000),cap=10)

    def test_necessary_carrier_bound_and_fixed_seed_clock(self):
        self.assertEqual(carrier_floor(4,1,F(1,1000)),15)
        self.assertEqual(carrier_floor(4,2,F(1,1000)),26)
        for K in [4,15,32,100]:
            r=SupportReactor(K,SupportRates(dissociation=F(2),release=F(3)))
            self.assertEqual(sum(a for _,a,_ in r.transitions((0,1,0))),5)
        self.assertEqual(F(9997,10000)**10>F(997,1000),True)


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