"""Exact FAISS search with explicit IDs, cosine metric and prefilter semantics."""
import json
from pathlib import Path
import tempfile
import faiss
import numpy as np


def normalized(vectors):
    x = np.asarray(vectors, dtype=np.float32)
    if x.ndim != 2 or x.shape[0] == 0 or x.shape[1] == 0 or not np.isfinite(x).all():
        raise ValueError("nonempty finite matrix required")
    norms = np.linalg.norm(x.astype(np.float64), axis=1)
    if np.any(norms == 0):
        raise ValueError("cosine is undefined for a zero vector")
    return np.ascontiguousarray(x / norms[:,None], dtype=np.float32)


class Store:
    def __init__(self, vectors, ids, tenants):
        self.vectors = normalized(vectors)
        if len(ids) != len(self.vectors) or len(tenants) != len(ids):
            raise ValueError("metadata length mismatch")
        if any(type(i) is not int or i < 0 or i > 2**63-1 for i in ids) or len(set(ids)) != len(ids):
            raise ValueError("unique nonnegative signed int64 IDs required")
        if any(not isinstance(t,str) or not t for t in tenants):
            raise ValueError("nonempty tenant strings required")
        self.ids, self.tenants = np.array(ids,dtype=np.int64), list(tenants)
        self.index = faiss.IndexIDMap2(faiss.IndexFlatIP(self.vectors.shape[1]))
        self.index.add_with_ids(self.vectors, self.ids)

    def search(self, query, k=2, tenant=None):
        if type(k) is not int or k < 1:
            raise ValueError("positive integer k required")
        q=normalized([query])
        if q.shape[1] != self.vectors.shape[1]:
            raise ValueError("query dimension mismatch")
        # Rebuild a small exact allowlisted index for this demonstration.
        # A FAISS index does not supply an application's authorization policy.
        mask=np.array([tenant is None or t == tenant for t in self.tenants])
        if not mask.any(): return []
        idx=self.index
        if tenant is not None:
            idx=faiss.IndexIDMap2(faiss.IndexFlatIP(q.shape[1]))
            idx.add_with_ids(self.vectors[mask],self.ids[mask])
        scores, ids=idx.search(q,min(k,int(mask.sum())))
        return [(int(i),float(s)) for i,s in zip(ids[0],scores[0]) if i != -1]

    def save(self, directory):
        path=Path(directory);path.mkdir(parents=True,exist_ok=True)
        faiss.write_index(self.index,str(path/'vectors.faiss'))
        (path/'metadata.json').write_text(json.dumps({'ids':self.ids.tolist(),'tenants':self.tenants}))

    @classmethod
    def load_trusted(cls, directory):
        # Only read locally produced, trusted files. FAISS does not validate hostile indexes.
        path=Path(directory); index=faiss.read_index(str(path/'vectors.faiss'))
        meta=json.loads((path/'metadata.json').read_text())
        vectors=np.vstack([index.reconstruct(int(i)) for i in meta['ids']])
        return cls(vectors,meta['ids'],meta['tenants'])


def fixture():
    return Store([[1,0],[.99,.1],[.8,.6],[.6,.8]],[101,102,201,202],['red','red','blue','blue'])
if __name__=='__main__':
    store=fixture(); global_hits=store.search([1,0],2)
    postfiltered=[i for i,s in global_hits if i in (201,202)]
    with tempfile.TemporaryDirectory() as d:
        store.save(d); restored=Store.load_trusted(d)
        print(json.dumps({'global_top2':[i for i,s in global_hits], 'postfilter_blue':postfiltered,
                          'prefilter_blue':[i for i,s in store.search([1,0],2,'blue')],
                          'round_trip_equal':restored.search([1,0]) == global_hits}))
