768 lines
29 KiB
Rust
768 lines
29 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>) -> crate::error::Result<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>) -> crate::error::Result<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>) -> crate::error::Result<serde_json::Value>;
|
|
}
|
|
|
|
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>) -> crate::error::Result<String> {
|
|
let state_clone = Arc::clone(&state);
|
|
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
|
|
let graph = state_clone.graph.cache.read().unwrap();
|
|
let data: Vec<_> = graph.entities.values().collect();
|
|
Ok(serde_json::to_string_pretty(&data)?)
|
|
})
|
|
.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>) -> crate::error::Result<String> {
|
|
let state_clone = Arc::clone(&state);
|
|
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
|
|
let graph = state_clone.graph.cache.read().unwrap();
|
|
let data = &graph.relations;
|
|
Ok(serde_json::to_string_pretty(&data)?)
|
|
})
|
|
.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>) -> crate::error::Result<String> {
|
|
let state_clone = Arc::clone(&state);
|
|
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
|
|
let tasks = state_clone.project.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)?)
|
|
})
|
|
.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>,
|
|
) -> crate::error::Result<serde_json::Value> {
|
|
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."
|
|
}
|
|
}
|
|
]
|
|
}))
|
|
}
|
|
}
|
|
|
|
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 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));
|
|
};
|
|
}
|
|
|
|
register_resource!(GraphEntitiesResource);
|
|
register_resource!(GraphRelationsResource);
|
|
register_resource!(TasksActiveResource);
|
|
|
|
register_prompt!(AnalyzeTechDebtPrompt);
|
|
struct TerminalHistoryResource;
|
|
#[async_trait]
|
|
impl McpResource for TerminalHistoryResource {
|
|
fn uri(&self) -> &'static str {
|
|
"memory://terminal/recent"
|
|
}
|
|
fn name(&self) -> &'static str {
|
|
"Terminal History"
|
|
}
|
|
fn description(&self) -> Option<&'static str> {
|
|
Some("Recent terminal execution history and exit codes")
|
|
}
|
|
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
|
|
let state_clone = Arc::clone(&state);
|
|
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
|
|
let items = state_clone.telemetry.terminal_history.cache.read().unwrap();
|
|
Ok(serde_json::to_string_pretty(&*items)?)
|
|
})
|
|
.await.unwrap()
|
|
}
|
|
}
|
|
struct PinnedFilesResource;
|
|
#[async_trait]
|
|
impl McpResource for PinnedFilesResource {
|
|
fn uri(&self) -> &'static str {
|
|
"memory://pinned_files"
|
|
}
|
|
fn name(&self) -> &'static str {
|
|
"Pinned Files"
|
|
}
|
|
fn description(&self) -> Option<&'static str> {
|
|
Some("Currently pinned files in the workspace")
|
|
}
|
|
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
|
|
let state_clone = Arc::clone(&state);
|
|
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
|
|
let items = state_clone.project.pinned_files.cache.read().unwrap();
|
|
Ok(serde_json::to_string_pretty(&*items)?)
|
|
})
|
|
.await.unwrap()
|
|
}
|
|
}
|
|
|
|
struct MilestonesResource;
|
|
#[async_trait]
|
|
impl McpResource for MilestonesResource {
|
|
fn uri(&self) -> &'static str {
|
|
"memory://milestones"
|
|
}
|
|
fn name(&self) -> &'static str {
|
|
"Milestones"
|
|
}
|
|
fn description(&self) -> Option<&'static str> {
|
|
Some("Project milestones and their status")
|
|
}
|
|
async fn read(&self, state: Arc<MemoryState>) -> crate::error::Result<String> {
|
|
let state_clone = Arc::clone(&state);
|
|
tokio::task::spawn_blocking(move || -> crate::error::Result<String> {
|
|
let items = state_clone.project.milestones.cache.read().unwrap();
|
|
Ok(serde_json::to_string_pretty(&*items)?)
|
|
})
|
|
.await.unwrap()
|
|
}
|
|
}
|
|
|
|
struct HandoffRoutinePrompt;
|
|
#[async_trait]
|
|
impl McpPrompt for HandoffRoutinePrompt {
|
|
fn name(&self) -> &'static str {
|
|
"handoff_routine"
|
|
}
|
|
fn description(&self) -> Option<&'static str> {
|
|
Some("Initiate the end-of-session handoff and standup report generation")
|
|
}
|
|
async fn get(
|
|
&self,
|
|
_args: Value,
|
|
_state: Arc<MemoryState>,
|
|
) -> crate::error::Result<serde_json::Value> {
|
|
Ok(serde_json::json!({
|
|
"messages": [
|
|
{
|
|
"role": "user",
|
|
"content": {
|
|
"type": "text",
|
|
"text": "I am logging off. Please invoke the DevOpsSRE subagent to generate a standup report and leave a handoff memo for the next session. Make sure to check active tasks and recent code changes."
|
|
}
|
|
}
|
|
]
|
|
}))
|
|
}
|
|
}
|
|
|
|
register_resource!(TerminalHistoryResource);
|
|
register_resource!(PinnedFilesResource);
|
|
register_resource!(MilestonesResource);
|
|
struct ArchiveRoutinePrompt;
|
|
#[async_trait]
|
|
impl McpPrompt for ArchiveRoutinePrompt {
|
|
fn name(&self) -> &'static str {
|
|
"archive_routine"
|
|
}
|
|
fn description(&self) -> Option<&'static str> {
|
|
Some("Compress old session summaries into a milestone retrospective")
|
|
}
|
|
async fn get(
|
|
&self,
|
|
_args: Value,
|
|
_state: Arc<MemoryState>,
|
|
) -> crate::error::Result<serde_json::Value> {
|
|
Ok(serde_json::json!({"
|
|
messages": [
|
|
{
|
|
"role": "user",
|
|
"content": {
|
|
"type": "text",
|
|
"text": "Please read the old session summaries, synthesize them into a dense Milestone Retrospective entity, and then use the appropriate tools to delete the old session summaries."
|
|
}
|
|
}
|
|
]
|
|
}))
|
|
}
|
|
}
|
|
register_prompt!(HandoffRoutinePrompt);
|
|
register_prompt!(ArchiveRoutinePrompt);
|
|
|
|
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::DeleteDecisionHandler);
|
|
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::DeleteContextWorkspaceHandler);
|
|
register!(workspaces::AddPrChecklistItemHandler);
|
|
register!(workspaces::GetPrChecklistHandler);
|
|
register!(workspaces::ClearPrChecklistHandler);
|
|
register!(vision::ReadClipboardHandler);
|
|
register!(vision::WriteClipboardHandler);
|
|
register!(git::GetActiveWorktreeContextHandler);
|
|
register!(logs::WatchProcessLogsHandler);
|
|
register!(logs::GetRecentLogsHandler);
|
|
register!(ast::ReadFileSkeletonHandler);
|
|
register!(vision::ToggleClipboardWatchModeHandler);
|
|
register!(ast::ReplaceAstNodeHandler);
|
|
register!(workspaces::ReadDirectoryArchitectureHandler);
|
|
register!(workspaces::SemanticCodeSearchHandler);
|
|
|
|
Self {
|
|
state,
|
|
tools,
|
|
resources,
|
|
prompts,
|
|
}
|
|
}
|
|
|
|
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 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": []
|
|
});
|
|
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("");
|
|
|
|
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.to_string())),
|
|
}
|
|
} 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 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.to_string())),
|
|
}
|
|
} 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: crate::error::Result<String> = if let Some(tool) = self.tools.get(name) {
|
|
tool.execute(args, self.state.clone()).await
|
|
} else {
|
|
Err(crate::error::AppError::Internal(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.to_string()}],
|
|
"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_resources_and_prompts_endpoints() {
|
|
let dir = tempdir().unwrap();
|
|
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
|
|
let handler = MemoryHandler::new(state);
|
|
|
|
// Test resources/list
|
|
let req_list_res = json!({
|
|
"jsonrpc": "2.0",
|
|
"id": 10,
|
|
"method": "resources/list",
|
|
"params": {}
|
|
});
|
|
let res_list = handler.handle_request(req_list_res).await.unwrap();
|
|
let resources_arr = res_list["result"]["resources"].as_array().unwrap();
|
|
assert!(
|
|
resources_arr
|
|
.iter()
|
|
.any(|r| r["uri"] == "memory://tasks/active")
|
|
);
|
|
assert!(
|
|
resources_arr
|
|
.iter()
|
|
.any(|r| r["uri"] == "memory://pinned_files")
|
|
);
|
|
|
|
// Test resources/read
|
|
let req_read_res = json!({
|
|
"jsonrpc": "2.0",
|
|
"id": 11,
|
|
"method": "resources/read",
|
|
"params": {
|
|
"uri": "memory://tasks/active"
|
|
}
|
|
});
|
|
let res_read = handler.handle_request(req_read_res).await.unwrap();
|
|
assert_eq!(
|
|
res_read["result"]["contents"][0]["uri"],
|
|
"memory://tasks/active"
|
|
);
|
|
assert!(
|
|
res_read["result"]["contents"][0]["text"]
|
|
.as_str()
|
|
.unwrap()
|
|
.contains("[]")
|
|
); // Empty tasks
|
|
|
|
// Test prompts/list
|
|
let req_list_prompts = json!({
|
|
"jsonrpc": "2.0",
|
|
"id": 12,
|
|
"method": "prompts/list",
|
|
"params": {}
|
|
});
|
|
let res_prompts = handler.handle_request(req_list_prompts).await.unwrap();
|
|
let prompts_arr = res_prompts["result"]["prompts"].as_array().unwrap();
|
|
assert!(prompts_arr.iter().any(|p| p["name"] == "handoff_routine"));
|
|
|
|
// Test prompts/get
|
|
let req_get_prompt = json!({
|
|
"jsonrpc": "2.0",
|
|
"id": 13,
|
|
"method": "prompts/get",
|
|
"params": {
|
|
"name": "handoff_routine",
|
|
"arguments": {}
|
|
}
|
|
});
|
|
let res_get = handler.handle_request(req_get_prompt).await.unwrap();
|
|
let messages = res_get["result"]["messages"].as_array().unwrap();
|
|
assert_eq!(messages[0]["role"], "user");
|
|
assert!(
|
|
messages[0]["content"]["text"]
|
|
.as_str()
|
|
.unwrap()
|
|
.contains("standup report")
|
|
);
|
|
}
|
|
|
|
#[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);
|
|
}
|
|
}
|