| | |
| | | import ( |
| | | "archive/tar" |
| | | "bytes" |
| | | "fmt" |
| | | "net/http" |
| | | "net/http/httptest" |
| | | "os" |
| | | "path/filepath" |
| | | "strconv" |
| | | "strings" |
| | | "testing" |
| | | ) |
| | | |
| | |
| | | } |
| | | 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 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 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 len(profile.UpgradeModelIDs) != 1 || profile.UpgradeModelIDs[0] != Qwen3ASRModelID { |
| | | t.Fatalf("Chinese upgrade models = %v, want [%s]", profile.UpgradeModelIDs, Qwen3ASRModelID) |
| | | } |
| | | } |
| | | |
| | | 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) |
| | | } |
| | | } |
| | | |
| | |
| | | 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 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 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{ |
| | |
| | | } |
| | | } |
| | | |
| | | 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 { |
| | |
| | | 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 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 tarArchive(t *testing.T, files map[string]string) []byte { |
| | | t.Helper() |
| | | var buf bytes.Buffer |