chore: rustfmt, clippy lints and code tidying
This commit is contained in:
1 parent
3716c3e698
commit
1752753fcc
19 files changed
+1273
-738
No files matched your search
@@ -1,8 +1,10 @@
|
||||
use rmcp::model::{InitializeResult, ServerCapabilities};
|
||||
|
||||
fn main() {
|
||||
let init = InitializeResult::new(
|
||||
ServerCapabilities::builder().enable_tools().build()
|
||||
).with_server_info(rmcp::model::Implementation::new("gemini-mcp-memory", "3.0.0"));
|
||||
let init = InitializeResult::new(ServerCapabilities::builder().enable_tools().build())
|
||||
.with_server_info(rmcp::model::Implementation::new(
|
||||
"gemini-mcp-memory",
|
||||
"3.0.0",
|
||||
));
|
||||
println!("{}", serde_json::to_string_pretty(&init).unwrap());
|
||||
}
|
||||
+533
-218
File diff suppressed because it is too large.
Load diff
+299
-235
@@ -1,4 +1,7 @@
|
||||
#![cfg_attr(not(target_os = "windows"), allow(dead_code, unused_imports, unreachable_code))]
|
||||
#![cfg_attr(
|
||||
not(target_os = "windows"),
|
||||
allow(dead_code, unused_imports, unreachable_code)
|
||||
)]
|
||||
|
||||
mod handlers;
|
||||
mod mcp;
|
||||
@@ -13,11 +16,11 @@ use crate::models::*;
|
||||
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, SystemTime, UNIX_EPOCH};
|
||||
use redb::ReadableTable;
|
||||
use tokio::time::sleep;
|
||||
|
||||
use clap::{Parser, Subcommand};
|
||||
@@ -83,11 +86,10 @@ enum GateCommands {
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async fn reconcile_worker(state: Arc<MemoryState>) {
|
||||
loop {
|
||||
sleep(Duration::from_secs(5)).await;
|
||||
|
||||
|
||||
let has_local = {
|
||||
let session = state.graph.read();
|
||||
!session.entities.is_empty() || !session.relations.is_empty()
|
||||
@@ -106,14 +108,18 @@ async fn reconcile_worker(state: Arc<MemoryState>) {
|
||||
let state_clone = state.clone();
|
||||
let _ = tokio::task::spawn_blocking(move || {
|
||||
state_clone.rebuild_index();
|
||||
}).await;
|
||||
})
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
use axum::{
|
||||
Json, Router,
|
||||
extract::{Query, State, ws::{WebSocket, Message}},
|
||||
extract::{
|
||||
Query, State,
|
||||
ws::{Message, WebSocket},
|
||||
},
|
||||
response::IntoResponse,
|
||||
routing::{get, post},
|
||||
};
|
||||
@@ -186,9 +192,11 @@ async fn gate_verify_handler(
|
||||
(axum::http::StatusCode::FORBIDDEN, msg).into_response()
|
||||
}
|
||||
}
|
||||
None => {
|
||||
(axum::http::StatusCode::NOT_FOUND, "Action not yet authorized (no gate record found).").into_response()
|
||||
}
|
||||
None => (
|
||||
axum::http::StatusCode::NOT_FOUND,
|
||||
"Action not yet authorized (no gate record found).",
|
||||
)
|
||||
.into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -229,18 +237,23 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
rt.block_on(async {
|
||||
tokio::spawn(reconcile_worker(Arc::clone(&state)));
|
||||
let app_state = Arc::new(AppState {
|
||||
handler: Arc::new(MemoryHandler { state: Arc::clone(&state) }),
|
||||
handler: Arc::new(MemoryHandler {
|
||||
state: Arc::clone(&state),
|
||||
}),
|
||||
clients: RwLock::new(HashMap::new()),
|
||||
next_id: AtomicUsize::new(1),
|
||||
});
|
||||
|
||||
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(
|
||||
"/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))
|
||||
@@ -260,65 +273,81 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
"/",
|
||||
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 = state_clone.get_full_graph();
|
||||
axum::Json(graph)
|
||||
}
|
||||
}))
|
||||
|
||||
.route("/api/tasks/{id}/complete", post({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move |axum::extract::Path(id): axum::extract::Path<String>| 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 = state_clone.tasks.read();
|
||||
axum::Json(tasks.clone())
|
||||
}
|
||||
}))
|
||||
.route("/api/sticky", get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move || async move {
|
||||
let sticky = state_clone.sticky.read();
|
||||
axum::Json(sticky.clone())
|
||||
}
|
||||
}))
|
||||
.route("/api/search", get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move |axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>| async move {
|
||||
if let Some(q) = params.get("q") {
|
||||
if let Ok(idx) = state_clone.search_index.read() {
|
||||
if let Ok(results) = idx.search(q, None) {
|
||||
let mut formatted_results = Vec::new();
|
||||
for (type_name, content) in results {
|
||||
formatted_results.push(serde_json::json!({
|
||||
"type_name": type_name,
|
||||
"content": content,
|
||||
"score": 1.0
|
||||
}));
|
||||
}
|
||||
return axum::Json(serde_json::json!({ "results": formatted_results }));
|
||||
}
|
||||
}
|
||||
.route(
|
||||
"/api/graph",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move || async move {
|
||||
let graph = state_clone.get_full_graph();
|
||||
axum::Json(graph)
|
||||
}
|
||||
axum::Json(serde_json::json!({ "results": [] }))
|
||||
}
|
||||
}))
|
||||
.route("/api/stats",
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/tasks/{id}/complete",
|
||||
post({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move |axum::extract::Path(id): axum::extract::Path<String>| 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 = state_clone.tasks.read();
|
||||
axum::Json(tasks.clone())
|
||||
}
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/sticky",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move || async move {
|
||||
let sticky = state_clone.sticky.read();
|
||||
axum::Json(sticky.clone())
|
||||
}
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/search",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move |axum::extract::Query(params): axum::extract::Query<
|
||||
std::collections::HashMap<String, String>,
|
||||
>| 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 (type_name, content) in results {
|
||||
formatted_results.push(serde_json::json!({
|
||||
"type_name": type_name,
|
||||
"content": content,
|
||||
"score": 1.0
|
||||
}));
|
||||
}
|
||||
return axum::Json(
|
||||
serde_json::json!({ "results": formatted_results }),
|
||||
);
|
||||
}
|
||||
axum::Json(serde_json::json!({ "results": [] }))
|
||||
}
|
||||
}),
|
||||
)
|
||||
.route(
|
||||
"/api/stats",
|
||||
get({
|
||||
let state_clone = app_state.handler.state.clone();
|
||||
move || async move {
|
||||
@@ -374,7 +403,7 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
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();
|
||||
|
||||
|
||||
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()
|
||||
@@ -386,29 +415,39 @@ fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::Error>>
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
|
||||
async fn ws_handler(
|
||||
ws: axum::extract::ws::WebSocketUpgrade,
|
||||
headers: axum::http::HeaderMap,
|
||||
_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)).into_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<AppState>, client_type: String) {
|
||||
let session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
|
||||
let (tx, mut rx) = mpsc::channel::<String>(100);
|
||||
|
||||
state.clients.write().unwrap().insert(session_id.clone(), tx.clone());
|
||||
state
|
||||
.clients
|
||||
.write()
|
||||
.unwrap()
|
||||
.insert(session_id.clone(), tx.clone());
|
||||
|
||||
let (mut sender, mut receiver) = socket.split();
|
||||
|
||||
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);
|
||||
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;
|
||||
@@ -426,75 +465,102 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
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::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);
|
||||
}
|
||||
} // 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);
|
||||
}
|
||||
if let Some(method) = payload.get("method").and_then(|m| m.as_str())
|
||||
&& 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
|
||||
);
|
||||
}
|
||||
} // 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)]
|
||||
pub struct NvimTelemetry {
|
||||
@@ -510,11 +576,13 @@ async fn nvim_telemetry_handler(
|
||||
axum::Json(payload): axum::Json<NvimTelemetry>,
|
||||
) -> 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());
|
||||
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 _ = std::fs::write(&win_path, &payload.session_id);
|
||||
|
||||
|
||||
let wsl_path = "\\\\wsl.localhost\\Ubuntu\\home\\riz\\.gemini\\active_nvim.txt";
|
||||
let _ = std::fs::write(wsl_path, &payload.session_id);
|
||||
}
|
||||
@@ -524,10 +592,10 @@ async fn nvim_telemetry_handler(
|
||||
"type": "nvim_telemetry",
|
||||
"data": payload
|
||||
});
|
||||
|
||||
|
||||
let msg_str = ws_msg.to_string();
|
||||
let clients = state.clients.read().unwrap().clone();
|
||||
for (_, tx) in clients.iter() {
|
||||
for tx in clients.values() {
|
||||
let _ = tx.send(msg_str.clone()).await;
|
||||
}
|
||||
|
||||
@@ -549,10 +617,10 @@ fn init_logging(app_name: &str) -> Option<tracing_appender::non_blocking::Worker
|
||||
});
|
||||
let log_dir = std::path::PathBuf::from(base_dir).join("logs");
|
||||
std::fs::create_dir_all(&log_dir).unwrap_or_default();
|
||||
|
||||
|
||||
let file_appender = tracing_appender::rolling::daily(log_dir, format!("{}.log", app_name));
|
||||
let (non_blocking, guard) = tracing_appender::non_blocking(file_appender);
|
||||
|
||||
|
||||
let _ = tracing_subscriber::fmt()
|
||||
.with_writer(non_blocking)
|
||||
.with_ansi(false)
|
||||
@@ -560,7 +628,7 @@ fn init_logging(app_name: &str) -> Option<tracing_appender::non_blocking::Worker
|
||||
.with_thread_ids(true)
|
||||
.with_thread_names(true)
|
||||
.try_init();
|
||||
|
||||
|
||||
Some(guard)
|
||||
}
|
||||
|
||||
@@ -611,96 +679,92 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
}
|
||||
}
|
||||
|
||||
let base_dir = std::env::var("MCP_MEMORY_STORE_DIR").unwrap_or_else(|_| {
|
||||
dirs::home_dir()
|
||||
.map(|mut h| {
|
||||
h.push(".gemini/mcp_memory");
|
||||
h.to_string_lossy().into_owned()
|
||||
})
|
||||
.unwrap_or_else(|| ".gemini/mcp_memory".into())
|
||||
});
|
||||
let base = PathBuf::from(base_dir);
|
||||
fs::create_dir_all(&base).expect("Failed to create store dir");
|
||||
|
||||
let base_dir = std::env::var("MCP_MEMORY_STORE_DIR").unwrap_or_else(|_| {
|
||||
dirs::home_dir()
|
||||
.map(|mut h| {
|
||||
h.push(".gemini/mcp_memory");
|
||||
h.to_string_lossy().into_owned()
|
||||
})
|
||||
.unwrap_or_else(|| ".gemini/mcp_memory".into())
|
||||
});
|
||||
let base = PathBuf::from(base_dir);
|
||||
fs::create_dir_all(&base).expect("Failed to create store dir");
|
||||
let redb_path = base.join("mcp_store.redb");
|
||||
let db = Arc::new(redb::Database::create(&redb_path).unwrap());
|
||||
|
||||
let redb_path = base.join("mcp_store.redb");
|
||||
let db = Arc::new(redb::Database::create(&redb_path).unwrap());
|
||||
|
||||
// Ensure table exists and migrate old JSON files
|
||||
// Ensure table exists and migrate old JSON files
|
||||
{
|
||||
let write_txn = db.begin_write().unwrap();
|
||||
{
|
||||
let write_txn = db.begin_write().unwrap();
|
||||
{
|
||||
let mut table = write_txn.open_table(crate::store::STORE_TABLE).unwrap();
|
||||
|
||||
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"),
|
||||
];
|
||||
let mut table = write_txn.open_table(crate::store::STORE_TABLE).unwrap();
|
||||
|
||||
for (key, file_name) in stores.iter() {
|
||||
if table.get(*key).unwrap().is_none() {
|
||||
let json_path = base.join(file_name);
|
||||
if json_path.exists() {
|
||||
if let Ok(data) = fs::read(&json_path) {
|
||||
if serde_json::from_slice::<serde_json::Value>(&data).is_ok() {
|
||||
table.insert(*key, data.as_slice()).unwrap();
|
||||
let _ = fs::rename(&json_path, json_path.with_extension("json.migrated"));
|
||||
}
|
||||
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).unwrap().is_none() {
|
||||
let json_path = base.join(file_name);
|
||||
if json_path.exists()
|
||||
&& let Ok(data) = fs::read(&json_path)
|
||||
&& 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"),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
write_txn.commit().unwrap();
|
||||
}
|
||||
write_txn.commit().unwrap();
|
||||
}
|
||||
|
||||
let state = Arc::new(MemoryState {
|
||||
graph: Store::new("knowledge_graph_master", db.clone()),
|
||||
base_dir: base.clone(),
|
||||
search_index: RwLock::new(crate::search::MemoryIndex::new(&base).unwrap()),
|
||||
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()),
|
||||
activity_tx: tokio::sync::broadcast::channel(100).0,
|
||||
});
|
||||
let state = Arc::new(MemoryState {
|
||||
graph: Store::new("knowledge_graph_master", db.clone()),
|
||||
base_dir: base.clone(),
|
||||
search_index: RwLock::new(crate::search::MemoryIndex::new(&base).unwrap()),
|
||||
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()),
|
||||
activity_tx: tokio::sync::broadcast::channel(100).0,
|
||||
});
|
||||
|
||||
state.rebuild_index();
|
||||
state.rebuild_index();
|
||||
|
||||
run_server(state)
|
||||
run_server(state)
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use schemars::JsonSchema;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use schemars::JsonSchema;
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct CodeChange {
|
||||
|
||||
+13
-8
@@ -28,7 +28,8 @@ impl MemoryIndex {
|
||||
|
||||
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 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
|
||||
@@ -72,7 +73,6 @@ impl MemoryIndex {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
pub fn commit(&self) -> tantivy::Result<()> {
|
||||
let mut writer = self.writer.lock().unwrap();
|
||||
writer.commit()?;
|
||||
@@ -113,9 +113,11 @@ impl MemoryIndex {
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("");
|
||||
if let Some(ns) = namespace
|
||||
&& doc_ns != ns && doc_ns != "global" {
|
||||
continue;
|
||||
}
|
||||
&& doc_ns != ns
|
||||
&& doc_ns != "global"
|
||||
{
|
||||
continue;
|
||||
}
|
||||
results.push((id, doc_type));
|
||||
}
|
||||
Ok(results)
|
||||
@@ -172,7 +174,10 @@ mod tests {
|
||||
status: "open".to_string(),
|
||||
created_at: 0,
|
||||
updated_at: 0,
|
||||
git_branch: None, acceptance_criteria: vec![], dependencies: vec![], parent_id: None,
|
||||
git_branch: None,
|
||||
acceptance_criteria: vec![],
|
||||
dependencies: vec![],
|
||||
parent_id: None,
|
||||
};
|
||||
index.index_task(&task).unwrap();
|
||||
|
||||
@@ -218,11 +223,11 @@ mod tests {
|
||||
fn test_search_malformed_query() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let index = MemoryIndex::new(temp_dir.path()).unwrap();
|
||||
|
||||
|
||||
// Malformed lucene query (unclosed parenthesis)
|
||||
let result = index.search("title: (unclosed", None);
|
||||
assert!(result.is_err());
|
||||
|
||||
|
||||
// Another malformed query (unclosed quote)
|
||||
let result2 = index.search("title: \"unclosed", None);
|
||||
assert!(result2.is_err());
|
||||
|
||||
+8
-6
@@ -33,19 +33,21 @@ pub struct MemoryState {
|
||||
impl MemoryState {
|
||||
pub fn unique_items<T: Eq + std::hash::Hash + Clone>(input: Vec<T>) -> Vec<T> {
|
||||
let mut keys = std::collections::HashSet::new();
|
||||
input.into_iter().filter(|entry| keys.insert(entry.clone())).collect()
|
||||
input
|
||||
.into_iter()
|
||||
.filter(|entry| keys.insert(entry.clone()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn broadcast_activity(&self, message: &str) {
|
||||
let payload = serde_json::json!({
|
||||
"type": "activity",
|
||||
"data": message
|
||||
}).to_string();
|
||||
})
|
||||
.to_string();
|
||||
let _ = self.activity_tx.send(payload);
|
||||
}
|
||||
|
||||
|
||||
|
||||
pub fn get_full_graph(&self) -> KnowledgeGraph {
|
||||
self.graph.read()
|
||||
}
|
||||
@@ -61,10 +63,10 @@ impl MemoryState {
|
||||
pub fn rebuild_index(&self) {
|
||||
if let Ok(new_idx) = MemoryIndex::new(&self.base_dir) {
|
||||
let graph = self.graph.read();
|
||||
for (_, e) in &graph.entities {
|
||||
for e in graph.entities.values() {
|
||||
let _ = new_idx.index_entity(e);
|
||||
}
|
||||
|
||||
|
||||
let tasks = self.tasks.read();
|
||||
for t in tasks {
|
||||
let _ = new_idx.index_task(&t);
|
||||
|
||||
+23
-13
@@ -22,13 +22,11 @@ impl<T: DeserializeOwned + Default + Serialize + Clone + Send + 'static> Store<T
|
||||
|
||||
fn load_from_db(key: &str, db: &Database) -> T {
|
||||
let read_txn = db.begin_read().unwrap();
|
||||
if let Ok(table) = read_txn.open_table(STORE_TABLE) {
|
||||
if let Ok(Some(value)) = table.get(key) {
|
||||
if let Ok(parsed) = serde_json::from_slice::<T>(value.value()) {
|
||||
if let Ok(table) = read_txn.open_table(STORE_TABLE)
|
||||
&& let Ok(Some(value)) = table.get(key)
|
||||
&& let Ok(parsed) = serde_json::from_slice::<T>(value.value()) {
|
||||
return parsed;
|
||||
}
|
||||
}
|
||||
}
|
||||
T::default()
|
||||
}
|
||||
|
||||
@@ -74,7 +72,7 @@ mod tests {
|
||||
async fn test_store_read_write() {
|
||||
let temp_file = NamedTempFile::new().unwrap();
|
||||
let db = Database::create(temp_file.path()).unwrap();
|
||||
|
||||
|
||||
let write_txn = db.begin_write().unwrap();
|
||||
{
|
||||
write_txn.open_table(STORE_TABLE).unwrap();
|
||||
@@ -94,18 +92,30 @@ mod tests {
|
||||
// Need to wait for spawn_blocking to finish
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
|
||||
assert_eq!(store.read(), TestData { name: "Hello".to_string(), value: 42 });
|
||||
assert_eq!(
|
||||
store.read(),
|
||||
TestData {
|
||||
name: "Hello".to_string(),
|
||||
value: 42
|
||||
}
|
||||
);
|
||||
|
||||
// Load again to verify persistence
|
||||
let store2 = Store::<TestData>::new("test_key", db.clone());
|
||||
assert_eq!(store2.read(), TestData { name: "Hello".to_string(), value: 42 });
|
||||
assert_eq!(
|
||||
store2.read(),
|
||||
TestData {
|
||||
name: "Hello".to_string(),
|
||||
value: 42
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
|
||||
async fn test_store_concurrency() {
|
||||
let temp_file = NamedTempFile::new().unwrap();
|
||||
let db = Database::create(temp_file.path()).unwrap();
|
||||
|
||||
|
||||
let write_txn = db.begin_write().unwrap();
|
||||
{
|
||||
write_txn.open_table(STORE_TABLE).unwrap();
|
||||
@@ -114,7 +124,7 @@ mod tests {
|
||||
|
||||
let db = Arc::new(db);
|
||||
let store = Arc::new(Store::<TestData>::new("concurrent_key", db.clone()));
|
||||
|
||||
|
||||
let mut handles = vec![];
|
||||
for _ in 0..50 {
|
||||
let s = store.clone();
|
||||
@@ -124,14 +134,14 @@ mod tests {
|
||||
});
|
||||
}));
|
||||
}
|
||||
|
||||
|
||||
for h in handles {
|
||||
h.await.unwrap();
|
||||
}
|
||||
|
||||
|
||||
// Wait for all blocking writes to flush
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(500)).await;
|
||||
|
||||
|
||||
assert_eq!(store.read().value, 50);
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user