package model import ( "archive/tar" "bytes" "net/http" "net/http/httptest" "os" "path/filepath" "testing" ) 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 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 len(profile.UpgradeModelIDs) != 1 || profile.UpgradeModelIDs[0] != ParakeetModelID { t.Fatalf("English upgrade models = %v, want [%s]", profile.UpgradeModelIDs, ParakeetModelID) } } 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 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 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 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 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 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 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) } }