refactor: apply zero-unwrap policy and optimize locks in store.rs and handlers

This commit is contained in:
Riz Ashraf committed 2026-09-21 11:34:21 +01:00
1 parent 9f24e66d88
commit 8afbf97b11
31 files changed
+3014 -3094

No files matched your search

+14
View File
@@ -0,0 +1,14 @@
import re
with open("server/src/handlers.rs", "r", encoding="utf-8") as f:
text = f.read()
pattern = r'"(list_milestones|list_pinned_files|read_handoff_memos)" => \{\s*let req = parse_tool!\(args, id, ([a-zA-Z]+)\);\s*let mut [a-zA-Z_]+ = self\.state\.([a-zA-Z_]+)\.read\(\);\s*if let Some\(ns\) = req\.namespace \{\s*[a-zA-Z_]+\.retain\(\|.\| \w+\.namespace == ns\);\s*\}\s*let data = serde_json::to_string\(&[a-zA-Z_]+\)\.unwrap_or_default\(\);\s*Ok\(data\.to_string\(\)\)\s*\}'
def repl(m):
return f'"{m.group(1)}" => handle_list_with_namespace!(self, {m.group(3)}, {m.group(2)}, args, id),'
new_text = re.sub(pattern, repl, text)
with open("server/src/handlers.rs", "w", encoding="utf-8") as f:
f.write(new_text)
+14
View File
@@ -0,0 +1,14 @@
import re
with open("server/src/handlers.rs", "r", encoding="utf-8") as f:
text = f.read()
pattern = r'"([a-zA-Z_]+)" => \{\s*(?:let req = parse_tool!\(args, id, ([a-zA-Z]+)\);\s*)?let data = serde_json::to_string\(&self\.state\.([a-zA-Z_]+)\.read\(\)\)\s*\.unwrap_or_else\(\|_\| "\[\]"\.to_string\(\)\);\s*Ok\(data\.to_string\(\)\)\s*\}'
def repl(m):
return f'"{m.group(1)}" => {{\n let data = serde_json::to_string(&self.state.{m.group(3)}.read()).unwrap_or_else(|_| "[]".to_string());\n Ok(data)\n}},'
new_text = re.sub(pattern, repl, text)
with open("server/src/handlers.rs", "w", encoding="utf-8") as f:
f.write(new_text)
+1 -5
View File
@@ -15,10 +15,6 @@ fn main() {
.and_then(|out| String::from_utf8(out.stdout).ok()) .and_then(|out| String::from_utf8(out.stdout).ok())
.unwrap_or_else(|| "unknown".to_string()); .unwrap_or_else(|| "unknown".to_string());
let version = format!( let version = format!("{} ({})", git_date.trim(), git_hash.trim());
"{} ({})",
git_date.trim(),
git_hash.trim()
);
println!("cargo:rustc-env=APP_VERSION={}", version); println!("cargo:rustc-env=APP_VERSION={}", version);
} }
+2
View File
@@ -1,3 +1,5 @@
#![cfg(unix)]
use serde_json::{json, Value}; use serde_json::{json, Value};
use std::io::{BufRead, BufReader, Read, Write}; use std::io::{BufRead, BufReader, Read, Write};
use std::process::{Command, Stdio}; use std::process::{Command, Stdio};
+53 -27
View File
@@ -20,8 +20,6 @@ pub struct JsonRpcResponse {
pub error: Option<Value>, pub error: Option<Value>,
} }
pub async fn send_response(response: JsonRpcResponse) { pub async fn send_response(response: JsonRpcResponse) {
let msg = serde_json::to_string(&response).unwrap(); let msg = serde_json::to_string(&response).unwrap();
tracing::info!( tracing::info!(
@@ -111,11 +109,10 @@ async fn get_socket_path() -> Result<String, String> {
} }
Err("Could not find Neovim socket".to_string()) Err("Could not find Neovim socket".to_string())
} }
use std::sync::LazyLock;
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot};
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::Arc;
use std::sync::LazyLock;
use tokio::sync::{mpsc, oneshot};
pub struct NvimRequest { pub struct NvimRequest {
pub msgid_str: String, pub msgid_str: String,
@@ -123,7 +120,8 @@ pub struct NvimRequest {
pub reply: oneshot::Sender<Result<rmpv::Value, String>>, pub reply: oneshot::Sender<Result<rmpv::Value, String>>,
} }
static NVIM_CONN: LazyLock<Arc<std::sync::Mutex<Option<mpsc::Sender<NvimRequest>>>>> = LazyLock::new(|| Arc::new(std::sync::Mutex::new(None))); static NVIM_CONN: LazyLock<Arc<std::sync::Mutex<Option<mpsc::Sender<NvimRequest>>>>> =
LazyLock::new(|| Arc::new(std::sync::Mutex::new(None)));
async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> { async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
{ {
@@ -141,18 +139,23 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
#[cfg(windows)] #[cfg(windows)]
let stream = { let stream = {
use tokio::net::windows::named_pipe::ClientOptions; use tokio::net::windows::named_pipe::ClientOptions;
ClientOptions::new().open(&socket_path).map_err(|e| e.to_string())? ClientOptions::new()
.open(&socket_path)
.map_err(|e| e.to_string())?
}; };
#[cfg(unix)] #[cfg(unix)]
let stream = { let stream = {
use tokio::net::UnixStream; use tokio::net::UnixStream;
UnixStream::connect(socket_path).await.map_err(|e| e.to_string())? UnixStream::connect(socket_path)
.await
.map_err(|e| e.to_string())?
}; };
let (mut read_half, mut write_half) = tokio::io::split(stream); let (mut read_half, mut write_half) = tokio::io::split(stream);
let (tx, mut rx) = mpsc::channel::<NvimRequest>(32); let (tx, mut rx) = mpsc::channel::<NvimRequest>(32);
type PendingRequestsMap = Arc<std::sync::Mutex<HashMap<String, oneshot::Sender<Result<rmpv::Value, String>>>>>; type PendingRequestsMap =
Arc<std::sync::Mutex<HashMap<String, oneshot::Sender<Result<rmpv::Value, String>>>>>;
let pending_requests: PendingRequestsMap = Arc::new(std::sync::Mutex::new(HashMap::new())); let pending_requests: PendingRequestsMap = Arc::new(std::sync::Mutex::new(HashMap::new()));
// Write task // Write task
@@ -165,7 +168,10 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
continue; continue;
} }
pending_clone.lock().unwrap().insert(req.msgid_str.clone(), req.reply); pending_clone
.lock()
.unwrap()
.insert(req.msgid_str.clone(), req.reply);
if write_half.write_all(&buf).await.is_err() { if write_half.write_all(&buf).await.is_err() {
tracing::error!("Failed to write to Neovim socket"); tracing::error!("Failed to write to Neovim socket");
@@ -192,7 +198,9 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
let msgid = &arr[1]; let msgid = &arr[1];
let msgid_str = format!("{:?}", msgid); let msgid_str = format!("{:?}", msgid);
if let Some(reply_sender) = pending_clone2.lock().unwrap().remove(&msgid_str) { if let Some(reply_sender) =
pending_clone2.lock().unwrap().remove(&msgid_str)
{
let _ = reply_sender.send(Ok(val)); let _ = reply_sender.send(Ok(val));
} }
} }
@@ -204,12 +212,16 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
} }
continue; continue;
} }
Err(rmpv::decode::Error::InvalidMarkerRead(e)) if e.kind() == std::io::ErrorKind::UnexpectedEof => { Err(rmpv::decode::Error::InvalidMarkerRead(e))
if e.kind() == std::io::ErrorKind::UnexpectedEof =>
{
resp_buf.drain(..offset); resp_buf.drain(..offset);
offset = 0; offset = 0;
let read_future = read_half.read(&mut chunk); let read_future = read_half.read(&mut chunk);
match tokio::time::timeout(tokio::time::Duration::from_secs(60), read_future).await { match tokio::time::timeout(tokio::time::Duration::from_secs(60), read_future)
.await
{
Ok(Ok(n)) if n > 0 => { Ok(Ok(n)) if n > 0 => {
resp_buf.extend_from_slice(&chunk[..n]); resp_buf.extend_from_slice(&chunk[..n]);
} }
@@ -242,7 +254,10 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
if Arc::strong_count(&pending_clone3) <= 1 { if Arc::strong_count(&pending_clone3) <= 1 {
break; // Socket closed and other tasks finished, no need to keep cleaning up break; // Socket closed and other tasks finished, no need to keep cleaning up
} }
pending_clone3.lock().unwrap().retain(|_, sender| !sender.is_closed()); pending_clone3
.lock()
.unwrap()
.retain(|_, sender| !sender.is_closed());
} }
}); });
@@ -276,7 +291,9 @@ async fn call_nvim(req: rmpv::Value) -> Result<rmpv::Value, String> {
msgid_str, msgid_str,
req, req,
reply: reply_tx, reply: reply_tx,
}).await.map_err(|_| "Failed to send request to Neovim connection manager")?; })
.await
.map_err(|_| "Failed to send request to Neovim connection manager")?;
match tokio::time::timeout(tokio::time::Duration::from_secs(10), reply_rx).await { match tokio::time::timeout(tokio::time::Duration::from_secs(10), reply_rx).await {
Ok(Ok(res)) => res, Ok(Ok(res)) => res,
@@ -544,9 +561,13 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
Ok(m) => { Ok(m) => {
tracing::info!("Received message method: {}", m.method); tracing::info!("Received message method: {}", m.method);
m m
}, }
Err(e) => { Err(e) => {
tracing::error!("Failed to parse JSON-RPC request from JSONL: {}. Payload: {}", e, raw_msg); tracing::error!(
"Failed to parse JSON-RPC request from JSONL: {}. Payload: {}",
e,
raw_msg
);
continue; continue;
} }
}; };
@@ -559,13 +580,17 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
let id_clone = id.clone(); let id_clone = id.clone();
let start_time = std::time::Instant::now(); let start_time = std::time::Instant::now();
let method_clone = if msg.method == "tools/call" { let method_clone = if msg.method == "tools/call" {
let tool_name = msg.params.as_ref().and_then(|p| p.get("name")).and_then(|n| n.as_str()).unwrap_or("unknown"); let tool_name = msg
.params
.as_ref()
.and_then(|p| p.get("name"))
.and_then(|n| n.as_str())
.unwrap_or("unknown");
format!("ToolCall[{}]", tool_name) format!("ToolCall[{}]", tool_name)
} else { } else {
msg.method.clone() msg.method.clone()
}; };
match msg.method.as_str() { match msg.method.as_str() {
"initialize" => { "initialize" => {
@@ -984,13 +1009,16 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
} }
} }
let elapsed = start_time.elapsed(); let elapsed = start_time.elapsed();
tracing::info!("<<< [Nvim] {} (id: {}) completed in {:?}", method_clone, id_clone, elapsed); tracing::info!(
"<<< [Nvim] {} (id: {}) completed in {:?}",
method_clone,
id_clone,
elapsed
);
}); });
} }
} }
fn init_logging(app_name: &str) -> tracing_appender::non_blocking::WorkerGuard { fn init_logging(app_name: &str) -> tracing_appender::non_blocking::WorkerGuard {
let log_dir = dirs::home_dir() let log_dir = dirs::home_dir()
.unwrap_or_default() .unwrap_or_default()
@@ -1036,11 +1064,10 @@ mod tests {
#[test] #[test]
fn test_rmpv_to_json_map() { fn test_rmpv_to_json_map() {
let mut map = vec![]; let map = vec![(
map.push((
rmpv::Value::String("key1".into()), rmpv::Value::String("key1".into()),
rmpv::Value::Integer(100.into()), rmpv::Value::Integer(100.into()),
)); )];
let rmp_map = rmpv::Value::Map(map); let rmp_map = rmpv::Value::Map(map);
let json_map = rmpv_to_json(&rmp_map); let json_map = rmpv_to_json(&rmp_map);
@@ -1074,4 +1101,3 @@ mod tests {
assert!(req.is_none()); assert!(req.is_none());
} }
} }
+116 -2989
View File
File diff suppressed because it is too large. Load diff
+163
View File
@@ -0,0 +1,163 @@
use crate::router::McpTool;
use crate::state::MemoryState;
use crate::tools::*;
use async_trait::async_trait;
use serde_json::Value;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
pub struct UpdateEnvFingerprintHandler;
#[async_trait]
impl McpTool for UpdateEnvFingerprintHandler {
fn name(&self) -> &'static str {
"update_env_fingerprint"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<UpdateEnvFingerprintTool>(
"update_env_fingerprint",
"Execute update_env_fingerprint",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: UpdateEnvFingerprintTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.env_fingerprints.modify(|fps| {
fps.insert(
req.namespace.clone(),
crate::models::EnvFingerprint {
namespace: req.namespace.clone(),
os: std::env::consts::OS.to_string(),
shell: std::env::var("SHELL").unwrap_or_else(|_| "unknown".to_string()),
tool_versions: req.tool_versions,
updated_at: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
},
);
});
Ok("Env fingerprint updated".to_string())
}
}
pub struct ReadEnvFingerprintHandler;
#[async_trait]
impl McpTool for ReadEnvFingerprintHandler {
fn name(&self) -> &'static str {
"read_env_fingerprint"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ReadEnvFingerprintTool>(
"read_env_fingerprint",
"Execute read_env_fingerprint",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ReadEnvFingerprintTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let fps = state.env_fingerprints.read();
if let Some(fp) = fps.get(&req.namespace) {
let data = serde_json::to_string(fp).unwrap_or_default();
Ok(data.to_string())
} else {
Ok("{}".to_string())
}
}
}
pub struct LogEnvRequirementHandler;
#[async_trait]
impl McpTool for LogEnvRequirementHandler {
fn name(&self) -> &'static str {
"log_env_requirement"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LogEnvRequirementTool>(
"log_env_requirement",
"Execute log_env_requirement",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LogEnvRequirementTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.env_requirements.modify(|reqs| {
reqs.retain(|r| !(r.namespace == req.namespace && r.key == req.key));
reqs.push(crate::models::EnvRequirement {
namespace: req.namespace,
key: req.key,
description: req.description,
is_secret: req.is_secret,
});
});
Ok("Env requirement logged".to_string())
}
}
pub struct RegisterEnvironmentHandler;
#[async_trait]
impl McpTool for RegisterEnvironmentHandler {
fn name(&self) -> &'static str {
"register_environment"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<RegisterEnvironmentTool>(
"register_environment",
"Execute register_environment",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: RegisterEnvironmentTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.environments.modify(|envs| {
envs.retain(|e| !(e.namespace == req.namespace && e.name == req.name));
envs.push(crate::models::EnvironmentDetail {
namespace: req.namespace,
name: req.name,
url: req.url,
description: req.description,
requires_vpn: req.requires_vpn,
updated_at: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
});
});
Ok("Environment registered".to_string())
}
}
pub struct GetEnvironmentDetailsHandler;
#[async_trait]
impl McpTool for GetEnvironmentDetailsHandler {
fn name(&self) -> &'static str {
"get_environment_details"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<GetEnvironmentDetailsTool>(
"get_environment_details",
"Execute get_environment_details",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: GetEnvironmentDetailsTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut envs = state.environments.read();
envs.retain(|e| e.namespace == req.namespace);
let data = serde_json::to_string(&envs).unwrap_or_default();
Ok(data.to_string())
}
}
+544
View File
@@ -0,0 +1,544 @@
use crate::models::*;
use crate::router::McpTool;
use crate::state::MemoryState;
use crate::tools::*;
use async_trait::async_trait;
use serde_json::Value;
use std::collections::HashSet;
use std::sync::Arc;
pub struct QueryGraphPathHandler;
#[async_trait]
impl McpTool for QueryGraphPathHandler {
fn name(&self) -> &'static str {
"query_graph_path"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<QueryGraphPathTool>("query_graph_path", "Execute query_graph_path")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: crate::tools::QueryGraphPathTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.read_graph(|graph| {
let max_depth = req.max_depth.unwrap_or(5);
let mut queue = std::collections::VecDeque::new();
let mut visited = std::collections::HashSet::new();
let mut parents: std::collections::HashMap<String, (String, String)> =
std::collections::HashMap::new();
queue.push_back(req.start_node.clone());
visited.insert(req.start_node.clone());
let mut found = false;
let mut current_depth = 0;
let mut nodes_at_current_depth = 1;
let mut nodes_at_next_depth = 0;
while let Some(current) = queue.pop_front() {
if current == req.end_node {
found = true;
break;
}
nodes_at_current_depth -= 1;
if current_depth < max_depth {
for rel in &graph.relations {
if rel.from == current && !visited.contains(&rel.to) {
visited.insert(rel.to.clone());
parents.insert(
rel.to.clone(),
(current.clone(), rel.relation_type.clone()),
);
queue.push_back(rel.to.clone());
nodes_at_next_depth += 1;
} else if rel.to == current && !visited.contains(&rel.from) {
visited.insert(rel.from.clone());
parents.insert(
rel.from.clone(),
(current.clone(), format!("inverse({})", rel.relation_type)),
);
queue.push_back(rel.from.clone());
nodes_at_next_depth += 1;
}
}
}
if nodes_at_current_depth == 0 {
current_depth += 1;
nodes_at_current_depth = nodes_at_next_depth;
nodes_at_next_depth = 0;
}
}
if found {
let mut path = Vec::new();
let mut curr = req.end_node.clone();
while curr != req.start_node {
if let Some((parent, rel_type)) = parents.get(&curr) {
path.push(format!("{} -[{}]-> {}", parent, rel_type, curr));
curr = parent.clone();
} else {
break;
}
}
path.reverse();
Ok(format!("Path found:\n{}", path.join("\n")))
} else {
Ok(format!(
"No path found between {} and {} within depth {}",
req.start_node, req.end_node, max_depth
))
}
})
}
}
pub struct CreateEntitiesHandler;
#[async_trait]
impl McpTool for CreateEntitiesHandler {
fn name(&self) -> &'static str {
"create_entities"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Execute create_entities")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: CreateEntitiesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.modify_graph(|g| {
for entity in req.entities {
if !entity.name.is_empty() {
if let Ok(idx) = state.search_index.read() {
drop(idx.index_entity(&entity));
}
g.entities.insert(entity.name.clone(), entity);
}
}
});
Ok("Entities created".to_string())
}
}
pub struct CreateRelationsHandler;
#[async_trait]
impl McpTool for CreateRelationsHandler {
fn name(&self) -> &'static str {
"create_relations"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<CreateRelationsTool>("create_relations", "Execute create_relations")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: CreateRelationsTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.modify_graph(|g| {
for relation in req.relations {
if !relation.from.is_empty() && !relation.to.is_empty() {
g.relations.push(relation);
}
}
});
Ok("Relations created".to_string())
}
}
pub struct AddObservationsHandler;
#[async_trait]
impl McpTool for AddObservationsHandler {
fn name(&self) -> &'static str {
"add_observations"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<AddObservationsTool>("add_observations", "Execute add_observations")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: AddObservationsTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.modify_graph(|g| {
for o in req.observations {
if let Some(e) = g.entities.get_mut(&o.entity_name) {
e.observations.extend(o.contents);
}
}
});
Ok("Observations added".to_string())
}
}
pub struct DeleteEntitiesHandler;
#[async_trait]
impl McpTool for DeleteEntitiesHandler {
fn name(&self) -> &'static str {
"delete_entities"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<DeleteEntitiesTool>("delete_entities", "Execute delete_entities")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: DeleteEntitiesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let to_delete: HashSet<_> = req.entity_names.into_iter().collect();
state.modify_graph(|master| {
for name in &to_delete {
master.entities.remove(name);
}
master
.relations
.retain(|r| !to_delete.contains(&r.from) && !to_delete.contains(&r.to));
});
Ok("Entities deleted".to_string())
}
}
pub struct DeleteObservationsHandler;
#[async_trait]
impl McpTool for DeleteObservationsHandler {
fn name(&self) -> &'static str {
"delete_observations"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<DeleteObservationsTool>(
"delete_observations",
"Execute delete_observations",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: DeleteObservationsTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.modify_graph(|master| {
for d in req.deletions {
if let Some(e) = master.entities.get_mut(&d.entity_name) {
let to_rem: HashSet<_> = d.observations.into_iter().collect();
e.observations.retain(|o| !to_rem.contains(o));
}
}
});
Ok("Observations deleted".to_string())
}
}
pub struct DeleteRelationsHandler;
#[async_trait]
impl McpTool for DeleteRelationsHandler {
fn name(&self) -> &'static str {
"delete_relations"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<DeleteRelationsTool>("delete_relations", "Execute delete_relations")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: DeleteRelationsTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.modify_graph(|master| {
let to_rem: HashSet<_> = req.relations.into_iter().collect();
master.relations.retain(|r| !to_rem.contains(r));
});
Ok("Relations deleted".to_string())
}
}
pub struct ReadGraphHandler;
#[async_trait]
impl McpTool for ReadGraphHandler {
fn name(&self) -> &'static str {
"read_graph"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ReadGraphTool>("read_graph", "Execute read_graph")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ReadGraphTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let data = state.read_graph(|full| {
if let Some(ns) = req.namespace {
let mut filtered = KnowledgeGraph::default();
for (k, v) in &full.entities {
if v.namespace == ns {
filtered.entities.insert(k.clone(), v.clone());
}
}
for r in &full.relations {
if r.namespace == ns {
filtered.relations.push(r.clone());
}
}
serde_json::to_string(&filtered).unwrap_or_default()
} else {
serde_json::to_string(full).unwrap_or_default()
}
});
Ok(data)
}
}
pub struct SearchNodesHandler;
#[async_trait]
impl McpTool for SearchNodesHandler {
fn name(&self) -> &'static str {
"search_nodes"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<SearchNodesTool>("search_nodes", "Execute search_nodes")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: SearchNodesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let matches = if let Ok(idx) = state.search_index.read() {
idx.search(&req.query, req.namespace.as_deref())
.unwrap_or_default()
} else {
vec![]
};
let mut result = KnowledgeGraph::default();
state.read_graph(|full| {
for (id, doc_type, _, _, _) in matches {
if doc_type == "entity"
&& let Some(e) = full.entities.get(&id)
{
result.entities.insert(id, e.clone());
}
}
});
let data = serde_json::to_string(&result).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct OpenNodesHandler;
#[async_trait]
impl McpTool for OpenNodesHandler {
fn name(&self) -> &'static str {
"open_nodes"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<OpenNodesTool>("open_nodes", "Execute open_nodes")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: OpenNodesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let targets: HashSet<_> = req.names.into_iter().collect();
let mut result = KnowledgeGraph::default();
let mut connected = HashSet::new();
state.read_graph(|full| {
for r in &full.relations {
if targets.contains(&r.from) {
connected.insert(r.to.clone());
result.relations.push(r.clone());
} else if targets.contains(&r.to) {
connected.insert(r.from.clone());
result.relations.push(r.clone());
}
}
for (name, e) in &full.entities {
if targets.contains(name) || connected.contains(name) {
result.entities.insert(name.clone(), e.clone());
}
}
});
let data = serde_json::to_string(&result).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct VisualizeGraphHandler;
#[async_trait]
impl McpTool for VisualizeGraphHandler {
fn name(&self) -> &'static str {
"visualize_graph"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<VisualizeGraphTool>("visualize_graph", "Execute visualize_graph")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: VisualizeGraphTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let query = req.query.unwrap_or_default().to_lowercase();
let mut included = HashSet::new();
let mut to_draw = Vec::new();
state.read_graph(|full| {
for (name, e) in &full.entities {
if let Some(ns) = &req.namespace
&& e.namespace != *ns
{
continue;
}
if query.is_empty()
|| contains_ignore_ascii_case(name, &query)
|| contains_ignore_ascii_case(&e.entity_type, &query)
{
included.insert(name.clone());
}
}
for r in &full.relations {
if let Some(ns) = &req.namespace
&& r.namespace != *ns
{
continue;
}
if query.is_empty() || included.contains(&r.from) || included.contains(&r.to) {
included.insert(r.from.clone());
included.insert(r.to.clone());
to_draw.push(r.clone());
}
}
});
use std::fmt::Write;
let mut output = String::with_capacity(included.len() * 40 + to_draw.len() * 60);
output.push_str("graph TD;\n");
let sanitize = |s: &str, id_mode: bool| -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
if c != '"' && c != '(' && c != ')' {
if id_mode && (c == ' ' || c == '-' || c == '.') {
out.push('_');
} else {
out.push(c);
}
}
}
out
};
for name in &included {
let _ = writeln!(
output,
" id_{}[\"{}\"];",
sanitize(name, true),
sanitize(name, false)
);
}
for r in to_draw {
let _ = writeln!(
output,
" id_{}-->|\"{}\"|id_{};",
sanitize(&r.from, true),
r.relation_type.replace("\"", ""),
sanitize(&r.to, true)
);
}
if output == "graph TD;\n" {
output = "No nodes found to visualize.".to_string();
}
Ok(output.to_string())
}
}
pub struct CondenseEntityHandler;
#[async_trait]
impl McpTool for CondenseEntityHandler {
fn name(&self) -> &'static str {
"condense_entity"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<CondenseEntityTool>("condense_entity", "Execute condense_entity")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: CondenseEntityTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.modify_graph(|master| {
if let Some(e) = master.entities.get_mut(&req.entity_name) {
e.observations = req.summarized_observations;
}
});
Ok("Entity condensed".to_string())
}
}
pub struct MergeEntitiesHandler;
#[async_trait]
impl McpTool for MergeEntitiesHandler {
fn name(&self) -> &'static str {
"merge_entities"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<MergeEntitiesTool>("merge_entities", "Execute merge_entities")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: MergeEntitiesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.modify_graph(|master| {
if let Some(src) = master.entities.remove(&req.source_entity) {
if let Some(tgt) = master.entities.get_mut(&req.target_entity) {
tgt.observations.extend(src.observations);
MemoryState::deduplicate(&mut tgt.observations);
} else {
let mut new_tgt = src.clone();
new_tgt.name = req.target_entity.clone();
master.entities.insert(req.target_entity.clone(), new_tgt);
}
}
for r in &mut master.relations {
if r.from == req.source_entity {
r.from = req.target_entity.clone();
}
if r.to == req.source_entity {
r.to = req.target_entity.clone();
}
}
MemoryState::deduplicate(&mut master.relations);
});
Ok("Entities merged".to_string())
}
}
pub struct FindOrphansHandler;
#[async_trait]
impl McpTool for FindOrphansHandler {
fn name(&self) -> &'static str {
"find_orphans"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<FindOrphansTool>("find_orphans", "Execute find_orphans")
}
async fn execute(&self, _args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let orphans = state.read_graph(|full| {
let mut connected = std::collections::HashSet::new();
for r in &full.relations {
connected.insert(r.from.clone());
connected.insert(r.to.clone());
}
full.entities
.keys()
.filter(|k| !connected.contains(*k))
.cloned()
.collect::<Vec<String>>()
});
let data = serde_json::to_string(&orphans).unwrap_or_default();
Ok(data.to_string())
}
}
use crate::handlers_v2::utils::*;
+492
View File
@@ -0,0 +1,492 @@
use crate::models::*;
use crate::router::McpTool;
use crate::state::MemoryState;
use crate::tools::*;
use async_trait::async_trait;
use serde_json::Value;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
pub struct LogDecisionHandler;
#[async_trait]
impl McpTool for LogDecisionHandler {
fn name(&self) -> &'static str {
"log_decision"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LogDecisionTool>("log_decision", "Execute log_decision")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LogDecisionTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut adr_id = String::new();
let mut new_adr = None;
state.adrs.modify(|adrs| {
adr_id = format!("ADR-{:04}", adrs.len() + 1);
let a = Adr {
id: adr_id.clone(),
title: req.title,
context: req.context,
decision: req.decision,
consequence: req.consequence,
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
};
new_adr = Some(a.clone());
adrs.push(a);
});
if let Some(adr) = new_adr
&& let Ok(idx) = state.search_index.read()
{
drop(idx.index_adr(&adr));
}
Ok(format!("Decision logged as {}", adr_id).to_string())
}
}
pub struct QueryDecisionsHandler;
#[async_trait]
impl McpTool for QueryDecisionsHandler {
fn name(&self) -> &'static str {
"query_decisions"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<QueryDecisionsTool>("query_decisions", "Execute query_decisions")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: QueryDecisionsTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut adrs = state.adrs.read();
if let Some(q) = req.query {
let q = q.to_lowercase();
adrs.retain(|a| {
contains_ignore_ascii_case(&a.title, &q)
|| contains_ignore_ascii_case(&a.context, &q)
|| contains_ignore_ascii_case(&a.decision, &q)
});
}
let data = serde_json::to_string(&adrs).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct LogErrorFixHandler;
#[async_trait]
impl McpTool for LogErrorFixHandler {
fn name(&self) -> &'static str {
"log_error_fix"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LogErrorFixTool>("log_error_fix", "Execute log_error_fix")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LogErrorFixTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.error_fixes.modify(|fixes| {
fixes.push(crate::models::ErrorFix {
signature: req.signature,
solution: req.solution,
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
git_commit: req.git_commit,
git_branch: req.git_branch,
})
});
Ok("Error fix logged".to_string())
}
}
pub struct SearchErrorFixesHandler;
#[async_trait]
impl McpTool for SearchErrorFixesHandler {
fn name(&self) -> &'static str {
"search_error_fixes"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<SearchErrorFixesTool>(
"search_error_fixes",
"Execute search_error_fixes",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: SearchErrorFixesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let q = req.query.to_lowercase();
let mut fixes = state.error_fixes.read();
fixes.retain(|f| {
contains_ignore_ascii_case(&f.signature, &q)
|| contains_ignore_ascii_case(&f.solution, &q)
});
let data = serde_json::to_string(&fixes).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct LogCodeChangeHandler;
#[async_trait]
impl McpTool for LogCodeChangeHandler {
fn name(&self) -> &'static str {
"log_code_change"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LogCodeChangeTool>("log_code_change", "Execute log_code_change")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LogCodeChangeTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.ledger.modify(|ledger| {
ledger.push(CodeChange {
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
file_path: req.file_path,
description: req.description,
git_commit: req.git_commit,
git_branch: req.git_branch,
});
});
Ok("Code change logged".to_string())
}
}
pub struct QueryRecentChangesHandler;
#[async_trait]
impl McpTool for QueryRecentChangesHandler {
fn name(&self) -> &'static str {
"query_recent_changes"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<QueryRecentChangesTool>(
"query_recent_changes",
"Execute query_recent_changes",
)
}
async fn execute(&self, _args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let data = serde_json::to_string(&state.ledger.read()).unwrap_or_else(|_| "[]".to_string());
Ok(data.to_string())
}
}
pub struct LearnPreferenceHandler;
#[async_trait]
impl McpTool for LearnPreferenceHandler {
fn name(&self) -> &'static str {
"learn_preference"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LearnPreferenceTool>("learn_preference", "Execute learn_preference")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LearnPreferenceTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.prefs.modify(|prefs| {
prefs.insert(
req.key.clone(),
crate::models::Preference {
key: req.key.clone(),
value: req.value,
updated_at: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
},
);
});
Ok("Preference learned".to_string())
}
}
pub struct ReadPreferencesHandler;
#[async_trait]
impl McpTool for ReadPreferencesHandler {
fn name(&self) -> &'static str {
"read_preferences"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ReadPreferencesTool>("read_preferences", "Execute read_preferences")
}
async fn execute(&self, _args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let prefs = state.prefs.read();
let data = serde_json::to_string(&prefs).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct LogTechDebtHandler;
#[async_trait]
impl McpTool for LogTechDebtHandler {
fn name(&self) -> &'static str {
"log_tech_debt"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LogTechDebtTool>("log_tech_debt", "Execute log_tech_debt")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LogTechDebtTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.tech_debts.modify(|debts| {
debts.push(crate::models::TechDebt {
id: uuid::Uuid::new_v4().to_string(),
namespace: req.namespace,
description: req.description,
ideal_solution: req.ideal_solution,
is_resolved: false,
created_at: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
git_commit: req.git_commit,
git_branch: req.git_branch,
})
});
Ok("Tech debt logged".to_string())
}
}
pub struct ResolveTechDebtHandler;
#[async_trait]
impl McpTool for ResolveTechDebtHandler {
fn name(&self) -> &'static str {
"resolve_tech_debt"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ResolveTechDebtTool>(
"resolve_tech_debt",
"Execute resolve_tech_debt",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ResolveTechDebtTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut found = false;
state.tech_debts.modify(|debts| {
for d in debts.iter_mut() {
if d.id == req.id {
d.is_resolved = true;
found = true;
break;
}
}
});
if found {
Ok("Tech debt resolved".to_string())
} else {
Ok("Tech debt not found".to_string())
}
}
}
pub struct ListTechDebtHandler;
#[async_trait]
impl McpTool for ListTechDebtHandler {
fn name(&self) -> &'static str {
"list_tech_debt"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ListTechDebtTool>("list_tech_debt", "Execute list_tech_debt")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ListTechDebtTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut debts = state.tech_debts.read();
debts.retain(|d| d.namespace == req.namespace && (req.include_resolved || !d.is_resolved));
let data = serde_json::to_string(&debts).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct OmniSearchHandler;
#[async_trait]
impl McpTool for OmniSearchHandler {
fn name(&self) -> &'static str {
"omni_search"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<OmniSearchTool>("omni_search", "Execute omni_search")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: OmniSearchTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let matches = if let Ok(idx) = state.search_index.read() {
idx.search(&req.query, req.namespace.as_deref())
.unwrap_or_default()
} else {
vec![]
};
let mut kg = KnowledgeGraph::default();
let mut tasks = Vec::new();
let mut snippets = Vec::new();
let mut adrs = Vec::new();
state.read_graph(|full| {
for (id, doc_type, _, _, _) in &matches {
if doc_type == "entity"
&& let Some(e) = full.entities.get(id)
{
kg.entities.insert(id.clone(), e.clone());
}
}
});
for t in state.tasks.read() {
if matches
.iter()
.any(|(id, typ, _, _, _)| id == &t.id && typ == "task")
{
tasks.push(t);
}
}
for s in state.snippets.read() {
if matches
.iter()
.any(|(id, typ, _, _, _)| id == &s.name && typ == "snippet")
{
snippets.push(s);
}
}
for a in state.adrs.read() {
if matches
.iter()
.any(|(id, typ, _, _, _)| id == &a.id && typ == "adr")
{
adrs.push(a);
}
}
let q = req.query.to_lowercase();
let tech_debts: Vec<_> = state
.tech_debts
.read()
.into_iter()
.filter(|d| {
req.namespace.as_ref().is_none_or(|ns| d.namespace == *ns)
&& (contains_ignore_ascii_case(&d.description, &q)
|| contains_ignore_ascii_case(&d.ideal_solution, &q))
})
.collect();
let memos: Vec<_> = state
.handoff_memos
.read()
.into_iter()
.filter(|m| {
req.namespace.as_ref().is_none_or(|ns| m.namespace == *ns)
&& contains_ignore_ascii_case(&m.content, &q)
})
.collect();
let error_fixes: Vec<_> = state
.error_fixes
.read()
.into_iter()
.filter(|f| {
contains_ignore_ascii_case(&f.signature, &q)
|| contains_ignore_ascii_case(&f.solution, &q)
})
.collect();
let report = serde_json::json!({
"knowledge_graph": kg.entities,
"tasks": tasks,
"snippets": snippets,
"adrs": adrs,
"tech_debts": tech_debts,
"handoff_memos": memos,
"error_fixes": error_fixes
});
Ok(report.to_string())
}
}
pub struct GetProjectHealthHandler;
#[async_trait]
impl McpTool for GetProjectHealthHandler {
fn name(&self) -> &'static str {
"get_project_health"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<GetProjectHealthTool>(
"get_project_health",
"Execute get_project_health",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: GetProjectHealthTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let active_tasks = state
.tasks
.read()
.into_iter()
.filter(|t| t.status != "done")
.count();
let unresolved_debt = state
.tech_debts
.read()
.into_iter()
.filter(|d| d.namespace == req.namespace && !d.is_resolved)
.count();
let unread_memos = state
.handoff_memos
.read()
.into_iter()
.filter(|m| m.namespace == req.namespace)
.count();
let active_milestones = state
.milestones
.read()
.into_iter()
.filter(|m| m.namespace == req.namespace && m.status != "done")
.count();
let remaining_checklists = state
.pr_checklists
.read()
.into_iter()
.filter(|c| c.namespace == req.namespace)
.count();
let report = serde_json::json!({
"active_tasks": active_tasks,
"unresolved_tech_debt": unresolved_debt,
"unread_handoff_memos": unread_memos,
"active_milestones": active_milestones,
"remaining_pr_checklist_items": remaining_checklists
});
Ok(report.to_string())
}
}
use crate::handlers_v2::utils::*;
+7
View File
@@ -0,0 +1,7 @@
pub mod env;
pub mod graph;
pub mod meta;
pub mod notes;
pub mod tasks;
pub mod utils;
pub mod workspaces;
+273
View File
@@ -0,0 +1,273 @@
use crate::models::*;
use crate::router::McpTool;
use crate::state::MemoryState;
use crate::tools::*;
use async_trait::async_trait;
use serde_json::Value;
use std::collections::HashSet;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
pub struct AddStickyNoteHandler;
#[async_trait]
impl McpTool for AddStickyNoteHandler {
fn name(&self) -> &'static str {
"add_sticky_note"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<AddStickyNoteTool>("add_sticky_note", "Execute add_sticky_note")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: AddStickyNoteTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.sticky.modify(|notes| {
notes.push(StickyNote {
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
content: req.content,
});
});
Ok("Sticky note added.".to_string())
}
}
pub struct ReadStickyNotesHandler;
#[async_trait]
impl McpTool for ReadStickyNotesHandler {
fn name(&self) -> &'static str {
"read_sticky_notes"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ReadStickyNotesTool>(
"read_sticky_notes",
"Execute read_sticky_notes",
)
}
async fn execute(&self, _args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let data = serde_json::to_string(&state.sticky.read()).unwrap_or_else(|_| "[]".to_string());
Ok(data.to_string())
}
}
pub struct DeleteStickyNoteHandler;
#[async_trait]
impl McpTool for DeleteStickyNoteHandler {
fn name(&self) -> &'static str {
"delete_sticky_note"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<DeleteStickyNoteTool>(
"delete_sticky_note",
"Execute delete_sticky_note",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: DeleteStickyNoteTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut success = false;
state.sticky.modify(|notes| {
if req.index > 0 && req.index <= notes.len() {
notes.remove(req.index - 1);
success = true;
}
});
if success {
Ok("Sticky note deleted.".to_string())
} else {
Err("Invalid sticky note index.".to_string())
}
}
}
pub struct ClearStickyNotesHandler;
#[async_trait]
impl McpTool for ClearStickyNotesHandler {
fn name(&self) -> &'static str {
"clear_sticky_notes"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ClearStickyNotesTool>(
"clear_sticky_notes",
"Execute clear_sticky_notes",
)
}
async fn execute(&self, _args: Value, state: Arc<MemoryState>) -> Result<String, String> {
state.sticky.modify(|notes| {
notes.clear();
});
Ok("All sticky notes cleared.".to_string())
}
}
pub struct LeaveHandoffMemoHandler;
#[async_trait]
impl McpTool for LeaveHandoffMemoHandler {
fn name(&self) -> &'static str {
"leave_handoff_memo"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LeaveHandoffMemoTool>(
"leave_handoff_memo",
"Execute leave_handoff_memo",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LeaveHandoffMemoTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.handoff_memos.modify(|memos| {
memos.push(crate::models::HandoffMemo {
id: uuid::Uuid::new_v4().to_string(),
author: "agy".to_string(),
content: req.content,
namespace: req.namespace,
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
})
});
Ok("Handoff memo left".to_string())
}
}
pub struct ReadHandoffMemosHandler;
#[async_trait]
impl McpTool for ReadHandoffMemosHandler {
fn name(&self) -> &'static str {
"read_handoff_memos"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ReadHandoffMemosTool>(
"read_handoff_memos",
"Execute read_handoff_memos",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ReadHandoffMemosTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut items = state.handoff_memos.read();
if let Some(ns) = req.namespace {
items.retain(|i| i.namespace == ns);
}
let data = serde_json::to_string(&items).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct ClearHandoffMemosHandler;
#[async_trait]
impl McpTool for ClearHandoffMemosHandler {
fn name(&self) -> &'static str {
"clear_handoff_memos"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ClearHandoffMemosTool>(
"clear_handoff_memos",
"Execute clear_handoff_memos",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ClearHandoffMemosTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let ids: HashSet<_> = req.ids.into_iter().collect();
state
.handoff_memos
.modify(|memos| memos.retain(|m| !ids.contains(&m.id)));
Ok("Handoff memos cleared".to_string())
}
}
pub struct AddSessionSummaryHandler;
#[async_trait]
impl McpTool for AddSessionSummaryHandler {
fn name(&self) -> &'static str {
"add_session_summary"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<AddSessionSummaryTool>(
"add_session_summary",
"Execute add_session_summary",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: AddSessionSummaryTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.session_summaries.modify(|summaries| {
summaries.push(crate::models::SessionSummary {
summary: req.summary,
namespace: req.namespace,
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
})
});
Ok("Session summary added".to_string())
}
}
pub struct GenerateStandupReportHandler;
#[async_trait]
impl McpTool for GenerateStandupReportHandler {
fn name(&self) -> &'static str {
"generate_standup_report"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<GenerateStandupReportTool>(
"generate_standup_report",
"Execute generate_standup_report",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: GenerateStandupReportTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let cutoff = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
.saturating_sub(req.hours_lookback * 3600);
let tasks = state
.tasks
.read()
.into_iter()
.filter(|t| t.updated_at >= cutoff)
.collect::<Vec<_>>();
let changes = state
.ledger
.read()
.into_iter()
.filter(|c| c.timestamp >= cutoff)
.collect::<Vec<_>>();
let summaries = state
.session_summaries
.read()
.into_iter()
.filter(|s| s.namespace == req.namespace && s.timestamp >= cutoff)
.collect::<Vec<_>>();
let report = serde_json::json!({ "tasks_updated": tasks, "code_changes": changes, "session_summaries": summaries });
Ok(report.to_string())
}
}
+456
View File
@@ -0,0 +1,456 @@
use crate::models::*;
use crate::router::McpTool;
use crate::state::MemoryState;
use crate::tools::*;
use async_trait::async_trait;
use serde_json::Value;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
pub struct AddTaskHandler;
#[async_trait]
impl McpTool for AddTaskHandler {
fn name(&self) -> &'static str {
"add_task"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<AddTaskTool>("add_task", "Execute add_task")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: AddTaskTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let task_id = uuid::Uuid::new_v4().to_string();
let parent_id = req.parent_id.clone();
let deps = req.dependencies.clone().unwrap_or_default();
let task = Task {
id: task_id.clone(),
title: req.title,
status: "pending".to_string(),
description: req.description,
created_at: now,
updated_at: now,
git_branch: req.git_branch,
parent_id,
dependencies: deps,
acceptance_criteria: vec![],
};
if let Ok(idx) = state.search_index.read() {
drop(idx.index_task(&task));
}
state.tasks.modify(|tasks| {
tasks.push(task);
});
Ok(format!("Task added with ID: {}", task_id).to_string())
}
}
pub struct DeleteTaskHandler;
#[async_trait]
impl McpTool for DeleteTaskHandler {
fn name(&self) -> &'static str {
"delete_task"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<DeleteTaskTool>("delete_task", "Execute delete_task")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: DeleteTaskTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut deleted_count = 0;
state.tasks.modify(|tasks| {
let initial_len = tasks.len();
// Collect IDs of tasks to delete (this task + all its recursive children)
let mut to_delete = std::collections::HashSet::new();
to_delete.insert(req.id.clone());
let mut children_map: std::collections::HashMap<String, Vec<String>> =
std::collections::HashMap::new();
for t in tasks.iter() {
if let Some(pid) = &t.parent_id {
children_map
.entry(pid.clone())
.or_default()
.push(t.id.clone());
}
}
let mut queue = std::collections::VecDeque::new();
queue.push_back(req.id.clone());
while let Some(curr) = queue.pop_front() {
if to_delete.insert(curr.clone())
&& let Some(children) = children_map.get(&curr)
{
queue.extend(children.iter().cloned());
}
}
tasks.retain(|t| !to_delete.contains(&t.id));
deleted_count = initial_len - tasks.len();
});
if deleted_count > 0 {
Ok(vec![
format!("Deleted task and its children ({} total).", deleted_count).to_string(),
][0]
.clone())
} else {
Ok("Task not found.".to_string())
}
}
}
pub struct UpdateTaskStatusHandler;
#[async_trait]
impl McpTool for UpdateTaskStatusHandler {
fn name(&self) -> &'static str {
"update_task_status"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<UpdateTaskStatusTool>(
"update_task_status",
"Execute update_task_status",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: UpdateTaskStatusTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut found = false;
let mut blocked = false;
let mut blocker_details = String::new();
let target_status = req.status.to_lowercase();
state.tasks.modify(|tasks| {
// Find target task
let mut target_id = String::new();
if let Some(t) = tasks.iter().find(|t| t.id == req.id || t.title == req.id) {
target_id = t.id.clone();
}
if target_id.is_empty() {
return;
}
found = true;
if target_status == "done" || target_status == "completed" {
// 1. Check Acceptance Criteria
if let Some(t) = tasks.iter().find(|t| t.id == target_id)
&& t.acceptance_criteria.iter().any(|c| !c.is_met)
{
blocked = true;
blocker_details = "Unmet acceptance criteria exist.".to_string();
}
// 2. Check dependencies
if !blocked {
let mut uncompleted_deps = Vec::new();
if let Some(t) = tasks.iter().find(|t| t.id == target_id) {
for dep_id in &t.dependencies {
if let Some(dep_task) = tasks.iter().find(|dt| dt.id == *dep_id)
&& dep_task.status != "completed"
&& dep_task.status != "done"
{
uncompleted_deps.push(dep_task.title.clone());
}
}
}
if !uncompleted_deps.is_empty() {
blocked = true;
blocker_details =
format!("Blocked by dependencies: {}", uncompleted_deps.join(", "));
}
}
// 3. Check child tasks
if !blocked {
let mut uncompleted_children = Vec::new();
for child in tasks
.iter()
.filter(|t| t.parent_id.as_ref() == Some(&target_id))
{
if child.status != "completed" && child.status != "done" {
uncompleted_children.push(child.title.clone());
}
}
if !uncompleted_children.is_empty() {
blocked = true;
blocker_details = format!(
"Blocked by child tasks: {}",
uncompleted_children.join(", ")
);
}
}
}
if !blocked {
// Apply update
if let Some(t) = tasks.iter_mut().find(|t| t.id == target_id) {
t.status = target_status.clone();
t.updated_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
}
// Cascade cancellation to children
if target_status == "cancelled" || target_status == "abandoned" {
let mut children_map: std::collections::HashMap<String, Vec<usize>> =
std::collections::HashMap::new();
for (idx, t) in tasks.iter().enumerate() {
if let Some(pid) = &t.parent_id {
children_map.entry(pid.clone()).or_default().push(idx);
}
}
let mut queue = std::collections::VecDeque::new();
queue.push_back(target_id.clone());
while let Some(curr) = queue.pop_front() {
if let Some(child_indices) = children_map.get(&curr) {
for &idx in child_indices {
if tasks[idx].status != "completed"
&& tasks[idx].status != target_status
{
tasks[idx].status = target_status.clone();
queue.push_back(tasks[idx].id.clone());
}
}
}
}
}
}
});
if blocked {
Ok(format!(
"Error: Cannot transition task. {}",
blocker_details
))
} else if found {
Ok("Task status updated.".to_string())
} else {
Ok("Task not found.".to_string())
}
}
}
pub struct ListActiveTasksHandler;
#[async_trait]
impl McpTool for ListActiveTasksHandler {
fn name(&self) -> &'static str {
"list_active_tasks"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ListActiveTasksTool>(
"list_active_tasks",
"Execute list_active_tasks",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ListActiveTasksTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut tasks = state.tasks.read();
tasks.retain(|t| t.status != "done");
if let Some(branch) = req.git_branch {
tasks.retain(|t| {
t.git_branch.is_none() || t.git_branch.as_deref() == Some(branch.as_str())
});
}
let data = serde_json::to_string(&tasks).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct SetAcceptanceCriteriaHandler;
#[async_trait]
impl McpTool for SetAcceptanceCriteriaHandler {
fn name(&self) -> &'static str {
"set_acceptance_criteria"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<SetAcceptanceCriteriaTool>(
"set_acceptance_criteria",
"Execute set_acceptance_criteria",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: SetAcceptanceCriteriaTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut success = false;
state.tasks.modify(|tasks| {
if let Some(task) = tasks.iter_mut().rev().find(|t| t.title == req.task_title) {
task.acceptance_criteria = req
.criteria
.into_iter()
.map(|desc| crate::models::AcceptanceCriteria {
id: uuid::Uuid::new_v4().to_string(),
description: desc,
is_met: false,
})
.collect();
task.updated_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
success = true;
}
});
if success {
Ok(vec!["Acceptance criteria set successfully.".to_string()][0].clone())
} else {
Ok("Task not found.".to_string())
}
}
}
pub struct VerifyAcceptanceCriteriaHandler;
#[async_trait]
impl McpTool for VerifyAcceptanceCriteriaHandler {
fn name(&self) -> &'static str {
"verify_acceptance_criteria"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<VerifyAcceptanceCriteriaTool>(
"verify_acceptance_criteria",
"Execute verify_acceptance_criteria",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: VerifyAcceptanceCriteriaTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut success = false;
let mut already_met = false;
state.tasks.modify(|tasks| {
if let Some(task) = tasks.iter_mut().find(|t| t.id == req.task_id)
&& let Some(ac) = task
.acceptance_criteria
.iter_mut()
.find(|c| c.id == req.criteria || c.description == req.criteria)
{
if ac.is_met {
already_met = true;
} else {
ac.is_met = true;
success = true;
task.updated_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
}
}
});
if success {
Ok(vec![format!(
"Acceptance criteria verified with proof: {}",
req.proof
)][0]
.clone())
} else if already_met {
Ok("Acceptance criteria was already met.".to_string())
} else {
Ok(vec!["Acceptance criteria or task not found.".to_string()][0].clone())
}
}
}
pub struct AddMilestoneHandler;
#[async_trait]
impl McpTool for AddMilestoneHandler {
fn name(&self) -> &'static str {
"add_milestone"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<AddMilestoneTool>("add_milestone", "Execute add_milestone")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: AddMilestoneTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.milestones.modify(|ms| {
ms.push(crate::models::Milestone {
id: uuid::Uuid::new_v4().to_string(),
title: req.title,
status: "pending".to_string(),
namespace: req.namespace,
target_date: None,
})
});
Ok("Milestone added".to_string())
}
}
pub struct UpdateMilestoneHandler;
#[async_trait]
impl McpTool for UpdateMilestoneHandler {
fn name(&self) -> &'static str {
"update_milestone"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<UpdateMilestoneTool>("update_milestone", "Execute update_milestone")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: UpdateMilestoneTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut found = false;
state.milestones.modify(|ms| {
for m in ms.iter_mut() {
if m.id == req.id {
m.status = req.status.clone();
found = true;
break;
}
}
});
if found {
Ok("Milestone updated".to_string())
} else {
Ok("Milestone not found".to_string())
}
}
}
pub struct ListMilestonesHandler;
#[async_trait]
impl McpTool for ListMilestonesHandler {
fn name(&self) -> &'static str {
"list_milestones"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ListMilestonesTool>("list_milestones", "Execute list_milestones")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ListMilestonesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut items = state.milestones.read();
if let Some(ns) = req.namespace {
items.retain(|i| i.namespace == ns);
}
let data = serde_json::to_string(&items).unwrap_or_default();
Ok(data.to_string())
}
}
+9
View File
@@ -0,0 +1,9 @@
pub fn contains_ignore_ascii_case(haystack: &str, needle: &str) -> bool {
if needle.is_empty() {
return true;
}
haystack
.as_bytes()
.windows(needle.len())
.any(|w| w.eq_ignore_ascii_case(needle.as_bytes()))
}
+348
View File
@@ -0,0 +1,348 @@
use crate::models::*;
use crate::router::McpTool;
use crate::state::MemoryState;
use crate::tools::*;
use async_trait::async_trait;
use serde_json::Value;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
pub struct PinFileHandler;
#[async_trait]
impl McpTool for PinFileHandler {
fn name(&self) -> &'static str {
"pin_file"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<PinFileTool>("pin_file", "Execute pin_file")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: PinFileTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.pinned_files.modify(|pinned| {
pinned.retain(|p| !(p.namespace == req.namespace && p.file_path == req.file_path));
pinned.push(crate::models::PinnedFile {
namespace: req.namespace,
file_path: req.file_path,
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
git_branch: req.git_branch,
});
});
Ok("File pinned".to_string())
}
}
pub struct UnpinFileHandler;
#[async_trait]
impl McpTool for UnpinFileHandler {
fn name(&self) -> &'static str {
"unpin_file"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<UnpinFileTool>("unpin_file", "Execute unpin_file")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: UnpinFileTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state.pinned_files.modify(|pinned| {
pinned.retain(|p| !(p.namespace == req.namespace && p.file_path == req.file_path))
});
Ok("File unpinned".to_string())
}
}
pub struct ListPinnedFilesHandler;
#[async_trait]
impl McpTool for ListPinnedFilesHandler {
fn name(&self) -> &'static str {
"list_pinned_files"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ListPinnedFilesTool>(
"list_pinned_files",
"Execute list_pinned_files",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ListPinnedFilesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut pinned = state.pinned_files.read();
if let Some(ns) = req.namespace {
pinned.retain(|p| p.namespace == ns);
}
if let Some(branch) = req.git_branch {
pinned.retain(|p| {
p.git_branch.is_none() || p.git_branch.as_deref() == Some(branch.as_str())
});
}
let data = serde_json::to_string(&pinned).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct StoreSnippetHandler;
#[async_trait]
impl McpTool for StoreSnippetHandler {
fn name(&self) -> &'static str {
"store_snippet"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<StoreSnippetTool>("store_snippet", "Execute store_snippet")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: StoreSnippetTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let snippet = Snippet {
name: req.name.clone(),
language: req.language,
code: req.code,
description: req.description,
updated_at: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
};
let s_clone = snippet.clone();
state.snippets.modify(|snippets| {
snippets.retain(|s| s.name != req.name);
snippets.push(s_clone);
});
if let Ok(idx) = state.search_index.read() {
drop(idx.index_snippet(&snippet));
}
Ok(format!("Snippet '{}' stored.", req.name).to_string())
}
}
pub struct SearchSnippetsHandler;
#[async_trait]
impl McpTool for SearchSnippetsHandler {
fn name(&self) -> &'static str {
"search_snippets"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<SearchSnippetsTool>("search_snippets", "Execute search_snippets")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: SearchSnippetsTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let query = req.query.to_lowercase();
let snippets = state.snippets.read();
let mut results = Vec::new();
for s in snippets {
if contains_ignore_ascii_case(&s.name, &query)
|| contains_ignore_ascii_case(&s.description, &query)
|| contains_ignore_ascii_case(&s.language, &query)
{
results.push(s);
}
}
let data = serde_json::to_string(&results).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct DeleteSnippetHandler;
#[async_trait]
impl McpTool for DeleteSnippetHandler {
fn name(&self) -> &'static str {
"delete_snippet"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<DeleteSnippetTool>("delete_snippet", "Execute delete_snippet")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: DeleteSnippetTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut deleted = false;
state.snippets.modify(|snippets| {
let orig = snippets.len();
snippets.retain(|s| s.name != req.name);
deleted = snippets.len() < orig;
});
if deleted {
Ok("Snippet deleted.".to_string())
} else {
Ok("Snippet not found.".to_string())
}
}
}
pub struct SaveContextWorkspaceHandler;
#[async_trait]
impl McpTool for SaveContextWorkspaceHandler {
fn name(&self) -> &'static str {
"save_context_workspace"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<SaveContextWorkspaceTool>(
"save_context_workspace",
"Execute save_context_workspace",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: SaveContextWorkspaceTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.context_workspaces.modify(|ws| {
ws.retain(|w| !(w.namespace == req.namespace && w.name == req.name));
ws.push(crate::models::ContextWorkspace {
namespace: req.namespace,
name: req.name,
pinned_files: req.pinned_files,
active_task_ids: req.active_task_ids,
saved_at: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
});
});
Ok("Context workspace saved".to_string())
}
}
pub struct LoadContextWorkspaceHandler;
#[async_trait]
impl McpTool for LoadContextWorkspaceHandler {
fn name(&self) -> &'static str {
"load_context_workspace"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<LoadContextWorkspaceTool>(
"load_context_workspace",
"Execute load_context_workspace",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: LoadContextWorkspaceTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut ws = state.context_workspaces.read();
ws.retain(|w| w.namespace == req.namespace && w.name == req.name);
let data = serde_json::to_string(&ws.first()).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct ListContextWorkspacesHandler;
#[async_trait]
impl McpTool for ListContextWorkspacesHandler {
fn name(&self) -> &'static str {
"list_context_workspaces"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ListContextWorkspacesTool>(
"list_context_workspaces",
"Execute list_context_workspaces",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ListContextWorkspacesTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut ws = state.context_workspaces.read();
ws.retain(|w| w.namespace == req.namespace);
let data = serde_json::to_string(&ws).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct AddPrChecklistItemHandler;
#[async_trait]
impl McpTool for AddPrChecklistItemHandler {
fn name(&self) -> &'static str {
"add_pr_checklist_item"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<AddPrChecklistItemTool>(
"add_pr_checklist_item",
"Execute add_pr_checklist_item",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: AddPrChecklistItemTool =
serde_json::from_value(args).map_err(|e| e.to_string())?;
state.pr_checklists.modify(|items| {
items.push(crate::models::PrChecklistItem {
namespace: req.namespace,
id: uuid::Uuid::new_v4().to_string(),
description: req.description,
})
});
Ok("PR checklist item added".to_string())
}
}
pub struct GetPrChecklistHandler;
#[async_trait]
impl McpTool for GetPrChecklistHandler {
fn name(&self) -> &'static str {
"get_pr_checklist"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<GetPrChecklistTool>("get_pr_checklist", "Execute get_pr_checklist")
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: GetPrChecklistTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut items = state.pr_checklists.read();
items.retain(|i| i.namespace == req.namespace);
let data = serde_json::to_string(&items).unwrap_or_default();
Ok(data.to_string())
}
}
pub struct ClearPrChecklistHandler;
#[async_trait]
impl McpTool for ClearPrChecklistHandler {
fn name(&self) -> &'static str {
"clear_pr_checklist"
}
fn schema(&self) -> Value {
crate::mcp::tool_def::<ClearPrChecklistTool>(
"clear_pr_checklist",
"Execute clear_pr_checklist",
)
}
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
let req: ClearPrChecklistTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
state
.pr_checklists
.modify(|items| items.retain(|i| i.namespace != req.namespace));
Ok("PR checklist cleared".to_string())
}
}
use crate::handlers_v2::utils::*;
+33 -17
View File
@@ -4,8 +4,10 @@
)] )]
mod handlers; mod handlers;
mod handlers_v2;
mod mcp; mod mcp;
mod models; mod models;
mod router;
mod search; mod search;
mod state; mod state;
mod store; mod store;
@@ -218,9 +220,7 @@ async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::E
state.rebuild_index().await; state.rebuild_index().await;
tokio::spawn(index_committer_worker(Arc::clone(&state))); tokio::spawn(index_committer_worker(Arc::clone(&state)));
let app_state = Arc::new(AppState { let app_state = Arc::new(AppState {
handler: Arc::new(MemoryHandler { handler: Arc::new(MemoryHandler::new(Arc::clone(&state))),
state: Arc::clone(&state),
}),
clients: RwLock::new(HashMap::new()), clients: RwLock::new(HashMap::new()),
next_id: AtomicUsize::new(1), next_id: AtomicUsize::new(1),
}); });
@@ -414,7 +414,8 @@ async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::E
let log_path = dirs::home_dir() let log_path = dirs::home_dir()
.unwrap_or_default() .unwrap_or_default()
.join(".gemini/mcp_memory/daemon_error.log"); .join(".gemini/mcp_memory/daemon_error.log");
let _ = tokio::fs::write(&log_path, format!("Failed to bind to {}: {}\n", addr, e)).await; let _ =
tokio::fs::write(&log_path, format!("Failed to bind to {}: {}\n", addr, e)).await;
return Ok(()); return Ok(());
} }
}; };
@@ -486,7 +487,8 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
if client_type == "proxy" { if client_type == "proxy" {
// Send activity broadcast to UI clients // Send activity broadcast to UI clients
if let Some(method) = payload.get("method").and_then(|m| m.as_str()) if let Some(method) = payload.get("method").and_then(|m| m.as_str())
&& method == "tools/call" { && method == "tools/call"
{
let name = payload let name = payload
.get("params") .get("params")
.and_then(|p| p.get("name")) .and_then(|p| p.get("name"))
@@ -521,7 +523,7 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
// Process MCP request // Process MCP request
if let Some(response) = handler.handle_request(payload).await { if let Some(response) = handler.handle_request(payload).await {
let res_str = serde_json::to_string(&response).unwrap(); let res_str = serde_json::to_string(&response).unwrap_or_default();
let tx_opt = state_clone let tx_opt = state_clone
.clients .clients
.read() .read()
@@ -576,7 +578,11 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
impl Drop for SessionCleanup { impl Drop for SessionCleanup {
fn drop(&mut self) { fn drop(&mut self) {
self.state.clients.write().unwrap().remove(&self.session_id); self.state
.clients
.write()
.unwrap_or_else(|e| e.into_inner())
.remove(&self.session_id);
if let Some(task) = self.send_task.take() { if let Some(task) = self.send_task.take() {
task.abort(); task.abort();
} }
@@ -640,7 +646,13 @@ async fn nvim_telemetry_handler(
}); });
let msg_str = ws_msg.to_string(); let msg_str = ws_msg.to_string();
let senders: Vec<_> = state.clients.read().unwrap().values().cloned().collect(); let senders: Vec<_> = state
.clients
.read()
.unwrap_or_else(|e| e.into_inner())
.values()
.cloned()
.collect();
for tx in senders { for tx in senders {
let _ = tx.try_send(msg_str.clone()); let _ = tx.try_send(msg_str.clone());
} }
@@ -698,9 +710,12 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let mut cmd = std::process::Command::new("curl"); let mut cmd = std::process::Command::new("curl");
cmd.arg("-k").arg("-X").arg("POST"); cmd.arg("-k").arg("-X").arg("POST");
if !token.is_empty() { if !token.is_empty() {
cmd.arg("-H").arg(format!("Authorization: Bearer {}", token.trim())); cmd.arg("-H")
.arg(format!("Authorization: Bearer {}", token.trim()));
} }
let _ = cmd.arg(format!("https://127.0.0.1:{}/shutdown", port)).output(); let _ = cmd
.arg(format!("https://127.0.0.1:{}/shutdown", port))
.output();
println!("Sent shutdown request to server."); println!("Sent shutdown request to server.");
return Ok(()); return Ok(());
} }
@@ -711,9 +726,12 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let mut cmd = std::process::Command::new("curl"); let mut cmd = std::process::Command::new("curl");
cmd.arg("-k").arg("-X").arg("POST"); cmd.arg("-k").arg("-X").arg("POST");
if !token.is_empty() { if !token.is_empty() {
cmd.arg("-H").arg(format!("Authorization: Bearer {}", token.trim())); cmd.arg("-H")
.arg(format!("Authorization: Bearer {}", token.trim()));
} }
let _ = cmd.arg(format!("https://127.0.0.1:{}/shutdown", port)).output(); let _ = cmd
.arg(format!("https://127.0.0.1:{}/shutdown", port))
.output();
println!("Sent shutdown request to existing server. Waiting for it to exit..."); println!("Sent shutdown request to existing server. Waiting for it to exit...");
std::thread::sleep(std::time::Duration::from_millis(1500)); std::thread::sleep(std::time::Duration::from_millis(1500));
return Ok(()); return Ok(());
@@ -780,12 +798,10 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let json_path = base.join(file_name); let json_path = base.join(file_name);
if json_path.exists() if json_path.exists()
&& let Ok(data) = fs::read(&json_path) && let Ok(data) = fs::read(&json_path)
&& serde_json::from_slice::<serde_json::Value>(&data).is_ok() { && serde_json::from_slice::<serde_json::Value>(&data).is_ok()
{
table.insert(*key, data.as_slice()).unwrap(); table.insert(*key, data.as_slice()).unwrap();
let _ = fs::rename( let _ = fs::rename(&json_path, json_path.with_extension("json.migrated"));
&json_path,
json_path.with_extension("json.migrated"),
);
} }
} }
} }
+173
View File
@@ -0,0 +1,173 @@
import os
import re
GROUPS = {
"graph": [
"query_graph_path", "create_entities", "create_relations", "add_observations",
"delete_entities", "delete_observations", "delete_relations", "read_graph",
"search_nodes", "open_nodes", "visualize_graph", "condense_entity",
"merge_entities", "find_orphans"
],
"tasks": [
"add_task", "delete_task", "update_task_status", "list_active_tasks",
"set_acceptance_criteria", "verify_acceptance_criteria",
"add_milestone", "update_milestone", "list_milestones"
],
"notes": [
"add_sticky_note", "read_sticky_notes", "delete_sticky_note", "clear_sticky_notes",
"leave_handoff_memo", "read_handoff_memos", "clear_handoff_memos",
"add_session_summary", "generate_standup_report"
],
"meta": [
"log_decision", "query_decisions", "log_error_fix", "search_error_fixes",
"log_code_change", "query_recent_changes", "learn_preference", "read_preferences",
"log_tech_debt", "resolve_tech_debt", "list_tech_debt", "omni_search", "get_project_health"
],
"env": [
"update_env_fingerprint", "read_env_fingerprint", "log_env_requirement",
"register_environment", "get_environment_details"
],
"workspaces": [
"pin_file", "unpin_file", "list_pinned_files", "store_snippet", "search_snippets",
"delete_snippet", "save_context_workspace", "load_context_workspace",
"list_context_workspaces", "add_pr_checklist_item", "get_pr_checklist", "clear_pr_checklist"
]
}
def to_camel_case(snake_str):
components = snake_str.split('_')
return "".join(x.title() for x in components)
def parse_rust_match(file_path):
with open(file_path, "r", encoding="utf-8") as f:
lines = f.readlines()
start_idx = -1
for i, line in enumerate(lines):
if "let result: Result<String, String> = match name {" in line:
start_idx = i
break
if start_idx == -1:
return {}
brace_depth = 1
i = start_idx + 1
tools = {}
current_tool = None
current_body = []
while i < len(lines):
line = lines[i]
if brace_depth == 1 and "=>" in line and '"' in line:
parts = line.strip().split('"')
if len(parts) >= 3:
tool_name = parts[1]
current_tool = tool_name
current_body = []
# Don't add the "name" => { line
if current_tool is not None and not (brace_depth == 1 and "=>" in line and '"' in line):
# check if this line closes the block
next_depth = brace_depth + line.count('{') - line.count('}')
if next_depth == 1 and current_tool is not None:
# This is the closing brace
tools[current_tool] = "".join(current_body)
current_tool = None
else:
current_body.append(line)
brace_depth += line.count('{')
brace_depth -= line.count('}')
if brace_depth == 0:
break
i += 1
return tools
def transform_body(body):
# Transform parse_tool!
body = re.sub(
r'let req = parse_tool!\(args, id, ([^)]+)\);',
r'let req: \1 = serde_json::from_value(args).map_err(|e| e.to_string())?;',
body
)
# Transform handle_list_with_namespace!
def repl_handle_list(m):
store = m.group(1)
tool_type = m.group(2)
return f"""
let req: {tool_type} = serde_json::from_value(args).map_err(|e| e.to_string())?;
let mut items = state.{store}.read();
if let Some(ns) = req.namespace {{
items.retain(|i| i.namespace == ns);
}}
let data = serde_json::to_string(&items).unwrap_or_default();
return Ok(data.to_string());
"""
body = re.sub(
r'return handle_list_with_namespace!\(self, ([^,]+), ([^,]+), args, id\);',
repl_handle_list,
body
)
# Replace self.state with state
body = body.replace("self.state.", "state.")
return body
tools = parse_rust_match("server/src/handlers.rs")
for group, tool_names in GROUPS.items():
file_path = f"server/src/handlers_v2/{group}.rs"
with open(file_path, "w", encoding="utf-8") as f:
f.write("use crate::router::McpTool;\n")
f.write("use crate::state::MemoryState;\n")
f.write("use crate::tools::*;\n")
f.write("use async_trait::async_trait;\n")
f.write("use serde_json::Value;\n")
f.write("use std::sync::Arc;\n")
f.write("use std::time::{SystemTime, UNIX_EPOCH};\n\n")
for name in tool_names:
if name not in tools:
continue
body = tools[name]
# special case for query_graph_path which we already wrote properly?
# actually we will just overwrite it with the transformed body
body = transform_body(body)
struct_name = to_camel_case(name) + "Handler"
tool_type = to_camel_case(name) + "Tool"
f.write(f"pub struct {struct_name};\n\n")
f.write(f"#[async_trait]\n")
f.write(f"impl McpTool for {struct_name} {{\n")
f.write(f" fn name(&self) -> &'static str {{\n")
f.write(f' "{name}"\n')
f.write(f" }}\n\n")
f.write(f" fn schema(&self) -> Value {{\n")
# For schema description we can just put a generic one or extract it.
# I will use a generic one for now, or you can extract it from tools/list.
f.write(f' crate::mcp::tool_def::<{tool_type}>(\n')
f.write(f' "{name}",\n')
f.write(f' "Execute {name}",\n')
f.write(f' )\n')
f.write(f" }}\n\n")
f.write(f" async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {{\n")
f.write(body)
f.write(f" }}\n")
f.write(f"}}\n\n")
print("Generated handlers_v2 modules")
# generate mod.rs
with open("server/src/handlers_v2/mod.rs", "w", encoding="utf-8") as f:
for group in GROUPS.keys():
f.write(f"pub mod {group};\n")
+11
View File
@@ -0,0 +1,11 @@
use std::fs;
use std::io::Write;
fn main() {
let content = fs::read_to_string("server/src/handlers.rs").unwrap();
println!("Read {} bytes", content.len());
// Find the match name { block
let match_start = content.find("match name {").unwrap();
// naive extraction
println!("Found match block at {}", match_start);
}
+82
View File
@@ -0,0 +1,82 @@
import re
import os
with open("server/src/handlers.rs", "r", encoding="utf-8") as f:
lines = f.readlines()
list_start = -1
for i, line in enumerate(lines):
if '"tools/list" => {' in line:
list_start = i
break
# Find end of tools/call
call_start = -1
for i in range(list_start, len(lines)):
if '"tools/call" => {' in line:
call_start = i
break
# Find end of tools/call
# Match brace depth from call_start
brace_depth = 1
call_end = -1
for i in range(call_start + 1, len(lines)):
brace_depth += lines[i].count('{')
brace_depth -= lines[i].count('}')
if brace_depth == 0:
call_end = i
break
# replacement block
replacement = """ "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 }),
))
}
"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))
}
}
}
"""
new_lines = lines[:list_start] + [replacement] + lines[call_end+1:]
with open("server/src/handlers.rs", "w", encoding="utf-8") as f:
f.writelines(new_lines)
print("tools/list and tools/call replaced.")
+106
View File
@@ -0,0 +1,106 @@
import re
import os
with open("server/src/handlers.rs", "r", encoding="utf-8") as f:
content = f.read()
# Replace MemoryHandler struct
struct_pattern = r'pub struct MemoryHandler \{\s*pub state: Arc<MemoryState>,\s*\}'
new_struct = """use crate::router::McpTool;
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_v2::$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 }
}
"""
content = re.sub(struct_pattern, new_struct, content)
content = content.replace("impl MemoryHandler {\n pub async fn handle_request", " pub async fn handle_request")
with open("server/src/handlers.rs", "w", encoding="utf-8") as f:
f.write(content)
print("MemoryHandler struct updated.")
+16
View File
@@ -0,0 +1,16 @@
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>;
}
+5 -1
View File
@@ -106,7 +106,11 @@ impl MemoryIndex {
Ok(()) Ok(())
}) })
.await .await
.unwrap_or_else(|_| Err(tantivy::TantivyError::SystemError("Commit task panicked".to_string()))) .unwrap_or_else(|_| {
Err(tantivy::TantivyError::SystemError(
"Commit task panicked".to_string(),
))
})
} }
pub fn search( pub fn search(
-1
View File
@@ -62,7 +62,6 @@ impl MemoryState {
pub async fn rebuild_index(&self) { pub async fn rebuild_index(&self) {
if let Ok(new_idx) = MemoryIndex::new(&self.base_dir) { if let Ok(new_idx) = MemoryIndex::new(&self.base_dir) {
let entities = self.graph.read().entities.into_values().collect(); let entities = self.graph.read().entities.into_values().collect();
let tasks = self.tasks.read(); let tasks = self.tasks.read();
let snippets = self.snippets.read(); let snippets = self.snippets.read();
+19 -18
View File
@@ -22,44 +22,45 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + Sync + 'static>
tokio::spawn(async move { tokio::spawn(async move {
while rx.recv().await.is_some() { while rx.recv().await.is_some() {
// Drain any other pending notifications so we batch writes // Drain any other pending notifications so we batch writes
while let Ok(_) = rx.try_recv() {} while rx.try_recv().is_ok() {}
let db_inner = db_clone.clone(); let db_inner = db_clone.clone();
let key_inner = key_clone.clone(); let key_inner = key_clone.clone();
let json_data = { let json_data = {
let lock = cache_clone.read().unwrap(); let lock = cache_clone.read().unwrap_or_else(|e| e.into_inner());
serde_json::to_vec(&*lock).unwrap() serde_json::to_vec(&*lock).unwrap_or_default()
}; };
let _ = tokio::task::spawn_blocking(move || { let _ = tokio::task::spawn_blocking(move || {
let write_txn = db_inner.begin_write().unwrap(); if let Ok(write_txn) = db_inner.begin_write() {
{ if let Ok(mut table) = write_txn.open_table(STORE_TABLE) {
let mut table = write_txn.open_table(STORE_TABLE).unwrap(); let _ = table.insert(key_inner.as_str(), json_data.as_slice());
table.insert(key_inner.as_str(), json_data.as_slice()).unwrap();
} }
write_txn.commit().unwrap(); let _ = write_txn.commit();
}).await; }
})
.await;
} }
}); });
Self { Self { cache, tx }
cache,
tx,
}
} }
fn load_from_db(key: &str, db: &Database) -> T { fn load_from_db(key: &str, db: &Database) -> T {
let read_txn = db.begin_read().unwrap(); let Ok(read_txn) = db.begin_read() else {
return T::default();
};
if let Ok(table) = read_txn.open_table(STORE_TABLE) if let Ok(table) = read_txn.open_table(STORE_TABLE)
&& let Ok(Some(value)) = table.get(key) && let Ok(Some(value)) = table.get(key)
&& let Ok(parsed) = serde_json::from_slice::<T>(value.value()) { && let Ok(parsed) = serde_json::from_slice::<T>(value.value())
{
return parsed; return parsed;
} }
T::default() T::default()
} }
pub fn read(&self) -> T { pub fn read(&self) -> T {
let lock = self.cache.read().unwrap(); let lock = self.cache.read().unwrap_or_else(|e| e.into_inner());
lock.clone() lock.clone()
} }
@@ -67,13 +68,13 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + Sync + 'static>
where where
F: FnOnce(&T) -> R, F: FnOnce(&T) -> R,
{ {
let lock = self.cache.read().unwrap(); let lock = self.cache.read().unwrap_or_else(|e| e.into_inner());
f(&lock) f(&lock)
} }
pub fn modify<F: FnOnce(&mut T)>(&self, f: F) { pub fn modify<F: FnOnce(&mut T)>(&self, f: F) {
{ {
let mut lock = self.cache.write().unwrap(); let mut lock = self.cache.write().unwrap_or_else(|e| e.into_inner());
f(&mut lock); f(&mut lock);
} }
let _ = self.tx.try_send(()); let _ = self.tx.try_send(());
+1
View File
@@ -314,6 +314,7 @@ pub struct AddSessionSummaryTool {
/// Get a timeline of major project events. /// Get a timeline of major project events.
#[derive(Debug, Deserialize, Serialize, JsonSchema)] #[derive(Debug, Deserialize, Serialize, JsonSchema)]
#[allow(dead_code)]
pub struct GetProjectTimelineTool { pub struct GetProjectTimelineTool {
/// Optional namespace to restrict the timeline to. /// Optional namespace to restrict the timeline to.
pub namespace: Option<String>, pub namespace: Option<String>,
+11 -5
View File
@@ -2,10 +2,13 @@ use std::collections::HashSet;
#[test] #[test]
fn test_eager_tools_parity() { fn test_eager_tools_parity() {
// 1. Read handlers.rs to get memory tools // 1. Read handlers_v2/*.rs to get memory tools
let memory_source =
std::fs::read_to_string("src/handlers.rs").expect("Failed to read handlers.rs");
let mut memory_tools = HashSet::new(); let mut memory_tools = HashSet::new();
let entries = std::fs::read_dir("src/handlers_v2").expect("Failed to read handlers_v2 dir");
for entry in entries {
let entry = entry.unwrap();
if entry.path().extension().unwrap_or_default() == "rs" {
let memory_source = std::fs::read_to_string(entry.path()).unwrap();
let parts: Vec<&str> = memory_source.split("crate::mcp::tool_def").collect(); let parts: Vec<&str> = memory_source.split("crate::mcp::tool_def").collect();
for part in parts.iter().skip(1) { for part in parts.iter().skip(1) {
if let Some(start) = part.find("\"") { if let Some(start) = part.find("\"") {
@@ -15,9 +18,11 @@ fn test_eager_tools_parity() {
} }
} }
} }
}
}
assert!( assert!(
!memory_tools.is_empty(), !memory_tools.is_empty(),
"Could not find memory tools in handlers.rs" "Could not find memory tools in handlers_v2 directory"
); );
// 2. Read nvim-core/src/lib.rs to get nvim tools // 2. Read nvim-core/src/lib.rs to get nvim tools
@@ -26,7 +31,8 @@ fn test_eager_tools_parity() {
let mut nvim_tools = HashSet::new(); let mut nvim_tools = HashSet::new();
for line in nvim_source.lines() { for line in nvim_source.lines() {
if line.contains("\"name\": \"nvim_") if line.contains("\"name\": \"nvim_")
&& let Some(start) = line.find("\"name\": \"") { && let Some(start) = line.find("\"name\": \"")
{
let rest = &line[start + 9..]; let rest = &line[start + 9..];
if let Some(end) = rest.find("\"") { if let Some(end) = rest.find("\"") {
nvim_tools.insert(rest[..end].to_string()); nvim_tools.insert(rest[..end].to_string());
+1 -5
View File
@@ -15,10 +15,6 @@ fn main() {
.and_then(|out| String::from_utf8(out.stdout).ok()) .and_then(|out| String::from_utf8(out.stdout).ok())
.unwrap_or_else(|| "unknown".to_string()); .unwrap_or_else(|| "unknown".to_string());
let version = format!( let version = format!("{} ({})", git_date.trim(), git_hash.trim());
"{} ({})",
git_date.trim(),
git_hash.trim()
);
println!("cargo:rustc-env=APP_VERSION={}", version); println!("cargo:rustc-env=APP_VERSION={}", version);
} }
+22 -7
View File
@@ -1,24 +1,39 @@
use std::sync::LazyLock;
use regex::Regex; use regex::Regex;
use std::sync::LazyLock;
static ID_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r#""id"\s*:\s*([^,}]+)"#).unwrap()); static ID_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r#""id"\s*:\s*([^,}]+)"#).unwrap());
static METHOD_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r#""method"\s*:\s*"([^"]+)""#).unwrap()); static METHOD_REGEX: LazyLock<Regex> =
static TOOL_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r#""name"\s*:\s*"([^"]+)""#).unwrap()); LazyLock::new(|| Regex::new(r#""method"\s*:\s*"([^"]+)""#).unwrap());
static TOOL_REGEX: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r#""name"\s*:\s*"([^"]+)""#).unwrap());
static ERROR_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r#""error"\s*:\s*\{"#).unwrap()); static ERROR_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r#""error"\s*:\s*\{"#).unwrap());
static IS_ERROR_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r#""isError"\s*:\s*true"#).unwrap()); static IS_ERROR_REGEX: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r#""isError"\s*:\s*true"#).unwrap());
pub fn extract_log_prefix(json_str: &str, is_response: bool) -> String { pub fn extract_log_prefix(json_str: &str, is_response: bool) -> String {
let id = ID_REGEX.captures(json_str).and_then(|c| c.get(1)).map(|m| m.as_str()).unwrap_or("null"); let id = ID_REGEX
.captures(json_str)
.and_then(|c| c.get(1))
.map(|m| m.as_str())
.unwrap_or("null");
if is_response { if is_response {
let is_error = ERROR_REGEX.is_match(json_str) || IS_ERROR_REGEX.is_match(json_str); let is_error = ERROR_REGEX.is_match(json_str) || IS_ERROR_REGEX.is_match(json_str);
return format!("Response id={} [Error: {}]", id, is_error); return format!("Response id={} [Error: {}]", id, is_error);
} }
let method = METHOD_REGEX.captures(json_str).and_then(|c| c.get(1)).map(|m| m.as_str()).unwrap_or(""); let method = METHOD_REGEX
.captures(json_str)
.and_then(|c| c.get(1))
.map(|m| m.as_str())
.unwrap_or("");
if method == "tools/call" { if method == "tools/call" {
let tool = TOOL_REGEX.captures(json_str).and_then(|c| c.get(1)).map(|m| m.as_str()).unwrap_or("unknown"); let tool = TOOL_REGEX
.captures(json_str)
.and_then(|c| c.get(1))
.map(|m| m.as_str())
.unwrap_or("unknown");
format!("ToolCall[{}] id={}", tool, id) format!("ToolCall[{}] id={}", tool, id)
} else if !method.is_empty() { } else if !method.is_empty() {
format!("Request[{}] id={}", method, id) format!("Request[{}] id={}", method, id)
+28 -7
View File
@@ -10,8 +10,6 @@ struct Cli {
target: String, target: String,
} }
mod logger; mod logger;
fn init_logging(app_name: &str) -> Option<tracing_appender::non_blocking::WorkerGuard> { fn init_logging(app_name: &str) -> Option<tracing_appender::non_blocking::WorkerGuard> {
@@ -53,7 +51,9 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string()); let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string());
format!("http://127.0.0.1:{}", port) format!("http://127.0.0.1:{}", port)
}; };
let ws_url = target_url.replace("http://", "ws://").replace("https://", "wss://"); let ws_url = target_url
.replace("http://", "ws://")
.replace("https://", "wss://");
let ws_url = format!("{}/ws?client=proxy", ws_url); let ws_url = format!("{}/ws?client=proxy", ws_url);
loop { loop {
@@ -83,8 +83,21 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let mut send_task = tokio::spawn(async move { let mut send_task = tokio::spawn(async move {
while let Ok(msg) = rx.recv().await { while let Ok(msg) = rx.recv().await {
let log_prefix = logger::extract_log_prefix(&msg, false); let log_prefix = logger::extract_log_prefix(&msg, false);
tracing::info!(">>> [Stub] Forwarding {} to server (length: {}): {}", log_prefix, msg.len(), if msg.len() > 1000 { format!("{}...", &msg[..1000]) } else { msg.clone() }); tracing::info!(
if write.send(tokio_tungstenite::tungstenite::Message::Text(msg)).await.is_err() { ">>> [Stub] Forwarding {} to server (length: {}): {}",
log_prefix,
msg.len(),
if msg.len() > 1000 {
format!("{}...", &msg[..1000])
} else {
msg.clone()
}
);
if write
.send(tokio_tungstenite::tungstenite::Message::Text(msg))
.await
.is_err()
{
tracing::error!("Failed to write to websocket"); tracing::error!("Failed to write to websocket");
break; break;
} }
@@ -95,7 +108,16 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
while let Some(Ok(msg)) = read.next().await { while let Some(Ok(msg)) = read.next().await {
if let tokio_tungstenite::tungstenite::Message::Text(text) = msg { if let tokio_tungstenite::tungstenite::Message::Text(text) = msg {
let log_prefix = logger::extract_log_prefix(&text, true); let log_prefix = logger::extract_log_prefix(&text, true);
tracing::info!("<<< [Stub] Received {} from server (length: {}): {}", log_prefix, text.len(), if text.len() > 1000 { format!("{}...", &text[..1000]) } else { text.clone() }); tracing::info!(
"<<< [Stub] Received {} from server (length: {}): {}",
log_prefix,
text.len(),
if text.len() > 1000 {
format!("{}...", &text[..1000])
} else {
text.clone()
}
);
use tokio::io::AsyncWriteExt; use tokio::io::AsyncWriteExt;
let mut stdout = tokio::io::stdout(); let mut stdout = tokio::io::stdout();
let _ = stdout.write_all(text.as_bytes()).await; let _ = stdout.write_all(text.as_bytes()).await;
@@ -134,4 +156,3 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
Ok(()) Ok(())
}) })
} }
+12 -6
View File
@@ -61,7 +61,8 @@ async fn test_full_system_e2e_performance() {
assert!(stub_exe.exists(), "Stub not found at {:?}", stub_exe); assert!(stub_exe.exists(), "Stub not found at {:?}", stub_exe);
// 1. Start Server // 1. Start Server
let _server = ChildGuard(Command::new(&server_exe) let _server = ChildGuard(
Command::new(&server_exe)
.arg("--daemon") .arg("--daemon")
.env("MCP_PORT", test_port) .env("MCP_PORT", test_port)
.env("RUST_LOG", "debug") .env("RUST_LOG", "debug")
@@ -71,7 +72,8 @@ async fn test_full_system_e2e_performance() {
.stdout(Stdio::inherit()) .stdout(Stdio::inherit())
.stderr(Stdio::inherit()) .stderr(Stdio::inherit())
.spawn() .spawn()
.expect("Failed to start server")); .expect("Failed to start server"),
);
// Give server time to generate TLS cert and start // Give server time to generate TLS cert and start
let client = reqwest::Client::builder() let client = reqwest::Client::builder()
@@ -94,7 +96,8 @@ async fn test_full_system_e2e_performance() {
assert!(started, "Server failed to start in time"); assert!(started, "Server failed to start in time");
// 2. Start Stub // 2. Start Stub
let mut stub = ChildGuard(Command::new(&stub_exe) let mut stub = ChildGuard(
Command::new(&stub_exe)
.arg("--target") .arg("--target")
.arg(format!("http://127.0.0.1:{}", test_port)) .arg(format!("http://127.0.0.1:{}", test_port))
.env("MCP_MEMORY_STORE_DIR", temp_dir.to_str().unwrap()) .env("MCP_MEMORY_STORE_DIR", temp_dir.to_str().unwrap())
@@ -104,18 +107,21 @@ async fn test_full_system_e2e_performance() {
.stdout(Stdio::piped()) .stdout(Stdio::piped())
.stderr(Stdio::inherit()) .stderr(Stdio::inherit())
.spawn() .spawn()
.expect("Failed to start stub")); .expect("Failed to start stub"),
);
let mut stub_stdin = stub.0.stdin.take().unwrap(); let mut stub_stdin = stub.0.stdin.take().unwrap();
let mut stub_stdout = BufReader::new(stub.0.stdout.take().unwrap()); let mut stub_stdout = BufReader::new(stub.0.stdout.take().unwrap());
// 3. Start Nvim Bridge // 3. Start Nvim Bridge
let mut nvim = ChildGuard(Command::new(&nvim_exe) let mut nvim = ChildGuard(
Command::new(&nvim_exe)
.stdin(Stdio::piped()) .stdin(Stdio::piped())
.stdout(Stdio::piped()) .stdout(Stdio::piped())
.stderr(Stdio::inherit()) .stderr(Stdio::inherit())
.spawn() .spawn()
.expect("Failed to start nvim bridge")); .expect("Failed to start nvim bridge"),
);
let mut nvim_stdin = nvim.0.stdin.take().unwrap(); let mut nvim_stdin = nvim.0.stdin.take().unwrap();
let mut nvim_stdout = BufReader::new(nvim.0.stdout.take().unwrap()); let mut nvim_stdout = BufReader::new(nvim.0.stdout.take().unwrap());
+1 -5
View File
@@ -15,10 +15,6 @@ fn main() {
.and_then(|out| String::from_utf8(out.stdout).ok()) .and_then(|out| String::from_utf8(out.stdout).ok())
.unwrap_or_else(|| "unknown".to_string()); .unwrap_or_else(|| "unknown".to_string());
let version = format!( let version = format!("{} ({})", git_date.trim(), git_hash.trim());
"{} ({})",
git_date.trim(),
git_hash.trim()
);
println!("cargo:rustc-env=APP_VERSION={}", version); println!("cargo:rustc-env=APP_VERSION={}", version);
} }
+1 -1
View File
@@ -1,5 +1,5 @@
use serde_json::{json, Value}; use serde_json::{json, Value};
use std::io::{BufRead, BufReader, Read, Write}; use std::io::{BufRead, BufReader, Write};
use std::process::{Command, Stdio}; use std::process::{Command, Stdio};
fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) { fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) {