feat: Add resilient MPSC queues to proxy stubs to fix thread leaks and message loss on reconnect

This commit is contained in:
Riz Ashraf committed 2026-09-09 04:32:07 +01:00
1 parent 51fc3692c8
commit 05e84e98bd
2 files changed
+164 -193

No files matched your search

+74 -78
View File
@@ -3,105 +3,101 @@ use tokio::io::AsyncBufReadExt;
use futures_util::StreamExt; use futures_util::StreamExt;
use std::sync::Arc; use std::sync::Arc;
use tokio::sync::RwLock; use tokio::sync::RwLock;
use tokio::sync::mpsc;
pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + Send + Sync>> { pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + Send + Sync>> {
let rt = tokio::runtime::Runtime::new().unwrap(); let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async { rt.block_on(async {
let client = reqwest::Client::builder().build().unwrap(); let (msg_tx, mut msg_rx) = mpsc::channel::<String>(100);
let sse_url = format!("{}/sse", target_url); let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1);
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);
}
};
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 || { tokio::task::spawn_blocking(move || {
let stdin = std::io::stdin(); let stdin = std::io::stdin();
let mut handle = stdin.lock(); let mut handle = stdin.lock();
let mut buffer = String::new(); let mut buffer = String::new();
while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) { while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) {
if bytes == 0 { break; } if bytes == 0 { break; }
let body = buffer.clone(); let _ = msg_tx.blocking_send(buffer.clone());
buffer.clear(); 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 target_url = target_url.to_string();
let tx_clone2 = tx.clone(); let post_url = Arc::new(RwLock::new(String::new()));
tokio::spawn(async move { let post_url_clone = Arc::clone(&post_url);
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 { let client = reqwest::Client::builder().build().unwrap();
if bytes == 0 { break; }
let trimmed = line.trim(); tokio::spawn(async move {
if trimmed.starts_with("event: message") { while let Some(msg) = msg_rx.recv().await {
is_message = true; loop {
is_endpoint = false; let url = post_url_clone.read().await.clone();
} else if trimmed.starts_with("event: endpoint") { if !url.is_empty() {
is_endpoint = true; let res = client.post(&url)
is_message = false; .header("Accept", "application/json, text/event-stream")
} else if trimmed.starts_with("data: ") { .header("Content-Type", "application/json")
if is_message { .body(msg.clone())
println!("{}", &trimmed[6..]); .send().await;
is_message = false;
} else if is_endpoint { if res.is_ok() {
let ep = &trimmed[6..]; break;
let mut p = post_url.write().await; }
*p = format!("{}{}", target_url, ep);
is_endpoint = false;
} }
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
} }
line.clear();
} }
let _ = tx_clone2.send(true).await; // Stream dropped (Leader dead)
}); });
let dropped = rx.recv().await.unwrap_or(true); loop {
Ok(dropped) 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);
}
}
}
}) })
} }
+90 -115
View File
@@ -4,6 +4,7 @@ use std::sync::Arc;
use tokio::io::AsyncBufReadExt; use tokio::io::AsyncBufReadExt;
use tokio::sync::RwLock; use tokio::sync::RwLock;
use tokio_util::io::StreamReader; use tokio_util::io::StreamReader;
use tokio::sync::mpsc;
#[derive(Parser)] #[derive(Parser)]
#[command(name = "mcp-memory-stub")] #[command(name = "mcp-memory-stub")]
@@ -14,142 +15,116 @@ struct Cli {
wake_cmd: Option<String>, wake_cmd: Option<String>,
} }
fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error>> { fn main() -> Result<(), Box<dyn std::error::Error>> {
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 client = reqwest::Client::builder().build()?; let (msg_tx, mut msg_rx) = mpsc::channel::<String>(100);
let sse_url = format!("{}/sse", target_url); let (shutdown_tx, mut shutdown_rx) = mpsc::channel::<()>(1);
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);
}
};
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 || { tokio::task::spawn_blocking(move || {
let stdin = std::io::stdin(); let stdin = std::io::stdin();
let mut handle = stdin.lock(); let mut handle = stdin.lock();
let mut buffer = String::new(); let mut buffer = String::new();
while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) { while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) {
if bytes == 0 { break; } if bytes == 0 { break; }
let body = buffer.clone(); let _ = msg_tx.blocking_send(buffer.clone());
buffer.clear(); buffer.clear();
eprintln!("[PROXY] Stdin read: {}", body.trim()); }
let client = client_clone.clone(); let _ = shutdown_tx.blocking_send(());
let url_arc = Arc::clone(&post_url_clone); });
tokio::spawn(async move { let target_url = cli.target;
let mut url = "".to_string(); let post_url = Arc::new(RwLock::new(String::new()));
for _ in 0..50 { let post_url_clone = Arc::clone(&post_url);
let u = url_arc.read().await.clone();
if u.contains("sessionId") { let client = reqwest::Client::builder().build()?;
url = u;
tokio::spawn(async move {
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; break;
} }
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
} }
if url.is_empty() { tokio::time::sleep(tokio::time::Duration::from_millis(500)).await;
url = url_arc.read().await.clone(); attempts += 1;
} if attempts % 10 == 0 {
eprintln!("[PROXY] POSTing to {}", url); eprintln!("[PROXY] Waiting for server to accept messages...");
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 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();
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;
} }
} }
line.clear();
} }
eprintln!("[PROXY] Stream closed");
let _ = tx_clone2.send(true).await;
}); });
let dropped = rx.recv().await.unwrap_or(true); let wake_cmd = cli.wake_cmd;
Ok(dropped) loop {
}) if shutdown_rx.try_recv().is_ok() {
}
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) => {
break; 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(())
})
} }