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

+58
View File
@@ -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();
}
}
+20 -1
View File
@@ -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]
+37
View File
@@ -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
View File
@@ -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
View File
@@ -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");
}
}
+23
View File
@@ -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");
}
}