safesight-edge/agent/internal/modelstore/modelstore.go
sladro 2a562514bc
Some checks are pending
CI / host-build (push) Waiting to run
CI / rk3588-cross-build (push) Waiting to run
修改agent功能,加入重启,准备测试
2026-01-10 11:06:39 +08:00

166 lines
3.6 KiB
Go

package modelstore
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"rk3588sys/agent/internal/files"
)
type Item struct {
Name string `json:"name"`
Sha256 string `json:"sha256"`
Path string `json:"path"`
Size int64 `json:"size"`
MtimeMS int64 `json:"mtime_ms"`
}
type Manifest struct {
Items []Item `json:"items"`
}
type Store struct {
ModelsDir string
MaxUploadBytes int64
}
func New(modelsDir string, maxUploadMB int) *Store {
max := int64(maxUploadMB) * 1024 * 1024
return &Store{ModelsDir: modelsDir, MaxUploadBytes: max}
}
func (s *Store) ManifestPath() string {
return filepath.Join(s.ModelsDir, "manifest.json")
}
func (s *Store) FilesDir() string {
return filepath.Join(s.ModelsDir, "files")
}
func (s *Store) Upload(name string, r io.Reader, contentLength int64, expectedSha256 string) (Item, error) {
if contentLength <= 0 {
return Item{}, errors.New("missing Content-Length")
}
if s.MaxUploadBytes > 0 && contentLength > s.MaxUploadBytes {
return Item{}, ErrPayloadTooLarge
}
filesDir := s.FilesDir()
if err := files.EnsureDir(filesDir, 0o755); err != nil {
return Item{}, err
}
tmpFile, err := os.CreateTemp(filesDir, ".tmp-*")
if err != nil {
return Item{}, fmt.Errorf("create temp: %w", err)
}
tmp := tmpFile.Name()
ok := false
defer func() {
_ = tmpFile.Close()
if !ok {
_ = os.Remove(tmp)
}
}()
h := sha256.New()
mw := io.MultiWriter(tmpFile, h)
if _, err := io.CopyN(mw, r, contentLength); err != nil {
return Item{}, fmt.Errorf("read body: %w", err)
}
if err := tmpFile.Chmod(0o644); err != nil {
return Item{}, fmt.Errorf("chmod temp: %w", err)
}
if err := tmpFile.Sync(); err != nil {
return Item{}, fmt.Errorf("fsync temp: %w", err)
}
if err := tmpFile.Close(); err != nil {
return Item{}, fmt.Errorf("close temp: %w", err)
}
sha := hex.EncodeToString(h.Sum(nil))
if expectedSha256 != "" && !strings.EqualFold(expectedSha256, sha) {
return Item{}, fmt.Errorf("sha256 mismatch")
}
finalName := fmt.Sprintf("%s__%s.rknn", name, sha)
finalPath := filepath.Join(filesDir, finalName)
if _, err := os.Stat(finalPath); err == nil {
_ = os.Remove(tmp)
} else {
if err := files.ReplaceFile(tmp, finalPath); err != nil {
return Item{}, err
}
}
stat, err := os.Stat(finalPath)
if err != nil {
return Item{}, fmt.Errorf("stat final: %w", err)
}
item := Item{
Name: name,
Sha256: sha,
Path: filepath.ToSlash(finalPath),
Size: stat.Size(),
MtimeMS: stat.ModTime().UnixMilli(),
}
if err := s.upsertManifest(item); err != nil {
return Item{}, err
}
ok = true
return item, nil
}
func (s *Store) List() (Manifest, error) {
var m Manifest
b, err := os.ReadFile(s.ManifestPath())
if err != nil {
if os.IsNotExist(err) {
return Manifest{Items: []Item{}}, nil
}
return Manifest{}, fmt.Errorf("read manifest: %w", err)
}
if err := json.Unmarshal(b, &m); err != nil {
return Manifest{}, fmt.Errorf("parse manifest: %w", err)
}
if m.Items == nil {
m.Items = []Item{}
}
return m, nil
}
func (s *Store) upsertManifest(item Item) error {
m, err := s.List()
if err != nil {
return err
}
found := false
for i := range m.Items {
if m.Items[i].Name == item.Name {
m.Items[i] = item
found = true
break
}
}
if !found {
m.Items = append(m.Items, item)
}
b, err := json.Marshal(m)
if err != nil {
return fmt.Errorf("marshal manifest: %w", err)
}
if err := files.WriteFileAtomic(s.ManifestPath(), append(b, '\n'), 0o644); err != nil {
return fmt.Errorf("write manifest: %w", err)
}
return nil
}
var ErrPayloadTooLarge = errors.New("payload too large")