import unittest,itertools
from dataclasses import replace
from fractions import Fraction as F
import z3
from activity import *
from solver import *
from elimination import *
from example import unit_problem


class ScientificTests(unittest.TestCase):
    def test_forced_irrational_jump_leastness_and_rounding(self):
        p=unit_problem();a=CycleAccelerator(p,F(1,100));r=a.solve()
        self.assertEqual(r['jumps'],1);self.assertTrue(a.audit_least(r['state']))
        self.assertTrue(truth(r['state']['u']==exact((5-z3.Sqrt(13))/10)))
        rounded=round_down(p,r['state'],F(1,100));self.assertTrue(p.check_rational(rounded,F(1,200)))
        self.assertFalse(p.check_rational(rounded,F(1,100)))
        for label in r['labels'].values():self.assertEqual(len(label['path']),len(set(label['path'])))

    def test_policy_omission_and_singleton_box(self):
        p=ActivityProblem({'x':ActivityBox(F(1,8),F(7,8)),'y':ActivityBox(F(1,16),F(7,8)),'z':ActivityBox(F(41,64),F(41,64))},(Core('x','y'),Core('x','z')))
        a=CycleAccelerator(p,F(1,64));r=a.solve();self.assertTrue(a.audit_least(r['state']))
        self.assertTrue(truth(r['state']['x']==q(F(3,4))))
        self.assertFalse(p.check_rational({'x':F(1,4),'y':F(9,64),'z':F(41,64)},F(1,64)))
        self.assertEqual(round_down(p,r['state'],F(1,64))['z'],F(41,64))

    def test_strict_zero_margin_and_aggregate_production_are_distinct(self):
        boxes={v:ActivityBox(F(1,100),F(9,10)) for v in 'ABCD'}
        p=ActivityProblem(boxes,(Core('A','B'),Core('B','C'),Core('A','D',F(2)),Core('D','C',F(2))))
        s,_=full_system(p,F(0));self.assertEqual(status(s),z3.sat)
        self.assertEqual(strict_decision(p)['status'],'infeasible')
        self.assertEqual(CycleAccelerator(p,F(1,10000)).solve()['status'],'infeasible')
        tri=windmill(1,F(1));state={'A':F(1,10),'B0':F(13,200),'C0':F(43,1000)}
        self.assertTrue(all(v>0 for v in tri.total_accounts(state)['species_production'].values()))
        self.assertEqual(strict_decision(tri)['status'],'infeasible')

    def test_exact_grid_elimination_against_exhaustive_lattice(self):
        for p in [unit_problem(),windmill(1),windmill(1,F(1))]:
            N=16;t=F(1,1000);result=GridElimination(p,t,N).solve();names=list(p.boxes)
            states=[dict(zip(names,values)) for values in itertools.product(*[[F(i,N) for i in range(N+1) if b.lower<=F(i,N)<=b.upper] for b in p.boxes.values()])]
            valid=[s for s in states if p.check_rational(s,t)]
            if not valid:self.assertEqual(result['status'],'grid_infeasible')
            else:
                self.assertEqual(result['status'],'grid_feasible')
                self.assertEqual(result['state'],{v:min(s[v] for s in valid) for v in names})
        nxt=ClosedUnionNext(((F(1,10),F(1,5)),(F(3,5),F(4,5))))
        self.assertEqual(nxt.value(F(3,10)),F(3,5));self.assertIsNone(nxt.value(F(9,10)))

    def test_robust_quantifiers_exact_tolerance_and_corner_minima(self):
        p=windmill(1);state={'A':F(1,10),'B0':F(13,200),'C0':F(43,1000)};radius=[];mins=[]
        for e in p.cores:
            x,y=state[e.tail],state[e.head];r=e.currents(x,y)
            radius += [r['tail_production']/(2*r['q']+r['p']),r['head_production']/(r['p']+r['q'])]
            lo,hi=e.uncertain_factor_activity_minima(x,y,F(1,20),F(1,10000));mins.extend([lo,hi])
            corner=[]
            for xx,yy,a,b in itertools.product([x-F(1,10000),x+F(1,10000)],[y-F(1,10000),y+F(1,10000)],[e.a*F(19,20),e.a*F(21,20)],[e.b*F(19,20),e.b*F(21,20)]):
                r=Core(e.tail,e.head,a,b).currents(xx,yy);corner.append(r)
            self.assertEqual(lo,min(r['tail_production'] for r in corner));self.assertEqual(hi,min(r['head_production'] for r in corner))
        self.assertEqual(min(radius),F(19,301));self.assertEqual(min(mins),F(1175221,2000000000))
        unit=unit_problem();edge=replace(unit.cores[0],ratio_lower=F(1),ratio_upper=F(2))
        self.assertEqual(strict_decision(ActivityProblem(unit.boxes,(edge,)))['status'],'infeasible')

    def test_capacity_weights_and_min_closed_feasibility(self):
        p=unit_problem();r=capacity_bracket(p,10);self.assertLessEqual(r['lower'],F(1,48));self.assertGreater(r['upper'],F(1,48))
        weighted=ActivityProblem(p.boxes,(replace(p.cores[0],lower_weight=F(1,1000),upper_weight=F(1,1000)),))
        r=capacity_bracket(weighted,10);self.assertLessEqual(r['lower'],F(1000,48));self.assertGreater(r['upper'],F(1000,48))
        a={'u':F(1,2),'v':F(17,48)};b={'u':F(3,4),'v':F(41,64)};m={v:min(a[v],b[v]) for v in a}
        self.assertTrue(all(p.check_rational(s,F(1,64)) for s in [a,b,m]))
        self.assertEqual(capacity_bracket(ActivityProblem({},()))['status'],'unbounded')

    def test_unknown_budgets_invalid_graph_and_source_accounts(self):
        r=CycleAccelerator(unit_problem(),F(1,100),max_jumps=0).solve();self.assertEqual(r['status'],'unknown')
        for fn in [lambda:ActivityBox(0,1),lambda:Core('a','a'),lambda:Core('a','b',lower_weight=0),lambda:ActivityProblem({'a':ActivityBox(F(1,10),1)},(Core('a','b'),)),lambda:ClosedUnionNext(((F(1,2),F(1,3)),))]:
            with self.assertRaises(ValueError):fn()
        p=windmill(3);s={v:F(1,10) if v=='A' else F(13,200) if v.startswith('B') else F(43,1000) for v in p.boxes};r=p.total_accounts(s)
        self.assertEqual(r['food'],3*F(5071,40000));self.assertEqual(r['total_internal_production'],r['food'])


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