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(
|
Some(crate::mcp::success(
|
||||||
id,
|
id,
|
||||||
serde_json::to_value(&init).unwrap(),
|
serde_json::to_value(&init).unwrap_or_default(),
|
||||||
))
|
))
|
||||||
}
|
}
|
||||||
"notifications/initialized" => None,
|
"notifications/initialized" => None,
|
||||||
|
|||||||
@@ -115,16 +115,19 @@ impl McpTool for CreateEntitiesHandler {
|
|||||||
|
|
||||||
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
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 req: CreateEntitiesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
||||||
|
let mut inserted = Vec::new();
|
||||||
state.modify_graph(|g| {
|
state.modify_graph(|g| {
|
||||||
for entity in req.entities {
|
for entity in req.entities {
|
||||||
if !entity.name.is_empty() {
|
if !entity.name.is_empty() {
|
||||||
if let Ok(idx) = state.search_index.read() {
|
inserted.push(entity.clone());
|
||||||
drop(idx.index_entity(&entity));
|
|
||||||
}
|
|
||||||
g.entities.insert(entity.name.clone(), entity);
|
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())
|
Ok("Entities created".to_string())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -41,10 +41,9 @@ impl McpTool for LogDecisionHandler {
|
|||||||
adrs.push(a);
|
adrs.push(a);
|
||||||
});
|
});
|
||||||
|
|
||||||
if let Some(adr) = new_adr
|
if let Some(adr) = new_adr {
|
||||||
&& let Ok(idx) = state.search_index.read()
|
let idx = state.search_index.read().unwrap_or_else(|e| e.into_inner()).clone();
|
||||||
{
|
let _ = idx.index_adr(&adr).await;
|
||||||
drop(idx.index_adr(&adr));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(format!("Decision logged as {}", adr_id).to_string())
|
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> {
|
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 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())
|
idx.search(&req.query, req.namespace.as_deref())
|
||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
} else {
|
|
||||||
vec![]
|
|
||||||
};
|
};
|
||||||
|
|
||||||
let kg_json = state.read_graph(|full| {
|
let kg_json = state.read_graph(|full| {
|
||||||
|
|||||||
@@ -42,9 +42,8 @@ impl McpTool for AddTaskHandler {
|
|||||||
dependencies: deps,
|
dependencies: deps,
|
||||||
acceptance_criteria: vec![],
|
acceptance_criteria: vec![],
|
||||||
};
|
};
|
||||||
if let Ok(idx) = state.search_index.read() {
|
let idx = state.search_index.read().unwrap_or_else(|e| e.into_inner()).clone();
|
||||||
drop(idx.index_task(&task));
|
let _ = idx.index_task(&task).await;
|
||||||
}
|
|
||||||
state.tasks.modify(|tasks| {
|
state.tasks.modify(|tasks| {
|
||||||
tasks.push(task);
|
tasks.push(task);
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -124,9 +124,8 @@ impl McpTool for StoreSnippetHandler {
|
|||||||
snippets.push(s_clone);
|
snippets.push(s_clone);
|
||||||
});
|
});
|
||||||
|
|
||||||
if let Ok(idx) = state.search_index.read() {
|
let idx = state.search_index.read().unwrap_or_else(|e| e.into_inner()).clone();
|
||||||
drop(idx.index_snippet(&snippet));
|
let _ = idx.index_snippet(&snippet).await;
|
||||||
}
|
|
||||||
|
|
||||||
Ok(format!("Snippet '{}' stored.", req.name).to_string())
|
Ok(format!("Snippet '{}' stored.", req.name).to_string())
|
||||||
}
|
}
|
||||||
|
|||||||
+19
-37
@@ -206,7 +206,7 @@ async fn gate_set_handler(
|
|||||||
reason: body.reason.clone(),
|
reason: body.reason.clone(),
|
||||||
timestamp: SystemTime::now()
|
timestamp: SystemTime::now()
|
||||||
.duration_since(UNIX_EPOCH)
|
.duration_since(UNIX_EPOCH)
|
||||||
.unwrap()
|
.unwrap_or_default()
|
||||||
.as_secs(),
|
.as_secs(),
|
||||||
};
|
};
|
||||||
app_state.handler.state.gates.modify(|gates| {
|
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");
|
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::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 {
|
let listener = match tokio::net::TcpListener::bind(&addr).await {
|
||||||
Ok(l) => l,
|
Ok(l) => l,
|
||||||
@@ -468,7 +468,7 @@ async fn ws_handler(
|
|||||||
.into_response()
|
.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 session_id = format!("{}", state.next_id.fetch_add(1, Ordering::SeqCst));
|
||||||
let (tx, mut rx) = mpsc::channel::<String>(100);
|
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);
|
tracing::trace!("Message content: {}", text);
|
||||||
if let Ok(payload) = serde_json::from_str::<serde_json::Value>(&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
|
// Process MCP request
|
||||||
if let Some(response) = handler.handle_request(payload).await {
|
if let Some(response) = handler.handle_request(payload).await {
|
||||||
let res_str = serde_json::to_string(&response).unwrap_or_default();
|
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 {
|
struct SessionCleanup {
|
||||||
session_id: String,
|
session_id: String,
|
||||||
state: Arc<AppState>,
|
state: Arc<AppState>,
|
||||||
send_task: Option<tokio::task::JoinHandle<()>>,
|
send_task: tokio::task::JoinHandle<()>,
|
||||||
recv_task: Option<tokio::task::JoinHandle<()>>,
|
recv_task: tokio::task::JoinHandle<()>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Drop for SessionCleanup {
|
impl Drop for SessionCleanup {
|
||||||
@@ -586,12 +572,8 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
|||||||
.write()
|
.write()
|
||||||
.unwrap_or_else(|e| e.into_inner())
|
.unwrap_or_else(|e| e.into_inner())
|
||||||
.remove(&self.session_id);
|
.remove(&self.session_id);
|
||||||
if let Some(task) = self.send_task.take() {
|
self.send_task.abort();
|
||||||
task.abort();
|
self.recv_task.abort();
|
||||||
}
|
|
||||||
if let Some(task) = self.recv_task.take() {
|
|
||||||
task.abort();
|
|
||||||
}
|
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"Websocket session {} closed and cleaned up",
|
"Websocket session {} closed and cleaned up",
|
||||||
self.session_id
|
self.session_id
|
||||||
@@ -602,15 +584,15 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
|||||||
let mut cleanup = SessionCleanup {
|
let mut cleanup = SessionCleanup {
|
||||||
session_id: session_id.clone(),
|
session_id: session_id.clone(),
|
||||||
state: Arc::clone(&state),
|
state: Arc::clone(&state),
|
||||||
send_task: Some(send_task),
|
send_task,
|
||||||
recv_task: Some(recv_task),
|
recv_task,
|
||||||
};
|
};
|
||||||
|
|
||||||
tokio::select! {
|
tokio::select! {
|
||||||
_ = cleanup.send_task.as_mut().unwrap() => {
|
_ = &mut cleanup.send_task => {
|
||||||
tracing::info!("Websocket send task finished for session {}", session_id);
|
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);
|
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 {
|
if !cli.daemon {
|
||||||
// Just spawn the daemon and exit. We no longer act as a proxy.
|
// Just spawn the daemon and exit. We no longer act as a proxy.
|
||||||
#[allow(clippy::zombie_processes)]
|
#[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")
|
.arg("--daemon")
|
||||||
.stdin(std::process::Stdio::null())
|
.stdin(std::process::Stdio::null())
|
||||||
.stdout(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 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
|
// 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![
|
let stores = vec![
|
||||||
("knowledge_graph_master", "knowledge_graph_master.json"),
|
("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() {
|
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);
|
let json_path = base.join(file_name);
|
||||||
if json_path.exists()
|
if json_path.exists()
|
||||||
&& let Ok(data) = fs::read(&json_path)
|
&& let Ok(data) = fs::read(&json_path)
|
||||||
&& serde_json::from_slice::<serde_json::Value>(&data).is_ok()
|
&& 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"));
|
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 _guard = rt.enter();
|
||||||
|
|
||||||
let state = Arc::new(MemoryState {
|
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 {
|
pub fn tool_def<T: JsonSchema>(name: &str, description: &str) -> serde_json::Value {
|
||||||
let schema = schemars::schema_for!(T);
|
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.
|
// MCP expects standard JSON schema. Schemars returns draft-07.
|
||||||
json!({
|
json!({
|
||||||
"name": name,
|
"name": name,
|
||||||
|
|||||||
@@ -30,11 +30,14 @@ impl MemoryIndex {
|
|||||||
let schema = schema_builder.build();
|
let schema = schema_builder.build();
|
||||||
|
|
||||||
let index_dir = store_dir.join("tantivy_index");
|
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)
|
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
|
let reader = index
|
||||||
.reader_builder()
|
.reader_builder()
|
||||||
.reload_policy(ReloadPolicy::OnCommitWithDelay)
|
.reload_policy(ReloadPolicy::OnCommitWithDelay)
|
||||||
|
|||||||
+3
-1
@@ -104,7 +104,9 @@ impl MemoryState {
|
|||||||
idx.add_adr_sync(a);
|
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;
|
let _ = new_idx.commit().await;
|
||||||
if let Ok(mut w) = self.search_index.write() {
|
if let Ok(mut w) = self.search_index.write() {
|
||||||
|
|||||||
Reference in new issue
Block a user