From 7999e66c9c78ada3666eb7cea7c22a352bb4cbf0 Mon Sep 17 00:00:00 2001
From: Ariver <shanghai3168@gmail.com>
Date: Wed, 24 Jun 2026 02:30:51 +0800
Subject: [PATCH] Add cancellable model downloads
---
privatevoice.src/internal/model/downloader.go | 68 +++++++-
privatevoice.src/frontend/src/components/onboarding/OnboardingView.svelte | 75 +++++++++
privatevoice.src/services/engine_service.go | 83 +++++++++-
privatevoice.src/frontend/src/components/settings/LanguagePage.svelte | 87 +++++++++-
privatevoice.src/frontend/src/lib/i18n/zh.json | 9 +
privatevoice.src/internal/model/model_test.go | 96 ++++++++++++
privatevoice.src/frontend/src/lib/i18n/en.json | 9 +
7 files changed, 391 insertions(+), 36 deletions(-)
diff --git a/privatevoice.src/frontend/src/components/onboarding/OnboardingView.svelte b/privatevoice.src/frontend/src/components/onboarding/OnboardingView.svelte
index 6535a9a..2b3a727 100755
--- a/privatevoice.src/frontend/src/components/onboarding/OnboardingView.svelte
+++ b/privatevoice.src/frontend/src/components/onboarding/OnboardingView.svelte
@@ -10,12 +10,18 @@
let progress = $state(0)
let failed = $state(false)
let downloadActive = $state(false)
+ let cancelling = $state(false)
+ let currentModelID = $state('')
onMount(() => {
Events.On('model:download-progress', (ev: any) => {
const data = ev?.data
if (data?.percent != null) {
+ if (data?.modelID) {
+ currentModelID = data.modelID
+ }
downloadActive = true
+ cancelling = false
failed = false
statusText = t('onboarding.syncModel')
progress = data.percent
@@ -29,6 +35,7 @@
const data = ev?.data
if (data?.status === 'ready') {
downloadActive = false
+ cancelling = false
failed = false
statusText = t('onboarding.complete')
progress = 100
@@ -41,6 +48,7 @@
}
} else if (data?.status === 'error') {
downloadActive = false
+ cancelling = false
statusText = t('onboarding.failed')
detailText = data?.error || ''
failed = true
@@ -52,7 +60,10 @@
async function triggerDownload() {
try {
+ const current: any = await Call.ByName('voicesnap/services.EngineService.GetCurrentModelStatus')
+ currentModelID = typeof current?.modelID === 'string' ? current.modelID : ''
downloadActive = true
+ cancelling = false
failed = false
statusText = t('onboarding.syncModel')
detailText = ''
@@ -62,10 +73,28 @@
detailText = t('onboarding.extracting')
} catch (err: any) {
downloadActive = false
- statusText = t('onboarding.failed')
- detailText = err?.message || String(err)
+ cancelling = false
+ if (isCancelError(err)) {
+ statusText = t('onboarding.cancelled')
+ detailText = t('onboarding.cancelledDetail')
+ } else {
+ statusText = t('onboarding.failed')
+ detailText = err?.message || String(err)
+ }
failed = true
}
+ }
+
+ function isCancelError(err: any): boolean {
+ const text = String(err?.message || err || '').toLowerCase()
+ return text.includes('context canceled') || text.includes('cancelled') || text.includes('canceled')
+ }
+
+ async function cancelDownload() {
+ if (!downloadActive || cancelling) return
+ cancelling = true
+ statusText = t('onboarding.cancelling')
+ await Call.ByName('voicesnap/services.EngineService.CancelModelDownload', currentModelID)
}
</script>
@@ -86,6 +115,18 @@
<!-- Progress bar -->
<div class="progress-track">
<div class="progress-fill" style="width: {progress}%" class:failed></div>
+ </div>
+
+ <div class="actions">
+ {#if downloadActive}
+ <button class="action cancel" onclick={cancelDownload} disabled={cancelling}>
+ {cancelling ? t('onboarding.cancelling') : t('onboarding.cancel')}
+ </button>
+ {:else if failed}
+ <button class="action primary" onclick={triggerDownload}>
+ {t('onboarding.retry')}
+ </button>
+ {/if}
</div>
</div>
</div>
@@ -162,4 +203,34 @@
.progress-fill.failed {
background: var(--color-red);
}
+
+ .actions {
+ margin-top: var(--spacing-lg);
+ min-height: 32px;
+ }
+
+ .action {
+ min-width: 96px;
+ height: 32px;
+ padding: 0 14px;
+ border: none;
+ border-radius: var(--radius-sm);
+ color: white;
+ font-size: var(--font-size-sm);
+ font-weight: 500;
+ cursor: pointer;
+ }
+
+ .action.primary {
+ background: var(--color-blue);
+ }
+
+ .action.cancel {
+ background: var(--color-red);
+ }
+
+ .action:disabled {
+ cursor: not-allowed;
+ opacity: 0.72;
+ }
</style>
diff --git a/privatevoice.src/frontend/src/components/settings/LanguagePage.svelte b/privatevoice.src/frontend/src/components/settings/LanguagePage.svelte
index 2bb1b72..24c435c 100644
--- a/privatevoice.src/frontend/src/components/settings/LanguagePage.svelte
+++ b/privatevoice.src/frontend/src/components/settings/LanguagePage.svelte
@@ -30,6 +30,8 @@
let languageOptions = $state<LanguageOption[]>([])
let modelOptions = $state<ModelOption[]>([])
let modelBusyID = $state('')
+ let modelBusyKind = $state<'download' | 'select' | ''>('')
+ let modelCancellingID = $state('')
let modelProgress = $state(0)
let modelError = $state('')
@@ -119,8 +121,8 @@
function modelActionLabel(option: ModelOption): string {
if (modelBusyID === option.modelID) {
- if (!option.installed && modelProgress > 0) {
- return `${Math.min(100, Math.max(0, modelProgress)).toFixed(0)}%`
+ if (modelBusyKind === 'download') {
+ return modelCancellingID === option.modelID ? t('settings.modelCancelling') : t('settings.modelCancel')
}
return t('settings.modelWorking')
}
@@ -129,9 +131,62 @@
return t('settings.modelDownload')
}
+ function modelMeta(option: ModelOption): string {
+ let meta = option.installed ? t('settings.modelInstalled') : t('settings.modelNotInstalled')
+ if (option.downloadSize) {
+ meta += ` · ${option.downloadSize}`
+ }
+ if (modelBusyID === option.modelID && modelBusyKind === 'download' && modelProgress > 0) {
+ const pct = Math.min(100, Math.max(0, modelProgress)).toFixed(0)
+ meta += ` · ${pct}%`
+ }
+ return meta
+ }
+
+ function isCancelError(err: any): boolean {
+ const text = String(err?.message || err || '').toLowerCase()
+ return text.includes('context canceled') || text.includes('cancelled') || text.includes('canceled')
+ }
+
+ function canCancelModel(option: ModelOption): boolean {
+ return modelBusyID === option.modelID && modelBusyKind === 'download'
+ }
+
+ async function onModelButton(option: ModelOption) {
+ if (canCancelModel(option)) {
+ await cancelModelDownload(option)
+ return
+ }
+ await onModelAction(option)
+ }
+
+ async function cancelModelDownload(option: ModelOption) {
+ if (!canCancelModel(option) || modelCancellingID) return
+ modelCancellingID = option.modelID
+ modelError = ''
+ try {
+ const cancelled: any = await Call.ByName('voicesnap/services.EngineService.CancelModelDownload', option.modelID)
+ if (!cancelled) {
+ modelBusyID = ''
+ modelBusyKind = ''
+ modelCancellingID = ''
+ modelProgress = 0
+ await loadModelOptions()
+ }
+ } catch (err: any) {
+ modelError = err?.message || String(err || t('settings.modelActionFailed'))
+ modelBusyID = ''
+ modelBusyKind = ''
+ modelCancellingID = ''
+ modelProgress = 0
+ await loadModelOptions()
+ }
+ }
+
async function onModelAction(option: ModelOption) {
if ((option.isCurrent && option.installed) || modelBusyID) return
modelBusyID = option.modelID
+ modelBusyKind = option.installed ? 'select' : 'download'
modelProgress = 0
modelError = ''
try {
@@ -144,11 +199,17 @@
}
await loadModelOptions()
} catch (err: any) {
- modelError = err?.message || String(err || t('settings.modelActionFailed'))
+ if (!isCancelError(err)) {
+ modelError = err?.message || String(err || t('settings.modelActionFailed'))
+ }
await loadModelOptions()
} finally {
- modelBusyID = ''
- modelProgress = 0
+ if (modelBusyID === option.modelID) {
+ modelBusyID = ''
+ modelBusyKind = ''
+ modelCancellingID = ''
+ modelProgress = 0
+ }
}
}
@@ -236,20 +297,16 @@
</span>
</div>
<span class="model-desc">{modelDescription(option)}</span>
- <span class="model-meta">
- {option.installed ? t('settings.modelInstalled') : t('settings.modelNotInstalled')}
- {#if option.downloadSize}
- · {option.downloadSize}
- {/if}
- </span>
+ <span class="model-meta">{modelMeta(option)}</span>
</div>
{#if option.isCurrent && option.installed}
<span class="model-current-status">{t('settings.modelCurrent')}</span>
{:else}
<button
class="model-action primary"
- disabled={!!modelBusyID}
- onclick={() => onModelAction(option)}
+ class:cancel={canCancelModel(option)}
+ disabled={!!modelBusyID && modelBusyID !== option.modelID}
+ onclick={() => onModelButton(option)}
>
{modelActionLabel(option)}
</button>
@@ -455,6 +512,10 @@
color: white;
}
+ .model-action.primary.cancel {
+ background: var(--color-red);
+ }
+
.model-action:disabled {
cursor: not-allowed;
opacity: 0.72;
diff --git a/privatevoice.src/frontend/src/lib/i18n/en.json b/privatevoice.src/frontend/src/lib/i18n/en.json
index 89ff807..6444716 100755
--- a/privatevoice.src/frontend/src/lib/i18n/en.json
+++ b/privatevoice.src/frontend/src/lib/i18n/en.json
@@ -66,6 +66,8 @@
"modelCurrent": "Current",
"modelUse": "Use",
"modelDownload": "Download",
+ "modelCancel": "Cancel",
+ "modelCancelling": "Cancelling",
"modelWorking": "Working",
"modelInstalled": "Installed",
"modelNotInstalled": "Not installed",
@@ -276,7 +278,12 @@
"optimizing": "Optimizing...",
"extracting": "Extracting model files...",
"complete": "Complete",
- "failed": "Setup failed"
+ "failed": "Setup failed",
+ "cancel": "Cancel",
+ "cancelling": "Cancelling",
+ "cancelled": "Cancelled",
+ "cancelledDetail": "Model download was cancelled. You can download it again later.",
+ "retry": "Download again"
},
"update": {
"title": "Update Available",
diff --git a/privatevoice.src/frontend/src/lib/i18n/zh.json b/privatevoice.src/frontend/src/lib/i18n/zh.json
index 423125c..c8890f0 100755
--- a/privatevoice.src/frontend/src/lib/i18n/zh.json
+++ b/privatevoice.src/frontend/src/lib/i18n/zh.json
@@ -66,6 +66,8 @@
"modelCurrent": "当前",
"modelUse": "使用",
"modelDownload": "下载",
+ "modelCancel": "取消",
+ "modelCancelling": "正在取消",
"modelWorking": "处理中",
"modelInstalled": "已安装",
"modelNotInstalled": "未安装",
@@ -276,7 +278,12 @@
"optimizing": "正在优化...",
"extracting": "正在解压模型文件...",
"complete": "完成",
- "failed": "设置失败"
+ "failed": "设置失败",
+ "cancel": "取消",
+ "cancelling": "正在取消",
+ "cancelled": "已取消",
+ "cancelledDetail": "模型下载已取消,可以稍后重新下载。",
+ "retry": "重新下载"
},
"update": {
"title": "发现新版本",
diff --git a/privatevoice.src/internal/model/downloader.go b/privatevoice.src/internal/model/downloader.go
index f77c66d..120fe61 100755
--- a/privatevoice.src/internal/model/downloader.go
+++ b/privatevoice.src/internal/model/downloader.go
@@ -1,6 +1,7 @@
package model
import (
+ "context"
"errors"
"fmt"
"io"
@@ -27,6 +28,10 @@
}
func DownloadProfile(profile ModelProfile, urls []string, modelsDir string, progress ProgressCallback) error {
+ return DownloadProfileWithContext(context.Background(), profile, urls, modelsDir, progress)
+}
+
+func DownloadProfileWithContext(ctx context.Context, profile ModelProfile, urls []string, modelsDir string, progress ProgressCallback) error {
if err := os.MkdirAll(modelsDir, 0755); err != nil {
return fmt.Errorf("failed to create models dir: %w", err)
}
@@ -34,10 +39,16 @@
var lastErr error
for i, url := range urls {
+ if err := ctx.Err(); err != nil {
+ return err
+ }
if url == "" {
continue
}
- if err := downloadAndInstallFromURL(profile, url, modelsDir, i+1, progress); err != nil {
+ if err := downloadAndInstallFromURL(ctx, profile, url, modelsDir, i+1, progress); err != nil {
+ if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
+ return err
+ }
lastErr = err
logger.Info("Model install from URL %d failed: %v", i+1, err)
continue
@@ -63,7 +74,7 @@
}
}
-func downloadAndInstallFromURL(profile ModelProfile, url, modelsDir string, attempt int, progress ProgressCallback) error {
+func downloadAndInstallFromURL(ctx context.Context, profile ModelProfile, url, modelsDir string, attempt int, progress ProgressCallback) error {
runID := fmt.Sprintf("%d-%d", time.Now().UnixNano(), attempt)
downloadDir := filepath.Join(modelsDir, ".downloads", profile.ID, runID)
extractDir := filepath.Join(downloadDir, "extract")
@@ -77,27 +88,42 @@
defer os.RemoveAll(stagingDir)
logger.Info("Downloading model from URL %d: %s", attempt, url)
- if err := downloadFile(url, archivePath, progress); err != nil {
+ if err := downloadFileWithContext(ctx, url, archivePath, progress); err != nil {
+ return err
+ }
+ if err := ctx.Err(); err != nil {
return err
}
logger.Info("Extracting model archive...")
- if err := extractArchive(archivePath, extractDir); err != nil {
+ if err := extractArchiveWithContext(ctx, archivePath, extractDir); err != nil {
return fmt.Errorf("extraction failed: %w", err)
+ }
+ if err := ctx.Err(); err != nil {
+ return err
}
sourceDir, err := findInstallSource(profile, extractDir)
if err != nil {
return err
}
+ if err := ctx.Err(); err != nil {
+ return err
+ }
if err := moveOrCopyDir(sourceDir, stagingDir); err != nil {
return fmt.Errorf("failed to stage model: %w", err)
+ }
+ if err := ctx.Err(); err != nil {
+ return err
}
validation := ValidateModelDir(profile, stagingDir)
if !validation.Valid {
return fmt.Errorf("downloaded model is incomplete: missing=%v problems=%v", validation.Missing, validation.Problems)
+ }
+ if err := ctx.Err(); err != nil {
+ return err
}
finalDir := filepath.Join(modelsDir, profile.InstallDirName)
@@ -105,15 +131,25 @@
}
func downloadFile(url, destPath string, progress ProgressCallback) error {
+ return downloadFileWithContext(context.Background(), url, destPath, progress)
+}
+
+func downloadFileWithContext(ctx context.Context, url, destPath string, progress ProgressCallback) error {
const maxAttempts = 6
var lastErr error
var lastPercent = -1.0
for attempt := 1; attempt <= maxAttempts; attempt++ {
+ if err := ctx.Err(); err != nil {
+ return err
+ }
downloaded := existingFileSize(destPath)
- err := downloadFileAttempt(url, destPath, downloaded, &lastPercent, progress)
+ err := downloadFileAttempt(ctx, url, destPath, downloaded, &lastPercent, progress)
if err == nil {
return nil
+ }
+ if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
+ return err
}
if isNonRetryableDownloadError(err) {
return err
@@ -121,15 +157,19 @@
lastErr = err
logger.Info("Download attempt %d/%d failed: %v", attempt, maxAttempts, err)
if attempt < maxAttempts {
- time.Sleep(time.Duration(attempt) * time.Second)
+ select {
+ case <-ctx.Done():
+ return ctx.Err()
+ case <-time.After(time.Duration(attempt) * time.Second):
+ }
}
}
return fmt.Errorf("download failed after %d attempts: %w", maxAttempts, lastErr)
}
-func downloadFileAttempt(url, destPath string, resumeFrom int64, lastPercent *float64, progress ProgressCallback) error {
- req, err := http.NewRequest(http.MethodGet, url, nil)
+func downloadFileAttempt(ctx context.Context, url, destPath string, resumeFrom int64, lastPercent *float64, progress ProgressCallback) error {
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nonRetryableDownloadError{err: err}
}
@@ -184,6 +224,9 @@
}
buf := make([]byte, 32*1024)
for {
+ if err := ctx.Err(); err != nil {
+ return err
+ }
n, readErr := resp.Body.Read(buf)
if n > 0 {
if _, writeErr := out.Write(buf[:n]); writeErr != nil {
@@ -252,9 +295,16 @@
}
func extractArchive(archivePath, destDir string) error {
- cmd := exec.Command("tar", "-xf", archivePath, "-C", destDir)
+ return extractArchiveWithContext(context.Background(), archivePath, destDir)
+}
+
+func extractArchiveWithContext(ctx context.Context, archivePath, destDir string) error {
+ cmd := exec.CommandContext(ctx, "tar", "-xf", archivePath, "-C", destDir)
output, err := cmd.CombinedOutput()
if err != nil {
+ if ctxErr := ctx.Err(); ctxErr != nil {
+ return ctxErr
+ }
return fmt.Errorf("tar extraction failed: %v, output: %s", err, string(output))
}
return nil
diff --git a/privatevoice.src/internal/model/model_test.go b/privatevoice.src/internal/model/model_test.go
index 795a0d0..d3f025e 100644
--- a/privatevoice.src/internal/model/model_test.go
+++ b/privatevoice.src/internal/model/model_test.go
@@ -3,6 +3,8 @@
import (
"archive/tar"
"bytes"
+ "context"
+ "errors"
"fmt"
"net/http"
"net/http/httptest"
@@ -10,7 +12,10 @@
"path/filepath"
"strconv"
"strings"
+ "sync"
+ "sync/atomic"
"testing"
+ "time"
)
func TestRegistryReturnsDefaultSenseVoice(t *testing.T) {
@@ -761,6 +766,97 @@
}
}
+func TestDownloadProfileCancelStopsWithoutFallbackOrState(t *testing.T) {
+ root := t.TempDir()
+ ctx, cancel := context.WithCancel(context.Background())
+ defer cancel()
+
+ var once sync.Once
+ var fallbackHits int32
+ started := make(chan struct{})
+ payload := bytes.Repeat([]byte("a"), 512*1024)
+ fallbackArchive := tarArchive(t, map[string]string{
+ "sensevoice/tokens.txt": "tokens",
+ "sensevoice/model.int8.onnx": "fallback-model",
+ })
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.URL.Path {
+ case "/slow":
+ once.Do(func() { close(started) })
+ w.Header().Set("Content-Length", strconv.Itoa(len(payload)))
+ w.WriteHeader(http.StatusOK)
+ flusher, _ := w.(http.Flusher)
+ for offset := 0; offset < len(payload); offset += 32 * 1024 {
+ select {
+ case <-r.Context().Done():
+ return
+ default:
+ }
+ end := offset + 32*1024
+ if end > len(payload) {
+ end = len(payload)
+ }
+ if _, err := w.Write(payload[offset:end]); err != nil {
+ return
+ }
+ if flusher != nil {
+ flusher.Flush()
+ }
+ time.Sleep(5 * time.Millisecond)
+ }
+ case "/fallback":
+ atomic.AddInt32(&fallbackHits, 1)
+ w.WriteHeader(http.StatusOK)
+ w.Write(fallbackArchive)
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer server.Close()
+
+ errCh := make(chan error, 1)
+ go func() {
+ errCh <- DownloadProfileWithContext(ctx, DefaultModelProfile(), []string{server.URL + "/slow", server.URL + "/fallback"}, root, func(percent float64, _, _ int64) {
+ if percent > 0 {
+ cancel()
+ }
+ })
+ }()
+
+ select {
+ case <-started:
+ case <-time.After(2 * time.Second):
+ t.Fatal("slow download did not start")
+ }
+
+ var err error
+ select {
+ case err = <-errCh:
+ case <-time.After(2 * time.Second):
+ t.Fatal("cancelled download did not return")
+ }
+ if !errors.Is(err, context.Canceled) {
+ t.Fatalf("download error = %v, want context.Canceled", err)
+ }
+ if got := atomic.LoadInt32(&fallbackHits); got != 0 {
+ t.Fatalf("fallback hits = %d, want 0", got)
+ }
+ state, err := LoadInstallStateFromRoot(root)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, ok := state.InstalledModels[DefaultModelID]; ok {
+ t.Fatal("cancelled download should not write installed model state")
+ }
+ resolved, err := ResolveModelInRoot(DefaultModelID, root)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if resolved.IsUsable() {
+ t.Fatal("cancelled download should not leave a usable model")
+ }
+}
+
func TestDownloadFileResumesExistingPartialFile(t *testing.T) {
payload := []byte("0123456789abcdefghijklmnopqrstuvwxyz")
diff --git a/privatevoice.src/services/engine_service.go b/privatevoice.src/services/engine_service.go
index b6c00b6..4dbfeb3 100755
--- a/privatevoice.src/services/engine_service.go
+++ b/privatevoice.src/services/engine_service.go
@@ -1,6 +1,7 @@
package services
import (
+ "context"
"fmt"
"sync"
"voicesnap/internal/config"
@@ -15,13 +16,16 @@
// EngineService provides engine status and model management to the frontend.
type EngineService struct {
- cfg *config.Config
- app *application.App
- initCallback func()
- mu sync.RWMutex
- status string
- hardwareInfo string
- lastError string
+ cfg *config.Config
+ app *application.App
+ initCallback func()
+ mu sync.RWMutex
+ status string
+ hardwareInfo string
+ lastError string
+ downloadMu sync.Mutex
+ downloadID string
+ downloadCancel context.CancelFunc
}
func NewEngineService(cfg *config.Config) *EngineService {
@@ -106,7 +110,13 @@
return fmt.Errorf("no download URL configured for model %s", profile.ID)
}
- err = model.DownloadProfile(profile, profile.DownloadURLs, paths.ModelsRoot(), func(percent float64, downloaded, total int64) {
+ ctx, finish, err := s.beginModelDownload(profile.ID)
+ if err != nil {
+ return err
+ }
+ defer finish()
+
+ err = model.DownloadProfileWithContext(ctx, profile, profile.DownloadURLs, paths.ModelsRoot(), func(percent float64, downloaded, total int64) {
if s.app != nil {
s.app.Event.Emit("model:download-progress", map[string]interface{}{
"percent": percent,
@@ -126,6 +136,24 @@
config.Save(s.cfg)
s.ReloadCurrentModel()
return nil
+}
+
+func (s *EngineService) CancelModelDownload(modelID string) bool {
+ s.downloadMu.Lock()
+ activeID := s.downloadID
+ if s.downloadCancel == nil || (modelID != "" && activeID != modelID) {
+ s.downloadMu.Unlock()
+ return false
+ }
+ s.downloadCancel()
+ s.downloadMu.Unlock()
+
+ if s.app != nil {
+ s.app.Event.Emit("model:download-cancelled", map[string]interface{}{
+ "modelID": activeID,
+ })
+ }
+ return true
}
func (s *EngineService) modelStatusMap(profile model.ModelProfile, current modelselection.CurrentModel) map[string]interface{} {
@@ -179,12 +207,21 @@
// DownloadModel downloads the ASR model with progress events.
func (s *EngineService) DownloadModel(primaryURL, fallbackURL string) error {
- err := model.Download(primaryURL, fallbackURL, paths.ModelsRoot(), func(percent float64, downloaded, total int64) {
+ profile := model.DefaultModelProfile()
+ ctx, finish, err := s.beginModelDownload(profile.ID)
+ if err != nil {
+ return err
+ }
+ defer finish()
+
+ err = model.DownloadProfileWithContext(ctx, profile, []string{primaryURL, fallbackURL}, paths.ModelsRoot(), func(percent float64, downloaded, total int64) {
if s.app != nil {
s.app.Event.Emit("model:download-progress", map[string]interface{}{
"percent": percent,
"downloaded": downloaded,
"total": total,
+ "modelID": profile.ID,
+ "modelName": profile.DisplayName,
})
}
})
@@ -204,7 +241,13 @@
return fmt.Errorf("no download URL configured for model %s", current.ModelID)
}
- err := model.DownloadProfile(current.Profile, current.Profile.DownloadURLs, paths.ModelsRoot(), func(percent float64, downloaded, total int64) {
+ ctx, finish, err := s.beginModelDownload(current.ModelID)
+ if err != nil {
+ return err
+ }
+ defer finish()
+
+ err = model.DownloadProfileWithContext(ctx, current.Profile, current.Profile.DownloadURLs, paths.ModelsRoot(), func(percent float64, downloaded, total int64) {
if s.app != nil {
s.app.Event.Emit("model:download-progress", map[string]interface{}{
"percent": percent,
@@ -252,6 +295,26 @@
return modelselection.Resolve(s.cfg, language.NewSystemDetector())
}
+func (s *EngineService) beginModelDownload(modelID string) (context.Context, func(), error) {
+ s.downloadMu.Lock()
+ defer s.downloadMu.Unlock()
+ if s.downloadCancel != nil {
+ return nil, nil, fmt.Errorf("model %s is already downloading", s.downloadID)
+ }
+ ctx, cancel := context.WithCancel(context.Background())
+ s.downloadID = modelID
+ s.downloadCancel = cancel
+ finish := func() {
+ s.downloadMu.Lock()
+ if s.downloadID == modelID {
+ s.downloadID = ""
+ s.downloadCancel = nil
+ }
+ s.downloadMu.Unlock()
+ }
+ return ctx, finish, nil
+}
+
func (s *EngineService) allowedModelProfile(modelID string) (model.ModelProfile, error) {
profile, err := model.GetModelProfile(modelID)
if err != nil {
--
Gitblit v1.9.3