#![cfg_attr( not(target_os = "windows"), allow(dead_code, unused_imports, unreachable_code) )] mod handlers; mod mcp; mod models; mod search; mod state; mod store; mod tools; use crate::handlers::MemoryHandler; use crate::models::*; use crate::state::MemoryState; use crate::store::Store; use redb::ReadableTable; use std::fs; use std::path::PathBuf; use std::sync::{Arc, RwLock}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use tokio::time::sleep; use clap::{Parser, Subcommand}; use std::collections::HashMap; #[derive(Parser)] #[command(author, version = env!("APP_VERSION"), about = "Antigravity MCP Memory Server", long_about = None)] struct Cli { #[command(subcommand)] command: Option, /// Target URL for the proxy to connect to (e.g., http://127.0.0.1:3000) #[arg(long)] target: Option, /// Run the server as a background daemon process (Windows only) #[arg(long)] daemon: bool, /// Send a shutdown request to the currently running server #[arg(long)] exit: bool, /// Send a shutdown request to the existing server and wait for it to exit #[arg(long)] restart: bool, } #[derive(Subcommand)] enum Commands { /// Manage authorization gates and verification for actions Gate { #[command(subcommand)] subcmd: GateCommands, }, } #[derive(Subcommand)] enum GateCommands { Set { #[arg(long)] action: String, #[arg(long)] target: String, #[arg(long)] namespace: Option, #[arg(short = 'p', long = "param")] params: Vec, #[arg(long, conflicts_with = "block")] authorize: bool, #[arg(long, conflicts_with = "authorize")] block: bool, #[arg(long)] reason: Option, }, Verify { #[arg(long)] action: String, #[arg(long)] target: String, #[arg(long)] namespace: Option, #[arg(short = 'p', long = "param")] params: Vec, #[arg(long)] consume: bool, }, } async fn index_committer_worker(state: Arc) { loop { sleep(Duration::from_secs(5)).await; // Periodically commit the search index to persist inline indexing operations if let Ok(idx) = state.search_index.read() { let _ = idx.commit(); } } } use axum::{ Json, Router, extract::{ Query, State, ws::{Message, WebSocket}, }, response::IntoResponse, routing::{get, post}, }; use futures_util::{SinkExt, StreamExt}; use std::sync::atomic::{AtomicUsize, Ordering}; use tokio::sync::mpsc; struct AppState { handler: Arc, clients: RwLock>>, next_id: AtomicUsize, } #[derive(serde::Deserialize)] struct GateVerifyReq { action: String, target: String, namespace: Option, #[serde(default)] params: HashMap, #[serde(default)] consume: bool, } #[derive(serde::Deserialize)] struct GateSetReq { action: String, target: String, namespace: Option, #[serde(default)] params: HashMap, authorize: Option, block: Option, reason: Option, } async fn gate_verify_handler( State(app_state): State>, Query(q): Query, ) -> axum::response::Response { let mut found = None; let mut to_remove = None; app_state.handler.state.gates.modify(|gates| { if let Some(idx) = gates.iter().position(|g| { g.action == q.action && g.target == q.target && g.namespace == q.namespace && g.params == q.params }) { found = Some(gates[idx].clone()); if q.consume { to_remove = Some(idx); } } if let Some(idx) = to_remove { gates.remove(idx); } }); match found { Some(record) => { if record.status == "authorized" { (axum::http::StatusCode::OK, "Authorized").into_response() } else { let msg = if let Some(r) = record.reason { format!("Action blocked. Reason: {}", r) } else { "Action blocked.".to_string() }; (axum::http::StatusCode::FORBIDDEN, msg).into_response() } } None => ( axum::http::StatusCode::NOT_FOUND, "Action not yet authorized (no gate record found).", ) .into_response(), } } async fn gate_set_handler( State(app_state): State>, Json(body): Json, ) -> axum::response::Response { let status = if body.block.unwrap_or(false) { "blocked".to_string() } else if body.authorize.unwrap_or(false) { "authorized".to_string() } else { "pending".to_string() }; let record = GateRecord { id: uuid::Uuid::new_v4().to_string(), action: body.action.clone(), target: body.target.clone(), namespace: body.namespace.clone(), params: body.params.clone(), status, reason: body.reason.clone(), timestamp: SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() .as_secs(), }; app_state.handler.state.gates.modify(|gates| { gates.retain(|g| !(g.action == record.action && g.target == record.target)); gates.push(record); }); (axum::http::StatusCode::OK, "Gate state updated.").into_response() } fn run_server(state: Arc) -> Result<(), Box> { let rt = tokio::runtime::Runtime::new().unwrap(); rt.block_on(async { tokio::spawn(index_committer_worker(Arc::clone(&state))); let app_state = Arc::new(AppState { handler: Arc::new(MemoryHandler { state: Arc::clone(&state), }), clients: RwLock::new(HashMap::new()), next_id: AtomicUsize::new(1), }); let app = Router::new() .route( "/api/version", get(|| async move { axum::Json(serde_json::json!({ "version": env!("APP_VERSION"), "git_hash": option_env!("GIT_HASH").unwrap_or("unknown") })) }), ) .route("/ws", get(ws_handler)) .route("/health", get(health_handler)) .route("/nvim/telemetry", post(nvim_telemetry_handler)) .route("/gate/verify", get(gate_verify_handler)) .route("/gate/set", post(gate_set_handler)) .route( "/shutdown", post( |headers: axum::http::HeaderMap, State(state): State>| async move { let token_path = state.handler.state.base_dir.join("admin.token"); let expected_token = std::fs::read_to_string(&token_path) .unwrap_or_default() .trim() .to_string(); let auth_header = headers .get(axum::http::header::AUTHORIZATION) .and_then(|h| h.to_str().ok()) .unwrap_or_default(); if expected_token.is_empty() || auth_header != format!("Bearer {}", expected_token) { return (axum::http::StatusCode::UNAUTHORIZED, "Unauthorized").into_response(); } std::thread::spawn(|| { tracing::info!( "Received shutdown request via /shutdown endpoint. Exiting process cleanly." ); std::thread::sleep(std::time::Duration::from_millis(100)); std::process::exit(0); }); (axum::http::StatusCode::OK, "Shutting down...").into_response() }, ), ) .route( "/", get(|| async move { axum::response::Html(include_str!("dashboard.html")) }), ) .route( "/api/graph", get({ let state_clone = app_state.handler.state.clone(); move || async move { let graph = state_clone.get_full_graph(); axum::Json(graph) } }), ) .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| { for t in tasks.iter_mut() { if t.id == id { t.status = "completed".to_string(); break; } } }); axum::Json(serde_json::json!({"status": "success"})) } }), ) .route( "/api/tasks", get({ let state_clone = app_state.handler.state.clone(); move || async move { let tasks = state_clone.tasks.read(); axum::Json(tasks.clone()) } }), ) .route( "/api/sticky", get({ let state_clone = app_state.handler.state.clone(); move || async move { let sticky = state_clone.sticky.read(); axum::Json(sticky.clone()) } }), ) .route( "/api/search", get({ let state_clone = app_state.handler.state.clone(); move |axum::extract::Query(params): axum::extract::Query< std::collections::HashMap, >| async move { if let Some(q) = params.get("q") && let Ok(idx) = state_clone.search_index.read() && let Ok(results) = idx.search(q, None) { let mut formatted_results = Vec::new(); for (id, doc_type, title, body, score) in results { formatted_results.push(serde_json::json!({ "id": id, "type_name": doc_type, "title": title, "content": body, "score": score })); } return axum::Json( serde_json::json!({ "results": formatted_results }), ); } axum::Json(serde_json::json!({ "results": [] })) } }), ) .route( "/api/stats", get({ let state_clone = app_state.handler.state.clone(); move || async move { let (entities, relations) = { let graph = state_clone.get_full_graph(); (graph.entities.len(), graph.relations.len()) }; let tasks = state_clone.tasks.read().len(); let snippets = state_clone.snippets.read().len(); let tech_debts = state_clone.tech_debts.read().len(); let adrs = state_clone.adrs.read().len(); let ledger = state_clone.ledger.read().len(); let sticky = state_clone.sticky.read().len(); let error_fixes = state_clone.error_fixes.read().len(); let pinned_files = state_clone.pinned_files.read().len(); let session_summaries = state_clone.session_summaries.read().len(); let handoff_memos = state_clone.handoff_memos.read().len(); let env_fingerprints = state_clone.env_fingerprints.read().len(); let env_requirements = state_clone.env_requirements.read().len(); let milestones = state_clone.milestones.read().len(); let environments = state_clone.environments.read().len(); let pr_checklists = state_clone.pr_checklists.read().len(); let gates = state_clone.gates.read().len(); let context_workspaces = state_clone.context_workspaces.read().len(); axum::Json(serde_json::json!({ "entities": entities, "relations": relations, "tasks": tasks, "snippets": snippets, "tech_debts": tech_debts, "adrs": adrs, "ledger": ledger, "sticky": sticky, "error_fixes": error_fixes, "pinned_files": pinned_files, "session_summaries": session_summaries, "handoff_memos": handoff_memos, "env_fingerprints": env_fingerprints, "env_requirements": env_requirements, "milestones": milestones, "environments": environments, "pr_checklists": pr_checklists, "gates": gates, "context_workspaces": context_workspaces })) } }), ) .with_state(app_state); tracing::info!("MCP Memory Server running on http://127.0.0.1:3000/sse"); let addr = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string()); let addr: std::net::SocketAddr = format!("127.0.0.1:{}", addr).parse().unwrap(); let listener = match tokio::net::TcpListener::bind(addr).await { Ok(l) => l, Err(e) => { let log_path = dirs::home_dir() .unwrap_or_default() .join(".gemini/mcp_memory/daemon_error.log"); let _ = std::fs::write(&log_path, format!("Failed to bind to {}: {}\n", addr, e)); return Ok(()); } }; if let Err(e) = axum::serve(listener, app.into_make_service()).await { let log_path = dirs::home_dir() .unwrap_or_default() .join(".gemini/mcp_memory/daemon_error.log"); let _ = std::fs::write(&log_path, format!("Server crashed: {}\n", e)); } Ok(()) }) } async fn ws_handler( ws: axum::extract::ws::WebSocketUpgrade, _headers: axum::http::HeaderMap, axum::extract::State(state): axum::extract::State>, axum::extract::Query(query): axum::extract::Query>, ) -> axum::response::Response { let client_type = query .get("client") .cloned() .unwrap_or_else(|| "unknown".to_string()); ws.on_upgrade(move |socket| handle_socket(socket, state, client_type)) .into_response() } 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, mut rx) = mpsc::channel::(100); state .clients .write() .unwrap() .insert(session_id.clone(), tx.clone()); let (mut sender, mut receiver) = socket.split(); let mut send_task = tokio::spawn(async move { while let Some(msg) = rx.recv().await { tracing::trace!( "Sending message to websocket (length: {}): {}", msg.len(), msg ); if sender.send(Message::Text(msg.into())).await.is_err() { tracing::error!("Failed to send message to websocket"); break; } } }); // Premature list_changed notification removed for MCP protocol compliance 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(msg_result) = receiver.next().await { match msg_result { Ok(Message::Text(text)) => { tracing::info!( "Received text message from websocket (length: {})", text.len() ); tracing::trace!("Message content: {}", text); 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()) && 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; } } } } // End if proxy // 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 { if let Err(e) = client_tx.send(res_str).await { tracing::error!( "Failed to send response to client channel for session {}: {}", session_id_clone, e ); } } else { tracing::warn!( "Could not find client_tx for session_id {} when trying to send response", session_id_clone ); } } } // End if let Ok(payload) else { tracing::warn!( "Failed to parse payload as JSON from websocket message: {}", text ); } } // End Ok(Message::Text(text)) Ok(other) => { tracing::info!("Received non-text message from websocket: {:?}", other); } Err(e) => { tracing::error!("Websocket receive error: {}", e); break; } } } tracing::info!( "Websocket receiver task ended for session {}", session_id_clone ); }); tokio::select! { _ = (&mut send_task) => { tracing::info!("Websocket send task finished for session {}", session_id); recv_task.abort(); }, _ = (&mut recv_task) => { tracing::info!("Websocket recv task finished for session {}", session_id); send_task.abort(); }, }; state.clients.write().unwrap().remove(&session_id); tracing::info!( "Websocket session {} closed and removed from state", session_id ); } #[derive(serde::Deserialize, serde::Serialize, Debug)] pub struct NvimTelemetry { pub session_id: String, pub event: String, pub file: Option, pub line: Option, pub col: Option, } async fn nvim_telemetry_handler( State(state): State>, axum::Json(payload): axum::Json, ) -> impl axum::response::IntoResponse { // 1. Update active_nvim.txt if FocusGained, BufEnter, or VimEnter if payload.event == "FocusGained" || payload.event == "BufEnter" || payload.event == "VimEnter" { let profile = std::env::var("USERPROFILE").unwrap_or_else(|_| "C:\\Users\\reazul.ashraf".into()); let win_path = format!("{}\\.gemini\\active_nvim.txt", profile); let _ = std::fs::write(&win_path, &payload.session_id); let wsl_path = "\\\\wsl.localhost\\Ubuntu\\home\\riz\\.gemini\\active_nvim.txt"; let _ = std::fs::write(wsl_path, &payload.session_id); } // 2. Broadcast to UI WebSockets let ws_msg = serde_json::json!({ "type": "nvim_telemetry", "data": payload }); let msg_str = ws_msg.to_string(); let clients = state.clients.read().unwrap().clone(); for tx in clients.values() { let _ = tx.send(msg_str.clone()).await; } axum::Json(serde_json::json!({"status": "ok"})) } async fn health_handler() -> &'static str { "OK" } fn init_logging(app_name: &str) -> Option { 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 file_appender = tracing_appender::rolling::daily(log_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_thread_ids(true) .with_thread_names(true) .try_init(); Some(guard) } fn main() -> Result<(), Box> { let _guard = init_logging("mcp-memory-server"); let cli = Cli::parse(); 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().into_owned() }) .unwrap_or_else(|| ".gemini/mcp_memory".into()) }); let base = PathBuf::from(base_dir); if cli.exit { let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string()); let token = std::fs::read_to_string(base.join("admin.token")).unwrap_or_default(); let mut cmd = std::process::Command::new("curl"); cmd.arg("-k").arg("-X").arg("POST"); if !token.is_empty() { cmd.arg("-H").arg(format!("Authorization: Bearer {}", token.trim())); } let _ = cmd.arg(format!("https://127.0.0.1:{}/shutdown", port)).output(); println!("Sent shutdown request to server."); return Ok(()); } if cli.restart { let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string()); let token = std::fs::read_to_string(base.join("admin.token")).unwrap_or_default(); let mut cmd = std::process::Command::new("curl"); cmd.arg("-k").arg("-X").arg("POST"); if !token.is_empty() { cmd.arg("-H").arg(format!("Authorization: Bearer {}", token.trim())); } let _ = cmd.arg(format!("https://127.0.0.1:{}/shutdown", port)).output(); println!("Sent shutdown request to existing server. Waiting for it to exit..."); std::thread::sleep(std::time::Duration::from_millis(1500)); return Ok(()); } #[cfg(target_os = "windows")] { use std::os::windows::process::CommandExt; if !cli.daemon { // 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(()); } } fs::create_dir_all(&base).expect("Failed to create store dir"); // Generate token let admin_token = uuid::Uuid::new_v4().to_string(); std::fs::write(base.join("admin.token"), &admin_token).expect("Failed to write admin token"); let redb_path = base.join("mcp_store.redb"); let db = Arc::new(redb::Database::create(&redb_path).unwrap()); // Ensure table exists and migrate old JSON files { let write_txn = db.begin_write().unwrap(); { let mut table = write_txn.open_table(crate::store::STORE_TABLE).unwrap(); let stores = vec![ ("knowledge_graph_master", "knowledge_graph_master.json"), ("audit_ledger", "audit_ledger.json"), ("sticky_notes", "sticky_notes.json"), ("tasks", "tasks.json"), ("snippets", "snippets.json"), ("adrs", "adrs.json"), ("preferences", "preferences.json"), ("error_fixes", "error_fixes.json"), ("pinned_files", "pinned_files.json"), ("session_summaries", "session_summaries.json"), ("handoff_memos", "handoff_memos.json"), ("env_fingerprints", "env_fingerprints.json"), ("env_requirements", "env_requirements.json"), ("milestones", "milestones.json"), ("environments", "environments.json"), ("pr_checklists", "pr_checklists.json"), ("tech_debts", "tech_debts.json"), ("gates", "gates.json"), ("context_workspaces", "context_workspaces.json"), ]; for (key, file_name) in stores.iter() { if table.get(*key).unwrap().is_none() { let json_path = base.join(file_name); if json_path.exists() && let Ok(data) = fs::read(&json_path) && serde_json::from_slice::(&data).is_ok() { table.insert(*key, data.as_slice()).unwrap(); let _ = fs::rename( &json_path, json_path.with_extension("json.migrated"), ); } } } } write_txn.commit().unwrap(); } let state = Arc::new(MemoryState { graph: Store::new("knowledge_graph_master", db.clone()), base_dir: base.clone(), search_index: RwLock::new(match crate::search::MemoryIndex::new(&base) { Ok(idx) => idx, Err(e) => { let log_path = dirs::home_dir() .unwrap_or_default() .join(".gemini/mcp_memory/daemon_error.log"); let _ = std::fs::write(&log_path, format!("Failed to create MemoryIndex: {}\n", e)); std::process::exit(1); } }), ledger: Store::new("audit_ledger", db.clone()), sticky: Store::new("sticky_notes", db.clone()), tasks: Store::new("tasks", db.clone()), snippets: Store::new("snippets", db.clone()), adrs: Store::new("adrs", db.clone()), prefs: Store::new("preferences", db.clone()), error_fixes: Store::new("error_fixes", db.clone()), pinned_files: Store::new("pinned_files", db.clone()), session_summaries: Store::new("session_summaries", db.clone()), handoff_memos: Store::new("handoff_memos", db.clone()), env_fingerprints: Store::new("env_fingerprints", db.clone()), env_requirements: Store::new("env_requirements", db.clone()), milestones: Store::new("milestones", db.clone()), environments: Store::new("environments", db.clone()), pr_checklists: Store::new("pr_checklists", db.clone()), tech_debts: Store::new("tech_debts", db.clone()), gates: Store::new("gates", db.clone()), context_workspaces: Store::new("context_workspaces", db.clone()), activity_tx: tokio::sync::broadcast::channel(100).0, }); state.rebuild_index(); run_server(state) }