Files
mcp-memory/stub/src/main.rs
T

181 lines
7.5 KiB
Rust

use clap::Parser;
use futures_util::{SinkExt, StreamExt};
#[derive(Parser)]
#[command(name = "mcp-memory-stub", author, version = env!("APP_VERSION"), about = "Antigravity MCP Memory Stub / Proxy", long_about = None)]
struct Cli {
/// Target URL for the stub to proxy messages to
#[arg(long, default_value = "http://localhost:3000")]
target: String,
}
mod logger;
fn init_logging(app_name: &str) -> Option<tracing_appender::non_blocking::WorkerGuard> {
let mut base_dir = dirs::home_dir().unwrap_or_else(|| std::path::PathBuf::from("."));
base_dir.push(".gemini/mcp_memory/logs");
std::fs::create_dir_all(&base_dir).unwrap_or_default();
let file_appender = tracing_appender::rolling::daily(base_dir, format!("{app_name}.log"));
let (non_blocking, guard) = tracing_appender::non_blocking(file_appender);
let _ = tracing_subscriber::fmt()
.with_writer(non_blocking)
.with_ansi(false)
.with_max_level(tracing::Level::TRACE)
.try_init();
Some(guard)
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _guard = init_logging("stub");
let cli = Cli::parse();
let rt = tokio::runtime::Runtime::new()?;
rt.block_on(async {
let (msg_tx, msg_rx) = async_channel::bounded::<String>(100);
let (shutdown_tx, mut shutdown_rx) = tokio::sync::watch::channel(false);
tokio::spawn(async move {
let mut stdin = tokio::io::BufReader::new(tokio::io::stdin());
while let Some(msg) = mcp_stdio::read_mcp_message(&mut stdin).await {
let _ = msg_tx.send(msg).await;
}
let _ = shutdown_tx.send(true);
});
let target_url = if cli.target == "http://localhost:3000" {
let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string());
format!("http://127.0.0.1:{port}")
} else {
cli.target
};
let ws_url = target_url
.replace("http://", "ws://")
.replace("https://", "wss://");
let ws_url = format!("{ws_url}/ws?client=proxy");
let mut retry_count = 0;
loop {
if *shutdown_rx.borrow() {
tracing::info!("Stub shutdown requested");
break;
}
tracing::info!("Attempting to connect to {}", ws_url);
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
let request = match ws_url.clone().into_client_request() {
Ok(req) => req,
Err(e) => {
tracing::error!("Failed to parse target URL {}: {}", ws_url, e);
tokio::time::sleep(tokio::time::Duration::from_millis(250)).await;
continue;
}
};
let connect_result = tokio::select! {
_ = shutdown_rx.wait_for(|&is_shutdown| is_shutdown) => {
tracing::info!("Shutdown received during connect");
break;
}
res = tokio::time::timeout(
tokio::time::Duration::from_secs(5),
tokio_tungstenite::connect_async(request)
) => res,
};
match connect_result {
Ok(Ok((ws_stream, _))) => {
retry_count = 0;
tracing::info!("Successfully connected to target server");
let (mut write, mut read) = ws_stream.split();
let rx = msg_rx.clone();
let mut send_task = tokio::spawn(async move {
while let Ok(msg) = rx.recv().await {
let log_prefix = logger::extract_log_prefix(&msg, false);
tracing::info!(
">>> [Stub] Forwarding {} to server (length: {}): {}{}",
log_prefix,
msg.len(),
&msg[..msg.floor_char_boundary(1000)],
if msg.len() > 1000 { "..." } else { "" }
);
if write
.send(tokio_tungstenite::tungstenite::Message::Text(msg.into()))
.await
.is_err()
{
tracing::warn!("Failed to write to websocket (connection closed)");
break;
}
}
});
let mut recv_task = tokio::spawn(async move {
use tokio::io::AsyncWriteExt;
let mut stdout = tokio::io::BufWriter::new(tokio::io::stdout());
while let Some(Ok(msg)) = read.next().await {
if let tokio_tungstenite::tungstenite::Message::Text(text) = msg {
let log_prefix = logger::extract_log_prefix(&text, true);
tracing::info!(
"<<< [Stub] Received {} from server (length: {}): {}{}",
log_prefix,
text.len(),
&text[..text.floor_char_boundary(1000)],
if text.len() > 1000 { "..." } else { "" }
);
let _ = stdout.write_all(text.as_bytes()).await;
let _ = stdout.write_all(b"\n").await;
let _ = stdout.flush().await;
}
}
tracing::info!("Websocket read loop exited");
});
tokio::select! {
_ = shutdown_rx.wait_for(|&is_shutdown| is_shutdown) => {
tracing::info!("Shutdown received while connected");
break;
}
_ = &mut send_task => {
tracing::error!("Send task exited");
recv_task.abort();
}
_ = &mut recv_task => {
tracing::error!("Recv task exited");
send_task.abort();
}
}
tokio::time::sleep(tokio::time::Duration::from_millis(250)).await;
}
Ok(Err(e)) => {
tracing::error!("Failed to connect via WSS: {}", e);
retry_count += 1;
if retry_count >= 20 {
tracing::error!("Max connection retries reached. Exiting proxy.");
break;
}
tokio::time::sleep(tokio::time::Duration::from_millis(250)).await;
}
Err(_) => {
tracing::error!("Connection attempt timed out");
retry_count += 1;
if retry_count >= 20 {
tracing::error!("Max connection retries reached. Exiting proxy.");
break;
}
tokio::time::sleep(tokio::time::Duration::from_millis(250)).await;
}
}
}
if retry_count >= 20 {
return Err("Failed to connect to the MCP memory server after multiple attempts. Is it running?".into());
}
Ok(())
})
}