refactor: Enforce strict typed JsonSchema for all MCP tools
This commit is contained in:
1 parent
686fea683d
commit
2d3aaed289
31 files changed
+2892
-308
No files matched your search
@@ -0,0 +1,823 @@
|
||||
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, ws::{WebSocketUpgrade, WebSocket, Message}},
|
||||
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()
|
||||
}
|
||||
|
||||
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: 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!("BUILD_DATE"),
|
||||
"git_hash": option_env!("GIT_HASH").unwrap_or("unknown")
|
||||
}))
|
||||
}))
|
||||
.route("/ws", get(ws_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/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();
|
||||
axum::Json(tasks.clone())
|
||||
}
|
||||
}))
|
||||
.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") {
|
||||
if let Ok(idx) = state_clone.search_index.read() {
|
||||
if let Ok(results) = idx.search(q, None) {
|
||||
let mut formatted_results = Vec::new();
|
||||
for (type_name, content) in results {
|
||||
formatted_results.push(serde_json::json!({
|
||||
"type_name": type_name,
|
||||
"content": content,
|
||||
"score": 1.0
|
||||
}));
|
||||
}
|
||||
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) = {
|
||||
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;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Background Garbage Collection for old tasks
|
||||
let state_gc = Arc::clone(&state);
|
||||
tokio::spawn(async move {
|
||||
loop {
|
||||
// Run every 24 hours
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(24 * 3600)).await;
|
||||
|
||||
let now = std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs();
|
||||
let fourteen_days = 14 * 24 * 3600;
|
||||
let cutoff = now.saturating_sub(fourteen_days);
|
||||
|
||||
state_gc.tasks.modify(|tasks| {
|
||||
let initial_len = tasks.len();
|
||||
tasks.retain(|task| {
|
||||
if task.status.to_lowercase() == "completed" && task.created_at < cutoff {
|
||||
false // remove
|
||||
} else {
|
||||
true // keep
|
||||
}
|
||||
});
|
||||
if tasks.len() < initial_len {
|
||||
eprintln!("GC: Removed {} old completed tasks", initial_len - tasks.len());
|
||||
}
|
||||
});
|
||||
}
|
||||
});
|
||||
|
||||
// Git Native Sync Background Task
|
||||
let state_git = Arc::clone(&state);
|
||||
tokio::spawn(async move {
|
||||
let repo_path = std::env::current_dir().unwrap_or_else(|_| ".".into());
|
||||
let mut last_commit_id = String::new();
|
||||
|
||||
loop {
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(30)).await;
|
||||
|
||||
if let Ok(repo) = git2::Repository::discover(&repo_path) {
|
||||
if let Ok(head) = repo.head() {
|
||||
if let Ok(commit) = head.peel_to_commit() {
|
||||
let current_id = commit.id().to_string();
|
||||
if current_id != last_commit_id && !last_commit_id.is_empty() {
|
||||
let msg = commit.message().unwrap_or("").to_string();
|
||||
let branch = head.shorthand().unwrap_or("unknown").to_string();
|
||||
|
||||
state_git.ledger.modify(|changes| {
|
||||
changes.push(crate::models::CodeChange {
|
||||
git_commit: Some(current_id.clone()),
|
||||
git_branch: Some(branch),
|
||||
description: format!("Auto-synced commit: {}", msg.trim()),
|
||||
timestamp: std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs(),
|
||||
file_path: "".to_string(),
|
||||
});
|
||||
});
|
||||
eprintln!("Git Sync: Logged new commit {}", current_id);
|
||||
|
||||
state_git.tasks.modify(|tasks| {
|
||||
for task in tasks.iter_mut() {
|
||||
if task.status != "completed" && msg.to_lowercase().contains(&task.title.to_lowercase()) {
|
||||
task.status = "completed".to_string();
|
||||
eprintln!("Git Sync: Auto-completed task '{}'", task.title);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
last_commit_id = current_id;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
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 ws_handler(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<Arc<AppState>>,
|
||||
Query(query): Query<std::collections::HashMap<String, String>>,
|
||||
) -> impl axum::response::IntoResponse {
|
||||
let client_type = query.get("client").cloned().unwrap_or_else(|| "unknown".to_string());
|
||||
ws.on_upgrade(move |socket| handle_socket(socket, state, client_type))
|
||||
}
|
||||
|
||||
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().insert(session_id.clone(), tx.clone());
|
||||
|
||||
let (mut sender, mut receiver) = socket.split();
|
||||
|
||||
let mut send_task = tokio::spawn(async move {
|
||||
while let Some(msg) = rx.recv().await {
|
||||
if sender.send(Message::Text(msg.into())).await.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
let handler = Arc::clone(&state.handler);
|
||||
let state_clone = Arc::clone(&state);
|
||||
let session_id_clone = session_id.clone();
|
||||
|
||||
let mut recv_task = tokio::spawn(async move {
|
||||
while let Some(Ok(Message::Text(text))) = receiver.next().await {
|
||||
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()) {
|
||||
if 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 clients_map = state_clone.clients.read().unwrap().clone();
|
||||
for (id, client_tx) in clients_map.iter() {
|
||||
if id != &session_id_clone {
|
||||
let _ = client_tx.send(event.to_string()).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Process MCP request
|
||||
if let Some(response) = handler.handle_request(payload).await {
|
||||
let res_str = serde_json::to_string(&response).unwrap();
|
||||
let tx_opt = state_clone.clients.read().unwrap().get(&session_id_clone).cloned();
|
||||
if let Some(client_tx) = tx_opt {
|
||||
let _ = client_tx.send(res_str).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
tokio::select! {
|
||||
_ = (&mut send_task) => recv_task.abort(),
|
||||
_ = (&mut recv_task) => send_task.abort(),
|
||||
};
|
||||
|
||||
state.clients.write().unwrap().remove(&session_id);
|
||||
}
|
||||
|
||||
async fn health_handler() -> &'static str {
|
||||
"OK"
|
||||
}
|
||||
|
||||
|
||||
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 {
|
||||
// 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(());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
// Linux no longer executes server logic natively due to workspace split
|
||||
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)
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in new issue
Block a user