import unittest
import numpy as np
from scipy.optimize import check_grad
from scipy.special import expit
from example import FIXTURE, fit, objective, validate

class RaschTests(unittest.TestCase):
    def test_order_and_identification(self):
        theta, b = fit(FIXTURE)
        self.assertTrue(np.all(np.diff(theta) < 0))
        self.assertTrue(np.all(np.diff(b) > 0))
        self.assertAlmostEqual(float(b.mean()), 0., places=12)
    def test_extreme_patterns_are_finite_because_of_penalty(self):
        theta, b = fit(FIXTURE)
        self.assertTrue(np.isfinite(theta).all())
        stronger, _ = fit(FIXTURE, penalty=10.)
        self.assertLess(np.linalg.norm(stronger), np.linalg.norm(theta))
    def test_location_invariance(self):
        t,b = fit(FIXTURE)
        np.testing.assert_allclose(expit(t[:,None]-b), expit((t+7)[:,None]-(b+7)))
    def test_gradient_including_last_item_constraint(self):
        y, mask = validate(FIXTURE)
        x = np.linspace(-.4,.6,sum(y.shape)-1)
        error = check_grad(lambda z: objective(z,y,mask,1.)[0],
                           lambda z: objective(z,y,mask,1.)[1], x)
        self.assertLess(error, 1e-5)
    def test_missing_is_not_incorrect(self):
        y=FIXTURE.astype(float); y[0,0]=np.nan
        self.assertTrue(np.isfinite(fit(y)[0]).all())
    def test_invalid_inputs(self):
        for y in ([[0,2],[1,0]], [[0,np.inf],[1,0]], [[np.nan,np.nan],[1,0]], [1,0]):
            with self.assertRaises(ValueError): fit(y)
        for p in (0,-1,np.inf,np.nan):
            with self.assertRaises(ValueError): fit(FIXTURE, p)
    def test_nonconvergence_is_not_result(self):
        with self.assertRaises(RuntimeError): fit(FIXTURE,maxiter=0)
if __name__ == '__main__': unittest.main()
