feat: Add resilient MPSC queues to proxy stubs to fix thread leaks and message loss on reconnect
This commit is contained in:
1 parent
51fc3692c8
commit
05e84e98bd
2 files changed
+108
-137
No files matched your search
+51
-55
@@ -3,73 +3,66 @@ 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;
|
let _ = shutdown_tx.blocking_send(());
|
||||||
}
|
|
||||||
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 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()));
|
||||||
|
let post_url_clone = Arc::clone(&post_url);
|
||||||
|
|
||||||
|
let client = reqwest::Client::builder().build().unwrap();
|
||||||
|
|
||||||
tokio::spawn(async move {
|
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 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 reader = tokio::io::BufReader::new(StreamReader::new(stream));
|
||||||
let mut line = String::new();
|
let mut line = String::new();
|
||||||
@@ -98,10 +91,13 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
|
|||||||
}
|
}
|
||||||
line.clear();
|
line.clear();
|
||||||
}
|
}
|
||||||
let _ = tx_clone2.send(true).await; // Stream dropped (Leader dead)
|
*post_url.write().await = String::new();
|
||||||
});
|
return Ok(true);
|
||||||
|
}
|
||||||
let dropped = rx.recv().await.unwrap_or(true);
|
Err(_) => {
|
||||||
Ok(dropped)
|
return Ok(true);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
+57
-82
@@ -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,85 +15,81 @@ 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);
|
});
|
||||||
|
|
||||||
|
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 {
|
tokio::spawn(async move {
|
||||||
let mut url = "".to_string();
|
while let Some(msg) = msg_rx.recv().await {
|
||||||
for _ in 0..50 {
|
let mut attempts = 0;
|
||||||
let u = url_arc.read().await.clone();
|
loop {
|
||||||
if u.contains("sessionId") {
|
let url = post_url_clone.read().await.clone();
|
||||||
url = u;
|
if !url.is_empty() {
|
||||||
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)
|
let res = client.post(&url)
|
||||||
.header("Accept", "application/json, text/event-stream")
|
.header("Accept", "application/json, text/event-stream")
|
||||||
.header("Content-Type", "application/json")
|
.header("Content-Type", "application/json")
|
||||||
.body(body)
|
.body(msg.clone())
|
||||||
.send().await;
|
.send().await;
|
||||||
eprintln!("[PROXY] POST response: {:?}", res.map(|r| r.status()));
|
|
||||||
});
|
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...");
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
eprintln!("[PROXY] Stdin closed");
|
|
||||||
let _ = tx_clone.blocking_send(false);
|
|
||||||
});
|
});
|
||||||
|
|
||||||
let target_url = target_url.to_string();
|
let wake_cmd = cli.wake_cmd;
|
||||||
let tx_clone2 = tx.clone();
|
loop {
|
||||||
tokio::spawn(async move {
|
if shutdown_rx.try_recv().is_ok() {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
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 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 reader = tokio::io::BufReader::new(StreamReader::new(stream));
|
||||||
let mut line = String::new();
|
let mut line = String::new();
|
||||||
let mut is_message = false;
|
let mut is_message = false;
|
||||||
let mut is_endpoint = false;
|
let mut is_endpoint = false;
|
||||||
|
|
||||||
while let Ok(bytes) = reader.read_line(&mut line).await {
|
while let Ok(bytes) = reader.read_line(&mut line).await {
|
||||||
if bytes == 0 { break; }
|
if bytes == 0 { break; }
|
||||||
let trimmed = line.trim();
|
let trimmed = line.trim();
|
||||||
eprintln!("[PROXY] Stream read: {}", trimmed);
|
|
||||||
if trimmed.starts_with("event: message") {
|
if trimmed.starts_with("event: message") {
|
||||||
is_message = true;
|
is_message = true;
|
||||||
is_endpoint = false;
|
is_endpoint = false;
|
||||||
@@ -107,29 +104,16 @@ fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error>> {
|
|||||||
let ep = &trimmed[6..];
|
let ep = &trimmed[6..];
|
||||||
let mut p = post_url.write().await;
|
let mut p = post_url.write().await;
|
||||||
*p = format!("{}{}", target_url, ep);
|
*p = format!("{}{}", target_url, ep);
|
||||||
eprintln!("[PROXY] Endpoint updated: {}", *p);
|
|
||||||
is_endpoint = false;
|
is_endpoint = false;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
line.clear();
|
line.clear();
|
||||||
}
|
}
|
||||||
eprintln!("[PROXY] Stream closed");
|
*post_url.write().await = String::new();
|
||||||
let _ = tx_clone2.send(true).await;
|
tokio::time::sleep(tokio::time::Duration::from_millis(500)).await;
|
||||||
});
|
}
|
||||||
|
Err(_) => {
|
||||||
let dropped = rx.recv().await.unwrap_or(true);
|
if let Some(ref cmd) = wake_cmd {
|
||||||
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();
|
let parts: Vec<&str> = cmd.split_whitespace().collect();
|
||||||
if !parts.is_empty() {
|
if !parts.is_empty() {
|
||||||
let _ = std::process::Command::new(parts[0])
|
let _ = std::process::Command::new(parts[0])
|
||||||
@@ -137,19 +121,10 @@ fn main() {
|
|||||||
.spawn();
|
.spawn();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
|
||||||
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;
|
|
||||||
}
|
|
||||||
Err(_) => {
|
|
||||||
std::thread::sleep(std::time::Duration::from_millis(1000));
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
Ok(())
|
||||||
|
})
|
||||||
}
|
}
|
||||||
Reference in new issue
Block a user