385 lines
14 KiB
Rust
385 lines
14 KiB
Rust
mod handlers;
|
|
mod models;
|
|
mod mcp;
|
|
mod state;
|
|
mod store;
|
|
mod tools;
|
|
mod search;
|
|
|
|
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 tokio::time::sleep;
|
|
|
|
use clap::{Parser, Subcommand};
|
|
use std::collections::HashMap;
|
|
|
|
#[derive(Parser)]
|
|
#[command(author, version, about, long_about = None)]
|
|
struct Cli {
|
|
#[command(subcommand)]
|
|
command: Option<Commands>,
|
|
#[arg(long)]
|
|
target: Option<String>,
|
|
#[arg(long)]
|
|
daemon: bool,
|
|
}
|
|
|
|
#[derive(Subcommand)]
|
|
enum Commands {
|
|
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| {});
|
|
state.rebuild_index();
|
|
}
|
|
|
|
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::{
|
|
extract::{State, Query},
|
|
response::sse::{Event, Sse},
|
|
routing::{get, post},
|
|
Json, Router,
|
|
};
|
|
use futures_util::stream::Stream;
|
|
use std::convert::Infallible;
|
|
use tokio::sync::mpsc;
|
|
use tokio_stream::wrappers::ReceiverStream;
|
|
use std::sync::atomic::{AtomicUsize, Ordering};
|
|
|
|
struct AppState {
|
|
handler: Arc<MemoryHandler>,
|
|
clients: RwLock<HashMap<String, mpsc::Sender<Result<Event, Infallible>>>>,
|
|
next_id: AtomicUsize,
|
|
}
|
|
|
|
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("/sse", get(sse_handler))
|
|
.route("/messages", post(message_handler))
|
|
.route("/health", get(health_handler))
|
|
.route("/dashboard", 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();
|
|
|
|
axum::Json(serde_json::json!({
|
|
"entities": entities,
|
|
"relations": relations,
|
|
"tasks": tasks,
|
|
"snippets": snippets,
|
|
"tech_debts": tech_debts,
|
|
"adrs": adrs
|
|
}))
|
|
}
|
|
}))
|
|
.with_state(app_state);
|
|
|
|
let listener = tokio::net::TcpListener::bind("0.0.0.0:3000").await.unwrap();
|
|
eprintln!("MCP Memory Server running on http://0.0.0.0:3000/sse");
|
|
axum::serve(listener, app).await.unwrap();
|
|
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();
|
|
|
|
#[cfg(target_os = "windows")]
|
|
{
|
|
use std::os::windows::process::CommandExt;
|
|
if !cli.daemon {
|
|
loop {
|
|
if let Err(_) = std::net::TcpListener::bind("0.0.0.0:3000") {
|
|
// 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
|
|
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 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().unwrap()),
|
|
ledger: Store::new(base.join("audit_ledger.json")),
|
|
sticky: Store::new(base.join("sticky_notes.json")),
|
|
tasks: Store::new(base.join("tasks.json")),
|
|
snippets: Store::new(base.join("snippets.json")),
|
|
adrs: Store::new(base.join("adrs.json")),
|
|
prefs: Store::new(base.join("preferences.json")),
|
|
error_fixes: Store::new(base.join("error_fixes.json")),
|
|
pinned_files: Store::new(base.join("pinned_files.json")),
|
|
session_summaries: Store::new(base.join("session_summaries.json")),
|
|
handoff_memos: Store::new(base.join("handoff_memos.json")),
|
|
env_fingerprints: Store::new(base.join("env_fingerprints.json")),
|
|
env_requirements: Store::new(base.join("env_requirements.json")),
|
|
milestones: Store::new(base.join("milestones.json")),
|
|
environments: Store::new(base.join("environments.json")),
|
|
pr_checklists: Store::new(base.join("pr_checklists.json")),
|
|
tech_debts: Store::new(base.join("tech_debts.json")),
|
|
gates: Store::new(base.join("gates.json")),
|
|
context_workspaces: Store::new(base.join("context_workspaces.json")),
|
|
});
|
|
|
|
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)
|
|
}
|