refactor(mcp): strip sse fallback architecture in favor of pure websockets and fix nvim NDJSON bug
This commit is contained in:
1 parent
0e29b12ac8
commit
3716c3e698
33 files changed
+2082
-1756
No files matched your search
+112
-229
@@ -83,98 +83,13 @@ enum GateCommands {
|
||||
},
|
||||
}
|
||||
|
||||
async fn garbage_collector_worker(state: Arc<MemoryState>) {
|
||||
loop {
|
||||
// Run every 6 hours
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(6 * 3600)).await;
|
||||
|
||||
let now = std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_secs();
|
||||
|
||||
// 1. Task GC (14 days)
|
||||
let fourteen_days = 14 * 24 * 3600;
|
||||
let task_cutoff = now.saturating_sub(fourteen_days);
|
||||
state.tasks.modify(|tasks| {
|
||||
let initial_len = tasks.len();
|
||||
tasks.retain(|task| !(task.status.to_lowercase() == "completed" && task.created_at < task_cutoff));
|
||||
if tasks.len() < initial_len {
|
||||
eprintln!("GC: Removed {} old completed tasks", initial_len - tasks.len());
|
||||
}
|
||||
});
|
||||
|
||||
// 2. Ledger GC (7 days or max 1000 items)
|
||||
state.ledger.modify(|ledger| {
|
||||
let seven_days = now.saturating_sub(7 * 24 * 3600);
|
||||
ledger.retain(|c| c.timestamp >= seven_days);
|
||||
if ledger.len() > 1000 {
|
||||
let excess = ledger.len() - 1000;
|
||||
ledger.drain(0..excess);
|
||||
}
|
||||
});
|
||||
|
||||
// 3. Sticky Notes GC (24 hours)
|
||||
state.sticky.modify(|notes| {
|
||||
notes.retain(|note| note.timestamp >= now.saturating_sub(24 * 3600));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async fn git_sync_worker(state: Arc<MemoryState>) {
|
||||
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;
|
||||
|
||||
let repo_path_clone = repo_path.clone();
|
||||
let commit_data = tokio::task::spawn_blocking(move || {
|
||||
if let Ok(repo) = git2::Repository::discover(&repo_path_clone) {
|
||||
if let Ok(head) = repo.head() {
|
||||
if let Ok(commit) = head.peel_to_commit() {
|
||||
let current_id = commit.id().to_string();
|
||||
let msg = commit.message().unwrap_or("").to_string();
|
||||
let branch = head.shorthand().unwrap_or("unknown").to_string();
|
||||
return Some((current_id, msg, branch));
|
||||
}
|
||||
}
|
||||
}
|
||||
None
|
||||
})
|
||||
.await
|
||||
.unwrap_or(None);
|
||||
|
||||
if let Some((current_id, msg, branch)) = commit_data {
|
||||
if current_id != last_commit_id && !last_commit_id.is_empty() {
|
||||
state.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(),
|
||||
});
|
||||
});
|
||||
tracing::info!("Git Sync: Logged new commit {}", current_id);
|
||||
|
||||
state.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();
|
||||
tracing::info!("Git Sync: Auto-completed task '{}'", task.title);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
last_commit_id = current_id;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn reconcile_worker(state: Arc<MemoryState>) {
|
||||
loop {
|
||||
sleep(Duration::from_secs(5)).await;
|
||||
|
||||
let has_local = {
|
||||
let session = state.session_graph.read().unwrap();
|
||||
let session = state.graph.read();
|
||||
!session.entities.is_empty() || !session.relations.is_empty()
|
||||
};
|
||||
|
||||
@@ -187,7 +102,7 @@ async fn reconcile_worker(state: Arc<MemoryState>) {
|
||||
.unwrap_or(false);
|
||||
|
||||
if has_local || has_files {
|
||||
state.apply_sync_write(|_master| {}).await;
|
||||
state.apply_sync_write(|_master| {});
|
||||
let state_clone = state.clone();
|
||||
let _ = tokio::task::spawn_blocking(move || {
|
||||
state_clone.rebuild_index();
|
||||
@@ -198,7 +113,7 @@ async fn reconcile_worker(state: Arc<MemoryState>) {
|
||||
|
||||
use axum::{
|
||||
Json, Router,
|
||||
extract::{Query, State, ws::{WebSocketUpgrade, WebSocket, Message}},
|
||||
extract::{Query, State, ws::{WebSocket, Message}},
|
||||
response::IntoResponse,
|
||||
routing::{get, post},
|
||||
};
|
||||
@@ -327,8 +242,6 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
}))
|
||||
}))
|
||||
.route("/ws", get(ws_handler))
|
||||
.route("/sse", get(sse_handler))
|
||||
.route("/messages", post(message_handler))
|
||||
.route("/health", get(health_handler))
|
||||
.route("/nvim/telemetry", post(nvim_telemetry_handler))
|
||||
.route("/gate/verify", get(gate_verify_handler))
|
||||
@@ -458,19 +371,12 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
)
|
||||
.with_state(app_state);
|
||||
|
||||
let listener = match tokio::net::TcpListener::bind(std::env::var("MCP_PORT").map(|p| format!("127.0.0.1:{}", p)).unwrap_or_else(|_| std::env::var("MCP_PORT").map(|p| format!("127.0.0.1:{}", p)).unwrap_or_else(|_| "127.0.0.1:3000".to_string()).to_string())).await {
|
||||
Ok(l) => l,
|
||||
Err(e) => {
|
||||
tracing::info!("Port 3000 is already in use ({}). Assuming server is already running and exiting gracefully.", e);
|
||||
std::process::exit(0);
|
||||
}
|
||||
};
|
||||
tracing::info!("MCP Memory Server running on http://127.0.0.1:3000/sse");
|
||||
let addr = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string());
|
||||
let addr: std::net::SocketAddr = format!("127.0.0.1:{}", addr).parse().unwrap();
|
||||
|
||||
tokio::spawn(garbage_collector_worker(Arc::clone(&state)));
|
||||
|
||||
tokio::spawn(git_sync_worker(Arc::clone(&state)));
|
||||
tracing::info!("MCP Memory Server running on http://127.0.0.1:3000/sse");
|
||||
if let Err(e) = axum::serve(listener, app).await {
|
||||
let listener = tokio::net::TcpListener::bind(addr).await.unwrap();
|
||||
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");
|
||||
@@ -480,56 +386,16 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct MsgQuery {
|
||||
session_id: String,
|
||||
}
|
||||
|
||||
async fn message_handler(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Query(q): Query<MsgQuery>,
|
||||
Json(payload): Json<serde_json::Value>,
|
||||
) -> impl axum::response::IntoResponse {
|
||||
let session_id = q.session_id;
|
||||
if let Some(response) = state.handler.handle_request(payload).await {
|
||||
let res_str = serde_json::to_string(&response).unwrap();
|
||||
let tx_opt = state.clients.read().unwrap().get(&session_id).cloned();
|
||||
if let Some(tx) = tx_opt {
|
||||
let _ = tx.send(res_str).await;
|
||||
}
|
||||
}
|
||||
(axum::http::StatusCode::ACCEPTED, "Accepted").into_response()
|
||||
}
|
||||
|
||||
async fn sse_handler(
|
||||
State(state): State<Arc<AppState>>,
|
||||
) -> axum::response::sse::Sse<impl tokio_stream::Stream<Item = Result<axum::response::sse::Event, std::convert::Infallible>>> {
|
||||
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
|
||||
let (tx, rx) = mpsc::channel::<String>(100);
|
||||
state.clients.write().unwrap().insert(session_id.clone(), tx.clone());
|
||||
|
||||
let endpoint = format!("/messages?session_id={}", session_id);
|
||||
let _ = tx.send(format!("endpoint|{}", endpoint)).await;
|
||||
|
||||
let rx_stream = tokio_stream::wrappers::ReceiverStream::new(rx);
|
||||
let event_stream = rx_stream.map(|msg| {
|
||||
if let Some(ep) = msg.strip_prefix("endpoint|") {
|
||||
Ok(axum::response::sse::Event::default().event("endpoint").data(ep))
|
||||
} else {
|
||||
Ok(axum::response::sse::Event::default().event("message").data(msg))
|
||||
}
|
||||
});
|
||||
|
||||
axum::response::sse::Sse::new(event_stream).keep_alive(axum::response::sse::KeepAlive::new())
|
||||
}
|
||||
|
||||
async fn ws_handler(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<Arc<AppState>>,
|
||||
Query(query): Query<std::collections::HashMap<String, String>>,
|
||||
) -> impl axum::response::IntoResponse {
|
||||
ws: axum::extract::ws::WebSocketUpgrade,
|
||||
headers: axum::http::HeaderMap,
|
||||
axum::extract::State(state): axum::extract::State<Arc<AppState>>,
|
||||
axum::extract::Query(query): axum::extract::Query<std::collections::HashMap<String, String>>,
|
||||
) -> axum::response::Response {
|
||||
let client_type = query.get("client").cloned().unwrap_or_else(|| "unknown".to_string());
|
||||
ws.on_upgrade(move |socket| handle_socket(socket, state, client_type))
|
||||
ws.on_upgrade(move |socket| handle_socket(socket, state, client_type)).into_response()
|
||||
}
|
||||
|
||||
async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: String) {
|
||||
@@ -542,71 +408,92 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
|
||||
let mut 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;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
if client_type == "proxy" {
|
||||
let tx_clone = tx.clone();
|
||||
tokio::spawn(async move {
|
||||
let notify = serde_json::json!({
|
||||
"jsonrpc": "2.0",
|
||||
"method": "notifications/tools/list_changed"
|
||||
});
|
||||
let _ = tx_clone.send(notify.to_string()).await;
|
||||
});
|
||||
}
|
||||
// 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 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;
|
||||
}
|
||||
}
|
||||
}
|
||||
while let Some(msg_result) = receiver.next().await {
|
||||
match msg_result {
|
||||
Ok(Message::Text(text)) => {
|
||||
tracing::info!("Received text message from websocket (length: {})", text.len());
|
||||
tracing::trace!("Message content: {}", text);
|
||||
if let Ok(payload) = serde_json::from_str::<serde_json::Value>(&text) {
|
||||
if client_type == "proxy" {
|
||||
// Send activity broadcast to UI clients
|
||||
if let Some(method) = payload.get("method").and_then(|m| m.as_str()) {
|
||||
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;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} // End if proxy
|
||||
|
||||
// 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 {
|
||||
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);
|
||||
}
|
||||
|
||||
// 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);
|
||||
}
|
||||
} // 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);
|
||||
});
|
||||
|
||||
tokio::select! {
|
||||
_ = (&mut send_task) => {
|
||||
tracing::info!("Websocket send task finished for session {}", session_id);
|
||||
recv_task.abort();
|
||||
},
|
||||
_ = (&mut recv_task) => {
|
||||
tracing::info!("Websocket recv task finished for session {}", session_id);
|
||||
send_task.abort();
|
||||
},
|
||||
};
|
||||
|
||||
state.clients.write().unwrap().remove(&session_id);
|
||||
tracing::info!("Websocket session {} closed and removed from state", session_id);
|
||||
}
|
||||
|
||||
|
||||
#[derive(serde::Deserialize, serde::Serialize, Debug)]
|
||||
@@ -670,35 +557,39 @@ fn init_logging(app_name: &str) -> Option<tracing_appender::non_blocking::Worker
|
||||
.with_writer(non_blocking)
|
||||
.with_ansi(false)
|
||||
.with_max_level(tracing::Level::INFO)
|
||||
.with_thread_ids(true)
|
||||
.with_thread_names(true)
|
||||
.try_init();
|
||||
|
||||
Some(guard)
|
||||
}
|
||||
|
||||
fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
let _guard = init_logging("server");
|
||||
let _guard = init_logging("mcp-memory-server");
|
||||
let cli = Cli::parse();
|
||||
|
||||
if cli.exit {
|
||||
if let Ok(mut stream) = std::net::TcpStream::connect(std::env::var("MCP_PORT").map(|p| format!("127.0.0.1:{}", p)).unwrap_or_else(|_| "127.0.0.1:3000".to_string())) {
|
||||
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",
|
||||
);
|
||||
}
|
||||
let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string());
|
||||
let _ = std::process::Command::new("curl")
|
||||
.arg("-k")
|
||||
.arg("-X")
|
||||
.arg("POST")
|
||||
.arg(format!("https://127.0.0.1:{}/shutdown", port))
|
||||
.output();
|
||||
println!("Sent shutdown request to server.");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if cli.restart {
|
||||
if let Ok(mut stream) = std::net::TcpStream::connect(std::env::var("MCP_PORT").map(|p| format!("127.0.0.1:{}", p)).unwrap_or_else(|_| "127.0.0.1:3000".to_string())) {
|
||||
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));
|
||||
}
|
||||
let port = std::env::var("MCP_PORT").unwrap_or_else(|_| "3000".to_string());
|
||||
let _ = std::process::Command::new("curl")
|
||||
.arg("-k")
|
||||
.arg("-X")
|
||||
.arg("POST")
|
||||
.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(());
|
||||
}
|
||||
|
||||
@@ -720,8 +611,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
|
||||
let base_dir = std::env::var("MCP_MEMORY_STORE_DIR").unwrap_or_else(|_| {
|
||||
dirs::home_dir()
|
||||
.map(|mut h| {
|
||||
@@ -742,7 +632,8 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
{
|
||||
let mut table = write_txn.open_table(crate::store::STORE_TABLE).unwrap();
|
||||
|
||||
let stores = [
|
||||
let stores = vec![
|
||||
("knowledge_graph_master", "knowledge_graph_master.json"),
|
||||
("audit_ledger", "audit_ledger.json"),
|
||||
("sticky_notes", "sticky_notes.json"),
|
||||
("tasks", "tasks.json"),
|
||||
@@ -770,6 +661,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
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();
|
||||
let _ = fs::rename(&json_path, json_path.with_extension("json.migrated"));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -780,10 +672,8 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
}
|
||||
|
||||
let state = Arc::new(MemoryState {
|
||||
master_path: base.join("knowledge_graph_master.json"),
|
||||
session_graph: RwLock::new(KnowledgeGraph::default()),
|
||||
graph: Store::new("knowledge_graph_master", db.clone()),
|
||||
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()),
|
||||
@@ -803,19 +693,12 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
tech_debts: Store::new("tech_debts", db.clone()),
|
||||
gates: Store::new("gates", db.clone()),
|
||||
context_workspaces: Store::new("context_workspaces", db.clone()),
|
||||
activity_tx: tokio::sync::broadcast::channel(100).0,
|
||||
});
|
||||
|
||||
state.recover_wal();
|
||||
state.rebuild_index();
|
||||
|
||||
run_server(state)
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
// Linux no longer executes server logic natively due to workspace split
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in new issue
Block a user