diff --git a/server/src/router.rs b/server/src/router.rs index 20e1211..a7d356c 100644 --- a/server/src/router.rs +++ b/server/src/router.rs @@ -15,24 +15,134 @@ pub trait McpTool: Send + Sync { async fn execute(&self, args: Value, state: Arc) -> Result; } - #[async_trait] pub trait McpResource: Send + Sync { fn uri(&self) -> &'static str; fn name(&self) -> &'static str; - fn description(&self) -> Option<&'static str> { None } - fn mime_type(&self) -> Option<&'static str> { Some("application/json") } + fn description(&self) -> Option<&'static str> { + None + } + fn mime_type(&self) -> Option<&'static str> { + Some("application/json") + } async fn read(&self, state: Arc) -> Result; } #[async_trait] pub trait McpPrompt: Send + Sync { fn name(&self) -> &'static str; - fn description(&self) -> Option<&'static str> { None } - fn arguments(&self) -> serde_json::Value { serde_json::json!([]) } + fn description(&self) -> Option<&'static str> { + None + } + fn arguments(&self) -> serde_json::Value { + serde_json::json!([]) + } async fn get(&self, args: Value, state: Arc) -> Result; } +struct GraphEntitiesResource; +#[async_trait] +impl McpResource for GraphEntitiesResource { + fn uri(&self) -> &'static str { + "memory://graph/entities" + } + fn name(&self) -> &'static str { + "Graph Entities" + } + fn description(&self) -> Option<&'static str> { + Some("All nodes and entities currently stored in the knowledge graph") + } + async fn read(&self, state: Arc) -> Result { + let state_clone = Arc::clone(&state); + tokio::task::spawn_blocking(move || { + let graph = state_clone.graph.cache.read().unwrap(); + let data: Vec<_> = graph.entities.values().collect(); + Ok(serde_json::to_string_pretty(&data).unwrap_or_default()) + }) + .await + .unwrap() + } +} + +struct GraphRelationsResource; +#[async_trait] +impl McpResource for GraphRelationsResource { + fn uri(&self) -> &'static str { + "memory://graph/relations" + } + fn name(&self) -> &'static str { + "Graph Relations" + } + fn description(&self) -> Option<&'static str> { + Some("All relationships between entities currently stored in the knowledge graph") + } + async fn read(&self, state: Arc) -> Result { + let state_clone = Arc::clone(&state); + tokio::task::spawn_blocking(move || { + let graph = state_clone.graph.cache.read().unwrap(); + let data = &graph.relations; + Ok(serde_json::to_string_pretty(&data).unwrap_or_default()) + }) + .await + .unwrap() + } +} + +struct TasksActiveResource; +#[async_trait] +impl McpResource for TasksActiveResource { + fn uri(&self) -> &'static str { + "memory://tasks/active" + } + fn name(&self) -> &'static str { + "Active Tasks" + } + fn description(&self) -> Option<&'static str> { + Some("List of currently active tasks") + } + async fn read(&self, state: Arc) -> Result { + let state_clone = Arc::clone(&state); + tokio::task::spawn_blocking(move || { + let tasks = state_clone.tasks.cache.read().unwrap(); + let data: Vec<_> = tasks + .iter() + .filter(|t| t.status != "completed" && t.status != "done") + .collect(); + Ok(serde_json::to_string_pretty(&data).unwrap_or_default()) + }) + .await + .unwrap() + } +} + +struct AnalyzeTechDebtPrompt; +#[async_trait] +impl McpPrompt for AnalyzeTechDebtPrompt { + fn name(&self) -> &'static str { + "analyze_tech_debt" + } + fn description(&self) -> Option<&'static str> { + Some("Analyze the project's current technical debt") + } + async fn get( + &self, + _args: Value, + _state: Arc, + ) -> Result { + Ok(serde_json::json!({ + "messages": [ + { + "role": "user", + "content": { + "type": "text", + "text": "Please read the current technical debt from the knowledge graph (or use the list_tech_debt tool) and provide a prioritization plan for addressing it." + } + } + ] + })) + } +} + pub struct MemoryHandler { pub state: Arc, pub tools: std::collections::HashMap>, @@ -42,9 +152,12 @@ pub struct MemoryHandler { impl MemoryHandler { pub fn new(state: Arc) -> Self { - let mut tools: std::collections::HashMap> = std::collections::HashMap::new(); - let mut resources: std::collections::HashMap> = std::collections::HashMap::new(); - let mut prompts: std::collections::HashMap> = std::collections::HashMap::new(); + let mut tools: std::collections::HashMap> = + std::collections::HashMap::new(); + let mut resources: std::collections::HashMap> = + std::collections::HashMap::new(); + let mut prompts: std::collections::HashMap> = + std::collections::HashMap::new(); macro_rules! register_resource { ($handler:ident) => { @@ -60,115 +173,71 @@ impl MemoryHandler { }; } - struct GraphEntitiesResource; - #[async_trait] - impl McpResource for GraphEntitiesResource { - fn uri(&self) -> &'static str { "memory://graph/entities" } - fn name(&self) -> &'static str { "Graph Entities" } - fn description(&self) -> Option<&'static str> { Some("All nodes and entities currently stored in the knowledge graph") } - async fn read(&self, state: Arc) -> Result { - let state_clone = Arc::clone(&state); - tokio::task::spawn_blocking(move || { - let graph = state_clone.graph.cache.read().unwrap(); - let data: Vec<_> = graph.entities.values().collect(); - Ok(serde_json::to_string_pretty(&data).unwrap_or_default()) - }).await.unwrap() - } - } - - struct GraphRelationsResource; - #[async_trait] - impl McpResource for GraphRelationsResource { - fn uri(&self) -> &'static str { "memory://graph/relations" } - fn name(&self) -> &'static str { "Graph Relations" } - fn description(&self) -> Option<&'static str> { Some("All relationships between entities currently stored in the knowledge graph") } - async fn read(&self, state: Arc) -> Result { - let state_clone = Arc::clone(&state); - tokio::task::spawn_blocking(move || { - let graph = state_clone.graph.cache.read().unwrap(); - let data = &graph.relations; - Ok(serde_json::to_string_pretty(&data).unwrap_or_default()) - }).await.unwrap() - } - } - - struct TasksActiveResource; - #[async_trait] - impl McpResource for TasksActiveResource { - fn uri(&self) -> &'static str { "memory://tasks/active" } - fn name(&self) -> &'static str { "Active Tasks" } - fn description(&self) -> Option<&'static str> { Some("List of currently active tasks") } - async fn read(&self, state: Arc) -> Result { - let state_clone = Arc::clone(&state); - tokio::task::spawn_blocking(move || { - let tasks = state_clone.tasks.cache.read().unwrap(); - let data: Vec<_> = tasks.iter().filter(|t| t.status != "completed" && t.status != "done").collect(); - Ok(serde_json::to_string_pretty(&data).unwrap_or_default()) - }).await.unwrap() - } - } - - struct AnalyzeTechDebtPrompt; - #[async_trait] - impl McpPrompt for AnalyzeTechDebtPrompt { - fn name(&self) -> &'static str { "analyze_tech_debt" } - fn description(&self) -> Option<&'static str> { Some("Analyze the project's current technical debt") } - async fn get(&self, _args: Value, _state: Arc) -> Result { - Ok(serde_json::json!({ - "messages": [ - { - "role": "user", - "content": { - "type": "text", - "text": "Please read the current technical debt from the knowledge graph (or use the list_tech_debt tool) and provide a prioritization plan for addressing it." - } - } - ] - })) - } - } - register_resource!(GraphEntitiesResource); register_resource!(GraphRelationsResource); register_resource!(TasksActiveResource); - + register_prompt!(AnalyzeTechDebtPrompt); struct PinnedFilesResource; #[async_trait] impl McpResource for PinnedFilesResource { - fn uri(&self) -> &'static str { "memory://pinned_files" } - fn name(&self) -> &'static str { "Pinned Files" } - fn description(&self) -> Option<&'static str> { Some("Currently pinned files in the workspace") } + fn uri(&self) -> &'static str { + "memory://pinned_files" + } + fn name(&self) -> &'static str { + "Pinned Files" + } + fn description(&self) -> Option<&'static str> { + Some("Currently pinned files in the workspace") + } async fn read(&self, state: Arc) -> Result { let state_clone = Arc::clone(&state); tokio::task::spawn_blocking(move || { let items = state_clone.pinned_files.cache.read().unwrap(); Ok(serde_json::to_string_pretty(&*items).unwrap_or_default()) - }).await.unwrap() + }) + .await + .unwrap() } } struct MilestonesResource; #[async_trait] impl McpResource for MilestonesResource { - fn uri(&self) -> &'static str { "memory://milestones" } - fn name(&self) -> &'static str { "Milestones" } - fn description(&self) -> Option<&'static str> { Some("Project milestones and their status") } + fn uri(&self) -> &'static str { + "memory://milestones" + } + fn name(&self) -> &'static str { + "Milestones" + } + fn description(&self) -> Option<&'static str> { + Some("Project milestones and their status") + } async fn read(&self, state: Arc) -> Result { let state_clone = Arc::clone(&state); tokio::task::spawn_blocking(move || { let items = state_clone.milestones.cache.read().unwrap(); Ok(serde_json::to_string_pretty(&*items).unwrap_or_default()) - }).await.unwrap() + }) + .await + .unwrap() } } struct HandoffRoutinePrompt; #[async_trait] impl McpPrompt for HandoffRoutinePrompt { - fn name(&self) -> &'static str { "handoff_routine" } - fn description(&self) -> Option<&'static str> { Some("Initiate the end-of-session handoff and standup report generation") } - async fn get(&self, _args: Value, _state: Arc) -> Result { + fn name(&self) -> &'static str { + "handoff_routine" + } + fn description(&self) -> Option<&'static str> { + Some("Initiate the end-of-session handoff and standup report generation") + } + async fn get( + &self, + _args: Value, + _state: Arc, + ) -> Result { Ok(serde_json::json!({ "messages": [ { @@ -187,7 +256,6 @@ impl MemoryHandler { register_resource!(MilestonesResource); register_prompt!(HandoffRoutinePrompt); - macro_rules! register { ($module:ident::$handler:ident) => { let h = crate::handlers::$module::$handler; @@ -265,7 +333,12 @@ impl MemoryHandler { register!(workspaces::GetPrChecklistHandler); register!(workspaces::ClearPrChecklistHandler); - Self { state, tools, resources, prompts } + Self { + state, + tools, + resources, + prompts, + } } pub async fn handle_request(&self, req: serde_json::Value) -> Option { @@ -328,24 +401,28 @@ impl MemoryHandler { )) } "resources/list" => { - let resources: Vec<_> = self.resources.values().map(|r| { - let mut obj = serde_json::json!({ - "uri": r.uri(), - "name": r.name(), - }); - if let Some(desc) = r.description() { - obj["description"] = serde_json::json!(desc); - } - if let Some(mime) = r.mime_type() { - obj["mimeType"] = serde_json::json!(mime); - } - obj - }).collect(); - + let resources: Vec<_> = self + .resources + .values() + .map(|r| { + let mut obj = serde_json::json!({ + "uri": r.uri(), + "name": r.name(), + }); + if let Some(desc) = r.description() { + obj["description"] = serde_json::json!(desc); + } + if let Some(mime) = r.mime_type() { + obj["mimeType"] = serde_json::json!(mime); + } + obj + }) + .collect(); + let payload = serde_json::json!({ "resources": resources }); Some(crate::mcp::success(id, payload)) } - + "resources/templates/list" => { let payload = serde_json::json!({ "resourceTemplates": [] @@ -355,7 +432,7 @@ 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(""); - + if let Some(resource) = self.resources.get(uri) { match resource.read(Arc::clone(&self.state)).await { Ok(text) => { @@ -368,44 +445,51 @@ impl MemoryHandler { }); Some(crate::mcp::success(id, payload)) } - Err(e) => Some(crate::mcp::error(id, -32603, &e)) + Err(e) => Some(crate::mcp::error(id, -32603, &e)), } } else { Some(crate::mcp::error(id, -32602, "Resource not found")) } } - + "prompts/list" => { - let prompts: Vec<_> = self.prompts.values().map(|p| { - let mut obj = serde_json::json!({ - "name": p.name(), - "arguments": p.arguments(), - }); - if let Some(desc) = p.description() { - obj["description"] = serde_json::json!(desc); - } - obj - }).collect(); - + let prompts: Vec<_> = self + .prompts + .values() + .map(|p| { + let mut obj = serde_json::json!({ + "name": p.name(), + "arguments": p.arguments(), + }); + if let Some(desc) = p.description() { + obj["description"] = serde_json::json!(desc); + } + obj + }) + .collect(); + let payload = serde_json::json!({ "prompts": prompts }); Some(crate::mcp::success(id, payload)) } - + "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(""); - let args = params.get("arguments").cloned().unwrap_or_else(|| serde_json::json!({})); - + let args = params + .get("arguments") + .cloned() + .unwrap_or_else(|| serde_json::json!({})); + if let Some(prompt) = self.prompts.get(name) { match prompt.get(args, Arc::clone(&self.state)).await { Ok(messages) => Some(crate::mcp::success(id, messages)), - Err(e) => Some(crate::mcp::error(id, -32603, &e)) + Err(e) => Some(crate::mcp::error(id, -32603, &e)), } } else { Some(crate::mcp::error(id, -32602, "Prompt not found")) } } - + "tools/call" => { let params = req.get("params").unwrap_or(&serde_json::Value::Null); let name = params.get("name").and_then(|n| n.as_str()).unwrap_or(""); @@ -488,6 +572,64 @@ mod tests { assert!(res_list["result"]["tools"].as_array().unwrap().len() > 10); } + #[tokio::test] + async fn test_resources_and_prompts_endpoints() { + let dir = tempdir().unwrap(); + let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap())); + let handler = MemoryHandler::new(state); + + // Test resources/list + let req_list_res = json!({ + "jsonrpc": "2.0", + "id": 10, + "method": "resources/list", + "params": {} + }); + let res_list = handler.handle_request(req_list_res).await.unwrap(); + let resources_arr = res_list["result"]["resources"].as_array().unwrap(); + assert!(resources_arr.iter().any(|r| r["uri"] == "memory://tasks/active")); + assert!(resources_arr.iter().any(|r| r["uri"] == "memory://pinned_files")); + + // Test resources/read + let req_read_res = json!({ + "jsonrpc": "2.0", + "id": 11, + "method": "resources/read", + "params": { + "uri": "memory://tasks/active" + } + }); + let res_read = handler.handle_request(req_read_res).await.unwrap(); + assert_eq!(res_read["result"]["contents"][0]["uri"], "memory://tasks/active"); + assert!(res_read["result"]["contents"][0]["text"].as_str().unwrap().contains("[]")); // Empty tasks + + // Test prompts/list + let req_list_prompts = json!({ + "jsonrpc": "2.0", + "id": 12, + "method": "prompts/list", + "params": {} + }); + let res_prompts = handler.handle_request(req_list_prompts).await.unwrap(); + let prompts_arr = res_prompts["result"]["prompts"].as_array().unwrap(); + assert!(prompts_arr.iter().any(|p| p["name"] == "handoff_routine")); + + // Test prompts/get + let req_get_prompt = json!({ + "jsonrpc": "2.0", + "id": 13, + "method": "prompts/get", + "params": { + "name": "handoff_routine", + "arguments": {} + } + }); + let res_get = handler.handle_request(req_get_prompt).await.unwrap(); + let messages = res_get["result"]["messages"].as_array().unwrap(); + assert_eq!(messages[0]["role"], "user"); + assert!(messages[0]["content"]["text"].as_str().unwrap().contains("standup report")); + } + #[tokio::test] async fn test_tool_call_success_and_error_responses() { let dir = tempdir().unwrap(); @@ -529,7 +671,12 @@ mod tests { assert_eq!(res_fail["id"], 3); // Semantic failures must explicitly return isError: true inside the result to halt the LLM assert_eq!(res_fail["result"]["isError"], true); - assert!(res_fail["result"]["content"][0]["text"].as_str().unwrap().contains("not found")); + assert!( + res_fail["result"]["content"][0]["text"] + .as_str() + .unwrap() + .contains("not found") + ); // 3. Test unknown JSON-RPC method returns JSON-RPC protocol error let req_unknown = json!({