Ariver
2026-07-02 5df683eba7c103010a7c110d97a5fd808b39fe46
privatevoice.src/internal/model/model_test.go
@@ -3,11 +3,19 @@
import (
   "archive/tar"
   "bytes"
   "context"
   "errors"
   "fmt"
   "net/http"
   "net/http/httptest"
   "os"
   "path/filepath"
   "strconv"
   "strings"
   "sync"
   "sync/atomic"
   "testing"
   "time"
)
func TestRegistryReturnsDefaultSenseVoice(t *testing.T) {
@@ -58,6 +66,54 @@
   }
}
func TestRegistryReturnsParakeetV3ForEuropeanLanguages(t *testing.T) {
   profile, err := GetModelProfile(ParakeetV3ModelID)
   if err != nil {
      t.Fatal(err)
   }
   if profile.BackendKind != BackendNemoTransducer {
      t.Fatalf("backend kind = %q, want %q", profile.BackendKind, BackendNemoTransducer)
   }
   for _, languageID := range []string{FrenchLanguageID, GermanLanguageID, SpanishLanguageID, ItalianLanguageID, PortugueseLanguageID} {
      if !stringSliceContains(profile.SupportedLanguageIDs, languageID) {
         t.Fatalf("parakeet v3 languages = %v, want %s", profile.SupportedLanguageIDs, languageID)
      }
      languageProfile, err := GetLanguageProfile(languageID)
      if err != nil {
         t.Fatal(err)
      }
      if languageProfile.DefaultModelID != ParakeetV3ModelID {
         t.Fatalf("%s default model = %q, want %q", languageID, languageProfile.DefaultModelID, ParakeetV3ModelID)
      }
   }
}
func TestRegistryReturnsJapaneseAndKoreanDefaults(t *testing.T) {
   tests := []struct {
      languageID string
      modelID    string
   }{
      {JapaneseLanguageID, JapaneseModelID},
      {KoreanLanguageID, KoreanModelID},
   }
   for _, tt := range tests {
      languageProfile, err := GetLanguageProfile(tt.languageID)
      if err != nil {
         t.Fatal(err)
      }
      if languageProfile.DefaultModelID != tt.modelID {
         t.Fatalf("%s default model = %q, want %q", tt.languageID, languageProfile.DefaultModelID, tt.modelID)
      }
      modelProfile, err := GetModelProfile(tt.modelID)
      if err != nil {
         t.Fatal(err)
      }
      if modelProfile.BackendKind != BackendTransducer {
         t.Fatalf("%s backend = %q, want %q", tt.modelID, modelProfile.BackendKind, BackendTransducer)
      }
   }
}
func TestRegistryReturnsQwen3ASRChinese(t *testing.T) {
   profile, err := GetModelProfile(Qwen3ASRModelID)
   if err != nil {
@@ -77,6 +133,25 @@
   }
}
func TestRegistryReturnsXASRZhEn960(t *testing.T) {
   profile, err := GetModelProfile(XASRZhEn960ModelID)
   if err != nil {
      t.Fatal(err)
   }
   if profile.BackendKind != BackendXASRStreaming {
      t.Fatalf("backend kind = %q, want %q", profile.BackendKind, BackendXASRStreaming)
   }
   if profile.InstallDirName != "x-asr-zh-en-960ms" {
      t.Fatalf("install dir = %q, want x-asr-zh-en-960ms", profile.InstallDirName)
   }
   if profile.Tier != "experimental" {
      t.Fatalf("tier = %q, want experimental", profile.Tier)
   }
   if len(profile.DownloadURLs) != 0 {
      t.Fatal("expected x-asr experiment to use manual/scripted install, not in-app tar download")
   }
}
func TestChineseLanguageOffersQwen3ASRUpgrade(t *testing.T) {
   profile, err := GetLanguageProfile(DefaultLanguageID)
   if err != nil {
@@ -85,8 +160,8 @@
   if profile.DefaultModelID != DefaultModelID {
      t.Fatalf("Chinese default model = %q, want %q", profile.DefaultModelID, DefaultModelID)
   }
   if len(profile.UpgradeModelIDs) != 1 || profile.UpgradeModelIDs[0] != Qwen3ASRModelID {
      t.Fatalf("Chinese upgrade models = %v, want [%s]", profile.UpgradeModelIDs, Qwen3ASRModelID)
   if !stringSliceContains(profile.UpgradeModelIDs, Qwen3ASRModelID) || !stringSliceContains(profile.UpgradeModelIDs, XASRZhEn960ModelID) {
      t.Fatalf("Chinese upgrade models = %v, want qwen3 and x-asr", profile.UpgradeModelIDs)
   }
}
@@ -98,8 +173,8 @@
   if profile.DefaultModelID != MoonshineModelID {
      t.Fatalf("English default model = %q, want %q", profile.DefaultModelID, MoonshineModelID)
   }
   if len(profile.UpgradeModelIDs) != 1 || profile.UpgradeModelIDs[0] != ParakeetModelID {
      t.Fatalf("English upgrade models = %v, want [%s]", profile.UpgradeModelIDs, ParakeetModelID)
   if !stringSliceContains(profile.UpgradeModelIDs, ParakeetModelID) || !stringSliceContains(profile.UpgradeModelIDs, XASRZhEn960ModelID) {
      t.Fatalf("English upgrade models = %v, want parakeet and x-asr", profile.UpgradeModelIDs)
   }
}
@@ -225,6 +300,61 @@
   }
}
func TestValidateJapaneseZipformerRequiredFiles(t *testing.T) {
   profile, err := GetModelProfile(JapaneseModelID)
   if err != nil {
      t.Fatal(err)
   }
   dir := t.TempDir()
   createValidZipformerSpecialized(t, dir)
   result := ValidateModelDir(profile, dir)
   if !result.Valid {
      t.Fatalf("expected valid japanese zipformer dir, missing=%v problems=%v", result.Missing, result.Problems)
   }
   for _, role := range []string{"encoder", "decoder", "joiner", "tokens"} {
      if result.Files[role] == "" {
         t.Fatalf("missing resolved file role %q", role)
      }
   }
}
func TestValidateKoreanZipformerRequiredFiles(t *testing.T) {
   profile, err := GetModelProfile(KoreanModelID)
   if err != nil {
      t.Fatal(err)
   }
   dir := t.TempDir()
   createValidZipformerSpecialized(t, dir)
   result := ValidateModelDir(profile, dir)
   if !result.Valid {
      t.Fatalf("expected valid korean zipformer dir, missing=%v problems=%v", result.Missing, result.Problems)
   }
   for _, role := range []string{"encoder", "decoder", "joiner", "tokens"} {
      if result.Files[role] == "" {
         t.Fatalf("missing resolved file role %q", role)
      }
   }
}
func TestValidateZipformerSpecializedRejectsMissingJoiner(t *testing.T) {
   profile, err := GetModelProfile(JapaneseModelID)
   if err != nil {
      t.Fatal(err)
   }
   dir := t.TempDir()
   createValidZipformerSpecialized(t, dir)
   if err := os.Remove(filepath.Join(dir, "joiner-epoch-99-avg-1.int8.onnx")); err != nil {
      t.Fatal(err)
   }
   result := ValidateModelDir(profile, dir)
   if result.Valid {
      t.Fatal("expected missing joiner to fail validation")
   }
}
func TestValidateQwen3ASRRequiredFiles(t *testing.T) {
   profile, err := GetModelProfile(Qwen3ASRModelID)
   if err != nil {
@@ -279,6 +409,42 @@
   result := ValidateModelDir(profile, dir)
   if result.Valid {
      t.Fatal("expected tokenizer file to fail directory validation")
   }
}
func TestValidateXASRZhEn960RequiredFiles(t *testing.T) {
   profile, err := GetModelProfile(XASRZhEn960ModelID)
   if err != nil {
      t.Fatal(err)
   }
   dir := t.TempDir()
   createValidXASRZhEn960(t, dir)
   result := ValidateModelDir(profile, dir)
   if !result.Valid {
      t.Fatalf("expected valid x-asr dir, missing=%v problems=%v", result.Missing, result.Problems)
   }
   for _, role := range []string{"encoder", "decoder", "joiner", "tokens"} {
      if result.Files[role] == "" {
         t.Fatalf("missing resolved file role %q", role)
      }
   }
}
func TestValidateXASRZhEn960RejectsMissingJoiner(t *testing.T) {
   profile, err := GetModelProfile(XASRZhEn960ModelID)
   if err != nil {
      t.Fatal(err)
   }
   dir := t.TempDir()
   createValidXASRZhEn960(t, dir)
   if err := os.Remove(filepath.Join(dir, "joiner-960ms.onnx")); err != nil {
      t.Fatal(err)
   }
   result := ValidateModelDir(profile, dir)
   if result.Valid {
      t.Fatal("expected missing joiner to fail validation")
   }
}
@@ -428,6 +594,18 @@
   }
   if resolved.Status != ModelInstalledIncomplete {
      t.Fatalf("status = %s, want %s", resolved.Status, ModelInstalledIncomplete)
   }
}
func TestHasAnyUsableModelInRoot(t *testing.T) {
   root := t.TempDir()
   if HasAnyUsableModelInRoot(root) {
      t.Fatal("expected no usable model in empty root")
   }
   createValidSenseVoice(t, filepath.Join(root, "sensevoice-zh"))
   if !HasAnyUsableModelInRoot(root) {
      t.Fatal("expected installed sensevoice to count as a usable model")
   }
}
@@ -588,6 +766,143 @@
   }
}
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")
   server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
      rangeHeader := r.Header.Get("Range")
      start := 0
      if strings.HasPrefix(rangeHeader, "bytes=") && strings.HasSuffix(rangeHeader, "-") {
         parsed, err := strconv.Atoi(strings.TrimSuffix(strings.TrimPrefix(rangeHeader, "bytes="), "-"))
         if err == nil {
            start = parsed
         }
      }
      if start > len(payload) {
         w.WriteHeader(http.StatusRequestedRangeNotSatisfiable)
         return
      }
      if start > 0 {
         w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, len(payload)-1, len(payload)))
         w.Header().Set("Content-Length", strconv.Itoa(len(payload)-start))
         w.WriteHeader(http.StatusPartialContent)
         w.Write(payload[start:])
         return
      }
      w.Header().Set("Content-Length", strconv.Itoa(len(payload)))
      w.WriteHeader(http.StatusOK)
      w.Write(payload)
   }))
   defer server.Close()
   dest := filepath.Join(t.TempDir(), "model_package")
   if err := os.WriteFile(dest, payload[:10], 0644); err != nil {
      t.Fatal(err)
   }
   if err := downloadFile(server.URL, dest, nil); err != nil {
      t.Fatal(err)
   }
   data, err := os.ReadFile(dest)
   if err != nil {
      t.Fatal(err)
   }
   if !bytes.Equal(data, payload) {
      t.Fatalf("downloaded payload = %q, want %q", string(data), string(payload))
   }
}
func createValidSenseVoice(t *testing.T, dir string) {
   t.Helper()
   if err := os.MkdirAll(dir, 0755); err != nil {
@@ -620,6 +935,17 @@
   writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens")
}
func createValidZipformerSpecialized(t *testing.T, dir string) {
   t.Helper()
   if err := os.MkdirAll(dir, 0755); err != nil {
      t.Fatal(err)
   }
   writeTestFile(t, filepath.Join(dir, "encoder-epoch-99-avg-1.int8.onnx"), "encoder")
   writeTestFile(t, filepath.Join(dir, "decoder-epoch-99-avg-1.onnx"), "decoder")
   writeTestFile(t, filepath.Join(dir, "joiner-epoch-99-avg-1.int8.onnx"), "joiner")
   writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens")
}
func createValidQwen3ASR(t *testing.T, dir string) {
   t.Helper()
   if err := os.MkdirAll(filepath.Join(dir, "tokenizer"), 0755); err != nil {
@@ -632,6 +958,26 @@
   writeTestFile(t, filepath.Join(dir, "tokenizer", "vocab.json"), "vocab")
}
func createValidXASRZhEn960(t *testing.T, dir string) {
   t.Helper()
   if err := os.MkdirAll(dir, 0755); err != nil {
      t.Fatal(err)
   }
   writeTestFile(t, filepath.Join(dir, "encoder-960ms.onnx"), "encoder")
   writeTestFile(t, filepath.Join(dir, "decoder-960ms.onnx"), "decoder")
   writeTestFile(t, filepath.Join(dir, "joiner-960ms.onnx"), "joiner")
   writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens")
}
func stringSliceContains(values []string, target string) bool {
   for _, value := range values {
      if value == target {
         return true
      }
   }
   return false
}
func tarArchive(t *testing.T, files map[string]string) []byte {
   t.Helper()
   var buf bytes.Buffer