diff --git a/agent/internal/httpapi/models_status_test.go b/agent/internal/httpapi/models_status_test.go index e42d5b1..046531c 100644 --- a/agent/internal/httpapi/models_status_test.go +++ b/agent/internal/httpapi/models_status_test.go @@ -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) diff --git a/agent/internal/modelstore/modelstore.go b/agent/internal/modelstore/modelstore.go index 99958cf..48e3cce 100644 --- a/agent/internal/modelstore/modelstore.go +++ b/agent/internal/modelstore/modelstore.go @@ -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 { diff --git a/agent/internal/modelstore/modelstore_test.go b/agent/internal/modelstore/modelstore_test.go index 2c40704..8dcf1ba 100644 --- a/agent/internal/modelstore/modelstore_test.go +++ b/agent/internal/modelstore/modelstore_test.go @@ -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 { diff --git a/agent/rk3588-agent_linux_arm64 b/agent/rk3588-agent_linux_arm64 index 9cc4ff2..056f380 100755 Binary files a/agent/rk3588-agent_linux_arm64 and b/agent/rk3588-agent_linux_arm64 differ