848 lines
32 KiB
Rust
848 lines
32 KiB
Rust
#![cfg_attr(
|
|
not(target_os = "windows"),
|
|
allow(dead_code, unused_imports, unreachable_code)
|
|
)]
|
|
|
|
mod handlers;
|
|
mod handlers_v2;
|
|
mod mcp;
|
|
mod models;
|
|
mod router;
|
|
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 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<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 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;
|
|
}
|
|
}
|
|
}
|
|
|
|
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<MemoryHandler>,
|
|
clients: RwLock<HashMap<String, mpsc::Sender<String>>>,
|
|
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()
|
|
}
|
|
|
|
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 = 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<Arc<AppState>>| async move {
|
|
let token_path = state.handler.state.base_dir.join("admin.token");
|
|
let expected_token = tokio::fs::read_to_string(&token_path)
|
|
.await
|
|
.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<String>| 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_with(|t| t.clone());
|
|
axum::Json(tasks)
|
|
}
|
|
}),
|
|
)
|
|
.route(
|
|
"/api/sticky",
|
|
get({
|
|
let state_clone = app_state.handler.state.clone();
|
|
move || async move {
|
|
let sticky = state_clone.sticky.read_with(|s| s.clone());
|
|
axum::Json(sticky)
|
|
}
|
|
}),
|
|
)
|
|
.route(
|
|
"/api/search",
|
|
get({
|
|
let state_clone = app_state.handler.state.clone();
|
|
move |axum::extract::Query(params): axum::extract::Query<
|
|
std::collections::HashMap<String, String>,
|
|
>| 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) = state_clone.read_graph(|g| (g.entities.len(), g.relations.len()));
|
|
let tasks = state_clone.tasks.read_with(|items| items.len());
|
|
let snippets = state_clone.snippets.read_with(|items| items.len());
|
|
let tech_debts = state_clone.tech_debts.read_with(|items| items.len());
|
|
let adrs = state_clone.adrs.read_with(|items| items.len());
|
|
|
|
let ledger = state_clone.ledger.read_with(|items| items.len());
|
|
let sticky = state_clone.sticky.read_with(|items| items.len());
|
|
let error_fixes = state_clone.error_fixes.read_with(|items| items.len());
|
|
let pinned_files = state_clone.pinned_files.read_with(|items| items.len());
|
|
let session_summaries = state_clone.session_summaries.read_with(|items| items.len());
|
|
let handoff_memos = state_clone.handoff_memos.read_with(|items| items.len());
|
|
let env_fingerprints = state_clone.env_fingerprints.read_with(|items| items.len());
|
|
let env_requirements = state_clone.env_requirements.read_with(|items| items.len());
|
|
let milestones = state_clone.milestones.read_with(|items| items.len());
|
|
let environments = state_clone.environments.read_with(|items| items.len());
|
|
let pr_checklists = state_clone.pr_checklists.read_with(|items| items.len());
|
|
let gates = state_clone.gates.read_with(|items| items.len());
|
|
let context_workspaces = state_clone.context_workspaces.read_with(|items| items.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 _ =
|
|
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(())
|
|
}
|
|
|
|
async fn ws_handler(
|
|
ws: axum::extract::ws::WebSocketUpgrade,
|
|
_headers: axum::http::HeaderMap,
|
|
axum::extract::State(state): axum::extract::State<Arc<AppState>>,
|
|
axum::extract::Query(query): axum::extract::Query<std::collections::HashMap<String, String>>,
|
|
) -> 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<AppState>, client_type: String) {
|
|
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
|
|
let (tx, mut rx) = mpsc::channel::<String>(100);
|
|
|
|
state
|
|
.clients
|
|
.write()
|
|
.unwrap_or_else(|e| e.into_inner())
|
|
.insert(session_id.clone(), tx.clone());
|
|
|
|
let (mut sender, mut receiver) = socket.split();
|
|
|
|
let 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 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::<serde_json::Value>(&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 senders: Vec<_> = state_clone
|
|
.clients
|
|
.read()
|
|
.unwrap_or_else(|e| e.into_inner())
|
|
.iter()
|
|
.filter_map(|(id, tx)| {
|
|
if id != &session_id_clone {
|
|
Some(tx.clone())
|
|
} else {
|
|
None
|
|
}
|
|
})
|
|
.collect();
|
|
|
|
for client_tx in senders {
|
|
let _ = client_tx.try_send(event.to_string());
|
|
}
|
|
}
|
|
} // End if proxy
|
|
|
|
// Process MCP request
|
|
if let Some(response) = handler.handle_request(payload).await {
|
|
let res_str = serde_json::to_string(&response).unwrap_or_default();
|
|
let tx_opt = state_clone
|
|
.clients
|
|
.read()
|
|
.unwrap_or_else(|e| e.into_inner())
|
|
.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
|
|
);
|
|
});
|
|
|
|
struct SessionCleanup {
|
|
session_id: String,
|
|
state: Arc<AppState>,
|
|
send_task: Option<tokio::task::JoinHandle<()>>,
|
|
recv_task: Option<tokio::task::JoinHandle<()>>,
|
|
}
|
|
|
|
impl Drop for SessionCleanup {
|
|
fn drop(&mut self) {
|
|
self.state
|
|
.clients
|
|
.write()
|
|
.unwrap_or_else(|e| e.into_inner())
|
|
.remove(&self.session_id);
|
|
if let Some(task) = self.send_task.take() {
|
|
task.abort();
|
|
}
|
|
if let Some(task) = self.recv_task.take() {
|
|
task.abort();
|
|
}
|
|
tracing::info!(
|
|
"Websocket session {} closed and cleaned up",
|
|
self.session_id
|
|
);
|
|
}
|
|
}
|
|
|
|
let mut cleanup = SessionCleanup {
|
|
session_id: session_id.clone(),
|
|
state: Arc::clone(&state),
|
|
send_task: Some(send_task),
|
|
recv_task: Some(recv_task),
|
|
};
|
|
|
|
tokio::select! {
|
|
_ = cleanup.send_task.as_mut().unwrap() => {
|
|
tracing::info!("Websocket send task finished for session {}", session_id);
|
|
},
|
|
_ = cleanup.recv_task.as_mut().unwrap() => {
|
|
tracing::info!("Websocket recv task finished for session {}", session_id);
|
|
},
|
|
};
|
|
// Drop guard automatically handles removal and aborts the other task.
|
|
}
|
|
|
|
#[derive(serde::Deserialize, serde::Serialize, Debug)]
|
|
pub struct NvimTelemetry {
|
|
pub session_id: String,
|
|
pub event: String,
|
|
pub file: Option<String>,
|
|
pub line: Option<i64>,
|
|
pub col: Option<i64>,
|
|
}
|
|
|
|
async fn nvim_telemetry_handler(
|
|
State(state): State<Arc<AppState>>,
|
|
axum::Json(payload): axum::Json<NvimTelemetry>,
|
|
) -> 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 _ = tokio::fs::write(&win_path, &payload.session_id).await;
|
|
|
|
let wsl_path = "\\\\wsl.localhost\\Ubuntu\\home\\riz\\.gemini\\active_nvim.txt";
|
|
let _ = tokio::fs::write(wsl_path, &payload.session_id).await;
|
|
}
|
|
|
|
// 2. Broadcast to UI WebSockets
|
|
let ws_msg = serde_json::json!({
|
|
"type": "nvim_telemetry",
|
|
"data": payload
|
|
});
|
|
|
|
let msg_str = ws_msg.to_string();
|
|
let senders: Vec<_> = state
|
|
.clients
|
|
.read()
|
|
.unwrap_or_else(|e| e.into_inner())
|
|
.values()
|
|
.cloned()
|
|
.collect();
|
|
for tx in senders {
|
|
let _ = tx.try_send(msg_str.clone());
|
|
}
|
|
|
|
axum::Json(serde_json::json!({"status": "ok"}))
|
|
}
|
|
|
|
async fn health_handler() -> &'static str {
|
|
"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!("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::<serde_json::Value>(&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 rt = tokio::runtime::Runtime::new().unwrap();
|
|
let _guard = rt.enter();
|
|
|
|
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,
|
|
});
|
|
|
|
rt.block_on(run_server(state))
|
|
}
|