chore: migrate store to redb

This commit is contained in:
Riz Ashraf committed 2026-09-10 10:39:33 +01:00
1 parent 7f0286dc40
commit 9a8b6e52e9
12 files changed
+965 -376

No files matched your search

Generated
+42 -1
View File
@@ -157,6 +157,15 @@ version = "0.22.1"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6"
[[package]]
name = "bincode"
version = "1.3.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b1f45e9417d87227c7a56d22e471c6206462cba514c7590c09aff4cf6d1ddcad"
dependencies = [
"serde",
]
[[package]] [[package]]
name = "bitflags" name = "bitflags"
version = "2.13.1" version = "2.13.1"
@@ -392,6 +401,20 @@ dependencies = [
"syn 3.0.5", "syn 3.0.5",
] ]
[[package]]
name = "dashmap"
version = "6.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e6361d5c062261c78a176addb82d4c821ae42bed6089de0e12603cd25de2059c"
dependencies = [
"cfg-if",
"crossbeam-utils",
"hashbrown 0.14.5",
"lock_api",
"once_cell",
"parking_lot_core",
]
[[package]] [[package]]
name = "datasketches" name = "datasketches"
version = "0.2.0" version = "0.2.0"
@@ -626,6 +649,12 @@ version = "0.3.4"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e4eba85ea1d0a966a983acd07deee566e67395d2d96b6fb39e62b5a833f1eb0b" checksum = "e4eba85ea1d0a966a983acd07deee566e67395d2d96b6fb39e62b5a833f1eb0b"
[[package]]
name = "hashbrown"
version = "0.14.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1"
[[package]] [[package]]
name = "hashbrown" name = "hashbrown"
version = "0.16.1" version = "0.16.1"
@@ -981,7 +1010,7 @@ version = "0.16.4"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f66e8d5d03f609abc3a39e6f08e4164ebf1447a732906d39eb9b99b7919ef39" checksum = "7f66e8d5d03f609abc3a39e6f08e4164ebf1447a732906d39eb9b99b7919ef39"
dependencies = [ dependencies = [
"hashbrown", "hashbrown 0.16.1",
] ]
[[package]] [[package]]
@@ -1008,10 +1037,13 @@ version = "0.1.0"
dependencies = [ dependencies = [
"async-trait", "async-trait",
"axum", "axum",
"bincode",
"clap", "clap",
"dashmap",
"dirs", "dirs",
"futures-util", "futures-util",
"glob", "glob",
"redb",
"reqwest", "reqwest",
"schemars", "schemars",
"serde", "serde",
@@ -1357,6 +1389,15 @@ dependencies = [
"crossbeam-utils", "crossbeam-utils",
] ]
[[package]]
name = "redb"
version = "4.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "de6c3b63e007e90ce536ec2ae4690826136a20ec8dbbbb400daef1bb999d2e36"
dependencies = [
"libc",
]
[[package]] [[package]]
name = "redox_syscall" name = "redox_syscall"
version = "0.5.18" version = "0.5.18"
+3
View File
@@ -6,10 +6,13 @@ edition = "2024"
[dependencies] [dependencies]
async-trait = "0.1.92" async-trait = "0.1.92"
axum = "0.8" axum = "0.8"
bincode = "1.3.3"
clap = { version = "4.6.6", features = ["derive"] } clap = { version = "4.6.6", features = ["derive"] }
dashmap = "6.2.1"
dirs = "6.0.0" dirs = "6.0.0"
futures-util = "0.3.34" futures-util = "0.3.34"
glob = "0.3.4" glob = "0.3.4"
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"] }
schemars = "1.2.2" schemars = "1.2.2"
serde = { version = "1.0.229", features = ["derive"] } serde = { version = "1.0.229", features = ["derive"] }
+2 -2
View File
@@ -5,13 +5,13 @@
<style> <style>
body { font-family: 'Segoe UI', Tahoma, Geneva, Verdana, sans-serif; padding: 40px; background-color: #f0f2f5; color: #333; } body { font-family: 'Segoe UI', Tahoma, Geneva, Verdana, sans-serif; padding: 40px; background-color: #f0f2f5; color: #333; }
h1 { color: #2c3e50; text-align: center; margin-bottom: 30px; } h1 { color: #2c3e50; text-align: center; margin-bottom: 30px; }
.grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(200px, 1fr)); gap: 20px; max-width: 800px; margin: 0 auto; } .grid { display: grid; grid-template-columns: repeat(4, 1fr); gap: 20px; max-width: 1000px; margin: 0 auto; }
.stat-card { background: white; padding: 20px; border-radius: 8px; box-shadow: 0 4px 6px rgba(0,0,0,0.05); text-align: center; transition: transform 0.2s ease; } .stat-card { background: white; padding: 20px; border-radius: 8px; box-shadow: 0 4px 6px rgba(0,0,0,0.05); text-align: center; transition: transform 0.2s ease; }
.stat-card:hover { transform: translateY(-5px); } .stat-card:hover { transform: translateY(-5px); }
.stat-value { font-size: 2.5em; font-weight: bold; color: #3498db; margin: 10px 0; } .stat-value { font-size: 2.5em; font-weight: bold; color: #3498db; margin: 10px 0; }
.stat-label { font-size: 1.1em; color: #7f8c8d; text-transform: uppercase; letter-spacing: 1px; } .stat-label { font-size: 1.1em; color: #7f8c8d; text-transform: uppercase; letter-spacing: 1px; }
.status-dot { display: inline-block; width: 10px; height: 10px; background-color: #2ecc71; border-radius: 50%; margin-right: 8px; box-shadow: 0 0 5px #2ecc71; } .status-dot { display: inline-block; width: 10px; height: 10px; background-color: #2ecc71; border-radius: 50%; margin-right: 8px; box-shadow: 0 0 5px #2ecc71; }
.header-bar { max-width: 800px; margin: 0 auto 20px; display: flex; justify-content: space-between; align-items: center; } .header-bar { max-width: 1000px; margin: 0 auto 20px; display: flex; justify-content: space-between; align-items: center; }
</style> </style>
</head> </head>
<body> <body>
+678 -282
View File
File diff suppressed because it is too large. Load diff
+143 -37
View File
@@ -1,10 +1,10 @@
mod handlers; mod handlers;
mod models;
mod mcp; mod mcp;
mod models;
mod search;
mod state; mod state;
mod store; mod store;
mod tools; mod tools;
mod search;
use crate::handlers::MemoryHandler; use crate::handlers::MemoryHandler;
use crate::models::*; use crate::models::*;
@@ -29,6 +29,10 @@ struct Cli {
target: Option<String>, target: Option<String>,
#[arg(long)] #[arg(long)]
daemon: bool, daemon: bool,
#[arg(long)]
exit: bool,
#[arg(long)]
restart: bool,
} }
#[derive(Subcommand)] #[derive(Subcommand)]
@@ -71,7 +75,6 @@ enum GateCommands {
}, },
} }
async fn reconcile_worker(state: Arc<MemoryState>) { async fn reconcile_worker(state: Arc<MemoryState>) {
loop { loop {
sleep(Duration::from_secs(5)).await; sleep(Duration::from_secs(5)).await;
@@ -104,21 +107,17 @@ async fn reconcile_worker(state: Arc<MemoryState>) {
} }
} }
use axum::{ use axum::{
extract::{State, Query}, Json, Router,
extract::{Query, State},
response::sse::{Event, Sse}, response::sse::{Event, Sse},
routing::{get, post}, routing::{get, post},
Json, Router,
}; };
use futures_util::stream::Stream; use futures_util::stream::Stream;
use std::convert::Infallible; use std::convert::Infallible;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::mpsc; use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream; use tokio_stream::wrappers::ReceiverStream;
use std::sync::atomic::{AtomicUsize, Ordering};
struct AppState { struct AppState {
handler: Arc<MemoryHandler>, handler: Arc<MemoryHandler>,
@@ -140,10 +139,23 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
.route("/sse", get(sse_handler)) .route("/sse", get(sse_handler))
.route("/messages", post(message_handler)) .route("/messages", post(message_handler))
.route("/health", get(health_handler)) .route("/health", get(health_handler))
.route("/", get(|| async move { .route(
axum::response::Html(include_str!("dashboard.html")) "/shutdown",
})) post(|| async move {
.route("/api/stats", get({ std::thread::spawn(|| {
std::thread::sleep(std::time::Duration::from_millis(100));
std::process::exit(0);
});
"Shutting down..."
}),
)
.route(
"/",
get(|| async move { axum::response::Html(include_str!("dashboard.html")) }),
)
.route(
"/api/stats",
get({
let state_clone = app_state.handler.state.clone(); let state_clone = app_state.handler.state.clone();
move || async move { move || async move {
let (entities, relations) = { let (entities, relations) = {
@@ -191,7 +203,8 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
"context_workspaces": context_workspaces "context_workspaces": context_workspaces
})) }))
} }
})) }),
)
.with_state(app_state); .with_state(app_state);
let mut retries = 0; let mut retries = 0;
@@ -199,16 +212,48 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
match tokio::net::TcpListener::bind("127.0.0.1:3000").await { match tokio::net::TcpListener::bind("127.0.0.1:3000").await {
Ok(l) => break l, Ok(l) => break l,
Err(e) => { Err(e) => {
// Check if it's already running and healthy
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::{Read, Write};
let _ = stream.write_all(
b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
let mut response = String::new();
let _ = stream.read_to_string(&mut response);
if response.contains("200 OK") {
// Already healthy! Just exit cleanly instead of panicking/retrying loop.
std::process::exit(0);
}
}
retries += 1; retries += 1;
if retries > 15 { if retries > 15 {
let log_path = dirs::home_dir().unwrap_or_default().join(".gemini/mcp_memory/daemon_fatal.log"); let log_path = dirs::home_dir()
let _ = std::fs::write(&log_path, format!("FATAL: Could not bind to 127.0.0.1:3000 after 15 seconds: {}\n", e)); .unwrap_or_default()
.join(".gemini/mcp_memory/daemon_fatal.log");
let _ = std::fs::write(
&log_path,
format!(
"FATAL: Could not bind to 127.0.0.1:3000 after 15 seconds: {}\n",
e
),
);
std::process::exit(1); std::process::exit(1);
} }
let log_path = dirs::home_dir().unwrap_or_default().join(".gemini/mcp_memory/daemon_error.log"); let log_path = dirs::home_dir()
if let Ok(mut file) = std::fs::OpenOptions::new().create(true).append(true).open(&log_path) { .unwrap_or_default()
.join(".gemini/mcp_memory/daemon_error.log");
if let Ok(mut file) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&log_path)
{
use std::io::Write; use std::io::Write;
let _ = writeln!(file, "Failed to bind to 127.0.0.1:3000 (attempt {}): {}. Retrying in 1s...", retries, e); let _ = writeln!(
file,
"Failed to bind to 127.0.0.1:3000 (attempt {}): {}. Retrying in 1s...",
retries, e
);
} }
tokio::time::sleep(std::time::Duration::from_secs(1)).await; tokio::time::sleep(std::time::Duration::from_secs(1)).await;
} }
@@ -216,7 +261,9 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
}; };
eprintln!("MCP Memory Server running on http://127.0.0.1:3000/sse"); eprintln!("MCP Memory Server running on http://127.0.0.1:3000/sse");
if let Err(e) = axum::serve(listener, app).await { if let Err(e) = axum::serve(listener, app).await {
let log_path = dirs::home_dir().unwrap_or_default().join(".gemini/mcp_memory/daemon_error.log"); let log_path = dirs::home_dir()
.unwrap_or_default()
.join(".gemini/mcp_memory/daemon_error.log");
let _ = std::fs::write(&log_path, format!("Server crashed: {}\n", e)); let _ = std::fs::write(&log_path, format!("Server crashed: {}\n", e));
} }
Ok(()) Ok(())
@@ -229,9 +276,17 @@ async fn sse_handler(
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst)); let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
let (tx, rx) = mpsc::channel::<Result<Event, Infallible>>(100); let (tx, rx) = mpsc::channel::<Result<Event, Infallible>>(100);
state.clients.write().unwrap().insert(session_id.clone(), tx.clone()); state
.clients
.write()
.unwrap()
.insert(session_id.clone(), tx.clone());
let _ = tx.send(Ok(Event::default().event("endpoint").data(format!("/messages?sessionId={}", session_id)))).await; let _ = tx
.send(Ok(Event::default()
.event("endpoint")
.data(format!("/messages?sessionId={}", session_id))))
.await;
let stream = ReceiverStream::new(rx); let stream = ReceiverStream::new(rx);
Sse::new(stream).keep_alive(axum::response::sse::KeepAlive::new()) Sse::new(stream).keep_alive(axum::response::sse::KeepAlive::new())
@@ -261,7 +316,9 @@ async fn message_handler(
let tx_opt = clients.clients.read().unwrap().get(&session_id).cloned(); let tx_opt = clients.clients.read().unwrap().get(&session_id).cloned();
if let Some(tx) = tx_opt { if let Some(tx) = tx_opt {
let data = serde_json::to_string(&response).unwrap(); let data = serde_json::to_string(&response).unwrap();
let _ = tx.send(Ok(Event::default().event("message").data(data))).await; let _ = tx
.send(Ok(Event::default().event("message").data(data)))
.await;
} }
} }
}); });
@@ -271,16 +328,38 @@ async fn message_handler(
mod proxy; mod proxy;
fn main() -> Result<(), Box<dyn std::error::Error>> { fn main() -> Result<(), Box<dyn std::error::Error>> {
let cli = Cli::parse(); let cli = Cli::parse();
if cli.exit {
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::Write;
let _ = stream.write_all(
b"POST /shutdown HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
}
println!("Sent shutdown request to server.");
return Ok(());
}
if cli.restart {
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::Write;
let _ = stream.write_all(
b"POST /shutdown HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
println!("Sent shutdown request to existing server. Waiting for it to exit...");
std::thread::sleep(std::time::Duration::from_millis(1500));
}
return Ok(());
}
#[cfg(target_os = "windows")] #[cfg(target_os = "windows")]
{ {
use std::os::windows::process::CommandExt; use std::os::windows::process::CommandExt;
if !cli.daemon { if !cli.daemon {
loop { loop {
if let Err(_) = std::net::TcpListener::bind("127.0.0.1:3000") { if std::net::TcpListener::bind("127.0.0.1:3000").is_err() {
// Port in use, become a stub proxy! // Port in use, become a stub proxy!
let target_url = cli.target.as_deref().unwrap_or("http://127.0.0.1:3000"); let target_url = cli.target.as_deref().unwrap_or("http://127.0.0.1:3000");
match proxy::run_proxy(target_url) { match proxy::run_proxy(target_url) {
@@ -293,7 +372,8 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
} }
} else { } else {
// Port is free. We must spawn the daemon, then loop again to become proxy // Port is free. We must spawn the daemon, then loop again to become proxy
std::process::Command::new(std::env::current_exe().unwrap()) #[allow(clippy::zombie_processes)]
let _ = std::process::Command::new(std::env::current_exe().unwrap())
.arg("--daemon") .arg("--daemon")
.stdin(std::process::Stdio::null()) .stdin(std::process::Stdio::null())
.stdout(std::process::Stdio::null()) .stdout(std::process::Stdio::null())
@@ -332,7 +412,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
session_graph: RwLock::new(KnowledgeGraph::default()), session_graph: RwLock::new(KnowledgeGraph::default()),
base_dir: base.clone(), base_dir: base.clone(),
master_cache: RwLock::new((KnowledgeGraph::default(), SystemTime::UNIX_EPOCH)), master_cache: RwLock::new((KnowledgeGraph::default(), SystemTime::UNIX_EPOCH)),
search_index: RwLock::new(crate::search::MemoryIndex::new().unwrap()), search_index: RwLock::new(crate::search::MemoryIndex::new(&base).unwrap()),
ledger: Store::new(base.join("audit_ledger.json")), ledger: Store::new(base.join("audit_ledger.json")),
sticky: Store::new(base.join("sticky_notes.json")), sticky: Store::new(base.join("sticky_notes.json")),
tasks: Store::new(base.join("tasks.json")), tasks: Store::new(base.join("tasks.json")),
@@ -357,10 +437,23 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
if let Some(command) = cli.command { if let Some(command) = cli.command {
match command { match command {
Commands::Gate { subcmd } => { Commands::Gate { subcmd } => match subcmd {
match subcmd { GateCommands::Set {
GateCommands::Set { action, target, namespace, params, authorize, block, reason } => { action,
let status = if authorize { "authorized".to_string() } else if block { "blocked".to_string() } else { "pending".to_string() }; target,
namespace,
params,
authorize,
block,
reason,
} => {
let status = if authorize {
"authorized".to_string()
} else if block {
"blocked".to_string()
} else {
"pending".to_string()
};
let mut param_map = HashMap::new(); let mut param_map = HashMap::new();
for p in params { for p in params {
if let Some((k, v)) = p.split_once('=') { if let Some((k, v)) = p.split_once('=') {
@@ -375,7 +468,10 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
params: param_map, params: param_map,
status, status,
reason, reason,
timestamp: SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_secs(), timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs(),
}; };
state.gates.modify(|gates| { state.gates.modify(|gates| {
gates.retain(|g| !(g.action == record.action && g.target == record.target)); gates.retain(|g| !(g.action == record.action && g.target == record.target));
@@ -384,7 +480,13 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
println!("Gate state updated."); println!("Gate state updated.");
std::process::exit(0); std::process::exit(0);
} }
GateCommands::Verify { action, target, namespace, params, consume } => { GateCommands::Verify {
action,
target,
namespace,
params,
consume,
} => {
let mut param_map = HashMap::new(); let mut param_map = HashMap::new();
for p in params { for p in params {
if let Some((k, v)) = p.split_once('=') { if let Some((k, v)) = p.split_once('=') {
@@ -394,7 +496,12 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let mut found = None; let mut found = None;
let mut to_remove = None; let mut to_remove = None;
state.gates.modify(|gates| { state.gates.modify(|gates| {
if let Some(idx) = gates.iter().position(|g| g.action == action && g.target == target && g.namespace == namespace && g.params == param_map) { if let Some(idx) = gates.iter().position(|g| {
g.action == action
&& g.target == target
&& g.namespace == namespace
&& g.params == param_map
}) {
found = Some(gates[idx].clone()); found = Some(gates[idx].clone());
if consume { if consume {
to_remove = Some(idx); to_remove = Some(idx);
@@ -424,8 +531,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
} }
} }
} }
} },
}
} }
} }
+1 -2
View File
@@ -165,8 +165,7 @@ pub struct ContextWorkspace {
pub saved_at: u64, pub saved_at: u64,
} }
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[derive(Default)]
pub struct GateRecord { pub struct GateRecord {
pub id: String, pub id: String,
pub action: String, pub action: String,
+23 -17
View File
@@ -1,9 +1,9 @@
use tokio_util::io::StreamReader;
use tokio::io::AsyncBufReadExt;
use futures_util::StreamExt; use futures_util::StreamExt;
use std::sync::Arc; use std::sync::Arc;
use tokio::io::AsyncBufReadExt;
use tokio::sync::RwLock; use tokio::sync::RwLock;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use tokio_util::io::StreamReader;
pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + Send + Sync>> { pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + Send + Sync>> {
let rt = tokio::runtime::Runtime::new().unwrap(); let rt = tokio::runtime::Runtime::new().unwrap();
@@ -16,7 +16,9 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
let mut handle = stdin.lock(); let mut handle = stdin.lock();
let mut buffer = String::new(); let mut buffer = String::new();
while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) { while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) {
if bytes == 0 { break; } if bytes == 0 {
break;
}
let _ = msg_tx.blocking_send(buffer.clone()); let _ = msg_tx.blocking_send(buffer.clone());
buffer.clear(); buffer.clear();
} }
@@ -34,11 +36,13 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
loop { loop {
let url = post_url_clone.read().await.clone(); let url = post_url_clone.read().await.clone();
if !url.is_empty() { if !url.is_empty() {
let res = client.post(&url) let res = client
.post(&url)
.header("Accept", "application/json, text/event-stream") .header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json") .header("Content-Type", "application/json")
.body(msg.clone()) .body(msg.clone())
.send().await; .send()
.await;
if res.is_ok() { if res.is_ok() {
break; break;
@@ -49,7 +53,6 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
} }
}); });
loop {
if shutdown_rx.try_recv().is_ok() { if shutdown_rx.try_recv().is_ok() {
return Ok(false); return Ok(false);
} }
@@ -57,13 +60,20 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
let sse_url = format!("{}/sse", target_url); let sse_url = format!("{}/sse", target_url);
let client = reqwest::Client::builder().build().unwrap(); let client = reqwest::Client::builder().build().unwrap();
match client.get(&sse_url).header("Accept", "text/event-stream").send().await { match client
.get(&sse_url)
.header("Accept", "text/event-stream")
.send()
.await
{
Ok(resp) => { Ok(resp) => {
if resp.status() == reqwest::StatusCode::GONE { if resp.status() == reqwest::StatusCode::GONE {
return Ok(true); return Ok(true);
} }
let stream = resp.bytes_stream().map(|res| res.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e))); let stream = resp.bytes_stream().map(|res| {
res.map_err(std::io::Error::other)
});
let mut reader = tokio::io::BufReader::new(StreamReader::new(stream)); let mut reader = tokio::io::BufReader::new(StreamReader::new(stream));
let mut line = String::new(); let mut line = String::new();
let mut is_message = false; let mut is_message = false;
@@ -85,14 +95,13 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
} else if trimmed.starts_with("event: endpoint") { } else if trimmed.starts_with("event: endpoint") {
is_endpoint = true; is_endpoint = true;
is_message = false; is_message = false;
} else if trimmed.starts_with("data: ") { } else if let Some(stripped) = trimmed.strip_prefix("data: ") {
if is_message { if is_message {
println!("{}", &trimmed[6..]); println!("{}", stripped);
is_message = false; is_message = false;
} else if is_endpoint { } else if is_endpoint {
let ep = &trimmed[6..];
let mut p = post_url.write().await; let mut p = post_url.write().await;
*p = format!("{}{}", target_url, ep); *p = format!("{}{}", target_url, stripped);
is_endpoint = false; is_endpoint = false;
} }
} }
@@ -104,12 +113,9 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
} }
} }
*post_url.write().await = String::new(); *post_url.write().await = String::new();
return Ok(true); Ok(true)
}
Err(_) => {
return Ok(true);
}
} }
Err(_) => Ok(true),
} }
}) })
} }
+37 -14
View File
@@ -1,7 +1,7 @@
use tantivy::schema::*; use crate::models::{Adr, Entity, Snippet, Task};
use tantivy::{doc, Index, IndexWriter, IndexReader, ReloadPolicy};
use std::sync::Mutex; use std::sync::Mutex;
use crate::models::{Entity, Task, Adr, Snippet}; use tantivy::schema::*;
use tantivy::{Index, IndexReader, IndexWriter, ReloadPolicy, doc};
pub struct MemoryIndex { pub struct MemoryIndex {
index: Index, index: Index,
@@ -17,7 +17,7 @@ pub struct MemoryIndex {
} }
impl MemoryIndex { impl MemoryIndex {
pub fn new() -> tantivy::Result<Self> { pub fn new(store_dir: &std::path::Path) -> tantivy::Result<Self> {
let mut schema_builder = Schema::builder(); let mut schema_builder = Schema::builder();
let id_field = schema_builder.add_text_field("id", STRING | STORED); let id_field = schema_builder.add_text_field("id", STRING | STORED);
let title_field = schema_builder.add_text_field("title", TEXT | STORED); let title_field = schema_builder.add_text_field("title", TEXT | STORED);
@@ -26,7 +26,10 @@ impl MemoryIndex {
let namespace_field = schema_builder.add_text_field("namespace", STRING | STORED); let namespace_field = schema_builder.add_text_field("namespace", STRING | STORED);
let schema = schema_builder.build(); let schema = schema_builder.build();
let index = Index::create_in_ram(schema.clone()); let index_dir = store_dir.join("tantivy_index");
std::fs::create_dir_all(&index_dir).unwrap();
let index = Index::open_in_dir(&index_dir).unwrap_or_else(|_| Index::create_in_dir(&index_dir, schema.clone()).unwrap());
let writer = index.writer(50_000_000)?; let writer = index.writer(50_000_000)?;
let reader = index let reader = index
.reader_builder() .reader_builder()
@@ -71,23 +74,43 @@ impl MemoryIndex {
Ok(()) Ok(())
} }
pub fn search(&self, query: &str, namespace: Option<&str>) -> tantivy::Result<Vec<(String, String)>> { pub fn search(
&self,
query: &str,
namespace: Option<&str>,
) -> tantivy::Result<Vec<(String, String)>> {
let searcher = self.reader.searcher(); let searcher = self.reader.searcher();
let query_parser = tantivy::query::QueryParser::for_index(&self.index, vec![self.title_field, self.body_field]); let query_parser = tantivy::query::QueryParser::for_index(
&self.index,
vec![self.title_field, self.body_field],
);
let q = query_parser.parse_query(query)?; let q = query_parser.parse_query(query)?;
let top_docs = searcher.search(&q, &tantivy::collector::TopDocs::with_limit(50).order_by_score())?; let top_docs = searcher.search(
&q,
&tantivy::collector::TopDocs::with_limit(50).order_by_score(),
)?;
let mut results = Vec::new(); let mut results = Vec::new();
for (_score, doc_address) in top_docs { for (_score, doc_address) in top_docs {
let retrieved_doc = searcher.doc::<tantivy::TantivyDocument>(doc_address)?; let retrieved_doc = searcher.doc::<tantivy::TantivyDocument>(doc_address)?;
let id = retrieved_doc.get_first(self.id_field).and_then(|v| v.as_str()).unwrap_or("").to_string(); let id = retrieved_doc
let doc_type = retrieved_doc.get_first(self.type_field).and_then(|v| v.as_str()).unwrap_or("").to_string(); .get_first(self.id_field)
let doc_ns = retrieved_doc.get_first(self.namespace_field).and_then(|v| v.as_str()).unwrap_or(""); .and_then(|v| v.as_str())
if let Some(ns) = namespace { .unwrap_or("")
if doc_ns != ns && doc_ns != "global" { .to_string();
let doc_type = retrieved_doc
.get_first(self.type_field)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let doc_ns = retrieved_doc
.get_first(self.namespace_field)
.and_then(|v| v.as_str())
.unwrap_or("");
if let Some(ns) = namespace
&& doc_ns != ns && doc_ns != "global" {
continue; continue;
} }
}
results.push((id, doc_type)); results.push((id, doc_type));
} }
Ok(results) Ok(results)
+8 -7
View File
@@ -1,6 +1,6 @@
use crate::models::*; use crate::models::*;
use crate::store::{Store, read_json_file, write_json_atomic};
use crate::search::MemoryIndex; use crate::search::MemoryIndex;
use crate::store::{Store, read_json_file, write_json_atomic};
use std::collections::{HashMap, HashSet}; use std::collections::{HashMap, HashSet};
use std::fs; use std::fs;
use std::path::PathBuf; use std::path::PathBuf;
@@ -100,13 +100,16 @@ impl MemoryState {
let mut session_graph = self.session_graph.write().unwrap(); let mut session_graph = self.session_graph.write().unwrap();
update_fn(&mut session_graph); update_fn(&mut session_graph);
let wal_path = self.base_dir.join("wal.jsonl"); let wal_path = self.base_dir.join("wal.jsonl");
if let Ok(payload) = serde_json::to_string(&*session_graph) { if let Ok(payload) = serde_json::to_string(&*session_graph)
if let Ok(mut file) = std::fs::OpenOptions::new().create(true).append(true).open(&wal_path) { && let Ok(mut file) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&wal_path)
{
use std::io::Write; use std::io::Write;
let _ = writeln!(file, "{}", payload); let _ = writeln!(file, "{}", payload);
} }
} }
}
pub fn apply_sync_write<F: FnOnce(&mut KnowledgeGraph)>(&self, update_fn: F) { pub fn apply_sync_write<F: FnOnce(&mut KnowledgeGraph)>(&self, update_fn: F) {
let lock_path = self.base_dir.join("master.lock"); let lock_path = self.base_dir.join("master.lock");
let mut attempts = 0; let mut attempts = 0;
@@ -142,7 +145,7 @@ impl MemoryState {
pub fn rebuild_index(&self) { pub fn rebuild_index(&self) {
if let Ok(new_idx) = MemoryIndex::new() { if let Ok(new_idx) = MemoryIndex::new() {
let full = self.get_full_graph(); let full = self.get_full_graph();
for (_, e) in &full.entities { for e in full.entities.values() {
let _ = new_idx.index_entity(e); let _ = new_idx.index_entity(e);
} }
for t in self.tasks.read() { for t in self.tasks.read() {
@@ -160,5 +163,3 @@ impl MemoryState {
} }
} }
} }
+5 -3
View File
@@ -1,4 +1,4 @@
use serde::{de::DeserializeOwned, Serialize}; use serde::{Serialize, de::DeserializeOwned};
use std::fs; use std::fs;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::RwLock; use std::sync::RwLock;
@@ -6,12 +6,14 @@ use std::time::SystemTime;
pub fn read_json_file<T: DeserializeOwned + Default>(path: &Path) -> T { pub fn read_json_file<T: DeserializeOwned + Default>(path: &Path) -> T {
if let Ok(data) = fs::read(path) if let Ok(data) = fs::read(path)
&& let Ok(parsed) = serde_json::from_slice(&data) { && let Ok(parsed) = serde_json::from_slice(&data)
{
return parsed; return parsed;
} }
let bak_path = path.with_extension("json.bak"); let bak_path = path.with_extension("json.bak");
if let Ok(data) = fs::read(&bak_path) if let Ok(data) = fs::read(&bak_path)
&& let Ok(parsed) = serde_json::from_slice(&data) { && let Ok(parsed) = serde_json::from_slice(&data)
{
let _ = fs::write(path, data); let _ = fs::write(path, data);
return parsed; return parsed;
} }
+3 -1
View File
@@ -78,7 +78,9 @@ pub struct UpdateTaskStatusTool {
pub status: String, pub status: String,
} }
#[derive(Debug, Deserialize, Serialize, JsonSchema)] #[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct ListActiveTasksTool { pub git_branch: Option<String>, } pub struct ListActiveTasksTool {
pub git_branch: Option<String>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)] #[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct StoreSnippetTool { pub struct StoreSnippetTool {
pub name: String, pub name: String,
+20 -10
View File
@@ -3,8 +3,8 @@ use futures_util::StreamExt;
use std::sync::Arc; use std::sync::Arc;
use tokio::io::AsyncBufReadExt; use tokio::io::AsyncBufReadExt;
use tokio::sync::RwLock; use tokio::sync::RwLock;
use tokio_util::io::StreamReader;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use tokio_util::io::StreamReader;
#[derive(Parser)] #[derive(Parser)]
#[command(name = "mcp-memory-stub")] #[command(name = "mcp-memory-stub")]
@@ -27,7 +27,9 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let mut handle = stdin.lock(); let mut handle = stdin.lock();
let mut buffer = String::new(); let mut buffer = String::new();
while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) { while let Ok(bytes) = std::io::BufRead::read_line(&mut handle, &mut buffer) {
if bytes == 0 { break; } if bytes == 0 {
break;
}
let _ = msg_tx.blocking_send(buffer.clone()); let _ = msg_tx.blocking_send(buffer.clone());
buffer.clear(); buffer.clear();
} }
@@ -46,11 +48,13 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
loop { loop {
let url = post_url_clone.read().await.clone(); let url = post_url_clone.read().await.clone();
if !url.is_empty() { if !url.is_empty() {
let res = client.post(&url) let res = client
.post(&url)
.header("Accept", "application/json, text/event-stream") .header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json") .header("Content-Type", "application/json")
.body(msg.clone()) .body(msg.clone())
.send().await; .send()
.await;
if res.is_ok() { if res.is_ok() {
break; break;
@@ -74,14 +78,21 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let sse_url = format!("{}/sse", target_url); let sse_url = format!("{}/sse", target_url);
let client = reqwest::Client::builder().build()?; let client = reqwest::Client::builder().build()?;
match client.get(&sse_url).header("Accept", "text/event-stream").send().await { match client
.get(&sse_url)
.header("Accept", "text/event-stream")
.send()
.await
{
Ok(resp) => { Ok(resp) => {
if resp.status() == reqwest::StatusCode::GONE { if resp.status() == reqwest::StatusCode::GONE {
eprintln!("[PROXY] Target gone, exiting."); eprintln!("[PROXY] Target gone, exiting.");
break; break;
} }
let stream = resp.bytes_stream().map(|res| res.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e))); let stream = resp.bytes_stream().map(|res| {
res.map_err(std::io::Error::other)
});
let mut reader = tokio::io::BufReader::new(StreamReader::new(stream)); let mut reader = tokio::io::BufReader::new(StreamReader::new(stream));
let mut line = String::new(); let mut line = String::new();
let mut is_message = false; let mut is_message = false;
@@ -103,14 +114,13 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
} else if trimmed.starts_with("event: endpoint") { } else if trimmed.starts_with("event: endpoint") {
is_endpoint = true; is_endpoint = true;
is_message = false; is_message = false;
} else if trimmed.starts_with("data: ") { } else if let Some(stripped) = trimmed.strip_prefix("data: ") {
if is_message { if is_message {
println!("{}", &trimmed[6..]); println!("{}", stripped);
is_message = false; is_message = false;
} else if is_endpoint { } else if is_endpoint {
let ep = &trimmed[6..];
let mut p = post_url.write().await; let mut p = post_url.write().await;
*p = format!("{}{}", target_url, ep); *p = format!("{}{}", target_url, stripped);
is_endpoint = false; is_endpoint = false;
} }
} }