From 05e84e98bd79b197032f6a049c82a757bc0b4c88 Mon Sep 17 00:00:00 2001 From: Riz Ashraf Date: Wed, 9 Sep 2026 04:32:07 +0100 Subject: [PATCH] feat: Add resilient MPSC queues to proxy stubs to fix thread leaks and message loss on reconnect --- server/src/proxy.rs | 162 +++++++++++++++++----------------- stub/src/main.rs | 205 +++++++++++++++++++------------------------- 2 files changed, 169 insertions(+), 198 deletions(-) diff --git a/server/src/proxy.rs b/server/src/proxy.rs index 4ab5e38..393df47 100644 --- a/server/src/proxy.rs +++ b/server/src/proxy.rs @@ -3,105 +3,101 @@ use tokio::io::AsyncBufReadExt; use futures_util::StreamExt; use std::sync::Arc; use tokio::sync::RwLock; +use tokio::sync::mpsc; pub fn run_proxy(target_url: &str) -> Result> { let rt = tokio::runtime::Runtime::new().unwrap(); rt.block_on(async { - let client = reqwest::Client::builder().build().unwrap(); - let sse_url = format!("{}/sse", target_url); - - eprintln!("[PROXY] Connecting to Streamable HTTP: {}", sse_url); - let resp = match client.get(&sse_url) - .header("Accept", "text/event-stream") - .send().await { - Ok(r) => r, - Err(e) => { - eprintln!("[PROXY] Connection failed: {}", e); - return Ok(true); - } - }; + let (msg_tx, mut msg_rx) = mpsc::channel::(100); + let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1); - if resp.status() == reqwest::StatusCode::GONE { - eprintln!("[PROXY] Streamable HTTP endpoint is gone."); - return Ok(true); - } - - let post_url = Arc::new(RwLock::new(format!("{}/messages", target_url))); // Will be updated by endpoint event - let post_url_clone = Arc::clone(&post_url); - - let (tx, mut rx) = tokio::sync::mpsc::channel(1); - - let tx_clone = tx.clone(); - let client_clone = client.clone(); tokio::task::spawn_blocking(move || { let stdin = std::io::stdin(); let mut handle = stdin.lock(); let mut buffer = String::new(); while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) { if bytes == 0 { break; } - let body = buffer.clone(); + let _ = msg_tx.blocking_send(buffer.clone()); buffer.clear(); - let client = client_clone.clone(); - let url_arc = Arc::clone(&post_url_clone); - - tokio::spawn(async move { - let mut url = "".to_string(); - for _ in 0..50 { - let u = url_arc.read().await.clone(); - if u.contains("sessionId") { - url = u; - break; - } - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - } - if url.is_empty() { - url = url_arc.read().await.clone(); - } - let _ = client.post(&url) - .header("Accept", "application/json, text/event-stream") - .header("Content-Type", "application/json") - .body(body) - .send().await; - }); } - let _ = tx_clone.blocking_send(false); // Stdin EOF + let _ = shutdown_tx.blocking_send(()); }); let target_url = target_url.to_string(); - let tx_clone2 = tx.clone(); - tokio::spawn(async move { - let stream = resp.bytes_stream().map(|res| res.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e))); - let mut reader = tokio::io::BufReader::new(StreamReader::new(stream)); - let mut line = String::new(); - let mut is_message = false; - let mut is_endpoint = false; - - while let Ok(bytes) = reader.read_line(&mut line).await { - if bytes == 0 { break; } - let trimmed = line.trim(); - if trimmed.starts_with("event: message") { - is_message = true; - is_endpoint = false; - } else if trimmed.starts_with("event: endpoint") { - is_endpoint = true; - is_message = false; - } else if trimmed.starts_with("data: ") { - if is_message { - println!("{}", &trimmed[6..]); - is_message = false; - } else if is_endpoint { - let ep = &trimmed[6..]; - let mut p = post_url.write().await; - *p = format!("{}{}", target_url, ep); - is_endpoint = false; - } - } - line.clear(); - } - let _ = tx_clone2.send(true).await; // Stream dropped (Leader dead) - }); + let post_url = Arc::new(RwLock::new(String::new())); + let post_url_clone = Arc::clone(&post_url); - let dropped = rx.recv().await.unwrap_or(true); - Ok(dropped) + let client = reqwest::Client::builder().build().unwrap(); + + tokio::spawn(async move { + while let Some(msg) = msg_rx.recv().await { + loop { + let url = post_url_clone.read().await.clone(); + if !url.is_empty() { + let res = client.post(&url) + .header("Accept", "application/json, text/event-stream") + .header("Content-Type", "application/json") + .body(msg.clone()) + .send().await; + + if res.is_ok() { + break; + } + } + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + } + } + }); + + loop { + if shutdown_rx.try_recv().is_ok() { + return Ok(false); + } + + let sse_url = format!("{}/sse", target_url); + let client = reqwest::Client::builder().build().unwrap(); + + match client.get(&sse_url).header("Accept", "text/event-stream").send().await { + Ok(resp) => { + if resp.status() == reqwest::StatusCode::GONE { + return Ok(true); + } + + let stream = resp.bytes_stream().map(|res| res.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e))); + let mut reader = tokio::io::BufReader::new(StreamReader::new(stream)); + let mut line = String::new(); + let mut is_message = false; + let mut is_endpoint = false; + + while let Ok(bytes) = reader.read_line(&mut line).await { + if bytes == 0 { break; } + let trimmed = line.trim(); + if trimmed.starts_with("event: message") { + is_message = true; + is_endpoint = false; + } else if trimmed.starts_with("event: endpoint") { + is_endpoint = true; + is_message = false; + } else if trimmed.starts_with("data: ") { + if is_message { + println!("{}", &trimmed[6..]); + is_message = false; + } else if is_endpoint { + let ep = &trimmed[6..]; + let mut p = post_url.write().await; + *p = format!("{}{}", target_url, ep); + is_endpoint = false; + } + } + line.clear(); + } + *post_url.write().await = String::new(); + return Ok(true); + } + Err(_) => { + return Ok(true); + } + } + } }) } diff --git a/stub/src/main.rs b/stub/src/main.rs index 8ca054e..acd6f34 100644 --- a/stub/src/main.rs +++ b/stub/src/main.rs @@ -4,6 +4,7 @@ use std::sync::Arc; use tokio::io::AsyncBufReadExt; use tokio::sync::RwLock; use tokio_util::io::StreamReader; +use tokio::sync::mpsc; #[derive(Parser)] #[command(name = "mcp-memory-stub")] @@ -14,142 +15,116 @@ struct Cli { wake_cmd: Option, } -fn run_proxy(target_url: &str) -> Result> { +fn main() -> Result<(), Box> { + let cli = Cli::parse(); let rt = tokio::runtime::Runtime::new()?; rt.block_on(async { - let client = reqwest::Client::builder().build()?; - let sse_url = format!("{}/sse", target_url); - - eprintln!("[PROXY] Connecting to Streamable HTTP: {}", sse_url); - let resp = match client.get(&sse_url) - .header("Accept", "text/event-stream") - .send().await { - Ok(r) => r, - Err(e) => { - eprintln!("[PROXY] Connection failed: {}", e); - return Ok(true); - } - }; + let (msg_tx, mut msg_rx) = mpsc::channel::(100); + let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1); - if resp.status() == reqwest::StatusCode::GONE { - eprintln!("[PROXY] Streamable HTTP endpoint is gone."); - return Ok(true); - } - - let post_url = Arc::new(RwLock::new(format!("{}/messages", target_url))); // Will be updated by endpoint event - let post_url_clone = Arc::clone(&post_url); - - let (tx, mut rx) = tokio::sync::mpsc::channel(1); - - let tx_clone = tx.clone(); - let client_clone = client.clone(); tokio::task::spawn_blocking(move || { let stdin = std::io::stdin(); let mut handle = stdin.lock(); let mut buffer = String::new(); while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) { if bytes == 0 { break; } - let body = buffer.clone(); + let _ = msg_tx.blocking_send(buffer.clone()); buffer.clear(); - eprintln!("[PROXY] Stdin read: {}", body.trim()); - let client = client_clone.clone(); - let url_arc = Arc::clone(&post_url_clone); - - tokio::spawn(async move { - let mut url = "".to_string(); - for _ in 0..50 { - let u = url_arc.read().await.clone(); - if u.contains("sessionId") { - url = u; - break; - } - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - } - if url.is_empty() { - url = url_arc.read().await.clone(); - } - eprintln!("[PROXY] POSTing to {}", url); - let res = client.post(&url) - .header("Accept", "application/json, text/event-stream") - .header("Content-Type", "application/json") - .body(body) - .send().await; - eprintln!("[PROXY] POST response: {:?}", res.map(|r| r.status())); - }); } - eprintln!("[PROXY] Stdin closed"); - let _ = tx_clone.blocking_send(false); + let _ = shutdown_tx.blocking_send(()); }); - let target_url = target_url.to_string(); - let tx_clone2 = tx.clone(); + let target_url = cli.target; + let post_url = Arc::new(RwLock::new(String::new())); + let post_url_clone = Arc::clone(&post_url); + + let client = reqwest::Client::builder().build()?; + tokio::spawn(async move { - let stream = resp.bytes_stream().map(|res| res.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e))); - let mut reader = tokio::io::BufReader::new(StreamReader::new(stream)); - let mut line = String::new(); - let mut is_message = false; - let mut is_endpoint = false; - while let Ok(bytes) = reader.read_line(&mut line).await { - if bytes == 0 { break; } - let trimmed = line.trim(); - eprintln!("[PROXY] Stream read: {}", trimmed); - if trimmed.starts_with("event: message") { - is_message = true; - is_endpoint = false; - } else if trimmed.starts_with("event: endpoint") { - is_endpoint = true; - is_message = false; - } else if trimmed.starts_with("data: ") { - if is_message { - println!("{}", &trimmed[6..]); - is_message = false; - } else if is_endpoint { - let ep = &trimmed[6..]; - let mut p = post_url.write().await; - *p = format!("{}{}", target_url, ep); - eprintln!("[PROXY] Endpoint updated: {}", *p); - is_endpoint = false; + while let Some(msg) = msg_rx.recv().await { + let mut attempts = 0; + loop { + let url = post_url_clone.read().await.clone(); + if !url.is_empty() { + let res = client.post(&url) + .header("Accept", "application/json, text/event-stream") + .header("Content-Type", "application/json") + .body(msg.clone()) + .send().await; + + if res.is_ok() { + break; + } + } + tokio::time::sleep(tokio::time::Duration::from_millis(500)).await; + attempts += 1; + if attempts % 10 == 0 { + eprintln!("[PROXY] Waiting for server to accept messages..."); } } - line.clear(); } - eprintln!("[PROXY] Stream closed"); - let _ = tx_clone2.send(true).await; }); - let dropped = rx.recv().await.unwrap_or(true); - Ok(dropped) - }) -} - -fn main() { - let cli = Cli::parse(); - loop { - if let Err(_) = std::net::TcpStream::connect(cli.target.replace("http://", "").replace("https://", "")) { - eprintln!("[PROXY] Target {} offline, sleeping", cli.target); - - if let Some(ref cmd) = cli.wake_cmd { - eprintln!("[PROXY] Executing wake command..."); - let parts: Vec<&str> = cmd.split_whitespace().collect(); - if !parts.is_empty() { - let _ = std::process::Command::new(parts[0]) - .args(&parts[1..]) - .spawn(); - } - } - - std::thread::sleep(std::time::Duration::from_secs(1)); - } - match run_proxy(&cli.target) { - Ok(true) => { - std::thread::sleep(std::time::Duration::from_millis(50)); - } - Ok(false) => { + let wake_cmd = cli.wake_cmd; + loop { + if shutdown_rx.try_recv().is_ok() { break; } - Err(_) => { - std::thread::sleep(std::time::Duration::from_millis(1000)); + + let sse_url = format!("{}/sse", target_url); + let client = reqwest::Client::builder().build()?; + + match client.get(&sse_url).header("Accept", "text/event-stream").send().await { + Ok(resp) => { + if resp.status() == reqwest::StatusCode::GONE { + eprintln!("[PROXY] Target gone, exiting."); + break; + } + + let stream = resp.bytes_stream().map(|res| res.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e))); + let mut reader = tokio::io::BufReader::new(StreamReader::new(stream)); + let mut line = String::new(); + let mut is_message = false; + let mut is_endpoint = false; + + while let Ok(bytes) = reader.read_line(&mut line).await { + if bytes == 0 { break; } + let trimmed = line.trim(); + if trimmed.starts_with("event: message") { + is_message = true; + is_endpoint = false; + } else if trimmed.starts_with("event: endpoint") { + is_endpoint = true; + is_message = false; + } else if trimmed.starts_with("data: ") { + if is_message { + println!("{}", &trimmed[6..]); + is_message = false; + } else if is_endpoint { + let ep = &trimmed[6..]; + let mut p = post_url.write().await; + *p = format!("{}{}", target_url, ep); + is_endpoint = false; + } + } + line.clear(); + } + *post_url.write().await = String::new(); + tokio::time::sleep(tokio::time::Duration::from_millis(500)).await; + } + Err(_) => { + if let Some(ref cmd) = wake_cmd { + let parts: Vec<&str> = cmd.split_whitespace().collect(); + if !parts.is_empty() { + let _ = std::process::Command::new(parts[0]) + .args(&parts[1..]) + .spawn(); + } + } + tokio::time::sleep(tokio::time::Duration::from_secs(1)).await; + } } } - } + Ok(()) + }) }