use std::collections::VecDeque;
use std::sync::Arc;
use ahash::HashMap;
use anyhow::{Result, bail};
use reblessive::tree::Stk;
use roaring::RoaringTreemap;
use tokio::sync::RwLock;
use crate::catalog::{Distance, HnswParams, TableId, VectorType};
use crate::ctx::{Context, FrozenContext};
use crate::err::Error;
use crate::idx::planner::ScanDirection;
use crate::idx::planner::iterators::KnnIteratorResult;
use crate::idx::trees::KnnCondFilter;
use crate::idx::trees::hnsw::cache::VectorCache;
use crate::idx::trees::hnsw::docs::{HnswDocs, VecDocs};
use crate::idx::trees::hnsw::filter::HnswTruthyDocumentFilter;
use crate::idx::trees::hnsw::flavor::HnswFlavor;
use crate::idx::trees::hnsw::{
ElementId, HnswRecordPendingUpdate, HnswSearch, VectorId, VectorPendingUpdate,
};
use crate::idx::trees::knn::KnnResultBuilder;
use crate::idx::trees::vector::{SerializedVector, SharedVector, Vector};
use crate::idx::{
IndexKeyBase, bump_compaction_generation, is_transaction_condition_not_met,
read_compaction_generation,
};
use crate::key::index::hr::HnswRecordPending;
use crate::kvs::{KVValue, Key, Transaction, Val};
use crate::val::{Number, RecordId, RecordIdKey, Value};
const HNSW_COMPACTION_MAX_PENDING_KEYS: usize = 1024;
const HNSW_COMPACTION_MAX_PENDING_BYTES: usize = 16 * 1024 * 1024;
struct CapturedPendingKey {
key: Key,
value: Val,
}
#[derive(Clone)]
struct PendingOperation {
id: VectorId,
old_vectors: Vec<SerializedVector>,
new_vectors: Vec<SerializedVector>,
}
pub(crate) struct HnswCompactionPlan {
generation: Option<u64>,
captured_keys: Vec<CapturedPendingKey>,
pending: Vec<PendingOperation>,
has_more: bool,
}
impl HnswCompactionPlan {
pub(crate) fn has_work(&self) -> bool {
!self.captured_keys.is_empty()
}
pub(crate) fn has_more(&self) -> bool {
self.has_more
}
#[cfg(test)]
pub(crate) fn len(&self) -> usize {
self.captured_keys.len()
}
}
struct PendingPlanBuilder {
generation: Option<u64>,
captured_keys: Vec<CapturedPendingKey>,
pending: Vec<PendingOperation>,
pending_by_id: HashMap<VectorId, usize>,
encoded_bytes: usize,
has_more: bool,
}
impl PendingPlanBuilder {
fn new(generation: Option<u64>) -> Self {
Self {
generation,
captured_keys: Vec::new(),
pending: Vec::new(),
pending_by_id: HashMap::default(),
encoded_bytes: 0,
has_more: false,
}
}
fn add(&mut self, key: Key, value: Val, pending: PendingOperation) -> bool {
if self.captured_keys.len() >= HNSW_COMPACTION_MAX_PENDING_KEYS
|| (!self.captured_keys.is_empty()
&& self.encoded_bytes + key.len() + value.len() > HNSW_COMPACTION_MAX_PENDING_BYTES)
{
self.has_more = true;
return false;
}
self.encoded_bytes += key.len() + value.len();
self.captured_keys.push(CapturedPendingKey {
key,
value,
});
self.add_pending(pending);
if self.captured_keys.len() >= HNSW_COMPACTION_MAX_PENDING_KEYS
|| self.encoded_bytes >= HNSW_COMPACTION_MAX_PENDING_BYTES
{
self.has_more = true;
}
true
}
fn add_pending(&mut self, pending: PendingOperation) {
if let Some(pos) = self.pending_by_id.get(&pending.id) {
self.pending[*pos].new_vectors = pending.new_vectors;
return;
}
let pos = self.pending.len();
self.pending_by_id.insert(pending.id.clone(), pos);
self.pending.push(pending);
}
fn into_plan(self) -> HnswCompactionPlan {
HnswCompactionPlan {
generation: self.generation,
captured_keys: self.captured_keys,
pending: self.pending,
has_more: self.has_more,
}
}
}
pub(crate) struct HnswIndex {
dim: usize,
distance: Distance,
table_id: TableId,
ikb: IndexKeyBase,
vector_type: VectorType,
vector_cache: VectorCache,
hnsw: RwLock<HnswFlavor>,
vec_docs: VecDocs,
}
pub(super) struct HnswContext<'a> {
pub(super) ctx: &'a FrozenContext,
pub(super) tx: Arc<Transaction>,
pub(super) ikb: IndexKeyBase,
pub(super) vec_docs: &'a VecDocs,
}
impl<'a> HnswContext<'a> {
pub(super) fn new(ctx: &'a FrozenContext, ikb: IndexKeyBase, vec_docs: &'a VecDocs) -> Self {
Self {
ctx,
tx: ctx.tx(),
ikb,
vec_docs,
}
}
}
impl HnswIndex {
pub(crate) async fn new(
vector_cache: VectorCache,
_tx: &Transaction,
ikb: IndexKeyBase,
tb: TableId,
p: &HnswParams,
) -> Result<Self> {
Ok(Self {
dim: p.dimension as usize,
vector_type: p.vector_type,
distance: p.distance.clone(),
table_id: tb,
hnsw: RwLock::new(HnswFlavor::new(tb, ikb.clone(), p, vector_cache.clone())?),
vec_docs: VecDocs::new(ikb.clone(), tb, vector_cache.clone(), p.use_hashed_vector),
vector_cache,
ikb,
})
}
fn content_to_vectors(&self, content: Vec<Value>) -> Result<Vec<SerializedVector>> {
let mut vectors = Vec::with_capacity(content.len());
for value in content.into_iter().filter(|v| !v.is_nullish()) {
let vector = SerializedVector::try_from_value(self.vector_type, self.dim, value)?;
Vector::check_expected_dimension(vector.dimension(), self.dim)?;
vectors.push(vector);
}
Ok(vectors)
}
pub(crate) async fn index(
&self,
ctx: &Context,
id: &RecordIdKey,
old_values: Option<Vec<Value>>,
new_values: Option<Vec<Value>>,
) -> Result<()> {
if old_values.is_none() && new_values.is_none() {
return Ok(());
}
let old_vectors = if let Some(v) = old_values {
self.content_to_vectors(v)?
} else {
vec![]
};
let new_vectors = if let Some(v) = new_values {
self.content_to_vectors(v)?
} else {
vec![]
};
let tx = ctx.tx();
let key = self.ikb.new_hr_key(id);
let pending = if let Some(mut pending) = tx.get(&key, None).await? {
pending.new_vectors = new_vectors;
pending
} else {
HnswRecordPendingUpdate {
doc_id: HnswDocs::get_doc_id(&self.ikb, &tx, id).await?,
old_vectors,
new_vectors,
}
};
tx.set(&key, &pending).await?;
Ok(())
}
fn append_pending_to_operation(pending: VectorPendingUpdate) -> PendingOperation {
PendingOperation {
id: pending.id,
old_vectors: pending.old_vectors,
new_vectors: pending.new_vectors,
}
}
fn record_pending_to_operation(
id: RecordIdKey,
pending: HnswRecordPendingUpdate,
) -> PendingOperation {
let id = if let Some(doc_id) = pending.doc_id {
VectorId::DocId(doc_id)
} else {
VectorId::RecordKey(Arc::new(id))
};
PendingOperation {
id,
old_vectors: pending.old_vectors,
new_vectors: pending.new_vectors,
}
}
pub(super) fn new_hnsw_context<'a>(&'a self, ctx: &'a FrozenContext) -> HnswContext<'a> {
HnswContext::new(ctx, self.ikb.clone(), &self.vec_docs)
}
pub(in crate::idx) async fn prepare_compaction(
ctx: &FrozenContext,
ikb: &IndexKeyBase,
) -> Result<HnswCompactionPlan> {
let tx = ctx.tx();
let generation = read_compaction_generation(&tx, &ikb.new_hg_key()).await?;
let mut builder = PendingPlanBuilder::new(generation);
let mut count = 0;
Self::collect_append_pending_for_plan(ctx, &tx, ikb, &mut builder, &mut count).await?;
if !builder.has_more {
Self::collect_record_pending_for_plan(ctx, &tx, ikb, &mut builder, &mut count).await?;
}
Ok(builder.into_plan())
}
async fn collect_append_pending_for_plan(
ctx: &FrozenContext,
tx: &Transaction,
ikb: &IndexKeyBase,
builder: &mut PendingPlanBuilder,
count: &mut usize,
) -> Result<()> {
let rng = ikb.new_hp_range()?;
let mut cursor = tx.open_vals_cursor(rng, ScanDirection::Forward, 0, None).await?;
loop {
let batch = cursor.next_batch(crate::kvs::NORMAL_BATCH_SIZE).await?;
if batch.is_empty() {
break;
}
let owned: Vec<(Vec<u8>, Vec<u8>)> =
batch.iter().map(|(k, v)| (k.to_vec(), v.to_vec())).collect();
for (key, value) in owned {
if ctx.is_done(Some(*count)).await? {
bail!(Error::QueryCancelled)
}
let pending = VectorPendingUpdate::kv_decode_value(&value, ())?;
let pending = Self::append_pending_to_operation(pending);
if !builder.add(key, value, pending) {
return Ok(());
}
*count += 1;
if builder.has_more {
return Ok(());
}
}
}
Ok(())
}
async fn collect_record_pending_for_plan(
ctx: &FrozenContext,
tx: &Transaction,
ikb: &IndexKeyBase,
builder: &mut PendingPlanBuilder,
count: &mut usize,
) -> Result<()> {
let rng = ikb.new_hr_range()?;
let mut cursor = tx.open_vals_cursor(rng, ScanDirection::Forward, 0, None).await?;
loop {
let batch = cursor.next_batch(crate::kvs::NORMAL_BATCH_SIZE).await?;
if batch.is_empty() {
break;
}
let owned: Vec<(Vec<u8>, Vec<u8>)> =
batch.iter().map(|(k, v)| (k.to_vec(), v.to_vec())).collect();
for (key, value) in owned {
if ctx.is_done(Some(*count)).await? {
bail!(Error::QueryCancelled)
}
let hr = HnswRecordPending::decode_key(&key)?;
let pending = HnswRecordPendingUpdate::kv_decode_value(&value, ())?;
let pending = Self::record_pending_to_operation(hr.id.into_owned(), pending);
if !builder.add(key, value, pending) {
return Ok(());
}
*count += 1;
if builder.has_more {
return Ok(());
}
}
}
Ok(())
}
pub(in crate::idx) async fn apply_compaction(
&self,
ctx: &FrozenContext,
plan: HnswCompactionPlan,
) -> Result<bool> {
let HnswCompactionPlan {
generation,
captured_keys,
pending,
has_more: _,
} = plan;
let tx = ctx.tx();
if captured_keys.is_empty() {
return Ok(false);
}
if !bump_compaction_generation(&tx, &self.ikb.new_hg_key(), generation).await? {
return Ok(false);
}
for captured in &captured_keys {
match tx.delc(&captured.key, Some(&captured.value)).await {
Ok(()) => {}
Err(e) if is_transaction_condition_not_met(&e) => return Ok(false),
Err(e) => return Err(e),
}
}
let mut hnsw = self.hnsw.write().await;
hnsw.check_state(ctx).await?;
let mut ctx = self.new_hnsw_context(ctx);
let mut docs = HnswDocs::new(&tx, self.ikb.clone()).await?;
for pending in pending {
self.apply_pending_operation(&mut ctx, &mut docs, &mut hnsw, pending).await?;
}
docs.finish(&tx).await?;
Ok(true)
}
#[cfg(test)]
pub(in crate::idx) async fn index_pendings(&self, ctx: &FrozenContext) -> Result<usize> {
let plan = Self::prepare_compaction(ctx, &self.ikb).await?;
let count = plan.len();
if self.apply_compaction(ctx, plan).await? {
Ok(count)
} else {
Ok(0)
}
}
async fn apply_pending_operation(
&self,
ctx: &mut HnswContext<'_>,
docs: &mut HnswDocs,
hnsw: &mut HnswFlavor,
pending: PendingOperation,
) -> Result<()> {
match pending.id {
VectorId::DocId(doc_id) => {
for vector in pending.old_vectors {
let vector = Vector::from(vector);
self.vec_docs.remove(ctx, &vector, doc_id, hnsw).await?;
}
if pending.new_vectors.is_empty() {
docs.remove(&ctx.tx, doc_id, self.table_id, &self.vector_cache).await?;
} else {
for vector in pending.new_vectors {
let vector = Vector::from(vector);
self.vec_docs.insert(ctx, vector, doc_id, hnsw).await?;
}
}
}
VectorId::RecordKey(id) => {
if !pending.new_vectors.is_empty() {
let doc_id = docs.resolve(&ctx.tx, &id).await?;
for vector in pending.new_vectors {
let vector = Vector::from(vector);
self.vec_docs.insert(ctx, vector, doc_id, hnsw).await?;
}
}
}
}
Ok(())
}
pub(crate) async fn check_state(&self, ctx: &FrozenContext) -> Result<()> {
{
let guard = self.hnsw.read().await;
if !guard.needs_state_reload(ctx).await? {
return Ok(());
}
}
let mut guard = self.hnsw.write().await;
if guard.needs_state_reload(ctx).await? {
guard.check_state(ctx).await?;
}
Ok(())
}
pub(crate) async fn knn_search(
&self,
ctx: &FrozenContext,
stk: &mut Stk,
pt: &[Number],
k: usize,
ef: usize,
cond_filter: Option<KnnCondFilter<'_>>,
) -> Result<VecDeque<KnnIteratorResult>> {
let compaction_generation =
read_compaction_generation(&ctx.tx(), &self.ikb.new_hg_key()).await?;
let mut filter = cond_filter.map(|f| {
HnswTruthyDocumentFilter::new(
f.opt,
self.ikb.clone(),
self.table_id,
self.vector_cache.clone(),
f.cond,
compaction_generation,
f.select_gate,
)
});
let vector: SharedVector = Vector::try_from_vector(self.vector_type, pt)?.into();
vector.check_dimension(self.dim)?;
let search = HnswSearch::new(vector, k, ef);
let ctx = self.new_hnsw_context(ctx);
let mut builder = KnnResultBuilder::new(k);
let pending_docs =
self.search_pendings(&ctx, stk, &search, &mut filter, &mut builder).await?;
self.search_graph(&ctx, stk, &search, pending_docs, &mut filter, &mut builder).await?;
let result = builder.collect();
let cache = if let Some(filter) = filter {
let cache = filter.release();
Some(cache)
} else {
None
};
let mut res_by_pos = vec![None; result.len()];
let mut doc_misses = Vec::new();
for (pos, (dist, id)) in result.into_iter().enumerate() {
let dist: f64 = dist.into();
if let Some(cache) = &cache
&& let Some(Some((rid, record))) = cache.get(&id)
{
res_by_pos[pos] = Some((Arc::clone(rid), dist, Some(Arc::clone(record))));
continue;
}
match id {
VectorId::DocId(doc_id) => {
doc_misses.push((pos, doc_id, dist));
}
VectorId::RecordKey(key) => {
let rid = RecordId::new(self.ikb.table().clone(), key.as_ref().clone());
res_by_pos[pos] = Some((Arc::new(rid), dist, None));
}
}
}
if !doc_misses.is_empty() {
let doc_ids: Vec<_> = doc_misses.iter().map(|(_, doc_id, _)| *doc_id).collect();
let rids = HnswDocs::get_things_batch(
&ctx.ikb,
self.table_id,
&self.vector_cache,
&ctx.tx,
&doc_ids,
compaction_generation,
)
.await?;
for ((pos, _, dist), rid) in doc_misses.into_iter().zip(rids) {
if let Some(rid) = rid {
res_by_pos[pos] = Some((rid, dist, None));
}
}
}
let mut res = VecDeque::with_capacity(res_by_pos.len());
res.extend(res_by_pos.into_iter().flatten());
Ok(res)
}
pub(super) async fn search_graph(
&self,
ctx: &HnswContext<'_>,
stk: &mut Stk,
search: &HnswSearch,
pending_docs: Option<RoaringTreemap>,
filter: &mut Option<HnswTruthyDocumentFilter<'_>>,
builder: &mut KnnResultBuilder,
) -> Result<()> {
let hnsw = self.hnsw.read().await;
if let Some(filter) = filter {
let neighbours = hnsw
.knn_search_with_filter(ctx, search, stk, filter, pending_docs.as_ref())
.await?;
self.add_graph_results(
&ctx.tx,
&hnsw,
neighbours,
pending_docs.as_ref(),
builder,
|evicted_docs| filter.expires(&evicted_docs),
)
.await
} else {
let neighbours = hnsw.knn_search(ctx, search, pending_docs.as_ref()).await?;
self.add_graph_results(
&ctx.tx,
&hnsw,
neighbours,
pending_docs.as_ref(),
builder,
|_| {},
)
.await
}
}
async fn search_pendings(
&self,
ctx: &HnswContext<'_>,
stk: &mut Stk,
search: &HnswSearch,
filter: &mut Option<HnswTruthyDocumentFilter<'_>>,
builder: &mut KnnResultBuilder,
) -> Result<Option<RoaringTreemap>> {
let mut all_existing_docs = RoaringTreemap::new();
let mut non_deleted_docs = HashMap::default();
self.collect_pending(ctx.ctx, &ctx.tx, |pending| {
if let VectorId::DocId(doc_id) = &pending.id {
all_existing_docs.insert(*doc_id);
};
if pending.new_vectors.is_empty() {
non_deleted_docs.remove(&pending.id);
} else {
non_deleted_docs.insert(pending.id, pending.new_vectors);
}
})
.await?;
if all_existing_docs.is_empty() && non_deleted_docs.is_empty() {
return Ok(None);
}
if let Some(filter) = filter.as_mut() {
let ids: Vec<VectorId> = non_deleted_docs.keys().cloned().collect();
filter.prefetch_records(ctx, &ids).await?;
}
for (id, vectors) in non_deleted_docs {
if let Some(filter) = filter
&& !filter.check_vector_id_truthy(ctx, stk, id.clone()).await?
{
continue;
}
for vector in vectors {
let vector = Vector::from(vector);
let d = self.distance.calculate(&search.pt, &vector);
if builder.check_add(d)
&& let Some(evicted_id) = builder.add_vector_id_result(d, id.clone())
&& let Some(filter) = filter
{
filter.expire(&evicted_id);
}
}
}
if all_existing_docs.is_empty() {
return Ok(None);
}
Ok(Some(all_existing_docs))
}
async fn collect_pending<F>(
&self,
ctx: &Context,
tx: &Transaction,
mut collector: F,
) -> Result<()>
where
F: FnMut(PendingOperation),
{
let rng = self.ikb.new_hp_range()?;
let mut cursor = tx.open_vals_cursor(rng, ScanDirection::Forward, 0, None).await?;
let mut count = 0;
loop {
let batch = cursor.next_batch(crate::kvs::NORMAL_BATCH_SIZE).await?;
if batch.is_empty() {
break;
}
for (_, v) in &batch {
if ctx.is_done(Some(count)).await? {
bail!(Error::QueryCancelled)
}
let pending = VectorPendingUpdate::kv_decode_value(v, ())?;
collector(Self::append_pending_to_operation(pending));
count += 1;
}
}
drop(cursor);
let rng = self.ikb.new_hr_range()?;
let mut cursor = tx.open_vals_cursor(rng, ScanDirection::Forward, 0, None).await?;
loop {
let batch = cursor.next_batch(crate::kvs::NORMAL_BATCH_SIZE).await?;
if batch.is_empty() {
break;
}
for (key, value) in &batch {
if ctx.is_done(Some(count)).await? {
bail!(Error::QueryCancelled)
}
let hr = HnswRecordPending::decode_key(key)?;
let pending = HnswRecordPendingUpdate::kv_decode_value(value, ())?;
collector(Self::record_pending_to_operation(hr.id.into_owned(), pending));
count += 1;
}
}
Ok(())
}
async fn add_graph_results<F>(
&self,
tx: &Transaction,
hnsw: &HnswFlavor,
neighbors: Vec<(f64, ElementId)>,
pending_docs: Option<&RoaringTreemap>,
builder: &mut KnnResultBuilder,
mut evicted_docs_func: F,
) -> Result<()>
where
F: FnMut(Vec<VectorId>),
{
for (e_dist, e_id) in neighbors {
if !builder.check_add(e_dist) {
continue;
}
let docs = if let Some(docs) = self.vec_docs.get_cached_doc_set(e_id).await {
Some(docs)
} else if let Some(v) = hnsw.get_vector(tx, &e_id).await? {
self.vec_docs.get_docs_by_element(tx, e_id, &v).await?
} else {
None
};
if let Some(docs) = docs {
let evicted_docs = if let Some(pending_docs) = pending_docs {
let mut evicted_docs = Vec::with_capacity(1);
for doc_id in docs.iter() {
if pending_docs.contains(doc_id) {
continue;
}
if let Some(evicted_id) =
builder.add_vector_id_result(e_dist, VectorId::DocId(doc_id))
{
evicted_docs.push(evicted_id);
}
}
evicted_docs
} else {
builder.add_graph_result(e_dist, &docs)
};
evicted_docs_func(evicted_docs);
}
}
Ok(())
}
#[cfg(test)]
pub(super) async fn check_hnsw_properties(&self, expected_count: usize) {
self.hnsw.read().await.check_hnsw_properties(expected_count).await
}
}