feat: complete 2.0.7.12 platform overhaul

This commit is contained in:
2026-08-16 19:33:03 +08:00
parent 73555cd04c
commit c9fa6f7a88
159 changed files with 13243 additions and 2539 deletions
@@ -1,6 +1,7 @@
package releases
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
@@ -8,8 +9,10 @@ import (
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"reflect"
"regexp"
"sort"
"strconv"
@@ -23,11 +26,19 @@ import (
)
type Service struct {
cfg *config.Config
store *db.Store
notices *notices.Service
hashMu sync.Mutex
hashes map[string]cachedFileHash
cfg *config.Config
store *db.Store
notices *notices.Service
hashMu sync.Mutex
hashes map[string]cachedFileHash
hashFile func(string) string
packageMu sync.RWMutex
packages []Package
reconcileMu sync.Mutex
callbackMu sync.RWMutex
onChange func()
stop chan struct{}
stopOnce sync.Once
}
type cachedFileHash struct {
@@ -72,16 +83,60 @@ var (
ErrUploadedPackageEmpty = errors.New("uploaded file is empty")
ErrUploadedPackageStorageFailed = errors.New("upload storage failed")
ErrUploadedPackageManifestFailed = errors.New("manifest update failed")
ErrUploadedPackageIndexFailed = errors.New("package index update failed")
ErrUnsupportedPackage = errors.New("unsupported package extension")
ErrUnsafePackageName = errors.New("unsafe package file name")
)
func NewService(cfg *config.Config, store *db.Store, noticeService ...*notices.Service) *Service {
service := &Service{cfg: cfg, store: store, hashes: map[string]cachedFileHash{}}
return newService(cfg, store, sha256File, noticeService...)
}
func newService(cfg *config.Config, store *db.Store, hashFile func(string) string, noticeService ...*notices.Service) *Service {
service := &Service{cfg: cfg, store: store, hashes: map[string]cachedFileHash{}, hashFile: hashFile, stop: make(chan struct{})}
if len(noticeService) > 0 {
service.notices = noticeService[0]
}
_ = service.Reconcile(context.Background())
return service
}
func (s *Service) Start(ctx context.Context) {
go func() {
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for {
select {
case <-ticker.C:
_ = s.Reconcile(context.Background())
case <-ctx.Done():
return
case <-s.stop:
return
}
}
}()
}
func (s *Service) Stop() {
s.stopOnce.Do(func() { close(s.stop) })
}
func (s *Service) SetChangeCallback(callback func()) {
s.callbackMu.Lock()
s.onChange = callback
s.callbackMu.Unlock()
}
func (s *Service) notifyChanged() {
s.callbackMu.RLock()
callback := s.onChange
s.callbackMu.RUnlock()
if callback != nil {
callback()
}
}
func (s *Service) LegacyUpdateInfo(r *http.Request) map[string]any {
payload := s.legacyUpdateBase()
manifest := s.Manifest(r)
@@ -102,7 +157,6 @@ func (s *Service) Manifest(r *http.Request) map[string]any {
}
payload["manifest_version"] = 2
payload["service_version"] = config.Version
payload["packages"] = packages
payload["modules"] = modules
payload["assets"] = []any{}
payload["generated_at"] = time.Now().UTC().Format(time.RFC3339)
@@ -121,23 +175,60 @@ func (s *Service) Manifest(r *http.Request) map[string]any {
setIfMissing(payload, "download_url", latestNotice.DownloadURL)
}
}
preferredVersion := manifestString(payload, "app_version")
packages = prioritizePackages(packages, preferredVersion)
payload["packages"] = packages
if len(packages) > 0 {
latest := packages[0]
payload["app_version"] = latest.Version
payload["download_url"] = latest.URL
payload["download_mirrors"] = []map[string]any{{
"id": "primary",
"name": "官方直连",
"url": latest.URL,
"type": "direct",
"sha256": latest.SHA256,
"enabled": true,
}}
payload["detected_product"] = latest.Name
if preferredVersion == "" {
preferredVersion = latest.Version
payload["app_version"] = preferredVersion
}
if sameVersion(latest.Version, preferredVersion) {
payload["download_url"] = latest.URL
payload["download_mirrors"] = []map[string]any{{
"id": "primary",
"name": "官方直连",
"url": latest.URL,
"type": "direct",
"sha256": latest.SHA256,
"enabled": true,
}}
payload["detected_product"] = latest.Name
}
}
return payload
}
func manifestString(payload map[string]any, key string) string {
value, _ := payload[key].(string)
return strings.TrimSpace(value)
}
func prioritizePackages(items []Package, preferredVersion string) []Package {
if strings.TrimSpace(preferredVersion) == "" {
return items
}
sort.SliceStable(items, func(i, j int) bool {
left := sameVersion(items[i].Version, preferredVersion)
right := sameVersion(items[j].Version, preferredVersion)
return left && !right
})
return items
}
func sameVersion(a, b string) bool {
a = strings.TrimSpace(a)
b = strings.TrimSpace(b)
if a == "" || b == "" {
return false
}
if a == b {
return true
}
return versionPattern.FindString(a) == a && versionPattern.FindString(b) == b && compareVersion(a, b) == 0
}
func (s *Service) PublishLegacyUpdateInfo(r *http.Request, actor string) error {
payload := s.LegacyUpdateInfo(r)
data, err := json.MarshalIndent(payload, "", " ")
@@ -162,48 +253,149 @@ func setIfMissing(payload map[string]any, key, value string) {
}
func (s *Service) ScanPackages(r *http.Request) []Package {
base := firstNonEmpty(strings.TrimRight(s.cfg.CDNBaseURL, "/"), requestBaseURL(r, s.cfg.BaseURL))
s.packageMu.RLock()
items := append([]Package(nil), s.packages...)
s.packageMu.RUnlock()
for index := range items {
items[index].URL = base + "/downloads/" + items[index].FileName
}
return items
}
// Reconcile keeps request-time manifest generation independent from package I/O.
// Hashes are reused from the persistent index whenever size and mtime still match.
func (s *Service) Reconcile(ctx context.Context) error {
s.reconcileMu.Lock()
defer s.reconcileMu.Unlock()
if err := os.MkdirAll(s.cfg.DownloadsDir, 0o750); err != nil {
return err
}
indexed := map[string]db.ReleasePackage{}
var firstErr error
if s.store != nil {
rows, err := s.store.ListReleasePackages()
if err != nil {
firstErr = err
} else {
for _, item := range rows {
indexed[item.FileName] = item
}
}
}
entries, err := os.ReadDir(s.cfg.DownloadsDir)
if err != nil {
return []Package{}
return err
}
base := firstNonEmpty(strings.TrimRight(s.cfg.CDNBaseURL, "/"), requestBaseURL(r, s.cfg.BaseURL))
items := []Package{}
items := make([]Package, 0, len(entries))
seen := map[string]bool{}
for _, entry := range entries {
if entry.IsDir() {
continue
if err := ctx.Err(); err != nil {
return err
}
name := entry.Name()
lower := strings.ToLower(name)
if !(strings.HasSuffix(lower, ".exe") || strings.HasSuffix(lower, ".msix") || strings.HasSuffix(lower, ".appinstaller") || strings.HasSuffix(lower, ".msi")) {
if entry.IsDir() || !isSupportedPackageName(entry.Name()) {
continue
}
info, err := entry.Info()
if err != nil {
if firstErr == nil {
firstErr = err
}
continue
}
version := detectVersion(name)
platform, arch := detectPlatform(name)
product := detectProduct(name)
url := base + "/downloads/" + name
items = append(items, Package{
ID: strings.ToLower(strings.ReplaceAll(product+"-"+platform+"-"+arch+"-"+version, " ", "-")),
Name: product,
Version: version,
Platform: platform,
Arch: arch,
URL: url,
SHA256: s.cachedSHA256(filepath.Join(s.cfg.DownloadsDir, name), info),
Size: info.Size(),
Required: strings.Contains(strings.ToLower(product), "ymhut"),
Enabled: true,
FileName: name,
UpdatedAt: info.ModTime().UTC().Format(time.RFC3339),
})
name := entry.Name()
seen[name] = true
modifiedAt := info.ModTime().UTC().Format(time.RFC3339Nano)
record, cached := indexed[name]
var item Package
if cached && record.SizeBytes == info.Size() && record.UpdatedAt == modifiedAt && strings.TrimSpace(record.SHA256) != "" {
item = packageFromRecord(record)
s.hashMu.Lock()
s.hashes[filepath.Join(s.cfg.DownloadsDir, name)] = cachedFileHash{size: info.Size(), modifiedAt: info.ModTime().UnixNano(), sha256: record.SHA256}
s.hashMu.Unlock()
} else {
version := detectVersion(name)
platform, arch := detectPlatform(name)
product := detectProduct(name)
hashValue := s.cachedSHA256(filepath.Join(s.cfg.DownloadsDir, name), info)
if strings.TrimSpace(hashValue) == "" {
if firstErr == nil {
firstErr = fmt.Errorf("hash package %q: empty SHA-256", name)
}
continue
}
item = Package{
ID: packageID(product, platform, arch, version),
Name: product,
Version: version,
Platform: platform,
Arch: arch,
SHA256: hashValue,
Size: info.Size(),
Required: strings.Contains(strings.ToLower(product), "ymhut"),
Enabled: true,
FileName: name,
UpdatedAt: modifiedAt,
}
if s.store != nil {
if _, err := s.store.UpsertReleasePackage(releasePackageRecord(item)); err != nil && firstErr == nil {
firstErr = err
}
}
}
items = append(items, item)
}
sort.Slice(items, func(i, j int) bool {
return compareVersion(items[i].Version, items[j].Version) > 0
})
return items
if s.store != nil {
for name := range indexed {
if !seen[name] {
if err := s.store.DeleteReleasePackage(name); err != nil && firstErr == nil {
firstErr = err
}
}
}
}
if firstErr != nil {
return firstErr
}
sort.Slice(items, func(i, j int) bool { return compareVersion(items[i].Version, items[j].Version) > 0 })
s.packageMu.Lock()
changed := !reflect.DeepEqual(s.packages, items)
s.packages = items
s.packageMu.Unlock()
if changed {
s.notifyChanged()
}
return firstErr
}
func packageFromRecord(item db.ReleasePackage) Package {
return Package{
ID: packageID(item.Product, item.Platform, item.Arch, item.Version), Name: item.Product, Version: item.Version,
Platform: item.Platform, Arch: item.Arch, URL: item.URL, SHA256: item.SHA256, Size: item.SizeBytes,
Required: strings.Contains(strings.ToLower(item.Product), "ymhut"), Enabled: item.Enabled,
FileName: item.FileName, UpdatedAt: item.UpdatedAt,
}
}
func releasePackageRecord(item Package) db.ReleasePackage {
return db.ReleasePackage{
Product: item.Name, Version: item.Version, Platform: item.Platform, Arch: item.Arch, FileName: item.FileName,
URL: item.URL, SHA256: item.SHA256, SizeBytes: item.Size, Enabled: item.Enabled, UpdatedAt: item.UpdatedAt,
}
}
func packageID(product, platform, arch, version string) string {
return strings.ToLower(strings.ReplaceAll(product+"-"+platform+"-"+arch+"-"+version, " ", "-"))
}
func isSupportedPackageName(name string) bool {
lower := strings.ToLower(name)
for _, suffix := range []string{".exe", ".msix", ".appinstaller", ".msi", ".zip", ".7z"} {
if strings.HasSuffix(lower, suffix) {
return true
}
}
return false
}
func (s *Service) cachedSHA256(path string, info os.FileInfo) string {
@@ -215,7 +407,7 @@ func (s *Service) cachedSHA256(path string, info os.FileInfo) string {
}
s.hashMu.Unlock()
value := sha256File(path)
value := s.hashFile(path)
s.hashMu.Lock()
s.hashes[path] = cachedFileHash{size: info.Size(), modifiedAt: modifiedAt, sha256: value}
for cachedPath := range s.hashes {
@@ -241,64 +433,28 @@ func (s *Service) SaveUploadedPackage(r *http.Request, reader io.Reader, opts Up
if err := os.MkdirAll(s.cfg.DownloadsDir, 0o750); err != nil {
return Package{}, fmt.Errorf("%w: %v", ErrUploadedPackageStorageFailed, err)
}
target := filepath.Join(s.cfg.DownloadsDir, name)
resolved, err := filepath.Abs(target)
if err != nil {
return Package{}, err
}
base, _ := filepath.Abs(s.cfg.DownloadsDir)
if resolved != base && !strings.HasPrefix(resolved, base+string(os.PathSeparator)) {
return Package{}, errors.New("path escape rejected")
}
tmp, err := os.CreateTemp(s.cfg.DownloadsDir, "."+name+".*.upload")
if err != nil {
return Package{}, err
return Package{}, fmt.Errorf("%w: %v", ErrUploadedPackageStorageFailed, err)
}
tmpName := tmp.Name()
defer os.Remove(tmpName)
hash := sha256.New()
written, err := io.Copy(tmp, io.TeeReader(reader, hash))
written, err := io.Copy(io.MultiWriter(tmp, hash), reader)
if err == nil {
err = tmp.Sync()
}
if closeErr := tmp.Close(); err == nil {
err = closeErr
}
if err != nil {
return Package{}, err
return Package{}, fmt.Errorf("%w: %v", ErrUploadedPackageStorageFailed, err)
}
if written <= 0 {
return Package{}, errors.New("uploaded file is empty")
return Package{}, ErrUploadedPackageEmpty
}
if err := os.Chmod(tmpName, 0o640); err != nil {
return Package{}, err
}
if err := os.Rename(tmpName, target); err != nil {
return Package{}, err
}
version := firstNonEmpty(opts.Version, detectVersion(name))
platform, arch := detectPlatform(name)
platform = firstNonEmpty(opts.Platform, platform)
arch = firstNonEmpty(opts.Arch, arch)
product := detectProduct(name)
pkg := Package{
ID: strings.ToLower(strings.ReplaceAll(product+"-"+platform+"-"+arch+"-"+version, " ", "-")),
Name: product,
Version: version,
Platform: platform,
Arch: arch,
URL: firstNonEmpty(strings.TrimRight(s.cfg.CDNBaseURL, "/"), requestBaseURL(r, s.cfg.BaseURL)) + "/downloads/" + name,
SHA256: hex.EncodeToString(hash.Sum(nil)),
Size: written,
Required: strings.Contains(strings.ToLower(product), "ymhut"),
Enabled: true,
FileName: name,
UpdatedAt: time.Now().UTC().Format(time.RFC3339),
}
if opts.UpdateManifest {
if err := s.updateLegacyManifest(pkg, opts); err != nil {
return Package{}, err
}
}
_ = s.store.InsertAudit(db.AuditLog{Actor: firstNonEmpty(actor, "admin"), Type: "release.package_uploaded", Target: name, Message: fmt.Sprintf("已上传发布包 %s%s%s", name, version, formatBytes(written))})
return pkg, nil
opts.FileName = name
return s.SavePreparedPackage(r, UploadedPackageFile{TempPath: tmpName, Size: written, SHA256: hex.EncodeToString(hash.Sum(nil))}, opts, actor)
}
func (s *Service) SavePreparedPackage(r *http.Request, uploaded UploadedPackageFile, opts UploadOptions, actor string) (Package, error) {
@@ -335,13 +491,26 @@ func (s *Service) SavePreparedPackage(r *http.Request, uploaded UploadedPackageF
_ = restorePackageBackup(target, backup)
return Package{}, fmt.Errorf("%w: %v", ErrUploadedPackageStorageFailed, err)
}
rollbackFile := func(cause error) error {
if rollbackErr := restorePackageBackup(target, backup); rollbackErr != nil {
return fmt.Errorf("%w; rollback failed: %v", cause, rollbackErr)
}
return cause
}
info, err := os.Stat(target)
if err != nil {
return Package{}, rollbackFile(fmt.Errorf("%w: %v", ErrUploadedPackageStorageFailed, err))
}
if info.Size() != uploaded.Size {
return Package{}, rollbackFile(fmt.Errorf("%w: uploaded size changed during commit", ErrUploadedPackageStorageFailed))
}
version := firstNonEmpty(opts.Version, detectVersion(name))
platform, arch := detectPlatform(name)
platform = firstNonEmpty(opts.Platform, platform)
arch = firstNonEmpty(opts.Arch, arch)
product := detectProduct(name)
pkg := Package{
ID: strings.ToLower(strings.ReplaceAll(product+"-"+platform+"-"+arch+"-"+version, " ", "-")),
ID: packageID(product, platform, arch, version),
Name: product,
Version: version,
Platform: platform,
@@ -352,23 +521,95 @@ func (s *Service) SavePreparedPackage(r *http.Request, uploaded UploadedPackageF
Required: strings.Contains(strings.ToLower(product), "ymhut"),
Enabled: true,
FileName: name,
UpdatedAt: time.Now().UTC().Format(time.RFC3339),
UpdatedAt: info.ModTime().UTC().Format(time.RFC3339Nano),
}
var previousManifest fileSnapshot
if opts.UpdateManifest {
previousManifest, err = captureFile(filepath.Join(s.cfg.UpdatePublicDir, "update-info.json"))
if err != nil {
return Package{}, rollbackFile(fmt.Errorf("%w: %v", ErrUploadedPackageManifestFailed, err))
}
if err := s.updateLegacyManifest(pkg, opts); err != nil {
if rollbackErr := restorePackageBackup(target, backup); rollbackErr != nil {
return Package{}, fmt.Errorf("%w: %v; rollback failed: %v", ErrUploadedPackageManifestFailed, err, rollbackErr)
return Package{}, rollbackFile(fmt.Errorf("%w: %v", ErrUploadedPackageManifestFailed, err))
}
}
if s.store != nil {
if _, err := s.store.UpsertReleasePackage(releasePackageRecord(pkg)); err != nil {
cause := fmt.Errorf("%w: %v", ErrUploadedPackageIndexFailed, err)
if opts.UpdateManifest {
if restoreErr := previousManifest.restore(filepath.Join(s.cfg.UpdatePublicDir, "update-info.json")); restoreErr != nil {
cause = fmt.Errorf("%w; manifest rollback failed: %v", cause, restoreErr)
}
}
return Package{}, fmt.Errorf("%w: %v", ErrUploadedPackageManifestFailed, err)
return Package{}, rollbackFile(cause)
}
}
if backup != "" {
_ = os.Remove(backup)
}
_ = s.store.InsertAudit(db.AuditLog{Actor: firstNonEmpty(actor, "admin"), Type: "release.package_uploaded", Target: name, Message: fmt.Sprintf("已上传发布包 %s%s%s", name, version, formatBytes(uploaded.Size))})
s.upsertPackageSnapshot(pkg)
if s.store != nil {
_ = s.store.InsertAudit(db.AuditLog{Actor: firstNonEmpty(actor, "admin"), Type: "release.package_uploaded", Target: name, Message: fmt.Sprintf("已上传发布包 %s%s%s", name, version, formatBytes(uploaded.Size))})
}
return pkg, nil
}
type fileSnapshot struct {
data []byte
mode os.FileMode
existed bool
}
func captureFile(path string) (fileSnapshot, error) {
info, err := os.Stat(path)
if errors.Is(err, os.ErrNotExist) {
return fileSnapshot{}, nil
}
if err != nil {
return fileSnapshot{}, err
}
data, err := os.ReadFile(path)
if err != nil {
return fileSnapshot{}, err
}
return fileSnapshot{data: data, mode: info.Mode(), existed: true}, nil
}
func (s fileSnapshot) restore(path string) error {
if !s.existed {
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
return err
}
return nil
}
if err := os.MkdirAll(filepath.Dir(path), 0o750); err != nil {
return err
}
mode := s.mode
if mode == 0 {
mode = 0o640
}
return os.WriteFile(path, s.data, mode)
}
func (s *Service) upsertPackageSnapshot(item Package) {
s.packageMu.Lock()
replaced := false
for index := range s.packages {
if s.packages[index].FileName == item.FileName {
s.packages[index] = item
replaced = true
break
}
}
if !replaced {
s.packages = append(s.packages, item)
}
sort.Slice(s.packages, func(i, j int) bool { return compareVersion(s.packages[i].Version, s.packages[j].Version) > 0 })
s.packageMu.Unlock()
s.notifyChanged()
}
func backupExistingPackage(target string) (string, error) {
if _, err := os.Stat(target); errors.Is(err, os.ErrNotExist) {
return "", nil
@@ -492,21 +733,46 @@ func atomicWrite(path string, data []byte) error {
func requestBaseURL(r *http.Request, fallback string) string {
if r != nil {
scheme := r.Header.Get("X-Forwarded-Proto")
if scheme == "" {
scheme := firstForwardedHeader(r.Header.Get("X-Forwarded-Proto"))
if scheme != "http" && scheme != "https" {
if r.TLS != nil {
scheme = "https"
} else {
scheme = "http"
}
}
if r.Host != "" {
return scheme + "://" + r.Host
host := firstForwardedHeader(r.Header.Get("X-Forwarded-Host"))
if !validForwardedHost(host) {
host = strings.TrimSpace(r.Host)
}
if validForwardedHost(host) {
return scheme + "://" + host
}
}
return strings.TrimRight(fallback, "/")
}
func firstForwardedHeader(value string) string {
return strings.ToLower(strings.TrimSpace(strings.Split(value, ",")[0]))
}
func validForwardedHost(value string) bool {
if value == "" || strings.TrimSpace(value) != value || strings.ContainsAny(value, "\\\\\r\n\t ") {
return false
}
parsed, err := url.Parse("http://" + value)
if err != nil || parsed.Host != value || parsed.User != nil || parsed.Hostname() == "" || parsed.Path != "" || parsed.RawQuery != "" || parsed.Fragment != "" {
return false
}
if port := parsed.Port(); port != "" {
value, err := strconv.Atoi(port)
if err != nil || value < 1 || value > 65535 {
return false
}
}
return true
}
var versionPattern = regexp.MustCompile(`\d+\.\d+\.\d+(?:\.\d+)?`)
func detectVersion(name string) string {
@@ -577,16 +843,13 @@ func sha256File(path string) string {
func safePackageName(name string) (string, error) {
original := strings.TrimSpace(name)
if original == "" || original == "." || original == ".." || strings.ContainsAny(original, `/\`) {
return "", errors.New("invalid filename")
return "", ErrUnsafePackageName
}
name = filepath.Base(original)
lower := strings.ToLower(name)
for _, suffix := range []string{".exe", ".msix", ".appinstaller", ".msi", ".zip", ".7z"} {
if strings.HasSuffix(lower, suffix) {
return name, nil
}
if isSupportedPackageName(name) {
return name, nil
}
return "", errors.New("unsupported package extension")
return "", ErrUnsupportedPackage
}
func formatBytes(value int64) string {
@@ -2,6 +2,7 @@ package releases
import (
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
@@ -12,6 +13,46 @@ import (
"ymhut-box/server/unified-management/internal/db"
)
func TestRequestBaseURLUsesTrustedForwardedValues(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "http://internal:33550/downloads/file.zip", nil)
req.Header.Set("X-Forwarded-Proto", "HTTPS, http")
req.Header.Set("X-Forwarded-Host", "updates.example.com:8443, internal:33550")
if got := requestBaseURL(req, "https://fallback.example.com/"); got != "https://updates.example.com:8443" {
t.Fatalf("requestBaseURL() = %q", got)
}
}
func TestRequestBaseURLRejectsInvalidForwardedValues(t *testing.T) {
tests := []struct {
name string
proto string
host string
request string
fallback string
want string
}{
{name: "invalid scheme", proto: "javascript", host: "updates.example.com", request: "https://internal:33550/path", want: "https://updates.example.com"},
{name: "invalid forwarded host", proto: "https", host: "user@evil.example", request: "http://internal:33550/path", want: "https://internal:33550"},
{name: "invalid forwarded port", proto: "https", host: "updates.example.com:99999", request: "http://internal:33550/path", want: "https://internal:33550"},
{name: "invalid request host", proto: "https", host: "evil.example/path", request: "http://internal/path", fallback: "https://fallback.example.com/", want: "https://fallback.example.com"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, test.request, nil)
req.Header.Set("X-Forwarded-Proto", test.proto)
req.Header.Set("X-Forwarded-Host", test.host)
if test.name == "invalid request host" {
req.Host = "bad host"
}
if got := requestBaseURL(req, test.fallback); got != test.want {
t.Fatalf("requestBaseURL() = %q, want %q", got, test.want)
}
})
}
}
func TestCompareVersion(t *testing.T) {
cases := []struct {
a string
@@ -67,6 +108,231 @@ func TestScanPackagesUsesCDNAndCachesHashes(t *testing.T) {
}
}
func TestReconcileListsEverySupportedPackageType(t *testing.T) {
dir := t.TempDir()
for _, name := range []string{"app.exe", "app.msix", "app.appinstaller", "app.msi", "app.zip", "app.7z"} {
if err := os.WriteFile(filepath.Join(dir, name), []byte(name), 0o640); err != nil {
t.Fatal(err)
}
}
if err := os.WriteFile(filepath.Join(dir, "notes.txt"), []byte("ignored"), 0o640); err != nil {
t.Fatal(err)
}
service := NewService(&config.Config{DownloadsDir: dir, BaseURL: "https://update.example"}, nil)
if got := len(service.ScanPackages(httptest.NewRequest("GET", "https://update.example/api/client/releases", nil))); got != 6 {
t.Fatalf("listed %d packages, want 6", got)
}
}
func TestReconcileReusesPersistentHashAfterRestart(t *testing.T) {
dir := t.TempDir()
cfg := &config.Config{
StorageDir: filepath.Join(dir, "storage"),
DownloadsDir: filepath.Join(dir, "downloads"),
BaseURL: "https://update.example",
Database: config.DatabaseConfig{
Provider: "sqlite",
SQLitePath: filepath.Join(dir, "storage", "unified.sqlite"),
HealthIntervalSec: 30,
},
}
if err := os.MkdirAll(cfg.DownloadsDir, 0o750); err != nil {
t.Fatal(err)
}
name := "YMhut_Box_2.1.0_x64.zip"
if err := os.WriteFile(filepath.Join(cfg.DownloadsDir, name), []byte("persistent package hash"), 0o640); err != nil {
t.Fatal(err)
}
store, err := db.Open(cfg)
if err != nil {
t.Fatal(err)
}
defer store.Close()
first := NewService(cfg, store)
if packages := first.ScanPackages(httptest.NewRequest("GET", "https://update.example", nil)); len(packages) != 1 || packages[0].SHA256 == "" {
t.Fatalf("initial reconcile did not index the package: %#v", packages)
}
hashCalls := 0
second := newService(cfg, store, func(path string) string {
hashCalls++
return sha256File(path)
})
if packages := second.ScanPackages(httptest.NewRequest("GET", "https://update.example", nil)); len(packages) != 1 {
t.Fatalf("restart snapshot contains %d packages, want 1", len(packages))
}
if hashCalls != 0 {
t.Fatalf("restart rehashed %d files, want 0", hashCalls)
}
}
func TestReconcileRehashesChangedFileAndRemovesDeletedIndex(t *testing.T) {
service, cfg, cleanup := newPreparedPackageTestService(t)
defer cleanup()
name := "YMhut_Box_2.2.0_x64.zip"
path := filepath.Join(cfg.DownloadsDir, name)
if err := os.MkdirAll(cfg.DownloadsDir, 0o750); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte("first"), 0o640); err != nil {
t.Fatal(err)
}
hashCalls := 0
service.hashFile = func(path string) string {
hashCalls++
return sha256File(path)
}
if err := service.Reconcile(t.Context()); err != nil {
t.Fatal(err)
}
first := service.ScanPackages(nil)
if len(first) != 1 || hashCalls != 1 {
t.Fatalf("first reconcile packages=%d hashCalls=%d", len(first), hashCalls)
}
if err := os.WriteFile(path, []byte("changed package bytes"), 0o640); err != nil {
t.Fatal(err)
}
if err := service.Reconcile(t.Context()); err != nil {
t.Fatal(err)
}
second := service.ScanPackages(nil)
if hashCalls != 2 || len(second) != 1 || second[0].SHA256 == first[0].SHA256 {
t.Fatalf("changed package was not rehashed: first=%#v second=%#v calls=%d", first, second, hashCalls)
}
if err := os.Remove(path); err != nil {
t.Fatal(err)
}
if err := service.Reconcile(t.Context()); err != nil {
t.Fatal(err)
}
if packages := service.ScanPackages(nil); len(packages) != 0 {
t.Fatalf("deleted package remains in snapshot: %#v", packages)
}
if _, ok, err := service.store.GetReleasePackage(name); err != nil || ok {
t.Fatalf("deleted package remains in index: ok=%v err=%v", ok, err)
}
}
func TestScanPackagesNeverHashesDuringRequest(t *testing.T) {
dir := t.TempDir()
name := "YMhut_Box_2.2.1_x64.7z"
if err := os.WriteFile(filepath.Join(dir, name), []byte("large package placeholder"), 0o640); err != nil {
t.Fatal(err)
}
service := NewService(&config.Config{DownloadsDir: dir, BaseURL: "https://update.example"}, nil)
service.hashFile = func(string) string {
t.Fatal("request-time package scan attempted to hash file content")
return ""
}
for range 10 {
if packages := service.ScanPackages(httptest.NewRequest("GET", "https://update.example/api/client/bootstrap", nil)); len(packages) != 1 {
t.Fatalf("snapshot contains %d packages, want 1", len(packages))
}
}
}
func TestManifestPrioritizesPublishedVersionOverNumericallyHigherHistoricalPackage(t *testing.T) {
dir := t.TempDir()
cfg := &config.Config{
DownloadsDir: filepath.Join(dir, "downloads"),
UpdatePublicDir: filepath.Join(dir, "update", "public"),
BaseURL: "https://update.example",
}
if err := os.MkdirAll(cfg.DownloadsDir, 0o750); err != nil {
t.Fatal(err)
}
if err := os.MkdirAll(cfg.UpdatePublicDir, 0o750); err != nil {
t.Fatal(err)
}
for _, name := range []string{
"YMhut_Box_WinUI_Setup_2.0.7.111.exe",
"YMhut_Box_WinUI_Setup_2.0.7.12.exe",
} {
if err := os.WriteFile(filepath.Join(cfg.DownloadsDir, name), []byte(name), 0o640); err != nil {
t.Fatal(err)
}
}
if err := os.WriteFile(
filepath.Join(cfg.UpdatePublicDir, "update-info.json"),
[]byte(`{"app_version":"2.0.7.12"}`),
0o640); err != nil {
t.Fatal(err)
}
service := NewService(cfg, nil)
manifest := service.Manifest(httptest.NewRequest(http.MethodGet, "https://update.example/api/client/releases", nil))
if got := manifestString(manifest, "app_version"); got != "2.0.7.12" {
t.Fatalf("manifest app_version = %q, want 2.0.7.12", got)
}
packages, ok := manifest["packages"].([]Package)
if !ok || len(packages) != 2 {
t.Fatalf("manifest packages = %#v", manifest["packages"])
}
if packages[0].Version != "2.0.7.12" {
t.Fatalf("first package version = %q, want published version 2.0.7.12", packages[0].Version)
}
if got := manifestString(manifest, "download_url"); !strings.Contains(got, "2.0.7.12.exe") {
t.Fatalf("manifest download_url = %q, want current package", got)
}
}
func TestManifestFallsBackToHighestPackageWithoutPublishedVersion(t *testing.T) {
dir := t.TempDir()
for _, name := range []string{"YMhut_Box_2.0.7.12.exe", "YMhut_Box_2.0.7.111.exe"} {
if err := os.WriteFile(filepath.Join(dir, name), []byte(name), 0o640); err != nil {
t.Fatal(err)
}
}
service := NewService(&config.Config{
DownloadsDir: dir,
UpdatePublicDir: filepath.Join(dir, "missing-public"),
BaseURL: "https://update.example",
}, nil)
manifest := service.Manifest(httptest.NewRequest(http.MethodGet, "https://update.example/api/client/releases", nil))
if got := manifestString(manifest, "app_version"); got != "2.0.7.111" {
t.Fatalf("manifest fallback app_version = %q, want highest indexed package 2.0.7.111", got)
}
}
func TestSavePreparedPackageRestoresFileAndManifestWhenIndexFails(t *testing.T) {
service, cfg, closeStore := newPreparedPackageTestService(t)
name := "YMhut_Box_WinUI_Setup_2.3.0_x64.exe"
if err := os.MkdirAll(cfg.DownloadsDir, 0o750); err != nil {
t.Fatal(err)
}
target := filepath.Join(cfg.DownloadsDir, name)
if err := os.WriteFile(target, []byte("old package"), 0o640); err != nil {
t.Fatal(err)
}
manifestPath := filepath.Join(cfg.UpdatePublicDir, "update-info.json")
if err := os.MkdirAll(cfg.UpdatePublicDir, 0o750); err != nil {
t.Fatal(err)
}
oldManifest := []byte("{\"app_version\":\"1.0.0\"}\n")
if err := os.WriteFile(manifestPath, oldManifest, 0o640); err != nil {
t.Fatal(err)
}
temp := filepath.Join(cfg.DownloadsDir, ".upload-new")
if err := os.WriteFile(temp, []byte("new package"), 0o640); err != nil {
t.Fatal(err)
}
closeStore()
_, err := service.SavePreparedPackage(
httptest.NewRequest("POST", "https://update.ymhut.cn/api/admin/releases/packages", nil),
UploadedPackageFile{TempPath: temp, Size: int64(len("new package")), SHA256: "abc123"},
UploadOptions{FileName: name, UpdateManifest: true},
"admin")
if !errors.Is(err, ErrUploadedPackageIndexFailed) {
t.Fatalf("got %v, want index failure", err)
}
if data, readErr := os.ReadFile(target); readErr != nil || string(data) != "old package" {
t.Fatalf("package rollback data=%q err=%v", data, readErr)
}
if data, readErr := os.ReadFile(manifestPath); readErr != nil || string(data) != string(oldManifest) {
t.Fatalf("manifest rollback data=%q err=%v", data, readErr)
}
}
func TestSaveUploadedPackageWritesFileAndUpdatesManifest(t *testing.T) {
dir := t.TempDir()
cfg := &config.Config{