feat: add updater package for in-place self-update

This commit is contained in:
Mathis Maquenne
2026-07-30 11:30:08 +02:00
parent 0e0cbdad5e
commit eb93d01c19
5 changed files with 498 additions and 0 deletions
+208
View File
@@ -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
}
+118
View File
@@ -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
}
+82
View File
@@ -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
}
+55
View File
@@ -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
}
+35
View File
@@ -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
}