diff --git a/src/daemon.rs b/src/daemon.rs index d45c247..5f8a4d2 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -183,7 +183,11 @@ struct Activity { struct State { active_activities: HashMap, - db: Arc>, +} + +pub struct DbConnections { + pub writer: Mutex, + pub reader: Mutex, } pub fn open_db(path: &PathBuf) -> Result { @@ -219,10 +223,9 @@ pub fn open_db(path: &PathBuf) -> Result { Ok(conn) } -pub async fn run_daemon(socket_path: PathBuf, db: Arc>) -> Result<()> { +pub async fn run_daemon(socket_path: PathBuf, db: Arc) -> Result<()> { let state = Arc::new(Mutex::new(State { active_activities: HashMap::new(), - db, })); if socket_path.exists() { @@ -237,15 +240,16 @@ pub async fn run_daemon(socket_path: PathBuf, db: Arc>) -> Res loop { let (stream, _) = listener.accept().await?; let state = Arc::clone(&state); + let db = Arc::clone(&db); tokio::spawn(async move { - if let Err(e) = handle_connection(stream, state).await { + if let Err(e) = handle_connection(stream, state, db).await { error!("Connection error: {}", e); } }); } } -async fn handle_connection(mut stream: UnixStream, state: Arc>) -> Result<()> { +async fn handle_connection(mut stream: UnixStream, state: Arc>, db: Arc) -> Result<()> { let (reader, mut writer) = stream.split(); let mut reader = BufReader::new(reader); let mut line = String::new(); @@ -255,23 +259,23 @@ async fn handle_connection(mut stream: UnixStream, state: Arc>) -> match serde_json::from_str::(line.trim()) { Ok(SocketMessage::Command(ClientCommand::GetStats { since })) => { - let db = state.lock().unwrap().db.clone(); - let stats = tokio::task::spawn_blocking(move || collect_stats(&db, since)) + let db = Arc::clone(&db); + let stats = tokio::task::spawn_blocking(move || collect_stats(&db.reader, since)) .await??; writer.write_all((serde_json::to_string(&stats)? + "\n").as_bytes()).await?; break; } Ok(SocketMessage::Command(ClientCommand::GetTrend { since, bucket, drv })) => { - let db = state.lock().unwrap().db.clone(); - let trend = tokio::task::spawn_blocking(move || collect_trend(&db, since, bucket, drv)) + let db = Arc::clone(&db); + let trend = tokio::task::spawn_blocking(move || collect_trend(&db.reader, since, bucket, drv)) .await??; writer.write_all((serde_json::to_string(&trend)? + "\n").as_bytes()).await?; break; } Ok(SocketMessage::Command(ClientCommand::Clean)) => { - let db = state.lock().unwrap().db.clone(); + let db = Arc::clone(&db); tokio::task::spawn_blocking(move || -> Result<()> { - let conn = db.lock().unwrap(); + let conn = db.writer.lock().unwrap(); conn.execute_batch("DELETE FROM events; VACUUM; PRAGMA wal_checkpoint(TRUNCATE);")?; Ok(()) }).await??; @@ -280,7 +284,7 @@ async fn handle_connection(mut stream: UnixStream, state: Arc>) -> break; } Ok(SocketMessage::Event(event)) => { - if let Err(e) = process_event(event, &state) { + if let Err(e) = process_event(event, &state, &db) { error!("Failed to process event: {}", e); } } @@ -290,7 +294,7 @@ async fn handle_connection(mut stream: UnixStream, state: Arc>) -> Ok(()) } -fn process_event(event: NixEvent, state: &Arc>) -> Result<()> { +fn process_event(event: NixEvent, state: &Arc>, db: &Arc) -> Result<()> { let mut s = state.lock().unwrap(); match event { @@ -302,14 +306,7 @@ fn process_event(event: NixEvent, state: &Arc>) -> Result<()> { fields.get(0).and_then(|v| v.as_str()).unwrap_or("").to_string() }; - info!( - id, - parent, - act_type = %act_type, - text = %text, - fields = ?fields, - "start" - ); + info!(id, parent, act_type = %act_type, text = %text, fields = ?fields, "start"); s.active_activities.insert(id, Activity { id, @@ -325,12 +322,7 @@ fn process_event(event: NixEvent, state: &Arc>) -> Result<()> { let res_type = ResultType::from(event_type); if res_type != ResultType::BuildLogLine && res_type != ResultType::Progress && res_type != ResultType::SetExpected { - info!( - id, - res_type = %res_type, - fields = ?fields, - "result" - ); + info!(id, res_type = %res_type, fields = ?fields, "result"); } if let Some(act) = s.active_activities.get_mut(&id) { @@ -348,12 +340,8 @@ fn process_event(event: NixEvent, state: &Arc>) -> Result<()> { let duration_ms = end_time.signed_duration_since(act.start_time).num_milliseconds(); info!( - id = act.id, - act_type = %act_type, - duration_ms, - total_bytes = act.total_bytes, - text = %act.text, - fields = ?act.fields, + id = act.id, act_type = %act_type, duration_ms, + total_bytes = act.total_bytes, text = %act.text, fields = ?act.fields, "stop" ); @@ -367,25 +355,20 @@ fn process_event(event: NixEvent, state: &Arc>) -> Result<()> { _ => None, }; - let conn = s.db.lock().unwrap(); - conn.execute( + drop(s); + db.writer.lock().unwrap().execute( "INSERT INTO events (nix_id, parent_id, event_type, text, drv_path, cache_url, start_time, end_time, duration_ms, total_bytes) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)", rusqlite::params![ - act.id as i64, - act.parent_id as i64, - act.event_type as i64, - act.text, - drv_path, - cache_url, - act.start_time.to_rfc3339(), - end_time.to_rfc3339(), - duration_ms, - act.total_bytes as i64, + act.id as i64, act.parent_id as i64, act.event_type as i64, + act.text, drv_path, cache_url, + act.start_time.to_rfc3339(), end_time.to_rfc3339(), + duration_ms, act.total_bytes as i64, ], ).context("Failed to insert event")?; } } } + Ok(()) } diff --git a/src/main.rs b/src/main.rs index 8b501fb..9636493 100644 --- a/src/main.rs +++ b/src/main.rs @@ -11,7 +11,7 @@ use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; use tokio::net::UnixStream; use tracing::{error, info}; -use daemon::{open_db, run_daemon}; +use daemon::{open_db, run_daemon, DbConnections}; use stats::{display_stats, display_trend, display_trend_test, output_csv_trend, BucketSize, Stats, Trend}; #[derive(Clone, clap::ValueEnum)] @@ -133,8 +133,13 @@ async fn main() -> Result<()> { anyhow::bail!("Daemon already running at {}", socket_path.display()); } - let conn = open_db(&db_path)?; - run_daemon(socket_path, Arc::new(Mutex::new(conn))).await.context("Daemon failed")? + let writer = open_db(&db_path)?; + let reader = open_db(&db_path)?; + let db = Arc::new(DbConnections { + writer: Mutex::new(writer), + reader: Mutex::new(reader), + }); + run_daemon(socket_path, db).await.context("Daemon failed")? } Commands::Stats { socket, days, months, years } => { let socket_path = socket.unwrap_or_else(|| { diff --git a/src/stats.rs b/src/stats.rs index 9cb5b04..f9e9a21 100644 --- a/src/stats.rs +++ b/src/stats.rs @@ -1,7 +1,7 @@ use anyhow::{Context, Result}; use rusqlite::Connection; use serde::{Deserialize, Serialize}; -use std::sync::{Arc, Mutex}; +use std::sync::Mutex; #[derive(Debug, Serialize, Deserialize)] pub struct Stats { @@ -29,7 +29,7 @@ pub struct CacheStat { pub count: i64, } -pub fn collect_stats(db: &Arc>, since: Option) -> Result { +pub fn collect_stats(db: &Mutex, since: Option) -> Result { let conn = db.lock().unwrap(); // SQL NULL makes the WHERE condition vacuously true, giving us "no filter". @@ -243,7 +243,7 @@ pub struct Trend { } pub fn collect_trend( - db: &Arc>, + db: &Mutex, since: Option, bucket: BucketSize, drv: Option,