feat: add updater package for in-place self-update
This commit is contained in:
@@ -0,0 +1,208 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
"github.com/mathismqn/godeez/internal/buildinfo"
|
||||
"github.com/mathismqn/godeez/internal/fileutil"
|
||||
)
|
||||
|
||||
var managedPrefixes = []string{
|
||||
"/nix/store",
|
||||
"/opt/homebrew",
|
||||
"/usr/local/Cellar",
|
||||
"/home/linuxbrew",
|
||||
"/snap",
|
||||
"/var/lib/flatpak",
|
||||
}
|
||||
|
||||
func resolveTarget() (string, error) {
|
||||
if buildinfo.IsDev() {
|
||||
return "", fmt.Errorf("development build cannot self-update. Install a release from https://github.com/%s/%s/releases",
|
||||
repoOwner, repoName)
|
||||
}
|
||||
|
||||
exe, err := os.Executable()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to locate the running binary: %w", err)
|
||||
}
|
||||
|
||||
target := exe
|
||||
if resolved, err := filepath.EvalSymlinks(exe); err == nil {
|
||||
target = resolved
|
||||
}
|
||||
|
||||
for _, prefix := range managedPrefixes {
|
||||
if strings.HasPrefix(target, prefix) {
|
||||
return "", fmt.Errorf("%s was installed by a package manager. Update it with that instead", target)
|
||||
}
|
||||
}
|
||||
|
||||
return target, nil
|
||||
}
|
||||
|
||||
func CheckUpdatable() error {
|
||||
_, err := resolveTarget()
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func checkWritable(dir string) error {
|
||||
f, err := os.CreateTemp(dir, tmpPattern)
|
||||
if err != nil {
|
||||
hint := "Re-run with sudo"
|
||||
if runtime.GOOS == "windows" {
|
||||
hint = "Re-run from an elevated prompt"
|
||||
}
|
||||
|
||||
return fmt.Errorf("cannot write to %s: %w. %s", dir, err, hint)
|
||||
}
|
||||
|
||||
name := f.Name()
|
||||
f.Close()
|
||||
os.Remove(name)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *Updater) Apply(ctx context.Context, release *Release) error {
|
||||
target, err := resolveTarget()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
dir := filepath.Dir(target)
|
||||
if err := checkWritable(dir); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
asset, err := release.assetForRuntime()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
want, err := u.fetchChecksum(ctx, release, asset.Name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
u.step("Downloading %s", asset.Name)
|
||||
tmp, sum, err := u.download(ctx, dir, asset)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer fileutil.DeleteFile(tmp)
|
||||
|
||||
u.step("Verifying checksum")
|
||||
if sum != want {
|
||||
return fmt.Errorf("checksum mismatch for %s: expected %s, got %s", asset.Name, want, sum)
|
||||
}
|
||||
|
||||
if err := os.Chmod(tmp, 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
u.step("Replacing %s", target)
|
||||
|
||||
return replaceBinary(target, tmp)
|
||||
}
|
||||
|
||||
func (u *Updater) fetchChecksum(ctx context.Context, release *Release, assetName string) (string, error) {
|
||||
asset, ok := release.asset(checksumsAsset)
|
||||
if !ok {
|
||||
return "", fmt.Errorf("release %s does not publish %s", release.TagName, checksumsAsset)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, apiTimeout)
|
||||
defer cancel()
|
||||
|
||||
body, err := u.get(ctx, asset.URL, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer body.Close()
|
||||
|
||||
return parseChecksums(io.LimitReader(body, maxResponseSize), assetName)
|
||||
}
|
||||
|
||||
func parseChecksums(r io.Reader, name string) (string, error) {
|
||||
scanner := bufio.NewScanner(r)
|
||||
for scanner.Scan() {
|
||||
fields := strings.Fields(scanner.Text())
|
||||
if len(fields) != 2 {
|
||||
continue
|
||||
}
|
||||
if strings.TrimPrefix(fields[1], "*") == name {
|
||||
return strings.ToLower(fields[0]), nil
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("no checksum listed for %s", name)
|
||||
}
|
||||
|
||||
func (u *Updater) download(ctx context.Context, dir string, asset Asset) (string, string, error) {
|
||||
body, err := u.get(ctx, asset.URL, nil)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
defer body.Close()
|
||||
|
||||
f, err := os.CreateTemp(dir, tmpPattern)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
tmp := f.Name()
|
||||
|
||||
hash := sha256.New()
|
||||
if _, err := io.Copy(io.MultiWriter(f, hash), body); err != nil {
|
||||
f.Close()
|
||||
fileutil.DeleteFile(tmp)
|
||||
|
||||
return "", "", fmt.Errorf("failed to download %s: %w", asset.Name, err)
|
||||
}
|
||||
if err := f.Close(); err != nil {
|
||||
fileutil.DeleteFile(tmp)
|
||||
|
||||
return "", "", err
|
||||
}
|
||||
|
||||
return tmp, hex.EncodeToString(hash.Sum(nil)), nil
|
||||
}
|
||||
|
||||
func replaceBinary(target, tmp string) error {
|
||||
if runtime.GOOS != "windows" {
|
||||
return os.Rename(tmp, target)
|
||||
}
|
||||
|
||||
old := target + ".old"
|
||||
os.Remove(old)
|
||||
|
||||
if err := os.Rename(target, old); err != nil {
|
||||
return fmt.Errorf("failed to move the current binary aside: %w", err)
|
||||
}
|
||||
|
||||
if err := os.Rename(tmp, target); err != nil {
|
||||
if rollbackErr := os.Rename(old, target); rollbackErr != nil {
|
||||
return fmt.Errorf("failed to install the new binary: %w. The previous one could not be restored from %s: %v",
|
||||
err, old, rollbackErr)
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to install the new binary: %w", err)
|
||||
}
|
||||
|
||||
os.Remove(old)
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/mathismqn/godeez/internal/buildinfo"
|
||||
"github.com/mathismqn/godeez/internal/fileutil"
|
||||
)
|
||||
|
||||
const noCheckEnv = "GODEEZ_NO_UPDATE_CHECK"
|
||||
|
||||
const (
|
||||
cacheTTL = 24 * time.Hour
|
||||
checkTimeout = 3 * time.Second
|
||||
)
|
||||
|
||||
type cacheEntry struct {
|
||||
CheckedAt time.Time `json:"checked_at"`
|
||||
LatestVersion string `json:"latest_version"`
|
||||
}
|
||||
|
||||
func cachePath() (string, error) {
|
||||
dir, err := os.UserCacheDir()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return filepath.Join(dir, "godeez", "update.json"), nil
|
||||
}
|
||||
|
||||
func readCache() (cacheEntry, bool) {
|
||||
path, err := cachePath()
|
||||
if err != nil {
|
||||
return cacheEntry{}, false
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return cacheEntry{}, false
|
||||
}
|
||||
|
||||
var entry cacheEntry
|
||||
if err := json.Unmarshal(data, &entry); err != nil {
|
||||
return cacheEntry{}, false
|
||||
}
|
||||
if entry.LatestVersion == "" || time.Since(entry.CheckedAt) > cacheTTL {
|
||||
return cacheEntry{}, false
|
||||
}
|
||||
|
||||
return entry, true
|
||||
}
|
||||
|
||||
func writeCache(version string) error {
|
||||
path, err := cachePath()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := fileutil.EnsureDir(filepath.Dir(path)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
data, err := json.Marshal(cacheEntry{CheckedAt: time.Now(), LatestVersion: version})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return os.WriteFile(path, data, 0644)
|
||||
}
|
||||
|
||||
func check(ctx context.Context) (string, error) {
|
||||
if entry, ok := readCache(); ok {
|
||||
return newerThanCurrent(entry.LatestVersion), nil
|
||||
}
|
||||
|
||||
release, err := New().Latest(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
latest := release.Version()
|
||||
_ = writeCache(latest)
|
||||
|
||||
return newerThanCurrent(latest), nil
|
||||
}
|
||||
|
||||
func newerThanCurrent(latest string) string {
|
||||
if IsNewer(buildinfo.Version(), latest) {
|
||||
return latest
|
||||
}
|
||||
|
||||
return ""
|
||||
}
|
||||
|
||||
func StartCheck(ctx context.Context) <-chan string {
|
||||
ch := make(chan string, 1)
|
||||
|
||||
if os.Getenv(noCheckEnv) != "" || buildinfo.IsDev() {
|
||||
close(ch)
|
||||
return ch
|
||||
}
|
||||
|
||||
go func() {
|
||||
defer close(ch)
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, checkTimeout)
|
||||
defer cancel()
|
||||
|
||||
if latest, err := check(ctx); err == nil && latest != "" {
|
||||
ch <- latest
|
||||
}
|
||||
}()
|
||||
|
||||
return ch
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"runtime"
|
||||
)
|
||||
|
||||
const (
|
||||
repoOwner = "mathismqn"
|
||||
repoName = "godeez"
|
||||
latestReleaseURL = "https://api.github.com/repos/" + repoOwner + "/" + repoName + "/releases/latest"
|
||||
checksumsAsset = "checksums.txt"
|
||||
maxResponseSize = 1 << 20
|
||||
)
|
||||
|
||||
var githubAPIHeaders = map[string]string{
|
||||
"Accept": "application/vnd.github+json",
|
||||
"X-GitHub-Api-Version": "2022-11-28",
|
||||
}
|
||||
|
||||
type Release struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Assets []Asset `json:"assets"`
|
||||
}
|
||||
|
||||
type Asset struct {
|
||||
Name string `json:"name"`
|
||||
URL string `json:"browser_download_url"`
|
||||
}
|
||||
|
||||
func (u *Updater) Latest(ctx context.Context) (*Release, error) {
|
||||
ctx, cancel := context.WithTimeout(ctx, apiTimeout)
|
||||
defer cancel()
|
||||
|
||||
body, err := u.get(ctx, latestReleaseURL, githubAPIHeaders)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer body.Close()
|
||||
|
||||
var release Release
|
||||
if err := json.NewDecoder(io.LimitReader(body, maxResponseSize)).Decode(&release); err != nil {
|
||||
return nil, fmt.Errorf("failed to decode release: %w", err)
|
||||
}
|
||||
if release.TagName == "" {
|
||||
return nil, fmt.Errorf("release has no tag name")
|
||||
}
|
||||
|
||||
return &release, nil
|
||||
}
|
||||
|
||||
func (r *Release) Version() string {
|
||||
return trimV(r.TagName)
|
||||
}
|
||||
|
||||
func (r *Release) asset(name string) (Asset, bool) {
|
||||
for _, a := range r.Assets {
|
||||
if a.Name == name {
|
||||
return a, true
|
||||
}
|
||||
}
|
||||
|
||||
return Asset{}, false
|
||||
}
|
||||
|
||||
func (r *Release) assetForRuntime() (Asset, error) {
|
||||
name := fmt.Sprintf("%s_%s_%s_%s", repoName, r.Version(), runtime.GOOS, runtime.GOARCH)
|
||||
if runtime.GOOS == "windows" {
|
||||
name += ".exe"
|
||||
}
|
||||
|
||||
asset, ok := r.asset(name)
|
||||
if !ok {
|
||||
return Asset{}, fmt.Errorf("release %s has no binary for %s/%s (expected %s)",
|
||||
r.TagName, runtime.GOOS, runtime.GOARCH, name)
|
||||
}
|
||||
|
||||
return asset, nil
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/mathismqn/godeez/internal/buildinfo"
|
||||
)
|
||||
|
||||
const (
|
||||
apiTimeout = 30 * time.Second
|
||||
tmpPattern = ".godeez-update-*"
|
||||
)
|
||||
|
||||
type Updater struct {
|
||||
client *http.Client
|
||||
Out io.Writer
|
||||
}
|
||||
|
||||
func New() *Updater {
|
||||
return &Updater{
|
||||
client: &http.Client{},
|
||||
Out: io.Discard,
|
||||
}
|
||||
}
|
||||
|
||||
func (u *Updater) step(format string, args ...any) {
|
||||
fmt.Fprintf(u.Out, format+"...\n", args...)
|
||||
}
|
||||
|
||||
func (u *Updater) get(ctx context.Context, url string, headers map[string]string) (io.ReadCloser, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("User-Agent", buildinfo.UserAgent())
|
||||
for name, value := range headers {
|
||||
req.Header.Set(name, value)
|
||||
}
|
||||
|
||||
resp, err := u.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
resp.Body.Close()
|
||||
|
||||
return nil, fmt.Errorf("unexpected status code %d from %s", resp.StatusCode, url)
|
||||
}
|
||||
|
||||
return resp.Body, nil
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package updater
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"golang.org/x/mod/semver"
|
||||
)
|
||||
|
||||
func trimV(v string) string {
|
||||
return strings.TrimPrefix(strings.TrimSpace(v), "v")
|
||||
}
|
||||
|
||||
func canonical(v string) string {
|
||||
v = strings.TrimSpace(v)
|
||||
if v == "" {
|
||||
return ""
|
||||
}
|
||||
if !strings.HasPrefix(v, "v") {
|
||||
v = "v" + v
|
||||
}
|
||||
if !semver.IsValid(v) {
|
||||
return ""
|
||||
}
|
||||
|
||||
return v
|
||||
}
|
||||
|
||||
func IsNewer(current, latest string) bool {
|
||||
c, l := canonical(current), canonical(latest)
|
||||
if c == "" || l == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
return semver.Compare(l, c) > 0
|
||||
}
|
||||
Reference in New Issue
Block a user