Files
mcp-memory/server/src/router.rs
T

411 lines
17 KiB
Rust

use crate::state::MemoryState;
use async_trait::async_trait;
use serde_json::Value;
use std::sync::Arc;
#[async_trait]
pub trait McpTool: Send + Sync {
/// The unique name of the tool
fn name(&self) -> &'static str;
/// The JSON schema for the tool
fn schema(&self) -> Value;
/// Execute the tool with the given arguments
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String>;
}
pub struct MemoryHandler {
pub state: Arc<MemoryState>,
pub tools: std::collections::HashMap<String, Box<dyn McpTool>>,
}
impl MemoryHandler {
pub fn new(state: Arc<MemoryState>) -> Self {
let mut tools: std::collections::HashMap<String, Box<dyn McpTool>> =
std::collections::HashMap::new();
macro_rules! register {
($module:ident::$handler:ident) => {
let h = crate::handlers::$module::$handler;
tools.insert(h.name().to_string(), Box::new(h));
};
}
register!(graph::QueryGraphPathHandler);
register!(graph::CreateEntitiesHandler);
register!(graph::CreateRelationsHandler);
register!(graph::AddObservationsHandler);
register!(graph::DeleteEntitiesHandler);
register!(graph::DeleteObservationsHandler);
register!(graph::DeleteRelationsHandler);
register!(graph::ReadGraphHandler);
register!(graph::SearchNodesHandler);
register!(graph::OpenNodesHandler);
register!(graph::VisualizeGraphHandler);
register!(graph::CondenseEntityHandler);
register!(graph::MergeEntitiesHandler);
register!(graph::FindOrphansHandler);
register!(tasks::AddTaskHandler);
register!(tasks::DeleteTaskHandler);
register!(tasks::UpdateTaskStatusHandler);
register!(tasks::ListActiveTasksHandler);
register!(tasks::SetAcceptanceCriteriaHandler);
register!(tasks::VerifyAcceptanceCriteriaHandler);
register!(tasks::AddMilestoneHandler);
register!(tasks::UpdateMilestoneHandler);
register!(tasks::ListMilestonesHandler);
register!(notes::AddStickyNoteHandler);
register!(notes::ReadStickyNotesHandler);
register!(notes::DeleteStickyNoteHandler);
register!(notes::ClearStickyNotesHandler);
register!(notes::LeaveHandoffMemoHandler);
register!(notes::ReadHandoffMemosHandler);
register!(notes::ClearHandoffMemosHandler);
register!(notes::AddSessionSummaryHandler);
register!(notes::GenerateStandupReportHandler);
register!(meta::LogDecisionHandler);
register!(meta::QueryDecisionsHandler);
register!(meta::LogErrorFixHandler);
register!(meta::SearchErrorFixesHandler);
register!(meta::LogCodeChangeHandler);
register!(meta::QueryRecentChangesHandler);
register!(meta::LearnPreferenceHandler);
register!(meta::ReadPreferencesHandler);
register!(meta::LogTechDebtHandler);
register!(meta::ResolveTechDebtHandler);
register!(meta::ListTechDebtHandler);
register!(meta::OmniSearchHandler);
register!(meta::GetProjectHealthHandler);
register!(env::UpdateEnvFingerprintHandler);
register!(env::ReadEnvFingerprintHandler);
register!(env::LogEnvRequirementHandler);
register!(env::RegisterEnvironmentHandler);
register!(env::GetEnvironmentDetailsHandler);
register!(workspaces::PinFileHandler);
register!(workspaces::UnpinFileHandler);
register!(workspaces::ListPinnedFilesHandler);
register!(workspaces::StoreSnippetHandler);
register!(workspaces::SearchSnippetsHandler);
register!(workspaces::DeleteSnippetHandler);
register!(workspaces::SaveContextWorkspaceHandler);
register!(workspaces::LoadContextWorkspaceHandler);
register!(workspaces::ListContextWorkspacesHandler);
register!(workspaces::AddPrChecklistItemHandler);
register!(workspaces::GetPrChecklistHandler);
register!(workspaces::ClearPrChecklistHandler);
Self { state, tools }
}
pub async fn handle_request(&self, req: serde_json::Value) -> Option<serde_json::Value> {
let id = req.get("id").cloned().unwrap_or(serde_json::Value::Null);
let id_clone = id.clone();
let method = req.get("method").and_then(|m| m.as_str()).unwrap_or("");
match method {
"server/discover" => {
let payload = serde_json::json!({
"resultType": "complete",
"ttlMs": 0,
"cacheScope": "public",
"supportedVersions": ["2026-07-28", "2025-11-25", "2025-06-18", "2025-03-26", "2024-11-05"],
"capabilities": {
"tools": serde_json::json!({}),
"resources": serde_json::json!({}),
"prompts": serde_json::json!({})
},
"_meta": {
"io.modelcontextprotocol/serverInfo": {
"name": "gemini-mcp-memory",
"version": "3.0.0"
}
}
});
Some(crate::mcp::success(id, payload))
}
"initialize" => {
let init = rmcp::model::InitializeResult::new(
rmcp::model::ServerCapabilities::builder()
.enable_tools()
.enable_resources()
.enable_prompts()
.build(),
)
.with_server_info(rmcp::model::Implementation::new(
"gemini-mcp-memory",
"3.0.0",
))
.with_instructions(include_str!("instructions.md"));
Some(crate::mcp::success(
id,
serde_json::to_value(&init).unwrap_or_default(),
))
}
"notifications/initialized" => None,
"tools/list" => {
let mut tools: Vec<serde_json::Value> =
self.tools.values().map(|t| t.schema()).collect();
tools.sort_by_key(|t| {
t.get("name")
.and_then(|n| n.as_str())
.unwrap_or("")
.to_string()
});
Some(crate::mcp::success(
id,
serde_json::json!({ "tools": tools }),
))
}
"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"
}
]
});
Some(crate::mcp::success(id, payload))
}
"resources/templates/list" => {
let payload = serde_json::json!({
"resourceTemplates": []
});
Some(crate::mcp::success(id, payload))
}
"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": []
}
]
});
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))
} 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("");
let args = params
.get("arguments")
.cloned()
.unwrap_or(serde_json::Value::Object(Default::default()));
self.state
.broadcast_activity(&format!("Agent executed tool: {}", name));
let result: Result<String, String> = if let Some(tool) = self.tools.get(name) {
tool.execute(args, self.state.clone()).await
} else {
Err(format!("Unknown tool: {}", name))
};
match result {
Ok(text) => {
let payload = serde_json::json!({
"content": [{"type": "text", "text": text}],
"isError": false
});
Some(crate::mcp::success(id_clone, payload))
}
Err(e) => {
tracing::error!("Tool {} failed: {}", name, e);
let payload = serde_json::json!({
"content": [{"type": "text", "text": e}],
"isError": true
});
Some(crate::mcp::success(id_clone, payload))
}
}
}
m if m.starts_with("notifications/") => None,
"ping" => Some(crate::mcp::success(id, serde_json::json!({}))),
_ => {
if id.is_null() {
None
} else {
Some(crate::mcp::error(
id,
-32601,
&format!("Method {} not found", method),
))
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use tempfile::tempdir;
#[tokio::test]
async fn test_memory_handler_tools_registration() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let handler = MemoryHandler::new(state);
// Assert some known tools are registered
assert!(handler.tools.contains_key("create_entities"));
assert!(handler.tools.contains_key("add_task"));
// Ensure we can fetch list of tools
let list_tools_req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list",
"params": {}
});
let res_list = handler.handle_request(list_tools_req).await.unwrap();
assert_eq!(res_list["jsonrpc"], "2.0");
assert_eq!(res_list["id"], 1);
assert!(res_list["result"]["tools"].as_array().unwrap().len() > 10);
}
#[tokio::test]
async fn test_tool_call_success_and_error_responses() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let handler = MemoryHandler::new(state);
// 1. Test successful tool call (e.g. read_graph)
let req_success = json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": {
"name": "read_graph",
"arguments": {}
}
});
let res_success = handler.handle_request(req_success).await.unwrap();
assert_eq!(res_success["jsonrpc"], "2.0");
assert_eq!(res_success["id"], 2);
// A successful tool call should return a result with isError: false
assert_eq!(res_success["result"]["isError"], false);
assert!(res_success["result"]["content"].as_array().is_some());
// 2. Test tool call semantic failure (e.g. updating a task that does not exist)
let req_fail = json!({
"jsonrpc": "2.0",
"id": 3,
"method": "tools/call",
"params": {
"name": "update_task_status",
"arguments": {
"id": "nonexistent_task_123",
"status": "in_progress"
}
}
});
let res_fail = handler.handle_request(req_fail).await.unwrap();
assert_eq!(res_fail["jsonrpc"], "2.0");
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"));
// 3. Test unknown JSON-RPC method returns JSON-RPC protocol error
let req_unknown = json!({
"jsonrpc": "2.0",
"id": 4,
"method": "unknown_method_xyz"
});
let res_unknown = handler.handle_request(req_unknown).await.unwrap();
assert!(res_unknown.get("error").is_some());
assert_eq!(res_unknown["error"]["code"], -32601);
}
}