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
+1954 -1365

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"
checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6"
[[package]]
name = "bincode"
version = "1.3.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b1f45e9417d87227c7a56d22e471c6206462cba514c7590c09aff4cf6d1ddcad"
dependencies = [
"serde",
]
[[package]]
name = "bitflags"
version = "2.13.1"
@@ -392,6 +401,20 @@ dependencies = [
"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]]
name = "datasketches"
version = "0.2.0"
@@ -626,6 +649,12 @@ version = "0.3.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e4eba85ea1d0a966a983acd07deee566e67395d2d96b6fb39e62b5a833f1eb0b"
[[package]]
name = "hashbrown"
version = "0.14.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1"
[[package]]
name = "hashbrown"
version = "0.16.1"
@@ -981,7 +1010,7 @@ version = "0.16.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f66e8d5d03f609abc3a39e6f08e4164ebf1447a732906d39eb9b99b7919ef39"
dependencies = [
"hashbrown",
"hashbrown 0.16.1",
]
[[package]]
@@ -1008,10 +1037,13 @@ version = "0.1.0"
dependencies = [
"async-trait",
"axum",
"bincode",
"clap",
"dashmap",
"dirs",
"futures-util",
"glob",
"redb",
"reqwest",
"schemars",
"serde",
@@ -1357,6 +1389,15 @@ dependencies = [
"crossbeam-utils",
]
[[package]]
name = "redb"
version = "4.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "de6c3b63e007e90ce536ec2ae4690826136a20ec8dbbbb400daef1bb999d2e36"
dependencies = [
"libc",
]
[[package]]
name = "redox_syscall"
version = "0.5.18"
+3
View File
@@ -6,10 +6,13 @@ edition = "2024"
[dependencies]
async-trait = "0.1.92"
axum = "0.8"
bincode = "1.3.3"
clap = { version = "4.6.6", features = ["derive"] }
dashmap = "6.2.1"
dirs = "6.0.0"
futures-util = "0.3.34"
glob = "0.3.4"
redb = "4.2.0"
reqwest = { version = "0.12", default-features = false, features = ["stream", "rustls-tls"] }
schemars = "1.2.2"
serde = { version = "1.0.229", features = ["derive"] }
+2 -2
View File
@@ -5,13 +5,13 @@
<style>
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; }
.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:hover { transform: translateY(-5px); }
.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; }
.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>
</head>
<body>
+1509 -1113
View File
File diff suppressed because it is too large. Load diff
+244 -138
View File
@@ -1,10 +1,10 @@
mod handlers;
mod models;
mod mcp;
mod models;
mod search;
mod state;
mod store;
mod tools;
mod search;
use crate::handlers::MemoryHandler;
use crate::models::*;
@@ -29,6 +29,10 @@ struct Cli {
target: Option<String>,
#[arg(long)]
daemon: bool,
#[arg(long)]
exit: bool,
#[arg(long)]
restart: bool,
}
#[derive(Subcommand)]
@@ -71,7 +75,6 @@ enum GateCommands {
},
}
async fn reconcile_worker(state: Arc<MemoryState>) {
loop {
sleep(Duration::from_secs(5)).await;
@@ -104,21 +107,17 @@ async fn reconcile_worker(state: Arc<MemoryState>) {
}
}
use axum::{
extract::{State, Query},
Json, Router,
extract::{Query, State},
response::sse::{Event, Sse},
routing::{get, post},
Json, Router,
};
use futures_util::stream::Stream;
use std::convert::Infallible;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use std::sync::atomic::{AtomicUsize, Ordering};
struct AppState {
handler: Arc<MemoryHandler>,
@@ -140,58 +139,72 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
.route("/sse", get(sse_handler))
.route("/messages", post(message_handler))
.route("/health", get(health_handler))
.route("/", get(|| async move {
axum::response::Html(include_str!("dashboard.html"))
}))
.route("/api/stats", get({
let state_clone = app_state.handler.state.clone();
move || async move {
let (entities, relations) = {
let graph = state_clone.get_full_graph();
(graph.entities.len(), graph.relations.len())
};
let tasks = state_clone.tasks.read().len();
let snippets = state_clone.snippets.read().len();
let tech_debts = state_clone.tech_debts.read().len();
let adrs = state_clone.adrs.read().len();
let ledger = state_clone.ledger.read().len();
let sticky = state_clone.sticky.read().len();
let error_fixes = state_clone.error_fixes.read().len();
let pinned_files = state_clone.pinned_files.read().len();
let session_summaries = state_clone.session_summaries.read().len();
let handoff_memos = state_clone.handoff_memos.read().len();
let env_fingerprints = state_clone.env_fingerprints.read().len();
let env_requirements = state_clone.env_requirements.read().len();
let milestones = state_clone.milestones.read().len();
let environments = state_clone.environments.read().len();
let pr_checklists = state_clone.pr_checklists.read().len();
let gates = state_clone.gates.read().len();
let context_workspaces = state_clone.context_workspaces.read().len();
.route(
"/shutdown",
post(|| async move {
std::thread::spawn(|| {
std::thread::sleep(std::time::Duration::from_millis(100));
std::process::exit(0);
});
"Shutting down..."
}),
)
.route(
"/",
get(|| async move { axum::response::Html(include_str!("dashboard.html")) }),
)
.route(
"/api/stats",
get({
let state_clone = app_state.handler.state.clone();
move || async move {
let (entities, relations) = {
let graph = state_clone.get_full_graph();
(graph.entities.len(), graph.relations.len())
};
let tasks = state_clone.tasks.read().len();
let snippets = state_clone.snippets.read().len();
let tech_debts = state_clone.tech_debts.read().len();
let adrs = state_clone.adrs.read().len();
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
}))
}
}))
let ledger = state_clone.ledger.read().len();
let sticky = state_clone.sticky.read().len();
let error_fixes = state_clone.error_fixes.read().len();
let pinned_files = state_clone.pinned_files.read().len();
let session_summaries = state_clone.session_summaries.read().len();
let handoff_memos = state_clone.handoff_memos.read().len();
let env_fingerprints = state_clone.env_fingerprints.read().len();
let env_requirements = state_clone.env_requirements.read().len();
let milestones = state_clone.milestones.read().len();
let environments = state_clone.environments.read().len();
let pr_checklists = state_clone.pr_checklists.read().len();
let gates = state_clone.gates.read().len();
let context_workspaces = state_clone.context_workspaces.read().len();
axum::Json(serde_json::json!({
"entities": entities,
"relations": relations,
"tasks": tasks,
"snippets": snippets,
"tech_debts": tech_debts,
"adrs": adrs,
"ledger": ledger,
"sticky": sticky,
"error_fixes": error_fixes,
"pinned_files": pinned_files,
"session_summaries": session_summaries,
"handoff_memos": handoff_memos,
"env_fingerprints": env_fingerprints,
"env_requirements": env_requirements,
"milestones": milestones,
"environments": environments,
"pr_checklists": pr_checklists,
"gates": gates,
"context_workspaces": context_workspaces
}))
}
}),
)
.with_state(app_state);
let mut retries = 0;
@@ -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 {
Ok(l) => break l,
Err(e) => {
// Check if it's already running and healthy
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::{Read, Write};
let _ = stream.write_all(
b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
let mut response = String::new();
let _ = stream.read_to_string(&mut response);
if response.contains("200 OK") {
// Already healthy! Just exit cleanly instead of panicking/retrying loop.
std::process::exit(0);
}
}
retries += 1;
if retries > 15 {
let log_path = dirs::home_dir().unwrap_or_default().join(".gemini/mcp_memory/daemon_fatal.log");
let _ = std::fs::write(&log_path, format!("FATAL: Could not bind to 127.0.0.1:3000 after 15 seconds: {}\n", e));
let log_path = dirs::home_dir()
.unwrap_or_default()
.join(".gemini/mcp_memory/daemon_fatal.log");
let _ = std::fs::write(
&log_path,
format!(
"FATAL: Could not bind to 127.0.0.1:3000 after 15 seconds: {}\n",
e
),
);
std::process::exit(1);
}
let log_path = dirs::home_dir().unwrap_or_default().join(".gemini/mcp_memory/daemon_error.log");
if let Ok(mut file) = std::fs::OpenOptions::new().create(true).append(true).open(&log_path) {
let log_path = dirs::home_dir()
.unwrap_or_default()
.join(".gemini/mcp_memory/daemon_error.log");
if let Ok(mut file) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&log_path)
{
use std::io::Write;
let _ = writeln!(file, "Failed to bind to 127.0.0.1:3000 (attempt {}): {}. Retrying in 1s...", retries, e);
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;
}
@@ -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");
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));
}
Ok(())
@@ -228,11 +275,19 @@ async fn sse_handler(
) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
let (tx, rx) = mpsc::channel::<Result<Event, Infallible>>(100);
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;
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 stream = ReceiverStream::new(rx);
Sse::new(stream).keep_alive(axum::response::sse::KeepAlive::new())
}
@@ -255,32 +310,56 @@ async fn message_handler(
let handler = Arc::clone(&state.handler);
let session_id = query.session_id.clone();
let clients = Arc::clone(&state);
tokio::spawn(async move {
if let Some(response) = handler.handle_request(payload).await {
let tx_opt = clients.clients.read().unwrap().get(&session_id).cloned();
if let Some(tx) = tx_opt {
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;
}
}
});
axum::http::StatusCode::ACCEPTED
}
mod proxy;
fn main() -> Result<(), Box<dyn std::error::Error>> {
let cli = Cli::parse();
if cli.exit {
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::Write;
let _ = stream.write_all(
b"POST /shutdown HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
}
println!("Sent shutdown request to server.");
return Ok(());
}
if cli.restart {
if let Ok(mut stream) = std::net::TcpStream::connect("127.0.0.1:3000") {
use std::io::Write;
let _ = stream.write_all(
b"POST /shutdown HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n",
);
println!("Sent shutdown request to existing server. Waiting for it to exit...");
std::thread::sleep(std::time::Duration::from_millis(1500));
}
return Ok(());
}
#[cfg(target_os = "windows")]
{
use std::os::windows::process::CommandExt;
if !cli.daemon {
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!
let target_url = cli.target.as_deref().unwrap_or("http://127.0.0.1:3000");
match proxy::run_proxy(target_url) {
@@ -293,7 +372,8 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
}
} else {
// 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")
.stdin(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()),
base_dir: base.clone(),
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")),
sticky: Store::new(base.join("sticky_notes.json")),
tasks: Store::new(base.join("tasks.json")),
@@ -357,75 +437,101 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
if let Some(command) = cli.command {
match command {
Commands::Gate { subcmd } => {
match subcmd {
GateCommands::Set { action, target, namespace, params, authorize, block, reason } => {
let status = if authorize { "authorized".to_string() } else if block { "blocked".to_string() } else { "pending".to_string() };
let mut param_map = HashMap::new();
for p in params {
if let Some((k, v)) = p.split_once('=') {
param_map.insert(k.to_string(), v.to_string());
}
Commands::Gate { subcmd } => match subcmd {
GateCommands::Set {
action,
target,
namespace,
params,
authorize,
block,
reason,
} => {
let status = if authorize {
"authorized".to_string()
} else if block {
"blocked".to_string()
} else {
"pending".to_string()
};
let mut param_map = HashMap::new();
for p in params {
if let Some((k, v)) = p.split_once('=') {
param_map.insert(k.to_string(), v.to_string());
}
let record = GateRecord {
id: uuid::Uuid::new_v4().to_string(),
action: action.clone(),
target: target.clone(),
namespace,
params: param_map,
status,
reason,
timestamp: SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_secs(),
};
state.gates.modify(|gates| {
gates.retain(|g| !(g.action == record.action && g.target == record.target));
gates.push(record);
});
println!("Gate state updated.");
std::process::exit(0);
}
GateCommands::Verify { action, target, namespace, params, consume } => {
let mut param_map = HashMap::new();
for p in params {
if let Some((k, v)) = p.split_once('=') {
param_map.insert(k.to_string(), v.to_string());
let record = GateRecord {
id: uuid::Uuid::new_v4().to_string(),
action: action.clone(),
target: target.clone(),
namespace,
params: param_map,
status,
reason,
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs(),
};
state.gates.modify(|gates| {
gates.retain(|g| !(g.action == record.action && g.target == record.target));
gates.push(record);
});
println!("Gate state updated.");
std::process::exit(0);
}
GateCommands::Verify {
action,
target,
namespace,
params,
consume,
} => {
let mut param_map = HashMap::new();
for p in params {
if let Some((k, v)) = p.split_once('=') {
param_map.insert(k.to_string(), v.to_string());
}
}
let mut found = None;
let mut to_remove = None;
state.gates.modify(|gates| {
if let Some(idx) = gates.iter().position(|g| {
g.action == action
&& g.target == target
&& g.namespace == namespace
&& g.params == param_map
}) {
found = Some(gates[idx].clone());
if consume {
to_remove = Some(idx);
}
}
let mut found = None;
let mut to_remove = None;
state.gates.modify(|gates| {
if let Some(idx) = gates.iter().position(|g| g.action == action && g.target == target && g.namespace == namespace && g.params == param_map) {
found = Some(gates[idx].clone());
if consume {
to_remove = Some(idx);
}
}
if let Some(idx) = to_remove {
gates.remove(idx);
}
});
match found {
Some(record) => {
if record.status == "authorized" {
std::process::exit(0);
if let Some(idx) = to_remove {
gates.remove(idx);
}
});
match found {
Some(record) => {
if record.status == "authorized" {
std::process::exit(0);
} else {
if let Some(r) = record.reason {
eprintln!("❌ Action blocked. Reason: {}", r);
} else {
if let Some(r) = record.reason {
eprintln!("❌ Action blocked. Reason: {}", r);
} else {
eprintln!("❌ Action blocked.");
}
std::process::exit(1);
eprintln!("❌ Action blocked.");
}
std::process::exit(1);
}
None => {
eprintln!("❌ Action not yet authorized (no gate record found).");
std::process::exit(2);
}
}
None => {
eprintln!("❌ Action not yet authorized (no gate record found).");
std::process::exit(2);
}
}
}
}
},
}
}
+1 -2
View File
@@ -165,8 +165,7 @@ pub struct ContextWorkspace {
pub saved_at: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[derive(Default)]
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct GateRecord {
pub id: String,
pub action: String,
+65 -59
View File
@@ -1,9 +1,9 @@
use tokio_util::io::StreamReader;
use tokio::io::AsyncBufReadExt;
use futures_util::StreamExt;
use std::sync::Arc;
use tokio::io::AsyncBufReadExt;
use tokio::sync::RwLock;
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>> {
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 buffer = String::new();
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());
buffer.clear();
}
@@ -26,7 +28,7 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
let target_url = target_url.to_string();
let post_url = Arc::new(RwLock::new(String::new()));
let post_url_clone = Arc::clone(&post_url);
let client = reqwest::Client::builder().build().unwrap();
tokio::spawn(async move {
@@ -34,14 +36,16 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
loop {
let url = post_url_clone.read().await.clone();
if !url.is_empty() {
let res = client.post(&url)
let res = client
.post(&url)
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.body(msg.clone())
.send().await;
.send()
.await;
if res.is_ok() {
break;
break;
}
}
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
@@ -49,67 +53,69 @@ pub fn run_proxy(target_url: &str) -> Result<bool, Box<dyn std::error::Error + S
}
});
loop {
if shutdown_rx.try_recv().is_ok() {
return Ok(false);
}
if shutdown_rx.try_recv().is_ok() {
return Ok(false);
}
let sse_url = format!("{}/sse", target_url);
let client = reqwest::Client::builder().build().unwrap();
match client.get(&sse_url).header("Accept", "text/event-stream").send().await {
Ok(resp) => {
if resp.status() == reqwest::StatusCode::GONE {
return Ok(true);
}
let stream = resp.bytes_stream().map(|res| res.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e)));
let mut reader = tokio::io::BufReader::new(StreamReader::new(stream));
let mut line = String::new();
let mut is_message = false;
let mut is_endpoint = false;
loop {
tokio::select! {
_ = shutdown_rx.recv() => {
return Ok(false);
}
res = reader.read_line(&mut line) => {
match res {
Ok(bytes) => {
if bytes == 0 { break; }
let trimmed = line.trim();
if trimmed.starts_with("event: message") {
is_message = true;
is_endpoint = false;
} else if trimmed.starts_with("event: endpoint") {
is_endpoint = true;
let sse_url = format!("{}/sse", target_url);
let client = reqwest::Client::builder().build().unwrap();
match client
.get(&sse_url)
.header("Accept", "text/event-stream")
.send()
.await
{
Ok(resp) => {
if resp.status() == reqwest::StatusCode::GONE {
return Ok(true);
}
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 line = String::new();
let mut is_message = false;
let mut is_endpoint = false;
loop {
tokio::select! {
_ = shutdown_rx.recv() => {
return Ok(false);
}
res = reader.read_line(&mut line) => {
match res {
Ok(bytes) => {
if bytes == 0 { break; }
let trimmed = line.trim();
if trimmed.starts_with("event: message") {
is_message = true;
is_endpoint = false;
} else if trimmed.starts_with("event: endpoint") {
is_endpoint = true;
is_message = false;
} else if let Some(stripped) = trimmed.strip_prefix("data: ") {
if is_message {
println!("{}", stripped);
is_message = false;
} else if trimmed.starts_with("data: ") {
if is_message {
println!("{}", &trimmed[6..]);
is_message = false;
} else if is_endpoint {
let ep = &trimmed[6..];
let mut p = post_url.write().await;
*p = format!("{}{}", target_url, ep);
is_endpoint = false;
}
} else if is_endpoint {
let mut p = post_url.write().await;
*p = format!("{}{}", target_url, stripped);
is_endpoint = false;
}
line.clear();
}
Err(_) => break,
line.clear();
}
Err(_) => break,
}
}
}
*post_url.write().await = String::new();
return Ok(true);
}
Err(_) => {
return Ok(true);
}
*post_url.write().await = String::new();
Ok(true)
}
Err(_) => Ok(true),
}
})
}
+41 -18
View File
@@ -1,13 +1,13 @@
use tantivy::schema::*;
use tantivy::{doc, Index, IndexWriter, IndexReader, ReloadPolicy};
use crate::models::{Adr, Entity, Snippet, Task};
use std::sync::Mutex;
use crate::models::{Entity, Task, Adr, Snippet};
use tantivy::schema::*;
use tantivy::{Index, IndexReader, IndexWriter, ReloadPolicy, doc};
pub struct MemoryIndex {
index: Index,
reader: IndexReader,
writer: Mutex<IndexWriter>,
// Schema fields
pub id_field: Field,
pub title_field: Field,
@@ -17,7 +17,7 @@ pub struct 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 id_field = schema_builder.add_text_field("id", STRING | STORED);
let title_field = schema_builder.add_text_field("title", TEXT | STORED);
@@ -25,14 +25,17 @@ impl MemoryIndex {
let type_field = schema_builder.add_text_field("type", STRING | STORED);
let namespace_field = schema_builder.add_text_field("namespace", STRING | STORED);
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 reader = index
.reader_builder()
.reload_policy(ReloadPolicy::OnCommitWithDelay)
.try_into()?;
Ok(Self {
index,
reader,
@@ -71,23 +74,43 @@ impl MemoryIndex {
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 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 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();
for (_score, doc_address) in top_docs {
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 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 {
if doc_ns != ns && doc_ns != "global" {
let id = retrieved_doc
.get_first(self.id_field)
.and_then(|v| v.as_str())
.unwrap_or("")
.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;
}
}
results.push((id, doc_type));
}
Ok(results)
+8 -7
View File
@@ -1,6 +1,6 @@
use crate::models::*;
use crate::store::{Store, read_json_file, write_json_atomic};
use crate::search::MemoryIndex;
use crate::store::{Store, read_json_file, write_json_atomic};
use std::collections::{HashMap, HashSet};
use std::fs;
use std::path::PathBuf;
@@ -100,12 +100,15 @@ impl MemoryState {
let mut session_graph = self.session_graph.write().unwrap();
update_fn(&mut session_graph);
let wal_path = self.base_dir.join("wal.jsonl");
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) {
if let Ok(payload) = serde_json::to_string(&*session_graph)
&& let Ok(mut file) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&wal_path)
{
use std::io::Write;
let _ = writeln!(file, "{}", payload);
}
}
}
pub fn apply_sync_write<F: FnOnce(&mut KnowledgeGraph)>(&self, update_fn: F) {
let lock_path = self.base_dir.join("master.lock");
@@ -142,7 +145,7 @@ impl MemoryState {
pub fn rebuild_index(&self) {
if let Ok(new_idx) = MemoryIndex::new() {
let full = self.get_full_graph();
for (_, e) in &full.entities {
for e in full.entities.values() {
let _ = new_idx.index_entity(e);
}
for t in self.tasks.read() {
@@ -160,5 +163,3 @@ impl MemoryState {
}
}
}
+10 -8
View File
@@ -1,4 +1,4 @@
use serde::{de::DeserializeOwned, Serialize};
use serde::{Serialize, de::DeserializeOwned};
use std::fs;
use std::path::{Path, PathBuf};
use std::sync::RwLock;
@@ -6,15 +6,17 @@ use std::time::SystemTime;
pub fn read_json_file<T: DeserializeOwned + Default>(path: &Path) -> T {
if let Ok(data) = fs::read(path)
&& let Ok(parsed) = serde_json::from_slice(&data) {
return parsed;
}
&& let Ok(parsed) = serde_json::from_slice(&data)
{
return parsed;
}
let bak_path = path.with_extension("json.bak");
if let Ok(data) = fs::read(&bak_path)
&& let Ok(parsed) = serde_json::from_slice(&data) {
let _ = fs::write(path, data);
return parsed;
}
&& let Ok(parsed) = serde_json::from_slice(&data)
{
let _ = fs::write(path, data);
return parsed;
}
T::default()
}
+3 -1
View File
@@ -78,7 +78,9 @@ pub struct UpdateTaskStatusTool {
pub status: String,
}
#[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)]
pub struct StoreSnippetTool {
pub name: String,
+26 -16
View File
@@ -3,8 +3,8 @@ use futures_util::StreamExt;
use std::sync::Arc;
use tokio::io::AsyncBufReadExt;
use tokio::sync::RwLock;
use tokio_util::io::StreamReader;
use tokio::sync::mpsc;
use tokio_util::io::StreamReader;
#[derive(Parser)]
#[command(name = "mcp-memory-stub")]
@@ -27,7 +27,9 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let mut handle = stdin.lock();
let mut buffer = String::new();
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());
buffer.clear();
}
@@ -37,7 +39,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let target_url = cli.target;
let post_url = Arc::new(RwLock::new(String::new()));
let post_url_clone = Arc::clone(&post_url);
let client = reqwest::Client::builder().build()?;
tokio::spawn(async move {
@@ -46,14 +48,16 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
loop {
let url = post_url_clone.read().await.clone();
if !url.is_empty() {
let res = client.post(&url)
let res = client
.post(&url)
.header("Accept", "application/json, text/event-stream")
.header("Content-Type", "application/json")
.body(msg.clone())
.send().await;
.send()
.await;
if res.is_ok() {
break;
break;
}
}
tokio::time::sleep(tokio::time::Duration::from_millis(500)).await;
@@ -73,20 +77,27 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
let sse_url = format!("{}/sse", target_url);
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) => {
if resp.status() == reqwest::StatusCode::GONE {
eprintln!("[PROXY] Target gone, exiting.");
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 line = String::new();
let mut is_message = false;
let mut is_endpoint = false;
loop {
tokio::select! {
_ = shutdown_rx.recv() => {
@@ -103,14 +114,13 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
} else if trimmed.starts_with("event: endpoint") {
is_endpoint = true;
is_message = false;
} else if trimmed.starts_with("data: ") {
} else if let Some(stripped) = trimmed.strip_prefix("data: ") {
if is_message {
println!("{}", &trimmed[6..]);
println!("{}", stripped);
is_message = false;
} else if is_endpoint {
let ep = &trimmed[6..];
let mut p = post_url.write().await;
*p = format!("{}{}", target_url, ep);
*p = format!("{}{}", target_url, stripped);
is_endpoint = false;
}
}