- Optimized memory allocation in router.rs by offloading JSON serialization to spawn_blocking and using references. - Prevented full graph duplication on startup in state.rs index rebuild. - Eliminated massive String allocations in QueryGraphPathHandler BFS loops. - Avoided temporary Strings in VisualizeGraphHandler via inline writing. - Fixed O(N) full-graph deduplication in MergeEntitiesHandler to scale efficiently.
324 lines
11 KiB
Rust
324 lines
11 KiB
Rust
#![cfg_attr(
|
|
not(target_os = "windows"),
|
|
allow(dead_code, unused_imports, unreachable_code)
|
|
)]
|
|
|
|
mod api;
|
|
pub mod db;
|
|
pub mod error;
|
|
mod handlers;
|
|
mod mcp;
|
|
mod models;
|
|
mod router;
|
|
mod search;
|
|
mod state;
|
|
mod store;
|
|
mod tools;
|
|
|
|
use crate::api::rest::GateSetReq;
|
|
use crate::router::MemoryHandler;
|
|
use crate::state::MemoryState;
|
|
use clap::{Parser, Subcommand};
|
|
use std::collections::HashMap;
|
|
use std::path::PathBuf;
|
|
use std::sync::atomic::AtomicUsize;
|
|
use std::sync::{Arc, RwLock};
|
|
use std::time::Duration;
|
|
use tokio::sync::mpsc;
|
|
|
|
#[derive(Parser)]
|
|
#[command(author, version = env!("APP_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>,
|
|
/// 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,
|
|
},
|
|
}
|
|
|
|
pub struct AppState {
|
|
handler: Arc<MemoryHandler>,
|
|
clients: RwLock<HashMap<String, mpsc::Sender<String>>>,
|
|
next_id: AtomicUsize,
|
|
}
|
|
|
|
async fn index_committer_worker(state: Arc<MemoryState>) {
|
|
loop {
|
|
tokio::time::sleep(Duration::from_secs(5)).await;
|
|
// Periodically commit the search index to persist inline indexing operations
|
|
let idx_opt = state.search_index.read().ok().map(|idx| idx.clone());
|
|
if let Some(idx) = idx_opt {
|
|
let _ = idx.commit().await;
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>> {
|
|
state.rebuild_index().await;
|
|
tokio::spawn(index_committer_worker(Arc::clone(&state)));
|
|
let app_state = Arc::new(AppState {
|
|
handler: Arc::new(MemoryHandler::new(Arc::clone(&state))),
|
|
clients: RwLock::new(HashMap::new()),
|
|
next_id: AtomicUsize::new(1),
|
|
});
|
|
|
|
let app_state_clone = Arc::clone(&app_state);
|
|
let mut rx = state.activity_tx.subscribe();
|
|
tokio::spawn(async move {
|
|
loop {
|
|
match rx.recv().await {
|
|
Ok(msg) => {
|
|
let senders: Vec<_> = app_state_clone
|
|
.clients
|
|
.read()
|
|
.unwrap_or_else(|e| e.into_inner())
|
|
.values()
|
|
.cloned()
|
|
.collect();
|
|
for client_tx in senders {
|
|
let _ = client_tx.try_send(msg.clone());
|
|
}
|
|
}
|
|
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
|
|
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue,
|
|
}
|
|
}
|
|
});
|
|
|
|
let app = api::setup::create_router(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()
|
|
.expect("Invalid bind address");
|
|
|
|
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 _ =
|
|
tokio::fs::write(&log_path, format!("Failed to bind to {}: {}\n", addr, e)).await;
|
|
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 _ = tokio::fs::write(&log_path, format!("Server crashed: {}\n", e)).await;
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
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 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<dyn std::error::Error>> {
|
|
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!("http://127.0.0.1:{}/shutdown", port))
|
|
.output();
|
|
|
|
if cli.restart {
|
|
std::thread::sleep(Duration::from_secs(2));
|
|
} else {
|
|
return Ok(());
|
|
}
|
|
}
|
|
|
|
if let Some(Commands::Gate { subcmd }) = cli.command {
|
|
let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string());
|
|
let rt = tokio::runtime::Runtime::new()?;
|
|
match subcmd {
|
|
GateCommands::Set {
|
|
action,
|
|
target,
|
|
namespace,
|
|
params,
|
|
authorize,
|
|
block,
|
|
reason,
|
|
} => {
|
|
let mut pmap = HashMap::new();
|
|
for p in params {
|
|
if let Some((k, v)) = p.split_once('=') {
|
|
pmap.insert(k.to_string(), v.to_string());
|
|
}
|
|
}
|
|
let req = GateSetReq {
|
|
action,
|
|
target,
|
|
namespace,
|
|
params: pmap,
|
|
authorize: if authorize { Some(true) } else { None },
|
|
block: if block { Some(true) } else { None },
|
|
reason,
|
|
};
|
|
rt.block_on(async {
|
|
let client = reqwest::Client::new();
|
|
let res = client
|
|
.post(format!("http://127.0.0.1:{}/gate/set", port))
|
|
.json(&req)
|
|
.send()
|
|
.await;
|
|
match res {
|
|
Ok(r) if r.status().is_success() => println!("Gate updated successfully"),
|
|
Ok(r) => println!("Failed to update gate: {}", r.status()),
|
|
Err(e) => println!("Error connecting to server: {}", e),
|
|
}
|
|
});
|
|
}
|
|
GateCommands::Verify {
|
|
action,
|
|
target,
|
|
namespace,
|
|
params: _,
|
|
consume,
|
|
} => {
|
|
let mut url = format!(
|
|
"http://127.0.0.1:{}/gate/verify?action={}&target={}&consume={}",
|
|
port, action, target, consume
|
|
);
|
|
if let Some(ns) = namespace {
|
|
url.push_str(&format!("&namespace={}", ns));
|
|
}
|
|
rt.block_on(async {
|
|
let res = reqwest::get(&url).await;
|
|
match res {
|
|
Ok(r) if r.status().is_success() => std::process::exit(0),
|
|
Ok(r) if r.status() == reqwest::StatusCode::FORBIDDEN => {
|
|
let text = r.text().await.unwrap_or_default();
|
|
eprintln!("{}", text);
|
|
std::process::exit(1);
|
|
}
|
|
Ok(r) if r.status() == reqwest::StatusCode::NOT_FOUND => {
|
|
eprintln!("Action not yet authorized.");
|
|
std::process::exit(2);
|
|
}
|
|
Ok(r) => {
|
|
eprintln!("Unexpected status: {}", r.status());
|
|
std::process::exit(3);
|
|
}
|
|
Err(e) => {
|
|
eprintln!("Error connecting to server: {}", e);
|
|
std::process::exit(4);
|
|
}
|
|
}
|
|
});
|
|
}
|
|
}
|
|
return Ok(());
|
|
}
|
|
|
|
let token = uuid::Uuid::new_v4().to_string();
|
|
std::fs::write(base.join("admin.token"), &token).unwrap_or_default();
|
|
|
|
let rt = tokio::runtime::Builder::new_multi_thread()
|
|
.enable_all()
|
|
.build()
|
|
.unwrap();
|
|
|
|
rt.block_on(async {
|
|
let state = Arc::new(MemoryState::new(&base.to_string_lossy()));
|
|
if let Err(e) = run_server(state).await {
|
|
tracing::error!("Server error: {}", e);
|
|
}
|
|
});
|
|
|
|
Ok(())
|
|
}
|