diff --git a/README.md b/README.md index f1bce35..82e4914 100644 --- a/README.md +++ b/README.md @@ -2,6 +2,8 @@ Easy CLI tool for servers managed by FlyWP. +Conforms to the FlyWP monitoring agent contract v0.2.1. + ## Installation ### Prerequisites @@ -99,6 +101,17 @@ fly --domain example.com wp plugin list --format=json All arguments after the WP-CLI command (or after the command for `fly exec`) go to that command unchanged, flags included. Put `--domain` before the command. To pass a flag as the first argument, put `--` before it, for example `fly wp -- --info`. +### Monitoring agent + +`fly agent run` is the FlyWP monitoring agent. It runs all the time under systemd (`fly-agent.service`, as the server user, not root), and FlyWP installs it. Each minute it measures CPU, load, memory, swap, disk and network traffic, and it sends the values and the server status (restart needed, waiting updates, OS, kernel, uptime) to FlyWP. It keeps unsent data on disk for up to 24 hours. FlyWP can update and restart the agent through it, without SSH. The agent does not need Docker. + +It reads `FLY_AGENT_URL` (https), `FLY_AGENT_TOKEN` and `FLY_AGENT_SERVER_ID` from `/etc/fly/agent.env`, and keeps its state in `STATE_DIRECTORY` (`/var/lib/fly-agent`). + +```bash +systemctl status fly-agent # is the agent running? +journalctl -u fly-agent -f # the agent log +``` + ### Global Commands A few helper commands to debug the server installation and start/stop all sites. diff --git a/agent_test.go b/agent_test.go index 3ad68ac..0a21e3e 100644 --- a/agent_test.go +++ b/agent_test.go @@ -1,20 +1,30 @@ package main -// End-to-end tests of "fly agent run": the configuration errors, the lock and -// a clean stop. The loop itself is tested in internal/agent with a fake clock. +// End-to-end tests of "fly agent run": the configuration errors, the lock, a +// clean stop and a real self-update. The loop itself is tested in +// internal/agent with a fake clock. import ( + "archive/tar" "bytes" + "compress/gzip" + "crypto/sha256" + "encoding/hex" "encoding/json" + "fmt" "net/http" "net/http/httptest" "os" "os/exec" + "path/filepath" + "runtime" "strings" "sync" "syscall" "testing" "time" + + "github.com/flywp/server-cli/internal/release" ) const testToken = "flyagt_0123456789abcdefghijABCDEFGHIJ" @@ -123,6 +133,171 @@ func TestAgentSendsAgentStartedToTheControlPlane(t *testing.T) { } } +// updateServer is a control plane with one open agent.update command. The +// command stays open until an event finishes it. +type updateServer struct { + mu sync.Mutex + args map[string]string + events []map[string]any + closed bool +} + +func (s *updateServer) ServeHTTP(w http.ResponseWriter, r *http.Request) { + s.mu.Lock() + defer s.mu.Unlock() + + switch r.URL.Path { + case "/agent/v1/events": + var body struct{ Events []map[string]any } + _ = json.NewDecoder(r.Body).Decode(&body) + for _, e := range body.Events { + s.events = append(s.events, e) + if e["command_id"] == "01JBX0000000000000000000E1" { + s.closed = true + } + } + _, _ = w.Write([]byte(`{"accepted": 1}`)) + case "/agent/v1/commands": + commands := []any{} + if !s.closed { + commands = append(commands, map[string]any{"id": "01JBX0000000000000000000E1", "verb": "agent.update", "args": s.args, "issued_at": "2026-09-22T10:00:00Z"}) + } + _ = json.NewEncoder(w).Encode(map[string]any{"commands": commands}) + default: + _, _ = w.Write([]byte(`{"accepted": 1, "rejected": [], "report_interval": 1}`)) + } +} + +// result returns the result event of the update command, or nil. +func (s *updateServer) result() map[string]any { + s.mu.Lock() + defer s.mu.Unlock() + for _, e := range s.events { + if e["command_id"] == "01JBX0000000000000000000E1" { + return e + } + } + return nil +} + +func TestAgentUpdatesItself(t *testing.T) { + environ := agentEnv(t) + + // The agent replaces its own binary, so it runs from a copy. + dir := t.TempDir() + exe := filepath.Join(dir, "fly") + data, err := os.ReadFile(flyBin) + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(exe, data, 0o755); err != nil { + t.Fatal(err) + } + + // The new release: fly built as v9.9.9, in the archive layout of a release. + newBin := filepath.Join(t.TempDir(), "fly") + build := exec.Command("go", "build", "-ldflags", "-X github.com/flywp/server-cli/internal/version.Version=v9.9.9", "-o", newBin, ".") + if out, err := build.CombinedOutput(); err != nil { + t.Fatalf("building the new release: %v\n%s", err, out) + } + tarball := releaseArchive(t, newBin, release.BinaryName(runtime.GOOS, runtime.GOARCH)) + sum := sha256.Sum256(tarball) + + files := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = w.Write(tarball) })) + defer files.Close() + cp := &updateServer{args: map[string]string{"url": files.URL + "/fly.tar.gz", "version": "v9.9.9", "sha256": hex.EncodeToString(sum[:])}} + srv := httptest.NewServer(cp) + defer srv.Close() + + // Work 3 seconds from now, not at the second of server id 17. + environ = append(environ, "FLY_AGENT_URL="+srv.URL, fmt.Sprintf("FLY_AGENT_SERVER_ID=%d", (time.Now().Second()+3)%60)) + + first := exec.Command(exe, "agent", "run") + first.Env = environ + stderr := &lockedBuffer{} + first.Stderr = stderr + if err := first.Start(); err != nil { + t.Fatal(err) + } + done := make(chan error, 1) + go func() { done <- first.Wait() }() + select { + case err := <-done: + if err != nil { + t.Fatalf("the agent exit after the update: %v, want 0. stderr:\n%s", err, stderr) + } + case <-time.After(30 * time.Second): + _ = first.Process.Kill() + t.Fatalf("the agent did not exit for the update within 30s. stderr:\n%s", stderr) + } + + if out := runFlyAt(t, exe, "version"); !strings.Contains(out, "v9.9.9") { + t.Fatalf("fly version = %q, want the new release v9.9.9 on disk", out) + } + if cp.result() != nil { + t.Error("the old process sent the result; the new process must send it") + } + + // systemd starts the new binary. It sends the result at once. + second := exec.Command(exe, "agent", "run") + second.Env = environ + second.Stderr = &lockedBuffer{} + if err := second.Start(); err != nil { + t.Fatal(err) + } + defer func() { _ = second.Process.Signal(syscall.SIGTERM); _ = second.Wait() }() + + deadline := time.Now().Add(10 * time.Second) + for cp.result() == nil && time.Now().Before(deadline) { + time.Sleep(50 * time.Millisecond) + } + e := cp.result() + if e == nil || e["name"] != "command.completed" { + t.Fatalf("result = %v, want command.completed", e) + } + if data, _ := e["data"].(map[string]any); data["version"] != "v9.9.9" { + t.Errorf("result data = %v, want version v9.9.9", e["data"]) + } +} + +// runFlyAt runs the fly binary at exe and returns its stdout. +func runFlyAt(t *testing.T, exe string, args ...string) string { + t.Helper() + out, err := exec.Command(exe, args...).Output() + if err != nil { + t.Fatalf("%s %v: %v", exe, args, err) + } + return string(out) +} + +// releaseArchive returns a tar.gz archive that holds the file bin as name. +func releaseArchive(t *testing.T, bin, name string) []byte { + t.Helper() + + data, err := os.ReadFile(bin) + if err != nil { + t.Fatal(err) + } + + var buf bytes.Buffer + gz := gzip.NewWriter(&buf) + tw := tar.NewWriter(gz) + if err := tw.WriteHeader(&tar.Header{Name: name, Mode: 0o755, Size: int64(len(data)), Typeflag: tar.TypeReg}); err != nil { + t.Fatal(err) + } + if _, err := tw.Write(data); err != nil { + t.Fatal(err) + } + if err := tw.Close(); err != nil { + t.Fatal(err) + } + if err := gz.Close(); err != nil { + t.Fatal(err) + } + + return buf.Bytes() +} + func TestAgentConfigErrors(t *testing.T) { tests := []struct { name string diff --git a/cmd/version.go b/cmd/version.go index f0f46a6..bef2ca2 100644 --- a/cmd/version.go +++ b/cmd/version.go @@ -5,7 +5,7 @@ import ( "fmt" "os" - "github.com/flywp/server-cli/internal/utils" + "github.com/flywp/server-cli/internal/release" "github.com/flywp/server-cli/internal/version" "github.com/spf13/cobra" ) @@ -30,7 +30,7 @@ var updateCmd = &cobra.Command{ return errors.New("the update command must be run as root, please run 'sudo fly update'") } - update, err := utils.CheckForUpdates(cmd.Context()) + update, err := release.CheckForUpdates(cmd.Context()) if err != nil { return fmt.Errorf("checking for updates: %w", err) } @@ -58,7 +58,7 @@ var updateCmd = &cobra.Command{ } fmt.Println("Updating...") - if err := utils.SelfUpdate(cmd.Context(), update.Release); err != nil { + if err := release.SelfUpdate(cmd.Context(), update.Release); err != nil { return fmt.Errorf("updating: %w", err) } diff --git a/internal/agent/agent.go b/internal/agent/agent.go index bfbcc14..d86b246 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -5,10 +5,12 @@ import ( "errors" "io/fs" "log/slog" + "os" "path/filepath" "time" "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/release" "github.com/flywp/server-cli/internal/statefile" "github.com/flywp/server-cli/internal/version" "github.com/oklog/ulid/v2" @@ -25,6 +27,7 @@ const ( type ControlPlane interface { PostMetrics(ctx context.Context, req *wire.MetricsRequest) (*wire.MetricsReply, error) PostEvents(ctx context.Context, req *wire.EventsRequest) (*wire.EventsReply, error) + PollCommands(ctx context.Context) (*wire.CommandsReply, error) } // Collector measures the server. @@ -46,6 +49,7 @@ type agent struct { cp ControlPlane collector Collector outbox *outbox + ledger *ledger // interval is the report interval, and pending is the number of samples // since the last report. @@ -59,12 +63,18 @@ type agent struct { // whether the last report sent all samples. eventsWait backoff metricsWait backoff + pollWait backoff samplesSent bool } // Run runs the agent until ctx is done. Only one agent can run with the same // state directory. A nil collector takes no samples. func Run(ctx context.Context, cfg Config, log *slog.Logger, collector Collector) error { + // A crash during an update can leave a download next to the binary. + if exe, err := os.Executable(); err == nil { + release.RemoveTemp(filepath.Dir(exe)) + } + return run(ctx, cfg, log, NewClient(cfg, nil), collector) } @@ -84,6 +94,7 @@ func run(ctx context.Context, cfg Config, log *slog.Logger, cp ControlPlane, col cp: cp, collector: collector, outbox: loadOutbox(cfg.StateDir, log), + ledger: loadLedger(cfg.StateDir, log), interval: loadInterval(cfg.StateDir, log), } log.Info("agent started", "version", version.Version, "offset", cfg.Offset(), "report_interval", a.interval) @@ -91,26 +102,35 @@ func run(ctx context.Context, cfg Config, log *slog.Logger, cp ControlPlane, col // Send the events at once, not at the next tick: after an update or a // restart, they hold the result of the command. a.addEvent(wire.EventAgentStarted, "", &wire.EventData{Version: version.Version}) + a.resolve() a.send(ctx, false) - a.loop(ctx) + if a.loop(ctx) { + // systemd starts the agent again (Restart=always), with the new + // binary after an update. + log.Info("agent exits for a command") + return nil + } log.Info("agent stopped") return nil } -// loop calls tick at the offset second of each minute until ctx is done. -func (a *agent) loop(ctx context.Context) { +// loop calls tick at the offset second of each minute until ctx is done. It +// returns true when a command ends the process. +func (a *agent) loop(ctx context.Context) (exit bool) { for { next := nextAfter(time.Now(), a.last, a.cfg.Offset()) timer := time.NewTimer(time.Until(next)) select { case <-ctx.Done(): timer.Stop() - return + return false case <-timer.C: a.last = next - a.tick(ctx, next) + if a.tick(ctx, next) { + return true + } } } } @@ -129,8 +149,9 @@ func nextAfter(now, last time.Time, offset time.Duration) time.Time { } // tick does the work of one minute: it takes a sample and, after each -// interval samples, sends a report. -func (a *agent) tick(ctx context.Context, now time.Time) { +// interval samples, sends a report and runs the new commands. It returns true +// when a command ends the process. +func (a *agent) tick(ctx context.Context, now time.Time) (exit bool) { a.log.Debug("tick", "at", now) if a.collector != nil { @@ -145,17 +166,32 @@ func (a *agent) tick(ctx context.Context, now time.Time) { a.pending++ if a.pending < a.interval { - return + return false } a.log.Debug("report") - a.send(ctx, true) + eventsSent := a.send(ctx, true) // Samples that could not go are tried again at the next tick, when their // wait allows it, not only after the next full interval. if a.samplesSent { a.pending = 0 } + + // Poll only when no event waits: the results of the commands that ran + // must reach the control plane first, so that it does not send them again. + if !eventsSent { + return false + } + if a.commands(ctx) { + return true + } + + // Send the results of the commands now, not at the next report. + if len(a.outbox.events) > 0 { + a.send(ctx, false) + } + return false } // addEvent puts an event in the queue. Its ID is made now, so that each resend diff --git a/internal/agent/commands.go b/internal/agent/commands.go new file mode 100644 index 0000000..3958db9 --- /dev/null +++ b/internal/agent/commands.go @@ -0,0 +1,285 @@ +package agent + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io/fs" + "log/slog" + "net/http" + "os" + "path/filepath" + "runtime" + "slices" + "strings" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/release" + "github.com/flywp/server-cli/internal/statefile" + "github.com/flywp/server-cli/internal/version" + "golang.org/x/mod/semver" +) + +const ( + // keepRan is how long the agent remembers a command that it ran. The + // control plane sends a command for at most 24 hours, and it keeps the + // event ids for 2 days. + keepRan = 48 * time.Hour + // updateTimeout limits the download and the install of an update. + updateTimeout = 5 * time.Minute +) + +// ranCommand is a command that the agent ran or started to run. +type ranCommand struct { + ID string `json:"id"` + Verb string `json:"verb"` + // Target is the version of an update. + Target string `json:"target_version,omitempty"` + RanAt time.Time `json:"ran_at"` + // ResultSent is false until the result event is in the queue. An + // update or a restart ends the process, so the next process sends it. + ResultSent bool `json:"result_sent"` +} + +// ledger is the list of the commands that the agent ran (ran.json). The poll +// sends each open command again until its result arrives, so the agent must +// remember a command to run it only one time. +type ledger struct { + path string + log *slog.Logger + entries []ranCommand +} + +func loadLedger(dir string, log *slog.Logger) *ledger { + l := &ledger{path: filepath.Join(dir, "ran.json"), log: log} + if err := statefile.Read(l.path, &l.entries); err != nil && !errors.Is(err, fs.ErrNotExist) { + log.Error("dropping the list of the commands that ran", "error", err) + l.entries = nil + } + return l +} + +func (l *ledger) has(id string) bool { + return slices.ContainsFunc(l.entries, func(c ranCommand) bool { return c.ID == id }) +} + +// add records c and saves the list. The command stays in the list in memory +// also when the save fails. +func (l *ledger) add(c ranCommand) error { + l.entries = append(l.entries, c) + return statefile.Write(l.path, l.entries) +} + +func (l *ledger) markSent(id string) { + for i := range l.entries { + if l.entries[i].ID == id { + l.entries[i].ResultSent = true + } + } + l.save() +} + +// prune forgets the commands that finished more than keepRan ago. +func (l *ledger) prune(now time.Time) { + n := len(l.entries) + l.entries = slices.DeleteFunc(l.entries, func(c ranCommand) bool { + return c.ResultSent && now.Sub(c.RanAt) > keepRan + }) + if len(l.entries) != n { + l.save() + } +} + +func (l *ledger) save() { + if err := statefile.Write(l.path, l.entries); err != nil { + l.log.Error("saving the list of the commands that ran", "error", err) + } +} + +// resolve puts the result of each command that the previous process started +// in the event queue: an update or a restart ends the process that runs it. +func (a *agent) resolve() { + for _, c := range slices.Clone(a.ledger.entries) { + if c.ResultSent { + continue + } + + if cmp, ok := compareVersions(version.Version, c.Target); c.Verb == wire.VerbUpdate && (!ok || cmp < 0) { + a.addEvent(wire.EventCommandFailed, c.ID, &wire.EventData{ + Version: version.Version, + Error: fmt.Sprintf("the agent runs %s, not %s, after the update", version.Version, c.Target), + }) + } else { + a.addEvent(wire.EventCommandCompleted, c.ID, &wire.EventData{Version: version.Version}) + } + a.ledger.markSent(c.ID) + } +} + +// commands polls for the open commands and runs the new ones, oldest first. +// It returns true when the process must exit: after an update or a restart. +// The next process does the other commands. +func (a *agent) commands(ctx context.Context) (exit bool) { + start := time.Now() + if a.pollWait.waiting(start) { + return false + } + + reply, err := a.cp.PollCommands(ctx) + if a.outcome(ctx, &a.pollWait, start, err, "the poll", 0) != sent { + return false + } + + a.ledger.prune(time.Now()) + for _, c := range reply.Commands { + if a.ledger.has(c.ID) { + continue + } + if a.runCommand(ctx, c) { + return true + } + } + + return false +} + +// runCommand runs one command. It returns true when the process must exit. +func (a *agent) runCommand(ctx context.Context, c wire.Command) (exit bool) { + log := a.log.With("command", c.ID, "verb", c.Verb) + + switch c.Verb { + case wire.VerbRestart: + // Record the command before the exit. Without the record, the next + // process gets the same command and restarts again. + if err := a.ledger.add(ranCommand{ID: c.ID, Verb: c.Verb, RanAt: time.Now()}); err != nil { + a.fail(c, fmt.Errorf("saving the list of the commands that ran: %w", err)) + return false + } + log.Info("exiting for a restart; systemd starts the agent again") + return true + + case wire.VerbUpdate: + return a.update(ctx, c, log) + + default: + log.Warn("not running a command with an unknown verb") + a.addEvent(wire.EventCommandUnknown, c.ID, nil) + a.record(c, "") + return false + } +} + +// update installs the release of an agent.update command. +func (a *agent) update(ctx context.Context, c wire.Command, log *slog.Logger) (exit bool) { + var args wire.UpdateArgs + if err := json.Unmarshal(c.Args, &args); err != nil || args.URL == "" || args.Version == "" || args.SHA256 == "" { + a.fail(c, errors.New("the arguments of agent.update are not valid: url, version and sha256 are necessary")) + return false + } + + // The agent never downgrades: a bad release is fixed with a newer one. + // The same version completes the command; an older target fails it, so + // the control plane sees that its version was not installed. + switch cmp, ok := compareVersions(version.Version, args.Version); { + case ok && cmp == 0: + log.Info("the agent already runs this version", "version", version.Version) + a.addEvent(wire.EventCommandCompleted, c.ID, &wire.EventData{Version: version.Version}) + a.record(c, args.Version) + return false + case ok && cmp > 0: + a.fail(c, fmt.Errorf("the agent runs %s, newer than %s; it does not downgrade", version.Version, args.Version)) + return false + } + + // Record the command before the update: after the exit, the next process + // sends the result. + if err := a.ledger.add(ranCommand{ID: c.ID, Verb: c.Verb, Target: args.Version, RanAt: time.Now()}); err != nil { + a.fail(c, fmt.Errorf("saving the list of the commands that ran: %w", err)) + return false + } + + ctx, cancel := context.WithTimeout(ctx, updateTimeout) + defer cancel() + if err := updateBinary(ctx, args); err != nil { + log.Error("the update failed; the old binary continues", "error", err) + a.addEvent(wire.EventCommandFailed, c.ID, &wire.EventData{Version: version.Version, Error: err.Error()}) + a.ledger.markSent(c.ID) + return false + } + + log.Info("installed the new binary; exiting so that systemd starts it", "version", args.Version) + return true +} + +// fail puts command.failed for c in the queue, and records c. +func (a *agent) fail(c wire.Command, err error) { + a.log.Error("the command failed", "command", c.ID, "verb", c.Verb, "error", err) + a.addEvent(wire.EventCommandFailed, c.ID, &wire.EventData{Version: version.Version, Error: err.Error()}) + a.record(c, "") +} + +// record adds c, with its result already in the queue, to the ledger. +func (a *agent) record(c wire.Command, target string) { + if a.ledger.has(c.ID) { + a.ledger.markSent(c.ID) + return + } + if err := a.ledger.add(ranCommand{ID: c.ID, Verb: c.Verb, Target: target, RanAt: time.Now(), ResultSent: true}); err != nil { + a.log.Error("saving the list of the commands that ran", "error", err) + } +} + +// updateBinary downloads the release of args, checks its sha256 and puts it +// in place of the running binary. Tests replace it. +var updateBinary = func(ctx context.Context, args wire.UpdateArgs) error { + exe, err := os.Executable() + if err != nil { + return err + } + if exe, err = filepath.EvalSymlinks(exe); err != nil { + return err + } + + // The download goes next to the binary, so that the rename in Install + // does not cross file systems. + archive, err := release.Download(ctx, args.URL, args.SHA256, filepath.Dir(exe)) + if err != nil { + return err + } + defer func() { _ = os.Remove(archive) }() + + return release.Install(archive, exe, release.BinaryName(runtime.GOOS, runtime.GOARCH)) +} + +// compareVersions compares the versions a and b (release tags, for example +// v0.2.1) with semver: -1, 0 or +1. ok is false when they have no order: +// a version that is not semver (for example "dev"), or two dev tags of one +// release (v0.2.0-dev.1a2b3c4 and v0.2.0-dev.9f8e7d6), whose commit hashes +// have no order. Equal strings always compare as 0. +func compareVersions(a, b string) (cmp int, ok bool) { + if a == b { + return 0, true + } + if !semver.IsValid(a) || !semver.IsValid(b) { + return 0, false + } + + pa, pb := semver.Prerelease(a), semver.Prerelease(b) + if pa != "" && pb != "" && strings.TrimSuffix(semver.Canonical(a), pa) == strings.TrimSuffix(semver.Canonical(b), pb) { + return 0, false + } + + return semver.Compare(a, b), true +} + +// PollCommands gets the open commands (contract section 5). +func (c *Client) PollCommands(ctx context.Context) (*wire.CommandsReply, error) { + var reply wire.CommandsReply + if err := c.do(ctx, http.MethodGet, "agent/v1/commands", nil, &reply); err != nil { + return nil, err + } + + return &reply, nil +} diff --git a/internal/agent/commands_test.go b/internal/agent/commands_test.go new file mode 100644 index 0000000..234cd5b --- /dev/null +++ b/internal/agent/commands_test.go @@ -0,0 +1,377 @@ +package agent + +import ( + "context" + "encoding/json" + "errors" + "log/slog" + "os" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/version" +) + +// setVersion sets the version of the running agent for one test. +func setVersion(t *testing.T, v string) { + t.Helper() + old := version.Version + version.Version = v + t.Cleanup(func() { version.Version = old }) +} + +// fakeUpdate replaces the download and the install of an update. +func fakeUpdate(t *testing.T, err error) *[]wire.UpdateArgs { + t.Helper() + var calls []wire.UpdateArgs + old := updateBinary + updateBinary = func(_ context.Context, args wire.UpdateArgs) error { + calls = append(calls, args) + return err + } + t.Cleanup(func() { updateBinary = old }) + return &calls +} + +func updateCommand(t *testing.T, id, v string) wire.Command { + t.Helper() + args, err := json.Marshal(wire.UpdateArgs{URL: "https://example.com/fly-linux-amd64.tar.gz", Version: v, SHA256: "9f86d081884c7d659a2feaa0c55ad015a3bf4f1b2b0b822cd15d6c15b0f00a08"}) + if err != nil { + t.Fatal(err) + } + return wire.Command{ID: id, Verb: wire.VerbUpdate, Args: args} +} + +// start runs an agent with server id 17 in the bubble until it exits or until +// d passes. It returns true when the agent exited for a command. +func start(t *testing.T, d time.Duration, dir string, cp *fakeCP) (exited bool) { + t.Helper() + + ctx, cancel := context.WithTimeout(context.Background(), d) + defer cancel() + if err := run(ctx, Config{ServerID: 17, StateDir: dir}, slog.New(&recorder{}), cp, &fakeCollector{}); err != nil { + t.Fatal(err) + } + return ctx.Err() == nil +} + +func eventNames(events []wire.Event) []string { + var names []string + for _, e := range events { + names = append(names, e.Name) + } + return names +} + +func TestRestartRunsOneTime(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + // The worst case: the command stays open after its result. + cp := &fakeCP{commands: []wire.Command{{ID: "01JBX0000000000000000000R1", Verb: wire.VerbRestart, Args: json.RawMessage(`{}`)}}, keepOpen: true} + + if !start(t, 5*time.Minute, dir, cp) { + t.Fatal("the agent did not exit for agent.restart") + } + if !time.Now().Equal(at(0)) { + t.Errorf("the agent exited at %s, want at the first report %s", time.Now(), at(0)) + } + if got := cp.results("01JBX0000000000000000000R1"); len(got) != 0 { + t.Errorf("results before the exit = %v, want none: the new process sends the result", eventNames(got)) + } + + // systemd starts the agent again. The command is still open, but the + // agent does not restart again. + if start(t, 5*time.Minute, dir, cp) { + t.Fatal("the new process exited again for the same agent.restart") + } + got := cp.results("01JBX0000000000000000000R1") + if len(got) != 1 || got[0].Name != wire.EventCommandCompleted { + t.Errorf("results = %v, want one command.completed", eventNames(got)) + } + }) +} + +func TestUnknownVerbIsNotRun(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + cp := &fakeCP{commands: []wire.Command{{ID: "01JBX0000000000000000000V1", Verb: "agent.shell", Args: json.RawMessage(`{"cmd":"rm -rf /"}`)}}, keepOpen: true} + calls := fakeUpdate(t, nil) + + if start(t, 3*time.Minute, t.TempDir(), cp) { + t.Fatal("the agent exited for an unknown verb") + } + got := cp.results("01JBX0000000000000000000V1") + if len(got) != 1 || got[0].Name != wire.EventCommandUnknown { + t.Errorf("results = %v, want one command.unknown in 3 polls", eventNames(got)) + } + if len(*calls) != 0 { + t.Errorf("updates = %d, want none", len(*calls)) + } + }) +} + +func TestUpdateToTheRunningVersionDoesNotDownload(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.3.0") + calls := fakeUpdate(t, nil) + cp := &fakeCP{commands: []wire.Command{updateCommand(t, "01JBX0000000000000000000A1", "v0.3.0")}} + + if start(t, 2*time.Minute, t.TempDir(), cp) { + t.Fatal("the agent exited for an update to its own version") + } + if len(*calls) != 0 { + t.Errorf("downloads = %d, want none", len(*calls)) + } + got := cp.results("01JBX0000000000000000000A1") + if len(got) != 1 || got[0].Name != wire.EventCommandCompleted || got[0].Data.Version != "v0.3.0" { + t.Errorf("results = %+v, want command.completed with v0.3.0", got) + } + }) +} + +func TestNoDowngrade(t *testing.T) { + tests := []struct { + name, running, target string + }{ + {"older release", "v0.3.0", "v0.2.0"}, + {"release over its dev tag", "v0.3.0", "v0.3.0-dev.1a2b3c4"}, + {"newer dev tag over an older release", "v0.3.0-dev.1a2b3c4", "v0.2.1"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, tt.running) + calls := fakeUpdate(t, nil) + cp := &fakeCP{commands: []wire.Command{updateCommand(t, "01JBX0000000000000000000K1", tt.target)}} + + if start(t, time.Minute, t.TempDir(), cp) { + t.Fatal("the agent exited: it must not downgrade") + } + if len(*calls) != 0 { + t.Errorf("updates = %+v, want none", *calls) + } + got := cp.results("01JBX0000000000000000000K1") + if len(got) != 1 || got[0].Name != wire.EventCommandFailed || !strings.Contains(got[0].Data.Error, "it does not downgrade") { + t.Errorf("results = %+v, want one command.failed that says why", got) + } + }) + }) + } +} + +func TestUpdateFromAnOlderOrUnorderedVersion(t *testing.T) { + tests := []struct { + name, running, target string + }{ + {"older release", "v0.2.0", "v0.3.0"}, + {"dev tag to its release", "v0.3.0-dev.1a2b3c4", "v0.3.0"}, + {"two dev tags of one release", "v0.3.0-dev.9f8e7d6", "v0.3.0-dev.1a2b3c4"}, + {"a build without a release tag", "dev", "v0.3.0"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, tt.running) + calls := fakeUpdate(t, nil) + cp := &fakeCP{commands: []wire.Command{updateCommand(t, "01JBX0000000000000000000K2", tt.target)}} + + if !start(t, time.Minute, t.TempDir(), cp) || len(*calls) != 1 { + t.Errorf("updates = %+v, want one update to %s", *calls, tt.target) + } + }) + }) + } +} + +func TestRestartIsNotRunWhenItsRecordCannotBeSaved(t *testing.T) { + if os.Geteuid() == 0 { + t.Skip("root can write to a read-only directory") + } + + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + cp := &fakeCP{commands: []wire.Command{{ID: "01JBX0000000000000000000S1", Verb: wire.VerbRestart, Args: json.RawMessage(`{}`)}}, keepOpen: true} + + done := make(chan bool) + go func() { done <- start(t, 3*time.Minute, dir, cp) }() + synctest.Wait() + + // The state directory becomes read-only after the start: the list of + // the commands that ran cannot be saved. + if err := os.Chmod(dir, 0o500); err != nil { + t.Fatal(err) + } + defer func() { _ = os.Chmod(dir, 0o700) }() + + if <-done { + t.Fatal("the agent exited for agent.restart without a saved record: the next process would restart again") + } + got := cp.results("01JBX0000000000000000000S1") + if len(got) == 0 || got[0].Name != wire.EventCommandFailed { + t.Errorf("results = %v, want command.failed", eventNames(got)) + } + }) +} + +func TestFailedUpdateKeepsTheOldBinary(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + calls := fakeUpdate(t, errors.New("the sha256 of the archive is abc, not 9f86")) + cp := &fakeCP{commands: []wire.Command{updateCommand(t, "01JBX0000000000000000000F1", "v0.3.0")}, keepOpen: true} + + if start(t, 3*time.Minute, t.TempDir(), cp) { + t.Fatal("the agent exited after a failed update") + } + if len(*calls) != 1 { + t.Errorf("downloads = %d, want 1: the agent does not try the same command again", len(*calls)) + } + got := cp.results("01JBX0000000000000000000F1") + if len(got) != 1 || got[0].Name != wire.EventCommandFailed || got[0].Data.Error == "" { + t.Errorf("results = %+v, want one command.failed with the error", got) + } + }) +} + +func TestUpdateExitsAndTheNewProcessReports(t *testing.T) { + tests := []struct { + name, newVersion, want string + }{ + {"new version runs", "v0.3.0", wire.EventCommandCompleted}, + {"a newer release runs", "v0.3.1", wire.EventCommandCompleted}, + // For example, the process stopped during the download. + {"old version still runs", "v0.2.0", wire.EventCommandFailed}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + setVersion(t, "v0.2.0") + calls := fakeUpdate(t, nil) + cp := &fakeCP{commands: []wire.Command{updateCommand(t, "01JBX0000000000000000000N1", "v0.3.0")}} + + if !start(t, 2*time.Minute, dir, cp) { + t.Fatal("the agent did not exit after the update") + } + if len(*calls) != 1 || (*calls)[0].Version != "v0.3.0" { + t.Fatalf("updates = %+v, want one to v0.3.0", *calls) + } + + // systemd starts the binary that is now on the disk. + version.Version = tt.newVersion + if start(t, 2*time.Minute, dir, cp) { + t.Fatal("the new process exited again") + } + got := cp.results("01JBX0000000000000000000N1") + if len(got) != 1 || got[0].Name != tt.want || got[0].Data.Version != tt.newVersion { + t.Errorf("results = %+v, want one %s with version %s", got, tt.want, tt.newVersion) + } + if len(*calls) != 1 { + t.Errorf("updates = %d, want 1", len(*calls)) + } + }) + }) + } +} + +func TestUpdateWithArgumentsThatAreNotValid(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + calls := fakeUpdate(t, nil) + cp := &fakeCP{commands: []wire.Command{{ID: "01JBX0000000000000000000B1", Verb: wire.VerbUpdate, Args: json.RawMessage(`{"url":"https://example.com"}`)}}} + + if start(t, time.Minute, t.TempDir(), cp) { + t.Fatal("the agent exited for an update that is not valid") + } + got := cp.results("01JBX0000000000000000000B1") + if len(got) != 1 || got[0].Name != wire.EventCommandFailed || len(*calls) != 0 { + t.Errorf("results = %v and %d downloads, want one command.failed and no download", eventNames(got), len(*calls)) + } + }) +} + +func TestCommandsAfterAnExitWaitForTheNextProcess(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + cp := &fakeCP{commands: []wire.Command{ + {ID: "01JBX0000000000000000000C1", Verb: wire.VerbRestart, Args: json.RawMessage(`{}`)}, + {ID: "01JBX0000000000000000000C2", Verb: "agent.unknown", Args: json.RawMessage(`{}`)}, + }} + + if !start(t, time.Minute, dir, cp) { + t.Fatal("the agent did not exit for agent.restart") + } + if got := cp.results("01JBX0000000000000000000C2"); len(got) != 0 { + t.Errorf("the second command ran before the exit: %v", eventNames(got)) + } + + start(t, 2*time.Minute, dir, cp) + if got := cp.results("01JBX0000000000000000000C2"); len(got) != 1 || got[0].Name != wire.EventCommandUnknown { + t.Errorf("results of the second command = %v, want command.unknown from the new process", eventNames(got)) + } + }) +} + +func TestNoPollWhileTheControlPlaneIsDown(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + cp := &fakeCP{eventsReply: func(int, *wire.EventsRequest) (*wire.EventsReply, error) { + return nil, &StatusError{StatusCode: 503} + }} + + start(t, 30*time.Second, t.TempDir(), cp) + if len(cp.pollAt) != 0 { + t.Errorf("polls = %d, want none while the events cannot be sent", len(cp.pollAt)) + } + }) +} + +func TestCompareVersions(t *testing.T) { + tests := []struct { + a, b string + cmp int + ok bool + }{ + {"v0.3.0", "v0.3.0", 0, true}, + {"v0.3.1", "v0.3.0", 1, true}, + {"v0.2.9", "v0.3.0", -1, true}, + {"v1.0.0", "v0.10.0", 1, true}, + {"v0.2.0-dev.1a2b3c4", "v0.2.0", -1, true}, + {"v0.3.0-dev.1a2b3c4", "v0.2.1", 1, true}, + {"v0.2.0-dev.1a2b3c4", "v0.2.0-dev.1a2b3c4", 0, true}, + // Dev tags of one release have no order. + {"v0.2.0-dev.9f8e7d6", "v0.2.0-dev.1a2b3c4", 0, false}, + {"dev", "v0.2.0", 0, false}, + {"v0.2.0", "0.2.0", 0, false}, + } + + for _, tt := range tests { + if cmp, ok := compareVersions(tt.a, tt.b); cmp != tt.cmp || ok != tt.ok { + t.Errorf("compareVersions(%q, %q) = %d, %v; want %d, %v", tt.a, tt.b, cmp, ok, tt.cmp, tt.ok) + } + } +} + +func TestLedgerForgetsOldCommands(t *testing.T) { + l := loadLedger(t.TempDir(), slog.New(&recorder{})) + now := time.Now() + for _, c := range []ranCommand{ + {ID: "old-done", RanAt: now.Add(-49 * time.Hour), ResultSent: true}, + {ID: "old-open", RanAt: now.Add(-49 * time.Hour)}, + {ID: "new-done", RanAt: now.Add(-time.Hour), ResultSent: true}, + } { + if err := l.add(c); err != nil { + t.Fatal(err) + } + } + + l.prune(now) + again := loadLedger(l.path[:len(l.path)-len("/ran.json")], slog.New(&recorder{})) + for id, want := range map[string]bool{"old-done": false, "old-open": true, "new-done": true} { + if again.has(id) != want { + t.Errorf("has(%s) = %v after prune, want %v", id, !want, want) + } + } +} diff --git a/internal/agent/fakes_test.go b/internal/agent/fakes_test.go index 79841e6..2e11734 100644 --- a/internal/agent/fakes_test.go +++ b/internal/agent/fakes_test.go @@ -19,6 +19,7 @@ type fakeCP struct { metricsAt []time.Time events []wire.EventsRequest eventsAt []time.Time + pollAt []time.Time // latency is the time of each metrics request. The request ends early // when its context ends, like a real HTTP request. @@ -27,6 +28,44 @@ type fakeCP struct { // The reply funcs get the number of the call, from 0. metricsReply func(call int, req *wire.MetricsRequest) (*wire.MetricsReply, error) eventsReply func(call int, req *wire.EventsRequest) (*wire.EventsReply, error) + // commands are the open commands of each poll. A command stays open + // until an event with its id arrives, unless keepOpen is true. + commands []wire.Command + keepOpen bool +} + +func (f *fakeCP) PollCommands(context.Context) (*wire.CommandsReply, error) { + f.mu.Lock() + defer f.mu.Unlock() + + f.pollAt = append(f.pollAt, time.Now()) + + var open []wire.Command + for _, c := range f.commands { + if f.keepOpen || len(f.resultsLocked(c.ID)) == 0 { + open = append(open, c) + } + } + return &wire.CommandsReply{Commands: open}, nil +} + +// results returns the events that the agent sent for the command id. +func (f *fakeCP) results(id string) []wire.Event { + f.mu.Lock() + defer f.mu.Unlock() + return f.resultsLocked(id) +} + +func (f *fakeCP) resultsLocked(id string) []wire.Event { + var out []wire.Event + for _, r := range f.events { + for _, e := range r.Events { + if e.CommandID == id { + out = append(out, e) + } + } + } + return out } func (f *fakeCP) PostMetrics(ctx context.Context, req *wire.MetricsRequest) (*wire.MetricsReply, error) { diff --git a/internal/agent/wire/wire.go b/internal/agent/wire/wire.go index d9e3d26..c144d49 100644 --- a/internal/agent/wire/wire.go +++ b/internal/agent/wire/wire.go @@ -2,7 +2,10 @@ // v0.2.1: the requests that the agent sends and the replies that it reads. package wire -import "time" +import ( + "encoding/json" + "time" +) // MetricsRequest is the body of POST /agent/v1/metrics (contract section 4). type MetricsRequest struct { @@ -89,3 +92,32 @@ type EventData struct { type EventsReply struct { Accepted int `json:"accepted"` } + +// CommandsReply is the reply to GET /agent/v1/commands (contract section 5): +// the open commands of the server, oldest first. +type CommandsReply struct { + Commands []Command `json:"commands"` +} + +// The verbs of the contract. The agent runs no other verb. +const ( + VerbUpdate = "agent.update" + VerbRestart = "agent.restart" +) + +// Command is a command from the control plane. It comes again on each poll +// until an event finishes it. +type Command struct { + ID string `json:"id"` + Verb string `json:"verb"` + Args json.RawMessage `json:"args"` + IssuedAt time.Time `json:"issued_at"` +} + +// UpdateArgs are the arguments of agent.update. SHA256 is the sha256 of the +// release archive at URL, and Version is its release tag. +type UpdateArgs struct { + URL string `json:"url"` + Version string `json:"version"` + SHA256 string `json:"sha256"` +} diff --git a/internal/release/download.go b/internal/release/download.go new file mode 100644 index 0000000..d71ea46 --- /dev/null +++ b/internal/release/download.go @@ -0,0 +1,133 @@ +package release + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "fmt" + "io" + "net" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + "time" +) + +// maxArchiveSize limits the size of a downloaded release archive. +const maxArchiveSize = 200 << 20 + +// downloadClient has no timeout of its own: the context of the caller limits +// the download, because a slow network can need some minutes. It follows a +// redirect (GitHub sends each download to its file host) only to https. +var downloadClient = &http.Client{ + CheckRedirect: func(req *http.Request, via []*http.Request) error { + if len(via) >= 10 { + return fmt.Errorf("too many redirects") + } + return checkURL(req.URL.String()) + }, +} + +// tempMaxAge is the age after which RemoveTemp removes a temporary file. A +// younger file can belong to an update that still runs. +const tempMaxAge = time.Hour + +// RemoveTemp removes the temporary files that an update left in dir after a +// crash, if they are older than one hour. +func RemoveTemp(dir string) { + for _, pattern := range []string{".fly-download-*", ".fly-update-*"} { + matches, _ := filepath.Glob(filepath.Join(dir, pattern)) + for _, m := range matches { + if info, err := os.Lstat(m); err == nil && info.Mode().IsRegular() && time.Since(info.ModTime()) > tempMaxAge { + _ = os.Remove(m) + } + } + } +} + +// Download fetches the release archive at rawURL into a temporary file in dir, +// and checks that its sha256 is wantSHA256 (hex). It returns the path of the +// file; the caller removes it. On any error, no file is left. +// +// The URL must be https. Plain http is accepted only for a loopback host, +// for tests and local development. +func Download(ctx context.Context, rawURL, wantSHA256, dir string) (path string, err error) { + if err := checkURL(rawURL); err != nil { + return "", err + } + want := strings.ToLower(strings.TrimSpace(wantSHA256)) + if len(want) != sha256.Size*2 { + return "", fmt.Errorf("the sha256 %q is not valid", wantSHA256) + } + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil) + if err != nil { + return "", err + } + resp, err := downloadClient.Do(req) + if err != nil { + return "", fmt.Errorf("downloading %s: %w", rawURL, err) + } + defer func() { _ = resp.Body.Close() }() + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("downloading %s: %s", rawURL, resp.Status) + } + + tmp, err := os.CreateTemp(dir, ".fly-download-*") + if err != nil { + return "", err + } + defer func() { + if err != nil { + _ = os.Remove(tmp.Name()) + } + }() + + hash := sha256.New() + n, err := io.Copy(io.MultiWriter(tmp, hash), io.LimitReader(resp.Body, maxArchiveSize+1)) + if cerr := tmp.Close(); err == nil { + err = cerr + } + if err != nil { + return "", fmt.Errorf("downloading %s: %w", rawURL, err) + } + if n > maxArchiveSize { + return "", fmt.Errorf("the archive at %s is larger than %d bytes", rawURL, maxArchiveSize) + } + + if got := hex.EncodeToString(hash.Sum(nil)); got != want { + return "", fmt.Errorf("the sha256 of %s is %s, not %s", rawURL, got, want) + } + + return tmp.Name(), nil +} + +// Install extracts the binary name from the archive at archivePath and puts it +// in place of exe. exe is never incomplete: see replaceBinary. +func Install(archivePath, exe, name string) error { + f, err := os.Open(archivePath) + if err != nil { + return err + } + defer func() { _ = f.Close() }() + + return replaceBinary(exe, f, name) +} + +func checkURL(rawURL string) error { + u, err := url.Parse(rawURL) + if err != nil || u.Host == "" { + return fmt.Errorf("the download URL %q is not valid", rawURL) + } + + switch host := u.Hostname(); { + case u.Scheme == "https": + return nil + case u.Scheme == "http" && (host == "localhost" || net.ParseIP(host).IsLoopback()): + return nil + } + + return fmt.Errorf("the download URL must be https, not %q", rawURL) +} diff --git a/internal/release/download_test.go b/internal/release/download_test.go new file mode 100644 index 0000000..eaac30d --- /dev/null +++ b/internal/release/download_test.go @@ -0,0 +1,147 @@ +package release + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func sum(data []byte) string { + h := sha256.Sum256(data) + return hex.EncodeToString(h[:]) +} + +func serveFile(t *testing.T, data []byte) string { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/fly-linux-amd64.tar.gz" { + http.NotFound(w, r) + return + } + _, _ = w.Write(data) + })) + t.Cleanup(srv.Close) + return srv.URL + "/fly-linux-amd64.tar.gz" +} + +func TestDownloadAndInstall(t *testing.T) { + data := archive(t, map[string]string{"fly-linux-amd64": "new"}).Bytes() + url := serveFile(t, data) + dir := t.TempDir() + exe := filepath.Join(dir, "fly") + if err := os.WriteFile(exe, []byte("old"), 0o755); err != nil { + t.Fatal(err) + } + + // The hash is compared without regard to case. + path, err := Download(context.Background(), url, strings.ToUpper(sum(data)), dir) + if err != nil { + t.Fatal(err) + } + defer func() { _ = os.Remove(path) }() + + if err := Install(path, exe, "fly-linux-amd64"); err != nil { + t.Fatal(err) + } + if got, _ := os.ReadFile(exe); string(got) != "new" { + t.Errorf("binary = %q, want the new binary", got) + } +} + +func TestDownloadErrorsLeaveNoFile(t *testing.T) { + data := archive(t, map[string]string{"fly-linux-amd64": "new"}).Bytes() + url := serveFile(t, data) + + tests := []struct { + name, url, sha256, want string + }{ + {"wrong sha256", url, sum([]byte("other")), "the sha256 of"}, + {"sha256 not valid", url, "abc", "is not valid"}, + {"not found", strings.Replace(url, "fly-linux-amd64", "missing", 1), sum(data), "404"}, + {"plain http to a remote host", "http://github.com/flywp/server-cli/fly.tar.gz", sum(data), "must be https"}, + {"not a URL", "://", sum(data), "not valid"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + path, err := Download(context.Background(), tt.url, tt.sha256, dir) + if err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("Download() = %q, %v; want an error with %q", path, err, tt.want) + } + if entries, _ := os.ReadDir(dir); len(entries) != 0 { + t.Errorf("Download() left %d files in the directory", len(entries)) + } + }) + } +} + +func TestDownloadLimitsTheSize(t *testing.T) { + big := make([]byte, maxArchiveSize+10) + url := serveFile(t, big) + dir := t.TempDir() + + if _, err := Download(context.Background(), url, sum(big), dir); err == nil || !strings.Contains(err.Error(), "larger than") { + t.Fatalf("Download() error = %v, want a size error", err) + } + if entries, _ := os.ReadDir(dir); len(entries) != 0 { + t.Errorf("Download() left %d files in the directory", len(entries)) + } +} + +func TestDownloadRefusesARedirectToPlainHTTP(t *testing.T) { + plain := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + t.Error("the download followed a redirect to plain http") + })) + defer plain.Close() + // A plain-http host that is not loopback, reached through a redirect. + remote := strings.Replace(plain.URL, "127.0.0.1", "localtest.invalid", 1) + + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, remote+"/fly.tar.gz", http.StatusFound) + })) + defer srv.Close() + + old := downloadClient.Transport + downloadClient.Transport = srv.Client().Transport + t.Cleanup(func() { downloadClient.Transport = old }) + + _, err := Download(context.Background(), srv.URL+"/fly.tar.gz", sum([]byte("x")), t.TempDir()) + if err == nil || !strings.Contains(err.Error(), "must be https") { + t.Fatalf("Download() error = %v, want a refusal of the http redirect", err) + } +} + +func TestRemoveTempKeepsYoungFiles(t *testing.T) { + dir := t.TempDir() + for _, name := range []string{".fly-download-1", ".fly-update-2", ".fly-download-new", "fly"} { + if err := os.WriteFile(filepath.Join(dir, name), []byte("x"), 0o600); err != nil { + t.Fatal(err) + } + } + old := time.Now().Add(-2 * time.Hour) + for _, name := range []string{".fly-download-1", ".fly-update-2"} { + if err := os.Chtimes(filepath.Join(dir, name), old, old); err != nil { + t.Fatal(err) + } + } + + RemoveTemp(dir) + + var names []string + entries, _ := os.ReadDir(dir) + for _, e := range entries { + names = append(names, e.Name()) + } + // A young file can belong to an update that still runs. + if strings.Join(names, ",") != ".fly-download-new,fly" { + t.Errorf("directory holds %v, want the young download and the binary only", names) + } +} diff --git a/internal/utils/version.go b/internal/release/release.go similarity index 86% rename from internal/utils/version.go rename to internal/release/release.go index 34f1a33..18b4192 100644 --- a/internal/utils/version.go +++ b/internal/release/release.go @@ -1,4 +1,6 @@ -package utils +// Package release finds, downloads, checks and installs fly releases. "fly +// update" and the agent command agent.update use it. +package release import ( "archive/tar" @@ -143,12 +145,12 @@ func SelfUpdate(ctx context.Context, release *GithubRelease) error { } defer func() { _ = resp.Body.Close() }() - return replaceBinary(exe, resp.Body, binaryName(runtime.GOOS, runtime.GOARCH)) + return replaceBinary(exe, resp.Body, BinaryName(runtime.GOOS, runtime.GOARCH)) } -// binaryName is the name of the binary in a release archive. Releases must +// BinaryName is the name of the binary in a release archive. Releases must // keep this name: installed versions of fly look for it. -func binaryName(goos, goarch string) string { +func BinaryName(goos, goarch string) string { return fmt.Sprintf("fly-%s-%s", goos, goarch) } @@ -159,7 +161,7 @@ func assetURL(release *GithubRelease, goos, goarch string) string { return "" } - expectedName := binaryName(goos, goarch) + ".tar.gz" + expectedName := BinaryName(goos, goarch) + ".tar.gz" for _, asset := range release.Assets { if asset.Name == expectedName { return asset.BrowserDownloadURL @@ -217,15 +219,31 @@ func writeBinary(exe string, r io.Reader) (err error) { _ = tmp.Close() return fmt.Errorf("writing new binary: %w", err) } - if err = tmp.Close(); err != nil { + // Change the open file, never the path: the directory can belong to an + // other user (the agent's ~fly/.fly/bin), who could put a link to a + // different file in place of the temporary file. + if err = tmp.Chmod(0o755); err != nil { + _ = tmp.Close() + return fmt.Errorf("making binary executable: %w", err) + } + // Put the binary on the disk before the rename: after a power loss, a + // renamed but empty binary would not start. + if err = tmp.Sync(); err != nil { + _ = tmp.Close() return fmt.Errorf("writing new binary: %w", err) } - if err = os.Chmod(tmp.Name(), 0o755); err != nil { - return fmt.Errorf("making binary executable: %w", err) + if err = tmp.Close(); err != nil { + return fmt.Errorf("writing new binary: %w", err) } if err = os.Rename(tmp.Name(), exe); err != nil { return fmt.Errorf("replacing binary: %w", err) } + // Sync the directory too, so that the rename survives a power loss. + if d, err := os.Open(filepath.Dir(exe)); err == nil { + _ = d.Sync() + _ = d.Close() + } + return nil } diff --git a/internal/utils/version_test.go b/internal/release/release_test.go similarity index 99% rename from internal/utils/version_test.go rename to internal/release/release_test.go index 91211b9..b9d8711 100644 --- a/internal/utils/version_test.go +++ b/internal/release/release_test.go @@ -1,4 +1,4 @@ -package utils +package release import ( "archive/tar"