use std::collections::HashMap;
use nodedb_fts::posting::Bm25Params;
use nodedb_fts::search::query_parser::parse_query;
use nodedb_types::Surrogate;
use crate::data::executor::core_loop::CoreLoop;
use crate::data::executor::handlers::transaction::overlay::{Staged, TxnOverlay};
use crate::engine::document::store::surrogate_to_doc_id;
use crate::types::{DatabaseId, TenantId, TxnId};
pub(in crate::data::executor) struct FtsMergeParams<'a> {
pub txn_id: TxnId,
pub database_id: DatabaseId,
pub tid: TenantId,
pub collection: &'a str,
pub query: &'a str,
pub top_k: usize,
}
impl CoreLoop {
pub(in crate::data::executor) fn merge_fts_overlay_into_results(
&self,
params: FtsMergeParams<'_>,
base_results: &mut Vec<(Surrogate, f32, bool)>,
) {
let FtsMergeParams {
txn_id,
database_id,
tid,
collection,
query,
top_k,
} = params;
let coll_key = (database_id, tid, collection.to_string());
self.touch_overlay(txn_id);
let Some(overlay) = self.txn_overlays.get(&txn_id) else {
return;
};
let Some((positive_terms, negative_terms)) =
self.analyze_query_terms(database_id.as_u64(), tid, collection, query)
else {
return;
};
if positive_terms.is_empty() {
remove_tombstoned(overlay, &coll_key, base_results);
return;
}
let config_key = (database_id, tid, collection.to_string());
let bm25_params = Bm25Params::default();
let ctx = self.staged_score_ctx(database_id, tid, collection, &config_key, &bm25_params);
let mut seen: HashMap<u32, usize> = base_results
.iter()
.enumerate()
.map(|(idx, (s, _, _))| (s.as_u32(), idx))
.collect();
for (surrogate, staged) in overlay.iter_for_collection(&coll_key) {
match staged {
Staged::Tombstone => {
if let Some(idx) = seen.remove(&surrogate) {
base_results.remove(idx);
reindex_after_removal(&mut seen, idx);
}
}
Staged::Put(body) => {
let score =
self.score_staged_fts_doc(&ctx, body, &positive_terms, &negative_terms);
match (score, seen.get(&surrogate).copied()) {
(Some(s), Some(idx)) => {
base_results[idx].1 = s;
base_results[idx].2 = false;
}
(Some(s), None) => {
seen.insert(surrogate, base_results.len());
base_results.push((Surrogate::new(surrogate), s, false));
}
(None, Some(idx)) => {
base_results.remove(idx);
seen.remove(&surrogate);
reindex_after_removal(&mut seen, idx);
}
(None, None) => {}
}
}
}
}
base_results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
base_results.truncate(top_k);
}
pub(in crate::data::executor) fn merge_fts_phrase_overlay_into_results(
&self,
params: FtsMergeParams<'_>,
terms: &[String],
base_results: &mut Vec<(Surrogate, f32, bool)>,
) {
let FtsMergeParams {
txn_id,
database_id,
tid,
collection,
query: _query,
top_k,
} = params;
let coll_key = (database_id, tid, collection.to_string());
self.touch_overlay(txn_id);
let Some(overlay) = self.txn_overlays.get(&txn_id) else {
return;
};
let db_u64 = database_id.as_u64();
let phrase_terms: Vec<String> = terms
.iter()
.map(|t| {
self.inverted
.analyze_for_collection(db_u64, tid, collection, t)
.ok()
.and_then(|tokens| tokens.into_iter().next())
.unwrap_or_else(|| t.clone())
})
.collect();
let config_key = (database_id, tid, collection.to_string());
let mut seen: HashMap<u32, usize> = base_results
.iter()
.enumerate()
.map(|(idx, (s, _, _))| (s.as_u32(), idx))
.collect();
for (surrogate, staged) in overlay.iter_for_collection(&coll_key) {
match staged {
Staged::Tombstone => {
if let Some(idx) = seen.remove(&surrogate) {
base_results.remove(idx);
reindex_after_removal(&mut seen, idx);
}
}
Staged::Put(body) => {
let score =
self.score_staged_phrase_doc(db_u64, &config_key, body, &phrase_terms);
match (score, seen.get(&surrogate).copied()) {
(Some(s), Some(idx)) => {
base_results[idx].1 = s;
base_results[idx].2 = false;
}
(Some(s), None) => {
seen.insert(surrogate, base_results.len());
base_results.push((Surrogate::new(surrogate), s, false));
}
(None, Some(idx)) => {
base_results.remove(idx);
seen.remove(&surrogate);
reindex_after_removal(&mut seen, idx);
}
(None, None) => {}
}
}
}
}
base_results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
base_results.truncate(top_k);
}
pub(in crate::data::executor) fn merge_fts_overlay_into_score_map(
&self,
params: FtsMergeParams<'_>,
base_results: &mut HashMap<Surrogate, f32>,
) {
let FtsMergeParams {
txn_id,
database_id,
tid,
collection,
query,
top_k: _top_k,
} = params;
let coll_key = (database_id, tid, collection.to_string());
self.touch_overlay(txn_id);
let Some(overlay) = self.txn_overlays.get(&txn_id) else {
return;
};
let Some((positive_terms, negative_terms)) =
self.analyze_query_terms(database_id.as_u64(), tid, collection, query)
else {
return;
};
let config_key = (database_id, tid, collection.to_string());
let bm25_params = Bm25Params::default();
let ctx = self.staged_score_ctx(database_id, tid, collection, &config_key, &bm25_params);
for (surrogate, staged) in overlay.iter_for_collection(&coll_key) {
match staged {
Staged::Tombstone => {
base_results.remove(&Surrogate::new(surrogate));
}
Staged::Put(body) if !positive_terms.is_empty() => {
match self.score_staged_fts_doc(&ctx, body, &positive_terms, &negative_terms) {
Some(s) => {
base_results.insert(Surrogate::new(surrogate), s);
}
None => {
base_results.remove(&Surrogate::new(surrogate));
}
}
}
Staged::Put(_) => {
base_results.remove(&Surrogate::new(surrogate));
}
}
}
}
pub(in crate::data::executor) fn merge_fts_rows_from_score_map(
&self,
params: FtsMergeParams<'_>,
rows: &mut Vec<(String, Vec<u8>)>,
score_map: &HashMap<Surrogate, f32>,
) {
let FtsMergeParams {
txn_id,
database_id,
tid,
collection,
..
} = params;
let coll_key = (database_id, tid, collection.to_string());
self.touch_overlay(txn_id);
let Some(overlay) = self.txn_overlays.get(&txn_id) else {
return;
};
let mut seen: std::collections::HashSet<u32> = rows
.iter()
.filter_map(|(k, _)| u32::from_str_radix(k, 16).ok())
.collect();
rows.retain_mut(|(row_key, body)| {
let Ok(surrogate) = u32::from_str_radix(row_key, 16) else {
return true;
};
match overlay.get(&coll_key, surrogate) {
Some(Staged::Tombstone) => false,
Some(Staged::Put(staged_body)) => {
*body = staged_body.clone();
score_map.contains_key(&Surrogate::new(surrogate))
}
None => true,
}
});
for (surrogate, staged) in overlay.iter_for_collection(&coll_key) {
if seen.contains(&surrogate) {
continue;
}
if let Staged::Put(body) = staged
&& score_map.contains_key(&Surrogate::new(surrogate))
{
rows.push((surrogate_to_doc_id(Surrogate::new(surrogate)), body.clone()));
seen.insert(surrogate);
}
}
}
fn analyze_query_terms(
&self,
database_id: u64,
tid: TenantId,
collection: &str,
query: &str,
) -> Option<(Vec<String>, Vec<String>)> {
let parsed = parse_query(query).ok()?;
let positive_terms = self
.inverted
.analyze_for_collection(database_id, tid, collection, &parsed.positive.join(" "))
.unwrap_or_default();
let negative_terms = self
.inverted
.analyze_for_collection(database_id, tid, collection, &parsed.negative.join(" "))
.unwrap_or_default();
Some((positive_terms, negative_terms))
}
}
fn remove_tombstoned(
overlay: &TxnOverlay,
coll_key: &(DatabaseId, TenantId, String),
base_results: &mut Vec<(Surrogate, f32, bool)>,
) {
let tombstoned: std::collections::HashSet<u32> = overlay
.iter_for_collection(coll_key)
.filter(|(_, staged)| matches!(staged, Staged::Tombstone))
.map(|(surrogate, _)| surrogate)
.collect();
if tombstoned.is_empty() {
return;
}
base_results.retain(|(s, _, _)| !tombstoned.contains(&s.as_u32()));
}
fn reindex_after_removal(seen: &mut HashMap<u32, usize>, removed_idx: usize) {
for idx in seen.values_mut() {
if *idx > removed_idx {
*idx -= 1;
}
}
}