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

958 lines
35 KiB
Rust

#![cfg_attr(
not(target_os = "windows"),
allow(dead_code, unused_imports, unreachable_code)
)]
pub mod api;
pub mod config;
pub mod db;
pub mod embedding;
pub mod error;
pub mod handlers;
pub mod indexer;
pub mod mcp;
pub mod models;
pub mod ollama;
pub mod router;
pub mod search;
pub mod state;
pub mod store;
pub mod tools;
pub mod watcher;
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)]
pub struct Cli {
#[command(subcommand)]
pub command: Option<Commands>,
/// Target URL for the proxy to connect to (e.g., http://127.0.0.1:3000)
#[arg(long)]
pub target: Option<String>,
/// Send a shutdown request to the currently running server
#[arg(long)]
pub exit: bool,
/// Send a shutdown request to the existing server and wait for it to exit
#[arg(long)]
pub restart: bool,
}
#[derive(Subcommand)]
pub enum Commands {
/// Manage authorization gates and verification for actions
Gate {
#[command(subcommand)]
subcmd: GateCommands,
},
}
#[derive(Subcommand)]
pub 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 {
pub handler: Arc<MemoryHandler>,
pub clients: RwLock<HashMap<String, mpsc::Sender<String>>>,
pub next_id: AtomicUsize,
pub shutdown_tx: std::sync::Mutex<Option<tokio::sync::oneshot::Sender<()>>>,
}
pub async fn ttl_sweeper_worker(state: Arc<MemoryState>) {
loop {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let mut next_expiry: Option<u64> = None;
state.project.tasks.read_with(|tasks| {
for t in tasks.iter() {
if let Some(exp) = t.expires_at
&& t.is_active() {
next_expiry = Some(next_expiry.map_or(exp, |curr| curr.min(exp)));
}
}
});
state.telemetry.handoff_memos.read_with(|memos| {
for m in memos.iter() {
if let Some(exp) = m.expires_at {
next_expiry = Some(next_expiry.map_or(exp, |curr| curr.min(exp)));
}
}
});
state.telemetry.session_summaries.read_with(|summaries| {
for s in summaries.iter() {
if let Some(exp) = s.expires_at {
next_expiry = Some(next_expiry.map_or(exp, |curr| curr.min(exp)));
}
}
});
state.env.gates.read_with(|gates| {
for g in gates.iter() {
if let Some(exp) = g.expires_at {
next_expiry = Some(next_expiry.map_or(exp, |curr| curr.min(exp)));
}
}
});
let sleep_duration = match next_expiry {
Some(exp) if exp > now => {
let diff = exp - now;
std::time::Duration::from_secs(diff.min(60).max(1))
}
Some(_) => std::time::Duration::from_millis(50),
None => std::time::Duration::from_secs(60),
};
tokio::select! {
_ = state.shutdown_notify.notified() => break,
_ = state.ttl_notify.notified() => {},
_ = tokio::time::sleep(sleep_duration) => {},
}
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let mut expired_tasks = Vec::new();
state.project.tasks.modify(|tasks| {
for t in tasks.iter_mut() {
if let Some(exp) = t.expires_at
&& exp <= now && t.is_active() {
t.status = "expired".to_string();
t.updated_at = now;
expired_tasks.push(t.id.clone());
}
}
});
for tid in expired_tasks {
state.record_activity(
"task_expired",
&format!("Task {} expired due to TTL", tid),
Some("expired"),
);
state.broadcast_task_event(crate::models::TaskEvent {
task_id: tid,
status: "expired".to_string(),
action: Some("ttl_expire".to_string()),
result: Some(serde_json::json!({ "status": "expired" })),
error: None,
timestamp: now,
session_id: None,
..Default::default()
});
}
state.telemetry.handoff_memos.modify(|memos| {
memos.retain(|m| m.expires_at.is_none_or(|exp| exp > now));
});
state.telemetry.session_summaries.modify(|summaries| {
summaries.retain(|s| s.expires_at.is_none_or(|exp| exp > now));
});
state.env.gates.modify(|gates| {
gates.retain(|g| g.expires_at.is_none_or(|exp| exp > now));
});
}
}
pub async fn index_committer_worker(state: Arc<MemoryState>) {
loop {
tokio::select! {
_ = state.shutdown_notify.notified() => break,
_ = state.index_commit_notify.notified() => {},
}
let idx = state.search_index.read().await.clone();
let _ = idx.commit().await;
}
}
pub async fn condense_graph_worker(state: Arc<MemoryState>) {
loop {
tokio::select! {
_ = state.shutdown_notify.notified() => break,
_ = state.condense_notify.notified() => {},
}
let threshold: usize = std::env::var("MCP_MEMORY_CONDENSE_THRESHOLD")
.unwrap_or_else(|_| "100".to_string())
.parse()
.unwrap_or(100);
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let snippet_condensation = state.code.snippets.read_with(|snippets| {
if snippets.len() > threshold {
let mut sorted = snippets.clone();
sorted.sort_by_key(|s| s.updated_at);
let to_remove = sorted.len() - (threshold / 2);
let removed: Vec<_> = sorted.into_iter().take(to_remove).collect();
let mut content = String::new();
let mut names = Vec::new();
for r in &removed {
content.push_str(&format!(
"Name: {}\nDesc: {}\nCode: {}\n",
r.name, r.description, r.code
));
names.push(r.name.clone());
}
Some((content, names))
} else {
None
}
});
if let Some((content, names)) = snippet_condensation
&& !content.is_empty() {
let name = format!("Snippet History {}", now);
state.modify_graph(|graph| {
graph.entities.insert(
name.clone(),
crate::models::Entity {
name: name.clone(),
entity_type: "Historical Summary".to_string(),
observations: vec![content],
namespace: crate::models::default_namespace(),
git_branch: None,
..Default::default()
},
);
});
let name_set: std::collections::HashSet<String> = names.into_iter().collect();
state.code.snippets.modify(|snippets| {
snippets.retain(|s| !name_set.contains(&s.name));
});
tracing::info!("Condensed snippets into Historical Summary.");
}
}
}
pub async fn memory_consolidation_worker(state: Arc<MemoryState>) {
let mut interval = tokio::time::interval(std::time::Duration::from_secs(300));
loop {
tokio::select! {
_ = state.shutdown_notify.notified() => break,
_ = interval.tick() => {},
}
let entities: Vec<_> = state.graph.read_with(|g| {
g.entities
.values()
.map(|e| (e.name.clone(), e.entity_type.clone()))
.collect()
});
if entities.len() < 2 {
continue;
}
let mut entity_summaries = String::new();
for (name, e_type) in entities.iter().take(50) {
entity_summaries.push_str(&format!("- [{}] {}\n", e_type, name));
}
let prompt = format!(
"Analyze the following list of entities and identify exactly TWO that represent the exact same concept or item but have slightly different names (e.g. 'auth_service' and 'AuthService'). Return ONLY a valid JSON array containing exactly two strings: the two names to merge. If no obvious duplicates exist, return an empty array []. Do not output any markdown formatting or extra text.\n\nEntities:\n{}",
entity_summaries
);
if let Ok(response) = state
.ollama
.generate(
&prompt,
None,
Some("You are a helpful JSON-only data deduplication assistant. Output only JSON."),
)
.await
{
let cleaned = response
.trim()
.trim_start_matches("```json")
.trim_start_matches("```")
.trim_end_matches("```")
.trim();
if let Ok(duplicates) = serde_json::from_str::<Vec<String>>(cleaned)
&& duplicates.len() == 2 {
let e1_name = &duplicates[0];
let e2_name = &duplicates[1];
if e1_name != e2_name {
tracing::info!(
"Memory Consolidation Daemon: Merging '{}' into '{}'",
e2_name,
e1_name
);
state.modify_graph(|g| {
if let Some(mut e2) = g.entities.remove(e2_name) {
if let Some(e1) = g.entities.get_mut(e1_name) {
e1.observations.append(&mut e2.observations);
} else {
g.entities.insert(e2_name.clone(), e2);
}
}
for rel in g.relations.iter_mut() {
if rel.from == *e2_name {
rel.from = e1_name.clone();
}
if rel.to == *e2_name {
rel.to = e1_name.clone();
}
}
});
}
}
}
}
}
pub async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>> {
let state_for_index = Arc::clone(&state);
tokio::spawn(async move {
state_for_index.rebuild_index().await;
tracing::info!("Index rebuild complete.");
});
crate::indexer::start_background_indexer(Arc::clone(&state)).await;
tokio::spawn(index_committer_worker(Arc::clone(&state)));
tokio::spawn(ttl_sweeper_worker(Arc::clone(&state)));
tokio::spawn(condense_graph_worker(Arc::clone(&state)));
tokio::spawn(memory_consolidation_worker(Arc::clone(&state)));
crate::watcher::spawn_watcher(Arc::clone(&state));
crate::handlers::vision::spawn_clipboard_listener(Arc::clone(&state));
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel();
let app_state = Arc::new(AppState {
handler: Arc::new(MemoryHandler::new(Arc::clone(&state))),
clients: RwLock::new(HashMap::new()),
next_id: AtomicUsize::new(1),
shutdown_tx: std::sync::Mutex::new(Some(shutdown_tx)),
});
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 mut closed_ids = Vec::new();
{
let clients_guard = app_state_clone
.clients
.read()
.unwrap_or_else(|e| e.into_inner());
for (id, tx) in clients_guard.iter() {
if tx.try_send(msg.clone()).is_err() && tx.is_closed() {
closed_ids.push(id.clone());
}
}
}
if !closed_ids.is_empty() {
let mut write_guard = app_state_clone
.clients
.write()
.unwrap_or_else(|e| e.into_inner());
for id in closed_ids {
write_guard.remove(&id);
}
}
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue,
}
}
});
let udp_state = Arc::clone(&app_state);
tokio::spawn(async move {
let port1 = std::env::var("MCP_UDP_PORT1").unwrap_or_else(|_| "3001".to_string());
if let Ok(socket) = tokio::net::UdpSocket::bind(format!("127.0.0.1:{}", port1)).await {
let socket = Arc::new(socket);
let socket_rx = socket.clone();
let mut subscribers: HashMap<(String, String), std::net::SocketAddr> = HashMap::new();
let mut buf = vec![0u8; 65536];
let mut event_rx = udp_state.handler.state.event_bus_tx.subscribe();
loop {
tokio::select! {
recv_res = socket_rx.recv_from(&mut buf) => {
if let Ok((len, addr)) = recv_res {
if let Ok(payload) = serde_json::from_slice::<crate::models::TerminalHistory>(&buf[..len]) {
udp_state
.handler
.state
.record_terminal_history(payload.clone());
let ws_msg = serde_json::json!({
"type": "terminal_telemetry",
"data": payload
});
let msg_str = ws_msg.to_string();
let senders: Vec<_> = udp_state
.clients
.read()
.unwrap_or_else(|e| e.into_inner())
.values()
.cloned()
.collect();
for tx in senders {
let _ = tx.try_send(msg_str.clone());
}
} else if let Ok(json_payload) = serde_json::from_slice::<serde_json::Value>(&buf[..len]) {
if json_payload.get("type").and_then(|t| t.as_str()) == Some("ping") {
let _ = socket.send_to(b"pong", addr).await;
} else if json_payload.get("type").and_then(|t| t.as_str()) == Some("gate_wait")
&& let (Some(action), Some(target)) = (
json_payload.get("action").and_then(|a| a.as_str()),
json_payload.get("target").and_then(|t| t.as_str())
) {
subscribers.insert((action.to_string(), target.to_string()), addr);
}
}
}
}
Ok(event) = event_rx.recv() => {
if event.topic == "gate:event"
&& let (Some(action), Some(target), Some(status)) = (
event.payload.get("action").and_then(|a| a.as_str()),
event.payload.get("target").and_then(|t| t.as_str()),
event.payload.get("status").and_then(|s| s.as_str()),
) {
if (status == "authorized" || status == "blocked")
&& let Some(addr) = subscribers.remove(&(action.to_string(), target.to_string())) {
let response = if status == "authorized" { b"APPROVED" } else { b"REJECTED" };
let _ = socket.send_to(response, addr).await;
}
}
}
}
}
}
});
let nvim_udp_state = Arc::clone(&app_state);
tokio::spawn(async move {
let port2 = std::env::var("MCP_UDP_PORT2").unwrap_or_else(|_| "3002".to_string());
if let Ok(socket) = tokio::net::UdpSocket::bind(format!("127.0.0.1:{}", port2)).await {
let mut buf = vec![0u8; 65536];
loop {
if let Ok((len, _addr)) = socket.recv_from(&mut buf).await
&& let Ok(payload) =
serde_json::from_slice::<crate::api::telemetry::NvimTelemetry>(&buf[..len])
{
if payload.event == "FocusGained"
|| payload.event == "BufEnter"
|| payload.event == "VimEnter"
{
let session = &payload.session_id;
let _is_unix_socket = session.starts_with('/') || session.starts_with('~');
if let Some(home) = dirs::home_dir() {
let nvim_dir = home.join(".gemini");
let nvim_txt_path = nvim_dir.join("active_nvim.txt");
let tmp_path =
nvim_dir.join(format!("active_nvim_{}.tmp", std::process::id()));
if tokio::fs::create_dir_all(&nvim_dir).await.is_ok()
&& tokio::fs::write(&tmp_path, session).await.is_ok()
{
let _ = tokio::fs::rename(&tmp_path, &nvim_txt_path).await;
}
}
}
let (tech_debts, adrs) = if let Some(ref f) = payload.file {
crate::api::telemetry::find_projected_knowledge(
&nvim_udp_state.handler.state,
f,
)
} else {
(Vec::new(), Vec::new())
};
let ws_msg = serde_json::json!({
"type": "nvim_telemetry",
"data": payload,
"tech_debts": tech_debts,
"adrs": adrs
});
let msg_str = ws_msg.to_string();
let senders: Vec<_> = nvim_udp_state
.clients
.read()
.unwrap_or_else(|e| e.into_inner())
.values()
.cloned()
.collect();
for tx in senders {
let _ = tx.try_send(msg_str.clone());
}
if payload.event == "BufWritePost"
&& let Some(ref file_path) = payload.file
{
let normalized_file = file_path.replace("\\", "/");
let topic = format!("nvim:save:{}", normalized_file);
let event = crate::state::GenericEvent {
topic,
session_id: Some(payload.session_id.clone()),
payload: serde_json::json!(&payload),
};
let _ = nvim_udp_state.handler.state.event_bus_tx.send(event);
}
if payload.event.starts_with("agent_") || payload.event.starts_with("diff_") {
let payload_val = serde_json::json!(&payload);
// 1. General event topic (e.g. nvim:ui:agent_prompt_response, nvim:ui:agent_diff_accepted)
let _ = nvim_udp_state.handler.state.event_bus_tx.send(
crate::state::GenericEvent {
topic: format!("nvim:ui:{}", payload.event),
session_id: Some(payload.session_id.clone()),
payload: payload_val.clone(),
},
);
// 2. Correlated request_id topic (e.g. nvim:ui:agent_prompt_response:REQ_ID)
if let Some(ref req_id) = payload.request_id {
let _ = nvim_udp_state.handler.state.event_bus_tx.send(
crate::state::GenericEvent {
topic: format!("nvim:ui:{}:{}", payload.event, req_id),
session_id: Some(payload.session_id.clone()),
payload: payload_val.clone(),
},
);
}
// 3. Correlated diff_id topics
if let Some(ref diff_id) = payload.diff_id {
let _ = nvim_udp_state.handler.state.event_bus_tx.send(
crate::state::GenericEvent {
topic: format!("nvim:ui:{}:{}", payload.event, diff_id),
session_id: Some(payload.session_id.clone()),
payload: payload_val.clone(),
},
);
// General diff decision topic
let _ = nvim_udp_state.handler.state.event_bus_tx.send(
crate::state::GenericEvent {
topic: format!("nvim:ui:diff_decision:{}", diff_id),
session_id: Some(payload.session_id.clone()),
payload: payload_val.clone(),
},
);
}
}
}
}
}
});
let app = api::setup::create_router(app_state);
let port_str = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string());
tracing::info!(
"MCP Memory Server running on http://127.0.0.1:{}/ws",
port_str
);
let addr: std::net::SocketAddr = format!("127.0.0.1:{}", port_str)
.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())
.with_graceful_shutdown(async move {
let _ = shutdown_rx.await;
})
.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;
}
tracing::info!("axum::serve graceful shutdown complete.");
Ok(())
}
pub 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)
}
pub fn run_cli() -> Result<(), Box<dyn std::error::Error>> {
crate::config::load_mcp_config_env();
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 rt = tokio::runtime::Runtime::new()?;
rt.block_on(async {
let client = reqwest::Client::builder().build().unwrap_or_default();
let mut req = client.post(format!("http://127.0.0.1:{}/shutdown", port));
if !token.trim().is_empty() {
req = req.header("Authorization", format!("Bearer {}", token.trim()));
}
let _ = req.send().await;
});
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(())
}
#[cfg(test)]
mod tests {
use super::*;
use clap::Parser;
#[test]
fn test_cli_parsing_default() {
let cli = Cli::try_parse_from(&["mcp-memory-server"]).unwrap();
assert!(cli.command.is_none());
assert!(!cli.exit);
assert!(!cli.restart);
}
#[test]
fn test_cli_parsing_exit_and_target() {
let cli = Cli::try_parse_from(&[
"mcp-memory-server",
"--exit",
"--target",
"http://localhost:3000",
])
.unwrap();
assert!(cli.exit);
assert_eq!(cli.target.as_deref(), Some("http://localhost:3000"));
}
#[test]
fn test_cli_parsing_gate_set() {
let cli = Cli::try_parse_from(&[
"mcp-memory-server",
"gate",
"set",
"--action",
"git_push",
"--target",
"master",
"--authorize",
"--reason",
"Approved by lead",
])
.unwrap();
if let Some(Commands::Gate {
subcmd:
GateCommands::Set {
action,
target,
authorize,
reason,
..
},
}) = cli.command
{
assert_eq!(action, "git_push");
assert_eq!(target, "master");
assert!(authorize);
assert_eq!(reason.as_deref(), Some("Approved by lead"));
} else {
panic!("Expected Gate Set subcommand");
}
}
#[test]
fn test_cli_parsing_gate_verify() {
let cli = Cli::try_parse_from(&[
"mcp-memory-server",
"gate",
"verify",
"--action",
"deploy",
"--target",
"prod",
"--consume",
])
.unwrap();
if let Some(Commands::Gate {
subcmd:
GateCommands::Verify {
action,
target,
consume,
..
},
}) = cli.command
{
assert_eq!(action, "deploy");
assert_eq!(target, "prod");
assert!(consume);
} else {
panic!("Expected Gate Verify subcommand");
}
}
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn test_init_logging_helper() {
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let temp_dir = tempfile::tempdir().unwrap();
unsafe {
std::env::set_var("MCP_MEMORY_STORE_DIR", temp_dir.path().to_str().unwrap());
}
let guard = init_logging("test_app");
assert!(guard.is_some());
}
#[tokio::test]
async fn test_background_workers_one_tick() {
let temp_dir = tempfile::tempdir().unwrap();
let state = Arc::new(MemoryState::new(temp_dir.path().to_str().unwrap()));
// Test worker functions by spawning them briefly
let handle1 = tokio::spawn(ttl_sweeper_worker(state.clone()));
let handle2 = tokio::spawn(index_committer_worker(state.clone()));
let handle3 = tokio::spawn(condense_graph_worker(state.clone()));
state.ttl_notify.notify_one();
state.index_commit_notify.notify_one();
state.condense_notify.notify_one();
tokio::time::sleep(Duration::from_millis(50)).await;
handle1.abort();
handle2.abort();
handle3.abort();
}
#[tokio::test]
async fn test_run_server_graceful_shutdown() {
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let temp_dir = tempfile::tempdir().unwrap();
let state = Arc::new(MemoryState::new(temp_dir.path().to_str().unwrap()));
// Bind to a free port to avoid conflicts
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
drop(listener);
unsafe {
std::env::set_var("MCP_PORT", port.to_string());
std::env::set_var("MCP_UDP_PORT1", (port + 1).to_string());
std::env::set_var("MCP_UDP_PORT2", (port + 2).to_string());
}
let server_handle = tokio::spawn(async move {
let _ = run_server(state).await;
});
tokio::time::sleep(Duration::from_millis(300)).await;
let _ = reqwest::Client::new()
.get(format!("http://127.0.0.1:{}/ping", port))
.send()
.await;
server_handle.abort();
}
}