chore: rustfmt, clippy lints and code tidying
This commit is contained in:
1 parent
3716c3e698
commit
1752753fcc
19 files changed
+890
-355
No files matched your search
+8
-3
@@ -2,19 +2,24 @@ use std::process::Command;
|
||||
|
||||
fn main() {
|
||||
let git_hash = Command::new("git")
|
||||
.args(&["rev-parse", "--short", "HEAD"])
|
||||
.args(["rev-parse", "--short", "HEAD"])
|
||||
.output()
|
||||
.ok()
|
||||
.and_then(|out| String::from_utf8(out.stdout).ok())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
|
||||
let git_date = Command::new("git")
|
||||
.args(&["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"])
|
||||
.args(["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"])
|
||||
.output()
|
||||
.ok()
|
||||
.and_then(|out| String::from_utf8(out.stdout).ok())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
|
||||
let version = format!("{} ({} {})", env!("CARGO_PKG_VERSION"), git_date.trim(), git_hash.trim());
|
||||
let version = format!(
|
||||
"{} ({} {})",
|
||||
env!("CARGO_PKG_VERSION"),
|
||||
git_date.trim(),
|
||||
git_hash.trim()
|
||||
);
|
||||
println!("cargo:rustc-env=APP_VERSION={}", version);
|
||||
}
|
||||
@@ -1,6 +1,5 @@
|
||||
use serde_json::{json, Value};
|
||||
use serde_json::Value;
|
||||
use std::io::{BufRead, BufReader, Read, Write};
|
||||
use std::process::{Command, Stdio};
|
||||
|
||||
fn send_message(stdin: &mut std::process::ChildStdin, msg: Value) {
|
||||
let s = serde_json::to_string(&msg).unwrap();
|
||||
@@ -46,7 +45,10 @@ fn test_mcp_initialization_and_tools_list() {
|
||||
let mut nvim_exe = std::env::current_exe().unwrap();
|
||||
nvim_exe.pop();
|
||||
nvim_exe.pop();
|
||||
nvim_exe.push(format!("mcp-memory-linux-nvim{}", std::env::consts::EXE_SUFFIX));
|
||||
nvim_exe.push(format!(
|
||||
"mcp-memory-linux-nvim{}",
|
||||
std::env::consts::EXE_SUFFIX
|
||||
));
|
||||
|
||||
let mut child = Command::new(&nvim_exe)
|
||||
.stdin(Stdio::piped())
|
||||
@@ -113,7 +115,9 @@ fn test_mcp_initialization_and_tools_list() {
|
||||
assert_eq!(tools_resp["jsonrpc"], "2.0");
|
||||
assert_eq!(tools_resp["id"], 2);
|
||||
|
||||
let tools = tools_resp["result"]["tools"].as_array().expect("result.tools must be an array");
|
||||
let tools = tools_resp["result"]["tools"]
|
||||
.as_array()
|
||||
.expect("result.tools must be an array");
|
||||
assert!(!tools.is_empty(), "Server must expose at least one tool");
|
||||
|
||||
let has_get_active_buffer = tools.iter().any(|t| t["name"] == "nvim_get_active_buffer");
|
||||
|
||||
+119
-58
@@ -20,7 +20,9 @@ pub struct JsonRpcResponse {
|
||||
pub error: Option<Value>,
|
||||
}
|
||||
|
||||
pub async fn read_message<R: tokio::io::AsyncRead + Unpin>(stdin: &mut BufReader<R>) -> Option<JsonRpcRequest> {
|
||||
pub async fn read_message<R: tokio::io::AsyncRead + Unpin>(
|
||||
stdin: &mut BufReader<R>,
|
||||
) -> Option<JsonRpcRequest> {
|
||||
let mut length = 0;
|
||||
loop {
|
||||
let mut line = String::new();
|
||||
@@ -32,7 +34,11 @@ pub async fn read_message<R: tokio::io::AsyncRead + Unpin>(stdin: &mut BufReader
|
||||
return match serde_json::from_str::<JsonRpcRequest>(line.trim_end()) {
|
||||
Ok(req) => Some(req),
|
||||
Err(e) => {
|
||||
tracing::error!("Failed to parse JSON-RPC request from JSONL: {}. Payload: {}", e, line);
|
||||
tracing::error!(
|
||||
"Failed to parse JSON-RPC request from JSONL: {}. Payload: {}",
|
||||
e,
|
||||
line
|
||||
);
|
||||
None
|
||||
}
|
||||
};
|
||||
@@ -58,7 +64,15 @@ pub async fn read_message<R: tokio::io::AsyncRead + Unpin>(stdin: &mut BufReader
|
||||
|
||||
pub async fn send_response(response: JsonRpcResponse) {
|
||||
let msg = serde_json::to_string(&response).unwrap();
|
||||
tracing::info!("Sending JSON-RPC response (id: {:?}): {}", response.id, if msg.len() > 500 { format!("{}...", &msg[..500]) } else { msg.clone() });
|
||||
tracing::info!(
|
||||
"Sending JSON-RPC response (id: {:?}): {}",
|
||||
response.id,
|
||||
if msg.len() > 500 {
|
||||
format!("{}...", &msg[..500])
|
||||
} else {
|
||||
msg.clone()
|
||||
}
|
||||
);
|
||||
// CRITICAL ARCHITECTURAL DECISION:
|
||||
// The MCP StdioTransport MUST use Newline-Delimited JSON (NDJSON).
|
||||
// Do NOT use LSP-style HTTP headers (e.g. Content-Length).
|
||||
@@ -75,12 +89,14 @@ pub async fn send_error(id: Value, code: i32, message: &str) {
|
||||
id,
|
||||
result: None,
|
||||
error: Some(serde_json::json!({"code": code, "message": message})),
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
#[cfg(windows)]
|
||||
async fn get_socket_path() -> Result<String, String> {
|
||||
let profile = std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into());
|
||||
let profile =
|
||||
std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into());
|
||||
let path = format!("{}\\.gemini\\active_nvim.txt", profile);
|
||||
|
||||
if let Ok(content) = std::fs::read_to_string(&path) {
|
||||
@@ -140,12 +156,20 @@ async fn call_nvim(req: rmpv::Value) -> Result<rmpv::Value, String> {
|
||||
use tokio::net::windows::named_pipe::ClientOptions;
|
||||
|
||||
let msgid = if let rmpv::Value::Array(ref arr) = req {
|
||||
if arr.len() > 1 { arr[1].clone() } else { rmpv::Value::Nil }
|
||||
} else { rmpv::Value::Nil };
|
||||
if arr.len() > 1 {
|
||||
arr[1].clone()
|
||||
} else {
|
||||
rmpv::Value::Nil
|
||||
}
|
||||
} else {
|
||||
rmpv::Value::Nil
|
||||
};
|
||||
|
||||
tracing::info!("Connecting to neovim pipe");
|
||||
let socket_path = get_socket_path().await?;
|
||||
let mut client = ClientOptions::new().open(&socket_path).map_err(|e| e.to_string())?;
|
||||
let mut client = ClientOptions::new()
|
||||
.open(&socket_path)
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let mut buf = Vec::new();
|
||||
rmpv::encode::write_value(&mut buf, &req).map_err(|e| e.to_string())?;
|
||||
@@ -163,18 +187,21 @@ async fn call_nvim(req: rmpv::Value) -> Result<rmpv::Value, String> {
|
||||
offset += cursor.position() as usize;
|
||||
|
||||
if let rmpv::Value::Array(ref arr) = val {
|
||||
if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) && arr[1] == msgid {
|
||||
if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) && arr[1] == msgid
|
||||
{
|
||||
tracing::info!("Received RPC response from neovim (msgid: {})", msgid);
|
||||
return Ok(val);
|
||||
}
|
||||
}
|
||||
continue;
|
||||
},
|
||||
}
|
||||
Err(_) => {
|
||||
let read_future = client.read(&mut chunk);
|
||||
match tokio::time::timeout(tokio::time::Duration::from_secs(5), read_future).await {
|
||||
Ok(Ok(n)) => {
|
||||
if n == 0 { return Err("Connection closed".into()); }
|
||||
if n == 0 {
|
||||
return Err("Connection closed".into());
|
||||
}
|
||||
resp_buf.extend_from_slice(&chunk[..n]);
|
||||
}
|
||||
Ok(Err(e)) => return Err(e.to_string()),
|
||||
@@ -193,12 +220,20 @@ async fn call_nvim(req: rmpv::Value) -> Result<rmpv::Value, String> {
|
||||
use tokio::net::UnixStream;
|
||||
|
||||
let msgid = if let rmpv::Value::Array(ref arr) = req {
|
||||
if arr.len() > 1 { arr[1].clone() } else { rmpv::Value::Nil }
|
||||
} else { rmpv::Value::Nil };
|
||||
if arr.len() > 1 {
|
||||
arr[1].clone()
|
||||
} else {
|
||||
rmpv::Value::Nil
|
||||
}
|
||||
} else {
|
||||
rmpv::Value::Nil
|
||||
};
|
||||
|
||||
tracing::info!("Connecting to neovim socket");
|
||||
let socket_path = get_socket_path().await?;
|
||||
let mut stream = UnixStream::connect(socket_path).await.map_err(|e| e.to_string())?;
|
||||
let mut stream = UnixStream::connect(socket_path)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
let mut buf = Vec::new();
|
||||
rmpv::encode::write_value(&mut buf, &req).map_err(|e| e.to_string())?;
|
||||
@@ -216,18 +251,21 @@ async fn call_nvim(req: rmpv::Value) -> Result<rmpv::Value, String> {
|
||||
offset += cursor.position() as usize;
|
||||
|
||||
if let rmpv::Value::Array(ref arr) = val {
|
||||
if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) && arr[1] == msgid {
|
||||
if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) && arr[1] == msgid
|
||||
{
|
||||
tracing::info!("Received RPC response from neovim (msgid: {})", msgid);
|
||||
return Ok(val);
|
||||
}
|
||||
}
|
||||
continue;
|
||||
},
|
||||
}
|
||||
Err(_) => {
|
||||
let read_future = stream.read(&mut chunk);
|
||||
match tokio::time::timeout(tokio::time::Duration::from_secs(5), read_future).await {
|
||||
Ok(Ok(n)) => {
|
||||
if n == 0 { return Err("Connection closed".into()); }
|
||||
if n == 0 {
|
||||
return Err("Connection closed".into());
|
||||
}
|
||||
resp_buf.extend_from_slice(&chunk[..n]);
|
||||
}
|
||||
Ok(Err(e)) => return Err(e.to_string()),
|
||||
@@ -300,9 +338,7 @@ async fn get_nvim_cursor() -> Result<String, String> {
|
||||
RmpValue::Integer(0.into()),
|
||||
RmpValue::Integer(3.into()), // msgid
|
||||
RmpValue::String("nvim_win_get_cursor".into()),
|
||||
RmpValue::Array(vec![
|
||||
RmpValue::Integer(0.into()),
|
||||
]),
|
||||
RmpValue::Array(vec![RmpValue::Integer(0.into())]),
|
||||
]);
|
||||
|
||||
let resp = call_nvim(req).await?;
|
||||
@@ -360,7 +396,8 @@ async fn get_nvim_visual_selection() -> Result<String, String> {
|
||||
|
||||
async fn set_nvim_diagnostics(line: i64, message: &str) -> Result<(), String> {
|
||||
let escaped_message = message.replace("\\", "\\\\").replace("\"", "\\\"");
|
||||
let lua_script = format!(r#"
|
||||
let lua_script = format!(
|
||||
r#"
|
||||
local ns = vim.api.nvim_create_namespace("gemini_diagnostics")
|
||||
local diagnostics = {{{{
|
||||
lnum = {} - 1,
|
||||
@@ -369,7 +406,9 @@ async fn set_nvim_diagnostics(line: i64, message: &str) -> Result<(), String> {
|
||||
message = "{}",
|
||||
}}}}
|
||||
vim.diagnostic.set(ns, 0, diagnostics, {{}})
|
||||
"#, line, escaped_message);
|
||||
"#,
|
||||
line, escaped_message
|
||||
);
|
||||
|
||||
use rmpv::Value as RmpValue;
|
||||
let req = RmpValue::Array(vec![
|
||||
@@ -404,7 +443,7 @@ fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value {
|
||||
} else {
|
||||
serde_json::Value::Null
|
||||
}
|
||||
},
|
||||
}
|
||||
rmpv::Value::F32(f) => serde_json::json!(f),
|
||||
rmpv::Value::F64(f) => serde_json::json!(f),
|
||||
rmpv::Value::String(s) => {
|
||||
@@ -413,11 +452,11 @@ fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value {
|
||||
} else {
|
||||
serde_json::Value::Null
|
||||
}
|
||||
},
|
||||
}
|
||||
rmpv::Value::Array(arr) => {
|
||||
let vec: Vec<serde_json::Value> = arr.iter().map(rmpv_to_json).collect();
|
||||
serde_json::Value::Array(vec)
|
||||
},
|
||||
}
|
||||
rmpv::Value::Map(map) => {
|
||||
let mut obj = serde_json::Map::new();
|
||||
for (k, v) in map {
|
||||
@@ -429,7 +468,7 @@ fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value {
|
||||
obj.insert(key_str, rmpv_to_json(v));
|
||||
}
|
||||
serde_json::Value::Object(obj)
|
||||
},
|
||||
}
|
||||
_ => serde_json::json!(format!("{:?}", val)),
|
||||
}
|
||||
}
|
||||
@@ -440,10 +479,7 @@ async fn execute_nvim_lua(code: &str) -> Result<String, String> {
|
||||
RmpValue::Integer(0.into()),
|
||||
RmpValue::Integer(6.into()), // msgid
|
||||
RmpValue::String("nvim_exec_lua".into()),
|
||||
RmpValue::Array(vec![
|
||||
RmpValue::String(code.into()),
|
||||
RmpValue::Array(vec![]),
|
||||
]),
|
||||
RmpValue::Array(vec![RmpValue::String(code.into()), RmpValue::Array(vec![])]),
|
||||
]);
|
||||
|
||||
let resp = call_nvim(req).await?;
|
||||
@@ -460,7 +496,12 @@ async fn execute_nvim_lua(code: &str) -> Result<String, String> {
|
||||
}
|
||||
pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
if std::env::args().any(|arg| arg == "--version") {
|
||||
println!("{} {} ({})", app_name, app_version, std::env::var("GIT_HASH").unwrap_or_else(|_| "unknown".to_string()));
|
||||
println!(
|
||||
"{} {} ({})",
|
||||
app_name,
|
||||
app_version,
|
||||
std::env::var("GIT_HASH").unwrap_or_else(|_| "unknown".to_string())
|
||||
);
|
||||
return;
|
||||
}
|
||||
let _guard = init_logging(app_name);
|
||||
@@ -471,7 +512,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
Some(m) => {
|
||||
tracing::info!("Received message method: {}", m.method);
|
||||
m
|
||||
},
|
||||
}
|
||||
None => {
|
||||
tracing::info!("Stdin closed, exiting loop");
|
||||
break;
|
||||
@@ -488,9 +529,14 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
match msg.method.as_str() {
|
||||
"initialize" => {
|
||||
let init = rmcp::model::InitializeResult::new(
|
||||
rmcp::model::ServerCapabilities::builder().enable_tools().build()
|
||||
rmcp::model::ServerCapabilities::builder()
|
||||
.enable_tools()
|
||||
.build(),
|
||||
)
|
||||
.with_server_info(rmcp::model::Implementation::new(app_name.clone(), app_version.clone()))
|
||||
.with_server_info(rmcp::model::Implementation::new(
|
||||
app_name.clone(),
|
||||
app_version.clone(),
|
||||
))
|
||||
.with_protocol_version(rmcp::model::ProtocolVersion::V_2024_11_05);
|
||||
|
||||
send_response(JsonRpcResponse {
|
||||
@@ -498,7 +544,8 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
id,
|
||||
result: Some(serde_json::to_value(init).unwrap()),
|
||||
error: None,
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
"notifications/initialized" => {}
|
||||
"tools/list" => {
|
||||
@@ -595,7 +642,10 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
|
||||
match name {
|
||||
"nvim_goto_line" => {
|
||||
if let (Some(file), Some(line)) = (args.get("file").and_then(|v| v.as_str()), args.get("line").and_then(|v| v.as_i64())) {
|
||||
if let (Some(file), Some(line)) = (
|
||||
args.get("file").and_then(|v| v.as_str()),
|
||||
args.get("line").and_then(|v| v.as_i64()),
|
||||
) {
|
||||
let escaped_file = file.replace("\\", "\\\\");
|
||||
let cmd = format!("e {} | {} | normal! zz", escaped_file, line);
|
||||
match send_nvim_command(&cmd).await {
|
||||
@@ -615,8 +665,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
send_error(id, -32602, "Missing 'file' or 'line'").await;
|
||||
}
|
||||
}
|
||||
"nvim_get_active_buffer" => {
|
||||
match get_nvim_active_buffer().await {
|
||||
"nvim_get_active_buffer" => match get_nvim_active_buffer().await {
|
||||
Ok(content) => {
|
||||
send_response(JsonRpcResponse {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
@@ -625,13 +674,12 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
"content": [{"type": "text", "text": content}]
|
||||
})),
|
||||
error: None,
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
Err(e) => send_error(id, -32603, &e).await,
|
||||
}
|
||||
}
|
||||
"nvim_get_cursor" => {
|
||||
match get_nvim_cursor().await {
|
||||
},
|
||||
"nvim_get_cursor" => match get_nvim_cursor().await {
|
||||
Ok(content) => {
|
||||
send_response(JsonRpcResponse {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
@@ -640,13 +688,12 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
"content": [{"type": "text", "text": content}]
|
||||
})),
|
||||
error: None,
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
Err(e) => send_error(id, -32603, &e).await,
|
||||
}
|
||||
}
|
||||
"nvim_get_visual_selection" => {
|
||||
match get_nvim_visual_selection().await {
|
||||
},
|
||||
"nvim_get_visual_selection" => match get_nvim_visual_selection().await {
|
||||
Ok(content) => {
|
||||
send_response(JsonRpcResponse {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
@@ -655,13 +702,16 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
"content": [{"type": "text", "text": content}]
|
||||
})),
|
||||
error: None,
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
Err(e) => send_error(id, -32603, &e).await,
|
||||
}
|
||||
}
|
||||
},
|
||||
"nvim_set_diagnostics" => {
|
||||
if let (Some(line), Some(message)) = (args.get("line").and_then(|v| v.as_i64()), args.get("message").and_then(|v| v.as_str())) {
|
||||
if let (Some(line), Some(message)) = (
|
||||
args.get("line").and_then(|v| v.as_i64()),
|
||||
args.get("message").and_then(|v| v.as_str()),
|
||||
) {
|
||||
match set_nvim_diagnostics(line, message).await {
|
||||
Ok(_) => {
|
||||
send_response(JsonRpcResponse {
|
||||
@@ -700,7 +750,8 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
"content": [{"type": "text", "text": result}]
|
||||
})),
|
||||
error: None,
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
Err(e) => send_error(id, -32603, &e).await,
|
||||
}
|
||||
@@ -730,7 +781,8 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
"content": [{"type": "text", "text": result}]
|
||||
})),
|
||||
error: None,
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
Err(e) => send_error(id, -32603, &e).await,
|
||||
}
|
||||
@@ -746,7 +798,8 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
"content": [{"type": "text", "text": result}]
|
||||
})),
|
||||
error: None,
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
Err(e) => send_error(id, -32603, &e).await,
|
||||
}
|
||||
@@ -770,7 +823,9 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
||||
}
|
||||
|
||||
fn init_logging(app_name: &str) -> tracing_appender::non_blocking::WorkerGuard {
|
||||
let log_dir = dirs::home_dir().unwrap_or_default().join(".gemini/mcp_memory/logs");
|
||||
let log_dir = dirs::home_dir()
|
||||
.unwrap_or_default()
|
||||
.join(".gemini/mcp_memory/logs");
|
||||
std::fs::create_dir_all(&log_dir).unwrap_or_default();
|
||||
|
||||
let file_appender = tracing_appender::rolling::daily(log_dir, format!("{}.log", app_name));
|
||||
@@ -794,7 +849,10 @@ mod tests {
|
||||
assert_eq!(rmpv_to_json(&rmpv::Value::Nil), serde_json::Value::Null);
|
||||
assert_eq!(rmpv_to_json(&rmpv::Value::Boolean(true)), json!(true));
|
||||
assert_eq!(rmpv_to_json(&rmpv::Value::Integer(42.into())), json!(42));
|
||||
assert_eq!(rmpv_to_json(&rmpv::Value::String("hello".into())), json!("hello"));
|
||||
assert_eq!(
|
||||
rmpv_to_json(&rmpv::Value::String("hello".into())),
|
||||
json!("hello")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -809,7 +867,10 @@ mod tests {
|
||||
#[test]
|
||||
fn test_rmpv_to_json_map() {
|
||||
let mut map = vec![];
|
||||
map.push((rmpv::Value::String("key1".into()), rmpv::Value::Integer(100.into())));
|
||||
map.push((
|
||||
rmpv::Value::String("key1".into()),
|
||||
rmpv::Value::Integer(100.into()),
|
||||
));
|
||||
let rmp_map = rmpv::Value::Map(map);
|
||||
|
||||
let json_map = rmpv_to_json(&rmp_map);
|
||||
|
||||
+8
-3
@@ -2,19 +2,24 @@ use std::process::Command;
|
||||
|
||||
fn main() {
|
||||
let git_hash = Command::new("git")
|
||||
.args(&["rev-parse", "--short", "HEAD"])
|
||||
.args(["rev-parse", "--short", "HEAD"])
|
||||
.output()
|
||||
.ok()
|
||||
.and_then(|out| String::from_utf8(out.stdout).ok())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
|
||||
let git_date = Command::new("git")
|
||||
.args(&["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"])
|
||||
.args(["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"])
|
||||
.output()
|
||||
.ok()
|
||||
.and_then(|out| String::from_utf8(out.stdout).ok())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
|
||||
let version = format!("{} ({} {})", env!("CARGO_PKG_VERSION"), git_date.trim(), git_hash.trim());
|
||||
let version = format!(
|
||||
"{} ({} {})",
|
||||
env!("CARGO_PKG_VERSION"),
|
||||
git_date.trim(),
|
||||
git_hash.trim()
|
||||
);
|
||||
println!("cargo:rustc-env=APP_VERSION={}", version);
|
||||
}
|
||||
@@ -1,8 +1,10 @@
|
||||
use rmcp::model::{InitializeResult, ServerCapabilities};
|
||||
|
||||
fn main() {
|
||||
let init = InitializeResult::new(
|
||||
ServerCapabilities::builder().enable_tools().build()
|
||||
).with_server_info(rmcp::model::Implementation::new("gemini-mcp-memory", "3.0.0"));
|
||||
let init = InitializeResult::new(ServerCapabilities::builder().enable_tools().build())
|
||||
.with_server_info(rmcp::model::Implementation::new(
|
||||
"gemini-mcp-memory",
|
||||
"3.0.0",
|
||||
));
|
||||
println!("{}", serde_json::to_string_pretty(&init).unwrap());
|
||||
}
|
||||
+466
-151
@@ -18,8 +18,6 @@ macro_rules! parse_tool {
|
||||
};
|
||||
}
|
||||
|
||||
|
||||
|
||||
use serde::de::DeserializeOwned;
|
||||
use std::collections::HashSet;
|
||||
use std::sync::Arc;
|
||||
@@ -57,92 +55,276 @@ impl MemoryHandler {
|
||||
}
|
||||
}
|
||||
});
|
||||
tracing::debug!("<<< [Server] Replying to server/discover with payload: {}", payload.to_string());
|
||||
tracing::debug!(
|
||||
"<<< [Server] Replying to server/discover with payload: {}",
|
||||
payload.to_string()
|
||||
);
|
||||
Some(crate::mcp::success(id, payload))
|
||||
}
|
||||
"initialize" => {
|
||||
let init = rmcp::model::InitializeResult::new(
|
||||
rmcp::model::ServerCapabilities::builder().enable_tools().build()
|
||||
).with_server_info(rmcp::model::Implementation::new("gemini-mcp-memory", "3.0.0"));
|
||||
rmcp::model::ServerCapabilities::builder()
|
||||
.enable_tools()
|
||||
.build(),
|
||||
)
|
||||
.with_server_info(rmcp::model::Implementation::new(
|
||||
"gemini-mcp-memory",
|
||||
"3.0.0",
|
||||
));
|
||||
tracing::debug!("<<< [Server] Replying to initialize with rmcp payload");
|
||||
Some(crate::mcp::success(id, serde_json::to_value(&init).unwrap()))
|
||||
Some(crate::mcp::success(
|
||||
id,
|
||||
serde_json::to_value(&init).unwrap(),
|
||||
))
|
||||
}
|
||||
|
||||
"notifications/initialized" => {
|
||||
None
|
||||
}
|
||||
"notifications/initialized" => None,
|
||||
"tools/list" => {
|
||||
let tools = vec![
|
||||
crate::mcp::tool_def::<crate::tools::QueryGraphPathTool>("query_graph_path", "Traverse the knowledge graph to find a path between two entities."),
|
||||
crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entities in the knowledge graph."),
|
||||
crate::mcp::tool_def::<CreateRelationsTool>("create_relations", "Create new relations between entities in the knowledge graph."),
|
||||
crate::mcp::tool_def::<AddObservationsTool>("add_observations", "Add new observations to existing entities in the knowledge graph."),
|
||||
crate::mcp::tool_def::<DeleteEntitiesTool>("delete_entities", "Delete entities from the knowledge graph."),
|
||||
crate::mcp::tool_def::<DeleteObservationsTool>("delete_observations", "Delete observations from existing entities."),
|
||||
crate::mcp::tool_def::<DeleteRelationsTool>("delete_relations", "Delete relations between entities."),
|
||||
crate::mcp::tool_def::<ReadGraphTool>("read_graph", "Read the entire knowledge graph."),
|
||||
crate::mcp::tool_def::<SearchNodesTool>("search_nodes", "Search for entities in the knowledge graph by name or type."),
|
||||
crate::mcp::tool_def::<OpenNodesTool>("open_nodes", "Open and retrieve full details of specific nodes in the knowledge graph."),
|
||||
crate::mcp::tool_def::<LogCodeChangeTool>("log_code_change", "Log a significant code change or refactor in the memory system."),
|
||||
crate::mcp::tool_def::<QueryRecentChangesTool>("query_recent_changes", "Query recently logged code changes."),
|
||||
crate::mcp::tool_def::<VisualizeGraphTool>("visualize_graph", "Generate a visual representation of the knowledge graph."),
|
||||
crate::mcp::tool_def::<AddStickyNoteTool>("add_sticky_note", "Add a sticky note for unstructured thoughts or reminders."),
|
||||
crate::mcp::tool_def::<ReadStickyNotesTool>("read_sticky_notes", "Read all active sticky notes."),
|
||||
crate::mcp::tool_def::<CondenseEntityTool>("condense_entity", "Condense or summarize an entity's observations to reduce size."),
|
||||
crate::mcp::tool_def::<AddTaskTool>("add_task", "Add a new task to the task tracker."),
|
||||
crate::mcp::tool_def::<UpdateTaskStatusTool>("update_task_status", "Update the status of an existing task."),
|
||||
crate::mcp::tool_def::<DeleteTaskTool>("delete_task", "Delete a task and all its children."),
|
||||
crate::mcp::tool_def::<ListActiveTasksTool>("list_active_tasks", "List all currently active tasks."),
|
||||
crate::mcp::tool_def::<SetAcceptanceCriteriaTool>("set_acceptance_criteria", "Define a strict checklist of acceptance criteria for a given task."),
|
||||
crate::mcp::tool_def::<VerifyAcceptanceCriteriaTool>("verify_acceptance_criteria", "Mark a previously defined acceptance criteria as met."),
|
||||
crate::mcp::tool_def::<StoreSnippetTool>("store_snippet", "Store a reusable code snippet."),
|
||||
crate::mcp::tool_def::<SearchSnippetsTool>("search_snippets", "Search through stored code snippets."),
|
||||
crate::mcp::tool_def::<DeleteSnippetTool>("delete_snippet", "Delete a stored code snippet."),
|
||||
crate::mcp::tool_def::<LogDecisionTool>("log_decision", "Log an architectural decision record (ADR)."),
|
||||
crate::mcp::tool_def::<QueryDecisionsTool>("query_decisions", "Query architectural decision records."),
|
||||
crate::mcp::tool_def::<MergeEntitiesTool>("merge_entities", "Merge two entities in the knowledge graph into one."),
|
||||
crate::mcp::tool_def::<FindOrphansTool>("find_orphans", "Find orphaned entities (entities without any relations) in the graph."),
|
||||
crate::mcp::tool_def::<LearnPreferenceTool>("learn_preference", "Record a user preference or behavior to adapt future interactions."),
|
||||
crate::mcp::tool_def::<ReadPreferencesTool>("read_preferences", "Read all learned user preferences."),
|
||||
crate::mcp::tool_def::<LogErrorFixTool>("log_error_fix", "Log a complex error and its fix for future reference."),
|
||||
crate::mcp::tool_def::<SearchErrorFixesTool>("search_error_fixes", "Search through previously logged error fixes."),
|
||||
crate::mcp::tool_def::<PinFileTool>("pin_file", "Pin a file to keep it explicitly in the context workspace."),
|
||||
crate::mcp::tool_def::<UnpinFileTool>("unpin_file", "Unpin a file from the context workspace."),
|
||||
crate::mcp::tool_def::<ListPinnedFilesTool>("list_pinned_files", "List all currently pinned files."),
|
||||
crate::mcp::tool_def::<AddSessionSummaryTool>("add_session_summary", "Add a summary of the current session."),
|
||||
crate::mcp::tool_def::<GetProjectTimelineTool>("get_project_timeline", "Get a timeline of major project events."),
|
||||
crate::mcp::tool_def::<LeaveHandoffMemoTool>("leave_handoff_memo", "Leave a memo for the next session or agent."),
|
||||
crate::mcp::tool_def::<ReadHandoffMemosTool>("read_handoff_memos", "Read pending handoff memos."),
|
||||
crate::mcp::tool_def::<ClearHandoffMemosTool>("clear_handoff_memos", "Clear handoff memos after reading."),
|
||||
crate::mcp::tool_def::<UpdateEnvFingerprintTool>("update_env_fingerprint", "Update the environment fingerprint (e.g., OS, tool versions)."),
|
||||
crate::mcp::tool_def::<ReadEnvFingerprintTool>("read_env_fingerprint", "Read the current environment fingerprint."),
|
||||
crate::mcp::tool_def::<LogEnvRequirementTool>("log_env_requirement", "Log a required tool or package for the environment."),
|
||||
crate::mcp::tool_def::<AddMilestoneTool>("add_milestone", "Add a new project milestone."),
|
||||
crate::mcp::tool_def::<UpdateMilestoneTool>("update_milestone", "Update the status of a project milestone."),
|
||||
crate::mcp::tool_def::<ListMilestonesTool>("list_milestones", "List all project milestones."),
|
||||
crate::mcp::tool_def::<crate::tools::QueryGraphPathTool>(
|
||||
"query_graph_path",
|
||||
"Traverse the knowledge graph to find a path between two entities.",
|
||||
),
|
||||
crate::mcp::tool_def::<CreateEntitiesTool>(
|
||||
"create_entities",
|
||||
"Create new entities in the knowledge graph.",
|
||||
),
|
||||
crate::mcp::tool_def::<CreateRelationsTool>(
|
||||
"create_relations",
|
||||
"Create new relations between entities in the knowledge graph.",
|
||||
),
|
||||
crate::mcp::tool_def::<AddObservationsTool>(
|
||||
"add_observations",
|
||||
"Add new observations to existing entities in the knowledge graph.",
|
||||
),
|
||||
crate::mcp::tool_def::<DeleteEntitiesTool>(
|
||||
"delete_entities",
|
||||
"Delete entities from the knowledge graph.",
|
||||
),
|
||||
crate::mcp::tool_def::<DeleteObservationsTool>(
|
||||
"delete_observations",
|
||||
"Delete observations from existing entities.",
|
||||
),
|
||||
crate::mcp::tool_def::<DeleteRelationsTool>(
|
||||
"delete_relations",
|
||||
"Delete relations between entities.",
|
||||
),
|
||||
crate::mcp::tool_def::<ReadGraphTool>(
|
||||
"read_graph",
|
||||
"Read the entire knowledge graph.",
|
||||
),
|
||||
crate::mcp::tool_def::<SearchNodesTool>(
|
||||
"search_nodes",
|
||||
"Search for entities in the knowledge graph by name or type.",
|
||||
),
|
||||
crate::mcp::tool_def::<OpenNodesTool>(
|
||||
"open_nodes",
|
||||
"Open and retrieve full details of specific nodes in the knowledge graph.",
|
||||
),
|
||||
crate::mcp::tool_def::<LogCodeChangeTool>(
|
||||
"log_code_change",
|
||||
"Log a significant code change or refactor in the memory system.",
|
||||
),
|
||||
crate::mcp::tool_def::<QueryRecentChangesTool>(
|
||||
"query_recent_changes",
|
||||
"Query recently logged code changes.",
|
||||
),
|
||||
crate::mcp::tool_def::<VisualizeGraphTool>(
|
||||
"visualize_graph",
|
||||
"Generate a visual representation of the knowledge graph.",
|
||||
),
|
||||
crate::mcp::tool_def::<AddStickyNoteTool>(
|
||||
"add_sticky_note",
|
||||
"Add a sticky note for unstructured thoughts or reminders.",
|
||||
),
|
||||
crate::mcp::tool_def::<ReadStickyNotesTool>(
|
||||
"read_sticky_notes",
|
||||
"Read all active sticky notes.",
|
||||
),
|
||||
crate::mcp::tool_def::<CondenseEntityTool>(
|
||||
"condense_entity",
|
||||
"Condense or summarize an entity's observations to reduce size.",
|
||||
),
|
||||
crate::mcp::tool_def::<AddTaskTool>(
|
||||
"add_task",
|
||||
"Add a new task to the task tracker.",
|
||||
),
|
||||
crate::mcp::tool_def::<UpdateTaskStatusTool>(
|
||||
"update_task_status",
|
||||
"Update the status of an existing task.",
|
||||
),
|
||||
crate::mcp::tool_def::<DeleteTaskTool>(
|
||||
"delete_task",
|
||||
"Delete a task and all its children.",
|
||||
),
|
||||
crate::mcp::tool_def::<ListActiveTasksTool>(
|
||||
"list_active_tasks",
|
||||
"List all currently active tasks.",
|
||||
),
|
||||
crate::mcp::tool_def::<SetAcceptanceCriteriaTool>(
|
||||
"set_acceptance_criteria",
|
||||
"Define a strict checklist of acceptance criteria for a given task.",
|
||||
),
|
||||
crate::mcp::tool_def::<VerifyAcceptanceCriteriaTool>(
|
||||
"verify_acceptance_criteria",
|
||||
"Mark a previously defined acceptance criteria as met.",
|
||||
),
|
||||
crate::mcp::tool_def::<StoreSnippetTool>(
|
||||
"store_snippet",
|
||||
"Store a reusable code snippet.",
|
||||
),
|
||||
crate::mcp::tool_def::<SearchSnippetsTool>(
|
||||
"search_snippets",
|
||||
"Search through stored code snippets.",
|
||||
),
|
||||
crate::mcp::tool_def::<DeleteSnippetTool>(
|
||||
"delete_snippet",
|
||||
"Delete a stored code snippet.",
|
||||
),
|
||||
crate::mcp::tool_def::<LogDecisionTool>(
|
||||
"log_decision",
|
||||
"Log an architectural decision record (ADR).",
|
||||
),
|
||||
crate::mcp::tool_def::<QueryDecisionsTool>(
|
||||
"query_decisions",
|
||||
"Query architectural decision records.",
|
||||
),
|
||||
crate::mcp::tool_def::<MergeEntitiesTool>(
|
||||
"merge_entities",
|
||||
"Merge two entities in the knowledge graph into one.",
|
||||
),
|
||||
crate::mcp::tool_def::<FindOrphansTool>(
|
||||
"find_orphans",
|
||||
"Find orphaned entities (entities without any relations) in the graph.",
|
||||
),
|
||||
crate::mcp::tool_def::<LearnPreferenceTool>(
|
||||
"learn_preference",
|
||||
"Record a user preference or behavior to adapt future interactions.",
|
||||
),
|
||||
crate::mcp::tool_def::<ReadPreferencesTool>(
|
||||
"read_preferences",
|
||||
"Read all learned user preferences.",
|
||||
),
|
||||
crate::mcp::tool_def::<LogErrorFixTool>(
|
||||
"log_error_fix",
|
||||
"Log a complex error and its fix for future reference.",
|
||||
),
|
||||
crate::mcp::tool_def::<SearchErrorFixesTool>(
|
||||
"search_error_fixes",
|
||||
"Search through previously logged error fixes.",
|
||||
),
|
||||
crate::mcp::tool_def::<PinFileTool>(
|
||||
"pin_file",
|
||||
"Pin a file to keep it explicitly in the context workspace.",
|
||||
),
|
||||
crate::mcp::tool_def::<UnpinFileTool>(
|
||||
"unpin_file",
|
||||
"Unpin a file from the context workspace.",
|
||||
),
|
||||
crate::mcp::tool_def::<ListPinnedFilesTool>(
|
||||
"list_pinned_files",
|
||||
"List all currently pinned files.",
|
||||
),
|
||||
crate::mcp::tool_def::<AddSessionSummaryTool>(
|
||||
"add_session_summary",
|
||||
"Add a summary of the current session.",
|
||||
),
|
||||
crate::mcp::tool_def::<GetProjectTimelineTool>(
|
||||
"get_project_timeline",
|
||||
"Get a timeline of major project events.",
|
||||
),
|
||||
crate::mcp::tool_def::<LeaveHandoffMemoTool>(
|
||||
"leave_handoff_memo",
|
||||
"Leave a memo for the next session or agent.",
|
||||
),
|
||||
crate::mcp::tool_def::<ReadHandoffMemosTool>(
|
||||
"read_handoff_memos",
|
||||
"Read pending handoff memos.",
|
||||
),
|
||||
crate::mcp::tool_def::<ClearHandoffMemosTool>(
|
||||
"clear_handoff_memos",
|
||||
"Clear handoff memos after reading.",
|
||||
),
|
||||
crate::mcp::tool_def::<UpdateEnvFingerprintTool>(
|
||||
"update_env_fingerprint",
|
||||
"Update the environment fingerprint (e.g., OS, tool versions).",
|
||||
),
|
||||
crate::mcp::tool_def::<ReadEnvFingerprintTool>(
|
||||
"read_env_fingerprint",
|
||||
"Read the current environment fingerprint.",
|
||||
),
|
||||
crate::mcp::tool_def::<LogEnvRequirementTool>(
|
||||
"log_env_requirement",
|
||||
"Log a required tool or package for the environment.",
|
||||
),
|
||||
crate::mcp::tool_def::<AddMilestoneTool>(
|
||||
"add_milestone",
|
||||
"Add a new project milestone.",
|
||||
),
|
||||
crate::mcp::tool_def::<UpdateMilestoneTool>(
|
||||
"update_milestone",
|
||||
"Update the status of a project milestone.",
|
||||
),
|
||||
crate::mcp::tool_def::<ListMilestonesTool>(
|
||||
"list_milestones",
|
||||
"List all project milestones.",
|
||||
),
|
||||
crate::mcp::tool_def::<GenerateStandupReportTool>(
|
||||
"generate_standup_report",
|
||||
"",
|
||||
),
|
||||
crate::mcp::tool_def::<RegisterEnvironmentTool>("register_environment", "Register details about a specific deployment environment."),
|
||||
crate::mcp::tool_def::<RegisterEnvironmentTool>(
|
||||
"register_environment",
|
||||
"Register details about a specific deployment environment.",
|
||||
),
|
||||
crate::mcp::tool_def::<GetEnvironmentDetailsTool>(
|
||||
"get_environment_details",
|
||||
"",
|
||||
),
|
||||
crate::mcp::tool_def::<AddPrChecklistItemTool>("add_pr_checklist_item", "Add an item to the PR checklist."),
|
||||
crate::mcp::tool_def::<GetPrChecklistTool>("get_pr_checklist", "Get the current PR checklist."),
|
||||
crate::mcp::tool_def::<ClearPrChecklistTool>("clear_pr_checklist", "Clear the PR checklist."),
|
||||
crate::mcp::tool_def::<LogTechDebtTool>("log_tech_debt", "Log identified technical debt."),
|
||||
crate::mcp::tool_def::<ResolveTechDebtTool>("resolve_tech_debt", "Mark a logged technical debt as resolved."),
|
||||
crate::mcp::tool_def::<ListTechDebtTool>("list_tech_debt", "List all unresolved technical debt."),
|
||||
crate::mcp::tool_def::<SaveContextWorkspaceTool>("save_context_workspace", "Save the current set of pinned files and context."),
|
||||
crate::mcp::tool_def::<LoadContextWorkspaceTool>("load_context_workspace", "Load a previously saved context workspace."),
|
||||
crate::mcp::tool_def::<AddPrChecklistItemTool>(
|
||||
"add_pr_checklist_item",
|
||||
"Add an item to the PR checklist.",
|
||||
),
|
||||
crate::mcp::tool_def::<GetPrChecklistTool>(
|
||||
"get_pr_checklist",
|
||||
"Get the current PR checklist.",
|
||||
),
|
||||
crate::mcp::tool_def::<ClearPrChecklistTool>(
|
||||
"clear_pr_checklist",
|
||||
"Clear the PR checklist.",
|
||||
),
|
||||
crate::mcp::tool_def::<LogTechDebtTool>(
|
||||
"log_tech_debt",
|
||||
"Log identified technical debt.",
|
||||
),
|
||||
crate::mcp::tool_def::<ResolveTechDebtTool>(
|
||||
"resolve_tech_debt",
|
||||
"Mark a logged technical debt as resolved.",
|
||||
),
|
||||
crate::mcp::tool_def::<ListTechDebtTool>(
|
||||
"list_tech_debt",
|
||||
"List all unresolved technical debt.",
|
||||
),
|
||||
crate::mcp::tool_def::<SaveContextWorkspaceTool>(
|
||||
"save_context_workspace",
|
||||
"Save the current set of pinned files and context.",
|
||||
),
|
||||
crate::mcp::tool_def::<LoadContextWorkspaceTool>(
|
||||
"load_context_workspace",
|
||||
"Load a previously saved context workspace.",
|
||||
),
|
||||
crate::mcp::tool_def::<ListContextWorkspacesTool>(
|
||||
"list_context_workspaces",
|
||||
"",
|
||||
),
|
||||
crate::mcp::tool_def::<OmniSearchTool>("omni_search", "Search across all memory sources (graph, tasks, snippets, ADRs, etc.) at once."),
|
||||
crate::mcp::tool_def::<GetProjectHealthTool>("get_project_health", "Get a synthesized health report of the project based on memory data."),
|
||||
crate::mcp::tool_def::<OmniSearchTool>(
|
||||
"omni_search",
|
||||
"Search across all memory sources (graph, tasks, snippets, ADRs, etc.) at once.",
|
||||
),
|
||||
crate::mcp::tool_def::<GetProjectHealthTool>(
|
||||
"get_project_health",
|
||||
"Get a synthesized health report of the project based on memory data.",
|
||||
),
|
||||
];
|
||||
Some(crate::mcp::success(
|
||||
id,
|
||||
@@ -157,7 +339,8 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Object(Default::default()));
|
||||
|
||||
self.state.broadcast_activity(&format!("Agent executed tool: {}", name));
|
||||
self.state
|
||||
.broadcast_activity(&format!("Agent executed tool: {}", name));
|
||||
|
||||
let result: Result<String, String> = match name {
|
||||
"query_graph_path" => {
|
||||
@@ -166,7 +349,8 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
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();
|
||||
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());
|
||||
@@ -186,12 +370,21 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
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()));
|
||||
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)));
|
||||
parents.insert(
|
||||
rel.from.clone(),
|
||||
(
|
||||
current.clone(),
|
||||
format!("inverse({})", rel.relation_type),
|
||||
),
|
||||
);
|
||||
queue.push_back(rel.from.clone());
|
||||
nodes_at_next_depth += 1;
|
||||
}
|
||||
@@ -215,7 +408,10 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
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))
|
||||
Ok(format!(
|
||||
"No path found between {} and {} within depth {}",
|
||||
req.start_node, req.end_node, max_depth
|
||||
))
|
||||
}
|
||||
}
|
||||
"create_entities" => {
|
||||
@@ -248,8 +444,7 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
let full = self.state.get_full_graph();
|
||||
self.state.write_to_local_delta(|g| {
|
||||
for o in req.observations {
|
||||
if let Some(full_e) = full.entities.get(&o.entity_name)
|
||||
{
|
||||
if let Some(full_e) = full.entities.get(&o.entity_name) {
|
||||
let mut e =
|
||||
g.entities.get(&o.entity_name).cloned().unwrap_or_else(
|
||||
|| Entity {
|
||||
@@ -284,8 +479,7 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
let req = parse_tool!(args.clone(), id, DeleteObservationsTool);
|
||||
self.state.apply_sync_write(|master| {
|
||||
for d in req.deletions {
|
||||
if let Some(e) = master.entities.get_mut(&d.entity_name)
|
||||
{
|
||||
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));
|
||||
}
|
||||
@@ -335,7 +529,8 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
let full = self.state.get_full_graph();
|
||||
for (id, doc_type) in matches {
|
||||
if doc_type == "entity"
|
||||
&& let Some(e) = full.entities.get(&id) {
|
||||
&& let Some(e) = full.entities.get(&id)
|
||||
{
|
||||
result.entities.insert(id, e.clone());
|
||||
}
|
||||
}
|
||||
@@ -502,7 +697,7 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
git_branch: req.git_branch,
|
||||
parent_id: parent_id,
|
||||
parent_id,
|
||||
dependencies: deps,
|
||||
acceptance_criteria: vec![],
|
||||
};
|
||||
@@ -527,21 +722,24 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
while added_new {
|
||||
added_new = false;
|
||||
for t in tasks.iter() {
|
||||
if let Some(pid) = &t.parent_id {
|
||||
if to_delete.contains(pid) && !to_delete.contains(&t.id) {
|
||||
if let Some(pid) = &t.parent_id
|
||||
&& to_delete.contains(pid) && !to_delete.contains(&t.id) {
|
||||
to_delete.insert(t.id.clone());
|
||||
added_new = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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())
|
||||
Ok(vec![
|
||||
format!("Deleted task and its children ({} total).", deleted_count)
|
||||
.to_string(),
|
||||
][0]
|
||||
.clone())
|
||||
} else {
|
||||
Ok(vec!["Task not found.".to_string()][0].clone())
|
||||
}
|
||||
@@ -556,20 +754,24 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
self.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) {
|
||||
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; }
|
||||
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) {
|
||||
if t.acceptance_criteria.iter().any(|c| !c.is_met) {
|
||||
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();
|
||||
}
|
||||
blocker_details =
|
||||
"Unmet acceptance criteria exist.".to_string();
|
||||
}
|
||||
|
||||
// 2. Check dependencies
|
||||
@@ -577,30 +779,41 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
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) {
|
||||
if dep_task.status != "completed" && dep_task.status != "done" {
|
||||
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(", "));
|
||||
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)) {
|
||||
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(", "));
|
||||
blocker_details = format!(
|
||||
"Blocked by child tasks: {}",
|
||||
uncompleted_children.join(", ")
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -609,7 +822,10 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
// 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().as_secs();
|
||||
t.updated_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs();
|
||||
}
|
||||
|
||||
// Cascade cancellation to children
|
||||
@@ -619,7 +835,9 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
while i < to_cancel.len() {
|
||||
let current_pid = to_cancel[i].clone();
|
||||
for t in tasks.iter_mut() {
|
||||
if t.parent_id.as_ref() == Some(¤t_pid) && t.status != "completed" {
|
||||
if t.parent_id.as_ref() == Some(¤t_pid)
|
||||
&& t.status != "completed"
|
||||
{
|
||||
t.status = target_status.clone();
|
||||
to_cancel.push(t.id.clone());
|
||||
}
|
||||
@@ -631,7 +849,13 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
});
|
||||
|
||||
if blocked {
|
||||
Ok(vec![format!("Error: Cannot transition task. {}", blocker_details)].into_iter().next().unwrap())
|
||||
Ok(vec![format!(
|
||||
"Error: Cannot transition task. {}",
|
||||
blocker_details
|
||||
)]
|
||||
.into_iter()
|
||||
.next()
|
||||
.unwrap())
|
||||
} else if found {
|
||||
Ok(vec!["Task status updated.".to_string()][0].clone())
|
||||
} else {
|
||||
@@ -655,18 +879,30 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
let req = parse_tool!(args.clone(), id, SetAcceptanceCriteriaTool);
|
||||
let mut success = false;
|
||||
self.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 {
|
||||
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().as_secs();
|
||||
})
|
||||
.collect();
|
||||
task.updated_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs();
|
||||
success = true;
|
||||
}
|
||||
});
|
||||
if success {
|
||||
Ok(vec!["Acceptance criteria set successfully.".to_string()][0].clone())
|
||||
Ok(
|
||||
vec!["Acceptance criteria set successfully.".to_string()][0]
|
||||
.clone(),
|
||||
)
|
||||
} else {
|
||||
Ok(vec!["Task not found.".to_string()][0].clone())
|
||||
}
|
||||
@@ -676,24 +912,37 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
let mut success = false;
|
||||
let mut already_met = false;
|
||||
self.state.tasks.modify(|tasks| {
|
||||
if let Some(task) = tasks.iter_mut().find(|t| t.id == req.task_id) {
|
||||
if let Some(ac) = task.acceptance_criteria.iter_mut().find(|c| c.id == req.criteria || c.description == req.criteria) {
|
||||
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().as_secs();
|
||||
}
|
||||
task.updated_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs();
|
||||
}
|
||||
}
|
||||
});
|
||||
if success {
|
||||
Ok(vec![format!("Acceptance criteria verified with proof: {}", req.proof)][0].clone())
|
||||
Ok(vec![format!(
|
||||
"Acceptance criteria verified with proof: {}",
|
||||
req.proof
|
||||
)][0]
|
||||
.clone())
|
||||
} else if already_met {
|
||||
Ok(vec!["Acceptance criteria was already met.".to_string()][0].clone())
|
||||
} else {
|
||||
Ok(vec!["Acceptance criteria or task not found.".to_string()][0].clone())
|
||||
Ok(
|
||||
vec!["Acceptance criteria or task not found.".to_string()][0]
|
||||
.clone(),
|
||||
)
|
||||
}
|
||||
}
|
||||
"store_snippet" => {
|
||||
@@ -1224,7 +1473,8 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
let full = self.state.get_full_graph();
|
||||
for (id, doc_type) in &matches {
|
||||
if doc_type == "entity"
|
||||
&& let Some(e) = full.entities.get(id) {
|
||||
&& let Some(e) = full.entities.get(id)
|
||||
{
|
||||
kg.entities.insert(id.clone(), e.clone());
|
||||
}
|
||||
}
|
||||
@@ -1345,23 +1595,19 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
};
|
||||
|
||||
match result {
|
||||
Ok(text) => {
|
||||
Some(crate::mcp::success(
|
||||
Ok(text) => Some(crate::mcp::success(
|
||||
id,
|
||||
serde_json::json!({
|
||||
"content": [{ "type": "text", "text": text }]
|
||||
}),
|
||||
))
|
||||
}
|
||||
Err(e) => {
|
||||
Some(crate::mcp::success(
|
||||
)),
|
||||
Err(e) => Some(crate::mcp::success(
|
||||
id,
|
||||
serde_json::json!({
|
||||
"isError": true,
|
||||
"content": [{ "type": "text", "text": e }]
|
||||
}),
|
||||
))
|
||||
}
|
||||
)),
|
||||
}
|
||||
}
|
||||
_ => {
|
||||
@@ -1381,13 +1627,19 @@ crate::mcp::tool_def::<CreateEntitiesTool>("create_entities", "Create new entiti
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::sync::Arc;
|
||||
use crate::state::MemoryState;
|
||||
use serde_json::json;
|
||||
use std::sync::Arc;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_handle_initialize() {
|
||||
let store_dir = std::env::temp_dir().join(format!("mcp_test_handlers_{}", std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs()));
|
||||
let store_dir = std::env::temp_dir().join(format!(
|
||||
"mcp_test_handlers_{}",
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs()
|
||||
));
|
||||
std::fs::create_dir_all(&store_dir).unwrap();
|
||||
let redb_path = store_dir.join("mcp_store.redb");
|
||||
let db = Arc::new(redb::Database::create(&redb_path).unwrap());
|
||||
@@ -1401,7 +1653,9 @@ mod tests {
|
||||
let state = Arc::new(MemoryState {
|
||||
base_dir: store_dir.clone(),
|
||||
graph: crate::store::Store::new("knowledge_graph_master", db.clone()),
|
||||
search_index: std::sync::RwLock::new(crate::search::MemoryIndex::new(&store_dir).unwrap()),
|
||||
search_index: std::sync::RwLock::new(
|
||||
crate::search::MemoryIndex::new(&store_dir).unwrap(),
|
||||
),
|
||||
ledger: crate::store::Store::new("audit_ledger", db.clone()),
|
||||
sticky: crate::store::Store::new("sticky_notes", db.clone()),
|
||||
tasks: crate::store::Store::new("tasks", db.clone()),
|
||||
@@ -1419,7 +1673,8 @@ mod tests {
|
||||
pr_checklists: crate::store::Store::new("pr_checklists", db.clone()),
|
||||
tech_debts: crate::store::Store::new("tech_debts", db.clone()),
|
||||
gates: crate::store::Store::new("gates", db.clone()),
|
||||
context_workspaces: crate::store::Store::new("context_workspaces", db.clone()), activity_tx: tokio::sync::broadcast::channel(100).0,
|
||||
context_workspaces: crate::store::Store::new("context_workspaces", db.clone()),
|
||||
activity_tx: tokio::sync::broadcast::channel(100).0,
|
||||
});
|
||||
let handler = MemoryHandler { state };
|
||||
|
||||
@@ -1437,7 +1692,10 @@ mod tests {
|
||||
}
|
||||
});
|
||||
|
||||
let response = handler.handle_request(req).await.expect("Expected a response");
|
||||
let response = handler
|
||||
.handle_request(req)
|
||||
.await
|
||||
.expect("Expected a response");
|
||||
|
||||
assert_eq!(response["id"], 1);
|
||||
assert!(response.get("result").is_some());
|
||||
@@ -1452,7 +1710,11 @@ mod tests {
|
||||
}
|
||||
|
||||
fn setup_test_handler(test_name: &str) -> MemoryHandler {
|
||||
let store_dir = std::env::temp_dir().join(format!("mcp_test_handlers_{}_{}", test_name, uuid::Uuid::new_v4()));
|
||||
let store_dir = std::env::temp_dir().join(format!(
|
||||
"mcp_test_handlers_{}_{}",
|
||||
test_name,
|
||||
uuid::Uuid::new_v4()
|
||||
));
|
||||
std::fs::create_dir_all(&store_dir).unwrap();
|
||||
let redb_path = store_dir.join("mcp_store.redb");
|
||||
let db = Arc::new(redb::Database::create(&redb_path).unwrap());
|
||||
@@ -1466,7 +1728,9 @@ mod tests {
|
||||
base_dir: store_dir.clone(),
|
||||
graph: crate::store::Store::new("knowledge_graph_master", db.clone()),
|
||||
|
||||
search_index: std::sync::RwLock::new(crate::search::MemoryIndex::new(&store_dir).unwrap()),
|
||||
search_index: std::sync::RwLock::new(
|
||||
crate::search::MemoryIndex::new(&store_dir).unwrap(),
|
||||
),
|
||||
ledger: crate::store::Store::new("audit_ledger", db.clone()),
|
||||
sticky: crate::store::Store::new("sticky_notes", db.clone()),
|
||||
tasks: crate::store::Store::new("tasks", db.clone()),
|
||||
@@ -1484,7 +1748,8 @@ mod tests {
|
||||
pr_checklists: crate::store::Store::new("pr_checklists", db.clone()),
|
||||
tech_debts: crate::store::Store::new("tech_debts", db.clone()),
|
||||
gates: crate::store::Store::new("gates", db.clone()),
|
||||
context_workspaces: crate::store::Store::new("context_workspaces", db.clone()), activity_tx: tokio::sync::broadcast::channel(100).0,
|
||||
context_workspaces: crate::store::Store::new("context_workspaces", db.clone()),
|
||||
activity_tx: tokio::sync::broadcast::channel(100).0,
|
||||
});
|
||||
MemoryHandler { state }
|
||||
}
|
||||
@@ -1500,15 +1765,26 @@ mod tests {
|
||||
"params": {}
|
||||
});
|
||||
|
||||
let response = handler.handle_request(req).await.expect("Expected a response");
|
||||
let response = handler
|
||||
.handle_request(req)
|
||||
.await
|
||||
.expect("Expected a response");
|
||||
assert_eq!(response["id"], 2);
|
||||
|
||||
let tools = response["result"]["tools"].as_array().expect("Tools must be an array");
|
||||
let tools = response["result"]["tools"]
|
||||
.as_array()
|
||||
.expect("Tools must be an array");
|
||||
assert!(!tools.is_empty());
|
||||
|
||||
// Verify a specific tool is registered
|
||||
let add_task_tool = tools.iter().find(|t| t["name"] == "add_task").expect("add_task tool missing");
|
||||
assert_eq!(add_task_tool["description"], "Add a new task to the task tracker.");
|
||||
let add_task_tool = tools
|
||||
.iter()
|
||||
.find(|t| t["name"] == "add_task")
|
||||
.expect("add_task tool missing");
|
||||
assert_eq!(
|
||||
add_task_tool["description"],
|
||||
"Add a new task to the task tracker."
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -1529,12 +1805,20 @@ mod tests {
|
||||
}
|
||||
});
|
||||
|
||||
let response = handler.handle_request(req).await.expect("Expected a response");
|
||||
let response = handler
|
||||
.handle_request(req)
|
||||
.await
|
||||
.expect("Expected a response");
|
||||
assert_eq!(response["id"], 3);
|
||||
|
||||
let content = &response["result"]["content"][0];
|
||||
assert_eq!(content["type"], "text");
|
||||
assert!(content["text"].as_str().unwrap().starts_with("Task added with ID: "));
|
||||
assert!(
|
||||
content["text"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.starts_with("Task added with ID: ")
|
||||
);
|
||||
|
||||
// Verify task was actually added to store
|
||||
let tasks = handler.state.tasks.read();
|
||||
@@ -1566,7 +1850,10 @@ mod tests {
|
||||
}
|
||||
});
|
||||
|
||||
let response = handler.handle_request(req).await.expect("Expected a response");
|
||||
let response = handler
|
||||
.handle_request(req)
|
||||
.await
|
||||
.expect("Expected a response");
|
||||
assert_eq!(response["id"], 4);
|
||||
|
||||
let content = &response["result"]["content"][0];
|
||||
@@ -1574,7 +1861,10 @@ mod tests {
|
||||
|
||||
// Verify entity was actually added to state
|
||||
let session_graph = handler.state.graph.read();
|
||||
let entity = session_graph.entities.get("MemoryHandler").expect("Entity should be in session graph");
|
||||
let entity = session_graph
|
||||
.entities
|
||||
.get("MemoryHandler")
|
||||
.expect("Entity should be in session graph");
|
||||
assert_eq!(entity.entity_type, "struct");
|
||||
assert_eq!(entity.observations, vec!["Handles MCP requests natively"]);
|
||||
assert_eq!(entity.namespace, "core".to_string());
|
||||
@@ -1599,7 +1889,10 @@ mod tests {
|
||||
}
|
||||
});
|
||||
|
||||
let response = handler.handle_request(req).await.expect("Expected a response");
|
||||
let response = handler
|
||||
.handle_request(req)
|
||||
.await
|
||||
.expect("Expected a response");
|
||||
assert_eq!(response["id"], 5);
|
||||
|
||||
let snippets = handler.state.snippets.read();
|
||||
@@ -1624,7 +1917,10 @@ mod tests {
|
||||
}
|
||||
});
|
||||
|
||||
let response = handler.handle_request(req).await.expect("Expected a response");
|
||||
let response = handler
|
||||
.handle_request(req)
|
||||
.await
|
||||
.expect("Expected a response");
|
||||
assert_eq!(response["id"], 6);
|
||||
|
||||
let notes = handler.state.sticky.read();
|
||||
@@ -1666,13 +1962,16 @@ mod tests {
|
||||
let handler = setup_test_handler("add_observations");
|
||||
// Pre-populate entity
|
||||
handler.state.graph.modify(|session| {
|
||||
session.entities.insert("NodeA".to_string(), crate::models::Entity {
|
||||
session.entities.insert(
|
||||
"NodeA".to_string(),
|
||||
crate::models::Entity {
|
||||
name: "NodeA".to_string(),
|
||||
entity_type: "class".to_string(),
|
||||
observations: vec!["Initial".to_string()],
|
||||
namespace: "".to_string(),
|
||||
git_branch: None,
|
||||
});
|
||||
},
|
||||
);
|
||||
});
|
||||
let req = json!({
|
||||
"jsonrpc": "2.0",
|
||||
@@ -1700,13 +1999,16 @@ mod tests {
|
||||
async fn test_handle_delete_entities() {
|
||||
let handler = setup_test_handler("delete_entities");
|
||||
handler.state.graph.modify(|session| {
|
||||
session.entities.insert("ToDelete".to_string(), crate::models::Entity {
|
||||
session.entities.insert(
|
||||
"ToDelete".to_string(),
|
||||
crate::models::Entity {
|
||||
name: "ToDelete".to_string(),
|
||||
entity_type: "var".to_string(),
|
||||
observations: vec![],
|
||||
namespace: "".to_string(),
|
||||
git_branch: None,
|
||||
});
|
||||
},
|
||||
);
|
||||
});
|
||||
// Force flush session to master
|
||||
handler.state.apply_sync_write(|_| {});
|
||||
@@ -1731,13 +2033,16 @@ mod tests {
|
||||
async fn test_handle_delete_observations() {
|
||||
let handler = setup_test_handler("delete_observations");
|
||||
handler.state.graph.modify(|session| {
|
||||
session.entities.insert("NodeA".to_string(), crate::models::Entity {
|
||||
session.entities.insert(
|
||||
"NodeA".to_string(),
|
||||
crate::models::Entity {
|
||||
name: "NodeA".to_string(),
|
||||
entity_type: "class".to_string(),
|
||||
observations: vec!["Keep".to_string(), "Drop".to_string()],
|
||||
namespace: "".to_string(),
|
||||
git_branch: None,
|
||||
});
|
||||
},
|
||||
);
|
||||
});
|
||||
handler.state.apply_sync_write(|_| {});
|
||||
let req = json!({
|
||||
@@ -1797,7 +2102,10 @@ mod tests {
|
||||
description: "".to_string(),
|
||||
created_at: 0,
|
||||
updated_at: 0,
|
||||
git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None,
|
||||
git_branch: None,
|
||||
acceptance_criteria: vec![],
|
||||
dependencies: vec![],
|
||||
parent_id: None,
|
||||
});
|
||||
tasks.push(crate::models::Task {
|
||||
id: "2".to_string(),
|
||||
@@ -1806,7 +2114,10 @@ mod tests {
|
||||
description: "".to_string(),
|
||||
created_at: 0,
|
||||
updated_at: 0,
|
||||
git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None,
|
||||
git_branch: None,
|
||||
acceptance_criteria: vec![],
|
||||
dependencies: vec![],
|
||||
parent_id: None,
|
||||
});
|
||||
});
|
||||
|
||||
@@ -1927,13 +2238,16 @@ mod tests {
|
||||
async fn test_handle_read_graph() {
|
||||
let handler = setup_test_handler("read_graph");
|
||||
handler.state.graph.modify(|session| {
|
||||
session.entities.insert("NodeA".to_string(), crate::models::Entity {
|
||||
session.entities.insert(
|
||||
"NodeA".to_string(),
|
||||
crate::models::Entity {
|
||||
name: "NodeA".to_string(),
|
||||
entity_type: "var".to_string(),
|
||||
observations: vec![],
|
||||
namespace: "".to_string(),
|
||||
git_branch: None,
|
||||
});
|
||||
},
|
||||
);
|
||||
});
|
||||
handler.state.apply_sync_write(|_| {});
|
||||
|
||||
@@ -1962,7 +2276,9 @@ mod tests {
|
||||
git_branch: None,
|
||||
};
|
||||
handler.state.graph.modify(|session| {
|
||||
session.entities.insert("UserRepository".to_string(), entity);
|
||||
session
|
||||
.entities
|
||||
.insert("UserRepository".to_string(), entity);
|
||||
});
|
||||
handler.state.apply_sync_write(|_| {});
|
||||
|
||||
@@ -2322,7 +2638,10 @@ mod tests {
|
||||
description: "".to_string(),
|
||||
created_at: 0,
|
||||
updated_at: 0,
|
||||
git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None,
|
||||
git_branch: None,
|
||||
acceptance_criteria: vec![],
|
||||
dependencies: vec![],
|
||||
parent_id: None,
|
||||
});
|
||||
});
|
||||
|
||||
@@ -2343,7 +2662,3 @@ mod tests {
|
||||
assert_eq!(tasks[0].status, "in_progress");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
+129
-65
@@ -1,4 +1,7 @@
|
||||
#![cfg_attr(not(target_os = "windows"), allow(dead_code, unused_imports, unreachable_code))]
|
||||
#![cfg_attr(
|
||||
not(target_os = "windows"),
|
||||
allow(dead_code, unused_imports, unreachable_code)
|
||||
)]
|
||||
|
||||
mod handlers;
|
||||
mod mcp;
|
||||
@@ -13,11 +16,11 @@ use crate::models::*;
|
||||
use crate::state::MemoryState;
|
||||
use crate::store::Store;
|
||||
|
||||
use redb::ReadableTable;
|
||||
use std::fs;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::{Arc, RwLock};
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
use redb::ReadableTable;
|
||||
use tokio::time::sleep;
|
||||
|
||||
use clap::{Parser, Subcommand};
|
||||
@@ -83,7 +86,6 @@ enum GateCommands {
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async fn reconcile_worker(state: Arc<MemoryState>) {
|
||||
loop {
|
||||
sleep(Duration::from_secs(5)).await;
|
||||
@@ -106,14 +108,18 @@ async fn reconcile_worker(state: Arc<MemoryState>) {
|
||||
let state_clone = state.clone();
|
||||
let _ = tokio::task::spawn_blocking(move || {
|
||||
state_clone.rebuild_index();
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
use axum::{
|
||||
Json, Router,
|
||||
extract::{Query, State, ws::{WebSocket, Message}},
|
||||
extract::{
|
||||
Query, State,
|
||||
ws::{Message, WebSocket},
|
||||
},
|
||||
response::IntoResponse,
|
||||
routing::{get, post},
|
||||
};
|
||||
@@ -186,9 +192,11 @@ async fn gate_verify_handler(
|
||||
(axum::http::StatusCode::FORBIDDEN, msg).into_response()
|
||||
}
|
||||
}
|
||||
None => {
|
||||
(axum::http::StatusCode::NOT_FOUND, "Action not yet authorized (no gate record found).").into_response()
|
||||
}
|
||||
None => (
|
||||
axum::http::StatusCode::NOT_FOUND,
|
||||
"Action not yet authorized (no gate record found).",
|
||||
)
|
||||
.into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -229,18 +237,23 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
rt.block_on(async {
|
||||
tokio::spawn(reconcile_worker(Arc::clone(&state)));
|
||||
let app_state = Arc::new(AppState {
|
||||
handler: Arc::new(MemoryHandler { state: Arc::clone(&state) }),
|
||||
handler: Arc::new(MemoryHandler {
|
||||
state: Arc::clone(&state),
|
||||
}),
|
||||
clients: RwLock::new(HashMap::new()),
|
||||
next_id: AtomicUsize::new(1),
|
||||
});
|
||||
|
||||
let app = Router::new()
|
||||
.route("/api/version", get(|| async move {
|
||||
.route(
|
||||
"/api/version",
|
||||
get(|| async move {
|
||||
axum::Json(serde_json::json!({
|
||||
"version": env!("APP_VERSION"),
|
||||
"git_hash": option_env!("GIT_HASH").unwrap_or("unknown")
|
||||
}))
|
||||
}))
|
||||
}),
|
||||
)
|
||||
.route("/ws", get(ws_handler))
|
||||
.route("/health", get(health_handler))
|
||||
.route("/nvim/telemetry", post(nvim_telemetry_handler))
|
||||
@@ -260,15 +273,19 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
"/",
|
||||
get(|| async move { axum::response::Html(include_str!("dashboard.html")) }),
|
||||
)
|
||||
.route("/api/graph", get({
|
||||
.route(
|
||||
"/api/graph",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move || async move {
|
||||
let graph = state_clone.get_full_graph();
|
||||
axum::Json(graph)
|
||||
}
|
||||
}))
|
||||
|
||||
.route("/api/tasks/{id}/complete", post({
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/tasks/{id}/complete",
|
||||
post({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move |axum::extract::Path(id): axum::extract::Path<String>| async move {
|
||||
state_clone.tasks.modify(|tasks| {
|
||||
@@ -281,28 +298,38 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
});
|
||||
axum::Json(serde_json::json!({"status": "success"}))
|
||||
}
|
||||
}))
|
||||
|
||||
.route("/api/tasks", get({
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/tasks",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move || async move {
|
||||
let tasks = state_clone.tasks.read();
|
||||
axum::Json(tasks.clone())
|
||||
}
|
||||
}))
|
||||
.route("/api/sticky", get({
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/sticky",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move || async move {
|
||||
let sticky = state_clone.sticky.read();
|
||||
axum::Json(sticky.clone())
|
||||
}
|
||||
}))
|
||||
.route("/api/search", get({
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/search",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move |axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>| async move {
|
||||
if let Some(q) = params.get("q") {
|
||||
if let Ok(idx) = state_clone.search_index.read() {
|
||||
if let Ok(results) = idx.search(q, None) {
|
||||
move |axum::extract::Query(params): axum::extract::Query<
|
||||
std::collections::HashMap<String, String>,
|
||||
>| async move {
|
||||
if let Some(q) = params.get("q")
|
||||
&& let Ok(idx) = state_clone.search_index.read()
|
||||
&& let Ok(results) = idx.search(q, None) {
|
||||
let mut formatted_results = Vec::new();
|
||||
for (type_name, content) in results {
|
||||
formatted_results.push(serde_json::json!({
|
||||
@@ -311,14 +338,16 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
"score": 1.0
|
||||
}));
|
||||
}
|
||||
return axum::Json(serde_json::json!({ "results": formatted_results }));
|
||||
}
|
||||
}
|
||||
return axum::Json(
|
||||
serde_json::json!({ "results": formatted_results }),
|
||||
);
|
||||
}
|
||||
axum::Json(serde_json::json!({ "results": [] }))
|
||||
}
|
||||
}))
|
||||
.route("/api/stats",
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/stats",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move || async move {
|
||||
@@ -386,29 +415,39 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
|
||||
async fn ws_handler(
|
||||
ws: axum::extract::ws::WebSocketUpgrade,
|
||||
headers: axum::http::HeaderMap,
|
||||
_headers: axum::http::HeaderMap,
|
||||
axum::extract::State(state): axum::extract::State<Arc<AppState>>,
|
||||
axum::extract::Query(query): axum::extract::Query<std::collections::HashMap<String, String>>,
|
||||
) -> axum::response::Response {
|
||||
let client_type = query.get("client").cloned().unwrap_or_else(|| "unknown".to_string());
|
||||
ws.on_upgrade(move |socket| handle_socket(socket, state, client_type)).into_response()
|
||||
let client_type = query
|
||||
.get("client")
|
||||
.cloned()
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
ws.on_upgrade(move |socket| handle_socket(socket, state, client_type))
|
||||
.into_response()
|
||||
}
|
||||
|
||||
async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: String) {
|
||||
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
|
||||
let (tx, mut rx) = mpsc::channel::<String>(100);
|
||||
|
||||
state.clients.write().unwrap().insert(session_id.clone(), tx.clone());
|
||||
state
|
||||
.clients
|
||||
.write()
|
||||
.unwrap()
|
||||
.insert(session_id.clone(), tx.clone());
|
||||
|
||||
let (mut sender, mut receiver) = socket.split();
|
||||
|
||||
let mut send_task = tokio::spawn(async move {
|
||||
while let Some(msg) = rx.recv().await {
|
||||
tracing::trace!("Sending message to websocket (length: {}): {}", msg.len(), msg);
|
||||
tracing::trace!(
|
||||
"Sending message to websocket (length: {}): {}",
|
||||
msg.len(),
|
||||
msg
|
||||
);
|
||||
if sender.send(Message::Text(msg.into())).await.is_err() {
|
||||
tracing::error!("Failed to send message to websocket");
|
||||
break;
|
||||
@@ -426,14 +465,21 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
while let Some(msg_result) = receiver.next().await {
|
||||
match msg_result {
|
||||
Ok(Message::Text(text)) => {
|
||||
tracing::info!("Received text message from websocket (length: {})", text.len());
|
||||
tracing::info!(
|
||||
"Received text message from websocket (length: {})",
|
||||
text.len()
|
||||
);
|
||||
tracing::trace!("Message content: {}", text);
|
||||
if let Ok(payload) = serde_json::from_str::<serde_json::Value>(&text) {
|
||||
if client_type == "proxy" {
|
||||
// Send activity broadcast to UI clients
|
||||
if let Some(method) = payload.get("method").and_then(|m| m.as_str()) {
|
||||
if method == "tools/call" {
|
||||
let name = payload.get("params").and_then(|p| p.get("name")).and_then(|n| n.as_str()).unwrap_or("unknown_tool");
|
||||
if let Some(method) = payload.get("method").and_then(|m| m.as_str())
|
||||
&& method == "tools/call" {
|
||||
let name = payload
|
||||
.get("params")
|
||||
.and_then(|p| p.get("name"))
|
||||
.and_then(|n| n.as_str())
|
||||
.unwrap_or("unknown_tool");
|
||||
let activity_msg = format!("Agent executed tool: {}", name);
|
||||
|
||||
let event = serde_json::json!({
|
||||
@@ -448,24 +494,39 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} // End if proxy
|
||||
|
||||
// Process MCP request
|
||||
if let Some(response) = handler.handle_request(payload).await {
|
||||
let res_str = serde_json::to_string(&response).unwrap();
|
||||
let tx_opt = state_clone.clients.read().unwrap().get(&session_id_clone).cloned();
|
||||
let tx_opt = state_clone
|
||||
.clients
|
||||
.read()
|
||||
.unwrap()
|
||||
.get(&session_id_clone)
|
||||
.cloned();
|
||||
if let Some(client_tx) = tx_opt {
|
||||
if let Err(e) = client_tx.send(res_str).await {
|
||||
tracing::error!("Failed to send response to client channel for session {}: {}", session_id_clone, e);
|
||||
tracing::error!(
|
||||
"Failed to send response to client channel for session {}: {}",
|
||||
session_id_clone,
|
||||
e
|
||||
);
|
||||
}
|
||||
} else {
|
||||
tracing::warn!("Could not find client_tx for session_id {} when trying to send response", session_id_clone);
|
||||
tracing::warn!(
|
||||
"Could not find client_tx for session_id {} when trying to send response",
|
||||
session_id_clone
|
||||
);
|
||||
}
|
||||
}
|
||||
} // End if let Ok(payload)
|
||||
}
|
||||
// End if let Ok(payload)
|
||||
else {
|
||||
tracing::warn!("Failed to parse payload as JSON from websocket message: {}", text);
|
||||
tracing::warn!(
|
||||
"Failed to parse payload as JSON from websocket message: {}",
|
||||
text
|
||||
);
|
||||
}
|
||||
} // End Ok(Message::Text(text))
|
||||
Ok(other) => {
|
||||
@@ -477,7 +538,10 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
}
|
||||
}
|
||||
}
|
||||
tracing::info!("Websocket receiver task ended for session {}", session_id_clone);
|
||||
tracing::info!(
|
||||
"Websocket receiver task ended for session {}",
|
||||
session_id_clone
|
||||
);
|
||||
});
|
||||
|
||||
tokio::select! {
|
||||
@@ -492,10 +556,12 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
};
|
||||
|
||||
state.clients.write().unwrap().remove(&session_id);
|
||||
tracing::info!("Websocket session {} closed and removed from state", session_id);
|
||||
tracing::info!(
|
||||
"Websocket session {} closed and removed from state",
|
||||
session_id
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
#[derive(serde::Deserialize, serde::Serialize, Debug)]
|
||||
pub struct NvimTelemetry {
|
||||
pub session_id: String,
|
||||
@@ -510,8 +576,10 @@ async fn nvim_telemetry_handler(
|
||||
axum::Json(payload): axum::Json<NvimTelemetry>,
|
||||
) -> impl axum::response::IntoResponse {
|
||||
// 1. Update active_nvim.txt if FocusGained, BufEnter, or VimEnter
|
||||
if payload.event == "FocusGained" || payload.event == "BufEnter" || payload.event == "VimEnter" {
|
||||
let profile = std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into());
|
||||
if payload.event == "FocusGained" || payload.event == "BufEnter" || payload.event == "VimEnter"
|
||||
{
|
||||
let profile =
|
||||
std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into());
|
||||
let win_path = format!("{}\\.gemini\\active_nvim.txt", profile);
|
||||
let _ = std::fs::write(&win_path, &payload.session_id);
|
||||
|
||||
@@ -527,7 +595,7 @@ async fn nvim_telemetry_handler(
|
||||
|
||||
let msg_str = ws_msg.to_string();
|
||||
let clients = state.clients.read().unwrap().clone();
|
||||
for (_, tx) in clients.iter() {
|
||||
for tx in clients.values() {
|
||||
let _ = tx.send(msg_str.clone()).await;
|
||||
}
|
||||
|
||||
@@ -611,7 +679,6 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
let base_dir = std::env::var("MCP_MEMORY_STORE_DIR").unwrap_or_else(|_| {
|
||||
dirs::home_dir()
|
||||
.map(|mut h| {
|
||||
@@ -657,13 +724,14 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
for (key, file_name) in stores.iter() {
|
||||
if table.get(*key).unwrap().is_none() {
|
||||
let json_path = base.join(file_name);
|
||||
if json_path.exists() {
|
||||
if let Ok(data) = fs::read(&json_path) {
|
||||
if serde_json::from_slice::<serde_json::Value>(&data).is_ok() {
|
||||
if json_path.exists()
|
||||
&& let Ok(data) = fs::read(&json_path)
|
||||
&& serde_json::from_slice::<serde_json::Value>(&data).is_ok() {
|
||||
table.insert(*key, data.as_slice()).unwrap();
|
||||
let _ = fs::rename(&json_path, json_path.with_extension("json.migrated"));
|
||||
}
|
||||
}
|
||||
let _ = fs::rename(
|
||||
&json_path,
|
||||
json_path.with_extension("json.migrated"),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -700,7 +768,3 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
|
||||
run_server(state)
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use schemars::JsonSchema;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use schemars::JsonSchema;
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct CodeChange {
|
||||
|
||||
@@ -28,7 +28,8 @@ impl MemoryIndex {
|
||||
|
||||
let index_dir = store_dir.join("tantivy_index");
|
||||
std::fs::create_dir_all(&index_dir).unwrap();
|
||||
let index = Index::open_in_dir(&index_dir).unwrap_or_else(|_| Index::create_in_dir(&index_dir, schema.clone()).unwrap());
|
||||
let index = Index::open_in_dir(&index_dir)
|
||||
.unwrap_or_else(|_| Index::create_in_dir(&index_dir, schema.clone()).unwrap());
|
||||
|
||||
let writer = index.writer(50_000_000)?;
|
||||
let reader = index
|
||||
@@ -72,7 +73,6 @@ impl MemoryIndex {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
pub fn commit(&self) -> tantivy::Result<()> {
|
||||
let mut writer = self.writer.lock().unwrap();
|
||||
writer.commit()?;
|
||||
@@ -113,7 +113,9 @@ impl MemoryIndex {
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("");
|
||||
if let Some(ns) = namespace
|
||||
&& doc_ns != ns && doc_ns != "global" {
|
||||
&& doc_ns != ns
|
||||
&& doc_ns != "global"
|
||||
{
|
||||
continue;
|
||||
}
|
||||
results.push((id, doc_type));
|
||||
@@ -172,7 +174,10 @@ mod tests {
|
||||
status: "open".to_string(),
|
||||
created_at: 0,
|
||||
updated_at: 0,
|
||||
git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None,
|
||||
git_branch: None,
|
||||
acceptance_criteria: vec![],
|
||||
dependencies: vec![],
|
||||
parent_id: None,
|
||||
};
|
||||
index.index_task(&task).unwrap();
|
||||
|
||||
|
||||
+7
-5
@@ -33,19 +33,21 @@ pub struct MemoryState {
|
||||
impl MemoryState {
|
||||
pub fn unique_items<T: Eq + std::hash::Hash + Clone>(input: Vec<T>) -> Vec<T> {
|
||||
let mut keys = std::collections::HashSet::new();
|
||||
input.into_iter().filter(|entry| keys.insert(entry.clone())).collect()
|
||||
input
|
||||
.into_iter()
|
||||
.filter(|entry| keys.insert(entry.clone()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn broadcast_activity(&self, message: &str) {
|
||||
let payload = serde_json::json!({
|
||||
"type": "activity",
|
||||
"data": message
|
||||
}).to_string();
|
||||
})
|
||||
.to_string();
|
||||
let _ = self.activity_tx.send(payload);
|
||||
}
|
||||
|
||||
|
||||
|
||||
pub fn get_full_graph(&self) -> KnowledgeGraph {
|
||||
self.graph.read()
|
||||
}
|
||||
@@ -61,7 +63,7 @@ impl MemoryState {
|
||||
pub fn rebuild_index(&self) {
|
||||
if let Ok(new_idx) = MemoryIndex::new(&self.base_dir) {
|
||||
let graph = self.graph.read();
|
||||
for (_, e) in &graph.entities {
|
||||
for e in graph.entities.values() {
|
||||
let _ = new_idx.index_entity(e);
|
||||
}
|
||||
|
||||
|
||||
+17
-7
@@ -22,13 +22,11 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + 'static> Store<T
|
||||
|
||||
fn load_from_db(key: &str, db: &Database) -> T {
|
||||
let read_txn = db.begin_read().unwrap();
|
||||
if let Ok(table) = read_txn.open_table(STORE_TABLE) {
|
||||
if let Ok(Some(value)) = table.get(key) {
|
||||
if let Ok(parsed) = serde_json::from_slice::<T>(value.value()) {
|
||||
if let Ok(table) = read_txn.open_table(STORE_TABLE)
|
||||
&& let Ok(Some(value)) = table.get(key)
|
||||
&& let Ok(parsed) = serde_json::from_slice::<T>(value.value()) {
|
||||
return parsed;
|
||||
}
|
||||
}
|
||||
}
|
||||
T::default()
|
||||
}
|
||||
|
||||
@@ -94,11 +92,23 @@ mod tests {
|
||||
// Need to wait for spawn_blocking to finish
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
|
||||
assert_eq!(store.read(), TestData { name: "Hello".to_string(), value: 42 });
|
||||
assert_eq!(
|
||||
store.read(),
|
||||
TestData {
|
||||
name: "Hello".to_string(),
|
||||
value: 42
|
||||
}
|
||||
);
|
||||
|
||||
// Load again to verify persistence
|
||||
let store2 = Store::<TestData>::new("test_key", db.clone());
|
||||
assert_eq!(store2.read(), TestData { name: "Hello".to_string(), value: 42 });
|
||||
assert_eq!(
|
||||
store2.read(),
|
||||
TestData {
|
||||
name: "Hello".to_string(),
|
||||
value: 42
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
|
||||
|
||||
+28
-13
@@ -3,37 +3,44 @@ use std::collections::HashSet;
|
||||
#[test]
|
||||
fn test_eager_tools_parity() {
|
||||
// 1. Read handlers.rs to get memory tools
|
||||
let memory_source = std::fs::read_to_string("src/handlers.rs").expect("Failed to read handlers.rs");
|
||||
let memory_source =
|
||||
std::fs::read_to_string("src/handlers.rs").expect("Failed to read handlers.rs");
|
||||
let mut memory_tools = HashSet::new();
|
||||
for line in memory_source.lines() {
|
||||
if line.contains("crate::mcp::tool_def") {
|
||||
if let Some(start) = line.find("(\"") {
|
||||
if line.contains("crate::mcp::tool_def")
|
||||
&& let Some(start) = line.find("(\"") {
|
||||
let rest = &line[start + 2..];
|
||||
if let Some(end) = rest.find("\"") {
|
||||
memory_tools.insert(rest[..end].to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
assert!(!memory_tools.is_empty(), "Could not find memory tools in handlers.rs");
|
||||
assert!(
|
||||
!memory_tools.is_empty(),
|
||||
"Could not find memory tools in handlers.rs"
|
||||
);
|
||||
|
||||
// 2. Read nvim-core/src/lib.rs to get nvim tools
|
||||
let nvim_source = std::fs::read_to_string("../nvim-core/src/lib.rs").expect("Failed to read nvim lib.rs");
|
||||
let nvim_source =
|
||||
std::fs::read_to_string("../nvim-core/src/lib.rs").expect("Failed to read nvim lib.rs");
|
||||
let mut nvim_tools = HashSet::new();
|
||||
for line in nvim_source.lines() {
|
||||
if line.contains("\"name\": \"nvim_") {
|
||||
if let Some(start) = line.find("\"name\": \"") {
|
||||
if line.contains("\"name\": \"nvim_")
|
||||
&& let Some(start) = line.find("\"name\": \"") {
|
||||
let rest = &line[start + 9..];
|
||||
if let Some(end) = rest.find("\"") {
|
||||
nvim_tools.insert(rest[..end].to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
assert!(!nvim_tools.is_empty(), "Could not find nvim tools in lib.rs");
|
||||
assert!(
|
||||
!nvim_tools.is_empty(),
|
||||
"Could not find nvim tools in lib.rs"
|
||||
);
|
||||
|
||||
// 3. Read Windows mcp_config.json
|
||||
let win_home = std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".to_string());
|
||||
let win_home =
|
||||
std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".to_string());
|
||||
let win_config_path = std::path::PathBuf::from(win_home).join(".gemini/config/mcp_config.json");
|
||||
if win_config_path.exists() {
|
||||
let config_str = std::fs::read_to_string(&win_config_path).unwrap();
|
||||
@@ -42,14 +49,22 @@ fn test_eager_tools_parity() {
|
||||
if let Some(eager) = config["mcpServers"]["memory"]["eagerTools"].as_array() {
|
||||
for tool in eager {
|
||||
let name = tool.as_str().unwrap();
|
||||
assert!(memory_tools.contains(name), "Windows config Memory tool '{}' not implemented in handlers.rs!", name);
|
||||
assert!(
|
||||
memory_tools.contains(name),
|
||||
"Windows config Memory tool '{}' not implemented in handlers.rs!",
|
||||
name
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(nvim_eager) = config["mcpServers"]["win-nvim"]["eagerTools"].as_array() {
|
||||
for tool in nvim_eager {
|
||||
let name = tool.as_str().unwrap();
|
||||
assert!(nvim_tools.contains(name), "Windows config Nvim tool '{}' not implemented in nvim-core!", name);
|
||||
assert!(
|
||||
nvim_tools.contains(name),
|
||||
"Windows config Nvim tool '{}' not implemented in nvim-core!",
|
||||
name
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+8
-3
@@ -2,19 +2,24 @@ use std::process::Command;
|
||||
|
||||
fn main() {
|
||||
let git_hash = Command::new("git")
|
||||
.args(&["rev-parse", "--short", "HEAD"])
|
||||
.args(["rev-parse", "--short", "HEAD"])
|
||||
.output()
|
||||
.ok()
|
||||
.and_then(|out| String::from_utf8(out.stdout).ok())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
|
||||
let git_date = Command::new("git")
|
||||
.args(&["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"])
|
||||
.args(["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"])
|
||||
.output()
|
||||
.ok()
|
||||
.and_then(|out| String::from_utf8(out.stdout).ok())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
|
||||
let version = format!("{} ({} {})", env!("CARGO_PKG_VERSION"), git_date.trim(), git_hash.trim());
|
||||
let version = format!(
|
||||
"{} ({} {})",
|
||||
env!("CARGO_PKG_VERSION"),
|
||||
git_date.trim(),
|
||||
git_hash.trim()
|
||||
);
|
||||
println!("cargo:rustc-env=APP_VERSION={}", version);
|
||||
}
|
||||
@@ -1,14 +1,15 @@
|
||||
use futures_util::StreamExt;
|
||||
use reqwest::Client;
|
||||
use std::env;
|
||||
use std::time::Duration;
|
||||
use futures_util::StreamExt;
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
tracing_subscriber::fmt::init();
|
||||
|
||||
let target = env::var("MCP_TARGET").unwrap_or_else(|_| "https://127.0.0.1:3000".to_string());
|
||||
let token = env::var("MCP_AUTH_TOKEN").unwrap_or_else(|_| "jP76lUJ5DtFRZmcvXH8LKdCTIkp29eAf".to_string());
|
||||
let token = env::var("MCP_AUTH_TOKEN")
|
||||
.unwrap_or_else(|_| "jP76lUJ5DtFRZmcvXH8LKdCTIkp29eAf".to_string());
|
||||
|
||||
tracing::info!("Starting skeletal client to {}", target);
|
||||
|
||||
@@ -20,10 +21,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
|
||||
tracing::info!("Connecting to SSE: {}", sse_url);
|
||||
|
||||
let res = client.get(&sse_url)
|
||||
.bearer_auth(&token)
|
||||
.send()
|
||||
.await?;
|
||||
let res = client.get(&sse_url).bearer_auth(&token).send().await?;
|
||||
|
||||
if !res.status().is_success() {
|
||||
tracing::error!("Failed to connect to SSE: {}", res.status());
|
||||
@@ -78,7 +76,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
tracing::info!("Sending test payload to {}", post_url);
|
||||
tracing::info!("Payload: {}", payload);
|
||||
|
||||
let post_res = client.post(&post_url)
|
||||
let post_res = client
|
||||
.post(&post_url)
|
||||
.bearer_auth(&token)
|
||||
.header("Content-Type", "application/json")
|
||||
.body(payload.to_string())
|
||||
|
||||
+2
-4
@@ -1,6 +1,6 @@
|
||||
use std::sync::Arc;
|
||||
use clap::Parser;
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use std::sync::Arc;
|
||||
use tokio::io::AsyncBufReadExt;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
@@ -94,7 +94,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
tracing::info!("Attempting to connect to {}", ws_url);
|
||||
|
||||
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
let mut request = match ws_url.clone().into_client_request() {
|
||||
let request = match ws_url.clone().into_client_request() {
|
||||
Ok(req) => req,
|
||||
Err(e) => {
|
||||
tracing::error!("Failed to parse target URL {}: {}", ws_url, e);
|
||||
@@ -165,5 +165,3 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
+29
-14
@@ -1,5 +1,5 @@
|
||||
use serde_json::{json, Value};
|
||||
use std::io::{BufRead, BufReader, Read, Write};
|
||||
use serde_json::{Value, json};
|
||||
use std::io::{BufRead, BufReader, Write};
|
||||
use std::process::{Command, Stdio};
|
||||
use std::time::Duration;
|
||||
|
||||
@@ -18,10 +18,15 @@ fn read_message(reader: &mut impl BufRead) -> Option<Value> {
|
||||
serde_json::from_str(line.trim()).ok()
|
||||
}
|
||||
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_full_system_e2e_performance() {
|
||||
let temp_dir = std::env::temp_dir().join(format!("mcp_e2e_{}", std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs()));
|
||||
let temp_dir = std::env::temp_dir().join(format!(
|
||||
"mcp_e2e_{}",
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs()
|
||||
));
|
||||
std::fs::create_dir_all(&temp_dir).unwrap();
|
||||
|
||||
let test_port = "3042"; // Use a distinct port
|
||||
@@ -47,10 +52,13 @@ async fn test_full_system_e2e_performance() {
|
||||
assert!(stub_exe.exists(), "Stub not found at {:?}", stub_exe);
|
||||
|
||||
// 1. Start Server
|
||||
let mut server = Command::new(&server_exe).arg("--daemon")
|
||||
.env("MCP_PORT", test_port).env("RUST_LOG", "debug")
|
||||
let mut server = Command::new(&server_exe)
|
||||
.arg("--daemon")
|
||||
.env("MCP_PORT", test_port)
|
||||
.env("RUST_LOG", "debug")
|
||||
.env("MCP_MEMORY_STORE_DIR", temp_dir.to_str().unwrap())
|
||||
.env("MCP_AUTH_TOKEN", test_auth_token).env("RUST_LOG", "debug")
|
||||
.env("MCP_AUTH_TOKEN", test_auth_token)
|
||||
.env("RUST_LOG", "debug")
|
||||
.stdout(Stdio::inherit())
|
||||
.stderr(Stdio::inherit())
|
||||
.spawn()
|
||||
@@ -63,12 +71,15 @@ async fn test_full_system_e2e_performance() {
|
||||
.unwrap();
|
||||
let mut started = false;
|
||||
for _ in 0..30 {
|
||||
if let Ok(resp) = client.get(format!("http://127.0.0.1:{}/health", test_port)).send().await {
|
||||
if resp.status().is_success() {
|
||||
if let Ok(resp) = client
|
||||
.get(format!("http://127.0.0.1:{}/health", test_port))
|
||||
.send()
|
||||
.await
|
||||
&& resp.status().is_success()
|
||||
{
|
||||
started = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(500)).await;
|
||||
}
|
||||
assert!(started, "Server failed to start in time");
|
||||
@@ -78,7 +89,8 @@ async fn test_full_system_e2e_performance() {
|
||||
.arg("--target")
|
||||
.arg(format!("http://127.0.0.1:{}", test_port))
|
||||
.env("MCP_MEMORY_STORE_DIR", temp_dir.to_str().unwrap())
|
||||
.env("MCP_AUTH_TOKEN", test_auth_token).env("RUST_LOG", "debug")
|
||||
.env("MCP_AUTH_TOKEN", test_auth_token)
|
||||
.env("RUST_LOG", "debug")
|
||||
.stdin(Stdio::piped())
|
||||
.stdout(Stdio::piped())
|
||||
.stderr(Stdio::inherit())
|
||||
@@ -121,7 +133,8 @@ async fn test_full_system_e2e_performance() {
|
||||
stub_stdin.flush().unwrap();
|
||||
}
|
||||
|
||||
let mut resp = read_message(&mut stub_stdout).expect("Failed to read rapid response from stub");
|
||||
let mut resp =
|
||||
read_message(&mut stub_stdout).expect("Failed to read rapid response from stub");
|
||||
while resp.get("id").is_none() || resp["id"].is_null() {
|
||||
resp = read_message(&mut stub_stdout).expect("Failed to read rapid response from stub");
|
||||
}
|
||||
@@ -149,9 +162,11 @@ async fn test_full_system_e2e_performance() {
|
||||
nvim_stdin.flush().unwrap();
|
||||
}
|
||||
|
||||
let mut resp = read_message(&mut nvim_stdout).expect("Failed to read rapid response from win-nvim");
|
||||
let mut resp =
|
||||
read_message(&mut nvim_stdout).expect("Failed to read rapid response from win-nvim");
|
||||
while resp.get("id").is_none() || resp["id"].is_null() {
|
||||
resp = read_message(&mut nvim_stdout).expect("Failed to read rapid response from win-nvim");
|
||||
resp = read_message(&mut nvim_stdout)
|
||||
.expect("Failed to read rapid response from win-nvim");
|
||||
}
|
||||
assert_eq!(resp["id"], i);
|
||||
}
|
||||
|
||||
@@ -10,7 +10,11 @@ fn get_stub_exe() -> std::path::PathBuf {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_stub_connection_refused() {
|
||||
let _ = std::process::Command::new("cargo").arg("build").arg("--bin").arg("mcp-memory-stub").status();
|
||||
let _ = std::process::Command::new("cargo")
|
||||
.arg("build")
|
||||
.arg("--bin")
|
||||
.arg("mcp-memory-stub")
|
||||
.status();
|
||||
|
||||
let target = "http://127.0.0.1:49999";
|
||||
|
||||
@@ -25,12 +29,20 @@ async fn test_stub_connection_refused() {
|
||||
let res = tokio::time::timeout(Duration::from_secs(5), child.wait()).await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
assert!(res.is_ok(), "Stub hung on connection refused! Took {:?}", elapsed);
|
||||
assert!(
|
||||
res.is_ok(),
|
||||
"Stub hung on connection refused! Took {:?}",
|
||||
elapsed
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_stub_handles_eof_cleanly() {
|
||||
let _ = std::process::Command::new("cargo").arg("build").arg("--bin").arg("mcp-memory-stub").status();
|
||||
let _ = std::process::Command::new("cargo")
|
||||
.arg("build")
|
||||
.arg("--bin")
|
||||
.arg("mcp-memory-stub")
|
||||
.status();
|
||||
|
||||
let target = "http://127.0.0.1:49998";
|
||||
|
||||
@@ -58,7 +70,11 @@ async fn test_stub_handles_eof_cleanly() {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_stub_sse_fallback_failure() {
|
||||
let _ = std::process::Command::new("cargo").arg("build").arg("--bin").arg("mcp-memory-stub").status();
|
||||
let _ = std::process::Command::new("cargo")
|
||||
.arg("build")
|
||||
.arg("--bin")
|
||||
.arg("mcp-memory-stub")
|
||||
.status();
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let local_addr = listener.local_addr().unwrap();
|
||||
@@ -84,5 +100,9 @@ async fn test_stub_sse_fallback_failure() {
|
||||
let res = tokio::time::timeout(Duration::from_secs(5), child.wait()).await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
assert!(res.is_ok(), "Stub hung on fallback failure! Took {:?}", elapsed);
|
||||
assert!(
|
||||
res.is_ok(),
|
||||
"Stub hung on fallback failure! Took {:?}",
|
||||
elapsed
|
||||
);
|
||||
}
|
||||
+8
-3
@@ -2,19 +2,24 @@ use std::process::Command;
|
||||
|
||||
fn main() {
|
||||
let git_hash = Command::new("git")
|
||||
.args(&["rev-parse", "--short", "HEAD"])
|
||||
.args(["rev-parse", "--short", "HEAD"])
|
||||
.output()
|
||||
.ok()
|
||||
.and_then(|out| String::from_utf8(out.stdout).ok())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
|
||||
let git_date = Command::new("git")
|
||||
.args(&["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"])
|
||||
.args(["log", "-1", "--format=%cd", "--date=format:%Y.%m.%d"])
|
||||
.output()
|
||||
.ok()
|
||||
.and_then(|out| String::from_utf8(out.stdout).ok())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
|
||||
let version = format!("{} ({} {})", env!("CARGO_PKG_VERSION"), git_date.trim(), git_hash.trim());
|
||||
let version = format!(
|
||||
"{} ({} {})",
|
||||
env!("CARGO_PKG_VERSION"),
|
||||
git_date.trim(),
|
||||
git_hash.trim()
|
||||
);
|
||||
println!("cargo:rustc-env=APP_VERSION={}", version);
|
||||
}
|
||||
@@ -112,7 +112,9 @@ fn test_mcp_initialization_and_tools_list() {
|
||||
assert_eq!(tools_resp["jsonrpc"], "2.0");
|
||||
assert_eq!(tools_resp["id"], 2);
|
||||
|
||||
let tools = tools_resp["result"]["tools"].as_array().expect("result.tools must be an array");
|
||||
let tools = tools_resp["result"]["tools"]
|
||||
.as_array()
|
||||
.expect("result.tools must be an array");
|
||||
assert!(!tools.is_empty(), "Server must expose at least one tool");
|
||||
|
||||
let has_get_active_buffer = tools.iter().any(|t| t["name"] == "nvim_get_active_buffer");
|
||||
@@ -137,7 +139,10 @@ fn test_mcp_initialization_and_tools_list() {
|
||||
|
||||
assert_eq!(call_resp["jsonrpc"], "2.0");
|
||||
assert_eq!(call_resp["id"], 3);
|
||||
assert!(call_resp.get("error").is_some(), "Expected an error response since Neovim shouldn't be running");
|
||||
assert!(
|
||||
call_resp.get("error").is_some(),
|
||||
"Expected an error response since Neovim shouldn't be running"
|
||||
);
|
||||
assert_eq!(call_resp["error"]["code"], -32603); // Internal Error
|
||||
|
||||
child.kill().expect("Failed to kill child");
|
||||
|
||||
Reference in new issue
Block a user