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/model_test.go | 96 ++++++++++++++++++++++++++++++++++++++++++++++++
1 files changed, 96 insertions(+), 0 deletions(-)
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")
--
Gitblit v1.9.3