From 6dde1cc3a35fff09fee60b6a1dec781215e6c65c Mon Sep 17 00:00:00 2001 From: Riz Ashraf Date: Sat, 12 Sep 2026 20:11:46 +0100 Subject: [PATCH] refactor: physically decouple server and stub architectures --- server/src/main.rs | 263 +++++++++++++++++++++++++++------------------ stub/Cargo.toml | 1 + stub/src/main.rs | 134 +++++++---------------- 3 files changed, 200 insertions(+), 198 deletions(-) diff --git a/server/src/main.rs b/server/src/main.rs index 6a1e980..adc3bfc 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -118,20 +118,17 @@ async fn reconcile_worker(state: Arc) { use axum::{ Json, Router, - extract::{Query, State}, - response::sse::{Event, Sse}, + extract::{Query, State, ws::{WebSocketUpgrade, WebSocket, Message}}, response::IntoResponse, routing::{get, post}, }; -use futures_util::stream::Stream; -use std::convert::Infallible; +use futures_util::{SinkExt, StreamExt}; use std::sync::atomic::{AtomicUsize, Ordering}; use tokio::sync::mpsc; -use tokio_stream::wrappers::ReceiverStream; struct AppState { handler: Arc, - clients: RwLock>>>, + clients: RwLock>>, next_id: AtomicUsize, } @@ -237,7 +234,7 @@ fn run_server(state: Arc) -> Result<(), Box> rt.block_on(async { tokio::spawn(reconcile_worker(Arc::clone(&state))); let app_state = Arc::new(AppState { - handler: Arc::new(MemoryHandler { state }), + handler: Arc::new(MemoryHandler { state: Arc::clone(&state) }), clients: RwLock::new(HashMap::new()), next_id: AtomicUsize::new(1), }); @@ -249,8 +246,7 @@ fn run_server(state: Arc) -> Result<(), Box> "git_hash": option_env!("GIT_HASH").unwrap_or("unknown") })) })) - .route("/sse", get(sse_handler)) - .route("/messages", post(message_handler)) + .route("/ws", get(ws_handler)) .route("/health", get(health_handler)) .route("/gate/verify", get(gate_verify_handler)) .route("/gate/set", post(gate_set_handler)) @@ -276,7 +272,7 @@ fn run_server(state: Arc) -> Result<(), Box> } })) - .route("/api/tasks/:id/complete", post({ + .route("/api/tasks/{id}/complete", post({ let state_clone = app_state.handler.state.clone(); move |axum::extract::Path(id): axum::extract::Path| async move { state_clone.tasks.modify(|tasks| { @@ -424,7 +420,78 @@ fn run_server(state: Arc) -> Result<(), Box> } } }; - eprintln!("MCP Memory Server running on http://127.0.0.1:3000/sse"); + + // Background Garbage Collection for old tasks + let state_gc = Arc::clone(&state); + tokio::spawn(async move { + loop { + // Run every 24 hours + tokio::time::sleep(tokio::time::Duration::from_secs(24 * 3600)).await; + + let now = std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs(); + let fourteen_days = 14 * 24 * 3600; + let cutoff = now.saturating_sub(fourteen_days); + + state_gc.tasks.modify(|tasks| { + let initial_len = tasks.len(); + tasks.retain(|task| { + if task.status.to_lowercase() == "completed" && task.created_at < cutoff { + false // remove + } else { + true // keep + } + }); + if tasks.len() < initial_len { + eprintln!("GC: Removed {} old completed tasks", initial_len - tasks.len()); + } + }); + } + }); + + // Git Native Sync Background Task + let state_git = Arc::clone(&state); + tokio::spawn(async move { + let repo_path = std::env::current_dir().unwrap_or_else(|_| ".".into()); + let mut last_commit_id = String::new(); + + loop { + tokio::time::sleep(tokio::time::Duration::from_secs(30)).await; + + if let Ok(repo) = git2::Repository::discover(&repo_path) { + if let Ok(head) = repo.head() { + if let Ok(commit) = head.peel_to_commit() { + let current_id = commit.id().to_string(); + if current_id != last_commit_id && !last_commit_id.is_empty() { + let msg = commit.message().unwrap_or("").to_string(); + let branch = head.shorthand().unwrap_or("unknown").to_string(); + + state_git.ledger.modify(|changes| { + changes.push(crate::models::CodeChange { + git_commit: Some(current_id.clone()), + git_branch: Some(branch), + description: format!("Auto-synced commit: {}", msg.trim()), + timestamp: std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs(), + file_path: "".to_string(), + }); + }); + eprintln!("Git Sync: Logged new commit {}", current_id); + + state_git.tasks.modify(|tasks| { + for task in tasks.iter_mut() { + if task.status != "completed" && msg.to_lowercase().contains(&task.title.to_lowercase()) { + task.status = "completed".to_string(); + eprintln!("Git Sync: Auto-completed task '{}'", task.title); + } + } + }); + } + last_commit_id = current_id; + } + } + } + } + }); +eprintln!("MCP Memory Server running on http://127.0.0.1:3000/sse"); if let Err(e) = axum::serve(listener, app).await { let log_path = dirs::home_dir() .unwrap_or_default() @@ -435,81 +502,84 @@ fn run_server(state: Arc) -> Result<(), Box> }) } -async fn sse_handler( +async fn ws_handler( + ws: WebSocketUpgrade, State(state): State>, -) -> Sse>> { + Query(query): Query>, +) -> impl axum::response::IntoResponse { + let client_type = query.get("client").cloned().unwrap_or_else(|| "unknown".to_string()); + ws.on_upgrade(move |socket| handle_socket(socket, state, client_type)) +} + +async fn handle_socket(socket: WebSocket, state: Arc, client_type: String) { let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst)); - let (tx, rx) = mpsc::channel::>(100); + let (tx, mut rx) = mpsc::channel::(100); - state - .clients - .write() - .unwrap() - .insert(session_id.clone(), tx.clone()); + state.clients.write().unwrap().insert(session_id.clone(), tx.clone()); - let _ = tx - .send(Ok(Event::default() - .event("endpoint") - .data(format!("/messages?sessionId={}", session_id)))) - .await; + let (mut sender, mut receiver) = socket.split(); - let stream = ReceiverStream::new(rx); - Sse::new(stream).keep_alive(axum::response::sse::KeepAlive::new()) + let mut send_task = tokio::spawn(async move { + while let Some(msg) = rx.recv().await { + if sender.send(Message::Text(msg.into())).await.is_err() { + break; + } + } + }); + + let handler = Arc::clone(&state.handler); + let state_clone = Arc::clone(&state); + let session_id_clone = session_id.clone(); + + let mut recv_task = tokio::spawn(async move { + while let Some(Ok(Message::Text(text))) = receiver.next().await { + if let Ok(payload) = serde_json::from_str::(&text) { + if client_type == "proxy" { + // Send activity broadcast to UI clients + if let Some(method) = payload.get("method").and_then(|m| m.as_str()) { + if method == "tools/call" { + let name = payload.get("params").and_then(|p| p.get("name")).and_then(|n| n.as_str()).unwrap_or("unknown_tool"); + let activity_msg = format!("Agent executed tool: {}", name); + + let event = serde_json::json!({ + "type": "activity", + "data": activity_msg + }); + + let clients_map = state_clone.clients.read().unwrap().clone(); + for (id, client_tx) in clients_map.iter() { + if id != &session_id_clone { + let _ = client_tx.send(event.to_string()).await; + } + } + } + } + + // Process MCP request + if let Some(response) = handler.handle_request(payload).await { + let res_str = serde_json::to_string(&response).unwrap(); + let tx_opt = state_clone.clients.read().unwrap().get(&session_id_clone).cloned(); + if let Some(client_tx) = tx_opt { + let _ = client_tx.send(res_str).await; + } + } + } + } + } + }); + + tokio::select! { + _ = (&mut send_task) => recv_task.abort(), + _ = (&mut recv_task) => send_task.abort(), + }; + + state.clients.write().unwrap().remove(&session_id); } async fn health_handler() -> &'static str { "OK" } -#[derive(serde::Deserialize)] -struct SessionQuery { - #[serde(rename = "sessionId")] - session_id: String, -} - -async fn message_handler( - State(state): State>, - Query(query): Query, - Json(payload): Json, -) -> axum::http::StatusCode { - let handler = Arc::clone(&state.handler); - let session_id = query.session_id.clone(); - let clients = Arc::clone(&state); - - let activity_msg = if let Some(method) = payload.get("method").and_then(|m| m.as_str()) { - if method == "tools/call" { - let name = payload.get("params").and_then(|p| p.get("name")).and_then(|n| n.as_str()).unwrap_or("unknown_tool"); - Some(format!("Agent executed tool: {}", name)) - } else { - None - } - } else { - None - }; - - if let Some(msg) = activity_msg { - let clients_map = clients.clients.read().unwrap().clone(); - for tx in clients_map.values() { - let _ = tx.send(Ok(Event::default().event("activity").data(msg.clone()))); - } - } - - tokio::spawn(async move { - if let Some(response) = handler.handle_request(payload).await { - let tx_opt = clients.clients.read().unwrap().get(&session_id).cloned(); - if let Some(tx) = tx_opt { - let data = serde_json::to_string(&response).unwrap(); - let _ = tx - .send(Ok(Event::default().event("message").data(data))) - .await; - } - } - }); - - axum::http::StatusCode::ACCEPTED -} - -mod proxy; fn main() -> Result<(), Box> { let cli = Cli::parse(); @@ -541,41 +611,23 @@ fn main() -> Result<(), Box> { { use std::os::windows::process::CommandExt; if !cli.daemon { - loop { - if std::net::TcpListener::bind("127.0.0.1:3000").is_err() { - // Port in use, become a stub proxy! - let target_url = cli.target.as_deref().unwrap_or("http://127.0.0.1:3000"); - match proxy::run_proxy(target_url) { - Ok(true) => { - std::thread::sleep(std::time::Duration::from_millis(50)); - continue; // Leader died, race to bind 3000 - } - Ok(false) => return Ok(()), // Stdin closed, user exited - Err(_) => std::thread::sleep(std::time::Duration::from_millis(1000)), - } - } else { - // Port is free. We must spawn the daemon, then loop again to become proxy - #[allow(clippy::zombie_processes)] - let _ = std::process::Command::new(std::env::current_exe().unwrap()) - .arg("--daemon") - .stdin(std::process::Stdio::null()) - .stdout(std::process::Stdio::null()) - .stderr(std::process::Stdio::null()) - .creation_flags(0x08000000) // CREATE_NO_WINDOW - .spawn() - .expect("Failed to spawn daemon"); - std::thread::sleep(std::time::Duration::from_millis(500)); - } - } + // Just spawn the daemon and exit. We no longer act as a proxy. + #[allow(clippy::zombie_processes)] + let _ = std::process::Command::new(std::env::current_exe().unwrap()) + .arg("--daemon") + .stdin(std::process::Stdio::null()) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .creation_flags(0x08000000) // CREATE_NO_WINDOW + .spawn() + .expect("Failed to spawn daemon"); + return Ok(()); } } #[cfg(not(target_os = "windows"))] { - // This shouldn't be executed on linux natively anymore due to workspace split, - // but keeping it as a fallback. - let target_url = cli.target.as_deref().unwrap_or("http://127.0.0.1:3000"); - let _ = proxy::run_proxy(target_url); + // Linux no longer executes server logic natively due to workspace split return Ok(()); } @@ -768,3 +820,4 @@ fn main() -> Result<(), Box> { } + diff --git a/stub/Cargo.toml b/stub/Cargo.toml index 1821b56..ac91716 100644 --- a/stub/Cargo.toml +++ b/stub/Cargo.toml @@ -9,3 +9,4 @@ reqwest = { version = "0.12", default-features = false, features = ["stream", "r tokio = { version = "1.53.1", features = ["full"] } tokio-util = { version = "0.7.19", features = ["io"] } futures-util = "0.3.34" +tokio-tungstenite = "0.21.0" diff --git a/stub/src/main.rs b/stub/src/main.rs index a7e1bc9..113c863 100644 --- a/stub/src/main.rs +++ b/stub/src/main.rs @@ -1,10 +1,8 @@ -use clap::Parser; -use futures_util::StreamExt; use std::sync::Arc; +use clap::Parser; +use futures_util::{SinkExt, StreamExt}; use tokio::io::AsyncBufReadExt; -use tokio::sync::RwLock; use tokio::sync::mpsc; -use tokio_util::io::StreamReader; #[derive(Parser)] #[command(name = "mcp-memory-stub", author, version, about = "Antigravity MCP Memory Stub / Proxy", long_about = None)] @@ -38,104 +36,56 @@ fn main() -> Result<(), Box> { }); let target_url = cli.target; - let post_url = Arc::new(RwLock::new(String::new())); - let post_url_proxy = Arc::clone(&post_url); - - let client = reqwest::Client::builder().build()?; - - tokio::spawn(async move { - while let Some(msg) = msg_rx.recv().await { - let post_url_clone = Arc::clone(&post_url_proxy); - let client = client.clone(); - tokio::spawn(async move { - 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 let Ok(resp) = res { if resp.status().is_success() { break; } } - } - tokio::time::sleep(tokio::time::Duration::from_millis(10)).await; - attempts += 1; - if attempts % 10 == 0 { - eprintln!("[PROXY] Waiting for server to accept messages..."); - } - } - }); - } - }); - + 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() { break; } - let sse_url = format!("{}/sse", target_url); - let client = reqwest::Client::builder().build()?; + match tokio_tungstenite::connect_async(&ws_url).await { + Ok((ws_stream, _)) => { + let (mut write, mut read) = ws_stream.split(); - 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(std::io::Error::other) - }); - 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; - - loop { - tokio::select! { - _ = shutdown_rx.recv() => { - return Ok(()); // Stdin closed, exit entirely - } - res = reader.read_line(&mut line) => { - match res { - Ok(bytes) => { - 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 let Some(stripped) = trimmed.strip_prefix("data: ") { - if is_message { - println!("{}", stripped); - is_message = false; - } else if is_endpoint { - let mut p = post_url.write().await; - *p = format!("{}{}", target_url, stripped); - is_endpoint = false; - } - } - line.clear(); + let msg_rx_clone = Arc::clone(&msg_rx); + let mut send_task = tokio::spawn(async move { + loop { + let mut rx = msg_rx_clone.lock().await; + match rx.recv().await { + Some(msg) => { + drop(rx); + if write.send(tokio_tungstenite::tungstenite::Message::Text(msg)).await.is_err() { + break; } - Err(_) => break, - } + }, + None => break, } } + }); + + 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 { + println!("{}", text); + } + } + }); + + tokio::select! { + _ = shutdown_rx.recv() => { + return Ok(()); // Stdin closed, exit entirely + } + _ = &mut send_task => { + recv_task.abort(); + } + _ = &mut recv_task => { + send_task.abort(); + } } - *post_url.write().await = String::new(); - tokio::time::sleep(tokio::time::Duration::from_millis(10)).await; + tokio::time::sleep(tokio::time::Duration::from_millis(50)).await; } Err(_) => { if let Some(ref cmd) = wake_cmd { @@ -153,5 +103,3 @@ fn main() -> Result<(), Box> { Ok(()) }) } - -