From 2e515a618c791da2cccbce012cdf898557650424 Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Thu, 6 Aug 2026 18:36:46 +0300 Subject: [PATCH 01/10] Add saving last used model --- .../interactive/components/app_settings.rs | 28 +++++++++++++++++++ .../src/interactive/components/application.rs | 4 +++ crates/cli/src/interactive/components/mod.rs | 2 ++ crates/cli/src/interactive/flows/models.rs | 17 +++++++++-- crates/cli/src/interactive/mod.rs | 21 +++++++++++++- crates/cli/src/main.rs | 14 +--------- 6 files changed, 70 insertions(+), 16 deletions(-) create mode 100644 crates/cli/src/interactive/components/app_settings.rs diff --git a/crates/cli/src/interactive/components/app_settings.rs b/crates/cli/src/interactive/components/app_settings.rs new file mode 100644 index 000000000..f1f1ce239 --- /dev/null +++ b/crates/cli/src/interactive/components/app_settings.rs @@ -0,0 +1,28 @@ +use serde::{Deserialize, Serialize}; +use uzu::settings::{SettingKind, Settings, SettingsError}; + +const SETTINGS_APP: &str = "app"; + +#[derive(Clone, Serialize, Deserialize, Default)] +pub struct AppSettings { + pub selected_model_id: Option, +} + +impl AppSettings { + pub fn load(settings: &Settings) -> Result { + let Some(raw) = settings.load(SettingKind::Config, SETTINGS_APP.to_string())? else { + return Ok(Self::default()); + }; + Ok(serde_json::from_str(&raw).unwrap_or_default()) + } + + pub fn save( + &self, + settings: &Settings, + ) -> Result<(), SettingsError> { + let raw = serde_json::to_string(self).map_err(|error| SettingsError::BackendError { + message: error.to_string(), + })?; + settings.save(SettingKind::Config, SETTINGS_APP.to_string(), Some(raw)) + } +} diff --git a/crates/cli/src/interactive/components/application.rs b/crates/cli/src/interactive/components/application.rs index 2b9107e6c..d30d026bf 100644 --- a/crates/cli/src/interactive/components/application.rs +++ b/crates/cli/src/interactive/components/application.rs @@ -9,6 +9,7 @@ use uzu::{ use crate::interactive::{ components::{ CommandInput, HistoryCell, HistoryCellType, Logo, ModelCapabilities, Preferences, SelectedModel, Theme, + app_settings::AppSettings, }, flows::{AuthFlow, ExitFlow, Flow, FlowEvent, FlowRegistry, ModelRegistriesFlow, SettingsFlow, ThemeFlow}, helpers::SYMBOL_COMMAND, @@ -23,6 +24,7 @@ pub struct ApplicationProps { pub settings: Option, pub theme: Option, pub preferences: Option, + pub app_settings: Option, pub model: Option, } @@ -38,6 +40,7 @@ pub struct ApplicationState { pub settings: Option, pub theme: Theme, pub preferences: Preferences, + pub app_settings: AppSettings, pub flow: Option>, pub history: Vec, pub registry: FlowRegistry, @@ -62,6 +65,7 @@ pub fn Application( settings: props.settings.clone(), theme: props.theme.clone().unwrap_or_default(), preferences: props.preferences.unwrap_or_default(), + app_settings: props.app_settings.clone().unwrap_or_default(), flow: None, history: Vec::new(), registry: FlowRegistry::default() diff --git a/crates/cli/src/interactive/components/mod.rs b/crates/cli/src/interactive/components/mod.rs index 9efc7f36f..b49b9e178 100644 --- a/crates/cli/src/interactive/components/mod.rs +++ b/crates/cli/src/interactive/components/mod.rs @@ -1,3 +1,4 @@ +mod app_settings; mod application; mod command_input; mod gradient; @@ -13,6 +14,7 @@ mod selector; mod text_input; mod theme; +pub use app_settings::AppSettings; pub use application::{Application, ApplicationState, ModelState}; pub use command_input::CommandInput; pub use gradient::Gradient; diff --git a/crates/cli/src/interactive/flows/models.rs b/crates/cli/src/interactive/flows/models.rs index a085bd390..7d3a964ea 100644 --- a/crates/cli/src/interactive/flows/models.rs +++ b/crates/cli/src/interactive/flows/models.rs @@ -108,14 +108,27 @@ fn Models( on_submit: move |index: usize| { let mut state = state; if let Some(model) = list.get(index) { - let summary = format!("Model: {}", model.name()); state.write().model_state = Some(ModelState { model: model.clone(), download_state: DownloadState::not_downloaded(0), session_state: None, capabilities: ModelCapabilities::default(), }); - on_event(FlowEvent::finish(summary)); + state.write().app_settings.selected_model_id = Some(model.identifier.clone()); + let settings_result = { + let state = state.read(); + match state.settings.as_ref() { + Some(settings) => state.app_settings.save(settings), + None => Ok(()), + } + }; + let result = match settings_result { + Ok(()) => format!("Model: {}", model.name()), + Err(error) => { + format!("Model: {}, unable to save preference: {}", model.name(), error) + }, + }; + on_event(FlowEvent::finish(result)); } }, ) diff --git a/crates/cli/src/interactive/mod.rs b/crates/cli/src/interactive/mod.rs index 650c04670..50850089d 100644 --- a/crates/cli/src/interactive/mod.rs +++ b/crates/cli/src/interactive/mod.rs @@ -12,6 +12,8 @@ use uzu::{ settings::SettingsError, }; +use crate::interactive::components::AppSettings; + #[derive(Debug, Clone, PartialEq, thiserror::Error)] #[non_exhaustive] pub enum CliError { @@ -56,6 +58,15 @@ impl CliApplication { Some(settings) => Preferences::load(settings)?, None => Preferences::default(), }; + let app_settings = match &settings { + Some(settings) => AppSettings::load(settings)?, + None => AppSettings::default(), + }; + + let mut selected_model = model; + if selected_model.is_none() { + selected_model = app_settings.selected_model_id.clone(); + } element! { Application( @@ -63,7 +74,8 @@ impl CliApplication { settings: settings, theme: Some(theme), preferences: Some(preferences), - model: model, + app_settings: Some(app_settings), + model: selected_model, ) } .render_loop() @@ -82,3 +94,10 @@ impl CliApplication { Ok(Self::new(engine)) } } + +pub async fn run_interactive(model: Option) -> anyhow::Result<()> { + let engine_config = EngineConfig::default().with_application_identifier("com.trymirai.cli".to_string()); + let application = CliApplication::create(engine_config).await?; + application.run_with_model(model).await?; + Ok(()) +} diff --git a/crates/cli/src/main.rs b/crates/cli/src/main.rs index 5564d6ab7..6c8ec4d65 100644 --- a/crates/cli/src/main.rs +++ b/crates/cli/src/main.rs @@ -55,20 +55,8 @@ async fn main() -> Result<()> { Some(Commands::Storage { download_manager, }) => storage::run(download_manager).await?, - None => run_interactive(cli.model).await?, + None => interactive::run_interactive(cli.model).await?, } Ok(()) } - -async fn run_interactive(model: Option) -> Result<()> { - use uzu::engine::EngineConfig; - - use crate::interactive::CliApplication; - - let engine_config = EngineConfig::default().with_application_identifier("com.trymirai.cli".to_string()); - let application = CliApplication::create(engine_config).await?; - application.run_with_model(model).await?; - - Ok(()) -} From 7b0eb70268b5f994c9d85945970023fce0d3e03c Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Thu, 6 Aug 2026 18:44:08 +0300 Subject: [PATCH 02/10] Show models registry if model is not set or were removed --- .../src/interactive/components/application.rs | 28 +++++++++++++++---- 1 file changed, 22 insertions(+), 6 deletions(-) diff --git a/crates/cli/src/interactive/components/application.rs b/crates/cli/src/interactive/components/application.rs index d30d026bf..45b0a67e6 100644 --- a/crates/cli/src/interactive/components/application.rs +++ b/crates/cli/src/interactive/components/application.rs @@ -84,10 +84,20 @@ pub fn Application( let mut state = state; async move { let Some(identifier) = initial_model else { + state.write().flow = Some(Box::new(ModelRegistriesFlow)); return; }; match engine.model(identifier.clone()).await { Ok(Some(model)) => { + let model_exists = !model.is_local() + || matches!( + engine.model_path(&model).await, + Some(path) if std::path::Path::new(&path).exists() + ); + if !model_exists { + state.write().flow = Some(Box::new(ModelRegistriesFlow)); + return; + } let summary = format!("Model: {}", model.name()); state.write().model_state = Some(ModelState { model, @@ -99,12 +109,18 @@ pub fn Application( result: summary, }); }, - Ok(None) => state.write().history.push(HistoryCellType::CommandResult { - result: format!("Unknown model: {}", identifier), - }), - Err(error) => state.write().history.push(HistoryCellType::CommandResult { - result: format!("Failed to load model {}: {}", identifier, error), - }), + Ok(None) => { + state.write().history.push(HistoryCellType::CommandResult { + result: format!("Unknown model: {}", identifier), + }); + state.write().flow = Some(Box::new(ModelRegistriesFlow)); + }, + Err(error) => { + state.write().history.push(HistoryCellType::CommandResult { + result: format!("Failed to load model {}: {}", identifier, error), + }); + state.write().flow = Some(Box::new(ModelRegistriesFlow)); + }, } } }); From 83d01376a24ed90ad099d3a6f262919be4fd20aa Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Fri, 7 Aug 2026 13:34:45 +0300 Subject: [PATCH 03/10] Fix checking if model exists --- crates/cli/src/interactive/components/application.rs | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/crates/cli/src/interactive/components/application.rs b/crates/cli/src/interactive/components/application.rs index 45b0a67e6..f39397123 100644 --- a/crates/cli/src/interactive/components/application.rs +++ b/crates/cli/src/interactive/components/application.rs @@ -89,7 +89,8 @@ pub fn Application( }; match engine.model(identifier.clone()).await { Ok(Some(model)) => { - let model_exists = !model.is_local() + let model_exists = model.is_downloadable() + || !model.is_local() || matches!( engine.model_path(&model).await, Some(path) if std::path::Path::new(&path).exists() From d1a095fd21bd423051db60df6251beac1b6609e1 Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Fri, 7 Aug 2026 13:42:35 +0300 Subject: [PATCH 04/10] Move CliApplication to separate file --- crates/cli/src/interactive/app.rs | 88 ++++++++++++++++++++++++++++ crates/cli/src/interactive/mod.rs | 97 ++----------------------------- 2 files changed, 93 insertions(+), 92 deletions(-) create mode 100644 crates/cli/src/interactive/app.rs diff --git a/crates/cli/src/interactive/app.rs b/crates/cli/src/interactive/app.rs new file mode 100644 index 000000000..8b56be1e4 --- /dev/null +++ b/crates/cli/src/interactive/app.rs @@ -0,0 +1,88 @@ +use std::io::IsTerminal; + +use iocraft::prelude::*; +use uzu::{ + engine::{Engine, EngineConfig, EngineError}, + settings::SettingsError, +}; + +use crate::interactive::components::{AppSettings, Application, Preferences, Theme}; + +#[derive(Debug, Clone, PartialEq, thiserror::Error)] +#[non_exhaustive] +pub enum CliError { + #[error(transparent)] + Engine(#[from] EngineError), + #[error(transparent)] + Settigs(#[from] SettingsError), + #[error("Rendering error: {message}")] + RenderingError { + message: String, + }, +} + +#[derive(Clone)] +pub struct CliApplication { + engine: Engine, +} + +impl CliApplication { + pub async fn create(config: EngineConfig) -> Result { + let engine = Engine::new(config).await?; + Ok(Self::new(engine)) + } + + pub fn new(engine: Engine) -> Self { + Self { + engine, + } + } + + pub async fn run_with_model( + &self, + model: Option, + ) -> Result<(), CliError> { + if !std::io::stdout().is_terminal() { + return Err(CliError::RenderingError { + message: "stdout is not a terminal".to_string(), + }); + } + + let settings = self.engine.settings().await.ok(); + let theme = match &settings { + Some(settings) => Theme::load(settings)?.unwrap_or_default(), + None => Theme::default(), + }; + let preferences = match &settings { + Some(settings) => Preferences::load(settings)?, + None => Preferences::default(), + }; + let app_settings = match &settings { + Some(settings) => AppSettings::load(settings)?, + None => AppSettings::default(), + }; + + let mut selected_model = model; + if selected_model.is_none() { + selected_model = app_settings.selected_model_id.clone(); + } + + element! { + Application( + engine: Some(self.engine.clone()), + settings: settings, + theme: Some(theme), + preferences: Some(preferences), + app_settings: Some(app_settings), + model: selected_model, + ) + } + .render_loop() + .await + .map_err(|error| CliError::RenderingError { + message: error.to_string(), + })?; + + Ok(()) + } +} diff --git a/crates/cli/src/interactive/mod.rs b/crates/cli/src/interactive/mod.rs index 50850089d..61ec7e652 100644 --- a/crates/cli/src/interactive/mod.rs +++ b/crates/cli/src/interactive/mod.rs @@ -1,100 +1,13 @@ +use uzu::engine::EngineConfig; + +use crate::interactive::app::CliApplication; + +mod app; mod components; mod flows; mod helpers; mod sessions; -use std::io::IsTerminal; - -use components::{Application, Preferences, Theme}; -use iocraft::prelude::*; -use uzu::{ - engine::{Engine, EngineConfig, EngineError}, - settings::SettingsError, -}; - -use crate::interactive::components::AppSettings; - -#[derive(Debug, Clone, PartialEq, thiserror::Error)] -#[non_exhaustive] -pub enum CliError { - #[error(transparent)] - Engine(#[from] EngineError), - #[error(transparent)] - Settigs(#[from] SettingsError), - #[error("Rendering error: {message}")] - RenderingError { - message: String, - }, -} - -#[derive(Clone)] -pub struct CliApplication { - engine: Engine, -} - -impl CliApplication { - pub fn new(engine: Engine) -> Self { - Self { - engine, - } - } - - pub async fn run_with_model( - &self, - model: Option, - ) -> Result<(), CliError> { - if !std::io::stdout().is_terminal() { - return Err(CliError::RenderingError { - message: "stdout is not a terminal".to_string(), - }); - } - - let settings = self.engine.settings().await.ok(); - let theme = match &settings { - Some(settings) => Theme::load(settings)?.unwrap_or_default(), - None => Theme::default(), - }; - let preferences = match &settings { - Some(settings) => Preferences::load(settings)?, - None => Preferences::default(), - }; - let app_settings = match &settings { - Some(settings) => AppSettings::load(settings)?, - None => AppSettings::default(), - }; - - let mut selected_model = model; - if selected_model.is_none() { - selected_model = app_settings.selected_model_id.clone(); - } - - element! { - Application( - engine: Some(self.engine.clone()), - settings: settings, - theme: Some(theme), - preferences: Some(preferences), - app_settings: Some(app_settings), - model: selected_model, - ) - } - .render_loop() - .await - .map_err(|error| CliError::RenderingError { - message: error.to_string(), - })?; - - Ok(()) - } -} - -impl CliApplication { - pub async fn create(config: EngineConfig) -> Result { - let engine = Engine::new(config).await?; - Ok(Self::new(engine)) - } -} - pub async fn run_interactive(model: Option) -> anyhow::Result<()> { let engine_config = EngineConfig::default().with_application_identifier("com.trymirai.cli".to_string()); let application = CliApplication::create(engine_config).await?; From 587da096713dbc59593a8112c3287a634524eadc Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Fri, 7 Aug 2026 15:45:54 +0300 Subject: [PATCH 05/10] Add list-models --- crates/cli/src/interactive/list.rs | 30 +++++++++++++++++++++++++ crates/cli/src/interactive/mod.rs | 35 ++++++++++++++++++++++++++++-- crates/cli/src/main.rs | 2 ++ 3 files changed, 65 insertions(+), 2 deletions(-) create mode 100644 crates/cli/src/interactive/list.rs diff --git a/crates/cli/src/interactive/list.rs b/crates/cli/src/interactive/list.rs new file mode 100644 index 000000000..62a754911 --- /dev/null +++ b/crates/cli/src/interactive/list.rs @@ -0,0 +1,30 @@ +use std::collections::HashSet; + +use shoji::types::model::Model; + +pub struct ModelFamily { + pub id: String, + pub name: String, +} + +pub fn get_families(models: &[Model]) -> Vec { + let mut families = Vec::::new(); + let mut ids_set = HashSet::::new(); + + for model in models.iter() { + if let Some(ref family) = model.family { + if let Some(ref properties) = model.properties { + let model_id = format!("{}:{}", family.identifier.split(":").last().unwrap(), properties.identifier); + if !ids_set.contains(&model_id) { + ids_set.insert(model_id.clone()); + families.push(ModelFamily { + id: model_id, + name: format!("{} {}", family.metadata.name, properties.metadata.name), + }) + } + } + } + } + + families +} diff --git a/crates/cli/src/interactive/mod.rs b/crates/cli/src/interactive/mod.rs index 61ec7e652..7fcc4f920 100644 --- a/crates/cli/src/interactive/mod.rs +++ b/crates/cli/src/interactive/mod.rs @@ -1,11 +1,17 @@ -use uzu::engine::EngineConfig; +use comfy_table::{ + ContentArrangement, Table, + modifiers::{UTF8_ROUND_CORNERS, UTF8_SOLID_INNER_BORDERS}, + presets::UTF8_FULL, +}; +use uzu::engine::{Engine, EngineConfig}; -use crate::interactive::app::CliApplication; +use crate::interactive::{app::CliApplication, list::get_families}; mod app; mod components; mod flows; mod helpers; +mod list; mod sessions; pub async fn run_interactive(model: Option) -> anyhow::Result<()> { @@ -14,3 +20,28 @@ pub async fn run_interactive(model: Option) -> anyhow::Result<()> { application.run_with_model(model).await?; Ok(()) } + +pub async fn run_list_models() -> anyhow::Result<()> { + let engine_config = EngineConfig::default().with_application_identifier("com.trymirai.cli".to_string()); + let engine = Engine::new(engine_config).await?; + let models = engine.models().await?; + if models.is_empty() { + return Err(anyhow::anyhow!("No models to run")); + } + + let families = get_families(&models); + let mut table = Table::new(); + table + .load_preset(UTF8_FULL) + .apply_modifier(UTF8_ROUND_CORNERS) + .apply_modifier(UTF8_SOLID_INNER_BORDERS) + .set_content_arrangement(ContentArrangement::Dynamic) + .set_header(vec!["Name", "ID"]); + + for family in &families { + table.add_row(vec![&family.name, &family.id]); + } + println!("{table}"); + + Ok(()) +} diff --git a/crates/cli/src/main.rs b/crates/cli/src/main.rs index 6c8ec4d65..a41b1bdbf 100644 --- a/crates/cli/src/main.rs +++ b/crates/cli/src/main.rs @@ -23,6 +23,7 @@ enum Commands { task_path: String, output_path: String, }, + ListModels, Server { #[arg(long, value_name = "MODEL")] model: String, @@ -47,6 +48,7 @@ async fn main() -> Result<()> { task_path, output_path, }) => bench::run_bench(model_path, task_path, output_path).await?, + Some(Commands::ListModels {}) => interactive::run_list_models().await?, Some(Commands::Server { model, port, From 5ca377c1885c4d8695a0a19ba35b33e898236049 Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Fri, 7 Aug 2026 15:58:57 +0300 Subject: [PATCH 06/10] Add list-checkpoints --- crates/cli/src/interactive/list.rs | 29 +++++++++++++++++++++++++++++ crates/cli/src/interactive/mod.rs | 30 +++++++++++++++++++++++++++++- crates/cli/src/main.rs | 8 ++++++++ 3 files changed, 66 insertions(+), 1 deletion(-) diff --git a/crates/cli/src/interactive/list.rs b/crates/cli/src/interactive/list.rs index 62a754911..4fda1139d 100644 --- a/crates/cli/src/interactive/list.rs +++ b/crates/cli/src/interactive/list.rs @@ -2,6 +2,35 @@ use std::collections::HashSet; use shoji::types::model::Model; +pub struct ModelCheckpoint { + pub id: String, + pub name: String, +} + +pub fn get_checkpoints( + models: &[Model], + model_id: &str, +) -> Vec { + models + .iter() + .filter(|model| { + let Some(family) = &model.family else { + return false; + }; + let Some(properties) = &model.properties else { + return false; + }; + + let family_id = family.identifier.rsplit(':').next().unwrap_or(&family.identifier); + format!("{family_id}:{}", properties.identifier) == model_id + }) + .map(|model| ModelCheckpoint { + id: model.identifier.clone(), + name: model.name(), + }) + .collect() +} + pub struct ModelFamily { pub id: String, pub name: String, diff --git a/crates/cli/src/interactive/mod.rs b/crates/cli/src/interactive/mod.rs index 7fcc4f920..3c8c042c9 100644 --- a/crates/cli/src/interactive/mod.rs +++ b/crates/cli/src/interactive/mod.rs @@ -5,7 +5,10 @@ use comfy_table::{ }; use uzu::engine::{Engine, EngineConfig}; -use crate::interactive::{app::CliApplication, list::get_families}; +use crate::interactive::{ + app::CliApplication, + list::{get_checkpoints, get_families}, +}; mod app; mod components; @@ -45,3 +48,28 @@ pub async fn run_list_models() -> anyhow::Result<()> { Ok(()) } + +pub async fn run_list_checkpoints(model_id: String) -> anyhow::Result<()> { + let engine_config = EngineConfig::default().with_application_identifier("com.trymirai.cli".to_string()); + let engine = Engine::new(engine_config).await?; + let models = engine.models().await?; + let checkpoints = get_checkpoints(&models, &model_id); + if checkpoints.is_empty() { + return Err(anyhow::anyhow!("No checkpoints found for model: {model_id}")); + } + + let mut table = Table::new(); + table + .load_preset(UTF8_FULL) + .apply_modifier(UTF8_ROUND_CORNERS) + .apply_modifier(UTF8_SOLID_INNER_BORDERS) + .set_content_arrangement(ContentArrangement::Dynamic) + .set_header(vec!["Name", "ID"]); + + for checkpoint in &checkpoints { + table.add_row(vec![&checkpoint.name, &checkpoint.id]); + } + println!("{table}"); + + Ok(()) +} diff --git a/crates/cli/src/main.rs b/crates/cli/src/main.rs index a41b1bdbf..b2971685e 100644 --- a/crates/cli/src/main.rs +++ b/crates/cli/src/main.rs @@ -23,6 +23,11 @@ enum Commands { task_path: String, output_path: String, }, + ListCheckpoints { + /// Model ID shown by `list-models`. + #[arg(value_name = "MODEL_ID")] + model_id: String, + }, ListModels, Server { #[arg(long, value_name = "MODEL")] @@ -48,6 +53,9 @@ async fn main() -> Result<()> { task_path, output_path, }) => bench::run_bench(model_path, task_path, output_path).await?, + Some(Commands::ListCheckpoints { + model_id, + }) => interactive::run_list_checkpoints(model_id).await?, Some(Commands::ListModels {}) => interactive::run_list_models().await?, Some(Commands::Server { model, From 2a2f21062a3c48b3e3ca71d8ba183e08fe9a3747 Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Fri, 7 Aug 2026 17:02:33 +0300 Subject: [PATCH 07/10] Add models resolving --- crates/cli/src/interactive/app.rs | 14 ++- crates/cli/src/interactive/mod.rs | 1 + crates/cli/src/interactive/model.rs | 137 ++++++++++++++++++++++ crates/cli/unit/interactive/model_test.rs | 126 ++++++++++++++++++++ 4 files changed, 273 insertions(+), 5 deletions(-) create mode 100644 crates/cli/src/interactive/model.rs create mode 100644 crates/cli/unit/interactive/model_test.rs diff --git a/crates/cli/src/interactive/app.rs b/crates/cli/src/interactive/app.rs index 8b56be1e4..27a72015d 100644 --- a/crates/cli/src/interactive/app.rs +++ b/crates/cli/src/interactive/app.rs @@ -6,7 +6,10 @@ use uzu::{ settings::SettingsError, }; -use crate::interactive::components::{AppSettings, Application, Preferences, Theme}; +use crate::interactive::{ + components::{AppSettings, Application, Preferences, Theme}, + model::resolve_model_id, +}; #[derive(Debug, Clone, PartialEq, thiserror::Error)] #[non_exhaustive] @@ -62,10 +65,11 @@ impl CliApplication { None => AppSettings::default(), }; - let mut selected_model = model; - if selected_model.is_none() { - selected_model = app_settings.selected_model_id.clone(); - } + let requested_model = model.or_else(|| app_settings.selected_model_id.clone()); + let selected_model = match requested_model { + Some(model) => resolve_model_id(&self.engine, model).await?, + None => None, + }; element! { Application( diff --git a/crates/cli/src/interactive/mod.rs b/crates/cli/src/interactive/mod.rs index 3c8c042c9..a35a04ef4 100644 --- a/crates/cli/src/interactive/mod.rs +++ b/crates/cli/src/interactive/mod.rs @@ -15,6 +15,7 @@ mod components; mod flows; mod helpers; mod list; +mod model; mod sessions; pub async fn run_interactive(model: Option) -> anyhow::Result<()> { diff --git a/crates/cli/src/interactive/model.rs b/crates/cli/src/interactive/model.rs new file mode 100644 index 000000000..580407599 --- /dev/null +++ b/crates/cli/src/interactive/model.rs @@ -0,0 +1,137 @@ +use shoji::types::model::{Model, ModelAccessibility, ModelReference}; +use sysinfo::System; +use uzu::engine::{Engine, EngineError}; + +pub async fn resolve_model_id( + engine: &Engine, + model: String, +) -> Result, EngineError> { + let models = engine.models().await?; + let mut system = System::new(); + system.refresh_memory(); + let resolved = models + .iter() + .find(|candidate| candidate.identifier == model) + .or_else(|| resolve_model_shorthand(&models, &model, system.total_memory())) + .map(|model| model.identifier.clone()) + .unwrap_or(model); + Ok(Some(resolved)) +} + +fn resolve_model_shorthand<'a>( + models: &'a [Model], + requested: &str, + memory_total: u64, +) -> Option<&'a Model> { + let mut candidates = models + .iter() + .filter(|model| model_shorthand_matches(model, requested)) + .filter(|model| checkpoint_size_bytes(model).is_some_and(|size| size <= memory_total)) + .collect::>(); + + candidates.sort_by(|left, right| { + quantization_priority(left) + .cmp(&quantization_priority(right)) + .then_with(|| model_parameter_count(right).cmp(&model_parameter_count(left))) + .then_with(|| quantization_bits(right).cmp(&quantization_bits(left))) + .then_with(|| left.identifier.cmp(&right.identifier)) + }); + candidates.into_iter().next() +} + +fn model_shorthand_matches( + model: &Model, + requested: &str, +) -> bool { + let Some(family) = &model.family else { + return false; + }; + let Some(properties) = &model.properties else { + return false; + }; + + let vendor = family.vendor.identifier.as_str(); + let family_id = family.identifier.strip_prefix(&format!("{vendor}:")).unwrap_or(&family.identifier); + let size = properties.identifier.as_str(); + let bits = model.quantization.as_ref().map(|quantization| quantization.bits_per_weight.to_string()); + + for include_vendor in [false, true] { + for include_size in [false, true] { + for include_quantization_vendor in [false, true] { + for include_quantization in [false, true] { + for include_bits in [false, true] { + let mut parts = Vec::with_capacity(6); + if include_vendor { + parts.push(vendor); + } + parts.push(family_id); + if include_size { + parts.push(size); + } + + if include_quantization_vendor || include_quantization || include_bits { + let Some(quantization) = &model.quantization else { + continue; + }; + if include_quantization_vendor { + parts.push(quantization.vendor.identifier.as_str()); + } + if include_quantization { + parts.push(quantization.method.as_str()); + } + if include_bits { + parts.push(bits.as_deref().expect("quantization bits are available")); + } + } + + if parts.join(":") == requested { + return true; + } + } + } + } + } + } + + false +} + +fn checkpoint_size_bytes(model: &Model) -> Option { + if let ModelAccessibility::Local { + reference: ModelReference::Mirai { + files, + .. + }, + } = &model.accessibility + && !files.is_empty() + { + return files.iter().try_fold(0_u64, |total, file| { + let size = u64::try_from(file.size).ok()?; + total.checked_add(size) + }); + } + + let parameters = u64::try_from(model.properties.as_ref()?.size).ok()?; + let bits = u64::from(model.quantization.as_ref().map_or(16, |quantization| quantization.bits_per_weight)); + Some(parameters.checked_mul(bits)?.div_ceil(8)) +} + +fn model_parameter_count(model: &Model) -> i64 { + model.properties.as_ref().map_or(0, |properties| properties.size) +} + +fn quantization_priority(model: &Model) -> u8 { + match model.quantization.as_ref().map(|quantization| quantization.method.as_str()) { + Some(method) if method.eq_ignore_ascii_case("mirai-m") => 0, + Some(method) if method.eq_ignore_ascii_case("mirai-s") => 1, + _ => 2, + } +} + +fn quantization_bits(model: &Model) -> u32 { + model.quantization.as_ref().map_or(16, |quantization| quantization.bits_per_weight) +} + +#[cfg(test)] +#[path = "../../unit/interactive/model_test.rs"] +mod tests; diff --git a/crates/cli/unit/interactive/model_test.rs b/crates/cli/unit/interactive/model_test.rs new file mode 100644 index 000000000..97aec0577 --- /dev/null +++ b/crates/cli/unit/interactive/model_test.rs @@ -0,0 +1,126 @@ +use shoji::types::{ + basic::{File, Metadata}, + model::{ModelFamily, ModelProperties, ModelQuantization, ModelRegistry, ModelVendor}, +}; + +use super::*; + +const GB: i64 = 1_000_000_000; + +#[test] +fn family_only_prefers_largest_mirai_model_that_fits() { + let models = vec![ + checkpoint("0.8b", 800_000_000, "mirai", "mirai-m", 4, GB), + checkpoint("4b", 4 * GB, "mirai", "mirai-m", 4, 4 * GB), + checkpoint("27b", 27 * GB, "mlx-community", "mlx", 4, 16 * GB), + ]; + + let resolved = resolve_model_shorthand(&models, "qwen3.5", 32 * GB as u64).unwrap(); + + assert_eq!(resolved.properties.as_ref().unwrap().identifier, "4b"); + assert_eq!(resolved.quantization.as_ref().unwrap().method, "mirai-m"); +} + +#[test] +fn family_only_uses_largest_checkpoint_when_no_mirai_model_exists() { + let models = vec![ + checkpoint("4b", 4 * GB, "mlx-community", "mlx", 8, 4 * GB), + checkpoint("27b", 27 * GB, "mlx-community", "mlx", 4, 16 * GB), + ]; + + let resolved = resolve_model_shorthand(&models, "qwen3.5", 32 * GB as u64).unwrap(); + + assert_eq!(resolved.properties.as_ref().unwrap().identifier, "27b"); +} + +#[test] +fn shorthand_prefers_mirai_m_then_mirai_s() { + let mlx = checkpoint("4b", 4 * GB, "mlx-community", "mlx", 8, 4 * GB); + let mirai_s = checkpoint("4b", 4 * GB, "mirai", "mirai-s", 4, 4 * GB); + let mirai_m = checkpoint("4b", 4 * GB, "mirai", "mirai-m", 4, 4 * GB); + + let models = vec![mlx.clone(), mirai_s.clone(), mirai_m.clone()]; + assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b", 8 * GB as u64), Some(&mirai_m)); + + let models = vec![mlx, mirai_s.clone()]; + assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b", 8 * GB as u64), Some(&mirai_s)); +} + +#[test] +fn shorthand_infers_family_and_quantization_vendors() { + let model = checkpoint("4b", 4 * GB, "mlx-community", "mlx", 8, 4 * GB); + let models = [model.clone()]; + + assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b:mlx:8", 8 * GB as u64), Some(&model)); + assert_eq!(resolve_model_shorthand(&models, "alibaba:qwen3.5:4b:mlx-community:mlx:8", 8 * GB as u64), Some(&model)); +} + +fn checkpoint( + size_id: &str, + parameters: i64, + quantization_vendor_id: &str, + quantization_method: &str, + bits_per_weight: u32, + checkpoint_size: i64, +) -> Model { + let vendor = ModelVendor { + identifier: "alibaba".to_string(), + metadata: metadata("Alibaba"), + }; + let quantization_vendor = ModelVendor { + identifier: quantization_vendor_id.to_string(), + metadata: metadata(quantization_vendor_id), + }; + Model { + identifier: format!( + "alibaba:qwen3.5:{size_id}:{quantization_vendor_id}:{quantization_method}:{bits_per_weight}" + ), + registry: ModelRegistry { + identifier: "mirai".to_string(), + metadata: metadata("Mirai"), + }, + backends: vec![], + family: Some(ModelFamily { + identifier: "alibaba:qwen3.5".to_string(), + vendor, + metadata: metadata("Qwen3.5"), + }), + properties: Some(ModelProperties { + identifier: size_id.to_string(), + size: parameters, + version: None, + metadata: metadata(size_id), + }), + quantization: Some(ModelQuantization { + identifier: format!("{quantization_vendor_id}:{quantization_method}:{bits_per_weight}"), + method: quantization_method.to_string(), + bits_per_weight, + vendor: quantization_vendor, + metadata: metadata(quantization_method), + }), + specializations: vec![], + accessibility: ModelAccessibility::Local { + reference: ModelReference::Mirai { + toolchain_version: "test".to_string(), + repository: None, + source_repository: None, + files: vec![File { + url: "https://example.com/model.safetensors".to_string(), + name: "model.safetensors".to_string(), + size: checkpoint_size, + hashes: vec![], + }], + }, + }, + encodings: vec![], + } +} + +fn metadata(name: &str) -> Metadata { + Metadata { + identifier: name.to_lowercase(), + name: name.to_string(), + description: None, + icons: vec![], + } +} From 8486f4bde87e285e6beebeea723b16a4de8485f8 Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Fri, 7 Aug 2026 17:26:31 +0300 Subject: [PATCH 08/10] Move app id to static var --- crates/cli/src/interactive/mod.rs | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/crates/cli/src/interactive/mod.rs b/crates/cli/src/interactive/mod.rs index a35a04ef4..8c30f4100 100644 --- a/crates/cli/src/interactive/mod.rs +++ b/crates/cli/src/interactive/mod.rs @@ -1,3 +1,4 @@ +use std::string::ToString; use comfy_table::{ ContentArrangement, Table, modifiers::{UTF8_ROUND_CORNERS, UTF8_SOLID_INNER_BORDERS}, @@ -18,15 +19,17 @@ mod list; mod model; mod sessions; +static APP_IDENTIFIER: String = "com.trymirai.cli".to_string(); + pub async fn run_interactive(model: Option) -> anyhow::Result<()> { - let engine_config = EngineConfig::default().with_application_identifier("com.trymirai.cli".to_string()); + let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER); let application = CliApplication::create(engine_config).await?; application.run_with_model(model).await?; Ok(()) } pub async fn run_list_models() -> anyhow::Result<()> { - let engine_config = EngineConfig::default().with_application_identifier("com.trymirai.cli".to_string()); + let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER); let engine = Engine::new(engine_config).await?; let models = engine.models().await?; if models.is_empty() { @@ -51,7 +54,7 @@ pub async fn run_list_models() -> anyhow::Result<()> { } pub async fn run_list_checkpoints(model_id: String) -> anyhow::Result<()> { - let engine_config = EngineConfig::default().with_application_identifier("com.trymirai.cli".to_string()); + let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER); let engine = Engine::new(engine_config).await?; let models = engine.models().await?; let checkpoints = get_checkpoints(&models, &model_id); From 795a169de9fbd8dc7f41777967fd6709ca42cc99 Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Fri, 7 Aug 2026 17:35:53 +0300 Subject: [PATCH 09/10] Handle known shorthands when no checkpoint fits --- crates/cli/src/interactive/app.rs | 4 +- crates/cli/src/interactive/mod.rs | 9 ++--- crates/cli/src/interactive/model.rs | 47 ++++++++++++++++------- crates/cli/unit/interactive/model_test.rs | 37 +++++++++++++++--- 4 files changed, 71 insertions(+), 26 deletions(-) diff --git a/crates/cli/src/interactive/app.rs b/crates/cli/src/interactive/app.rs index 27a72015d..ba79aeda4 100644 --- a/crates/cli/src/interactive/app.rs +++ b/crates/cli/src/interactive/app.rs @@ -8,7 +8,7 @@ use uzu::{ use crate::interactive::{ components::{AppSettings, Application, Preferences, Theme}, - model::resolve_model_id, + model::{ModelResolutionError, resolve_model_id}, }; #[derive(Debug, Clone, PartialEq, thiserror::Error)] @@ -17,6 +17,8 @@ pub enum CliError { #[error(transparent)] Engine(#[from] EngineError), #[error(transparent)] + ModelResolution(#[from] ModelResolutionError), + #[error(transparent)] Settigs(#[from] SettingsError), #[error("Rendering error: {message}")] RenderingError { diff --git a/crates/cli/src/interactive/mod.rs b/crates/cli/src/interactive/mod.rs index 8c30f4100..dffd62450 100644 --- a/crates/cli/src/interactive/mod.rs +++ b/crates/cli/src/interactive/mod.rs @@ -1,4 +1,3 @@ -use std::string::ToString; use comfy_table::{ ContentArrangement, Table, modifiers::{UTF8_ROUND_CORNERS, UTF8_SOLID_INNER_BORDERS}, @@ -19,17 +18,17 @@ mod list; mod model; mod sessions; -static APP_IDENTIFIER: String = "com.trymirai.cli".to_string(); +const APP_IDENTIFIER: &str = "com.trymirai.cli"; pub async fn run_interactive(model: Option) -> anyhow::Result<()> { - let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER); + let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER.to_string()); let application = CliApplication::create(engine_config).await?; application.run_with_model(model).await?; Ok(()) } pub async fn run_list_models() -> anyhow::Result<()> { - let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER); + let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER.to_string()); let engine = Engine::new(engine_config).await?; let models = engine.models().await?; if models.is_empty() { @@ -54,7 +53,7 @@ pub async fn run_list_models() -> anyhow::Result<()> { } pub async fn run_list_checkpoints(model_id: String) -> anyhow::Result<()> { - let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER); + let engine_config = EngineConfig::default().with_application_identifier(APP_IDENTIFIER.to_string()); let engine = Engine::new(engine_config).await?; let models = engine.models().await?; let checkpoints = get_checkpoints(&models, &model_id); diff --git a/crates/cli/src/interactive/model.rs b/crates/cli/src/interactive/model.rs index 580407599..cd124bb39 100644 --- a/crates/cli/src/interactive/model.rs +++ b/crates/cli/src/interactive/model.rs @@ -2,19 +2,31 @@ use shoji::types::model::{Model, ModelAccessibility, ModelReference}; use sysinfo::System; use uzu::engine::{Engine, EngineError}; +#[derive(Debug, Clone, PartialEq, thiserror::Error)] +pub enum ModelResolutionError { + #[error(transparent)] + Engine(#[from] EngineError), + #[error("model `{model}` has checkpoints, but none fit in {memory_total} bytes of total memory")] + InsufficientMemory { + model: String, + memory_total: u64, + }, +} + pub async fn resolve_model_id( engine: &Engine, model: String, -) -> Result, EngineError> { +) -> Result, ModelResolutionError> { let models = engine.models().await?; let mut system = System::new(); system.refresh_memory(); - let resolved = models - .iter() - .find(|candidate| candidate.identifier == model) - .or_else(|| resolve_model_shorthand(&models, &model, system.total_memory())) - .map(|model| model.identifier.clone()) - .unwrap_or(model); + let resolved = if let Some(model) = models.iter().find(|candidate| candidate.identifier == model) { + model.identifier.clone() + } else if let Some(model) = resolve_model_shorthand(&models, &model, system.total_memory())? { + model.identifier.clone() + } else { + model + }; Ok(Some(resolved)) } @@ -22,12 +34,19 @@ fn resolve_model_shorthand<'a>( models: &'a [Model], requested: &str, memory_total: u64, -) -> Option<&'a Model> { - let mut candidates = models - .iter() - .filter(|model| model_shorthand_matches(model, requested)) - .filter(|model| checkpoint_size_bytes(model).is_some_and(|size| size <= memory_total)) - .collect::>(); +) -> Result, ModelResolutionError> { + let mut candidates = models.iter().filter(|model| model_shorthand_matches(model, requested)).collect::>(); + if candidates.is_empty() { + return Ok(None); + } + + candidates.retain(|model| checkpoint_size_bytes(model).is_some_and(|size| size <= memory_total)); + if candidates.is_empty() { + return Err(ModelResolutionError::InsufficientMemory { + model: requested.to_string(), + memory_total, + }); + } candidates.sort_by(|left, right| { quantization_priority(left) @@ -36,7 +55,7 @@ fn resolve_model_shorthand<'a>( .then_with(|| quantization_bits(right).cmp(&quantization_bits(left))) .then_with(|| left.identifier.cmp(&right.identifier)) }); - candidates.into_iter().next() + Ok(candidates.into_iter().next()) } fn model_shorthand_matches( diff --git a/crates/cli/unit/interactive/model_test.rs b/crates/cli/unit/interactive/model_test.rs index 97aec0577..4a0ea8a91 100644 --- a/crates/cli/unit/interactive/model_test.rs +++ b/crates/cli/unit/interactive/model_test.rs @@ -15,7 +15,7 @@ fn family_only_prefers_largest_mirai_model_that_fits() { checkpoint("27b", 27 * GB, "mlx-community", "mlx", 4, 16 * GB), ]; - let resolved = resolve_model_shorthand(&models, "qwen3.5", 32 * GB as u64).unwrap(); + let resolved = resolve_model_shorthand(&models, "qwen3.5", 32 * GB as u64).unwrap().unwrap(); assert_eq!(resolved.properties.as_ref().unwrap().identifier, "4b"); assert_eq!(resolved.quantization.as_ref().unwrap().method, "mirai-m"); @@ -28,7 +28,7 @@ fn family_only_uses_largest_checkpoint_when_no_mirai_model_exists() { checkpoint("27b", 27 * GB, "mlx-community", "mlx", 4, 16 * GB), ]; - let resolved = resolve_model_shorthand(&models, "qwen3.5", 32 * GB as u64).unwrap(); + let resolved = resolve_model_shorthand(&models, "qwen3.5", 32 * GB as u64).unwrap().unwrap(); assert_eq!(resolved.properties.as_ref().unwrap().identifier, "27b"); } @@ -40,10 +40,10 @@ fn shorthand_prefers_mirai_m_then_mirai_s() { let mirai_m = checkpoint("4b", 4 * GB, "mirai", "mirai-m", 4, 4 * GB); let models = vec![mlx.clone(), mirai_s.clone(), mirai_m.clone()]; - assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b", 8 * GB as u64), Some(&mirai_m)); + assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b", 8 * GB as u64), Ok(Some(&mirai_m))); let models = vec![mlx, mirai_s.clone()]; - assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b", 8 * GB as u64), Some(&mirai_s)); + assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b", 8 * GB as u64), Ok(Some(&mirai_s))); } #[test] @@ -51,8 +51,33 @@ fn shorthand_infers_family_and_quantization_vendors() { let model = checkpoint("4b", 4 * GB, "mlx-community", "mlx", 8, 4 * GB); let models = [model.clone()]; - assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b:mlx:8", 8 * GB as u64), Some(&model)); - assert_eq!(resolve_model_shorthand(&models, "alibaba:qwen3.5:4b:mlx-community:mlx:8", 8 * GB as u64), Some(&model)); + assert_eq!(resolve_model_shorthand(&models, "qwen3.5:4b:mlx:8", 8 * GB as u64), Ok(Some(&model))); + assert_eq!( + resolve_model_shorthand(&models, "alibaba:qwen3.5:4b:mlx-community:mlx:8", 8 * GB as u64), + Ok(Some(&model)) + ); +} + +#[test] +fn valid_shorthand_returns_insufficient_memory_when_no_checkpoint_fits() { + let models = [checkpoint("4b", 4 * GB, "mirai", "mirai-m", 4, 4 * GB)]; + + let result = resolve_model_shorthand(&models, "qwen3.5:4b", GB as u64); + + assert_eq!( + result, + Err(ModelResolutionError::InsufficientMemory { + model: "qwen3.5:4b".to_string(), + memory_total: GB as u64, + }) + ); +} + +#[test] +fn unknown_shorthand_remains_unresolved() { + let models = [checkpoint("4b", 4 * GB, "mirai", "mirai-m", 4, 4 * GB)]; + + assert_eq!(resolve_model_shorthand(&models, "unknown", GB as u64), Ok(None)); } fn checkpoint( From 9a8ce0adb4cfee67de413fe510d7c04a5edc3dc1 Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Fri, 7 Aug 2026 20:52:35 +0300 Subject: [PATCH 10/10] Refactoring --- crates/cli/src/interactive/list.rs | 20 ++++++++++---------- crates/cli/src/main.rs | 2 +- 2 files changed, 11 insertions(+), 11 deletions(-) diff --git a/crates/cli/src/interactive/list.rs b/crates/cli/src/interactive/list.rs index 4fda1139d..ddb76267d 100644 --- a/crates/cli/src/interactive/list.rs +++ b/crates/cli/src/interactive/list.rs @@ -41,16 +41,16 @@ pub fn get_families(models: &[Model]) -> Vec { let mut ids_set = HashSet::::new(); for model in models.iter() { - if let Some(ref family) = model.family { - if let Some(ref properties) = model.properties { - let model_id = format!("{}:{}", family.identifier.split(":").last().unwrap(), properties.identifier); - if !ids_set.contains(&model_id) { - ids_set.insert(model_id.clone()); - families.push(ModelFamily { - id: model_id, - name: format!("{} {}", family.metadata.name, properties.metadata.name), - }) - } + if let Some(ref family) = model.family + && let Some(ref properties) = model.properties + { + let model_id = format!("{}:{}", family.identifier.split(":").last().unwrap(), properties.identifier); + if !ids_set.contains(&model_id) { + ids_set.insert(model_id.clone()); + families.push(ModelFamily { + id: model_id, + name: format!("{} {}", family.metadata.name, properties.metadata.name), + }) } } } diff --git a/crates/cli/src/main.rs b/crates/cli/src/main.rs index b2971685e..a1b2bf370 100644 --- a/crates/cli/src/main.rs +++ b/crates/cli/src/main.rs @@ -56,7 +56,7 @@ async fn main() -> Result<()> { Some(Commands::ListCheckpoints { model_id, }) => interactive::run_list_checkpoints(model_id).await?, - Some(Commands::ListModels {}) => interactive::run_list_models().await?, + Some(Commands::ListModels) => interactive::run_list_models().await?, Some(Commands::Server { model, port,