use ailake_core::{AilakeResult, Centroid, RowId, VectorStoragePolicy};
use ailake_index::{
HnswBuilder, HnswConfig, HnswSerializer, IvfPqCodebook, IvfPqConfig, IvfPqIndex,
IvfPqSerializer,
};
use ailake_parquet::ParquetVectorWriter;
use ailake_vec::compute_centroid_and_radius;
use arrow_array::RecordBatch;
use bytes::{BufMut, Bytes, BytesMut};
use crate::footer::{
parquet_footer_start, AilakeHeader, AilakeTrailer, DistanceMetric, Precision,
AILAKE_FORMAT_VERSION, AILK_FTS_HEADER_SIZE, AILK_FTS_MAGIC, FLAG_INDEX_IVF_PQ, HEADER_SIZE,
KV_FTS_OFFSET, TRAILER_SIZE,
};
#[derive(Debug, Clone)]
pub enum IndexType {
Hnsw(HnswConfig),
IvfPq(IvfPqConfig),
Auto,
}
impl Default for IndexType {
fn default() -> Self {
IndexType::Hnsw(HnswConfig::default())
}
}
pub struct VectorColumnBatch<'a> {
pub policy: &'a VectorStoragePolicy,
pub embeddings: &'a [Vec<f32>],
}
pub struct AilakeFileWriter {
policy: VectorStoragePolicy,
index_type: IndexType,
shared_codebook: Option<std::sync::Arc<IvfPqCodebook>>,
fts_config: Option<ailake_fts::FtsConfig>,
prebuilt_fts_blob: Option<Vec<u8>>,
}
impl AilakeFileWriter {
pub fn new(policy: VectorStoragePolicy) -> Self {
Self {
policy,
index_type: IndexType::default(),
shared_codebook: None,
fts_config: None,
prebuilt_fts_blob: None,
}
}
pub fn with_fts(mut self, config: ailake_fts::FtsConfig) -> Self {
self.fts_config = Some(config);
self
}
pub fn with_prebuilt_fts_blob(mut self, blob: Vec<u8>) -> Self {
self.prebuilt_fts_blob = Some(blob);
self
}
pub fn with_shared_ivf_codebook(mut self, codebook: std::sync::Arc<IvfPqCodebook>) -> Self {
self.shared_codebook = Some(codebook);
self
}
pub fn with_hnsw_config(mut self, config: HnswConfig) -> Self {
self.index_type = IndexType::Hnsw(config);
self
}
pub fn with_ivf_pq(mut self, config: IvfPqConfig) -> Self {
self.index_type = IndexType::IvfPq(config);
self
}
pub fn with_index_type(mut self, index_type: IndexType) -> Self {
self.index_type = index_type;
self
}
pub fn with_auto_index(mut self) -> Self {
self.index_type = IndexType::Auto;
self
}
pub fn write_with_prebuilt_hnsw(
&self,
batch: &RecordBatch,
embeddings: &[Vec<f32>],
hnsw: &ailake_index::HnswIndex,
) -> AilakeResult<Bytes> {
use ailake_core::AilakeError;
let parquet_writer = ParquetVectorWriter::new(self.policy.clone());
let (parquet_v1, record_count) = parquet_writer.write_batch(batch, embeddings)?;
let footer_start = parquet_footer_start(&parquet_v1)?;
let index_bytes = HnswSerializer::to_bytes(hnsw)?;
let ailk_section = build_ailk_section_from_index_bytes(
&self.policy,
embeddings,
record_count,
footer_start as u64,
&index_bytes,
0u16, )?;
let kv_val = footer_start.to_string();
let kv_refs: &[(&str, &str)] = &[("ailake.footer_offset", kv_val.as_str())];
let (parquet_v2, _) = parquet_writer.write_batch_with_kv(batch, embeddings, kv_refs)?;
let footer_start_v2 = parquet_footer_start(&parquet_v2).map_err(|e| {
AilakeError::Parquet(format!(
"footer_start unstable in write_with_prebuilt_hnsw: {e}"
))
})?;
debug_assert_eq!(
footer_start, footer_start_v2,
"footer_start must be stable across KV injection"
);
let footer_len_v2 = parquet_v2.len() - footer_start_v2;
let mut out = BytesMut::with_capacity(footer_start + ailk_section.len() + footer_len_v2);
out.put_slice(&parquet_v1[..footer_start]);
drop(parquet_v1);
out.put(ailk_section);
out.put_slice(&parquet_v2[footer_start_v2..]);
Ok(out.freeze())
}
pub fn write_parquet_only(
&self,
batch: &RecordBatch,
embeddings: &[Vec<f32>],
) -> AilakeResult<Bytes> {
let parquet_writer = ParquetVectorWriter::new(self.policy.clone());
let (bytes, _) = parquet_writer.write_batch(batch, embeddings)?;
Ok(bytes)
}
pub fn write(&self, batch: &RecordBatch, embeddings: &[Vec<f32>]) -> AilakeResult<Bytes> {
let col = VectorColumnBatch {
policy: &self.policy,
embeddings,
};
self.write_multi(batch, &[col])
}
pub fn write_single_pass(
&self,
batch: &RecordBatch,
embeddings: &[Vec<f32>],
) -> AilakeResult<Bytes> {
let col = VectorColumnBatch {
policy: &self.policy,
embeddings,
};
self.write_multi_single_pass(batch, &[col])
}
pub fn write_multi_single_pass(
&self,
batch: &RecordBatch,
columns: &[VectorColumnBatch<'_>],
) -> AilakeResult<Bytes> {
use ailake_core::AilakeError;
if columns.is_empty() {
return Err(AilakeError::InvalidArgument(
"write_multi_single_pass requires at least one vector column".into(),
));
}
let primary = &columns[0];
let parquet_writer = ParquetVectorWriter::new(primary.policy.clone());
let (parquet_bytes, record_count) =
parquet_writer.write_batch(batch, primary.embeddings)?;
let footer_start = parquet_footer_start(&parquet_bytes)?;
let mut ailk_sections: Vec<Bytes> = Vec::with_capacity(columns.len());
let mut current_offset = footer_start as u64;
for col in columns.iter() {
let section = build_ailk_section(
col.policy,
col.embeddings,
record_count,
current_offset,
&self.index_type,
self.shared_codebook.as_deref(),
)?;
current_offset += section.len() as u64;
ailk_sections.push(section);
}
let total_ailk: usize = ailk_sections.iter().map(|s| s.len()).sum();
let mut out = BytesMut::with_capacity(parquet_bytes.len() + total_ailk);
out.put_slice(&parquet_bytes[..footer_start]);
for section in ailk_sections {
out.put(section);
}
out.put_slice(&parquet_bytes[footer_start..]);
Ok(out.freeze())
}
pub fn write_multi(
&self,
batch: &RecordBatch,
columns: &[VectorColumnBatch<'_>],
) -> AilakeResult<Bytes> {
use ailake_core::AilakeError;
if columns.is_empty() {
return Err(AilakeError::InvalidArgument(
"write_multi requires at least one vector column".into(),
));
}
let primary = &columns[0];
let parquet_writer = ParquetVectorWriter::new(primary.policy.clone());
let (parquet_v1, record_count) = parquet_writer.write_batch(batch, primary.embeddings)?;
let footer_start = parquet_footer_start(&parquet_v1)?;
let mut ailk_sections: Vec<Bytes> = Vec::with_capacity(columns.len());
let mut kv_owned: Vec<(String, String)> = Vec::with_capacity(columns.len());
let mut current_offset = footer_start as u64;
for (i, col) in columns.iter().enumerate() {
let section = build_ailk_section(
col.policy,
col.embeddings,
record_count,
current_offset,
&self.index_type,
self.shared_codebook.as_deref(),
)?;
let kv_key = if i == 0 {
"ailake.footer_offset".to_string()
} else {
format!("ailake.{}.footer_offset", col.policy.column_name)
};
kv_owned.push((kv_key, current_offset.to_string()));
current_offset += section.len() as u64;
ailk_sections.push(section);
}
let fts_blob: Option<Vec<u8>> = self.prebuilt_fts_blob.clone().or_else(|| {
self.fts_config
.as_ref()
.and_then(|cfg| ailake_fts::build_fts_blob_from_batch(cfg, batch).ok())
});
let fts_section: Option<Bytes> = fts_blob.map(|blob| {
let blob_len = blob.len() as u64;
let mut sec = BytesMut::with_capacity(AILK_FTS_HEADER_SIZE + blob.len());
sec.put_slice(&AILK_FTS_MAGIC);
sec.put_slice(&1u16.to_le_bytes()); sec.put_slice(&0u16.to_le_bytes()); sec.put_slice(&blob_len.to_le_bytes());
sec.put_slice(&blob);
sec.freeze()
});
if let Some(ref sec) = fts_section {
let fts_abs_offset = current_offset; kv_owned.push((KV_FTS_OFFSET.to_string(), fts_abs_offset.to_string()));
let _ = sec; }
let kv_refs: Vec<(&str, &str)> = kv_owned
.iter()
.map(|(k, v)| (k.as_str(), v.as_str()))
.collect();
let (parquet_v2, _) =
parquet_writer.write_batch_with_kv(batch, primary.embeddings, &kv_refs)?;
let footer_start_v2 = parquet_footer_start(&parquet_v2)?;
debug_assert_eq!(
footer_start, footer_start_v2,
"footer_start must be stable across KV injection (row groups unchanged)"
);
let total_vector_ailk: usize = ailk_sections.iter().map(|s| s.len()).sum();
let total_fts: usize = fts_section.as_ref().map_or(0, |s| s.len());
let footer_len_v2 = parquet_v2.len() - footer_start_v2;
let total = footer_start + total_vector_ailk + total_fts + footer_len_v2;
let mut out = BytesMut::with_capacity(total);
out.put_slice(&parquet_v1[..footer_start]);
drop(parquet_v1);
for section in ailk_sections {
out.put(section);
}
if let Some(fts_sec) = fts_section {
out.put(fts_sec);
}
out.put_slice(&parquet_v2[footer_start_v2..]);
Ok(out.freeze())
}
}
fn build_ailk_section_from_index_bytes(
policy: &VectorStoragePolicy,
embeddings: &[Vec<f32>],
record_count: u64,
ailk_abs_offset: u64,
index_bytes: &[u8],
flags: u16,
) -> AilakeResult<Bytes> {
let norm_storage: Vec<Vec<f32>>;
let (emb_for_centroid, centroid_metric) =
if policy.pre_normalize && policy.metric == ailake_core::VectorMetric::Cosine {
norm_storage = embeddings
.iter()
.map(|v| ailake_vec::normalize_l2(v))
.collect();
(
norm_storage.as_slice(),
ailake_core::VectorMetric::NormalizedCosine,
)
} else {
(embeddings, policy.metric)
};
let centroid = compute_centroid_and_radius(emb_for_centroid, centroid_metric);
let centroid_bytes = encode_centroid(¢roid);
let centroid_offset = HEADER_SIZE as u64;
let centroid_len = centroid_bytes.len() as u64;
let index_offset_in_ailk = centroid_offset + centroid_len;
let index_len = index_bytes.len() as u64;
let ailk_total_len = HEADER_SIZE as u64 + centroid_len + index_len + TRAILER_SIZE as u64;
let header = AilakeHeader {
format_version: AILAKE_FORMAT_VERSION,
flags,
dim: policy.dim,
precision: Precision::from(policy.precision),
distance_metric: DistanceMetric::from(policy.metric),
record_count,
centroid_offset,
centroid_len,
hnsw_offset: index_offset_in_ailk,
hnsw_len: index_len,
};
let trailer = AilakeTrailer {
footer_offset: ailk_abs_offset,
footer_len: ailk_total_len,
format_version: AILAKE_FORMAT_VERSION,
flags,
};
let mut buf = BytesMut::with_capacity(ailk_total_len as usize);
buf.put_slice(&header.to_bytes());
buf.put_slice(¢roid_bytes);
buf.put_slice(index_bytes);
buf.put_slice(&trailer.to_bytes());
Ok(buf.freeze())
}
fn build_ailk_section(
policy: &VectorStoragePolicy,
embeddings: &[Vec<f32>],
record_count: u64,
ailk_abs_offset: u64,
index_type: &IndexType,
shared_codebook: Option<&IvfPqCodebook>,
) -> AilakeResult<Bytes> {
let norm_storage: Vec<Vec<f32>>;
let (embeddings, hnsw_metric) =
if policy.pre_normalize && policy.metric == ailake_core::VectorMetric::Cosine {
norm_storage = embeddings
.iter()
.map(|v| ailake_vec::normalize_l2(v))
.collect();
(
norm_storage.as_slice(),
ailake_core::VectorMetric::NormalizedCosine,
)
} else {
(embeddings, policy.metric)
};
let centroid: Centroid = compute_centroid_and_radius(embeddings, hnsw_metric);
let centroid_bytes = encode_centroid(¢roid);
let resolved: IndexType;
let index_type = if matches!(index_type, IndexType::Auto) {
let profile = ailake_index::HardwareProfile::detect();
resolved = if profile.recommend_ivf_pq(embeddings.len()) {
IndexType::IvfPq(ailake_index::IvfPqConfig::for_dataset(
policy.dim as usize,
embeddings.len(),
))
} else {
IndexType::Hnsw(ailake_index::HnswConfig::default())
};
&resolved
} else {
index_type
};
let (index_bytes, flags) = match index_type {
IndexType::Hnsw(hnsw_config) => {
let config = HnswConfig {
m: policy.hnsw_m.map(|v| v as usize).unwrap_or(hnsw_config.m),
ef_construction: policy
.hnsw_ef_construction
.map(|v| v as usize)
.unwrap_or(hnsw_config.ef_construction),
max_elements: hnsw_config.max_elements,
};
let mut builder = HnswBuilder::new(policy.dim, hnsw_metric, config);
for (i, v) in embeddings.iter().enumerate() {
builder.insert(RowId::new(i as u64), v.clone());
}
let index = builder.build();
(HnswSerializer::to_bytes(&index)?, 0u16)
}
IndexType::IvfPq(ivf_config) => {
let row_ids: Vec<RowId> = (0..embeddings.len() as u64).map(RowId::new).collect();
let index = if let Some(cb) = shared_codebook {
IvfPqIndex::build_with_codebook(&row_ids, embeddings, cb)?
} else {
ailake_index::IvfPqIndex::train(
&row_ids,
embeddings,
policy.metric,
ivf_config.clone(),
)?
};
(IvfPqSerializer::to_bytes(&index)?, FLAG_INDEX_IVF_PQ)
}
IndexType::Auto => unreachable!("Auto resolved above"),
};
let centroid_offset = HEADER_SIZE as u64;
let centroid_len = centroid_bytes.len() as u64;
let index_offset_in_ailk = centroid_offset + centroid_len;
let index_len = index_bytes.len() as u64;
let ailk_total_len = HEADER_SIZE as u64 + centroid_len + index_len + TRAILER_SIZE as u64;
let header = AilakeHeader {
format_version: AILAKE_FORMAT_VERSION,
flags,
dim: policy.dim,
precision: Precision::from(policy.precision),
distance_metric: DistanceMetric::from(policy.metric),
record_count,
centroid_offset,
centroid_len,
hnsw_offset: index_offset_in_ailk,
hnsw_len: index_len,
};
let trailer = AilakeTrailer {
footer_offset: ailk_abs_offset,
footer_len: ailk_total_len,
format_version: AILAKE_FORMAT_VERSION,
flags,
};
let mut buf = BytesMut::with_capacity(ailk_total_len as usize);
buf.put_slice(&header.to_bytes());
buf.put_slice(¢roid_bytes);
buf.put_slice(&index_bytes);
buf.put_slice(&trailer.to_bytes());
Ok(buf.freeze())
}
fn encode_centroid(c: &Centroid) -> Vec<u8> {
let mut bytes = Vec::with_capacity(c.values.len() * 4 + 4);
for &v in &c.values {
bytes.extend_from_slice(&v.to_le_bytes());
}
bytes.extend_from_slice(&c.radius.to_le_bytes());
bytes
}
#[cfg(test)]
mod tests {
use super::*;
use ailake_core::{VectorMetric, VectorPrecision};
use arrow_array::{Int32Array, RecordBatch};
use arrow_schema::{DataType, Field, Schema};
use std::sync::Arc;
fn make_policy(dim: u32) -> VectorStoragePolicy {
VectorStoragePolicy {
column_name: "embedding".to_string(),
dim,
metric: VectorMetric::Cosine,
precision: VectorPrecision::F16,
pq: None,
keep_raw_for_reranking: true,
pre_normalize: false,
hnsw_m: None,
hnsw_ef_construction: None,
ivf_residual: false,
embedding_model: None,
modality: None,
partition_by: None,
partition_value: None,
partition_column_type: None,
partition_fields: vec![],
}
}
#[test]
fn write_single_pass_valid_parquet_and_ailk() {
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
let batch =
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 2, 3]))]).unwrap();
let embs: Vec<Vec<f32>> = (0..3).map(|_| vec![0.1, 0.2, 0.3, 0.4]).collect();
let writer = AilakeFileWriter::new(make_policy(4));
let file = writer.write_single_pass(&batch, &embs).unwrap();
assert_eq!(&file[..4], b"PAR1");
assert_eq!(&file[file.len() - 4..], b"PAR1");
assert!(file.windows(4).any(|w| w == b"AILK"));
}
#[test]
fn write_single_pass_reader_bootstrap_from_trailer() {
use crate::reader::AilakeFileReader;
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
let batch =
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![10, 20, 30]))])
.unwrap();
let embs: Vec<Vec<f32>> = (0..3).map(|i| vec![i as f32, 0.0, 0.0, 0.0]).collect();
let writer = AilakeFileWriter::new(make_policy(4));
let file_bytes = writer.write_single_pass(&batch, &embs).unwrap();
let reader = AilakeFileReader::new(file_bytes, "embedding", 4);
assert!(
reader.is_ailake_file(),
"single-pass file must be recognised as AI-Lake file via trailer bootstrap"
);
let header = reader.read_header().expect("must read AILK header");
assert_eq!(header.dim, 4);
assert_eq!(header.record_count, 3);
}
#[test]
fn write_and_write_single_pass_same_index() {
use crate::reader::AilakeFileReader;
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
let batch = RecordBatch::try_new(
schema,
vec![Arc::new(Int32Array::from(vec![1, 2, 3, 4, 5]))],
)
.unwrap();
let embs: Vec<Vec<f32>> = vec![
vec![1.0, 0.0, 0.0, 0.0],
vec![0.0, 1.0, 0.0, 0.0],
vec![0.0, 0.0, 1.0, 0.0],
vec![0.0, 0.0, 0.0, 1.0],
vec![0.5, 0.5, 0.0, 0.0],
];
let policy = make_policy(4);
let writer = AilakeFileWriter::new(policy);
let bytes_two_pass = writer.write(&batch, &embs).unwrap();
let bytes_single_pass = writer.write_single_pass(&batch, &embs).unwrap();
let query = vec![1.0f32, 0.0, 0.0, 0.0];
let reader_tp = AilakeFileReader::new(bytes_two_pass, "embedding", 4);
let reader_sp = AilakeFileReader::new(bytes_single_pass, "embedding", 4);
let idx_tp = reader_tp.load_index().unwrap();
let idx_sp = reader_sp.load_index().unwrap();
let res_tp = idx_tp.search(&query, 1, 50);
let res_sp = idx_sp.search(&query, 1, 50);
assert_eq!(res_tp[0].0, res_sp[0].0, "nearest neighbour must match");
}
#[test]
fn write_ends_with_par1() {
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
let batch =
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 2, 3]))]).unwrap();
let embs: Vec<Vec<f32>> = (0..3).map(|_| vec![0.1, 0.2, 0.3, 0.4]).collect();
let writer = AilakeFileWriter::new(make_policy(4));
let file = writer.write(&batch, &embs).unwrap();
assert_eq!(&file[file.len() - 4..], b"PAR1");
assert_eq!(&file[..4], b"PAR1");
assert!(file.windows(4).any(|w| w == b"AILK"));
}
#[test]
fn write_multi_two_columns() {
use ailake_core::{VectorMetric, VectorPrecision};
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
let batch =
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 2, 3]))]).unwrap();
let embs: Vec<Vec<f32>> = (0..3).map(|i| vec![i as f32, 0.0, 0.0, 0.0]).collect();
let ctx_embs: Vec<Vec<f32>> = (0..3).map(|i| vec![0.0, i as f32, 0.0, 0.0]).collect();
let policy1 = make_policy(4);
let policy2 = VectorStoragePolicy {
column_name: "context_embedding".to_string(),
dim: 4,
metric: VectorMetric::Cosine,
precision: VectorPrecision::F16,
pq: None,
keep_raw_for_reranking: true,
pre_normalize: false,
hnsw_m: None,
hnsw_ef_construction: None,
ivf_residual: false,
embedding_model: None,
modality: None,
partition_by: None,
partition_value: None,
partition_column_type: None,
partition_fields: vec![],
};
let writer = AilakeFileWriter::new(policy1.clone());
let file = writer
.write_multi(
&batch,
&[
VectorColumnBatch {
policy: &policy1,
embeddings: &embs,
},
VectorColumnBatch {
policy: &policy2,
embeddings: &ctx_embs,
},
],
)
.unwrap();
assert_eq!(&file[..4], b"PAR1");
assert_eq!(&file[file.len() - 4..], b"PAR1");
let ailk_count = file.windows(4).filter(|w| *w == b"AILK").count();
assert!(
ailk_count >= 2,
"expected >= 2 AILK markers, got {ailk_count}"
);
}
}