use std::collections::HashMap;
use std::num::NonZeroUsize;
use std::sync::{Arc, Mutex};
use lru::LruCache;
use rayon::iter::{IntoParallelIterator, ParallelIterator};
use scryer_db::{ScryerDb, SourceFile, Symbol};
use super::bm25::{Bm25Builder, Bm25Index};
use super::tokenize::{tokenize_identifier, tokenize_text};
pub type IndexKey = (u64, Option<u64>);
const SYMBOL_FIELD_WEIGHTS: [f32; 5] = [3.0, 2.0, 1.0, 1.0, 1.0];
const DOC_SNIPPET_CHARS: usize = 200;
#[derive(Debug, Clone)]
pub struct SymbolDoc {
pub id: u64,
pub file_id: u64,
pub kind: String,
pub name: String,
pub qualified_name: String,
pub signature: String,
pub doc_snippet: String,
pub start_line: u32,
pub end_line: u32,
pub dependency_package_id: Option<u64>,
pub file_path: String,
pub is_test: bool,
}
const TEST_DIRS: [&str; 5] = ["tests", "test", "__tests__", "benches", "examples"];
pub fn is_test_symbol(file_path: &str, qualified_name: &str) -> bool {
let mut parts = file_path.split(['/', '\\']).peekable();
let mut dirs_are_tests = false;
let mut file = "";
while let Some(part) = parts.next() {
if parts.peek().is_some() {
dirs_are_tests |= TEST_DIRS.contains(&part);
} else {
file = part;
}
}
let (stem, extension) = file.rsplit_once('.').unwrap_or((file, ""));
let test_file = match extension {
"py" => stem.starts_with("test_") || stem.ends_with("_test") || stem == "conftest",
"ts" | "tsx" | "js" | "jsx" | "mts" | "cts" => {
stem.ends_with(".test") || stem.ends_with(".spec")
}
_ => false,
};
dirs_are_tests
|| test_file
|| qualified_name.contains("::tests::")
|| qualified_name.contains("::test::")
}
impl SymbolDoc {
pub fn parent_path(&self) -> &str {
self.qualified_name
.rsplit_once("::")
.or_else(|| self.qualified_name.rsplit_once('.'))
.map(|(parent, _)| parent)
.unwrap_or("")
}
}
#[derive(Debug)]
pub struct SymbolIndex {
pub generation: u64,
pub index: Bm25Index<SymbolDoc>,
}
pub struct SearchIndexCache {
generations: Mutex<HashMap<u64, u64>>,
entries: tokio::sync::Mutex<LruCache<IndexKey, Arc<SymbolIndex>>>,
build_locks: Mutex<HashMap<IndexKey, Arc<tokio::sync::Mutex<()>>>>,
}
impl Default for SearchIndexCache {
fn default() -> Self {
Self::new(16)
}
}
impl SearchIndexCache {
pub fn new(capacity: usize) -> Self {
Self {
generations: Mutex::default(),
entries: tokio::sync::Mutex::new(LruCache::new(
NonZeroUsize::new(capacity.max(1)).expect("capacity is non-zero"),
)),
build_locks: Mutex::default(),
}
}
pub fn generation(&self, project_id: u64) -> u64 {
let gens = self.generations.lock().unwrap_or_else(|e| e.into_inner());
gens.get(&project_id).copied().unwrap_or(0)
}
pub fn bump(&self, project_id: u64) {
let mut gens = self.generations.lock().unwrap_or_else(|e| e.into_inner());
*gens.entry(project_id).or_insert(0) += 1;
}
pub async fn invalidate(&self, project_id: u64) {
self.bump(project_id);
let mut entries = self.entries.lock().await;
let stale: Vec<IndexKey> = entries
.iter()
.map(|(k, _)| *k)
.filter(|k| k.0 == project_id)
.collect();
for key in stale {
entries.pop(&key);
}
}
async fn lookup(&self, key: IndexKey, generation: u64) -> Option<Arc<SymbolIndex>> {
let mut entries = self.entries.lock().await;
entries
.get(&key)
.filter(|e| e.generation == generation)
.cloned()
}
pub async fn get_or_build(
&self,
db: &ScryerDb,
rayon_pool: &Arc<rayon::ThreadPool>,
key: IndexKey,
) -> anyhow::Result<Arc<SymbolIndex>> {
if let Some(hit) = self.lookup(key, self.generation(key.0)).await {
return Ok(hit);
}
let build_lock = {
let mut locks = self.build_locks.lock().unwrap_or_else(|e| e.into_inner());
Arc::clone(locks.entry(key).or_default())
};
let _build_guard = build_lock.lock().await;
let generation = self.generation(key.0);
if let Some(hit) = self.lookup(key, generation).await {
return Ok(hit);
}
let timer = std::time::Instant::now();
let (symbols, paths) = load_rows(db, key).await?;
let load_ms = timer.elapsed().as_millis();
let pool = Arc::clone(rayon_pool);
let index =
tokio::task::spawn_blocking(move || pool.install(|| build_index(symbols, paths)))
.await?;
tracing::debug!(
"search index built project={} package={:?} symbols={} load={}ms total={}ms",
key.0,
key.1,
index.len(),
load_ms,
timer.elapsed().as_millis()
);
let built = Arc::new(SymbolIndex { generation, index });
self.entries.lock().await.put(key, Arc::clone(&built));
Ok(built)
}
}
async fn load_rows(
db: &ScryerDb,
(project_id, package_id): IndexKey,
) -> anyhow::Result<(Vec<Symbol>, HashMap<u64, String>)> {
let mut guard = db.lock().await;
let (symbols, files) = match package_id {
Some(pkg) => {
let symbols = Symbol::filter(
Symbol::fields()
.project_id()
.eq(project_id)
.and(Symbol::fields().dependency_package_id().eq(Some(pkg))),
)
.exec(&mut *guard)
.await?;
let files = SourceFile::filter(
SourceFile::fields()
.project_id()
.eq(project_id)
.and(SourceFile::fields().dependency_package_id().eq(Some(pkg))),
)
.exec(&mut *guard)
.await?;
(symbols, files)
}
None => {
let symbols = Symbol::filter(Symbol::fields().project_id().eq(project_id))
.exec(&mut *guard)
.await?;
let files = SourceFile::filter(SourceFile::fields().project_id().eq(project_id))
.exec(&mut *guard)
.await?;
(symbols, files)
}
};
let paths = files.into_iter().map(|f| (f.id, f.path)).collect();
Ok((symbols, paths))
}
fn build_index(symbols: Vec<Symbol>, paths: HashMap<u64, String>) -> Bm25Index<SymbolDoc> {
let docs: Vec<(SymbolDoc, Vec<Vec<String>>)> = symbols
.into_par_iter()
.filter(|s| s.kind != "reexport")
.map(|s| {
let doc = SymbolDoc {
is_test: false,
id: s.id,
file_id: s.file_id,
file_path: paths.get(&s.file_id).cloned().unwrap_or_default(),
doc_snippet: s
.docstring
.as_deref()
.map(|d| d.chars().take(DOC_SNIPPET_CHARS).collect())
.unwrap_or_default(),
kind: s.kind,
name: s.name,
qualified_name: s.qualified_name,
signature: s.signature,
start_line: s.start_line,
end_line: s.end_line,
dependency_package_id: s.dependency_package_id,
};
let doc = SymbolDoc {
is_test: is_test_symbol(&doc.file_path, &doc.qualified_name),
..doc
};
let fields = vec![
tokenize_identifier(&doc.name),
tokenize_text(doc.parent_path()),
tokenize_text(&doc.signature),
tokenize_text(s.docstring.as_deref().unwrap_or("")),
tokenize_text(&doc.file_path),
];
(doc, fields)
})
.collect();
let mut builder = Bm25Builder::new(&SYMBOL_FIELD_WEIGHTS);
for (doc, fields) in docs {
builder.add(doc, &fields);
}
builder.build()
}