import tempfile
import unittest
import faiss
import numpy as np
from example import Store, fixture, normalized
class VectorTests(unittest.TestCase):
    def test_cosine_and_squared_l2(self):
        x=normalized([[1,0],[.8,.6],[0,1]]); q=normalized([[1,0]])
        ip=faiss.IndexFlatIP(2); ip.add(x)
        l2=faiss.IndexFlatL2(2); l2.add(x)
        scores, ids=ip.search(q,3); distance, other=l2.search(q,3)
        np.testing.assert_array_equal(ids,other)
        np.testing.assert_allclose(distance,2-2*scores,atol=1e-6)
    def test_filter_before_limit(self):
        s=fixture()
        self.assertEqual([i for i,_ in s.search([1,0],2)],[101,102])
        self.assertEqual([i for i,_ in s.search([1,0],2,'blue')],[201,202])
        self.assertEqual(s.search([1,0],2,'missing'),[])
    def test_persistence_and_large_k(self):
        s=fixture()
        with tempfile.TemporaryDirectory() as d:
            s.save(d); self.assertEqual(Store.load_trusted(d).search([1,0],100),s.search([1,0],100))
        self.assertEqual(len(s.search([1,0],100)),4)
    def test_invalid_vectors(self):
        for x in ([[0,0]],[[float('nan'),1]],[[float('inf'),1]],[],[1,2]):
            with self.assertRaises(ValueError): normalized(x)
    def test_invalid_metadata_and_query(self):
        for ids in ([1,1],[-1,2],[True,2],[2**63,2]):
            with self.assertRaises(ValueError): Store([[1,0],[0,1]],ids,['a','a'])
        for q,k in (([1],2),([1,0],0),([1,0],True)):
            with self.assertRaises(ValueError): fixture().search(q,k)
if __name__=='__main__': unittest.main()
