refactor: apply 5-pass audit optimizations across mcp-memory codebase

This commit is contained in:
Riz Ashraf committed 2026-10-06 06:05:38 +01:00
1 parent 924b6d09fa
commit 5bd8b1587a
43 files changed
+1866 -1658

No files matched your search

+62 -16
View File
@@ -38,8 +38,12 @@ pub async fn send_response(response: JsonRpcResponse) {
payload.extend_from_slice(msg.as_bytes());
payload.push(b'\n');
let mut stdout = tokio::io::stdout();
let _ = stdout.write_all(&payload).await;
let _ = stdout.flush().await;
if let Err(e) = stdout.write_all(&payload).await {
tracing::error!("Failed to write response payload to stdout: {}", e);
}
if let Err(e) = stdout.flush().await {
tracing::error!("Failed to flush stdout: {}", e);
}
}
macro_rules! handle_lua_result {
@@ -177,6 +181,8 @@ static RPC_SEMAPHORE: LazyLock<Arc<tokio::sync::Semaphore>> =
LazyLock::new(|| Arc::new(tokio::sync::Semaphore::new(100)));
static NVIM_CONN: LazyLock<Arc<std::sync::Mutex<Option<mpsc::Sender<NvimRequest>>>>> =
LazyLock::new(|| Arc::new(std::sync::Mutex::new(None)));
static PENDING_REQUESTS: LazyLock<dashmap::DashMap<u64, oneshot::Sender<Result<rmpv::Value, String>>>> =
LazyLock::new(dashmap::DashMap::new);
#[derive(Debug, Default, Clone)]
pub struct NvimState {
@@ -293,12 +299,8 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
let (mut read_half, mut write_half) = tokio::io::split(stream);
let (tx, mut rx) = mpsc::channel::<NvimRequest>(32);
type PendingRequestsMap =
Arc<dashmap::DashMap<u64, oneshot::Sender<Result<rmpv::Value, String>>>>;
let pending_requests: PendingRequestsMap = Arc::new(dashmap::DashMap::new());
// Write task
let pending_clone = Arc::clone(&pending_requests);
tokio::spawn(async move {
while let Some(req) = rx.recv().await {
let mut buf = Vec::new();
@@ -307,11 +309,11 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
continue;
}
pending_clone.insert(req.msgid, req.reply);
PENDING_REQUESTS.insert(req.msgid, req.reply);
if write_half.write_all(&buf).await.is_err() {
tracing::error!("Failed to write to Neovim socket");
if let Some((_, sender)) = pending_clone.remove(&req.msgid) {
if let Some((_, sender)) = PENDING_REQUESTS.remove(&req.msgid) {
let _ = sender.send(Err("Connection closed during write".to_string()));
}
break;
@@ -320,7 +322,6 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
});
// Read task
let pending_clone2 = Arc::clone(&pending_requests);
tokio::spawn(async move {
use bytes::{Buf, BytesMut};
let mut resp_buf = BytesMut::with_capacity(65536);
@@ -338,7 +339,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
rmpv::Value::Integer(i) => i.as_u64().unwrap_or(0),
_ => 0,
};
if let Some((_, reply_sender)) = pending_clone2.remove(&msgid) {
if let Some((_, reply_sender)) = PENDING_REQUESTS.remove(&msgid) {
let _ = reply_sender.send(Ok(val));
}
} else if arr.len() >= 3
@@ -382,9 +383,9 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
}
// Cleanup pending requests on disconnect
let keys: Vec<_> = pending_clone2.iter().map(|kv| *kv.key()).collect();
let keys: Vec<_> = PENDING_REQUESTS.iter().map(|kv| *kv.key()).collect();
for k in keys {
if let Some((_, sender)) = pending_clone2.remove(&k) {
if let Some((_, sender)) = PENDING_REQUESTS.remove(&k) {
let _ = sender.send(Err("Connection closed".to_string()));
}
}
@@ -469,8 +470,14 @@ async fn call_nvim(req: rmpv::Value) -> Result<rmpv::Value, String> {
match tokio::time::timeout(tokio::time::Duration::from_secs(30), reply_rx).await {
Ok(Ok(res)) => res,
Ok(Err(_)) => Err("Response channel dropped".to_string()),
Err(_) => Err("Timeout waiting for Neovim response".to_string()),
Ok(Err(_)) => {
PENDING_REQUESTS.remove(&msgid);
Err("Response channel dropped".to_string())
}
Err(_) => {
PENDING_REQUESTS.remove(&msgid);
Err("Timeout waiting for Neovim response".to_string())
}
}
}
@@ -1311,12 +1318,19 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
let args_rmp = json_to_rmpv(args);
let code = "
local args = ...
local target_file = (args.file and args.file ~= '' and args.file ~= 'null') and args.file
or (args.file_path and args.file_path ~= '' and args.file_path ~= 'null') and args.file_path
or (args.path and args.path ~= '' and args.path ~= 'null') and args.path
or (args.target and args.target ~= '' and args.target ~= 'null') and args.target
if not target_file or target_file == vim.NIL then
return 'Error: No valid file path provided'
end
local curr_buf = vim.api.nvim_get_current_buf()
local is_empty_unnamed = vim.api.nvim_buf_get_name(curr_buf) == ''
and vim.api.nvim_buf_get_option(curr_buf, 'modified') == false
and vim.api.nvim_buf_line_count(curr_buf) <= 1
and (vim.api.nvim_buf_get_lines(curr_buf, 0, 1, false)[1] or '') == ''
vim.cmd('edit ' .. vim.fn.fnameescape(args.file))
vim.cmd('edit ' .. vim.fn.fnameescape(target_file))
local new_buf = vim.api.nvim_get_current_buf()
if is_empty_unnamed and curr_buf ~= new_buf and vim.api.nvim_buf_is_valid(curr_buf) then
pcall(vim.api.nvim_buf_delete, curr_buf, { force = true })
@@ -1324,7 +1338,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
if args.filetype and args.filetype ~= '' then
vim.bo.filetype = args.filetype
end
return 'Opened file ' .. args.file
return 'Opened file ' .. target_file
";
handle_lua_result!(id, execute_nvim_lua_with_args(code, vec![args_rmp]))
}
@@ -1841,6 +1855,38 @@ mod tests {
let ext_val = rmpv::Value::Ext(1, vec![10, 20]);
assert_eq!(rmpv_to_json(&ext_val), serde_json::json!("Ext(1, [10, 20])"));
}
#[test]
fn test_mock_nvim_get_api_info_response() {
let channel_id = rmpv::Value::Integer(42.into());
let api_metadata = vec![
(
rmpv::Value::String("version".into()),
rmpv::Value::Map(vec![
(rmpv::Value::String("major".into()), rmpv::Value::Integer(0.into())),
(rmpv::Value::String("minor".into()), rmpv::Value::Integer(10.into())),
(rmpv::Value::String("patch".into()), rmpv::Value::Integer(0.into())),
]),
),
(
rmpv::Value::String("functions".into()),
rmpv::Value::Array(vec![]),
),
];
let api_info_res = rmpv::Value::Array(vec![channel_id, rmpv::Value::Map(api_metadata)]);
if let rmpv::Value::Array(arr) = &api_info_res {
assert_eq!(arr.len(), 2);
let chan = arr[0].as_i64().unwrap();
assert_eq!(chan, 42);
let json_res = rmpv_to_json(&api_info_res);
assert_eq!(json_res[0], json!(42));
assert_eq!(json_res[1]["version"]["minor"], json!(10));
} else {
panic!("Expected array response for nvim_get_api_info");
}
}
}