
import unittest
import math
from basic_kinematics import *

class EpsTest(unittest.TestCase):

    # Overwrite "assertEqual" so that we can use epsilon for
    # comparison of quantities calculated with round-off errors
    def assertEqual(self, first, second, msg=None, eps=None):
        if eps is None:
            unittest.TestCase.failUnlessEqual(self, first, second, msg)
        else:
            if abs(first - second) > eps:
                raise self.failureException, \
                      (msg or '%s != %s' % (`first`, `second`))

    assertEquals = failUnlessEqual = assertEqual


class V3Test(EpsTest):

    def test_norm2(self):
        self.assertEqual(V3(1, -2, 3).norm2(), 14.0)

    def test_norm(self):
        v = V3(3, 4, -5)
        self.assertEqual(v.norm(), 7.071, eps=0.001)
        self.assertEqual(v.norm(), abs(v))

    def test_eq(self):
        self.assertEqual(V3(3, 4, -5), V3(3, 4, -5))

    def test_neq(self):
        v1 = V3(3, 4, -5)
        self.assertNotEqual(V3(4, 4, -5), v1)
        self.assertNotEqual(V3(3, 5, -5), v1)
        self.assertNotEqual(V3(3, 4, 5),  v1)

    def test_subtraction(self):
        self.assertEqual(V3(3, 4, -5) - V3(6, 2, 7), V3(-3, 2, -12))

    def test_addition(self):
        self.assertEqual(V3(3, 4, -5) + V3(6, 2, 7), V3(9, 6, 2))

    def test_mul(self):
        self.assertEqual(V3(3, 4, -5)*3, V3(9, 12, -15))
        self.assertEqual(3*V3(3, 4, -5), V3(9, 12, -15))
        self.assertEqual(V3(3, 4, -5)*-3, V3(-9, -12, 15))

    def test_div(self):
        self.assertEqual(V3(9, 12, -15)/3, V3(3, 4, -5))
        self.assertRaises(ZeroDivisionError, V3(9, 12, -15).__div__, 0.0)
        self.assertRaises(ZeroDivisionError, V3(9, 12, -15).__div__, 0)

    def test_direction(self):
        v = V3(3, 4, -5)
        dir = v.direction()
        self.assertEqual(abs(dir), 1.0, eps=1.0e-15)
        self.assertEqual(dir, v/abs(v), eps=1.0e-15)

    def test_eta(self):
        v = V3(3, 4, -5)
        theta = v.theta()
        eta = v.eta()
        self.assertEqual(eta, -math.log(math.tan(theta/2.0)), eps=1.0e-12)

    def test_phi(self):
        v = V3(3, 4, -5)
        self.assertEqual(v.phi(), math.atan2(v.y, v.x))

    def test_accum_plus(self):
        v = V3()
        v += V3(3, 4, -5)
        v += V3(2, -1, 7)
        self.assertEqual(v, (V3(5, 3, 2)))

    def test_accum_minus(self):
        v = V3()
        v -= V3(3, 4, -5)
        v -= V3(2, -1, 7)
        self.assertEqual(v, (V3(-5, -3, -2)))

    def test_accum_mul(self):
        v = V3(3, 4, -5)
        v *= 3
        self.assertEqual(v, (V3(9, 12, -15)))

    def test_accum_div(self):
        v = V3(9, 12, -15)
        v /= 3
        self.assertEqual(v, V3(3, 4, -5))

    def test_str(self):
        self.assertEqual(str(V3(1, 2, 3)), "{1.0, 2.0, 3.0}")

    def test_bool(self):
        self.assertEqual(bool(V3(1, 2, 3)), True)
        self.assertEqual(bool(V3(0, 0, 0)), False)
        self.assertEqual(bool(V3()), False)

    def test_unary(self):
        v = V3(1, 2, 3)
        self.assertEqual(v, +v)
        self.assertEqual(v * -1.0, -v)


# Function for comparing 4-momenta
def sim4(p1, p2, eps):
    return abs(p1.e - p2.e) <= eps and abs(p1.p - p2.p) <= eps


