refactor: apply zero-unwrap policy and optimize locks in store.rs and handlers

This commit is contained in:
Riz Ashraf committed 2026-09-21 11:34:21 +01:00
1 parent 9f24e66d88
commit 8afbf97b11
33 files changed
+3127 -3207

No files matched your search

+64 -38
View File
@@ -20,8 +20,6 @@ pub struct JsonRpcResponse {
pub error: Option<Value>,
}
pub async fn send_response(response: JsonRpcResponse) {
let msg = serde_json::to_string(&response).unwrap();
tracing::info!(
@@ -111,11 +109,10 @@ async fn get_socket_path() -> Result<String, String> {
}
Err("Could not find Neovim socket".to_string())
}
use std::sync::LazyLock;
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot};
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::LazyLock;
use tokio::sync::{mpsc, oneshot};
pub struct NvimRequest {
pub msgid_str: String,
@@ -123,7 +120,8 @@ pub struct NvimRequest {
pub reply: oneshot::Sender<Result<rmpv::Value, String>>,
}
static NVIM_CONN: LazyLock<Arc<std::sync::Mutex<Option<mpsc::Sender<NvimRequest>>>>> = LazyLock::new(|| Arc::new(std::sync::Mutex::new(None)));
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> {
{
@@ -141,18 +139,23 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
#[cfg(windows)]
let stream = {
use tokio::net::windows::named_pipe::ClientOptions;
ClientOptions::new().open(&socket_path).map_err(|e| e.to_string())?
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())?
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::<NvimRequest>(32);
type PendingRequestsMap = Arc<std::sync::Mutex<HashMap<String, oneshot::Sender<Result<rmpv::Value, String>>>>>;
type PendingRequestsMap =
Arc<std::sync::Mutex<HashMap<String, oneshot::Sender<Result<rmpv::Value, String>>>>>;
let pending_requests: PendingRequestsMap = Arc::new(std::sync::Mutex::new(HashMap::new()));
// Write task
@@ -164,9 +167,12 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
let _ = req.reply.send(Err(e.to_string()));
continue;
}
pending_clone.lock().unwrap().insert(req.msgid_str.clone(), req.reply);
pending_clone
.lock()
.unwrap()
.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;
@@ -191,8 +197,10 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
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().unwrap().remove(&msgid_str) {
if let Some(reply_sender) =
pending_clone2.lock().unwrap().remove(&msgid_str)
{
let _ = reply_sender.send(Ok(val));
}
}
@@ -204,12 +212,16 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
}
continue;
}
Err(rmpv::decode::Error::InvalidMarkerRead(e)) if e.kind() == std::io::ErrorKind::UnexpectedEof => {
Err(rmpv::decode::Error::InvalidMarkerRead(e))
if e.kind() == std::io::ErrorKind::UnexpectedEof =>
{
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 {
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]);
}
@@ -225,7 +237,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
}
}
}
// Cleanup pending requests on disconnect
let mut pending = pending_clone2.lock().unwrap();
for (_, sender) in pending.drain() {
@@ -242,7 +254,10 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
if Arc::strong_count(&pending_clone3) <= 1 {
break; // Socket closed and other tasks finished, no need to keep cleaning up
}
pending_clone3.lock().unwrap().retain(|_, sender| !sender.is_closed());
pending_clone3
.lock()
.unwrap()
.retain(|_, sender| !sender.is_closed());
}
});
@@ -267,17 +282,19 @@ async fn call_nvim(req: rmpv::Value) -> Result<rmpv::Value, String> {
} 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")?;
})
.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()),
@@ -544,9 +561,13 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
Ok(m) => {
tracing::info!("Received message method: {}", m.method);
m
},
}
Err(e) => {
tracing::error!("Failed to parse JSON-RPC request from JSONL: {}. Payload: {}", e, raw_msg);
tracing::error!(
"Failed to parse JSON-RPC request from JSONL: {}. Payload: {}",
e,
raw_msg
);
continue;
}
};
@@ -559,13 +580,17 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
let id_clone = id.clone();
let start_time = std::time::Instant::now();
let method_clone = if msg.method == "tools/call" {
let tool_name = msg.params.as_ref().and_then(|p| p.get("name")).and_then(|n| n.as_str()).unwrap_or("unknown");
let tool_name = msg
.params
.as_ref()
.and_then(|p| p.get("name"))
.and_then(|n| n.as_str())
.unwrap_or("unknown");
format!("ToolCall[{}]", tool_name)
} else {
msg.method.clone()
};
match msg.method.as_str() {
"initialize" => {
@@ -774,7 +799,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
Err(e) => send_error(id, -32603, &e).await,
}
}
"nvim_open_file" => {
let json_str = serde_json::to_string(args).unwrap().replace("\\", "\\\\").replace("'", "\\'");
let code = format!("
@@ -899,7 +924,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
local buf = args.buf_id or vim.api.nvim_get_current_buf()
local group = args.group or 'IncSearch'
local ns = vim.api.nvim_create_namespace('antigravity_highlight')
if args.clear_only then
vim.api.nvim_buf_clear_namespace(buf, ns, 0, -1)
return 'Cleared highlights'
@@ -909,7 +934,7 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
for i = args.start_line - 1, args.end_line - 1 do
pcall(vim.api.nvim_buf_add_highlight, buf, ns, group, i, 0, -1)
end
local duration = args.duration_ms or 5000
if duration > 0 then
vim.defer_fn(function()
@@ -984,13 +1009,16 @@ pub async fn run_mcp_loop(app_name: &str, app_version: &str) {
}
}
let elapsed = start_time.elapsed();
tracing::info!("<<< [Nvim] {} (id: {}) completed in {:?}", method_clone, id_clone, elapsed);
tracing::info!(
"<<< [Nvim] {} (id: {}) completed in {:?}",
method_clone,
id_clone,
elapsed
);
});
}
}
fn init_logging(app_name: &str) -> tracing_appender::non_blocking::WorkerGuard {
let log_dir = dirs::home_dir()
.unwrap_or_default()
@@ -1036,11 +1064,10 @@ mod tests {
#[test]
fn test_rmpv_to_json_map() {
let mut map = vec![];
map.push((
let map = vec![(
rmpv::Value::String("key1".into()),
rmpv::Value::Integer(100.into()),
));
)];
let rmp_map = rmpv::Value::Map(map);
let json_map = rmpv_to_json(&rmp_map);
@@ -1074,4 +1101,3 @@ mod tests {
assert!(req.is_none());
}
}