Add model management and macOS integration
This commit is contained in:
316
src/app.rs
316
src/app.rs
@@ -3,14 +3,13 @@ mod view;
|
||||
pub(crate) use view::app_theme;
|
||||
|
||||
use crate::database::{AppPreferences, Database, ProjectWithSessions};
|
||||
use crate::model::{self, DownloadOutcome, DownloadProgress, ModelChoice};
|
||||
use iced::futures::{SinkExt, Stream};
|
||||
use iced::{Subscription, Task, keyboard};
|
||||
use crate::model::{self, DownloadOutcome, DownloadProgress, ManagedArtifactId, ModelChoice};
|
||||
use iced::{Size, Subscription, Task, keyboard, window};
|
||||
use rfd::AsyncFileDialog;
|
||||
use std::fs;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
|
||||
use std::sync::mpsc::{self, TryRecvError};
|
||||
use std::thread;
|
||||
use std::time::{Duration, Instant};
|
||||
@@ -37,6 +36,11 @@ impl PreferenceDraft {
|
||||
}
|
||||
|
||||
pub(crate) struct App {
|
||||
main_window: window::Id,
|
||||
pub(super) model_manager_window: Option<window::Id>,
|
||||
pub(super) pending_model_delete: Option<ManagedArtifactId>,
|
||||
#[cfg(target_os = "macos")]
|
||||
_native_menu: Option<crate::native_menu::NativeMenu>,
|
||||
database: Option<Database>,
|
||||
projects: Vec<ProjectWithSessions>,
|
||||
preferences: AppPreferences,
|
||||
@@ -56,18 +60,25 @@ pub(crate) struct App {
|
||||
pub(super) enum ModelDownload {
|
||||
Idle,
|
||||
Active(ActiveDownload),
|
||||
Complete(ModelChoice, DownloadProgress),
|
||||
Failed(ModelChoice, String, DownloadProgress),
|
||||
Complete(ManagedArtifactId, ModelOperation, DownloadProgress),
|
||||
Failed(ManagedArtifactId, String, DownloadProgress),
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub(super) enum ModelOperation {
|
||||
Download,
|
||||
Validate,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(super) struct ActiveDownload {
|
||||
model: ModelChoice,
|
||||
dspark_enabled: bool,
|
||||
artifact: ManagedArtifactId,
|
||||
operation: ModelOperation,
|
||||
progress: DownloadProgress,
|
||||
sampled_at: Instant,
|
||||
sampled_bytes: u64,
|
||||
bytes_per_second: f64,
|
||||
verified_bytes: Arc<AtomicU64>,
|
||||
cancel: Arc<AtomicBool>,
|
||||
result: mpsc::Receiver<Result<DownloadOutcome, String>>,
|
||||
stopping: bool,
|
||||
@@ -75,13 +86,22 @@ pub(super) struct ActiveDownload {
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) enum Message {
|
||||
Noop,
|
||||
OpenPreferences,
|
||||
OpenModelManager,
|
||||
ModelManagerOpened(window::Id),
|
||||
WindowOpened(window::Id),
|
||||
WindowClosed(window::Id),
|
||||
DismissPanel,
|
||||
PreferenceModelChanged(ModelChoice),
|
||||
PreferenceDsparkChanged(bool),
|
||||
PreferenceTimeoutChanged(String),
|
||||
SavePreferences,
|
||||
DownloadSelectedModel,
|
||||
DownloadArtifact(ManagedArtifactId),
|
||||
ValidateArtifact(ManagedArtifactId),
|
||||
DeleteArtifact(ManagedArtifactId),
|
||||
ConfirmDeleteArtifact,
|
||||
CancelDeleteArtifact,
|
||||
StopModelDownload,
|
||||
DownloadProgressTick,
|
||||
ChooseProjectFolder,
|
||||
@@ -97,16 +117,21 @@ pub(crate) enum Message {
|
||||
}
|
||||
|
||||
impl App {
|
||||
pub(crate) fn load() -> Self {
|
||||
pub(crate) fn load(main_window: window::Id) -> Self {
|
||||
let path = application_support_path().join("data.sqlite3");
|
||||
match Database::open(&path) {
|
||||
Ok(mut database) => match (database.load_projects(), database.load_preferences()) {
|
||||
(Ok(projects), Ok(preferences)) => {
|
||||
let preference_draft = match PreferenceDraft::from_saved(&preferences) {
|
||||
Ok(draft) => draft,
|
||||
Err(error) => return Self::failed(error),
|
||||
Err(error) => return Self::failed(error, main_window),
|
||||
};
|
||||
Self {
|
||||
main_window,
|
||||
model_manager_window: None,
|
||||
pending_model_delete: None,
|
||||
#[cfg(target_os = "macos")]
|
||||
_native_menu: None,
|
||||
database: Some(database),
|
||||
projects,
|
||||
preferences,
|
||||
@@ -122,17 +147,22 @@ impl App {
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
(Err(error), _) | (_, Err(error)) => Self::failed(error),
|
||||
(Err(error), _) | (_, Err(error)) => Self::failed(error, main_window),
|
||||
},
|
||||
Err(error) => Self::failed(error),
|
||||
Err(error) => Self::failed(error, main_window),
|
||||
}
|
||||
}
|
||||
|
||||
fn failed(error: String) -> Self {
|
||||
fn failed(error: String, main_window: window::Id) -> Self {
|
||||
let preferences = AppPreferences::default();
|
||||
let preference_draft = PreferenceDraft::from_saved(&preferences)
|
||||
.expect("default preferences must use a supported model");
|
||||
Self {
|
||||
main_window,
|
||||
model_manager_window: None,
|
||||
pending_model_delete: None,
|
||||
#[cfg(target_os = "macos")]
|
||||
_native_menu: None,
|
||||
database: None,
|
||||
projects: Vec::new(),
|
||||
preferences,
|
||||
@@ -151,7 +181,36 @@ impl App {
|
||||
|
||||
pub(crate) fn update(&mut self, message: Message) -> Task<Message> {
|
||||
match message {
|
||||
Message::Noop => {}
|
||||
Message::OpenPreferences => self.open_preferences(),
|
||||
Message::OpenModelManager => return self.open_model_manager(),
|
||||
Message::ModelManagerOpened(id) => {
|
||||
if self.model_manager_window == Some(id) {
|
||||
return window::gain_focus(id);
|
||||
}
|
||||
}
|
||||
Message::WindowOpened(id) =>
|
||||
{
|
||||
#[cfg(target_os = "macos")]
|
||||
if id == self.main_window && self._native_menu.is_none() {
|
||||
match crate::native_menu::install() {
|
||||
Ok(menu) => self._native_menu = Some(menu),
|
||||
Err(error) => {
|
||||
self.error =
|
||||
Some(format!("Could not install the application menu: {error}"))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Message::WindowClosed(id) => {
|
||||
if id == self.main_window {
|
||||
return iced::exit();
|
||||
}
|
||||
if self.model_manager_window == Some(id) {
|
||||
self.model_manager_window = None;
|
||||
self.pending_model_delete = None;
|
||||
}
|
||||
}
|
||||
Message::DismissPanel => {
|
||||
if self.preferences_open {
|
||||
self.preferences_open = false;
|
||||
@@ -179,41 +238,27 @@ impl App {
|
||||
self.preference_error = None;
|
||||
}
|
||||
Message::SavePreferences => self.save_preferences(),
|
||||
Message::DownloadSelectedModel => {
|
||||
if matches!(self.model_download, ModelDownload::Active(_)) {
|
||||
return Task::none();
|
||||
}
|
||||
let model = self.preference_draft.model;
|
||||
let dspark_enabled =
|
||||
model.supports_dspark() && self.preference_draft.dspark_enabled;
|
||||
let models_path = models_path();
|
||||
let progress = model::download_progress(model, dspark_enabled, &models_path);
|
||||
let cancel = Arc::new(AtomicBool::new(false));
|
||||
let worker_cancel = Arc::clone(&cancel);
|
||||
let (result_sender, result_receiver) = mpsc::channel();
|
||||
if let Err(error) = thread::Builder::new()
|
||||
.name("model-download".to_owned())
|
||||
.spawn(move || {
|
||||
let result =
|
||||
model::download(model, dspark_enabled, &models_path, &worker_cancel);
|
||||
let _ = result_sender.send(result);
|
||||
})
|
||||
{
|
||||
self.error = Some(format!("Could not start model download: {error}"));
|
||||
return Task::none();
|
||||
}
|
||||
self.model_download = ModelDownload::Active(ActiveDownload {
|
||||
model,
|
||||
dspark_enabled,
|
||||
sampled_at: Instant::now(),
|
||||
sampled_bytes: progress.downloaded,
|
||||
progress,
|
||||
bytes_per_second: 0.0,
|
||||
cancel: Arc::clone(&cancel),
|
||||
result: result_receiver,
|
||||
stopping: false,
|
||||
});
|
||||
Message::DownloadArtifact(artifact) => {
|
||||
self.start_model_operation(artifact, ModelOperation::Download)
|
||||
}
|
||||
Message::ValidateArtifact(artifact) => {
|
||||
self.start_model_operation(artifact, ModelOperation::Validate)
|
||||
}
|
||||
Message::DeleteArtifact(artifact) => self.pending_model_delete = Some(artifact),
|
||||
Message::ConfirmDeleteArtifact => {
|
||||
if let Some(artifact) = self.pending_model_delete.take() {
|
||||
match model::delete_managed_artifact(artifact, &models_path()) {
|
||||
Ok(()) => {
|
||||
self.model_download = ModelDownload::Idle;
|
||||
self.error = None;
|
||||
}
|
||||
Err(error) => {
|
||||
self.error = Some(format!("Could not delete {artifact}: {error}"))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Message::CancelDeleteArtifact => self.pending_model_delete = None,
|
||||
Message::StopModelDownload => {
|
||||
if let ModelDownload::Active(download) = &mut self.model_download {
|
||||
download.stopping = true;
|
||||
@@ -290,15 +335,105 @@ impl App {
|
||||
}
|
||||
|
||||
pub(crate) fn subscription(&self) -> Subscription<Message> {
|
||||
let shortcuts = keyboard::on_key_press(shortcut);
|
||||
let mut subscriptions = vec![
|
||||
keyboard::on_key_press(shortcut),
|
||||
window::close_requests().map(Message::WindowClosed),
|
||||
window::close_events().map(Message::WindowClosed),
|
||||
];
|
||||
#[cfg(target_os = "macos")]
|
||||
subscriptions.push(iced::time::every(Duration::from_millis(50)).map(|_| {
|
||||
match crate::native_menu::next_event() {
|
||||
Some(crate::native_menu::NativeMenuEvent::Preferences) => Message::OpenPreferences,
|
||||
Some(crate::native_menu::NativeMenuEvent::ModelManager) => {
|
||||
Message::OpenModelManager
|
||||
}
|
||||
None => Message::Noop,
|
||||
}
|
||||
}));
|
||||
if matches!(self.model_download, ModelDownload::Active(_)) {
|
||||
Subscription::batch([
|
||||
shortcuts,
|
||||
Subscription::run(download_ticks).map(|_| Message::DownloadProgressTick),
|
||||
])
|
||||
} else {
|
||||
shortcuts
|
||||
subscriptions.push(
|
||||
iced::time::every(Duration::from_secs(1)).map(|_| Message::DownloadProgressTick),
|
||||
);
|
||||
}
|
||||
Subscription::batch(subscriptions)
|
||||
}
|
||||
|
||||
pub(crate) fn title(&self, id: window::Id) -> String {
|
||||
if self.model_manager_window == Some(id) {
|
||||
"Model Manager — DS4Server".to_owned()
|
||||
} else {
|
||||
"DS4Server".to_owned()
|
||||
}
|
||||
}
|
||||
|
||||
fn open_model_manager(&mut self) -> Task<Message> {
|
||||
if let Some(id) = self.model_manager_window {
|
||||
return window::gain_focus(id);
|
||||
}
|
||||
let (id, open) = window::open(window::Settings {
|
||||
size: Size::new(760.0, 560.0),
|
||||
min_size: Some(Size::new(620.0, 420.0)),
|
||||
icon: Some(app_icon()),
|
||||
..Default::default()
|
||||
});
|
||||
self.model_manager_window = Some(id);
|
||||
open.map(Message::ModelManagerOpened)
|
||||
}
|
||||
|
||||
fn start_model_operation(&mut self, artifact: ManagedArtifactId, operation: ModelOperation) {
|
||||
if matches!(self.model_download, ModelDownload::Active(_)) {
|
||||
return;
|
||||
}
|
||||
let models_path = models_path();
|
||||
let progress = match operation {
|
||||
ModelOperation::Download => model::artifact_download_progress(artifact, &models_path),
|
||||
ModelOperation::Validate => model::artifact_verification_progress(artifact, 0),
|
||||
};
|
||||
let cancel = Arc::new(AtomicBool::new(false));
|
||||
let worker_cancel = Arc::clone(&cancel);
|
||||
let verified_bytes = Arc::new(AtomicU64::new(0));
|
||||
let worker_verified_bytes = Arc::clone(&verified_bytes);
|
||||
let (result_sender, result_receiver) = mpsc::channel();
|
||||
let thread_name = match operation {
|
||||
ModelOperation::Download => "model-download",
|
||||
ModelOperation::Validate => "model-validation",
|
||||
};
|
||||
if let Err(error) = thread::Builder::new()
|
||||
.name(thread_name.to_owned())
|
||||
.spawn(move || {
|
||||
let result = match operation {
|
||||
ModelOperation::Download => model::download_managed_artifact(
|
||||
artifact,
|
||||
&models_path,
|
||||
&worker_cancel,
|
||||
&worker_verified_bytes,
|
||||
),
|
||||
ModelOperation::Validate => model::validate_managed_artifact(
|
||||
artifact,
|
||||
&models_path,
|
||||
&worker_cancel,
|
||||
&worker_verified_bytes,
|
||||
),
|
||||
};
|
||||
let _ = result_sender.send(result);
|
||||
})
|
||||
{
|
||||
self.error = Some(format!("Could not start {thread_name}: {error}"));
|
||||
return;
|
||||
}
|
||||
self.pending_model_delete = None;
|
||||
self.model_download = ModelDownload::Active(ActiveDownload {
|
||||
artifact,
|
||||
operation,
|
||||
sampled_at: Instant::now(),
|
||||
sampled_bytes: progress.completed(),
|
||||
progress,
|
||||
bytes_per_second: 0.0,
|
||||
verified_bytes,
|
||||
cancel,
|
||||
result: result_receiver,
|
||||
stopping: false,
|
||||
});
|
||||
}
|
||||
|
||||
fn open_preferences(&mut self) {
|
||||
@@ -443,11 +578,27 @@ impl App {
|
||||
return;
|
||||
};
|
||||
let now = Instant::now();
|
||||
let progress =
|
||||
model::download_progress(download.model, download.dspark_enabled, &models_path());
|
||||
let verified = download.verified_bytes.load(Ordering::Relaxed);
|
||||
let mut progress = match download.operation {
|
||||
ModelOperation::Download => {
|
||||
model::artifact_download_progress(download.artifact, &models_path())
|
||||
}
|
||||
ModelOperation::Validate => {
|
||||
model::artifact_verification_progress(download.artifact, verified)
|
||||
}
|
||||
};
|
||||
if download.operation == ModelOperation::Download
|
||||
&& let Some(verification) = &mut progress.verification
|
||||
{
|
||||
verification.verified = verified.min(verification.total);
|
||||
}
|
||||
let elapsed = now.duration_since(download.sampled_at).as_secs_f64();
|
||||
let transferred = progress.downloaded.saturating_sub(download.sampled_bytes);
|
||||
if transferred > 0 && elapsed > 0.0 {
|
||||
let completed = progress.completed();
|
||||
let phase_changed = progress.phase != download.progress.phase;
|
||||
let transferred = completed.saturating_sub(download.sampled_bytes);
|
||||
if phase_changed {
|
||||
download.bytes_per_second = 0.0;
|
||||
} else if transferred > 0 && elapsed > 0.0 {
|
||||
let current = transferred as f64 / elapsed;
|
||||
download.bytes_per_second = if download.bytes_per_second == 0.0 {
|
||||
current
|
||||
@@ -457,7 +608,7 @@ impl App {
|
||||
}
|
||||
download.progress = progress;
|
||||
download.sampled_at = now;
|
||||
download.sampled_bytes = download.progress.downloaded;
|
||||
download.sampled_bytes = completed;
|
||||
|
||||
let result = match download.result.try_recv() {
|
||||
Ok(result) => Some(result),
|
||||
@@ -466,12 +617,14 @@ impl App {
|
||||
Some(Err("Download worker stopped unexpectedly.".into()))
|
||||
}
|
||||
};
|
||||
let model = download.model;
|
||||
let artifact = download.artifact;
|
||||
let operation = download.operation;
|
||||
let progress = download.progress.clone();
|
||||
if let Some(result) = result {
|
||||
match result {
|
||||
Ok(DownloadOutcome::Complete) => {
|
||||
self.model_download = ModelDownload::Complete(model, progress);
|
||||
let progress = model::artifact_download_progress(artifact, &models_path());
|
||||
self.model_download = ModelDownload::Complete(artifact, operation, progress);
|
||||
self.error = None;
|
||||
}
|
||||
Ok(DownloadOutcome::Stopped) => {
|
||||
@@ -480,7 +633,7 @@ impl App {
|
||||
}
|
||||
Err(error) => {
|
||||
self.error = Some(error.clone());
|
||||
self.model_download = ModelDownload::Failed(model, error, progress);
|
||||
self.model_download = ModelDownload::Failed(artifact, error, progress);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -495,20 +648,12 @@ impl Drop for App {
|
||||
}
|
||||
}
|
||||
|
||||
fn download_ticks() -> impl Stream<Item = ()> {
|
||||
iced::stream::channel(1, |mut output| async move {
|
||||
loop {
|
||||
std::thread::sleep(Duration::from_secs(1));
|
||||
if output.send(()).await.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn shortcut(key: keyboard::Key, modifiers: keyboard::Modifiers) -> Option<Message> {
|
||||
match key.as_ref() {
|
||||
keyboard::Key::Character(",") if modifiers.command() => Some(Message::OpenPreferences),
|
||||
keyboard::Key::Character("m") if modifiers.command() && modifiers.shift() => {
|
||||
Some(Message::OpenModelManager)
|
||||
}
|
||||
keyboard::Key::Named(keyboard::key::Named::Escape) => Some(Message::DismissPanel),
|
||||
_ => None,
|
||||
}
|
||||
@@ -527,6 +672,24 @@ fn models_path() -> PathBuf {
|
||||
application_support_path().join("models")
|
||||
}
|
||||
|
||||
pub(crate) fn app_icon() -> window::Icon {
|
||||
let decoder = png::Decoder::new(std::io::Cursor::new(include_bytes!(
|
||||
"../assets/app-icon.png"
|
||||
)));
|
||||
let mut reader = decoder
|
||||
.read_info()
|
||||
.expect("bundled application icon must be valid PNG");
|
||||
let mut rgba = vec![0; reader.output_buffer_size()];
|
||||
let info = reader
|
||||
.next_frame(&mut rgba)
|
||||
.expect("bundled application icon must decode");
|
||||
assert_eq!(info.color_type, png::ColorType::Rgba);
|
||||
assert_eq!(info.bit_depth, png::BitDepth::Eight);
|
||||
rgba.truncate(info.buffer_size());
|
||||
window::icon::from_rgba(rgba, info.width, info.height)
|
||||
.expect("bundled application icon dimensions must be valid")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -538,6 +701,11 @@ mod tests {
|
||||
keyboard::Modifiers::COMMAND,
|
||||
);
|
||||
assert!(matches!(message, Some(Message::OpenPreferences)));
|
||||
let message = shortcut(
|
||||
keyboard::Key::Character("m".into()),
|
||||
keyboard::Modifiers::COMMAND | keyboard::Modifiers::SHIFT,
|
||||
);
|
||||
assert!(matches!(message, Some(Message::OpenModelManager)));
|
||||
assert!(ModelChoice::DeepSeekV4Flash.supports_dspark());
|
||||
assert!(!ModelChoice::DeepSeekV4Pro.supports_dspark());
|
||||
assert!(!ModelChoice::Glm52.supports_dspark());
|
||||
|
||||
Reference in New Issue
Block a user