use crate::embed::{active_dims, active_provider_name, embed_code_dispatch};
use crate::index::symbols::SymbolQuery;
use crate::index::vector::{CodeVectorIndex, cosine_distance_to_similarity};
use anyhow::{Context, Result};
use rayon::prelude::*;
use rusqlite::Connection;
use std::collections::{HashMap, HashSet};
use std::sync::{LazyLock, RwLock};
const RRF_K: f32 = 60.0;
static VECTOR_CACHE: LazyLock<RwLock<Option<CodeVectorIndex>>> =
LazyLock::new(|| RwLock::new(None));
static PROJECT_ID_CACHE: LazyLock<RwLock<HashMap<i64, HashSet<i64>>>> =
LazyLock::new(|| RwLock::new(HashMap::new()));
#[derive(Debug, Clone, serde::Serialize)]
pub struct BrainResult {
pub symbol_id: i64,
pub name: String,
pub kind: String,
pub file: String,
pub line: i64,
pub signature: String,
pub score: f32,
pub signals: Vec<String>,
}
fn vector_index_path() -> std::path::PathBuf {
crate::data_dir::cora_data_dir().join("cora_index.usearch")
}
fn check_dimension_compat(vi_path: &std::path::Path, expected_dims: usize) -> bool {
let Ok(file) = std::fs::File::open(vi_path) else {
return false;
};
let Ok(metadata) = file.metadata() else {
return false;
};
if metadata.len() == 0 {
return false; }
let result = std::panic::catch_unwind(|| {
use std::io::{Read, Seek, SeekFrom};
let mut file = file;
let _ = file.seek(SeekFrom::Start(0));
let mut buffer = Vec::new();
let _ = file.read_to_end(&mut buffer);
if buffer.is_empty() {
return None;
}
let index = usearch::Index::restore_from_buffer(&buffer).ok()?;
Some(index.dimensions())
});
match result {
Ok(Some(disk_dims)) if disk_dims != expected_dims => {
tracing::warn!(
"⚠ Vector index dimension mismatch: on-disk={disk_dims}, current={expected_dims} ({})",
active_provider_name()
);
tracing::warn!(
" The vector index was built with a different embedding backend. \
Run `cora index` to re-index with the current backend."
);
if let Err(e) = std::fs::remove_file(vi_path) {
tracing::warn!(" Failed to remove stale index: {e}");
false
} else {
let keys_path = vi_path.with_extension("keys");
let _ = std::fs::remove_file(&keys_path);
tracing::info!(" Removed stale vector index — will create fresh on next index");
true
}
}
Ok(Some(_)) => {
false
}
_ => {
false
}
}
}
pub fn embed_project(conn: &Connection, project_id: i64) -> Result<usize> {
let vi_path = vector_index_path();
let active = active_dims();
let dims_stale = if vi_path.exists() {
check_dimension_compat(&vi_path, active)
} else {
false
};
let mut cache = VECTOR_CACHE.write().unwrap();
let vi = if let Some(ref mut cached) = *cache {
cached
} else {
let vi = CodeVectorIndex::load_or_create(&vi_path, active).context("load vector index")?;
*cache = Some(vi);
cache.as_mut().unwrap()
};
if dims_stale || (vi.is_dirty() && vi.is_empty()) {
let cleared = conn.execute("UPDATE symbols SET embed_fingerprint = NULL", [])?;
tracing::info!(
cleared,
"vector index rebuilt from empty — all projects re-embed on their next index run"
);
}
let mut stmt = conn.prepare(
"SELECT id, name, kind, signature, embed_fingerprint \
FROM symbols WHERE project_id = ?1",
)?;
let rows: Vec<(i64, String, String, String, Option<String>)> = stmt
.query_map(rusqlite::params![project_id], |row| {
Ok((
row.get(0)?,
row.get(1)?,
row.get(2)?,
row.get(3)?,
row.get(4)?,
))
})?
.filter_map(|r| r.ok())
.collect();
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let changed: Vec<(i64, String)> = rows
.iter()
.filter_map(|(sym_id, name, _kind, signature, stored_fp)| {
let text = if signature.is_empty() || signature == name {
name.clone()
} else {
format!("{name} {signature}")
};
let mut hasher = DefaultHasher::new();
text.hash(&mut hasher);
let current_fp = format!("{:016x}", hasher.finish());
if stored_fp.as_deref() == Some(¤t_fp) {
None } else {
Some((*sym_id, text))
}
})
.collect();
let total_symbols = rows.len();
let skipped = total_symbols - changed.len();
if skipped > 0 {
tracing::info!(
"Incremental embed: {total_symbols} total, {skipped} unchanged (skipped), {} changed (re-embedding)",
changed.len()
);
}
let t_compute = std::time::Instant::now();
let embedded: Vec<(i64, Vec<f32>)> = changed
.par_iter()
.map(|(sym_id, text)| {
let vec = embed_code_dispatch(text);
(*sym_id, vec)
})
.collect();
let compute_ms = t_compute.elapsed().as_millis();
let t_insert = std::time::Instant::now();
let mut count = 0;
let mut new_ids: HashSet<i64> = HashSet::with_capacity(rows.len());
for (sym_id, _, _, _, _) in &rows {
new_ids.insert(*sym_id);
}
for (sym_id, vec) in &embedded {
vi.insert(*sym_id, vec).context("insert symbol embedding")?;
count += 1;
}
let insert_ms = t_insert.elapsed().as_millis();
tracing::debug!(
"embed_compute={}ms, usearch_insert={}ms, re-embedded={}, total={}, skipped={}, dims={}, provider={}",
compute_ms,
insert_ms,
count,
total_symbols,
skipped,
active,
active_provider_name()
);
if vi.is_dirty() {
vi.save().context("save vector index")?;
}
let mut update_fp = conn.prepare("UPDATE symbols SET embed_fingerprint = ?2 WHERE id = ?1")?;
for (sym_id, text) in &changed {
let mut hasher = DefaultHasher::new();
text.hash(&mut hasher);
let fp = format!("{:016x}", hasher.finish());
update_fp.execute(rusqlite::params![sym_id, fp])?;
}
PROJECT_ID_CACHE
.write()
.unwrap()
.insert(project_id, new_ids);
let provider = active_provider_name();
let tier = if provider.contains("pretrained") {
"pretrained"
} else {
"static"
};
conn.execute(
"UPDATE projects SET embedding_tier = ?3, embedding_dims = ?1, \
embedding_provider = ?4, last_embedded_at = datetime('now') WHERE id = ?2",
rusqlite::params![active, project_id, tier, provider],
)?;
tracing::info!(
"Embedded {count}/{total_symbols} symbols for project {project_id} ({skipped} unchanged, provider={provider}, dims={active})",
);
Ok(count)
}
pub fn brain_search(
conn: &Connection,
project_id: i64,
query: &str,
limit: usize,
) -> Result<Vec<BrainResult>> {
let limit = limit.min(50);
let fetch_limit = limit * 2;
let fts_hits = fts5_search(conn, project_id, query, fetch_limit);
let vec_hits = vector_search(conn, project_id, query, fetch_limit);
let graph_hits = graph_proximity_search(conn, project_id, &fts_hits, fetch_limit);
let mut fused: HashMap<i64, (f32, Vec<String>)> = HashMap::new();
for (rank, (id, _score)) in fts_hits.iter().enumerate() {
let rrf = 1.0 / (RRF_K + (rank as f32 + 1.0));
let entry = fused.entry(*id).or_insert((0.0, Vec::new()));
entry.0 += rrf;
entry.1.push("fts".into());
}
for (rank, (id, _sim)) in vec_hits.iter().enumerate() {
let rrf = 1.0 / (RRF_K + (rank as f32 + 1.0));
let entry = fused.entry(*id).or_insert((0.0, Vec::new()));
entry.0 += rrf;
entry.1.push("vector".into());
}
for (rank, (id, _depth)) in graph_hits.iter().enumerate() {
let rrf = 1.0 / (RRF_K + (rank as f32 + 1.0));
let entry = fused.entry(*id).or_insert((0.0, Vec::new()));
entry.0 += rrf;
entry.1.push("graph".into());
}
let mut ranked: Vec<_> = fused.into_iter().collect();
ranked.sort_by(|a, b| {
b.1.0
.partial_cmp(&a.1.0)
.unwrap_or(std::cmp::Ordering::Equal)
});
ranked.truncate(limit);
let results = batch_get_symbols(conn, &ranked);
Ok(results)
}
fn fts5_search(conn: &Connection, project_id: i64, query: &str, limit: usize) -> Vec<(i64, f64)> {
let sq = SymbolQuery::text(query);
match crate::index::search(conn, project_id, &sq) {
Ok(results) => results
.into_iter()
.take(limit)
.map(|r| (r.symbol.id, r.score))
.collect(),
Err(e) => {
tracing::warn!("FTS5 search error: {e}");
Vec::new()
}
}
}
pub fn vector_index_needs_rebuild() -> bool {
if crate::index::vector::current_vector_store() != crate::index::vector::VectorStoreKind::Vecq {
return false;
}
let path = vector_index_path().with_extension("vecq");
path.exists() && crate::index::vector::vecq_file_needs_rebuild(&path, active_dims())
}
fn ensure_vector_cache() -> std::sync::RwLockReadGuard<'static, Option<CodeVectorIndex>> {
{
let cache = VECTOR_CACHE.read().unwrap();
if cache.is_some() {
return cache;
}
}
let mut cache = VECTOR_CACHE.write().unwrap();
if cache.is_none() {
let dims = active_dims();
match CodeVectorIndex::load_or_create(&vector_index_path(), dims) {
Ok(vi) => *cache = Some(vi),
Err(e) => tracing::warn!("vector index load failed (search degrades to FTS): {e}"),
}
}
drop(cache);
VECTOR_CACHE.read().unwrap()
}
fn vector_search(conn: &Connection, project_id: i64, query: &str, limit: usize) -> Vec<(i64, f32)> {
let cache = ensure_vector_cache();
let vi = match cache.as_ref() {
Some(v) if !v.is_empty() => v,
_ => return Vec::new(),
};
let vec = embed_code_dispatch(query);
if vi.dims() != vec.len() {
return Vec::new();
}
let over_fetch = (limit * 5).max(50);
let raw = vi.search(&vec, over_fetch);
let project_ids = {
let cache = PROJECT_ID_CACHE.read().unwrap();
cache.get(&project_id).cloned()
};
let project_ids = match project_ids {
Some(ids) => ids,
None => {
let mut stmt = match conn.prepare("SELECT id FROM symbols WHERE project_id = ?1") {
Ok(s) => s,
Err(_) => return Vec::new(),
};
let rows: Vec<i64> = stmt
.query_map([project_id], |r| r.get(0))
.ok()
.map(|rows| rows.filter_map(|r| r.ok()).collect())
.unwrap_or_default();
let ids: HashSet<i64> = rows.into_iter().collect();
PROJECT_ID_CACHE
.write()
.unwrap()
.insert(project_id, ids.clone());
ids
}
};
raw.into_iter()
.filter(|(sym_id, _)| project_ids.contains(sym_id))
.take(limit)
.map(|(sym_id, dist)| (sym_id, cosine_distance_to_similarity(dist)))
.collect()
}
fn graph_proximity_search(
conn: &Connection,
project_id: i64,
fts_results: &[(i64, f64)],
limit: usize,
) -> Vec<(i64, f32)> {
if fts_results.is_empty() {
return Vec::new();
}
let top_id = fts_results[0].0;
let top_name: String = match conn.query_row(
"SELECT name FROM symbols WHERE id = ?1",
rusqlite::params![top_id],
|row| row.get(0),
) {
Ok(n) => n,
Err(_) => return Vec::new(),
};
let Ok(mut stmt) = conn.prepare(
"SELECT DISTINCT s.id FROM symbols s
JOIN edges e ON (e.target = s.name OR e.source = s.name)
WHERE (e.source = ?1 OR e.target = ?1) AND e.project_id = ?2
AND s.project_id = ?2 AND s.id != ?3
LIMIT ?4",
) else {
return Vec::new();
};
let ids: Vec<(i64, f32)> = stmt
.query_map(
rusqlite::params![top_name, project_id, top_id, limit],
|row| row.get(0),
)
.ok()
.map(|rows| {
rows.filter_map(|r| r.ok())
.enumerate()
.map(|(i, id)| (id, 1.0 / (i as f32 + 2.0)))
.collect()
})
.unwrap_or_default();
ids
}
fn batch_get_symbols(conn: &Connection, ranked: &[(i64, (f32, Vec<String>))]) -> Vec<BrainResult> {
if ranked.is_empty() {
return Vec::new();
}
let placeholders = ranked.iter().map(|_| "?").collect::<Vec<_>>().join(", ");
let sql = format!(
"SELECT id, name, kind, file, line, signature \
FROM symbols WHERE id IN ({placeholders})"
);
let ids: Vec<i64> = ranked.iter().map(|(id, _)| *id).collect();
let params: Vec<&dyn rusqlite::types::ToSql> = ids
.iter()
.map(|id| id as &dyn rusqlite::types::ToSql)
.collect();
let mut stmt = match conn.prepare(&sql) {
Ok(s) => s,
Err(e) => {
tracing::warn!("batch_get_symbols prepare error: {e}");
return Vec::new();
}
};
let mut score_map: HashMap<i64, (f32, Vec<String>)> = HashMap::with_capacity(ranked.len());
for (id, (score, signals)) in ranked {
score_map.insert(*id, (*score, signals.clone()));
}
let results: Vec<BrainResult> = stmt
.query_map(params.as_slice(), |row| {
Ok((
row.get::<_, i64>(0)?,
row.get::<_, String>(1)?,
row.get::<_, String>(2)?,
row.get::<_, String>(3)?,
row.get::<_, i64>(4)?,
row.get::<_, String>(5)?,
))
})
.ok()
.map(|rows| {
rows.filter_map(|r| r.ok())
.filter_map(|(id, name, kind, file, line, signature)| {
score_map.remove(&id).map(|(score, mut signals)| {
signals.sort();
signals.dedup();
BrainResult {
symbol_id: id,
name,
kind,
file,
line,
signature,
score,
signals,
}
})
})
.collect()
})
.unwrap_or_default();
results
}