use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
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::SegmentMeta;
use tantivy::indexer::{LogMergePolicy, MergeCandidate, MergePolicy, NoMergePolicy, 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};
#[derive(Debug)]
struct ArcMergePolicy(Arc<dyn MergePolicy>);
impl MergePolicy for ArcMergePolicy {
fn compute_merge_candidates(&self, segments: &[SegmentMeta]) -> Vec<MergeCandidate> {
self.0.compute_merge_candidates(segments)
}
}
pub const CASS_SCHEMA_VERSION: &str = "v8";
pub const CASS_SCHEMA_HASH: &str =
"tantivy-schema-v8-hyphen-cjk-bigrams-bounded-content-prefix-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 b = *text.as_bytes().get(offset)?;
if b < 128 {
return Some((b as char, offset + 1));
}
let ch = text[offset..].chars().next()?;
Some((ch, offset + ch.len_utf8()))
}
#[doc(hidden)]
#[must_use]
pub fn next_char_from_slow(text: &str, offset: usize) -> Option<(char, usize)> {
let ch = text[offset..].chars().next()?;
Some((ch, offset + ch.len_utf8()))
}
#[doc(hidden)]
#[must_use]
pub fn cass_char_walk_fast(text: &str) -> u64 {
let mut acc = 0u64;
let mut off = 0usize;
while let Some((ch, next)) = next_char_from(text, off) {
acc = acc.wrapping_add(ch as u64);
off = next;
}
acc
}
#[doc(hidden)]
#[must_use]
pub fn cass_char_walk_slow(text: &str) -> u64 {
let mut acc = 0u64;
let mut off = 0usize;
while let Some((ch, next)) = next_char_from_slow(text, off) {
acc = acc.wrapping_add(ch as u64);
off = next;
}
acc
}
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}' )
}
#[doc(hidden)]
#[must_use]
pub fn cass_cjk_collect_fast(text: &str) -> Option<Vec<char>> {
let mut chars = text.chars();
let first = chars.next()?;
if !is_cjk(first) {
return None;
}
let mut buf: Vec<char> = Vec::with_capacity(text.len() / 3 + 1);
buf.push(first);
for c in chars {
if !is_cjk(c) {
return None;
}
buf.push(c);
}
if buf.len() < 2 {
return None;
}
Some(buf)
}
#[doc(hidden)]
#[must_use]
pub fn cass_cjk_collect_slow(text: &str) -> Option<Vec<char>> {
if text.is_empty() || !text.chars().all(is_cjk) {
return None;
}
let chars: Vec<char> = text.chars().collect();
if chars.len() < 2 {
return None;
}
Some(chars)
}
#[cfg(feature = "bench-internals")]
#[doc(hidden)]
pub fn cass_cjk_bigrams_staged(text: &str, token: &Token, pending: &mut Vec<Token>) {
let Some(chars) = cass_cjk_collect_fast(text) else {
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,
});
}
pending.extend(bigrams);
}
#[cfg(feature = "bench-internals")]
#[doc(hidden)]
pub fn cass_cjk_bigrams_direct(text: &str, token: &Token, pending: &mut Vec<Token>) {
let Some(chars) = cass_cjk_collect_fast(text) else {
return;
};
pending.reserve(chars.len() - 1);
for i in (0..chars.len() - 1).rev() {
let mut bigram = String::with_capacity(8);
bigram.push(chars[i]);
bigram.push(chars[i + 1]);
pending.push(Token {
text: bigram,
position: token.position,
offset_from: token.offset_from,
offset_to: token.offset_to,
position_length: token.position_length,
});
}
}
#[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();
let Some(chars) = cass_cjk_collect_fast(&token.text) else {
return;
};
self.pending.reserve(chars.len() - 1);
for i in (0..chars.len() - 1).rev() {
let mut bigram = String::with_capacity(8);
bigram.push(chars[i]);
bigram.push(chars[i + 1]);
self.pending.push(Token {
text: bigram,
position: token.position,
offset_from: token.offset_from,
offset_to: token.offset_to,
position_length: token.position_length,
});
}
}
}
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>,
}
#[cfg(feature = "tantivy-oracle")]
#[must_use]
pub fn cass_document_identity(document: &tantivy::TantivyDocument, fields: &CassFields) -> String {
use tantivy::schema::Value as _;
let source_id = document
.get_first(fields.source_id)
.and_then(|value| value.as_str())
.unwrap_or_default();
let msg_idx = document
.get_first(fields.msg_idx)
.and_then(|value| value.as_u64())
.unwrap_or_default();
cass_document_identity_parts(source_id, msg_idx)
}
#[cfg(feature = "tantivy-oracle")]
#[must_use]
pub fn cass_document_identity_parts(source_id: &str, msg_idx: u64) -> String {
format!("{source_id}#{msg_idx}")
}
#[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 {
#[cfg(feature = "tantivy-oracle")]
#[doc(hidden)]
pub fn in_memory_single_threaded_oracle() -> SearchResult<Self> {
let mut index = Index::create_in_ram(cass_build_schema());
cass_ensure_tokenizer(&mut index);
let fields = cass_fields_from_schema(&index.schema())?;
let writer_config = cass_writer_config_for_parallelism(1);
let writer = index
.writer_with_num_threads(writer_config.num_threads, writer_config.heap_size_bytes)
.map_err(tantivy_err)?;
Ok(Self {
index,
writer,
fields,
})
}
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)
}
#[cfg(feature = "tantivy-oracle")]
pub fn cass_oracle_observe_query(
&self,
raw_query: &str,
filters: &CassQueryFilters,
limit: usize,
tie_expansion_limit: usize,
) -> SearchResult<crate::OracleQueryObservation> {
cass_oracle_preflight(raw_query, filters, limit, tie_expansion_limit)?;
let fields = self.fields();
let tokens = cass_parse_boolean_query_bounded(raw_query, MAX_CASS_ORACLE_TOKENS).map_err(
|actual| cass_oracle_invalid_bound("token_count", actual, MAX_CASS_ORACLE_TOKENS),
)?;
let has_boolean_operators = cass_tokens_have_boolean_operators(&tokens);
let query = cass_fail_closed_tantivy_query(
raw_query,
cass_try_build_tantivy_query_from_tokens(
tokens,
has_boolean_operators,
filters,
&fields,
cass_regex_query_cached,
),
);
self.cass_oracle_observe_built_query(&*query, limit, tie_expansion_limit)
}
#[cfg(feature = "tantivy-oracle")]
pub fn cass_oracle_observe_query_profile(
&self,
raw_query: &str,
filters: &CassQueryFilters,
limit: usize,
tie_expansion_limit: usize,
) -> SearchResult<CassOracleProfileObservation> {
self.cass_oracle_observe_query_profile_with_regex_factory(
raw_query,
filters,
limit,
tie_expansion_limit,
cass_regex_query_cached,
)
}
#[cfg(feature = "tantivy-oracle")]
fn cass_oracle_observe_query_profile_with_regex_factory(
&self,
raw_query: &str,
filters: &CassQueryFilters,
limit: usize,
tie_expansion_limit: usize,
regex_query_factory: CassRegexQueryFactory,
) -> SearchResult<CassOracleProfileObservation> {
cass_oracle_preflight(raw_query, filters, limit, tie_expansion_limit)?;
let tokens = cass_parse_boolean_query_bounded(raw_query, MAX_CASS_ORACLE_TOKENS).map_err(
|actual| cass_oracle_invalid_bound("token_count", actual, MAX_CASS_ORACLE_TOKENS),
)?;
let has_boolean_operators = cass_tokens_have_boolean_operators(&tokens);
let fields = self.fields();
let outcome = cass_try_build_tantivy_query_from_tokens(
tokens.clone(),
has_boolean_operators,
filters,
&fields,
regex_query_factory,
)
.and_then(|query| self.cass_oracle_observe_built_query(&*query, limit, tie_expansion_limit))
.map_or_else(
CassOracleProfileOutcome::Error,
CassOracleProfileOutcome::Success,
);
Ok(CassOracleProfileObservation {
sanitized_query: cass_sanitize_query(raw_query),
tokens,
has_boolean_operators,
filters: filters.clone(),
outcome,
})
}
#[cfg(feature = "tantivy-oracle")]
fn cass_oracle_observe_built_query(
&self,
query: &dyn Query,
limit: usize,
tie_expansion_limit: usize,
) -> SearchResult<crate::OracleQueryObservation> {
let searcher = self.reader()?.searcher();
let fields = self.fields();
let fetch_limit = if limit == 0 {
0
} else {
limit.checked_add(tie_expansion_limit).ok_or_else(|| {
cass_oracle_invalid_bound(
"expanded_fetch_hits",
usize::MAX,
MAX_CASS_ORACLE_FETCH_HITS,
)
})?
};
cass_oracle_bound(
"expanded_fetch_hits",
fetch_limit,
MAX_CASS_ORACLE_FETCH_HITS,
)?;
let search_result = crate::execute_query_with_offset(&searcher, query, fetch_limit, 0)?;
let mut materialized = Vec::with_capacity(search_result.hits.len());
let mut retained_identity_bytes = 0_usize;
for hit in search_result.hits {
let document = crate::load_doc(&searcher, hit.doc_address)?;
let doc_id = cass_document_identity(&document, &fields);
cass_oracle_bound(
"document_identity_bytes",
doc_id.len(),
MAX_CASS_ORACLE_DOCUMENT_ID_BYTES,
)?;
retained_identity_bytes = retained_identity_bytes.saturating_add(doc_id.len());
cass_oracle_bound(
"aggregate_document_identity_bytes",
retained_identity_bytes,
MAX_CASS_ORACLE_DOCUMENT_ID_AGGREGATE_BYTES,
)?;
materialized.push(crate::OracleRankedHit {
doc_id,
score_bits: hit.bm25_score.to_bits(),
rank: hit.rank,
segment_ord: hit.doc_address.segment_ord,
segment_doc_id: hit.doc_address.doc_id,
snippet: None,
});
}
let top_len = limit.min(materialized.len());
let cutoff_bits = top_len
.checked_sub(1)
.and_then(|index| materialized.get(index))
.map(|hit| hit.score_bits);
let mut cutoff_tie_group = Vec::new();
if let Some(cutoff) = cutoff_bits {
for hit in &materialized {
if f32::from_bits(hit.score_bits)
.total_cmp(&f32::from_bits(cutoff))
.is_eq()
{
retained_identity_bytes =
retained_identity_bytes.saturating_add(hit.doc_id.len());
cass_oracle_bound(
"aggregate_document_identity_bytes",
retained_identity_bytes,
MAX_CASS_ORACLE_DOCUMENT_ID_AGGREGATE_BYTES,
)?;
cutoff_tie_group.push(hit.clone());
}
}
}
let cutoff_tie_complete = cutoff_bits.is_none_or(|cutoff| {
search_result.total_count <= fetch_limit
|| materialized.last().is_none_or(|last| {
!f32::from_bits(last.score_bits)
.total_cmp(&f32::from_bits(cutoff))
.is_eq()
})
});
materialized.truncate(top_len);
let doc_count =
usize::try_from(searcher.num_docs()).map_err(|_| SearchError::SubsystemError {
subsystem: "tantivy",
source: "current Tantivy reader document count does not fit usize".into(),
})?;
Ok(crate::OracleQueryObservation {
hits: materialized,
cutoff_tie_group,
cutoff_tie_complete,
total_count: search_result.total_count,
doc_count,
})
}
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 force_merge_bounded(&mut self, batch_size: usize) -> SearchResult<usize> {
let batch_size = batch_size.max(2);
let prior = self.writer.get_merge_policy();
self.writer.set_merge_policy(Box::new(NoMergePolicy));
let mut merges = 0usize;
let outcome: SearchResult<usize> = (|| {
loop {
let segment_ids = self.index.searchable_segment_ids().map_err(tantivy_err)?;
if segment_ids.len() <= 1 {
return Ok(merges);
}
let take = segment_ids.len().min(batch_size);
self.writer
.merge(&segment_ids[..take])
.wait()
.map_err(tantivy_err)?;
merges += 1;
}
})();
self.writer
.set_merge_policy(Box::new(ArcMergePolicy(prior)));
if matches!(outcome, Ok(n) if n > 0) {
let now_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |d| i64::try_from(d.as_millis()).unwrap_or(i64::MAX));
LAST_MERGE_TS.store(now_ms, Ordering::Relaxed);
}
outcome
}
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()) {
if word.is_ascii() {
let upper = word.len().min(MAX_NGRAM_INDICES - 1);
for end in 2..=upper {
cass_push_prefix_term(&mut ngrams, &word[..end]);
}
continue;
}
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
}
#[doc(hidden)]
#[must_use]
pub fn cass_generate_edge_ngrams_slow(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 cut = content.len();
let mut count = 0usize;
#[allow(clippy::explicit_counter_loop)]
for (byte_idx, _) in content.char_indices() {
if count == max_chars {
cut = byte_idx;
break;
}
count += 1;
}
let truncated = cut < content.len();
let mut out = String::with_capacity(cut + if truncated { '…'.len_utf8() } else { 0 });
out.push_str(&content[..cut]);
if truncated {
out.push('…');
}
out
}
#[doc(hidden)]
#[must_use]
pub fn cass_build_preview_slow(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 CONTENT_PREFIX_MAX_BYTES: usize = 4 * 1024;
let prefix_source = cass_prefix_source(content, CONTENT_PREFIX_MAX_BYTES);
(
cass_generate_edge_ngrams(prefix_source),
cass_build_preview(content, PREVIEW_MAX_CHARS),
)
}
fn cass_prefix_source(content: &str, max_bytes: usize) -> &str {
if content.len() <= max_bytes {
return content;
}
let mut end = max_bytes;
while !content.is_char_boundary(end) {
end -= 1;
}
&content[..end]
}
#[doc(hidden)]
#[must_use]
pub fn cass_prefix_source_slow(content: &str, max_bytes: usize) -> &str {
if content.len() <= max_bytes {
return content;
}
let mut end = 0usize;
for (byte_idx, _) in content.char_indices() {
if byte_idx > max_bytes {
break;
}
end = byte_idx;
}
&content[..end]
}
#[doc(hidden)]
#[must_use]
pub fn cass_prefix_source_fast_bench(content: &str, max_bytes: usize) -> usize {
cass_prefix_source(content, max_bytes).len()
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub enum CassSourceFilter {
#[default]
All,
Local,
Remote,
SourceId(String),
}
#[derive(Debug, Clone, PartialEq, Eq, 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,
}
#[cfg(feature = "tantivy-oracle")]
#[derive(Debug)]
pub struct CassOracleProfileObservation {
pub sanitized_query: String,
pub tokens: Vec<CassQueryToken>,
pub has_boolean_operators: bool,
pub filters: CassQueryFilters,
pub outcome: CassOracleProfileOutcome,
}
#[cfg(feature = "tantivy-oracle")]
#[derive(Debug)]
pub enum CassOracleProfileOutcome {
Success(crate::OracleQueryObservation),
Error(SearchError),
}
#[cfg(feature = "tantivy-oracle")]
pub const MAX_CASS_ORACLE_RAW_QUERY_BYTES: usize = 1_048_576;
#[cfg(feature = "tantivy-oracle")]
pub const MAX_CASS_ORACLE_TOKENS: usize = 20_000;
#[cfg(feature = "tantivy-oracle")]
pub const MAX_CASS_ORACLE_FILTER_VALUES: usize = 4_096;
#[cfg(feature = "tantivy-oracle")]
pub const MAX_CASS_ORACLE_FILTER_VALUE_BYTES: usize = 4_096;
#[cfg(feature = "tantivy-oracle")]
pub const MAX_CASS_ORACLE_FILTER_BYTES: usize = 1_048_576;
#[cfg(feature = "tantivy-oracle")]
pub const MAX_CASS_ORACLE_FETCH_HITS: usize = 100_000;
#[cfg(feature = "tantivy-oracle")]
pub const MAX_CASS_ORACLE_DOCUMENT_ID_BYTES: usize = 1_024;
#[cfg(feature = "tantivy-oracle")]
pub const MAX_CASS_ORACLE_DOCUMENT_ID_AGGREGATE_BYTES: usize = 16 * 1_024 * 1_024;
#[cfg(feature = "tantivy-oracle")]
fn cass_oracle_preflight(
raw_query: &str,
filters: &CassQueryFilters,
limit: usize,
tie_expansion_limit: usize,
) -> SearchResult<()> {
cass_oracle_bound(
"raw_query_bytes",
raw_query.len(),
MAX_CASS_ORACLE_RAW_QUERY_BYTES,
)?;
cass_oracle_bound("limit", limit, MAX_CASS_ORACLE_FETCH_HITS)?;
cass_oracle_bound(
"tie_expansion_limit",
tie_expansion_limit,
MAX_CASS_ORACLE_FETCH_HITS,
)?;
let expanded_fetch = limit.saturating_add(tie_expansion_limit);
cass_oracle_bound(
"expanded_fetch_hits",
expanded_fetch,
MAX_CASS_ORACLE_FETCH_HITS,
)?;
let source_value = match &filters.source_filter {
CassSourceFilter::SourceId(value) => Some(value.as_str()),
CassSourceFilter::All | CassSourceFilter::Local | CassSourceFilter::Remote => None,
};
let filter_count = filters
.agents
.len()
.checked_add(filters.workspaces.len())
.and_then(|count| count.checked_add(usize::from(source_value.is_some())))
.unwrap_or(usize::MAX);
cass_oracle_bound(
"filter_value_count",
filter_count,
MAX_CASS_ORACLE_FILTER_VALUES,
)?;
let mut aggregate_bytes = 0_usize;
for value in filters
.agents
.iter()
.chain(&filters.workspaces)
.map(String::as_str)
.chain(source_value)
{
cass_oracle_bound(
"filter_value_bytes",
value.len(),
MAX_CASS_ORACLE_FILTER_VALUE_BYTES,
)?;
aggregate_bytes = aggregate_bytes.saturating_add(value.len());
cass_oracle_bound(
"aggregate_filter_bytes",
aggregate_bytes,
MAX_CASS_ORACLE_FILTER_BYTES,
)?;
}
Ok(())
}
#[cfg(feature = "tantivy-oracle")]
fn cass_oracle_bound(field: &'static str, actual: usize, limit: usize) -> SearchResult<()> {
if actual > limit {
return Err(cass_oracle_invalid_bound(field, actual, limit));
}
Ok(())
}
#[cfg(feature = "tantivy-oracle")]
fn cass_oracle_invalid_bound(field: &'static str, actual: usize, limit: usize) -> SearchError {
SearchError::InvalidConfig {
field: format!("cass_oracle.{field}"),
value: actual.to_string(),
reason: format!("must be no greater than {limit}"),
}
}
#[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(".*");
}
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(".*");
}
Some(regex)
}
_ => None,
}
}
}
#[must_use]
pub fn cass_parse_boolean_query(query: &str) -> Vec<CassQueryToken> {
match cass_parse_boolean_query_with_sink(query, |tokens, token| {
tokens.push(token);
Ok::<(), std::convert::Infallible>(())
}) {
Ok(tokens) => tokens,
Err(never) => match never {},
}
}
#[cfg(feature = "tantivy-oracle")]
fn cass_parse_boolean_query_bounded(
query: &str,
max_tokens: usize,
) -> Result<Vec<CassQueryToken>, usize> {
cass_parse_boolean_query_with_sink(query, |tokens, token| {
cass_push_query_token(tokens, token, max_tokens)
})
}
fn cass_parse_boolean_query_with_sink<Error, PushToken>(
query: &str,
mut push_token: PushToken,
) -> Result<Vec<CassQueryToken>, Error>
where
PushToken: FnMut(&mut Vec<CassQueryToken>, CassQueryToken) -> Result<(), Error>,
{
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() {
push_token(
&mut tokens,
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() {
push_token(&mut tokens, CassQueryToken::Phrase(phrase))?;
}
}
'&' if chars.peek() == Some(&'&') => {
chars.next();
if !current_word.is_empty() {
push_token(
&mut tokens,
CassQueryToken::Term(std::mem::take(&mut current_word)),
)?;
}
push_token(&mut tokens, CassQueryToken::And)?;
}
'|' if chars.peek() == Some(&'|') => {
chars.next();
if !current_word.is_empty() {
push_token(
&mut tokens,
CassQueryToken::Term(std::mem::take(&mut current_word)),
)?;
}
push_token(&mut tokens, CassQueryToken::Or)?;
}
'-' if current_word.is_empty() => {
push_token(&mut tokens, CassQueryToken::Not)?;
}
' ' | '\t' | '\n' => {
if !current_word.is_empty() {
let word = std::mem::take(&mut current_word);
let token = cass_classify_query_word(word);
push_token(&mut tokens, token)?;
}
}
_ => current_word.push(c),
}
}
if !current_word.is_empty() {
let token = cass_classify_query_word(current_word);
push_token(&mut tokens, token)?;
}
Ok(tokens)
}
fn cass_classify_query_word(word: String) -> CassQueryToken {
if word.eq_ignore_ascii_case("AND") {
CassQueryToken::And
} else if word.eq_ignore_ascii_case("OR") {
CassQueryToken::Or
} else if word.eq_ignore_ascii_case("NOT") {
CassQueryToken::Not
} else {
CassQueryToken::Term(word)
}
}
#[cfg(feature = "tantivy-oracle")]
fn cass_push_query_token(
tokens: &mut Vec<CassQueryToken>,
token: CassQueryToken,
max_tokens: usize,
) -> Result<(), usize> {
if tokens.len() >= max_tokens {
return Err(tokens.len().saturating_add(1));
}
tokens.push(token);
Ok(())
}
#[must_use]
pub fn cass_has_boolean_operators(query: &str) -> bool {
let tokens = cass_parse_boolean_query(query);
cass_tokens_have_boolean_operators(&tokens)
}
fn cass_tokens_have_boolean_operators(tokens: &[CassQueryToken]) -> bool {
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))),
}
}
type CassRegexQueryFactory = fn(Field, &str) -> SearchResult<RegexQuery>;
fn cass_build_term_query_clauses(
pattern: &CassWildcardPattern,
fields: &CassFields,
regex_query_factory: CassRegexQueryFactory,
) -> SearchResult<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 Ok(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 Ok(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() {
let content_query = regex_query_factory(fields.content, ®ex_pattern)?;
let title_query = regex_query_factory(fields.title, ®ex_pattern)?;
shoulds.push((Occur::Should, Box::new(content_query)));
shoulds.push((Occur::Should, Box::new(title_query)));
}
}
}
Ok(shoulds)
}
fn cass_build_compound_term_query(
parts: &[String],
fields: &CassFields,
regex_query_factory: CassRegexQueryFactory,
) -> SearchResult<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, regex_query_factory)?;
if !term_shoulds.is_empty() {
subqueries.push(Box::new(BooleanQuery::new(term_shoulds)));
}
}
Ok(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,
regex_query_factory: CassRegexQueryFactory,
) -> SearchResult<Option<Box<dyn Query>>> {
if terms.is_empty() {
return Ok(None);
}
if terms.len() == 1 {
return cass_build_compound_term_query(terms, fields, regex_query_factory);
}
if terms.iter().any(|t| contains_cjk(t)) {
return cass_build_compound_term_query(terms, fields, regex_query_factory);
}
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))));
}
Ok(Some(Box::new(BooleanQuery::new(shoulds))))
}
fn cass_build_boolean_query_clauses(
tokens: &[CassQueryToken],
fields: &CassFields,
regex_query_factory: CassRegexQueryFactory,
) -> SearchResult<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, regex_query_factory)?;
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, regex_query_factory)?;
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);
if !clauses.is_empty() && clauses.iter().all(|(occur, _)| *occur == Occur::MustNot) {
clauses.insert(0, (Occur::Must, Box::new(AllQuery)));
}
Ok(clauses)
}
fn cass_match_none_query() -> Box<dyn Query> {
Box::new(BooleanQuery::new(vec![
(Occur::Must, Box::new(AllQuery)),
(Occur::MustNot, Box::new(AllQuery)),
]))
}
fn cass_fail_closed_tantivy_query(
raw_query: &str,
result: SearchResult<Box<dyn Query>>,
) -> Box<dyn Query> {
result.unwrap_or_else(|error| {
warn!(
query_bytes = raw_query.len(),
error = %error,
"CASS query construction failed; returning an explicit match-none query"
);
cass_match_none_query()
})
}
#[must_use]
pub fn cass_build_tantivy_query(
raw_query: &str,
filters: &CassQueryFilters,
fields: &CassFields,
) -> Box<dyn Query> {
let tokens = cass_parse_boolean_query(raw_query);
let has_boolean_operators = !tokens.is_empty() && cass_has_boolean_operators(raw_query);
cass_fail_closed_tantivy_query(
raw_query,
cass_try_build_tantivy_query_from_tokens(
tokens,
has_boolean_operators,
filters,
fields,
cass_regex_query_cached,
),
)
}
#[cfg(any(test, feature = "bench-internals"))]
#[doc(hidden)]
#[must_use]
pub fn cass_build_tantivy_query_single_parse(
raw_query: &str,
filters: &CassQueryFilters,
fields: &CassFields,
) -> Box<dyn Query> {
let tokens = cass_parse_boolean_query(raw_query);
let has_boolean_operators = cass_tokens_have_boolean_operators(&tokens);
cass_fail_closed_tantivy_query(
raw_query,
cass_try_build_tantivy_query_from_tokens(
tokens,
has_boolean_operators,
filters,
fields,
cass_regex_query_cached,
),
)
}
fn cass_try_build_tantivy_query_from_tokens(
tokens: Vec<CassQueryToken>,
has_boolean_operators: bool,
filters: &CassQueryFilters,
fields: &CassFields,
regex_query_factory: CassRegexQueryFactory,
) -> SearchResult<Box<dyn Query>> {
let mut clauses: Vec<(Occur, Box<dyn Query>)> = Vec::new();
if tokens.is_empty() {
clauses.push((Occur::Must, Box::new(AllQuery)));
} else if has_boolean_operators {
clauses.extend(cass_build_boolean_query_clauses(
&tokens,
fields,
regex_query_factory,
)?);
} 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, regex_query_factory)?
{
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)),
));
}
}
Ok(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::*;
use tantivy::schema::Value as _;
fn fields() -> CassFields {
let schema = cass_build_schema();
cass_fields_from_schema(&schema).expect("cass fields")
}
fn scored_cass_results(
searcher: &tantivy::Searcher,
fields: &CassFields,
raw_query: &str,
filters: &CassQueryFilters,
) -> std::collections::BTreeMap<u64, f32> {
let query = cass_build_tantivy_query(raw_query, filters, fields);
searcher
.search(
&*query,
&tantivy::collector::TopDocs::with_limit(64).order_by_score(),
)
.expect("search CASS Boolean fixture")
.into_iter()
.map(|(score, address)| {
let document: tantivy::TantivyDocument = searcher
.doc(address)
.expect("load CASS Boolean fixture document");
let msg_idx = document
.get_first(fields.msg_idx)
.and_then(|value| value.as_u64())
.expect("stored CASS fixture msg_idx");
(msg_idx, score)
})
.collect()
}
fn cass_result_ids(
searcher: &tantivy::Searcher,
fields: &CassFields,
raw_query: &str,
filters: &CassQueryFilters,
) -> Vec<u64> {
scored_cass_results(searcher, fields, raw_query, filters)
.into_keys()
.collect()
}
fn cass_result_fixture() -> (tempfile::TempDir, CassTantivyIndex) {
let directory = tempfile::tempdir().expect("temporary CASS result index directory");
let mut index =
CassTantivyIndex::open_or_create_with_writer_parallelism(directory.path(), 1)
.expect("create CASS result index");
let documents = [
CassDocumentRef {
agent: "claude",
workspace: Some("/alpha"),
workspace_original: Some("/alpha"),
source_path: "/fixture/0.jsonl",
msg_idx: 0,
created_at: Some(100),
title: Some("Current Alpha"),
content: "alpha active",
source_id: "local-a",
origin_kind: "local",
origin_host: None,
conversation_id: None,
},
CassDocumentRef {
agent: "claude",
workspace: Some("/alpha"),
workspace_original: Some("/alpha"),
source_path: "/fixture/1.jsonl",
msg_idx: 1,
created_at: Some(200),
title: Some("Legacy Alpha"),
content: "alpha deprecated",
source_id: "local-b",
origin_kind: "local",
origin_host: None,
conversation_id: None,
},
CassDocumentRef {
agent: "codex",
workspace: Some("/beta"),
workspace_original: Some("/beta"),
source_path: "/fixture/2.jsonl",
msg_idx: 2,
created_at: Some(300),
title: Some("Current Beta"),
content: "beta active",
source_id: "remote-a",
origin_kind: "ssh",
origin_host: Some("worker-a"),
conversation_id: None,
},
CassDocumentRef {
agent: "claude",
workspace: Some("/gamma"),
workspace_original: Some("/gamma"),
source_path: "/fixture/3.jsonl",
msg_idx: 3,
created_at: Some(400),
title: Some("Current Gamma"),
content: "gamma active",
source_id: "remote-b",
origin_kind: "ssh",
origin_host: Some("worker-b"),
conversation_id: None,
},
CassDocumentRef {
agent: "codex",
workspace: Some("/delta"),
workspace_original: Some("/delta"),
source_path: "/fixture/4.jsonl",
msg_idx: 4,
created_at: Some(500),
title: Some("Currentish Omega"),
content: "xalpha archived",
source_id: "local-c",
origin_kind: "local",
origin_host: None,
conversation_id: None,
},
CassDocumentRef {
agent: "gemini",
workspace: Some("/delta"),
workspace_original: Some("/delta"),
source_path: "/fixture/5.jsonl",
msg_idx: 5,
created_at: Some(600),
title: Some("Middle Marker"),
content: "alphabeta retained",
source_id: "local-d",
origin_kind: "local",
origin_host: None,
conversation_id: None,
},
];
index
.add_cass_document_refs(&documents)
.expect("index CASS result documents");
index.commit().expect("commit CASS result documents");
(directory, index)
}
#[cfg(feature = "tantivy-oracle")]
#[test]
fn cass_oracle_observation_carries_contract_identity_and_full_tie_evidence() {
let (_directory, index) = cass_result_fixture();
let filters = CassQueryFilters::default();
let observation = index
.cass_oracle_observe_query("alpha", &filters, 10, 8)
.expect("observe CASS term query");
assert_eq!(
observation
.hits
.iter()
.map(|hit| hit.doc_id.as_str())
.collect::<std::collections::BTreeSet<_>>(),
["local-a#0", "local-b#1", "local-d#5"]
.into_iter()
.collect(),
"identity must be the contract's source_id#msg_idx rendering, \
over the prefix-expanded match set",
);
assert_eq!(observation.total_count, 3);
assert_eq!(observation.doc_count, 6, "live document count");
assert!(
observation.hits.iter().all(|hit| hit.snippet.is_none()),
"the CASS activation profile is snippet-free",
);
assert!(
observation
.hits
.iter()
.enumerate()
.all(|(index, hit)| hit.rank == index),
"ranks must be the native fetched order",
);
let profiled_filters = CassQueryFilters {
agents: vec!["claude".to_owned()],
source_filter: CassSourceFilter::Local,
..CassQueryFilters::default()
};
let profile = index
.cass_oracle_observe_query_profile(
"c++ \"alpha active\" OR -deprecated",
&profiled_filters,
10,
8,
)
.expect("bounded profile observation");
assert_eq!(
profile.sanitized_query,
"c \"alpha active\" OR -deprecated"
);
assert_eq!(
profile.tokens,
vec![
CassQueryToken::Term("c++".to_owned()),
CassQueryToken::Phrase("alpha active".to_owned()),
CassQueryToken::Or,
CassQueryToken::Not,
CassQueryToken::Term("deprecated".to_owned()),
]
);
assert!(profile.has_boolean_operators);
assert_eq!(profile.filters, profiled_filters);
let profile_retrieval = match profile.outcome {
CassOracleProfileOutcome::Success(retrieval) => retrieval,
CassOracleProfileOutcome::Error(error) => {
panic!("fallible CASS profile unexpectedly failed: {error:?}")
}
};
assert_eq!(
profile_retrieval,
index
.cass_oracle_observe_query(
"c++ \"alpha active\" OR -deprecated",
&profiled_filters,
10,
8,
)
.expect("legacy CASS observation for parity")
);
let browse = index
.cass_oracle_observe_query("", &filters, 10, 8)
.expect("observe CASS filter-only browse");
assert_eq!(
browse.hits.len(),
6,
"empty CASS query must browse, not short-circuit: {:?}",
browse
.hits
.iter()
.map(|hit| &hit.doc_id)
.collect::<Vec<_>>(),
);
assert_eq!(browse.total_count, 6);
let filtered = index
.cass_oracle_observe_query(
"",
&CassQueryFilters {
source_filter: CassSourceFilter::SourceId("local-a".to_owned()),
..CassQueryFilters::default()
},
10,
8,
)
.expect("observe CASS structured filter");
assert_eq!(
filtered
.hits
.iter()
.map(|hit| hit.doc_id.as_str())
.collect::<Vec<_>>(),
vec!["local-a#0"],
);
let truncated = index
.cass_oracle_observe_query("alpha", &filters, 1, 8)
.expect("observe CASS truncated query");
assert_eq!(truncated.hits.len(), 1, "hits truncate to the requested k");
assert_eq!(
truncated.total_count, 3,
"total_count stays independent of k",
);
assert!(
truncated.cutoff_tie_complete,
"the whole match set was fetched, so the tie group is complete",
);
assert!(
truncated
.cutoff_tie_group
.iter()
.all(|hit| hit.score_bits == truncated.hits[0].score_bits),
"the tie group holds exactly the cutoff-score rows",
);
}
#[cfg(feature = "tantivy-oracle")]
#[test]
fn cass_oracle_profile_rejects_hostile_bounds_before_clone_or_search() {
let (_directory, index) = cass_result_fixture();
let oversized_raw = "x".repeat(MAX_CASS_ORACLE_RAW_QUERY_BYTES + 1);
assert!(matches!(
index.cass_oracle_observe_query_profile(
&oversized_raw,
&CassQueryFilters::default(),
10,
8,
),
Err(SearchError::InvalidConfig { field, .. })
if field == "cass_oracle.raw_query_bytes"
));
let too_many_tokens = "x ".repeat(MAX_CASS_ORACLE_TOKENS + 1);
assert!(matches!(
index.cass_oracle_observe_query_profile(
&too_many_tokens,
&CassQueryFilters::default(),
10,
8,
),
Err(SearchError::InvalidConfig { field, .. })
if field == "cass_oracle.token_count"
));
let aggregate_filters = CassQueryFilters {
agents: vec![
"x".repeat(MAX_CASS_ORACLE_FILTER_VALUE_BYTES);
MAX_CASS_ORACLE_FILTER_BYTES / MAX_CASS_ORACLE_FILTER_VALUE_BYTES + 1
],
..CassQueryFilters::default()
};
assert!(matches!(
index.cass_oracle_observe_query_profile("x", &aggregate_filters, 10, 8),
Err(SearchError::InvalidConfig { field, .. })
if field == "cass_oracle.aggregate_filter_bytes"
));
assert!(matches!(
index.cass_oracle_observe_query_profile(
"x",
&CassQueryFilters::default(),
MAX_CASS_ORACLE_FETCH_HITS,
1,
),
Err(SearchError::InvalidConfig { field, .. })
if field == "cass_oracle.expanded_fetch_hits"
));
}
#[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_boolean_operator_detection_matches_parsed_tokens() {
for (query, expected) in [
("", false),
("auth token cache", false),
("auth AND token", true),
("auth && token", true),
("auth OR token", true),
("auth || token", true),
("NOT legacy", true),
("-legacy", true),
("\"exact phrase\"", true),
("search 搜索", false),
] {
let tokens = cass_parse_boolean_query(query);
assert_eq!(
cass_tokens_have_boolean_operators(&tokens),
expected,
"token-slice detection mismatch for {query:?}"
);
assert_eq!(
cass_has_boolean_operators(query),
expected,
"public detection mismatch for {query:?}"
);
}
}
#[test]
fn cass_query_builder_single_parse_matches_operator_reparse() {
let fields = fields();
let filter_sets = [
CassQueryFilters::default(),
CassQueryFilters {
agents: vec!["claude".to_string(), "codex".to_string()],
workspaces: vec!["/data/projects/frankensearch".to_string()],
created_from: Some(1_700_000_000_000),
created_to: Some(1_800_000_000_000),
source_filter: CassSourceFilter::Remote,
},
];
let queries = [
"",
"auth token cache",
"auth AND token",
"auth && token OR cache",
"\"exact phrase\" OR cache",
"NOT legacy",
"-legacy identifier",
"AND",
"auth OR",
"\"unterminated phrase",
"search 搜索 NOT stale",
];
for filters in &filter_sets {
for query in queries {
let legacy = cass_build_tantivy_query(query, filters, &fields);
let single_parse = cass_build_tantivy_query_single_parse(query, filters, &fields);
assert_eq!(
format!("{legacy:?}"),
format!("{single_parse:?}"),
"query tree changed for {query:?} with filters {filters:?}"
);
}
}
}
#[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 next_char_from_ascii_matches_decode() {
for text in [
"",
"hello world",
"bd-q3fy ID_42",
"éclair café",
"日本語 mixed 한국어 text",
"a\u{300}b\u{1F600}c", "\u{4E00}\u{4E01}xyz",
] {
let mut off = 0usize;
loop {
let fast = next_char_from(text, off);
let slow = next_char_from_slow(text, off);
assert_eq!(fast, slow, "mismatch in {text:?} at offset {off}");
match fast {
Some((_, next)) => off = next,
None => break,
}
}
}
for text in ["hello world code_id-42", "日本語 mixed", ""] {
assert_eq!(
cass_char_walk_fast(text),
cass_char_walk_slow(text),
"{text:?}"
);
}
}
#[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_generate_edge_ngrams_matches_slow() {
for text in [
"",
"x",
"hi",
"hello world",
"bd-q3fy ID_42 snake_case",
"abcdefghijklmnopqrstuvwxyz longwordovercap",
"éclair café naïve",
"日本語 hello 한국어 world",
"aB3 café x9 déjà vu 12345678901234567890123",
] {
assert_eq!(
cass_generate_edge_ngrams(text),
cass_generate_edge_ngrams_slow(text),
"mismatch for {text:?}"
);
}
}
#[test]
fn cass_prefix_source_matches_slow() {
let ascii = "abcdefghij".repeat(50); let uni = "aéb日cé".repeat(50); for content in ["", "x", "abc", ascii.as_str(), uni.as_str()] {
for max_bytes in [0usize, 1, 2, 3, 4, 5, 7, 100, 499, 500, 501, 10_000] {
let fast = cass_prefix_source(content, max_bytes);
let slow = cass_prefix_source_slow(content, max_bytes);
assert_eq!(
fast,
slow,
"content.len()={} max_bytes={max_bytes}",
content.len()
);
}
}
}
#[test]
fn cass_build_preview_matches_slow() {
let ascii = "the quick brown fox jumps over the lazy dog ".repeat(20);
let uni = "café 日本語 éclair 한국어 naïve ".repeat(20);
for content in ["", "x", "hello world", ascii.as_str(), uni.as_str()] {
for max_chars in [0usize, 1, 3, 4, 5, 10, 50, 400, 100_000] {
assert_eq!(
cass_build_preview(content, max_chars),
cass_build_preview_slow(content, max_chars),
"content.len()={} max_chars={max_chars}",
content.len()
);
}
}
}
#[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_build_content_prefix_caps_long_message_bodies() {
let in_prefix_window = "alpha ".repeat(900);
let after_prefix_window = "omega ".repeat(900);
let content = format!("{in_prefix_window}{after_prefix_window}");
let (prefix, preview) = cass_build_content_prefix_and_preview(&content);
assert_eq!(preview, cass_build_preview(&content, 400));
assert!(prefix.contains("alpha"));
assert!(
!prefix.contains("omega"),
"content_prefix should not edge-ngram the whole message body"
);
assert_eq!(
prefix,
cass_generate_edge_ngrams(cass_prefix_source(&content, 4 * 1024))
);
}
#[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 cass_standalone_negation_matches_complement_and_is_score_neutral() {
let (_directory, index) = cass_result_fixture();
let reader = index.reader().expect("open CASS Boolean reader");
reader.reload().expect("reload CASS Boolean reader");
let searcher = reader.searcher();
let fields = index.fields();
let all_scores = scored_cass_results(&searcher, &fields, "", &CassQueryFilters::default());
for raw_query in ["NOT deprecated", "-deprecated", "NOT NOT deprecated"] {
let actual =
scored_cass_results(&searcher, &fields, raw_query, &CassQueryFilters::default());
assert_eq!(
actual.keys().copied().collect::<Vec<_>>(),
vec![0, 2, 3, 4, 5],
"standalone complement result set for {raw_query:?}"
);
for (msg_idx, score) in actual {
assert_eq!(
score.to_bits(),
all_scores[&msg_idx].to_bits(),
"MustNot changed the AllQuery score for {raw_query:?} document {msg_idx}"
);
}
}
let claude_filter = CassQueryFilters {
agents: vec!["claude".to_owned()],
..CassQueryFilters::default()
};
let filtered_all_scores = scored_cass_results(&searcher, &fields, "", &claude_filter);
for raw_query in ["NOT deprecated", "-deprecated"] {
let actual = scored_cass_results(&searcher, &fields, raw_query, &claude_filter);
assert_eq!(
actual.keys().copied().collect::<Vec<_>>(),
vec![0, 3],
"filtered standalone complement result set for {raw_query:?}"
);
for (msg_idx, score) in actual {
assert_eq!(
score.to_bits(),
filtered_all_scores[&msg_idx].to_bits(),
"MustNot changed the filtered score for {raw_query:?} document {msg_idx}"
);
}
}
let positive_scores =
scored_cass_results(&searcher, &fields, "alpha", &CassQueryFilters::default());
let mixed = scored_cass_results(
&searcher,
&fields,
"alpha AND NOT deprecated",
&CassQueryFilters::default(),
);
assert_eq!(mixed.keys().copied().collect::<Vec<_>>(), vec![0, 5]);
for (msg_idx, score) in mixed {
assert_eq!(
score.to_bits(),
positive_scores[&msg_idx].to_bits(),
"positive AND NOT gained a score-producing AllQuery anchor"
);
}
let all_negative = scored_cass_results(
&searcher,
&fields,
"-deprecated -beta",
&CassQueryFilters::default(),
);
assert_eq!(
all_negative.keys().copied().collect::<Vec<_>>(),
vec![0, 3, 4, 5]
);
for (msg_idx, score) in all_negative {
assert_eq!(score.to_bits(), all_scores[&msg_idx].to_bits());
}
}
#[test]
fn cass_regex_globs_match_whole_terms_and_fail_closed() {
assert_eq!(
CassWildcardPattern::Suffix("legacy".to_owned())
.to_regex()
.as_deref(),
Some(".*legacy")
);
assert_eq!(
CassWildcardPattern::Complex("a*ha".to_owned())
.to_regex()
.as_deref(),
Some("a.*ha")
);
assert_eq!(
CassWildcardPattern::Complex("*urr*nt".to_owned())
.to_regex()
.as_deref(),
Some(".*urr.*nt")
);
let (_directory, index) = cass_result_fixture();
let reader = index.reader().expect("open CASS regex reader");
reader.reload().expect("reload CASS regex reader");
let searcher = reader.searcher();
let fields = index.fields();
let no_filters = CassQueryFilters::default();
for (raw_query, expected) in [
("*legacy", vec![1]),
("*deprecated", vec![1]),
("*alpha", vec![0, 1, 4]),
("*alpha*", vec![0, 1, 4, 5]),
("a*ha", vec![0, 1]),
("*urr*nt", vec![0, 2, 3]),
("*pha*b*", vec![5]),
("*legacy *deprecated", vec![1]),
] {
assert_eq!(
cass_result_ids(&searcher, &fields, raw_query, &no_filters),
expected,
"regex-backed CASS result set for {raw_query:?}"
);
}
let claude_filter = CassQueryFilters {
agents: vec!["claude".to_owned()],
..CassQueryFilters::default()
};
assert_eq!(
cass_result_ids(&searcher, &fields, "*active", &claude_filter),
vec![0, 3]
);
fn force_regex_construction_failure(
field: Field,
_pattern: &str,
) -> SearchResult<RegexQuery> {
cass_regex_query_cached(field, "(")
}
let raw_query = "NOT *legacy";
let failed_build = cass_try_build_tantivy_query_from_tokens(
cass_parse_boolean_query(raw_query),
true,
&CassQueryFilters::default(),
&fields,
force_regex_construction_failure,
);
assert!(
failed_build.is_err(),
"the forced invalid regex must reach the root builder"
);
let fail_closed = cass_fail_closed_tantivy_query(raw_query, failed_build);
let hits = searcher
.search(&*fail_closed, &tantivy::collector::DocSetCollector)
.expect("search explicit match-none fallback");
assert!(
hits.is_empty(),
"a negated regex construction failure must not widen to AllQuery"
);
#[cfg(feature = "tantivy-oracle")]
{
let profile = index
.cass_oracle_observe_query_profile_with_regex_factory(
raw_query,
&CassQueryFilters::default(),
10,
8,
force_regex_construction_failure,
)
.expect("bounded forced-failure profile");
assert_eq!(
profile.tokens,
vec![
CassQueryToken::Not,
CassQueryToken::Term("*legacy".to_owned())
]
);
assert!(matches!(
profile.outcome,
CassOracleProfileOutcome::Error(_)
));
}
}
#[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 cass_cjk_collect_fast_matches_slow() {
let cases = [
"", "a", "hello world", "搜", "搜索", "搜索引擎", "検索とうきょう", "한국어", "搜a", "a搜", "搜 索", "𠀀𠀁", ];
for c in cases {
assert_eq!(
cass_cjk_collect_fast(c),
cass_cjk_collect_slow(c),
"fast/slow parity for {c:?}"
);
}
}
#[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(255)),
&format!("{} keep", "X".repeat(257)),
&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 cass_normalize_and_limit_pins_inclusive_shipping_boundary() {
let shipping_tokens = |text: &str| {
let mut analyzer = TextAnalyzer::builder(CassTokenizer::default())
.filter(HyphenDecompose)
.filter(CjkBigramDecompose)
.filter(CassNormalizeAndLimit)
.build();
let mut stream = analyzer.token_stream(text);
let mut tokens = Vec::new();
while stream.advance() {
let token = stream.token();
tokens.push((
token.text.clone(),
token.offset_from,
token.offset_to,
token.position,
token.position_length,
));
}
tokens
};
let tantivy_tokens = |text: &str| {
let mut analyzer = TextAnalyzer::builder(CassTokenizer::default())
.filter(HyphenDecompose)
.filter(CjkBigramDecompose)
.filter(tantivy::tokenizer::LowerCaser)
.filter(tantivy::tokenizer::RemoveLongFilter::limit(256))
.build();
let mut stream = analyzer.token_stream(text);
let mut tokens = Vec::new();
while stream.advance() {
let token = stream.token();
tokens.push((
token.text.clone(),
token.offset_from,
token.offset_to,
token.position,
token.position_length,
));
}
tokens
};
for length in [255, 256, 257] {
let input = format!("{} keep", "X".repeat(length));
let kept = vec![
("x".repeat(length), 0, length, 0, 1),
("keep".to_owned(), length + 1, length + 5, 1, 1),
];
let dropped = vec![("keep".to_owned(), length + 1, length + 5, 1, 1)];
let shipping_actual = shipping_tokens(&input);
let tantivy_actual = tantivy_tokens(&input);
assert_eq!(
shipping_actual.as_slice(),
if length <= 256 {
kept.as_slice()
} else {
dropped.as_slice()
},
"shipping CASS boundary at {length} bytes"
);
assert_eq!(
tantivy_actual.as_slice(),
if length < 256 {
kept.as_slice()
} else {
dropped.as_slice()
},
"Tantivy RemoveLongFilter boundary at {length} bytes"
);
}
}
#[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).order_by_score(),
)
.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).order_by_score(),
)
.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).order_by_score(),
)
.expect("search");
assert!(
!results.is_empty(),
"CJK search for '검색' should find the indexed Korean document"
);
}
#[test]
fn force_merge_bounded_consolidates_segments() {
let dir = tempfile::TempDir::new().expect("temp dir");
let mut idx = CassTantivyIndex::open_or_create(dir.path()).expect("create");
idx.writer.set_merge_policy(Box::new(NoMergePolicy));
let make_doc = |i: u64| CassDocument {
agent: "claude".to_string(),
workspace: None,
workspace_original: None,
source_path: format!("/tmp/test/{i}"),
msg_idx: i,
created_at: Some(1_700_000_000_000_i64.saturating_add(i.cast_signed())),
title: Some(format!("doc {i}")),
content: format!("bounded merge segment number {i}"),
source_id: "local".to_string(),
origin_kind: "local".to_string(),
origin_host: None,
conversation_id: None,
};
for i in 0..5u64 {
idx.add_cass_documents(&[make_doc(i)]).expect("add");
idx.commit().expect("commit");
}
let before = idx.index.searchable_segment_ids().expect("segments").len();
assert!(before > 1, "expected multiple segments, got {before}");
let merges = idx.force_merge_bounded(2).expect("bounded merge");
assert!(merges >= 1, "expected at least one bounded merge pass");
let after = idx.index.searchable_segment_ids().expect("segments").len();
assert!(
after <= 1,
"bounded merge should consolidate to <=1 segment, got {after}"
);
let reader = idx.reader().expect("reader");
reader.reload().expect("reload");
let searcher = reader.searcher();
let fields = idx.fields();
let query = cass_build_tantivy_query("bounded", &CassQueryFilters::default(), &fields);
let results = searcher
.search(
&query,
&tantivy::collector::TopDocs::with_limit(10).order_by_score(),
)
.expect("search");
assert_eq!(
results.len(),
5,
"all documents should remain searchable after bounded merge"
);
}
#[test]
fn force_merge_bounded_progressively_collapses_many_segments() {
let dir = tempfile::TempDir::new().expect("temp dir");
let mut idx = CassTantivyIndex::open_or_create(dir.path()).expect("create");
idx.writer.set_merge_policy(Box::new(NoMergePolicy));
for i in 0..100u64 {
idx.add_cass_documents(&[CassDocument {
agent: "claude".to_string(),
workspace: None,
workspace_original: None,
source_path: format!("/tmp/test100/{i}"),
msg_idx: i,
created_at: Some(1_700_000_000_000_i64.saturating_add(i.cast_signed())),
title: Some(format!("doc {i}")),
content: format!("scale test segment {i}"),
source_id: "local".to_string(),
origin_kind: "local".to_string(),
origin_host: None,
conversation_id: None,
}])
.expect("add");
idx.commit().expect("commit");
}
let before = idx.index.searchable_segment_ids().expect("segments").len();
assert!(
before >= 50,
"expected many segments before bounded merge, got {before}"
);
let merges = idx.force_merge_bounded(8).expect("bounded merge");
assert!(
merges >= 13,
"expected at least 13 bounded passes for 100 segments / batch 8, got {merges}"
);
let after = idx.index.searchable_segment_ids().expect("segments").len();
assert!(
after <= 1,
"bounded merge should collapse 100 segments to <=1, got {after}"
);
}
#[test]
fn force_merge_bounded_restores_prior_merge_policy() {
let dir = tempfile::TempDir::new().expect("temp dir");
let mut idx = CassTantivyIndex::open_or_create(dir.path()).expect("create");
idx.writer.set_merge_policy(Box::new(NoMergePolicy));
for i in 0..3u64 {
idx.add_cass_documents(&[CassDocument {
agent: "claude".to_string(),
workspace: None,
workspace_original: None,
source_path: format!("/tmp/restore/{i}"),
msg_idx: i,
created_at: Some(1_700_000_000_000_i64.saturating_add(i.cast_signed())),
title: Some(format!("restore {i}")),
content: format!("merge policy restore probe {i}"),
source_id: "local".to_string(),
origin_kind: "local".to_string(),
origin_host: None,
conversation_id: None,
}])
.expect("add");
idx.commit().expect("commit");
}
idx.configure_bulk_load_merge_policy();
let prior_dbg = format!("{:?}", idx.writer.get_merge_policy());
assert!(
prior_dbg.contains("min_num_segments: 256"),
"sanity: prior policy should be the bulk-load LogMergePolicy, got {prior_dbg}"
);
let _ = idx.force_merge_bounded(2).expect("bounded merge");
let after_dbg = format!("{:?}", idx.writer.get_merge_policy());
assert!(
after_dbg.contains("min_num_segments: 256"),
"force_merge_bounded dropped the bulk-load tuning: after = {after_dbg}"
);
assert!(
!after_dbg.contains("NoMergePolicy"),
"force_merge_bounded left NoMergePolicy in place: after = {after_dbg}"
);
}
}