diff --git a/Cargo.lock b/Cargo.lock index 4d85626..940600b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -302,21 +302,6 @@ dependencies = [ "piper", ] -[[package]] -name = "bread-onnx" -version = "0.7.2" -source = "git+https://git.breadway.dev/Breadway/bread-ecosystem?tag=v0.7.2#30517f161724132cdeb658c04cf5e490be07ee73" -dependencies = [ - "anyhow", - "bread-utils", - "hex", - "ort", - "sha2", - "tokenizers", - "tracing", - "ureq", -] - [[package]] name = "bread-screenshots" version = "0.7.2" @@ -408,12 +393,12 @@ name = "breadpad-shared" version = "0.5.3" dependencies = [ "anyhow", - "bread-onnx", "bread-theme", "chrono", "dirs 5.0.1", "gtk4", "ical", + "ndarray 0.16.1", "ort", "regex", "reqwest", @@ -679,18 +664,6 @@ dependencies = [ "windows-sys 0.59.0", ] -[[package]] -name = "console" -version = "0.16.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4fe5f465a4f6fee88fad41b85d990f84c835335e85b5d9e6e63e0d06d28cba7c" -dependencies = [ - "encode_unicode", - "libc", - "unicode-width", - "windows-sys 0.61.2", -] - [[package]] name = "convert_case" version = "0.10.0" @@ -791,12 +764,6 @@ dependencies = [ "typenum", ] -[[package]] -name = "daachorse" -version = "1.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6f55d7153ba3b507595872a3874803f07a8a81d1e888abed8e5db7da0597d6e2" - [[package]] name = "darling" version = "0.20.11" @@ -1027,9 +994,6 @@ name = "esaxx-rs" version = "0.1.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d817e038c30374a4bcb22f94d0a8a0e216958d4c3dcde369b1439fec4bdda6e6" -dependencies = [ - "cc", -] [[package]] name = "event-listener" @@ -1053,9 +1017,9 @@ dependencies = [ [[package]] name = "fancy-regex" -version = "0.17.0" +version = "0.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "72cf461f865c862bb7dc573f643dd6a2b6842f7c30b07882b56bd148cc2761b8" +checksum = "6e24cb5a94bcae1e5408b0effca5cd7172ea3c5755049c5f3af4cd283a165298" dependencies = [ "bit-set", "regex-automata", @@ -1566,7 +1530,7 @@ checksum = "629d8f3bbeda9d148036d6b0de0a3ab947abd08ce90626327fc3547a49d59d97" dependencies = [ "dirs 6.0.0", "http", - "indicatif 0.17.11", + "indicatif", "libc", "log", "rand 0.9.5", @@ -1863,26 +1827,13 @@ version = "0.17.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "183b3088984b400f4cfac3620d5e076c84da5364016b4f49473de574b2586235" dependencies = [ - "console 0.15.11", + "console", "number_prefix", "portable-atomic", "unicode-width", "web-time", ] -[[package]] -name = "indicatif" -version = "0.18.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9433806cd6b4ec1aba79c021c7e4c58fb4c3b9977c085062e611ac929998fb0c" -dependencies = [ - "console 0.16.4", - "portable-atomic", - "unicode-width", - "unit-prefix", - "web-time", -] - [[package]] name = "ipnet" version = "2.12.1" @@ -2126,6 +2077,21 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "ndarray" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "882ed72dce9365842bf196bdeedf5055305f11fc8c03dee7bb0194a6cad34841" +dependencies = [ + "matrixmultiply", + "num-complex", + "num-integer", + "num-traits", + "portable-atomic", + "portable-atomic-util", + "rawpointer", +] + [[package]] name = "ndarray" version = "0.17.2" @@ -2218,28 +2184,6 @@ version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" -[[package]] -name = "onig" -version = "6.5.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0cc3cbf698f9438986c11a880c90a6d04b9de27575afd28bbf45b154b6c709e2" -dependencies = [ - "bitflags", - "libc", - "once_cell", - "onig_sys", -] - -[[package]] -name = "onig_sys" -version = "69.9.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1e68317604e77e53b85896388e1a803c1d21b74c899ec9e5e1112db90735edd7" -dependencies = [ - "cc", - "pkg-config", -] - [[package]] name = "option-ext" version = "0.2.0" @@ -2263,7 +2207,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4336a1e2b38848325241c72889086886004e589b7c74f335e60a8e8db5138a0b" dependencies = [ "libloading", - "ndarray", + "ndarray 0.17.2", "ort-sys", "smallvec", "tracing", @@ -2994,17 +2938,6 @@ dependencies = [ "digest", ] -[[package]] -name = "sha2" -version = "0.10.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" -dependencies = [ - "cfg-if", - "cpufeatures 0.2.17", - "digest", -] - [[package]] name = "sharded-slab" version = "0.1.7" @@ -3274,25 +3207,23 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokenizers" -version = "0.23.1" +version = "0.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "44e5bea67576e04b6ff8564c5d9e09c2ef0cf476502245f2f120e497769d3112" +checksum = "a620b996116a59e184c2fa2dfd8251ea34a36d0a514758c6f966386bd2e03476" dependencies = [ "ahash", + "aho-corasick", "compact_str", - "daachorse", "dary_heap", "derive_builder", "esaxx-rs", "fancy-regex", "getrandom 0.3.4", "hf-hub", - "indicatif 0.18.6", "itertools", "log", "macro_rules_attribute", "monostate", - "onig", "paste", "rand 0.9.5", "rayon", @@ -3607,12 +3538,6 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e" -[[package]] -name = "unit-prefix" -version = "0.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "81e544489bf3d8ef66c953931f56617f423cd4b5494be343d9b9d3dda037b9a3" - [[package]] name = "untrusted" version = "0.9.0" diff --git a/Cargo.toml b/Cargo.toml index b9508be..84621db 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -24,10 +24,9 @@ chrono = { version = "0.4", features = ["serde"] } rrule = "0.12" tokio = { version = "1", features = ["full"] } zbus = { version = "4", default-features = false, features = ["tokio"] } -# WHY: bread-onnx's Provider enum references every EP type, so those ort -# features must be on in the consumer even if breadpad only requests MIGraphX. -ort = { version = "2.0.0-rc.12", default-features = false, features = ["std", "tracing", "api-24", "migraphx", "cuda", "openvino", "vitis", "load-dynamic"] } -tokenizers = { version = "0.23", default-features = false, features = ["http", "fancy-regex"] } +ort = { version = "2.0.0-rc.12", default-features = false, features = ["std", "ndarray", "tracing", "api-24", "rocm", "load-dynamic"] } +ndarray = "0.16" +tokenizers = { version = "0.21", default-features = false, features = ["http", "fancy-regex"] } gtk4 = { version = "0.11", features = ["v4_12"] } gtk4-layer-shell = "0.8" hyprland = "0.4.0-beta.3" diff --git a/breadman/Cargo.toml b/breadman/Cargo.toml index d47a586..e81b7e6 100644 --- a/breadman/Cargo.toml +++ b/breadman/Cargo.toml @@ -15,7 +15,7 @@ breadpad-shared = { path = "../breadpad-shared" } bread-screenshots = { git = "https://git.breadway.dev/Breadway/bread-ecosystem", tag = "v0.7.2" } # Shared `--screenshot` pair validation + settle delay (`screenshot_cli`). bread-utils = { git = "https://git.breadway.dev/Breadway/bread-ecosystem", tag = "v0.7.2" } -# `adw` implies `gtk` (`chip`, `set_chip_active`, `adw::init`). +# `adw` implies `gtk`. Local chip/init shims remain in src/theme_widgets.rs. bread-theme = { git = "https://git.breadway.dev/Breadway/bread-ecosystem", tag = "v0.7.4", features = ["adw"] } libadwaita = { version = "0.9", features = ["v1_7"] } anyhow.workspace = true diff --git a/breadman/src/editor.rs b/breadman/src/editor.rs index e39b11c..bea9803 100644 --- a/breadman/src/editor.rs +++ b/breadman/src/editor.rs @@ -62,16 +62,16 @@ pub fn open_editor( let type_row = libadwaita::ActionRow::builder().title("Type").build(); let type_pill_box = gtk4::Box::builder().orientation(gtk4::Orientation::Horizontal).spacing(4).valign(gtk4::Align::Center).build(); let selected_type: Rc> = Rc::new(RefCell::new(note.note_type.as_str().to_string())); - let type_pills: Vec<(gtk4::Button, &'static str)> = NoteType::all_builtin().iter().map(|&name| (bread_theme::gtk::chip(name), name)).collect(); + let type_pills: Vec<(gtk4::Button, &'static str)> = NoteType::all_builtin().iter().map(|&name| (crate::theme_widgets::chip(name), name)).collect(); for (btn, name) in &type_pills { - bread_theme::gtk::set_chip_active(btn, *name == selected_type.borrow().as_str()); + crate::theme_widgets::set_chip_active(btn, *name == selected_type.borrow().as_str()); let sel = selected_type.clone(); let name = *name; let all_btns: Vec = type_pills.iter().map(|(b, _)| b.clone()).collect(); btn.connect_clicked(move |clicked| { *sel.borrow_mut() = name.to_string(); - for b in &all_btns { bread_theme::gtk::set_chip_active(b, false); } - bread_theme::gtk::set_chip_active(clicked, true); + for b in &all_btns { crate::theme_widgets::set_chip_active(b, false); } + crate::theme_widgets::set_chip_active(clicked, true); }); type_pill_box.append(btn); } diff --git a/breadman/src/main.rs b/breadman/src/main.rs index efd5b3f..0a4f810 100644 --- a/breadman/src/main.rs +++ b/breadman/src/main.rs @@ -14,6 +14,7 @@ use std::sync::Arc; mod editor; mod screenshot; +mod theme_widgets; mod views; // ── Args ───────────────────────────────────────────────────────────────────── @@ -362,7 +363,7 @@ fn build_app_window( // Needed once before constructing any adw:: widget (see views::settings) — // also forces dark mode, since bread-theme's palette is a fixed dark base // regardless of the system GTK preference. - bread_theme::adw::init(); + theme_widgets::init_adw(); let store = Arc::new(Store::new()?); let notes = store.load_all()?; @@ -644,17 +645,17 @@ fn show_add_note_window(parent: >k4::ApplicationWindow, state: AppState, prese let selected_type: Rc> = Rc::new(RefCell::new(preselect.clone())); let chips: Vec<(gtk4::Button, NoteType)> = NoteType::all_builtin() .iter() - .map(|&name| (bread_theme::gtk::chip(name), NoteType::from_str(name))) + .map(|&name| (theme_widgets::chip(name), NoteType::from_str(name))) .collect(); for (btn, nt) in &chips { - bread_theme::gtk::set_chip_active(btn, *nt == preselect); + theme_widgets::set_chip_active(btn, *nt == preselect); let sel = selected_type.clone(); let nt_c = nt.clone(); let all_btns: Vec = chips.iter().map(|(b, _)| b.clone()).collect(); btn.connect_clicked(move |clicked| { *sel.borrow_mut() = nt_c.clone(); - for b in &all_btns { bread_theme::gtk::set_chip_active(b, false); } - bread_theme::gtk::set_chip_active(clicked, true); + for b in &all_btns { theme_widgets::set_chip_active(b, false); } + theme_widgets::set_chip_active(clicked, true); }); chip_box.append(btn); } diff --git a/breadman/src/theme_widgets.rs b/breadman/src/theme_widgets.rs new file mode 100644 index 0000000..a5221e6 --- /dev/null +++ b/breadman/src/theme_widgets.rs @@ -0,0 +1,23 @@ +//! Local stand-ins for `bread_theme::gtk::{chip, set_chip_active}` and +//! `bread_theme::adw::init`, which are not on bread-theme v0.7.1. + +use gtk4::prelude::*; + +pub fn chip(label: &str) -> gtk4::Button { + gtk4::Button::builder().label(label).css_classes(["chip"]).build() +} + +pub fn set_chip_active(chip: &impl IsA, active: bool) { + if active { + chip.add_css_class("active"); + } else { + chip.remove_css_class("active"); + } +} + +/// Initializes libadwaita and forces dark mode (bread-theme's palette is a +/// fixed dark base regardless of the system GTK preference). +pub fn init_adw() { + libadwaita::init().expect("failed to initialize libadwaita"); + libadwaita::StyleManager::default().set_color_scheme(libadwaita::ColorScheme::ForceDark); +} diff --git a/breadman/src/views/settings.rs b/breadman/src/views/settings.rs index 6090274..114edd1 100644 --- a/breadman/src/views/settings.rs +++ b/breadman/src/views/settings.rs @@ -92,10 +92,10 @@ pub fn build(cfg: &Config, on_save: impl Fn(Config) + 'static) -> gtk4::Scrolled let selected_type: Rc> = Rc::new(RefCell::new(cfg.settings.default_type.clone())); let type_pills: Vec<(gtk4::Button, &'static str)> = NoteType::all_builtin() .iter() - .map(|&name| (bread_theme::gtk::chip(name), name)) + .map(|&name| (crate::theme_widgets::chip(name), name)) .collect(); for (btn, name) in &type_pills { - bread_theme::gtk::set_chip_active(btn, *name == selected_type.borrow().as_str()); + crate::theme_widgets::set_chip_active(btn, *name == selected_type.borrow().as_str()); type_pill_box.append(btn); } general_list.append(&field_row("Default type", None, &type_pill_box)); @@ -260,8 +260,8 @@ pub fn build(cfg: &Config, on_save: impl Fn(Config) + 'static) -> gtk4::Scrolled let all_btns: Vec = type_pills.iter().map(|(b, _)| b.clone()).collect(); btn.connect_clicked(move |clicked| { *sel.borrow_mut() = name.to_string(); - for b in &all_btns { bread_theme::gtk::set_chip_active(b, false); } - bread_theme::gtk::set_chip_active(clicked, true); + for b in &all_btns { crate::theme_widgets::set_chip_active(b, false); } + crate::theme_widgets::set_chip_active(clicked, true); apply_now(); }); } diff --git a/breadpad-shared/Cargo.toml b/breadpad-shared/Cargo.toml index a56c187..30a1467 100644 --- a/breadpad-shared/Cargo.toml +++ b/breadpad-shared/Cargo.toml @@ -20,7 +20,7 @@ tokio.workspace = true zbus.workspace = true ort.workspace = true tokenizers.workspace = true -bread-onnx = { git = "https://git.breadway.dev/Breadway/bread-ecosystem", tag = "v0.7.2" } +ndarray.workspace = true toml.workspace = true dirs.workspace = true regex.workspace = true diff --git a/breadpad-shared/src/classifier.rs b/breadpad-shared/src/classifier.rs index 35c2650..e752a31 100644 --- a/breadpad-shared/src/classifier.rs +++ b/breadpad-shared/src/classifier.rs @@ -2,8 +2,6 @@ use crate::ai::OllamaClient; use crate::config::OllamaConfig; use crate::parser::parse_rule_based; use crate::types::{ClassificationResult, NoteType}; -use bread_onnx::{build_session, Provider}; -use ort::session::builder::GraphOptimizationLevel; use std::path::PathBuf; /// Minimum Tier 1 confidence needed to skip Tier 2 entirely. @@ -18,7 +16,7 @@ pub enum ExecutionProvider { impl ExecutionProvider { pub fn as_str(&self) -> &str { match self { - ExecutionProvider::Gpu => "MIGraphX (iGPU)", + ExecutionProvider::Gpu => "ROCm (iGPU)", ExecutionProvider::Cpu => "CPU", } } @@ -109,7 +107,9 @@ impl Classifier { // ── Tier 2 ─────────────────────────────────────────────────────────── // ONNX model classifies the type only; Tier 1's time/rrule/body are kept. - let tier2 = if let (Some(session), Some(tokenizer)) = (&mut self.session, &self.tokenizer) { + let tier2 = if let (Some(session), Some(tokenizer)) = + (&mut self.session, &self.tokenizer) + { match run_onnx(session, tokenizer, text) { Ok(r) => { tracing::debug!("Tier 2: {:?} conf={:.2}", r.note_type, r.confidence); @@ -163,18 +163,9 @@ impl Classifier { // entailment score across all five passes. const HYPOTHESES: [(&str, &str); 5] = [ ("This note is a task or action item to complete.", "todo"), - ( - "This note is a reminder with a specific time or deadline.", - "reminder", - ), - ( - "This note is an idea, suggestion, or creative thought.", - "idea", - ), - ( - "This note is a general observation or piece of information.", - "note", - ), + ("This note is a reminder with a specific time or deadline.", "reminder"), + ("This note is an idea, suggestion, or creative thought.", "idea"), + ("This note is a general observation or piece of information.", "note"), ("This note is a question that needs an answer.", "question"), ]; @@ -193,17 +184,15 @@ fn run_onnx( .map_err(|e| anyhow::anyhow!("tokenize: {}", e))?; let ids: Vec = encoding.get_ids().iter().map(|&x| x as i64).collect(); - let mask: Vec = encoding - .get_attention_mask() - .iter() - .map(|&x| x as i64) - .collect(); + let mask: Vec = encoding.get_attention_mask().iter().map(|&x| x as i64).collect(); let len = ids.len(); - let ids_tensor = ort::value::Tensor::::from_array((vec![1i64, len as i64], ids)) - .map_err(|e| anyhow::anyhow!("ids tensor: {}", e))?; - let mask_tensor = ort::value::Tensor::::from_array((vec![1i64, len as i64], mask)) - .map_err(|e| anyhow::anyhow!("mask tensor: {}", e))?; + let ids_tensor = ort::value::Tensor::::from_array( + (vec![1i64, len as i64], ids) + ).map_err(|e| anyhow::anyhow!("ids tensor: {}", e))?; + let mask_tensor = ort::value::Tensor::::from_array( + (vec![1i64, len as i64], mask) + ).map_err(|e| anyhow::anyhow!("mask tensor: {}", e))?; let inputs = ort::inputs![ "input_ids" => ids_tensor, @@ -218,7 +207,10 @@ fn run_onnx( .map_err(|e| anyhow::anyhow!("extract logits: {}", e))?; let (_, logits_slice) = logits; - entailment_scores[i] = logits_slice.get(ENTAILMENT_IDX).copied().unwrap_or(0.0); + entailment_scores[i] = logits_slice + .get(ENTAILMENT_IDX) + .copied() + .unwrap_or(0.0); } let best_idx = entailment_scores @@ -251,30 +243,42 @@ fn softmax_single(logits: &[f32], idx: usize) -> f32 { exps[idx] / sum } -fn try_load_session(path: &std::path::Path) -> (Option, ExecutionProvider) { - // WHY: distro onnxruntime-rocm is MIGraphX, not classic ROCm; bread-onnx - // appends CPU so a missing GPU EP does not disable Tier 2. - match build_session( - path, - GraphOptimizationLevel::Level3, - &[Provider::MiGraphX { device_id: 0 }], - ) { +fn try_load_session( + path: &std::path::Path, +) -> (Option, ExecutionProvider) { + // Try ROCm (iGPU) first, fall back to CPU. + let rocm_available = { + use ort::execution_providers::ExecutionProvider as _; + ort::ep::ROCm::default().is_available().unwrap_or(false) + }; + if rocm_available { + match build_onnx_session(path, ort::ep::ROCm::default().build()) { + Ok(s) => { + tracing::info!("ONNX session loaded (ROCm iGPU)"); + return (Some(s), ExecutionProvider::Gpu); + } + Err(e) => tracing::debug!("ROCm EP unavailable: {}; trying CPU", e), + } + } + match build_onnx_session(path, ort::ep::CPU::default().build()) { Ok(s) => { - tracing::info!("ONNX session loaded (MIGraphX, CPU fallback)"); - (Some(s), ExecutionProvider::Gpu) + tracing::info!("ONNX session loaded (CPU)"); + (Some(s), ExecutionProvider::Cpu) } Err(e) => { - tracing::debug!("MIGraphX session failed: {}; trying CPU", e); - match build_session(path, GraphOptimizationLevel::Level3, &[Provider::Cpu]) { - Ok(s) => { - tracing::info!("ONNX session loaded (CPU)"); - (Some(s), ExecutionProvider::Cpu) - } - Err(e) => { - tracing::warn!("failed to load ONNX session: {}; Tier 2 disabled", e); - (None, ExecutionProvider::Cpu) - } - } + tracing::warn!("failed to load ONNX session: {}; Tier 2 disabled", e); + (None, ExecutionProvider::Cpu) } } } + +fn build_onnx_session( + path: &std::path::Path, + ep: ort::ep::ExecutionProviderDispatch, +) -> anyhow::Result { + let mut builder = ort::session::Session::builder() + .map_err(|e| anyhow::anyhow!("builder: {}", e))? + .with_execution_providers([ep]) + .map_err(|e| anyhow::anyhow!("ep: {}", e))?; + builder.commit_from_file(path).map_err(|e| anyhow::anyhow!("load: {}", e)) +} diff --git a/breadpad-shared/tests/classifier.rs b/breadpad-shared/tests/classifier.rs index 7f0294a..3b5527c 100644 --- a/breadpad-shared/tests/classifier.rs +++ b/breadpad-shared/tests/classifier.rs @@ -1,24 +1,17 @@ use breadpad_shared::classifier::{Classifier, ExecutionProvider}; use breadpad_shared::types::NoteType; use chrono::Timelike; -use std::path::PathBuf; -/// Rule-based path only — a present `~/.local/share/breadpad/model` must not -/// change these assertions. fn cl() -> Classifier { - Classifier::load_with_paths( - "08:00", - PathBuf::from("/nonexistent/classifier.onnx"), - PathBuf::from("/nonexistent/tokenizer.json"), - ) + Classifier::load("08:00") } #[test] fn active_provider_is_valid() { // The active provider depends on the host: a machine with the ONNX model present and - // a working MIGraphX iGPU loads `Gpu`, otherwise `Cpu`. Either is valid — but when no + // a working ROCm iGPU loads `Gpu`, otherwise `Cpu`. Either is valid — but when no // model is available we must be on CPU (no session => no GPU EP in use). - let c = Classifier::load("08:00"); + let c = cl(); assert!(matches!( c.active_provider, ExecutionProvider::Cpu | ExecutionProvider::Gpu @@ -56,28 +49,19 @@ fn classify_reminder_via_fallback() { #[test] fn classify_idea_via_fallback() { let mut c = cl(); - assert_eq!( - c.classify("what if we added a calendar view").note_type, - NoteType::Idea - ); + assert_eq!(c.classify("what if we added a calendar view").note_type, NoteType::Idea); } #[test] fn classify_question_via_fallback() { let mut c = cl(); - assert_eq!( - c.classify("why does this fail?").note_type, - NoteType::Question - ); + assert_eq!(c.classify("why does this fail?").note_type, NoteType::Question); } #[test] fn classify_note_via_fallback() { let mut c = cl(); - assert_eq!( - c.classify("meeting went well today").note_type, - NoteType::Note - ); + assert_eq!(c.classify("meeting went well today").note_type, NoteType::Note); } #[test] @@ -90,11 +74,7 @@ fn classify_recurrence_via_fallback() { #[test] fn classify_custom_morning_time() { - let mut c = Classifier::load_with_paths( - "07:15", - PathBuf::from("/nonexistent/classifier.onnx"), - PathBuf::from("/nonexistent/tokenizer.json"), - ); + let mut c = Classifier::load("07:15"); let r = c.classify("sync tomorrow morning"); let t = r.time.expect("should have a time for tomorrow morning"); let local: chrono::DateTime = t.into(); @@ -134,16 +114,12 @@ fn classify_returns_cleaned_body() { let mut c = cl(); let r = c.classify("call mum at 6pm"); assert!(r.body.contains("call mum"), "body: {}", r.body); - assert!( - !r.body.contains("6pm"), - "time phrase should be stripped from body: {}", - r.body - ); + assert!(!r.body.contains("6pm"), "time phrase should be stripped from body: {}", r.body); } #[test] fn model_path_points_to_expected_location() { - let c = Classifier::load("08:00"); + let c = cl(); assert!( c.model_path.to_str().unwrap().contains("breadpad"), "model path: {:?}", diff --git a/breadpad-shared/tests/pipeline.rs b/breadpad-shared/tests/pipeline.rs index 603a199..84cc850 100644 --- a/breadpad-shared/tests/pipeline.rs +++ b/breadpad-shared/tests/pipeline.rs @@ -9,18 +9,12 @@ use breadpad_shared::classifier::Classifier; use breadpad_shared::store::Store; use breadpad_shared::types::{Note, NoteType}; use chrono::Timelike; -use std::path::PathBuf; use tempfile::TempDir; // Mirrors commit_note() in breadpad/src/main.rs. // `user_type` is the type the user selected in the chip row (default = NoteType::Note). fn capture(store: &Store, text: &str, user_type: NoteType) -> Note { - // WHY: pipeline tests cover classify→save→reload, not a host ONNX model. - let mut classifier = Classifier::load_with_paths( - "08:00", - PathBuf::from("/nonexistent/classifier.onnx"), - PathBuf::from("/nonexistent/tokenizer.json"), - ); + let mut classifier = Classifier::load("08:00"); let result = classifier.classify(text); let mut note = Note::new(text.into(), user_type.clone(), None); @@ -67,11 +61,7 @@ fn todo_note_appears_in_store() { #[test] fn idea_note_appears_in_store() { let (dir, store) = setup(); - capture( - &store, - "what if we added dark mode", - NoteType::from_str("note"), - ); + capture(&store, "what if we added dark mode", NoteType::from_str("note")); let notes = breadman_store(&dir).load_all().unwrap(); assert_eq!(notes.len(), 1); @@ -81,11 +71,7 @@ fn idea_note_appears_in_store() { #[test] fn question_note_appears_in_store() { let (dir, store) = setup(); - capture( - &store, - "why does the cache miss on cold start?", - NoteType::from_str("note"), - ); + capture(&store, "why does the cache miss on cold start?", NoteType::from_str("note")); let notes = breadman_store(&dir).load_all().unwrap(); assert_eq!(notes.len(), 1); @@ -111,10 +97,7 @@ fn reminder_has_time_set() { let notes = breadman_store(&dir).load_all().unwrap(); assert_eq!(notes[0].note_type, NoteType::Reminder); - assert!( - notes[0].time.is_some(), - "reminder should have a scheduled time" - ); + assert!(notes[0].time.is_some(), "reminder should have a scheduled time"); let local: chrono::DateTime = notes[0].time.unwrap().into(); assert_eq!(local.hour(), 18); } @@ -125,21 +108,14 @@ fn reminder_body_has_time_stripped() { capture(&store, "call mum at 6pm", NoteType::from_str("note")); let notes = breadman_store(&dir).load_all().unwrap(); - assert!( - !notes[0].body.contains("6pm"), - "time phrase should be removed from body" - ); + assert!(!notes[0].body.contains("6pm"), "time phrase should be removed from body"); assert!(notes[0].body.contains("call mum")); } #[test] fn in_duration_reminder_has_time() { let (dir, store) = setup(); - capture( - &store, - "check on the build in 30 minutes", - NoteType::from_str("note"), - ); + capture(&store, "check on the build in 30 minutes", NoteType::from_str("note")); let notes = breadman_store(&dir).load_all().unwrap(); assert_eq!(notes[0].note_type, NoteType::Reminder); @@ -151,11 +127,7 @@ fn in_duration_reminder_has_time() { #[test] fn recurring_reminder_has_rrule() { let (dir, store) = setup(); - capture( - &store, - "standup every monday at 9am", - NoteType::from_str("note"), - ); + capture(&store, "standup every monday at 9am", NoteType::from_str("note")); let notes = breadman_store(&dir).load_all().unwrap(); assert_eq!(notes[0].note_type, NoteType::Reminder); @@ -167,20 +139,11 @@ fn recurring_reminder_has_rrule() { #[test] fn daily_reminder_has_rrule() { let (dir, store) = setup(); - capture( - &store, - "drink water every day at 8am", - NoteType::from_str("note"), - ); + capture(&store, "drink water every day at 8am", NoteType::from_str("note")); let notes = breadman_store(&dir).load_all().unwrap(); assert_eq!(notes[0].note_type, NoteType::Reminder); - assert!(notes[0] - .rrule - .as_ref() - .unwrap() - .as_str() - .contains("FREQ=DAILY")); + assert!(notes[0].rrule.as_ref().unwrap().as_str().contains("FREQ=DAILY")); } // ---- user-forced type is respected ---- @@ -192,11 +155,7 @@ fn user_selected_type_overrides_classifier() { capture(&store, "fix the login bug", NoteType::Idea); let notes = breadman_store(&dir).load_all().unwrap(); - assert_eq!( - notes[0].note_type, - NoteType::Idea, - "user chip selection should win over classifier" - ); + assert_eq!(notes[0].note_type, NoteType::Idea, "user chip selection should win over classifier"); } #[test] @@ -214,11 +173,7 @@ fn user_selected_reminder_overrides_classifier() { fn three_notes_all_visible_to_breadman() { let (dir, store) = setup(); capture(&store, "buy milk", NoteType::from_str("note")); - capture( - &store, - "what if we rewrote in Zig", - NoteType::from_str("note"), - ); + capture(&store, "what if we rewrote in Zig", NoteType::from_str("note")); capture(&store, "team standup went well", NoteType::from_str("note")); let notes = breadman_store(&dir).load_all().unwrap();