Skip to main content

ailake_file/
reader.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2use crate::footer::Precision;
3use ailake_core::{AilakeError, AilakeResult, Centroid, VectorMetric};
4use ailake_index::{AnyIndex, HnswIndex, IvfPqSerializer, MmapLoader};
5use ailake_parquet::ParquetVectorReader;
6use arrow_array::RecordBatch;
7use bytes::Bytes;
8
9use crate::footer::{AilakeHeader, DistanceMetric, FLAG_INDEX_IVF_PQ, HEADER_SIZE};
10
11pub struct AilakeFileReader {
12    bytes: Bytes,
13    vector_column: String,
14    #[allow(dead_code)]
15    dim: u32,
16}
17
18impl AilakeFileReader {
19    pub fn new(bytes: Bytes, vector_column: &str, dim: u32) -> Self {
20        Self {
21            bytes,
22            vector_column: vector_column.to_string(),
23            dim,
24        }
25    }
26
27    /// Returns the absolute byte offset of the primary AILK section.
28    /// Reads `ailake.footer_offset` from the Parquet footer key-value metadata.
29    pub fn ailk_offset(&self) -> AilakeResult<u64> {
30        let reader = ParquetVectorReader::new(self.bytes.clone(), &self.vector_column);
31        let val = reader
32            .kv_metadata("ailake.footer_offset")?
33            .ok_or(AilakeError::NotAnAilakeFile)?;
34        val.parse::<u64>().map_err(|_| AilakeError::NotAnAilakeFile)
35    }
36
37    /// Returns the absolute byte offset of the AILK section for a named vector column.
38    ///
39    /// For additional columns tries `ailake.{column}.footer_offset` first,
40    /// then falls back to `ailake.footer_offset` (primary / single-column files).
41    pub fn ailk_offset_for_column(&self, column: &str) -> AilakeResult<u64> {
42        let reader = ParquetVectorReader::new(self.bytes.clone(), column);
43        let col_key = format!("ailake.{column}.footer_offset");
44        if let Some(val) = reader.kv_metadata(&col_key)? {
45            return val.parse::<u64>().map_err(|_| AilakeError::NotAnAilakeFile);
46        }
47        let val = reader
48            .kv_metadata("ailake.footer_offset")?
49            .ok_or(AilakeError::NotAnAilakeFile)?;
50        val.parse::<u64>().map_err(|_| AilakeError::NotAnAilakeFile)
51    }
52
53    /// Returns true if the file contains an embedded AILK section.
54    pub fn is_ailake_file(&self) -> bool {
55        self.ailk_offset().is_ok()
56    }
57
58    /// Parse the 64-byte AI-Lake header from the embedded AILK section.
59    pub fn read_header(&self) -> AilakeResult<AilakeHeader> {
60        let offset = self.ailk_offset()? as usize;
61        if offset + HEADER_SIZE > self.bytes.len() {
62            return Err(AilakeError::NotAnAilakeFile);
63        }
64        let header_bytes: &[u8; HEADER_SIZE] = self.bytes[offset..offset + HEADER_SIZE]
65            .try_into()
66            .map_err(|_| AilakeError::NotAnAilakeFile)?;
67        AilakeHeader::from_bytes(header_bytes)
68    }
69
70    /// Read centroid + radius from the AILK section.
71    pub fn get_centroid(&self) -> AilakeResult<Centroid> {
72        let ailk_start = self.ailk_offset()? as usize;
73        let header = self.read_header()?;
74        let centroid_start = ailk_start + header.centroid_offset as usize;
75        let centroid_end = centroid_start + header.centroid_len as usize;
76
77        if centroid_end > self.bytes.len() {
78            return Err(AilakeError::NotAnAilakeFile);
79        }
80
81        let centroid_data = &self.bytes[centroid_start..centroid_end];
82        let dim = header.dim as usize;
83        let expected_len = dim * 4 + 4;
84        if centroid_data.len() != expected_len {
85            return Err(AilakeError::InvalidCentroidLength {
86                expected_dim: header.dim,
87                actual: centroid_data.len(),
88            });
89        }
90
91        let values: Vec<f32> = centroid_data[..dim * 4]
92            .chunks_exact(4)
93            .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
94            .collect();
95        let radius = f32::from_le_bytes(centroid_data[dim * 4..].try_into().unwrap());
96        let metric = distance_metric_to_vector_metric(header.distance_metric);
97
98        Ok(Centroid {
99            values,
100            radius,
101            metric,
102        })
103    }
104
105    /// Load the HNSW index from the primary AILK section.
106    pub fn load_index(&self) -> AilakeResult<HnswIndex> {
107        self.load_index_for_column(&self.vector_column.clone())
108    }
109
110    /// Load the HNSW index for a specific vector column.
111    ///
112    /// Works for both single-column files (falls back to primary AILK) and
113    /// multi-column files written with `AilakeFileWriter::write_multi`.
114    pub fn load_index_for_column(&self, column: &str) -> AilakeResult<HnswIndex> {
115        let ailk_start = self.ailk_offset_for_column(column)? as usize;
116
117        if ailk_start + HEADER_SIZE > self.bytes.len() {
118            return Err(AilakeError::NotAnAilakeFile);
119        }
120        let header_bytes: &[u8; HEADER_SIZE] = self.bytes[ailk_start..ailk_start + HEADER_SIZE]
121            .try_into()
122            .map_err(|_| AilakeError::NotAnAilakeFile)?;
123        let header = AilakeHeader::from_bytes(header_bytes)?;
124
125        let hnsw_start = ailk_start + header.hnsw_offset as usize;
126        let hnsw_end = hnsw_start + header.hnsw_len as usize;
127
128        if hnsw_end > self.bytes.len() {
129            return Err(AilakeError::NotAnAilakeFile);
130        }
131        let mut idx = MmapLoader::from_bytes(&self.bytes[hnsw_start..hnsw_end])?;
132        if header.precision == Precision::F16 {
133            idx.quantize_to_f16();
134        }
135        Ok(idx)
136    }
137
138    /// Load primary index as `AnyIndex`, dispatching on header flags.
139    pub fn load_any_index(&self) -> AilakeResult<AnyIndex> {
140        self.load_any_index_for_column(&self.vector_column.clone())
141    }
142
143    /// Load index for a specific vector column as `AnyIndex`.
144    pub fn load_any_index_for_column(&self, column: &str) -> AilakeResult<AnyIndex> {
145        let ailk_start = self.ailk_offset_for_column(column)? as usize;
146
147        if ailk_start + HEADER_SIZE > self.bytes.len() {
148            return Err(AilakeError::NotAnAilakeFile);
149        }
150        let header_bytes: &[u8; HEADER_SIZE] = self.bytes[ailk_start..ailk_start + HEADER_SIZE]
151            .try_into()
152            .map_err(|_| AilakeError::NotAnAilakeFile)?;
153        let header = AilakeHeader::from_bytes(header_bytes)?;
154
155        let index_start = ailk_start + header.hnsw_offset as usize;
156        let index_end = index_start + header.hnsw_len as usize;
157
158        if index_end > self.bytes.len() {
159            return Err(AilakeError::NotAnAilakeFile);
160        }
161        let index_bytes = &self.bytes[index_start..index_end];
162
163        if header.flags & FLAG_INDEX_IVF_PQ != 0 {
164            let idx = IvfPqSerializer::from_bytes(index_bytes)?;
165            Ok(AnyIndex::IvfPq(idx))
166        } else {
167            let mut idx = MmapLoader::from_bytes(index_bytes)?;
168            if header.precision == Precision::F16 {
169                idx.quantize_to_f16();
170            }
171            Ok(AnyIndex::Hnsw(idx))
172        }
173    }
174
175    /// Read the Parquet section (tabular data + decoded embeddings).
176    /// The full file is valid Parquet; the AILK section is invisible to standard readers.
177    pub fn read_parquet(&self) -> AilakeResult<(RecordBatch, Vec<Vec<f32>>)> {
178        let reader = ParquetVectorReader::new(self.bytes.clone(), &self.vector_column);
179        reader.read_all()
180    }
181
182    /// Verify the positional invariant: Parquet record_count == HNSW node_count.
183    pub fn verify_integrity(&self) -> AilakeResult<()> {
184        let header = self.read_header()?;
185        let index = self.load_index()?;
186        let reader = ParquetVectorReader::new(self.bytes.clone(), &self.vector_column);
187        let parquet_count = reader.record_count()?;
188
189        if parquet_count != index.node_count() {
190            return Err(AilakeError::RowCountMismatch {
191                parquet: parquet_count,
192                hnsw: index.node_count(),
193            });
194        }
195        if parquet_count != header.record_count {
196            return Err(AilakeError::RowCountMismatch {
197                parquet: parquet_count,
198                hnsw: header.record_count,
199            });
200        }
201        Ok(())
202    }
203}
204
205fn distance_metric_to_vector_metric(dm: DistanceMetric) -> VectorMetric {
206    match dm {
207        DistanceMetric::Cosine => VectorMetric::Cosine,
208        DistanceMetric::Euclidean => VectorMetric::Euclidean,
209        DistanceMetric::DotProduct => VectorMetric::DotProduct,
210        DistanceMetric::NormalizedCosine => VectorMetric::NormalizedCosine,
211    }
212}
213
214#[cfg(test)]
215mod tests {
216    use super::*;
217    use crate::writer::AilakeFileWriter;
218    use ailake_core::{VectorMetric, VectorPrecision, VectorStoragePolicy};
219    use arrow_array::{Int32Array, RecordBatch};
220    use arrow_schema::{DataType, Field, Schema};
221    use std::sync::Arc;
222
223    fn make_policy(dim: u32) -> VectorStoragePolicy {
224        VectorStoragePolicy {
225            column_name: "embedding".to_string(),
226            dim,
227            metric: VectorMetric::Cosine,
228            precision: VectorPrecision::F16,
229            pq: None,
230            keep_raw_for_reranking: true,
231            pre_normalize: false,
232            hnsw_m: None,
233            hnsw_ef_construction: None,
234        }
235    }
236
237    fn write_file(rows: usize, dim: u32) -> Bytes {
238        let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
239        let ids: Vec<i32> = (0..rows as i32).collect();
240        let batch = RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(ids))]).unwrap();
241        let embs: Vec<Vec<f32>> = (0..rows)
242            .map(|i| {
243                let mut v = vec![0.0f32; dim as usize];
244                v[i % dim as usize] = 1.0;
245                v
246            })
247            .collect();
248        AilakeFileWriter::new(make_policy(dim))
249            .write(&batch, &embs)
250            .unwrap()
251    }
252
253    #[test]
254    fn is_ailake_file() {
255        let file = write_file(3, 4);
256        let reader = AilakeFileReader::new(file, "embedding", 4);
257        assert!(reader.is_ailake_file());
258    }
259
260    #[test]
261    fn integrity_check_passes() {
262        let file = write_file(10, 8);
263        let reader = AilakeFileReader::new(file, "embedding", 8);
264        reader.verify_integrity().unwrap();
265    }
266
267    #[test]
268    fn centroid_has_correct_dim() {
269        let file = write_file(5, 4);
270        let reader = AilakeFileReader::new(file, "embedding", 4);
271        let centroid = reader.get_centroid().unwrap();
272        assert_eq!(centroid.values.len(), 4);
273    }
274
275    #[test]
276    fn search_finds_nearest() {
277        let dim = 4u32;
278        let file = write_file(4, dim);
279        let reader = AilakeFileReader::new(file, "embedding", dim);
280        let index = reader.load_index().unwrap();
281        let query = vec![1.0f32, 0.0, 0.0, 0.0];
282        let results = index.search(&query, 1, 50);
283        assert_eq!(results.len(), 1);
284        assert_eq!(results[0].0, ailake_core::RowId::new(0));
285    }
286
287    #[test]
288    fn parquet_read_returns_tabular_data() {
289        let file = write_file(3, 4);
290        let reader = AilakeFileReader::new(file, "embedding", 4);
291        let (batch, embs) = reader.read_parquet().unwrap();
292        assert_eq!(batch.num_rows(), 3);
293        assert_eq!(embs.len(), 3);
294    }
295}