fix(nvim): resolve concurrent msgid clashes and eliminate continuous buffer memory leak
This commit is contained in:
1 parent
f8d98a77fd
commit
8e10950fc0
2 files changed
+75
-19
No files matched your search
+31
-19
@@ -21,7 +21,7 @@ pub struct JsonRpcResponse {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn send_response(response: JsonRpcResponse) {
|
pub async fn send_response(response: JsonRpcResponse) {
|
||||||
let msg = serde_json::to_string(&response).unwrap();
|
let msg = serde_json::to_string(&response).unwrap_or_else(|_| "{}".to_string());
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"Sending JSON-RPC response (id: {:?}): {}",
|
"Sending JSON-RPC response (id: {:?}): {}",
|
||||||
response.id,
|
response.id,
|
||||||
@@ -119,12 +119,14 @@ pub struct NvimRequest {
|
|||||||
pub reply: oneshot::Sender<Result<rmpv::Value, String>>,
|
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>>>>> =
|
static NVIM_CONN: LazyLock<Arc<std::sync::Mutex<Option<mpsc::Sender<NvimRequest>>>>> =
|
||||||
LazyLock::new(|| Arc::new(std::sync::Mutex::new(None)));
|
LazyLock::new(|| Arc::new(std::sync::Mutex::new(None)));
|
||||||
|
|
||||||
async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
|
async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
|
||||||
{
|
{
|
||||||
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 let Some(sender) = conn_lock.as_ref() {
|
||||||
if !sender.is_closed() {
|
if !sender.is_closed() {
|
||||||
return Ok(sender.clone());
|
return Ok(sender.clone());
|
||||||
@@ -198,7 +200,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
|
|||||||
let msgid_str = format!("{msgid:?}");
|
let msgid_str = format!("{msgid:?}");
|
||||||
|
|
||||||
if let Some(reply_sender) =
|
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));
|
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
|
// 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() {
|
for (_, sender) in pending.drain() {
|
||||||
let _ = sender.send(Err("Connection closed".to_string()));
|
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 let Some(existing_sender) = conn_lock.as_ref() {
|
||||||
if !existing_sender.is_closed() {
|
if !existing_sender.is_closed() {
|
||||||
// Another task established the connection while we were waiting
|
// 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> {
|
async fn send_nvim_command(cmd: &str) -> Result<(), String> {
|
||||||
use rmpv::Value as RmpValue;
|
use rmpv::Value as RmpValue;
|
||||||
|
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
|
||||||
let req = RmpValue::Array(vec![
|
let req = RmpValue::Array(vec![
|
||||||
RmpValue::Integer(0.into()),
|
RmpValue::Integer(0.into()),
|
||||||
RmpValue::Integer(1.into()), // msgid
|
RmpValue::Integer(msgid.into()), // msgid
|
||||||
RmpValue::String("nvim_command".into()),
|
RmpValue::String("nvim_command".into()),
|
||||||
RmpValue::Array(vec![RmpValue::String(cmd.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> {
|
async fn get_nvim_active_buffer() -> Result<String, String> {
|
||||||
use rmpv::Value as RmpValue;
|
use rmpv::Value as RmpValue;
|
||||||
|
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
|
||||||
let req = RmpValue::Array(vec![
|
let req = RmpValue::Array(vec![
|
||||||
RmpValue::Integer(0.into()),
|
RmpValue::Integer(0.into()),
|
||||||
RmpValue::Integer(2.into()), // msgid
|
RmpValue::Integer(msgid.into()), // msgid
|
||||||
RmpValue::String("nvim_buf_get_lines".into()),
|
RmpValue::String("nvim_buf_get_lines".into()),
|
||||||
RmpValue::Array(vec![
|
RmpValue::Array(vec![
|
||||||
RmpValue::Integer(0.into()),
|
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> {
|
async fn get_nvim_cursor() -> Result<String, String> {
|
||||||
use rmpv::Value as RmpValue;
|
use rmpv::Value as RmpValue;
|
||||||
|
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
|
||||||
let req = RmpValue::Array(vec![
|
let req = RmpValue::Array(vec![
|
||||||
RmpValue::Integer(0.into()),
|
RmpValue::Integer(0.into()),
|
||||||
RmpValue::Integer(3.into()), // msgid
|
RmpValue::Integer(msgid.into()), // msgid
|
||||||
RmpValue::String("nvim_win_get_cursor".into()),
|
RmpValue::String("nvim_win_get_cursor".into()),
|
||||||
RmpValue::Array(vec![RmpValue::Integer(0.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;
|
use rmpv::Value as RmpValue;
|
||||||
|
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
|
||||||
let req = RmpValue::Array(vec![
|
let req = RmpValue::Array(vec![
|
||||||
RmpValue::Integer(0.into()),
|
RmpValue::Integer(0.into()),
|
||||||
RmpValue::Integer(4.into()), // msgid
|
RmpValue::Integer(msgid.into()), // msgid
|
||||||
RmpValue::String("nvim_exec_lua".into()),
|
RmpValue::String("nvim_exec_lua".into()),
|
||||||
RmpValue::Array(vec![
|
RmpValue::Array(vec![
|
||||||
RmpValue::String(lua_script.into()),
|
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;
|
use rmpv::Value as RmpValue;
|
||||||
|
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
|
||||||
let req = RmpValue::Array(vec![
|
let req = RmpValue::Array(vec![
|
||||||
RmpValue::Integer(0.into()),
|
RmpValue::Integer(0.into()),
|
||||||
RmpValue::Integer(5.into()), // msgid
|
RmpValue::Integer(msgid.into()), // msgid
|
||||||
RmpValue::String("nvim_exec_lua".into()),
|
RmpValue::String("nvim_exec_lua".into()),
|
||||||
RmpValue::Array(vec![
|
RmpValue::Array(vec![
|
||||||
RmpValue::String(lua_script.into()),
|
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> {
|
async fn execute_nvim_lua(code: &str) -> Result<String, String> {
|
||||||
use rmpv::Value as RmpValue;
|
use rmpv::Value as RmpValue;
|
||||||
|
let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst);
|
||||||
let req = RmpValue::Array(vec![
|
let req = RmpValue::Array(vec![
|
||||||
RmpValue::Integer(0.into()),
|
RmpValue::Integer(0.into()),
|
||||||
RmpValue::Integer(6.into()), // msgid
|
RmpValue::Integer(msgid.into()), // msgid
|
||||||
RmpValue::String("nvim_exec_lua".into()),
|
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![])]),
|
||||||
]);
|
]);
|
||||||
@@ -796,7 +804,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
"nvim_open_file" => {
|
"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!("
|
let code = format!("
|
||||||
local args = vim.json.decode('{json_str}')
|
local args = vim.json.decode('{json_str}')
|
||||||
vim.cmd('edit ' .. vim.fn.fnameescape(args.file))
|
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" => {
|
"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!("
|
let code = format!("
|
||||||
local args = vim.json.decode('{json_str}')
|
local args = vim.json.decode('{json_str}')
|
||||||
local buf = vim.api.nvim_create_buf(true, true)
|
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" => {
|
"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!("
|
let code = format!("
|
||||||
local args = vim.json.decode('{json_str}')
|
local args = vim.json.decode('{json_str}')
|
||||||
local buf = args.buf_id or vim.api.nvim_get_current_buf()
|
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" => {
|
"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!("
|
let code = format!("
|
||||||
local args = vim.json.decode('{json_str}')
|
local args = vim.json.decode('{json_str}')
|
||||||
local cmd = args.direction == 'horizontal' and 'split' or 'vsplit'
|
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" => {
|
"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!("
|
let code = format!("
|
||||||
local args = vim.json.decode('{json_str}')
|
local args = vim.json.decode('{json_str}')
|
||||||
local buf = args.buf_id or vim.api.nvim_get_current_buf()
|
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" => {
|
"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!("
|
let code = format!("
|
||||||
local args = vim.json.decode('{json_str}')
|
local args = vim.json.decode('{json_str}')
|
||||||
local items = args.items or {{}}
|
local items = args.items or {{}}
|
||||||
@@ -913,7 +921,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
"nvim_highlight_lines" => {
|
"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!("
|
let code = format!("
|
||||||
local args = vim.json.decode('{json_str}')
|
local args = vim.json.decode('{json_str}')
|
||||||
local buf = args.buf_id or vim.api.nvim_get_current_buf()
|
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" => {
|
"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!("
|
let code = format!("
|
||||||
local args = vim.json.decode('{json_str}')
|
local args = vim.json.decode('{json_str}')
|
||||||
local msg = vim.fn.execute('messages')
|
local msg = vim.fn.execute('messages')
|
||||||
@@ -1096,3 +1104,7 @@ mod tests {
|
|||||||
assert!(req.is_none());
|
assert!(req.is_none());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@@ -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()
|
||||||
Reference in new issue
Block a user