diff --git a/server/src/main.rs b/server/src/main.rs index 61ae35e..bf29c44 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -15,28 +15,34 @@ 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, long_about = None)] +#[command(author, version, about = "Antigravity MCP Memory Server", long_about = None)] struct Cli { #[command(subcommand)] command: Option, + /// Target URL for the proxy to connect to (e.g., http://127.0.0.1:3000) #[arg(long)] target: Option, + /// 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, @@ -111,6 +117,7 @@ use axum::{ Json, Router, extract::{Query, State}, response::sse::{Event, Sse}, + response::IntoResponse, routing::{get, post}, }; use futures_util::stream::Stream; @@ -125,6 +132,103 @@ struct AppState { next_id: AtomicUsize, } +#[derive(serde::Deserialize)] +struct GateVerifyReq { + action: String, + target: String, + namespace: Option, + #[serde(default)] + params: HashMap, + #[serde(default)] + consume: bool, +} + +#[derive(serde::Deserialize)] +struct GateSetReq { + action: String, + target: String, + namespace: Option, + #[serde(default)] + params: HashMap, + authorize: Option, + block: Option, + reason: Option, +} + +async fn gate_verify_handler( + State(app_state): State>, + Query(q): Query, +) -> 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>, + Json(body): Json, +) -> 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) -> Result<(), Box> { let rt = tokio::runtime::Runtime::new().unwrap(); rt.block_on(async { @@ -139,6 +243,8 @@ fn run_server(state: Arc) -> Result<(), Box> .route("/sse", get(sse_handler)) .route("/messages", post(message_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 { @@ -410,10 +516,46 @@ fn main() -> Result<(), Box> { let redb_path = base.join("mcp_store.redb"); let db = Arc::new(redb::Database::create(&redb_path).unwrap()); - // Ensure table exists + // Ensure table exists and migrate old JSON files { let write_txn = db.begin_write().unwrap(); - let _ = write_txn.open_table(crate::store::STORE_TABLE); + { + 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::(&data).is_ok() { + table.insert(*key, data.as_slice()).unwrap(); + } + } + } + } + } + } write_txn.commit().unwrap(); }