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_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 state: Arc<MemoryState>,
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 {
pub fn new(state: Arc<MemoryState>) -> Self {
let mut tools: std::collections::HashMap<String, Box<dyn McpTool>> =
std::collections::HashMap::new();
let mut tools: std::collections::HashMap<String, Box<dyn McpTool>> = 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 {
($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<serde_json::Value> {
@@ -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 uri = params.get("uri").and_then(|u| u.as_str()).unwrap_or("");
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": []
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("");
let args = params.get("arguments").cloned().unwrap_or_else(|| serde_json::json!({}));
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))
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("");