From 4bae4e9c08fd2d26a137493a00d64fb34e3d84c5 Mon Sep 17 00:00:00 2001 From: Riz Ashraf Date: Sun, 27 Sep 2026 08:10:09 +0100 Subject: [PATCH] Test coverage and headless nvim bug fixes --- nvim-core/src/headless.rs | 13 +++ nvim-core/src/lib.rs | 1 + server/src/api/events.rs | 58 +++++++++++++ server/src/handlers/graph.rs | 21 ++++- server/src/handlers/meta.rs | 37 +++++++++ server/src/search.rs | 22 ++++- server/src/state.rs | 71 +++++++++------- server/src/tools.rs | 23 ++++++ win-nvim/tests/integration_test.rs | 128 +++++++++++++++++++++++------ 9 files changed, 316 insertions(+), 58 deletions(-) diff --git a/nvim-core/src/headless.rs b/nvim-core/src/headless.rs index fc5be28..313f190 100644 --- a/nvim-core/src/headless.rs +++ b/nvim-core/src/headless.rs @@ -21,6 +21,9 @@ pub async fn spawn_headless_nvim() -> Result { .arg("--headless") .arg("--listen") .arg(&socket_name) + .stdin(std::process::Stdio::piped()) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) .kill_on_drop(true) .spawn() .map_err(|e| format!("Failed to spawn headless Neovim: {}", e))?; @@ -40,3 +43,13 @@ pub async fn spawn_headless_nvim() -> Result { Ok(socket_name) } + +pub async fn kill_headless_nvim() { + let child_to_kill = { + let mut proc_lock = HEADLESS_PROC.lock().unwrap_or_else(|e| e.into_inner()); + proc_lock.take() + }; + if let Some(mut child) = child_to_kill { + let _ = child.kill().await; + } +} diff --git a/nvim-core/src/lib.rs b/nvim-core/src/lib.rs index b25ac51..9d81ce8 100644 --- a/nvim-core/src/lib.rs +++ b/nvim-core/src/lib.rs @@ -726,6 +726,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) { m } else { tracing::info!("Stdin closed, exiting loop"); + headless::kill_headless_nvim().await; break; }; diff --git a/server/src/api/events.rs b/server/src/api/events.rs index eea514d..e32e172 100644 --- a/server/src/api/events.rs +++ b/server/src/api/events.rs @@ -42,3 +42,61 @@ pub async fn post_event_handler( let _ = state.handler.state.event_bus_tx.send(event); axum::Json(serde_json::json!({"status": "ok"})) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::state::MemoryState; + use crate::router::MemoryHandler; + use tempfile::tempdir; + use axum::extract::Query; + use axum::extract::State; + use std::collections::HashMap; + use std::sync::atomic::AtomicUsize; + use std::sync::RwLock; + + #[tokio::test] + async fn test_events_wait_and_post() { + let dir = tempdir().unwrap(); + let mem_state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); + let app_state = Arc::new(AppState { + handler: Arc::new(MemoryHandler::new(mem_state.clone())), + clients: RwLock::new(HashMap::new()), + next_id: AtomicUsize::new(1), + }); + + // Start wait_for_event in a background task + let app_state_clone = app_state.clone(); + let mut params = HashMap::new(); + params.insert("topic".to_string(), "test_topic".to_string()); + params.insert("session_id".to_string(), "123".to_string()); + + let wait_task = tokio::spawn(async move { + let res = wait_for_event_handler(State(app_state_clone), Query(params)).await; + // axum::Json is returned, we need to extract it somehow, but just returning is enough for testing + res + }); + + // Yield slightly to ensure the wait task has subscribed + tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; + + // Post an event that shouldn't match + let unmatched_event = GenericEvent { + topic: "wrong_topic".to_string(), + session_id: Some("123".to_string()), + payload: serde_json::json!({}), + }; + post_event_handler(State(app_state.clone()), axum::Json(unmatched_event)).await; + + // Post the matching event + let matched_event = GenericEvent { + topic: "test_topic".to_string(), + session_id: Some("123".to_string()), + payload: serde_json::json!({"foo": "bar"}), + }; + post_event_handler(State(app_state.clone()), axum::Json(matched_event)).await; + + // Wait for the wait task to complete + let _ = wait_task.await.unwrap(); + } +} diff --git a/server/src/handlers/graph.rs b/server/src/handlers/graph.rs index 5d2731e..ee0698f 100644 --- a/server/src/handlers/graph.rs +++ b/server/src/handlers/graph.rs @@ -147,7 +147,16 @@ impl McpTool for CreateRelationsHandler { } async fn execute(&self, args: Value, state: Arc) -> Result { - let req: CreateRelationsTool = serde_json::from_value(args).map_err(|e| e.to_string())?; + let req: CreateRelationsTool = match serde_json::from_value(args.clone()) { + Ok(r) => r, + Err(e) => { + let err_msg = e.to_string(); + if err_msg.contains("missing field `from`") || err_msg.contains("missing field `to`") || err_msg.contains("missing field `relation_type`") { + return Err(format!("Schema error: {}. Note that the relation schema strictly uses 'from', 'to', and 'relation_type' (not 'source', 'target', or 'relationType'). Please correct your tool call arguments.", err_msg)); + } + return Err(err_msg); + } + }; let mut missing_nodes = std::collections::HashSet::new(); state.modify_graph(|g| { for relation in req.relations { @@ -682,6 +691,16 @@ mod tests { }); let res = handler.execute(args, state.clone()).await.unwrap(); assert_eq!(res, "Relations created"); + + // Test semantic LLM schema feedback (User request) + let bad_args = json!({ + "relations": [ + {"source": "A", "target": "B", "relationType": "knows"} + ] + }); + let err_res = handler.execute(bad_args, state.clone()).await.unwrap_err(); + assert!(err_res.contains("Schema error:")); + assert!(err_res.contains("strictly uses 'from', 'to', and 'relation_type'")); } #[tokio::test] diff --git a/server/src/handlers/meta.rs b/server/src/handlers/meta.rs index 1e2e7ce..93b44d8 100644 --- a/server/src/handlers/meta.rs +++ b/server/src/handlers/meta.rs @@ -344,6 +344,7 @@ impl McpTool for OmniSearchHandler { idx.search(&req.query, req.namespace.as_deref()) .unwrap_or_default() }; + println!("OMNI SEARCH MATCHES: {:?}", matches); let kg_json = state.read_graph(|full| { let mut kg_entities = std::collections::HashMap::new(); @@ -660,4 +661,40 @@ mod tests { .await .unwrap(); } + + #[tokio::test] + async fn test_omni_search() { + let dir = tempfile::tempdir().unwrap(); + let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); + + let task = crate::models::Task { + id: "omni-1".to_string(), + title: "Omni Task".to_string(), + description: "Testing omni search functionality".to_string(), + status: "open".to_string(), + created_at: 0, + updated_at: 0, + git_branch: None, + parent_id: None, + dependencies: vec![], + acceptance_criteria: vec![], + }; + + { + state.tasks.modify(|t| { + t.push(task.clone()); + }); + } + + state.rebuild_index().await; + state.get_search_index().reader.reload().unwrap(); + + let omni = OmniSearchHandler; + let omni_res = omni + .execute(json!({"query": "Omni"}), state.clone()) + .await + .unwrap(); + println!("OMNI RES: {}", omni_res); + assert!(omni_res.contains("omni-1"), "omni search should return results containing the task id"); + } } diff --git a/server/src/search.rs b/server/src/search.rs index 85bb1e9..da69dcb 100644 --- a/server/src/search.rs +++ b/server/src/search.rs @@ -7,8 +7,8 @@ pub type SearchResultTuple = (String, String, String, String, f32); #[derive(Clone)] pub struct MemoryIndex { - index: Index, - reader: IndexReader, + pub index: Index, + pub reader: IndexReader, writer: Arc>, needs_commit: Arc, @@ -244,18 +244,32 @@ impl MemoryIndex { self.type_field => "entity", self.namespace_field => e.namespace.as_str() )); + self.needs_commit.store(true, std::sync::atomic::Ordering::SeqCst); + } + } + + pub fn delete_all(&self) { + if let Ok(mut writer) = self.writer.lock() { + let _ = writer.delete_all_documents(); + self.needs_commit.store(true, std::sync::atomic::Ordering::SeqCst); } } pub fn add_task_sync(&self, t: &Task) { + println!("add_task_sync called for task: {}", t.id); if let Ok(writer) = self.writer.lock() { - let _ = writer.add_document(doc!( + let res = writer.add_document(doc!( self.id_field => t.id.as_str(), self.title_field => t.title.as_str(), self.body_field => t.description.as_str(), self.type_field => "task", self.namespace_field => "global" )); + println!("Writer add_document returned id/result"); + self.needs_commit.store(true, std::sync::atomic::Ordering::SeqCst); + println!("Needs_commit set to true in add_task_sync"); + } else { + println!("Failed to acquire writer lock in add_task_sync"); } } @@ -268,6 +282,7 @@ impl MemoryIndex { self.type_field => "snippet", self.namespace_field => "global" )); + self.needs_commit.store(true, std::sync::atomic::Ordering::SeqCst); } } @@ -280,6 +295,7 @@ impl MemoryIndex { self.type_field => "adr", self.namespace_field => "global" )); + self.needs_commit.store(true, std::sync::atomic::Ordering::SeqCst); } } } diff --git a/server/src/state.rs b/server/src/state.rs index 53e6ef7..8ae1d3c 100644 --- a/server/src/state.rs +++ b/server/src/state.rs @@ -128,37 +128,40 @@ impl MemoryState { } pub async fn rebuild_index(self: &Arc) { - if let Ok(new_idx) = MemoryIndex::new(&self.base_dir) { - let idx = new_idx.clone(); + let idx = self.search_index.read().unwrap().clone(); + idx.delete_all(); - let entities: Vec<_> = self.graph.read_with(|g| g.entities.values().cloned().collect()); - let tasks = self.tasks.read_with(|t| t.clone()); - let snippets = self.snippets.read_with(|s| s.clone()); - let adrs = self.adrs.read_with(|a| a.clone()); + let entities: Vec<_> = self.graph.read_with(|g| g.entities.values().cloned().collect()); + let tasks = self.tasks.read_with(|t| t.clone()); + let snippets = self.snippets.read_with(|s| s.clone()); + let adrs = self.adrs.read_with(|a| a.clone()); - tokio::task::spawn_blocking(move || { - for e in entities { - idx.add_entity_sync(&e); - } - for task in tasks { - idx.add_task_sync(&task); - } - for snippet in snippets { - idx.add_snippet_sync(&snippet); - } - for adr in adrs { - idx.add_adr_sync(&adr); - } - }) - .await - .unwrap_or_else(|e| { - tracing::error!("Failed to join tantivy index rebuild thread: {}", e); - }); + println!("rebuild_index: found {} entities, {} tasks", entities.len(), tasks.len()); - let _ = new_idx.commit().await; - if let Ok(mut w) = self.search_index.write() { - *w = new_idx; + let idx_clone = idx.clone(); + tokio::task::spawn_blocking(move || { + println!("spawn_blocking started in rebuild_index"); + for e in entities { + idx_clone.add_entity_sync(&e); } + for task in tasks { + idx_clone.add_task_sync(&task); + } + for snippet in snippets { + idx_clone.add_snippet_sync(&snippet); + } + for adr in adrs { + idx_clone.add_adr_sync(&adr); + } + }) + .await + .unwrap_or_else(|e| { + tracing::error!("Failed to join tantivy index rebuild thread: {}", e); + }); + + let _ = idx.commit().await; + if let Ok(mut w) = self.search_index.write() { + *w = idx; } } } @@ -204,7 +207,17 @@ mod tests { // Check search index initialization let idx = arc_state.search_index.read().unwrap(); - // Just verify we can read it without panic - assert!(idx.search("Test", None).is_ok()); + // Force reload reader to ensure it sees the commit made by rebuild_index + idx.reader.reload().unwrap(); + println!("Index reader doc count: {}", idx.reader.searcher().num_docs()); + + let all_docs = idx.search("Test", None).expect("Search failed"); + println!("All docs for 'Test': {:?}", all_docs); + + // Verify the task added synchronously is actually searchable + let results = idx.search("Test", None).expect("Search failed"); + assert_eq!(results.len(), 1, "Expected exactly 1 search result"); + assert_eq!(results[0].0, "123", "Expected the result to be the task we just added"); + assert_eq!(results[0].1, "task", "Expected document type to be task"); } } diff --git a/server/src/tools.rs b/server/src/tools.rs index f75f2f7..106f634 100644 --- a/server/src/tools.rs +++ b/server/src/tools.rs @@ -562,3 +562,26 @@ pub struct VerifyAcceptanceCriteriaTool { pub criteria: String, pub proof: String, } + +#[cfg(test)] +mod tests { + use super::*; + use schemars::schema_for; + + #[test] + fn test_schema_extraction_includes_descriptions() { + let schema = schema_for!(SetAcceptanceCriteriaTool); + let schema_json = serde_json::to_value(&schema).unwrap(); + + let desc = schema_json.get("description").and_then(|d| d.as_str()).unwrap_or(""); + assert!(desc.contains("Define a strict checklist of acceptance criteria"), "Schema should include struct docstring as description"); + + let schema2 = schema_for!(LogCodeChangeTool); + let schema2_json = serde_json::to_value(&schema2).unwrap(); + let props = schema2_json.get("properties").expect("Missing properties"); + + let file_path_prop = props.get("file_path").expect("Missing file_path property"); + let field_desc = file_path_prop.get("description").and_then(|d| d.as_str()).unwrap_or(""); + assert!(field_desc.contains("The path of the file that was changed"), "Schema should include field docstring as description"); + } +} diff --git a/win-nvim/tests/integration_test.rs b/win-nvim/tests/integration_test.rs index 31aa101..51287ba 100644 --- a/win-nvim/tests/integration_test.rs +++ b/win-nvim/tests/integration_test.rs @@ -1,5 +1,5 @@ use serde_json::{Value, json}; -use std::io::{BufRead, BufReader, Write}; +use std::io::{BufRead, BufReader, Write, Read}; use std::process::{Command, Stdio}; fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) { @@ -24,7 +24,14 @@ fn test_mcp_initialization_and_tools_list() { nvim_exe.pop(); nvim_exe.push("mcp-memory-win-nvim.exe"); + let temp_dir = std::env::temp_dir().join("win-nvim-test"); + std::fs::create_dir_all(&temp_dir).unwrap(); + let temp_gemini = temp_dir.join(".gemini"); + std::fs::create_dir_all(&temp_gemini).unwrap(); + let mut child = Command::new(&nvim_exe) + .env("USERPROFILE", temp_dir.to_str().unwrap()) + .env("HOME", temp_dir.to_str().unwrap()) .stdin(Stdio::piped()) .stdout(Stdio::piped()) .stderr(Stdio::piped()) @@ -94,35 +101,106 @@ fn test_mcp_initialization_and_tools_list() { let has_get_active_buffer = tools.iter().any(|t| t["name"] == "nvim_get_active_buffer"); assert!(has_get_active_buffer, "Missing nvim_get_active_buffer tool"); - // 3. Test negative scenario: tools/call when Neovim is not running - // Since Neovim is not guaranteed to be running on the test agent's system, - // calling a Neovim-specific tool should gracefully return a JSON-RPC error. - let call_req = json!({ + let tool_names = vec![ + "nvim_goto_line", + "nvim_get_active_buffer", + "nvim_get_cursor", + "nvim_get_visual_selection", + "nvim_set_diagnostics", + "nvim_set_extmark", + "nvim_list_buffers", + "nvim_list_windows", + "nvim_get_active_window", + "nvim_set_active_window", + "nvim_get_diagnostics", + "nvim_open_file", + "nvim_open_buffer", + "nvim_close_buffer", + "nvim_close_window", + "nvim_split_window", + "nvim_reload_buffer", + "nvim_save_buffer", + "nvim_set_quickfix", + "nvim_highlight_lines", + "nvim_get_messages", + "nvim_get_viewport", + "nvim_read_file", + "nvim_search_file", + "nvim_execute_lua", + "nvim_get_server_info", + "unknown_tool", + ]; + + for (i, tool_name) in tool_names.iter().enumerate() { + let req_id = i + 100; + let call_req = json!({ + "jsonrpc": "2.0", + "method": "tools/call", + "params": { + "name": tool_name, + "arguments": { + "file": "test.txt", + "line": 10, + "code": "return 1", + "lua_code": "return 1", + "message": "test msg", + "text": "test", + "hl_group": "Error", + "win_id": 1, + "buf_id": 1, + "name": "test", + "start_line": 1, + "end_line": 2, + "force": true, + "group": "test", + "pattern": "foo" + } + }, + "id": req_id + }); + + send_message(&mut stdin, call_req); + let call_resp = read_message(&mut stdout).expect("Failed to read tools/call response"); + println!("Response for {}: {}", tool_name, serde_json::to_string(&call_resp).unwrap()); + assert_eq!(call_resp["jsonrpc"], "2.0"); + assert_eq!(call_resp["id"], req_id); + } + + for (i, tool_name) in tool_names.iter().enumerate() { + let req_id = i + 200; + let call_req = json!({ + "jsonrpc": "2.0", + "method": "tools/call", + "params": { + "name": tool_name, + "arguments": {} + }, + "id": req_id + }); + + send_message(&mut stdin, call_req); + let call_resp = read_message(&mut stdout).expect("Failed to read tools/call response"); + assert_eq!(call_resp["jsonrpc"], "2.0"); + assert_eq!(call_resp["id"], req_id); + } + + let call_req_no_id = json!({ "jsonrpc": "2.0", "method": "tools/call", "params": { - "name": "nvim_get_active_buffer", + "name": "nvim_goto_line", "arguments": {} - }, - "id": 3 + } }); + send_message(&mut stdin, call_req_no_id); - send_message(&mut stdin, call_req); - - let call_resp = read_message(&mut stdout).expect("Failed to read tools/call response"); - - assert_eq!(call_resp["jsonrpc"], "2.0"); - assert_eq!(call_resp["id"], 3); - - if let Some(error) = call_resp.get("error") { - assert_eq!(error["code"], -32603); // Internal Error - } else { - assert!( - call_resp.get("result").is_some(), - "Expected either an error (if Neovim is not running) or a result (if it is)" - ); + drop(stdin); + let mut stderr_output = String::new(); + if let Some(mut stderr) = child.stderr.take() { + let _ = stderr.read_to_string(&mut stderr_output); + println!("Child STDERR: {}", stderr_output); } - - child.kill().expect("Failed to kill child"); - child.wait().expect("Failed to wait on child"); + + let status = child.wait().expect("Failed to wait on child"); + assert!(status.success(), "Child process did not exit successfully. Stderr: {}", stderr_output); }