diff --git a/server/src/router.rs b/server/src/router.rs index 6e533e8..14b97ab 100644 --- a/server/src/router.rs +++ b/server/src/router.rs @@ -15,15 +15,124 @@ 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") } + 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!([]) } + async fn get(&self, args: Value, state: Arc) -> Result; +} + pub struct MemoryHandler { pub state: Arc, pub tools: std::collections::HashMap>, + pub resources: std::collections::HashMap>, + pub prompts: std::collections::HashMap>, } impl MemoryHandler { pub fn new(state: Arc) -> Self { - let mut tools: 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) => { + let h = $handler; + resources.insert(h.uri().to_string(), Box::new(h)); + }; + } + + macro_rules! register_prompt { + ($handler:ident) => { + let h = $handler; + prompts.insert(h.name().to_string(), Box::new(h)); + }; + } + + 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); macro_rules! register { ($module:ident::$handler:ident) => { @@ -102,7 +211,7 @@ impl MemoryHandler { register!(workspaces::GetPrChecklistHandler); register!(workspaces::ClearPrChecklistHandler); - Self { state, tools } + Self { state, tools, resources, prompts } } pub async fn handle_request(&self, req: serde_json::Value) -> Option { @@ -165,30 +274,24 @@ impl MemoryHandler { )) } "resources/list" => { - let payload = serde_json::json!({ - "resources": [ - { - "uri": "memory://graph/entities", - "name": "Graph Entities", - "mimeType": "application/json", - "description": "All nodes and entities currently stored in the knowledge graph" - }, - { - "uri": "memory://graph/relations", - "name": "Graph Relations", - "mimeType": "application/json", - "description": "All edge relationships between entities in the knowledge graph" - }, - { - "uri": "memory://tasks/active", - "name": "Active Tasks", - "mimeType": "application/json", - "description": "All currently active or uncompleted tracking tasks" - } - ] - }); + 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": [] @@ -197,83 +300,58 @@ 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("") - .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 = 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 = 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 = state_clone.tasks.cache.read().unwrap(); - let data: Vec<_> = tasks - .iter() - .filter(|t| t.status != "completed" && t.status != "done") - .collect(); - 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!({ - "contents": [{ - "uri": uri, - "mimeType": "application/json", - "text": text - }] - }); - Some(crate::mcp::success(id, payload)) - } - "prompts/list" => { - let payload = serde_json::json!({ - "prompts": [ - { - "name": "analyze_tech_debt", - "description": "Analyze the project's current technical debt", - "arguments": [] + 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) => { + let payload = serde_json::json!({ + "contents": [{ + "uri": uri, + "mimeType": resource.mime_type().unwrap_or("application/json"), + "text": text + }] + }); + Some(crate::mcp::success(id, payload)) } - ] - }); + 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 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(""); - - if name == "analyze_tech_debt" { - let payload = 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." - } - } - ] - }); - Some(crate::mcp::success(id, payload)) + 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)) + } } 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("");