class FourMomentumTest(EpsTest):

    def test_time_like(self):
        self.assertRaises(ValueError, FourMomentum, V3(3, 4, -5), -10)
        self.assertRaises(ValueError, FourMomentum, 1, V3(3, 4, -5))

    def test_eq(self):
        self.assertEqual(FourMomentum(V3(3, 4, -5), 1), \
                         FourMomentum(V3(3, 4, -5), 1))
        self.assertEqual(FourMomentum(10, V3(3, 4, -5)), \
                         FourMomentum(10, V3(3, 4, -5)))
        self.assertEqual(FourMomentum(V3(0, 0, 0), 0, 1), \
                         FourMomentum(V3(0, 0, 0), 0, -1))
        self.assertNotEqual(FourMomentum(V3(3, 4, -5), 1, -1), \
                            FourMomentum(V3(3, 4, -5), 1, 1))

    def test_subtraction(self):
        self.assertEqual(sim4(FourMomentum(10, V3(3, 4, -5)) - \
                              FourMomentum(5,  V3(2, 1, -3)), \
                              FourMomentum(5,  V3(1, 3, -2)), 1.0e-14), 1)

    def test_addition(self):
        self.assertEqual(sim4(FourMomentum(10, V3(3, 4, -5)) + \
                              FourMomentum(5,  V3(2, 1, -3)), \
                              FourMomentum(15, V3(5, 5, -8)), 1.0e-14), 1)

    def test_mul(self):
        self.assertEqual(FourMomentum(10, V3(3, 4, -5))*3, \
                         FourMomentum(30, V3(9, 12, -15)))
        self.assertEqual(3 * FourMomentum(10, V3(3, 4, -5)), \
                         FourMomentum(30, V3(9, 12, -15)))
        self.assertEqual(FourMomentum(10, V3(3, 4, -5))*-3, \
                         FourMomentum(-30, V3(-9, -12, 15)))

    def test_div(self):
        self.assertEqual(FourMomentum(30, V3(9, 12, -15))/3,
                         FourMomentum(10, V3(3, 4, -5)))
        self.assertEqual(FourMomentum(30, V3(9, 12, -15))/-3,
                         FourMomentum(-10, V3(-3, -4, 5)))
        self.assertRaises(ZeroDivisionError, \
                          FourMomentum(30, V3(9, 12, -15)).__div__, 0.0)

    def test_accum_plus(self):
        v = FourMomentum()
        v += FourMomentum(10, V3(3, 4, -5))
        v += FourMomentum(8,  V3(2.5, 0.5, -1))
        self.assertEqual(v, FourMomentum(18, V3(5.5, 4.5, -6)))

    def test_accum_minus(self):
        v = FourMomentum()
        v -= FourMomentum(10, V3(3, 4, -5))
        v -= FourMomentum(8,  V3(2.5, 0.5, -1))
        self.assertEqual(v, FourMomentum(-18, V3(-5.5, -4.5, 6)))

    def test_accum_mul(self):
        v = FourMomentum(10, V3(3, 4, -5))
        v *= 3
        self.assertEqual(v, FourMomentum(30, V3(9, 12, -15)))

    def test_accum_div(self):
        v = FourMomentum(30, V3(9, 12, -15))
        v /= 3
        self.assertEqual(v, FourMomentum(10, V3(3, 4, -5)))

    def test_str(self):
        self.assertEqual(str(FourMomentum(V3(1, 2, 3), 4)), \
                         "{+, {1.0, 2.0, 3.0}, 4.0}")

    def test_bool(self):
        self.assertEqual(bool(FourMomentum(1, V3())), True)
        self.assertEqual(bool(FourMomentum(0, V3())), False)
        self.assertEqual(bool(FourMomentum()), False)

    def test_unary(self):
        v = FourMomentum(10, V3(3, 4, -5))
        self.assertEqual(v, +v)
        self.assertEqual(v * -1.0, -v)


# Coordinate axes for rotation tests
xaxis  = V3(1, 0, 0)
yaxis  = V3(0, 1, 0)
zaxis  = V3(0, 0, 1)
pxaxis = FourMomentum(xaxis, 1.65)
pyaxis = FourMomentum(yaxis, 1.65)
pzaxis = FourMomentum(zaxis, 1.65)


