diff --git a/Cargo.lock b/Cargo.lock index adc394c..c394e45 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1399,6 +1399,7 @@ dependencies = [ "serde_json", "tantivy", "tempfile", + "thiserror 2.0.20", "tokio", "tokio-stream", "tokio-util", diff --git a/nvim-core/src/lib.rs b/nvim-core/src/lib.rs index f5257f6..dc181d6 100644 --- a/nvim-core/src/lib.rs +++ b/nvim-core/src/lib.rs @@ -23,13 +23,10 @@ pub struct JsonRpcResponse { pub async fn send_response(response: JsonRpcResponse) { let msg = serde_json::to_string(&response).unwrap_or_else(|_| "{}".to_string()); tracing::info!( - "Sending JSON-RPC response (id: {:?}): {}", + "Sending JSON-RPC response (id: {:?}): {}{}", response.id, - if msg.len() > 500 { - format!("{}...", &msg[..500]) - } else { - msg.clone() - } + &msg[..std::cmp::min(msg.len(), 500)], + if msg.len() > 500 { "..." } else { "" } ); // CRITICAL ARCHITECTURAL DECISION: // The MCP StdioTransport MUST use Newline-Delimited JSON (NDJSON). @@ -208,16 +205,20 @@ async fn get_nvim_connection() -> Result, String> { } } } - // Trim buffer if it gets too large - if offset > 1024 * 1024 { + if offset == resp_buf.len() { + resp_buf.clear(); + offset = 0; + } else if offset > 1024 * 1024 { resp_buf.drain(..offset); offset = 0; } continue; } - Err(rmpv::decode::Error::InvalidMarkerRead(e)) - if e.kind() == std::io::ErrorKind::UnexpectedEof => - { + Err(e) if match &e { + rmpv::decode::Error::InvalidMarkerRead(io_err) => io_err.kind() == std::io::ErrorKind::UnexpectedEof, + rmpv::decode::Error::InvalidDataRead(io_err) => io_err.kind() == std::io::ErrorKind::UnexpectedEof, + _ => false, + } => { resp_buf.drain(..offset); offset = 0; @@ -305,86 +306,68 @@ async fn call_nvim(req: rmpv::Value) -> Result { } } -async fn send_nvim_command(cmd: &str) -> Result<(), String> { +async fn call_nvim_method(method: &str, args: Vec) -> Result { use rmpv::Value as RmpValue; let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst); let req = RmpValue::Array(vec![ RmpValue::Integer(0.into()), - RmpValue::Integer(msgid.into()), // msgid - RmpValue::String("nvim_command".into()), - RmpValue::Array(vec![RmpValue::String(cmd.into())]), + RmpValue::Integer(msgid.into()), + RmpValue::String(method.into()), + RmpValue::Array(args), ]); let resp = call_nvim(req).await?; - if let RmpValue::Array(arr) = resp { + if let RmpValue::Array(mut arr) = resp { + if arr.len() < 4 { + return Err("Invalid response length".to_string()); + } if !arr[2].is_nil() { return Err(format!("Neovim error: {:?}", arr[2])); } - return Ok(()); + return Ok(arr.swap_remove(3)); } - Err("Invalid response".to_string()) + Err("Invalid response format".to_string()) +} + +async fn send_nvim_command(cmd: &str) -> Result<(), String> { + call_nvim_method("nvim_command", vec![rmpv::Value::String(cmd.into())]).await?; + Ok(()) } async fn get_nvim_active_buffer() -> Result { - use rmpv::Value as RmpValue; - let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst); - let req = RmpValue::Array(vec![ - RmpValue::Integer(0.into()), - RmpValue::Integer(msgid.into()), // msgid - RmpValue::String("nvim_buf_get_lines".into()), - RmpValue::Array(vec![ - RmpValue::Integer(0.into()), - RmpValue::Integer(0.into()), - RmpValue::Integer((-1).into()), - RmpValue::Boolean(true), - ]), - ]); - - let resp = call_nvim(req).await?; - if let RmpValue::Array(arr) = resp { - if !arr[2].is_nil() { - return Err(format!("Neovim error: {:?}", arr[2])); - } - if let RmpValue::Array(lines) = &arr[3] { - let mut text = String::new(); - for line in lines { - if let RmpValue::String(s) = line { - if let Some(s) = s.as_str() { - text.push_str(s); - text.push('\n'); - } + let result = call_nvim_method("nvim_buf_get_lines", vec![ + rmpv::Value::Integer(0.into()), + rmpv::Value::Integer(0.into()), + rmpv::Value::Integer((-1).into()), + rmpv::Value::Boolean(true), + ]).await?; + + if let rmpv::Value::Array(lines) = result { + let mut text = String::new(); + for line in lines { + if let rmpv::Value::String(s) = line { + if let Some(s) = s.as_str() { + text.push_str(s); + text.push('\n'); } } - return Ok(text); } + return Ok(text); } - Err("Invalid response".to_string()) + Err("Invalid response format".to_string()) } async fn get_nvim_cursor() -> Result { - use rmpv::Value as RmpValue; - let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst); - let req = RmpValue::Array(vec![ - RmpValue::Integer(0.into()), - RmpValue::Integer(msgid.into()), // msgid - RmpValue::String("nvim_win_get_cursor".into()), - RmpValue::Array(vec![RmpValue::Integer(0.into())]), - ]); - - let resp = call_nvim(req).await?; - if let RmpValue::Array(arr) = resp { - if !arr[2].is_nil() { - return Err(format!("Neovim error: {:?}", arr[2])); - } - if let RmpValue::Array(pos) = &arr[3] { - if pos.len() == 2 { - if let (RmpValue::Integer(row), RmpValue::Integer(col)) = (&pos[0], &pos[1]) { - return Ok(format!("Line: {row}, Column: {col}")); - } + let result = call_nvim_method("nvim_win_get_cursor", vec![rmpv::Value::Integer(0.into())]).await?; + + if let rmpv::Value::Array(pos) = result { + if pos.len() == 2 { + if let (rmpv::Value::Integer(row), rmpv::Value::Integer(col)) = (&pos[0], &pos[1]) { + return Ok(format!("Line: {row}, Column: {col}")); } } } - Err("Invalid response".to_string()) + Err("Invalid response format".to_string()) } async fn get_nvim_visual_selection() -> Result { @@ -399,30 +382,17 @@ async fn get_nvim_visual_selection() -> Result { end "#; - use rmpv::Value as RmpValue; - let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst); - let req = RmpValue::Array(vec![ - RmpValue::Integer(0.into()), - RmpValue::Integer(msgid.into()), // msgid - RmpValue::String("nvim_exec_lua".into()), - RmpValue::Array(vec![ - RmpValue::String(lua_script.into()), - RmpValue::Array(vec![]), - ]), - ]); - - let resp = call_nvim(req).await?; - if let RmpValue::Array(arr) = resp { - if !arr[2].is_nil() { - return Err(format!("Neovim error: {:?}", arr[2])); - } - if let RmpValue::String(s) = &arr[3] { - if let Some(text) = s.as_str() { - return Ok(text.to_string()); - } + let result = call_nvim_method("nvim_exec_lua", vec![ + rmpv::Value::String(lua_script.into()), + rmpv::Value::Array(vec![]), + ]).await?; + + if let rmpv::Value::String(s) = result { + if let Some(text) = s.as_str() { + return Ok(text.to_string()); } } - Err("Invalid response".to_string()) + Err("Invalid response format".to_string()) } async fn set_nvim_diagnostics(line: i64, message: &str) -> Result<(), String> { @@ -440,26 +410,12 @@ async fn set_nvim_diagnostics(line: i64, message: &str) -> Result<(), String> { "# ); - use rmpv::Value as RmpValue; - let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst); - let req = RmpValue::Array(vec![ - RmpValue::Integer(0.into()), - RmpValue::Integer(msgid.into()), // msgid - RmpValue::String("nvim_exec_lua".into()), - RmpValue::Array(vec![ - RmpValue::String(lua_script.into()), - RmpValue::Array(vec![]), - ]), - ]); - - let resp = call_nvim(req).await?; - if let RmpValue::Array(arr) = resp { - if !arr[2].is_nil() { - return Err(format!("Neovim error: {:?}", arr[2])); - } - return Ok(()); - } - Err("Invalid response".to_string()) + call_nvim_method("nvim_exec_lua", vec![ + rmpv::Value::String(lua_script.into()), + rmpv::Value::Array(vec![]), + ]).await?; + + Ok(()) } fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value { @@ -505,26 +461,11 @@ fn rmpv_to_json(val: &rmpv::Value) -> serde_json::Value { } async fn execute_nvim_lua(code: &str) -> Result { - use rmpv::Value as RmpValue; - let msgid = NEXT_MSGID.fetch_add(1, Ordering::SeqCst); - let req = RmpValue::Array(vec![ - RmpValue::Integer(0.into()), - RmpValue::Integer(msgid.into()), // msgid - RmpValue::String("nvim_exec_lua".into()), - RmpValue::Array(vec![RmpValue::String(code.into()), RmpValue::Array(vec![])]), - ]); - - let resp = call_nvim(req).await?; - if let RmpValue::Array(arr) = resp { - if !arr[2].is_nil() { - return Err(format!("Neovim error: {:?}", arr[2])); - } - if arr.len() > 3 { - return Ok(serde_json::to_string_pretty(&rmpv_to_json(&arr[3])).unwrap_or_default()); - } - return Ok(String::new()); - } - Err("Invalid response".to_string()) + let result = call_nvim_method("nvim_exec_lua", vec![ + rmpv::Value::String(code.into()), + rmpv::Value::Array(vec![]), + ]).await?; + Ok(serde_json::to_string_pretty(&rmpv_to_json(&result)).unwrap_or_default()) } macro_rules! send_text_result { diff --git a/server/Cargo.toml b/server/Cargo.toml index 0088c2c..079ca00 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -15,7 +15,7 @@ glob = "0.3.4" git2 = { version = "0.19.0", default-features = false } notify = "6.1.1" redb = "4.2.0" -reqwest = { version = "0.12", default-features = false, features = ["stream", "rustls-tls"] } +reqwest = { version = "0.12", default-features = false, features = ["stream", "rustls-tls", "json"] } schemars = "1.2.2" serde = { version = "1.0.229", features = ["derive"] } serde_json = "1.0.151" @@ -28,6 +28,7 @@ tracing-subscriber = "0.3.23" uuid = { version = "1.26.0", features = ["v4"] } tracing-appender = "0.2.5" rmcp = { version = "3.4.0", features = ["server"] } +thiserror = "2.0.20" [build-dependencies] chrono = "0.4.45" @@ -35,8 +36,3 @@ winres = "0.1.12" [dev-dependencies] tempfile = "3.27.0" - -[[bin]] -name = "test_rmcp" -path = "src/bin_test.rs" - diff --git a/server/src/api/mod.rs b/server/src/api/mod.rs new file mode 100644 index 0000000..6d09d65 --- /dev/null +++ b/server/src/api/mod.rs @@ -0,0 +1,4 @@ +pub mod rest; +pub mod setup; +pub mod telemetry; +pub mod ws; diff --git a/server/src/api/rest.rs b/server/src/api/rest.rs new file mode 100644 index 0000000..a5f5500 --- /dev/null +++ b/server/src/api/rest.rs @@ -0,0 +1,107 @@ +use crate::AppState; +use crate::error::AppError; +use crate::models::GateRecord; +use axum::{ + Json, + extract::{Query, State}, + response::IntoResponse, +}; +use std::collections::HashMap; +use std::sync::Arc; + +#[derive(serde::Deserialize, serde::Serialize)] +pub struct GateVerifyReq { + pub action: String, + pub target: String, + pub namespace: Option, + #[serde(default)] + pub params: HashMap, + #[serde(default)] + pub consume: bool, +} + +#[derive(serde::Deserialize, serde::Serialize)] +pub struct GateSetReq { + pub action: String, + pub target: String, + pub namespace: Option, + #[serde(default)] + pub params: HashMap, + pub authorize: Option, + pub block: Option, + pub reason: Option, +} + +pub async fn gate_verify_handler( + State(app_state): State>, + Query(q): Query, +) -> Result { + 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" { + Ok((axum::http::StatusCode::OK, "Authorized")) + } else { + let msg = if let Some(r) = record.reason { + format!("Action blocked. Reason: {}", r) + } else { + "Action blocked.".to_string() + }; + Err(AppError::Forbidden(msg)) + } + } + None => Err(AppError::NotFound("Action not yet authorized (no gate record found).".to_string())), + } +} + +pub async fn gate_set_handler( + State(app_state): State>, + Json(body): Json, +) -> Result { + 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: crate::handlers::utils::now_secs(), + }; + app_state.handler.state.gates.modify(|gates| { + gates.retain(|g| !(g.action == record.action && g.target == record.target)); + gates.push(record); + }); + + Ok((axum::http::StatusCode::OK, "Gate state updated.")) +} + +pub async fn health_handler() -> &'static str { + "OK" +} diff --git a/server/src/api/routes.txt b/server/src/api/routes.txt new file mode 100644 index 0000000..3ff9c34 --- /dev/null +++ b/server/src/api/routes.txt @@ -0,0 +1,301 @@ + 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: crate::handlers::utils::now_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) -> Result<(), Box> { + 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_state_clone = Arc::clone(&app_state); + let mut rx = state.activity_tx.subscribe(); + tokio::spawn(async move { + while let Ok(msg) = rx.recv().await { + let senders: Vec<_> = app_state_clone + .clients + .read() + .unwrap_or_else(|e| e.into_inner()) + .values() + .cloned() + .collect(); + for client_tx in senders { + let _ = client_tx.try_send(msg.clone()); + } + } + }); + + 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>| 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_json = state_clone.read_graph(|g| serde_json::to_string(g).unwrap_or_else(|_| "{}".to_string())); + ([(axum::http::header::CONTENT_TYPE, "application/json")], graph_json) + } + }), + ) + .route( + "/api/tasks/{id}/complete", + post({ + let state_clone = app_state.handler.state.clone(); + move |axum::extract::Path(id): axum::extract::Path| 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_json = state_clone.tasks.read_with(|t| serde_json::to_string(t).unwrap_or_else(|_| "[]".to_string())); + ([(axum::http::header::CONTENT_TYPE, "application/json")], tasks_json) + } + }), + ) + .route( + "/api/sticky", + get({ + let state_clone = app_state.handler.state.clone(); + move || async move { + let sticky_json = state_clone.sticky.read_with(|s| serde_json::to_string(s).unwrap_or_else(|_| "[]".to_string())); + ([(axum::http::header::CONTENT_TYPE, "application/json")], sticky_json) + } + }), + ) + .route( + "/api/search", + get({ + let state_clone = app_state.handler.state.clone(); + move |axum::extract::Query(params): axum::extract::Query< + std::collections::HashMap, + >| 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/activity", + get({ + let state_clone = app_state.handler.state.clone(); + move || async move { + let activities_json = state_clone.recent_activities.read_with(|a| serde_json::to_string(a).unwrap_or_else(|_| "[]".to_string())); + ([(axum::http::header::CONTENT_TYPE, "application/json")], activities_json) + } + }), + ) + .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() + .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()).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; + } diff --git a/server/src/api/setup.rs b/server/src/api/setup.rs new file mode 100644 index 0000000..d33f035 --- /dev/null +++ b/server/src/api/setup.rs @@ -0,0 +1,199 @@ +use crate::AppState; +use crate::api::rest::{gate_set_handler, gate_verify_handler, health_handler}; +use crate::api::telemetry::nvim_telemetry_handler; +use crate::api::ws::ws_handler; +use axum::{ + Router, + extract::State, + response::IntoResponse, + routing::{get, post}, +}; +use std::sync::Arc; + +pub fn create_router(app_state: Arc) -> Router { + 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>| 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_json = state_clone.read_graph(|g| serde_json::to_string(g).unwrap_or_else(|_| "{}".to_string())); + ([(axum::http::header::CONTENT_TYPE, "application/json")], graph_json) + } + }), + ) + .route( + "/api/tasks/{id}/complete", + post({ + let state_clone = app_state.handler.state.clone(); + move |axum::extract::Path(id): axum::extract::Path| 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_json = state_clone.tasks.read_with(|t| serde_json::to_string(t).unwrap_or_else(|_| "[]".to_string())); + ([(axum::http::header::CONTENT_TYPE, "application/json")], tasks_json) + } + }), + ) + .route( + "/api/sticky", + get({ + let state_clone = app_state.handler.state.clone(); + move || async move { + let sticky_json = state_clone.sticky.read_with(|s| serde_json::to_string(s).unwrap_or_else(|_| "[]".to_string())); + ([(axum::http::header::CONTENT_TYPE, "application/json")], sticky_json) + } + }), + ) + .route( + "/api/search", + get({ + let state_clone = app_state.handler.state.clone(); + move |axum::extract::Query(params): axum::extract::Query< + std::collections::HashMap, + >| 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/activity", + get({ + let state_clone = app_state.handler.state.clone(); + move || async move { + let activities_json = state_clone.recent_activities.read_with(|a| serde_json::to_string(a).unwrap_or_else(|_| "[]".to_string())); + ([(axum::http::header::CONTENT_TYPE, "application/json")], activities_json) + } + }), + ) + .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) +} diff --git a/server/src/api/telemetry.rs b/server/src/api/telemetry.rs new file mode 100644 index 0000000..dc88e7a --- /dev/null +++ b/server/src/api/telemetry.rs @@ -0,0 +1,49 @@ +use crate::AppState; +use axum::extract::State; +use std::sync::Arc; + +#[derive(serde::Deserialize, serde::Serialize, Debug)] +pub struct NvimTelemetry { + pub session_id: String, + pub event: String, + pub file: Option, + pub line: Option, + pub col: Option, +} + +pub async fn nvim_telemetry_handler( + State(state): State>, + axum::Json(payload): axum::Json, +) -> 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"})) +} diff --git a/server/src/api/ws.rs b/server/src/api/ws.rs new file mode 100644 index 0000000..c8f4c50 --- /dev/null +++ b/server/src/api/ws.rs @@ -0,0 +1,144 @@ +use crate::AppState; +use axum::extract::{ + Query, State, + ws::{Message, WebSocket}, +}; +use axum::response::IntoResponse; +use futures_util::{SinkExt, StreamExt}; +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::Ordering; +use tokio::sync::mpsc; + +pub async fn ws_handler( + ws: axum::extract::ws::WebSocketUpgrade, + _headers: axum::http::HeaderMap, + State(state): State>, + Query(query): Query>, +) -> 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() +} + +pub struct SessionCleanup { + pub session_id: String, + pub state: Arc, + pub send_task: tokio::task::JoinHandle<()>, + pub recv_task: tokio::task::JoinHandle<()>, +} + +impl Drop for SessionCleanup { + fn drop(&mut self) { + tracing::info!("Dropping session {}", self.session_id); + self.state + .clients + .write() + .unwrap_or_else(|e| e.into_inner()) + .remove(&self.session_id); + self.send_task.abort(); + self.recv_task.abort(); + } +} + +pub async fn handle_socket(socket: WebSocket, state: Arc, _client_type: String) { + let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst)); + let (tx, mut rx) = mpsc::channel::(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; + } + } + }); + + 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::(&text) { + // 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 + ); + } + } + } else { + tracing::warn!( + "Failed to parse payload as JSON from websocket message: {}", + text + ); + } + } + Ok(other) => { + tracing::info!("Received non-text message from websocket: {:?}", other); + } + Err(e) => { + tracing::error!("Websocket receive error: {}", e); + break; + } + } + } + }); + + let mut cleanup = SessionCleanup { + session_id: session_id.clone(), + state: Arc::clone(&state), + send_task, + recv_task, + }; + + tokio::select! { + _ = &mut cleanup.send_task => { + tracing::info!("Websocket send task finished for session {}", session_id); + }, + _ = &mut cleanup.recv_task => { + tracing::info!("Websocket recv task finished for session {}", session_id); + }, + }; +} diff --git a/server/src/db.rs b/server/src/db.rs new file mode 100644 index 0000000..fed58f6 --- /dev/null +++ b/server/src/db.rs @@ -0,0 +1,64 @@ +use std::path::Path; +use std::sync::Arc; +use redb::{Database, ReadableTable}; +use crate::store::STORE_TABLE; + +pub fn init_redb(base: &Path) -> Arc { + let redb_path = base.join("mcp_store.redb"); + let db = Arc::new(redb::Database::create(&redb_path).expect("Failed to create redb database")); + + // Ensure the table exists and migrate legacy JSON files + let write_txn = db.begin_write().expect("Failed to begin write txn on redb"); + { + let mut table = write_txn + .open_table(STORE_TABLE) + .expect("Failed to open STORE_TABLE"); + + 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) + .expect("Failed to read from table") + .is_none() + { + let json_path = base.join(file_name); + if json_path.exists() + && let Ok(data) = std::fs::read(&json_path) + && serde_json::from_slice::(&data).is_ok() + { + table + .insert(*key, data.as_slice()) + .expect("Failed to insert migrated data"); + let _ = std::fs::rename( + &json_path, + json_path.with_extension("json.migrated"), + ); + } + } + } + } + write_txn.commit().expect("Failed to commit db migration"); + + db +} diff --git a/server/src/error.rs b/server/src/error.rs new file mode 100644 index 0000000..5a98846 --- /dev/null +++ b/server/src/error.rs @@ -0,0 +1,39 @@ +use axum::{ + http::StatusCode, + response::{IntoResponse, Response}, + Json, +}; +use serde_json::json; +use thiserror::Error; + +#[derive(Error, Debug)] +pub enum AppError { + #[error("Not Found: {0}")] + NotFound(String), + + #[error("Forbidden: {0}")] + Forbidden(String), + + #[error("Internal Server Error: {0}")] + Internal(String), + + #[error("Bad Request: {0}")] + BadRequest(String), +} + +impl IntoResponse for AppError { + fn into_response(self) -> Response { + let (status, error_message) = match &self { + AppError::NotFound(msg) => (StatusCode::NOT_FOUND, msg.clone()), + AppError::Forbidden(msg) => (StatusCode::FORBIDDEN, msg.clone()), + AppError::Internal(msg) => (StatusCode::INTERNAL_SERVER_ERROR, msg.clone()), + AppError::BadRequest(msg) => (StatusCode::BAD_REQUEST, msg.clone()), + }; + + let body = Json(json!({ + "error": error_message, + })); + + (status, body).into_response() + } +} diff --git a/server/src/handlers_v2/env.rs b/server/src/handlers/env.rs similarity index 97% rename from server/src/handlers_v2/env.rs rename to server/src/handlers/env.rs index 6ef407d..575fc4a 100644 --- a/server/src/handlers_v2/env.rs +++ b/server/src/handlers/env.rs @@ -31,7 +31,7 @@ impl McpTool for UpdateEnvFingerprintHandler { os: std::env::consts::OS.to_string(), shell: std::env::var("SHELL").unwrap_or_else(|_| "unknown".to_string()), tool_versions: req.tool_versions, - updated_at: crate::handlers_v2::utils::now_secs(), + updated_at: crate::handlers::utils::now_secs(), }, ); }); @@ -125,7 +125,7 @@ impl McpTool for RegisterEnvironmentHandler { url: req.url, description: req.description, requires_vpn: req.requires_vpn, - updated_at: crate::handlers_v2::utils::now_secs(), + updated_at: crate::handlers::utils::now_secs(), }); }); Ok("Environment registered".to_string()) diff --git a/server/src/handlers_v2/graph.rs b/server/src/handlers/graph.rs similarity index 99% rename from server/src/handlers_v2/graph.rs rename to server/src/handlers/graph.rs index 6fdb8f8..927cb38 100644 --- a/server/src/handlers_v2/graph.rs +++ b/server/src/handlers/graph.rs @@ -576,4 +576,4 @@ impl McpTool for FindOrphansHandler { } } -use crate::handlers_v2::utils::*; +use crate::handlers::utils::*; diff --git a/server/src/handlers_v2/meta.rs b/server/src/handlers/meta.rs similarity index 97% rename from server/src/handlers_v2/meta.rs rename to server/src/handlers/meta.rs index 85f158b..6db96fa 100644 --- a/server/src/handlers_v2/meta.rs +++ b/server/src/handlers/meta.rs @@ -36,7 +36,7 @@ impl McpTool for LogDecisionHandler { context: req.context, decision: req.decision, consequence: req.consequence, - timestamp: crate::handlers_v2::utils::now_secs(), + timestamp: crate::handlers::utils::now_secs(), }; drop(idx.index_adr(&a)); @@ -98,7 +98,7 @@ impl McpTool for LogErrorFixHandler { fixes.push(crate::models::ErrorFix { signature: req.signature, solution: req.solution, - timestamp: crate::handlers_v2::utils::now_secs(), + timestamp: crate::handlers::utils::now_secs(), git_commit: req.git_commit, git_branch: req.git_branch, }) @@ -155,7 +155,7 @@ impl McpTool for LogCodeChangeHandler { let req: LogCodeChangeTool = serde_json::from_value(args).map_err(|e| e.to_string())?; state.ledger.modify(|ledger| { ledger.push(CodeChange { - timestamp: crate::handlers_v2::utils::now_secs(), + timestamp: crate::handlers::utils::now_secs(), file_path: req.file_path, description: req.description, git_commit: req.git_commit, @@ -209,7 +209,7 @@ impl McpTool for LearnPreferenceHandler { crate::models::Preference { key: req.key.clone(), value: req.value, - updated_at: crate::handlers_v2::utils::now_secs(), + updated_at: crate::handlers::utils::now_secs(), }, ); }); @@ -258,7 +258,7 @@ impl McpTool for LogTechDebtHandler { description: req.description, ideal_solution: req.ideal_solution, is_resolved: false, - created_at: crate::handlers_v2::utils::now_secs(), + created_at: crate::handlers::utils::now_secs(), git_commit: req.git_commit, git_branch: req.git_branch, }) @@ -501,4 +501,4 @@ impl McpTool for GetProjectHealthHandler { } } -use crate::handlers_v2::utils::*; +use crate::handlers::utils::*; diff --git a/server/src/handlers_v2/mod.rs b/server/src/handlers/mod.rs similarity index 100% rename from server/src/handlers_v2/mod.rs rename to server/src/handlers/mod.rs diff --git a/server/src/handlers_v2/notes.rs b/server/src/handlers/notes.rs similarity index 96% rename from server/src/handlers_v2/notes.rs rename to server/src/handlers/notes.rs index bf1cfee..d59dbfd 100644 --- a/server/src/handlers_v2/notes.rs +++ b/server/src/handlers/notes.rs @@ -23,7 +23,7 @@ impl McpTool for AddStickyNoteHandler { let req: AddStickyNoteTool = serde_json::from_value(args).map_err(|e| e.to_string())?; state.sticky.modify(|notes| { notes.push(StickyNote { - timestamp: crate::handlers_v2::utils::now_secs(), + timestamp: crate::handlers::utils::now_secs(), content: req.content, }); }); @@ -132,7 +132,7 @@ impl McpTool for LeaveHandoffMemoHandler { author: "agy".to_string(), content: req.content, namespace: req.namespace, - timestamp: crate::handlers_v2::utils::now_secs(), + timestamp: crate::handlers::utils::now_secs(), }) }); Ok("Handoff memo left".to_string()) @@ -219,7 +219,7 @@ impl McpTool for AddSessionSummaryHandler { summaries.push(crate::models::SessionSummary { summary: req.summary, namespace: req.namespace, - timestamp: crate::handlers_v2::utils::now_secs(), + timestamp: crate::handlers::utils::now_secs(), }) }); Ok("Session summary added".to_string()) @@ -245,7 +245,7 @@ impl McpTool for GenerateStandupReportHandler { let req: GenerateStandupReportTool = serde_json::from_value(args).map_err(|e| e.to_string())?; let cutoff = - crate::handlers_v2::utils::now_secs().saturating_sub(req.hours_lookback * 3600); + crate::handlers::utils::now_secs().saturating_sub(req.hours_lookback * 3600); let report_str = state.tasks.read_with(|items| { state.ledger.read_with(|changes| { diff --git a/server/src/handlers_v2/tasks.rs b/server/src/handlers/tasks.rs similarity index 98% rename from server/src/handlers_v2/tasks.rs rename to server/src/handlers/tasks.rs index 0551cb7..f2a3767 100644 --- a/server/src/handlers_v2/tasks.rs +++ b/server/src/handlers/tasks.rs @@ -20,7 +20,7 @@ impl McpTool for AddTaskHandler { async fn execute(&self, args: Value, state: Arc) -> Result { let req: AddTaskTool = serde_json::from_value(args).map_err(|e| e.to_string())?; - let now = crate::handlers_v2::utils::now_secs(); + let now = crate::handlers::utils::now_secs(); let task_id = uuid::Uuid::new_v4().to_string(); let deps = req.dependencies.unwrap_or_default(); @@ -214,7 +214,7 @@ impl McpTool for UpdateTaskStatusHandler { if !blocked { // Apply update tasks[target_idx].status = target_status.clone(); - tasks[target_idx].updated_at = crate::handlers_v2::utils::now_secs(); + tasks[target_idx].updated_at = crate::handlers::utils::now_secs(); // Cascade cancellation to children if target_status == "cancelled" || target_status == "abandoned" { @@ -337,7 +337,7 @@ impl McpTool for SetAcceptanceCriteriaHandler { is_met: false, }) .collect(); - task.updated_at = crate::handlers_v2::utils::now_secs(); + task.updated_at = crate::handlers::utils::now_secs(); success = true; } }); @@ -381,7 +381,7 @@ impl McpTool for VerifyAcceptanceCriteriaHandler { } else { ac.is_met = true; success = true; - task.updated_at = crate::handlers_v2::utils::now_secs(); + task.updated_at = crate::handlers::utils::now_secs(); } } }); diff --git a/server/src/handlers_v2/utils.rs b/server/src/handlers/utils.rs similarity index 100% rename from server/src/handlers_v2/utils.rs rename to server/src/handlers/utils.rs diff --git a/server/src/handlers_v2/workspaces.rs b/server/src/handlers/workspaces.rs similarity index 98% rename from server/src/handlers_v2/workspaces.rs rename to server/src/handlers/workspaces.rs index 35cff51..413c12c 100644 --- a/server/src/handlers_v2/workspaces.rs +++ b/server/src/handlers/workspaces.rs @@ -25,7 +25,7 @@ impl McpTool for PinFileHandler { pinned.push(crate::models::PinnedFile { namespace: req.namespace, file_path: req.file_path, - timestamp: crate::handlers_v2::utils::now_secs(), + timestamp: crate::handlers::utils::now_secs(), git_branch: req.git_branch, }); }); @@ -115,7 +115,7 @@ impl McpTool for StoreSnippetHandler { language: req.language, code: req.code, description: req.description, - updated_at: crate::handlers_v2::utils::now_secs(), + updated_at: crate::handlers::utils::now_secs(), }; let idx = state @@ -223,7 +223,7 @@ impl McpTool for SaveContextWorkspaceHandler { name: req.name, pinned_files: req.pinned_files, active_task_ids: req.active_task_ids, - saved_at: crate::handlers_v2::utils::now_secs(), + saved_at: crate::handlers::utils::now_secs(), }); }); Ok("Context workspace saved".to_string()) @@ -363,4 +363,4 @@ impl McpTool for ClearPrChecklistHandler { } } -use crate::handlers_v2::utils::*; +use crate::handlers::utils::*; diff --git a/server/src/main.rs b/server/src/main.rs index 27c55b0..0717353 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -3,29 +3,28 @@ allow(dead_code, unused_imports, unreachable_code) )] +mod api; mod handlers; -mod handlers_v2; mod mcp; mod models; mod router; mod search; +pub mod db; +pub mod error; mod state; mod store; mod tools; -use crate::handlers::MemoryHandler; -use crate::models::*; +use crate::api::rest::GateSetReq; +use crate::router::MemoryHandler; 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; - 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)] @@ -87,6 +86,12 @@ enum GateCommands { }, } +pub struct AppState { + handler: Arc, + clients: RwLock>>, + next_id: AtomicUsize, +} + async fn index_committer_worker(state: Arc) { loop { tokio::time::sleep(Duration::from_secs(5)).await; @@ -98,121 +103,6 @@ async fn index_committer_worker(state: Arc) { } } -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, - clients: RwLock>>, - 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: crate::handlers_v2::utils::now_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) -> Result<(), Box> { state.rebuild_index().await; tokio::spawn(index_committer_worker(Arc::clone(&state))); @@ -239,191 +129,7 @@ async fn run_server(state: Arc) -> Result<(), Box>| 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_json = state_clone.read_graph(|g| serde_json::to_string(g).unwrap_or_else(|_| "{}".to_string())); - ([(axum::http::header::CONTENT_TYPE, "application/json")], graph_json) - } - }), - ) - .route( - "/api/tasks/{id}/complete", - post({ - let state_clone = app_state.handler.state.clone(); - move |axum::extract::Path(id): axum::extract::Path| 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_json = state_clone.tasks.read_with(|t| serde_json::to_string(t).unwrap_or_else(|_| "[]".to_string())); - ([(axum::http::header::CONTENT_TYPE, "application/json")], tasks_json) - } - }), - ) - .route( - "/api/sticky", - get({ - let state_clone = app_state.handler.state.clone(); - move || async move { - let sticky_json = state_clone.sticky.read_with(|s| serde_json::to_string(s).unwrap_or_else(|_| "[]".to_string())); - ([(axum::http::header::CONTENT_TYPE, "application/json")], sticky_json) - } - }), - ) - .route( - "/api/search", - get({ - let state_clone = app_state.handler.state.clone(); - move |axum::extract::Query(params): axum::extract::Query< - std::collections::HashMap, - >| 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/activity", - get({ - let state_clone = app_state.handler.state.clone(); - move || async move { - let activities_json = state_clone.recent_activities.read_with(|a| serde_json::to_string(a).unwrap_or_else(|_| "[]".to_string())); - ([(axum::http::header::CONTENT_TYPE, "application/json")], activities_json) - } - }), - ) - .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); + let app = api::setup::create_router(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()); @@ -451,201 +157,6 @@ async fn run_server(state: Arc) -> Result<(), Box>, - axum::extract::Query(query): axum::extract::Query>, -) -> 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, _client_type: String) { - let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst)); - let (tx, mut rx) = mpsc::channel::(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::(&text) { - // 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, - send_task: tokio::task::JoinHandle<()>, - recv_task: 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); - self.send_task.abort(); - self.recv_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, - recv_task, - }; - - tokio::select! { - _ = &mut cleanup.send_task => { - tracing::info!("Websocket send task finished for session {}", session_id); - }, - _ = &mut cleanup.recv_task => { - 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, - pub line: Option, - pub col: Option, -} - -async fn nvim_telemetry_handler( - State(state): State>, - axum::Json(payload): axum::Json, -) -> 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 { let base_dir = std::env::var("MCP_MEMORY_STORE_DIR").unwrap_or_else(|_| { dirs::home_dir() @@ -696,148 +207,131 @@ fn main() -> Result<(), Box> { .arg(format!("Authorization: Bearer {}", token.trim())); } let _ = cmd - .arg(format!("https://127.0.0.1:{}/shutdown", port)) + .arg(format!("http://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().expect("Failed to get current executable path"), - ) - .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"); + if cli.restart { + std::thread::sleep(Duration::from_secs(2)); + } else { 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).expect("Failed to create redb database")); - - // Ensure table exists and migrate old JSON files - { - let write_txn = db.begin_write().expect("Failed to begin write txn on redb"); - { - let mut table = write_txn - .open_table(crate::store::STORE_TABLE) - .expect("Failed to open STORE_TABLE"); - - 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) - .expect("Failed to read from table") - .is_none() - { - let json_path = base.join(file_name); - if json_path.exists() - && let Ok(data) = fs::read(&json_path) - && serde_json::from_slice::(&data).is_ok() - { - table - .insert(*key, data.as_slice()) - .expect("Failed to insert migrated data"); - let _ = fs::rename(&json_path, json_path.with_extension("json.migrated")); + 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); + } + } + }); } } - write_txn.commit().expect("Failed to commit db migration"); + return Ok(()); } - let rt = tokio::runtime::Runtime::new().expect("Failed to create tokio runtime"); - let _guard = rt.enter(); + #[cfg(target_os = "windows")] + if cli.daemon { + let exe = std::env::current_exe()?; + std::process::Command::new("powershell") + .args([ + "-WindowStyle", + "Hidden", + "-Command", + &format!( + "Start-Process -FilePath '{}' -WindowStyle Hidden", + exe.display() + ), + ]) + .spawn()?; + return Ok(()); + } - 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()), - recent_activities: Store::new("recent_activities", db.clone()), - activity_tx: tokio::sync::broadcast::channel(100).0, + 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); + } }); - rt.block_on(run_server(state)) + Ok(()) } diff --git a/server/src/router.rs b/server/src/router.rs index 2520791..527c390 100644 --- a/server/src/router.rs +++ b/server/src/router.rs @@ -14,3 +14,197 @@ pub trait McpTool: Send + Sync { /// Execute the tool with the given arguments async fn execute(&self, args: Value, state: Arc) -> Result; } + +pub struct MemoryHandler { + pub state: Arc, + pub tools: std::collections::HashMap>, +} + +impl MemoryHandler { + pub fn new(state: Arc) -> Self { + let mut tools: std::collections::HashMap> = + std::collections::HashMap::new(); + + macro_rules! register { + ($module:ident::$handler:ident) => { + let h = crate::handlers::$module::$handler; + tools.insert(h.name().to_string(), Box::new(h)); + }; + } + + register!(graph::QueryGraphPathHandler); + register!(graph::CreateEntitiesHandler); + register!(graph::CreateRelationsHandler); + register!(graph::AddObservationsHandler); + register!(graph::DeleteEntitiesHandler); + register!(graph::DeleteObservationsHandler); + register!(graph::DeleteRelationsHandler); + register!(graph::ReadGraphHandler); + register!(graph::SearchNodesHandler); + register!(graph::OpenNodesHandler); + register!(graph::VisualizeGraphHandler); + register!(graph::CondenseEntityHandler); + register!(graph::MergeEntitiesHandler); + register!(graph::FindOrphansHandler); + + register!(tasks::AddTaskHandler); + register!(tasks::DeleteTaskHandler); + register!(tasks::UpdateTaskStatusHandler); + register!(tasks::ListActiveTasksHandler); + register!(tasks::SetAcceptanceCriteriaHandler); + register!(tasks::VerifyAcceptanceCriteriaHandler); + register!(tasks::AddMilestoneHandler); + register!(tasks::UpdateMilestoneHandler); + register!(tasks::ListMilestonesHandler); + + register!(notes::AddStickyNoteHandler); + register!(notes::ReadStickyNotesHandler); + register!(notes::DeleteStickyNoteHandler); + register!(notes::ClearStickyNotesHandler); + register!(notes::LeaveHandoffMemoHandler); + register!(notes::ReadHandoffMemosHandler); + register!(notes::ClearHandoffMemosHandler); + register!(notes::AddSessionSummaryHandler); + register!(notes::GenerateStandupReportHandler); + + register!(meta::LogDecisionHandler); + register!(meta::QueryDecisionsHandler); + register!(meta::LogErrorFixHandler); + register!(meta::SearchErrorFixesHandler); + register!(meta::LogCodeChangeHandler); + register!(meta::QueryRecentChangesHandler); + register!(meta::LearnPreferenceHandler); + register!(meta::ReadPreferencesHandler); + register!(meta::LogTechDebtHandler); + register!(meta::ResolveTechDebtHandler); + register!(meta::ListTechDebtHandler); + register!(meta::OmniSearchHandler); + register!(meta::GetProjectHealthHandler); + + register!(env::UpdateEnvFingerprintHandler); + register!(env::ReadEnvFingerprintHandler); + register!(env::LogEnvRequirementHandler); + register!(env::RegisterEnvironmentHandler); + register!(env::GetEnvironmentDetailsHandler); + + register!(workspaces::PinFileHandler); + register!(workspaces::UnpinFileHandler); + register!(workspaces::ListPinnedFilesHandler); + register!(workspaces::StoreSnippetHandler); + register!(workspaces::SearchSnippetsHandler); + register!(workspaces::DeleteSnippetHandler); + register!(workspaces::SaveContextWorkspaceHandler); + register!(workspaces::LoadContextWorkspaceHandler); + register!(workspaces::ListContextWorkspacesHandler); + register!(workspaces::AddPrChecklistItemHandler); + register!(workspaces::GetPrChecklistHandler); + register!(workspaces::ClearPrChecklistHandler); + + Self { state, tools } + } + + pub async fn handle_request(&self, req: serde_json::Value) -> Option { + let id = req.get("id").cloned().unwrap_or(serde_json::Value::Null); + let id_clone = id.clone(); + let method = req.get("method").and_then(|m| m.as_str()).unwrap_or(""); + + match method { + "server/discover" => { + let payload = serde_json::json!({ + "resultType": "complete", + "ttlMs": 0, + "cacheScope": "public", + "supportedVersions": ["2026-07-28", "2025-11-25", "2025-06-18", "2025-03-26", "2024-11-05"], + "capabilities": { + "tools": serde_json::json!({}) + }, + "_meta": { + "io.modelcontextprotocol/serverInfo": { + "name": "gemini-mcp-memory", + "version": "3.0.0" + } + } + }); + Some(crate::mcp::success(id, payload)) + } + "initialize" => { + let init = rmcp::model::InitializeResult::new( + rmcp::model::ServerCapabilities::builder() + .enable_tools() + .build(), + ) + .with_server_info(rmcp::model::Implementation::new( + "gemini-mcp-memory", + "3.0.0", + )); + Some(crate::mcp::success( + id, + serde_json::to_value(&init).unwrap_or_default(), + )) + } + "notifications/initialized" => None, + "tools/list" => { + let mut tools: Vec = + self.tools.values().map(|t| t.schema()).collect(); + tools.sort_by_key(|t| { + t.get("name") + .and_then(|n| n.as_str()) + .unwrap_or("") + .to_string() + }); + Some(crate::mcp::success( + id, + serde_json::json!({ "tools": tools }), + )) + } + "tools/call" => { + let params = req.get("params").unwrap_or(&serde_json::Value::Null); + let name = params.get("name").and_then(|n| n.as_str()).unwrap_or(""); + let args = params + .get("arguments") + .cloned() + .unwrap_or(serde_json::Value::Object(Default::default())); + + self.state + .broadcast_activity(&format!("Agent executed tool: {}", name)); + + let result: Result = if let Some(tool) = self.tools.get(name) { + tool.execute(args, self.state.clone()).await + } else { + Err(format!("Unknown tool: {}", name)) + }; + + match result { + Ok(text) => { + let payload = serde_json::json!({ + "content": [{"type": "text", "text": text}], + "isError": false + }); + Some(crate::mcp::success(id_clone, payload)) + } + Err(e) => { + tracing::error!("Tool {} failed: {}", name, e); + let payload = serde_json::json!({ + "content": [{"type": "text", "text": e}], + "isError": true + }); + Some(crate::mcp::success(id_clone, payload)) + } + } + } + m if m.starts_with("notifications/") => None, + "ping" => Some(crate::mcp::success(id, serde_json::json!({}))), + _ => { + if id.is_null() { + None + } else { + Some(crate::mcp::error( + id, + -32601, + &format!("Method {} not found", method), + )) + } + } + } + } +} diff --git a/server/src/state.rs b/server/src/state.rs index 7eb3008..28694ff 100644 --- a/server/src/state.rs +++ b/server/src/state.rs @@ -32,6 +32,49 @@ pub struct MemoryState { } impl MemoryState { + pub fn new(base_dir_str: &str) -> Self { + let base = std::path::PathBuf::from(base_dir_str); + std::fs::create_dir_all(&base).expect("Failed to create store dir"); + + let db = crate::db::init_redb(&base); + + Self { + 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()), + recent_activities: Store::new("recent_activities", db.clone()), + activity_tx: tokio::sync::broadcast::channel(100).0, + } + } + pub fn deduplicate(input: &mut Vec) { let mut keys = std::collections::HashSet::new(); input.retain(|entry| keys.insert(entry.clone())); diff --git a/stub/src/main.rs b/stub/src/main.rs index f615935..7277890 100644 --- a/stub/src/main.rs +++ b/stub/src/main.rs @@ -95,14 +95,11 @@ fn main() -> Result<(), Box> { while let Ok(msg) = rx.recv().await { let log_prefix = logger::extract_log_prefix(&msg, false); tracing::info!( - ">>> [Stub] Forwarding {} to server (length: {}): {}", + ">>> [Stub] Forwarding {} to server (length: {}): {}{}", log_prefix, msg.len(), - if msg.len() > 1000 { - format!("{}...", &msg[..1000]) - } else { - msg.clone() - } + &msg[..std::cmp::min(msg.len(), 1000)], + if msg.len() > 1000 { "..." } else { "" } ); if write .send(tokio_tungstenite::tungstenite::Message::Text(msg)) @@ -120,14 +117,11 @@ fn main() -> Result<(), Box> { if let tokio_tungstenite::tungstenite::Message::Text(text) = msg { let log_prefix = logger::extract_log_prefix(&text, true); tracing::info!( - "<<< [Stub] Received {} from server (length: {}): {}", + "<<< [Stub] Received {} from server (length: {}): {}{}", log_prefix, text.len(), - if text.len() > 1000 { - format!("{}...", &text[..1000]) - } else { - text.clone() - } + &text[..std::cmp::min(text.len(), 1000)], + if text.len() > 1000 { "..." } else { "" } ); use tokio::io::AsyncWriteExt; let mut stdout = tokio::io::stdout();