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

701 lines
26 KiB
Rust

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 std::fs;
use std::path::PathBuf;
use std::sync::{Arc, RwLock};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use redb::ReadableTable;
use tokio::time::sleep;
use clap::{Parser, Subcommand};
use std::collections::HashMap;
#[derive(Parser)]
#[command(author, version, about = "Antigravity MCP Memory Server", long_about = None)]
struct Cli {
#[command(subcommand)]
command: Option<Commands>,
/// Target URL for the proxy to connect to (e.g., http://127.0.0.1:3000)
#[arg(long)]
target: Option<String>,
/// 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<String>,
#[arg(short = 'p', long = "param")]
params: Vec<String>,
#[arg(long, conflicts_with = "block")]
authorize: bool,
#[arg(long, conflicts_with = "authorize")]
block: bool,
#[arg(long)]
reason: Option<String>,
},
Verify {
#[arg(long)]
action: String,
#[arg(long)]
target: String,
#[arg(long)]
namespace: Option<String>,
#[arg(short = 'p', long = "param")]
params: Vec<String>,
#[arg(long)]
consume: bool,
},
}
async fn reconcile_worker(state: Arc<MemoryState>) {
loop {
sleep(Duration::from_secs(5)).await;
let pattern = format!("{}/delta_*.json", state.base_dir.display());
let has_local = {
let session = state.session_graph.read().unwrap();
!session.entities.is_empty() || !session.relations.is_empty()
};
let has_files = glob::glob(&pattern).map(|p| p.count() > 0).unwrap_or(false);
if has_local || has_files {
state.apply_sync_write(|_master| {}).await;
let state_clone = state.clone();
let _ = tokio::task::spawn_blocking(move || {
state_clone.rebuild_index();
}).await;
}
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
state.ledger.modify(|ledger| {
let seven_days = now.saturating_sub(7 * 24 * 60 * 60);
ledger.retain(|c| c.timestamp >= seven_days);
if ledger.len() > 1000 {
let excess = ledger.len() - 1000;
ledger.drain(0..excess);
}
});
state.sticky.modify(|notes| {
notes.retain(|note| note.timestamp >= now.saturating_sub(24 * 60 * 60));
});
}
}
use axum::{
Json, Router,
extract::{Query, State},
response::sse::{Event, Sse},
response::IntoResponse,
routing::{get, post},
};
use futures_util::stream::Stream;
use std::convert::Infallible;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
struct AppState {
handler: Arc<MemoryHandler>,
clients: RwLock<HashMap<String, mpsc::Sender<Result<Event, Infallible>>>>,
next_id: AtomicUsize,
}
#[derive(serde::Deserialize)]
struct GateVerifyReq {
action: String,
target: String,
namespace: Option<String>,
#[serde(default)]
params: HashMap<String, String>,
#[serde(default)]
consume: bool,
}
#[derive(serde::Deserialize)]
struct GateSetReq {
action: String,
target: String,
namespace: Option<String>,
#[serde(default)]
params: HashMap<String, String>,
authorize: Option<bool>,
block: Option<bool>,
reason: Option<String>,
}
async fn gate_verify_handler(
State(app_state): State<Arc<AppState>>,
Query(q): Query<GateVerifyReq>,
) -> 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<Arc<AppState>>,
Json(body): Json<GateSetReq>,
) -> 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<MemoryState>) -> Result<(), Box<dyn std::error::Error>> {
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
tokio::spawn(reconcile_worker(Arc::clone(&state)));
let app_state = Arc::new(AppState {
handler: Arc::new(MemoryHandler { 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!("CARGO_PKG_VERSION"),
"git_hash": option_env!("GIT_HASH").unwrap_or("unknown")
}))
}))
.route("/sse", get(sse_handler))
.route("/messages", post(message_handler))
.route("/health", get(health_handler))
.route("/gate/verify", get(gate_verify_handler))
.route("/gate/set", post(gate_set_handler))
.route(
"/shutdown",
post(|| async move {
std::thread::spawn(|| {
std::thread::sleep(std::time::Duration::from_millis(100));
std::process::exit(0);
});
"Shutting down..."
}),
)
.route(
"/",
get(|| async move { axum::response::Html(include_str!("dashboard.html")) }),
)
.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);
let mut retries = 0;
let listener = loop {
match tokio::net::TcpListener::bind("127.0.0.1:3000").await {
Ok(l) => break l,
Err(e) => {
// Check if it's already running and healthy
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::{Read, Write};
let _ = stream.write_all(
b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
let mut response = String::new();
let _ = stream.read_to_string(&mut response);
if response.contains("200 OK") {
// Already healthy! Just exit cleanly instead of panicking/retrying loop.
std::process::exit(0);
}
}
retries += 1;
if retries > 15 {
let log_path = dirs::home_dir()
.unwrap_or_default()
.join(".gemini/mcp_memory/daemon_fatal.log");
let _ = std::fs::write(
&log_path,
format!(
"FATAL: Could not bind to 127.0.0.1:3000 after 15 seconds: {}\n",
e
),
);
std::process::exit(1);
}
let log_path = dirs::home_dir()
.unwrap_or_default()
.join(".gemini/mcp_memory/daemon_error.log");
if let Ok(mut file) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&log_path)
{
use std::io::Write;
let _ = writeln!(
file,
"Failed to bind to 127.0.0.1:3000 (attempt {}): {}. Retrying in 1s...",
retries, e
);
}
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
}
}
};
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()
.join(".gemini/mcp_memory/daemon_error.log");
let _ = std::fs::write(&log_path, format!("Server crashed: {}\n", e));
}
Ok(())
})
}
async fn sse_handler(
State(state): State<Arc<AppState>>,
) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
let (tx, rx) = mpsc::channel::<Result<Event, Infallible>>(100);
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 stream = ReceiverStream::new(rx);
Sse::new(stream).keep_alive(axum::response::sse::KeepAlive::new())
}
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<Arc<AppState>>,
Query(query): Query<SessionQuery>,
Json(payload): Json<serde_json::Value>,
) -> axum::http::StatusCode {
let handler = Arc::clone(&state.handler);
let session_id = query.session_id.clone();
let clients = Arc::clone(&state);
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<dyn std::error::Error>> {
let cli = Cli::parse();
if cli.exit {
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::Write;
let _ = stream.write_all(
b"POST /shutdown HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
}
println!("Sent shutdown request to server.");
return Ok(());
}
if cli.restart {
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::Write;
let _ = stream.write_all(
b"POST /shutdown HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
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 {
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));
}
}
}
}
#[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);
return Ok(());
}
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);
fs::create_dir_all(&base).expect("Failed to create store dir");
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 = [
("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() {
if let Ok(data) = fs::read(&json_path) {
if serde_json::from_slice::<serde_json::Value>(&data).is_ok() {
table.insert(*key, data.as_slice()).unwrap();
}
}
}
}
}
}
write_txn.commit().unwrap();
}
let state = Arc::new(MemoryState {
master_path: base.join("knowledge_graph_master.json"),
session_graph: RwLock::new(KnowledgeGraph::default()),
base_dir: base.clone(),
master_cache: RwLock::new((KnowledgeGraph::default(), SystemTime::UNIX_EPOCH)),
search_index: RwLock::new(crate::search::MemoryIndex::new(&base).unwrap()),
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()),
});
state.rebuild_index();
if let Some(command) = cli.command {
match command {
Commands::Gate { subcmd } => match subcmd {
GateCommands::Set {
action,
target,
namespace,
params,
authorize,
block,
reason,
} => {
let status = if authorize {
"authorized".to_string()
} else if block {
"blocked".to_string()
} else {
"pending".to_string()
};
let mut param_map = HashMap::new();
for p in params {
if let Some((k, v)) = p.split_once('=') {
param_map.insert(k.to_string(), v.to_string());
}
}
let record = GateRecord {
id: uuid::Uuid::new_v4().to_string(),
action: action.clone(),
target: target.clone(),
namespace,
params: param_map,
status,
reason,
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs(),
};
state.gates.modify(|gates| {
gates.retain(|g| !(g.action == record.action && g.target == record.target));
gates.push(record);
});
println!("Gate state updated.");
std::process::exit(0);
}
GateCommands::Verify {
action,
target,
namespace,
params,
consume,
} => {
let mut param_map = HashMap::new();
for p in params {
if let Some((k, v)) = p.split_once('=') {
param_map.insert(k.to_string(), v.to_string());
}
}
let mut found = None;
let mut to_remove = None;
state.gates.modify(|gates| {
if let Some(idx) = gates.iter().position(|g| {
g.action == action
&& g.target == target
&& g.namespace == namespace
&& g.params == param_map
}) {
found = Some(gates[idx].clone());
if consume {
to_remove = Some(idx);
}
}
if let Some(idx) = to_remove {
gates.remove(idx);
}
});
match found {
Some(record) => {
if record.status == "authorized" {
std::process::exit(0);
} else {
if let Some(r) = record.reason {
eprintln!("❌ Action blocked. Reason: {}", r);
} else {
eprintln!("❌ Action blocked.");
}
std::process::exit(1);
}
}
None => {
eprintln!("❌ Action not yet authorized (no gate record found).");
std::process::exit(2);
}
}
}
},
}
}
run_server(state)
}