# /// script
# requires-python = ">=3.12,<3.13"
# dependencies = ["pylance==11.0.0", "pyarrow==25.0.1", "numpy==2.4.3"]
# ///
import json
import sys
import tempfile
from pathlib import Path

import lance
import numpy as np
import pyarrow as pa

vectors = np.array([[0,0],[1,0],[0,2],[3,1],[2,3],[5,0],[4,4],[6,2],[7,5],[8,1],[9,3],[10,6]], dtype=np.float32)
centroids = np.array([[1,1],[5,2],[9,4]], dtype=np.float32)
assignments = np.argmin(np.sum((vectors[:,None,:]-centroids[None,:,:])**2,axis=2),axis=1)
points = [[2.25,1.125],[6.25,3.5],[9.125,2.25],[4.125,1.25],[5.25,4.125],[7.125,2.25]]
groups = ['A','A','B','A','B','A','B','B','A','B','A','B']
exact_queries = []
queries = []
with tempfile.TemporaryDirectory(prefix='columnar-ivf-') as directory:
    dataset = lance.write_dataset(pa.table({'id':list(range(1,13)), 'group':groups, 'vector':pa.FixedSizeListArray.from_arrays(pa.array(vectors.flatten()),2)}), Path(directory)/'table', data_storage_version='2.1')
    dataset.create_index('vector','IVF_FLAT',num_partitions=3,ivf_centroids=centroids)
    assert dataset.list_indices()[0]['type'] == 'IVF_FLAT'
    for point in points:
        distances = np.sum((vectors-np.array(point,dtype=np.float32))**2,axis=1)
        exact = np.argsort(distances,kind='stable')[:3]
        for group in ['all','A','B']:
            eligible = np.array([i for i in range(12) if group=='all' or groups[i]==group])
            nearest = eligible[np.argsort(distances[eligible],kind='stable')[:3]]
            rows = dataset.to_table(filter=None if group=='all' else f"group = '{group}'", nearest={'column':'vector','q':point,'k':3,'use_index':False}, prefilter=True).to_pylist()
            assert [row['id'] for row in rows] == (nearest+1).tolist()
            assert [row['_distance'] for row in rows] == distances[nearest].tolist()
            exact_queries.append({'point':point,'group':group,'rows':rows})
        centroid_order = np.argsort(np.sum((centroids-np.array(point,dtype=np.float32))**2,axis=1),kind='stable')
        for probes in [1,2,3]:
            selected = centroid_order[:probes]
            candidates = np.flatnonzero(np.isin(assignments,selected))
            expected = candidates[np.argsort(distances[candidates],kind='stable')[:3]]
            actual = dataset.to_table(nearest={'column':'vector','q':point,'k':3,'nprobes':probes}).to_pylist()
            assert [r['id'] for r in actual] == (expected+1).tolist()
            assert [r['_distance'] for r in actual] == distances[expected].tolist()
            queries.append({'point':point,'probes':probes,'partitions':selected.tolist(),'candidates':(candidates+1).tolist(),'ids':(expected+1).tolist(),'distances':distances[expected].tolist(),'exactIds':(exact+1).tolist(),'recall':len(set(expected)&set(exact))/3})
result = {'sdk':lance.__version__,'index':'IVF_FLAT','centroids':centroids.tolist(),'vectors':vectors.tolist(),'assignments':assignments.tolist(),'queries':queries,'points':points,'exactQueries':exact_queries}
repo = next((parent for parent in Path(__file__).resolve().parents if (parent/'explorer/src/lib/depth').is_dir()), None)
if repo is not None:
    targets = [repo/'explorer/src/lib/depth/lance-index.json',repo/'explorer/static/lance/index-proof.json']
else:
    targets = [Path('index-proof.json')]
for target in targets:
    if '--verify' in sys.argv:
        assert json.loads(target.read_text()) == result
    else:
        target.write_text(json.dumps(result,indent=2)+'\n')
print('Verified all eighteen real Lance IVF_FLAT searches against explicit partitions and NumPy exact neighbors.')
