use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicI64, Ordering};
use std::sync::{OnceLock, RwLock};
use std::time::{SystemTime, UNIX_EPOCH};
use frankensearch_core::error::{SearchError, SearchResult};
use tantivy::indexer::LogMergePolicy;
use tantivy::indexer::UserOperation;
use tantivy::query::{
AllQuery, BooleanQuery, Occur, PhraseQuery, Query, RangeQuery, RegexQuery, TermQuery,
};
use tantivy::schema::IndexRecordOption;
use tantivy::schema::{
FAST, Field, INDEXED, STORED, STRING, Schema, TextFieldIndexing, TextOptions,
};
#[cfg(test)]
use tantivy::tokenizer::RegexTokenizer;
use tantivy::tokenizer::{TextAnalyzer, Token, TokenFilter, TokenStream, Tokenizer};
use tantivy::{Index, IndexReader, IndexWriter, ReloadPolicy, Term, doc};
use tracing::{debug, info, warn};
pub const CASS_SCHEMA_VERSION: &str = "v7";
pub const CASS_SCHEMA_HASH: &str = "tantivy-schema-v7-hyphen-cjk-bigrams-prefix-basic-prefix-tokenizer-preview-stored-content-external";
#[derive(Clone, Default)]
struct CassTokenizer {
token: Token,
}
struct CassTokenStream<'a> {
text: &'a str,
cursor: usize,
token: &'a mut Token,
}
#[inline]
fn is_cass_tokenizer_cjk(c: char) -> bool {
matches!(
c,
'\u{4E00}'..='\u{9FFF}'
| '\u{3400}'..='\u{4DBF}'
| '\u{3040}'..='\u{309F}'
| '\u{30A0}'..='\u{30FF}'
| '\u{AC00}'..='\u{D7AF}'
| '\u{3100}'..='\u{312F}'
| '\u{3300}'..='\u{33FF}'
| '\u{F900}'..='\u{FAFF}'
| '\u{20000}'..='\u{2A6DF}'
)
}
#[inline]
fn next_char_from(text: &str, offset: usize) -> Option<(char, usize)> {
let ch = text[offset..].chars().next()?;
Some((ch, offset + ch.len_utf8()))
}
impl Tokenizer for CassTokenizer {
type TokenStream<'a> = CassTokenStream<'a>;
fn token_stream<'a>(&'a mut self, text: &'a str) -> Self::TokenStream<'a> {
self.token.reset();
CassTokenStream {
text,
cursor: 0,
token: &mut self.token,
}
}
}
impl CassTokenStream<'_> {
fn scan_ascii_token(&self, mut cursor: usize) -> usize {
let mut end = cursor;
let mut last_was_ascii_alnum = false;
while let Some((ch, next_cursor)) = next_char_from(self.text, cursor) {
if ch.is_ascii_alphanumeric() {
end = next_cursor;
cursor = next_cursor;
last_was_ascii_alnum = true;
continue;
}
if ch == '-'
&& last_was_ascii_alnum
&& let Some((next_ch, _)) = next_char_from(self.text, next_cursor)
&& next_ch.is_ascii_alphanumeric()
{
end = next_cursor;
cursor = next_cursor;
last_was_ascii_alnum = false;
continue;
}
break;
}
end
}
fn scan_cjk_token(&self, mut cursor: usize) -> usize {
let mut end = cursor;
while let Some((ch, next_cursor)) = next_char_from(self.text, cursor) {
if !is_cass_tokenizer_cjk(ch) {
break;
}
end = next_cursor;
cursor = next_cursor;
}
end
}
}
impl TokenStream for CassTokenStream<'_> {
fn advance(&mut self) -> bool {
self.token.text.clear();
self.token.position = self.token.position.wrapping_add(1);
while let Some((ch, next_cursor)) = next_char_from(self.text, self.cursor) {
if ch.is_ascii_alphanumeric() {
let offset_from = self.cursor;
let offset_to = self.scan_ascii_token(self.cursor);
self.token.offset_from = offset_from;
self.token.offset_to = offset_to;
self.token.text.push_str(&self.text[offset_from..offset_to]);
self.cursor = offset_to;
return true;
}
if is_cass_tokenizer_cjk(ch) {
let offset_from = self.cursor;
let offset_to = self.scan_cjk_token(next_cursor);
self.token.offset_from = offset_from;
self.token.offset_to = offset_to;
self.token.text.push_str(&self.text[offset_from..offset_to]);
self.cursor = offset_to;
return true;
}
self.cursor = next_cursor;
}
false
}
fn token(&self) -> &Token {
self.token
}
fn token_mut(&mut self) -> &mut Token {
self.token
}
}
#[derive(Clone)]
pub struct HyphenDecompose;
impl TokenFilter for HyphenDecompose {
type Tokenizer<T: Tokenizer> = HyphenDecomposeFilter<T>;
fn transform<T: Tokenizer>(self, tokenizer: T) -> HyphenDecomposeFilter<T> {
HyphenDecomposeFilter {
inner: tokenizer,
pending: Vec::new(),
}
}
}
#[derive(Clone)]
pub struct HyphenDecomposeFilter<T> {
inner: T,
pending: Vec<Token>,
}
impl<T: Tokenizer> Tokenizer for HyphenDecomposeFilter<T> {
type TokenStream<'a> = HyphenDecomposeStream<'a, T::TokenStream<'a>>;
fn token_stream<'a>(&'a mut self, text: &'a str) -> Self::TokenStream<'a> {
self.pending.clear();
HyphenDecomposeStream {
tail: self.inner.token_stream(text),
pending: &mut self.pending,
}
}
}
pub struct HyphenDecomposeStream<'a, T> {
tail: T,
pending: &'a mut Vec<Token>,
}
impl<T: TokenStream> HyphenDecomposeStream<'_, T> {
fn decompose(&mut self) {
let token = self.tail.token();
if !token.text.contains('-') {
return;
}
let parts: Vec<&str> = token.text.split('-').filter(|s| !s.is_empty()).collect();
if parts.len() < 2 {
return;
}
for &part in parts.iter().rev() {
self.pending.push(Token {
text: part.to_owned(),
position: token.position,
offset_from: token.offset_from,
offset_to: token.offset_to,
position_length: token.position_length,
});
}
self.pending.push(token.clone());
}
}
impl<T: TokenStream> TokenStream for HyphenDecomposeStream<'_, T> {
fn advance(&mut self) -> bool {
self.pending.pop();
if !self.pending.is_empty() {
return true;
}
if !self.tail.advance() {
return false;
}
self.decompose();
true
}
fn token(&self) -> &Token {
self.pending.last().unwrap_or_else(|| self.tail.token())
}
fn token_mut(&mut self) -> &mut Token {
self.pending
.last_mut()
.unwrap_or_else(|| self.tail.token_mut())
}
}
#[inline]
fn is_cjk(c: char) -> bool {
matches!(c,
'\u{4E00}'..='\u{9FFF}' | '\u{3400}'..='\u{4DBF}' | '\u{3040}'..='\u{309F}' | '\u{30A0}'..='\u{30FF}' | '\u{AC00}'..='\u{D7AF}' | '\u{3100}'..='\u{312F}' | '\u{3300}'..='\u{33FF}' | '\u{F900}'..='\u{FAFF}' | '\u{20000}'..='\u{2A6DF}' )
}
#[derive(Clone)]
pub struct CjkBigramDecompose;
impl TokenFilter for CjkBigramDecompose {
type Tokenizer<T: Tokenizer> = CjkBigramDecomposeFilter<T>;
fn transform<T: Tokenizer>(self, tokenizer: T) -> CjkBigramDecomposeFilter<T> {
CjkBigramDecomposeFilter {
inner: tokenizer,
pending: Vec::new(),
}
}
}
#[derive(Clone)]
pub struct CjkBigramDecomposeFilter<T> {
inner: T,
pending: Vec<Token>,
}
impl<T: Tokenizer> Tokenizer for CjkBigramDecomposeFilter<T> {
type TokenStream<'a> = CjkBigramDecomposeStream<'a, T::TokenStream<'a>>;
fn token_stream<'a>(&'a mut self, text: &'a str) -> Self::TokenStream<'a> {
self.pending.clear();
CjkBigramDecomposeStream {
tail: self.inner.token_stream(text),
pending: &mut self.pending,
}
}
}
pub struct CjkBigramDecomposeStream<'a, T> {
tail: T,
pending: &'a mut Vec<Token>,
}
impl<T: TokenStream> CjkBigramDecomposeStream<'_, T> {
fn decompose_cjk(&mut self) {
let token = self.tail.token();
if token.text.is_empty() || !token.text.chars().all(is_cjk) {
return;
}
let chars: Vec<char> = token.text.chars().collect();
if chars.len() == 1 {
return;
}
let mut bigrams: Vec<Token> = Vec::with_capacity(chars.len());
for i in (0..chars.len() - 1).rev() {
let mut bigram = String::with_capacity(8);
bigram.push(chars[i]);
bigram.push(chars[i + 1]);
bigrams.push(Token {
text: bigram,
position: token.position,
offset_from: token.offset_from,
offset_to: token.offset_to,
position_length: token.position_length,
});
}
self.pending.extend(bigrams);
}
}
impl<T: TokenStream> TokenStream for CjkBigramDecomposeStream<'_, T> {
fn advance(&mut self) -> bool {
self.pending.pop();
if !self.pending.is_empty() {
return true;
}
if !self.tail.advance() {
return false;
}
self.decompose_cjk();
true
}
fn token(&self) -> &Token {
self.pending.last().unwrap_or_else(|| self.tail.token())
}
fn token_mut(&mut self) -> &mut Token {
self.pending
.last_mut()
.unwrap_or_else(|| self.tail.token_mut())
}
}
#[derive(Clone, Copy, Default)]
pub struct CassNormalizeAndLimit;
impl TokenFilter for CassNormalizeAndLimit {
type Tokenizer<T: Tokenizer> = CassNormalizeAndLimitFilter<T>;
fn transform<T: Tokenizer>(self, tokenizer: T) -> CassNormalizeAndLimitFilter<T> {
CassNormalizeAndLimitFilter { inner: tokenizer }
}
}
#[derive(Clone)]
pub struct CassNormalizeAndLimitFilter<T> {
inner: T,
}
impl<T: Tokenizer> Tokenizer for CassNormalizeAndLimitFilter<T> {
type TokenStream<'a> = CassNormalizeAndLimitStream<T::TokenStream<'a>>;
fn token_stream<'a>(&'a mut self, text: &'a str) -> Self::TokenStream<'a> {
CassNormalizeAndLimitStream {
tail: self.inner.token_stream(text),
}
}
}
pub struct CassNormalizeAndLimitStream<T> {
tail: T,
}
impl<T: TokenStream> TokenStream for CassNormalizeAndLimitStream<T> {
fn advance(&mut self) -> bool {
while self.tail.advance() {
let token = self.tail.token_mut();
if token.text.len() > 256 {
continue;
}
token.text.make_ascii_lowercase();
return true;
}
false
}
fn token(&self) -> &Token {
self.tail.token()
}
fn token_mut(&mut self) -> &mut Token {
self.tail.token_mut()
}
}
const MERGE_COOLDOWN_MS: i64 = 300_000;
const MERGE_SEGMENT_THRESHOLD: usize = 4;
const CASS_MAX_WRITER_THREADS: usize = 32;
const CASS_MIN_WRITER_HEAP_BYTES: usize = 256 * 1024 * 1024;
const CASS_WRITER_HEAP_PER_THREAD_BYTES: usize = 128 * 1024 * 1024;
const CASS_BULK_LOAD_MIN_SEGMENTS_PER_MERGE: usize = 256;
const CASS_PARALLEL_ADD_MIN_DOCS: usize = 2_048;
const CASS_PARALLEL_ADD_TARGET_BATCH_DOCS: usize = 512;
const CASS_PARALLEL_ADD_MAX_BATCHES: usize = 64;
static LAST_MERGE_TS: AtomicI64 = AtomicI64::new(0);
const CASS_REGEX_QUERY_CACHE_CAP: usize = 128;
static CASS_REGEX_QUERY_CACHE: OnceLock<RwLock<HashMap<Field, HashMap<String, RegexQuery>>>> =
OnceLock::new();
fn tantivy_err<E>(err: E) -> SearchError
where
E: std::error::Error + Send + Sync + 'static,
{
SearchError::SubsystemError {
subsystem: "tantivy",
source: Box::new(err),
}
}
pub fn cass_regex_query_cached(field: Field, pattern: &str) -> SearchResult<RegexQuery> {
let cache = CASS_REGEX_QUERY_CACHE.get_or_init(|| RwLock::new(HashMap::new()));
{
let guard = cache
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(field_cache) = guard.get(&field)
&& let Some(q) = field_cache.get(pattern)
{
return Ok(q.clone());
}
}
let query = RegexQuery::from_pattern(pattern, field).map_err(tantivy_err)?;
let mut guard = cache
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let total_entries: usize = guard.values().map(HashMap::len).sum();
if total_entries >= CASS_REGEX_QUERY_CACHE_CAP {
guard.clear();
}
guard
.entry(field)
.or_default()
.insert(pattern.to_string(), query.clone());
drop(guard);
Ok(query)
}
pub fn cass_regex_query_uncached(field: Field, pattern: &str) -> SearchResult<RegexQuery> {
RegexQuery::from_pattern(pattern, field).map_err(tantivy_err)
}
#[must_use]
pub fn cass_schema_hash_matches(stored: &str) -> bool {
stored == CASS_SCHEMA_HASH
}
#[derive(Clone, Copy, Debug)]
pub struct CassFields {
pub agent: Field,
pub workspace: Field,
pub workspace_original: Field,
pub source_path: Field,
pub msg_idx: Field,
pub created_at: Field,
pub title: Field,
pub content: Field,
pub title_prefix: Field,
pub content_prefix: Field,
pub preview: Field,
pub source_id: Field,
pub origin_kind: Field,
pub origin_host: Field,
pub conversation_id: Option<Field>,
}
#[derive(Debug, Clone)]
pub struct CassMergeStatus {
pub segment_count: usize,
pub last_merge_ts: i64,
pub ms_since_last_merge: i64,
pub merge_threshold: usize,
pub cooldown_ms: i64,
}
impl CassMergeStatus {
#[must_use]
pub const fn should_merge(&self) -> bool {
self.segment_count >= self.merge_threshold
&& (self.ms_since_last_merge < 0 || self.ms_since_last_merge >= self.cooldown_ms)
}
}
#[derive(Debug, Clone)]
pub struct CassDocument {
pub agent: String,
pub workspace: Option<String>,
pub workspace_original: Option<String>,
pub source_path: String,
pub msg_idx: u64,
pub created_at: Option<i64>,
pub title: Option<String>,
pub content: String,
pub source_id: String,
pub origin_kind: String,
pub origin_host: Option<String>,
pub conversation_id: Option<i64>,
}
#[derive(Debug, Clone, Copy)]
pub struct CassDocumentRef<'a> {
pub agent: &'a str,
pub workspace: Option<&'a str>,
pub workspace_original: Option<&'a str>,
pub source_path: &'a str,
pub msg_idx: u64,
pub created_at: Option<i64>,
pub title: Option<&'a str>,
pub content: &'a str,
pub source_id: &'a str,
pub origin_kind: &'a str,
pub origin_host: Option<&'a str>,
pub conversation_id: Option<i64>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct CassWriterConfig {
num_threads: usize,
heap_size_bytes: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct CassAddPlan {
parallel_batches: bool,
batch_docs: usize,
batch_count: usize,
}
fn cass_writer_config() -> CassWriterConfig {
let available_parallelism = std::thread::available_parallelism()
.map(std::num::NonZeroUsize::get)
.unwrap_or(1);
cass_writer_config_for_parallelism(available_parallelism)
}
fn cass_writer_config_for_parallelism(available_parallelism: usize) -> CassWriterConfig {
let max_threads = std::env::var("CASS_TANTIVY_MAX_WRITER_THREADS")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(CASS_MAX_WRITER_THREADS);
let num_threads = available_parallelism.clamp(1, max_threads);
let heap_size_bytes = num_threads
.saturating_mul(CASS_WRITER_HEAP_PER_THREAD_BYTES)
.max(CASS_MIN_WRITER_HEAP_BYTES);
CassWriterConfig {
num_threads,
heap_size_bytes,
}
}
fn cass_parallel_add_min_docs() -> usize {
std::env::var("CASS_TANTIVY_PARALLEL_ADD_MIN_DOCS")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(CASS_PARALLEL_ADD_MIN_DOCS)
}
fn cass_parallel_add_target_batch_docs() -> usize {
std::env::var("CASS_TANTIVY_PARALLEL_ADD_BATCH_DOCS")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(CASS_PARALLEL_ADD_TARGET_BATCH_DOCS)
}
fn cass_parallel_add_max_batches() -> usize {
std::env::var("CASS_TANTIVY_PARALLEL_ADD_MAX_BATCHES")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(CASS_PARALLEL_ADD_MAX_BATCHES)
}
fn cass_parallel_add_plan(
doc_count: usize,
available_parallelism: usize,
target_batch_docs: usize,
min_parallel_docs: usize,
max_batches: usize,
) -> CassAddPlan {
if doc_count == 0 {
return CassAddPlan {
parallel_batches: false,
batch_docs: 1,
batch_count: 0,
};
}
let target_batch_docs = target_batch_docs.max(1);
let max_batches = max_batches.max(1);
let parallelism_batches = available_parallelism.max(1).saturating_mul(2).max(1);
let batch_count = doc_count
.div_ceil(target_batch_docs)
.min(max_batches)
.min(parallelism_batches)
.max(1);
let batch_docs = doc_count.div_ceil(batch_count).max(1);
CassAddPlan {
parallel_batches: doc_count >= min_parallel_docs && batch_count > 1,
batch_docs,
batch_count,
}
}
pub struct CassTantivyIndex {
index: Index,
writer: IndexWriter,
fields: CassFields,
}
impl CassTantivyIndex {
pub fn open_or_create(path: &Path) -> SearchResult<Self> {
let available_parallelism = std::thread::available_parallelism()
.map(std::num::NonZeroUsize::get)
.unwrap_or(1);
Self::open_or_create_with_writer_parallelism(path, available_parallelism)
}
pub fn open_or_create_with_writer_parallelism(
path: &Path,
available_parallelism: usize,
) -> SearchResult<Self> {
std::fs::create_dir_all(path).map_err(tantivy_err)?;
let meta_path = path.join("schema_hash.json");
let needs_rebuild = if meta_path.exists()
&& let Ok(meta) = std::fs::read_to_string(&meta_path)
&& let Ok(json) = serde_json::from_str::<serde_json::Value>(&meta)
&& json.get("schema_hash").and_then(|v| v.as_str()) == Some(CASS_SCHEMA_HASH)
{
false
} else {
true
};
if needs_rebuild {
if let Err(e) = std::fs::remove_dir_all(path) {
if e.kind() != std::io::ErrorKind::NotFound {
return Err(SearchError::from(e));
}
}
std::fs::create_dir_all(path).map_err(tantivy_err)?;
}
let mut index = if path.join("meta.json").exists() && !needs_rebuild {
match Index::open_in_dir(path) {
Ok(idx) => idx,
Err(e) => {
warn!(
error = %e,
"failed to open existing cass-compatible index; rebuilding"
);
if let Err(e) = std::fs::remove_dir_all(path) {
if e.kind() != std::io::ErrorKind::NotFound {
return Err(SearchError::from(e));
}
}
std::fs::create_dir_all(path).map_err(tantivy_err)?;
Index::create_in_dir(path, cass_build_schema()).map_err(tantivy_err)?
}
}
} else {
Index::create_in_dir(path, cass_build_schema()).map_err(tantivy_err)?
};
cass_ensure_tokenizer(&mut index);
std::fs::write(
&meta_path,
format!("{{\"schema_hash\":\"{CASS_SCHEMA_HASH}\"}}"),
)
.map_err(tantivy_err)?;
let actual_schema = index.schema();
let writer_config = cass_writer_config_for_parallelism(available_parallelism.max(1));
debug!(
tantivy_writer_threads = writer_config.num_threads,
tantivy_writer_heap_mb = writer_config.heap_size_bytes / (1024 * 1024),
"opening cass-compatible tantivy writer"
);
let writer = index
.writer_with_num_threads(writer_config.num_threads, writer_config.heap_size_bytes)
.map_err(tantivy_err)?;
let fields = cass_fields_from_schema(&actual_schema)?;
Ok(Self {
index,
writer,
fields,
})
}
#[must_use]
pub const fn fields(&self) -> CassFields {
self.fields
}
pub fn reader(&self) -> SearchResult<IndexReader> {
self.index.reader().map_err(tantivy_err)
}
pub fn delete_all(&mut self) -> SearchResult<()> {
self.writer.delete_all_documents().map_err(tantivy_err)?;
Ok(())
}
pub fn commit(&mut self) -> SearchResult<()> {
self.writer.commit().map_err(tantivy_err)?;
Ok(())
}
pub fn configure_bulk_load_merge_policy(&mut self) {
let mut merge_policy = LogMergePolicy::default();
merge_policy.set_min_num_segments(CASS_BULK_LOAD_MIN_SEGMENTS_PER_MERGE);
self.writer.set_merge_policy(Box::new(merge_policy));
debug!(
min_num_segments = CASS_BULK_LOAD_MIN_SEGMENTS_PER_MERGE,
"configured cass bulk-load merge policy"
);
}
#[must_use]
pub fn segment_count(&self) -> usize {
self.index
.searchable_segment_ids()
.map_or(0, |ids| ids.len())
}
#[must_use]
pub fn merge_status(&self) -> CassMergeStatus {
let last_merge_ts = LAST_MERGE_TS.load(Ordering::Relaxed);
let now_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |d| i64::try_from(d.as_millis()).unwrap_or(i64::MAX));
let ms_since_last = if last_merge_ts > 0 {
now_ms - last_merge_ts
} else {
-1
};
CassMergeStatus {
segment_count: self.segment_count(),
last_merge_ts,
ms_since_last_merge: ms_since_last,
merge_threshold: MERGE_SEGMENT_THRESHOLD,
cooldown_ms: MERGE_COOLDOWN_MS,
}
}
pub fn optimize_if_idle(&mut self) -> SearchResult<bool> {
let segment_ids = self.index.searchable_segment_ids().map_err(tantivy_err)?;
let segment_count = segment_ids.len();
if segment_count < MERGE_SEGMENT_THRESHOLD {
debug!(
segments = segment_count,
threshold = MERGE_SEGMENT_THRESHOLD,
"skipping merge: below threshold"
);
return Ok(false);
}
let now_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |d| i64::try_from(d.as_millis()).unwrap_or(i64::MAX));
let last_merge = LAST_MERGE_TS.load(Ordering::Relaxed);
if last_merge > 0 && (now_ms - last_merge) < MERGE_COOLDOWN_MS {
debug!(
ms_since_last = now_ms - last_merge,
cooldown = MERGE_COOLDOWN_MS,
"skipping merge: cooldown active"
);
return Ok(false);
}
info!(
segments = segment_count,
"starting cass-compatible segment merge"
);
let _merge_future = self.writer.merge(&segment_ids);
LAST_MERGE_TS.store(now_ms, Ordering::Relaxed);
Ok(true)
}
pub fn force_merge(&mut self) -> SearchResult<()> {
let segment_ids = self.index.searchable_segment_ids().map_err(tantivy_err)?;
if segment_ids.is_empty() {
return Ok(());
}
let now_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |d| i64::try_from(d.as_millis()).unwrap_or(i64::MAX));
let merge_future = self.writer.merge(&segment_ids);
match merge_future.wait() {
Ok(_) => {
LAST_MERGE_TS.store(now_ms, Ordering::Relaxed);
Ok(())
}
Err(err) => Err(tantivy_err(err)),
}
}
pub fn add_cass_documents(&mut self, docs: &[CassDocument]) -> SearchResult<()> {
use rayon::prelude::*;
const PARALLEL_PREP_THRESHOLD: usize = 8;
let fields = self.fields;
let writer_parallelism = cass_writer_config().num_threads.max(1);
let add_plan = cass_parallel_add_plan(
docs.len(),
writer_parallelism,
cass_parallel_add_target_batch_docs(),
cass_parallel_add_min_docs(),
cass_parallel_add_max_batches(),
);
debug!(
docs = docs.len(),
batch_docs = add_plan.batch_docs,
batch_count = add_plan.batch_count,
parallel_batches = add_plan.parallel_batches,
writer_parallelism,
"submitting cass-compatible documents to tantivy writer"
);
if add_plan.parallel_batches {
let writer = &self.writer;
docs.par_chunks(add_plan.batch_docs)
.try_for_each(|chunk| -> SearchResult<()> {
writer
.run(
chunk
.iter()
.map(|cass_doc| build_cass_tantivy_document(fields, cass_doc))
.map(UserOperation::Add),
)
.map_err(tantivy_err)?;
Ok(())
})?;
} else {
let prepared: Vec<tantivy::TantivyDocument> = if docs.len() < PARALLEL_PREP_THRESHOLD {
docs.iter()
.map(|cass_doc| build_cass_tantivy_document(fields, cass_doc))
.collect()
} else {
docs.par_iter()
.map(|cass_doc| build_cass_tantivy_document(fields, cass_doc))
.collect()
};
self.writer
.run(prepared.into_iter().map(UserOperation::Add))
.map_err(tantivy_err)?;
}
Ok(())
}
pub fn add_cass_document_refs(&mut self, docs: &[CassDocumentRef<'_>]) -> SearchResult<()> {
use rayon::prelude::*;
const PARALLEL_PREP_THRESHOLD: usize = 8;
let fields = self.fields;
let writer_parallelism = cass_writer_config().num_threads.max(1);
let add_plan = cass_parallel_add_plan(
docs.len(),
writer_parallelism,
cass_parallel_add_target_batch_docs(),
cass_parallel_add_min_docs(),
cass_parallel_add_max_batches(),
);
debug!(
docs = docs.len(),
batch_docs = add_plan.batch_docs,
batch_count = add_plan.batch_count,
parallel_batches = add_plan.parallel_batches,
writer_parallelism,
"submitting borrowed cass-compatible documents to tantivy writer"
);
if add_plan.parallel_batches {
let writer = &self.writer;
docs.par_chunks(add_plan.batch_docs)
.try_for_each(|chunk| -> SearchResult<()> {
writer
.run(
chunk
.iter()
.copied()
.map(|cass_doc| build_cass_tantivy_document_ref(fields, cass_doc))
.map(UserOperation::Add),
)
.map_err(tantivy_err)?;
Ok(())
})?;
} else {
let prepared: Vec<tantivy::TantivyDocument> = if docs.len() < PARALLEL_PREP_THRESHOLD {
docs.iter()
.copied()
.map(|cass_doc| build_cass_tantivy_document_ref(fields, cass_doc))
.collect()
} else {
docs.par_iter()
.copied()
.map(|cass_doc| build_cass_tantivy_document_ref(fields, cass_doc))
.collect()
};
self.writer
.run(prepared.into_iter().map(UserOperation::Add))
.map_err(tantivy_err)?;
}
Ok(())
}
}
fn build_cass_tantivy_document(
fields: CassFields,
cass_doc: &CassDocument,
) -> tantivy::TantivyDocument {
let mut d = doc! {
fields.agent => cass_doc.agent.clone(),
fields.source_path => cass_doc.source_path.clone(),
fields.msg_idx => cass_doc.msg_idx,
fields.content => cass_doc.content.clone(),
fields.source_id => cass_doc.source_id.clone(),
fields.origin_kind => cass_doc.origin_kind.clone(),
};
if let Some(host) = &cass_doc.origin_host
&& !host.is_empty()
{
d.add_text(fields.origin_host, host);
}
if let Some(field) = fields.conversation_id
&& let Some(conversation_id) = cass_doc.conversation_id
{
d.add_i64(field, conversation_id);
}
if let Some(workspace) = &cass_doc.workspace {
d.add_text(fields.workspace, workspace);
}
if let Some(workspace_original) = &cass_doc.workspace_original {
d.add_text(fields.workspace_original, workspace_original);
}
if let Some(ts) = cass_doc.created_at {
d.add_i64(fields.created_at, ts);
}
if let Some(title) = &cass_doc.title {
d.add_text(fields.title, title);
d.add_text(fields.title_prefix, cass_generate_edge_ngrams(title));
}
let (content_prefix, preview) = cass_build_content_prefix_and_preview(&cass_doc.content);
d.add_text(fields.content_prefix, content_prefix);
d.add_text(fields.preview, preview);
d
}
fn build_cass_tantivy_document_ref(
fields: CassFields,
cass_doc: CassDocumentRef<'_>,
) -> tantivy::TantivyDocument {
let mut d = doc! {
fields.agent => cass_doc.agent,
fields.source_path => cass_doc.source_path,
fields.msg_idx => cass_doc.msg_idx,
fields.content => cass_doc.content,
fields.source_id => cass_doc.source_id,
fields.origin_kind => cass_doc.origin_kind,
};
if let Some(host) = cass_doc.origin_host
&& !host.is_empty()
{
d.add_text(fields.origin_host, host);
}
if let Some(field) = fields.conversation_id
&& let Some(conversation_id) = cass_doc.conversation_id
{
d.add_i64(field, conversation_id);
}
if let Some(workspace) = cass_doc.workspace {
d.add_text(fields.workspace, workspace);
}
if let Some(workspace_original) = cass_doc.workspace_original {
d.add_text(fields.workspace_original, workspace_original);
}
if let Some(ts) = cass_doc.created_at {
d.add_i64(fields.created_at, ts);
}
if let Some(title) = cass_doc.title {
d.add_text(fields.title, title);
d.add_text(fields.title_prefix, cass_generate_edge_ngrams(title));
}
let (content_prefix, preview) = cass_build_content_prefix_and_preview(cass_doc.content);
d.add_text(fields.content_prefix, content_prefix);
d.add_text(fields.preview, preview);
d
}
#[must_use]
pub fn cass_build_schema() -> Schema {
let mut schema_builder = Schema::builder();
let indexed_text = TextOptions::default().set_indexing_options(
TextFieldIndexing::default()
.set_tokenizer("hyphen_normalize")
.set_index_option(tantivy::schema::IndexRecordOption::WithFreqsAndPositions),
);
let stored_indexed_text = indexed_text.clone().set_stored();
let prefix_text = TextOptions::default().set_indexing_options(
TextFieldIndexing::default()
.set_tokenizer("prefix_normalize")
.set_index_option(tantivy::schema::IndexRecordOption::Basic),
);
schema_builder.add_text_field("agent", STRING | STORED);
schema_builder.add_text_field("workspace", STRING | STORED);
schema_builder.add_text_field("workspace_original", STORED);
schema_builder.add_text_field("source_path", STORED);
schema_builder.add_u64_field("msg_idx", INDEXED | STORED);
schema_builder.add_i64_field("created_at", INDEXED | STORED | FAST);
schema_builder.add_text_field("title", stored_indexed_text);
schema_builder.add_text_field("content", indexed_text);
schema_builder.add_text_field("title_prefix", prefix_text.clone());
schema_builder.add_text_field("content_prefix", prefix_text);
schema_builder.add_text_field("preview", STORED);
schema_builder.add_text_field("source_id", STRING | STORED);
schema_builder.add_text_field("origin_kind", STRING | STORED);
schema_builder.add_text_field("origin_host", STRING | STORED);
schema_builder.add_i64_field("conversation_id", STORED);
schema_builder.build()
}
pub fn cass_fields_from_schema(schema: &Schema) -> SearchResult<CassFields> {
let get = |name: &str| {
schema
.get_field(name)
.map_err(|_| SearchError::InvalidConfig {
field: "schema".to_string(),
value: name.to_string(),
reason: format!("schema missing required field `{name}`"),
})
};
Ok(CassFields {
agent: get("agent")?,
workspace: get("workspace")?,
workspace_original: get("workspace_original")?,
source_path: get("source_path")?,
msg_idx: get("msg_idx")?,
created_at: get("created_at")?,
title: get("title")?,
content: get("content")?,
title_prefix: get("title_prefix")?,
content_prefix: get("content_prefix")?,
preview: get("preview")?,
source_id: get("source_id")?,
origin_kind: get("origin_kind")?,
origin_host: get("origin_host")?,
conversation_id: schema.get_field("conversation_id").ok(),
})
}
pub fn cass_open_search_reader(
index_path: &Path,
reload_policy: ReloadPolicy,
) -> SearchResult<(IndexReader, CassFields)> {
let mut index = Index::open_in_dir(index_path).map_err(tantivy_err)?;
cass_ensure_tokenizer(&mut index);
let schema = index.schema();
let fields = cass_fields_from_schema(&schema)?;
let reader = index
.reader_builder()
.reload_policy(reload_policy)
.try_into()
.map_err(tantivy_err)?;
if let Err(e) = reader.reload() {
warn!(error = %e, "index reader reload failed — searches may serve stale results");
}
Ok((reader, fields))
}
pub fn cass_index_dir(base: &Path) -> SearchResult<PathBuf> {
let dir = base.join("index").join(CASS_SCHEMA_VERSION);
std::fs::create_dir_all(&dir).map_err(tantivy_err)?;
Ok(dir)
}
pub fn cass_ensure_tokenizer(index: &mut Index) {
let analyzer = TextAnalyzer::builder(CassTokenizer::default())
.filter(HyphenDecompose)
.filter(CjkBigramDecompose)
.filter(CassNormalizeAndLimit)
.build();
index.tokenizers().register("hyphen_normalize", analyzer);
let prefix_analyzer = TextAnalyzer::builder(CassTokenizer::default())
.filter(CjkBigramDecompose)
.filter(CassNormalizeAndLimit)
.build();
index
.tokenizers()
.register("prefix_normalize", prefix_analyzer);
}
fn cass_push_prefix_term(out: &mut String, term: &str) {
if !out.is_empty() {
out.push(' ');
}
out.push_str(term);
}
#[must_use]
pub fn cass_generate_edge_ngrams(text: &str) -> String {
const MAX_NGRAM_INDICES: usize = 21;
let mut ngrams = String::with_capacity(text.len() * 2);
for word in text.split(|c: char| !c.is_alphanumeric()) {
let mut indices = [0usize; MAX_NGRAM_INDICES];
let mut index_count = 0usize;
for (i, _) in word.char_indices() {
if index_count == MAX_NGRAM_INDICES {
break;
}
indices[index_count] = i;
index_count += 1;
}
if index_count < MAX_NGRAM_INDICES {
indices[index_count] = word.len();
index_count += 1;
}
if index_count < 3 {
continue;
}
for &end_idx in &indices[2..index_count] {
cass_push_prefix_term(&mut ngrams, &word[..end_idx]);
}
}
ngrams
}
#[must_use]
pub fn cass_build_preview(content: &str, max_chars: usize) -> String {
let mut out = String::new();
let mut chars = content.chars();
for _ in 0..max_chars {
if let Some(ch) = chars.next() {
out.push(ch);
} else {
return out;
}
}
if chars.next().is_some() {
out.push('…');
}
out
}
#[must_use]
fn cass_build_content_prefix_and_preview(content: &str) -> (String, String) {
const PREVIEW_MAX_CHARS: usize = 400;
const MAX_NGRAM_INDICES: usize = 21;
let mut ngrams = String::with_capacity(content.len() * 2);
let mut preview = String::with_capacity(content.len().min(PREVIEW_MAX_CHARS + 8));
let mut preview_chars = 0usize;
let mut preview_truncated = false;
let mut word_indices = [0usize; MAX_NGRAM_INDICES];
let mut word_index_count = 0usize;
let mut word_start = 0usize;
let mut in_word = false;
for (byte_idx, ch) in content.char_indices() {
if preview_chars < PREVIEW_MAX_CHARS {
preview.push(ch);
preview_chars += 1;
} else {
preview_truncated = true;
}
if ch.is_alphanumeric() {
if !in_word {
in_word = true;
word_start = byte_idx;
word_indices[0] = 0;
word_index_count = 1;
} else if word_index_count < MAX_NGRAM_INDICES {
word_indices[word_index_count] = byte_idx - word_start;
word_index_count += 1;
}
continue;
}
if in_word {
if word_index_count < MAX_NGRAM_INDICES {
word_indices[word_index_count] = byte_idx - word_start;
word_index_count += 1;
}
if word_index_count >= 3 {
for &end_idx in &word_indices[2..word_index_count] {
cass_push_prefix_term(&mut ngrams, &content[word_start..word_start + end_idx]);
}
}
in_word = false;
}
}
if in_word {
if word_index_count < MAX_NGRAM_INDICES {
word_indices[word_index_count] = content.len() - word_start;
word_index_count += 1;
}
if word_index_count >= 3 {
for &end_idx in &word_indices[2..word_index_count] {
cass_push_prefix_term(&mut ngrams, &content[word_start..word_start + end_idx]);
}
}
}
if preview_truncated {
preview.push('…');
}
(ngrams, preview)
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub enum CassSourceFilter {
#[default]
All,
Local,
Remote,
SourceId(String),
}
#[derive(Debug, Clone, Default)]
pub struct CassQueryFilters {
pub agents: Vec<String>,
pub workspaces: Vec<String>,
pub created_from: Option<i64>,
pub created_to: Option<i64>,
pub source_filter: CassSourceFilter,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CassQueryToken {
Term(String),
Phrase(String),
And,
Or,
Not,
}
#[must_use]
pub fn cass_sanitize_query(raw: &str) -> String {
raw.chars()
.map(|c| {
if c.is_alphanumeric() || c == '*' || c == '"' || c == '-' {
c
} else {
' '
}
})
.collect()
}
#[must_use]
fn cass_escape_regex(s: &str) -> String {
let mut escaped = String::with_capacity(s.len() * 2);
for c in s.chars() {
match c {
'\\' | '.' | '+' | '*' | '?' | '(' | ')' | '[' | ']' | '{' | '}' | '|' | '^' | '$' => {
escaped.push('\\');
escaped.push(c);
}
_ => escaped.push(c),
}
}
escaped
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CassWildcardPattern {
Exact(String),
Prefix(String),
Suffix(String),
Substring(String),
Complex(String),
}
impl CassWildcardPattern {
#[must_use]
pub fn parse(term: &str) -> Self {
let starts_with_star = term.starts_with('*');
let ends_with_star = term.ends_with('*');
let core = term.trim_matches('*').to_lowercase();
if core.is_empty() {
return Self::Exact(String::new());
}
if core.contains('*') {
return Self::Complex(term.to_lowercase());
}
match (starts_with_star, ends_with_star) {
(true, true) => Self::Substring(core),
(true, false) => Self::Suffix(core),
(false, true) => Self::Prefix(core),
(false, false) => Self::Exact(core),
}
}
#[must_use]
pub fn to_regex(&self) -> Option<String> {
match self {
Self::Suffix(core) => Some(format!(".*{}$", cass_escape_regex(core))),
Self::Substring(core) => Some(format!(".*{}.*", cass_escape_regex(core))),
Self::Complex(full_term) => {
let mut regex = String::with_capacity(full_term.len() * 2 + 2);
if full_term.starts_with('*') {
regex.push_str(".*");
} else {
regex.push('^');
}
let trimmed_start = full_term.trim_start_matches('*');
let trimmed = trimmed_start.trim_end_matches('*');
for c in trimmed.chars() {
if c == '*' {
regex.push_str(".*");
} else {
match c {
'\\' | '.' | '+' | '?' | '(' | ')' | '[' | ']' | '{' | '}' | '|'
| '^' | '$' => {
regex.push('\\');
regex.push(c);
}
_ => regex.push(c),
}
}
}
if full_term.ends_with('*') {
regex.push_str(".*");
} else {
regex.push('$');
}
Some(regex)
}
_ => None,
}
}
}
#[must_use]
pub fn cass_parse_boolean_query(query: &str) -> Vec<CassQueryToken> {
let mut tokens = Vec::new();
let mut chars = query.chars().peekable();
let mut current_word = String::new();
while let Some(c) = chars.next() {
match c {
'"' => {
if !current_word.is_empty() {
tokens.push(CassQueryToken::Term(std::mem::take(&mut current_word)));
}
let mut phrase = String::new();
while let Some(&next) = chars.peek() {
if next == '"' {
chars.next();
break;
}
if let Some(c) = chars.next() {
phrase.push(c);
}
}
if !phrase.is_empty() {
tokens.push(CassQueryToken::Phrase(phrase));
}
}
'&' if chars.peek() == Some(&'&') => {
chars.next();
if !current_word.is_empty() {
tokens.push(CassQueryToken::Term(std::mem::take(&mut current_word)));
}
tokens.push(CassQueryToken::And);
}
'|' if chars.peek() == Some(&'|') => {
chars.next();
if !current_word.is_empty() {
tokens.push(CassQueryToken::Term(std::mem::take(&mut current_word)));
}
tokens.push(CassQueryToken::Or);
}
'-' if current_word.is_empty() => {
tokens.push(CassQueryToken::Not);
}
' ' | '\t' | '\n' => {
if !current_word.is_empty() {
let word = std::mem::take(&mut current_word);
let upper = word.to_ascii_uppercase();
match upper.as_str() {
"AND" => tokens.push(CassQueryToken::And),
"OR" => tokens.push(CassQueryToken::Or),
"NOT" => tokens.push(CassQueryToken::Not),
_ => tokens.push(CassQueryToken::Term(word)),
}
}
}
_ => current_word.push(c),
}
}
if !current_word.is_empty() {
let upper = current_word.to_ascii_uppercase();
match upper.as_str() {
"AND" => tokens.push(CassQueryToken::And),
"OR" => tokens.push(CassQueryToken::Or),
"NOT" => tokens.push(CassQueryToken::Not),
_ => tokens.push(CassQueryToken::Term(current_word)),
}
}
tokens
}
#[must_use]
pub fn cass_has_boolean_operators(query: &str) -> bool {
let tokens = cass_parse_boolean_query(query);
tokens.iter().any(|t| {
matches!(
t,
CassQueryToken::And
| CassQueryToken::Or
| CassQueryToken::Not
| CassQueryToken::Phrase(_)
)
})
}
#[must_use]
fn cass_normalize_term_parts(raw: &str) -> Vec<String> {
cass_sanitize_query(raw)
.split_whitespace()
.map(str::to_owned)
.collect()
}
#[must_use]
fn cass_normalize_phrase_terms(raw: &str) -> Vec<String> {
cass_sanitize_query(raw)
.split_whitespace()
.map(|s| s.trim_matches('*').to_lowercase())
.filter(|s| !s.is_empty())
.collect()
}
fn cass_flush_pending_or_group(
pending_or_group: &mut Vec<Box<dyn Query>>,
clauses: &mut Vec<(Occur, Box<dyn Query>)>,
) {
if pending_or_group.is_empty() {
return;
}
let or_clauses: Vec<_> = std::mem::take(pending_or_group)
.into_iter()
.map(|query| (Occur::Should, query))
.collect();
clauses.push((Occur::Must, Box::new(BooleanQuery::new(or_clauses))));
}
fn cass_lift_must_clause_into_or_group(
clauses: &mut Vec<(Occur, Box<dyn Query>)>,
pending_or_group: &mut Vec<Box<dyn Query>>,
) {
let can_pull = clauses
.last()
.is_some_and(|(occ, _)| *occ == Occur::Must || *occ == Occur::MustNot);
if !can_pull {
return;
}
if let Some((occur, last_query)) = clauses.pop() {
let lifted_query = if occur == Occur::MustNot {
Box::new(BooleanQuery::new(vec![
(Occur::Must, Box::new(AllQuery)),
(Occur::MustNot, last_query),
]))
} else {
last_query
};
pending_or_group.push(lifted_query);
}
}
fn cass_wrap_negated_clause(query: Box<dyn Query>) -> Box<dyn Query> {
Box::new(BooleanQuery::new(vec![
(Occur::Must, Box::new(AllQuery)),
(Occur::MustNot, query),
]))
}
fn cass_apply_query_token(
query: Box<dyn Query>,
next_occur: Occur,
in_or_sequence: &mut bool,
just_saw_or: &mut bool,
pending_or_group: &mut Vec<Box<dyn Query>>,
clauses: &mut Vec<(Occur, Box<dyn Query>)>,
) {
if *in_or_sequence && *just_saw_or {
if pending_or_group.is_empty() {
cass_lift_must_clause_into_or_group(clauses, pending_or_group);
}
let pushed_query = if next_occur == Occur::MustNot {
cass_wrap_negated_clause(query)
} else {
query
};
pending_or_group.push(pushed_query);
} else {
cass_flush_pending_or_group(pending_or_group, clauses);
*in_or_sequence = false;
clauses.push((next_occur, query));
}
*just_saw_or = false;
}
#[inline]
fn contains_cjk(s: &str) -> bool {
s.chars().any(is_cjk)
}
fn cjk_bigrams(s: &str) -> Vec<String> {
let chars: Vec<char> = s.chars().filter(|c| is_cjk(*c)).collect();
if chars.len() <= 1 {
return chars.iter().map(|c| c.to_string()).collect();
}
(0..chars.len() - 1)
.map(|i| {
let mut b = String::with_capacity(8);
b.push(chars[i]);
b.push(chars[i + 1]);
b
})
.collect()
}
#[inline]
fn cass_term_query_fields(fields: &CassFields) -> [(Field, IndexRecordOption); 4] {
[
(fields.title, IndexRecordOption::WithFreqsAndPositions),
(fields.content, IndexRecordOption::WithFreqsAndPositions),
(fields.title_prefix, IndexRecordOption::Basic),
(fields.content_prefix, IndexRecordOption::Basic),
]
}
fn cass_build_cjk_term_query(bigrams: &[String], fields: &CassFields) -> Option<Box<dyn Query>> {
if bigrams.is_empty() {
return None;
}
let mut bigram_musts: Vec<(Occur, Box<dyn Query>)> = Vec::new();
for bigram in bigrams {
let mut field_shoulds: Vec<(Occur, Box<dyn Query>)> = Vec::new();
for (field, index_record_option) in cass_term_query_fields(fields) {
field_shoulds.push((
Occur::Should,
Box::new(TermQuery::new(
Term::from_field_text(field, bigram),
index_record_option,
)),
));
}
bigram_musts.push((Occur::Must, Box::new(BooleanQuery::new(field_shoulds))));
}
match bigram_musts.len() {
0 => None,
1 => bigram_musts.pop().map(|(_, q)| q),
_ => Some(Box::new(BooleanQuery::new(bigram_musts))),
}
}
fn cass_build_term_query_clauses(
pattern: &CassWildcardPattern,
fields: &CassFields,
) -> Vec<(Occur, Box<dyn Query>)> {
let mut shoulds: Vec<(Occur, Box<dyn Query>)> = Vec::new();
match pattern {
CassWildcardPattern::Exact(term) | CassWildcardPattern::Prefix(term) => {
if term.is_empty() {
return shoulds;
}
if contains_cjk(term) {
let bigrams = cjk_bigrams(term);
if let Some(q) = cass_build_cjk_term_query(&bigrams, fields) {
shoulds.push((Occur::Should, q));
}
return shoulds;
}
for (field, index_record_option) in cass_term_query_fields(fields) {
shoulds.push((
Occur::Should,
Box::new(TermQuery::new(
Term::from_field_text(field, term),
index_record_option,
)),
));
}
}
CassWildcardPattern::Suffix(_)
| CassWildcardPattern::Substring(_)
| CassWildcardPattern::Complex(_) => {
if let Some(regex_pattern) = pattern.to_regex() {
if let Ok(rq) = cass_regex_query_cached(fields.content, ®ex_pattern) {
shoulds.push((Occur::Should, Box::new(rq)));
}
if let Ok(rq) = cass_regex_query_cached(fields.title, ®ex_pattern) {
shoulds.push((Occur::Should, Box::new(rq)));
}
}
}
}
shoulds
}
fn cass_build_compound_term_query(parts: &[String], fields: &CassFields) -> Option<Box<dyn Query>> {
let mut subqueries: Vec<Box<dyn Query>> = Vec::new();
for part in parts {
let pattern = CassWildcardPattern::parse(part);
let term_shoulds = cass_build_term_query_clauses(&pattern, fields);
if !term_shoulds.is_empty() {
subqueries.push(Box::new(BooleanQuery::new(term_shoulds)));
}
}
match subqueries.len() {
0 => None,
1 => subqueries.pop(),
_ => {
let musts = subqueries.into_iter().map(|q| (Occur::Must, q)).collect();
Some(Box::new(BooleanQuery::new(musts)))
}
}
}
fn cass_build_phrase_query(terms: &[String], fields: &CassFields) -> Option<Box<dyn Query>> {
if terms.is_empty() {
return None;
}
if terms.len() == 1 {
return cass_build_compound_term_query(terms, fields);
}
if terms.iter().any(|t| contains_cjk(t)) {
return cass_build_compound_term_query(terms, fields);
}
let mut shoulds: Vec<(Occur, Box<dyn Query>)> = Vec::new();
for field in [fields.title, fields.content] {
let phrase_terms = terms
.iter()
.map(|t| Term::from_field_text(field, t))
.collect::<Vec<_>>();
shoulds.push((Occur::Should, Box::new(PhraseQuery::new(phrase_terms))));
}
Some(Box::new(BooleanQuery::new(shoulds)))
}
fn cass_build_boolean_query_clauses(
tokens: &[CassQueryToken],
fields: &CassFields,
) -> Vec<(Occur, Box<dyn Query>)> {
let mut clauses: Vec<(Occur, Box<dyn Query>)> = Vec::new();
let mut pending_or_group: Vec<Box<dyn Query>> = Vec::new();
let mut next_occur = Occur::Must;
let mut in_or_sequence = false;
let mut just_saw_or = false;
for token in tokens {
match token {
CassQueryToken::And => {
cass_flush_pending_or_group(&mut pending_or_group, &mut clauses);
in_or_sequence = false;
just_saw_or = false;
next_occur = Occur::Must;
}
CassQueryToken::Or => {
in_or_sequence = true;
just_saw_or = true;
}
CassQueryToken::Not => {
if just_saw_or {
just_saw_or = true;
} else {
cass_flush_pending_or_group(&mut pending_or_group, &mut clauses);
in_or_sequence = false;
just_saw_or = false;
}
next_occur = Occur::MustNot;
}
CassQueryToken::Term(term) => {
let parts = cass_normalize_term_parts(term);
let term_query = cass_build_compound_term_query(&parts, fields);
let Some(term_query) = term_query else {
continue;
};
cass_apply_query_token(
term_query,
next_occur,
&mut in_or_sequence,
&mut just_saw_or,
&mut pending_or_group,
&mut clauses,
);
next_occur = Occur::Must;
}
CassQueryToken::Phrase(phrase) => {
let terms = cass_normalize_phrase_terms(phrase);
let phrase_query = cass_build_phrase_query(&terms, fields);
let Some(phrase_query) = phrase_query else {
continue;
};
cass_apply_query_token(
phrase_query,
next_occur,
&mut in_or_sequence,
&mut just_saw_or,
&mut pending_or_group,
&mut clauses,
);
next_occur = Occur::Must;
}
}
}
cass_flush_pending_or_group(&mut pending_or_group, &mut clauses);
clauses
}
#[must_use]
pub fn cass_build_tantivy_query(
raw_query: &str,
filters: &CassQueryFilters,
fields: &CassFields,
) -> Box<dyn Query> {
let mut clauses: Vec<(Occur, Box<dyn Query>)> = Vec::new();
let tokens = cass_parse_boolean_query(raw_query);
if tokens.is_empty() {
clauses.push((Occur::Must, Box::new(AllQuery)));
} else if cass_has_boolean_operators(raw_query) {
clauses.extend(cass_build_boolean_query_clauses(&tokens, fields));
} else {
for token in tokens {
if let CassQueryToken::Term(term_str) = token {
let parts = cass_normalize_term_parts(&term_str);
if let Some(term_query) = cass_build_compound_term_query(&parts, fields) {
clauses.push((Occur::Must, term_query));
}
}
}
}
if !filters.agents.is_empty() {
let terms = filters
.agents
.iter()
.map(|agent| {
(
Occur::Should,
Box::new(TermQuery::new(
Term::from_field_text(fields.agent, agent),
IndexRecordOption::Basic,
)) as Box<dyn Query>,
)
})
.collect();
clauses.push((Occur::Must, Box::new(BooleanQuery::new(terms))));
}
if !filters.workspaces.is_empty() {
let terms = filters
.workspaces
.iter()
.map(|ws| {
(
Occur::Should,
Box::new(TermQuery::new(
Term::from_field_text(fields.workspace, ws),
IndexRecordOption::Basic,
)) as Box<dyn Query>,
)
})
.collect();
clauses.push((Occur::Must, Box::new(BooleanQuery::new(terms))));
}
if filters.created_from.is_some() || filters.created_to.is_some() {
use std::ops::Bound::{Included, Unbounded};
let lower = filters.created_from.map_or(Unbounded, |v| {
Included(Term::from_field_i64(fields.created_at, v))
});
let upper = filters.created_to.map_or(Unbounded, |v| {
Included(Term::from_field_i64(fields.created_at, v))
});
let range = RangeQuery::new(lower, upper);
clauses.push((Occur::Must, Box::new(range)));
}
match &filters.source_filter {
CassSourceFilter::All => {}
CassSourceFilter::Local => {
let term = Term::from_field_text(fields.origin_kind, "local");
clauses.push((
Occur::Must,
Box::new(TermQuery::new(term, IndexRecordOption::Basic)),
));
}
CassSourceFilter::Remote => {
let term = Term::from_field_text(fields.origin_kind, "ssh");
clauses.push((
Occur::Must,
Box::new(TermQuery::new(term, IndexRecordOption::Basic)),
));
}
CassSourceFilter::SourceId(source_id) => {
let term = Term::from_field_text(fields.source_id, source_id);
clauses.push((
Occur::Must,
Box::new(TermQuery::new(term, IndexRecordOption::Basic)),
));
}
}
match clauses.len() {
0 => Box::new(AllQuery),
1 => {
if let Some((occur, query_box)) = clauses.pop() {
match occur {
Occur::Must => query_box,
_ => Box::new(BooleanQuery::new(vec![(occur, query_box)])),
}
} else {
Box::new(AllQuery)
}
}
_ => Box::new(BooleanQuery::new(clauses)),
}
}
#[cfg(test)]
mod cass_query_tests {
use super::*;
fn fields() -> CassFields {
let schema = cass_build_schema();
cass_fields_from_schema(&schema).expect("cass fields")
}
#[test]
fn cass_sanitize_query_preserves_wildcards_quotes_and_hyphens() {
let out = cass_sanitize_query("c++ \"hello-world\" *config*");
assert!(out.contains('"'));
assert!(out.contains('*'));
assert!(out.contains("hello-world"));
}
#[test]
fn cass_writer_config_scales_with_parallelism() {
assert_eq!(
cass_writer_config_for_parallelism(1),
CassWriterConfig {
num_threads: 1,
heap_size_bytes: 256 * 1024 * 1024,
}
);
assert_eq!(
cass_writer_config_for_parallelism(2),
CassWriterConfig {
num_threads: 2,
heap_size_bytes: 256 * 1024 * 1024,
}
);
assert_eq!(
cass_writer_config_for_parallelism(4),
CassWriterConfig {
num_threads: 4,
heap_size_bytes: 512 * 1024 * 1024,
}
);
assert_eq!(
cass_writer_config_for_parallelism(8),
CassWriterConfig {
num_threads: 8,
heap_size_bytes: 1024 * 1024 * 1024,
}
);
assert_eq!(
cass_writer_config_for_parallelism(32),
CassWriterConfig {
num_threads: 32,
heap_size_bytes: 4 * 1024 * 1024 * 1024,
}
);
assert_eq!(
cass_writer_config_for_parallelism(64),
CassWriterConfig {
num_threads: 32,
heap_size_bytes: 4 * 1024 * 1024 * 1024,
}
);
}
#[test]
fn cass_parallel_add_plan_enables_parallel_submission_for_large_batches() {
assert_eq!(
cass_parallel_add_plan(4_096, 8, 512, 2_048, 64),
CassAddPlan {
parallel_batches: true,
batch_docs: 512,
batch_count: 8,
}
);
assert_eq!(
cass_parallel_add_plan(200_000, 32, 512, 2_048, 64),
CassAddPlan {
parallel_batches: true,
batch_docs: 3_125,
batch_count: 64,
}
);
assert_eq!(
cass_parallel_add_plan(1_024, 8, 512, 2_048, 64),
CassAddPlan {
parallel_batches: false,
batch_docs: 512,
batch_count: 2,
}
);
}
#[test]
fn cass_generate_edge_ngrams_emits_expected_prefixes() {
assert_eq!(
cass_generate_edge_ngrams("hello world"),
"he hel hell hello wo wor worl world"
);
assert_eq!(
cass_generate_edge_ngrams("éclair"),
"éc écl écla éclai éclair"
);
assert_eq!(cass_generate_edge_ngrams("x"), "");
}
#[test]
fn cass_generate_edge_ngrams_caps_prefixes_at_twenty_chars() {
assert_eq!(
cass_generate_edge_ngrams("abcdefghijklmnopqrstuvwxy"),
"ab abc abcd abcde abcdef abcdefg abcdefgh abcdefghi abcdefghij abcdefghijk abcdefghijkl abcdefghijklm abcdefghijklmn abcdefghijklmno abcdefghijklmnop abcdefghijklmnopq abcdefghijklmnopqr abcdefghijklmnopqrs abcdefghijklmnopqrst"
);
}
#[test]
fn cass_build_preview_preserves_existing_behavior() {
assert_eq!(cass_build_preview("", 0), "");
assert_eq!(cass_build_preview("hello", 0), "…");
assert_eq!(cass_build_preview("hello", 10), "hello");
assert_eq!(cass_build_preview("hello world", 5), "hello…");
assert_eq!(cass_build_preview("éclair", 3), "écl…");
}
#[test]
fn cass_prefix_fields_store_basic_without_freqs_or_positions() {
let schema = cass_build_schema();
for field_name in ["title_prefix", "content_prefix"] {
let field = schema.get_field(field_name).unwrap();
let field_entry = schema.get_field_entry(field);
assert_eq!(
field_entry.field_type().get_index_record_option(),
Some(IndexRecordOption::Basic),
"unexpected index record option for {field_name}"
);
let tantivy::schema::FieldType::Str(text_options) = field_entry.field_type() else {
panic!("{field_name} should be a text field");
};
assert_eq!(
text_options
.get_indexing_options()
.expect("prefix field indexing options")
.tokenizer(),
"prefix_normalize",
"prefix field should use the cheaper generated-prefix analyzer"
);
}
}
#[test]
fn cass_prefix_analyzer_matches_full_analyzer_for_generated_prefix_terms() {
for text in [
"",
"Hello, happy tax payer!",
"bd-q3fy foo_bar baz-qux",
"abc-123 -- def",
"Hello搜索World",
"foo搜索-barあいう123",
"caf\u{00E9} 𠀀 token",
"multi---dash and trailing- hyphen",
] {
let prefix_terms = cass_generate_edge_ngrams(text);
let mut full = TextAnalyzer::builder(CassTokenizer::default())
.filter(HyphenDecompose)
.filter(CjkBigramDecompose)
.filter(CassNormalizeAndLimit)
.build();
let mut prefix = TextAnalyzer::builder(CassTokenizer::default())
.filter(CjkBigramDecompose)
.filter(CassNormalizeAndLimit)
.build();
let mut full_stream = full.token_stream(&prefix_terms);
let mut full_tokens = Vec::new();
while full_stream.advance() {
let token = full_stream.token();
full_tokens.push((
token.text.clone(),
token.offset_from,
token.offset_to,
token.position,
token.position_length,
));
}
let mut prefix_stream = prefix.token_stream(&prefix_terms);
let mut prefix_tokens = Vec::new();
while prefix_stream.advance() {
let token = prefix_stream.token();
prefix_tokens.push((
token.text.clone(),
token.offset_from,
token.offset_to,
token.position,
token.position_length,
));
}
assert_eq!(
prefix_tokens, full_tokens,
"prefix analyzer changed generated edge-ngram tokens for {text:?}"
);
}
}
#[test]
fn cass_preview_field_is_stored_only() {
let schema = cass_build_schema();
let field = schema.get_field("preview").unwrap();
let field_entry = schema.get_field_entry(field);
assert!(field_entry.is_stored(), "preview should stay stored");
assert_eq!(
field_entry.field_type().get_index_record_option(),
None,
"preview should not be indexed"
);
}
#[test]
fn cass_content_field_is_indexed_not_stored() {
let schema = cass_build_schema();
let field = schema.get_field("content").unwrap();
let field_entry = schema.get_field_entry(field);
assert!(field_entry.is_indexed(), "content must stay indexed");
assert!(
!field_entry.is_stored(),
"content should hydrate from canonical storage instead of Tantivy stored fields"
);
assert_eq!(
field_entry.field_type().get_index_record_option(),
Some(IndexRecordOption::WithFreqsAndPositions),
"content should keep full positional indexing"
);
}
#[test]
fn cass_build_content_prefix_and_preview_matches_existing_helpers() {
let samples = [
"",
"hello world",
"éclair crème brûlée",
"foo_bar baz-qux 12345",
"你好 世界 from cass",
&"alpha beta gamma ".repeat(64),
];
for sample in samples {
let (prefix, preview) = cass_build_content_prefix_and_preview(sample);
assert_eq!(prefix, cass_generate_edge_ngrams(sample));
assert_eq!(preview, cass_build_preview(sample, 400));
}
}
#[test]
fn cass_sanitize_query_splits_on_underscores() {
let out = cass_sanitize_query("hello_world");
assert_eq!(out, "hello world");
}
#[test]
fn cass_build_query_empty_returns_allquery() {
let f = fields();
let q = cass_build_tantivy_query("", &CassQueryFilters::default(), &f);
assert!(format!("{q:?}").to_ascii_lowercase().contains("allquery"));
}
#[test]
fn cass_build_query_applies_agent_filter() {
let f = fields();
let filters = CassQueryFilters {
agents: vec!["claude".to_string(), "codex".to_string()],
..CassQueryFilters::default()
};
let q = cass_build_tantivy_query("auth", &filters, &f);
let dbg = format!("{q:?}");
assert!(
dbg.contains("BooleanQuery"),
"expected boolean query: {dbg}"
);
}
#[test]
fn is_cjk_detects_chinese_characters() {
assert!(is_cjk('\u{4E00}')); assert!(is_cjk('\u{641C}')); assert!(is_cjk('\u{7D22}')); assert!(!is_cjk('a'));
assert!(!is_cjk('1'));
}
#[test]
fn is_cjk_detects_japanese_hiragana_katakana() {
assert!(is_cjk('\u{3042}')); assert!(is_cjk('\u{30A2}')); }
#[test]
fn is_cjk_detects_korean_hangul() {
assert!(is_cjk('\u{AC00}')); assert!(is_cjk('\u{D558}')); }
#[test]
fn cjk_bigram_decompose_produces_bigrams() {
let regex_tok = RegexTokenizer::new(
r"[\u4E00-\u9FFF\u3400-\u4DBF\u3040-\u309F\u30A0-\u30FF\uAC00-\uD7AF\u3100-\u312F\u3300-\u33FF\uF900-\uFAFF]+"
).unwrap();
let mut analyzer = TextAnalyzer::builder(regex_tok)
.filter(CjkBigramDecompose)
.build();
let mut stream = analyzer.token_stream("搜索引擎");
let mut tokens = Vec::new();
while stream.advance() {
tokens.push(stream.token().text.clone());
}
assert_eq!(tokens, vec!["搜索", "索引", "引擎"]);
}
#[test]
fn cjk_bigram_single_char_emits_unigram() {
let regex_tok = RegexTokenizer::new(r"[\u4E00-\u9FFF]+").unwrap();
let mut analyzer = TextAnalyzer::builder(regex_tok)
.filter(CjkBigramDecompose)
.build();
let mut stream = analyzer.token_stream("搜");
let mut tokens = Vec::new();
while stream.advance() {
tokens.push(stream.token().text.clone());
}
assert_eq!(tokens, vec!["搜"]);
}
#[test]
fn cjk_bigram_passes_through_ascii() {
let regex_tok =
RegexTokenizer::new(r"[a-zA-Z0-9]+(?:-[a-zA-Z0-9]+)*|[\u4E00-\u9FFF]+").unwrap();
let mut analyzer = TextAnalyzer::builder(regex_tok)
.filter(CjkBigramDecompose)
.build();
let mut stream = analyzer.token_stream("hello");
let mut tokens = Vec::new();
while stream.advance() {
tokens.push(stream.token().text.clone());
}
assert_eq!(tokens, vec!["hello"]);
}
#[test]
fn cjk_mixed_text_tokenizes_both() {
let regex_tok = RegexTokenizer::new(
r"[a-zA-Z0-9]+(?:-[a-zA-Z0-9]+)*|[\u4E00-\u9FFF\u3040-\u309F\u30A0-\u30FF\uAC00-\uD7AF]+"
).unwrap();
let mut analyzer = TextAnalyzer::builder(regex_tok)
.filter(CjkBigramDecompose)
.filter(CassNormalizeAndLimit)
.build();
let mut stream = analyzer.token_stream("Hello搜索World");
let mut tokens = Vec::new();
while stream.advance() {
tokens.push(stream.token().text.clone());
}
assert_eq!(tokens, vec!["hello", "搜索", "world"]);
}
#[test]
fn cass_tokenizer_matches_legacy_regex_boundaries() {
let regex_pattern = r"[a-zA-Z0-9]+(?:-[a-zA-Z0-9]+)*|[\u4E00-\u9FFF\u3400-\u4DBF\u3040-\u309F\u30A0-\u30FF\uAC00-\uD7AF\u3100-\u312F\u3300-\u33FF\uF900-\uFAFF\U00020000-\U0002A6DF]+";
let mut regex = RegexTokenizer::new(regex_pattern).unwrap();
let mut custom = CassTokenizer::default();
for text in [
"",
"Hello, happy tax payer!",
"bd-q3fy foo_bar baz-qux",
"abc-123 -- def",
"Hello搜索World",
"foo搜索-barあいう123",
"caf\u{00E9} 𠀀 token",
"multi---dash and trailing- hyphen",
] {
let mut regex_stream = regex.token_stream(text);
let mut regex_tokens = Vec::new();
while regex_stream.advance() {
let token = regex_stream.token();
regex_tokens.push((
token.text.clone(),
token.offset_from,
token.offset_to,
token.position,
));
}
let mut custom_stream = custom.token_stream(text);
let mut custom_tokens = Vec::new();
while custom_stream.advance() {
let token = custom_stream.token();
custom_tokens.push((
token.text.clone(),
token.offset_from,
token.offset_to,
token.position,
));
}
assert_eq!(custom_tokens, regex_tokens, "token mismatch for {text:?}");
}
}
#[test]
fn cass_tokenizer_keeps_extension_b_cjk() {
let mut tokenizer = CassTokenizer::default();
let mut stream = tokenizer.token_stream("𠀀搜索 test");
let mut tokens = Vec::new();
while stream.advance() {
let token = stream.token();
tokens.push(token.text.clone());
}
assert_eq!(tokens, vec!["𠀀搜索".to_string(), "test".to_string()]);
}
#[test]
fn cass_normalize_and_limit_matches_legacy_pipeline() {
for text in [
"",
"Hello, happy tax payer!",
"BD-Q3FY foo_bar BAZ-QUX",
"abc-123 -- def",
"Hello搜索World",
"foo搜索-barあいう123",
"caf\u{00E9} 𠀀 token",
"multi---dash and trailing- hyphen",
&format!("{} keep", "X".repeat(300)),
] {
let mut legacy = TextAnalyzer::builder(CassTokenizer::default())
.filter(HyphenDecompose)
.filter(CjkBigramDecompose)
.filter(tantivy::tokenizer::LowerCaser)
.filter(tantivy::tokenizer::RemoveLongFilter::limit(256))
.build();
let mut optimized = TextAnalyzer::builder(CassTokenizer::default())
.filter(HyphenDecompose)
.filter(CjkBigramDecompose)
.filter(CassNormalizeAndLimit)
.build();
let mut legacy_stream = legacy.token_stream(text);
let mut legacy_tokens = Vec::new();
while legacy_stream.advance() {
let token = legacy_stream.token();
legacy_tokens.push((
token.text.clone(),
token.offset_from,
token.offset_to,
token.position,
token.position_length,
));
}
let mut optimized_stream = optimized.token_stream(text);
let mut optimized_tokens = Vec::new();
while optimized_stream.advance() {
let token = optimized_stream.token();
optimized_tokens.push((
token.text.clone(),
token.offset_from,
token.offset_to,
token.position,
token.position_length,
));
}
assert_eq!(
optimized_tokens, legacy_tokens,
"normalized analyzer mismatch for {text:?}"
);
}
}
#[test]
fn cjk_sanitize_query_preserves_chinese() {
let out = cass_sanitize_query("搜索引擎 test");
assert!(out.contains("搜索引擎"));
assert!(out.contains("test"));
}
#[test]
fn cjk_index_and_search_roundtrip() {
let dir = tempfile::TempDir::new().expect("temp dir");
let mut idx = CassTantivyIndex::open_or_create(dir.path()).expect("create");
let doc = CassDocument {
agent: "claude".to_string(),
workspace: None,
workspace_original: None,
source_path: "/tmp/test".to_string(),
msg_idx: 0,
created_at: Some(1_700_000_000_000),
title: Some("搜索引擎测试".to_string()),
content: "这是一个搜索引擎的测试".to_string(),
conversation_id: None,
source_id: "local".to_string(),
origin_kind: "local".to_string(),
origin_host: None,
};
idx.add_cass_documents(&[doc]).expect("add");
idx.commit().expect("commit");
let reader = idx.reader().expect("reader");
reader.reload().expect("reload");
let searcher = reader.searcher();
let fields = idx.fields();
let query = cass_build_tantivy_query("搜索", &CassQueryFilters::default(), &fields);
let results = searcher
.search(&query, &tantivy::collector::TopDocs::with_limit(10))
.expect("search");
assert!(
!results.is_empty(),
"CJK search for '搜索' should find the indexed Chinese document"
);
}
#[test]
fn japanese_index_and_search_roundtrip() {
let dir = tempfile::TempDir::new().expect("temp dir");
let mut idx = CassTantivyIndex::open_or_create(dir.path()).expect("create");
let doc = CassDocument {
agent: "claude".to_string(),
workspace: None,
workspace_original: None,
source_path: "/tmp/test".to_string(),
msg_idx: 0,
created_at: Some(1_700_000_000_000),
title: Some("テスト".to_string()),
content: "これはテストです".to_string(),
conversation_id: None,
source_id: "local".to_string(),
origin_kind: "local".to_string(),
origin_host: None,
};
idx.add_cass_documents(&[doc]).expect("add");
idx.commit().expect("commit");
let reader = idx.reader().expect("reader");
reader.reload().expect("reload");
let searcher = reader.searcher();
let fields = idx.fields();
let query = cass_build_tantivy_query("テスト", &CassQueryFilters::default(), &fields);
let results = searcher
.search(&query, &tantivy::collector::TopDocs::with_limit(10))
.expect("search");
assert!(
!results.is_empty(),
"CJK search for 'テスト' should find the indexed Japanese document"
);
}
#[test]
fn korean_index_and_search_roundtrip() {
let dir = tempfile::TempDir::new().expect("temp dir");
let mut idx = CassTantivyIndex::open_or_create(dir.path()).expect("create");
let doc = CassDocument {
agent: "claude".to_string(),
workspace: None,
workspace_original: None,
source_path: "/tmp/test".to_string(),
msg_idx: 0,
created_at: Some(1_700_000_000_000),
title: Some("검색엔진".to_string()),
content: "한국어 검색엔진 테스트".to_string(),
conversation_id: None,
source_id: "local".to_string(),
origin_kind: "local".to_string(),
origin_host: None,
};
idx.add_cass_documents(&[doc]).expect("add");
idx.commit().expect("commit");
let reader = idx.reader().expect("reader");
reader.reload().expect("reload");
let searcher = reader.searcher();
let fields = idx.fields();
let query = cass_build_tantivy_query("검색", &CassQueryFilters::default(), &fields);
let results = searcher
.search(&query, &tantivy::collector::TopDocs::with_limit(10))
.expect("search");
assert!(
!results.is_empty(),
"CJK search for '검색' should find the indexed Korean document"
);
}
}