package model
|
|
import (
|
"archive/tar"
|
"bytes"
|
"fmt"
|
"net/http"
|
"net/http/httptest"
|
"os"
|
"path/filepath"
|
"strconv"
|
"strings"
|
"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 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 {
|
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 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 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 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)
|
}
|
}
|