class RotationTest(EpsTest):

    def test_x_90(self):
        scale = 1.2345
        r = Rotation(xaxis, math.pi/2.0)
        self.assertEqual(r(xaxis*scale), xaxis*scale)
        self.assertEqual(r(yaxis*scale), zaxis*scale, eps=1.0e-15)
        self.assertEqual(r(zaxis*scale), -yaxis*scale, eps=1.0e-15)
        self.assertEqual(sim4(r(pxaxis*scale), pxaxis*scale, 0.0), 1)
        self.assertEqual(sim4(r(pyaxis*scale), pzaxis*scale, 1.0e-15), 1)
        self.assertEqual(sim4(r(pzaxis*scale), \
                              FourMomentum(-pyaxis.p*scale, pyaxis.m*scale), \
                              1.0e-15), 1)

    def test_y_90(self):
        scale = 1.3345
        r = Rotation(yaxis, math.pi/2.0)
        self.assertEqual(r(xaxis*scale), -zaxis*scale, eps=1.0e-15)
        self.assertEqual(r(yaxis*scale), yaxis*scale)
        self.assertEqual(r(zaxis*scale), xaxis*scale, eps=1.0e-15)
        self.assertEqual(sim4(r(pxaxis*scale), \
                              FourMomentum(-pzaxis.p*scale, pzaxis.m*scale), \
                              1.0e-15), 1)
        self.assertEqual(sim4(r(pyaxis*scale), pyaxis*scale, 0.0), 1)
        self.assertEqual(sim4(r(pzaxis*scale), pxaxis*scale, 1.0e-15), 1)

    def test_z_90(self):
        scale = 1.9657
        r = Rotation(zaxis, math.pi/2.0)
        self.assertEqual(r(xaxis*scale), yaxis*scale, eps=1.0e-15)
        self.assertEqual(r(yaxis*scale), -xaxis*scale, eps=1.0e-15)
        self.assertEqual(r(zaxis*scale), zaxis*scale)
        self.assertEqual(sim4(r(pxaxis*scale), pyaxis*scale, 1.0e-15), 1)
        self.assertEqual(sim4(r(pyaxis*scale), \
                              FourMomentum(-pxaxis.p*scale, pxaxis.m*scale), \
                              1.0e-15), 1)
        self.assertEqual(sim4(r(pzaxis*scale), pzaxis*scale, 0.0), 1)

    def test_diag_120(self):
        scale = 2.34
        r = Rotation(xaxis+yaxis+zaxis, math.pi*2.0/3.0)
        self.assertEqual(r(xaxis*scale), yaxis*scale, eps=1.0e-15)
        self.assertEqual(r(yaxis*scale), zaxis*scale, eps=1.0e-15)
        self.assertEqual(r(zaxis*scale), xaxis*scale, eps=1.0e-15)
        self.assertEqual(sim4(r(pxaxis*scale), pyaxis*scale, 1.0e-15), 1)
        self.assertEqual(sim4(r(pyaxis*scale), pzaxis*scale, 1.0e-15), 1)
        self.assertEqual(sim4(r(pzaxis*scale), pxaxis*scale, 1.0e-15), 1)

    def test_eq(self):
        v  = V3(-2, 3, 1)
        self.assertEqual(Rotation(v, 1), Rotation(v, 1))
        self.assertEqual(Rotation(v, 1), Rotation(-v, -1))        
        self.assertEqual(Rotation(v, 1).quat, \
                         -Rotation(v, 1+2.0*math.pi).quat, eps=1.0e-15)

    def test_inverse(self):
        v  = V3(-2, 3, 1)
        r1 = Rotation(V3(1, 2, 3), 1)
        r2 = Rotation(V3(-6, 7, 8), 0.4)
        self.assertEqual(v, (~r1)(r1(v)), eps=1.0e-15)
        self.assertEqual(v, (~r2)(r2(v)), eps=1.0e-15)

    def test_mul(self):
        v  = V3(-2, 3, 1)
        r1 = Rotation(V3(1, 2, 3), 1)
        r2 = Rotation(V3(-6, 7, 8), 0.4)
        r3 = r1 * r2
        self.assertEqual(r1(r2(v)), r3(v), eps=1.0e-15)

    def _test_ra_dec(self, axis, theta):
        r1 = Rotation(axis, theta)
        xrot = r1(xaxis)
        zrot = r1(zaxis)
        r2 = Rotation(xrot.ra(), xrot.dec(), zrot.ra(), zrot.dec())
        self.assertEqual(min(abs(r1.quat-r2.quat),abs(r1.quat+r2.quat)), \
                         0.0, eps=1.0e-12)

    def test_ra_dec(self):
        self._test_ra_dec(V3(0.1, 2.2, 1.4), 0.1)
        self._test_ra_dec(V3(3.2, 0.5, 4.4), 0.56)
        self._test_ra_dec(V3(0.1, 2.2, 1.4), 0.0)
        self._test_ra_dec(V3(0.1, -2.2, 1.4), 4.5)
        self._test_ra_dec(V3(2.3, 4.5, -7.1), math.pi)
        self._test_ra_dec(V3(-7.1, 2.3, 4.5), math.pi)
        self._test_ra_dec(V3(4.5, -7.1, 2.3), math.pi)
        self._test_ra_dec(V3(1, 0, 0), math.pi)
        self._test_ra_dec(V3(0, 1, 0), math.pi)
        self._test_ra_dec(V3(0, 0, 1), math.pi)
        self._test_ra_dec(V3(1, 0, 0), math.pi/2)
        self._test_ra_dec(V3(0, 1, 0), math.pi/2)
        self._test_ra_dec(V3(0, 0, 1), math.pi/2)
        self._test_ra_dec(V3(1, 0, 0), math.pi/4)
        self._test_ra_dec(V3(0, 1, 0), math.pi/4)
        self._test_ra_dec(V3(0, 0, 1), math.pi/4)
        self._test_ra_dec(V3(1, 0, 0), 0)
        self._test_ra_dec(V3(0, 1, 0), 0)
        self._test_ra_dec(V3(0, 0, 1), 0)



