diff --git a/nvim-core/src/lib.rs b/nvim-core/src/lib.rs index 70d5304..6f1b1b3 100644 --- a/nvim-core/src/lib.rs +++ b/nvim-core/src/lib.rs @@ -111,10 +111,128 @@ async fn get_socket_path() -> Result { } Err("Could not find Neovim socket".to_string()) } -#[cfg(windows)] -async fn call_nvim(req: rmpv::Value) -> Result { - use tokio::net::windows::named_pipe::ClientOptions; +use std::sync::LazyLock; +use std::sync::Arc; +use tokio::sync::{Mutex, mpsc, oneshot}; +use std::collections::HashMap; + +pub struct NvimRequest { + pub msgid_str: String, + pub req: rmpv::Value, + pub reply: oneshot::Sender>, +} + +static NVIM_CONN: LazyLock>>>> = LazyLock::new(|| Arc::new(Mutex::new(None))); + +async fn get_nvim_connection() -> Result, String> { + let mut conn_lock = NVIM_CONN.lock().await; + if let Some(sender) = conn_lock.as_ref() { + if !sender.is_closed() { + return Ok(sender.clone()); + } + } + + tracing::info!("Establishing new persistent connection to Neovim"); + let socket_path = get_socket_path().await?; + + #[cfg(windows)] + let stream = { + use tokio::net::windows::named_pipe::ClientOptions; + ClientOptions::new().open(&socket_path).map_err(|e| e.to_string())? + }; + + #[cfg(unix)] + let stream = { + use tokio::net::UnixStream; + UnixStream::connect(socket_path).await.map_err(|e| e.to_string())? + }; + + let (mut read_half, mut write_half) = tokio::io::split(stream); + let (tx, mut rx) = mpsc::channel::(32); + let pending_requests: Arc>>>> = Arc::new(Mutex::new(HashMap::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(); + if let Err(e) = rmpv::encode::write_value(&mut buf, &req.req) { + let _ = req.reply.send(Err(e.to_string())); + continue; + } + + pending_clone.lock().await.insert(req.msgid_str.clone(), req.reply); + + if write_half.write_all(&buf).await.is_err() { + tracing::error!("Failed to write to Neovim socket"); + break; + } + } + }); + + // Read task + let pending_clone2 = Arc::clone(&pending_requests); + tokio::spawn(async move { + let mut resp_buf = Vec::new(); + let mut chunk = vec![0u8; 8192]; + let mut offset = 0; + + loop { + let mut cursor = std::io::Cursor::new(&resp_buf[offset..]); + match rmpv::decode::read_value(&mut cursor) { + Ok(val) => { + offset += cursor.position() as usize; + + if let rmpv::Value::Array(ref arr) = val { + if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) { + let msgid = &arr[1]; + let msgid_str = format!("{:?}", msgid); + + if let Some(reply_sender) = pending_clone2.lock().await.remove(&msgid_str) { + let _ = reply_sender.send(Ok(val)); + } + } + } + // Trim buffer if it gets too large + if offset > 1024 * 1024 { + resp_buf.drain(..offset); + offset = 0; + } + continue; + } + Err(_) => { + if offset > 0 { + resp_buf.drain(..offset); + offset = 0; + } + + let read_future = read_half.read(&mut chunk); + match tokio::time::timeout(tokio::time::Duration::from_secs(60), read_future).await { + Ok(Ok(n)) if n > 0 => { + resp_buf.extend_from_slice(&chunk[..n]); + } + _ => { + tracing::error!("Neovim socket read loop closed or timeout"); + break; + } + } + } + } + } + + // Cleanup pending requests on disconnect + let mut pending = pending_clone2.lock().await; + for (_, sender) in pending.drain() { + let _ = sender.send(Err("Connection closed".to_string())); + } + }); + + *conn_lock = Some(tx.clone()); + Ok(tx) +} + +async fn call_nvim(req: rmpv::Value) -> Result { let msgid = if let rmpv::Value::Array(ref arr) = req { if arr.len() > 1 { arr[1].clone() @@ -124,120 +242,24 @@ async fn call_nvim(req: rmpv::Value) -> Result { } 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 buf = Vec::new(); - rmpv::encode::write_value(&mut buf, &req).map_err(|e| e.to_string())?; - tracing::info!("Sending RPC request to neovim (msgid: {})", msgid); - client.write_all(&buf).await.map_err(|e| e.to_string())?; - - let mut resp_buf = Vec::new(); - let mut chunk = vec![0u8; 8192]; - let mut offset = 0; - - loop { - let mut cursor = std::io::Cursor::new(&resp_buf[offset..]); - match rmpv::decode::read_value(&mut cursor) { - Ok(val) => { - 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 - { - 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()); - } - resp_buf.extend_from_slice(&chunk[..n]); - } - Ok(Err(e)) => return Err(e.to_string()), - Err(_) => { - tracing::error!("Timeout waiting for Neovim response (msgid: {})", msgid); - return Err("Timeout waiting for Neovim response".into()); - } - } - } - } + + let msgid_str = format!("{:?}", msgid); + let tx = get_nvim_connection().await?; + let (reply_tx, reply_rx) = oneshot::channel(); + + tx.send(NvimRequest { + msgid_str, + req, + reply: reply_tx, + }).await.map_err(|_| "Failed to send request to Neovim connection manager")?; + + match tokio::time::timeout(tokio::time::Duration::from_secs(10), reply_rx).await { + Ok(Ok(res)) => res, + Ok(Err(_)) => Err("Response channel dropped".to_string()), + Err(_) => Err("Timeout waiting for Neovim response".to_string()), } } -#[cfg(unix)] -async fn call_nvim(req: rmpv::Value) -> Result { - 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 - }; - - 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 buf = Vec::new(); - rmpv::encode::write_value(&mut buf, &req).map_err(|e| e.to_string())?; - tracing::info!("Sending RPC request to neovim (msgid: {})", msgid); - stream.write_all(&buf).await.map_err(|e| e.to_string())?; - - let mut resp_buf = Vec::new(); - let mut chunk = vec![0u8; 8192]; - let mut offset = 0; - - loop { - let mut cursor = std::io::Cursor::new(&resp_buf[offset..]); - match rmpv::decode::read_value(&mut cursor) { - Ok(val) => { - 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 - { - 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()); - } - resp_buf.extend_from_slice(&chunk[..n]); - } - Ok(Err(e)) => return Err(e.to_string()), - Err(_) => { - tracing::error!("Timeout waiting for Neovim response (msgid: {})", msgid); - return Err("Timeout waiting for Neovim response".into()); - } - } - } - } - } -} async fn send_nvim_command(cmd: &str) -> Result<(), String> { use rmpv::Value as RmpValue; let req = RmpValue::Array(vec![ diff --git a/patch_nvim.py b/patch_nvim.py new file mode 100644 index 0000000..64a3aa8 --- /dev/null +++ b/patch_nvim.py @@ -0,0 +1,168 @@ +import re +import sys + +with open('C:\\Users\\reazul.ashraf\\workspace\\rust\\mcp-memory\\nvim-core\\src\\lib.rs', 'r', encoding='utf-8') as f: + content = f.read() + +start_pattern = r'#\[cfg\(windows\)\]\nasync fn call_nvim' +end_pattern = r'async fn send_nvim_command' + +start_idx = re.search(start_pattern, content).start() +end_idx = re.search(end_pattern, content).start() + +new_code = """use std::sync::LazyLock; +use std::sync::Arc; +use tokio::sync::{Mutex, mpsc, oneshot}; +use std::collections::HashMap; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; + +pub struct NvimRequest { + pub msgid_str: String, + pub req: rmpv::Value, + pub reply: oneshot::Sender>, +} + +static NVIM_CONN: LazyLock>>>> = LazyLock::new(|| Arc::new(Mutex::new(None))); + +async fn get_nvim_connection() -> Result, String> { + let mut conn_lock = NVIM_CONN.lock().await; + if let Some(sender) = conn_lock.as_ref() { + if !sender.is_closed() { + return Ok(sender.clone()); + } + } + + tracing::info!("Establishing new persistent connection to Neovim"); + let socket_path = get_socket_path().await?; + + #[cfg(windows)] + let stream = { + use tokio::net::windows::named_pipe::ClientOptions; + ClientOptions::new().open(&socket_path).map_err(|e| e.to_string())? + }; + + #[cfg(unix)] + let stream = { + use tokio::net::UnixStream; + UnixStream::connect(socket_path).await.map_err(|e| e.to_string())? + }; + + let (mut read_half, mut write_half) = tokio::io::split(stream); + let (tx, mut rx) = mpsc::channel::(32); + let pending_requests: Arc>>>> = Arc::new(Mutex::new(HashMap::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(); + if let Err(e) = rmpv::encode::write_value(&mut buf, &req.req) { + let _ = req.reply.send(Err(e.to_string())); + continue; + } + + pending_clone.lock().await.insert(req.msgid_str.clone(), req.reply); + + if write_half.write_all(&buf).await.is_err() { + tracing::error!("Failed to write to Neovim socket"); + break; + } + } + }); + + // Read task + let pending_clone2 = Arc::clone(&pending_requests); + tokio::spawn(async move { + let mut resp_buf = Vec::new(); + let mut chunk = vec![0u8; 8192]; + let mut offset = 0; + + loop { + let mut cursor = std::io::Cursor::new(&resp_buf[offset..]); + match rmpv::decode::read_value(&mut cursor) { + Ok(val) => { + offset += cursor.position() as usize; + + if let rmpv::Value::Array(ref arr) = val { + if arr.len() >= 4 && arr[0] == rmpv::Value::Integer(1.into()) { + let msgid = &arr[1]; + let msgid_str = format!("{:?}", msgid); + + if let Some(reply_sender) = pending_clone2.lock().await.remove(&msgid_str) { + let _ = reply_sender.send(Ok(val)); + } + } + } + // Trim buffer if it gets too large + if offset > 1024 * 1024 { + resp_buf.drain(..offset); + offset = 0; + } + continue; + } + Err(_) => { + if offset > 0 { + resp_buf.drain(..offset); + offset = 0; + } + + let read_future = read_half.read(&mut chunk); + match tokio::time::timeout(tokio::time::Duration::from_secs(60), read_future).await { + Ok(Ok(n)) if n > 0 => { + resp_buf.extend_from_slice(&chunk[..n]); + } + _ => { + tracing::error!("Neovim socket read loop closed or timeout"); + break; + } + } + } + } + } + + // Cleanup pending requests on disconnect + let mut pending = pending_clone2.lock().await; + for (_, sender) in pending.drain() { + let _ = sender.send(Err("Connection closed".to_string())); + } + }); + + *conn_lock = Some(tx.clone()); + Ok(tx) +} + +async fn call_nvim(req: rmpv::Value) -> Result { + 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 + }; + + let msgid_str = format!("{:?}", msgid); + let tx = get_nvim_connection().await?; + let (reply_tx, reply_rx) = oneshot::channel(); + + tx.send(NvimRequest { + msgid_str, + req, + reply: reply_tx, + }).await.map_err(|_| "Failed to send request to Neovim connection manager")?; + + match tokio::time::timeout(tokio::time::Duration::from_secs(10), reply_rx).await { + Ok(Ok(res)) => res, + Ok(Err(_)) => Err("Response channel dropped".to_string()), + Err(_) => Err("Timeout waiting for Neovim response".to_string()), + } +} + +""" + +new_content = content[:start_idx] + new_code + content[end_idx:] +with open('C:\\Users\\reazul.ashraf\\workspace\\rust\\mcp-memory\\nvim-core\\src\\lib.rs', 'w', encoding='utf-8') as f: + f.write(new_content) + +print("Patched!")