Test coverage and headless nvim bug fixes

This commit is contained in:
Riz Ashraf committed 2026-09-27 08:10:09 +01:00
1 parent bffc8896f1
commit 4bae4e9c08
9 files changed
+316 -58

No files matched your search

+13
View File
@@ -21,6 +21,9 @@ pub async fn spawn_headless_nvim() -> Result<String, String> {
.arg("--headless") .arg("--headless")
.arg("--listen") .arg("--listen")
.arg(&socket_name) .arg(&socket_name)
.stdin(std::process::Stdio::piped())
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.kill_on_drop(true) .kill_on_drop(true)
.spawn() .spawn()
.map_err(|e| format!("Failed to spawn headless Neovim: {}", e))?; .map_err(|e| format!("Failed to spawn headless Neovim: {}", e))?;
@@ -40,3 +43,13 @@ pub async fn spawn_headless_nvim() -> Result<String, String> {
Ok(socket_name) 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;
}
}
+1
View File
@@ -726,6 +726,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
m m
} else { } else {
tracing::info!("Stdin closed, exiting loop"); tracing::info!("Stdin closed, exiting loop");
headless::kill_headless_nvim().await;
break; break;
}; };
+58
View File
@@ -42,3 +42,61 @@ pub async fn post_event_handler(
let _ = state.handler.state.event_bus_tx.send(event); let _ = state.handler.state.event_bus_tx.send(event);
axum::Json(serde_json::json!({"status": "ok"})) 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<GenericEvent> 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();
}
}
+20 -1
View File
@@ -147,7 +147,16 @@ impl McpTool for CreateRelationsHandler {
} }
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> { async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
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(); let mut missing_nodes = std::collections::HashSet::new();
state.modify_graph(|g| { state.modify_graph(|g| {
for relation in req.relations { for relation in req.relations {
@@ -682,6 +691,16 @@ mod tests {
}); });
let res = handler.execute(args, state.clone()).await.unwrap(); let res = handler.execute(args, state.clone()).await.unwrap();
assert_eq!(res, "Relations created"); 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] #[tokio::test]
+37
View File
@@ -344,6 +344,7 @@ impl McpTool for OmniSearchHandler {
idx.search(&req.query, req.namespace.as_deref()) idx.search(&req.query, req.namespace.as_deref())
.unwrap_or_default() .unwrap_or_default()
}; };
println!("OMNI SEARCH MATCHES: {:?}", matches);
let kg_json = state.read_graph(|full| { let kg_json = state.read_graph(|full| {
let mut kg_entities = std::collections::HashMap::new(); let mut kg_entities = std::collections::HashMap::new();
@@ -660,4 +661,40 @@ mod tests {
.await .await
.unwrap(); .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");
}
} }
+19 -3
View File
@@ -7,8 +7,8 @@ pub type SearchResultTuple = (String, String, String, String, f32);
#[derive(Clone)] #[derive(Clone)]
pub struct MemoryIndex { pub struct MemoryIndex {
index: Index, pub index: Index,
reader: IndexReader, pub reader: IndexReader,
writer: Arc<Mutex<IndexWriter>>, writer: Arc<Mutex<IndexWriter>>,
needs_commit: Arc<std::sync::atomic::AtomicBool>, needs_commit: Arc<std::sync::atomic::AtomicBool>,
@@ -244,18 +244,32 @@ impl MemoryIndex {
self.type_field => "entity", self.type_field => "entity",
self.namespace_field => e.namespace.as_str() 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) { pub fn add_task_sync(&self, t: &Task) {
println!("add_task_sync called for task: {}", t.id);
if let Ok(writer) = self.writer.lock() { 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.id_field => t.id.as_str(),
self.title_field => t.title.as_str(), self.title_field => t.title.as_str(),
self.body_field => t.description.as_str(), self.body_field => t.description.as_str(),
self.type_field => "task", self.type_field => "task",
self.namespace_field => "global" 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.type_field => "snippet",
self.namespace_field => "global" 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.type_field => "adr",
self.namespace_field => "global" self.namespace_field => "global"
)); ));
self.needs_commit.store(true, std::sync::atomic::Ordering::SeqCst);
} }
} }
} }
+42 -29
View File
@@ -128,37 +128,40 @@ impl MemoryState {
} }
pub async fn rebuild_index(self: &Arc<Self>) { pub async fn rebuild_index(self: &Arc<Self>) {
if let Ok(new_idx) = MemoryIndex::new(&self.base_dir) { let idx = self.search_index.read().unwrap().clone();
let idx = new_idx.clone(); idx.delete_all();
let entities: Vec<_> = self.graph.read_with(|g| g.entities.values().cloned().collect()); let entities: Vec<_> = self.graph.read_with(|g| g.entities.values().cloned().collect());
let tasks = self.tasks.read_with(|t| t.clone()); let tasks = self.tasks.read_with(|t| t.clone());
let snippets = self.snippets.read_with(|s| s.clone()); let snippets = self.snippets.read_with(|s| s.clone());
let adrs = self.adrs.read_with(|a| a.clone()); let adrs = self.adrs.read_with(|a| a.clone());
tokio::task::spawn_blocking(move || { println!("rebuild_index: found {} entities, {} tasks", entities.len(), tasks.len());
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);
});
let _ = new_idx.commit().await; let idx_clone = idx.clone();
if let Ok(mut w) = self.search_index.write() { tokio::task::spawn_blocking(move || {
*w = new_idx; 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 // Check search index initialization
let idx = arc_state.search_index.read().unwrap(); let idx = arc_state.search_index.read().unwrap();
// Just verify we can read it without panic // Force reload reader to ensure it sees the commit made by rebuild_index
assert!(idx.search("Test", None).is_ok()); 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");
} }
} }
+23
View File
@@ -562,3 +562,26 @@ pub struct VerifyAcceptanceCriteriaTool {
pub criteria: String, pub criteria: String,
pub proof: 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");
}
}
+103 -25
View File
@@ -1,5 +1,5 @@
use serde_json::{Value, json}; use serde_json::{Value, json};
use std::io::{BufRead, BufReader, Write}; use std::io::{BufRead, BufReader, Write, Read};
use std::process::{Command, Stdio}; use std::process::{Command, Stdio};
fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) { 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.pop();
nvim_exe.push("mcp-memory-win-nvim.exe"); 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) let mut child = Command::new(&nvim_exe)
.env("USERPROFILE", temp_dir.to_str().unwrap())
.env("HOME", temp_dir.to_str().unwrap())
.stdin(Stdio::piped()) .stdin(Stdio::piped())
.stdout(Stdio::piped()) .stdout(Stdio::piped())
.stderr(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"); 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"); assert!(has_get_active_buffer, "Missing nvim_get_active_buffer tool");
// 3. Test negative scenario: tools/call when Neovim is not running let tool_names = vec![
// Since Neovim is not guaranteed to be running on the test agent's system, "nvim_goto_line",
// calling a Neovim-specific tool should gracefully return a JSON-RPC error. "nvim_get_active_buffer",
let call_req = json!({ "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", "jsonrpc": "2.0",
"method": "tools/call", "method": "tools/call",
"params": { "params": {
"name": "nvim_get_active_buffer", "name": "nvim_goto_line",
"arguments": {} "arguments": {}
}, }
"id": 3
}); });
send_message(&mut stdin, call_req_no_id);
send_message(&mut stdin, call_req); drop(stdin);
let mut stderr_output = String::new();
let call_resp = read_message(&mut stdout).expect("Failed to read tools/call response"); if let Some(mut stderr) = child.stderr.take() {
let _ = stderr.read_to_string(&mut stderr_output);
assert_eq!(call_resp["jsonrpc"], "2.0"); println!("Child STDERR: {}", stderr_output);
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)"
);
} }
child.kill().expect("Failed to kill child"); let status = child.wait().expect("Failed to wait on child");
child.wait().expect("Failed to wait on child"); assert!(status.success(), "Child process did not exit successfully. Stderr: {}", stderr_output);
} }