From ea34fb1b99765a52da5c881c7cc14b1fd30a5c89 Mon Sep 17 00:00:00 2001 From: jakka Date: Mon, 29 Sep 2025 18:13:20 +0300 Subject: rewrote everything in async. waiting for turso to support actual sql features --- src/db.rs | 10 +++--- src/files.rs | 112 ++++++++++++++++++++++++++++++++++------------------------- src/main.rs | 10 ++---- 3 files changed, 72 insertions(+), 60 deletions(-) (limited to 'src') diff --git a/src/db.rs b/src/db.rs index 420fdb7..a899ac3 100644 --- a/src/db.rs +++ b/src/db.rs @@ -19,7 +19,7 @@ const DEDUPE_DB: &str = "DELETE FROM flacs WHERE rowid NOT IN (SELECT MAX(rowid) FROM flacs GROUP BY path)"; const GET_MODTIME: &str = "SELECT modtime FROM flacs WHERE path = ?1"; -pub(crate) async fn init_db(path: Option<&Path>) -> Result { +pub(crate) async fn init_db(path: Option<&PathBuf>) -> Result { let db = if let Some(file) = path { turso::Builder::new_local(file.to_str().unwrap()) .build() @@ -111,13 +111,12 @@ pub(crate) async fn get_toencode_files(conn: &Connection) -> Result } pub(crate) async fn get_toencode_number(conn: &Connection) -> Result { - Ok(conn - .query(TOENCODE_NUMBER, ()) + conn.query(TOENCODE_NUMBER, ()) .await? .next() .await? .unwrap() - .get::(0)?) + .get::(0) } pub(crate) async fn get_modtime(conn: &Connection, file: &Path) -> Result { @@ -213,6 +212,7 @@ mod tests { } std::fs::remove_file(dbname).unwrap(); assert!(counter == 0) - }); + }) + .await; } } diff --git a/src/files.rs b/src/files.rs index 71310c4..12794ab 100644 --- a/src/files.rs +++ b/src/files.rs @@ -8,7 +8,7 @@ use std::{ fmt::Display, path::{Path, PathBuf}, sync::{ - Arc, Mutex, + Arc, atomic::{AtomicBool, AtomicUsize, Ordering}, }, thread::{self, sleep}, @@ -130,21 +130,19 @@ pub fn index_files_recursively( } pub async fn reencode_files( - conn: Connection, + conn: &Connection, handler: Arc, threads: usize, ) -> Result<()> { #[cfg(not(test))] let bar = ProgressBar::with_draw_target( - Some(db::get_toencode_number(&conn).await?), + Some(db::get_toencode_number(conn).await?), ProgressDrawTarget::stdout_with_hz(60), ) .with_style(ProgressStyle::with_template(BAR_TEMPLATE)?.progress_chars("#>-")) .with_message("Reencoding"); - let mut files = db::get_toencode_files(&conn).await?.into_iter(); - - let lock = Arc::new(Mutex::new(conn)); + let mut files = db::get_toencode_files(conn).await?.into_iter(); let thread_counter = Arc::new(AtomicUsize::new(0)); @@ -165,17 +163,19 @@ pub async fn reencode_files( thread_counter.fetch_add(1, Ordering::Relaxed); let handler = handler.clone(); - let lock = lock.clone(); #[cfg(not(test))] let bar = bar.clone(); let thread_counter = thread_counter.clone(); + let conn = conn.clone(); - s.spawn(async move || { + s.spawn(move || { match handle_encode(&file, handler) { Err(error) => eprintln!("{}", FileError::new(&file, error)), Ok(false) => { - if let Err(error) = db::update_file(&lock.lock().unwrap(), &file).await { - eprintln!("{}", FileError::new(&file, error)); + if let Err(error) = + smol::block_on(async { db::update_file(&conn, &file).await }) + { + eprintln!("{}", FileError::new(&file, error)) } #[cfg(not(test))] bar.inc(1) @@ -198,8 +198,8 @@ pub async fn reencode_files( Ok(()) } -pub fn clean_files(conn: &Connection, handler: Arc) -> Result<()> { - let files = conn.init_clean_files()?; +pub async fn clean_files(conn: &Connection, handler: Arc) -> Result<()> { + let files = db::init_clean_files(conn).await?; #[cfg(not(test))] let spinner = ProgressBar::with_draw_target(None, ProgressDrawTarget::stdout_with_hz(60)) @@ -210,17 +210,20 @@ pub fn clean_files(conn: &Connection, handler: Arc) -> Result<()> { files.iter().for_each(|file| { #[allow(clippy::collapsible_if)] if handler.load(Ordering::SeqCst) && !file.exists() { - if let Err(error) = conn.remove_file(file) { + #[cfg(not(test))] + let spinner = spinner.clone(); + if let Err(error) = smol::block_on(async { db::remove_file(conn, file).await }) { eprintln!("{}", FileError::new(file, error)) }; #[cfg(not(test))] spinner.inc(1); } }); + #[cfg(not(test))] spinner.finish(); - conn.vacuum()?; + db::vacuum(conn).await?; Ok(()) } @@ -228,52 +231,65 @@ pub fn clean_files(conn: &Connection, handler: Arc) -> Result<()> { #[cfg(test)] mod tests { use super::*; + use macro_rules_attribute::apply; + use smol_macros::{Executor, test}; - #[test] - fn test_index_lots_of_files() { + #[apply(test!)] + async fn test_index_lots_of_files(ex: &Executor<'_>) { let dbname = PathBuf::from("temp3.db"); let handler = Arc::new(AtomicBool::new(true)); - let conn = Connection::new(Some(&dbname)).unwrap(); - index_files_recursively(Path::new("./testfiles"), &conn, handler).unwrap(); - std::fs::remove_file(dbname).unwrap(); + ex.spawn(async { + let db = db::init_db(Some(&dbname)).await.unwrap(); + let conn = db.connect().unwrap(); + index_files_recursively(Path::new("./testfiles"), &conn, handler).unwrap(); + std::fs::remove_file(dbname).unwrap(); + }) + .await } - #[test] - fn test_clean_files() { + #[apply(test!)] + async fn test_clean_files(ex: &Executor<'_>) { let dbname = PathBuf::from("temp4.db"); let handler = Arc::new(AtomicBool::new(true)); - let conn = Connection::new(Some(&dbname)).unwrap(); - let filenames = [ - "./samples/16bit.flac", - "./samples/24bit.flac", - "./samples/32bit.flac", - "./samples/nonexisting.flac", - ]; - std::fs::copy("./samples/32bit.flac", "./samples/nonexisting.flac").unwrap(); - for file in filenames { - let filename = PathBuf::from(file); - conn.insert_file(&filename).unwrap(); - } + ex.spawn(async { + let db = db::init_db(Some(&dbname)).await.unwrap(); + let conn = db.connect().unwrap(); + let filenames = [ + "./samples/16bit.flac", + "./samples/24bit.flac", + "./samples/32bit.flac", + "./samples/nonexisting.flac", + ]; + std::fs::copy("./samples/32bit.flac", "./samples/nonexisting.flac").unwrap(); + for file in filenames { + let filename = PathBuf::from(file); + db::insert_file(&conn, &filename).await.unwrap(); + } - std::fs::remove_file("./samples/nonexisting.flac").unwrap(); + std::fs::remove_file("./samples/nonexisting.flac").unwrap(); - clean_files(&conn, handler).unwrap(); - let counter = conn.init_clean_files().unwrap().len(); - std::fs::remove_file(dbname).unwrap(); - assert!(counter == 3) + clean_files(&conn, handler).await.unwrap(); + let counter = db::init_clean_files(&conn).await.unwrap().len(); + std::fs::remove_file(dbname).unwrap(); + assert!(counter == 3) + }) + .await; } - #[test] - fn test_reencode_lots_of_files() { + #[apply(test!)] + async fn test_reencode_lots_of_files(ex: &Executor<'_>) { let dbname = PathBuf::from("temp5.db"); let handler = Arc::new(AtomicBool::new(true)); - let conn = Connection::new(Some(&dbname)).unwrap(); - let temp = handler.clone(); - index_files_recursively(Path::new("./testfiles"), &conn, temp).unwrap(); - println!("\n{}", conn.get_toencode_number().unwrap()); - reencode_files(conn, handler, 4).unwrap(); - let conn = Connection::new(Some(&dbname)).unwrap(); - println!("\n{}", conn.get_toencode_number().unwrap()); - std::fs::remove_file(dbname).unwrap(); + ex.spawn(async { + let db = db::init_db(Some(&dbname)).await.unwrap(); + let conn = db.connect().unwrap(); + let temp = handler.clone(); + index_files_recursively(Path::new("./testfiles"), &conn, temp).unwrap(); + println!("\n{}", db::get_toencode_number(&conn).await.unwrap()); + reencode_files(&conn, handler, 4).await.unwrap(); + println!("\n{}", db::get_toencode_number(&conn).await.unwrap()); + std::fs::remove_file(dbname).unwrap(); + }) + .await; } } diff --git a/src/main.rs b/src/main.rs index 2ef2491..e9543f8 100644 --- a/src/main.rs +++ b/src/main.rs @@ -91,11 +91,7 @@ fn main() -> Result<()> { })?; smol::block_on(async move { - let path = if let Some(path) = args.get_one::("db") { - Some(path.as_path()) - } else { - None - }; + let path = args.get_one::("db"); let db = db::init_db(path).await?; if path.is_none() && !args.get_flag("clean") && !args.get_flag("doit") { @@ -111,13 +107,13 @@ fn main() -> Result<()> { if args.get_flag("clean") { let handler = running.clone(); - files::clean_files(&db.connect()?, handler)?; + files::clean_files(&db.connect()?, handler).await?; } if args.get_flag("doit") { let hanlder = running.clone(); let threads = *args.get_one::("threads").unwrap(); - files::reencode_files(db.connect()?, hanlder, threads).await?; + files::reencode_files(&db.connect()?, hanlder, threads).await?; } Ok::<(), anyhow::Error>(()) -- cgit v1.3.1