class LorentzBoostTest(EpsTest):

    def test_consistency(self):
        p = FourMomentum(10, V3(3, 4, -5))
        b = LorentzBoost(p)
        boosted = b(FourMomentum(V3(), p.m))
        self.assertEqual(boosted.p * -1.0, p.p, eps=1.0e-14)
        self.assertEqual(boosted.m, p.m)
        self.assertEqual(boosted.esign, p.esign)

    def test_inverse(self):
        t = FourMomentum(7, V3(-1, 2, 3))
        p = FourMomentum(10, V3(3, 4, -5))
        b = LorentzBoost(p)
        self.assertEqual(sim4(t, (~b)(b(t)), 1.0e-14), 1)

    def test_magnitude_x(self):
        p1 = FourMomentum(5, V3(4, 0, 0))
        p2 = FourMomentum(6, V3(1, 2, 3))
        b  = LorentzBoost(p1)
        b1 = b(p2)
        newe = p1.gamma() * (p2.e  - p2.p.x * p1.beta())
        newp = p1.gamma() * (p2.p.x - p2.e * p1.beta())
        self.assertEqual(b1.p.x, newp, eps=1.0e-14)
        self.assertEqual(b1.e, newe, eps=1.0e-14)
        self.assertEqual(b1.p.y, p2.p.y)
        self.assertEqual(b1.p.z, p2.p.z)

    def test_magnitude_y(self):
        p1 = FourMomentum(3, V3(0, 0.5, 0))
        p2 = FourMomentum(6, V3(1, 2, 3))
        b  = LorentzBoost(p1)
        b1 = b(p2)
        newe = p1.gamma() * (p2.e  - p2.p.y * p1.beta())
        newp = p1.gamma() * (p2.p.y - p2.e * p1.beta())
        self.assertEqual(b1.p.y, newp, eps=1.0e-14)
        self.assertEqual(b1.e, newe, eps=1.0e-14)
        self.assertEqual(b1.p.x, p2.p.x)
        self.assertEqual(b1.p.z, p2.p.z)

    def test_magnitude_z(self):
        p1 = FourMomentum(3, V3(0, 0, 1.5))
        p2 = FourMomentum(6, V3(1, 2, 3))
        b  = LorentzBoost(p1)
        b1 = b(p2)
        newe = p1.gamma() * (p2.e  - p2.p.z * p1.beta())
        newp = p1.gamma() * (p2.p.z - p2.e * p1.beta())
        self.assertEqual(b1.p.z, newp, eps=1.0e-14)
        self.assertEqual(b1.e, newe, eps=1.0e-14)
        self.assertEqual(b1.p.x, p2.p.x)
        self.assertEqual(b1.p.y, p2.p.y)

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