perf: eliminate tokio::sync::Mutex across async boundaries in stub and nvim-core

This commit is contained in:
Riz Ashraf committed 2026-09-21 08:32:13 +01:00
1 parent 573c9586fd
commit 5715625220
5 files changed
+59 -20

No files matched your search

Generated
+48
View File
@@ -85,6 +85,18 @@ dependencies = [
"rustversion", "rustversion",
] ]
[[package]]
name = "async-channel"
version = "2.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "924ed96dd52d1b75e9c1a3e6275715fd320f5f9439fb5a4a11fa51f4221158d2"
dependencies = [
"concurrent-queue",
"event-listener-strategy",
"futures-core",
"pin-project-lite",
]
[[package]] [[package]]
name = "async-trait" name = "async-trait"
version = "0.1.92" version = "0.1.92"
@@ -356,6 +368,15 @@ version = "1.0.5"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570"
[[package]]
name = "concurrent-queue"
version = "2.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4ca0197aee26d1ae37445ee532fefce43251d24cc7c166799f4d46817f1d3973"
dependencies = [
"crossbeam-utils",
]
[[package]] [[package]]
name = "core-foundation-sys" name = "core-foundation-sys"
version = "0.8.7" version = "0.8.7"
@@ -604,6 +625,26 @@ dependencies = [
"windows-sys 0.61.2", "windows-sys 0.61.2",
] ]
[[package]]
name = "event-listener"
version = "5.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5a23add41df1562121a9393cb065eab5146a1242410f23a644851e90cfd669d2"
dependencies = [
"parking",
"pin-project-lite",
]
[[package]]
name = "event-listener-strategy"
version = "0.5.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8be9f3dfaaffdae2972880079a491a1a8bb7cbed0b8dd7a347f668b4150a3b93"
dependencies = [
"event-listener",
"pin-project-lite",
]
[[package]] [[package]]
name = "fastdivide" name = "fastdivide"
version = "0.4.2" version = "0.4.2"
@@ -1372,6 +1413,7 @@ dependencies = [
name = "mcp-memory-stub" name = "mcp-memory-stub"
version = "0.1.0" version = "0.1.0"
dependencies = [ dependencies = [
"async-channel",
"clap", "clap",
"dirs 7.0.0", "dirs 7.0.0",
"futures-util", "futures-util",
@@ -1591,6 +1633,12 @@ dependencies = [
"stable_deref_trait", "stable_deref_trait",
] ]
[[package]]
name = "parking"
version = "2.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f38d5652c16fde515bb1ecef450ab0f6a219d619a7274976324d5e377f7dceba"
[[package]] [[package]]
name = "parking_lot" name = "parking_lot"
version = "0.12.5" version = "0.12.5"
+6 -6
View File
@@ -150,8 +150,8 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
let (mut read_half, mut write_half) = tokio::io::split(stream); let (mut read_half, mut write_half) = tokio::io::split(stream);
let (tx, mut rx) = mpsc::channel::<NvimRequest>(32); let (tx, mut rx) = mpsc::channel::<NvimRequest>(32);
type PendingRequestsMap = Arc<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(Mutex::new(HashMap::new())); let pending_requests: PendingRequestsMap = Arc::new(std::sync::Mutex::new(HashMap::new()));
// Write task // Write task
let pending_clone = Arc::clone(&pending_requests); let pending_clone = Arc::clone(&pending_requests);
@@ -163,7 +163,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
continue; continue;
} }
pending_clone.lock().await.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() { if write_half.write_all(&buf).await.is_err() {
tracing::error!("Failed to write to Neovim socket"); tracing::error!("Failed to write to Neovim socket");
@@ -190,7 +190,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
let msgid = &arr[1]; let msgid = &arr[1];
let msgid_str = format!("{:?}", msgid); let msgid_str = format!("{:?}", msgid);
if let Some(reply_sender) = pending_clone2.lock().await.remove(&msgid_str) { if let Some(reply_sender) = pending_clone2.lock().unwrap().remove(&msgid_str) {
let _ = reply_sender.send(Ok(val)); let _ = reply_sender.send(Ok(val));
} }
} }
@@ -225,7 +225,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().await; let mut pending = pending_clone2.lock().unwrap();
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()));
} }
@@ -240,7 +240,7 @@ async fn get_nvim_connection() -> Result<mpsc::Sender<NvimRequest>, String> {
if Arc::strong_count(&pending_clone3) <= 1 { if Arc::strong_count(&pending_clone3) <= 1 {
break; // Socket closed and other tasks finished, no need to keep cleaning up break; // Socket closed and other tasks finished, no need to keep cleaning up
} }
pending_clone3.lock().await.retain(|_, sender| !sender.is_closed()); pending_clone3.lock().unwrap().retain(|_, sender| !sender.is_closed());
} }
}); });
+1
View File
@@ -18,6 +18,7 @@ dirs = "7.0.0"
serde_json = "1.0.151" serde_json = "1.0.151"
mcp-stdio = { version = "0.1.0", path = "../mcp-stdio" } mcp-stdio = { version = "0.1.0", path = "../mcp-stdio" }
regex = "1.13.1" regex = "1.13.1"
async-channel = "2.5.0"
+3 -13
View File
@@ -1,6 +1,5 @@
use clap::Parser; use clap::Parser;
use futures_util::{SinkExt, StreamExt}; use futures_util::{SinkExt, StreamExt};
use std::sync::Arc;
use tokio::sync::mpsc; use tokio::sync::mpsc;
#[derive(Parser)] #[derive(Parser)]
@@ -37,7 +36,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let cli = Cli::parse(); let cli = Cli::parse();
let rt = tokio::runtime::Runtime::new()?; let rt = tokio::runtime::Runtime::new()?;
rt.block_on(async { rt.block_on(async {
let (msg_tx, msg_rx) = mpsc::channel::<String>(100); let (msg_tx, msg_rx) = async_channel::bounded::<String>(100);
let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1); let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1);
tokio::spawn(async move { tokio::spawn(async move {
@@ -56,7 +55,6 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
}; };
let ws_url = target_url.replace("http://", "ws://").replace("https://", "wss://"); let ws_url = target_url.replace("http://", "ws://").replace("https://", "wss://");
let ws_url = format!("{}/ws?client=proxy", ws_url); let ws_url = format!("{}/ws?client=proxy", ws_url);
let msg_rx = Arc::new(tokio::sync::Mutex::new(msg_rx));
loop { loop {
if shutdown_rx.try_recv().is_ok() { if shutdown_rx.try_recv().is_ok() {
@@ -81,23 +79,15 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
tracing::info!("Successfully connected to target server"); tracing::info!("Successfully connected to target server");
let (mut write, mut read) = ws_stream.split(); let (mut write, mut read) = ws_stream.split();
let msg_rx_clone = Arc::clone(&msg_rx); let rx = msg_rx.clone();
let mut send_task = tokio::spawn(async move { let mut send_task = tokio::spawn(async move {
loop { while let Ok(msg) = rx.recv().await {
let mut rx = msg_rx_clone.lock().await;
match rx.recv().await {
Some(msg) => {
drop(rx);
let log_prefix = logger::extract_log_prefix(&msg, false); let log_prefix = logger::extract_log_prefix(&msg, false);
tracing::info!(">>> [Stub] Forwarding {} to server (length: {}): {}", log_prefix, msg.len(), if msg.len() > 1000 { format!("{}...", &msg[..1000]) } else { msg.clone() }); tracing::info!(">>> [Stub] Forwarding {} to server (length: {}): {}", log_prefix, msg.len(), if msg.len() > 1000 { format!("{}...", &msg[..1000]) } else { msg.clone() });
if write.send(tokio_tungstenite::tungstenite::Message::Text(msg)).await.is_err() { if write.send(tokio_tungstenite::tungstenite::Message::Text(msg)).await.is_err() {
tracing::error!("Failed to write to websocket"); tracing::error!("Failed to write to websocket");
break; break;
} }
},
None => break,
}
} }
}); });
+1 -1
View File
@@ -61,7 +61,7 @@ async fn test_full_system_e2e_performance() {
assert!(stub_exe.exists(), "Stub not found at {:?}", stub_exe); assert!(stub_exe.exists(), "Stub not found at {:?}", stub_exe);
// 1. Start Server // 1. Start Server
let mut server = ChildGuard(Command::new(&server_exe) let _server = ChildGuard(Command::new(&server_exe)
.arg("--daemon") .arg("--daemon")
.env("MCP_PORT", test_port) .env("MCP_PORT", test_port)
.env("RUST_LOG", "debug") .env("RUST_LOG", "debug")