use ahash::HashSet;
use anyhow::{Result, bail};
use reblessive::tree::Stk;
use revision::revisioned;
use roaring::RoaringTreemap;
use serde::{Deserialize, Serialize};
use crate::ctx::Context;
use crate::err::Error;
use crate::idx::IndexKeyBase;
use crate::idx::planner::ScanDirection;
use crate::idx::trees::dynamicset::DynamicSet;
use crate::idx::trees::graph::UndirectedGraph;
use crate::idx::trees::hnsw::filter::HnswTruthyDocumentFilter;
use crate::idx::trees::hnsw::heuristic::Heuristic;
use crate::idx::trees::hnsw::index::HnswContext;
use crate::idx::trees::hnsw::{ElementId, HnswElements, HnswSearch, VectorId};
use crate::idx::trees::knn::{DoublePriorityQueue, Ids64};
use crate::idx::trees::vector::SharedVector;
use crate::key::index::hn::HnswNode;
use crate::kvs::Transaction;
#[revisioned(revision = 1)]
#[derive(Default, Debug, Serialize, Deserialize)]
pub(super) struct LayerState {
pub(super) version: u64,
pub(super) chunks: u32,
}
#[derive(Debug)]
pub(super) struct HnswLayer<S>
where
S: DynamicSet,
{
ikb: IndexKeyBase,
level: u16,
graph: UndirectedGraph<S>,
m_max: usize,
}
impl<S> HnswLayer<S>
where
S: DynamicSet,
{
pub(super) fn new(ikb: IndexKeyBase, level: usize, m_max: usize) -> Self {
Self {
ikb,
level: level as u16,
graph: UndirectedGraph::new(m_max + 1),
m_max,
}
}
pub(super) fn m_max(&self) -> usize {
self.m_max
}
pub(super) fn get_edges(&self, e_id: ElementId) -> Option<&S> {
self.graph.get_edges(e_id)
}
pub(super) async fn add_empty_node(
&mut self,
tx: &Transaction,
node: ElementId,
st: &mut LayerState,
) -> Result<bool> {
if !self.graph.add_empty_node(node) {
return Ok(false);
}
self.save_nodes(tx, st, &[node]).await?;
Ok(true)
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn search_single(
&self,
ctx: &HnswContext<'_>,
elements: &HnswElements,
pt: &SharedVector,
ep_dist: f64,
ep_id: ElementId,
ef: usize,
pending_docs: Option<&RoaringTreemap>,
) -> Result<DoublePriorityQueue> {
let visited = HashSet::from_iter([ep_id]);
let candidates = DoublePriorityQueue::from(ep_dist, ep_id);
let w = if pending_docs.is_some() {
let mut w = DoublePriorityQueue::default();
if let Some(ep_pt) = elements.get_vector(&ctx.tx, &ep_id).await?
&& !Self::are_all_docs_in_pending(ctx, ep_id, &ep_pt, pending_docs).await?
{
w.push(ep_dist, ep_id);
}
w
} else {
candidates.clone()
};
self.search(ctx, elements, pt, candidates, visited, w, ef, pending_docs).await
}
pub(super) async fn search_single_with_ignore(
&self,
ctx: &HnswContext<'_>,
elements: &HnswElements,
pt: &SharedVector,
ignore_id: ElementId,
ef: usize,
) -> Result<Option<ElementId>> {
let visited = HashSet::from_iter([ignore_id]);
let mut candidates = DoublePriorityQueue::default();
if let Some(dist) = elements.get_distance(&ctx.tx, pt, &ignore_id).await? {
candidates.push(dist, ignore_id);
}
let w = DoublePriorityQueue::default();
let q = self.search(ctx, elements, pt, candidates, visited, w, ef, None).await?;
Ok(q.peek_first().map(|(_, e_id)| e_id))
}
#[expect(clippy::too_many_arguments)]
pub(super) async fn search_single_with_filter(
&self,
ctx: &HnswContext<'_>,
stk: &mut Stk,
elements: &HnswElements,
search: &HnswSearch,
ep_dist: f64,
ep_id: ElementId,
filter: &mut HnswTruthyDocumentFilter<'_>,
pending_docs: Option<&RoaringTreemap>,
) -> Result<DoublePriorityQueue> {
let visited = HashSet::from_iter([ep_id]);
let candidates = DoublePriorityQueue::from(ep_dist, ep_id);
let mut w = DoublePriorityQueue::default();
Self::add_if_truthy(
ctx,
stk,
search.ef,
&mut w,
&search.pt,
ep_dist,
ep_id,
filter,
pending_docs,
)
.await?;
self.search_with_filter(
ctx,
stk,
elements,
search,
candidates,
visited,
w,
filter,
pending_docs,
)
.await
}
pub(super) async fn search_multi(
&self,
ctx: &HnswContext<'_>,
elements: &HnswElements,
pt: &SharedVector,
candidates: DoublePriorityQueue,
ef: usize,
) -> Result<DoublePriorityQueue> {
let w = candidates.clone();
let visited = w.to_set();
self.search(ctx, elements, pt, candidates, visited, w, ef, None).await
}
pub(super) async fn search_multi_with_ignore(
&self,
ctx: &HnswContext<'_>,
elements: &HnswElements,
pt: &SharedVector,
ignore_ids: Vec<ElementId>,
efc: usize,
) -> Result<DoublePriorityQueue> {
let mut candidates = DoublePriorityQueue::default();
for id in &ignore_ids {
if let Some(dist) = elements.get_distance(&ctx.tx, pt, id).await? {
candidates.push(dist, *id);
}
}
let visited = HashSet::from_iter(ignore_ids);
let w = DoublePriorityQueue::default();
self.search(ctx, elements, pt, candidates, visited, w, efc, None).await
}
#[expect(clippy::too_many_arguments)]
pub(super) async fn search(
&self,
ctx: &HnswContext<'_>,
elements: &HnswElements,
q: &SharedVector,
mut candidates: DoublePriorityQueue, mut visited: HashSet<ElementId>, mut w: DoublePriorityQueue,
ef: usize,
pending_docs: Option<&RoaringTreemap>,
) -> Result<DoublePriorityQueue> {
let mut fq_dist = w.peek_last_dist().unwrap_or(f64::MAX);
while let Some((cq_dist, doc)) = candidates.pop_first() {
if cq_dist > fq_dist {
break;
}
if let Some(neighbourhood) = self.graph.get_edges(doc) {
for &e_id in neighbourhood.iter() {
if !visited.insert(e_id) {
continue;
}
if let Some(e_pt) = elements.get_vector(&ctx.tx, &e_id).await? {
let e_dist = elements.distance(&e_pt, q);
if e_dist < fq_dist || w.len() < ef {
if Self::are_all_docs_in_pending(ctx, e_id, &e_pt, pending_docs).await?
{
continue;
}
candidates.push(e_dist, e_id);
w.push(e_dist, e_id);
if w.len() > ef {
w.pop_last();
}
fq_dist = w.peek_last_dist().unwrap_or(f64::MAX);
}
}
}
}
}
Ok(w)
}
#[expect(clippy::too_many_arguments)]
pub(super) async fn search_with_filter(
&self,
ctx: &HnswContext<'_>,
stk: &mut Stk,
elements: &HnswElements,
search: &HnswSearch,
mut candidates: DoublePriorityQueue,
mut visited: HashSet<ElementId>,
mut w: DoublePriorityQueue,
filter: &mut HnswTruthyDocumentFilter<'_>,
pending_docs: Option<&RoaringTreemap>,
) -> Result<DoublePriorityQueue> {
let mut f_dist = w.peek_last_dist().unwrap_or(f64::MAX);
while let Some((dist, doc)) = candidates.pop_first() {
if dist > f_dist {
break;
}
if let Some(neighbourhood) = self.graph.get_edges(doc) {
self.prefetch_neighbourhood_records(
ctx,
elements,
search,
neighbourhood,
&visited,
w.len(),
f_dist,
filter,
pending_docs,
)
.await?;
for &e_id in neighbourhood.iter() {
if !visited.insert(e_id) {
continue;
}
if let Some(e_pt) = elements.get_vector(&ctx.tx, &e_id).await? {
let e_dist = elements.distance(&e_pt, &search.pt);
if e_dist < f_dist || w.len() < search.ef {
candidates.push(e_dist, e_id);
if Self::add_if_truthy(
ctx,
stk,
search.ef,
&mut w,
&e_pt,
e_dist,
e_id,
filter,
pending_docs,
)
.await?
{
f_dist = w.peek_last_dist().expect("w is non-empty"); }
}
}
}
}
}
Ok(w)
}
#[expect(clippy::too_many_arguments)]
async fn prefetch_neighbourhood_records(
&self,
ctx: &HnswContext<'_>,
elements: &HnswElements,
search: &HnswSearch,
neighbourhood: &S,
visited: &HashSet<ElementId>,
w_len: usize,
f_dist: f64,
filter: &mut HnswTruthyDocumentFilter<'_>,
pending_docs: Option<&RoaringTreemap>,
) -> Result<()> {
let mut ids: Vec<VectorId> = Vec::new();
for &e_id in neighbourhood.iter() {
if visited.contains(&e_id) {
continue;
}
let Some(e_pt) = elements.get_vector(&ctx.tx, &e_id).await? else {
continue;
};
let e_dist = elements.distance(&e_pt, &search.pt);
if !(e_dist < f_dist || w_len < search.ef) {
continue;
}
let Some(docs) = ctx.vec_docs.get_docs_by_element(&ctx.tx, e_id, &e_pt).await? else {
continue;
};
if let Some(pending_docs) = pending_docs
&& Self::check_all_docs_in_pending(&docs, pending_docs)
{
continue;
}
for doc_id in docs.iter() {
ids.push(VectorId::DocId(doc_id));
}
}
filter.prefetch_records(ctx, &ids).await
}
#[expect(clippy::too_many_arguments)]
pub(super) async fn add_if_truthy(
ctx: &HnswContext<'_>,
stk: &mut Stk,
efc: usize,
w: &mut DoublePriorityQueue,
e_pt: &SharedVector,
e_dist: f64,
e_id: ElementId,
filter: &mut HnswTruthyDocumentFilter<'_>,
pending_docs: Option<&RoaringTreemap>,
) -> Result<bool> {
if let Some(docs) = ctx.vec_docs.get_docs_by_element(&ctx.tx, e_id, e_pt).await? {
if let Some(pending_docs) = pending_docs
&& Self::check_all_docs_in_pending(&docs, pending_docs)
{
return Ok(false);
}
if filter.check_any_doc_truthy(ctx, stk, docs).await? {
w.push(e_dist, e_id);
if w.len() > efc {
w.pop_last();
}
return Ok(true);
}
}
Ok(false)
}
fn check_all_docs_in_pending(docs: &Ids64, pending_docs: &RoaringTreemap) -> bool {
if pending_docs.is_empty() {
return false;
}
for doc_id in docs.iter() {
if !pending_docs.contains(doc_id) {
return false;
}
}
true
}
async fn are_all_docs_in_pending(
search_ctx: &HnswContext<'_>,
e_id: ElementId,
e_pt: &SharedVector,
pending_docs: Option<&RoaringTreemap>,
) -> Result<bool> {
let Some(pending_docs) = pending_docs else {
return Ok(false);
};
if pending_docs.is_empty() {
return Ok(false);
}
if let Some(docs) =
search_ctx.vec_docs.get_docs_by_element(&search_ctx.tx, e_id, e_pt).await?
{
for doc_id in docs.iter() {
if !pending_docs.contains(doc_id) {
return Ok(false);
}
}
}
Ok(true)
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn insert(
&mut self,
ctx: &HnswContext<'_>,
st: &mut LayerState,
elements: &HnswElements,
heuristic: &Heuristic,
efc: usize,
(q_id, q_pt): (ElementId, &SharedVector),
mut eps: DoublePriorityQueue,
) -> Result<DoublePriorityQueue> {
let w;
let mut neighbors = self.graph.new_edges();
{
w = self.search_multi(ctx, elements, q_pt, eps, efc).await?;
eps = w.clone();
heuristic.select(&ctx.tx, elements, self, q_id, q_pt, w, None, &mut neighbors).await?;
};
let neighbors = self.graph.add_node_and_bidirectional_edges(q_id, neighbors);
for e_id in &neighbors {
if let Some(e_conn) = self.graph.get_edges(*e_id) {
if e_conn.len() > self.m_max
&& let Some(e_pt) = elements.get_vector(&ctx.tx, e_id).await?
{
let e_c = self.build_priority_list(&ctx.tx, elements, *e_id, e_conn).await?;
let mut e_new_conn = self.graph.new_edges();
heuristic
.select(&ctx.tx, elements, self, *e_id, &e_pt, e_c, None, &mut e_new_conn)
.await?;
#[cfg(debug_assertions)]
assert!(!e_new_conn.contains(e_id));
self.graph.set_node(*e_id, e_new_conn);
}
} else {
#[cfg(debug_assertions)]
unreachable!("Element: {}", e_id);
}
}
let mut changed_nodes = Vec::with_capacity(neighbors.len() + 1);
changed_nodes.push(q_id);
changed_nodes.extend_from_slice(&neighbors);
self.save_nodes(&ctx.tx, st, &changed_nodes).await?;
Ok(eps)
}
async fn build_priority_list(
&self,
tx: &Transaction,
elements: &HnswElements,
e_id: ElementId,
neighbors: &S,
) -> Result<DoublePriorityQueue> {
let mut w = DoublePriorityQueue::default();
if let Some(e_pt) = elements.get_vector(tx, &e_id).await? {
for n_id in neighbors.iter() {
if let Some(n_pt) = elements.get_vector(tx, n_id).await? {
let dist = elements.distance(&e_pt, &n_pt);
w.push(dist, *n_id);
}
}
}
Ok(w)
}
pub(super) async fn remove(
&mut self,
ctx: &HnswContext<'_>,
st: &mut LayerState,
elements: &HnswElements,
heuristic: &Heuristic,
e_id: ElementId,
efc: usize,
) -> Result<bool> {
if let Some(f_ids) = self.graph.remove_node_and_bidirectional_edges(e_id) {
let mut changed_nodes = Vec::with_capacity(f_ids.len());
for &q_id in f_ids.iter() {
if let Some(q_pt) = elements.get_vector(&ctx.tx, &q_id).await? {
let c = self
.search_multi_with_ignore(ctx, elements, &q_pt, vec![q_id, e_id], efc)
.await?;
let mut q_new_conn = self.graph.new_edges();
heuristic
.select(
&ctx.tx,
elements,
self,
q_id,
&q_pt,
c,
Some(e_id),
&mut q_new_conn,
)
.await?;
#[cfg(debug_assertions)]
{
assert!(
!q_new_conn.contains(&q_id),
"!q_new_conn.contains(&q_id) - q_id: {q_id} - f_ids: {q_new_conn:?}"
);
assert!(
!q_new_conn.contains(&e_id),
"!q_new_conn.contains(&e_id) - e_id: {e_id} - f_ids: {q_new_conn:?}"
);
assert!(q_new_conn.len() <= self.m_max);
}
self.graph.set_node(q_id, q_new_conn);
changed_nodes.push(q_id);
}
}
self.delete_node(&ctx.tx, e_id).await?;
self.save_nodes(&ctx.tx, st, &changed_nodes).await?;
Ok(true)
} else {
Ok(false)
}
}
async fn save_nodes(
&self,
tx: &Transaction,
st: &mut LayerState,
nodes: &[ElementId],
) -> Result<()> {
for &node_id in nodes {
if let Some(val) = self.graph.node_to_val(node_id) {
let key = self.ikb.new_hn_key(self.level, node_id);
tx.set(&key, &val).await?;
}
}
st.version += 1;
Ok(())
}
async fn delete_node(&self, tx: &Transaction, node_id: ElementId) -> Result<()> {
let key = self.ikb.new_hn_key(self.level, node_id);
tx.del(&key).await?;
Ok(())
}
pub(super) async fn load(
&mut self,
ctx: &Context,
tx: &Transaction,
st: &mut LayerState,
) -> Result<bool> {
self.graph.clear();
if st.chunks > 0 {
let mut val = Vec::new();
for i in 0..st.chunks {
let key = self.ikb.new_hl_key(self.level, i);
let chunk =
tx.get(&key, None).await?.ok_or_else(|| Error::unreachable("Missing chunk"))?;
val.extend(chunk);
}
self.graph.lecacy_reload(&val)?;
}
let range = self.ikb.new_hn_layer_range(self.level)?;
let mut count = 0;
let mut cursor = tx.open_vals_cursor(range, ScanDirection::Forward, 0, None).await?;
loop {
let batch = cursor.next_batch(crate::kvs::NORMAL_BATCH_SIZE).await?;
if batch.is_empty() {
break;
}
for (k, v) in &batch {
if ctx.is_done(Some(count)).await? {
bail!(Error::QueryCancelled);
}
let key = HnswNode::decode_key(k)?;
self.graph.load_node(key.node, v);
count += 1;
}
}
drop(cursor);
if st.chunks > 0 && tx.writeable() {
for &node_id in &self.graph.node_ids() {
if let Some(node_val) = self.graph.node_to_val(node_id) {
let key = self.ikb.new_hn_key(self.level, node_id);
tx.set(&key, &node_val).await?;
}
}
let hl_range = self.ikb.new_hl_layer_range(self.level)?;
tx.delr(hl_range).await?;
st.chunks = 0;
return Ok(true);
}
Ok(false)
}
}
#[cfg(test)]
impl<S> HnswLayer<S>
where
S: DynamicSet,
{
pub(in crate::idx::trees::hnsw) async fn check_props(&self, elements: &HnswElements) {
let elements_len = elements.len().await;
assert!(self.graph.len() <= elements_len, "{} - {}", self.graph.len(), elements_len);
for (e_id, f_ids) in self.graph.nodes() {
assert!(
f_ids.len() <= self.m_max,
"Foreign list e_id: {e_id} - len = len({}) <= m_layer({})",
self.m_max,
f_ids.len(),
);
assert!(!f_ids.contains(e_id), "!f_ids.contains(e_id) - el: {e_id} - f_ids: {f_ids:?}");
assert!(
elements.contains(*e_id).await,
"h.elements.contains_key(e_id) - el: {e_id} - f_ids: {f_ids:?}"
);
}
}
}