refactor: apply 5-pass audit optimizations across mcp-memory codebase
This commit is contained in:
1 parent
924b6d09fa
commit
5bd8b1587a
43 files changed
+1866
-1658
No files matched your search
+62
-16
@@ -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");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in new issue
Block a user