From 56f2c89bdc257794006efa3f93557ddbd92dbc70 Mon Sep 17 00:00:00 2001 From: ajianaz Date: Mon, 27 Jul 2026 08:16:23 +0700 Subject: [PATCH] fix(brain): filter vector search by project_id (#382) vector_search() in brain.rs searched the global usearch index without filtering by project_id, causing cross-project symbols to leak into brain search results even when the correct project_id was resolved. Fix: over-fetch from vector index (5x limit, min 50), then filter results against the set of symbol IDs belonging to the target project. Validated with real use case: - Before: brain search from cora-code returned symbols from uteke-core, nginjen-core (different projects in the same global index) - After: brain search returns only cora-code symbols Fixes #382 --- src/index/brain.rs | 29 ++++++++++++++++++++++++----- 1 file changed, 24 insertions(+), 5 deletions(-) diff --git a/src/index/brain.rs b/src/index/brain.rs index 639e07a..cd71d10 100644 --- a/src/index/brain.rs +++ b/src/index/brain.rs @@ -92,7 +92,7 @@ pub fn brain_search( let fetch_limit = limit * 2; let fts_hits = fts5_search(conn, project_id, query, fetch_limit); - let vec_hits = vector_search(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); // ── RRF Fusion ────────────────────────────────────────────────────── @@ -168,8 +168,9 @@ fn fts5_search(conn: &Connection, project_id: i64, query: &str, limit: usize) -> } } -/// usearch vector search → (symbol_id, cosine_similarity) pairs. -fn vector_search(query: &str, limit: usize) -> Vec<(i64, f32)> { +/// usearch vector search → (symbol_id, cosine_similarity) pairs, filtered to project. +/// Over-fetches from global vector index then filters by project_id via DB lookup. +fn vector_search(conn: &Connection, project_id: i64, query: &str, limit: usize) -> Vec<(i64, f32)> { let vi_path = vector_index_path(); if !vi_path.exists() { return Vec::new(); @@ -189,8 +190,26 @@ fn vector_search(query: &str, limit: usize) -> Vec<(i64, f32)> { let embedding = embed_code(query); let vec: Vec = embedding.as_slice().iter().map(|&v| v as f32).collect(); - vi.search(&vec, limit) - .into_iter() + // Over-fetch to compensate for post-filter by project_id. + let over_fetch = (limit * 5).max(50); + let raw = vi.search(&vec, over_fetch); + + // Build a set of symbol IDs belonging to this project. + let project_ids: std::collections::HashSet = { + let mut stmt = conn + .prepare("SELECT id FROM symbols WHERE project_id = ?1") + .unwrap(); + let rows: Vec = stmt + .query_map([project_id], |r| r.get(0)) + .unwrap() + .filter_map(|r| r.ok()) + .collect(); + rows.into_iter().collect() + }; + + 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() }