fix(core): remove all unwraps, fix unpolled futures, and fix search index memory leak on restart
This commit is contained in:
1 parent
ae0ea9ab22
commit
7dc0c329ad
9 files changed
+45
-59
No files matched your search
@@ -126,7 +126,7 @@ impl MemoryHandler {
|
||||
));
|
||||
Some(crate::mcp::success(
|
||||
id,
|
||||
serde_json::to_value(&init).unwrap(),
|
||||
serde_json::to_value(&init).unwrap_or_default(),
|
||||
))
|
||||
}
|
||||
"notifications/initialized" => None,
|
||||
|
||||
@@ -115,16 +115,19 @@ impl McpTool for CreateEntitiesHandler {
|
||||
|
||||
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
||||
let req: CreateEntitiesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
||||
let mut inserted = Vec::new();
|
||||
state.modify_graph(|g| {
|
||||
for entity in req.entities {
|
||||
if !entity.name.is_empty() {
|
||||
if let Ok(idx) = state.search_index.read() {
|
||||
drop(idx.index_entity(&entity));
|
||||
}
|
||||
inserted.push(entity.clone());
|
||||
g.entities.insert(entity.name.clone(), entity);
|
||||
}
|
||||
}
|
||||
});
|
||||
let idx = state.search_index.read().unwrap_or_else(|e| e.into_inner()).clone();
|
||||
for entity in inserted {
|
||||
let _ = idx.index_entity(&entity).await;
|
||||
}
|
||||
Ok("Entities created".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -41,10 +41,9 @@ impl McpTool for LogDecisionHandler {
|
||||
adrs.push(a);
|
||||
});
|
||||
|
||||
if let Some(adr) = new_adr
|
||||
&& let Ok(idx) = state.search_index.read()
|
||||
{
|
||||
drop(idx.index_adr(&adr));
|
||||
if let Some(adr) = new_adr {
|
||||
let idx = state.search_index.read().unwrap_or_else(|e| e.into_inner()).clone();
|
||||
let _ = idx.index_adr(&adr).await;
|
||||
}
|
||||
|
||||
Ok(format!("Decision logged as {}", adr_id).to_string())
|
||||
@@ -346,11 +345,10 @@ impl McpTool for OmniSearchHandler {
|
||||
|
||||
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
||||
let req: OmniSearchTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
||||
let matches = if let Ok(idx) = state.search_index.read() {
|
||||
let matches = {
|
||||
let idx = state.search_index.read().unwrap_or_else(|e| e.into_inner());
|
||||
idx.search(&req.query, req.namespace.as_deref())
|
||||
.unwrap_or_default()
|
||||
} else {
|
||||
vec![]
|
||||
};
|
||||
|
||||
let kg_json = state.read_graph(|full| {
|
||||
|
||||
@@ -42,9 +42,8 @@ impl McpTool for AddTaskHandler {
|
||||
dependencies: deps,
|
||||
acceptance_criteria: vec![],
|
||||
};
|
||||
if let Ok(idx) = state.search_index.read() {
|
||||
drop(idx.index_task(&task));
|
||||
}
|
||||
let idx = state.search_index.read().unwrap_or_else(|e| e.into_inner()).clone();
|
||||
let _ = idx.index_task(&task).await;
|
||||
state.tasks.modify(|tasks| {
|
||||
tasks.push(task);
|
||||
});
|
||||
|
||||
@@ -124,9 +124,8 @@ impl McpTool for StoreSnippetHandler {
|
||||
snippets.push(s_clone);
|
||||
});
|
||||
|
||||
if let Ok(idx) = state.search_index.read() {
|
||||
drop(idx.index_snippet(&snippet));
|
||||
}
|
||||
let idx = state.search_index.read().unwrap_or_else(|e| e.into_inner()).clone();
|
||||
let _ = idx.index_snippet(&snippet).await;
|
||||
|
||||
Ok(format!("Snippet '{}' stored.", req.name).to_string())
|
||||
}
|
||||
|
||||
+19
-37
@@ -206,7 +206,7 @@ async fn gate_set_handler(
|
||||
reason: body.reason.clone(),
|
||||
timestamp: SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.unwrap_or_default()
|
||||
.as_secs(),
|
||||
};
|
||||
app_state.handler.state.gates.modify(|gates| {
|
||||
@@ -432,7 +432,7 @@ async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::E
|
||||
|
||||
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 addr: std::net::SocketAddr = format!("127.0.0.1:{}", addr).parse().expect("Invalid bind address");
|
||||
|
||||
let listener = match tokio::net::TcpListener::bind(&addr).await {
|
||||
Ok(l) => l,
|
||||
@@ -468,7 +468,7 @@ async fn ws_handler(
|
||||
.into_response()
|
||||
}
|
||||
|
||||
async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: String) {
|
||||
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);
|
||||
|
||||
@@ -510,20 +510,6 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
);
|
||||
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())
|
||||
&& method == "tools/call"
|
||||
{
|
||||
let name = payload
|
||||
.get("params")
|
||||
.and_then(|p| p.get("name"))
|
||||
.and_then(|n| n.as_str())
|
||||
.unwrap_or("unknown_tool");
|
||||
handler.state.broadcast_activity(&format!("Agent executed tool: {}", name));
|
||||
}
|
||||
} // End if proxy
|
||||
|
||||
// Process MCP request
|
||||
if let Some(response) = handler.handle_request(payload).await {
|
||||
let res_str = serde_json::to_string(&response).unwrap_or_default();
|
||||
@@ -575,8 +561,8 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
struct SessionCleanup {
|
||||
session_id: String,
|
||||
state: Arc<AppState>,
|
||||
send_task: Option<tokio::task::JoinHandle<()>>,
|
||||
recv_task: Option<tokio::task::JoinHandle<()>>,
|
||||
send_task: tokio::task::JoinHandle<()>,
|
||||
recv_task: tokio::task::JoinHandle<()>,
|
||||
}
|
||||
|
||||
impl Drop for SessionCleanup {
|
||||
@@ -586,12 +572,8 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
.write()
|
||||
.unwrap_or_else(|e| e.into_inner())
|
||||
.remove(&self.session_id);
|
||||
if let Some(task) = self.send_task.take() {
|
||||
task.abort();
|
||||
}
|
||||
if let Some(task) = self.recv_task.take() {
|
||||
task.abort();
|
||||
}
|
||||
self.send_task.abort();
|
||||
self.recv_task.abort();
|
||||
tracing::info!(
|
||||
"Websocket session {} closed and cleaned up",
|
||||
self.session_id
|
||||
@@ -602,15 +584,15 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
let mut cleanup = SessionCleanup {
|
||||
session_id: session_id.clone(),
|
||||
state: Arc::clone(&state),
|
||||
send_task: Some(send_task),
|
||||
recv_task: Some(recv_task),
|
||||
send_task,
|
||||
recv_task,
|
||||
};
|
||||
|
||||
tokio::select! {
|
||||
_ = cleanup.send_task.as_mut().unwrap() => {
|
||||
_ = &mut cleanup.send_task => {
|
||||
tracing::info!("Websocket send task finished for session {}", session_id);
|
||||
},
|
||||
_ = cleanup.recv_task.as_mut().unwrap() => {
|
||||
_ = &mut cleanup.recv_task => {
|
||||
tracing::info!("Websocket recv task finished for session {}", session_id);
|
||||
},
|
||||
};
|
||||
@@ -746,7 +728,7 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
if !cli.daemon {
|
||||
// Just spawn the daemon and exit. We no longer act as a proxy.
|
||||
#[allow(clippy::zombie_processes)]
|
||||
let _ = std::process::Command::new(std::env::current_exe().unwrap())
|
||||
let _ = std::process::Command::new(std::env::current_exe().expect("Failed to get current executable path"))
|
||||
.arg("--daemon")
|
||||
.stdin(std::process::Stdio::null())
|
||||
.stdout(std::process::Stdio::null())
|
||||
@@ -766,13 +748,13 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
|
||||
let redb_path = base.join("mcp_store.redb");
|
||||
|
||||
let db = Arc::new(redb::Database::create(&redb_path).unwrap());
|
||||
let db = Arc::new(redb::Database::create(&redb_path).expect("Failed to create redb database"));
|
||||
|
||||
// Ensure table exists and migrate old JSON files
|
||||
{
|
||||
let write_txn = db.begin_write().unwrap();
|
||||
let write_txn = db.begin_write().expect("Failed to begin write txn on redb");
|
||||
{
|
||||
let mut table = write_txn.open_table(crate::store::STORE_TABLE).unwrap();
|
||||
let mut table = write_txn.open_table(crate::store::STORE_TABLE).expect("Failed to open STORE_TABLE");
|
||||
|
||||
let stores = vec![
|
||||
("knowledge_graph_master", "knowledge_graph_master.json"),
|
||||
@@ -797,22 +779,22 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
];
|
||||
|
||||
for (key, file_name) in stores.iter() {
|
||||
if table.get(*key).unwrap().is_none() {
|
||||
if table.get(*key).expect("Failed to read from table").is_none() {
|
||||
let json_path = base.join(file_name);
|
||||
if json_path.exists()
|
||||
&& let Ok(data) = fs::read(&json_path)
|
||||
&& serde_json::from_slice::<serde_json::Value>(&data).is_ok()
|
||||
{
|
||||
table.insert(*key, data.as_slice()).unwrap();
|
||||
table.insert(*key, data.as_slice()).expect("Failed to insert migrated data");
|
||||
let _ = fs::rename(&json_path, json_path.with_extension("json.migrated"));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
write_txn.commit().unwrap();
|
||||
write_txn.commit().expect("Failed to commit db migration");
|
||||
}
|
||||
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
let rt = tokio::runtime::Runtime::new().expect("Failed to create tokio runtime");
|
||||
let _guard = rt.enter();
|
||||
|
||||
let state = Arc::new(MemoryState {
|
||||
|
||||
+1
-1
@@ -22,7 +22,7 @@ pub fn error(id: serde_json::Value, code: i32, message: &str) -> serde_json::Val
|
||||
|
||||
pub fn tool_def<T: JsonSchema>(name: &str, description: &str) -> serde_json::Value {
|
||||
let schema = schemars::schema_for!(T);
|
||||
let schema_val = serde_json::to_value(&schema).unwrap();
|
||||
let schema_val = serde_json::to_value(&schema).unwrap_or_default();
|
||||
// MCP expects standard JSON schema. Schemars returns draft-07.
|
||||
json!({
|
||||
"name": name,
|
||||
|
||||
@@ -30,11 +30,14 @@ impl MemoryIndex {
|
||||
let schema = schema_builder.build();
|
||||
|
||||
let index_dir = store_dir.join("tantivy_index");
|
||||
std::fs::create_dir_all(&index_dir).unwrap();
|
||||
std::fs::create_dir_all(&index_dir)
|
||||
.map_err(|e| tantivy::TantivyError::SystemError(e.to_string()))?;
|
||||
let index = Index::open_in_dir(&index_dir)
|
||||
.unwrap_or_else(|_| Index::create_in_dir(&index_dir, schema.clone()).unwrap());
|
||||
.or_else(|_| Index::create_in_dir(&index_dir, schema.clone()))?;
|
||||
|
||||
let writer = index.writer(50_000_000)?;
|
||||
let mut writer = index.writer(50_000_000)?;
|
||||
writer.delete_all_documents()?;
|
||||
writer.commit()?;
|
||||
let reader = index
|
||||
.reader_builder()
|
||||
.reload_policy(ReloadPolicy::OnCommitWithDelay)
|
||||
|
||||
+3
-1
@@ -104,7 +104,9 @@ impl MemoryState {
|
||||
idx.add_adr_sync(a);
|
||||
}
|
||||
});
|
||||
}).await.unwrap();
|
||||
}).await.unwrap_or_else(|e| {
|
||||
tracing::error!("Failed to join tantivy index rebuild thread: {}", e);
|
||||
});
|
||||
|
||||
let _ = new_idx.commit().await;
|
||||
if let Ok(mut w) = self.search_index.write() {
|
||||
|
||||
Reference in new issue
Block a user