Files
mcp-memory/server/src/proxy.rs
T

94 lines
3.7 KiB
Rust

use rust_mcp_sdk::error::SdkResult;
use tokio_util::io::StreamReader;
use tokio::io::AsyncBufReadExt;
use futures_util::StreamExt;
use std::sync::Arc;
use tokio::sync::RwLock;
pub fn run_proxy(target_url: &str) -> SdkResult<bool> {
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);
let resp = match client.get(&sse_url).send().await {
Ok(r) => r,
Err(_) => return Ok(true), // Connection failed (Leader is dead)
};
let post_url = Arc::new(RwLock::new(format!("{}/messages", target_url)));
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();
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("Content-Type", "application/json").body(body).send().await;
});
}
let _ = tx_clone.blocking_send(false); // Stdin EOF
});
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 dropped = rx.recv().await.unwrap_or(true);
Ok(dropped)
})
}