use super::{codec_deser, codec_ser, MAX_CODEC_BYTES};
use crate::datatypes::values::Value;
use crate::graph::algorithms::hnsw::HnswIndex;
use crate::graph::index_freshness::IndexFreshness;
use crate::graph::schema::DirGraph;
use crate::graph::storage::GraphRead;
use crate::serde_codec;
use flate2::read::GzDecoder;
use flate2::write::GzEncoder;
use flate2::Compression;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::fs::File;
use std::io::{self, BufReader, BufWriter, Read, Write};
pub(super) const VECTOR_INDEX_MAGIC: &[u8; 8] = b"KGLVIDX1";
const VECTOR_INDEX_FORMAT_VERSION: u32 = 3;
struct HeldIndex<'a> {
node_type: &'a String,
embedding_property: &'a String,
guard: crate::graph::schema::HnswRead<'a>,
watermark: u32,
limit: usize,
dirty: Vec<u32>,
}
#[derive(Serialize)]
struct PersistedVectorIndexRef<'a> {
node_type: &'a str,
embedding_property: &'a str,
index: &'a HnswIndex,
watermark: u32,
limit: usize,
dirty: Vec<u32>,
}
#[derive(Serialize, Deserialize)]
struct PersistedVectorIndex {
node_type: String,
embedding_property: String,
index: HnswIndex,
watermark: u32,
limit: usize,
dirty: Vec<u32>,
}
pub(super) fn encode_vector_indexes(graph: &DirGraph) -> io::Result<Option<Vec<u8>>> {
let mut stores: Vec<_> = graph.embeddings.iter().collect();
stores.sort_unstable_by(|a, b| a.0.cmp(b.0));
let held: Vec<HeldIndex<'_>> = stores
.into_iter()
.filter_map(|((nt, prop), store)| {
let (watermark, limit, dirty) = store.freshness_state().persisted_parts();
Some(HeldIndex {
node_type: nt,
embedding_property: prop,
guard: store.index_read()?,
watermark,
limit,
dirty,
})
})
.collect();
if held.is_empty() {
return Ok(None);
}
let entries: Vec<PersistedVectorIndexRef<'_>> = held
.iter()
.map(|held| PersistedVectorIndexRef {
node_type: held.node_type.as_str(),
embedding_property: held.embedding_property.as_str(),
index: &held.guard,
watermark: held.watermark,
limit: held.limit,
dirty: held.dirty.clone(),
})
.collect();
let body = codec_ser(serde_codec::CodecVersion::PostcardV1, &entries)?;
let mut payload = Vec::with_capacity(12 + body.len());
payload.extend_from_slice(VECTOR_INDEX_MAGIC);
payload.extend_from_slice(&VECTOR_INDEX_FORMAT_VERSION.to_le_bytes());
payload.extend_from_slice(&body);
Ok(Some(payload))
}
pub(super) fn decode_vector_indexes(payload: &[u8], graph: &mut DirGraph) {
if payload.len() < 12 || &payload[..8] != VECTOR_INDEX_MAGIC {
return;
}
let ver = u32::from_le_bytes([payload[8], payload[9], payload[10], payload[11]]);
if ver != VECTOR_INDEX_FORMAT_VERSION {
return; }
let codec = serde_codec::CodecVersion::PostcardV1;
let entries: Vec<PersistedVectorIndex> =
match codec_deser(codec, &payload[12..], (payload.len() - 12) as u64) {
Ok(e) => e,
Err(_) => return,
};
for entry in entries {
let key = (entry.node_type, entry.embedding_property);
let Some(store) = graph.embeddings.get_mut(&key) else {
continue;
};
let shape_ok = entry
.index
.validate_for_store(&store.data, &store.norms, store.dimension)
.is_ok();
if !shape_ok
|| entry.watermark as usize != entry.index.len()
|| entry.dirty.iter().any(|slot| *slot >= entry.watermark)
{
continue;
}
let freshness = IndexFreshness::restored(entry.watermark, entry.limit, &entry.dirty);
store.attach_persisted_index(entry.index, freshness);
}
}
const KGLE_MAGIC: [u8; 4] = *b"KGLE";
const KGLE_VERSION: u32 = 3;
#[derive(Serialize, Deserialize)]
struct ExportedEmbeddingStore {
node_type: String,
text_column: String, dimension: usize,
metric: Option<String>,
model_id: Option<String>,
entries: Vec<(Value, Vec<f32>, Option<u64>)>,
}
pub enum EmbeddingExportFilter {
Types(Vec<String>),
TypeProperties(HashMap<String, Vec<String>>),
}
pub struct ExportStats {
pub stores: usize,
pub embeddings: usize,
}
pub struct ImportStats {
pub stores: usize,
pub imported: usize,
pub skipped: usize,
pub dropped_stores: usize,
}
fn decode_embedding_file_payload(
buf: &[u8],
version: u32,
) -> io::Result<Vec<ExportedEmbeddingStore>> {
if version < KGLE_VERSION {
return Err(super::pre_014_bincode_error(
format!(".kgle embedding file v{version}").as_str(),
));
}
if buf.len() < 9 {
return Err(io::Error::other(
"Embedding file v3 is truncated before its codec tag.",
));
}
let codec = serde_codec::CodecVersion::from_tag(buf[8])
.map_err(|e| io::Error::other(format!("Invalid .kgle codec tag: {e}")))?;
let decoder = GzDecoder::new(&buf[9..]);
let mut bounded = decoder.take(MAX_CODEC_BYTES.saturating_add(1));
let mut payload = Vec::new();
bounded.read_to_end(&mut payload)?;
if payload.len() as u64 > MAX_CODEC_BYTES {
return Err(io::Error::other(format!(
"Decompressed embedding payload exceeds the {MAX_CODEC_BYTES} byte limit"
)));
}
codec_deser(codec, &payload, payload.capacity() as u64)
.map_err(|e| io::Error::other(format!("Failed to deserialize embedding data: {e}")))
}
pub fn export_embeddings_to_file(
graph: &DirGraph,
path: &str,
filter: Option<&EmbeddingExportFilter>,
) -> io::Result<ExportStats> {
let _arena_guard = graph.graph.begin_query();
let mut exported_stores: Vec<ExportedEmbeddingStore> = Vec::new();
let mut total_embeddings = 0usize;
let mut stores_sorted: Vec<_> = graph.embeddings.iter().collect();
stores_sorted.sort_unstable_by(|a, b| a.0.cmp(b.0));
for ((node_type, store_name), store) in stores_sorted {
let text_column =
crate::graph::embeddings::text_column_of(store_name).unwrap_or(store_name.as_str());
if let Some(f) = filter {
match f {
EmbeddingExportFilter::Types(types) => {
if !types.iter().any(|t| t == node_type) {
continue;
}
}
EmbeddingExportFilter::TypeProperties(map) => {
match map.get(node_type) {
None => continue, Some(props) if !props.is_empty() => {
if !props.iter().any(|p| p == text_column) {
continue;
}
}
Some(_) => {} }
}
}
}
let mut entries: Vec<(Value, Vec<f32>, Option<u64>)> = Vec::with_capacity(store.len());
for &node_index in &store.slot_to_node {
if let Some(node) = graph
.graph
.node_view(petgraph::graph::NodeIndex::new(node_index))
{
if let Some(embedding) = store.get_embedding(node_index) {
let hash = store.text_hashes.get(&node_index).copied();
entries.push((node.id().into_owned(), embedding.to_vec(), hash));
}
}
}
total_embeddings += entries.len();
exported_stores.push(ExportedEmbeddingStore {
node_type: node_type.clone(),
text_column: text_column.to_string(),
dimension: store.dimension,
metric: store.metric.clone(),
model_id: store.model_id.clone(),
entries,
});
}
let file = File::create(path)?;
let mut writer = BufWriter::new(file);
writer.write_all(&KGLE_MAGIC)?;
writer.write_all(&KGLE_VERSION.to_le_bytes())?;
writer.write_all(&[serde_codec::CodecVersion::PostcardV1.tag()])?;
let payload = codec_ser(serde_codec::CodecVersion::PostcardV1, &exported_stores)
.map_err(|e| io::Error::other(format!("Failed to serialize embeddings: {e}")))?;
let mut gz = GzEncoder::new(&mut writer, Compression::new(3));
gz.write_all(&payload)?;
gz.finish()?;
writer.flush()?;
Ok(ExportStats {
stores: exported_stores.len(),
embeddings: total_embeddings,
})
}
pub fn import_embeddings_from_file(graph: &mut DirGraph, path: &str) -> io::Result<ImportStats> {
let file = File::open(path)?;
let mut reader = BufReader::new(file);
let mut buf = Vec::new();
reader.read_to_end(&mut buf)?;
if buf.len() < 8 {
return Err(io::Error::other(
"File is too small to be a valid .kgle file.",
));
}
if buf[..4] != KGLE_MAGIC {
return Err(io::Error::other(
"Not a valid .kgle file (bad magic bytes).",
));
}
let version = u32::from_le_bytes([buf[4], buf[5], buf[6], buf[7]]);
if version > KGLE_VERSION {
return Err(io::Error::other(format!(
"Embedding file version {} is newer than supported version {}. Please upgrade kglite.",
version, KGLE_VERSION,
)));
}
let exported_stores = decode_embedding_file_payload(&buf, version)?;
let mut total_imported = 0usize;
let mut total_skipped = 0usize;
let mut stores_count = 0usize;
let mut dropped_stores = 0usize;
for exported in exported_stores {
graph.build_id_index(&exported.node_type);
let mut store = crate::graph::schema::EmbeddingStore::new(exported.dimension);
store.metric = exported.metric.clone();
store.model_id = exported.model_id.clone();
store
.data
.reserve(exported.entries.len() * exported.dimension);
let mut imported = 0usize;
let mut skipped = 0usize;
for (id, vec, hash) in &exported.entries {
match graph.lookup_by_id(&exported.node_type, id) {
Some(node_idx) => {
store.set_embedding(node_idx.index(), vec);
if let Some(h) = hash {
store.set_text_hash(node_idx.index(), *h);
}
imported += 1;
}
None => {
skipped += 1;
}
}
}
if imported > 0 {
graph.set_embedding_store(&exported.node_type, &exported.text_column, store);
stores_count += 1;
} else if !exported.entries.is_empty() {
dropped_stores += 1;
}
total_imported += imported;
total_skipped += skipped;
}
Ok(ImportStats {
stores: stores_count,
imported: total_imported,
skipped: total_skipped,
dropped_stores,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn fixture_store() -> ExportedEmbeddingStore {
ExportedEmbeddingStore {
node_type: "Doc".to_string(),
text_column: "summary".to_string(),
dimension: 2,
metric: Some("cosine".to_string()),
model_id: Some("fixture".to_string()),
entries: vec![(Value::UniqueId(7), vec![0.25, 0.75], Some(99))],
}
}
fn embedding_file(version: u32, codec_tag: Option<u8>, payload: &[u8]) -> Vec<u8> {
let mut compressed = GzEncoder::new(Vec::new(), Compression::new(3));
compressed.write_all(payload).unwrap();
let compressed = compressed.finish().unwrap();
let mut bytes = Vec::new();
bytes.extend_from_slice(&KGLE_MAGIC);
bytes.extend_from_slice(&version.to_le_bytes());
if let Some(tag) = codec_tag {
bytes.push(tag);
}
bytes.extend_from_slice(&compressed);
bytes
}
#[test]
fn a_payload_whose_watermark_disagrees_with_its_topology_is_skipped() {
use crate::graph::algorithms::hnsw::{HnswMetric, HnswParams};
use crate::graph::schema::EmbeddingStore;
let mut store = EmbeddingStore::new(2);
for slot in 0..4 {
store.set_embedding(slot, &[slot as f32, 1.0]);
}
let index = HnswIndex::build(
&store.data,
&store.norms,
2,
HnswMetric::Cosine,
HnswParams::default(),
3,
);
let entry = PersistedVectorIndexRef {
node_type: "Doc",
embedding_property: "vec_emb",
index: &index,
watermark: 4 + 1,
limit: 1000,
dirty: Vec::new(),
};
let body = codec_ser(serde_codec::CodecVersion::PostcardV1, &vec![entry]).unwrap();
let mut payload = Vec::new();
payload.extend_from_slice(VECTOR_INDEX_MAGIC);
payload.extend_from_slice(&VECTOR_INDEX_FORMAT_VERSION.to_le_bytes());
payload.extend_from_slice(&body);
let mut graph = DirGraph::new();
graph
.embeddings
.insert(("Doc".to_string(), "vec_emb".to_string()), store);
decode_vector_indexes(&payload, &mut graph);
assert!(
!graph.embeddings[&("Doc".to_string(), "vec_emb".to_string())].has_index(),
"a watermark ahead of the topology must be refused"
);
}
#[test]
fn pre_014_embedding_payload_is_rejected() {
let stores = vec![fixture_store()];
let old_payload = codec_ser(serde_codec::CodecVersion::PostcardV1, &stores).unwrap();
let old = embedding_file(2, None, &old_payload);
let error = decode_embedding_file_payload(&old, 2).err().unwrap();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
assert!(error.to_string().contains("pre-0.14"));
}
#[test]
fn postcard_v3_embedding_payload_decodes() {
let stores = vec![fixture_store()];
let postcard_payload = codec_ser(serde_codec::CodecVersion::PostcardV1, &stores).unwrap();
let current = embedding_file(
3,
Some(serde_codec::CodecVersion::PostcardV1.tag()),
&postcard_payload,
);
let decoded = decode_embedding_file_payload(¤t, 3).unwrap();
assert_eq!(decoded.len(), 1);
assert_eq!(decoded[0].node_type, "Doc");
assert_eq!(decoded[0].text_column, "summary");
assert_eq!(decoded[0].dimension, 2);
assert_eq!(decoded[0].metric.as_deref(), Some("cosine"));
assert_eq!(decoded[0].model_id.as_deref(), Some("fixture"));
assert_eq!(decoded[0].entries[0].2, Some(99));
}
#[test]
fn postcard_v3_embedding_payload_requires_its_codec_tag() {
let truncated = [b'K', b'G', b'L', b'E', 3, 0, 0, 0];
assert!(decode_embedding_file_payload(&truncated, 3)
.err()
.unwrap()
.to_string()
.contains("codec tag"));
let invalid = [b'K', b'G', b'L', b'E', 3, 0, 0, 0, 99];
assert!(decode_embedding_file_payload(&invalid, 3)
.err()
.unwrap()
.to_string()
.contains("Invalid .kgle codec tag"));
}
}