From 408b3ea06febf534b6683849fe3794ed3ad79f88 Mon Sep 17 00:00:00 2001 From: Rahul Bansal Date: Sat, 21 Feb 2026 11:09:26 +0530 Subject: [PATCH] feat: add self-update command with periodic update hints --- cmd/picoclaw/cmd_agent.go | 20 ++ cmd/picoclaw/cmd_update.go | 82 +++++++ cmd/picoclaw/main.go | 3 + pkg/update/update.go | 442 +++++++++++++++++++++++++++++++++++++ pkg/update/update_test.go | 298 +++++++++++++++++++++++++ 5 files changed, 845 insertions(+) create mode 100644 cmd/picoclaw/cmd_update.go create mode 100644 pkg/update/update.go create mode 100644 pkg/update/update_test.go diff --git a/cmd/picoclaw/cmd_agent.go b/cmd/picoclaw/cmd_agent.go index 69584c961..712c82029 100644 --- a/cmd/picoclaw/cmd_agent.go +++ b/cmd/picoclaw/cmd_agent.go @@ -19,6 +19,7 @@ import ( "github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/providers" "github.com/sipeed/picoclaw/pkg/tools" + "github.com/sipeed/picoclaw/pkg/update" ) func agentCmd() { @@ -66,6 +67,12 @@ func agentCmd() { } } + // Check for updates in the background (non-blocking, at most once per 24h) + updateHint := make(chan string, 1) + go func() { + updateHint <- update.CheckHint(version) + }() + cfg, err := loadConfig() if err != nil { fmt.Printf("Error loading config: %v\n", err) @@ -109,6 +116,17 @@ func agentCmd() { "skills_available": startupInfo["skills"].(map[string]any)["available"], }) + // Print update hint if available (non-blocking select) + printUpdateHint := func() { + select { + case v := <-updateHint: + if v != "" { + fmt.Printf("\nA new version (%s) is available. Run 'picoclaw update' to upgrade.\n\n", v) + } + default: + } + } + if message != "" { ctx := context.Background() spin := newSpinner("Thinking...") @@ -120,7 +138,9 @@ func agentCmd() { os.Exit(1) } fmt.Printf("\n%s %s\n", logo, response) + printUpdateHint() } else { + printUpdateHint() fmt.Printf("%s Interactive mode (Ctrl+C to exit)\n\n", logo) interactiveMode(agentLoop, sessionKey) } diff --git a/cmd/picoclaw/cmd_update.go b/cmd/picoclaw/cmd_update.go new file mode 100644 index 000000000..2b4fe4c10 --- /dev/null +++ b/cmd/picoclaw/cmd_update.go @@ -0,0 +1,82 @@ +// PicoClaw - Ultra-lightweight personal AI agent +// License: MIT + +package main + +import ( + "fmt" + "os" + + "github.com/sipeed/picoclaw/pkg/update" +) + +func updateCmd() { + // Handle flags + checkOnly := false + for _, arg := range os.Args[2:] { + switch arg { + case "--help", "-h": + updateHelp() + return + case "--check", "-c": + checkOnly = true + } + } + + fmt.Printf("%s Checking for updates...\n", logo) + + release, err := update.CheckLatest() + if err != nil { + fmt.Printf("Error checking for updates: %v\n", err) + os.Exit(1) + } + + latestVersion := release.TagName + + if !update.IsNewer(version, latestVersion) { + fmt.Printf("Already up to date! (current: %s, latest: %s)\n", version, latestVersion) + return + } + + fmt.Printf("New version available: %s (current: %s)\n", latestVersion, version) + fmt.Printf("Release: %s\n", release.HTMLURL) + + if checkOnly { + return + } + + assetURL, err := update.FindAssetURL(release) + if err != nil { + fmt.Printf("Error: %v\n", err) + fmt.Printf("You can download manually from: %s\n", release.HTMLURL) + os.Exit(1) + } + + fmt.Printf("\nUpdate to %s? (y/n): ", latestVersion) + var confirm string + fmt.Scanln(&confirm) + + if confirm != "y" && confirm != "Y" { + fmt.Println("Update cancelled.") + return + } + + if err := update.DownloadAndReplace(assetURL, os.Stdout); err != nil { + fmt.Printf("Update failed: %v\n", err) + os.Exit(1) + } + + fmt.Printf("\n%s Updated to %s!\n", logo, latestVersion) +} + +func updateHelp() { + fmt.Println("Check for updates and self-update picoclaw") + fmt.Println() + fmt.Println("Usage: picoclaw update [flags]") + fmt.Println() + fmt.Println("Checks GitHub for the latest release and offers to update in-place.") + fmt.Println() + fmt.Println("Flags:") + fmt.Println(" -c, --check Check only, don't prompt to update") + fmt.Println(" -h, --help Show this help") +} diff --git a/cmd/picoclaw/main.go b/cmd/picoclaw/main.go index 9d6b0490e..f8faf3377 100644 --- a/cmd/picoclaw/main.go +++ b/cmd/picoclaw/main.go @@ -120,6 +120,8 @@ func main() { sessionsCmd() case "skills": skillsCmd() + case "update": + updateCmd() case "version", "--version", "-v": printVersion() case "--help", "-h": @@ -146,6 +148,7 @@ func printHelp() { fmt.Println(" skills Manage skills (install, list, remove)") fmt.Println(" sessions Manage sessions (list, show, delete, clear)") fmt.Println(" doctor Diagnose and fix common problems (--fix to auto-repair)") + fmt.Println(" update Check for updates and self-update") fmt.Println(" version Show version information") } diff --git a/pkg/update/update.go b/pkg/update/update.go new file mode 100644 index 000000000..562c84e78 --- /dev/null +++ b/pkg/update/update.go @@ -0,0 +1,442 @@ +// PicoClaw - Ultra-lightweight personal AI agent +// License: MIT + +package update + +import ( + "archive/tar" + "archive/zip" + "compress/gzip" + "encoding/json" + "fmt" + "io" + "net/http" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "time" +) + +const ( + // GitHubRepo is the upstream repository for release checks + GitHubRepo = "sipeed/picoclaw" + + // releaseAPIURL is the GitHub API endpoint for the latest release + releaseAPIURL = "https://api.github.com/repos/" + GitHubRepo + "/releases/latest" + + // checkInterval is how often the periodic hint checks for updates + checkInterval = 24 * time.Hour + + // httpTimeout is the timeout for GitHub API and download requests + httpTimeout = 30 * time.Second +) + +// httpClient is the HTTP client used for all requests. Tests can replace it. +var httpClient = &http.Client{Timeout: httpTimeout} + +// ReleaseInfo holds information about a GitHub release +type ReleaseInfo struct { + TagName string `json:"tag_name"` + Assets []ReleaseAsset `json:"assets"` + HTMLURL string `json:"html_url"` +} + +// ReleaseAsset holds information about a release asset +type ReleaseAsset struct { + Name string `json:"name"` + BrowserDownloadURL string `json:"browser_download_url"` +} + +// checkCache stores the last update check result +type checkCache struct { + LastCheck time.Time `json:"last_check"` + LatestVersion string `json:"latest_version"` + HTMLURL string `json:"html_url"` +} + +// CheckLatest fetches the latest release info from GitHub +func CheckLatest() (*ReleaseInfo, error) { + req, err := http.NewRequest("GET", releaseAPIURL, nil) + if err != nil { + return nil, fmt.Errorf("creating request: %w", err) + } + req.Header.Set("Accept", "application/vnd.github.v3+json") + req.Header.Set("User-Agent", "picoclaw-updater") + + resp, err := httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("fetching latest release: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("GitHub API returned HTTP %d", resp.StatusCode) + } + + var release ReleaseInfo + if err := json.NewDecoder(resp.Body).Decode(&release); err != nil { + return nil, fmt.Errorf("parsing release info: %w", err) + } + + return &release, nil +} + +// IsNewer returns true if the remote version is newer than the current version. +// Both versions should be semver strings, optionally prefixed with "v". +func IsNewer(current, remote string) bool { + curParts := parseSemver(current) + remParts := parseSemver(remote) + + if curParts == nil || remParts == nil { + return false + } + + for i := 0; i < 3; i++ { + if remParts[i] > curParts[i] { + return true + } + if remParts[i] < curParts[i] { + return false + } + } + return false +} + +// parseSemver extracts [major, minor, patch] from a version string. +// Accepts "v1.2.3", "1.2.3", "v0.1.2-42-gabcdef", etc. +func parseSemver(v string) []int { + v = strings.TrimPrefix(v, "v") + + // Strip anything after a hyphen (pre-release/build metadata) + if idx := strings.Index(v, "-"); idx != -1 { + v = v[:idx] + } + + parts := strings.Split(v, ".") + if len(parts) != 3 { + return nil + } + + nums := make([]int, 3) + for i, p := range parts { + n, err := strconv.Atoi(p) + if err != nil { + return nil + } + nums[i] = n + } + return nums +} + +// AssetName returns the expected asset filename for the current platform. +// Asset naming convention: picoclaw_{OS}_{Arch}.tar.gz (or .zip for Windows) +func AssetName() string { + osName := normalizeOS(runtime.GOOS) + archName := normalizeArch(runtime.GOARCH) + + ext := "tar.gz" + if runtime.GOOS == "windows" { + ext = "zip" + } + return fmt.Sprintf("picoclaw_%s_%s.%s", osName, archName, ext) +} + +func normalizeOS(goos string) string { + switch goos { + case "darwin": + return "Darwin" + case "linux": + return "Linux" + case "windows": + return "Windows" + case "freebsd": + return "Freebsd" + default: + // Title case the first letter + return strings.ToUpper(goos[:1]) + goos[1:] + } +} + +func normalizeArch(goarch string) string { + switch goarch { + case "amd64": + return "x86_64" + case "arm64": + return "arm64" + case "arm": + return "armv6" + case "riscv64": + return "riscv64" + case "mips64": + return "mips64" + case "s390x": + return "s390x" + default: + return goarch + } +} + +// FindAssetURL finds the download URL for the current platform in a release +func FindAssetURL(release *ReleaseInfo) (string, error) { + want := AssetName() + for _, asset := range release.Assets { + if asset.Name == want { + return asset.BrowserDownloadURL, nil + } + } + return "", fmt.Errorf("no release asset found for %s", want) +} + +// DownloadAndReplace downloads the release asset and replaces the current binary. +// It writes progress to the provided writer. +func DownloadAndReplace(downloadURL string, progress io.Writer) error { + // 1. Download the archive + fmt.Fprintf(progress, "Downloading %s...\n", filepath.Base(downloadURL)) + resp, err := httpClient.Get(downloadURL) + if err != nil { + return fmt.Errorf("downloading release: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("download returned HTTP %d", resp.StatusCode) + } + + tmpDir, err := os.MkdirTemp("", "picoclaw-update-*") + if err != nil { + return fmt.Errorf("creating temp directory: %w", err) + } + defer os.RemoveAll(tmpDir) + + archivePath := filepath.Join(tmpDir, filepath.Base(downloadURL)) + archiveFile, err := os.Create(archivePath) + if err != nil { + return fmt.Errorf("creating temp file: %w", err) + } + + written, err := io.Copy(archiveFile, resp.Body) + archiveFile.Close() + if err != nil { + return fmt.Errorf("saving download: %w", err) + } + fmt.Fprintf(progress, "Downloaded %.1f MB\n", float64(written)/1024/1024) + + // 2. Extract the binary + fmt.Fprintf(progress, "Extracting...\n") + binaryPath, err := extractBinary(archivePath, tmpDir) + if err != nil { + return fmt.Errorf("extracting archive: %w", err) + } + + // 3. Replace the current binary + currentBinary, err := os.Executable() + if err != nil { + return fmt.Errorf("finding current binary: %w", err) + } + currentBinary, err = filepath.EvalSymlinks(currentBinary) + if err != nil { + return fmt.Errorf("resolving binary path: %w", err) + } + + fmt.Fprintf(progress, "Replacing %s...\n", currentBinary) + if err := replaceBinary(binaryPath, currentBinary); err != nil { + return fmt.Errorf("replacing binary: %w", err) + } + + // 4. Recreate pico symlink if it exists next to the binary + picoSymlink := filepath.Join(filepath.Dir(currentBinary), "pico") + if target, err := os.Readlink(picoSymlink); err == nil { + // Only recreate if it was pointing at a picoclaw binary + if strings.Contains(target, "picoclaw") { + os.Remove(picoSymlink) + os.Symlink(filepath.Base(currentBinary), picoSymlink) + } + } + + return nil +} + +// extractBinary extracts the picoclaw binary from a tar.gz or zip archive +func extractBinary(archivePath, destDir string) (string, error) { + if strings.HasSuffix(archivePath, ".zip") { + return extractFromZip(archivePath, destDir) + } + return extractFromTarGz(archivePath, destDir) +} + +func extractFromTarGz(archivePath, destDir string) (string, error) { + f, err := os.Open(archivePath) + if err != nil { + return "", err + } + defer f.Close() + + gzr, err := gzip.NewReader(f) + if err != nil { + return "", fmt.Errorf("opening gzip: %w", err) + } + defer gzr.Close() + + tr := tar.NewReader(gzr) + for { + header, err := tr.Next() + if err == io.EOF { + break + } + if err != nil { + return "", fmt.Errorf("reading tar: %w", err) + } + + name := filepath.Base(header.Name) + if name == "picoclaw" || name == "picoclaw.exe" { + outPath := filepath.Join(destDir, name) + outFile, err := os.OpenFile(outPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o755) + if err != nil { + return "", err + } + if _, err := io.Copy(outFile, tr); err != nil { + outFile.Close() + return "", err + } + outFile.Close() + return outPath, nil + } + } + return "", fmt.Errorf("picoclaw binary not found in archive") +} + +func extractFromZip(archivePath, destDir string) (string, error) { + r, err := zip.OpenReader(archivePath) + if err != nil { + return "", err + } + defer r.Close() + + for _, f := range r.File { + name := filepath.Base(f.Name) + if name == "picoclaw" || name == "picoclaw.exe" { + rc, err := f.Open() + if err != nil { + return "", err + } + outPath := filepath.Join(destDir, name) + outFile, err := os.OpenFile(outPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o755) + if err != nil { + rc.Close() + return "", err + } + if _, err := io.Copy(outFile, rc); err != nil { + outFile.Close() + rc.Close() + return "", err + } + outFile.Close() + rc.Close() + return outPath, nil + } + } + return "", fmt.Errorf("picoclaw binary not found in archive") +} + +// replaceBinary atomically replaces the old binary with the new one +func replaceBinary(newPath, oldPath string) error { + // Rename old binary to .bak + bakPath := oldPath + ".bak" + if err := os.Rename(oldPath, bakPath); err != nil { + return fmt.Errorf("backing up current binary: %w", err) + } + + // Copy new binary into place (can't rename across filesystems) + src, err := os.Open(newPath) + if err != nil { + // Restore backup + os.Rename(bakPath, oldPath) + return fmt.Errorf("opening new binary: %w", err) + } + defer src.Close() + + dst, err := os.OpenFile(oldPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o755) + if err != nil { + os.Rename(bakPath, oldPath) + return fmt.Errorf("writing new binary: %w", err) + } + + if _, err := io.Copy(dst, src); err != nil { + dst.Close() + os.Rename(bakPath, oldPath) + return fmt.Errorf("copying new binary: %w", err) + } + dst.Close() + + // Remove backup + os.Remove(bakPath) + return nil +} + +// --- Periodic hint support --- + +// cacheFilePath returns the path to the update check cache file. +// Tests can override this variable. +var cacheFilePath = func() string { + home, _ := os.UserHomeDir() + return filepath.Join(home, ".picoclaw", "last_update_check.json") +} + +// CheckHint checks if a newer version is available, using a cached result +// if the last check was within the check interval. Returns the latest version +// string if an update is available, or empty string if up-to-date or on error. +// This is designed to be called from a goroutine and never block. +func CheckHint(currentVersion string) string { + // 1. Try to read cache + cache, err := loadCache() + if err == nil && time.Since(cache.LastCheck) < checkInterval { + // Cache is fresh, use it + if IsNewer(currentVersion, cache.LatestVersion) { + return cache.LatestVersion + } + return "" + } + + // 2. Cache is stale or missing, check GitHub + release, err := CheckLatest() + if err != nil { + return "" + } + + // 3. Update cache + newCache := checkCache{ + LastCheck: time.Now(), + LatestVersion: release.TagName, + HTMLURL: release.HTMLURL, + } + saveCache(&newCache) + + if IsNewer(currentVersion, release.TagName) { + return release.TagName + } + return "" +} + +func loadCache() (*checkCache, error) { + data, err := os.ReadFile(cacheFilePath()) + if err != nil { + return nil, err + } + var cache checkCache + if err := json.Unmarshal(data, &cache); err != nil { + return nil, err + } + return &cache, nil +} + +func saveCache(cache *checkCache) { + data, err := json.Marshal(cache) + if err != nil { + return + } + dir := filepath.Dir(cacheFilePath()) + os.MkdirAll(dir, 0o755) + os.WriteFile(cacheFilePath(), data, 0o644) +} diff --git a/pkg/update/update_test.go b/pkg/update/update_test.go new file mode 100644 index 000000000..d953c6bd3 --- /dev/null +++ b/pkg/update/update_test.go @@ -0,0 +1,298 @@ +package update + +import ( + "archive/tar" + "compress/gzip" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "testing" + "time" +) + +func TestParseSemver(t *testing.T) { + tests := []struct { + name string + input string + want []int + }{ + {name: "plain", input: "1.2.3", want: []int{1, 2, 3}}, + {name: "v-prefix", input: "v0.1.2", want: []int{0, 1, 2}}, + {name: "with-prerelease", input: "v0.1.2-42-gabcdef", want: []int{0, 1, 2}}, + {name: "with-dirty", input: "v0.1.2-dirty", want: []int{0, 1, 2}}, + {name: "invalid-empty", input: "", want: nil}, + {name: "invalid-two-parts", input: "1.2", want: nil}, + {name: "invalid-letters", input: "v1.x.3", want: nil}, + {name: "dev", input: "dev", want: nil}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := parseSemver(tt.input) + if tt.want == nil { + if got != nil { + t.Errorf("parseSemver(%q) = %v, want nil", tt.input, got) + } + return + } + if got == nil { + t.Fatalf("parseSemver(%q) = nil, want %v", tt.input, tt.want) + } + for i := 0; i < 3; i++ { + if got[i] != tt.want[i] { + t.Errorf("parseSemver(%q)[%d] = %d, want %d", tt.input, i, got[i], tt.want[i]) + } + } + }) + } +} + +func TestIsNewer(t *testing.T) { + tests := []struct { + name string + current string + remote string + want bool + }{ + {name: "newer-patch", current: "v0.1.2", remote: "v0.1.3", want: true}, + {name: "newer-minor", current: "v0.1.2", remote: "v0.2.0", want: true}, + {name: "newer-major", current: "v0.1.2", remote: "v1.0.0", want: true}, + {name: "same", current: "v0.1.2", remote: "v0.1.2", want: false}, + {name: "older", current: "v0.2.0", remote: "v0.1.9", want: false}, + {name: "current-with-metadata", current: "v0.1.2-42-gabcdef", remote: "v0.1.3", want: true}, + {name: "same-with-metadata", current: "v0.1.2-42-gabcdef", remote: "v0.1.2", want: false}, + {name: "invalid-current", current: "dev", remote: "v0.1.2", want: false}, + {name: "invalid-remote", current: "v0.1.2", remote: "dev", want: false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := IsNewer(tt.current, tt.remote) + if got != tt.want { + t.Errorf("IsNewer(%q, %q) = %v, want %v", tt.current, tt.remote, got, tt.want) + } + }) + } +} + +func TestNormalizeOS(t *testing.T) { + tests := []struct { + input string + want string + }{ + {"darwin", "Darwin"}, + {"linux", "Linux"}, + {"windows", "Windows"}, + {"freebsd", "Freebsd"}, + } + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + got := normalizeOS(tt.input) + if got != tt.want { + t.Errorf("normalizeOS(%q) = %q, want %q", tt.input, got, tt.want) + } + }) + } +} + +func TestNormalizeArch(t *testing.T) { + tests := []struct { + input string + want string + }{ + {"amd64", "x86_64"}, + {"arm64", "arm64"}, + {"arm", "armv6"}, + {"riscv64", "riscv64"}, + } + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + got := normalizeArch(tt.input) + if got != tt.want { + t.Errorf("normalizeArch(%q) = %q, want %q", tt.input, got, tt.want) + } + }) + } +} + +func TestCheckLatest(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + release := ReleaseInfo{ + TagName: "v1.0.0", + HTMLURL: "https://github.com/sipeed/picoclaw/releases/tag/v1.0.0", + Assets: []ReleaseAsset{ + {Name: "picoclaw_Darwin_arm64.tar.gz", BrowserDownloadURL: "https://example.com/picoclaw_Darwin_arm64.tar.gz"}, + }, + } + json.NewEncoder(w).Encode(release) + })) + defer server.Close() + + // Override the HTTP client to use test server + origClient := httpClient + httpClient = server.Client() + defer func() { httpClient = origClient }() + + // Override releaseAPIURL by testing CheckLatest indirectly via FindAssetURL + // For direct testing, we test the JSON parsing + resp, err := httpClient.Get(server.URL) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer resp.Body.Close() + + var release ReleaseInfo + if err := json.NewDecoder(resp.Body).Decode(&release); err != nil { + t.Fatalf("unexpected decode error: %v", err) + } + + if release.TagName != "v1.0.0" { + t.Errorf("TagName = %q, want %q", release.TagName, "v1.0.0") + } +} + +func TestFindAssetURL(t *testing.T) { + assetName := AssetName() + release := &ReleaseInfo{ + Assets: []ReleaseAsset{ + {Name: "picoclaw_Linux_x86_64.tar.gz", BrowserDownloadURL: "https://example.com/linux"}, + {Name: assetName, BrowserDownloadURL: "https://example.com/match"}, + {Name: "picoclaw_Windows_x86_64.zip", BrowserDownloadURL: "https://example.com/windows"}, + }, + } + + url, err := FindAssetURL(release) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if url != "https://example.com/match" { + t.Errorf("FindAssetURL() = %q, want %q", url, "https://example.com/match") + } + + // Test missing asset + release.Assets = []ReleaseAsset{ + {Name: "picoclaw_UnknownOS_unknownarch.tar.gz", BrowserDownloadURL: "https://example.com/nope"}, + } + _, err = FindAssetURL(release) + if err == nil { + t.Error("expected error for missing asset, got nil") + } +} + +func TestExtractFromTarGz(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "update-test-*") + if err != nil { + t.Fatalf("creating temp dir: %v", err) + } + defer os.RemoveAll(tmpDir) + + // Create a test tar.gz with a picoclaw binary + archivePath := filepath.Join(tmpDir, "test.tar.gz") + binaryContent := []byte("#!/bin/sh\necho hello\n") + + f, err := os.Create(archivePath) + if err != nil { + t.Fatalf("creating archive: %v", err) + } + gw := gzip.NewWriter(f) + tw := tar.NewWriter(gw) + + // Add a non-binary file first + if err := tw.WriteHeader(&tar.Header{Name: "README.md", Size: 6, Mode: 0o644}); err != nil { + t.Fatalf("writing tar header: %v", err) + } + tw.Write([]byte("readme")) + + // Add the binary + if err := tw.WriteHeader(&tar.Header{Name: "picoclaw", Size: int64(len(binaryContent)), Mode: 0o755}); err != nil { + t.Fatalf("writing tar header: %v", err) + } + tw.Write(binaryContent) + + tw.Close() + gw.Close() + f.Close() + + // Extract + destDir := filepath.Join(tmpDir, "extracted") + os.MkdirAll(destDir, 0o755) + + binPath, err := extractFromTarGz(archivePath, destDir) + if err != nil { + t.Fatalf("extractFromTarGz: %v", err) + } + + data, err := os.ReadFile(binPath) + if err != nil { + t.Fatalf("reading extracted binary: %v", err) + } + if string(data) != string(binaryContent) { + t.Errorf("extracted content = %q, want %q", string(data), string(binaryContent)) + } +} + +func TestReplaceBinary(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "replace-test-*") + if err != nil { + t.Fatalf("creating temp dir: %v", err) + } + defer os.RemoveAll(tmpDir) + + oldPath := filepath.Join(tmpDir, "picoclaw") + newPath := filepath.Join(tmpDir, "picoclaw-new") + + os.WriteFile(oldPath, []byte("old"), 0o755) + os.WriteFile(newPath, []byte("new"), 0o755) + + if err := replaceBinary(newPath, oldPath); err != nil { + t.Fatalf("replaceBinary: %v", err) + } + + data, err := os.ReadFile(oldPath) + if err != nil { + t.Fatalf("reading replaced binary: %v", err) + } + if string(data) != "new" { + t.Errorf("replaced binary content = %q, want %q", string(data), "new") + } + + // Backup should be cleaned up + if _, err := os.Stat(oldPath + ".bak"); !os.IsNotExist(err) { + t.Error("backup file was not cleaned up") + } +} + +func TestCacheRoundTrip(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "cache-test-*") + if err != nil { + t.Fatalf("creating temp dir: %v", err) + } + defer os.RemoveAll(tmpDir) + + // Override home dir for cache path + origCacheFilePath := cacheFilePath + cacheFilePath = func() string { + return filepath.Join(tmpDir, "last_update_check.json") + } + defer func() { cacheFilePath = origCacheFilePath }() + + cache := &checkCache{ + LastCheck: time.Now().Truncate(time.Second), + LatestVersion: "v1.2.3", + HTMLURL: "https://example.com", + } + + saveCache(cache) + + loaded, err := loadCache() + if err != nil { + t.Fatalf("loadCache: %v", err) + } + + if loaded.LatestVersion != cache.LatestVersion { + t.Errorf("LatestVersion = %q, want %q", loaded.LatestVersion, cache.LatestVersion) + } + if loaded.HTMLURL != cache.HTMLURL { + t.Errorf("HTMLURL = %q, want %q", loaded.HTMLURL, cache.HTMLURL) + } +}