Scan model directory for installed model status

This commit is contained in:
tian 2026-05-05 12:12:49 +08:00
parent 452d344b34
commit 85f3ca54a3
4 changed files with 104 additions and 19 deletions

View File

@ -14,17 +14,13 @@ import (
func TestHandleModelsStatusReturnsInstalledModels(t *testing.T) {
dir := t.TempDir()
store := modelstore.New(dir, 8)
itemDir := store.FilesDir()
if err := os.MkdirAll(itemDir, 0o755); err != nil {
if err := os.MkdirAll(dir, 0o755); err != nil {
t.Fatalf("MkdirAll: %v", err)
}
modelPath := filepath.Join(itemDir, "face_det_scrfd_500m_640_rk3588__abc.rknn")
modelPath := filepath.Join(dir, "face_det_scrfd_500m_640_rk3588.rknn")
if err := os.WriteFile(modelPath, []byte("abc"), 0o644); err != nil {
t.Fatalf("WriteFile: %v", err)
}
if err := os.WriteFile(store.ManifestPath(), []byte(`{"items":[{"name":"face_det_scrfd_500m_640_rk3588","sha256":"abc","path":"`+filepath.ToSlash(modelPath)+`","size":3,"mtime_ms":1}]}`), 0o644); err != nil {
t.Fatalf("WriteFile manifest: %v", err)
}
srv := &Server{store: store}
req := httptest.NewRequest(http.MethodGet, "/v1/models/status", nil)

View File

@ -9,6 +9,7 @@ import (
"io"
"os"
"path/filepath"
"sort"
"strings"
"rk3588sys/agent/internal/files"
@ -146,25 +147,80 @@ func (s *Store) List() (Manifest, error) {
}
func (s *Store) ListInstalledModels() ([]InstalledModel, error) {
manifest, err := s.List()
paths, err := s.installedModelPaths()
if err != nil {
return nil, err
}
items := make([]InstalledModel, 0, len(manifest.Items))
for _, item := range manifest.Items {
fileName := filepath.Base(filepath.FromSlash(item.Path))
items := make([]InstalledModel, 0, len(paths))
for _, path := range paths {
stat, err := os.Stat(path)
if err != nil {
return nil, fmt.Errorf("stat model %q: %w", path, err)
}
sha, err := fileSHA256(path)
if err != nil {
return nil, fmt.Errorf("hash model %q: %w", path, err)
}
fileName := filepath.Base(path)
items = append(items, InstalledModel{
Name: item.Name,
Name: modelNameFromFileName(fileName),
FileName: fileName,
Sha256: item.Sha256,
Path: item.Path,
Size: item.Size,
MtimeMS: item.MtimeMS,
Sha256: sha,
Path: filepath.ToSlash(path),
Size: stat.Size(),
MtimeMS: stat.ModTime().UnixMilli(),
})
}
sort.Slice(items, func(i, j int) bool {
return items[i].FileName < items[j].FileName
})
return items, nil
}
func (s *Store) installedModelPaths() ([]string, error) {
patterns := []string{
filepath.Join(s.ModelsDir, "*.rknn"),
filepath.Join(s.FilesDir(), "*.rknn"),
}
seen := map[string]struct{}{}
paths := make([]string, 0)
for _, pattern := range patterns {
matches, err := filepath.Glob(pattern)
if err != nil {
return nil, fmt.Errorf("glob models %q: %w", pattern, err)
}
for _, match := range matches {
if _, ok := seen[match]; ok {
continue
}
seen[match] = struct{}{}
paths = append(paths, match)
}
}
return paths, nil
}
func fileSHA256(path string) (string, error) {
f, err := os.Open(path)
if err != nil {
return "", err
}
defer f.Close()
h := sha256.New()
if _, err := io.Copy(h, f); err != nil {
return "", err
}
return hex.EncodeToString(h.Sum(nil)), nil
}
func modelNameFromFileName(fileName string) string {
base := strings.TrimSuffix(fileName, filepath.Ext(fileName))
if name, _, ok := strings.Cut(base, "__"); ok {
return name
}
return base
}
func (s *Store) upsertManifest(item Item) error {
m, err := s.List()
if err != nil {

View File

@ -1,12 +1,48 @@
package modelstore
import (
"crypto/sha256"
"encoding/hex"
"os"
"path/filepath"
"testing"
)
func TestListInstalledModelsUsesManifestEntries(t *testing.T) {
func TestListInstalledModelsScansModelsDirWithoutManifest(t *testing.T) {
dir := t.TempDir()
store := New(dir, 8)
if err := os.MkdirAll(store.ModelsDir, 0o755); err != nil {
t.Fatalf("MkdirAll: %v", err)
}
content := []byte("abc")
modelPath := filepath.Join(store.ModelsDir, "face_det_scrfd_500m_640_rk3588.rknn")
if err := os.WriteFile(modelPath, content, 0o644); err != nil {
t.Fatalf("WriteFile: %v", err)
}
if err := os.WriteFile(filepath.Join(store.ModelsDir, "readme.txt"), []byte("ignore"), 0o644); err != nil {
t.Fatalf("WriteFile readme: %v", err)
}
items, err := store.ListInstalledModels()
if err != nil {
t.Fatalf("ListInstalledModels: %v", err)
}
if len(items) != 1 {
t.Fatalf("unexpected installed models: %#v", items)
}
expectedSHA := sha256.Sum256(content)
if items[0].Name != "face_det_scrfd_500m_640_rk3588" {
t.Fatalf("unexpected model name: %#v", items[0])
}
if items[0].FileName != "face_det_scrfd_500m_640_rk3588.rknn" {
t.Fatalf("unexpected file name: %#v", items[0])
}
if items[0].Sha256 != hex.EncodeToString(expectedSHA[:]) {
t.Fatalf("unexpected sha256: %#v", items[0])
}
}
func TestListInstalledModelsScansUploadedFilesDir(t *testing.T) {
dir := t.TempDir()
store := New(dir, 8)
if err := os.MkdirAll(store.FilesDir(), 0o755); err != nil {
@ -16,9 +52,6 @@ func TestListInstalledModelsUsesManifestEntries(t *testing.T) {
if err := os.WriteFile(modelPath, []byte("abc"), 0o644); err != nil {
t.Fatalf("WriteFile: %v", err)
}
if err := os.WriteFile(store.ManifestPath(), []byte(`{"items":[{"name":"face_det_scrfd_500m_640_rk3588","sha256":"abc","path":"`+filepath.ToSlash(modelPath)+`","size":3,"mtime_ms":1}]}`), 0o644); err != nil {
t.Fatalf("WriteFile manifest: %v", err)
}
items, err := store.ListInstalledModels()
if err != nil {

Binary file not shown.