From 37003be62095e4204d010ea3a4a3cbd5cb0e54e4 Mon Sep 17 00:00:00 2001 From: Riz Ashraf Date: Tue, 22 Sep 2026 22:01:38 +0100 Subject: [PATCH] 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. --- server/src/api/rest.rs | 15 ++- server/src/api/setup.rs | 16 +-- server/src/api/ws.rs | 10 +- server/src/db.rs | 1 - server/src/error.rs | 9 +- server/src/handlers/env.rs | 27 +++-- server/src/handlers/graph.rs | 169 ++++++++++++++++++++---------- server/src/handlers/meta.rs | 67 +++++++++--- server/src/handlers/notes.rs | 42 +++++--- server/src/handlers/tasks.rs | 73 +++++++++---- server/src/handlers/utils.rs | 1 - server/src/handlers/workspaces.rs | 59 ++++++++--- server/src/main.rs | 1 - server/src/mcp.rs | 5 +- server/src/models.rs | 2 - server/src/router.rs | 52 +++++---- server/src/state.rs | 48 +++++---- 17 files changed, 399 insertions(+), 198 deletions(-) diff --git a/server/src/api/rest.rs b/server/src/api/rest.rs index cf7a45c..63894bc 100644 --- a/server/src/api/rest.rs +++ b/server/src/api/rest.rs @@ -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 { diff --git a/server/src/api/setup.rs b/server/src/api/setup.rs index 098572f..9dbbbf5 100644 --- a/server/src/api/setup.rs +++ b/server/src/api/setup.rs @@ -201,16 +201,16 @@ pub fn create_router(app_state: Arc) -> 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() { @@ -221,15 +221,15 @@ mod tests { clients: RwLock::new(HashMap::new()), next_id: AtomicUsize::new(1), }); - + let app = create_router(app_state); - + // Test health endpoint let request = Request::builder() .uri("/health") .body(Body::empty()) .unwrap(); - + let response = app.oneshot(request).await.unwrap(); assert_eq!(response.status(), 200); } diff --git a/server/src/api/ws.rs b/server/src/api/ws.rs index 4d80617..1908d44 100644 --- a/server/src/api/ws.rs +++ b/server/src/api/ws.rs @@ -146,10 +146,10 @@ pub async fn handle_socket(socket: WebSocket, state: Arc, _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 {}); diff --git a/server/src/db.rs b/server/src/db.rs index 6378ba1..2a5320f 100644 --- a/server/src/db.rs +++ b/server/src/db.rs @@ -59,4 +59,3 @@ pub fn init_redb(base: &Path) -> Arc { db } - diff --git a/server/src/error.rs b/server/src/error.rs index 1c234f8..75e8656 100644 --- a/server/src/error.rs +++ b/server/src/error.rs @@ -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() { @@ -76,15 +76,14 @@ mod tests { fn test_app_error_display() { let err = AppError::NotFound("test".into()); assert_eq!(err.to_string(), "Not Found: test"); - + let err = AppError::Forbidden("test".into()); assert_eq!(err.to_string(), "Forbidden: test"); - + let err = AppError::Internal("test".into()); assert_eq!(err.to_string(), "Internal Server Error: test"); - + let err = AppError::BadRequest("test".into()); assert_eq!(err.to_string(), "Bad Request: test"); } } - diff --git a/server/src/handlers/env.rs b/server/src/handlers/env.rs index 4a477bf..79ce160 100644 --- a/server/src/handlers/env.rs +++ b/server/src/handlers/env.rs @@ -164,14 +164,14 @@ 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() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let update_handler = UpdateEnvFingerprintHandler; let args = json!({ "namespace": "global", @@ -179,12 +179,15 @@ mod tests { "rustc": "1.70.0" } }); - + let res = update_handler.execute(args, state.clone()).await.unwrap(); 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")); } @@ -193,7 +196,7 @@ mod tests { async fn test_env_details() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + // Ensure namespace is present in test setup state.environments.modify(|e| { e.push(crate::models::EnvironmentDetail { @@ -207,8 +210,11 @@ 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()); } } diff --git a/server/src/handlers/graph.rs b/server/src/handlers/graph.rs index 3e27e9c..4c11f57 100644 --- a/server/src/handlers/graph.rs +++ b/server/src/handlers/graph.rs @@ -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) { - path.push(format!("{} -[{}]-> {}", parent, rel_type, 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) -> Result { 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(); } - } - MemoryState::deduplicate(&mut master.relations); + + if r.from == req.target_entity || r.to == req.target_entity { + seen.insert(r.clone()) + } else { + true + } + }); }); Ok("Entities merged".to_string()) } @@ -581,24 +578,24 @@ 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() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let create_handler = CreateEntitiesHandler; let args = json!({ "entities": [ {"name": "Alice", "entityType": "Person", "observations": ["Likes Bob"]} ] }); - + let res = create_handler.execute(args, state.clone()).await.unwrap(); assert_eq!(res, "Entities created"); - + // Ensure graph contains the entity state.graph.read_with(|g| { assert!(g.entities.contains_key("Alice")); @@ -610,11 +607,29 @@ mod tests { async fn test_create_relations() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + // 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; @@ -631,10 +646,19 @@ mod tests { async fn test_observations_and_reads() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + // 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")); } @@ -663,7 +696,7 @@ mod tests { async fn test_advanced_graph_operations() { let dir = tempfile::tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let create_handler = CreateEntitiesHandler; let args_ent = json!({ "entities": [ @@ -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; diff --git a/server/src/handlers/meta.rs b/server/src/handlers/meta.rs index de4b6a4..87a57c2 100644 --- a/server/src/handlers/meta.rs +++ b/server/src/handlers/meta.rs @@ -512,14 +512,14 @@ 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() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let handler = LogErrorFixHandler; let args = json!({ "signature": "IndexOutOfBounds", @@ -528,7 +528,7 @@ mod tests { "git_commit": "abcdef", "git_branch": "main" }); - + let res = handler.execute(args, state.clone()).await.unwrap(); assert!(res.contains("Error fix logged")); } @@ -537,10 +537,10 @@ mod tests { async fn test_project_health() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let handler = GetProjectHealthHandler; let args = json!({"namespace": "global"}); - + let res = handler.execute(args, state.clone()).await.unwrap(); assert!(res.contains("unresolved_tech_debt")); } @@ -549,7 +549,7 @@ mod tests { async fn test_log_decision_and_tech_debt() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let decision_handler = LogDecisionHandler; let args_dec = json!({ "title": "Architecture", @@ -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 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(); } } diff --git a/server/src/handlers/notes.rs b/server/src/handlers/notes.rs index a0631c3..dd33957 100644 --- a/server/src/handlers/notes.rs +++ b/server/src/handlers/notes.rs @@ -264,24 +264,27 @@ 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() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let add_handler = AddStickyNoteHandler; let args = json!({ "content": "Buy milk", }); - + let res = add_handler.execute(args, state.clone()).await.unwrap(); 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")); } @@ -297,19 +303,22 @@ mod tests { async fn test_handoff_and_summaries() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let handoff_handler = LeaveHandoffMemoHandler; let args = json!({ "content": "Finished implementing graph tests", "author": "Antigravity", "namespace": "global" }); - + let res = handoff_handler.execute(args, state.clone()).await.unwrap(); 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()); } } diff --git a/server/src/handlers/tasks.rs b/server/src/handlers/tasks.rs index 49fa2e3..48d390f 100644 --- a/server/src/handlers/tasks.rs +++ b/server/src/handlers/tasks.rs @@ -491,26 +491,29 @@ 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() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let add_handler = AddTaskHandler; let args = json!({ "title": "Fix the hyperdrive", "description": "It's making a strange noise", "acceptance_criteria": ["Stop the noise", "Reach lightspeed"], }); - + let res = add_handler.execute(args, state.clone()).await.unwrap(); 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")); } @@ -518,10 +521,16 @@ mod tests { async fn test_update_task_status() { let dir = tempdir().unwrap(); 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(); @@ -532,9 +541,12 @@ mod tests { }); let res3 = update_handler.execute(args, state.clone()).await.unwrap(); 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)); } @@ -542,7 +554,7 @@ mod tests { async fn test_milestones_and_criteria() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + // Add Milestone let add_milestone = AddMilestoneHandler; let args_ms = json!({ @@ -555,7 +567,7 @@ mod tests { }); let res1 = add_milestone.execute(args_ms, state.clone()).await.unwrap(); assert!(res1.contains("Milestone added")); - + // Fetch milestone ID from state directly to update let ms_id = state.milestones.read_with(|ms| ms[0].id.clone()); @@ -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).")); } } - diff --git a/server/src/handlers/utils.rs b/server/src/handlers/utils.rs index e73e824..9ce58d1 100644 --- a/server/src/handlers/utils.rs +++ b/server/src/handlers/utils.rs @@ -36,4 +36,3 @@ mod tests { assert!(t2 >= t1); } } - diff --git a/server/src/handlers/workspaces.rs b/server/src/handlers/workspaces.rs index 08bf7d4..6ad5843 100644 --- a/server/src/handlers/workspaces.rs +++ b/server/src/handlers/workspaces.rs @@ -368,14 +368,14 @@ 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() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let save_handler = SaveContextWorkspaceHandler; let args = json!({ "name": "wsl-session", @@ -383,12 +383,15 @@ mod tests { "pinned_files": ["src/main.rs"], "active_task_ids": ["123"] }); - + let res = save_handler.execute(args, state.clone()).await.unwrap(); 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")); } @@ -397,7 +400,7 @@ mod tests { async fn test_snippets_and_pr_checklists() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); - + let store_handler = StoreSnippetHandler; let args_snip = json!({ "name": "init_db", @@ -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"); } } diff --git a/server/src/main.rs b/server/src/main.rs index 44608c6..9e3d473 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -321,4 +321,3 @@ fn main() -> Result<(), Box> { Ok(()) } - diff --git a/server/src/mcp.rs b/server/src/mcp.rs index 25b5a50..1f6f67d 100644 --- a/server/src/mcp.rs +++ b/server/src/mcp.rs @@ -36,7 +36,7 @@ mod tests { use super::*; use schemars::JsonSchema; use serde::Serialize; - + #[derive(JsonSchema, Serialize)] struct DummyTool { name: String, @@ -66,11 +66,10 @@ mod tests { assert_eq!(def["name"], "dummy"); assert_eq!(def["description"], "A dummy tool"); assert!(def["inputSchema"].is_object()); - + let schema = def["inputSchema"].as_object().unwrap(); let properties = schema["properties"].as_object().unwrap(); assert!(properties.contains_key("name")); assert!(properties.contains_key("age")); } } - diff --git a/server/src/models.rs b/server/src/models.rs index f20cc43..12c55f3 100644 --- a/server/src/models.rs +++ b/server/src/models.rs @@ -196,5 +196,3 @@ pub struct GateRecord { pub reason: Option, pub timestamp: u64, } - - diff --git a/server/src/router.rs b/server/src/router.rs index cdc0412..326ae3a 100644 --- a/server/src/router.rs +++ b/server/src/router.rs @@ -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 text = match uri { + let uri = params + .get("uri") + .and_then(|u| u.as_str()) + .unwrap_or("") + .to_string(); + + 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!({ @@ -245,7 +254,7 @@ impl MemoryHandler { "prompts/get" => { let params = req.get("params").unwrap_or(&serde_json::Value::Null); let name = params.get("name").and_then(|n| n.as_str()).unwrap_or(""); - + if name == "analyze_tech_debt" { let payload = serde_json::json!({ "messages": [ @@ -318,19 +327,19 @@ 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() { let dir = tempdir().unwrap(); let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); let handler = MemoryHandler::new(state); - + // Assert some known tools are registered assert!(handler.tools.contains_key("create_entities")); assert!(handler.tools.contains_key("add_task")); - + // Ensure we can fetch list of tools let list_tools_req = json!({ "jsonrpc": "2.0", @@ -338,11 +347,10 @@ mod tests { "method": "tools/list", "params": {} }); - + let res_list = handler.handle_request(list_tools_req).await.unwrap(); assert_eq!(res_list["jsonrpc"], "2.0"); assert_eq!(res_list["id"], 1); assert!(res_list["result"]["tools"].as_array().unwrap().len() > 10); } } - diff --git a/server/src/state.rs b/server/src/state.rs index 6c1df49..bd530dc 100644 --- a/server/src/state.rs +++ b/server/src/state.rs @@ -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| { @@ -168,10 +170,10 @@ mod tests { async fn test_memory_state_initialization() { let dir = tempdir().unwrap(); let state = MemoryState::new(dir.path().to_str().unwrap()); - + // Ensure state fields are properly initialized assert_eq!(state.base_dir, dir.path()); - + // Write a test value state.tasks.modify(|tasks| { tasks.push(Task { @@ -187,17 +189,17 @@ mod tests { parent_id: None, }); }); - + // Ensure it is saved state.tasks.read_with(|tasks| { assert_eq!(tasks.len(), 1); assert_eq!(tasks[0].id, "123"); }); - + // Test rebuild index let arc_state = Arc::new(state); arc_state.rebuild_index().await; - + // Check search index initialization let idx = arc_state.search_index.read().unwrap(); // Just verify we can read it without panic