refactor: implement McpResource and McpPrompt traits to handle MCP endpoints generically

This commit is contained in:
Riz Ashraf committed 2026-09-27 21:47:52 +01:00
1 parent 41461a41ef
commit 45577a99c4
1 file changed
+168 -90
+168 -90
View File
@@ -15,15 +15,124 @@ pub trait McpTool: Send + Sync {
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String>; async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String>;
} }
#[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<MemoryState>) -> Result<String, String>;
}
#[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<MemoryState>) -> Result<serde_json::Value, String>;
}
pub struct MemoryHandler { pub struct MemoryHandler {
pub state: Arc<MemoryState>, pub state: Arc<MemoryState>,
pub tools: std::collections::HashMap<String, Box<dyn McpTool>>, pub tools: std::collections::HashMap<String, Box<dyn McpTool>>,
pub resources: std::collections::HashMap<String, Box<dyn McpResource>>,
pub prompts: std::collections::HashMap<String, Box<dyn McpPrompt>>,
} }
impl MemoryHandler { impl MemoryHandler {
pub fn new(state: Arc<MemoryState>) -> Self { pub fn new(state: Arc<MemoryState>) -> Self {
let mut tools: std::collections::HashMap<String, Box<dyn McpTool>> = let mut tools: std::collections::HashMap<String, Box<dyn McpTool>> = std::collections::HashMap::new();
std::collections::HashMap::new(); let mut resources: std::collections::HashMap<String, Box<dyn McpResource>> = std::collections::HashMap::new();
let mut prompts: std::collections::HashMap<String, Box<dyn McpPrompt>> = 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<MemoryState>) -> Result<String, String> {
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<MemoryState>) -> Result<String, String> {
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<MemoryState>) -> Result<String, String> {
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<MemoryState>) -> Result<serde_json::Value, String> {
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 { macro_rules! register {
($module:ident::$handler:ident) => { ($module:ident::$handler:ident) => {
@@ -102,7 +211,7 @@ impl MemoryHandler {
register!(workspaces::GetPrChecklistHandler); register!(workspaces::GetPrChecklistHandler);
register!(workspaces::ClearPrChecklistHandler); register!(workspaces::ClearPrChecklistHandler);
Self { state, tools } Self { state, tools, resources, prompts }
} }
pub async fn handle_request(&self, req: serde_json::Value) -> Option<serde_json::Value> { pub async fn handle_request(&self, req: serde_json::Value) -> Option<serde_json::Value> {
@@ -165,30 +274,24 @@ impl MemoryHandler {
)) ))
} }
"resources/list" => { "resources/list" => {
let payload = serde_json::json!({ let resources: Vec<_> = self.resources.values().map(|r| {
"resources": [ let mut obj = serde_json::json!({
{ "uri": r.uri(),
"uri": "memory://graph/entities", "name": r.name(),
"name": "Graph Entities", });
"mimeType": "application/json", if let Some(desc) = r.description() {
"description": "All nodes and entities currently stored in the knowledge graph" obj["description"] = serde_json::json!(desc);
}, }
{ if let Some(mime) = r.mime_type() {
"uri": "memory://graph/relations", obj["mimeType"] = serde_json::json!(mime);
"name": "Graph Relations", }
"mimeType": "application/json", obj
"description": "All edge relationships between entities in the knowledge graph" }).collect();
},
{ let payload = serde_json::json!({ "resources": resources });
"uri": "memory://tasks/active",
"name": "Active Tasks",
"mimeType": "application/json",
"description": "All currently active or uncompleted tracking tasks"
}
]
});
Some(crate::mcp::success(id, payload)) Some(crate::mcp::success(id, payload))
} }
"resources/templates/list" => { "resources/templates/list" => {
let payload = serde_json::json!({ let payload = serde_json::json!({
"resourceTemplates": [] "resourceTemplates": []
@@ -197,83 +300,58 @@ impl MemoryHandler {
} }
"resources/read" => { "resources/read" => {
let params = req.get("params").unwrap_or(&serde_json::Value::Null); let params = req.get("params").unwrap_or(&serde_json::Value::Null);
let uri = params let uri = params.get("uri").and_then(|u| u.as_str()).unwrap_or("");
.get("uri")
.and_then(|u| u.as_str())
.unwrap_or("")
.to_string();
let state_clone = Arc::clone(&self.state); if let Some(resource) = self.resources.get(uri) {
let uri_clone = uri.clone(); match resource.read(Arc::clone(&self.state)).await {
let text = match tokio::task::spawn_blocking(move || match uri_clone.as_str() { Ok(text) => {
"memory://graph/entities" => { let payload = serde_json::json!({
let graph = state_clone.graph.cache.read().unwrap(); "contents": [{
let data: Vec<_> = graph.entities.values().collect(); "uri": uri,
Some(serde_json::to_string_pretty(&data).unwrap_or_default()) "mimeType": resource.mime_type().unwrap_or("application/json"),
} "text": text
"memory://graph/relations" => { }]
let graph = state_clone.graph.cache.read().unwrap(); });
let data = &graph.relations; Some(crate::mcp::success(id, payload))
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": []
} }
] 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)) Some(crate::mcp::success(id, payload))
} }
"prompts/get" => { "prompts/get" => {
let params = req.get("params").unwrap_or(&serde_json::Value::Null); 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 name = params.get("name").and_then(|n| n.as_str()).unwrap_or("");
let args = params.get("arguments").cloned().unwrap_or_else(|| serde_json::json!({}));
if name == "analyze_tech_debt" { if let Some(prompt) = self.prompts.get(name) {
let payload = serde_json::json!({ match prompt.get(args, Arc::clone(&self.state)).await {
"messages": [ Ok(messages) => Some(crate::mcp::success(id, messages)),
{ Err(e) => Some(crate::mcp::error(id, -32603, &e))
"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))
} else { } else {
Some(crate::mcp::error(id, -32602, "Prompt not found")) Some(crate::mcp::error(id, -32602, "Prompt not found"))
} }
} }
"tools/call" => { "tools/call" => {
let params = req.get("params").unwrap_or(&serde_json::Value::Null); 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 name = params.get("name").and_then(|n| n.as_str()).unwrap_or("");