Add resumable background model downloads
This commit is contained in:
227
src/app.rs
227
src/app.rs
@@ -3,59 +3,24 @@ 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 rfd::AsyncFileDialog;
|
||||
use std::fmt;
|
||||
use std::fs;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::mpsc::{self, TryRecvError};
|
||||
use std::thread;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
const APP_ID: &str = "DS4Server.rfc1437.de";
|
||||
const MODEL_CHOICES: [ModelChoice; 3] = [
|
||||
ModelChoice::DeepSeekV4Flash,
|
||||
ModelChoice::DeepSeekV4Pro,
|
||||
ModelChoice::Glm52,
|
||||
];
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
|
||||
pub(crate) enum ModelChoice {
|
||||
#[default]
|
||||
DeepSeekV4Flash,
|
||||
DeepSeekV4Pro,
|
||||
Glm52,
|
||||
}
|
||||
|
||||
impl ModelChoice {
|
||||
fn id(self) -> &'static str {
|
||||
match self {
|
||||
Self::DeepSeekV4Flash => "deepseek-v4-flash",
|
||||
Self::DeepSeekV4Pro => "deepseek-v4-pro",
|
||||
Self::Glm52 => "glm-5.2",
|
||||
}
|
||||
}
|
||||
|
||||
fn from_id(id: &str) -> Option<Self> {
|
||||
MODEL_CHOICES.into_iter().find(|model| model.id() == id)
|
||||
}
|
||||
|
||||
fn supports_dflash(self) -> bool {
|
||||
self == Self::DeepSeekV4Flash
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for ModelChoice {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str(match self {
|
||||
Self::DeepSeekV4Flash => "DeepSeek V4 Flash",
|
||||
Self::DeepSeekV4Pro => "DeepSeek V4 Pro",
|
||||
Self::Glm52 => "GLM 5.2",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct PreferenceDraft {
|
||||
model: ModelChoice,
|
||||
dflash_enabled: bool,
|
||||
dspark_enabled: bool,
|
||||
idle_timeout_minutes: String,
|
||||
}
|
||||
|
||||
@@ -65,7 +30,7 @@ impl PreferenceDraft {
|
||||
.ok_or_else(|| format!("Unsupported model: {}", preferences.selected_model))?;
|
||||
Ok(Self {
|
||||
model,
|
||||
dflash_enabled: model.supports_dflash() && preferences.dflash_enabled,
|
||||
dspark_enabled: model.supports_dspark() && preferences.dspark_enabled,
|
||||
idle_timeout_minutes: preferences.idle_timeout_minutes.to_string(),
|
||||
})
|
||||
}
|
||||
@@ -83,17 +48,42 @@ pub(crate) struct App {
|
||||
choosing_folder: bool,
|
||||
pending_project_path: Option<PathBuf>,
|
||||
project_name_input: String,
|
||||
model_download: ModelDownload,
|
||||
error: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(super) enum ModelDownload {
|
||||
Idle,
|
||||
Active(ActiveDownload),
|
||||
Complete(ModelChoice, DownloadProgress),
|
||||
Failed(ModelChoice, String, DownloadProgress),
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(super) struct ActiveDownload {
|
||||
model: ModelChoice,
|
||||
dspark_enabled: bool,
|
||||
progress: DownloadProgress,
|
||||
sampled_at: Instant,
|
||||
sampled_bytes: u64,
|
||||
bytes_per_second: f64,
|
||||
cancel: Arc<AtomicBool>,
|
||||
result: mpsc::Receiver<Result<DownloadOutcome, String>>,
|
||||
stopping: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) enum Message {
|
||||
OpenPreferences,
|
||||
DismissPanel,
|
||||
PreferenceModelChanged(ModelChoice),
|
||||
PreferenceDflashChanged(bool),
|
||||
PreferenceDsparkChanged(bool),
|
||||
PreferenceTimeoutChanged(String),
|
||||
SavePreferences,
|
||||
DownloadSelectedModel,
|
||||
StopModelDownload,
|
||||
DownloadProgressTick,
|
||||
ChooseProjectFolder,
|
||||
ProjectFolderPicked(Option<PathBuf>),
|
||||
ProjectNameChanged(String),
|
||||
@@ -128,6 +118,7 @@ impl App {
|
||||
choosing_folder: false,
|
||||
pending_project_path: None,
|
||||
project_name_input: String::new(),
|
||||
model_download: ModelDownload::Idle,
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
@@ -153,6 +144,7 @@ impl App {
|
||||
choosing_folder: false,
|
||||
pending_project_path: None,
|
||||
project_name_input: String::new(),
|
||||
model_download: ModelDownload::Idle,
|
||||
error: Some(format!("Could not open the project database: {error}")),
|
||||
}
|
||||
}
|
||||
@@ -172,14 +164,14 @@ impl App {
|
||||
}
|
||||
Message::PreferenceModelChanged(model) => {
|
||||
self.preference_draft.model = model;
|
||||
if !model.supports_dflash() {
|
||||
self.preference_draft.dflash_enabled = false;
|
||||
if !model.supports_dspark() {
|
||||
self.preference_draft.dspark_enabled = false;
|
||||
}
|
||||
self.preference_error = None;
|
||||
}
|
||||
Message::PreferenceDflashChanged(enabled) => {
|
||||
self.preference_draft.dflash_enabled =
|
||||
self.preference_draft.model.supports_dflash() && enabled;
|
||||
Message::PreferenceDsparkChanged(enabled) => {
|
||||
self.preference_draft.dspark_enabled =
|
||||
self.preference_draft.model.supports_dspark() && enabled;
|
||||
self.preference_error = None;
|
||||
}
|
||||
Message::PreferenceTimeoutChanged(value) => {
|
||||
@@ -187,6 +179,48 @@ 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::StopModelDownload => {
|
||||
if let ModelDownload::Active(download) = &mut self.model_download {
|
||||
download.stopping = true;
|
||||
download.cancel.store(true, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
Message::DownloadProgressTick => self.update_download_progress(),
|
||||
Message::ChooseProjectFolder => {
|
||||
self.choosing_folder = true;
|
||||
return Task::perform(
|
||||
@@ -256,7 +290,15 @@ impl App {
|
||||
}
|
||||
|
||||
pub(crate) fn subscription(&self) -> Subscription<Message> {
|
||||
keyboard::on_key_press(shortcut)
|
||||
let shortcuts = keyboard::on_key_press(shortcut);
|
||||
if matches!(self.model_download, ModelDownload::Active(_)) {
|
||||
Subscription::batch([
|
||||
shortcuts,
|
||||
Subscription::run(download_ticks).map(|_| Message::DownloadProgressTick),
|
||||
])
|
||||
} else {
|
||||
shortcuts
|
||||
}
|
||||
}
|
||||
|
||||
fn open_preferences(&mut self) {
|
||||
@@ -289,11 +331,11 @@ impl App {
|
||||
}
|
||||
|
||||
let model = self.preference_draft.model;
|
||||
let dflash_enabled = model.supports_dflash() && self.preference_draft.dflash_enabled;
|
||||
let dspark_enabled = model.supports_dspark() && self.preference_draft.dspark_enabled;
|
||||
let Some(database) = &mut self.database else {
|
||||
return;
|
||||
};
|
||||
match database.update_preferences(model.id(), dflash_enabled, idle_timeout_minutes) {
|
||||
match database.update_preferences(model.id(), dspark_enabled, idle_timeout_minutes) {
|
||||
Ok(preferences) => {
|
||||
self.preferences = preferences;
|
||||
self.preference_draft = PreferenceDraft::from_saved(&self.preferences)
|
||||
@@ -395,6 +437,73 @@ impl App {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn update_download_progress(&mut self) {
|
||||
let ModelDownload::Active(download) = &mut self.model_download else {
|
||||
return;
|
||||
};
|
||||
let now = Instant::now();
|
||||
let progress =
|
||||
model::download_progress(download.model, download.dspark_enabled, &models_path());
|
||||
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 current = transferred as f64 / elapsed;
|
||||
download.bytes_per_second = if download.bytes_per_second == 0.0 {
|
||||
current
|
||||
} else {
|
||||
download.bytes_per_second * 0.75 + current * 0.25
|
||||
};
|
||||
}
|
||||
download.progress = progress;
|
||||
download.sampled_at = now;
|
||||
download.sampled_bytes = download.progress.downloaded;
|
||||
|
||||
let result = match download.result.try_recv() {
|
||||
Ok(result) => Some(result),
|
||||
Err(TryRecvError::Empty) => None,
|
||||
Err(TryRecvError::Disconnected) => {
|
||||
Some(Err("Download worker stopped unexpectedly.".into()))
|
||||
}
|
||||
};
|
||||
let model = download.model;
|
||||
let progress = download.progress.clone();
|
||||
if let Some(result) = result {
|
||||
match result {
|
||||
Ok(DownloadOutcome::Complete) => {
|
||||
self.model_download = ModelDownload::Complete(model, progress);
|
||||
self.error = None;
|
||||
}
|
||||
Ok(DownloadOutcome::Stopped) => {
|
||||
self.model_download = ModelDownload::Idle;
|
||||
self.error = None;
|
||||
}
|
||||
Err(error) => {
|
||||
self.error = Some(error.clone());
|
||||
self.model_download = ModelDownload::Failed(model, error, progress);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for App {
|
||||
fn drop(&mut self) {
|
||||
if let ModelDownload::Active(download) = &self.model_download {
|
||||
download.cancel.store(true, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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> {
|
||||
@@ -414,19 +523,23 @@ fn application_support_path() -> PathBuf {
|
||||
.join(APP_ID)
|
||||
}
|
||||
|
||||
fn models_path() -> PathBuf {
|
||||
application_support_path().join("models")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn preferences_shortcut_and_dflash_support_are_explicit() {
|
||||
fn preferences_shortcut_and_dspark_support_are_explicit() {
|
||||
let message = shortcut(
|
||||
keyboard::Key::Character(",".into()),
|
||||
keyboard::Modifiers::COMMAND,
|
||||
);
|
||||
assert!(matches!(message, Some(Message::OpenPreferences)));
|
||||
assert!(ModelChoice::DeepSeekV4Flash.supports_dflash());
|
||||
assert!(!ModelChoice::DeepSeekV4Pro.supports_dflash());
|
||||
assert!(!ModelChoice::Glm52.supports_dflash());
|
||||
assert!(ModelChoice::DeepSeekV4Flash.supports_dspark());
|
||||
assert!(!ModelChoice::DeepSeekV4Pro.supports_dspark());
|
||||
assert!(!ModelChoice::Glm52.supports_dspark());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user