refactor: apply zero-unwrap policy and optimize locks in store.rs and handlers
This commit is contained in:
1 parent
9f24e66d88
commit
8afbf97b11
33 files changed
+3127
-3207
No files matched your search
+80
-64
@@ -4,8 +4,10 @@
|
||||
)]
|
||||
|
||||
mod handlers;
|
||||
mod handlers_v2;
|
||||
mod mcp;
|
||||
mod models;
|
||||
mod router;
|
||||
mod search;
|
||||
mod state;
|
||||
mod store;
|
||||
@@ -218,14 +220,12 @@ async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::E
|
||||
state.rebuild_index().await;
|
||||
tokio::spawn(index_committer_worker(Arc::clone(&state)));
|
||||
let app_state = Arc::new(AppState {
|
||||
handler: Arc::new(MemoryHandler {
|
||||
state: Arc::clone(&state),
|
||||
}),
|
||||
clients: RwLock::new(HashMap::new()),
|
||||
next_id: AtomicUsize::new(1),
|
||||
});
|
||||
handler: Arc::new(MemoryHandler::new(Arc::clone(&state))),
|
||||
clients: RwLock::new(HashMap::new()),
|
||||
next_id: AtomicUsize::new(1),
|
||||
});
|
||||
|
||||
let app = Router::new()
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/api/version",
|
||||
get(|| async move {
|
||||
@@ -404,27 +404,28 @@ async fn run_server(state: Arc<MemoryState>) -> Result<(), Box<dyn std::error::E
|
||||
)
|
||||
.with_state(app_state);
|
||||
|
||||
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();
|
||||
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 = match tokio::net::TcpListener::bind(&addr).await {
|
||||
Ok(l) => l,
|
||||
Err(e) => {
|
||||
let log_path = dirs::home_dir()
|
||||
.unwrap_or_default()
|
||||
.join(".gemini/mcp_memory/daemon_error.log");
|
||||
let _ = tokio::fs::write(&log_path, format!("Failed to bind to {}: {}\n", addr, e)).await;
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
if let Err(e) = axum::serve(listener, app.into_make_service()).await {
|
||||
let listener = match tokio::net::TcpListener::bind(&addr).await {
|
||||
Ok(l) => l,
|
||||
Err(e) => {
|
||||
let log_path = dirs::home_dir()
|
||||
.unwrap_or_default()
|
||||
.join(".gemini/mcp_memory/daemon_error.log");
|
||||
let _ = tokio::fs::write(&log_path, format!("Server crashed: {}\n", e)).await;
|
||||
let _ =
|
||||
tokio::fs::write(&log_path, format!("Failed to bind to {}: {}\n", addr, e)).await;
|
||||
return Ok(());
|
||||
}
|
||||
Ok(())
|
||||
};
|
||||
if let Err(e) = axum::serve(listener, app.into_make_service()).await {
|
||||
let log_path = dirs::home_dir()
|
||||
.unwrap_or_default()
|
||||
.join(".gemini/mcp_memory/daemon_error.log");
|
||||
let _ = tokio::fs::write(&log_path, format!("Server crashed: {}\n", e)).await;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn ws_handler(
|
||||
@@ -486,42 +487,43 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
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");
|
||||
let activity_msg = format!("Agent executed tool: {}", name);
|
||||
&& 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 event = serde_json::json!({
|
||||
"type": "activity",
|
||||
"data": activity_msg
|
||||
});
|
||||
|
||||
let senders: Vec<_> = state_clone
|
||||
.clients
|
||||
.read()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.filter_map(|(id, tx)| {
|
||||
if id != &session_id_clone {
|
||||
Some(tx.clone())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
let senders: Vec<_> = state_clone
|
||||
.clients
|
||||
.read()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.filter_map(|(id, tx)| {
|
||||
if id != &session_id_clone {
|
||||
Some(tx.clone())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
for client_tx in senders {
|
||||
let _ = client_tx.try_send(event.to_string());
|
||||
}
|
||||
for client_tx in senders {
|
||||
let _ = client_tx.try_send(event.to_string());
|
||||
}
|
||||
}
|
||||
} // 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 res_str = serde_json::to_string(&response).unwrap_or_default();
|
||||
let tx_opt = state_clone
|
||||
.clients
|
||||
.read()
|
||||
@@ -576,7 +578,11 @@ async fn handle_socket(socket: WebSocket, state: Arc<AppState>, client_type: Str
|
||||
|
||||
impl Drop for SessionCleanup {
|
||||
fn drop(&mut self) {
|
||||
self.state.clients.write().unwrap().remove(&self.session_id);
|
||||
self.state
|
||||
.clients
|
||||
.write()
|
||||
.unwrap_or_else(|e| e.into_inner())
|
||||
.remove(&self.session_id);
|
||||
if let Some(task) = self.send_task.take() {
|
||||
task.abort();
|
||||
}
|
||||
@@ -640,7 +646,13 @@ async fn nvim_telemetry_handler(
|
||||
});
|
||||
|
||||
let msg_str = ws_msg.to_string();
|
||||
let senders: Vec<_> = state.clients.read().unwrap().values().cloned().collect();
|
||||
let senders: Vec<_> = state
|
||||
.clients
|
||||
.read()
|
||||
.unwrap_or_else(|e| e.into_inner())
|
||||
.values()
|
||||
.cloned()
|
||||
.collect();
|
||||
for tx in senders {
|
||||
let _ = tx.try_send(msg_str.clone());
|
||||
}
|
||||
@@ -698,9 +710,12 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
let mut cmd = std::process::Command::new("curl");
|
||||
cmd.arg("-k").arg("-X").arg("POST");
|
||||
if !token.is_empty() {
|
||||
cmd.arg("-H").arg(format!("Authorization: Bearer {}", token.trim()));
|
||||
cmd.arg("-H")
|
||||
.arg(format!("Authorization: Bearer {}", token.trim()));
|
||||
}
|
||||
let _ = cmd.arg(format!("https://127.0.0.1:{}/shutdown", port)).output();
|
||||
let _ = cmd
|
||||
.arg(format!("https://127.0.0.1:{}/shutdown", port))
|
||||
.output();
|
||||
println!("Sent shutdown request to server.");
|
||||
return Ok(());
|
||||
}
|
||||
@@ -711,9 +726,12 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
let mut cmd = std::process::Command::new("curl");
|
||||
cmd.arg("-k").arg("-X").arg("POST");
|
||||
if !token.is_empty() {
|
||||
cmd.arg("-H").arg(format!("Authorization: Bearer {}", token.trim()));
|
||||
cmd.arg("-H")
|
||||
.arg(format!("Authorization: Bearer {}", token.trim()));
|
||||
}
|
||||
let _ = cmd.arg(format!("https://127.0.0.1:{}/shutdown", port)).output();
|
||||
let _ = cmd
|
||||
.arg(format!("https://127.0.0.1:{}/shutdown", port))
|
||||
.output();
|
||||
println!("Sent shutdown request to existing server. Waiting for it to exit...");
|
||||
std::thread::sleep(std::time::Duration::from_millis(1500));
|
||||
return Ok(());
|
||||
@@ -780,13 +798,11 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
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"),
|
||||
);
|
||||
}
|
||||
&& 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"));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in new issue
Block a user