| | |
| | | import ( |
| | | "archive/tar" |
| | | "bytes" |
| | | "context" |
| | | "errors" |
| | | "fmt" |
| | | "net/http" |
| | | "net/http/httptest" |
| | |
| | | "path/filepath" |
| | | "strconv" |
| | | "strings" |
| | | "sync" |
| | | "sync/atomic" |
| | | "testing" |
| | | "time" |
| | | ) |
| | | |
| | | func TestRegistryReturnsDefaultSenseVoice(t *testing.T) { |
| | |
| | | } |
| | | } |
| | | |
| | | 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") |
| | | |