refactor: implement McpResource and McpPrompt traits to handle MCP endpoints generically
This commit is contained in:
1 parent
41461a41ef
commit
45577a99c4
1 file changed
+170
-92
+170
-92
@@ -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 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("");
|
||||
|
||||
Reference in new issue
Block a user