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