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 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)
|
}
|
}
|
|
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 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 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 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 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 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
|
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)
|
}
|
}
|