use crate::embedding::colbert::ColbertEmbedder;
use crate::types::{DENSE_TOMBSTONE, ScoredChunkId};
use anyhow::Result;
use next_plaid::{MmapIndex, SearchParameters};
use std::path::Path;
const DEFAULT_N_IVF_PROBE: usize = 32;
const DEFAULT_CENTROID_THRESHOLD: f32 = 0.4;
fn resolve_plaid_probe<F>(env: F) -> usize
where
F: Fn(&str) -> Option<String>,
{
env("SEMANTEX_PLAID_PROBE")
.and_then(|v| v.parse().ok())
.unwrap_or(DEFAULT_N_IVF_PROBE)
}
fn resolve_plaid_centroid_threshold<F>(env: F) -> Option<f32>
where
F: Fn(&str) -> Option<String>,
{
let Some(raw) = env("SEMANTEX_PLAID_CENTROID_THRESHOLD") else {
return Some(DEFAULT_CENTROID_THRESHOLD);
};
let trimmed = raw.trim();
if matches!(
trimmed.to_ascii_lowercase().as_str(),
"none" | "off" | "disabled" | "0"
) {
return None;
}
if let Ok(v) = trimmed.parse::<f32>() {
Some(v)
} else {
tracing::warn!(
env_value = %raw,
"SEMANTEX_PLAID_CENTROID_THRESHOLD={raw} unparseable; using default {DEFAULT_CENTROID_THRESHOLD}"
);
Some(DEFAULT_CENTROID_THRESHOLD)
}
}
fn translate_chunk_subset_to_doc_subset(doc_to_chunk: &[u64], chunk_id_subset: &[u64]) -> Vec<i64> {
let chunk_set: std::collections::HashSet<u64> = chunk_id_subset.iter().copied().collect();
doc_to_chunk
.iter()
.enumerate()
.filter_map(|(doc_idx, &cid)| {
if cid != DENSE_TOMBSTONE && chunk_set.contains(&cid) {
Some(doc_idx as i64)
} else {
None
}
})
.collect()
}
pub struct PlaidSearcher {
index: MmapIndex,
doc_to_chunk: Vec<u64>,
}
impl PlaidSearcher {
pub fn open(index_dir: &Path, mapping_path: &Path) -> Result<Self> {
let index = MmapIndex::load(&index_dir.to_string_lossy())?;
let mapping_bytes = std::fs::read(mapping_path)?;
let doc_to_chunk: Vec<u64> = postcard::from_bytes(&mapping_bytes)?;
Ok(Self {
index,
doc_to_chunk,
})
}
pub fn search(
&self,
embedder: &ColbertEmbedder,
query: &str,
top_k: usize,
) -> Result<Vec<ScoredChunkId>> {
let query_emb = embedder.encode_query(query)?;
let n_ivf_probe = resolve_plaid_probe(|k| std::env::var(k).ok());
let centroid_score_threshold = resolve_plaid_centroid_threshold(|k| std::env::var(k).ok());
let params = SearchParameters {
top_k,
n_ivf_probe,
centroid_score_threshold,
..Default::default()
};
let results = self.index.search(&query_emb, ¶ms, None)?;
let mut scored: Vec<ScoredChunkId> = results
.passage_ids
.iter()
.zip(results.scores.iter())
.filter_map(|(&doc_id, &score)| {
let doc_idx = doc_id as usize;
self.doc_to_chunk
.get(doc_idx)
.filter(|&&cid| cid != DENSE_TOMBSTONE)
.map(|&chunk_id| ScoredChunkId::new(chunk_id, score))
})
.collect();
scored.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
Ok(scored)
}
pub fn search_with_subset(
&self,
embedder: &ColbertEmbedder,
query: &str,
top_k: usize,
chunk_id_subset: Option<&[u64]>,
) -> Result<Vec<ScoredChunkId>> {
let query_emb = embedder.encode_query(query)?;
let n_ivf_probe = resolve_plaid_probe(|k| std::env::var(k).ok());
let centroid_score_threshold = resolve_plaid_centroid_threshold(|k| std::env::var(k).ok());
let params = SearchParameters {
top_k,
n_ivf_probe,
centroid_score_threshold,
..Default::default()
};
let plaid_subset: Option<Vec<i64>> = chunk_id_subset
.map(|chunks| translate_chunk_subset_to_doc_subset(&self.doc_to_chunk, chunks));
if matches!(plaid_subset.as_ref(), Some(s) if s.is_empty()) {
return Ok(Vec::new());
}
let results = self
.index
.search(&query_emb, ¶ms, plaid_subset.as_deref())?;
let mut scored: Vec<ScoredChunkId> = results
.passage_ids
.iter()
.zip(results.scores.iter())
.filter_map(|(&doc_id, &score)| {
let doc_idx = doc_id as usize;
self.doc_to_chunk
.get(doc_idx)
.filter(|&&cid| cid != DENSE_TOMBSTONE)
.map(|&chunk_id| ScoredChunkId::new(chunk_id, score))
})
.collect();
scored.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
Ok(scored)
}
pub fn doc_to_chunk(&self) -> &[u64] {
&self.doc_to_chunk
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn resolve_plaid_probe_uses_default_when_env_unset() {
let probe = resolve_plaid_probe(|_| None);
assert_eq!(probe, DEFAULT_N_IVF_PROBE);
assert_eq!(probe, 32, "DEFAULT_N_IVF_PROBE must be 32 per v0.4 spec");
}
#[test]
fn resolve_plaid_probe_honors_env_override() {
let probe = resolve_plaid_probe(|k| {
if k == "SEMANTEX_PLAID_PROBE" {
Some("8".to_string())
} else {
None
}
});
assert_eq!(probe, 8);
}
#[test]
fn resolve_plaid_probe_ignores_unparseable_env() {
let probe = resolve_plaid_probe(|k| {
if k == "SEMANTEX_PLAID_PROBE" {
Some("not-a-number".to_string())
} else {
None
}
});
assert_eq!(probe, DEFAULT_N_IVF_PROBE);
}
#[test]
fn resolve_plaid_centroid_threshold_uses_default_when_env_unset() {
let thr = resolve_plaid_centroid_threshold(|_| None);
assert_eq!(thr, Some(DEFAULT_CENTROID_THRESHOLD));
}
#[test]
fn resolve_plaid_centroid_threshold_honors_env_override() {
let thr = resolve_plaid_centroid_threshold(|k| {
if k == "SEMANTEX_PLAID_CENTROID_THRESHOLD" {
Some("0.25".to_string())
} else {
None
}
});
assert!(matches!(thr, Some(v) if (v - 0.25).abs() < 1e-6));
}
#[test]
fn resolve_plaid_centroid_threshold_supports_explicit_none() {
for token in &["none", "NONE", " None ", "off", "OFF", "disabled", "0"] {
let thr = resolve_plaid_centroid_threshold(|k| {
if k == "SEMANTEX_PLAID_CENTROID_THRESHOLD" {
Some((*token).to_string())
} else {
None
}
});
assert!(
thr.is_none(),
"SEMANTEX_PLAID_CENTROID_THRESHOLD={token:?} must opt out (None), got {thr:?}",
);
}
}
#[test]
fn resolve_plaid_centroid_threshold_unparseable_falls_back_to_default() {
let thr = resolve_plaid_centroid_threshold(|k| {
if k == "SEMANTEX_PLAID_CENTROID_THRESHOLD" {
Some("not-a-float".to_string())
} else {
None
}
});
assert_eq!(thr, Some(DEFAULT_CENTROID_THRESHOLD));
}
#[test]
fn translate_subset_emits_positional_doc_ids() {
let d2c: Vec<u64> = vec![100, 200, 300, 400, 500];
let subset = [200u64, 400, 500];
let docs = translate_chunk_subset_to_doc_subset(&d2c, &subset);
assert_eq!(docs, vec![1i64, 3, 4]);
}
#[test]
fn translate_subset_skips_unmapped_chunks() {
let d2c: Vec<u64> = vec![10, 20, 30];
let subset = [20u64, 999, 30];
let docs = translate_chunk_subset_to_doc_subset(&d2c, &subset);
assert_eq!(docs, vec![1i64, 2]);
}
#[test]
fn translate_subset_empty_chunks_yields_empty_docs() {
let d2c: Vec<u64> = vec![10, 20, 30];
let docs = translate_chunk_subset_to_doc_subset(&d2c, &[]);
assert!(docs.is_empty());
}
#[test]
fn translate_subset_preserves_doc_order_not_subset_order() {
let d2c: Vec<u64> = vec![100, 200, 300, 400, 500];
let subset = [500u64, 100];
let docs = translate_chunk_subset_to_doc_subset(&d2c, &subset);
assert_eq!(docs, vec![0i64, 4], "doc IDs must be ascending positional");
}
#[test]
fn translate_subset_skips_tombstone_positions() {
let d2c: Vec<u64> = vec![100, DENSE_TOMBSTONE, 300, DENSE_TOMBSTONE, 500];
let subset = [100u64, DENSE_TOMBSTONE, 300, 500];
let docs = translate_chunk_subset_to_doc_subset(&d2c, &subset);
assert_eq!(
docs,
vec![0i64, 2, 4],
"tombstone positions must be skipped even when the sentinel \
appears in the subset",
);
}
}