| | |
| | | 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) { |
| | |
| | | } |
| | | if len(profile.DownloadURLs) == 0 { |
| | | t.Fatal("expected parakeet download URL") |
| | | } |
| | | } |
| | | |
| | | 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 TestRegistryReturnsCantoneseSenseVoiceYueDefault(t *testing.T) { |
| | | languageProfile, err := GetLanguageProfile(CantoneseLanguageID) |
| | | if err != nil { |
| | | t.Fatal(err) |
| | | } |
| | | if languageProfile.NativeName != "粤语" { |
| | | t.Fatalf("Cantonese native name = %q, want 粤语", languageProfile.NativeName) |
| | | } |
| | | if languageProfile.DefaultModelID != SenseVoiceYueModelID { |
| | | t.Fatalf("Cantonese default model = %q, want %q", languageProfile.DefaultModelID, SenseVoiceYueModelID) |
| | | } |
| | | if len(languageProfile.UpgradeModelIDs) != 0 { |
| | | t.Fatalf("Cantonese upgrade models = %v, want none", languageProfile.UpgradeModelIDs) |
| | | } |
| | | |
| | | profile, err := GetModelProfile(SenseVoiceYueModelID) |
| | | if err != nil { |
| | | t.Fatal(err) |
| | | } |
| | | if profile.BackendKind != BackendSenseVoice { |
| | | t.Fatalf("backend kind = %q, want %q", profile.BackendKind, BackendSenseVoice) |
| | | } |
| | | if profile.LanguageParam != "yue" { |
| | | t.Fatalf("language param = %q, want yue", profile.LanguageParam) |
| | | } |
| | | if profile.InstallDirName != "sensevoice-yue-2025-09-09" { |
| | | t.Fatalf("install dir = %q, want sensevoice-yue-2025-09-09", profile.InstallDirName) |
| | | } |
| | | if !stringSliceContains(profile.SupportedLanguageIDs, CantoneseLanguageID) { |
| | | t.Fatalf("supported languages = %v, want %s", profile.SupportedLanguageIDs, CantoneseLanguageID) |
| | | } |
| | | if len(profile.DownloadURLs) == 0 { |
| | | t.Fatal("expected SenseVoice Yue download URL") |
| | | } |
| | | } |
| | | |
| | |
| | | } |
| | | } |
| | | |
| | | func TestLanguageProfileOrderPlacesCantoneseAfterChinese(t *testing.T) { |
| | | profiles := ListLanguageProfiles() |
| | | if len(profiles) < 3 { |
| | | t.Fatalf("language profiles len = %d, want at least 3", len(profiles)) |
| | | } |
| | | if profiles[0].ID != DefaultLanguageID { |
| | | t.Fatalf("first language = %q, want %q", profiles[0].ID, DefaultLanguageID) |
| | | } |
| | | if profiles[1].ID != CantoneseLanguageID { |
| | | t.Fatalf("second language = %q, want %q", profiles[1].ID, CantoneseLanguageID) |
| | | } |
| | | if profiles[2].ID != EnglishLanguageID { |
| | | t.Fatalf("third language = %q, want %q", profiles[2].ID, EnglishLanguageID) |
| | | } |
| | | } |
| | | |
| | | func TestNormalizeLanguageIDAcceptsCantoneseAlias(t *testing.T) { |
| | | if got := NormalizeLanguageID("yue"); got != CantoneseLanguageID { |
| | | t.Fatalf("NormalizeLanguageID(yue) = %q, want %q", got, CantoneseLanguageID) |
| | | } |
| | | if got := NormalizeLanguageID(CantoneseLanguageID); got != CantoneseLanguageID { |
| | | t.Fatalf("NormalizeLanguageID(%s) = %q", CantoneseLanguageID, got) |
| | | } |
| | | } |
| | | |
| | | func TestValidateSenseVoiceAnyOfAndAllOf(t *testing.T) { |
| | | dir := t.TempDir() |
| | | writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens") |
| | |
| | | result := ValidateModelDir(DefaultModelProfile(), dir) |
| | | if result.Valid { |
| | | t.Fatal("expected zero-byte model file to fail validation") |
| | | } |
| | | } |
| | | |
| | | func TestValidateSenseVoiceYueRequiredFiles(t *testing.T) { |
| | | profile, err := GetModelProfile(SenseVoiceYueModelID) |
| | | if err != nil { |
| | | t.Fatal(err) |
| | | } |
| | | dir := t.TempDir() |
| | | createValidSenseVoice(t, dir) |
| | | |
| | | result := ValidateModelDir(profile, dir) |
| | | if !result.Valid { |
| | | t.Fatalf("expected valid SenseVoice Yue dir, missing=%v problems=%v", result.Missing, result.Problems) |
| | | } |
| | | if got := filepath.Base(result.Files["model"]); got != "model.int8.onnx" { |
| | | t.Fatalf("model file = %q, want model.int8.onnx", got) |
| | | } |
| | | if result.Files["tokens"] == "" { |
| | | t.Fatal("missing resolved tokens file") |
| | | } |
| | | } |
| | | |
| | | func TestValidateSenseVoiceYueRejectsFp32OnlyModel(t *testing.T) { |
| | | profile, err := GetModelProfile(SenseVoiceYueModelID) |
| | | if err != nil { |
| | | t.Fatal(err) |
| | | } |
| | | dir := t.TempDir() |
| | | writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens") |
| | | writeTestFile(t, filepath.Join(dir, "model.onnx"), "fp32") |
| | | |
| | | result := ValidateModelDir(profile, dir) |
| | | if result.Valid { |
| | | t.Fatal("expected SenseVoice Yue to require model.int8.onnx") |
| | | } |
| | | } |
| | | |
| | |
| | | dir := t.TempDir() |
| | | createValidParakeet(t, dir) |
| | | if err := os.Remove(filepath.Join(dir, "joiner.int8.onnx")); err != nil { |
| | | t.Fatal(err) |
| | | } |
| | | |
| | | result := ValidateModelDir(profile, dir) |
| | | if result.Valid { |
| | | t.Fatal("expected missing joiner to fail validation") |
| | | } |
| | | } |
| | | |
| | | 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) |
| | | } |
| | | |
| | |
| | | } |
| | | } |
| | | |
| | | 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") |
| | | } |
| | | } |
| | | |
| | | func TestInstallStateRoundTrip(t *testing.T) { |
| | | root := t.TempDir() |
| | | if err := UpdateInstalledModelState(root, DefaultModelProfile(), "sensevoice-zh", "new"); err != nil { |
| | |
| | | } |
| | | } |
| | | |
| | | 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") |
| | | |
| | |
| | | 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 { |