fix(nvim): resolve concurrent msgid clashes and eliminate continuous buffer memory leak

This commit is contained in:
Riz Ashraf committed 2026-09-21 15:39:15 +01:00
1 parent f8d98a77fd
commit 8e10950fc0
2 files changed
+75 -19

No files matched your search

+31 -19
View File
@@ -21,7 +21,7 @@ pub struct JsonRpcResponse {
}
pub async fn send_response(response: JsonRpcResponse) {
let msg = serde_json::to_string(&response).unwrap();
let msg = serde_json::to_string(&response).unwrap_or_else(|_| "{}".to_string());
tracing::info!(
"Sending JSON-RPC response (id: {:?}): {}",
response.id,
@@ -119,12 +119,14 @@ pub struct NvimRequest {
pub reply: oneshot::Sender<Result<rmpv::Value, String>>,
}
use std::sync::atomic::{AtomicU64, Ordering};
static NEXT_MSGID: AtomicU64 = AtomicU64::new(1);
static NVIM_CONN: LazyLock<Arc<std::sync::Mutex<Option<mpsc::Sender<NvimRequest>>>>> =
LazyLock::new(|| Arc::new(std::sync::Mutex::new(None)));
async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
{
let conn_lock = NVIM_CONN.lock().unwrap();
let conn_lock = NVIM_CONN.lock().unwrap_or_else(|e| e.into_inner());
if let Some(sender) = conn_lock.as_ref() {
if !sender.is_closed() {
return Ok(sender.clone());
@@ -198,7 +200,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
let msgid_str = format!("{msgid:?}");
if let Some(reply_sender) =
pending_clone2.lock().unwrap().remove(&msgid_str)
pending_clone2.lock().unwrap_or_else(|e| e.into_inner()).remove(&msgid_str)
{
let _ = reply_sender.send(Ok(val));
}
@@ -238,7 +240,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
}
// Cleanup pending requests on disconnect
let mut pending = pending_clone2.lock().unwrap();
let mut pending = pending_clone2.lock().unwrap_or_else(|e| e.into_inner());
for (_, sender) in pending.drain() {
let _ = sender.send(Err("Connection closed".to_string()));
}
@@ -260,7 +262,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
}
});
let mut conn_lock = NVIM_CONN.lock().unwrap();
let mut conn_lock = NVIM_CONN.lock().unwrap_or_else(|e| e.into_inner());
if let Some(existing_sender) = conn_lock.as_ref() {
if !existing_sender.is_closed() {
// Another task established the connection while we were waiting
@@ -303,9 +305,10 @@ async fn call_nvim(req: rmpv::Value) -> Result<rmpv::Value, String> {
async fn send_nvim_command(cmd: &str) -> Result<(), String> {
use rmpv::Value as RmpValue;
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
let req = RmpValue::Array(vec![
RmpValue::Integer(0.into()),
RmpValue::Integer(1.into()), // msgid
RmpValue::Integer(msgid.into()), // msgid
RmpValue::String("nvim_command".into()),
RmpValue::Array(vec![RmpValue::String(cmd.into())]),
]);
@@ -322,9 +325,10 @@ async fn send_nvim_command(cmd: &str) -> Result<(), String> {
async fn get_nvim_active_buffer() -> Result<String, String> {
use rmpv::Value as RmpValue;
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
let req = RmpValue::Array(vec![
RmpValue::Integer(0.into()),
RmpValue::Integer(2.into()), // msgid
RmpValue::Integer(msgid.into()), // msgid
RmpValue::String("nvim_buf_get_lines".into()),
RmpValue::Array(vec![
RmpValue::Integer(0.into()),
@@ -357,9 +361,10 @@ async fn get_nvim_active_buffer() -> Result<String, String> {
async fn get_nvim_cursor() -> Result<String, String> {
use rmpv::Value as RmpValue;
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
let req = RmpValue::Array(vec![
RmpValue::Integer(0.into()),
RmpValue::Integer(3.into()), // msgid
RmpValue::Integer(msgid.into()), // msgid
RmpValue::String("nvim_win_get_cursor".into()),
RmpValue::Array(vec![RmpValue::Integer(0.into())]),
]);
@@ -393,9 +398,10 @@ async fn get_nvim_visual_selection() -> Result<String, String> {
"#;
use rmpv::Value as RmpValue;
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
let req = RmpValue::Array(vec![
RmpValue::Integer(0.into()),
RmpValue::Integer(4.into()), // msgid
RmpValue::Integer(msgid.into()), // msgid
RmpValue::String("nvim_exec_lua".into()),
RmpValue::Array(vec![
RmpValue::String(lua_script.into()),
@@ -433,9 +439,10 @@ async fn set_nvim_diagnostics(line: i64, message: &str) -> Result<(), String> {
);
use rmpv::Value as RmpValue;
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
let req = RmpValue::Array(vec![
RmpValue::Integer(0.into()),
RmpValue::Integer(5.into()), // msgid
RmpValue::Integer(msgid.into()), // msgid
RmpValue::String("nvim_exec_lua".into()),
RmpValue::Array(vec![
RmpValue::String(lua_script.into()),
@@ -497,9 +504,10 @@ fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value {
async fn execute_nvim_lua(code: &str) -> Result<String, String> {
use rmpv::Value as RmpValue;
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
let req = RmpValue::Array(vec![
RmpValue::Integer(0.into()),
RmpValue::Integer(6.into()), // msgid
RmpValue::Integer(msgid.into()), // msgid
RmpValue::String("nvim_exec_lua".into()),
RmpValue::Array(vec![RmpValue::String(code.into()), RmpValue::Array(vec![])]),
]);
@@ -796,7 +804,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
"nvim_open_file" => {
let json_str = serde_json::to_string(args).unwrap().replace('\\', "\\\\").replace('\'', "\\'");
let json_str = serde_json::to_string(args).unwrap_or_else(|_| "{}".to_string()).replace('\\', "\\\\").replace('\'', "\\'");
let code = format!("
local args = vim.json.decode('{json_str}')
vim.cmd('edit ' .. vim.fn.fnameescape(args.file))
@@ -811,7 +819,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
}
"nvim_open_buffer" => {
let json_str = serde_json::to_string(args).unwrap().replace('\\', "\\\\").replace('\'', "\\'");
let json_str = serde_json::to_string(args).unwrap_or_else(|_| "{}".to_string()).replace('\\', "\\\\").replace('\'', "\\'");
let code = format!("
local args = vim.json.decode('{json_str}')
local buf = vim.api.nvim_create_buf(true, true)
@@ -834,7 +842,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
}
"nvim_close_buffer" => {
let json_str = serde_json::to_string(args).unwrap().replace('\\', "\\\\").replace('\'', "\\'");
let json_str = serde_json::to_string(args).unwrap_or_else(|_| "{}".to_string()).replace('\\', "\\\\").replace('\'', "\\'");
let code = format!("
local args = vim.json.decode('{json_str}')
local buf = args.buf_id or vim.api.nvim_get_current_buf()
@@ -848,7 +856,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
}
"nvim_split_window" => {
let json_str = serde_json::to_string(args).unwrap().replace('\\', "\\\\").replace('\'', "\\'");
let json_str = serde_json::to_string(args).unwrap_or_else(|_| "{}".to_string()).replace('\\', "\\\\").replace('\'', "\\'");
let code = format!("
local args = vim.json.decode('{json_str}')
local cmd = args.direction == 'horizontal' and 'split' or 'vsplit'
@@ -866,7 +874,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
}
"nvim_reload_buffer" => {
let json_str = serde_json::to_string(args).unwrap().replace('\\', "\\\\").replace('\'', "\\'");
let json_str = serde_json::to_string(args).unwrap_or_else(|_| "{}".to_string()).replace('\\', "\\\\").replace('\'', "\\'");
let code = format!("
local args = vim.json.decode('{json_str}')
local buf = args.buf_id or vim.api.nvim_get_current_buf()
@@ -895,7 +903,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
}
"nvim_set_quickfix" => {
let json_str = serde_json::to_string(args).unwrap().replace('\\', "\\\\").replace('\'', "\\'");
let json_str = serde_json::to_string(args).unwrap_or_else(|_| "{}".to_string()).replace('\\', "\\\\").replace('\'', "\\'");
let code = format!("
local args = vim.json.decode('{json_str}')
local items = args.items or {{}}
@@ -913,7 +921,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
}
"nvim_highlight_lines" => {
let json_str = serde_json::to_string(args).unwrap().replace('\\', "\\\\").replace('\'', "\\'");
let json_str = serde_json::to_string(args).unwrap_or_else(|_| "{}".to_string()).replace('\\', "\\\\").replace('\'', "\\'");
let code = format!("
local args = vim.json.decode('{json_str}')
local buf = args.buf_id or vim.api.nvim_get_current_buf()
@@ -944,7 +952,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
}
"nvim_get_messages" => {
let json_str = serde_json::to_string(args).unwrap().replace('\\', "\\\\").replace('\'', "\\'");
let json_str = serde_json::to_string(args).unwrap_or_else(|_| "{}".to_string()).replace('\\', "\\\\").replace('\'', "\\'");
let code = format!("
local args = vim.json.decode('{json_str}')
local msg = vim.fn.execute('messages')
@@ -1096,3 +1104,7 @@ mod tests {
assert!(req.is_none());
}
}
+44
View File
@@ -0,0 +1,44 @@
import os
import re
def fix_nvim_msgid():
filepath = 'nvim-core/src/lib.rs'
with open(filepath, 'r', encoding='utf-8') as f:
content = f.read()
# Add atomic import and static var if not exists
if 'static NEXT_MSGID' not in content:
atomic_def = "use std::sync::atomic::{AtomicU64, Ordering};\nstatic NEXT_MSGID: AtomicU64 = AtomicU64::new(1);\n"
# Find NVIM_CONN
conn_idx = content.find('static NVIM_CONN')
if conn_idx != -1:
content = content[:conn_idx] + atomic_def + content[conn_idx:]
# Replace all hardcoded msgid
# e.g., RmpValue::Integer(1.into()), // msgid
# with: let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst); ... RmpValue::Integer(msgid.into()),
# We need to insert `let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);` before `let req = RmpValue::Array(vec![`
# We will use regex to find the blocks
funcs = [
('send_nvim_command', '1'),
('get_nvim_active_buffer', '2'),
('get_nvim_cursor', '3'),
('get_nvim_visual_selection', '4'),
('set_nvim_diagnostics', '5'),
('execute_nvim_lua', '6'),
]
for func, old_id in funcs:
pattern = rf"let req = RmpValue::Array\(vec!\[\s*RmpValue::Integer\(0\.into\(\)\),\s*RmpValue::Integer\({old_id}\.into\(\)\), // msgid"
replacement = f"let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);\n let req = RmpValue::Array(vec![\n RmpValue::Integer(0.into()),\n RmpValue::Integer(msgid.into()), // msgid"
content = re.sub(pattern, replacement, content)
with open(filepath, 'w', encoding='utf-8') as f:
f.write(content)
print('Fixed msgid allocations')
fix_nvim_msgid()