diff --git a/cmd/vox/main.go b/cmd/vox/main.go index 5c50c82..b8ff12d 100644 --- a/cmd/vox/main.go +++ b/cmd/vox/main.go @@ -30,11 +30,11 @@ import ( "vox/internal/inject" "vox/internal/pipeline" "vox/internal/prompt" + "vox/internal/sttmodel" "vox/internal/transcribe" "vox/internal/ui" "vox/internal/userconfig" "vox/internal/vocab" - "vox/internal/whispermodel" "vox/internal/whisperserver" ) @@ -160,13 +160,13 @@ func run() { } // Resolve selected whisper model + start the embedded whisper-server child. - selectedModel, ok := whispermodel.ByID(cfg.ModelID) + selectedModel, ok := sttmodel.ByID(cfg.ModelID) if !ok { - selectedModel, _ = whispermodel.ByID(whispermodel.DefaultID) + selectedModel, _ = sttmodel.ByID(sttmodel.DefaultID) } - if !whispermodel.IsInstalled(selectedModel) { + if !sttmodel.IsInstalled(selectedModel) { fmt.Printf("Selected model %q is not installed, downloading...\n", selectedModel.ID) - if err := whispermodel.Download(ctx, selectedModel, nil); err != nil { + if err := sttmodel.Download(ctx, selectedModel, nil); err != nil { fmt.Fprintf(os.Stderr, "Error downloading model %q: %v\n", selectedModel.ID, err) os.Exit(1) } @@ -176,7 +176,7 @@ func run() { fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) } - modelPath, err := whispermodel.Path(selectedModel) + modelPath, err := sttmodel.ResolvePath(selectedModel) if err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) @@ -365,13 +365,13 @@ func shutdownWatcher(cancel context.CancelFunc, logger *slog.Logger, recorder *a } func buildModelPresets() []ui.ModelPreset { - models := whispermodel.All() + models := sttmodel.All() presets := make([]ui.ModelPreset, 0, len(models)) for _, m := range models { presets = append(presets, ui.ModelPreset{ ID: m.ID, Label: m.Label, - Installed: whispermodel.IsInstalled(m), + Installed: sttmodel.IsInstalled(m), }) } return presets @@ -379,9 +379,9 @@ func buildModelPresets() []ui.ModelPreset { func buildModelRemovePresets() []ui.ModelRemovePreset { var out []ui.ModelRemovePreset - nInstalled := whispermodel.InstalledCount() - for _, m := range whispermodel.All() { - if whispermodel.IsInstalled(m) { + nInstalled := sttmodel.InstalledCount() + for _, m := range sttmodel.All() { + if sttmodel.IsInstalled(m) { label := m.Label if nInstalled <= 1 { label = m.Label + " (required)" @@ -402,9 +402,9 @@ func refreshModelMenus() { ui.SetModelRemovePresets(buildModelRemovePresets()) } -func activateWhisperModel(ctx context.Context, whisperClient *transcribe.Client, whisperSrv *whisperserver.Server, model whispermodel.Model) error { - if !whispermodel.IsInstalled(model) { - err := whispermodel.Download(ctx, model, func(downloaded, total int64) { +func activateWhisperModel(ctx context.Context, whisperClient *transcribe.Client, whisperSrv *whisperserver.Server, model sttmodel.Model) error { + if !sttmodel.IsInstalled(model) { + err := sttmodel.Download(ctx, model, func(downloaded, total int64) { if total > 0 { pct := (downloaded * 100) / total ui.SetStatusLine(fmt.Sprintf("Status: Downloading %s (%d%%)…", model.ID, pct)) @@ -416,7 +416,7 @@ func activateWhisperModel(ctx context.Context, whisperClient *transcribe.Client, return err } } - path, err := whispermodel.Path(model) + path, err := sttmodel.ResolvePath(model) if err != nil { return err } @@ -430,19 +430,19 @@ func activateWhisperModel(ctx context.Context, whisperClient *transcribe.Client, return switchErr } -func pickFallbackModel(excludeID string) (whispermodel.Model, error) { +func pickFallbackModel(excludeID string) (sttmodel.Model, error) { // Prefer any installed model that isn't the one being excluded. - for _, m := range whispermodel.All() { + for _, m := range sttmodel.All() { if m.ID == excludeID { continue } - if whispermodel.IsInstalled(m) { + if sttmodel.IsInstalled(m) { return m, nil } } // No installed fallback found — return an error rather than silently // triggering a download for a non-installed model. - return whispermodel.Model{}, fmt.Errorf("no installed fallback model (excluding %s)", excludeID) + return sttmodel.Model{}, fmt.Errorf("no installed fallback model (excluding %s)", excludeID) } func whisperLogPath() string { @@ -513,7 +513,7 @@ func modelChangeWatcher(ctx context.Context, logger *slog.Logger, client *transc return case modelID := <-ui.OnModelChange(): modelMu.Lock() - model, ok := whispermodel.ByID(modelID) + model, ok := sttmodel.ByID(modelID) if !ok { logger.Warn("unknown model selected", "id", modelID) modelMu.Unlock() @@ -552,12 +552,12 @@ func modelDeleteWatcher(ctx context.Context, logger *slog.Logger, whisperClient return case modelID := <-ui.OnModelDelete(): modelMu.Lock() - model, ok := whispermodel.ByID(modelID) - if !ok || !whispermodel.IsInstalled(model) { + model, ok := sttmodel.ByID(modelID) + if !ok || !sttmodel.IsInstalled(model) { modelMu.Unlock() continue } - if whispermodel.InstalledCount() <= 1 { + if sttmodel.InstalledCount() <= 1 { logger.Info("refusing remove: last downloaded model", "id", modelID) ui.SetStatusLine("Status: Idle") fmt.Printf("Can't remove your only downloaded model.\n") @@ -592,7 +592,7 @@ func modelDeleteWatcher(ctx context.Context, logger *slog.Logger, whisperClient } } - if err := whispermodel.Remove(model); err != nil { + if err := sttmodel.Remove(model); err != nil { logger.Warn("remove model file", "id", model.ID, "error", err) ui.SetStatusLine("Status: Remove failed") } else { @@ -889,8 +889,8 @@ func handleStopAndProcess( fmt.Println("Ready!") } -// transcribeStage returns a pipeline stage that sends audio to the Whisper API. -func transcribeStage(client *transcribe.Client, opts transcribe.TranscribeOptions) pipeline.Stage { +// transcribeStage returns a pipeline stage that converts recorded audio to text. +func transcribeStage(client transcribe.Transcriber, opts transcribe.TranscribeOptions) pipeline.Stage { return func(ctx context.Context, r *pipeline.Result) error { text, err := client.Transcribe(ctx, r.RawAudio, opts) if err != nil { diff --git a/cmd/vox/model_test.go b/cmd/vox/model_test.go index 5b3d409..95ccf11 100644 --- a/cmd/vox/model_test.go +++ b/cmd/vox/model_test.go @@ -4,26 +4,32 @@ import ( "os" "testing" - "vox/internal/whispermodel" + "path/filepath" + + "vox/internal/sttmodel" ) func TestPickFallbackModel_MultipleInstalled(t *testing.T) { t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + t.Setenv("VOX_MODEL_DIR", t.TempDir()) // Install two models: tiny.en and base.en. - tiny, ok := whispermodel.ByID("tiny.en") + tiny, ok := sttmodel.ByID("tiny.en") if !ok { t.Fatal("ByID tiny.en") } - base, ok := whispermodel.ByID("base.en") + base, ok := sttmodel.ByID("base.en") if !ok { t.Fatal("ByID base.en") } - for _, m := range []whispermodel.Model{tiny, base} { - p, err := whispermodel.Path(m) + for _, m := range []sttmodel.Model{tiny, base} { + p, err := sttmodel.Path(m) if err != nil { t.Fatalf("Path %s: %v", m.ID, err) } + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatalf("MkdirAll %s: %v", m.ID, err) + } if err := os.WriteFile(p, []byte("x"), 0o600); err != nil { t.Fatalf("WriteFile %s: %v", m.ID, err) } @@ -50,16 +56,20 @@ func TestPickFallbackModel_MultipleInstalled(t *testing.T) { func TestPickFallbackModel_OnlyOneInstalled(t *testing.T) { t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + t.Setenv("VOX_MODEL_DIR", t.TempDir()) // Install only base.en. - base, ok := whispermodel.ByID("base.en") + base, ok := sttmodel.ByID("base.en") if !ok { t.Fatal("ByID base.en") } - p, err := whispermodel.Path(base) + p, err := sttmodel.Path(base) if err != nil { t.Fatalf("Path: %v", err) } + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatalf("MkdirAll: %v", err) + } if err := os.WriteFile(p, []byte("x"), 0o600); err != nil { t.Fatalf("WriteFile: %v", err) } @@ -73,6 +83,7 @@ func TestPickFallbackModel_OnlyOneInstalled(t *testing.T) { func TestPickFallbackModel_NoneInstalled(t *testing.T) { t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + t.Setenv("VOX_MODEL_DIR", t.TempDir()) _, err := pickFallbackModel("base.en") if err == nil { @@ -82,16 +93,20 @@ func TestPickFallbackModel_NoneInstalled(t *testing.T) { func TestBuildModelRemovePresets_LastModelProtected(t *testing.T) { t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + t.Setenv("VOX_MODEL_DIR", t.TempDir()) // Install one model. - base, ok := whispermodel.ByID("base.en") + base, ok := sttmodel.ByID("base.en") if !ok { t.Fatal("ByID base.en") } - p, err := whispermodel.Path(base) + p, err := sttmodel.Path(base) if err != nil { t.Fatalf("Path: %v", err) } + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatalf("MkdirAll: %v", err) + } if err := os.WriteFile(p, []byte("x"), 0o600); err != nil { t.Fatalf("WriteFile: %v", err) } @@ -110,17 +125,21 @@ func TestBuildModelRemovePresets_LastModelProtected(t *testing.T) { func TestBuildModelRemovePresets_MultipleRemovable(t *testing.T) { t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + t.Setenv("VOX_MODEL_DIR", t.TempDir()) // Install two models. for _, id := range []string{"tiny.en", "base.en"} { - m, ok := whispermodel.ByID(id) + m, ok := sttmodel.ByID(id) if !ok { t.Fatalf("ByID %s", id) } - p, err := whispermodel.Path(m) + p, err := sttmodel.Path(m) if err != nil { t.Fatalf("Path %s: %v", id, err) } + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatalf("MkdirAll %s: %v", m.ID, err) + } if err := os.WriteFile(p, []byte("x"), 0o600); err != nil { t.Fatalf("WriteFile %s: %v", id, err) } diff --git a/cmd/vox/stages_test.go b/cmd/vox/stages_test.go index b0dd086..c866ab7 100644 --- a/cmd/vox/stages_test.go +++ b/cmd/vox/stages_test.go @@ -5,8 +5,39 @@ import ( "testing" "vox/internal/pipeline" + "vox/internal/transcribe" ) +type fakeTranscriber struct { + text string + err error + got []byte +} + +func (f *fakeTranscriber) Transcribe(_ context.Context, wav []byte, _ transcribe.TranscribeOptions) (string, error) { + f.got = wav + return f.text, f.err +} + +func TestTranscribeStageUsesInterface(t *testing.T) { + ft := &fakeTranscriber{text: "hello world"} + stage := transcribeStage(ft, transcribe.TranscribeOptions{}) + + r := &pipeline.Result{RawAudio: []byte("fake-wav")} + if err := stage(context.Background(), r); err != nil { + t.Fatalf("stage: %v", err) + } + if r.RawText != "hello world" { + t.Errorf("RawText = %q, want %q", r.RawText, "hello world") + } + if r.OutputText != "hello world" { + t.Errorf("OutputText = %q, want %q", r.OutputText, "hello world") + } + if string(ft.got) != "fake-wav" { + t.Errorf("transcriber got %q, want %q", ft.got, "fake-wav") + } +} + func TestFilterBlankStage_EmptyText(t *testing.T) { stage := filterBlankStage() r := &pipeline.Result{RawText: ""} diff --git a/internal/config/config.go b/internal/config/config.go index 0574542..c4e894b 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -6,7 +6,7 @@ import ( "strings" "vox/internal/hotkey" - "vox/internal/whispermodel" + "vox/internal/sttmodel" ) // Config holds runtime configuration for the vox dictation tool. @@ -63,8 +63,8 @@ func Load() Config { if modelID == "" { modelID = prefs.Model } - if _, ok := whispermodel.ByID(modelID); !ok { - modelID = whispermodel.DefaultID + if _, ok := sttmodel.ByID(modelID); !ok { + modelID = sttmodel.DefaultID } return Config{ diff --git a/internal/sttmodel/archive.go b/internal/sttmodel/archive.go new file mode 100644 index 0000000..ebabe59 --- /dev/null +++ b/internal/sttmodel/archive.go @@ -0,0 +1,135 @@ +package sttmodel + +import ( + "archive/tar" + "compress/bzip2" + "fmt" + "io" + "os" + "path/filepath" + "strings" +) + +// maxArchiveEntryBytes caps a single extracted file to guard against +// decompression bombs. The largest real entry is the ~500 MiB encoder. +const maxArchiveEntryBytes = 2 * 1024 * 1024 * 1024 // 2 GiB + +// extractArchive unpacks a tar archive (optionally bzip2-compressed) into +// destDir. Entries are flattened by one level: the archive's single top-level +// directory is stripped, so "model-name/encoder.onnx" lands at +// "/encoder.onnx". +func extractArchive(archivePath, destDir string) error { + f, err := os.Open(archivePath) + if err != nil { + return fmt.Errorf("open archive: %w", err) + } + defer f.Close() + + var r io.Reader = f + compressed, err := isBzip2(f) + if err != nil { + return err + } + if compressed { + r = bzip2.NewReader(f) + } + + if err := os.MkdirAll(destDir, 0o755); err != nil { + return fmt.Errorf("create dest dir: %w", err) + } + + tr := tar.NewReader(r) + for { + hdr, err := tr.Next() + if err == io.EOF { + return nil + } + if err != nil { + return fmt.Errorf("read tar: %w", err) + } + + rel := stripTopLevel(hdr.Name) + if rel == "" { + continue + } + target, err := safeJoin(destDir, rel) + if err != nil { + return err + } + + switch hdr.Typeflag { + case tar.TypeDir: + if err := os.MkdirAll(target, 0o755); err != nil { + return fmt.Errorf("mkdir %s: %w", rel, err) + } + case tar.TypeReg: + if hdr.Size > maxArchiveEntryBytes { + return fmt.Errorf("archive entry %s too large (%d bytes)", rel, hdr.Size) + } + if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { + return fmt.Errorf("mkdir parent of %s: %w", rel, err) + } + out, err := os.OpenFile(target, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o644) + if err != nil { + return fmt.Errorf("create %s: %w", rel, err) + } + if _, err := io.Copy(out, io.LimitReader(tr, maxArchiveEntryBytes)); err != nil { + out.Close() + return fmt.Errorf("write %s: %w", rel, err) + } + if err := out.Close(); err != nil { + return fmt.Errorf("close %s: %w", rel, err) + } + default: + // Skip symlinks, devices, and anything else. Model archives only + // contain regular files and directories. + } + } +} + +// isBzip2 sniffs the magic bytes and rewinds the file. +func isBzip2(f *os.File) (bool, error) { + var magic [3]byte + n, err := f.Read(magic[:]) + if err != nil && n == 0 { + return false, fmt.Errorf("read magic: %w", err) + } + if _, err := f.Seek(0, io.SeekStart); err != nil { + return false, fmt.Errorf("rewind: %w", err) + } + return n == 3 && magic[0] == 'B' && magic[1] == 'Z' && magic[2] == 'h', nil +} + +// stripTopLevel removes the archive's single root directory component. +func stripTopLevel(name string) string { + name = strings.TrimPrefix(filepath.Clean(name), "./") + i := strings.Index(name, "/") + if i < 0 { + return "" // the root directory entry itself + } + return name[i+1:] +} + +// safeJoin joins rel onto base, refusing any path that escapes base. +func safeJoin(base, rel string) (string, error) { + target := filepath.Join(base, rel) + cleanBase := filepath.Clean(base) + string(os.PathSeparator) + if !strings.HasPrefix(target, cleanBase) { + return "", fmt.Errorf("archive entry escapes destination: %q", rel) + } + return target, nil +} + +// verifyFiles checks that every required file exists and is non-empty. +func verifyFiles(root string, files []string) error { + for _, f := range files { + st, err := os.Stat(filepath.Join(root, f)) + if err != nil { + return fmt.Errorf("missing required file %s: %w", f, err) + } + if st.Size() == 0 { + return fmt.Errorf("required file %s is empty", f) + } + } + return nil +} diff --git a/internal/sttmodel/archive_real_test.go b/internal/sttmodel/archive_real_test.go new file mode 100644 index 0000000..8904463 --- /dev/null +++ b/internal/sttmodel/archive_real_test.go @@ -0,0 +1,110 @@ +//go:build smoke_manual + +package sttmodel + +import ( + "context" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strconv" + "testing" +) + +// TestRealBzip2ArchiveEndToEnd closes the coverage gap called out in Task 5: +// unit tests build uncompressed tars (Go cannot write bzip2), so the +// bzip2.NewReader branch of extractArchive is exercised by nothing. This +// serves the genuine sherpa-onnx release archive over HTTP and drives the +// full path: download -> checksum -> bzip2 detect -> extract -> flatten one +// level -> verify required files -> atomic install. +// +// Requires VOX_TEST_BZ2 to point at a real .tar.bz2 release archive: +// +// curl -L -o /tmp/v3.tar.bz2 \ +// https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-nemo-parakeet-tdt-0.6b-v3-int8.tar.bz2 +// VOX_TEST_BZ2=/tmp/v3.tar.bz2 go test -tags smoke_manual ./internal/sttmodel/ -run RealBzip2 -v +// +// Tagged out of the default suite because it needs a 465 MiB archive and +// takes ~30s. +func TestRealBzip2ArchiveEndToEnd(t *testing.T) { + src := os.Getenv("VOX_TEST_BZ2") + if src == "" { + t.Skip("VOX_TEST_BZ2 unset") + } + data, err := os.ReadFile(src) + if err != nil { + t.Skipf("cannot read %s: %v", src, err) + } + if len(data) < 3 || data[0] != 'B' || data[1] != 'Z' || data[2] != 'h' { + t.Fatalf("%s is not bzip2 (magic %q)", src, data[:3]) + } + + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + + // Set Content-Length explicitly. Without it, net/http chunk-encodes a body + // this large and ContentLength is -1, which is a property of httptest, not + // of the real GitHub release (verified: it sends Content-Length). + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Length", strconv.Itoa(len(data))) + _, _ = w.Write(data) + })) + defer srv.Close() + + m, ok := ByID("parakeet-v3") + if !ok { + t.Fatal("parakeet-v3 not in catalog") + } + m.URL = srv.URL // real bytes, real checksum from the catalog + + var lastPct int64 + var lastDone, sawTotal int64 + if err := Download(context.Background(), m, func(done, total int64) { + lastDone, sawTotal = done, total + if total > 0 { + lastPct = done * 100 / total + } + }); err != nil { + t.Fatalf("Download: %v", err) + } + + // Byte count is the assertion that always holds; the percentage only + // exists when the server advertises a length. + if lastDone != int64(len(data)) { + t.Errorf("final progress bytes = %d, want %d", lastDone, len(data)) + } + if sawTotal > 0 && lastPct != 100 { + t.Errorf("final progress = %d%%, want 100%%", lastPct) + } + if !IsInstalled(m) { + t.Fatal("model should be installed after bzip2 extraction") + } + + root, _ := Path(m) + for _, f := range m.Files { + st, err := os.Stat(filepath.Join(root, f)) + if err != nil { + t.Errorf("missing %s: %v", f, err) + continue + } + if st.Size() == 0 { + t.Errorf("%s is empty", f) + } + t.Logf("extracted %-22s %10d bytes", f, st.Size()) + } + + // The archive nests everything under a top-level dir; stripTopLevel must + // have flattened exactly one level, not zero and not two. + if _, err := os.Stat(filepath.Join(root, m.Dirname)); err == nil { + t.Error("top-level directory was not flattened (double nesting)") + } + // Staging must not survive a successful install. + if _, err := os.Stat(root + ".staging"); !os.IsNotExist(err) { + t.Error("staging directory left behind") + } + if _, err := os.Stat(root + ".part"); !os.IsNotExist(err) { + t.Error(".part file left behind") + } +} diff --git a/internal/sttmodel/catalog.go b/internal/sttmodel/catalog.go new file mode 100644 index 0000000..d9b0450 --- /dev/null +++ b/internal/sttmodel/catalog.go @@ -0,0 +1,292 @@ +// Package sttmodel is the catalog of speech-to-text models vox can download +// and run. It covers both whisper.cpp GGML models (a single .bin file) and +// Parakeet models (a directory of ONNX files extracted from a tar.bz2). +package sttmodel + +import ( + "fmt" + "os" + "path/filepath" + "slices" + "strings" +) + +// Engine identifies which inference backend runs a model. +type Engine string + +const ( + EngineWhisper Engine = "whisper" + EngineParakeet Engine = "parakeet" +) + +// DefaultID is the model vox uses when no override exists in flags, env, or +// prefs. Whisper stays the default until benchmarks justify switching. +const DefaultID = "base.en" + +// Model describes one downloadable STT model. +type Model struct { + ID string + Engine Engine + Label string + URL string + Checksum string // SHA-256 of the downloaded artifact + SizeMB int + + // Filename is set for single-file models (whisper GGML). Empty for archives. + Filename string + + // Dirname is the extracted directory name for archive models. Empty for + // single-file models. + Dirname string + + // Files lists paths (relative to Dirname) that must exist for an archive + // model to count as installed. Empty for single-file models. + Files []string +} + +// IsArchive reports whether the model ships as an archive that extracts into a +// directory, rather than a single downloadable file. +func (m Model) IsArchive() bool { return m.Dirname != "" } + +// Whisper checksums are SHA-256 from HuggingFace Git LFS pointers: +// https://huggingface.co/ggerganov/whisper.cpp/tree/main +// +// Parakeet checksums are SHA-256 of the sherpa-onnx release archives, verified +// by download on 2026-08-19. +var catalog = []Model{ + { + ID: "tiny.en", + Engine: EngineWhisper, + Label: "Whisper Tiny (English, ~75 MiB)", + Filename: "ggml-tiny.en.bin", + URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-tiny.en.bin", + Checksum: "921e4cf8686fdd993dcd081a5da5b6c365bfde1162e72b08d75ac75289920b1f", + SizeMB: 75, + }, + { + ID: "base.en", + Engine: EngineWhisper, + Label: "Whisper Base (English, ~142 MiB, default)", + Filename: "ggml-base.en.bin", + URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-base.en.bin", + Checksum: "a03779c86df3323075f5e796cb2ce5029f00ec8869eee3fdfb897afe36c6d002", + SizeMB: 142, + }, + { + ID: "small.en", + Engine: EngineWhisper, + Label: "Whisper Small (English, ~466 MiB)", + Filename: "ggml-small.en.bin", + URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-small.en.bin", + Checksum: "c6138d6d58ecc8322097e0f987c32f1be8bb0a18532a3f88f734d1bbf9c41e5d", + SizeMB: 466, + }, + { + ID: "medium.en", + Engine: EngineWhisper, + Label: "Whisper Medium (English, ~1.5 GiB)", + Filename: "ggml-medium.en.bin", + URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-medium.en.bin", + Checksum: "cc37e93478338ec7700281a7ac30a10128929eb8f427dda2e865faa8f6da4356", + SizeMB: 1536, + }, + { + ID: "large-v3-turbo", + Engine: EngineWhisper, + Label: "Whisper Large v3 Turbo (Multilingual, ~1.5 GiB)", + Filename: "ggml-large-v3-turbo.bin", + URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-large-v3-turbo.bin", + Checksum: "1fc70f774d38eb169993ac391eea357ef47c88757ef72ee5943879b7e8e2bc69", + SizeMB: 1536, + }, + { + ID: "parakeet-v2", + Engine: EngineParakeet, + Label: "Parakeet v2 (English, ~460 MiB)", + Dirname: "sherpa-onnx-nemo-parakeet-tdt-0.6b-v2-int8", + URL: "https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-nemo-parakeet-tdt-0.6b-v2-int8.tar.bz2", + // Verified 2026-08-19 (482468385 bytes). + Checksum: "157c157bc51155e03e37d2466522a3a737dd9c72bb25f36eb18912964161e1ad", + SizeMB: 460, + Files: []string{ + "encoder.int8.onnx", + "decoder.int8.onnx", + "joiner.int8.onnx", + "tokens.txt", + }, + }, + { + ID: "parakeet-v3", + Engine: EngineParakeet, + Label: "Parakeet v3 (25 European languages, ~465 MiB)", + Dirname: "sherpa-onnx-nemo-parakeet-tdt-0.6b-v3-int8", + URL: "https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-nemo-parakeet-tdt-0.6b-v3-int8.tar.bz2", + // Verified 2026-08-19 (487170055 bytes). + Checksum: "5793d0fd397c5778d2cf2126994d58e9d56b1be7c04d13c7a15bb1b4eafb16bf", + SizeMB: 465, + Files: []string{ + "encoder.int8.onnx", + "decoder.int8.onnx", + "joiner.int8.onnx", + "tokens.txt", + }, + }, +} + +// All returns a copy of the catalog. +func All() []Model { return slices.Clone(catalog) } + +// ByID resolves a model by its stable ID. +func ByID(id string) (Model, bool) { + for _, m := range catalog { + if m.ID == id { + return m, true + } + } + return Model{}, false +} + +// ByEngine returns every catalog model for one engine, in catalog order. +func ByEngine(e Engine) []Model { + var out []Model + for _, m := range catalog { + if m.Engine == e { + out = append(out, m) + } + } + return out +} + +// RootDir is the base directory for all downloaded models. +// VOX_MODEL_DIR overrides it, primarily for tests. +func RootDir() (string, error) { + if p := os.Getenv("VOX_MODEL_DIR"); p != "" { + return p, nil + } + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + return filepath.Join(home, ".local", "share", "vox", "models"), nil +} + +// legacyWhisperDir is where vox stored GGML models before the sttmodel split. +// Honored so existing installs do not re-download multi-gigabyte models. +func legacyWhisperDir() (string, error) { + if p := os.Getenv("WHISPER_MODEL_DIR"); p != "" { + return p, nil + } + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + return filepath.Join(home, ".local", "share", "whisper-cpp"), nil +} + +// Path is the canonical location for a model: a file path for single-file +// models, a directory path for archive models. New downloads always land here. +func Path(m Model) (string, error) { + root, err := RootDir() + if err != nil { + return "", err + } + dir := filepath.Join(root, string(m.Engine)) + if m.IsArchive() { + return filepath.Join(dir, m.Dirname), nil + } + return filepath.Join(dir, m.Filename), nil +} + +// legacyPath is the pre-migration location, or "" if the model has none. +func legacyPath(m Model) (string, error) { + if m.Engine != EngineWhisper { + return "", nil + } + dir, err := legacyWhisperDir() + if err != nil { + return "", err + } + return filepath.Join(dir, m.Filename), nil +} + +// ResolvePath returns where the model actually lives on disk. It prefers the +// canonical path and falls back to the legacy whisper directory so existing +// installs keep working without a re-download. +func ResolvePath(m Model) (string, error) { + p, err := Path(m) + if err != nil { + return "", err + } + if existsComplete(m, p) { + return p, nil + } + lp, err := legacyPath(m) + if err != nil { + return "", err + } + if lp != "" && existsComplete(m, lp) { + return lp, nil + } + return p, nil +} + +// existsComplete reports whether a fully usable model exists at root. +func existsComplete(m Model, root string) bool { + if m.IsArchive() { + if len(m.Files) == 0 { + return false // archive with no listed files is never complete + } + for _, f := range m.Files { + st, err := os.Stat(filepath.Join(root, f)) + if err != nil || st.Size() == 0 { + return false + } + } + return true + } + st, err := os.Stat(root) + return err == nil && st.Size() > 0 +} + +// IsInstalled reports whether the model is usable from either location. +func IsInstalled(m Model) bool { + p, err := ResolvePath(m) + if err != nil { + return false + } + return existsComplete(m, p) +} + +// InstalledCount returns how many catalog models are present on disk. +func InstalledCount() int { + n := 0 + for _, m := range catalog { + if IsInstalled(m) { + n++ + } + } + return n +} + +// InstalledCountForEngine returns how many models of one engine are on disk. +func InstalledCountForEngine(e Engine) int { + n := 0 + for _, m := range ByEngine(e) { + if IsInstalled(m) { + n++ + } + } + return n +} + +func validateChecksum(gotHex, expected string) error { + if expected == "" { + return fmt.Errorf("missing checksum") + } + got := strings.ToLower(strings.TrimSpace(gotHex)) + want := strings.ToLower(strings.TrimSpace(expected)) + if got != want { + return fmt.Errorf("checksum mismatch: got %s, want %s", got, want) + } + return nil +} diff --git a/internal/sttmodel/catalog_test.go b/internal/sttmodel/catalog_test.go new file mode 100644 index 0000000..e88980e --- /dev/null +++ b/internal/sttmodel/catalog_test.go @@ -0,0 +1,221 @@ +package sttmodel + +import ( + "os" + "path/filepath" + "testing" +) + +func TestByIDKnownModels(t *testing.T) { + for _, id := range []string{"base.en", "tiny.en", "parakeet-v2", "parakeet-v3"} { + if _, ok := ByID(id); !ok { + t.Errorf("ByID(%q) not found", id) + } + } + if _, ok := ByID("nope"); ok { + t.Error("ByID(\"nope\") should not be found") + } +} + +func TestDefaultIDIsWhisperBase(t *testing.T) { + if DefaultID != "base.en" { + t.Errorf("DefaultID = %q, want base.en", DefaultID) + } + m, ok := ByID(DefaultID) + if !ok { + t.Fatal("default model not in catalog") + } + if m.Engine != EngineWhisper { + t.Errorf("default engine = %q, want %q", m.Engine, EngineWhisper) + } +} + +func TestParakeetModelsAreArchives(t *testing.T) { + m, _ := ByID("parakeet-v2") + if m.Engine != EngineParakeet { + t.Errorf("engine = %q, want %q", m.Engine, EngineParakeet) + } + if !m.IsArchive() { + t.Error("parakeet-v2 should be an archive model") + } + want := []string{"encoder.int8.onnx", "decoder.int8.onnx", "joiner.int8.onnx", "tokens.txt"} + if len(m.Files) != len(want) { + t.Fatalf("Files = %v, want %v", m.Files, want) + } + if m.Checksum != "157c157bc51155e03e37d2466522a3a737dd9c72bb25f36eb18912964161e1ad" { + t.Errorf("unexpected v2 checksum %q", m.Checksum) + } +} + +func TestWhisperModelsAreSingleFile(t *testing.T) { + m, _ := ByID("base.en") + if m.IsArchive() { + t.Error("base.en should not be an archive model") + } + if m.Filename != "ggml-base.en.bin" { + t.Errorf("Filename = %q", m.Filename) + } +} + +func TestPathAndEngineDirs(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + + w, _ := ByID("base.en") + gotW, err := Path(w) + if err != nil { + t.Fatal(err) + } + wantW := filepath.Join(dir, "whisper", "ggml-base.en.bin") + if gotW != wantW { + t.Errorf("whisper path = %q, want %q", gotW, wantW) + } + + p, _ := ByID("parakeet-v2") + gotP, err := Path(p) + if err != nil { + t.Fatal(err) + } + wantP := filepath.Join(dir, "parakeet", "sherpa-onnx-nemo-parakeet-tdt-0.6b-v2-int8") + if gotP != wantP { + t.Errorf("parakeet path = %q, want %q", gotP, wantP) + } +} + +func TestIsInstalledSingleFile(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + // Isolate the legacy dir too: without this the whisper fallback finds a + // real ~/.local/share/whisper-cpp install on a developer machine and the + // "empty dir" assertion below fails for reasons unrelated to the code. + t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + m, _ := ByID("base.en") + + if IsInstalled(m) { + t.Error("should not be installed in empty dir") + } + p, _ := Path(m) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + if !IsInstalled(m) { + t.Error("should be installed after writing file") + } +} + +func TestIsInstalledArchiveRequiresAllFiles(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + m, _ := ByID("parakeet-v2") + root, _ := Path(m) + if err := os.MkdirAll(root, 0o755); err != nil { + t.Fatal(err) + } + + // Partial extraction: only some required files present. + for _, f := range m.Files[:2] { + if err := os.WriteFile(filepath.Join(root, f), []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + } + if IsInstalled(m) { + t.Error("partial extraction should not count as installed") + } + + for _, f := range m.Files[2:] { + if err := os.WriteFile(filepath.Join(root, f), []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + } + if !IsInstalled(m) { + t.Error("complete extraction should count as installed") + } +} + +func TestLegacyWhisperDirFallback(t *testing.T) { + newDir := t.TempDir() + legacy := t.TempDir() + t.Setenv("VOX_MODEL_DIR", newDir) + t.Setenv("WHISPER_MODEL_DIR", legacy) + + m, _ := ByID("base.en") + legacyPath := filepath.Join(legacy, "ggml-base.en.bin") + if err := os.WriteFile(legacyPath, []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + + if !IsInstalled(m) { + t.Error("legacy-installed model should be detected") + } + got, err := ResolvePath(m) + if err != nil { + t.Fatal(err) + } + if got != legacyPath { + t.Errorf("ResolvePath = %q, want legacy %q", got, legacyPath) + } +} + +func TestRemoveAtLegacyPath(t *testing.T) { + newDir := t.TempDir() + legacy := t.TempDir() + t.Setenv("VOX_MODEL_DIR", newDir) + t.Setenv("WHISPER_MODEL_DIR", legacy) + + m, _ := ByID("base.en") + legacyFile := filepath.Join(legacy, "ggml-base.en.bin") + if err := os.WriteFile(legacyFile, []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + if !IsInstalled(m) { + t.Fatal("model should be installed at legacy path") + } + if err := Remove(m); err != nil { + t.Fatalf("Remove: %v", err) + } + if IsInstalled(m) { + t.Error("model should no longer be installed after Remove") + } +} + +func TestInstalledCountRespectsEngine(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + + if got := InstalledCount(); got != 0 { + t.Errorf("InstalledCount = %d, want 0 in an empty dir", got) + } + + // Install one whisper model. + w, _ := ByID("base.en") + wp, _ := Path(w) + if err := os.MkdirAll(filepath.Dir(wp), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(wp, []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + + if got := InstalledCount(); got != 1 { + t.Errorf("InstalledCount = %d, want 1", got) + } + if got := InstalledCountForEngine(EngineWhisper); got != 1 { + t.Errorf("whisper installed = %d, want 1", got) + } + if got := InstalledCountForEngine(EngineParakeet); got != 0 { + t.Errorf("parakeet installed = %d, want 0", got) + } +} + +func TestByEngine(t *testing.T) { + if got := len(ByEngine(EngineParakeet)); got != 2 { + t.Errorf("parakeet models = %d, want 2", got) + } + if got := len(ByEngine(EngineWhisper)); got != 5 { + t.Errorf("whisper models = %d, want 5", got) + } +} diff --git a/internal/sttmodel/download.go b/internal/sttmodel/download.go new file mode 100644 index 0000000..474cbf2 --- /dev/null +++ b/internal/sttmodel/download.go @@ -0,0 +1,156 @@ +package sttmodel + +import ( + "context" + "crypto/sha256" + "fmt" + "io" + "net/http" + "os" + "path/filepath" + "sync" + "time" +) + +// maxDownloadBytes caps any single download. +const maxDownloadBytes = 4 * 1024 * 1024 * 1024 // 4 GiB + +var downloadLocks sync.Map // dest path -> *sync.Mutex + +var httpClient = &http.Client{ + Timeout: 30 * time.Minute, + CheckRedirect: func(req *http.Request, via []*http.Request) error { + if len(via) >= 10 { + return fmt.Errorf("too many redirects") + } + if req.URL.Scheme != "https" { + return fmt.Errorf("refusing non-HTTPS redirect to %s", req.URL) + } + return nil + }, +} + +// Download fetches a model and installs it at its canonical path. Single-file +// models are written directly; archive models are extracted and their required +// files verified. onProgress may be nil. +// +// Download is a no-op if the model is already installed. +func Download(ctx context.Context, m Model, onProgress func(downloaded, total int64)) error { + dest, err := Path(m) + if err != nil { + return err + } + + mu := downloadLock(dest) + mu.Lock() + defer mu.Unlock() + + // Re-check under the lock: a concurrent caller may have finished. + if IsInstalled(m) { + return nil + } + + if err := os.MkdirAll(filepath.Dir(dest), 0o755); err != nil { + return fmt.Errorf("create model dir: %w", err) + } + + tmp := dest + ".part" + defer os.Remove(tmp) + + sum, err := fetchToFile(ctx, m.URL, tmp, onProgress) + if err != nil { + return err + } + if err := validateChecksum(sum, m.Checksum); err != nil { + return err + } + + if !m.IsArchive() { + return os.Rename(tmp, dest) + } + + // Extract into a staging dir, verify, then swap into place atomically. + staging := dest + ".staging" + os.RemoveAll(staging) + defer os.RemoveAll(staging) + + if err := extractArchive(tmp, staging); err != nil { + return err + } + if err := verifyFiles(staging, m.Files); err != nil { + return fmt.Errorf("archive for %s is incomplete: %w", m.ID, err) + } + os.RemoveAll(dest) + if err := os.Rename(staging, dest); err != nil { + return fmt.Errorf("install extracted model: %w", err) + } + return nil +} + +// fetchToFile downloads url into path and returns the hex SHA-256 of the body. +func fetchToFile(ctx context.Context, url, path string, onProgress func(int64, int64)) (string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return "", fmt.Errorf("create request: %w", err) + } + resp, err := httpClient.Do(req) + if err != nil { + return "", fmt.Errorf("download: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("download returned status %d", resp.StatusCode) + } + + f, err := os.Create(path) + if err != nil { + return "", fmt.Errorf("create temp file: %w", err) + } + defer f.Close() + + hasher := sha256.New() + pw := &progressWriter{ + fn: onProgress, + total: resp.ContentLength, + interval: 250 * time.Millisecond, + } + if _, err := io.Copy(io.MultiWriter(f, hasher, pw), io.LimitReader(resp.Body, maxDownloadBytes)); err != nil { + return "", fmt.Errorf("write download: %w", err) + } + pw.flush() + + if err := f.Sync(); err != nil { + return "", fmt.Errorf("sync: %w", err) + } + return fmt.Sprintf("%x", hasher.Sum(nil)), nil +} + +func downloadLock(dest string) *sync.Mutex { + v, _ := downloadLocks.LoadOrStore(dest, &sync.Mutex{}) + return v.(*sync.Mutex) +} + +// progressWriter reports bytes written, throttled to interval. +type progressWriter struct { + fn func(downloaded, total int64) + total int64 + written int64 + lastSent time.Time + interval time.Duration +} + +func (p *progressWriter) Write(b []byte) (int, error) { + p.written += int64(len(b)) + if p.fn != nil && time.Since(p.lastSent) >= p.interval { + p.lastSent = time.Now() + p.fn(p.written, p.total) + } + return len(b), nil +} + +// flush emits a final progress callback so consumers always see completion. +func (p *progressWriter) flush() { + if p.fn != nil { + p.fn(p.written, p.total) + } +} diff --git a/internal/sttmodel/download_test.go b/internal/sttmodel/download_test.go new file mode 100644 index 0000000..0fce7e7 --- /dev/null +++ b/internal/sttmodel/download_test.go @@ -0,0 +1,406 @@ +package sttmodel + +import ( + "archive/tar" + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + "sync/atomic" + "testing" + "time" +) + +func sha256Hex(b []byte) string { + sum := sha256.Sum256(b) + return hex.EncodeToString(sum[:]) +} + +// buildTestArchive builds an in-memory .tar.bz2-shaped payload for tests. +// bzip2 compression is not available in the stdlib for writing, so tests use +// an uncompressed tar; extractArchive sniffs the format. +func buildTestArchive(t *testing.T, root string, files map[string]string) []byte { + t.Helper() + var buf bytes.Buffer + tw := tar.NewWriter(&buf) + if err := tw.WriteHeader(&tar.Header{ + Name: root + "/", + Typeflag: tar.TypeDir, + Mode: 0o755, + }); err != nil { + t.Fatal(err) + } + for name, body := range files { + if err := tw.WriteHeader(&tar.Header{ + Name: root + "/" + name, + Mode: 0o644, + Size: int64(len(body)), + }); err != nil { + t.Fatal(err) + } + if _, err := tw.Write([]byte(body)); err != nil { + t.Fatal(err) + } + } + if err := tw.Close(); err != nil { + t.Fatal(err) + } + return buf.Bytes() +} + +func TestDownloadSingleFileVerifiesChecksum(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", filepath.Join(dir, "legacy-empty")) + + payload := []byte("fake ggml model bytes") + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write(payload) + })) + defer srv.Close() + + m := Model{ + ID: "test-single", + Engine: EngineWhisper, + Filename: "test.bin", + URL: srv.URL, + Checksum: sha256Hex(payload), + } + + var lastPct int64 + if err := Download(context.Background(), m, func(done, total int64) { + if total > 0 { + lastPct = done * 100 / total + } + }); err != nil { + t.Fatalf("Download: %v", err) + } + if !IsInstalled(m) { + t.Error("model should be installed after download") + } + if lastPct != 100 { + t.Errorf("final progress = %d%%, want 100%%", lastPct) + } + + p, _ := Path(m) + got, _ := os.ReadFile(p) + if string(got) != string(payload) { + t.Errorf("content mismatch") + } +} + +func TestDownloadRejectsBadChecksum(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", filepath.Join(dir, "legacy-empty")) + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("actual bytes")) + })) + defer srv.Close() + + m := Model{ + ID: "test-bad", + Engine: EngineWhisper, + Filename: "bad.bin", + URL: srv.URL, + Checksum: sha256Hex([]byte("different bytes")), + } + + if err := Download(context.Background(), m, nil); err == nil { + t.Fatal("expected checksum error, got nil") + } + if IsInstalled(m) { + t.Error("failed download must not leave an installed model") + } + p, _ := Path(m) + if _, err := os.Stat(p + ".part"); !os.IsNotExist(err) { + t.Error("temp .part file should be cleaned up") + } +} + +func TestDownloadArchiveExtractsAndVerifies(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", filepath.Join(dir, "legacy-empty")) + + archive := buildTestArchive(t, "test-model", map[string]string{ + "encoder.int8.onnx": "enc", + "decoder.int8.onnx": "dec", + "joiner.int8.onnx": "join", + "tokens.txt": "tok", + }) + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write(archive) + })) + defer srv.Close() + + m := Model{ + ID: "test-archive", + Engine: EngineParakeet, + Dirname: "test-model", + URL: srv.URL, + Checksum: sha256Hex(archive), + Files: []string{"encoder.int8.onnx", "decoder.int8.onnx", "joiner.int8.onnx", "tokens.txt"}, + } + + if err := Download(context.Background(), m, nil); err != nil { + t.Fatalf("Download: %v", err) + } + if !IsInstalled(m) { + t.Fatal("archive model should be installed") + } + root, _ := Path(m) + got, err := os.ReadFile(filepath.Join(root, "tokens.txt")) + if err != nil { + t.Fatal(err) + } + if string(got) != "tok" { + t.Errorf("tokens.txt = %q, want %q", got, "tok") + } +} + +func TestDownloadArchiveRejectsMissingFiles(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", filepath.Join(dir, "legacy-empty")) + + // Archive is missing joiner.int8.onnx. + archive := buildTestArchive(t, "incomplete", map[string]string{ + "encoder.int8.onnx": "enc", + "decoder.int8.onnx": "dec", + "tokens.txt": "tok", + }) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write(archive) + })) + defer srv.Close() + + m := Model{ + ID: "test-incomplete", + Engine: EngineParakeet, + Dirname: "incomplete", + URL: srv.URL, + Checksum: sha256Hex(archive), + Files: []string{"encoder.int8.onnx", "decoder.int8.onnx", "joiner.int8.onnx", "tokens.txt"}, + } + + if err := Download(context.Background(), m, nil); err == nil { + t.Fatal("expected missing-file error, got nil") + } + if IsInstalled(m) { + t.Error("incomplete extraction must not count as installed") + } + root, _ := Path(m) + if _, err := os.Stat(root); !os.IsNotExist(err) { + t.Error("incomplete extraction directory should be removed") + } +} + +func TestConcurrentDownloadsRunOnce(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", filepath.Join(dir, "legacy-empty")) + + payload := []byte("model bytes") + var hits int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + atomic.AddInt32(&hits, 1) + time.Sleep(50 * time.Millisecond) // widen the race window + _, _ = w.Write(payload) + })) + defer srv.Close() + + m := Model{ + ID: "test-concurrent", + Engine: EngineWhisper, + Filename: "concurrent.bin", + URL: srv.URL, + Checksum: sha256Hex(payload), + } + + const callers = 5 + var wg sync.WaitGroup + errs := make(chan error, callers) + for i := 0; i < callers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + errs <- Download(context.Background(), m, nil) + }() + } + wg.Wait() + close(errs) + + for err := range errs { + if err != nil { + t.Errorf("Download: %v", err) + } + } + // The per-destination lock plus the re-check inside it must collapse + // five callers into one fetch. More than one means the guard is broken + // and users on a slow link would pull 460 MiB twice. + if got := atomic.LoadInt32(&hits); got != 1 { + t.Errorf("server hits = %d, want 1", got) + } + if !IsInstalled(m) { + t.Error("model should be installed") + } +} + +func TestDownloadCleansUpOnContextCancel(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", filepath.Join(dir, "legacy-empty")) + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Length", "1000000") + w.WriteHeader(http.StatusOK) + for i := 0; i < 100; i++ { + _, _ = w.Write(make([]byte, 10000)) + w.(http.Flusher).Flush() + time.Sleep(10 * time.Millisecond) + } + })) + defer srv.Close() + + m := Model{ + ID: "test-cancel", + Engine: EngineWhisper, + Filename: "cancel.bin", + URL: srv.URL, + Checksum: sha256Hex([]byte("irrelevant")), + } + + ctx, cancel := context.WithTimeout(context.Background(), 60*time.Millisecond) + defer cancel() + + if err := Download(ctx, m, nil); err == nil { + t.Fatal("expected an error from the cancelled download") + } + if IsInstalled(m) { + t.Error("cancelled download must not leave an installed model") + } + p, _ := Path(m) + if _, err := os.Stat(p + ".part"); !os.IsNotExist(err) { + t.Error("temp .part file should be cleaned up after cancellation") + } +} + +func TestRemoveArchiveModel(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", filepath.Join(dir, "legacy-empty")) + + m, _ := ByID("parakeet-v2") + root, _ := Path(m) + if err := os.MkdirAll(root, 0o755); err != nil { + t.Fatal(err) + } + for _, f := range m.Files { + if err := os.WriteFile(filepath.Join(root, f), []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + } + if !IsInstalled(m) { + t.Fatal("setup failed") + } + if err := Remove(m); err != nil { + t.Fatalf("Remove: %v", err) + } + if IsInstalled(m) { + t.Error("model should be gone after Remove") + } +} + +// TestStripTopLevelAbsorbsOneLevel documents that stripTopLevel's +// Clean-then-drop-first-component already neutralises a single "..", so +// "root/../../escape.txt" lands harmlessly inside destDir rather than beside +// it. safeJoin is the backstop for anything deeper. +func TestStripTopLevelAbsorbsOneLevel(t *testing.T) { + if got := stripTopLevel("root/../../escape.txt"); got != "escape.txt" { + t.Errorf("stripTopLevel = %q, want %q", got, "escape.txt") + } + if got := stripTopLevel("model/encoder.int8.onnx"); got != "encoder.int8.onnx" { + t.Errorf("stripTopLevel = %q, want %q", got, "encoder.int8.onnx") + } + if got := stripTopLevel("model/"); got != "" { + t.Errorf("stripTopLevel of root entry = %q, want empty", got) + } +} + +// TestExtractArchiveRejectsPathTraversal covers the safeJoin guard directly: +// a malicious archive entry must not be able to write outside destDir. +func TestExtractArchiveRejectsPathTraversal(t *testing.T) { + tmp := t.TempDir() + + var buf bytes.Buffer + tw := tar.NewWriter(&buf) + body := "pwned" + // Three "..' segments: one is consumed by the top-level strip, the rest + // must be caught by safeJoin. + if err := tw.WriteHeader(&tar.Header{ + Name: "root/../../../escape.txt", + Mode: 0o644, + Size: int64(len(body)), + }); err != nil { + t.Fatal(err) + } + if _, err := tw.Write([]byte(body)); err != nil { + t.Fatal(err) + } + if err := tw.Close(); err != nil { + t.Fatal(err) + } + + archivePath := filepath.Join(tmp, "evil.tar") + if err := os.WriteFile(archivePath, buf.Bytes(), 0o644); err != nil { + t.Fatal(err) + } + + dest := filepath.Join(tmp, "nested", "dest") + if err := extractArchive(archivePath, dest); err == nil { + t.Fatal("expected extractArchive to reject a traversal entry") + } + if _, err := os.Stat(filepath.Join(tmp, "escape.txt")); !os.IsNotExist(err) { + t.Error("traversal entry escaped the destination directory") + } +} + +func TestRemoveIsIdempotent(t *testing.T) { + dir := t.TempDir() + t.Setenv("VOX_MODEL_DIR", dir) + t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) + + // Removing a model that was never installed must not error, for either + // shape. Callers should not have to guard with IsInstalled. + single, _ := ByID("base.en") + if err := Remove(single); err != nil { + t.Errorf("Remove(uninstalled single-file) = %v, want nil", err) + } + archive, _ := ByID("parakeet-v2") + if err := Remove(archive); err != nil { + t.Errorf("Remove(uninstalled archive) = %v, want nil", err) + } + + // And removing twice after a real install is also fine. + p, _ := Path(single) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, []byte("x"), 0o644); err != nil { + t.Fatal(err) + } + if err := Remove(single); err != nil { + t.Fatalf("first Remove: %v", err) + } + if err := Remove(single); err != nil { + t.Errorf("second Remove = %v, want nil", err) + } +} diff --git a/internal/sttmodel/remove.go b/internal/sttmodel/remove.go new file mode 100644 index 0000000..361c493 --- /dev/null +++ b/internal/sttmodel/remove.go @@ -0,0 +1,34 @@ +package sttmodel + +import "os" + +// Remove deletes an installed model from disk, including any leftover +// temp artifacts. Archive models remove their whole directory. +// +// Remove is idempotent: removing a model that is not installed is not an +// error. os.RemoveAll already behaves this way for the archive branch, and +// the single-file branch matches it so callers do not have to guard with +// IsInstalled to avoid a spurious failure. +func Remove(m Model) error { + p, err := ResolvePath(m) + if err != nil { + return err + } + // Clean temp artifacts at the resolved path (where the model lives). + os.Remove(p + ".part") + os.RemoveAll(p + ".staging") + // Also clean temp artifacts at the canonical path if it differs + // (Download creates them there, but ResolvePath may return the legacy path). + cp, cpErr := Path(m) + if cpErr == nil && cp != p { + os.Remove(cp + ".part") + os.RemoveAll(cp + ".staging") + } + if m.IsArchive() { + return os.RemoveAll(p) + } + if err := os.Remove(p); err != nil && !os.IsNotExist(err) { + return err + } + return nil +} diff --git a/internal/transcribe/transcriber.go b/internal/transcribe/transcriber.go new file mode 100644 index 0000000..076f700 --- /dev/null +++ b/internal/transcribe/transcriber.go @@ -0,0 +1,20 @@ +package transcribe + +import "context" + +// Transcriber converts recorded audio into text. Implementations may call a +// local HTTP server (whisper.cpp), run inference in-process (Parakeet via +// sherpa-onnx), or stub the call out in tests. +// +// wavData is a complete 16 kHz mono 16-bit WAV payload as produced by +// internal/audio. +// +// TranscribeOptions is the seam for engine-specific knobs. Not every option +// applies to every engine: Parakeet ignores InitialPrompt (it has no +// equivalent of whisper's prompt conditioning) and ignores Language, because +// the chosen model fixes the language set. If a future engine needs a runtime +// language hint (Parakeet v3 covers 25 languages but auto-detects), add it +// here rather than widening the interface. +type Transcriber interface { + Transcribe(ctx context.Context, wavData []byte, opts TranscribeOptions) (string, error) +} diff --git a/internal/transcribe/transcriber_test.go b/internal/transcribe/transcriber_test.go new file mode 100644 index 0000000..ce629b2 --- /dev/null +++ b/internal/transcribe/transcriber_test.go @@ -0,0 +1,29 @@ +package transcribe + +import ( + "context" + "testing" +) + +// stubTranscriber is a minimal Transcriber used to prove the interface is +// satisfiable by something other than *Client. +type stubTranscriber struct{ text string } + +func (s stubTranscriber) Transcribe(_ context.Context, _ []byte, _ TranscribeOptions) (string, error) { + return s.text, nil +} + +func TestClientSatisfiesTranscriber(t *testing.T) { + var _ Transcriber = (*Client)(nil) +} + +func TestStubSatisfiesTranscriber(t *testing.T) { + var tr Transcriber = stubTranscriber{text: "hello"} + got, err := tr.Transcribe(context.Background(), []byte("ignored"), TranscribeOptions{}) + if err != nil { + t.Fatalf("Transcribe: %v", err) + } + if got != "hello" { + t.Errorf("got %q, want %q", got, "hello") + } +} diff --git a/internal/whispermodel/catalog.go b/internal/whispermodel/catalog.go deleted file mode 100644 index 563cade..0000000 --- a/internal/whispermodel/catalog.go +++ /dev/null @@ -1,137 +0,0 @@ -package whispermodel - -import ( - "fmt" - "os" - "path/filepath" - "slices" - "strings" -) - -// DefaultID is the model Vox uses when no override exists in env or prefs. -const DefaultID = "base.en" - -// Model describes one downloadable whisper.cpp GGML model. -type Model struct { - ID string - Label string - Filename string - URL string - Checksum string - SizeMB int -} - -// All checksums are SHA-256, sourced from HuggingFace Git LFS pointers: -// https://huggingface.co/ggerganov/whisper.cpp/tree/main -var catalog = []Model{ - { - ID: "tiny.en", - Label: "Tiny (English, ~75 MiB)", - Filename: "ggml-tiny.en.bin", - URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-tiny.en.bin", - Checksum: "921e4cf8686fdd993dcd081a5da5b6c365bfde1162e72b08d75ac75289920b1f", - SizeMB: 75, - }, - { - ID: "base.en", - Label: "Base (English, ~142 MiB, default)", - Filename: "ggml-base.en.bin", - URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-base.en.bin", - Checksum: "a03779c86df3323075f5e796cb2ce5029f00ec8869eee3fdfb897afe36c6d002", - SizeMB: 142, - }, - { - ID: "small.en", - Label: "Small (English, ~466 MiB)", - Filename: "ggml-small.en.bin", - URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-small.en.bin", - Checksum: "c6138d6d58ecc8322097e0f987c32f1be8bb0a18532a3f88f734d1bbf9c41e5d", - SizeMB: 466, - }, - { - ID: "medium.en", - Label: "Medium (English, ~1.5 GiB)", - Filename: "ggml-medium.en.bin", - URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-medium.en.bin", - Checksum: "cc37e93478338ec7700281a7ac30a10128929eb8f427dda2e865faa8f6da4356", - SizeMB: 1536, - }, - { - ID: "large-v3-turbo", - Label: "Large v3 Turbo (Multilingual, ~1.5 GiB)", - Filename: "ggml-large-v3-turbo.bin", - URL: "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-large-v3-turbo.bin", - Checksum: "1fc70f774d38eb169993ac391eea357ef47c88757ef72ee5943879b7e8e2bc69", - SizeMB: 1536, - }, -} - -// All returns a copy of the supported model catalog. -func All() []Model { - return slices.Clone(catalog) -} - -// ByID resolves a model by stable ID (for prefs/env values). -func ByID(id string) (Model, bool) { - for _, m := range catalog { - if m.ID == id { - return m, true - } - } - return Model{}, false -} - -// ModelDir is the download/cache directory for GGML model files. -// Honors WHISPER_MODEL_DIR for parity with Makefile behavior. -func ModelDir() (string, error) { - if p := os.Getenv("WHISPER_MODEL_DIR"); p != "" { - return p, nil - } - home, err := os.UserHomeDir() - if err != nil { - return "", err - } - return filepath.Join(home, ".local", "share", "whisper-cpp"), nil -} - -// Path returns the target path for the model file under ModelDir. -func Path(m Model) (string, error) { - dir, err := ModelDir() - if err != nil { - return "", err - } - return filepath.Join(dir, m.Filename), nil -} - -// IsInstalled returns true if the model file exists and is non-empty. -func IsInstalled(m Model) bool { - path, err := Path(m) - if err != nil { - return false - } - st, err := os.Stat(path) - return err == nil && st.Size() > 0 -} - -// InstalledCount returns how many catalog models are present on disk. -func InstalledCount() int { - n := 0 - for _, m := range catalog { - if IsInstalled(m) { - n++ - } - } - return n -} - -func validateChecksum(gotHex, expected string) error { - if expected == "" { - return fmt.Errorf("missing checksum") - } - got := strings.ToLower(strings.TrimSpace(gotHex)) - want := strings.ToLower(strings.TrimSpace(expected)) - if got != want { - return fmt.Errorf("checksum mismatch: got %s, want %s", got, want) - } - return nil -} diff --git a/internal/whispermodel/download.go b/internal/whispermodel/download.go deleted file mode 100644 index 7a88bb9..0000000 --- a/internal/whispermodel/download.go +++ /dev/null @@ -1,168 +0,0 @@ -package whispermodel - -import ( - "context" - "crypto/sha1" - "crypto/sha256" - "encoding/hex" - "fmt" - "hash" - "io" - "net/http" - "os" - "path/filepath" - "strings" - "sync" - "time" -) - -var downloadLocks sync.Map // map[string]*sync.Mutex - -// maxDownloadBytes caps the response body to prevent a compromised CDN from -// filling disk. Set to 2x the largest catalog model (~1.5 GiB) plus margin. -const maxDownloadBytes = 4 * 1024 * 1024 * 1024 // 4 GiB - -// downloadHTTPClient is a dedicated client with a generous but finite timeout -// for model downloads. Using http.DefaultClient risks no timeout and inherits -// application-global transport changes. -var downloadHTTPClient = &http.Client{ - Timeout: 30 * time.Minute, // large models on slow connections - CheckRedirect: func(req *http.Request, via []*http.Request) error { - if req.URL.Scheme != "https" { - return fmt.Errorf("refusing non-HTTPS redirect to %s", req.URL) - } - if len(via) >= 10 { - return fmt.Errorf("too many redirects") - } - return nil - }, -} - -// Download fetches a model file to disk, verifies checksum, and atomically -// installs it at the final destination path. Existing valid files are kept. -func Download(ctx context.Context, m Model, onProgress func(downloaded, total int64)) error { - dest, err := Path(m) - if err != nil { - return err - } - lock := downloadLock(dest) - lock.Lock() - defer lock.Unlock() - - if IsInstalled(m) { - return nil - } - if err := os.MkdirAll(filepath.Dir(dest), 0o755); err != nil { - return err - } - - req, err := http.NewRequestWithContext(ctx, http.MethodGet, m.URL, nil) - if err != nil { - return fmt.Errorf("create request: %w", err) - } - resp, err := downloadHTTPClient.Do(req) - if err != nil { - return fmt.Errorf("download model: %w", err) - } - defer resp.Body.Close() - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("download model: unexpected status %d", resp.StatusCode) - } - - tmp := dest + ".part" - f, err := os.Create(tmp) - if err != nil { - return err - } - closed := false - defer func() { - if !closed { - _ = f.Close() - } - }() - defer func() { - _ = os.Remove(tmp) - }() - - hasher, checksumName, err := checksumHash(m.Checksum) - if err != nil { - return err - } - pw := &progressWriter{ - fn: onProgress, - total: resp.ContentLength, - interval: 250 * time.Millisecond, - } - // Limit the response body to prevent unbounded disk writes. - limited := io.LimitReader(resp.Body, maxDownloadBytes) - w := io.MultiWriter(f, hasher, pw) - if _, err := io.Copy(w, limited); err != nil { - return fmt.Errorf("write model file: %w", err) - } - pw.flush() - - got := hex.EncodeToString(hasher.Sum(nil)) - if err := validateChecksum(got, m.Checksum); err != nil { - return fmt.Errorf("verify %s: %w", checksumName, err) - } - if err := f.Sync(); err != nil { - return fmt.Errorf("sync model file: %w", err) - } - if err := f.Close(); err != nil { - return err - } - closed = true - if err := os.Rename(tmp, dest); err != nil { - return err - } - return nil -} - -func downloadLock(dest string) *sync.Mutex { - lock, _ := downloadLocks.LoadOrStore(dest, &sync.Mutex{}) - return lock.(*sync.Mutex) -} - -func checksumHash(sum string) (hash.Hash, string, error) { - switch len(strings.TrimSpace(sum)) { - case 40: - return sha1.New(), "sha1", nil - case 64: - return sha256.New(), "sha256", nil - default: - return nil, "", fmt.Errorf("unsupported checksum length %d", len(strings.TrimSpace(sum))) - } -} - -type progressWriter struct { - fn func(downloaded, total int64) - total int64 - written int64 - lastSent time.Time - interval time.Duration -} - -func (p *progressWriter) Write(b []byte) (int, error) { - n := len(b) - p.written += int64(n) - p.maybeSend() - return n, nil -} - -func (p *progressWriter) maybeSend() { - if p.fn == nil { - return - } - if time.Since(p.lastSent) < p.interval { - return - } - p.lastSent = time.Now() - p.fn(p.written, p.total) -} - -func (p *progressWriter) flush() { - if p.fn == nil { - return - } - p.fn(p.written, p.total) -} diff --git a/internal/whispermodel/remove.go b/internal/whispermodel/remove.go deleted file mode 100644 index 7eac173..0000000 --- a/internal/whispermodel/remove.go +++ /dev/null @@ -1,17 +0,0 @@ -package whispermodel - -import "os" - -// Remove deletes the on-disk model file and any in-progress partial download -// (.part) for this catalog entry. Missing files are not an error. -func Remove(m Model) error { - path, err := Path(m) - if err != nil { - return err - } - _ = os.Remove(path + ".part") - if err := os.Remove(path); err != nil && !os.IsNotExist(err) { - return err - } - return nil -} diff --git a/internal/whispermodel/whispermodel_test.go b/internal/whispermodel/whispermodel_test.go deleted file mode 100644 index e95c093..0000000 --- a/internal/whispermodel/whispermodel_test.go +++ /dev/null @@ -1,235 +0,0 @@ -package whispermodel - -import ( - "context" - "crypto/sha256" - "encoding/hex" - "net/http" - "net/http/httptest" - "os" - "path/filepath" - "sync" - "sync/atomic" - "testing" - "time" -) - -func TestByID(t *testing.T) { - if _, ok := ByID(DefaultID); !ok { - t.Fatalf("ByID(%q) not found", DefaultID) - } - if _, ok := ByID("nope"); ok { - t.Fatal("ByID(nope) unexpectedly found model") - } -} - -func TestInstalledCount(t *testing.T) { - t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) - if got := InstalledCount(); got != 0 { - t.Fatalf("InstalledCount with empty dir = %d, want 0", got) - } - tiny, ok := ByID("tiny.en") - if !ok { - t.Fatal("ByID tiny.en") - } - pathTiny, err := Path(tiny) - if err != nil { - t.Fatalf("Path: %v", err) - } - if err := os.WriteFile(pathTiny, []byte("x"), 0o600); err != nil { - t.Fatalf("WriteFile: %v", err) - } - if got := InstalledCount(); got != 1 { - t.Fatalf("InstalledCount one catalog model = %d, want 1", got) - } - base, ok := ByID("base.en") - if !ok { - t.Fatal("ByID base.en") - } - pathBase, err := Path(base) - if err != nil { - t.Fatalf("Path base: %v", err) - } - if err := os.WriteFile(pathBase, []byte("y"), 0o600); err != nil { - t.Fatalf("WriteFile base: %v", err) - } - if got := InstalledCount(); got != 2 { - t.Fatalf("InstalledCount two catalog models = %d, want 2", got) - } -} - -func TestPathAndIsInstalled(t *testing.T) { - t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) - m := Model{Filename: "ggml-foo.bin"} - path, err := Path(m) - if err != nil { - t.Fatalf("Path: %v", err) - } - if IsInstalled(m) { - t.Fatal("IsInstalled true before file exists") - } - if err := os.WriteFile(path, []byte("x"), 0o600); err != nil { - t.Fatalf("WriteFile: %v", err) - } - if !IsInstalled(m) { - t.Fatal("IsInstalled false after file write") - } -} - -func TestRemove(t *testing.T) { - t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) - m := Model{Filename: "ggml-foo.bin"} - path, err := Path(m) - if err != nil { - t.Fatalf("Path: %v", err) - } - part := path + ".part" - if err := os.WriteFile(path, []byte("installed"), 0o600); err != nil { - t.Fatalf("WriteFile: %v", err) - } - if err := os.WriteFile(part, []byte("partial"), 0o600); err != nil { - t.Fatalf("WriteFile partial: %v", err) - } - if err := Remove(m); err != nil { - t.Fatalf("Remove: %v", err) - } - if _, err := os.Stat(path); !os.IsNotExist(err) { - t.Fatal("expected model file gone after Remove") - } - if _, err := os.Stat(part); !os.IsNotExist(err) { - t.Fatal("expected .part file gone after Remove") - } - if err := Remove(m); err != nil { - t.Fatalf("Remove idempotent: %v", err) - } -} - -func TestDownloadSuccess(t *testing.T) { - t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) - payload := []byte("vox-test-model") - sum := sha256.Sum256(payload) - checksum := hex.EncodeToString(sum[:]) - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write(payload) - })) - defer srv.Close() - - m := Model{ - ID: "test", - Filename: "ggml-test.bin", - URL: srv.URL, - Checksum: checksum, - } - var updates int - if err := Download(context.Background(), m, func(downloaded, total int64) { - updates++ - if downloaded < 0 { - t.Errorf("downloaded < 0: %d", downloaded) - } - _ = total - }); err != nil { - t.Fatalf("Download: %v", err) - } - if updates == 0 { - t.Fatal("expected at least one progress update") - } - path, _ := Path(m) - got, err := os.ReadFile(path) - if err != nil { - t.Fatalf("ReadFile: %v", err) - } - if string(got) != string(payload) { - t.Fatalf("file contents mismatch: got %q want %q", string(got), string(payload)) - } -} - -func TestDownloadChecksumMismatch(t *testing.T) { - t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("bad")) - })) - defer srv.Close() - - m := Model{ - ID: "test", - Filename: "ggml-test.bin", - URL: srv.URL, - Checksum: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", - } - err := Download(context.Background(), m, nil) - if err == nil { - t.Fatal("expected checksum mismatch") - } - path, _ := Path(m) - if _, statErr := os.Stat(path); !os.IsNotExist(statErr) { - t.Fatalf("expected no final file after mismatch, stat err = %v", statErr) - } - part := path + ".part" - if _, statErr := os.Stat(part); !os.IsNotExist(statErr) { - t.Fatalf("expected no .part file after mismatch, stat err = %v", statErr) - } -} - -func TestModelDirDefault(t *testing.T) { - t.Setenv("WHISPER_MODEL_DIR", "") - home := t.TempDir() - t.Setenv("HOME", home) - dir, err := ModelDir() - if err != nil { - t.Fatalf("ModelDir: %v", err) - } - want := filepath.Join(home, ".local", "share", "whisper-cpp") - if dir != want { - t.Fatalf("ModelDir = %q, want %q", dir, want) - } -} - -func TestDownloadConcurrentSameModelIsSerialized(t *testing.T) { - t.Setenv("WHISPER_MODEL_DIR", t.TempDir()) - payload := []byte("vox-concurrent-model") - sum := sha256.Sum256(payload) - checksum := hex.EncodeToString(sum[:]) - - var requests atomic.Int32 - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - requests.Add(1) - time.Sleep(50 * time.Millisecond) // increase overlap window - w.WriteHeader(http.StatusOK) - _, _ = w.Write(payload) - })) - defer srv.Close() - - m := Model{ - ID: "test-concurrent", - Filename: "ggml-test-concurrent.bin", - URL: srv.URL, - Checksum: checksum, - } - - var wg sync.WaitGroup - errs := make(chan error, 2) - wg.Add(2) - for i := 0; i < 2; i++ { - go func() { - defer wg.Done() - errs <- Download(context.Background(), m, nil) - }() - } - wg.Wait() - close(errs) - - for err := range errs { - if err != nil { - t.Fatalf("concurrent download returned error: %v", err) - } - } - - if got := requests.Load(); got != 1 { - t.Fatalf("download requests = %d, want 1", got) - } - if !IsInstalled(m) { - t.Fatal("model should be installed after concurrent download") - } -}