refactor(mcp): strip sse fallback architecture in favor of pure websockets and fix nvim NDJSON bug

This commit is contained in:
Riz Ashraf committed 2026-09-17 15:26:22 +01:00
1 parent 0e29b12ac8
commit 3716c3e698
33 files changed
+2082 -1756

No files matched your search

+35 -37
View File
@@ -10,9 +10,6 @@ struct Cli {
/// Target URL for the stub to proxy messages to
#[arg(long, default_value = "http://localhost:3000")]
target: String,
/// Optional command to execute if the target server is unreachable
#[arg(long)]
wake_cmd: Option<String>,
}
async fn read_mcp_message(stdin: &mut tokio::io::BufReader<tokio::io::Stdin>) -> Option<String> {
@@ -26,6 +23,11 @@ async fn read_mcp_message(stdin: &mut tokio::io::BufReader<tokio::io::Stdin>) ->
return None;
}
tracing::info!("Read {} bytes from stdin: {:?}", bytes_read, line);
if line.starts_with('{') {
return Some(line.trim_end().to_string());
}
let line = line.trim_end();
if line.is_empty() {
break;
@@ -46,24 +48,17 @@ async fn read_mcp_message(stdin: &mut tokio::io::BufReader<tokio::io::Stdin>) ->
}
fn init_logging(app_name: &str) -> Option<tracing_appender::non_blocking::WorkerGuard> {
let base_dir = std::env::var("MCP_MEMORY_STORE_DIR").unwrap_or_else(|_| {
dirs::home_dir()
.map(|mut h| {
h.push(".gemini/mcp_memory");
h.to_string_lossy().to_string()
})
.unwrap_or_else(|| ".gemini/mcp_memory".into())
});
let log_dir = std::path::PathBuf::from(base_dir).join("logs");
std::fs::create_dir_all(&log_dir).unwrap_or_default();
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(log_dir, format!("{}.log", app_name));
let file_appender = tracing_appender::rolling::daily(base_dir, format!("{}.log", app_name));
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::INFO)
.with_max_level(tracing::Level::TRACE)
.try_init();
Some(guard)
@@ -89,7 +84,6 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let ws_url = target_url.replace("http://", "ws://").replace("https://", "wss://");
let ws_url = format!("{}/ws?client=proxy", ws_url);
let msg_rx = Arc::new(tokio::sync::Mutex::new(msg_rx));
let wake_cmd = cli.wake_cmd;
loop {
if shutdown_rx.try_recv().is_ok() {
@@ -98,7 +92,18 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
}
tracing::info!("Attempting to connect to {}", ws_url);
match tokio_tungstenite::connect_async(&ws_url).await {
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
let mut 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_secs(1)).await;
continue;
}
};
match tokio_tungstenite::connect_async(request).await {
Ok((ws_stream, _)) => {
tracing::info!("Successfully connected to target server");
let (mut write, mut read) = ws_stream.split();
@@ -110,7 +115,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
match rx.recv().await {
Some(msg) => {
drop(rx);
tracing::info!("Forwarding message to target server");
tracing::info!("Forwarding message to target server (length: {}): {}", msg.len(), if msg.len() > 1000 { format!("{}...", &msg[..1000]) } else { msg.clone() });
if write.send(tokio_tungstenite::tungstenite::Message::Text(msg)).await.is_err() {
tracing::error!("Failed to write to websocket");
break;
@@ -124,25 +129,25 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let mut recv_task = tokio::spawn(async move {
while let Some(Ok(msg)) = read.next().await {
if let tokio_tungstenite::tungstenite::Message::Text(text) = msg {
tracing::info!("Received message from target server, proxying to stdout");
let payload = format!("Content-Length: {}\r\n\r\n{}", text.len(), text);
use std::io::Write;
let mut stdout = std::io::stdout();
let _ = stdout.write_all(payload.as_bytes());
let _ = stdout.flush();
tracing::info!("Received message from target server (length: {}): {}", text.len(), if text.len() > 1000 { format!("{}...", &text[..1000]) } else { text.clone() });
let payload = format!("{}\n", text);
use tokio::io::AsyncWriteExt;
let mut stdout = tokio::io::stdout();
let _ = stdout.write_all(payload.as_bytes()).await;
let _ = stdout.flush().await;
}
}
tracing::error!("Websocket read loop exited");
});
tokio::select! {
_ = shutdown_rx.recv() => {
_ = shutdown_rx.recv() => { tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
tracing::info!("Shutdown received while connected");
return Ok(()); // Stdin closed, exit entirely
}
_ = &mut send_task => {
tracing::error!("Send task exited");
recv_task.abort();
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; recv_task.abort();
}
_ = &mut recv_task => {
tracing::error!("Recv task exited");
@@ -152,16 +157,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
}
Err(e) => {
tracing::error!("Failed to connect to target server: {}", e);
if let Some(ref cmd) = wake_cmd {
tracing::info!("Executing wake command: {}", cmd);
let parts: Vec<&str> = cmd.split_whitespace().collect();
if !parts.is_empty() {
let _ = std::process::Command::new(parts[0])
.args(&parts[1..])
.spawn();
}
}
tracing::error!("Failed to connect via WSS: {}", e);
tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
}
}
@@ -169,3 +165,5 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
Ok(())
})
}