diff --git a/internal/updater/apply.go b/internal/updater/apply.go new file mode 100644 index 0000000..9a4c549 --- /dev/null +++ b/internal/updater/apply.go @@ -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 +} diff --git a/internal/updater/check.go b/internal/updater/check.go new file mode 100644 index 0000000..a68ed27 --- /dev/null +++ b/internal/updater/check.go @@ -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 +} diff --git a/internal/updater/release.go b/internal/updater/release.go new file mode 100644 index 0000000..252a986 --- /dev/null +++ b/internal/updater/release.go @@ -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 +} diff --git a/internal/updater/updater.go b/internal/updater/updater.go new file mode 100644 index 0000000..0a70aee --- /dev/null +++ b/internal/updater/updater.go @@ -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 +} diff --git a/internal/updater/version.go b/internal/updater/version.go new file mode 100644 index 0000000..b1ee293 --- /dev/null +++ b/internal/updater/version.go @@ -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 +}