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

import lance
import numpy as np

folder = Path(__file__).resolve().parent / 'sample'
fixture = json.loads((folder / 'index.json').read_text())
for item in fixture['artifacts']:
    assert hashlib.sha256((folder / item['path']).read_bytes()).hexdigest() == item['sha256']
dataset = lance.dataset(folder)
assert dataset.to_table().to_pylist() == fixture['rows']
vectors = np.array([row['vector'] for row in fixture['rows']],dtype=np.float64)
groups = np.array([row['group'] for row in fixture['rows']])
for query in fixture['queries']:
    group = query['group']
    result = dataset.to_table(nearest={'column':'vector','q':query['point'],'k':3,'use_index':False},filter=None if group=='all' else f"`group` = '{group}'",prefilter=True).to_pylist()
    assert result == query['rows']
    distances = np.sum((vectors-np.array(query['point']))**2,axis=1)
    candidates = np.arange(len(vectors)) if group=='all' else np.flatnonzero(groups==group)
    order = candidates[np.argsort(distances[candidates])][:3]
    assert [row['id'] for row in result] == (order+1).tolist()
    np.testing.assert_array_equal([row['_distance'] for row in result],distances[order])
print('All hashes, source rows and nine exact searches verified against NumPy.')
