package model 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) { profile := DefaultModelProfile() if profile.ID != DefaultModelID { t.Fatalf("default model id = %q, want %q", profile.ID, DefaultModelID) } if profile.BackendKind != BackendSenseVoice { t.Fatalf("backend kind = %q, want %q", profile.BackendKind, BackendSenseVoice) } if _, err := GetModelProfile("missing"); err == nil { t.Fatal("expected unknown model id to fail") } } func TestRegistryReturnsMoonshineEnglish(t *testing.T) { profile, err := GetModelProfile(MoonshineModelID) if err != nil { t.Fatal(err) } if profile.BackendKind != BackendMoonshine { t.Fatalf("backend kind = %q, want %q", profile.BackendKind, BackendMoonshine) } if profile.InstallDirName != "moonshine-en" { t.Fatalf("install dir = %q, want moonshine-en", profile.InstallDirName) } if len(profile.DownloadURLs) == 0 { t.Fatal("expected moonshine download URL") } } func TestRegistryReturnsParakeetEnglish(t *testing.T) { profile, err := GetModelProfile(ParakeetModelID) if err != nil { t.Fatal(err) } if profile.BackendKind != BackendNemoTransducer { t.Fatalf("backend kind = %q, want %q", profile.BackendKind, BackendNemoTransducer) } if profile.InstallDirName != "parakeet-en" { t.Fatalf("install dir = %q, want parakeet-en", profile.InstallDirName) } if profile.Tier != "advanced" { t.Fatalf("tier = %q, want advanced", profile.Tier) } 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 TestRegistryReturnsQwen3ASRChinese(t *testing.T) { profile, err := GetModelProfile(Qwen3ASRModelID) if err != nil { t.Fatal(err) } if profile.BackendKind != BackendQwen3ASR { t.Fatalf("backend kind = %q, want %q", profile.BackendKind, BackendQwen3ASR) } if profile.InstallDirName != "qwen3-asr-0.6b" { t.Fatalf("install dir = %q, want qwen3-asr-0.6b", profile.InstallDirName) } if profile.Tier != "advanced" { t.Fatalf("tier = %q, want advanced", profile.Tier) } if len(profile.DownloadURLs) == 0 { t.Fatal("expected qwen3-asr download URL") } } 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 { t.Fatal(err) } if profile.DefaultModelID != DefaultModelID { t.Fatalf("Chinese default model = %q, want %q", profile.DefaultModelID, DefaultModelID) } if !stringSliceContains(profile.UpgradeModelIDs, Qwen3ASRModelID) || !stringSliceContains(profile.UpgradeModelIDs, XASRZhEn960ModelID) { t.Fatalf("Chinese upgrade models = %v, want qwen3 and x-asr", profile.UpgradeModelIDs) } } func TestEnglishLanguageDefaultsToMoonshine(t *testing.T) { profile, err := GetLanguageProfile(EnglishLanguageID) if err != nil { t.Fatal(err) } if profile.DefaultModelID != MoonshineModelID { t.Fatalf("English default model = %q, want %q", profile.DefaultModelID, MoonshineModelID) } if !stringSliceContains(profile.UpgradeModelIDs, ParakeetModelID) || !stringSliceContains(profile.UpgradeModelIDs, XASRZhEn960ModelID) { t.Fatalf("English upgrade models = %v, want parakeet and x-asr", profile.UpgradeModelIDs) } } 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") writeTestFile(t, filepath.Join(dir, "model.onnx"), "model") result := ValidateModelDir(DefaultModelProfile(), dir) if !result.Valid { t.Fatalf("expected valid model dir, missing=%v problems=%v", result.Missing, result.Problems) } if got := filepath.Base(result.Files["model"]); got != "model.onnx" { t.Fatalf("model file = %q, want model.onnx", got) } } func TestValidateSenseVoicePrefersInt8(t *testing.T) { dir := t.TempDir() writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens") writeTestFile(t, filepath.Join(dir, "model.onnx"), "model") writeTestFile(t, filepath.Join(dir, "model.int8.onnx"), "int8") result := ValidateModelDir(DefaultModelProfile(), dir) if !result.Valid { t.Fatalf("expected valid model 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) } } func TestValidateSenseVoiceRejectsMissingTokens(t *testing.T) { dir := t.TempDir() writeTestFile(t, filepath.Join(dir, "model.int8.onnx"), "model") result := ValidateModelDir(DefaultModelProfile(), dir) if result.Valid { t.Fatal("expected missing tokens to fail validation") } } func TestValidateSenseVoiceRejectsZeroByteFile(t *testing.T) { dir := t.TempDir() writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens") writeTestFile(t, filepath.Join(dir, "model.int8.onnx"), "") 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") } } func TestValidateMoonshineRequiredFiles(t *testing.T) { profile, err := GetModelProfile(MoonshineModelID) if err != nil { t.Fatal(err) } dir := t.TempDir() createValidMoonshine(t, dir) result := ValidateModelDir(profile, dir) if !result.Valid { t.Fatalf("expected valid moonshine dir, missing=%v problems=%v", result.Missing, result.Problems) } for _, role := range []string{"preprocessor", "encoder", "uncached_decoder", "cached_decoder", "tokens"} { if result.Files[role] == "" { t.Fatalf("missing resolved file role %q", role) } } } func TestValidateMoonshineRejectsMissingFile(t *testing.T) { profile, err := GetModelProfile(MoonshineModelID) if err != nil { t.Fatal(err) } dir := t.TempDir() createValidMoonshine(t, dir) if err := os.Remove(filepath.Join(dir, "cached_decode.int8.onnx")); err != nil { t.Fatal(err) } result := ValidateModelDir(profile, dir) if result.Valid { t.Fatal("expected missing cached decoder to fail validation") } } func TestValidateParakeetRequiredFiles(t *testing.T) { profile, err := GetModelProfile(ParakeetModelID) if err != nil { t.Fatal(err) } dir := t.TempDir() createValidParakeet(t, dir) result := ValidateModelDir(profile, dir) if !result.Valid { t.Fatalf("expected valid parakeet 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 TestValidateParakeetRejectsMissingFile(t *testing.T) { profile, err := GetModelProfile(ParakeetModelID) if err != nil { t.Fatal(err) } 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) } 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 { t.Fatal(err) } dir := t.TempDir() createValidQwen3ASR(t, dir) result := ValidateModelDir(profile, dir) if !result.Valid { t.Fatalf("expected valid qwen3-asr dir, missing=%v problems=%v", result.Missing, result.Problems) } for _, role := range []string{"conv_frontend", "encoder", "decoder", "tokenizer", "tokenizer_merges", "tokenizer_vocab"} { if result.Files[role] == "" { t.Fatalf("missing resolved file role %q", role) } } if info, err := os.Stat(result.Files["tokenizer"]); err != nil || !info.IsDir() { t.Fatalf("tokenizer path = %q, want directory, statErr=%v", result.Files["tokenizer"], err) } } func TestValidateQwen3ASRRejectsMissingTokenizerFile(t *testing.T) { profile, err := GetModelProfile(Qwen3ASRModelID) if err != nil { t.Fatal(err) } dir := t.TempDir() createValidQwen3ASR(t, dir) if err := os.Remove(filepath.Join(dir, "tokenizer", "vocab.json")); err != nil { t.Fatal(err) } result := ValidateModelDir(profile, dir) if result.Valid { t.Fatal("expected missing tokenizer vocab to fail validation") } } func TestValidateQwen3ASRRejectsTokenizerFileInsteadOfDirectory(t *testing.T) { profile, err := GetModelProfile(Qwen3ASRModelID) if err != nil { t.Fatal(err) } dir := t.TempDir() createValidQwen3ASR(t, dir) if err := os.RemoveAll(filepath.Join(dir, "tokenizer")); err != nil { t.Fatal(err) } writeTestFile(t, filepath.Join(dir, "tokenizer"), "not a directory") 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") } } func TestResolveModelUsesLegacyDirWithoutState(t *testing.T) { root := t.TempDir() createValidSenseVoice(t, filepath.Join(root, "sensevoice")) resolved, err := ResolveModelInRoot(DefaultModelID, root) if err != nil { t.Fatal(err) } if !resolved.IsUsable() { t.Fatalf("expected usable model, status=%s missing=%v problems=%v", resolved.Status, resolved.Missing, resolved.Problems) } if resolved.SourceDirKind != "legacy" { t.Fatalf("source dir kind = %q, want legacy", resolved.SourceDirKind) } } func TestResolveModelPrefersNewDirWithoutState(t *testing.T) { root := t.TempDir() createValidSenseVoice(t, filepath.Join(root, "sensevoice")) createValidSenseVoice(t, filepath.Join(root, "sensevoice-zh")) resolved, err := ResolveModelInRoot(DefaultModelID, root) if err != nil { t.Fatal(err) } if resolved.SourceDirKind != "new" { t.Fatalf("source dir kind = %q, want new", resolved.SourceDirKind) } if got := filepath.Base(resolved.RootDir); got != "sensevoice-zh" { t.Fatalf("root dir = %q, want sensevoice-zh", got) } } func TestResolveModelInvalidStateFallsBackToLegacy(t *testing.T) { root := t.TempDir() writeTestFile(t, filepath.Join(root, "state.json"), "{invalid json") createValidSenseVoice(t, filepath.Join(root, "sensevoice")) resolved, err := ResolveModelInRoot(DefaultModelID, root) if err != nil { t.Fatal(err) } if resolved.Status != ModelStateInvalidFallbackFound { t.Fatalf("status = %s, want %s", resolved.Status, ModelStateInvalidFallbackFound) } if resolved.SourceDirKind != "legacy" { t.Fatalf("source dir kind = %q, want legacy", resolved.SourceDirKind) } } func TestResolveModelStatePointsToBrokenDirFallsBackToLegacy(t *testing.T) { root := t.TempDir() broken := filepath.Join(root, "broken") if err := os.MkdirAll(broken, 0755); err != nil { t.Fatal(err) } writeTestFile(t, filepath.Join(broken, "tokens.txt"), "tokens") createValidSenseVoice(t, filepath.Join(root, "sensevoice")) state := NewInstallState() state.InstalledModels[DefaultModelID] = InstalledModelState{ Path: "broken", SourceDirKind: "new", } if err := SaveInstallStateToRoot(root, state); err != nil { t.Fatal(err) } resolved, err := ResolveModelInRoot(DefaultModelID, root) if err != nil { t.Fatal(err) } if resolved.Status != ModelStateInvalidFallbackFound { t.Fatalf("status = %s, want %s", resolved.Status, ModelStateInvalidFallbackFound) } if resolved.SourceDirKind != "legacy" { t.Fatalf("source dir kind = %q, want legacy", resolved.SourceDirKind) } } func TestResolveModelRejectsStatePathOutsideModelsRoot(t *testing.T) { root := t.TempDir() outside := t.TempDir() createValidSenseVoice(t, filepath.Join(outside, "outside-model")) createValidSenseVoice(t, filepath.Join(root, "sensevoice-zh")) state := NewInstallState() state.InstalledModels[DefaultModelID] = InstalledModelState{ Path: filepath.Join(outside, "outside-model"), SourceDirKind: "new", } if err := SaveInstallStateToRoot(root, state); err != nil { t.Fatal(err) } resolved, err := ResolveModelInRoot(DefaultModelID, root) if err != nil { t.Fatal(err) } if resolved.Status != ModelStateInvalidFallbackFound { t.Fatalf("status = %s, want %s", resolved.Status, ModelStateInvalidFallbackFound) } if got := filepath.Base(resolved.RootDir); got != "sensevoice-zh" { t.Fatalf("root dir = %q, want sensevoice-zh", got) } } func TestResolveModelRejectsStatePathEscapingModelsRoot(t *testing.T) { root := t.TempDir() createValidSenseVoice(t, filepath.Join(root, "sensevoice")) state := NewInstallState() state.InstalledModels[DefaultModelID] = InstalledModelState{ Path: "../outside-model", SourceDirKind: "new", } if err := SaveInstallStateToRoot(root, state); err != nil { t.Fatal(err) } resolved, err := ResolveModelInRoot(DefaultModelID, root) if err != nil { t.Fatal(err) } if resolved.Status != ModelStateInvalidFallbackFound { t.Fatalf("status = %s, want %s", resolved.Status, ModelStateInvalidFallbackFound) } if resolved.SourceDirKind != "legacy" { t.Fatalf("source dir kind = %q, want legacy", resolved.SourceDirKind) } } func TestResolveModelIncompleteDir(t *testing.T) { root := t.TempDir() dir := filepath.Join(root, "sensevoice-zh") if err := os.MkdirAll(dir, 0755); err != nil { t.Fatal(err) } writeTestFile(t, filepath.Join(dir, "model.int8.onnx"), "model") resolved, err := ResolveModelInRoot(DefaultModelID, root) if err != nil { t.Fatal(err) } 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") } } func TestInstallStateRoundTrip(t *testing.T) { root := t.TempDir() if err := UpdateInstalledModelState(root, DefaultModelProfile(), "sensevoice-zh", "new"); err != nil { t.Fatal(err) } state, err := LoadInstallStateFromRoot(root) if err != nil { t.Fatal(err) } if state.SelectedModelID != DefaultModelID { t.Fatalf("selected model = %q, want %q", state.SelectedModelID, DefaultModelID) } installed := state.InstalledModels[DefaultModelID] if installed.Path != "sensevoice-zh" { t.Fatalf("installed path = %q, want sensevoice-zh", installed.Path) } } func TestDownloadFailureDoesNotRemoveLegacyModel(t *testing.T) { root := t.TempDir() legacyDir := filepath.Join(root, "sensevoice") createValidSenseVoice(t, legacyDir) err := DownloadProfile(DefaultModelProfile(), []string{"://bad-url"}, root, nil) if err == nil { t.Fatal("expected download failure") } result := ValidateModelDir(DefaultModelProfile(), legacyDir) if !result.Valid { t.Fatalf("legacy model was damaged, missing=%v problems=%v", result.Missing, result.Problems) } } func TestDownloadQwen3ASRFailureDoesNotRemoveSenseVoice(t *testing.T) { root := t.TempDir() senseVoiceDir := filepath.Join(root, "sensevoice-zh") createValidSenseVoice(t, senseVoiceDir) profile, err := GetModelProfile(Qwen3ASRModelID) if err != nil { t.Fatal(err) } err = DownloadProfile(profile, []string{"://bad-url"}, root, nil) if err == nil { t.Fatal("expected download failure") } result := ValidateModelDir(DefaultModelProfile(), senseVoiceDir) if !result.Valid { t.Fatalf("sensevoice model was damaged, missing=%v problems=%v", result.Missing, result.Problems) } } func TestDownloadProfileRemovesStalePartialDownload(t *testing.T) { root := t.TempDir() stalePath := filepath.Join(root, ".downloads", DefaultModelID, "old-run", "model_package") writeTestFile(t, stalePath, "partial") err := DownloadProfile(DefaultModelProfile(), []string{"://bad-url"}, root, nil) if err == nil { t.Fatal("expected download failure") } if _, statErr := os.Stat(filepath.Dir(stalePath)); !os.IsNotExist(statErr) { t.Fatalf("stale download run still exists or stat failed: %v", statErr) } } func TestDownloadProfileFallsBackWhenPrimaryArchiveInvalid(t *testing.T) { root := t.TempDir() validArchive := 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 "/bad": w.WriteHeader(http.StatusOK) w.Write([]byte("not a tar archive")) case "/fallback": w.WriteHeader(http.StatusOK) w.Write(validArchive) default: http.NotFound(w, r) } })) defer server.Close() err := DownloadProfile(DefaultModelProfile(), []string{server.URL + "/bad", server.URL + "/fallback"}, root, nil) if err != nil { t.Fatal(err) } resolved, err := ResolveModelInRoot(DefaultModelID, root) if err != nil { t.Fatal(err) } if !resolved.IsUsable() { t.Fatalf("expected usable downloaded model, status=%s missing=%v problems=%v", resolved.Status, resolved.Missing, resolved.Problems) } if got := filepath.Base(resolved.RootDir); got != "sensevoice-zh" { t.Fatalf("root dir = %q, want sensevoice-zh", got) } } func TestDownloadProfileReplacesExistingNewDir(t *testing.T) { root := t.TempDir() finalDir := filepath.Join(root, "sensevoice-zh") createValidSenseVoice(t, finalDir) writeTestFile(t, filepath.Join(finalDir, "model.int8.onnx"), "old-model") validArchive := tarArchive(t, map[string]string{ "sensevoice/tokens.txt": "tokens", "sensevoice/model.int8.onnx": "new-model", }) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) w.Write(validArchive) })) defer server.Close() if err := DownloadProfile(DefaultModelProfile(), []string{server.URL}, root, nil); err != nil { t.Fatal(err) } data, err := os.ReadFile(filepath.Join(finalDir, "model.int8.onnx")) if err != nil { t.Fatal(err) } if string(data) != "new-model" { t.Fatalf("model contents = %q, want new-model", string(data)) } } func TestDownloadProfileBadArchiveDoesNotWriteInstalledState(t *testing.T) { root := t.TempDir() server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) w.Write([]byte("not a tar archive")) })) defer server.Close() if err := DownloadProfile(DefaultModelProfile(), []string{server.URL}, root, nil); err == nil { t.Fatal("expected bad archive to fail") } state, err := LoadInstallStateFromRoot(root) if err != nil { t.Fatal(err) } if _, ok := state.InstalledModels[DefaultModelID]; ok { t.Fatal("bad archive should not write installed model state") } } 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 { t.Fatal(err) } writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens") writeTestFile(t, filepath.Join(dir, "model.int8.onnx"), "model") } func createValidMoonshine(t *testing.T, dir string) { t.Helper() if err := os.MkdirAll(dir, 0755); err != nil { t.Fatal(err) } writeTestFile(t, filepath.Join(dir, "preprocess.onnx"), "preprocessor") writeTestFile(t, filepath.Join(dir, "encode.int8.onnx"), "encoder") writeTestFile(t, filepath.Join(dir, "uncached_decode.int8.onnx"), "uncached") writeTestFile(t, filepath.Join(dir, "cached_decode.int8.onnx"), "cached") writeTestFile(t, filepath.Join(dir, "tokens.txt"), "tokens") } func createValidParakeet(t *testing.T, dir string) { t.Helper() if err := os.MkdirAll(dir, 0755); err != nil { t.Fatal(err) } writeTestFile(t, filepath.Join(dir, "encoder.int8.onnx"), "encoder") writeTestFile(t, filepath.Join(dir, "decoder.int8.onnx"), "decoder") writeTestFile(t, filepath.Join(dir, "joiner.int8.onnx"), "joiner") 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 { t.Fatal(err) } writeTestFile(t, filepath.Join(dir, "conv_frontend.onnx"), "conv") writeTestFile(t, filepath.Join(dir, "encoder.int8.onnx"), "encoder") writeTestFile(t, filepath.Join(dir, "decoder.int8.onnx"), "decoder") writeTestFile(t, filepath.Join(dir, "tokenizer", "merges.txt"), "merges") 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 tw := tar.NewWriter(&buf) for name, content := range files { data := []byte(content) header := &tar.Header{ Name: name, Mode: 0644, Size: int64(len(data)), } if err := tw.WriteHeader(header); err != nil { t.Fatal(err) } if _, err := tw.Write(data); err != nil { t.Fatal(err) } } if err := tw.Close(); err != nil { t.Fatal(err) } return buf.Bytes() } func writeTestFile(t *testing.T, path, content string) { t.Helper() if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { t.Fatal(err) } if err := os.WriteFile(path, []byte(content), 0644); err != nil { t.Fatal(err) } }