Refactor: Migrate unwrap calls to AppError in MCP handlers

This commit is contained in:
Riz Ashraf committed 2026-09-30 21:02:50 +01:00
1 parent 4e1a633dbd
commit 0e866f2465
12 files changed
+276 -326

No files matched your search

+43 -49
View File
@@ -12,7 +12,7 @@ pub trait McpTool: Send + Sync {
fn schema(&self) -> Value;
/// Execute the tool with the given arguments
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String>;
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<String>;
}
#[async_trait]
@@ -25,7 +25,7 @@ pub trait McpResource: Send + Sync {
fn mime_type(&self) -> Option<&'static str> {
Some("application/json")
}
async fn read(&self, state: Arc<MemoryState>) -> Result<String, String>;
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String>;
}
#[async_trait]
@@ -37,7 +37,7 @@ pub trait McpPrompt: Send + Sync {
fn arguments(&self) -> serde_json::Value {
serde_json::json!([])
}
async fn get(&self, args: Value, state: Arc<MemoryState>) -> Result<serde_json::Value, String>;
async fn get(&self, args: Value, state: Arc<MemoryState>) -> crate::error::Result<serde_json::Value>;
}
struct GraphEntitiesResource;
@@ -52,15 +52,14 @@ impl McpResource for GraphEntitiesResource {
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> {
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
let state_clone = Arc::clone(&state);
tokio::task::spawn_blocking(move || {
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
let graph = state_clone.graph.cache.read().unwrap();
let data: Vec<_> = graph.entities.values().collect();
serde_json::to_string_pretty(&data).map_err(|e| e.to_string())
Ok(serde_json::to_string_pretty(&data)?)
})
.await
.unwrap()
.await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?
}
}
@@ -76,15 +75,14 @@ impl McpResource for GraphRelationsResource {
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> {
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
let state_clone = Arc::clone(&state);
tokio::task::spawn_blocking(move || {
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
let graph = state_clone.graph.cache.read().unwrap();
let data = &graph.relations;
serde_json::to_string_pretty(&data).map_err(|e| e.to_string())
Ok(serde_json::to_string_pretty(&data)?)
})
.await
.unwrap()
.await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?
}
}
@@ -100,18 +98,17 @@ impl McpResource for TasksActiveResource {
fn description(&self) -> Option<&'static str> {
Some("List of currently active tasks")
}
async fn read(&self, state: Arc<MemoryState>) -> Result<String, String> {
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
let state_clone = Arc::clone(&state);
tokio::task::spawn_blocking(move || {
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
let tasks = state_clone.tasks.cache.read().unwrap();
let data: Vec<_> = tasks
.iter()
.filter(|t| t.status != "completed" && t.status != "done")
.collect();
serde_json::to_string_pretty(&data).map_err(|e| e.to_string())
Ok(serde_json::to_string_pretty(&data)?)
})
.await
.unwrap()
.await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?
}
}
@@ -128,7 +125,7 @@ impl McpPrompt for AnalyzeTechDebtPrompt {
&self,
_args: Value,
_state: Arc<MemoryState>,
) -> Result<serde_json::Value, String> {
) -> crate::error::Result<serde_json::Value> {
Ok(serde_json::json!({
"messages": [
{
@@ -190,14 +187,13 @@ impl MemoryHandler {
fn description(&self) -> Option<&'static str> {
Some("Recent terminal execution history and exit codes")
}
async fn read(&self, state: Arc<MemoryState>) -> Result<String, String> {
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
let state_clone = Arc::clone(&state);
tokio::task::spawn_blocking(move || {
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
let items = state_clone.terminal_history.cache.read().unwrap();
serde_json::to_string_pretty(&*items).map_err(|e| e.to_string())
Ok(serde_json::to_string_pretty(&*items)?)
})
.await
.unwrap()
.await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?
}
}
struct PinnedFilesResource;
@@ -212,14 +208,13 @@ impl MemoryHandler {
fn description(&self) -> Option<&'static str> {
Some("Currently pinned files in the workspace")
}
async fn read(&self, state: Arc<MemoryState>) -> Result<String, String> {
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
let state_clone = Arc::clone(&state);
tokio::task::spawn_blocking(move || {
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
let items = state_clone.pinned_files.cache.read().unwrap();
serde_json::to_string_pretty(&*items).map_err(|e| e.to_string())
Ok(serde_json::to_string_pretty(&*items)?)
})
.await
.unwrap()
.await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?
}
}
@@ -235,14 +230,13 @@ impl MemoryHandler {
fn description(&self) -> Option<&'static str> {
Some("Project milestones and their status")
}
async fn read(&self, state: Arc<MemoryState>) -> Result<String, String> {
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
let state_clone = Arc::clone(&state);
tokio::task::spawn_blocking(move || {
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
let items = state_clone.milestones.cache.read().unwrap();
serde_json::to_string_pretty(&*items).map_err(|e| e.to_string())
Ok(serde_json::to_string_pretty(&*items)?)
})
.await
.unwrap()
.await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?
}
}
@@ -259,7 +253,7 @@ impl MemoryHandler {
&self,
_args: Value,
_state: Arc<MemoryState>,
) -> Result<serde_json::Value, String> {
) -> crate::error::Result<serde_json::Value> {
Ok(serde_json::json!({
"messages": [
{
@@ -290,7 +284,7 @@ impl MemoryHandler {
&self,
_args: Value,
_state: Arc<MemoryState>,
) -> Result<serde_json::Value, String> {
) -> crate::error::Result<serde_json::Value> {
Ok(serde_json::json!({"
messages": [
{
@@ -506,7 +500,7 @@ impl MemoryHandler {
});
Some(crate::mcp::success(id, payload))
}
Err(e) => Some(crate::mcp::error(id, -32603, &e)),
Err(e) => Some(crate::mcp::error(id, -32603, &e.to_string())),
}
} else {
Some(crate::mcp::error(id, -32602, "Resource not found"))
@@ -544,7 +538,7 @@ impl MemoryHandler {
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)),
Err(e) => Some(crate::mcp::error(id, -32603, &e.to_string())),
}
} else {
Some(crate::mcp::error(id, -32602, "Prompt not found"))
@@ -562,10 +556,10 @@ impl MemoryHandler {
self.state
.broadcast_activity(&format!("Agent executed tool: {}", name));
let result: Result<String, String> = if let Some(tool) = self.tools.get(name) {
let result: crate::error::Result<String> = if let Some(tool) = self.tools.get(name) {
tool.execute(args, self.state.clone()).await
} else {
Err(format!("Unknown tool: {}", name))
Err(crate::error::AppError::Internal(format!("Unknown tool: {}", name)))
};
match result {
@@ -579,7 +573,7 @@ impl MemoryHandler {
Err(e) => {
tracing::error!("Tool {} failed: {}", name, e);
let payload = serde_json::json!({
"content": [{"type": "text", "text": e}],
"content": [{"type": "text", "text": e.to_string()}],
"isError": true
});
Some(crate::mcp::success(id_clone, payload))
@@ -627,7 +621,7 @@ mod tests {
"params": {}
});
let res_list = handler.handle_request(list_tools_req).await.unwrap();
let res_list = handler.handle_request(list_tools_req).await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
assert_eq!(res_list["jsonrpc"], "2.0");
assert_eq!(res_list["id"], 1);
assert!(res_list["result"]["tools"].as_array().unwrap().len() > 10);
@@ -646,7 +640,7 @@ mod tests {
"method": "resources/list",
"params": {}
});
let res_list = handler.handle_request(req_list_res).await.unwrap();
let res_list = handler.handle_request(req_list_res).await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
let resources_arr = res_list["result"]["resources"].as_array().unwrap();
assert!(
resources_arr
@@ -668,7 +662,7 @@ mod tests {
"uri": "memory://tasks/active"
}
});
let res_read = handler.handle_request(req_read_res).await.unwrap();
let res_read = handler.handle_request(req_read_res).await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
assert_eq!(
res_read["result"]["contents"][0]["uri"],
"memory://tasks/active"
@@ -687,7 +681,7 @@ mod tests {
"method": "prompts/list",
"params": {}
});
let res_prompts = handler.handle_request(req_list_prompts).await.unwrap();
let res_prompts = handler.handle_request(req_list_prompts).await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
let prompts_arr = res_prompts["result"]["prompts"].as_array().unwrap();
assert!(prompts_arr.iter().any(|p| p["name"] == "handoff_routine"));
@@ -701,7 +695,7 @@ mod tests {
"arguments": {}
}
});
let res_get = handler.handle_request(req_get_prompt).await.unwrap();
let res_get = handler.handle_request(req_get_prompt).await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
let messages = res_get["result"]["messages"].as_array().unwrap();
assert_eq!(messages[0]["role"], "user");
assert!(
@@ -728,7 +722,7 @@ mod tests {
"arguments": {}
}
});
let res_success = handler.handle_request(req_success).await.unwrap();
let res_success = handler.handle_request(req_success).await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
assert_eq!(res_success["jsonrpc"], "2.0");
assert_eq!(res_success["id"], 2);
// A successful tool call should return a result with isError: false
@@ -748,7 +742,7 @@ mod tests {
}
}
});
let res_fail = handler.handle_request(req_fail).await.unwrap();
let res_fail = handler.handle_request(req_fail).await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
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
@@ -766,7 +760,7 @@ mod tests {
"id": 4,
"method": "unknown_method_xyz"
});
let res_unknown = handler.handle_request(req_unknown).await.unwrap();
let res_unknown = handler.handle_request(req_unknown).await.map_err(|e| crate::error::AppError::Internal(e.to_string()))?;
assert!(res_unknown.get("error").is_some());
assert_eq!(res_unknown["error"]["code"], -32601);
}