refactor: apply rust best practices and fix memory optimizations

- Optimized memory allocation in router.rs by offloading JSON serialization to spawn_blocking and using references.
- Prevented full graph duplication on startup in state.rs index rebuild.
- Eliminated massive String allocations in QueryGraphPathHandler BFS loops.
- Avoided temporary Strings in VisualizeGraphHandler via inline writing.
- Fixed O(N) full-graph deduplication in MergeEntitiesHandler to scale efficiently.
This commit is contained in:
Riz Ashraf committed 2026-09-22 22:01:38 +01:00
1 parent 251757f8fc
commit 37003be620
17 files changed
+342 -141

No files matched your search

+11 -4
View File
@@ -113,8 +113,8 @@ mod tests {
use super::*;
use crate::router::MemoryHandler;
use crate::state::MemoryState;
use std::sync::atomic::AtomicUsize;
use std::sync::RwLock;
use std::sync::atomic::AtomicUsize;
use tempfile::tempdir;
#[tokio::test]
@@ -137,7 +137,9 @@ mod tests {
block: None,
reason: None,
};
let res_set = gate_set_handler(State(app_state.clone()), Json(set_req)).await.unwrap();
let res_set = gate_set_handler(State(app_state.clone()), Json(set_req))
.await
.unwrap();
assert_eq!(res_set.into_response().status(), axum::http::StatusCode::OK);
// Verify the gate (and consume it)
@@ -148,8 +150,13 @@ mod tests {
params: HashMap::new(),
consume: true,
};
let res_verify = gate_verify_handler(State(app_state.clone()), Query(verify_req)).await.unwrap();
assert_eq!(res_verify.into_response().status(), axum::http::StatusCode::OK);
let res_verify = gate_verify_handler(State(app_state.clone()), Query(verify_req))
.await
.unwrap();
assert_eq!(
res_verify.into_response().status(),
axum::http::StatusCode::OK
);
// Verify again should fail since it was consumed
let verify_req2 = GateVerifyReq {
+5 -5
View File
@@ -201,16 +201,16 @@ pub fn create_router(app_state: Arc<AppState>) -> Router {
#[cfg(test)]
mod tests {
use super::*;
use crate::AppState;
use crate::router::MemoryHandler;
use crate::state::MemoryState;
use crate::AppState;
use axum::http::Request;
use axum::body::Body;
use std::sync::atomic::AtomicUsize;
use std::sync::RwLock;
use axum::http::Request;
use std::collections::HashMap;
use tower::ServiceExt;
use std::sync::RwLock;
use std::sync::atomic::AtomicUsize;
use tempfile::tempdir;
use tower::ServiceExt;
#[tokio::test]
async fn test_create_router_health() {
+7 -3
View File
@@ -146,10 +146,10 @@ pub async fn handle_socket(socket: WebSocket, state: Arc<AppState>, _client_type
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicUsize;
use std::sync::RwLock;
use crate::router::MemoryHandler;
use crate::state::MemoryState;
use std::sync::RwLock;
use std::sync::atomic::AtomicUsize;
use tempfile::tempdir;
#[tokio::test]
@@ -163,7 +163,11 @@ mod tests {
});
// Insert a dummy client
app_state.clients.write().unwrap().insert("test-session".to_string(), tokio::sync::mpsc::channel(1).0);
app_state
.clients
.write()
.unwrap()
.insert("test-session".to_string(), tokio::sync::mpsc::channel(1).0);
let send_task = tokio::spawn(async {});
let recv_task = tokio::spawn(async {});
-1
View File
@@ -59,4 +59,3 @@ pub fn init_redb(base: &Path) -> Arc<Database> {
db
}
+1 -2
View File
@@ -41,8 +41,8 @@ impl IntoResponse for AppError {
#[cfg(test)]
mod tests {
use super::*;
use axum::response::IntoResponse;
use axum::http::StatusCode;
use axum::response::IntoResponse;
#[test]
fn test_app_error_not_found() {
@@ -87,4 +87,3 @@ mod tests {
assert_eq!(err.to_string(), "Bad Request: test");
}
}
+14 -5
View File
@@ -164,8 +164,8 @@ impl McpTool for GetEnvironmentDetailsHandler {
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_env_fingerprint() {
@@ -184,7 +184,10 @@ mod tests {
assert_eq!(res, "Env fingerprint updated");
let read_handler = ReadEnvFingerprintHandler;
let res2 = read_handler.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res2 = read_handler
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res2.contains("rustc"));
assert!(res2.contains("1.70.0"));
}
@@ -207,7 +210,10 @@ mod tests {
});
let handler = GetEnvironmentDetailsHandler;
let res = handler.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res = handler
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res.contains("global"));
}
@@ -241,8 +247,11 @@ mod tests {
assert_eq!(res2, "Environment registered");
let get_handler = GetEnvironmentDetailsHandler;
let res3 = get_handler.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res3 = get_handler
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res3.contains("prod.local"));
assert!(res3.len() > 0);
assert!(!res3.is_empty());
}
}
+105 -48
View File
@@ -31,7 +31,7 @@ impl McpTool for QueryGraphPathHandler {
let max_depth = req.max_depth.unwrap_or(5);
let mut queue: std::collections::VecDeque<&str> = std::collections::VecDeque::new();
let mut visited: std::collections::HashSet<&str> = std::collections::HashSet::new();
let mut parents: std::collections::HashMap<&str, (&str, std::borrow::Cow<'_, str>)> =
let mut parents: std::collections::HashMap<&str, (&str, &str, bool)> =
std::collections::HashMap::new();
queue.push_back(req.start_node.as_str());
@@ -54,10 +54,7 @@ impl McpTool for QueryGraphPathHandler {
visited.insert(rel.to.as_str());
parents.insert(
rel.to.as_str(),
(
current,
std::borrow::Cow::Borrowed(rel.relation_type.as_str()),
),
(current, rel.relation_type.as_str(), false),
);
queue.push_back(rel.to.as_str());
nodes_at_next_depth += 1;
@@ -65,13 +62,7 @@ impl McpTool for QueryGraphPathHandler {
visited.insert(rel.from.as_str());
parents.insert(
rel.from.as_str(),
(
current,
std::borrow::Cow::Owned(format!(
"inverse({})",
rel.relation_type
)),
),
(current, rel.relation_type.as_str(), true),
);
queue.push_back(rel.from.as_str());
nodes_at_next_depth += 1;
@@ -89,8 +80,12 @@ impl McpTool for QueryGraphPathHandler {
let mut path = Vec::new();
let mut curr = req.end_node.as_str();
while curr != req.start_node {
if let Some((parent, rel_type)) = parents.get(&curr) {
if let Some((parent, rel_type, is_inverse)) = parents.get(&curr) {
if *is_inverse {
path.push(format!("{} -[inverse({})]-> {}", parent, rel_type, curr));
} else {
path.push(format!("{} -[{}]-> {}", parent, rel_type, curr));
}
curr = parent;
} else {
break;
@@ -406,7 +401,6 @@ impl McpTool for VisualizeGraphHandler {
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: VisualizeGraphTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let query = req.query.unwrap_or_default();
use std::fmt::Write;
let mut output = state.read_graph(|full| {
let mut included: HashSet<&str> = HashSet::new();
let mut to_draw = Vec::new();
@@ -444,36 +438,33 @@ impl McpTool for VisualizeGraphHandler {
let mut out = String::with_capacity(included.len() * 40 + to_draw.len() * 60);
out.push_str("graph TD;\n");
let sanitize = |s: &str, id_mode: bool| -> String {
let mut o = String::with_capacity(s.len());
let sanitize_to = |out_str: &mut String, s: &str, id_mode: bool| {
for c in s.chars() {
if c != '"' && c != '(' && c != ')' {
if id_mode && (c == ' ' || c == '-' || c == '.') {
o.push('_');
out_str.push('_');
} else {
o.push(c);
out_str.push(c);
}
}
}
o
};
for name in &included {
let _ = writeln!(
out,
" id_{}[\"{}\"];",
sanitize(name, true),
sanitize(name, false)
);
out.push_str(" id_");
sanitize_to(&mut out, name, true);
out.push_str("[\"");
sanitize_to(&mut out, name, false);
out.push_str("\"];\n");
}
for r in to_draw {
let _ = writeln!(
out,
" id_{}-->|\"{}\"|id_{};",
sanitize(&r.from, true),
r.relation_type.replace("\"", ""),
sanitize(&r.to, true)
);
out.push_str(" id_");
sanitize_to(&mut out, &r.from, true);
out.push_str("-->|\"");
out.push_str(&r.relation_type.replace("\"", ""));
out.push_str("\"|id_");
sanitize_to(&mut out, &r.to, true);
out.push_str(";\n");
}
out
});
@@ -532,15 +523,21 @@ impl McpTool for MergeEntitiesHandler {
master.entities.insert(req.target_entity.clone(), new_tgt);
}
}
for r in &mut master.relations {
let mut seen = std::collections::HashSet::new();
master.relations.retain_mut(|r| {
if r.from == req.source_entity {
r.from = req.target_entity.clone();
}
if r.to == req.source_entity {
r.to = req.target_entity.clone();
}
if r.from == req.target_entity || r.to == req.target_entity {
seen.insert(r.clone())
} else {
true
}
MemoryState::deduplicate(&mut master.relations);
});
});
Ok("Entities merged".to_string())
}
@@ -581,8 +578,8 @@ use crate::handlers::utils::*;
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_create_and_read_entities() {
@@ -613,8 +610,26 @@ mod tests {
// Needs entities first
state.graph.modify(|g| {
g.entities.insert("A".to_string(), crate::models::Entity { name: "A".to_string(), entity_type: "Node".to_string(), observations: vec![], namespace: "global".to_string(), git_branch: None });
g.entities.insert("B".to_string(), crate::models::Entity { name: "B".to_string(), entity_type: "Node".to_string(), observations: vec![], namespace: "global".to_string(), git_branch: None });
g.entities.insert(
"A".to_string(),
crate::models::Entity {
name: "A".to_string(),
entity_type: "Node".to_string(),
observations: vec![],
namespace: "global".to_string(),
git_branch: None,
},
);
g.entities.insert(
"B".to_string(),
crate::models::Entity {
name: "B".to_string(),
entity_type: "Node".to_string(),
observations: vec![],
namespace: "global".to_string(),
git_branch: None,
},
);
});
let handler = CreateRelationsHandler;
@@ -634,7 +649,16 @@ mod tests {
// Need entity first
state.graph.modify(|g| {
g.entities.insert("A".to_string(), crate::models::Entity { name: "A".to_string(), entity_type: "Node".to_string(), observations: vec![], namespace: "global".to_string(), git_branch: None });
g.entities.insert(
"A".to_string(),
crate::models::Entity {
name: "A".to_string(),
entity_type: "Node".to_string(),
observations: vec![],
namespace: "global".to_string(),
git_branch: None,
},
);
});
let add_obs = AddObservationsHandler;
@@ -647,15 +671,24 @@ mod tests {
assert_eq!(res1, "Observations added");
let read_graph = ReadGraphHandler;
let res2 = read_graph.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res2 = read_graph
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res2.contains("Obs 1"));
assert!(res2.contains("Obs 2"));
let del_entity = DeleteEntitiesHandler;
let res4 = del_entity.execute(json!({"entityNames": ["A"]}), state.clone()).await.unwrap();
let res4 = del_entity
.execute(json!({"entityNames": ["A"]}), state.clone())
.await
.unwrap();
assert_eq!(res4, "Entities deleted");
let res5 = read_graph.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res5 = read_graph
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(!res5.contains("A"));
}
@@ -671,7 +704,10 @@ mod tests {
{"name": "Y", "entityType": "File", "observations": ["Obs Y"], "namespace": "global"}
]
});
create_handler.execute(args_ent, state.clone()).await.unwrap();
create_handler
.execute(args_ent, state.clone())
.await
.unwrap();
let rel_handler = CreateRelationsHandler;
let args_rel = json!({
@@ -682,24 +718,45 @@ mod tests {
rel_handler.execute(args_rel, state.clone()).await.unwrap();
let read_handler = ReadGraphHandler;
let res_read = read_handler.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res_read = read_handler
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res_read.contains("X"));
assert!(res_read.contains("depends_on"));
let open_handler = OpenNodesHandler;
let res_open = open_handler.execute(json!({"names": ["X"]}), state.clone()).await.unwrap();
let res_open = open_handler
.execute(json!({"names": ["X"]}), state.clone())
.await
.unwrap();
assert!(res_open.contains("Y"));
let viz_handler = VisualizeGraphHandler;
let res_viz = viz_handler.execute(json!({"query": "X"}), state.clone()).await.unwrap();
assert!(res_viz.len() > 0);
let res_viz = viz_handler
.execute(json!({"query": "X"}), state.clone())
.await
.unwrap();
assert!(!res_viz.is_empty());
let condense = CondenseEntityHandler;
let res_cond = condense.execute(json!({"entityName": "X", "summarized_observations": ["X condensed"]}), state.clone()).await.unwrap();
let res_cond = condense
.execute(
json!({"entityName": "X", "summarized_observations": ["X condensed"]}),
state.clone(),
)
.await
.unwrap();
assert_eq!(res_cond, "Entity condensed");
let merge = MergeEntitiesHandler;
let res_merge = merge.execute(json!({"sourceEntity": "X", "targetEntity": "Y"}), state.clone()).await.unwrap();
let res_merge = merge
.execute(
json!({"sourceEntity": "X", "targetEntity": "Y"}),
state.clone(),
)
.await
.unwrap();
assert_eq!(res_merge, "Entities merged");
let orphans = FindOrphansHandler;
+43 -10
View File
@@ -512,8 +512,8 @@ use crate::handlers::utils::*;
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_log_error_fix() {
@@ -557,7 +557,10 @@ mod tests {
"decision": "Use SQLite",
"consequence": "Simple",
});
let res1 = decision_handler.execute(args_dec, state.clone()).await.unwrap();
let res1 = decision_handler
.execute(args_dec, state.clone())
.await
.unwrap();
assert_eq!(res1, "Decision logged as ADR-0001");
let debt_handler = LogTechDebtHandler;
@@ -571,11 +574,20 @@ mod tests {
"git_branch": "main",
"namespace": "global"
});
let res2 = debt_handler.execute(args_debt, state.clone()).await.unwrap();
let res2 = debt_handler
.execute(args_debt, state.clone())
.await
.unwrap();
assert_eq!(res2, "Tech debt logged");
let list_debt = ListTechDebtHandler;
let res3 = list_debt.execute(json!({"namespace": "global", "include_resolved": false}), state.clone()).await.unwrap();
let res3 = list_debt
.execute(
json!({"namespace": "global", "include_resolved": false}),
state.clone(),
)
.await
.unwrap();
assert!(res3.contains("Hardcoded path"));
let pref_handler = LearnPreferenceHandler;
@@ -583,7 +595,10 @@ mod tests {
"key": "formatting",
"value": "use spaces",
});
let res4 = pref_handler.execute(args_pref, state.clone()).await.unwrap();
let res4 = pref_handler
.execute(args_pref, state.clone())
.await
.unwrap();
assert_eq!(res4, "Preference learned");
let read_pref = ReadPreferencesHandler;
@@ -603,10 +618,16 @@ mod tests {
"git_commit": "def",
"git_branch": "main"
});
code_handler.execute(args_code, state.clone()).await.unwrap();
code_handler
.execute(args_code, state.clone())
.await
.unwrap();
let query_changes = QueryRecentChangesHandler;
let res_changes = query_changes.execute(json!({}), state.clone()).await.unwrap();
let res_changes = query_changes
.execute(json!({}), state.clone())
.await
.unwrap();
assert!(res_changes.contains("main.rs"));
let debt_handler = LogTechDebtHandler;
@@ -620,15 +641,27 @@ mod tests {
"git_branch": "main",
"namespace": "global"
});
debt_handler.execute(args_debt, state.clone()).await.unwrap();
debt_handler
.execute(args_debt, state.clone())
.await
.unwrap();
// resolve it
let list_debt = ListTechDebtHandler;
let debt_list = list_debt.execute(json!({"namespace": "global", "include_resolved": false}), state.clone()).await.unwrap();
let debt_list = list_debt
.execute(
json!({"namespace": "global", "include_resolved": false}),
state.clone(),
)
.await
.unwrap();
let uuid_start = debt_list.find("id\":\"").unwrap() + 5;
let uuid = &debt_list[uuid_start..uuid_start + 36];
let resolve_debt = ResolveTechDebtHandler;
resolve_debt.execute(json!({"id": uuid}), state.clone()).await.unwrap();
resolve_debt
.execute(json!({"id": uuid}), state.clone())
.await
.unwrap();
}
}
+25 -7
View File
@@ -264,8 +264,8 @@ impl McpTool for GenerateStandupReportHandler {
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_notes_lifecycle() {
@@ -281,7 +281,10 @@ mod tests {
assert!(res.contains("Sticky note added"));
let read_handler = ReadStickyNotesHandler;
let res2 = read_handler.execute(json!({}), state.clone()).await.unwrap();
let res2 = read_handler
.execute(json!({}), state.clone())
.await
.unwrap();
assert!(res2.contains("Buy milk"));
let delete_handler = DeleteStickyNoteHandler;
@@ -289,7 +292,10 @@ mod tests {
let res3 = delete_handler.execute(args2, state.clone()).await.unwrap();
assert_eq!(res3, "Sticky note deleted.");
let res4 = read_handler.execute(json!({}), state.clone()).await.unwrap();
let res4 = read_handler
.execute(json!({}), state.clone())
.await
.unwrap();
assert!(!res4.contains("Buy milk"));
}
@@ -309,7 +315,10 @@ mod tests {
assert_eq!(res, "Handoff memo left");
let read_handoff = ReadHandoffMemosHandler;
let res2 = read_handoff.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res2 = read_handoff
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res2.contains("Finished implementing graph tests"));
let summary_handler = AddSessionSummaryHandler;
@@ -317,11 +326,20 @@ mod tests {
"summary": "Completed a bunch of tests",
"namespace": "global"
});
let res3 = summary_handler.execute(args_sum, state.clone()).await.unwrap();
let res3 = summary_handler
.execute(args_sum, state.clone())
.await
.unwrap();
assert_eq!(res3, "Session summary added");
let standup_handler = GenerateStandupReportHandler;
let res4 = standup_handler.execute(json!({"namespace": "global", "hours_lookback": 24}), state.clone()).await.unwrap();
assert!(res4.len() > 0);
let res4 = standup_handler
.execute(
json!({"namespace": "global", "hours_lookback": 24}),
state.clone(),
)
.await
.unwrap();
assert!(!res4.is_empty());
}
}
+48 -11
View File
@@ -491,8 +491,8 @@ impl McpTool for ListMilestonesHandler {
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_add_task_and_list() {
@@ -510,7 +510,10 @@ mod tests {
assert!(res.contains("Task added with ID:"));
let list_handler = ListActiveTasksHandler;
let res2 = list_handler.execute(json!({}), state.clone()).await.unwrap();
let res2 = list_handler
.execute(json!({}), state.clone())
.await
.unwrap();
assert!(res2.contains("Fix the hyperdrive"));
}
@@ -520,7 +523,13 @@ mod tests {
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let add_handler = AddTaskHandler;
let res = add_handler.execute(json!({"title": "Test", "description": "test"}), state.clone()).await.unwrap();
let res = add_handler
.execute(
json!({"title": "Test", "description": "test"}),
state.clone(),
)
.await
.unwrap();
let id_start = res.find("ID: ").unwrap() + 4;
let task_id = res[id_start..].trim();
@@ -534,7 +543,10 @@ mod tests {
assert_eq!(res3, "Task status updated.");
let list_handler = ListActiveTasksHandler;
let res4 = list_handler.execute(json!({}), state.clone()).await.unwrap();
let res4 = list_handler
.execute(json!({}), state.clone())
.await
.unwrap();
assert!(!res4.contains(task_id));
}
@@ -570,13 +582,22 @@ mod tests {
// List Milestones
let list_ms = ListMilestonesHandler;
let res3 = list_ms.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res3 = list_ms
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res3.contains("completed"));
assert!(res3.contains("Release 1.0"));
// Task Acceptance Criteria
let add_task = AddTaskHandler;
let res_task = add_task.execute(json!({"title": "Test", "description": "desc"}), state.clone()).await.unwrap();
let res_task = add_task
.execute(
json!({"title": "Test", "description": "desc"}),
state.clone(),
)
.await
.unwrap();
let task_id = res_task[res_task.find("ID: ").unwrap() + 4..].trim();
let set_ac = SetAcceptanceCriteriaHandler;
@@ -604,15 +625,31 @@ mod tests {
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let add_task = AddTaskHandler;
let parent = add_task.execute(json!({"title": "Parent", "description": "p"}), state.clone()).await.unwrap();
let parent_id = parent[parent.find("ID: ").unwrap() + 4..].trim().to_string();
let parent = add_task
.execute(
json!({"title": "Parent", "description": "p"}),
state.clone(),
)
.await
.unwrap();
let parent_id = parent[parent.find("ID: ").unwrap() + 4..]
.trim()
.to_string();
let child = add_task.execute(json!({"title": "Child", "description": "c", "parent_id": parent_id}), state.clone()).await.unwrap();
let child = add_task
.execute(
json!({"title": "Child", "description": "c", "parent_id": parent_id}),
state.clone(),
)
.await
.unwrap();
let _child_id = child[child.find("ID: ").unwrap() + 4..].trim().to_string();
let del_task = DeleteTaskHandler;
let res_del = del_task.execute(json!({"id": parent_id}), state.clone()).await.unwrap();
let res_del = del_task
.execute(json!({"id": parent_id}), state.clone())
.await
.unwrap();
assert!(res_del.contains("Deleted task and its children (2 total)."));
}
}
-1
View File
@@ -36,4 +36,3 @@ mod tests {
assert!(t2 >= t1);
}
}
+42 -9
View File
@@ -368,8 +368,8 @@ use crate::handlers::utils::*;
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_workspace_lifecycle() {
@@ -388,7 +388,10 @@ mod tests {
assert_eq!(res, "Context workspace saved");
let list_handler = ListContextWorkspacesHandler;
let res2 = list_handler.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res2 = list_handler
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res2.contains("wsl-session"));
assert!(res2.contains("src/main.rs"));
}
@@ -406,11 +409,20 @@ mod tests {
"code": "SELECT 1;",
"namespace": "global"
});
let res1 = store_handler.execute(args_snip, state.clone()).await.unwrap();
let res1 = store_handler
.execute(args_snip, state.clone())
.await
.unwrap();
assert_eq!(res1, "Snippet 'init_db' stored.");
let search_handler = SearchSnippetsHandler;
let _res2 = search_handler.execute(json!({"query": "SELECT", "namespace": "global"}), state.clone()).await.unwrap();
let _res2 = search_handler
.execute(
json!({"query": "SELECT", "namespace": "global"}),
state.clone(),
)
.await
.unwrap();
// Skip assertion since it requires index rebuild
let pr_handler = AddPrChecklistItemHandler;
@@ -422,25 +434,46 @@ mod tests {
assert_eq!(res3, "PR checklist item added");
let get_pr = GetPrChecklistHandler;
let res4 = get_pr.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res4 = get_pr
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res4.contains("Check coverage"));
// Pin lifecycle
let pin = PinFileHandler;
let res5 = pin.execute(json!({"file_path": "src/lib.rs", "namespace": "global"}), state.clone()).await.unwrap();
let res5 = pin
.execute(
json!({"file_path": "src/lib.rs", "namespace": "global"}),
state.clone(),
)
.await
.unwrap();
assert_eq!(res5, "File pinned");
let list_pins = ListPinnedFilesHandler;
let res6 = list_pins.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res6 = list_pins
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert!(res6.contains("src/lib.rs"));
let unpin = UnpinFileHandler;
let res7 = unpin.execute(json!({"file_path": "src/lib.rs", "namespace": "global"}), state.clone()).await.unwrap();
let res7 = unpin
.execute(
json!({"file_path": "src/lib.rs", "namespace": "global"}),
state.clone(),
)
.await
.unwrap();
assert_eq!(res7, "File unpinned");
// Clear PR
let clear_pr = ClearPrChecklistHandler;
let res8 = clear_pr.execute(json!({"namespace": "global"}), state.clone()).await.unwrap();
let res8 = clear_pr
.execute(json!({"namespace": "global"}), state.clone())
.await
.unwrap();
assert_eq!(res8, "PR checklist cleared");
}
}
-1
View File
@@ -321,4 +321,3 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
Ok(())
}
-1
View File
@@ -73,4 +73,3 @@ mod tests {
assert!(properties.contains_key("age"));
}
}
-2
View File
@@ -196,5 +196,3 @@ pub struct GateRecord {
pub reason: Option<String>,
pub timestamp: u64,
}
+25 -17
View File
@@ -195,30 +195,39 @@ impl MemoryHandler {
}
"resources/read" => {
let params = req.get("params").unwrap_or(&serde_json::Value::Null);
let uri = params.get("uri").and_then(|u| u.as_str()).unwrap_or("");
let uri = params
.get("uri")
.and_then(|u| u.as_str())
.unwrap_or("")
.to_string();
let text = match uri {
let state_clone = Arc::clone(&self.state);
let uri_clone = uri.clone();
let text = match tokio::task::spawn_blocking(move || match uri_clone.as_str() {
"memory://graph/entities" => {
let graph = self.state.graph.cache.read().unwrap();
let data: Vec<_> = graph.entities.values().cloned().collect();
serde_json::to_string_pretty(&data).unwrap_or_default()
let graph = state_clone.graph.cache.read().unwrap();
let data: Vec<_> = graph.entities.values().collect();
Some(serde_json::to_string_pretty(&data).unwrap_or_default())
}
"memory://graph/relations" => {
let graph = self.state.graph.cache.read().unwrap();
let data = graph.relations.clone();
serde_json::to_string_pretty(&data).unwrap_or_default()
let graph = state_clone.graph.cache.read().unwrap();
let data = &graph.relations;
Some(serde_json::to_string_pretty(&data).unwrap_or_default())
}
"memory://tasks/active" => {
let tasks = self.state.tasks.cache.read().unwrap();
let data: Vec<_> = tasks.iter()
let tasks = state_clone.tasks.cache.read().unwrap();
let data: Vec<_> = tasks
.iter()
.filter(|t| t.status != "completed" && t.status != "done")
.cloned()
.collect();
serde_json::to_string_pretty(&data).unwrap_or_default()
}
_ => {
return Some(crate::mcp::error(id, -32602, "Resource not found"));
Some(serde_json::to_string_pretty(&data).unwrap_or_default())
}
_ => None,
})
.await
{
Ok(Some(text)) => text,
_ => return Some(crate::mcp::error(id, -32602, "Resource not found")),
};
let payload = serde_json::json!({
@@ -318,8 +327,8 @@ impl MemoryHandler {
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_memory_handler_tools_registration() {
@@ -345,4 +354,3 @@ mod tests {
assert!(res_list["result"]["tools"].as_array().unwrap().len() > 10);
}
}
+16 -14
View File
@@ -124,27 +124,29 @@ impl MemoryState {
let idx = new_idx.clone();
tokio::task::spawn_blocking(move || {
let entities: Vec<_> = state
.graph
.read_with(|g| g.entities.values().cloned().collect());
for e in entities {
idx.add_entity_sync(&e);
state.graph.read_with(|g| {
for e in g.entities.values() {
idx.add_entity_sync(e);
}
});
let tasks = state.tasks.read_with(|t| t.clone());
for t in tasks {
idx.add_task_sync(&t);
state.tasks.read_with(|t| {
for task in t {
idx.add_task_sync(task);
}
});
let snippets = state.snippets.read_with(|s| s.clone());
for s in snippets {
idx.add_snippet_sync(&s);
state.snippets.read_with(|s| {
for snippet in s {
idx.add_snippet_sync(snippet);
}
});
let adrs = state.adrs.read_with(|a| a.clone());
for a in adrs {
idx.add_adr_sync(&a);
state.adrs.read_with(|a| {
for adr in a {
idx.add_adr_sync(adr);
}
});
})
.await
.unwrap_or_else(|e| {