From 39e26a1b9f9e52ecf2b8474bbb98d5fbd2331193 Mon Sep 17 00:00:00 2001 From: Breadway Date: Sun, 23 Aug 2026 14:42:57 +0800 Subject: [PATCH] Switch classifier sessions onto bread-onnx v0.7.2 Replace local ort Session::builder/ROCm wiring with bread_onnx::build_session (MIGraphX + CPU fallback, same pin as breadmill). Drop the breadman chip/init_adw shim now that bread-theme v0.7.4 exports those helpers. --- Cargo.lock | 123 ++++++++++++++++++++++------ Cargo.toml | 7 +- breadman/Cargo.toml | 2 +- breadman/src/editor.rs | 8 +- breadman/src/main.rs | 11 ++- breadman/src/theme_widgets.rs | 23 ------ breadman/src/views/settings.rs | 8 +- breadpad-shared/Cargo.toml | 2 +- breadpad-shared/src/classifier.rs | 98 +++++++++++----------- breadpad-shared/tests/classifier.rs | 42 ++++++++-- breadpad-shared/tests/pipeline.rs | 67 ++++++++++++--- 11 files changed, 254 insertions(+), 137 deletions(-) delete mode 100644 breadman/src/theme_widgets.rs diff --git a/Cargo.lock b/Cargo.lock index 940600b..4d85626 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -302,6 +302,21 @@ 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" @@ -393,12 +408,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", @@ -664,6 +679,18 @@ 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" @@ -764,6 +791,12 @@ 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" @@ -994,6 +1027,9 @@ 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" @@ -1017,9 +1053,9 @@ dependencies = [ [[package]] name = "fancy-regex" -version = "0.14.0" +version = "0.17.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6e24cb5a94bcae1e5408b0effca5cd7172ea3c5755049c5f3af4cd283a165298" +checksum = "72cf461f865c862bb7dc573f643dd6a2b6842f7c30b07882b56bd148cc2761b8" dependencies = [ "bit-set", "regex-automata", @@ -1530,7 +1566,7 @@ checksum = "629d8f3bbeda9d148036d6b0de0a3ab947abd08ce90626327fc3547a49d59d97" dependencies = [ "dirs 6.0.0", "http", - "indicatif", + "indicatif 0.17.11", "libc", "log", "rand 0.9.5", @@ -1827,13 +1863,26 @@ version = "0.17.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "183b3088984b400f4cfac3620d5e076c84da5364016b4f49473de574b2586235" dependencies = [ - "console", + "console 0.15.11", "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" @@ -2077,21 +2126,6 @@ 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" @@ -2184,6 +2218,28 @@ 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" @@ -2207,7 +2263,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4336a1e2b38848325241c72889086886004e589b7c74f335e60a8e8db5138a0b" dependencies = [ "libloading", - "ndarray 0.17.2", + "ndarray", "ort-sys", "smallvec", "tracing", @@ -2938,6 +2994,17 @@ 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" @@ -3207,23 +3274,25 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokenizers" -version = "0.21.4" +version = "0.23.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a620b996116a59e184c2fa2dfd8251ea34a36d0a514758c6f966386bd2e03476" +checksum = "44e5bea67576e04b6ff8564c5d9e09c2ef0cf476502245f2f120e497769d3112" 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", @@ -3538,6 +3607,12 @@ 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 84621db..b9508be 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -24,9 +24,10 @@ chrono = { version = "0.4", features = ["serde"] } rrule = "0.12" tokio = { version = "1", features = ["full"] } zbus = { version = "4", default-features = false, features = ["tokio"] } -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"] } +# 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"] } 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 e81b7e6..d47a586 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`. Local chip/init shims remain in src/theme_widgets.rs. +# `adw` implies `gtk` (`chip`, `set_chip_active`, `adw::init`). 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 bea9803..e39b11c 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| (crate::theme_widgets::chip(name), name)).collect(); + let type_pills: Vec<(gtk4::Button, &'static str)> = NoteType::all_builtin().iter().map(|&name| (bread_theme::gtk::chip(name), name)).collect(); for (btn, name) in &type_pills { - crate::theme_widgets::set_chip_active(btn, *name == selected_type.borrow().as_str()); + bread_theme::gtk::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 { crate::theme_widgets::set_chip_active(b, false); } - crate::theme_widgets::set_chip_active(clicked, true); + for b in &all_btns { bread_theme::gtk::set_chip_active(b, false); } + bread_theme::gtk::set_chip_active(clicked, true); }); type_pill_box.append(btn); } diff --git a/breadman/src/main.rs b/breadman/src/main.rs index 0a4f810..efd5b3f 100644 --- a/breadman/src/main.rs +++ b/breadman/src/main.rs @@ -14,7 +14,6 @@ use std::sync::Arc; mod editor; mod screenshot; -mod theme_widgets; mod views; // ── Args ───────────────────────────────────────────────────────────────────── @@ -363,7 +362,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. - theme_widgets::init_adw(); + bread_theme::adw::init(); let store = Arc::new(Store::new()?); let notes = store.load_all()?; @@ -645,17 +644,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| (theme_widgets::chip(name), NoteType::from_str(name))) + .map(|&name| (bread_theme::gtk::chip(name), NoteType::from_str(name))) .collect(); for (btn, nt) in &chips { - theme_widgets::set_chip_active(btn, *nt == preselect); + bread_theme::gtk::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 { theme_widgets::set_chip_active(b, false); } - theme_widgets::set_chip_active(clicked, true); + for b in &all_btns { bread_theme::gtk::set_chip_active(b, false); } + bread_theme::gtk::set_chip_active(clicked, true); }); chip_box.append(btn); } diff --git a/breadman/src/theme_widgets.rs b/breadman/src/theme_widgets.rs deleted file mode 100644 index a5221e6..0000000 --- a/breadman/src/theme_widgets.rs +++ /dev/null @@ -1,23 +0,0 @@ -//! 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 114edd1..6090274 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| (crate::theme_widgets::chip(name), name)) + .map(|&name| (bread_theme::gtk::chip(name), name)) .collect(); for (btn, name) in &type_pills { - crate::theme_widgets::set_chip_active(btn, *name == selected_type.borrow().as_str()); + bread_theme::gtk::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 { crate::theme_widgets::set_chip_active(b, false); } - crate::theme_widgets::set_chip_active(clicked, true); + for b in &all_btns { bread_theme::gtk::set_chip_active(b, false); } + bread_theme::gtk::set_chip_active(clicked, true); apply_now(); }); } diff --git a/breadpad-shared/Cargo.toml b/breadpad-shared/Cargo.toml index 30a1467..a56c187 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 -ndarray.workspace = true +bread-onnx = { git = "https://git.breadway.dev/Breadway/bread-ecosystem", tag = "v0.7.2" } toml.workspace = true dirs.workspace = true regex.workspace = true diff --git a/breadpad-shared/src/classifier.rs b/breadpad-shared/src/classifier.rs index e752a31..35c2650 100644 --- a/breadpad-shared/src/classifier.rs +++ b/breadpad-shared/src/classifier.rs @@ -2,6 +2,8 @@ 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. @@ -16,7 +18,7 @@ pub enum ExecutionProvider { impl ExecutionProvider { pub fn as_str(&self) -> &str { match self { - ExecutionProvider::Gpu => "ROCm (iGPU)", + ExecutionProvider::Gpu => "MIGraphX (iGPU)", ExecutionProvider::Cpu => "CPU", } } @@ -107,9 +109,7 @@ 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,9 +163,18 @@ 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"), ]; @@ -184,15 +193,17 @@ 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, @@ -207,10 +218,7 @@ 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 @@ -243,42 +251,30 @@ fn softmax_single(logits: &[f32], idx: usize) -> f32 { exps[idx] / sum } -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()) { +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 }], + ) { Ok(s) => { - tracing::info!("ONNX session loaded (CPU)"); - (Some(s), ExecutionProvider::Cpu) + tracing::info!("ONNX session loaded (MIGraphX, CPU fallback)"); + (Some(s), ExecutionProvider::Gpu) } Err(e) => { - tracing::warn!("failed to load ONNX session: {}; Tier 2 disabled", e); - (None, ExecutionProvider::Cpu) + 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) + } + } } } } - -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 3b5527c..7f0294a 100644 --- a/breadpad-shared/tests/classifier.rs +++ b/breadpad-shared/tests/classifier.rs @@ -1,17 +1,24 @@ 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("08:00") + Classifier::load_with_paths( + "08:00", + PathBuf::from("/nonexistent/classifier.onnx"), + PathBuf::from("/nonexistent/tokenizer.json"), + ) } #[test] fn active_provider_is_valid() { // The active provider depends on the host: a machine with the ONNX model present and - // a working ROCm iGPU loads `Gpu`, otherwise `Cpu`. Either is valid — but when no + // a working MIGraphX 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 = cl(); + let c = Classifier::load("08:00"); assert!(matches!( c.active_provider, ExecutionProvider::Cpu | ExecutionProvider::Gpu @@ -49,19 +56,28 @@ 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] @@ -74,7 +90,11 @@ fn classify_recurrence_via_fallback() { #[test] fn classify_custom_morning_time() { - let mut c = Classifier::load("07:15"); + let mut c = Classifier::load_with_paths( + "07:15", + PathBuf::from("/nonexistent/classifier.onnx"), + PathBuf::from("/nonexistent/tokenizer.json"), + ); 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(); @@ -114,12 +134,16 @@ 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 = cl(); + let c = Classifier::load("08:00"); 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 84cc850..603a199 100644 --- a/breadpad-shared/tests/pipeline.rs +++ b/breadpad-shared/tests/pipeline.rs @@ -9,12 +9,18 @@ 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 { - let mut classifier = Classifier::load("08:00"); + // 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 result = classifier.classify(text); let mut note = Note::new(text.into(), user_type.clone(), None); @@ -61,7 +67,11 @@ 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); @@ -71,7 +81,11 @@ 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); @@ -97,7 +111,10 @@ 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); } @@ -108,14 +125,21 @@ 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); @@ -127,7 +151,11 @@ 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); @@ -139,11 +167,20 @@ 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 ---- @@ -155,7 +192,11 @@ 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] @@ -173,7 +214,11 @@ 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();