Test coverage and headless nvim bug fixes
This commit is contained in:
1 parent
bffc8896f1
commit
4bae4e9c08
9 files changed
+315
-57
No files matched your search
@@ -21,6 +21,9 @@ pub async fn spawn_headless_nvim() -> Result<String, String> {
|
||||
.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<String, String> {
|
||||
|
||||
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;
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
@@ -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<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();
|
||||
}
|
||||
}
|
||||
@@ -147,7 +147,16 @@ impl McpTool for CreateRelationsHandler {
|
||||
}
|
||||
|
||||
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();
|
||||
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]
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
+19
-3
@@ -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<Mutex<IndexWriter>>,
|
||||
needs_commit: Arc<std::sync::atomic::AtomicBool>,
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+42
-29
@@ -128,37 +128,40 @@ impl MemoryState {
|
||||
}
|
||||
|
||||
pub async fn rebuild_index(self: &Arc<Self>) {
|
||||
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");
|
||||
}
|
||||
}
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
Reference in new issue
Block a user