use std::path::Path;
use serde::Serialize;
use crate::index::IndexResult;
#[cfg(any(feature = "cli", feature = "mcp", test))]
use crate::kit::{AsyncKit, AsyncReady, IndexerModule};
#[cfg(any(feature = "cli", feature = "mcp", feature = "lsp", test))]
use crate::service::error::CodeNexusError;
#[cfg(any(feature = "cli", feature = "mcp", feature = "lsp", test))]
use crate::storage::{QualityChecker, Repository};
#[cfg(any(feature = "cli", feature = "mcp"))]
use crate::service::error::{kit_not_initialized, to_api_error, wrap_kit_error};
#[cfg(any(feature = "cli", feature = "mcp"))]
use crate::service::runtime::kit;
#[cfg(any(feature = "cli", feature = "mcp"))]
use crate::storage::StorageConfig;
#[cfg(feature = "lsp")]
use crate::lsp::LspProvider;
#[cfg(feature = "cli")]
use sdforge::forge;
#[cfg(any(feature = "cli", feature = "mcp"))]
use sdforge::prelude::ApiError;
#[derive(Debug, Clone, Serialize, PartialEq)]
pub struct IndexOutput {
pub project_id: String,
pub files_indexed: usize,
pub files_skipped: usize,
pub nodes_created: usize,
pub edges_created: usize,
pub duration_ms: u64,
}
impl From<IndexResult> for IndexOutput {
fn from(r: IndexResult) -> Self {
Self {
project_id: r.project_id,
files_indexed: r.files_indexed,
files_skipped: r.files_skipped,
nodes_created: r.nodes_created,
edges_created: r.edges_created,
duration_ms: r.duration_ms,
}
}
}
#[cfg(feature = "lsp")]
const MAX_LSP_PROVIDERS: usize = 8;
#[cfg(feature = "lsp")]
fn build_lsp_providers() -> Vec<(&'static str, Box<dyn LspProvider>)> {
use crate::lsp::{
ClangdClient, FortlsClient, GoplsClient, JdtlsClient, PyrightClient, RustAnalyzerClient,
TypeScriptLanguageClient,
};
let providers: Vec<(&'static str, Box<dyn LspProvider>)> = vec![
("rs", Box::new(RustAnalyzerClient::new())),
("py", Box::new(PyrightClient::new())),
("c", Box::new(ClangdClient::new())),
("cpp", Box::new(ClangdClient::new())),
("go", Box::new(GoplsClient::new())),
("ts", Box::new(TypeScriptLanguageClient::new())),
("f90", Box::new(FortlsClient::new())),
("java", Box::new(JdtlsClient::new())),
];
assert!(
providers.len() <= MAX_LSP_PROVIDERS,
"MAX_LSP_PROVIDERS={} but {} providers wired; raise the const or remove a provider",
MAX_LSP_PROVIDERS,
providers.len()
);
providers
}
#[cfg(feature = "lsp")]
fn build_symbol_queries(project: &str) -> [String; 2] {
use crate::storage::schema::escape_cypher_string;
let proj = escape_cypher_string(project);
[
format!(
"MATCH (n:Function) WHERE n.project = '{proj}' \
AND n.filePath IS NOT NULL AND n.startLine IS NOT NULL \
RETURN n.id AS id, n.filePath AS filePath, n.startLine AS startLine;"
),
format!(
"MATCH (n:Method) WHERE n.project = '{proj}' \
AND n.filePath IS NOT NULL AND n.startLine IS NOT NULL \
RETURN n.id AS id, n.filePath AS filePath, n.startLine AS startLine;"
),
]
}
#[cfg(feature = "lsp")]
fn select_provider_for_ext<'a>(
ext_map: &std::collections::HashMap<&'static str, &'a dyn LspProvider>,
ext: &str,
) -> Option<&'a dyn LspProvider> {
ext_map.get(ext).copied()
}
#[cfg(feature = "lsp")]
fn build_semantic_type_update(id: &str, project: &str, text: &str) -> String {
use crate::storage::schema::escape_cypher_string;
format!(
"MATCH (n {{id: '{id}', project: '{proj}'}}) \
SET n.semantic_type = '{sem}';",
id = escape_cypher_string(id),
proj = escape_cypher_string(project),
sem = escape_cypher_string(text),
)
}
#[cfg(feature = "lsp")]
const LSP_HOVER_BATCH_SIZE: usize = 500;
#[cfg(feature = "lsp")]
fn build_batch_semantic_type_update(batch: &[(String, String)], project: &str) -> String {
use crate::storage::schema::escape_cypher_string;
use std::fmt::Write as _;
if batch.is_empty() {
return String::new();
}
let proj = escape_cypher_string(project);
let estimated: usize = batch
.iter()
.map(|(id, sem)| id.len() + sem.len() + 16)
.sum::<usize>()
+ batch.len()
+ 80; let mut rows = String::with_capacity(estimated);
for (i, (id, sem)) in batch.iter().enumerate() {
if i > 0 {
rows.push_str(", ");
}
let _ = write!(
rows,
"{{id: '{}', sem: '{}'}}",
escape_cypher_string(id),
escape_cypher_string(sem),
);
}
format!(
"UNWIND [{rows}] AS row \
MATCH (n {{id: row.id, project: '{proj}'}}) \
SET n.semantic_type = row.sem;"
)
}
#[cfg(feature = "lsp")]
fn flush_semantic_type_batch<F>(
exec: F,
batch: &mut Vec<(String, String)>,
project: &str,
enhanced: &mut u32,
skipped: &mut u32,
) where
F: Fn(&str) -> crate::storage::Result<()>,
{
if batch.is_empty() {
return;
}
let processed = batch.len();
let stmt = build_batch_semantic_type_update(batch, project);
if stmt.is_empty() {
debug_assert!(false, "non-empty batch must produce non-empty statement");
batch.clear();
return;
}
match exec(&stmt) {
Ok(()) => {
*enhanced += processed as u32;
}
Err(e) => {
eprintln!(
"[warn] P-01 batch UNWIND failed ({e:?}), falling back to per-row execute for {processed} symbols"
);
for (id, sem) in batch.iter() {
let row_stmt = build_semantic_type_update(id, project, sem);
if exec(&row_stmt).is_ok() {
*enhanced += 1;
} else {
*skipped += 1;
}
}
}
}
batch.clear();
}
#[cfg(feature = "lsp")]
fn start_active_providers_parallel(
providers: &[(&'static str, Box<dyn LspProvider>)],
active_exts: &std::collections::HashSet<&'static str>,
workspace: &Path,
) -> u32 {
use crate::lsp::LspError;
let active: Vec<(&'static str, &dyn LspProvider)> = providers
.iter()
.filter(|(ext, _)| active_exts.contains(ext))
.map(|(ext, p)| (*ext, p.as_ref()))
.collect();
if active.is_empty() {
return 0;
}
let results: Vec<(&'static str, Result<(), LspError>)> = std::thread::scope(|s| {
let handles: Vec<_> = active
.iter()
.map(|(ext, provider)| {
s.spawn(move || {
let res = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
provider.start(workspace)
}));
match res {
Ok(outer) => (*ext, outer),
Err(payload) => {
let msg = payload
.downcast_ref::<String>()
.cloned()
.or_else(|| {
payload
.downcast_ref::<&'static str>()
.map(|s| s.to_string())
})
.unwrap_or_else(|| "unknown panic payload".to_string());
(
*ext,
Err(LspError::ServerStart(format!("provider panicked: {msg}"))),
)
}
}
})
})
.collect();
handles
.into_iter()
.map(|h| h.join().expect("scope thread panicked unexpectedly"))
.collect()
});
let mut started_count: u32 = 0;
for (ext, result) in results {
match result {
Ok(()) => started_count += 1,
Err(LspError::ServerStart(msg)) => {
eprintln!("[warn] LSP server start failed for {ext} (degrading): {msg}");
}
Err(other) => {
eprintln!("[warn] LSP server start failed for {ext} (degrading): {other:?}");
}
}
}
started_count
}
#[cfg(feature = "lsp")]
fn extract_lsp_row_fields(row: &[serde_json::Value]) -> Option<(String, String, u64)> {
let id = row.first().and_then(|v| v.as_str())?.to_string();
let file_path_str = row.get(1).and_then(|v| v.as_str())?.to_string();
let start_line = row.get(2).and_then(|v| v.as_u64())?;
Some((id, file_path_str, start_line))
}
#[cfg(feature = "lsp")]
fn resolve_abs_file_path(workspace: &Path, file_path_str: &str) -> std::path::PathBuf {
let file_path = Path::new(file_path_str);
if file_path.is_absolute() {
file_path.to_path_buf()
} else {
workspace.join(file_path)
}
}
#[cfg(feature = "lsp")]
struct PathInterner {
map: std::collections::HashMap<std::path::PathBuf, std::sync::Arc<std::path::PathBuf>>,
}
#[cfg(feature = "lsp")]
impl PathInterner {
#[must_use]
fn new() -> Self {
Self {
map: std::collections::HashMap::new(),
}
}
fn intern(&mut self, path: std::path::PathBuf) -> std::sync::Arc<std::path::PathBuf> {
if let Some(arc) = self.map.get(&path) {
return arc.clone();
}
let arc = std::sync::Arc::new(path);
self.map.insert((*arc).clone(), arc.clone());
arc
}
}
#[cfg(feature = "lsp")]
#[allow(clippy::result_large_err)]
fn enhance_with_lsp(
workspace: &Path,
repo: &Repository,
project: &str,
) -> Result<(), CodeNexusError> {
use crate::lsp::LspError;
let providers = build_lsp_providers();
let queries = build_symbol_queries(project);
let mut rows = Vec::new();
for q in &queries {
let r = repo.connection().query(q).map_err(|e| {
CodeNexusError::Storage(crate::storage::StorageError::Query(e.to_string()))
})?;
rows.extend(r);
}
let provider_exts_set: std::collections::HashSet<&'static str> =
providers.iter().map(|(e, _)| *e).collect();
let ext_map: std::collections::HashMap<&'static str, &dyn LspProvider> = providers
.iter()
.map(|(ext, p)| (*ext, p.as_ref()))
.collect();
let mut interner = PathInterner::new();
let mut entries: Vec<(String, std::sync::Arc<std::path::PathBuf>, u32)> =
Vec::with_capacity(rows.len());
let mut skipped: u32 = 0;
for row in &rows {
let Some((id, file_path_str, start_line)) = extract_lsp_row_fields(row) else {
skipped += 1;
continue;
};
let abs_file = interner.intern(resolve_abs_file_path(workspace, &file_path_str));
let line = u32::try_from(start_line).unwrap_or(0);
entries.push((id, abs_file, line));
}
let active_exts = collect_active_extensions(&entries, &provider_exts_set);
let started_count = start_active_providers_parallel(&providers, &active_exts, workspace);
eprintln!(
"[info] LSP enhancement: started {started_count} of {} configured servers for {} symbols",
providers.len(),
entries.len()
);
let mut enhanced: u32 = 0;
let mut batch: Vec<(String, String)> = Vec::with_capacity(LSP_HOVER_BATCH_SIZE);
if !entries.is_empty() {
for (id, abs_file, line) in &entries {
let ext = abs_file.extension().and_then(|e| e.to_str()).unwrap_or("");
let Some(client) = select_provider_for_ext(&ext_map, ext) else {
skipped += 1;
continue;
};
match client.hover(abs_file, *line, 0) {
Ok(Some(hover)) => {
if let Some(text) = crate::lsp::extract_hover_text(&hover) {
batch.push((id.clone(), text));
if batch.len() >= LSP_HOVER_BATCH_SIZE {
flush_semantic_type_batch(
|q| repo.connection().execute(q),
&mut batch,
project,
&mut enhanced,
&mut skipped,
);
}
} else {
skipped += 1;
}
}
Ok(None) => {
skipped += 1;
}
Err(LspError::Timeout(_)) | Err(LspError::Communication(_)) => {
skipped += 1;
}
Err(LspError::ServerStart(_)) => {
skipped += 1;
}
Err(LspError::NotImplemented(_)) => {
skipped += 1;
}
}
}
if !batch.is_empty() {
flush_semantic_type_batch(
|q| repo.connection().execute(q),
&mut batch,
project,
&mut enhanced,
&mut skipped,
);
}
}
eprintln!("[info] LSP enhancement: {enhanced} symbol(s) enhanced, {skipped} skipped");
for (_ext, provider) in &providers {
let _ = provider.shutdown();
}
Ok(())
}
#[cfg(feature = "lsp")]
fn collect_active_extensions<'a>(
entries: &[(String, std::sync::Arc<std::path::PathBuf>, u32)],
provider_exts: &std::collections::HashSet<&'a str>,
) -> std::collections::HashSet<&'a str> {
let mut active: std::collections::HashSet<&'a str> = std::collections::HashSet::new();
for (_, abs_file, _) in entries {
let ext = abs_file.extension().and_then(|e| e.to_str()).unwrap_or("");
if let Some(&provider_ext) = provider_exts.get(ext) {
active.insert(provider_ext);
}
}
active
}
#[cfg(any(feature = "cli", feature = "mcp", test))]
#[allow(clippy::result_large_err)]
#[cfg_attr(not(feature = "lsp"), allow(unused_variables))]
pub(crate) fn index_core(
kit: &AsyncKit<AsyncReady>,
db_path: &Path,
path: &str,
name: &str,
force: bool,
lsp: bool,
ram_first: bool,
) -> Result<IndexOutput, CodeNexusError> {
let path_ref = Path::new(path);
let indexer = kit.require::<IndexerModule>()?;
let result = if ram_first {
indexer.index_ram_first(path_ref, name, force)?
} else {
indexer.index(path_ref, name, force)?
};
let fresh_repo = Repository::open(db_path)?;
let checker = QualityChecker::new(&fresh_repo);
let dq_report = checker.run_all()?;
if !dq_report.is_clean() {
eprintln!("Data quality violations found:");
for violation in &dq_report.violations {
eprintln!(
" [{}] {} (project: {})",
violation.rule,
violation.message,
violation.project.as_deref().unwrap_or("N/A")
);
}
}
if let Err(err) = fresh_repo.connection().execute("CHECKPOINT;") {
eprintln!("[warn] post-quality-check checkpoint failed: {err}");
}
#[cfg(feature = "lsp")]
if lsp {
if let Err(err) = enhance_with_lsp(path_ref, &fresh_repo, name) {
eprintln!("[warn] LSP enhancement aborted: {err}");
}
}
Ok(IndexOutput::from(result))
}
#[cfg(feature = "cli")]
#[forge(
name = "index",
version = "0.3.5",
description = "Index a codebase into the knowledge graph.",
cli = true
)]
async fn index(
path: String,
name: String,
force: bool,
lsp: bool,
embed: bool,
ram_first: bool,
fresh: bool,
) -> Result<(), ApiError> {
if embed {
eprintln!(
"[warn] --embed flag is deprecated; embedding is controlled by \
the `embed` cargo feature (rebuild with --features embed to enable)"
);
}
let _ = fresh;
let kit = kit().ok_or_else(kit_not_initialized)?;
let storage_config = kit
.config::<StorageConfig>()
.map_err(|e| wrap_kit_error("Failed to resolve storage config", e))?;
let db_path = storage_config.db_path.clone();
let output = index_core(&kit, &db_path, &path, &name, force, lsp, ram_first)
.map_err(|e| to_api_error(e, "index_error"))?;
let json = serde_json::to_string(&output)
.map_err(|e| to_api_error(CodeNexusError::from(e), "index_error"))?;
println!("{json}");
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "lsp")]
mod path_interner_tests {
use super::*;
use std::path::PathBuf;
use std::sync::Arc;
#[test]
fn intern_returns_same_arc_for_same_path() {
let mut interner = PathInterner::new();
let a = interner.intern(PathBuf::from("src/lib.rs"));
let b = interner.intern(PathBuf::from("src/lib.rs"));
assert!(
Arc::ptr_eq(&a, &b),
"same path must return pointer-equal Arc"
);
}
#[test]
fn intern_returns_different_arc_for_different_path() {
let mut interner = PathInterner::new();
let a = interner.intern(PathBuf::from("src/a.rs"));
let b = interner.intern(PathBuf::from("src/b.rs"));
assert!(!Arc::ptr_eq(&a, &b), "different paths must not share Arc");
}
#[test]
fn intern_increments_strong_count() {
let mut interner = PathInterner::new();
let first = interner.intern(PathBuf::from("src/lib.rs"));
let second = interner.intern(PathBuf::from("src/lib.rs"));
let third = interner.intern(PathBuf::from("src/lib.rs"));
assert!(Arc::ptr_eq(&first, &second));
assert!(Arc::ptr_eq(&second, &third));
let count = Arc::strong_count(&third);
assert!(
count >= 3,
"strong_count after 3 interns of same path should be >= 3 (map + first + second + third - 1 shared), got {count}"
);
}
#[test]
fn intern_handles_empty_path() {
let mut interner = PathInterner::new();
let arc = interner.intern(PathBuf::from(""));
assert_eq!(arc.as_path(), std::path::Path::new(""));
}
#[test]
fn intern_handles_unicode_path() {
let mut interner = PathInterner::new();
let arc = interner.intern(PathBuf::from("src/䏿–‡/文件.rs"));
assert_eq!(arc.as_path(), std::path::Path::new("src/䏿–‡/文件.rs"));
let arc2 = interner.intern(PathBuf::from("src/䏿–‡/文件.rs"));
assert!(Arc::ptr_eq(&arc, &arc2));
}
#[test]
fn path_interner_is_send_sync() {
fn _assert_send_sync<T: Send + Sync>() {}
_assert_send_sync::<PathInterner>();
}
}
#[test]
fn index_output_from_index_result_maps_all_fields() {
let result = IndexResult {
project_id: "p1".into(),
files_indexed: 10,
files_skipped: 2,
nodes_created: 100,
edges_created: 50,
duration_ms: 5000,
};
let output = IndexOutput::from(result);
assert_eq!(output.project_id, "p1");
assert_eq!(output.files_indexed, 10);
assert_eq!(output.files_skipped, 2);
assert_eq!(output.nodes_created, 100);
assert_eq!(output.edges_created, 50);
assert_eq!(output.duration_ms, 5000);
}
#[test]
fn index_output_from_handles_zero_values() {
let result = IndexResult {
project_id: "".into(),
files_indexed: 0,
files_skipped: 0,
nodes_created: 0,
edges_created: 0,
duration_ms: 0,
};
let output = IndexOutput::from(result);
assert_eq!(output.project_id, "");
assert_eq!(output.files_indexed, 0);
assert_eq!(output.nodes_created, 0);
assert_eq!(output.edges_created, 0);
assert_eq!(output.duration_ms, 0);
}
#[test]
fn index_output_serializes_to_json() {
let output = IndexOutput {
project_id: "p1".into(),
files_indexed: 10,
files_skipped: 2,
nodes_created: 100,
edges_created: 50,
duration_ms: 5000,
};
let json = serde_json::to_string(&output).unwrap();
assert!(json.contains("\"project_id\":\"p1\""));
assert!(json.contains("\"files_indexed\":10"));
assert!(json.contains("\"files_skipped\":2"));
assert!(json.contains("\"nodes_created\":100"));
assert!(json.contains("\"edges_created\":50"));
assert!(json.contains("\"duration_ms\":5000"));
}
#[test]
fn index_output_from_preserves_large_values() {
let result = IndexResult {
project_id: "uuid-v7-12345".into(),
files_indexed: usize::MAX,
files_skipped: usize::MAX,
nodes_created: usize::MAX,
edges_created: usize::MAX,
duration_ms: u64::MAX,
};
let output = IndexOutput::from(result);
assert_eq!(output.files_indexed, usize::MAX);
assert_eq!(output.duration_ms, u64::MAX);
}
#[cfg(feature = "lang-rust")]
#[test]
fn index_core_indexes_rust_project() {
use crate::kit::{build_kit, KitBootstrapConfig};
use std::fs;
use tempfile::TempDir;
let src_dir = TempDir::new().unwrap();
let src_path = src_dir.path();
fs::write(
src_path.join("main.rs"),
"fn main() { println!(\"hello\"); }\n",
)
.unwrap();
let db_dir = TempDir::new().unwrap();
let db_path = db_dir.path().join("index_testdb");
let config = KitBootstrapConfig::new(db_path.clone());
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(build_kit(&config))
.expect("build_kit");
let output = index_core(
&kit,
&db_path,
src_path.to_str().unwrap(),
"test_project",
false,
false,
false,
)
.expect("index should succeed");
assert!(!output.project_id.is_empty());
assert!(output.files_indexed >= 1);
assert!(output.nodes_created > 0);
}
#[cfg(feature = "lang-rust")]
#[test]
fn index_core_force_reindexes_unchanged_files() {
use crate::kit::{build_kit, KitBootstrapConfig};
use std::fs;
use tempfile::TempDir;
let src_dir = TempDir::new().unwrap();
let src_path = src_dir.path();
fs::write(
src_path.join("lib.rs"),
"pub fn add(a: i32, b: i32) -> i32 { a + b }\n",
)
.unwrap();
let db_dir = TempDir::new().unwrap();
let db_path = db_dir.path().join("index_force_testdb");
let config = KitBootstrapConfig::new(db_path.clone());
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(build_kit(&config))
.expect("build_kit");
let output1 = index_core(
&kit,
&db_path,
src_path.to_str().unwrap(),
"force_project",
false,
false,
false,
)
.expect("first index should succeed");
assert!(output1.files_indexed >= 1);
let output2 = index_core(
&kit,
&db_path,
src_path.to_str().unwrap(),
"force_project",
true,
false,
false,
)
.expect("forced reindex should succeed");
assert!(output2.files_indexed >= 1);
}
#[cfg(feature = "lang-rust")]
#[test]
fn index_core_handles_empty_directory() {
use crate::kit::{build_kit, KitBootstrapConfig};
use tempfile::TempDir;
let src_dir = TempDir::new().unwrap();
let src_path = src_dir.path();
let db_dir = TempDir::new().unwrap();
let db_path = db_dir.path().join("index_empty_testdb");
let config = KitBootstrapConfig::new(db_path.clone());
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(build_kit(&config))
.expect("build_kit");
let output = index_core(
&kit,
&db_path,
src_path.to_str().unwrap(),
"empty_project",
false,
false,
false,
)
.expect("index should succeed on empty dir");
assert_eq!(output.files_indexed, 0);
assert_eq!(output.nodes_created, 0);
}
#[cfg(feature = "lang-rust")]
#[test]
fn index_core_with_ram_first_indexes_rust_project() {
use crate::kit::{build_kit, KitBootstrapConfig};
use std::fs;
use tempfile::TempDir;
let src_dir = TempDir::new().unwrap();
let src_path = src_dir.path();
fs::write(
src_path.join("lib.rs"),
"pub fn add(a: i32, b: i32) -> i32 { a + b }\n",
)
.unwrap();
let db_dir = TempDir::new().unwrap();
let db_path = db_dir.path().join("index_ram_first_testdb");
let config = KitBootstrapConfig::new(db_path.clone());
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(build_kit(&config))
.expect("build_kit");
let output = index_core(
&kit,
&db_path,
src_path.to_str().unwrap(),
"ram_first_project",
false,
false,
true,
)
.expect("ram_first index should succeed");
assert!(!output.project_id.is_empty());
assert!(output.files_indexed >= 1);
assert!(output.nodes_created > 0);
}
#[cfg(feature = "lang-rust")]
#[test]
fn index_core_with_ram_first_and_force_reindexes() {
use crate::kit::{build_kit, KitBootstrapConfig};
use std::fs;
use tempfile::TempDir;
let src_dir = TempDir::new().unwrap();
let src_path = src_dir.path();
fs::write(
src_path.join("main.rs"),
"fn main() { println!(\"hello\"); }\n",
)
.unwrap();
let db_dir = TempDir::new().unwrap();
let db_path = db_dir.path().join("index_ram_force_testdb");
let config = KitBootstrapConfig::new(db_path.clone());
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(build_kit(&config))
.expect("build_kit");
let output1 = index_core(
&kit,
&db_path,
src_path.to_str().unwrap(),
"ram_force_project",
false,
false,
true,
)
.expect("first ram_first index should succeed");
assert!(output1.files_indexed >= 1);
let output2 = index_core(
&kit,
&db_path,
src_path.to_str().unwrap(),
"ram_force_project",
true,
false,
true,
)
.expect("forced ram_first reindex should succeed");
assert!(output2.files_indexed >= 1);
}
#[cfg(feature = "lsp")]
#[test]
fn build_lsp_providers_returns_8_providers() {
let providers = build_lsp_providers();
assert_eq!(providers.len(), 8);
let exts: Vec<&str> = providers.iter().map(|(e, _)| *e).collect();
assert!(exts.contains(&"rs"));
assert!(exts.contains(&"py"));
assert!(exts.contains(&"go"));
assert!(exts.contains(&"java"));
}
#[cfg(feature = "lsp")]
#[test]
fn build_symbol_queries_contains_project_name() {
let queries = build_symbol_queries("demo");
assert!(queries[0].contains("demo"));
assert!(queries[1].contains("demo"));
}
#[cfg(feature = "lsp")]
#[test]
fn build_symbol_queries_escapes_special_chars() {
let queries = build_symbol_queries("demo' OR '1'='1");
assert!(!queries[0].contains("' OR '1'='1"));
}
#[cfg(feature = "lsp")]
#[test]
fn select_provider_for_ext_returns_some_for_rs() {
let providers = build_lsp_providers();
let ext_map: std::collections::HashMap<&'static str, &dyn LspProvider> = providers
.iter()
.map(|(ext, p)| (*ext, p.as_ref()))
.collect();
let provider = select_provider_for_ext(&ext_map, "rs").expect("rs must map to a provider");
let _ = provider.shutdown();
drop(providers);
}
#[cfg(feature = "lsp")]
#[test]
fn select_provider_for_ext_returns_none_for_unknown() {
let providers = build_lsp_providers();
let ext_map: std::collections::HashMap<&'static str, &dyn LspProvider> = providers
.iter()
.map(|(ext, p)| (*ext, p.as_ref()))
.collect();
assert!(
select_provider_for_ext(&ext_map, "md").is_none(),
"md must NOT map to any provider"
);
assert!(
select_provider_for_ext(&ext_map, "toml").is_none(),
"toml must NOT map to any provider"
);
assert!(
select_provider_for_ext(&ext_map, "unknown").is_none(),
"unknown ext must NOT map to any provider"
);
for (_, p) in &providers {
let _ = p.shutdown();
}
}
#[cfg(feature = "lsp")]
#[test]
fn build_semantic_type_update_escapes_id() {
let update = build_semantic_type_update("id'with'quotes", "demo", "type_text");
assert!(update.contains(r"id\'with\'quotes"));
}
#[cfg(feature = "lsp")]
#[test]
fn extract_lsp_row_fields_returns_all_fields_when_valid() {
let row = vec![
serde_json::Value::String("sym_id".into()),
serde_json::Value::String("/src/main.rs".into()),
serde_json::Value::Number(serde_json::Number::from(42u64)),
];
let (id, file_path, start_line) =
extract_lsp_row_fields(&row).expect("valid row should extract");
assert_eq!(id, "sym_id");
assert_eq!(file_path, "/src/main.rs");
assert_eq!(start_line, 42);
}
#[cfg(feature = "lsp")]
#[test]
fn extract_lsp_row_fields_returns_none_when_id_missing() {
let row = vec![
serde_json::Value::Null,
serde_json::Value::String("/src/main.rs".into()),
serde_json::Value::Number(serde_json::Number::from(1u64)),
];
assert!(extract_lsp_row_fields(&row).is_none());
}
#[cfg(feature = "lsp")]
#[test]
fn extract_lsp_row_fields_returns_none_when_file_path_missing() {
let row = vec![
serde_json::Value::String("sym_id".into()),
serde_json::Value::Null,
serde_json::Value::Number(serde_json::Number::from(1u64)),
];
assert!(extract_lsp_row_fields(&row).is_none());
}
#[cfg(feature = "lsp")]
#[test]
fn extract_lsp_row_fields_returns_none_when_start_line_missing() {
let row = vec![
serde_json::Value::String("sym_id".into()),
serde_json::Value::String("/src/main.rs".into()),
serde_json::Value::Null,
];
assert!(extract_lsp_row_fields(&row).is_none());
}
#[cfg(feature = "lsp")]
#[test]
fn extract_lsp_row_fields_returns_none_for_empty_row() {
let row: Vec<serde_json::Value> = vec![];
assert!(extract_lsp_row_fields(&row).is_none());
}
#[cfg(feature = "lsp")]
#[test]
fn extract_lsp_row_fields_returns_none_when_id_not_string() {
let row = vec![
serde_json::Value::Number(serde_json::Number::from(123u64)),
serde_json::Value::String("/src/main.rs".into()),
serde_json::Value::Number(serde_json::Number::from(1u64)),
];
assert!(extract_lsp_row_fields(&row).is_none());
}
#[cfg(feature = "lsp")]
#[test]
fn extract_lsp_row_fields_returns_none_when_start_line_not_number() {
let row = vec![
serde_json::Value::String("sym_id".into()),
serde_json::Value::String("/src/main.rs".into()),
serde_json::Value::String("not_a_number".into()),
];
assert!(extract_lsp_row_fields(&row).is_none());
}
#[cfg(feature = "lsp")]
#[test]
fn resolve_abs_file_path_returns_absolute_unchanged() {
let workspace = std::path::Path::new("/workspace");
let abs = resolve_abs_file_path(workspace, "/src/main.rs");
assert_eq!(abs, std::path::PathBuf::from("/src/main.rs"));
}
#[cfg(feature = "lsp")]
#[test]
fn resolve_abs_file_path_joins_relative_with_workspace() {
let workspace = std::path::Path::new("/workspace");
let abs = resolve_abs_file_path(workspace, "src/main.rs");
assert_eq!(abs, std::path::PathBuf::from("/workspace/src/main.rs"));
}
#[cfg(feature = "lsp")]
#[test]
fn resolve_abs_file_path_handles_empty_string() {
let workspace = std::path::Path::new("/workspace");
let abs = resolve_abs_file_path(workspace, "");
assert_eq!(abs, std::path::PathBuf::from("/workspace"));
}
#[cfg(feature = "lsp")]
#[test]
fn resolve_abs_file_path_handles_relative_with_dots() {
let workspace = std::path::Path::new("/workspace");
let abs = resolve_abs_file_path(workspace, "./src/main.rs");
assert_eq!(abs, std::path::PathBuf::from("/workspace/./src/main.rs"));
}
#[cfg(feature = "lsp")]
#[test]
fn collect_active_extensions_empty_entries_returns_empty_set() {
let provider_exts: std::collections::HashSet<&str> =
["rs", "py", "c", "cpp", "go", "ts", "f90", "java"]
.into_iter()
.collect();
let active = collect_active_extensions(&[], &provider_exts);
assert!(
active.is_empty(),
"empty entries must produce empty active set"
);
}
#[cfg(feature = "lsp")]
#[test]
fn collect_active_extensions_pure_rust_returns_only_rs() {
let provider_exts: std::collections::HashSet<&str> =
["rs", "py", "c", "cpp", "go", "ts", "f90", "java"]
.into_iter()
.collect();
let entries = vec![
(
"id1".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/main.rs")),
1,
),
(
"id2".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/lib.rs")),
10,
),
(
"id3".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/mod.rs")),
20,
),
];
let active = collect_active_extensions(&entries, &provider_exts);
assert_eq!(
active.len(),
1,
"pure Rust repo must activate only rust-analyzer"
);
assert!(active.contains("rs"));
assert!(!active.contains("py"));
assert!(!active.contains("c"));
assert!(!active.contains("cpp"));
assert!(!active.contains("go"));
assert!(!active.contains("ts"));
assert!(!active.contains("f90"));
assert!(!active.contains("java"));
}
#[cfg(feature = "lsp")]
#[test]
fn collect_active_extensions_pure_python_returns_only_py() {
let provider_exts: std::collections::HashSet<&str> =
["rs", "py", "c", "cpp", "go", "ts", "f90", "java"]
.into_iter()
.collect();
let entries = vec![
(
"id1".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/main.py")),
1,
),
(
"id2".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/utils.py")),
5,
),
];
let active = collect_active_extensions(&entries, &provider_exts);
assert_eq!(
active.len(),
1,
"pure Python repo must activate only pyright"
);
assert!(active.contains("py"));
assert!(!active.contains("rs"));
}
#[cfg(feature = "lsp")]
#[test]
fn collect_active_extensions_mixed_rs_py_returns_both() {
let provider_exts: std::collections::HashSet<&str> =
["rs", "py", "c", "cpp", "go", "ts", "f90", "java"]
.into_iter()
.collect();
let entries = vec![
(
"id1".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/main.rs")),
1,
),
(
"id2".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/script.py")),
5,
),
];
let active = collect_active_extensions(&entries, &provider_exts);
assert_eq!(
active.len(),
2,
"mixed repo must activate exactly the union"
);
assert!(active.contains("rs"));
assert!(active.contains("py"));
assert!(!active.contains("java"));
}
#[cfg(feature = "lsp")]
#[test]
fn collect_active_extensions_unknown_extension_skipped() {
let provider_exts: std::collections::HashSet<&str> =
["rs", "py", "c", "cpp", "go", "ts", "f90", "java"]
.into_iter()
.collect();
let entries = vec![
(
"id1".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/README.md")),
1,
),
(
"id2".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/Cargo.toml")),
1,
),
];
let active = collect_active_extensions(&entries, &provider_exts);
assert!(
active.is_empty(),
"unknown extensions (.md, .toml) must NOT activate any LSP server; got: {active:?}"
);
}
#[cfg(feature = "lsp")]
#[test]
fn collect_active_extensions_cpp_and_c_both_active() {
let provider_exts: std::collections::HashSet<&str> =
["rs", "py", "c", "cpp", "go", "ts", "f90", "java"]
.into_iter()
.collect();
let entries = vec![
(
"id1".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/main.c")),
1,
),
(
"id2".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/src/util.cpp")),
5,
),
];
let active = collect_active_extensions(&entries, &provider_exts);
assert_eq!(active.len(), 2);
assert!(active.contains("c"));
assert!(active.contains("cpp"));
}
#[cfg(feature = "lsp")]
#[test]
fn collect_active_extensions_dedupes_repeated_extensions() {
let provider_exts: std::collections::HashSet<&str> =
["rs", "py", "c", "cpp", "go", "ts", "f90", "java"]
.into_iter()
.collect();
let entries: Vec<(String, std::sync::Arc<std::path::PathBuf>, u32)> = (0..1000)
.map(|i| {
(
format!("id{i}"),
std::sync::Arc::new(std::path::PathBuf::from("/src/mod.rs")),
i,
)
})
.collect();
let active = collect_active_extensions(&entries, &provider_exts);
assert_eq!(active.len(), 1);
assert!(active.contains("rs"));
}
#[cfg(feature = "lsp")]
#[test]
fn collect_active_extensions_all_eight_exts_active_for_polyglot_repo() {
let provider_exts: std::collections::HashSet<&str> =
["rs", "py", "c", "cpp", "go", "ts", "f90", "java"]
.into_iter()
.collect();
let entries = vec![
(
"id1".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/a.rs")),
1,
),
(
"id2".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/b.py")),
1,
),
(
"id3".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/c.c")),
1,
),
(
"id4".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/d.cpp")),
1,
),
(
"id5".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/e.go")),
1,
),
(
"id6".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/f.ts")),
1,
),
(
"id7".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/g.f90")),
1,
),
(
"id8".to_string(),
std::sync::Arc::new(std::path::PathBuf::from("/h.java")),
1,
),
];
let active = collect_active_extensions(&entries, &provider_exts);
assert_eq!(
active.len(),
8,
"polyglot repo must activate all 8 providers"
);
}
#[serial_test::serial(kit_init)]
#[cfg(feature = "cli")]
#[test]
fn index_wrapper_fails_when_kit_not_initialized() {
use crate::service::runtime::reset_kit_for_testing;
reset_kit_for_testing();
let rt = tokio::runtime::Runtime::new().expect("runtime");
let result = rt.block_on(index(
"/nonexistent/path".to_string(),
"test_project".to_string(),
false,
false,
false,
false,
false,
));
assert!(result.is_err(), "wrapper should fail without kit");
reset_kit_for_testing();
}
#[serial_test::serial(kit_init)]
#[cfg(all(feature = "cli", feature = "lang-rust"))]
#[test]
fn index_wrapper_succeeds_via_init_kit() {
use crate::kit::{build_kit, KitBootstrapConfig};
use crate::service::runtime::{init_kit, reset_kit_for_testing};
use std::fs;
use tempfile::TempDir;
reset_kit_for_testing();
let src_dir = TempDir::new().unwrap();
fs::write(src_dir.path().join("main.rs"), "fn main() {}\n").unwrap();
let db_dir = TempDir::new().unwrap();
let db_path = db_dir.path().join("wrapper_testdb");
let config = KitBootstrapConfig::new(db_path.clone());
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(build_kit(&config))
.expect("build_kit");
init_kit(kit).expect("init_kit");
let rt = tokio::runtime::Runtime::new().expect("runtime");
let result = rt.block_on(index(
src_dir.path().to_str().unwrap().to_string(),
"wrapper_test".to_string(),
true,
false,
false,
false,
false,
));
assert!(result.is_ok(), "wrapper should succeed: {:?}", result.err());
reset_kit_for_testing();
}
#[serial_test::serial(kit_init)]
#[cfg(all(feature = "cli", feature = "lang-rust"))]
#[test]
fn index_wrapper_with_embed_true_emits_deprecation_warning() {
use crate::kit::{build_kit, KitBootstrapConfig};
use crate::service::runtime::{init_kit, reset_kit_for_testing};
use std::fs;
use tempfile::TempDir;
reset_kit_for_testing();
let src_dir = TempDir::new().unwrap();
fs::write(src_dir.path().join("main.rs"), "fn main() {}\n").unwrap();
let db_dir = TempDir::new().unwrap();
let db_path = db_dir.path().join("embed_warn_testdb");
let config = KitBootstrapConfig::new(db_path.clone());
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(build_kit(&config))
.expect("build_kit");
init_kit(kit).expect("init_kit");
let rt = tokio::runtime::Runtime::new().expect("runtime");
let result = rt.block_on(index(
src_dir.path().to_str().unwrap().to_string(),
"embed_warn_test".to_string(),
true,
false,
true,
false,
false,
));
assert!(
result.is_ok(),
"wrapper should succeed even with embed=true: {:?}",
result.err()
);
reset_kit_for_testing();
}
#[cfg(feature = "lang-rust")]
#[test]
fn index_core_reports_dq_violations_when_present() {
use crate::kit::{build_kit, KitBootstrapConfig};
use crate::storage::Repository;
use std::fs;
use tempfile::TempDir;
let src_dir = TempDir::new().unwrap();
let src_path = src_dir.path();
fs::write(src_path.join("main.rs"), "fn main() {}\n").unwrap();
let db_dir = TempDir::new().unwrap();
let db_path = db_dir.path().join("dq_violation_testdb");
let config = KitBootstrapConfig::new(db_path.clone());
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(build_kit(&config))
.expect("build_kit");
let repo = Repository::open(&db_path).expect("open repo for injection");
repo.connection()
.execute(
"CREATE (:File {id: 'bad_file', project: 'ghost_proj', name: 'bad.rs', \
filePath: '/bad.rs', language: 'rust', hash: '', lineCount: 0});",
)
.expect("insert bad file");
let output = index_core(
&kit,
&db_path,
src_path.to_str().unwrap(),
"test_project",
false,
false,
false,
)
.expect("index should succeed even with DQ violations");
assert!(output.files_indexed >= 1);
}
#[serial_test::serial(kit_init)]
#[cfg(feature = "cli")]
#[test]
fn index_wrapper_fails_when_storage_config_not_registered() {
use crate::kit::AsyncKit;
use crate::service::runtime::{init_kit, reset_kit_for_testing};
use sdforge::prelude::ApiError;
reset_kit_for_testing();
let kit = tokio::runtime::Runtime::new()
.unwrap()
.block_on(AsyncKit::new().build())
.expect("build empty kit");
init_kit(kit).expect("init_kit");
let rt = tokio::runtime::Runtime::new().expect("runtime");
let result = rt.block_on(index(
"/nonexistent/path".to_string(),
"test_project".to_string(),
false,
false,
false,
false,
false,
));
assert!(result.is_err(), "wrapper should fail without StorageConfig");
match result.unwrap_err() {
ApiError::Internal { message, .. } => {
assert!(
message.contains("storage config"),
"error should mention storage config: {message}"
);
}
other => panic!("expected Internal error, got {other:?}"),
}
reset_kit_for_testing();
}
#[cfg(feature = "lsp")]
mod batch_update_tests {
use super::*;
use crate::storage::StorageError;
#[test]
fn build_batch_empty_returns_empty_string() {
let batch: Vec<(String, String)> = Vec::new();
let stmt = build_batch_semantic_type_update(&batch, "proj");
assert!(stmt.is_empty(), "empty batch must return empty string");
}
#[test]
fn build_batch_single_entry_generates_valid_unwind_statement() {
let batch = vec![(String::from("sym1"), String::from("fn foo() -> i32"))];
let stmt = build_batch_semantic_type_update(&batch, "my-project");
assert!(
stmt.starts_with("UNWIND ["),
"must start with UNWIND [: got {stmt}"
);
assert!(
stmt.contains("{id: 'sym1', sem: 'fn foo() -> i32'}"),
"must contain the entry map: got {stmt}"
);
assert!(
stmt.contains("] AS row"),
"must close list and alias as row: got {stmt}"
);
assert!(
stmt.contains("MATCH (n {id: row.id, project: 'my-project'})"),
"must MATCH by row.id + project: got {stmt}"
);
assert!(
stmt.contains("SET n.semantic_type = row.sem"),
"must SET semantic_type from row.sem: got {stmt}"
);
assert!(
stmt.ends_with(';'),
"statement must end with semicolon: got {stmt}"
);
}
#[test]
fn build_batch_500_entries_contains_500_id_sem_maps() {
let batch: Vec<(String, String)> = (0..500)
.map(|i| (format!("sym_{i}"), format!("type_{i}")))
.collect();
let stmt = build_batch_semantic_type_update(&batch, "proj");
let count = stmt.matches("{id: '").count();
assert_eq!(count, 500, "expected 500 entry maps, got {count}");
}
#[test]
fn build_batch_escapes_single_quotes_in_id() {
let batch = vec![(String::from("a'b"), String::from("type"))];
let stmt = build_batch_semantic_type_update(&batch, "proj");
assert!(
stmt.contains("{id: 'a\\'b', sem: 'type'}"),
"id single quote must be escaped: got {stmt}"
);
}
#[test]
fn build_batch_escapes_single_quotes_in_sem() {
let batch = vec![(String::from("sym"), String::from("fn it's() -> ()"))];
let stmt = build_batch_semantic_type_update(&batch, "proj");
assert!(
stmt.contains("{id: 'sym', sem: 'fn it\\'s() -> ()'}"),
"sem single quote must be escaped: got {stmt}"
);
}
#[test]
fn build_batch_escapes_single_quotes_in_project() {
let batch = vec![(String::from("sym"), String::from("type"))];
let stmt = build_batch_semantic_type_update(&batch, "proj'ect");
assert!(
stmt.contains("project: 'proj\\'ect'"),
"project single quote must be escaped: got {stmt}"
);
}
#[test]
fn build_batch_escapes_backslashes_in_id_and_sem() {
let batch = vec![(String::from("a\\b"), String::from("c\\d"))];
let stmt = build_batch_semantic_type_update(&batch, "proj");
assert!(
stmt.contains("{id: 'a\\\\b', sem: 'c\\\\d'}"),
"backslashes must be doubled: got {stmt}"
);
}
#[test]
fn flush_batch_empty_does_not_call_exec_and_returns_zero() {
let mut batch: Vec<(String, String)> = Vec::new();
let mut enhanced: u32 = 0;
let mut skipped: u32 = 0;
let call_count = std::cell::Cell::new(0u32);
let exec = |_q: &str| -> crate::storage::Result<()> {
call_count.set(call_count.get() + 1);
Ok(())
};
flush_semantic_type_batch(exec, &mut batch, "proj", &mut enhanced, &mut skipped);
assert_eq!(
call_count.get(),
0,
"exec must not be called on empty batch"
);
assert_eq!(enhanced, 0);
assert_eq!(skipped, 0);
assert!(batch.is_empty(), "batch must remain empty");
}
#[test]
fn flush_batch_500_entries_calls_exec_once_when_batch_succeeds() {
let mut batch: Vec<(String, String)> = (0..500)
.map(|i| (format!("sym_{i}"), format!("type_{i}")))
.collect();
let mut enhanced: u32 = 0;
let mut skipped: u32 = 0;
let call_count = std::cell::Cell::new(0u32);
let exec = |_q: &str| -> crate::storage::Result<()> {
call_count.set(call_count.get() + 1);
Ok(())
};
flush_semantic_type_batch(exec, &mut batch, "proj", &mut enhanced, &mut skipped);
assert_eq!(
call_count.get(),
1,
"exec must be called exactly once on batch success"
);
assert_eq!(enhanced, 500, "all 500 entries counted as enhanced");
assert_eq!(skipped, 0);
assert!(batch.is_empty(), "batch must be cleared after flush");
}
#[test]
fn flush_batch_falls_back_to_per_row_when_batch_exec_fails() {
let mut batch: Vec<(String, String)> = (0..10)
.map(|i| (format!("sym_{i}"), format!("type_{i}")))
.collect();
let mut enhanced: u32 = 0;
let mut skipped: u32 = 0;
let call_count = std::cell::Cell::new(0u32);
let exec = |_q: &str| -> crate::storage::Result<()> {
let n = call_count.get();
call_count.set(n + 1);
if n == 0 {
Err(StorageError::Query(String::from(
"UNWIND syntax not supported",
)))
} else {
let row_idx = (n - 1) as usize;
if row_idx.is_multiple_of(2) {
Ok(())
} else {
Err(StorageError::Query(String::from("per-row fail")))
}
}
};
flush_semantic_type_batch(exec, &mut batch, "proj", &mut enhanced, &mut skipped);
assert_eq!(
call_count.get(),
11,
"must call batch once then per-row 10 times"
);
assert_eq!(enhanced, 5, "5 even-indexed per-row calls succeed");
assert_eq!(skipped, 5, "5 odd-indexed per-row calls fail");
assert!(batch.is_empty(), "batch must be cleared after flush");
}
#[test]
fn flush_batch_full_failure_increments_skipped_only() {
let mut batch: Vec<(String, String)> = (0..3)
.map(|i| (format!("sym_{i}"), format!("type_{i}")))
.collect();
let mut enhanced: u32 = 99;
let mut skipped: u32 = 99;
let exec = |_q: &str| -> crate::storage::Result<()> {
Err(StorageError::Query(String::from("always fail")))
};
flush_semantic_type_batch(exec, &mut batch, "proj", &mut enhanced, &mut skipped);
assert_eq!(enhanced, 99, "no per-row success");
assert_eq!(skipped, 102, "3 per-row failures increment skipped from 99");
assert!(batch.is_empty(), "batch cleared");
}
}
#[cfg(feature = "lsp")]
mod parallel_start_tests {
use super::*;
use crate::lsp::{LspError, LspProvider};
use lsp_types::{Hover, Location};
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
struct MockLspProvider {
start_count: Arc<AtomicU32>,
start_delay_ms: u64,
should_panic: bool,
should_fail: bool,
}
impl MockLspProvider {
fn new(start_delay_ms: u64) -> Self {
Self {
start_count: Arc::new(AtomicU32::new(0)),
start_delay_ms,
should_panic: false,
should_fail: false,
}
}
fn with_panic(mut self) -> Self {
self.should_panic = true;
self
}
fn with_failure(mut self) -> Self {
self.should_fail = true;
self
}
}
impl LspProvider for MockLspProvider {
fn start(&self, _workspace: &Path) -> Result<(), LspError> {
self.start_count.fetch_add(1, Ordering::SeqCst);
if self.should_panic {
panic!("MockLspProvider: simulated panic");
}
if self.start_delay_ms > 0 {
std::thread::sleep(Duration::from_millis(self.start_delay_ms));
}
if self.should_fail {
return Err(LspError::ServerStart(String::from(
"MockLspProvider: simulated failure",
)));
}
Ok(())
}
fn definition(
&self,
_file: &Path,
_line: u32,
_col: u32,
) -> Result<Option<Location>, LspError> {
Ok(None)
}
fn type_definition(
&self,
_file: &Path,
_line: u32,
_col: u32,
) -> Result<Option<Location>, LspError> {
Ok(None)
}
fn hover(
&self,
_file: &Path,
_line: u32,
_col: u32,
) -> Result<Option<Hover>, LspError> {
Ok(None)
}
fn shutdown(&self) -> Result<(), LspError> {
Ok(())
}
}
type MockProvidersResult = (
Vec<(&'static str, Box<dyn LspProvider>)>,
Vec<Arc<AtomicU32>>,
);
fn build_mock_providers(delays_ms: &[u64]) -> MockProvidersResult {
const EXTS: [&str; 8] = ["rs", "py", "c", "cpp", "go", "ts", "f90", "java"];
let mut providers: Vec<(&'static str, Box<dyn LspProvider>)> = Vec::new();
let mut counts: Vec<Arc<AtomicU32>> = Vec::new();
for (i, &delay) in delays_ms.iter().enumerate() {
let ext = EXTS[i];
let mock = MockLspProvider::new(delay);
counts.push(mock.start_count.clone());
providers.push((ext, Box::new(mock) as Box<dyn LspProvider>));
}
while providers.len() < 8 {
let ext = EXTS[providers.len()];
let mock = MockLspProvider::new(0);
providers.push((ext, Box::new(mock) as Box<dyn LspProvider>));
}
(providers, counts)
}
#[test]
fn parallel_start_completes_faster_than_serial() {
let (providers, _counts) = build_mock_providers(&[50, 50, 50, 50, 50, 50, 50, 50]);
let active_exts: std::collections::HashSet<&'static str> =
providers.iter().map(|(ext, _)| *ext).collect();
let workspace = PathBuf::from("/tmp/mock-workspace");
let start = Instant::now();
let started = start_active_providers_parallel(&providers, &active_exts, &workspace);
let elapsed = start.elapsed();
assert_eq!(started, 8, "all 8 providers should start successfully");
assert!(
elapsed < Duration::from_millis(200),
"parallel start must complete in < 200 ms, got {elapsed:?}"
);
}
#[test]
fn parallel_start_calls_start_on_all_active_providers() {
let delays = [10, 10, 10, 10, 10, 10, 10, 10];
const EXTS: [&str; 8] = ["rs", "py", "c", "cpp", "go", "ts", "f90", "java"];
let mut providers: Vec<(&'static str, Box<dyn LspProvider>)> = Vec::new();
let mut counts: Vec<Arc<AtomicU32>> = Vec::new();
for (i, &delay) in delays.iter().enumerate() {
let mock = MockLspProvider::new(delay);
counts.push(mock.start_count.clone());
providers.push((EXTS[i], Box::new(mock) as Box<dyn LspProvider>));
}
let active_exts: std::collections::HashSet<&'static str> =
EXTS.iter().copied().collect();
let workspace = PathBuf::from("/tmp/mock-workspace");
let started = start_active_providers_parallel(&providers, &active_exts, &workspace);
assert_eq!(started, 8);
for (i, count) in counts.iter().enumerate() {
assert_eq!(
count.load(Ordering::SeqCst),
1,
"provider {} start_count must be 1",
EXTS[i]
);
}
}
#[test]
fn parallel_start_skips_inactive_providers() {
const EXTS: [&str; 8] = ["rs", "py", "c", "cpp", "go", "ts", "f90", "java"];
let mut providers: Vec<(&'static str, Box<dyn LspProvider>)> = Vec::new();
let mut counts: Vec<Arc<AtomicU32>> = Vec::new();
for ext in EXTS.iter() {
let mock = MockLspProvider::new(5);
counts.push(mock.start_count.clone());
providers.push((*ext, Box::new(mock) as Box<dyn LspProvider>));
}
let active_exts: std::collections::HashSet<&'static str> =
std::collections::HashSet::from(["rs"]);
let workspace = PathBuf::from("/tmp/mock-workspace");
let started = start_active_providers_parallel(&providers, &active_exts, &workspace);
assert_eq!(started, 1, "only 1 provider should be started");
assert_eq!(
counts[0].load(Ordering::SeqCst),
1,
"rs provider called once"
);
for i in 1..8 {
assert_eq!(
counts[i].load(Ordering::SeqCst),
0,
"provider {} must NOT be started (not in active_exts)",
EXTS[i]
);
}
}
#[test]
fn parallel_start_panic_does_not_block_other_providers() {
const EXTS: [&str; 8] = ["rs", "py", "c", "cpp", "go", "ts", "f90", "java"];
let mut providers: Vec<(&'static str, Box<dyn LspProvider>)> = Vec::new();
let mut counts: Vec<Arc<AtomicU32>> = Vec::new();
for (i, ext) in EXTS.iter().enumerate() {
let mut mock = MockLspProvider::new(10);
if i == 2 {
mock = mock.with_panic();
}
counts.push(mock.start_count.clone());
providers.push((*ext, Box::new(mock) as Box<dyn LspProvider>));
}
let active_exts: std::collections::HashSet<&'static str> =
EXTS.iter().copied().collect();
let workspace = PathBuf::from("/tmp/mock-workspace");
let started = start_active_providers_parallel(&providers, &active_exts, &workspace);
assert_eq!(
started, 7,
"7 of 8 providers should start; panicking one is degraded"
);
assert_eq!(
counts[2].load(Ordering::SeqCst),
1,
"panicking provider was called"
);
for (i, count) in counts.iter().enumerate() {
if i == 2 {
continue;
}
assert_eq!(
count.load(Ordering::SeqCst),
1,
"provider {} called once",
EXTS[i]
);
}
}
#[test]
fn parallel_start_failure_does_not_block_other_providers() {
const EXTS: [&str; 8] = ["rs", "py", "c", "cpp", "go", "ts", "f90", "java"];
let mut providers: Vec<(&'static str, Box<dyn LspProvider>)> = Vec::new();
let mut counts: Vec<Arc<AtomicU32>> = Vec::new();
for (i, ext) in EXTS.iter().enumerate() {
let mut mock = MockLspProvider::new(5);
if i == 1 {
mock = mock.with_failure();
}
counts.push(mock.start_count.clone());
providers.push((*ext, Box::new(mock) as Box<dyn LspProvider>));
}
let active_exts: std::collections::HashSet<&'static str> =
EXTS.iter().copied().collect();
let workspace = PathBuf::from("/tmp/mock-workspace");
let started = start_active_providers_parallel(&providers, &active_exts, &workspace);
assert_eq!(
started, 7,
"7 of 8 providers should start; failing one is degraded"
);
for (i, count) in counts.iter().enumerate() {
assert_eq!(
count.load(Ordering::SeqCst),
1,
"provider {} called once",
EXTS[i]
);
}
}
#[test]
fn parallel_start_empty_active_exts_returns_zero() {
const EXTS: [&str; 8] = ["rs", "py", "c", "cpp", "go", "ts", "f90", "java"];
let mut providers: Vec<(&'static str, Box<dyn LspProvider>)> = Vec::new();
let mut counts: Vec<Arc<AtomicU32>> = Vec::new();
for ext in EXTS.iter() {
let mock = MockLspProvider::new(5);
counts.push(mock.start_count.clone());
providers.push((*ext, Box::new(mock) as Box<dyn LspProvider>));
}
let active_exts: std::collections::HashSet<&'static str> =
std::collections::HashSet::new();
let workspace = PathBuf::from("/tmp/mock-workspace");
let started = start_active_providers_parallel(&providers, &active_exts, &workspace);
assert_eq!(
started, 0,
"no providers should be started on empty active_exts"
);
for (i, count) in counts.iter().enumerate() {
assert_eq!(
count.load(Ordering::SeqCst),
0,
"provider {} must not be called",
EXTS[i]
);
}
}
}
}