diff --git a/.claude/skills/gh-stack/SKILL.md b/.claude/skills/gh-stack/SKILL.md new file mode 100644 index 0000000..7d5634d --- /dev/null +++ b/.claude/skills/gh-stack/SKILL.md @@ -0,0 +1,182 @@ +--- +description: | + Manages stacked PRs and splits multi-part work into reviewable branches with gh-stack. Use for stack creation, viewing, edits, push, submit, sync, rebase, merge, or checkout; when asked to split or isolate work for review; whenever a user mentions a stack, branch layers, dependent PRs, or gh stack; or when a stack is checked out. +metadata: + author: github + github-path: skills/gh-stack + github-pinned: v0.1.1 + github-ref: refs/tags/v0.1.1 + github-repo: https://github.com/github/gh-stack + github-tree-sha: 6bfe555e89a6b264e09d5ece7c4040226458c8e5 + version: 0.1.0 +name: gh-stack +--- +# gh-stack + +`gh stack` is a [GitHub CLI](https://cli.github.com/) extension for stacked branches and pull +requests. A stack is an ordered chain of branches rooted on a trunk, where each branch has one PR +based on the branch below it, so a reviewer sees only that layer's diff. + +`gh stack` prints a stack trunk-first, left to right: + +``` +(main) <- auth <- api <- frontend +``` + +Left is the **bottom**, right is the **top**. `auth` is based on `main` and merges first; +`frontend` merges last. `up` moves toward the top, away from trunk; `down` moves toward it. +Foundational work belongs at the bottom, code that depends on it above. For how to choose the +layers, read `references/stack-design.md`. + +## Setup + +```bash +gh extension install github/gh-stack +git config rerere.enabled true # remember conflict resolutions +git config remote.pushDefault origin # required if the repo has more than one remote +``` + +## Non-interactive use + +`gh stack` branches on whether **stdout is a TTY**. Piped, most commands error cleanly or print +static text; under a PTY the same commands open a prompt or a full-screen TUI and block forever. +Agent harnesses differ, so always pass the flags below instead of relying on that detection. + +**Multiple remotes:** never run `push`, `submit`, `sync`, `rebase`, or `link` without +`--remote ` unless `remote.pushDefault` is configured. `checkout` and `trunk` have no +`--remote` flag and require the config. + +| Always run | Never run bare | Why | +|---|---|---| +| `gh stack view --json` | `gh stack view` | opens a TUI under a PTY | +| `gh stack submit --auto` | `gh stack submit` | prompts for a title per new PR | +| `gh stack merge --yes` | `gh pr merge` | `gh pr merge` cannot merge a stack | +| `gh stack init ...` | `gh stack init` | prompts for branch names | +| `gh stack add ` | `gh stack add` | prompts for a name, and fails even when piped | +| `gh stack checkout ` | `gh stack checkout` | opens a selection menu | +| `gh stack up` / `down` / `top` / `bottom` | `gh stack switch` | `switch` is menu-only | +| — | `gh stack modify` | TUI-only, no non-interactive path | + +- `view --short` is safe in both modes, but it is formatted for humans. Use `--json` to parse. +- **`checkout ` when a different local stack already covers those branches** cannot be forced. + Run `gh stack unstack --local` first (this keeps the stack on GitHub), then retry. + +## Branch placement + +- **Starting multi-part work:** create the stack before writing files. Do not implement every + concern on trunk and split it later. Put one dependent concern in each layer, bottom to top. +- **Editing an existing stack:** check out the layer that owns the change before editing. Never + commit a lower layer's concern on the current top branch. Run `gh stack view --json`; if + ownership is unclear, inspect `git log --all -- `. Then check out the owner, edit, commit, + rebase upstack, and return to top. + +```bash +gh stack down # or: gh stack checkout api +git add ... && git commit -m "Add get-user endpoint" +gh stack rebase --upstack # replay every branch above onto the change +gh stack top # return to where you were +gh stack push +``` + +## Core loop + +```bash +gh stack init auth # create the stack and check out its branch +git add ... && git commit -m "Add auth middleware" +gh stack add api # next layer, branched from the current one +git add ... && git commit -m "Add API routes" +gh stack submit --auto # push every branch and open draft PRs +gh stack view --json # confirm +``` + +Add `--open` to `submit` to create PRs ready for review instead of drafts. Branch names are +verbatim — `gh stack add refactor/foo` creates `refactor/foo`. + +## Staying in sync + +```bash +gh stack sync # fetch, reconcile with GitHub, rebase, push, refresh PR state +gh stack sync --prune # also delete local branches for merged PRs +``` + +Pruning never happens without `--prune` when non-interactive. If the local and remote stacks have +diverged, `sync` prints both chains, makes no changes, and exits 0 with `Sync aborted` — see +`references/troubleshooting.md`. + +## Merging + +Scope the merge with an argument: + +```bash +gh stack merge 42 --yes # PR #42 plus every unmerged PR below it +gh stack merge 7 --yes # every unmerged PR in stack #7 +gh stack merge 42 --yes --squash # or --merge, --rebase, --merge-method +``` + +Pass a PR number to merge that PR and every unmerged PR below it, or a stack number to merge every +unmerged PR in that stack. The operation is all-or-nothing: if any PR in that set cannot merge, +none do. + +Without a method flag the last-used method is reused. If the base branch uses a merge queue, the +stack is queued instead and the queue picks the method, ignoring any flag you passed with a +warning; queued PRs may land in separate groups. + +## Reading state + +`gh stack view --json` writes JSON to **stdout**. Status messages go to **stderr** — do not parse +them, branch on exit codes instead. + +``` +trunk string +currentBranch string +branches[] name, head, base, isCurrent, isMerged, isQueued, needsRebase +branches[].pr number, url, state ("OPEN" | "MERGED" | "QUEUED"); absent when no PR exists +``` + +`base` is the saved SHA of the parent branch that this branch was last known to contain. It may be +older than the parent's current tip. `needsRebase` is true when the current parent tip is no longer +an ancestor of the branch. + +## Exit codes + +| Code | Meaning | Recovery | +|---|---|---| +| 0 | Success | — | +| 1 | Generic error | Read stderr | +| 2 | Not in a stack | `gh stack init`, or `gh stack checkout ` | +| 3 | Rebase conflict | Follow the Exit 3 recovery below | +| 4 | GitHub API failure | Check `gh auth status`, retry | +| 5 | Invalid arguments | Fix the invocation; see ` --help` | +| 6 | Disambiguation required | Branch is in several stacks; check out a non-shared branch | +| 7 | Rebase already in progress | `gh stack rebase --continue` or `--abort` | +| 8 | Stack file locked | Another `gh stack` process is writing; retry after ~5s | +| 9 | Stacked PRs unavailable | Not enabled on the repository; tell the user | +| 10 | Modify recovery required | `gh stack modify --abort` | + +**Exit 3 recovery:** + +- After `gh stack rebase`: resolve the files, run `git add`, then + `gh stack rebase --continue`; use `gh stack rebase --abort` to restore the stack. +- After `gh stack sync`: the stack has already been restored. Run `gh stack rebase` to recreate the + conflict, then resolve and continue as above. + +## Constraints + +- Stacks are strictly linear: one parent, at most one child. Use separate stacks for parallel work. +- There is no non-interactive reorder or removal. Errors may suggest `gh stack modify`, but it is + TUI-only — restructure with `unstack` then `init` instead. +- PR titles and bodies are auto-generated. Use `gh pr edit` afterwards to change them. + +## More detail + +`gh stack --help` is authoritative for flags and arguments. Note that +`gh stack help ` does **not** work — it prints the top-level help. + +Open the reference whose trigger matches the task; no need to preload all three. + +- `references/stack-design.md` — read before creating a stack, when deciding how many layers to + use, what belongs in each one, or whether work belongs in a new stack. +- `references/commands.md` — read when a command fails unexpectedly or you need its preconditions, + side effects, atomicity, or ordering guarantees. +- `references/troubleshooting.md` — read on a rebase conflict, after a squash-merge, on local and + remote divergence, when restructuring a stack, or when driving stacks from another tool. diff --git a/.claude/skills/gh-stack/references/commands.md b/.claude/skills/gh-stack/references/commands.md new file mode 100644 index 0000000..fcb219a --- /dev/null +++ b/.claude/skills/gh-stack/references/commands.md @@ -0,0 +1,177 @@ +# Command behavior + +`gh stack --help` is authoritative for flags and arguments. (`gh stack help ` only prints the top-level help.) This file only covers behavior `--help` does not +explain: preconditions, side effects, atomicity, and failure modes. + +## Contents + +- [init](#init) +- [add](#add) +- [push](#push) +- [submit](#submit) +- [link](#link) +- [sync](#sync) +- [rebase](#rebase) +- [view](#view) +- [checkout](#checkout) +- [unstack](#unstack) +- [merge](#merge) +- [Navigation](#navigation) + +## init + +Creates the stack and checks out the **last** branch in the list, so a single `init` can lay down +the whole chain: `gh stack init auth api frontend`. + +`init` processes branch arguments from bottom to top. Existing branches are adopted. If the first +branch does not exist, it is created from the trunk; each later new branch is created from the +branch immediately before it. There is no separate adopt mode — existence decides. `--base` +selects a non-default trunk. + +`init` also enables `git rerere`. Under a TTY the first run in a repo asks for confirmation; set +`git config rerere.enabled true` beforehand to skip it. + +## add + +- **Must run from the top branch** of the stack (or the trunk when the stack is still empty). + Anywhere else it exits **5** with `can only add branches on top of the stack`. Run `gh stack top` + first. +- **Uncommitted changes carry over.** Without `-Am`, `add` does not touch the working tree, so + staged and unstaged changes follow you onto the new branch. Commit or stash first for a clean start. +- **`add -Am` commits in place when the current branch has no commits yet** — for example + immediately after `init` — instead of creating a branch. This is deliberate: the first layer + usually needs its content before a second layer exists. +- `-A` and `-u` are mutually exclusive, and both require `-m`. + +## push + +Pushes every active (non-merged, non-queued) branch in one multi-ref push with per-branch +`--force-with-lease`. + +**Not atomic.** Some branches may update while another is rejected. A rejection means that branch +moved on the remote; fix that branch and rerun — rerunning is safe and skips what already landed. + +`push` never creates or updates pull requests. Use `submit` for that. + +## submit + +Pushes each active branch, then creates a PR for every branch that lacks one, basing it on the +first non-merged ancestor, then links them into a Stack on GitHub. + +- **Not atomic.** Branches are pushed sequentially with per-branch `--force-with-lease`. If a later + push is rejected, earlier pushes and PR updates stand. Fix the rejection and rerun the same command. +- **A fully merged stack cannot be extended.** When every PR in the current stack is already merged, + `submit` forks the remaining unmerged branches into a **new** stack rooted at the trunk and creates + it on GitHub, leaving the merged stack untouched. +- **Title generation with `--auto`:** a branch with a single commit uses that commit's subject as + the title and its body as the PR body. A branch with multiple commits humanizes the branch name + (hyphens and underscores become spaces). There is no flag for a custom title or body; use + `gh pr edit` afterwards. +- `--open` marks new *and existing* PRs ready for review; without it new PRs are drafts. +- Requires stacked PRs to be enabled on the repository. If not, `submit` exits **9** when + non-interactive (under a TTY it offers to create ordinary unstacked PRs instead). + +## link + +Creates or updates a stack on GitHub **without any local tracking state**. This is the path for +branches managed by another tool or living in another worktree — see `troubleshooting.md`. + +- Arguments are given bottom to top. Each is a branch name or a PR number; a numeric argument is + tried as a PR number first and falls back to a branch name. +- **A numeric first argument is treated as a stack number only when a stack with that number + exists.** In that case the remaining arguments are appended to the top of that stack and you do + not re-list its current PRs: `gh stack link 7 feature-c`. Arguments already in the stack are + skipped; arguments belonging to a different stack are rejected. +- Branch arguments are pushed automatically (non-force, atomic). Missing PRs are created with + auto-generated titles and correctly chained bases; existing PRs with a wrong base are corrected. +- Stack membership is **additive only** — `link` never removes a PR from a stack. + +## sync + +The routine command. Steps, in order: + +1. **Fetch** from the remote. +2. **Reconcile with the GitHub stack.** PRs added to the stack on github.com are pulled down and + appended locally. On divergence, aborts when non-interactive (see `troubleshooting.md`). +3. **Fast-forward the trunk.** Skipped when already current; warns when diverged. +4. **Cascade rebase when needed.** This runs if the trunk moved, a stack branch was fast-forwarded + from its remote, or a branch no longer contains its expected parent. Merged PRs are handled + automatically. On conflict, **all branches are restored** to their pre-rebase state and the + command exits **3**. +5. **Push** all active branches, atomically. +6. **Refresh PR state** from GitHub. +7. **Sync the stack object** — link open PRs into a stack, additively. Only when two or more PRs + exist. `sync` never opens PRs; that is `submit`. +8. **Prune** local branches for merged PRs, only when `--prune` is passed in a non-interactive + environment. + +## rebase + +Pulls from the remote and cascade-rebases. Use it when `sync` reported a conflict or when you need +to rebase only part of the stack. + +- `--upstack` rebases from the current branch to the top. This is what you run after editing a + lower layer. +- `--downstack` rebases from the trunk to the current branch. +- `--no-trunk` skips fetching and the trunk rebase entirely, aligning stack branches with each + other only. +- `--continue` after staging resolutions; `--abort` restores every branch. +- A merged PR is detected automatically and replayed with `--onto` against the correct target, so a + squash-merged parent does not produce spurious conflicts. +- Starting a rebase while one is in progress exits **7**. + +## view + +- `--json` writes the machine-readable payload to stdout. Its schema is in `SKILL.md`. +- Bare `view` opens a full-screen TUI when stdout is a TTY, and prints static text when piped. +- `--short` prints a compact one-line-per-branch summary and never opens the TUI, but it is + formatted for humans; parse `--json` instead. +- `view` refreshes PR state from GitHub as a side effect, best-effort — it does not fail when the + API is unreachable. + +## checkout + +Accepts a stack number, PR number, PR URL, or branch name. + +- A bare number resolves as a **stack number first**, then a PR number, then a branch name. +- Stack numbers, PR numbers, and PR URLs fetch from GitHub, pull the branches down, and set the + stack up locally. +- If a local stack already exists over those branches with a different composition, `checkout` + cannot be forced past it. Run `gh stack unstack --local` first, then retry. +- `checkout` has no flags. It relies on `remote.pushDefault` when several remotes exist. + +## unstack + +Removes the stack **grouping** only. It never deletes pull requests or branches. + +- With no argument it targets the active stack — the one containing the current branch — removing + it on GitHub and locally. +- With a stack number it works from anywhere in the repository, tracked locally or not, via the API. + Local tracking is also removed when present. +- `--local` removes local tracking only and never contacts GitHub. Combining `--local` with a stack + number that is not tracked locally is an error. +- An unknown stack number exits **2**. + +## merge + +- Scope with an argument: pass a PR number to merge that PR and every unmerged PR below it in the + stack, or pass a stack number to merge every unmerged PR in that stack. +- **All-or-nothing.** If any PR in that exact merge set cannot be merged, none are, and the reason + is reported. +- The method comes from `--squash`, `--rebase`, `--merge`, or `--merge-method `. Without + one, the last-used method is reused. +- Only basic PR state is checked before merging: open and not a draft. Bypassing merge requirements + is not supported for stacks. +- **A merge queue on the base branch overrides everything.** The stack is added to the queue rather + than merged; the queue chooses the method and any method flag you passed is ignored with a + warning. Queued PRs are submitted together but land as the queue processes them, so they may merge + in separate groups rather than all at once. +- `gh pr merge` cannot merge a stack. Always use `gh stack merge`. + +## Navigation + +`up`, `down`, `top`, `bottom`, and `trunk` are always non-interactive. `up` and `down` accept a +count (`gh stack up 3`). Movement clamps at the stack bounds, and merged branches are skipped when +navigating from an active branch, so `bottom` lands on the lowest *unmerged* branch. + +`gh stack switch` is a selection menu with no non-interactive path. Use the commands above instead. diff --git a/.claude/skills/gh-stack/references/stack-design.md b/.claude/skills/gh-stack/references/stack-design.md new file mode 100644 index 0000000..f64543a --- /dev/null +++ b/.claude/skills/gh-stack/references/stack-design.md @@ -0,0 +1,97 @@ +# Designing a stack + +How to decide what goes in each layer. Read this before running `gh stack init`. + +## Contents + +- [Plan the layers before writing code](#plan-the-layers-before-writing-code) +- [Branch naming](#branch-naming) +- [Staging changes deliberately](#staging-changes-deliberately) +- [When to add a layer](#when-to-add-a-layer) +- [One stack, one story](#one-stack-one-story) + +## Plan the layers before writing code + +A stack is a dependency chain. If code in one layer depends on code in another, the dependency must +live in the same branch or a lower one. That constraint is much cheaper to satisfy by planning than +by restructuring later, because there is no non-interactive in-place reorder — fixing the order +means`unstack` and `init` again. + +Decide the layers first, then write code into them: + +``` +(main) <- todo-app/models <- todo-app/api <- todo-app/frontend <- todo-app/integration +``` + +- `todo-app/models` — shared types and schema +- `todo-app/api` — routes that use the models +- `todo-app/frontend` — components that call the routes +- `todo-app/integration` — tests exercising the whole feature + +This is illustrative. Infer the stack topic and layer names from the actual task; do not reuse +`todo-app` or these layer names literally. + +The failure mode to avoid is writing everything on one branch and trying to split it afterwards. +If a task is large enough to warrant a stack, create the stack at the start. + +## Branch naming + +Prefer a shared topic prefix plus the layer's concern: +`/` — for example, `billing/schema`, `billing/api`, `billing/ui`. +This keeps related branches recognizable without using generic names that could belong to any +stack. **User and repository branch naming conventions take precedence; follow them instead.** + +Names are used exactly as given — nothing is prepended or transformed, and slashes are kept, so +`gh stack add refactor/foo` creates a branch literally named `refactor/foo`. + +If you pass `-m` without a branch name, the name is generated from the commit message in +date-and-slug form (for example `03-24-add_api_routes`). Prefer naming the branch yourself. + +## Staging changes deliberately + +Use `git add` and `git commit` directly rather than the `add -Am` shortcut. The point is control +over which changes land in which branch. With several modified files in the working tree, stage the +subset that belongs to the current layer, commit it, then create the next branch and stage the rest +there: + +```bash +git add internal/models/user.go internal/models/session.go +git commit -m "Add user and session models" + +gh stack add api-routes +git add internal/api/routes.go internal/api/handlers.go +git commit -m "Add user API routes" +``` + +Multiple commits per branch are fine. What matters is that every commit in a branch serves the same +concern, and that a change belonging to a different concern goes in a different branch. + +Note that `gh stack add ` without `-Am` does not touch the working tree, so uncommitted +changes carry over to the new branch. Commit or stash first if you want the new layer to start clean. + +## When to add a layer + +Add a branch when you start a **different concern that depends on what you have built so far**. +Signals: + +- Moving from backend to frontend, or from core logic to tests or documentation +- The next changes have a different reviewer audience +- The current branch's diff is already large enough to review on its own + +A layer that cannot be described in one sentence is usually two layers. + +## One stack, one story + +A stack should read as a coherent progression: a reviewer walks the PRs bottom to top and sees the +feature being built. + +**Use a single stack** when every branch serves the same feature or project, even if the layers span +different concerns. + +**Start a separate stack** for unrelated work — a different feature, an unrelated bug fix, an +independent refactor. Do not mix efforts into one stack just because you happened to work on both. +Use `gh stack init` for the new effort, or `gh stack checkout ` to move between existing +stacks. + +A trivial incidental fix can ride along in the current stack. Once it grows into its own project, it +deserves its own stack. diff --git a/.claude/skills/gh-stack/references/troubleshooting.md b/.claude/skills/gh-stack/references/troubleshooting.md new file mode 100644 index 0000000..fc97b41 --- /dev/null +++ b/.claude/skills/gh-stack/references/troubleshooting.md @@ -0,0 +1,159 @@ +# Troubleshooting and recovery + +## Contents + +- [Rebase conflicts (exit 3)](#rebase-conflicts-exit-3) +- [After a squash merge](#after-a-squash-merge) +- [Local and remote stacks have diverged](#local-and-remote-stacks-have-diverged) +- [Restructuring a stack](#restructuring-a-stack) +- [Branch belongs to several stacks (exit 6)](#branch-belongs-to-several-stacks-exit-6) +- [Driving stacks from another tool or worktree](#driving-stacks-from-another-tool-or-worktree) +- [Stack file is locked (exit 8)](#stack-file-is-locked-exit-8) +- [An interrupted modify session (exit 10)](#an-interrupted-modify-session-exit-10) + +## Rebase conflicts (exit 3) + +`rebase` and `sync` both exit 3 on conflict. `sync` restores every branch to its pre-rebase state +first, so a failed `sync` leaves nothing half-applied; a failed `rebase` stops mid-flight and waits. + +```bash +gh stack rebase +# exit 3 — conflicted paths are listed on stderr +git add +gh stack rebase --continue # repeat if the next branch also conflicts +``` + +`gh stack rebase --abort` restores every branch in the stack, not just the current one. + +Because `init` enables `git rerere`, a conflict you resolve once is replayed automatically the next +time the same conflict appears — which is common, since a change low in the stack is rebased through +every branch above it. Without `rerere`, repeated conflicts may need manual resolution on each +affected layer. + +## After a squash merge + +A squash merge replaces the branch's commits with one new commit, so the originals no longer exist +in the trunk's history and an ordinary rebase would try to replay them again. + +`gh stack sync` detects this and rebases with `--onto` against the correct target, skipping the +merged branch: + +```bash +gh stack sync +gh stack view --json # merged branch reports "isMerged": true, "state": "MERGED" +``` + +No manual action is needed. If the replay conflicts, `sync` restores all branches and exits 3. +Run `gh stack rebase` to rerun the rebase, which will stop at the conflict and allow you to resolve +and then `--continue` until complete. Use `gh stack sync --prune` to also delete local branches for +merged PRs. + +## Local and remote stacks have diverged + +Divergence means the local stack and the stack on GitHub changed in different ways — for example +branches were added locally while a PR was added to the stack on github.com. + +When non-interactive, `sync` prints both chains, changes nothing, and exits **0** with +`Sync aborted`. Success here does not mean the sync happened; check for that message, or re-run +`gh stack view --json` and compare. + +Two resolution paths: + +- **Keep the remote version.** Drop local tracking and pull the stack back down. + + ```bash + gh stack unstack --local # keeps the stack on GitHub + gh stack checkout # or a PR number + ``` + +- **Keep the local version.** Remove the grouping on GitHub, then recreate it from local state. + + ```bash + gh stack unstack # removes the grouping; PRs and branches survive + gh stack submit --auto + ``` + +Neither path deletes pull requests or branches. +Remote unstacking leaves PRs that are merging (auto-merge enabled) or are queued (in a merge queue) +stacked. If needed, clear that state before retrying. + +## Restructuring a stack + +There is no non-interactive reorder, rename, or removal. `add` run from the wrong branch suggests +`gh stack modify`, but that is TUI-only. Tear the stack down and rebuild it instead: + +```bash +gh stack unstack # removes local tracking and the GitHub grouping +# Rename or drop branches, and rewrite ancestry as needed. +gh stack init --base main branch-1 branch-2 branch-3 +gh stack submit --auto # re-link on GitHub +``` + +`init` adopts branches that already exist, so the rebuild reuses them rather than creating new ones. +Existing PRs survive. Once Git ancestry is correct, `submit` updates their base branches and +re-links the stack on GitHub. + +Changing metadata does **not** change Git ancestry. Reorder commits first, then rebuild the stack. +For example, to change `main <- models <- migration <- ui` into +`main <- migration <- models <- ui`: + +```bash +old_models=$(git rev-parse models) +old_migration=$(git rev-parse migration) +git rebase --onto main "$old_models" migration +git rebase --onto migration main models +git rebase --onto models "$old_migration" ui +gh stack unstack +gh stack init --base main migration models ui +``` + +The first rebase moves migration-only commits onto trunk, the second replays model commits above +them, and the third replays UI-only commits above models. Preserve the old boundary SHAs before +moving any branch. For a different reorder, identify each layer's range with +`git log ..`, then replay the ranges bottom to top. + +## Branch belongs to several stacks (exit 6) + +Commands exit 6 when the current branch cannot identify a single stack — typically because it is the +trunk of more than one stack. There is no flag to disambiguate. + +```bash +gh stack checkout +``` + +Then rerun. Commands that take an explicit stack number (`merge 7`, `unstack 7`) sidestep the +problem entirely, since they do not infer the stack from the current branch. + +## Driving stacks from another tool or worktree + +`gh stack link` creates and updates stacks purely through the API, with no local tracking state. +Use it when branches are managed by jj, Sapling, git-town, a separate worktree, or any workflow +where the local `.git/gh-stack` file would be wrong or absent. + +```bash +gh stack link branch-a branch-b branch-c # bottom to top +gh stack link --base develop --open a b c # non-default trunk, ready for review +gh stack link 10 20 30 # by PR number +gh stack link 7 feature-d # append to existing stack #7 +``` + +Because `link` writes no local state, the local navigation commands (`up`, `down`, `top`, `bottom`) +will not work on the result. Use `gh stack checkout ` if you later want local tracking. + +## Stack file is locked (exit 8) + +Another `gh stack` process holds the exclusive lock on `.git/gh-stack.lock`. The lock times out +after about five seconds, so wait and retry. A persistent exit 8 means another process still holds +the lock; identify and stop that process before retrying. + +## An interrupted modify session (exit 10) + +`gh stack modify` is TUI-only and should never be invoked by an agent. If a repository is left in +this state by someone else, restore it: + +```bash +gh stack modify --abort +``` + +Related: `submit` also detects a pending modify state, and under a TTY asks before overwriting the +stack on GitHub with local state. diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 838445f..c8ac895 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -1,70 +1,67 @@ -name: Build and Release +name: Release on: push: tags: - - 'v*.*.*' # Matches tags like v1.0.0, v2.1.3, etc. + - 'v*.*.*' # Matches tags like v1.0.0, v2.1.3, v0.2.0-rc.1 + +permissions: + contents: write # create the release and upload its assets jobs: - build: + release: runs-on: ubuntu-latest steps: - name: Check out the repository - uses: actions/checkout@v3 - - - name: Set up Go - uses: actions/setup-go@v4 + uses: actions/checkout@v7 with: - go-version: 1.22 + fetch-depth: 0 # the tag check needs main, and git describe needs the tags - - name: Run build script - run: bash ./build.sh - - - name: Upload build artifacts - uses: actions/upload-artifact@v4 - with: - name: build-artifacts - path: build/*.tar.gz - - release: - needs: build - runs-on: ubuntu-latest + - name: Check that a release tag is on main + run: | + # A pre-release tag (for example v0.2.0-dev.1a2b3c4 from make + # dev-release) can come from any branch: /releases/latest never + # returns a pre-release, so installed CLIs do not update to it. + if [[ "$GITHUB_REF_NAME" == *-* ]]; then + echo "Pre-release tag $GITHUB_REF_NAME: any branch is allowed." + exit 0 + fi + git fetch --no-tags origin main + if ! git merge-base --is-ancestor "$GITHUB_SHA" origin/main; then + echo "::error::Tag $GITHUB_REF_NAME is not on main. Release tags must be on the release branch." + exit 1 + fi - steps: - - name: Download build artifacts - uses: actions/download-artifact@v4 + - name: Set up Go + uses: actions/setup-go@v7 with: - name: build-artifacts - path: build + go-version: '1.27' + check-latest: true - - name: Create GitHub Release - id: create_release - uses: actions/create-release@v1 - env: - GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} - with: - tag_name: ${{ github.ref }} - release_name: Release ${{ github.ref }} - draft: false - prerelease: false + - name: Run the checks + run: make check - - name: Upload Release Assets for amd64 - uses: actions/upload-release-asset@v1 - env: - GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} - with: - upload_url: ${{ steps.create_release.outputs.upload_url }} - asset_path: build/fly-linux-amd64.tar.gz - asset_name: fly-linux-amd64.tar.gz - asset_content_type: application/gzip + # Installed CLIs download fly-linux-.tar.gz and look for the binary + # fly-linux- in it. Do not change these asset names. + - name: Build the release archives + run: make release VERSION="$GITHUB_REF_NAME" - - name: Upload Release Assets for arm64 - uses: actions/upload-release-asset@v1 + - name: Create the GitHub release env: - GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} - with: - upload_url: ${{ steps.create_release.outputs.upload_url }} - asset_path: build/fly-linux-arm64.tar.gz - asset_name: fly-linux-arm64.tar.gz - asset_content_type: application/gzip \ No newline at end of file + GH_TOKEN: ${{ github.token }} + run: | + # A pre-release (for example v0.2.0-rc.1) is not returned by + # /releases/latest, so installed CLIs do not update to it. + prerelease=() + if [[ "$GITHUB_REF_NAME" == *-* ]]; then + prerelease=(--prerelease) + fi + gh release create "$GITHUB_REF_NAME" \ + build/fly-linux-amd64.tar.gz \ + build/fly-linux-arm64.tar.gz \ + build/checksums.txt \ + --title "$GITHUB_REF_NAME" \ + --generate-notes \ + --verify-tag \ + "${prerelease[@]}" diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..813253f --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,43 @@ +name: CI + +on: + push: + branches: [develop, main] + # All pull requests: stacked pull requests target each other, not develop. + pull_request: + +permissions: + contents: read + +# A new push cancels the older run for the same branch or pull request. +concurrency: + group: ci-${{ github.ref }} + cancel-in-progress: true + +jobs: + check: + runs-on: ubuntu-latest + + steps: + - name: Check out the repository + uses: actions/checkout@v7 + with: + fetch-depth: 0 # git describe needs the tags for the version + + - name: Set up Go + uses: actions/setup-go@v7 + with: + go-version: '1.27' + check-latest: true + + - name: Run the checks + run: make check + + # make check runs as the runner user, so the tests that need root skip. + # They cover "sudo fly update": the owner of the binary, and a link that + # a hostile owner of the binary folder puts in place of the new file. + - name: Run the tests that need root + run: sudo env "PATH=$PATH" "GOCACHE=$(go env GOCACHE)" "GOMODCACHE=$(go env GOMODCACHE)" go test -count=1 -run 'TestReplaceBinary' ./internal/release/ + + - name: Build the release archives + run: make release diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..57e30c2 --- /dev/null +++ b/Makefile @@ -0,0 +1,121 @@ +# Developer commands for fly. Run "make help" for the list. +# Keep this file compatible with GNU Make 3.81, the version on macOS. + +.PHONY: build test vet lint vuln fmt fmt-check check release dev-version dev-release release-key sign-release clean help + +BINARY := fly +PKG := github.com/flywp/server-cli + +# VERSION defaults to the git version. Override it for a release: +# make release VERSION=v0.2.0 +VERSION ?= $(shell git describe --tags --always --dirty 2>/dev/null || echo dev) +COMMIT := $(shell git rev-parse HEAD 2>/dev/null || echo unknown) +# The date of the commit, not of the build: the same commit gives the same +# binary on any computer, so make sign-release can build a release again and +# compare it with the one that CI published. +BUILD_DATE := $(shell TZ=UTC git log -1 --date=format-local:%Y-%m-%d --format=%cd 2>/dev/null || date -u +%Y-%m-%d) + +LDFLAGS := -X $(PKG)/internal/version.Version=$(VERSION) \ + -X $(PKG)/internal/version.CommitHash=$(COMMIT) \ + -X $(PKG)/internal/version.BuildDate=$(BUILD_DATE) + +# Dev pre-releases are named -dev., for +# example v0.2.0-dev.1a2b3c4. DEV_BASE defaults to the minor version after the +# latest release tag; override it, for example make dev-release DEV_BASE=v0.1.2. +DEV_BASE ?= $(shell latest=$$(git describe --tags --abbrev=0 --exclude '*-*' 2>/dev/null || echo v0.0.0); \ + echo "$$latest" | awk -F. '{ sub(/^v/, "", $$1); printf "v%d.%d.0", $$1, $$2 + 1 }') +DEV_VERSION = $(DEV_BASE)-dev.$(shell git rev-parse --short=7 HEAD) + +# Release platforms. Installed CLIs download fly--.tar.gz and look +# for the binary fly-- in it, so do not change these names. +RELEASE_PLATFORMS := linux/amd64 linux/arm64 + +# Release archives hold the binary as root:root, not as the user who built +# it: an installer that keeps the owner must not give the binary to a user. +# GNU tar (CI) and bsdtar (macOS) name the options differently. +TAR_OWNER := $(shell tar --version 2>/dev/null | grep -q GNU && echo '--owner=0 --group=0 --numeric-owner' || echo '--uid 0 --gid 0 --numeric-owner') + +# "go run pkg@version" builds each tool with the Go version of this module, +# so a tool cannot be older than go.mod. +GOLANGCI_LINT := go run github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.13.2 +GOVULNCHECK := go run golang.org/x/vuln/cmd/govulncheck@v1.8.0 + +build: ## Build bin/fly for this platform + go build -trimpath -ldflags "$(LDFLAGS)" -o bin/$(BINARY) . + +test: ## Run the tests with the race detector + go test ./... -race -count=1 + +vet: ## Run go vet + go vet ./... + +lint: ## Run golangci-lint + $(GOLANGCI_LINT) run + +vuln: ## Scan for known vulnerabilities + $(GOVULNCHECK) ./... + +fmt: ## Format the code + gofmt -w . + +fmt-check: ## Fail if the code is not formatted + @files=$$(gofmt -l .); if [ -n "$$files" ]; then echo "Run make fmt for:"; echo "$$files"; exit 1; fi + +check: fmt-check vet lint test vuln ## Run all checks (the gate before each merge) + +release: ## Build the static release archives and checksums.txt in build/ + rm -rf build + mkdir -p build + @for platform in $(RELEASE_PLATFORMS); do \ + os=$${platform%/*}; arch=$${platform#*/}; out=$(BINARY)-$$os-$$arch; \ + echo "Building $$out ($(VERSION))"; \ + CGO_ENABLED=0 GOOS=$$os GOARCH=$$arch go build -trimpath -ldflags "-s -w $(LDFLAGS)" -o build/$$out . || exit 1; \ + COPYFILE_DISABLE=1 tar $(TAR_OWNER) -czf build/$$out.tar.gz -C build $$out || exit 1; \ + done + cd build && (command -v sha256sum >/dev/null 2>&1 && sha256sum *.tar.gz || shasum -a 256 *.tar.gz) > checksums.txt + +dev-version: ## Print the dev pre-release version of HEAD + @echo $(DEV_VERSION) + +# The Release workflow publishes the tag as a pre-release. A tag on a commit +# whose release workflow does not know pre-releases would become the latest +# release and reach every installed CLI, so dev-release refuses such commits. +dev-release: ## Tag HEAD as a dev pre-release and push the tag (CI publishes it) + @set -e; \ + if [ -n "$$(git status --porcelain --untracked-files=no)" ]; then \ + echo "Commit or stash your changes first: the pre-release is built from the commit."; exit 1; fi; \ + if [ -z "$$(git branch -r --contains HEAD)" ]; then \ + echo "Push this commit to a branch first."; exit 1; fi; \ + if ! git show HEAD:.github/workflows/build.yml | grep -q -- '--prerelease'; then \ + echo "The release workflow at this commit does not publish pre-releases. Do not tag it."; exit 1; fi; \ + git tag -a "$(DEV_VERSION)" -m "Dev pre-release $(DEV_VERSION)"; \ + git push origin "$(DEV_VERSION)"; \ + echo "Pushed $(DEV_VERSION). The Release workflow publishes it as a pre-release:"; \ + echo " https://github.com/flywp/server-cli/releases/tag/$(DEV_VERSION)" + +# The agent installs a release by itself only when checksums.txt has a +# signature from a key that is kept outside GitHub. Make the key one time, put +# the printed public key line in internal/release/keys.go, and keep the +# private key in a password manager, with a backup. Never commit it. +# COMMENT names the key, in the key file and next to its line in keys.go. +COMMENT ?= server-cli release key for flywp +release-key: ## Make the release signing key: make release-key KEY= [COMMENT=...] + @test -n "$(KEY)" || { echo "Usage: make release-key KEY= [COMMENT=\"...\"]"; exit 1; } + go run ./tools/releasesign keygen -out "$(KEY)" -comment "$(COMMENT)" + +# Run this after the Release workflow publishes VERSION, on your own +# computer. tools/sign-release.sh builds the release again from your local +# tag, checks that GitHub serves the same binaries, signs checksums.txt, +# checks the signature with the keys of the tag, and uploads +# checksums.txt.sig. Agents install the release 24 hours after the +# signature. KEY=- reads the key from stdin. +sign-release: ## Sign a published release: make sign-release VERSION=v0.2.1 KEY= + @if [ "$(origin VERSION)" != "command line" ] || [ -z "$(KEY)" ]; then \ + echo "Usage: make sign-release VERSION=v0.2.1 KEY="; exit 1; fi + @tools/sign-release.sh "$(VERSION)" "$(KEY)" + +clean: ## Remove bin/ and build/ + rm -rf bin/ build/ + +help: ## Show the targets + @grep -E '^[a-z-]+:.*## ' $(MAKEFILE_LIST) | awk -F':.*## ' '{printf " %-12s %s\n", $$1, $$2}' diff --git a/README.md b/README.md index f146ce4..e271a39 100644 --- a/README.md +++ b/README.md @@ -1,112 +1,78 @@ -# server-cli +# fly — the FlyWP server CLI -Easy CLI tool for servers managed by FlyWP. +`fly` manages the sites on a server that [FlyWP](https://flywp.com) provisions. It also runs the +FlyWP monitoring agent. -## Installation - -### Prerequisites - -- [Docker](https://www.docker.com/get-started) -- [Docker Compose](https://docs.docker.com/compose/install/) - -### Quick Install - -You can easily install the `fly` CLI tool using the following command. This will download and run the `install.sh` script, which will automatically detect your operating system and architecture, download the latest release, and install it to `/usr/local/bin`: +## Install ```bash -curl -sL https://raw.githubusercontent.com/flywp/server-cli/main/install.sh | sudo bash +curl -fsSL https://raw.githubusercontent.com/flywp/server-cli/main/install.sh | sudo bash ``` -
- -Manual Installation - -### Manual Installation - -If you prefer to manually download and install the binary, follow these steps: - -1. Download the precompiled binaries from the [Releases](https://github.com/flywp/server-cli/releases) page. Choose the version suitable for your operating system and architecture. - -1. Download the [latest tarball]((https://github.com/flywp/server-cli/releases)) for your platform: - - ```bash - wget https://github.com/flywp/server-cli/releases/download/v0.1.0/fly-linux-amd64.tar.gz - ``` - -2. Extract the tarball: - ```bash - tar -xzf fly-linux-amd64.tar.gz - ``` - -3. Move the binary to a directory in your PATH: - ```bash - sudo mv fly-linux-amd64 /usr/local/bin/fly - ``` +The script installs the latest release for your architecture (Linux amd64 or arm64) to +`/usr/local/bin/fly`. It checks the download against `checksums.txt` first. -4. Verify the installation: - ```bash - fly version - ``` +To update later, run `sudo fly update`. -
- -## Usage - -### Base Docker Compose - -FlyWP has a base Docker Compose configuration for running MySQL, Redis, Ofelia, and Nginx Proxy that are shared for all sites hosted on the server. The base Docker Compose must be started before a site can be created. +
+Install by hand ```bash -fly base start # starts the base services (mysql, redis, nginx-proxy) -fly base stop # stops the base services -fly base restart # restarts the base services +arch=amd64 # or arm64 +base=https://github.com/flywp/server-cli/releases/latest/download +curl -fsSLO "$base/fly-linux-$arch.tar.gz" && curl -fsSLO "$base/checksums.txt" +sha256sum -c --ignore-missing checksums.txt +tar -xzf "fly-linux-$arch.tar.gz" && sudo install -m 0755 "fly-linux-$arch" /usr/local/bin/fly +fly version ``` -### Site Operations - -You can run the following commands from anywhere inside a site folder or by specifying the domain name. +
-```bash -fly start --domain example.com # starts the website -fly stop --domain example.com # stops the website -fly restart --domain example.com # restarts the website -fly wp --domain example.com # execute WP-CLI commands -fly logs --domain example.com # view logs from all containers or a single one -fly restart --domain example.com # restart a container -fly exec --domain example.com # execute commands inside a container. Default: "php" -``` +## Use -Or run the commands from within the site directory without specifying the domain: +The site commands need Docker and Docker Compose. Run them inside a site folder, or name the site +with `--domain`. ```bash -fly start # starts the website -fly stop # stops the website -fly restart # restarts the website -fly wp # execute WP-CLI commands -fly logs # view logs from all containers or a single one -fly restart # restart a container -fly exec # execute commands inside a container. Default: "php" +fly base start|stop|restart # the shared services: MySQL, Redis, Ofelia, Nginx Proxy +fly start|stop|restart # the site +fly restart # one container of the site +fly wp # WP-CLI, for example: fly wp plugin list +fly exec [container] # a command in a container (default: PHP) +fly logs [-f] [container] # the logs of the site, or of one container + +fly sites start|stop|restart # all sites +fly status # the state of the server and its services +fly update # install the latest release +fly version ``` -### WP-CLI +Put `--domain` before the command: `fly --domain example.com wp plugin list`. Everything after +the command goes to it unchanged. To pass a flag as the first argument, put `--` before it: +`fly wp -- --info`. -**wp-cli**: To access `wp-cli`, use the following command from anywhere in the website folder or specify the domain name. The CLI will find the appropriate WordPress folder to execute the `wp` command. +## Monitoring agent + +FlyWP installs `fly agent run` as the systemd service `fly-agent`. Each minute, it sends the health +of the server and of each site to FlyWP. It runs as the server user, not as root, and it updates +itself only to signed releases. ```bash -fly wp --domain example.com +systemctl status fly-agent +journalctl -u fly-agent -f ``` -### Global Commands +See [docs/monitoring-agent.md](docs/monitoring-agent.md). -A few helper commands to debug the server installation and start/stop all sites. +## Documentation -```bash -fly status # shows the status of the system -fly sites start # starts all sites -fly sites stop # stops all sites -fly sites restart # stops and starts all sites -``` +| | | +|---|---| +| [Monitoring agent](docs/monitoring-agent.md) | What it measures, its settings and files, updates, security | +| [Development](docs/development.md) | Build, tests, code layout | +| [Releasing](docs/releasing.md) | Release, signing, dev pre-releases, rollback | +| [Design decisions](docs/decisions.md) | Why the agent works the way it does | ## License -This project is licensed under the MIT License. See the [LICENSE](LICENSE) file for details. +MIT. See [LICENSE](LICENSE). diff --git a/agent_test.go b/agent_test.go new file mode 100644 index 0000000..ab03ddc --- /dev/null +++ b/agent_test.go @@ -0,0 +1,460 @@ +package main + +// 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/ed25519" + "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" + +// agentEnv is the environment of a valid agent with its own state directory. +// PATH holds no docker command: the agent must not need Docker. +func agentEnv(t *testing.T) []string { + t.Helper() + + if os.Geteuid() == 0 { + t.Skip("fly refuses to run as root") + } + + return []string{ + "PATH=" + t.TempDir(), + "HOME=" + t.TempDir(), + "FLY_AGENT_URL=http://127.0.0.1:9", + "FLY_AGENT_TOKEN=" + testToken, + "FLY_AGENT_SERVER_ID=17", + "STATE_DIRECTORY=" + t.TempDir(), + // A test never asks the real GitHub for a release. + "FLY_AGENT_AUTO_UPDATE=off", + } +} + +// without returns environ without the variable key. +func without(environ []string, key string) []string { + var out []string + for _, kv := range environ { + if !strings.HasPrefix(kv, key+"=") { + out = append(out, kv) + } + } + return out +} + +// lockedBuffer is a bytes.Buffer that a child process and the test can use +// at the same time. +type lockedBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (b *lockedBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +func (b *lockedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +// startAgent starts "fly agent run" and waits until it logs that it started. +func startAgent(t *testing.T, environ []string) (*exec.Cmd, *lockedBuffer) { + t.Helper() + + cmd := exec.Command(flyBin, "agent", "run") + cmd.Env = environ + stderr := &lockedBuffer{} + cmd.Stderr = stderr + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = cmd.Process.Kill() }) + + deadline := time.Now().Add(10 * time.Second) + for !strings.Contains(stderr.String(), "agent started") { + if time.Now().After(deadline) { + t.Fatalf("the agent did not start within 10s. stderr:\n%s", stderr) + } + time.Sleep(20 * time.Millisecond) + } + + return cmd, stderr +} + +func TestAgentSendsAgentStartedToTheControlPlane(t *testing.T) { + type request struct { + path, auth string + body map[string][]map[string]any + } + got := make(chan request, 10) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + req := request{path: r.URL.Path, auth: r.Header.Get("Authorization")} + _ = json.NewDecoder(r.Body).Decode(&req.body) + got <- req + _, _ = w.Write([]byte(`{"accepted": 1}`)) + })) + defer srv.Close() + + environ := append(agentEnv(t), "FLY_AGENT_URL="+srv.URL) + cmd, stderr := startAgent(t, environ) + defer func() { _ = cmd.Process.Signal(syscall.SIGTERM); _ = cmd.Wait() }() + + select { + case req := <-got: + if req.path != "/agent/v1/events" || req.auth != "Bearer "+testToken { + t.Errorf("request to %s with Authorization %q, want /agent/v1/events with the token", req.path, req.auth) + } + if events := req.body["events"]; len(events) != 1 || events[0]["name"] != "agent.started" { + t.Errorf("events = %v, want agent.started", events) + } + case <-time.After(10 * time.Second): + t.Fatalf("the control plane got no request within 10s. stderr:\n%s", stderr) + } +} + +// 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"]) + } +} + +func TestAgentInstallsASignedReleaseByItself(t *testing.T) { + // A release has archives for Linux only. + if runtime.GOOS != "linux" { + t.Skip("the releases have binaries for Linux only") + } + environ := agentEnv(t) + + pub, priv, err := ed25519.GenerateKey(nil) + if err != nil { + t.Fatal(err) + } + + // The new release: fly built as v9.9.9, signed 25 hours ago. + 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) + } + archiveName := release.BinaryName(runtime.GOOS, runtime.GOARCH) + ".tar.gz" + tarball := releaseArchive(t, newBin, release.BinaryName(runtime.GOOS, runtime.GOARCH)) + sum := sha256.Sum256(tarball) + checksums := []byte(hex.EncodeToString(sum[:]) + " " + archiveName + "\n") + sigFile := release.Sign(priv, "v9.9.9", time.Now().Add(-25*time.Hour), checksums) + + var github *httptest.Server + github = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/releases/latest": + assets := []map[string]string{} + for _, name := range []string{archiveName, release.ChecksumsAsset, release.SignatureAsset} { + assets = append(assets, map[string]string{"name": name, "browser_download_url": github.URL + "/download/" + name}) + } + _ = json.NewEncoder(w).Encode(map[string]any{"tag_name": "v9.9.9", "assets": assets}) + case "/download/" + archiveName: + _, _ = w.Write(tarball) + case "/download/" + release.ChecksumsAsset: + _, _ = w.Write(checksums) + case "/download/" + release.SignatureAsset: + _, _ = w.Write(sigFile) + default: + http.NotFound(w, r) + } + })) + defer github.Close() + + // The running agent: an older release that asks this GitHub and trusts + // the test key. + dir := t.TempDir() + exe := filepath.Join(dir, "fly") + ldflags := strings.Join([]string{ + "-X github.com/flywp/server-cli/internal/version.Version=v0.0.1", + "-X github.com/flywp/server-cli/internal/release.GithubAPI=" + github.URL + "/releases/latest", + "-X github.com/flywp/server-cli/internal/release.trustedKeys=" + release.KeyLine(pub), + }, " ") + if out, err := exec.Command("go", "build", "-ldflags", ldflags, "-o", exe, ".").CombinedOutput(); err != nil { + t.Fatalf("building the old release: %v\n%s", err, out) + } + + srv := httptest.NewServer(&updateServer{closed: true}) + defer srv.Close() + + // Work 3 seconds from now. The agent never checked, so it checks at its + // first tick. + environ = append(environ, "FLY_AGENT_URL="+srv.URL, fmt.Sprintf("FLY_AGENT_SERVER_ID=%d", (time.Now().Second()+3)%60), "FLY_AGENT_AUTO_UPDATE=on") + + agent := exec.Command(exe, "agent", "run") + agent.Env = environ + stderr := &lockedBuffer{} + agent.Stderr = stderr + if err := agent.Start(); err != nil { + t.Fatal(err) + } + done := make(chan error, 1) + go func() { done <- agent.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(70 * time.Second): + _ = agent.Process.Kill() + t.Fatalf("the agent did not install the release within 70s. 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. stderr:\n%s", out, stderr) + } + if !strings.Contains(stderr.String(), "installed the new release") { + t.Errorf("stderr does not log the install:\n%s", stderr) + } +} + +// 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 + environ func([]string) []string + want string + }{ + {"no token", func(e []string) []string { return without(e, "FLY_AGENT_TOKEN") }, "FLY_AGENT_TOKEN is not set"}, + {"no state directory", func(e []string) []string { return without(e, "STATE_DIRECTORY") }, "STATE_DIRECTORY is not set"}, + {"plain http", func(e []string) []string { return append(e, "FLY_AGENT_URL=http://example.com") }, "must be an https URL"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + res := runFly(t, t.TempDir(), tt.environ(agentEnv(t)), "agent", "run") + + if res.code != 1 { + t.Errorf("exit code = %d, want 1", res.code) + } + if !strings.Contains(res.stderr, tt.want) { + t.Errorf("stderr = %q, want it to contain %q", res.stderr, tt.want) + } + }) + } +} + +func TestAgentRejectsArguments(t *testing.T) { + res := runFly(t, t.TempDir(), agentEnv(t), "agent", "run", "extra") + if res.code != 1 || !strings.Contains(res.stderr, "unknown command") && !strings.Contains(res.stderr, "accepts 0 arg") { + t.Errorf("fly agent run extra: exit %d, stderr %q; want a usage error", res.code, res.stderr) + } +} + +func TestAgentRunsWithoutDockerAndStopsOnSIGTERM(t *testing.T) { + environ := agentEnv(t) + cmd, stderr := startAgent(t, environ) + + // Only one agent can use the state directory. + second := runFly(t, t.TempDir(), environ, "agent", "run") + if second.code != 1 || !strings.Contains(second.stderr, "a different agent is running") { + t.Errorf("second agent: exit %d, stderr %q; want exit 1 and a lock error", second.code, second.stderr) + } + + if err := cmd.Process.Signal(syscall.SIGTERM); err != nil { + t.Fatal(err) + } + + done := make(chan error, 1) + go func() { done <- cmd.Wait() }() + select { + case err := <-done: + if err != nil { + t.Errorf("agent exit after SIGTERM: %v, want exit status 0. stderr:\n%s", err, stderr) + } + case <-time.After(10 * time.Second): + t.Fatal("the agent did not stop within 10s after SIGTERM") + } + + out := stderr.String() + if !strings.Contains(out, "agent stopped") { + t.Errorf("stderr = %q, want an \"agent stopped\" line", out) + } + if strings.Contains(out, testToken) { + t.Error("the agent log contains the token") + } +} diff --git a/build.sh b/build.sh deleted file mode 100755 index 6a60336..0000000 --- a/build.sh +++ /dev/null @@ -1,45 +0,0 @@ -#!/bin/bash - -set -e - -# Set variables -CLI_NAME="fly" -REPO_NAME="flywp/server-cli" -VERSION=$(git describe --tags --always --dirty) -COMMIT_HASH=$(git rev-parse HEAD) -BUILD_DATE=$(date -u +"%Y-%m-%d") -LDFLAGS="-X github.com/${REPO_NAME}/internal/version.Version=${VERSION} -X github.com/${REPO_NAME}/internal/version.CommitHash=${COMMIT_HASH} -X github.com/${REPO_NAME}/internal/version.BuildDate=${BUILD_DATE}" - -# Build function -build() { - local GOOS=$1 - local GOARCH=$2 - local OUTPUT="${CLI_NAME}-${GOOS}-${GOARCH}" - - echo "Building for ${GOOS}/${GOARCH}..." - GOOS=${GOOS} GOARCH=${GOARCH} go build -ldflags "${LDFLAGS}" -o "build/${OUTPUT}" . - echo "Done building ${OUTPUT}" - - create_archive "${OUTPUT}" -} - -# Create tar.gz archive function -create_archive() { - local OUTPUT=$1 - - echo "Creating archive for ${OUTPUT}..." - tar -czvf "build/${OUTPUT}.tar.gz" -C build "${OUTPUT}" - echo "Done creating archive ${OUTPUT}.tar.gz" -} - -# Clean build directory -rm -rf build - -# Create build directory if not exists -mkdir -p build - -# Build for different platforms -build linux amd64 -build linux arm64 - -echo "All builds completed!" diff --git a/cmd/agent.go b/cmd/agent.go new file mode 100644 index 0000000..9819443 --- /dev/null +++ b/cmd/agent.go @@ -0,0 +1,48 @@ +package cmd + +import ( + "log/slog" + "os" + "os/signal" + "syscall" + + "github.com/flywp/server-cli/internal/agent" + "github.com/flywp/server-cli/internal/metrics" + "github.com/spf13/cobra" +) + +var agentCmd = &cobra.Command{ + Use: "agent", + Short: "Run the FlyWP monitoring agent", +} + +// agentRunCmd does not need Docker: the agent must also report when Docker +// is down. +var agentRunCmd = &cobra.Command{ + Use: "run", + Short: "Run the monitoring agent until it is stopped", + Long: `Run the FlyWP monitoring agent until it is stopped. systemd starts this +command (fly-agent.service). The agent reads FLY_AGENT_URL, FLY_AGENT_TOKEN, +FLY_AGENT_SERVER_ID and STATE_DIRECTORY from the environment. + +Once a day the agent installs a newer signed release by itself. +FLY_AGENT_AUTO_UPDATE=off stops this.`, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, args []string) error { + cfg, err := agent.ConfigFromEnv(os.Getenv) + if err != nil { + return err + } + + ctx, stop := signal.NotifyContext(cmd.Context(), syscall.SIGTERM, os.Interrupt) + defer stop() + + log := slog.New(slog.NewTextHandler(os.Stderr, nil)) + return agent.Run(ctx, cfg, log, metrics.New("/", cfg.StateDir, log)) + }, +} + +func init() { + agentCmd.AddCommand(agentRunCmd) + rootCmd.AddCommand(agentCmd) +} diff --git a/cmd/base.go b/cmd/base.go index 031c8f7..1bc7088 100644 --- a/cmd/base.go +++ b/cmd/base.go @@ -1,7 +1,7 @@ package cmd import ( - "os" + "fmt" "github.com/fatih/color" "github.com/flywp/server-cli/internal/docker" @@ -18,50 +18,51 @@ var baseCmd = &cobra.Command{ var baseStartCmd = &cobra.Command{ Use: "start", Short: "Start base services", - Run: func(cmd *cobra.Command, args []string) { + RunE: func(cmd *cobra.Command, args []string) error { if err := docker.RunCompose(baseCompose, "up", "-d"); err != nil { - color.Red("Error starting base services: %v", err) - os.Exit(1) - return + return fmt.Errorf("starting base services: %w", err) } if err := docker.RunCompose(baseCompose, "ps"); err != nil { - color.Red("Error checking status of base services: %v", err) + return fmt.Errorf("checking status of base services: %w", err) } color.Green("Base services started successfully") + return nil }, } var baseStopCmd = &cobra.Command{ Use: "stop", Short: "Stop base services", - Run: func(cmd *cobra.Command, args []string) { + RunE: func(cmd *cobra.Command, args []string) error { if err := docker.RunCompose(baseCompose, "down"); err != nil { - color.Red("Error stopping base services:", err) + return fmt.Errorf("stopping base services: %w", err) } color.Green("Base services stopped successfully") + return nil }, } var baseRestartCmd = &cobra.Command{ Use: "restart", Short: "Restart base services", - Run: func(cmd *cobra.Command, args []string) { + RunE: func(cmd *cobra.Command, args []string) error { if err := docker.RunCompose(baseCompose, "down"); err != nil { - color.Red("Error stopping base services:", err) + return fmt.Errorf("stopping base services: %w", err) } if err := docker.RunCompose(baseCompose, "up", "-d"); err != nil { - color.Red("Error starting base services:", err) + return fmt.Errorf("starting base services: %w", err) } if err := docker.RunCompose(baseCompose, "ps"); err != nil { - color.Red("Error checking status of base services:", err) + return fmt.Errorf("checking status of base services: %w", err) } color.Green("Base services restarted successfully") + return nil }, } @@ -69,6 +70,7 @@ func init() { baseCmd.AddCommand(baseStartCmd) baseCmd.AddCommand(baseStopCmd) baseCmd.AddCommand(baseRestartCmd) + requireDocker(baseStartCmd, baseStopCmd, baseRestartCmd) rootCmd.AddCommand(baseCmd) } diff --git a/cmd/global.go b/cmd/global.go index d045ca8..4ca707e 100644 --- a/cmd/global.go +++ b/cmd/global.go @@ -1,9 +1,9 @@ package cmd import ( + "errors" "fmt" "os" - "os/exec" "path/filepath" "strings" @@ -60,18 +60,13 @@ var statusCmd = &cobra.Command{ color.Green("Nginx directory exists") } - // check if docker is installed - if output, err := exec.Command("docker", "version", "--format", "{{.Server.Version}}").CombinedOutput(); err != nil { - color.Red("Docker is not installed") - } else { - color.Green("Docker is installed, version: %s", strings.TrimSpace(string(output))) - } - - // check if docker is running - if _, err := exec.Command("docker", "version").CombinedOutput(); err != nil { - color.Red("Docker is not running") - } else { - color.Green("Docker is running") + // check the docker CLI, the compose plugin and the daemon separately + for _, s := range docker.Status(cmd.Context()) { + if s.Err != nil { + color.Red("%s is not available: %s", s.Part, s.Err.Detail) + } else { + color.Green("%s is available: %s", s.Part, s.Info) + } } }, } @@ -84,96 +79,75 @@ var sitesCmd = &cobra.Command{ var sitesStartCmd = &cobra.Command{ Use: "start", Short: "Start all sites", - Run: func(cmd *cobra.Command, args []string) { - startAllSites() + RunE: func(cmd *cobra.Command, args []string) error { + return forEachSite("Starting", "up", "-d") }, } var sitesStopCmd = &cobra.Command{ Use: "stop", Short: "Stop all sites", - Run: func(cmd *cobra.Command, args []string) { - stopAllSites() + RunE: func(cmd *cobra.Command, args []string) error { + return forEachSite("Stopping", "down") }, } var restartSitesCmd = &cobra.Command{ Use: "restart", Short: "Restart all sites", - Run: func(cmd *cobra.Command, args []string) { - stopAllSites() - startAllSites() + RunE: func(cmd *cobra.Command, args []string) error { + // Start the sites even if some of them failed to stop. + stopErr := forEachSite("Stopping", "down") + startErr := forEachSite("Starting", "up", "-d") + return errors.Join(stopErr, startErr) }, } -func startAllSites() { - sitesDir := "/home/fly" - foundSite := false +// sitesDir holds one directory per site, each with its own docker-compose.yml. +var sitesDir = "/home/fly" +// forEachSite runs docker compose with args for every site in sitesDir. +// It continues after a failed site and returns all failures together. +func forEachSite(verb string, args ...string) error { entries, err := os.ReadDir(sitesDir) if err != nil { - color.Red("Error reading directory %s: %v\n", sitesDir, err) - return + return fmt.Errorf("reading sites directory: %w", err) } + var errs []error + foundSite := false for _, entry := range entries { - if entry.IsDir() { - // Skip hidden directories - if strings.HasPrefix(entry.Name(), ".") { - continue - } - - path := filepath.Join(sitesDir, entry.Name()) - composePath := filepath.Join(path, "docker-compose.yml") - if _, err := os.Stat(composePath); err == nil { - color.Yellow("Starting site in %s\n", filepath.Base(path)) - docker.RunCompose(composePath, "up", "-d") - foundSite = true - } + // Skip files and hidden directories + if !entry.IsDir() || strings.HasPrefix(entry.Name(), ".") { + continue } - } - if !foundSite { - fmt.Println("No sites found to start.") - } -} - -func stopAllSites() { - sitesDir := "/home/fly" - foundSite := false - - entries, err := os.ReadDir(sitesDir) - if err != nil { - color.Red("Error reading directory %s: %v\n", sitesDir, err) - return - } - - for _, entry := range entries { - if entry.IsDir() { - // Skip hidden directories - if strings.HasPrefix(entry.Name(), ".") { - continue - } + composePath := filepath.Join(sitesDir, entry.Name(), "docker-compose.yml") + if _, err := os.Stat(composePath); err != nil { + continue + } - path := filepath.Join(sitesDir, entry.Name()) - composePath := filepath.Join(path, "docker-compose.yml") - if _, err := os.Stat(composePath); err == nil { - color.Yellow("Stopping site in %s\n", filepath.Base(path)) - docker.RunCompose(composePath, "down") - foundSite = true - } + foundSite = true + color.Yellow("%s site in %s", verb, entry.Name()) + if err := docker.RunCompose(composePath, args...); err != nil { + // %v, not %w: a failed child process must not hide this summary. + // The verb tells restart's stop failures from its start failures. + errs = append(errs, fmt.Errorf("%s %s: %v", strings.ToLower(verb), entry.Name(), err)) } } if !foundSite { - fmt.Println("No sites found to stop.") + fmt.Println("No sites found.") } + + return errors.Join(errs...) } func init() { sitesCmd.AddCommand(sitesStartCmd) sitesCmd.AddCommand(sitesStopCmd) sitesCmd.AddCommand(restartSitesCmd) + requireDocker(sitesStartCmd, sitesStopCmd, restartSitesCmd) rootCmd.AddCommand(sitesCmd) rootCmd.AddCommand(statusCmd) diff --git a/cmd/global_test.go b/cmd/global_test.go new file mode 100644 index 0000000..72b5b11 --- /dev/null +++ b/cmd/global_test.go @@ -0,0 +1,84 @@ +package cmd + +import ( + "errors" + "os/exec" + "path/filepath" + "strings" + "testing" + + "github.com/flywp/server-cli/internal/testutil" +) + +func TestForEachSiteContinuesAfterFailure(t *testing.T) { + fake := testutil.NewFakeDocker(t) + fake.Install(t) + t.Setenv(testutil.EnvExit, "3") + + root := testutil.TempDir(t) + testutil.WriteSite(t, filepath.Join(root, "a.example.com"), "php") + testutil.WriteSite(t, filepath.Join(root, "b.example.com"), "php") + testutil.WriteSite(t, filepath.Join(root, ".fly"), "mysql") // hidden: skipped + + old := sitesDir + sitesDir = root + t.Cleanup(func() { sitesDir = old }) + + err := forEachSite("Starting", "up", "-d") + if err == nil { + t.Fatal("forEachSite() = nil, want an error for the failed sites") + } + + // The summary must be printed, so it must not unwrap to the child's exit status. + var exitErr *exec.ExitError + if errors.As(err, &exitErr) { + t.Errorf("forEachSite() error unwraps to *exec.ExitError, want a plain summary: %v", err) + } + for _, site := range []string{"a.example.com", "b.example.com"} { + if !strings.Contains(err.Error(), "starting "+site) { + t.Errorf("error %q does not name the phase and site %s", err, site) + } + } + + if calls := fake.Calls(t); len(calls) != 2 { + t.Errorf("docker calls = %q, want one call per site", calls) + } +} + +func TestSitesRestartNamesThePhaseOfEachFailure(t *testing.T) { + fake := testutil.NewFakeDocker(t) + fake.Install(t) + t.Setenv(testutil.EnvExit, "3") + + root := testutil.TempDir(t) + testutil.WriteSite(t, filepath.Join(root, "broken.example.com"), "php") + + old := sitesDir + sitesDir = root + t.Cleanup(func() { sitesDir = old }) + + // The site fails to stop and then fails to start: two different failures. + err := restartSitesCmd.RunE(restartSitesCmd, nil) + if err == nil { + t.Fatal("sites restart = nil, want an error for the failed site") + } + + got := strings.Split(err.Error(), "\n") + want := []string{ + "stopping broken.example.com: exit status 3", + "starting broken.example.com: exit status 3", + } + if strings.Join(got, "\n") != strings.Join(want, "\n") { + t.Errorf("error lines = %q, want %q", got, want) + } +} + +func TestForEachSiteMissingDirectory(t *testing.T) { + old := sitesDir + sitesDir = filepath.Join(t.TempDir(), "missing") + t.Cleanup(func() { sitesDir = old }) + + if err := forEachSite("Starting", "up", "-d"); err == nil { + t.Fatal("forEachSite() = nil, want an error for a missing sites directory") + } +} diff --git a/cmd/root.go b/cmd/root.go index d9aca14..2013d85 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -1,9 +1,14 @@ package cmd import ( + "errors" "fmt" + "io" "os" + "os/exec" + "github.com/fatih/color" + "github.com/flywp/server-cli/internal/docker" "github.com/spf13/cobra" ) @@ -11,31 +16,71 @@ var rootCmd = &cobra.Command{ Use: "fly", Short: "Fly CLI for managing WordPress sites", Long: `A CLI tool for managing WordPress sites using Docker and custom commands.`, + // Execute reports errors itself, so that errors from child processes + // are not printed twice and usage is not printed for runtime errors. + SilenceErrors: true, + SilenceUsage: true, PersistentPreRunE: func(cmd *cobra.Command, args []string) error { if cmd.Name() != "update" && os.Geteuid() == 0 { return fmt.Errorf("you should not run this command as root") } + if cmd.Annotations[requiresAnnotation] == "docker" { + return docker.Check(cmd.Context()) + } + return nil }, } -// Execute the root command. -func Execute() { - err := rootCmd.Execute() - if err != nil { - os.Exit(1) +// requiresAnnotation names what a command needs to run. The root command +// checks it before the command runs. +const requiresAnnotation = "requires" + +// requireDocker marks cmds as commands that need Docker. +func requireDocker(cmds ...*cobra.Command) { + for _, c := range cmds { + if c.Annotations == nil { + c.Annotations = map[string]string{} + } + c.Annotations[requiresAnnotation] = "docker" } } -func init() { - // Here you will define your flags and configuration settings. - // Cobra supports persistent flags, which, if defined here, - // will be global for your application. +// Execute runs the root command and exits with the resulting status code. +func Execute() { + os.Exit(exitCode(rootCmd.Execute(), os.Stderr)) +} - // rootCmd.PersistentFlags().StringVar(&cfgFile, "config", "", "config file (default is $HOME/.server-cli.yaml)") +// exitCode reports err on stderr and returns the exit status for it. +func exitCode(err error, stderr io.Writer) int { + if err == nil { + return 0 + } - // Cobra also supports local flags, which will only run - // when this action is called directly. - // rootCmd.Flags().BoolP("toggle", "t", false, "Help message for toggle") + // Docker is down or not installed: show one warning, not a raw error. + var unavailable *docker.UnavailableError + if errors.As(err, &unavailable) { + _, _ = color.New(color.FgYellow).Fprintf(stderr, "Warning: %v\n", unavailable) + return unavailable.ExitCode() + } + + // A child process (docker compose, wp-cli) has already reported its own + // error, so pass its exit status through without printing anything. + var exitErr *exec.ExitError + if errors.As(err, &exitErr) { + if code := exitErr.ExitCode(); code > 0 { + return code + } + return 1 + } + + _, _ = color.New(color.FgRed).Fprintf(stderr, "Error: %v\n", err) + return 1 +} + +func init() { + rootCmd.SetFlagErrorFunc(func(cmd *cobra.Command, err error) error { + return fmt.Errorf("%w\nRun '%s --help' for usage", err, cmd.CommandPath()) + }) } diff --git a/cmd/root_test.go b/cmd/root_test.go new file mode 100644 index 0000000..e57126d --- /dev/null +++ b/cmd/root_test.go @@ -0,0 +1,38 @@ +package cmd + +import ( + "bytes" + "errors" + "fmt" + "os/exec" + "strings" + "testing" +) + +func TestExitCode(t *testing.T) { + childErr := exec.Command("sh", "-c", "exit 7").Run() + + tests := []struct { + name string + err error + wantCode int + wantStderr string + }{ + {name: "success", err: nil, wantCode: 0}, + {name: "child exit status passes through silently", err: childErr, wantCode: 7}, + {name: "wrapped child exit status", err: fmt.Errorf("starting site: %w", childErr), wantCode: 7}, + {name: "other error is printed", err: errors.New("boom"), wantCode: 1, wantStderr: "Error: boom\n"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var stderr bytes.Buffer + if got := exitCode(tt.err, &stderr); got != tt.wantCode { + t.Errorf("exitCode() = %d, want %d", got, tt.wantCode) + } + if got := stderr.String(); !strings.Contains(got, tt.wantStderr) || (tt.wantStderr == "" && got != "") { + t.Errorf("stderr = %q, want %q", got, tt.wantStderr) + } + }) + } +} diff --git a/cmd/site.go b/cmd/site.go index 01cbd8f..0cf3818 100644 --- a/cmd/site.go +++ b/cmd/site.go @@ -1,9 +1,10 @@ package cmd import ( - "os" + "errors" + "fmt" + "slices" - "github.com/fatih/color" "github.com/flywp/server-cli/internal/docker" "github.com/flywp/server-cli/internal/utils" "github.com/spf13/cobra" @@ -12,19 +13,47 @@ import ( // Define domain flag as a global variable var domain string +var errNoSite = errors.New(`no docker-compose.yml file found + +You are not inside a site directory. +Please run this command from inside a site directory, e.g: + cd ~/example.com + fly start + +Or specify the domain name: + fly start --domain example.com`) + +// siteComposePath returns the compose file of the site selected by --domain +// or by the current directory. +func siteComposePath() (string, error) { + composePath, err := utils.FindComposeFile(domain) + if errors.Is(err, utils.ErrComposeNotFound) && domain == "" { + return "", errNoSite + } + + return composePath, err +} + var wpCmd = &cobra.Command{ - Use: "wp", + Use: "wp [wp-cli command] [args...]", Short: "Run wp-cli commands", - Run: func(cmd *cobra.Command, args []string) { - composePath := utils.FindComposeFile(domain) - if composePath == "" { - utils.ShowNoComposeError() - return - } + Long: `Run wp-cli commands in the site's PHP container. - if err := docker.RunWPCLI(composePath, args); err != nil { - color.Red("Error running wp-cli: %s", err) +All arguments after the first wp-cli word go to wp-cli unchanged, flags included. +Put --domain before the wp-cli command: + + fly --domain example.com wp plugin list --format=json + +To pass a flag as the first argument, put -- before it: + + fly wp -- --info`, + RunE: func(cmd *cobra.Command, args []string) error { + composePath, err := siteComposePath() + if err != nil { + return err } + + return docker.RunWPCLI(composePath, args) }, } @@ -32,16 +61,16 @@ var startCmd = &cobra.Command{ Use: "start", Short: "Start the site", Long: "Start the Docker container for the site", - Run: func(cmd *cobra.Command, args []string) { - composePath := utils.FindComposeFile(domain) - if composePath == "" { - utils.ShowNoComposeError() - return + RunE: func(cmd *cobra.Command, args []string) error { + composePath, err := siteComposePath() + if err != nil { + return err } if err := docker.RunCompose(composePath, "up", "-d"); err != nil { - color.Red("Error starting container:", err) + return fmt.Errorf("starting site: %w", err) } + return nil }, } @@ -49,16 +78,16 @@ var stopCmd = &cobra.Command{ Use: "stop", Short: "Stop the site", Long: "Stop the Docker container for the site", - Run: func(cmd *cobra.Command, args []string) { - composePath := utils.FindComposeFile(domain) - if composePath == "" { - utils.ShowNoComposeError() - return + RunE: func(cmd *cobra.Command, args []string) error { + composePath, err := siteComposePath() + if err != nil { + return err } if err := docker.RunCompose(composePath, "down"); err != nil { - color.Red("Error stopping container:", err) + return fmt.Errorf("stopping site: %w", err) } + return nil }, } @@ -67,80 +96,89 @@ var restartCmd = &cobra.Command{ Short: "Restart the site or a specific container", Long: "Restart the Docker containers for the site. Optionally, specify a container to restart only that container.", Args: cobra.MaximumNArgs(1), - Run: func(cmd *cobra.Command, args []string) { - composePath := utils.FindComposeFile(domain) - if composePath == "" { - utils.ShowNoComposeError() - return + RunE: func(cmd *cobra.Command, args []string) error { + composePath, err := siteComposePath() + if err != nil { + return err } - if len(args) == 1 { - containerName := args[0] - if err := docker.RunCompose(composePath, "restart", containerName); err != nil { - color.Red("Error restarting container:", err) - } - } else { - if err := docker.RunCompose(composePath, "restart"); err != nil { - color.Red("Error restarting Docker Compose setup:", err) - } + if err := docker.RunCompose(composePath, append([]string{"restart"}, args...)...); err != nil { + return fmt.Errorf("restarting site: %w", err) } + return nil }, } +// execServices are the services that fly exec accepts as its first argument. +var execServices = []string{"php", "nginx", "openlitespeed"} + +// splitService splits the arguments of fly exec into a service name and a +// command. The service is empty when the first argument is not a service. +func splitService(args []string) (service string, command []string) { + if len(args) > 0 && slices.Contains(execServices, args[0]) { + return args[0], args[1:] + } + + return "", args +} + var execCmd = &cobra.Command{ - Use: "exec", + Use: "exec [service] command [args...]", Short: "Execute a command in the Docker container", - Run: func(cmd *cobra.Command, args []string) { - composePath := utils.FindComposeFile(domain) - if composePath == "" { - utils.ShowNoComposeError() - return + Long: `Execute a command in a Docker container of the site. + +If the first argument is php, nginx or openlitespeed, the command runs in that +service. Otherwise it runs in the site's PHP service (php or openlitespeed). +All arguments after the first one go to the command unchanged, flags included. +Put --domain before the command: + + fly --domain example.com exec php ls -la`, + Args: cobra.MinimumNArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + composePath, err := siteComposePath() + if err != nil { + return err } - if len(args) == 0 { - color.Yellow("No command provided") - return + service, command := splitService(args) + if service == "" { + if service, err = docker.DefaultService(composePath); err != nil { + return err + } } - // if the next argument is "php", "nginx" or "litespeed", use it as the service name - // otherwise, use "php" as the default service name - composeArgs := []string{"exec"} - if args[0] == "php" || args[0] == "nginx" || args[0] == "litespeed" { - composeArgs = append(composeArgs, args[0]) - args = args[1:] - } else { - composeArgs = append(composeArgs, "php") + if len(command) == 0 { + return fmt.Errorf("no command given for service %q", service) } - composeArgs = append(composeArgs, args...) - - if err := docker.RunCompose(composePath, composeArgs...); err != nil { - color.Red("Error executing command: %v\n", err) - } + return docker.RunCompose(composePath, append([]string{"exec", service}, command...)...) }, } +var ( + logsFollow bool + logsTail string +) + var logsCmd = &cobra.Command{ - Use: "logs", + Use: "logs [service...]", Short: "Show logs of the Docker container", Long: `Show logs of Docker container(s). If no container is specified, it shows logs for all containers.`, - Args: cobra.MaximumNArgs(1), - Run: func(cmd *cobra.Command, args []string) { - composePath := utils.FindComposeFile(domain) - if composePath == "" { - utils.ShowNoComposeError() - os.Exit(1) + RunE: func(cmd *cobra.Command, args []string) error { + composePath, err := siteComposePath() + if err != nil { + return err } composeArgs := []string{"logs"} - if len(args) == 1 { - composeArgs = append(composeArgs, args[0]) + if logsFollow { + composeArgs = append(composeArgs, "--follow") } - - if err := docker.RunCompose(composePath, composeArgs...); err != nil { - color.Red("Error showing logs: %v\n", err) - os.Exit(1) + if logsTail != "" { + composeArgs = append(composeArgs, "--tail", logsTail) } + + return docker.RunCompose(composePath, append(composeArgs, args...)...) }, } @@ -148,10 +186,18 @@ func init() { // Add domain flag to rootCmd rootCmd.PersistentFlags().StringVar(&domain, "domain", "", "Specify domain for executing commands in a specific site") + // Flags after the first argument belong to wp-cli or to the command. + wpCmd.Flags().SetInterspersed(false) + execCmd.Flags().SetInterspersed(false) + + logsCmd.Flags().BoolVarP(&logsFollow, "follow", "f", false, "Follow log output") + logsCmd.Flags().StringVar(&logsTail, "tail", "", "Number of lines to show from the end of the logs") + rootCmd.AddCommand(wpCmd) rootCmd.AddCommand(startCmd) rootCmd.AddCommand(stopCmd) rootCmd.AddCommand(restartCmd) rootCmd.AddCommand(execCmd) rootCmd.AddCommand(logsCmd) + requireDocker(wpCmd, startCmd, stopCmd, restartCmd, execCmd, logsCmd) } diff --git a/cmd/site_test.go b/cmd/site_test.go new file mode 100644 index 0000000..5e39234 --- /dev/null +++ b/cmd/site_test.go @@ -0,0 +1,27 @@ +package cmd + +import ( + "slices" + "testing" +) + +func TestSplitService(t *testing.T) { + tests := []struct { + args []string + wantService string + wantCommand []string + }{ + {args: []string{"php", "ls", "-la"}, wantService: "php", wantCommand: []string{"ls", "-la"}}, + {args: []string{"nginx", "nginx", "-t"}, wantService: "nginx", wantCommand: []string{"nginx", "-t"}}, + {args: []string{"openlitespeed", "ls"}, wantService: "openlitespeed", wantCommand: []string{"ls"}}, + {args: []string{"ls", "-la"}, wantService: "", wantCommand: []string{"ls", "-la"}}, + {args: []string{"php"}, wantService: "php", wantCommand: []string{}}, + } + + for _, tt := range tests { + service, command := splitService(tt.args) + if service != tt.wantService || !slices.Equal(command, tt.wantCommand) { + t.Errorf("splitService(%q) = %q, %q, want %q, %q", tt.args, service, command, tt.wantService, tt.wantCommand) + } + } +} diff --git a/cmd/version.go b/cmd/version.go index e95e7f9..a32680a 100644 --- a/cmd/version.go +++ b/cmd/version.go @@ -1,11 +1,18 @@ package cmd import ( + "context" + "errors" "fmt" + "io" "os" + "os/signal" + "syscall" - "github.com/flywp/server-cli/internal/utils" + "github.com/flywp/server-cli/internal/release" + "github.com/flywp/server-cli/internal/service" "github.com/flywp/server-cli/internal/version" + "github.com/mattn/go-isatty" "github.com/spf13/cobra" ) @@ -24,46 +31,104 @@ var versionCmd = &cobra.Command{ var updateCmd = &cobra.Command{ Use: "update", Short: "Update fly-cli to the latest version", - Run: func(cmd *cobra.Command, args []string) { + Long: `Update fly-cli to the latest release. When fly already runs the latest +release, the command does nothing. On a server with the monitoring agent, the +command also restarts the agent, so that the agent runs the new binary.`, + RunE: func(cmd *cobra.Command, args []string) error { if os.Geteuid() != 0 { - fmt.Println("Error: The update command must be run as root.") - fmt.Println("Please run 'sudo fly update'") - return + return errors.New("the update command must be run as root, please run 'sudo fly update'") } - latestVersion, hasUpdate, err := utils.CheckForUpdates() + // Ctrl-C stops the download, and the partial download is removed. + ctx, stop := signal.NotifyContext(cmd.Context(), os.Interrupt, syscall.SIGTERM) + defer stop() + + update, err := release.CheckForUpdates(ctx) if err != nil { - fmt.Println("Error checking for updates:", err) - return + return fmt.Errorf("checking for updates: %w", err) } - if !hasUpdate { + latest := update.Release.TagName + switch { + case !update.Comparable: + fmt.Printf("This is not a release build (version %s). Latest release: %s\n", version.Version, latest) + case !update.Available: fmt.Println("You are already running the latest version.") - return + return restartStaleAgent(ctx) + default: + fmt.Printf("New version available: %s\n", latest) } - fmt.Printf("New version available: %s\n", latestVersion) - if !yesFlag { - fmt.Print("Do you want to update? (y/n): ") - var response string - fmt.Scanln(&response) - if response != "y" && response != "Y" { + ok, err := confirm(os.Stdin, os.Stdout, latest) + if err != nil { + return err + } + if !ok { fmt.Println("Update cancelled.") - return + return nil } } fmt.Println("Updating...") - if err := utils.SelfUpdate(); err != nil { - fmt.Println("Error updating:", err) - } else { - fmt.Println("Update successful. Please restart fly cli.") - os.Exit(0) + if err := release.SelfUpdate(ctx, update.Release); err != nil { + return fmt.Errorf("updating: %w", err) } + fmt.Printf("Updated to %s.\n", latest) + + return restartAgent(ctx) }, } +// confirm asks whether to install latest. Without a terminal nobody can +// answer, so it returns an error: a script must not read "cancelled" as done. +func confirm(in *os.File, out io.Writer, latest string) (bool, error) { + if !isatty.IsTerminal(in.Fd()) && !isatty.IsCygwinTerminal(in.Fd()) { + return false, errors.New("there is no terminal to confirm the update: run 'sudo fly update --yes'") + } + + _, _ = fmt.Fprintf(out, "Do you want to install %s? (y/n): ", latest) + var response string + // An empty or unreadable answer cancels the update. + _, _ = fmt.Fscanln(in, &response) + + return response == "y" || response == "Y", nil +} + +// restartAgent restarts the monitoring agent, if the server has it, so that it +// runs the new binary. +func restartAgent(ctx context.Context) error { + if !service.Installed() { + return nil + } + + fmt.Println("Restarting the monitoring agent...") + if err := service.Restart(ctx); err != nil { + return fmt.Errorf("the update is installed, but the monitoring agent did not restart: %w", err) + } + + return nil +} + +// restartStaleAgent restarts the monitoring agent when it still runs a binary +// that an earlier update replaced. +func restartStaleAgent(ctx context.Context) error { + if !service.Installed() { + return nil + } + + stale, err := service.Stale(ctx) + if err != nil { + return fmt.Errorf("checking the monitoring agent: %w", err) + } + if !stale { + return nil + } + + fmt.Println("The monitoring agent runs an older binary.") + return restartAgent(ctx) +} + func init() { updateCmd.Flags().BoolVarP(&yesFlag, "yes", "y", false, "Automatically answer yes to update confirmation") rootCmd.AddCommand(versionCmd) diff --git a/cmd/version_test.go b/cmd/version_test.go new file mode 100644 index 0000000..09b29ca --- /dev/null +++ b/cmd/version_test.go @@ -0,0 +1,117 @@ +package cmd + +import ( + "context" + "io" + "os" + "path/filepath" + "strconv" + "strings" + "testing" + + "github.com/flywp/server-cli/internal/service" +) + +func TestConfirmNeedsATerminal(t *testing.T) { + r, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + defer func() { _ = r.Close() }() + _, _ = w.WriteString("y\n") + _ = w.Close() + + // A script that pipes "y" still needs --yes: the pipe is not a terminal. + ok, err := confirm(r, io.Discard, "v0.2.0") + if ok || err == nil || !strings.Contains(err.Error(), "--yes") { + t.Errorf("confirm() = %v, %v; want an error that tells to use --yes", ok, err) + } +} + +// fakeAgentService makes the agent unit exist and puts a fake systemctl in +// PATH. It returns the log of the systemctl calls. +func fakeAgentService(t *testing.T, installed bool) string { + return fakeAgentServiceExit(t, installed, 0) +} + +// fakeAgentServiceExit is fakeAgentService with a systemctl that exits with +// exit. +func fakeAgentServiceExit(t *testing.T, installed bool, exit int) string { + t.Helper() + + unit := filepath.Join(t.TempDir(), "fly-agent.service") + if installed { + if err := os.WriteFile(unit, []byte("[Unit]\n"), 0o644); err != nil { + t.Fatal(err) + } + } + old := service.UnitPath + service.UnitPath = unit + t.Cleanup(func() { service.UnitPath = old }) + + dir := t.TempDir() + log := filepath.Join(t.TempDir(), "calls") + script := "#!/bin/sh\nprintf '%s\\n' \"$*\" >> " + log + "\n[ \"$1\" = show ] && echo 0\n[ " + strconv.Itoa(exit) + " -ne 0 ] && echo 'Failed to connect to bus' >&2\nexit " + strconv.Itoa(exit) + "\n" + if err := os.WriteFile(filepath.Join(dir, "systemctl"), []byte(script), 0o755); err != nil { + t.Fatal(err) + } + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + + return log +} + +func systemctlCalls(t *testing.T, log string) string { + t.Helper() + data, err := os.ReadFile(log) + if err != nil && !os.IsNotExist(err) { + t.Fatal(err) + } + return strings.TrimSpace(string(data)) +} + +func TestRestartAgentAfterAnUpdate(t *testing.T) { + log := fakeAgentService(t, true) + if err := restartAgent(context.Background()); err != nil { + t.Fatal(err) + } + if got := systemctlCalls(t, log); got != "try-restart fly-agent" { + t.Errorf("systemctl calls = %q, want try-restart fly-agent", got) + } +} + +func TestNoRestartWithoutTheAgent(t *testing.T) { + log := fakeAgentService(t, false) + if err := restartAgent(context.Background()); err != nil { + t.Fatal(err) + } + if err := restartStaleAgent(context.Background()); err != nil { + t.Fatal(err) + } + if got := systemctlCalls(t, log); got != "" { + t.Errorf("systemctl calls = %q, want none on a server without the agent", got) + } +} + +func TestNoRestartWhenTheAgentDoesNotRun(t *testing.T) { + // The fake systemctl shows MainPID 0: the agent does not run, so systemd + // starts the binary on the disk. + log := fakeAgentService(t, true) + if err := restartStaleAgent(context.Background()); err != nil { + t.Fatal(err) + } + if got := systemctlCalls(t, log); got != "show --property=MainPID --value fly-agent" { + t.Errorf("systemctl calls = %q, want only the check", got) + } +} + +func TestASystemctlFailureShowsItsMessage(t *testing.T) { + fakeAgentServiceExit(t, true, 1) + + for name, f := range map[string]func(context.Context) error{"restartAgent": restartAgent, "restartStaleAgent": restartStaleAgent} { + var stderr strings.Builder + code := exitCode(f(context.Background()), &stderr) + if code != 1 || !strings.Contains(stderr.String(), "Failed to connect to bus") { + t.Errorf("%s: exit %d, stderr %q; want exit 1 and the systemctl message", name, code, stderr.String()) + } + } +} diff --git a/docs/decisions.md b/docs/decisions.md new file mode 100644 index 0000000..5035c2c --- /dev/null +++ b/docs/decisions.md @@ -0,0 +1,90 @@ +# Design decisions + +Each decision of the monitoring agent, with its reason. Change one only when new evidence +contradicts the reason. + +## Shape + +- **One binary.** The agent is `fly agent run`, part of `fly`, not a separate program. One release, + one install, one update path. +- **One loop.** One goroutine wakes every 10 seconds, reads the server, and at the tick of each + minute makes the sample and, every few minutes, sends. Only the disk walk of the sites runs apart. + The work is small, and one loop is easy to reason about after a crash. +- **Plain JSON files for state**, written to a temporary file, synced and renamed. No database: the + queues are small, and a crash must never leave half a file. Event ids are ULIDs, made when the + event is queued, so a resend keeps its id and FlyWP can drop duplicates. +- **HTTPS only**, except for a loopback host, which tests and local development need. Redirects are + refused: a redirect could turn a POST into a GET, or send the token over plain http. + +## Measurements + +- **Peaks, not only averages.** A minute average hides a 10-second spike: 10 seconds at 100% CPU and + 50 seconds idle average to about 17%. So the agent reads every 10 seconds and sends the peak + window of the minute next to the minute value. The minute values stay as they were, because alerts + read them: a sustained value is an incident, a 10-second peak is not. +- **No peak of load or disk space.** The kernel already smooths the load over a minute, and disk + space changes over minutes, not seconds. +- **Pressure (PSI)**, from the `total=` counters, not the kernel's `avg10`/`avg60` (moving averages + that do not match the windows of a sample). "How full" is not "under pressure": memory that makes + tasks wait is a problem at any percent. +- **Disk activity, not a busy percent.** An SSD or a cloud disk serves many requests at once, so + "100% busy" does not mean full. The agent sends bytes and operations, and I/O pressure shows the + waiting. Only whole hardware disks count: partitions, loop, device-mapper and RAID devices would + count the same I/O twice. +- **Network:** only interfaces with a hardware device (and no `master`), else the default-route + interface. Virtual interfaces would count the same traffic twice. +- **Short windows are dropped.** A reading less than 5 seconds from its neighbour is left out, and a + first minute shorter than 10 seconds after a fresh start sends no sample. A window of a few milliseconds, or the load of + the start itself, would show as a false spike. +- **`null`, not 0,** when a value is unknown: after a reboot, a counter that went back, or a missing + kernel feature. A 0 would draw a false dip. +- **Waiting updates** come from `apt-check`, once an hour and not in the send budget. A failed count + keeps the last one. + +## Sites + +- **A site is a Docker Compose project** whose folder is directly in the home folder of the server + user. That is where FlyWP puts sites, and the Compose label gives the folder without guessing. +- **CPU and memory from cgroup v2**, not `docker stats`: the agent reads files and makes no extra + Docker requests. A container counts for CPU only when both ticks saw the same container **and the + same cgroup**: `docker restart` keeps the container id but starts a new cgroup whose CPU time + begins at 0. +- **Only two Docker requests** (`GET /version`, `GET /containers/json`), with a 5-second timeout. + The socket is powerful; the agent uses the least of it. +- **Disk use once an hour, in the background, at idle I/O priority**, with a 5-minute limit for each + folder. A walk of a large site must never delay a sample or slow the sites. + +## Sending + +- **Keep data until FlyWP accepts it**, up to 24 hours of samples. A 400 drops the batch (it will + never be accepted); a 401, 429, 5xx or network error keeps it. +- **Values are checked before they are queued.** One bad value makes FlyWP refuse the whole batch, + so the agent clamps each value to the range of the contract first. +- **Samples and events have separate backoffs.** A failing event must not stop the metrics. + +## Updates + +- **The agent never downgrades.** An update to an older release fails with the reason. A bad + release is fixed with a newer one. This keeps agents that updated themselves from being moved back. +- **Two update paths, and both stay:** + - FlyWP's `agent.update` with a pinned sha256. It is fast, and it is the recovery path: it does + not read the signature. + - The daily automatic update, which needs a signature. Without it, a person who controls the GitHub + repository alone could put code on every server. +- **Offline ed25519 signatures.** The maintainer signs `checksums.txt` on their own computer after + rebuilding the release from their local tag. The key is never on GitHub or in CI. The public keys + are compiled in. +- **A 24-hour wait after the signature**, then each server at its own time of day. A bad release can + be stopped in that time by marking it as a pre-release. +- **A signature does not expire.** +- **Why the command path must stay:** if the key leaks, the release that removes it would itself + need a signature, and only the leaked key could make one. The pinned-sha256 path is the way out. +- **`fly update` stays.** It restarts the agent, and it keeps the owner of the binary, so the agent + can still replace it. + +## Security + +- **The agent never runs as root.** It runs as the server user under systemd. +- **Root never writes into the server user's folder.** When `install.sh` or `fly update` replaces + the agent's binary, the file is written by the server user or changed through the open file, never + by a path that the server user could swap for a link. diff --git a/docs/development.md b/docs/development.md new file mode 100644 index 0000000..42e5318 --- /dev/null +++ b/docs/development.md @@ -0,0 +1,60 @@ +# Development + +Go 1.27 or later (`go.mod` selects the toolchain). `make help` lists every target. + +```bash +make build # bin/fly, with the version from git +make test # go test ./... -race +make check # fmt-check, vet, lint, test and govulncheck: the gate before each merge +make release # static linux/amd64 and linux/arm64 archives and checksums.txt in build/ +``` + +CI runs `make check` and `make release` on each pull request and on each push to `develop` and +`main`. + +## Tests on Linux + +The agent reads `/proc`, `/sys` and cgroups, so some tests build only on Linux +(`*_linux_test.go`, for example the test against the real `/proc`). On macOS, `make check` skips +them. Before you push a change to `internal/metrics` or `internal/agent`, run the tests on Linux as +a non-root user: + +```bash +docker run --rm -u "$(id -u):$(id -g)" -e HOME=/tmp -v "$PWD":/src -w /src golang:1.27 go test ./... +``` + +A few tests need root (the owner of the binary after `sudo fly update`). CI runs them with `sudo`. + +## Run the agent locally + +The agent accepts plain `http://` for a loopback host, so you can point it at a local control +plane: + +```bash +make build +STATE_DIRECTORY=$(mktemp -d) FLY_AGENT_URL=http://127.0.0.1:8080 \ +FLY_AGENT_TOKEN=test FLY_AGENT_SERVER_ID=1 FLY_AGENT_AUTO_UPDATE=off bin/fly agent run +``` + +## Reproducible builds + +`make release` gives the same bytes for the same commit and Go version, on any computer: the build +date is the commit date, paths are trimmed, and the archives are written with owner 0:0. +`make sign-release` depends on this. See [releasing.md](releasing.md). + +## Code layout + +| Path | | +|---|---| +| `cmd/` | The commands (cobra). `cmd/agent.go` is `fly agent run`. | +| `internal/agent` | The agent: settings, lock, schedule, queues, sending, commands, auto-update. | +| `internal/agent/wire` | The JSON types of the contract. | +| `internal/metrics` | The Linux measurements: `/proc`, PSI, disks, Docker, sites. | +| `internal/dockerapi` | A minimal Docker client over the unix socket (two GET requests). | +| `internal/release` | Release download, sha256 check, binary replacement, signatures, trusted keys. | +| `internal/service` | The `fly-agent` systemd unit: restart after an update. | +| `internal/statefile` | JSON files written atomically. | +| `internal/docker` | Docker Compose calls and the Docker checks of the site commands. | +| `internal/utils` | Finds the site folder from the current folder or `--domain`. | +| `tools/releasesign`, `tools/sign-release.sh` | Key generation and release signing. | +| `install.sh` | The installer. | diff --git a/docs/monitoring-agent.md b/docs/monitoring-agent.md new file mode 100644 index 0000000..32a0951 --- /dev/null +++ b/docs/monitoring-agent.md @@ -0,0 +1,165 @@ +# The monitoring agent + +`fly agent run` is the FlyWP monitoring agent. It measures the server each minute and sends the +values to FlyWP. FlyWP shows them as charts and alerts, and can restart or update the agent without +SSH. + +The agent follows the FlyWP monitoring agent contract **v0.5.0**. The contract defines the requests, +the fields and the rules on both sides. + +## How it runs + +| | | +|---|---| +| Service | `/etc/systemd/system/fly-agent.service`, `Restart=always`. FlyWP installs it. | +| User | The server user (for example `fly`), never root. | +| Binary | `~/.fly/bin/fly` of the server user. `/usr/local/bin/fly` is a link to it. | +| Settings | `/etc/fly/agent.env` (root, mode 0600) | +| State | `/var/lib/fly-agent` (`STATE_DIRECTORY` of systemd) | +| Log | `journalctl -u fly-agent` | + +Only one agent can run with one state folder. The agent does not need Docker. + +### Settings + +| Variable | | +|---|---| +| `FLY_AGENT_URL` | The FlyWP control plane. It must be `https://`. Plain `http://` is allowed only for a loopback host (tests and local development). | +| `FLY_AGENT_TOKEN` | The token of this server. The agent never logs it. | +| `FLY_AGENT_SERVER_ID` | The id of the server in FlyWP. It also sets the second of the minute at which the agent works, so that servers do not all send at the same time. | +| `FLY_AGENT_AUTO_UPDATE` | `off` stops the automatic updates on this server. See [Updates](#updates). | + +After a change, run `sudo systemctl restart fly-agent`. + +### State files + +| File | | +|---|---| +| `samples.json` | Samples not yet accepted by FlyWP. At most 24 hours (1440 samples); then the oldest go. | +| `events.json` | Events not yet accepted (at most 1000). | +| `ran.json` | The commands that ran, so that no command runs two times. | +| `counters.json` | The last counters, so that a restart does not lose a minute of traffic. | +| `state.json` | The report interval, and the time of the last release check. | +| `lock` | Stops a second agent. | + +Each file is written to a temporary file first, then renamed, so a crash never leaves half a file. + +## What it sends + +The agent reads the server every 10 seconds. Each minute, it sends one sample with the value of the +minute and, where it makes sense, the **peak** of the minute. A peak shows a short spike that the +average of the minute hides. + +| Group | Values | +|---|---| +| CPU | use (%) and its peak, load (1 min) | +| Memory | used and total, swap used and total, with the peaks of the used values | +| Disk space | used and total of `/` (the same as `df`) | +| Network | bytes in and out of the physical interfaces, and the peak bytes each second | +| Pressure (PSI) | the share of time that tasks waited for CPU, memory or disk I/O, and its peak. Not sent when the kernel has no PSI. | +| Disk activity | bytes and operations read and written on the physical disks, and the peaks each second | +| Sites | the CPU, memory and disk use of each site (see below) | + +With the samples, the agent sends the **status** of the server: restart needed, waiting updates +(and security updates), OS, kernel, uptime, architecture, CPU count, Docker state (`running`, +`not_running`, `not_installed`) and Docker version. + +### Sites + +A site is a Docker Compose project whose folder is directly in the home folder of the server user. +For each site, the agent sends: + +- **CPU:** the CPU time of its containers in the minute, as a share of all CPUs. +- **Memory:** the memory of its containers, without the page cache that the kernel can free (the + same as `docker stats`). +- **Disk:** the size of its folder. The agent measures it at most once an hour, in the background, at + the lowest I/O priority. A folder that the server user cannot list is skipped, so the value can be + lower than the real use (for example, the database files in `~/.fly` belong to the container + user). + +The agent reads Docker through its socket with two requests only: `GET /version` and +`GET /containers/json`. It reads CPU and memory from cgroup v2. Without Docker or cgroup v2, it +sends no site values. + +### Gaps and null values + +- A value that the agent cannot measure is `null`, not 0. +- The first sample after a fresh start has no traffic (`net_counters_reset` is true) and no + pressure or disk activity values: they need two readings. A restart within 90 seconds keeps + them. +- When the first minute after a fresh start is shorter than 10 seconds, the agent sends no sample + for it. That short minute is mostly the load of the start itself. +- A reading that fails skips that minute. The agent never stops for a measurement error. + +## Sending + +- The agent sends every 1 to 10 minutes. FlyWP sets the interval in each reply. +- When FlyWP cannot be reached, the agent keeps the data on disk and tries again: after 1, 2, 4, 8, + then every 10 minutes. It keeps measuring meanwhile. After a 401 it tries each 5 minutes; after a + 429 it waits as FlyWP asks. +- When FlyWP accepts the data again, the agent logs one line: "the control plane accepts the + requests again". +- Samples are sent even while events fail, and the other way around. + +## Commands from FlyWP + +| Command | What the agent does | +|---|---| +| `agent.restart` | Exits; systemd starts it again. | +| `agent.update` | Downloads the version that FlyWP names, checks its sha256 against the value that FlyWP sends, replaces the binary and exits. | + +Each command runs at most one time, also after a crash. The result goes back as an event +(`command.completed` or `command.failed`). The agent never installs an older version: an update to +an older release fails with the reason. + +## Updates + +There are three ways to a new version. All keep the binary owned by the server user and restart the +agent. + +1. **FlyWP sends `agent.update`** with the version and its sha256. This path does not need a + signature, so it also works when the release key is lost. +2. **The agent updates itself.** Once a day, at a time set by the server id, it checks the latest + release on GitHub. It installs it only when: + - the release is newer than the running version; + - `checksums.txt.sig` has a valid signature from a FlyWP release key that this build trusts; + - the signature is more than 24 hours old. + + Each server installs at its own time of day. A dev pre-release updates itself too. A build + without a release version (`make build` between tags, `go build`) never does. + `FLY_AGENT_AUTO_UPDATE=off` stops this path on one server. +3. **`sudo fly update`** or `install.sh` install the latest release after checking its sha256 in + `checksums.txt`. + +## Security + +- The agent runs as the server user. It never runs as root, and it refuses to start as root. +- It sends data only to `FLY_AGENT_URL`, over HTTPS, and it does not follow redirects. Downloads + are HTTPS too. +- The token is never logged, and an error never shows it. +- Every download is checked against a sha256 before it replaces the binary. +- The release signing key is kept offline, never on GitHub. A person who controls the GitHub + repository alone cannot make every agent install their code. See [releasing.md](releasing.md). + +## Troubleshooting + +| Log line | Meaning | +|---|---| +| "must be an https URL" | `FLY_AGENT_URL` is plain http. | +| "does not accept the token" | The token in `/etc/fly/agent.env` does not match FlyWP. The agent keeps its data and tries each 5 minutes. | +| "asks the agent to wait" | FlyWP rate-limits the agent (429). | +| "a different agent is running" | A second `fly agent run` uses the same state folder. | +| "no sample for this minute" | The first minute after a start was too short. Normal. | +| "some files of a site cannot be read" | The disk walk skipped folders; the site value is lower than the real use. Logged once per folder after each start. | +| "a new release installs after its wait" | Normal: a signed release waits 24 hours. | +| "a new release waits for its signature" | The latest release is not signed yet. | +| "auto-update is off: this build has no release version" | A local build. Install a release or a dev pre-release. | +| "auto-update is off: the agent cannot write the folder of its binary" | The binary is not in a folder of the server user. Install again with `install.sh`. | + +To remove the agent from a server: + +```bash +sudo systemctl disable --now fly-agent +sudo rm -rf /etc/systemd/system/fly-agent.service /etc/fly /var/lib/fly-agent +sudo systemctl daemon-reload +``` diff --git a/docs/releasing.md b/docs/releasing.md new file mode 100644 index 0000000..b912dc7 --- /dev/null +++ b/docs/releasing.md @@ -0,0 +1,107 @@ +# Releasing + +`develop` is the development branch. `main` is the release branch: a release tag must be on `main`. + +## A release + +1. **Merge `develop` into `main`** with a pull request and a **merge commit** (not a squash, so the + two branches keep a common history). +2. **Tag and push at once:** + + ```bash + git checkout main && git pull --ff-only + git tag -a v0.2.0 -m "v0.2.0" + git push origin v0.2.0 + ``` + + Push the tag right after the merge: `install.sh` from `main` installs only a release with + `checksums.txt`, so until the release exists it stops. + + The Release workflow checks that the tag is on `main`, runs `make check`, builds the archives + with `make release`, and publishes the release with `fly-linux-amd64.tar.gz`, + `fly-linux-arm64.tar.gz` and `checksums.txt`. Do not rename these files: installed CLIs look for + them. + +3. **Sign it**, on your own computer (see [Signing](#signing)): + + ```bash + make sign-release VERSION=v0.2.0 KEY= # KEY=- reads the key from stdin + ``` + +4. **Watch it for 24 hours.** Agents install a signed release by themselves only 24 hours after the + signature, each at its own time of day. First update one test server at once with + `sudo fly update --yes`, and watch its log. +5. **To stop a bad release** within those 24 hours, mark it as a pre-release on GitHub. Agents no + longer see it. After that, fix forward: the agent never downgrades. + +A tag with a suffix, such as `v0.3.0-rc.1`, becomes a pre-release. `fly update`, `install.sh` and +the agents ignore pre-releases. + +## Signing + +The signature says: "this release is the code of my tag". Agents check it before they update +themselves. The key never goes to GitHub, so control of the GitHub repository alone is not enough to +reach every server. + +`make sign-release` signs only what it can build again: + +1. It shows the commit of your **local** tag and asks you to type the tag. Review that commit first: + the signature is your approval. +2. It downloads the archives and `checksums.txt`, and checks the archives against the sums. +3. It builds the release again from your local tag, with the Go version of the CI build, and + compares the binaries byte for byte. A binary holds its commit, so this also proves that CI built + your tag. +4. It signs `checksums.txt`, checks the signature with the keys of the tag, and uploads + `checksums.txt.sig`. + +If an archive on GitHub was swapped, or the tag moved, step 2 or 3 stops before anything is signed. +`UPLOAD=0` does every check and writes the signature without uploading it, which is useful for a +rehearsal on a dev pre-release. + +### Keys + +- `make release-key KEY=` makes a key pair. It prints the public key line + for `internal/release/keys.go`. `COMMENT="..."` names the key. +- Keep the private key in a password manager or vault, with a backup. Never commit it, and never put + it in a GitHub secret or CI. +- `keys.go` can list more than one key. One valid signature from any listed key is enough. To + replace a key, list both for one release, then remove the old one. +- A change to `keys.go` takes effect only in the binaries of the next release. + +**If the key is lost or leaked:** make a new key, put its line in `keys.go` (remove the leaked +line), and release. Ship that release through a FlyWP `agent.update`: that path checks the sha256 +that FlyWP sends and does not read the signature. This is why the command path must stay. + +## Dev pre-releases + +To test a branch on a real server before it merges: + +```bash +make dev-version # prints the tag, for example v0.2.0-dev.1a2b3c4 +make dev-release # tags the pushed commit; CI publishes a pre-release +``` + +The version is the next minor version after the latest release, plus the short commit hash +(`DEV_BASE=v0.2.1` overrides the first part). `make dev-release` refuses uncommitted changes and +commits that are not pushed. Pre-release tags can come from any branch. + +Install one on a test server by hand: + +```bash +tag=v0.2.0-dev.1a2b3c4 arch=amd64 +base=https://github.com/flywp/server-cli/releases/download/$tag +curl -fsSLO "$base/fly-linux-$arch.tar.gz" && curl -fsSLO "$base/checksums.txt" +sha256sum -c --ignore-missing checksums.txt +tar -xzf "fly-linux-$arch.tar.gz" && sudo install -m 0755 "fly-linux-$arch" /usr/local/bin/fly +``` + +On a server that runs the agent, FlyWP can install a dev pre-release with `agent.update` instead. + +## Roll back and stop + +| Need | Do | +|---|---| +| Stop a signed release before agents take it | Within 24 hours of the signature, mark it as a pre-release on GitHub. | +| Fix a bad release | Release a newer, fixed version. The agent never downgrades. | +| Stop automatic updates on one server | Add `FLY_AGENT_AUTO_UPDATE=off` to `/etc/fly/agent.env` and restart `fly-agent`. FlyWP updates still work. | +| Stop the agent on one server | `sudo systemctl disable --now fly-agent` | diff --git a/go.mod b/go.mod index f9fb119..b8b9ef0 100644 --- a/go.mod +++ b/go.mod @@ -1,18 +1,21 @@ module github.com/flywp/server-cli -go 1.22.0 +go 1.27.0 -require github.com/spf13/cobra v1.8.1 +toolchain go1.27.1 + +require ( + github.com/fatih/color v1.19.0 + github.com/mattn/go-isatty v0.0.24 + github.com/oklog/ulid/v2 v2.1.2 + github.com/spf13/cobra v1.10.2 + golang.org/x/mod v0.41.0 + golang.org/x/sys v0.48.0 + gopkg.in/yaml.v2 v2.4.0 +) require ( - github.com/cpuguy83/go-md2man/v2 v2.0.4 // indirect - github.com/fatih/color v1.17.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect - github.com/mattn/go-colorable v0.1.13 // indirect - github.com/mattn/go-isatty v0.0.20 // indirect - github.com/russross/blackfriday/v2 v2.1.0 // indirect - github.com/spf13/pflag v1.0.5 // indirect - golang.org/x/sys v0.18.0 // indirect - gopkg.in/yaml.v2 v2.4.0 // indirect - gopkg.in/yaml.v3 v3.0.1 // indirect + github.com/mattn/go-colorable v0.1.15 // indirect + github.com/spf13/pflag v1.0.10 // indirect ) diff --git a/go.sum b/go.sum index 6fd6aa7..567b99e 100644 --- a/go.sum +++ b/go.sum @@ -1,26 +1,27 @@ -github.com/cpuguy83/go-md2man/v2 v2.0.4 h1:wfIWP927BUkWJb2NmU/kNDYIBTh/ziUX91+lVfRxZq4= -github.com/cpuguy83/go-md2man/v2 v2.0.4/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o= -github.com/fatih/color v1.17.0 h1:GlRw1BRJxkpqUCBKzKOw098ed57fEsKeNjpTe3cSjK4= -github.com/fatih/color v1.17.0/go.mod h1:YZ7TlrGPkiz6ku9fK3TLD/pl3CpsiFyu8N92HLgmosI= +github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= +github.com/fatih/color v1.19.0 h1:Zp3PiM21/9Ld6FzSKyL5c/BULoe/ONr9KlbYVOfG8+w= +github.com/fatih/color v1.19.0/go.mod h1:zNk67I0ZUT1bEGsSGyCZYZNrHuTkJJB+r6Q9VuMi0LE= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= -github.com/mattn/go-colorable v0.1.13 h1:fFA4WZxdEF4tXPZVKMLwD8oUnCTTo08duU7wxecdEvA= -github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg= -github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= -github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= -github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= -github.com/russross/blackfriday/v2 v2.1.0 h1:JIOH55/0cWyOuilr9/qlrm0BSXldqnqwMsf35Ld67mk= +github.com/mattn/go-colorable v0.1.15 h1:+u9SLTRGnXv73cEsnsmoZBom+dMU88B2M0aDcWy0/jY= +github.com/mattn/go-colorable v0.1.15/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8= +github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI= +github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= +github.com/oklog/ulid/v2 v2.1.2 h1:IEclFb9JNvzYA6MW2SCxbLzcHTVsfqm3PrqGQJH5zec= +github.com/oklog/ulid/v2 v2.1.2/go.mod h1:rcEKHmBBKfef9DhnvX7y1HZBYxjXb0cP5ExxNsTT1QQ= +github.com/pborman/getopt v0.0.0-20170112200414-7148bc3a4c30/go.mod h1:85jBQOZwpVEaDAr341tbn15RS4fCAsIst0qp7i8ex1o= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= -github.com/spf13/cobra v1.8.1 h1:e5/vxKd/rZsfSJMUX1agtjeTDf+qv1/JdBF8gg5k9ZM= -github.com/spf13/cobra v1.8.1/go.mod h1:wHxEcudfqmLYa8iTfL+OuZPbBZkmvliBWKIezN3kD9Y= -github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA= -github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= -golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.18.0 h1:DBdB3niSjOA/O0blCZBqDefyWNYveAYMNF1Wum0DYQ4= -golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= +github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU= +github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4= +github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= +github.com/spf13/pflag v1.0.10 h1:4EBh2KAYBwaONj6b2Ye1GiHfwjqyROoF4RwYO+vPwFk= +github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= +go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= +golang.org/x/mod v0.41.0 h1:qJmnOUb4YB+FsEuM3HcWucdZASCPGhsX6uljO6pog0c= +golang.org/x/mod v0.41.0/go.mod h1:Ek9pY8RKWXwsWvd3rQiHYtMqkjSUV+s1Rj7j4H5Ur6o= +golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo= +golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY= gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ= -gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= -gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/install.sh b/install.sh index a587a84..8b02e1b 100644 --- a/install.sh +++ b/install.sh @@ -1,4 +1,14 @@ #!/bin/bash +# +# Install the latest release of fly. Run it as root: +# curl -fsSL https://raw.githubusercontent.com/flywp/server-cli/main/install.sh | sudo bash + +set -euo pipefail + +REPO="flywp/server-cli" +TARGET="/usr/local/bin/fly" +AGENT_UNIT="/etc/systemd/system/fly-agent.service" +TEMP_DIR="" error_exit() { echo -e "\033[31mERROR: $1\033[0m" >&2 @@ -13,212 +23,183 @@ info_msg() { echo -e "\033[34m$1\033[0m" } -warning_msg() { - echo -e "\033[33m$1\033[0m" +cleanup() { + if [ -n "$TEMP_DIR" ]; then + rm -rf "$TEMP_DIR" + fi } +trap cleanup EXIT -# Check if script is running with sudo privileges check_sudo() { if [ "$(id -u)" -ne 0 ]; then error_exit "This script requires sudo privileges. Please run with sudo." fi } -# Check for required commands and install if missing -check_dependencies() { - - # Check for Perl regex support - if ! echo "test" | grep -P "test" &> /dev/null; then - warning_msg "Perl regex support not detected. Installing..." - - # Install perl-compatible grep for Ubuntu - apt-get install -y -qq grep - - # Check again after installation - if ! echo "test" | grep -P "test" &> /dev/null; then - warning_msg "Perl regex support still not available. Using alternative parsing method." - USE_PERL_REGEX=false - else - USE_PERL_REGEX=true - fi - else - USE_PERL_REGEX=true - fi -} - # Determine OS and architecture determine_platform() { info_msg "Detecting system platform..." - - # Check if running on Ubuntu - if [ -f /etc/os-release ]; then - . /etc/os-release - if [[ "$ID" != "ubuntu" ]]; then - error_exit "This script only supports Ubuntu. Detected OS: $ID" - fi - info_msg "Detected Ubuntu version: $VERSION_ID" - else + + if [ ! -f /etc/os-release ]; then error_exit "Cannot detect OS. This script only supports Ubuntu." fi - - OS=$(uname -s | tr '[:upper:]' '[:lower:]') - ARCH=$(uname -m) - - info_msg "Detected architecture: $ARCH" - - if [ "$ARCH" == "x86_64" ]; then - ARCH="amd64" - elif [[ "$ARCH" == "aarch64" || "$ARCH" == "arm64" ]]; then - ARCH="arm64" - else - error_exit "Unsupported architecture: $ARCH. Only amd64 and arm64 are supported." + # shellcheck disable=SC1091 + . /etc/os-release + if [ "${ID:-}" != "ubuntu" ]; then + error_exit "This script only supports Ubuntu. Detected OS: ${ID:-unknown}" fi - + info_msg "Detected Ubuntu version: ${VERSION_ID:-unknown}" + + OS=$(uname -s | tr '[:upper:]' '[:lower:]') + case "$(uname -m)" in + x86_64) ARCH="amd64" ;; + aarch64 | arm64) ARCH="arm64" ;; + *) error_exit "Unsupported architecture: $(uname -m). Only amd64 and arm64 are supported." ;; + esac + + NAME="fly-${OS}-${ARCH}" info_msg "Using OS: $OS, Architecture: $ARCH" } -# Get latest release from GitHub API +# Get the tag of the latest release from the GitHub API get_release_info() { info_msg "Fetching latest release information from GitHub..." - - # Use a temporary file for the API response - GITHUB_API_RESPONSE=$(mktemp) - - # Add a user-agent to avoid rate limiting - if ! curl -s -L -H "User-Agent: FlyWP-Installer" \ - https://api.github.com/repos/flywp/server-cli/releases/latest \ - -o "$GITHUB_API_RESPONSE"; then - error_exit "Failed to access GitHub API. Please check your internet connection." - fi - - # Check for rate limiting - if grep -q "API rate limit exceeded" "$GITHUB_API_RESPONSE"; then - error_exit "GitHub API rate limit exceeded. Please try again later or use a GitHub token." - fi - - # Extract tag name with more robust methods - if [ "$USE_PERL_REGEX" = true ]; then - TAG_NAME=$(grep -oP '"tag_name":\s*"\K[^"]+' "$GITHUB_API_RESPONSE") - else - TAG_NAME=$(grep '"tag_name"' "$GITHUB_API_RESPONSE" | sed -E 's/.*"tag_name":\s*"([^"]+)".*/\1/') - fi - - if [ -z "$TAG_NAME" ]; then - # Fallback to a more basic approach - TAG_NAME=$(grep "tag_name" "$GITHUB_API_RESPONSE" | cut -d'"' -f4) - fi - + + local response="$TEMP_DIR/release.json" + local status + status=$(curl -sSL --connect-timeout 10 --max-time 60 --retry 3 -H "User-Agent: FlyWP-Installer" -o "$response" -w '%{http_code}' \ + "https://api.github.com/repos/${REPO}/releases/latest") || + error_exit "Failed to access the GitHub API. Please check your internet connection." + + case "$status" in + 200) ;; + 403 | 429) error_exit "GitHub API rate limit exceeded. Please try again later." ;; + *) error_exit "The GitHub API answered with HTTP $status." ;; + esac + + # The API answers with JSON on one line. Take the first tag_name. + TAG_NAME=$(grep -o '"tag_name"[[:space:]]*:[[:space:]]*"[^"]*"' "$response" | head -n 1 | sed 's/.*"\([^"]*\)"$/\1/') || true if [ -z "$TAG_NAME" ]; then error_exit "Failed to determine the latest release version." fi - + + DOWNLOAD_BASE="https://github.com/${REPO}/releases/download/${TAG_NAME}" info_msg "Latest release version: $TAG_NAME" - - # Since we know the exact format of the release assets, construct the URL directly - DOWNLOAD_URL="https://github.com/flywp/server-cli/releases/download/${TAG_NAME}/fly-linux-${ARCH}.tar.gz" - - # Clean up - rm -f "$GITHUB_API_RESPONSE" - - info_msg "Download URL: $DOWNLOAD_URL" } -# Download and verify the release download_release() { - info_msg "Creating temporary directory..." - TEMP_DIR=$(mktemp -d) - if [ ! -d "$TEMP_DIR" ]; then - error_exit "Failed to create temporary directory." - fi - - DOWNLOAD_FILE="$TEMP_DIR/fly-$OS-$ARCH.tar.gz" - - info_msg "Downloading latest release..." - if ! curl -s -L -o "$DOWNLOAD_FILE" "$DOWNLOAD_URL"; then - rm -rf "$TEMP_DIR" - error_exit "Failed to download the release file." - fi - - # Verify the downloaded file - if [ ! -s "$DOWNLOAD_FILE" ]; then - rm -rf "$TEMP_DIR" - error_exit "Downloaded file is empty or corrupted." - fi - + info_msg "Downloading ${NAME}.tar.gz..." + curl -fsSL --connect-timeout 10 --max-time 300 --retry 3 -o "$TEMP_DIR/${NAME}.tar.gz" "${DOWNLOAD_BASE}/${NAME}.tar.gz" || + error_exit "Failed to download ${DOWNLOAD_BASE}/${NAME}.tar.gz." info_msg "Download completed successfully." } -# Extract and install -install_binary() { - info_msg "Extracting $DOWNLOAD_FILE..." - if ! tar -xzf "$DOWNLOAD_FILE" -C "$TEMP_DIR"; then - rm -rf "$TEMP_DIR" - error_exit "Failed to extract the archive." +# Examine the download with the checksum file of the release. The file has +# one line for each archive: " ". +verify_download() { + info_msg "Verifying the download with checksums.txt..." + + curl -fsSL --connect-timeout 10 --max-time 60 --retry 3 -o "$TEMP_DIR/checksums.txt" "${DOWNLOAD_BASE}/checksums.txt" || + error_exit "Failed to download checksums.txt of ${TAG_NAME}. The download cannot be checked, so it is not installed." + + # The same rules as fly update: an optional "*" (binary mode) before the + # name, CRLF line ends. Two different sums for the file fail the check. + if ! (cd "$TEMP_DIR" && awk -v f="${NAME}.tar.gz" '{ sub(/\r$/, "") } $2 == f || $2 == "*" f { print $1 " " f }' checksums.txt | sha256sum -c --status -); then + error_exit "The checksum of ${NAME}.tar.gz does not agree with checksums.txt. The download is not installed." fi - - # Look for the binary file - BINARY_FILE=$(find "$TEMP_DIR" -type f -executable | head -n 1) - - if [ -z "$BINARY_FILE" ]; then - # Fallback to expected name pattern - BINARY_FILE="$TEMP_DIR/fly-$OS-$ARCH" - - if [ ! -f "$BINARY_FILE" ]; then - # Try finding any file that might be the binary - BINARY_FILE=$(find "$TEMP_DIR" -type f -name "fly*" | head -n 1) - fi - - if [ -z "$BINARY_FILE" ]; then - rm -rf "$TEMP_DIR" - error_exit "Could not find the executable in the extracted archive." - fi + + info_msg "Checksum verified." +} + +install_binary() { + # Extract only the binary, and do not keep the owner from the archive: + # a binary that root runs must not belong to a different user. + tar --no-same-owner -xzf "$TEMP_DIR/${NAME}.tar.gz" -C "$TEMP_DIR" "$NAME" || + error_exit "The archive does not contain ${NAME}." + if [ ! -f "$TEMP_DIR/$NAME" ] || [ -L "$TEMP_DIR/$NAME" ]; then + error_exit "${NAME} in the archive is not a regular file." fi - - info_msg "Installing to /usr/local/bin/fly..." - - if ! mv "$BINARY_FILE" /usr/local/bin/fly; then - rm -rf "$TEMP_DIR" - error_exit "Failed to move the binary to /usr/local/bin/fly. Check your permissions." + + if [ -f "$AGENT_UNIT" ]; then + install_for_agent + else + install_for_root fi - - if ! chmod +x /usr/local/bin/fly; then - rm -rf "$TEMP_DIR" - error_exit "Failed to make the binary executable." + + command -v fly >/dev/null || error_exit "Installation failed: 'fly' command not found in PATH." + + success_msg "Installation of fly ${TAG_NAME} completed successfully!" + info_msg "Verify with 'fly version'" +} + +# Without the monitoring agent, fly is a root file in /usr/local/bin. A link +# or a file that is there is replaced; a link is never followed. +install_for_root() { + info_msg "Installing to ${TARGET}..." + + local tmp + tmp=$(mktemp "$(dirname "$TARGET")/.fly-update-XXXXXX") || + error_exit "Failed to create a temporary file next to ${TARGET}. Check your permissions." + if ! { cat "$TEMP_DIR/$NAME" >"$tmp" && chmod 0755 "$tmp" && chown 0:0 "$tmp"; }; then + rm -f "$tmp" + error_exit "Failed to write the binary next to ${TARGET}." + fi + # A rename is atomic: a running fly never sees a partial binary. -T + # replaces a link itself, not the file it points to. + mv -Tf "$tmp" "$TARGET" || { rm -f "$tmp"; error_exit "Failed to install the binary to ${TARGET}."; } +} + +# With the monitoring agent, the binary is ~/.fly/bin/fly, the command +# of fly-agent.service, and /usr/local/bin/fly is a link to it. The folder +# belongs to the agent user, so root does not write in it: the user could put +# a link to a root file there. The agent user writes the new binary, and root +# only gives it the file on stdin. +install_for_agent() { + local user home bin + user=$(sed -n 's/^User=//p' "$AGENT_UNIT" | tail -n 1) + bin=$(sed -n 's/^ExecStart=\([^ ]*\).*/\1/p' "$AGENT_UNIT" | tail -n 1) + + if [ -z "$user" ] || [ "$user" = "root" ]; then + error_exit "${AGENT_UNIT} has no User= for the monitoring agent." fi - - # Clean up - rm -rf "$TEMP_DIR" - - # Verify installation - if ! command -v fly &> /dev/null; then - error_exit "Installation failed: 'fly' command not found in PATH." + home=$(getent passwd "$user" | cut -d: -f6) || true + if [ -z "$home" ] || [ "$bin" != "$home/.fly/bin/fly" ]; then + error_exit "${AGENT_UNIT} does not run ${home:-~$user}/.fly/bin/fly. The layout is not known, so fly is not installed." fi - - success_msg "Installation completed successfully!" - info_msg "Verify with 'fly version'" + + info_msg "Installing to ${bin} as ${user}..." + # shellcheck disable=SC2016 # $1 is for the inner shell + runuser -u "$user" -- sh -c 'set -e + tmp=$(mktemp "$(dirname "$1")/.fly-update-XXXXXX") + trap '"'"'rm -f "$tmp"'"'"' EXIT + cat >"$tmp" + chmod 0755 "$tmp" + mv -f "$tmp" "$1" + trap - EXIT' sh "$bin" <"$TEMP_DIR/$NAME" || + error_exit "Failed to install the binary to ${bin} as ${user}." + + # The CLI and the agent use the same binary. The old installer put a + # separate file here; replace it with the link. + ln -sfn "$bin" "${TARGET}.link-$$" && mv -Tf "${TARGET}.link-$$" "$TARGET" || + error_exit "fly is installed in ${bin}, but ${TARGET} could not be linked to it." + + # try-restart: an agent that an administrator stopped stays stopped. + info_msg "Restarting the monitoring agent..." + systemctl try-restart fly-agent || + error_exit "fly is installed, but the monitoring agent did not restart." } main() { echo "===== FlyWP Server CLI Installer =====" - - # Check for sudo access + check_sudo - - # Determine OS and architecture (and check for Ubuntu) determine_platform - - # Check for Perl regex support (only essential dependency check) - check_dependencies - - # Get release information + + TEMP_DIR=$(mktemp -d) get_release_info - - # Download the release download_release - - # Install the binary + verify_download install_binary } -# Run the main function -main \ No newline at end of file +main diff --git a/internal/agent/agent.go b/internal/agent/agent.go new file mode 100644 index 0000000..ffce100 --- /dev/null +++ b/internal/agent/agent.go @@ -0,0 +1,349 @@ +package agent + +import ( + "context" + "crypto/ed25519" + "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" +) + +// The report interval is the number of samples in one report. The control +// plane sets it in each reply. +const ( + minReportInterval = 1 + maxReportInterval = 10 +) + +// ControlPlane is the part of the control plane that the agent uses. +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) +} + +// ErrNoSample is the error of a minute that has no sample, for example a +// first minute that is too short to measure. It is not a problem. +var ErrNoSample = errors.New("no sample for this minute") + +// Collector measures the server. +type Collector interface { + // Read takes a reading between two ticks, for the peaks of the minute. + Read(now time.Time) + // Sample measures the minute that ends at now. It returns an error that + // wraps ErrNoSample when the minute has no sample. + Sample(now time.Time) (wire.Sample, error) + // Status describes the server now. + Status(ctx context.Context) wire.Status +} + +// state is the part of the agent state that is not a queue. +type state struct { + ReportInterval int `json:"report_interval"` + // LastUpdateCheck is the time of the last check for a new release. + LastUpdateCheck time.Time `json:"last_update_check,omitzero"` +} + +type agent struct { + cfg Config + log *slog.Logger + cp ControlPlane + collector Collector + outbox *outbox + ledger *ledger + + // interval is the report interval, and pending is the number of samples + // since the last report. + interval int + pending int + + // last is the time of the last tick. + last time.Time + + // autoUpdate is true when the agent checks for a new release each day, + // with keys. lastUpdateCheck is the time of the last check, and + // nextUpdateCheck the time of the next one. + autoUpdate bool + keys map[string]ed25519.PublicKey + lastUpdateCheck time.Time + nextUpdateCheck time.Time + + // The waits of the events and the metrics requests after a failure, and + // 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) +} + +func run(ctx context.Context, cfg Config, log *slog.Logger, cp ControlPlane, collector Collector) error { + unlock, err := lock(cfg.StateDir) + if err != nil { + return err + } + defer unlock() + + // A crash during a write can leave a temporary file. + statefile.RemoveTemp(cfg.StateDir) + + saved := loadState(cfg.StateDir, log) + a := &agent{ + cfg: cfg, + log: log, + cp: cp, + collector: collector, + outbox: loadOutbox(cfg.StateDir, log), + ledger: loadLedger(cfg.StateDir, log), + interval: saved.ReportInterval, + lastUpdateCheck: saved.LastUpdateCheck, + } + log.Info("agent started", "version", version.Version, "offset", cfg.Offset(), "report_interval", a.interval) + for _, w := range cfg.Warnings { + log.Warn(w) + } + a.startAutoUpdate(time.Now()) + + // 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) + + if a.loop(ctx) { + // systemd starts the agent again (Restart=always), with the new + // binary after an update. + log.Info("agent exits; systemd starts it again") + return nil + } + log.Info("agent stopped") + + return nil +} + +// loop calls tick at the offset second of each minute until ctx is done. +// Between two ticks, it reads the server each 10 seconds for the peaks of the +// minute. It returns true when a command ends the process. +func (a *agent) loop(ctx context.Context) (exit bool) { + for { + now := time.Now() + next := nextAfter(now, a.last, a.cfg.Offset()) + wake := next + if a.collector != nil { + if r := nextReading(now, a.cfg.Offset()); r.Before(next) { + wake = r + } + } + + timer := time.NewTimer(time.Until(wake)) + select { + case <-ctx.Done(): + timer.Stop() + return false + case <-timer.C: + if wake.Before(next) { + // The reading has the time at which it ran, so that a late + // timer or a step of the wall clock does not change the + // length of a window. + a.collector.Read(time.Now()) + // A reading that ran late must not skip the tick after it. + if time.Now().Before(next) { + continue + } + } + a.last = next + if a.tick(ctx, next) { + return true + } + } + } +} + +// readEvery is the time between two readings of the server (contract v0.4.0). +const readEvery = 10 * time.Second + +// nextReading returns the next reading after now: offset past a full minute, +// and each 10 seconds after it. A tick is also a reading time: the loop then +// runs the tick instead. After a step back of the wall clock, a reading can +// come again; the collector ignores a reading that is not newer than its last. +func nextReading(now time.Time, offset time.Duration) time.Time { + t := now.Truncate(readEvery).Add(offset % readEvery) + for !t.After(now) { + t = t.Add(readEvery) + } + + return t +} + +// maxStepBack is the largest step back of the wall clock after which the loop +// still skips the minute that already ran. A larger step means that the clock +// was wrong before (for example a VM that booted with its clock ahead, then +// NTP): the loop follows the new clock at once, or it would wait, silent, for +// the size of the step. A minute that runs two times is harmless: the control +// plane keeps one sample for each minute. +const maxStepBack = 2 * time.Minute + +// nextAfter returns the next tick after now. After a small step back of the +// wall clock, it never returns a tick at or before last. The timer runs on the +// monotonic clock, but the tick times come from the wall clock: when the wall +// clock steps back, the timer fires before the tick time, and without last +// the same tick would run two times. +func nextAfter(now, last time.Time, offset time.Duration) time.Time { + next := nextTick(now, offset) + if !last.IsZero() && !next.After(last) && last.Sub(now) < maxStepBack { + next = nextTick(last, offset) + } + + return next +} + +// tick does the work of one minute: it takes a sample and, after each +// interval samples, sends a report and runs the new commands. One time each +// day it checks for a new release. It returns true when a command or an +// update 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 { + // The reading has the time at which it ran; the sample has the time + // of its tick, for its minute on the control plane. + s, err := a.collector.Sample(time.Now()) + switch { + case errors.Is(err, ErrNoSample): + a.log.Info("no sample for this minute", "reason", err) + case err != nil: + a.log.Warn("skipping the sample of this minute", "error", err) + default: + s.RecordedAt = now + a.outbox.addSample(cleanSample(s)) + } + } + + if a.report(ctx) { + return true + } + + // On each tick, not only after a report: a long report interval or a + // control plane that is down must not move the check. + if a.autoUpdate && !now.Before(a.nextUpdateCheck) { + return a.checkForUpdate(ctx, now) + } + return false +} + +// report sends a report after each interval samples, and runs the new +// commands. It returns true when a command ends the process. +func (a *agent) report(ctx context.Context) (exit bool) { + a.pending++ + if a.pending < a.interval { + return false + } + + a.log.Debug("report") + 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 +// has the same ID. +func (a *agent) addEvent(name, commandID string, data *wire.EventData) { + a.outbox.addEvent(cleanEvent(wire.Event{ + ID: ulid.Make().String(), + Name: name, + CommandID: commandID, + At: time.Now().UTC(), + Data: data, + })) +} + +// setInterval applies the report interval of a reply. +func (a *agent) setInterval(n int) { + if n < minReportInterval || n > maxReportInterval || n == a.interval { + return + } + + a.log.Info("report interval changed", "from", a.interval, "to", n) + a.interval = n + a.saveState() +} + +// saveState writes all of the state: a write of one field must not remove +// an other. +func (a *agent) saveState() { + s := state{ReportInterval: a.interval, LastUpdateCheck: a.lastUpdateCheck} + if err := statefile.Write(filepath.Join(a.cfg.StateDir, "state.json"), s); err != nil { + a.log.Error("saving the state", "error", err) + } +} + +// nextTick returns the first time after now that is offset past a full minute. +// It comes from the wall clock, so a slow tick moves to the next minute and +// never runs two times in one minute. +func nextTick(now time.Time, offset time.Duration) time.Time { + t := now.Truncate(time.Minute).Add(offset) + if !t.After(now) { + t = t.Add(time.Minute) + } + + return t +} + +// loadState returns the saved state. The report interval is the one of the +// last reply, or 1. +func loadState(dir string, log *slog.Logger) state { + var s state + err := statefile.Read(filepath.Join(dir, "state.json"), &s) + switch { + case errors.Is(err, fs.ErrNotExist): + return state{ReportInterval: minReportInterval} + case err != nil: + log.Warn("ignoring the saved state", "error", err) + return state{ReportInterval: minReportInterval} + } + + s.ReportInterval = clampInterval(s.ReportInterval) + return s +} + +func clampInterval(n int) int { + return min(max(n, minReportInterval), maxReportInterval) +} diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go new file mode 100644 index 0000000..0f3a81a --- /dev/null +++ b/internal/agent/agent_test.go @@ -0,0 +1,319 @@ +package agent + +import ( + "context" + "log/slog" + "path/filepath" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/flywp/server-cli/internal/statefile" +) + +// recorder is a slog handler that keeps the message and the time of each +// record, so that tests can see when the agent did its work. +type recorder struct { + mu sync.Mutex + records []slog.Record +} + +func (r *recorder) Enabled(context.Context, slog.Level) bool { return true } +func (r *recorder) WithAttrs([]slog.Attr) slog.Handler { return r } +func (r *recorder) WithGroup(string) slog.Handler { return r } + +func (r *recorder) Handle(_ context.Context, rec slog.Record) error { + r.mu.Lock() + defer r.mu.Unlock() + r.records = append(r.records, rec) + return nil +} + +// recordsOf returns the records with message msg. +func (r *recorder) recordsOf(msg string) []slog.Record { + r.mu.Lock() + defer r.mu.Unlock() + + var out []slog.Record + for _, rec := range r.records { + if rec.Message == msg { + out = append(out, rec) + } + } + return out +} + +// attr returns the value of the attribute key of a record, as a string, or +// nil. +func attr(rec slog.Record, key string) any { + var v any + rec.Attrs(func(a slog.Attr) bool { + if a.Key == key { + v = a.Value.String() + return false + } + return true + }) + return v +} + +// times returns the times of the records with message msg. +func (r *recorder) times(msg string) []time.Time { + r.mu.Lock() + defer r.mu.Unlock() + + var ts []time.Time + for _, rec := range r.records { + if rec.Message == msg { + ts = append(ts, rec.Time) + } + } + return ts +} + +func TestNextTick(t *testing.T) { + base := time.Date(2026, 9, 22, 10, 0, 0, 0, time.UTC) + tests := []struct { + now time.Time + offset time.Duration + want time.Time + }{ + {base, 17 * time.Second, base.Add(17 * time.Second)}, + {base.Add(10 * time.Second), 17 * time.Second, base.Add(17 * time.Second)}, + // At the offset itself, the next tick is one minute later. + {base.Add(17 * time.Second), 17 * time.Second, base.Add(77 * time.Second)}, + {base.Add(30 * time.Second), 17 * time.Second, base.Add(77 * time.Second)}, + {base.Add(59*time.Second + 999*time.Millisecond), 0, base.Add(time.Minute)}, + } + + for _, tt := range tests { + if got := nextTick(tt.now, tt.offset); !got.Equal(tt.want) { + t.Errorf("nextTick(%s, %v) = %s, want %s", tt.now.Format(time.TimeOnly), tt.offset, got.Format(time.TimeOnly), tt.want.Format(time.TimeOnly)) + } + } +} + +func TestNextAfterAClockStepBack(t *testing.T) { + base := time.Date(2026, 9, 22, 10, 0, 0, 0, time.UTC) + last := base.Add(17 * time.Second) + + // The wall clock stepped back 2 s after the tick at :17, so the timer + // fired at :15 wall time. The next tick must be in the next minute. + if got, want := nextAfter(base.Add(15*time.Second), last, 17*time.Second), base.Add(77*time.Second); !got.Equal(want) { + t.Errorf("nextAfter() = %s, want %s: never the same tick two times", got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } + // Without a step, the last tick does not change the result. + if got, want := nextAfter(base.Add(30*time.Second), last, 17*time.Second), base.Add(77*time.Second); !got.Equal(want) { + t.Errorf("nextAfter() = %s, want %s", got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } + if got, want := nextAfter(base, time.Time{}, 17*time.Second), last; !got.Equal(want) { + t.Errorf("nextAfter() without a last tick = %s, want %s", got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } +} + +func TestNextAfterALargeClockStepBack(t *testing.T) { + base := time.Date(2026, 9, 22, 10, 0, 0, 0, time.UTC) + offset := 17 * time.Second + + tests := []struct { + name string + step time.Duration + want time.Time + }{ + // A step of 90 s is still small: skip to the minute after the last tick. + {"90 seconds", 90 * time.Second, base.Add(time.Minute + offset)}, + // The clock was one hour ahead. Follow the new clock: the next tick + // comes in less than one minute, not in one hour. + {"one hour", time.Hour, base.Add(-time.Hour + time.Minute + offset)}, + {"one year", 365 * 24 * time.Hour, base.Add(-365*24*time.Hour + time.Minute + offset)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + last := base.Add(offset) + now := last.Add(-tt.step).Add(time.Second) + got := nextAfter(now, last, offset) + if !got.Equal(tt.want) { + t.Errorf("nextAfter() = %s, want %s", got, tt.want) + } + if wait := got.Sub(now); wait <= 0 || wait > tt.step+2*time.Minute { + t.Errorf("the loop waits %v after a step back of %v", wait, tt.step) + } + }) + } +} + +func TestLoopTicksAtTheOffsetAndReportsEachInterval(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + if err := statefile.Write(filepath.Join(dir, "state.json"), state{ReportInterval: 2}); err != nil { + t.Fatal(err) + } + + rec := &recorder{} + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error) + go func() { + done <- run(ctx, Config{ServerID: 17, StateDir: dir}, slog.New(rec), &fakeCP{}, nil) + }() + + // The bubble starts at 2000-01-01 00:00:00 UTC. Let 4 minutes pass. + time.Sleep(4 * time.Minute) + cancel() + if err := <-done; err != nil { + t.Fatal(err) + } + + start := time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC) + ticks := rec.times("tick") + if len(ticks) != 4 { + t.Fatalf("ticks at %v, want 4 ticks", ticks) + } + for i, got := range ticks { + if want := start.Add(time.Duration(i)*time.Minute + 17*time.Second); !got.Equal(want) { + t.Errorf("tick %d at %s, want %s", i, got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } + } + + // Report interval 2: a report after tick 2 and tick 4. + if reports := rec.times("report"); len(reports) != 2 || !reports[0].Equal(ticks[1]) || !reports[1].Equal(ticks[3]) { + t.Errorf("reports at %v, want at ticks 2 and 4 (%v)", reports, ticks) + } + if len(rec.times("agent stopped")) != 1 { + t.Error("no \"agent stopped\" log record") + } + }) +} + +func TestNextReading(t *testing.T) { + base := time.Date(2026, 9, 22, 10, 0, 0, 0, time.UTC) + tests := []struct { + now time.Time + offset time.Duration + want time.Time + }{ + {base, 17 * time.Second, base.Add(7 * time.Second)}, + {base.Add(7 * time.Second), 17 * time.Second, base.Add(17 * time.Second)}, + {base.Add(18 * time.Second), 17 * time.Second, base.Add(27 * time.Second)}, + {base.Add(58 * time.Second), 17 * time.Second, base.Add(67 * time.Second)}, + {base.Add(59*time.Second + 999*time.Millisecond), 0, base.Add(time.Minute)}, + {base.Add(3 * time.Second), 43 * time.Second, base.Add(3 * time.Second).Add(10 * time.Second)}, + } + + for _, tt := range tests { + if got := nextReading(tt.now, tt.offset); !got.Equal(tt.want) { + t.Errorf("nextReading(%s, %v) = %s, want %s", tt.now.Format(time.TimeOnly), tt.offset, got.Format(time.TimeOnly), tt.want.Format(time.TimeOnly)) + } + } +} + +func TestLoopReadsEach10SecondsBetweenTheTicks(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + collector := &fakeCollector{} + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error) + go func() { + done <- run(ctx, Config{ServerID: 17, StateDir: t.TempDir()}, slog.New(slog.DiscardHandler), &fakeCP{}, collector) + }() + + time.Sleep(2 * time.Minute) + cancel() + if err := <-done; err != nil { + t.Fatal(err) + } + + // The ticks are at :17. The readings are at :07, :27, :37, :47 and :57: + // the loop never reads at a tick, because the tick takes the reading. + start := time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC) + var want []time.Time + for s := 7 * time.Second; s < 2*time.Minute; s += 10 * time.Second { + if s%time.Minute != 17*time.Second { + want = append(want, start.Add(s)) + } + } + got := collector.readTimes() + if len(got) != len(want) { + t.Fatalf("readings at %v, want %v", got, want) + } + for i := range want { + if !got[i].Equal(want[i]) { + t.Errorf("reading %d at %s, want %s", i, got[i].Format(time.TimeOnly), want[i].Format(time.TimeOnly)) + } + } + if collector.n != 2 { + t.Errorf("%d samples, want 2 (at 0:17 and 1:17)", collector.n) + } + }) +} + +func TestASlowReadingDoesNotSkipTheTick(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + // Each reading takes 10 s: the reading at :07 ends at the tick. + collector := &fakeCollector{readTime: 10 * time.Second} + rec := &recorder{} + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error) + go func() { + done <- run(ctx, Config{ServerID: 17, StateDir: t.TempDir()}, slog.New(rec), &fakeCP{}, collector) + }() + + time.Sleep(2*time.Minute + 30*time.Second) + cancel() + if err := <-done; err != nil { + t.Fatal(err) + } + + start := time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC) + ticks := rec.times("tick") + if len(ticks) != 3 { + t.Fatalf("ticks at %v, want 3 ticks", ticks) + } + for i, got := range ticks { + if want := start.Add(time.Duration(i)*time.Minute + 17*time.Second); !got.Equal(want) { + t.Errorf("tick %d at %s, want %s", i, got.Format(time.TimeOnly), want.Format(time.TimeOnly)) + } + } + }) +} + +func TestRunRefusesASecondAgent(t *testing.T) { + dir := t.TempDir() + unlock, err := lock(dir) + if err != nil { + t.Fatal(err) + } + defer unlock() + + if err := run(context.Background(), Config{StateDir: dir}, slog.New(&recorder{}), &fakeCP{}, nil); err == nil { + t.Fatal("Run() = nil, want an error while a different agent holds the lock") + } +} + +func TestLoadStateInterval(t *testing.T) { + tests := []struct { + name string + saved *state + want int + }{ + {"no state", nil, 1}, + {"saved", &state{ReportInterval: 5}, 5}, + {"too low", &state{ReportInterval: 0}, 1}, + {"too high", &state{ReportInterval: 60}, 10}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + if tt.saved != nil { + if err := statefile.Write(filepath.Join(dir, "state.json"), tt.saved); err != nil { + t.Fatal(err) + } + } + + if got := loadState(dir, slog.New(&recorder{})).ReportInterval; got != tt.want { + t.Errorf("loadState().ReportInterval = %d, want %d", got, tt.want) + } + }) + } +} diff --git a/internal/agent/autoupdate.go b/internal/agent/autoupdate.go new file mode 100644 index 0000000..60d2754 --- /dev/null +++ b/internal/agent/autoupdate.go @@ -0,0 +1,165 @@ +package agent + +import ( + "context" + "crypto/ed25519" + "errors" + "os" + "path/filepath" + "runtime" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/release" + "github.com/flywp/server-cli/internal/version" + "golang.org/x/mod/semver" + "golang.org/x/sys/unix" +) + +const ( + // autoUpdateWait is the time between the signature of a release and its + // install by the agents. In this time the maintainer can mark a bad + // release as a pre-release: /releases/latest then stops returning it. + autoUpdateWait = 24 * time.Hour + + // checkGap is the shortest time between two checks. With it, a clock + // step cannot cause two checks in one day. + checkGap = 12 * time.Hour +) + +// The parts of an update by release. Tests replace them. +var ( + latestRelease = release.LatestRelease + signedChecksum = func(ctx context.Context, rel *release.GithubRelease, keys map[string]ed25519.PublicKey, now time.Time) (*release.Signed, error) { + return release.SignedChecksum(ctx, rel, runtime.GOOS, runtime.GOARCH, keys, now) + } + releaseKeys = release.TrustedKeys + binaryDirWritable = func() bool { + exe, err := os.Executable() + if err != nil { + return false + } + if exe, err = filepath.EvalSymlinks(exe); err != nil { + return false + } + return unix.Access(filepath.Dir(exe), unix.W_OK) == nil + } +) + +// updateSlot is the time after 00:00 UTC at which the server checks for a new +// release each day. It spreads the installs of a release over one day, so +// that the first servers show a problem before the others get the release. +func updateSlot(serverID int64) time.Duration { + return time.Duration(serverID%(24*60)) * time.Minute +} + +// nextCheck returns the time of the next check for a new release: the first +// slot that comes more than checkGap after the last check. An agent that +// never checked checks at once. +func nextCheck(last time.Time, slot time.Duration) time.Time { + if last.IsZero() { + return time.Time{} + } + + after := last.Add(checkGap) + // Truncate counts from the zero time, which is a midnight UTC. + t := after.UTC().Truncate(24 * time.Hour).Add(slot) + if !t.After(after) { + t = t.Add(24 * time.Hour) + } + + return t +} + +// startAutoUpdate decides at start whether the agent checks for new releases, +// and when it checks first. +func (a *agent) startAutoUpdate(now time.Time) { + switch keys, err := releaseKeys(); { + case !a.cfg.AutoUpdate: + a.log.Info("auto-update is off", "key", EnvAutoUpdate) + return + case !semver.IsValid(version.Version) || release.IsLocalBuild(version.Version): + // A build without a release tag (go build gives "dev") has no + // place in the order of the releases. A make build of a commit + // after a tag (v0.2.0-3-gabcdef1) sorts before that tag, so the + // release would replace code that is newer. + a.log.Info("auto-update is off: this build has no release version", "version", version.Version) + return + case err != nil: + a.log.Error("auto-update is off: the release keys of this build are not valid", "error", err) + return + case len(keys) == 0: + a.log.Warn("auto-update is off: this build trusts no release key") + return + case !binaryDirWritable(): + a.log.Warn("auto-update is off: the agent cannot write the folder of its binary") + return + default: + a.keys = keys + } + + // A check time after now comes from a clock that was wrong. Without + // this, the agent could wait for a day that is years away. + if a.lastUpdateCheck.After(now) { + a.log.Warn("the time of the last release check is in the future; checking again", "last", a.lastUpdateCheck) + a.lastUpdateCheck = time.Time{} + } + + a.autoUpdate = true + a.nextUpdateCheck = nextCheck(a.lastUpdateCheck, updateSlot(a.cfg.ServerID)) +} + +// checkForUpdate installs the latest release when it is newer, has a valid +// signature and waited autoUpdateWait after the signature. It returns true +// when it installed the release: the process must exit, so that systemd +// starts the new binary. +func (a *agent) checkForUpdate(ctx context.Context, now time.Time) (exit bool) { + // Save the time first: a crash or a failure waits for the next day, and + // does not check at each restart. + a.lastUpdateCheck = now + a.nextUpdateCheck = nextCheck(now, updateSlot(a.cfg.ServerID)) + a.saveState() + + log := a.log.With("version", version.Version, "next_check", a.nextUpdateCheck) + + ctx, cancel := context.WithTimeout(ctx, updateTimeout) + defer cancel() + + rel, err := latestRelease(ctx) + if err != nil { + log.Warn("checking for a new release failed", "error", err) + return false + } + log = log.With("release", rel.TagName) + + // The agent never downgrades. A version without an order (a dev build) + // does not update by itself. + if cmp, ok := compareVersions(version.Version, rel.TagName); !ok || cmp >= 0 { + log.Info("checked for a new release: none to install") + return false + } + + signed, err := signedChecksum(ctx, rel, a.keys, now) + switch { + case errors.Is(err, release.ErrNotSigned): + log.Info("a new release waits for its signature") + return false + case err != nil: + log.Error("not installing a new release", "error", err) + return false + } + + if ready := signed.SignedAt.Add(autoUpdateWait); now.Before(ready) { + log.Info("a new release installs after its wait", "signed_at", signed.SignedAt, "ready_at", ready) + return false + } + + log.Info("installing a new release") + if err := updateBinary(ctx, wire.UpdateArgs{URL: signed.URL, Version: rel.TagName, SHA256: signed.SHA256}); err != nil { + log.Error("the update failed; the old binary continues", "error", err) + return false + } + + log.Info("installed the new release; exiting so that systemd starts it") + return true +} diff --git a/internal/agent/autoupdate_test.go b/internal/agent/autoupdate_test.go new file mode 100644 index 0000000..e4d2d11 --- /dev/null +++ b/internal/agent/autoupdate_test.go @@ -0,0 +1,402 @@ +package agent + +import ( + "context" + "crypto/ed25519" + "errors" + "log/slog" + "path/filepath" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/release" + "github.com/flywp/server-cli/internal/statefile" +) + +func TestUpdateSlot(t *testing.T) { + tests := map[int64]time.Duration{ + 0: 0, + 17: 17 * time.Minute, + 1439: 1439 * time.Minute, + 1445: 5 * time.Minute, + } + for id, want := range tests { + if got := updateSlot(id); got != want { + t.Errorf("updateSlot(%d) = %v, want %v", id, got, want) + } + } +} + +func TestNextCheck(t *testing.T) { + day := time.Date(2026, 9, 23, 0, 0, 0, 0, time.UTC) + slot := 17 * time.Minute + tests := []struct { + name string + last time.Time + want time.Time + }{ + {"never checked", time.Time{}, time.Time{}}, + {"checked at the slot", day.Add(slot), day.Add(24*time.Hour + slot)}, + {"checked late in the day", day.Add(20 * time.Hour), day.Add(48*time.Hour + slot)}, + {"checked just before the slot", day.Add(slot - time.Minute), day.Add(24*time.Hour + slot)}, + // Exactly 12 hours before a slot is not more than 12 hours. + {"checked 12 hours before the slot", day.Add(slot - 12*time.Hour), day.Add(24*time.Hour + slot)}, + {"a zone that is not UTC", day.Add(slot).In(time.FixedZone("UTC+6", 6*3600)), day.Add(24*time.Hour + slot)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := nextCheck(tt.last, slot); !got.Equal(tt.want) { + t.Errorf("nextCheck(%v) = %v, want %v", tt.last, got, tt.want) + } + }) + } +} + +// fakeReleases replaces GitHub and the signature check for one test. +type fakeReleases struct { + mu sync.Mutex + tag string + latestErr error + signedAt time.Time + signErr error + + checks []time.Time + signed int +} + +func useFakeReleases(t *testing.T, f *fakeReleases) { + t.Helper() + + oldLatest, oldSigned, oldKeys, oldWritable := latestRelease, signedChecksum, releaseKeys, binaryDirWritable + t.Cleanup(func() { + latestRelease, signedChecksum, releaseKeys, binaryDirWritable = oldLatest, oldSigned, oldKeys, oldWritable + }) + + latestRelease = func(context.Context) (*release.GithubRelease, error) { + f.mu.Lock() + defer f.mu.Unlock() + f.checks = append(f.checks, time.Now()) + if f.latestErr != nil { + return nil, f.latestErr + } + return &release.GithubRelease{TagName: f.tag}, nil + } + signedChecksum = func(context.Context, *release.GithubRelease, map[string]ed25519.PublicKey, time.Time) (*release.Signed, error) { + f.mu.Lock() + defer f.mu.Unlock() + f.signed++ + if f.signErr != nil { + return nil, f.signErr + } + return &release.Signed{URL: "https://example.com/fly-linux-amd64.tar.gz", SHA256: "abc123", SignedAt: f.signedAt}, nil + } + releaseKeys = func() (map[string]ed25519.PublicKey, error) { + return map[string]ed25519.PublicKey{"test": make(ed25519.PublicKey, ed25519.PublicKeySize)}, nil + } + binaryDirWritable = func() bool { return true } +} + +func (f *fakeReleases) checkTimes() []time.Time { + f.mu.Lock() + defer f.mu.Unlock() + return f.checks +} + +// setReportInterval saves the report interval n in dir, as a reply would. +func setReportInterval(t *testing.T, dir string, n int) { + t.Helper() + if err := statefile.Write(filepath.Join(dir, "state.json"), state{ReportInterval: n}); err != nil { + t.Fatal(err) + } +} + +// runAuto runs an agent with server id 17 (slot 00:17 UTC) and auto-update +// on, until it exits or until d passes. It returns true when it exited. +func runAuto(t *testing.T, d time.Duration, dir string, cp *fakeCP, c Collector) (exited bool, rec *recorder) { + t.Helper() + + rec = &recorder{} + ctx, cancel := context.WithTimeout(context.Background(), d) + defer cancel() + cfg := Config{ServerID: 17, StateDir: dir, AutoUpdate: true} + if err := run(ctx, cfg, slog.New(rec), cp, c); err != nil { + t.Fatal(err) + } + return ctx.Err() == nil, rec +} + +func TestAutoUpdateChecksOneTimeEachDay(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + f := &fakeReleases{tag: "v0.1.1"} + useFakeReleases(t, f) + calls := fakeUpdate(t, nil) + dir := t.TempDir() + setReportInterval(t, dir, 10) + + exited, rec := runAuto(t, 72*time.Hour+time.Minute, dir, &fakeCP{}, nil) + + // The first check at the first tick, which is not a report tick. Then + // one check each day at the slot of the server. + want := []time.Time{ + at(0), + bubbleStart.Add(24*time.Hour + 17*time.Minute + 17*time.Second), + bubbleStart.Add(48*time.Hour + 17*time.Minute + 17*time.Second), + } + if got := f.checkTimes(); !equalTimes(got, want) { + t.Errorf("checks at %v, want %v", got, want) + } + // The release is older: the agent does not fetch its signature. + if exited || f.signed != 0 || len(*calls) != 0 { + t.Errorf("exited %v, %d signature checks, %d installs; want none for an older release", exited, f.signed, len(*calls)) + } + if len(rec.times("checked for a new release: none to install")) != 3 { + t.Error("want one log record for each check") + } + }) +} + +func TestAutoUpdateInstallsAfterTheWait(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + // Signed one hour before the start: the first check is too early. + f := &fakeReleases{tag: "v0.2.1", signedAt: bubbleStart.Add(-time.Hour)} + useFakeReleases(t, f) + calls := fakeUpdate(t, nil) + dir := t.TempDir() + setReportInterval(t, dir, 10) + + exited, rec := runAuto(t, 72*time.Hour, dir, &fakeCP{}, nil) + + if !exited { + t.Fatal("the agent did not exit after the update") + } + if len(rec.times("a new release installs after its wait")) != 1 { + t.Error("want one \"installs after its wait\" record, at the first check") + } + // The second check, on the second day, installs the release. + checks := f.checkTimes() + if len(checks) != 2 || !checks[1].Equal(bubbleStart.Add(24*time.Hour+17*time.Minute+17*time.Second)) { + t.Errorf("checks at %v, want the install at the second check", checks) + } + want := wire.UpdateArgs{URL: "https://example.com/fly-linux-amd64.tar.gz", Version: "v0.2.1", SHA256: "abc123"} + if len(*calls) != 1 || (*calls)[0] != want { + t.Errorf("installs = %+v, want one install of %+v", *calls, want) + } + }) +} + +func TestAutoUpdateDoesNotInstall(t *testing.T) { + old := bubbleStart.Add(-48 * time.Hour) + tests := []struct { + name string + running string + releases *fakeReleases + wantLog string + wantSign int + }{ + {"the same version", "v0.2.1", &fakeReleases{tag: "v0.2.1", signedAt: old}, "checked for a new release: none to install", 0}, + {"an older release", "v0.3.0", &fakeReleases{tag: "v0.2.1", signedAt: old}, "checked for a new release: none to install", 0}, + {"no signature", "v0.2.0", &fakeReleases{tag: "v0.2.1", signErr: release.ErrNotSigned}, "a new release waits for its signature", 1}, + {"a bad signature", "v0.2.0", &fakeReleases{tag: "v0.2.1", signErr: errors.New("the signature is not correct")}, "not installing a new release", 1}, + {"GitHub is down", "v0.2.0", &fakeReleases{latestErr: errors.New("unexpected response: 502")}, "checking for a new release failed", 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, tt.running) + f := tt.releases + useFakeReleases(t, f) + calls := fakeUpdate(t, nil) + dir := t.TempDir() + setReportInterval(t, dir, 10) + + // Two hours: a failure waits for the next day, not for the + // next tick. + exited, rec := runAuto(t, 2*time.Hour, dir, &fakeCP{}, nil) + + if exited || len(*calls) != 0 { + t.Errorf("exited %v with %d installs, want no install", exited, len(*calls)) + } + if n := len(f.checkTimes()); n != 1 { + t.Errorf("%d checks in two hours, want 1", n) + } + if f.signed != tt.wantSign { + t.Errorf("%d signature checks, want %d", f.signed, tt.wantSign) + } + if len(rec.times(tt.wantLog)) != 1 { + t.Errorf("no %q log record", tt.wantLog) + } + }) + }) + } +} + +func TestAutoUpdateFailedInstallKeepsTheAgent(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + useFakeReleases(t, &fakeReleases{tag: "v0.2.1", signedAt: bubbleStart.Add(-48 * time.Hour)}) + calls := fakeUpdate(t, errors.New("the sha256 of the download does not agree")) + + exited, rec := runAuto(t, 5*time.Minute, t.TempDir(), &fakeCP{}, nil) + + if exited || len(*calls) != 1 { + t.Errorf("exited %v after %d installs, want one failed install and no exit", exited, len(*calls)) + } + if len(rec.times("the update failed; the old binary continues")) != 1 { + t.Error("no log record for the failed update") + } + }) +} + +func TestAutoUpdateIsOff(t *testing.T) { + tests := []struct { + name string + setup func(t *testing.T, cfg *Config) + wantLog string + }{ + {"by the switch", func(_ *testing.T, cfg *Config) { cfg.AutoUpdate = false }, "auto-update is off"}, + {"for a dev build", func(t *testing.T, _ *Config) { setVersion(t, "dev") }, "auto-update is off: this build has no release version"}, + {"for a build after a tag", func(t *testing.T, _ *Config) { setVersion(t, "v0.2.0-3-gabcdef1") }, "auto-update is off: this build has no release version"}, + {"for a build with changes", func(t *testing.T, _ *Config) { setVersion(t, "v0.2.0-dirty") }, "auto-update is off: this build has no release version"}, + {"without a trusted key", func(*testing.T, *Config) { + releaseKeys = func() (map[string]ed25519.PublicKey, error) { return nil, nil } + }, "auto-update is off: this build trusts no release key"}, + {"with keys that are not valid", func(*testing.T, *Config) { + releaseKeys = func() (map[string]ed25519.PublicKey, error) { return nil, errors.New("bad key") } + }, "auto-update is off: the release keys of this build are not valid"}, + {"when the binary folder is not writable", func(*testing.T, *Config) { + binaryDirWritable = func() bool { return false } + }, "auto-update is off: the agent cannot write the folder of its binary"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + f := &fakeReleases{tag: "v0.2.1", signedAt: bubbleStart.Add(-48 * time.Hour)} + useFakeReleases(t, f) + cfg := Config{ServerID: 17, StateDir: t.TempDir(), AutoUpdate: true} + tt.setup(t, &cfg) + + rec := &recorder{} + ctx, cancel := context.WithTimeout(context.Background(), 48*time.Hour) + defer cancel() + if err := run(ctx, cfg, slog.New(rec), &fakeCP{}, nil); err != nil { + t.Fatal(err) + } + + if n := len(f.checkTimes()); n != 0 { + t.Errorf("%d checks, want none: the agent must not call GitHub", n) + } + if len(rec.times(tt.wantLog)) != 1 { + t.Errorf("no %q log record", tt.wantLog) + } + }) + }) + } +} + +func TestAutoUpdateCheckTimeSurvivesARestart(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + f := &fakeReleases{tag: "v0.1.1"} + useFakeReleases(t, f) + dir := t.TempDir() + + runAuto(t, 5*time.Minute, dir, &fakeCP{}, nil) + runAuto(t, 5*time.Minute, dir, &fakeCP{}, nil) + + // The second process knows the check of the first one. + if n := len(f.checkTimes()); n != 1 { + t.Errorf("%d checks in two runs on one day, want 1", n) + } + }) +} + +func TestAutoUpdateCheckTimeInTheFuture(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + f := &fakeReleases{tag: "v0.1.1"} + useFakeReleases(t, f) + dir := t.TempDir() + // A clock that was years ahead saved this time. + saved := state{ReportInterval: 10, LastUpdateCheck: bubbleStart.AddDate(5, 0, 0)} + if err := statefile.Write(filepath.Join(dir, "state.json"), saved); err != nil { + t.Fatal(err) + } + + runAuto(t, 2*time.Minute, dir, &fakeCP{}, nil) + + if got := f.checkTimes(); len(got) != 1 || !got[0].Equal(at(0)) { + t.Errorf("checks at %v, want one check at the first tick", got) + } + }) +} + +func TestAutoUpdateWhileTheEventsFail(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + f := &fakeReleases{tag: "v0.1.1"} + useFakeReleases(t, f) + cp := &fakeCP{eventsReply: func(int, *wire.EventsRequest) (*wire.EventsReply, error) { + return nil, errors.New("connection refused") + }} + + // Report interval 1: each tick is a report tick, and each report + // stops before the poll. + runAuto(t, 2*time.Minute, t.TempDir(), cp, nil) + + if got := f.checkTimes(); len(got) != 1 || !got[0].Equal(at(0)) { + t.Errorf("checks at %v, want a check at the first tick", got) + } + }) +} + +func TestNewReportIntervalKeepsTheCheckTime(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + setVersion(t, "v0.2.0") + useFakeReleases(t, &fakeReleases{tag: "v0.1.1"}) + dir := t.TempDir() + // The first reply sets 3, before the check. A later reply sets 5, + // after the check: that write must keep the check time. + cp := &fakeCP{metricsReply: func(call int, req *wire.MetricsRequest) (*wire.MetricsReply, error) { + n := 5 + if call == 0 { + n = 3 + } + return &wire.MetricsReply{Accepted: len(req.Samples), ReportInterval: n}, nil + }} + + runAuto(t, 5*time.Minute, dir, cp, &fakeCollector{}) + + var s state + if err := statefile.Read(filepath.Join(dir, "state.json"), &s); err != nil { + t.Fatal(err) + } + if s.ReportInterval != 5 || !s.LastUpdateCheck.Equal(at(0)) { + t.Errorf("state = %+v, want interval 5 and the check at %v", s, at(0)) + } + }) +} + +func TestConfigWarningsAreLogged(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + rec := &recorder{} + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + cfg := Config{ServerID: 17, StateDir: t.TempDir(), Warnings: []string{"FLY_AGENT_AUTO_UPDATE=\"of\" is not on or off; auto-update stays on"}} + if err := run(ctx, cfg, slog.New(rec), &fakeCP{}, nil); err != nil { + t.Fatal(err) + } + + if len(rec.times(cfg.Warnings[0])) != 1 { + t.Error("the warning of the configuration is not in the log") + } + }) +} diff --git a/internal/agent/clean.go b/internal/agent/clean.go new file mode 100644 index 0000000..fdfe535 --- /dev/null +++ b/internal/agent/clean.go @@ -0,0 +1,161 @@ +package agent + +import ( + "math" + "time" + "unicode/utf8" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/oklog/ulid/v2" +) + +// The value ranges of the contract (section 4, "Value ranges", and section 6). +// A value outside its range makes the control plane refuse the whole request +// with a 400, and the agent then drops up to 240 samples. So the agent keeps +// each value in its range before the value goes into a queue. +const ( + maxLoad = 999999.99 + maxVersionLen = 32 + maxStatusTextLen = 255 + maxArchLen = 16 + maxCPUCount = 4096 + maxSites = 1000 + maxDirectoryLen = 255 + maxEventNameLen = 64 + maxErrorLen = 2000 +) + +func cleanSample(s wire.Sample) wire.Sample { + s.RecordedAt = s.RecordedAt.UTC().Truncate(time.Second) + s.CPUPercent = clamp(s.CPUPercent, 0, 100) + s.Load1 = clamp(s.Load1, 0, maxLoad) + for _, v := range []*uint64{ + &s.MemoryUsedBytes, &s.MemoryTotalBytes, &s.SwapUsedBytes, &s.SwapTotalBytes, + &s.DiskUsedBytes, &s.DiskTotalBytes, &s.NetInBytes, &s.NetOutBytes, + } { + *v = clampInt(*v) + } + for _, v := range []**float64{ + &s.CPUMaxPercent, &s.CPUPressurePercent, &s.CPUPressureMaxPercent, &s.MemoryPressurePercent, + &s.MemoryPressureMaxPercent, &s.IOPressurePercent, &s.IOPressureMaxPercent, + } { + *v = clampPtr(*v, 0, 100) + } + for _, v := range []**uint64{ + &s.MemoryUsedMaxBytes, &s.SwapUsedMaxBytes, &s.NetInMaxBytesPerSecond, &s.NetOutMaxBytesPerSecond, + &s.DiskReadBytes, &s.DiskWriteBytes, &s.DiskReadOps, &s.DiskWriteOps, + &s.DiskReadMaxBytesPerSecond, &s.DiskWriteMaxBytesPerSecond, &s.DiskReadMaxOpsPerSecond, &s.DiskWriteMaxOpsPerSecond, + } { + *v = clampIntPtr(*v) + } + s.Sites = cleanSites(s.Sites) + return s +} + +// cleanSites keeps at most maxSites items. It drops an item whose directory +// is empty or too long: a cut directory would match an other site. A nil +// slice stays nil, and an empty slice stays empty: they mean different things. +func cleanSites(sites []wire.Site) []wire.Site { + if sites == nil { + return nil + } + + out := make([]wire.Site, 0, min(len(sites), maxSites)) + for _, site := range sites { + if len(out) == maxSites { + break + } + if n := utf8.RuneCountInString(site.Directory); n == 0 || n > maxDirectoryLen { + continue + } + site.CPUPercent = clampPtr(site.CPUPercent, 0, 100) + site.MemoryUsedBytes = clampIntPtr(site.MemoryUsedBytes) + site.DiskUsedBytes = clampIntPtr(site.DiskUsedBytes) + out = append(out, site) + } + return out +} + +func cleanStatus(s wire.Status) wire.Status { + s.OS = truncate(s.OS, maxStatusTextLen) + s.Kernel = truncate(s.Kernel, maxStatusTextLen) + s.Arch = truncate(s.Arch, maxArchLen) + s.UpdatesTotal = clampIntPtr(s.UpdatesTotal) + s.UpdatesSecurity = clampIntPtr(s.UpdatesSecurity) + s.UptimeSeconds = clampInt(s.UptimeSeconds) + if s.CPUCount != nil && (*s.CPUCount < 1 || *s.CPUCount > maxCPUCount) { + s.CPUCount = nil + } + if s.DockerStatus != nil { + switch *s.DockerStatus { + case wire.DockerRunning, wire.DockerNotRunning, wire.DockerNotInstalled: + default: + s.DockerStatus = nil + } + } + if s.DockerVersion != nil { + v := truncate(*s.DockerVersion, maxVersionLen) + s.DockerVersion = &v + } + return s +} + +func cleanEvent(e wire.Event) wire.Event { + // A command id that is not a ULID makes a 400, and a 400 drops all the + // events of the request. Without the id, the event changes nothing. + if e.CommandID != "" { + if _, err := ulid.ParseStrict(e.CommandID); err != nil { + e.CommandID = "" + } + } + e.Name = truncate(e.Name, maxEventNameLen) + if e.Data != nil { + d := *e.Data + d.Version = truncate(d.Version, maxVersionLen) + d.Error = truncate(d.Error, maxErrorLen) + e.Data = &d + } + return e +} + +// clampInt keeps v in the range of a signed 64-bit integer: the control plane +// (PHP) cannot hold a larger integer, and fails the request with a 500. The +// agent would then send the same sample again for 24 hours. +func clampInt(v uint64) uint64 { + return min(v, math.MaxInt64) +} + +// clampIntPtr is clampInt for a value that can be "not known" (nil). +func clampIntPtr(v *uint64) *uint64 { + if v == nil { + return nil + } + c := clampInt(*v) + return &c +} + +// clamp keeps v between lo and hi. NaN becomes lo. +func clamp(v, lo, hi float64) float64 { + if math.IsNaN(v) { + return lo + } + return min(max(v, lo), hi) +} + +// clampPtr is clamp for a value that can be "not known" (nil). +func clampPtr(v *float64, lo, hi float64) *float64 { + if v == nil { + return nil + } + c := clamp(*v, lo, hi) + return &c +} + +// truncate cuts s to at most n characters. The control plane counts +// characters, not bytes. +func truncate(s string, n int) string { + if utf8.RuneCountInString(s) <= n { + return s + } + return string([]rune(s)[:n]) +} diff --git a/internal/agent/client.go b/internal/agent/client.go new file mode 100644 index 0000000..83f112d --- /dev/null +++ b/internal/agent/client.go @@ -0,0 +1,139 @@ +package agent + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "strconv" + "strings" + "time" + + "github.com/flywp/server-cli/internal/version" +) + +// requestTimeout limits each request to the control plane. +const requestTimeout = 20 * time.Second + +// maxReplySize limits the reply body that the agent reads. +const maxReplySize = 1 << 20 + +// maxRetryAfter limits the wait that a Retry-After header can ask for, so +// that one bad reply cannot stop the sends for a long time. +const maxRetryAfter = time.Hour + +// Client sends requests to the control plane. +type Client struct { + base *url.URL + token string + http *http.Client +} + +// NewClient returns a client for the control plane of cfg. A nil hc uses a +// client with the default timeout. The client never follows a redirect: see +// noRedirects. +func NewClient(cfg Config, hc *http.Client) *Client { + if hc == nil { + hc = &http.Client{Timeout: requestTimeout} + } + c := *hc + c.CheckRedirect = noRedirects + + return &Client{base: cfg.URL, token: cfg.Token, http: &c} +} + +// noRedirects makes a redirect a reply like any other status that is not 200. +// A followed redirect replays a POST as a GET without the body, so a 200 to +// that GET would drop samples that were never stored. It can also send the +// token over plain http. +func noRedirects(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse +} + +// maxErrorBody limits the part of an error reply that the agent keeps for +// its log. +const maxErrorBody = 1 << 10 + +// StatusError is a reply from the control plane that is not 200 OK. +type StatusError struct { + StatusCode int + // RetryAfter is the wait that the Retry-After header asks for, or 0. + RetryAfter time.Duration + // Body is the start of the reply, for example the validation errors of + // a 400. + Body string +} + +func (e *StatusError) Error() string { + return fmt.Sprintf("control plane replied %d %s", e.StatusCode, http.StatusText(e.StatusCode)) +} + +// do sends in as JSON (no body when in is nil) to path and decodes a 200 reply +// into out (no decoding when out is nil). Another status is a *StatusError. +func (c *Client) do(ctx context.Context, method, path string, in, out any) error { + var body io.Reader + if in != nil { + data, err := json.Marshal(in) + if err != nil { + return err + } + body = bytes.NewReader(data) + } + + req, err := http.NewRequestWithContext(ctx, method, c.base.JoinPath(path).String(), body) + if err != nil { + return err + } + req.Header.Set("Authorization", "Bearer "+c.token) + req.Header.Set("User-Agent", "fly/"+version.Version) + req.Header.Set("Accept", "application/json") + if in != nil { + req.Header.Set("Content-Type", "application/json") + } + + resp, err := c.http.Do(req) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(io.LimitReader(resp.Body, maxErrorBody)) + return &StatusError{ + StatusCode: resp.StatusCode, + RetryAfter: retryAfter(resp.Header.Get("Retry-After"), time.Now()), + Body: strings.ToValidUTF8(string(body), "?"), + } + } + + if out == nil { + return nil + } + if err := json.NewDecoder(io.LimitReader(resp.Body, maxReplySize)).Decode(out); err != nil { + return fmt.Errorf("reading the reply of %s: %w", path, err) + } + + return nil +} + +// retryAfter reads a Retry-After value: a number of seconds or an HTTP date. +// The wait is at most maxRetryAfter. +func retryAfter(v string, now time.Time) time.Duration { + if v == "" { + return 0 + } + if s, err := strconv.ParseInt(v, 10, 64); err == nil { + if s <= 0 { + return 0 + } + return time.Duration(min(s, int64(maxRetryAfter/time.Second))) * time.Second + } + if t, err := http.ParseTime(v); err == nil { + return min(max(t.Sub(now), 0), maxRetryAfter) + } + + return 0 +} diff --git a/internal/agent/client_test.go b/internal/agent/client_test.go new file mode 100644 index 0000000..a0fe988 --- /dev/null +++ b/internal/agent/client_test.go @@ -0,0 +1,168 @@ +package agent + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + "time" + + "github.com/flywp/server-cli/internal/version" +) + +func testClient(t *testing.T, handler http.HandlerFunc) *Client { + t.Helper() + + srv := httptest.NewTLSServer(handler) + t.Cleanup(srv.Close) + + u, err := url.Parse(srv.URL + "/base") + if err != nil { + t.Fatal(err) + } + + return NewClient(Config{URL: u, Token: testToken}, srv.Client()) +} + +func TestClientSendsTheContractHeaders(t *testing.T) { + var got *http.Request + var body map[string]int + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + got = r + _ = json.NewDecoder(r.Body).Decode(&body) + _, _ = w.Write([]byte(`{"report_interval": 5}`)) + }) + + var reply struct { + ReportInterval int `json:"report_interval"` + } + if err := c.do(context.Background(), http.MethodPost, "agent/v1/metrics", map[string]int{"n": 1}, &reply); err != nil { + t.Fatal(err) + } + + if got.URL.Path != "/base/agent/v1/metrics" { + t.Errorf("path = %q, want the contract path under the base URL", got.URL.Path) + } + checks := map[string]string{ + "Authorization": "Bearer " + testToken, + "User-Agent": "fly/" + version.Version, + "Accept": "application/json", + "Content-Type": "application/json", + } + for header, want := range checks { + if v := got.Header.Get(header); v != want { + t.Errorf("%s = %q, want %q", header, v, want) + } + } + if body["n"] != 1 { + t.Errorf("body = %v, want the JSON of the request", body) + } + if reply.ReportInterval != 5 { + t.Errorf("reply = %+v, want the decoded JSON", reply) + } +} + +func TestClientWithoutBody(t *testing.T) { + var got *http.Request + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + got = r + _, _ = w.Write([]byte(`{}`)) + }) + + if err := c.do(context.Background(), http.MethodGet, "agent/v1/commands", nil, nil); err != nil { + t.Fatal(err) + } + if got.Method != http.MethodGet || got.Header.Get("Content-Type") != "" { + t.Errorf("request = %s with Content-Type %q, want GET without a body", got.Method, got.Header.Get("Content-Type")) + } +} + +func TestClientStatusError(t *testing.T) { + tests := []struct { + name string + code int + retryAfter string + want time.Duration + }{ + {"unauthorized", http.StatusUnauthorized, "", 0}, + {"throttled", http.StatusTooManyRequests, "30", 30 * time.Second}, + {"unavailable", http.StatusServiceUnavailable, "", 0}, + {"bad gateway", http.StatusBadGateway, "", 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + if tt.retryAfter != "" { + w.Header().Set("Retry-After", tt.retryAfter) + } + w.WriteHeader(tt.code) + _, _ = w.Write([]byte(`{"message":"The samples.0.cpu_percent field must be between 0 and 100."}`)) + }) + + err := c.do(context.Background(), http.MethodPost, "agent/v1/events", map[string]int{}, nil) + var statusErr *StatusError + if !errors.As(err, &statusErr) { + t.Fatalf("do() error = %v, want a *StatusError", err) + } + if statusErr.StatusCode != tt.code || statusErr.RetryAfter != tt.want { + t.Errorf("StatusError = %+v, want code %d and wait %v", statusErr, tt.code, tt.want) + } + if !strings.Contains(statusErr.Body, "cpu_percent") { + t.Errorf("StatusError.Body = %q, want the reply for the log", statusErr.Body) + } + }) + } +} + +func TestRetryAfter(t *testing.T) { + now := time.Date(2026, 9, 22, 10, 0, 0, 0, time.UTC) + tests := map[string]time.Duration{ + "": 0, + "120": 2 * time.Minute, + "-5": 0, + "soon": 0, + "Tue, 22 Sep 2026 10:01:30 GMT": 90 * time.Second, + "Tue, 22 Sep 2026 09:00:00 GMT": 0, + // One bad reply must not stop the sends for years. + "31536000": time.Hour, + "10000000000": time.Hour, + "20000000000": time.Hour, + "Tue, 22 Sep 2027 10:00:00 GMT": time.Hour, + } + + for v, want := range tests { + if got := retryAfter(v, now); got != want { + t.Errorf("retryAfter(%q) = %v, want %v", v, got, want) + } + } +} + +func TestClientDoesNotFollowRedirects(t *testing.T) { + var targetHits int + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + targetHits++ + _, _ = w.Write([]byte(`{"report_interval": 3}`)) + })) + defer target.Close() + + for _, code := range []int{http.StatusMovedPermanently, http.StatusFound, http.StatusTemporaryRedirect, http.StatusPermanentRedirect} { + c := testClient(t, func(w http.ResponseWriter, r *http.Request) { + // For example a proxy during a deploy, or a redirect to plain http. + http.Redirect(w, r, target.URL+r.URL.Path, code) + }) + + err := c.do(context.Background(), http.MethodPost, "agent/v1/metrics", map[string]int{"samples": 5}, nil) + var statusErr *StatusError + if !errors.As(err, &statusErr) || statusErr.StatusCode != code { + t.Errorf("do() after a %d = %v, want a *StatusError with %d: the samples must stay in the queue", code, err, code) + } + } + if targetHits != 0 { + t.Errorf("the redirect target got %d requests, want none (the token must not go there)", targetHits) + } +} 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/config.go b/internal/agent/config.go new file mode 100644 index 0000000..afcbddf --- /dev/null +++ b/internal/agent/config.go @@ -0,0 +1,168 @@ +// Package agent is the FlyWP monitoring agent: the long-running mode of fly +// that "fly agent run" starts. It follows the FlyWP monitoring agent +// contract v0.5.0. +package agent + +import ( + "errors" + "fmt" + "net" + "net/url" + "os" + "strconv" + "strings" + "time" +) + +// The environment keys that the FlyWP installer writes to /etc/fly/agent.env, +// and the key that systemd sets for StateDirectory=. EnvAutoUpdate is +// optional: a person adds it to stop the updates by release on one server. +const ( + EnvURL = "FLY_AGENT_URL" + EnvToken = "FLY_AGENT_TOKEN" + EnvServerID = "FLY_AGENT_SERVER_ID" + EnvStateDir = "STATE_DIRECTORY" + EnvAutoUpdate = "FLY_AGENT_AUTO_UPDATE" +) + +// Config is the configuration of the agent. +type Config struct { + // URL is the base URL of the control plane. The agent adds the paths + // of the contract to it. + URL *url.URL + // Token is the bearer token of the agent. It is opaque: the agent sends + // it and never parses it. Never log it. + Token string + // ServerID sets the second of the minute at which the agent works. The + // agent never sends it. + ServerID int64 + // StateDir keeps the files that must survive a restart. + StateDir string + // AutoUpdate lets the agent install a new signed release by itself, one + // time each day. The command agent.update works in both cases. + AutoUpdate bool + // Warnings are the problems of optional keys. They do not stop the + // agent: the agent writes them to its log at start. + Warnings []string +} + +// Offset is the time after each full minute at which the agent works. It +// spreads the reports of the fleet over the minute, and it does not change +// between restarts. +func (c Config) Offset() time.Duration { + return time.Duration(c.ServerID%60) * time.Second +} + +// ConfigFromEnv reads the configuration from the environment. The error names +// each key that is not set or not valid. +func ConfigFromEnv(getenv func(string) string) (Config, error) { + var cfg Config + var errs []error + + u, err := parseURL(getenv(EnvURL)) + if err != nil { + errs = append(errs, err) + } + cfg.URL = u + + cfg.Token = getenv(EnvToken) + switch { + case cfg.Token == "": + errs = append(errs, fmt.Errorf("%s is not set", EnvToken)) + case strings.ContainsFunc(cfg.Token, func(r rune) bool { return r < '!' || r > '~' }): + // The token is opaque, but it goes in a header: only printable ASCII + // (the format is flyagt_ and base62) is safe there. + errs = append(errs, fmt.Errorf("%s may contain only printable ASCII characters, without spaces", EnvToken)) + } + + if v := getenv(EnvServerID); v == "" { + errs = append(errs, fmt.Errorf("%s is not set", EnvServerID)) + } else if id, err := strconv.ParseInt(v, 10, 64); err != nil || id < 0 { + errs = append(errs, fmt.Errorf("%s must be an integer of 0 or more, not %q", EnvServerID, v)) + } else { + cfg.ServerID = id + } + + dir, err := stateDir(getenv(EnvStateDir)) + if err != nil { + errs = append(errs, err) + } + cfg.StateDir = dir + + // A typo in an optional key must not stop the metrics. + on, ok := parseSwitch(getenv(EnvAutoUpdate)) + if !ok { + cfg.Warnings = append(cfg.Warnings, fmt.Sprintf("%s=%q is not on or off; auto-update stays on", EnvAutoUpdate, getenv(EnvAutoUpdate))) + } + cfg.AutoUpdate = on + + return cfg, errors.Join(errs...) +} + +// parseSwitch reads an optional on or off value. An empty value is on. ok is +// false for a value that is not known; that value is also on. +func parseSwitch(v string) (on, ok bool) { + switch strings.ToLower(strings.TrimSpace(v)) { + case "", "on", "true", "1", "yes": + return true, true + case "off", "false", "0", "no": + return false, true + default: + return true, false + } +} + +// parseURL accepts an https URL. It also accepts http for a loopback host, +// for tests and local development: plain http never leaves the machine. The +// URL is a base URL only: no user, password, query or fragment. An error +// never shows a password. +func parseURL(v string) (*url.URL, error) { + if v == "" { + return nil, fmt.Errorf("%s is not set", EnvURL) + } + + u, err := url.Parse(strings.TrimRight(v, "/")) + if err != nil || u.Host == "" { + return nil, fmt.Errorf("%s is not a valid URL", EnvURL) + } + + switch { + case u.User != nil: + return nil, fmt.Errorf("%s must not contain a user or a password: %s", EnvURL, u.Redacted()) + case u.RawQuery != "" || u.ForceQuery || u.Fragment != "": + return nil, fmt.Errorf("%s must not contain a query or a fragment: %s", EnvURL, u.Redacted()) + case u.Scheme == "https": + case u.Scheme == "http" && isLoopback(u.Hostname()): + default: + return nil, fmt.Errorf("%s must be an https URL, not %s", EnvURL, u.Redacted()) + } + + return u, nil +} + +func isLoopback(host string) bool { + if strings.EqualFold(host, "localhost") { + return true + } + ip := net.ParseIP(host) + return ip != nil && ip.IsLoopback() +} + +// stateDir returns the first directory of STATE_DIRECTORY. systemd separates +// the directories with ":" when a unit has more than one. +func stateDir(v string) (string, error) { + if v == "" { + return "", fmt.Errorf("%s is not set (systemd sets it for StateDirectory=)", EnvStateDir) + } + + dir, _, _ := strings.Cut(v, ":") + info, err := os.Stat(dir) + if err != nil { + return "", fmt.Errorf("%s: %w", EnvStateDir, err) + } + if !info.IsDir() { + return "", fmt.Errorf("%s is not a directory: %s", EnvStateDir, dir) + } + + return dir, nil +} diff --git a/internal/agent/config_test.go b/internal/agent/config_test.go new file mode 100644 index 0000000..510bcc6 --- /dev/null +++ b/internal/agent/config_test.go @@ -0,0 +1,187 @@ +package agent + +import ( + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +const testToken = "flyagt_0123456789abcdefghijABCDEFGHIJ" + +func validEnv(t *testing.T) map[string]string { + t.Helper() + return map[string]string{ + EnvURL: "https://app.flywp.com", + EnvToken: testToken, + EnvServerID: "59", + EnvStateDir: t.TempDir(), + } +} + +func getenv(env map[string]string) func(string) string { + return func(k string) string { return env[k] } +} + +func TestConfigFromEnv(t *testing.T) { + env := validEnv(t) + env[EnvURL] = "https://app.flywp.com/" + + cfg, err := ConfigFromEnv(getenv(env)) + if err != nil { + t.Fatal(err) + } + + if got := cfg.URL.String(); got != "https://app.flywp.com" { + t.Errorf("URL = %q, want the trailing slash removed", got) + } + if cfg.Token != testToken || cfg.ServerID != 59 || cfg.StateDir != env[EnvStateDir] { + t.Errorf("ConfigFromEnv() = %+v", cfg) + } + if got := cfg.Offset(); got != 59*time.Second { + t.Errorf("Offset() = %v, want 59s", got) + } +} + +func TestConfigOffsetWraps(t *testing.T) { + if got := (Config{ServerID: 125}).Offset(); got != 5*time.Second { + t.Errorf("Offset() = %v, want 5s (125 %% 60)", got) + } +} + +func TestConfigStateDirectoryList(t *testing.T) { + env := validEnv(t) + first := env[EnvStateDir] + env[EnvStateDir] = first + ":" + t.TempDir() + + cfg, err := ConfigFromEnv(getenv(env)) + if err != nil { + t.Fatal(err) + } + if cfg.StateDir != first { + t.Errorf("StateDir = %q, want the first directory %q", cfg.StateDir, first) + } +} + +func TestConfigFromEnvErrors(t *testing.T) { + file := filepath.Join(t.TempDir(), "file") + if err := os.WriteFile(file, nil, 0o600); err != nil { + t.Fatal(err) + } + + tests := []struct { + name, key, value, want string + }{ + {"no URL", EnvURL, "", "FLY_AGENT_URL is not set"}, + {"plain http", EnvURL, "http://app.flywp.com", "must be an https URL"}, + {"other scheme", EnvURL, "ftp://app.flywp.com", "must be an https URL"}, + {"no host", EnvURL, "https://", "not a valid URL"}, + {"no token", EnvToken, "", "FLY_AGENT_TOKEN is not set"}, + {"token with newline", EnvToken, "flyagt_abc\nX-Other: 1", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"token with a space", EnvToken, "flyagt_abc def", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"token that is not ASCII", EnvToken, "flyagt_é", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"token with invalid UTF-8", EnvToken, "flyagt_\x85", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"token with a zero-width space", EnvToken, "flyagt_\u200b", "FLY_AGENT_TOKEN may contain only printable ASCII"}, + {"URL with a query", EnvURL, "https://app.flywp.com?x=1", "must not contain a query"}, + {"URL with a fragment", EnvURL, "https://app.flywp.com#x", "must not contain a query or a fragment"}, + {"no server id", EnvServerID, "", "FLY_AGENT_SERVER_ID is not set"}, + {"negative server id", EnvServerID, "-1", "FLY_AGENT_SERVER_ID must be an integer"}, + {"text server id", EnvServerID, "abc", "FLY_AGENT_SERVER_ID must be an integer"}, + {"no state directory", EnvStateDir, "", "STATE_DIRECTORY is not set"}, + {"missing state directory", EnvStateDir, filepath.Join(t.TempDir(), "missing"), "STATE_DIRECTORY"}, + {"state directory is a file", EnvStateDir, file, "STATE_DIRECTORY is not a directory"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + env := validEnv(t) + env[tt.key] = tt.value + + _, err := ConfigFromEnv(getenv(env)) + if err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("ConfigFromEnv() error = %v, want it to contain %q", err, tt.want) + } + if strings.Contains(err.Error(), testToken) { + t.Errorf("error %q contains the token", err) + } + }) + } +} + +func TestConfigFromEnvNamesEachProblem(t *testing.T) { + _, err := ConfigFromEnv(getenv(map[string]string{})) + if err == nil { + t.Fatal("ConfigFromEnv() = nil, want an error") + } + for _, key := range []string{EnvURL, EnvToken, EnvServerID, EnvStateDir} { + if !strings.Contains(err.Error(), key) { + t.Errorf("error %q does not name %s", err, key) + } + } +} + +func TestConfigAcceptsHTTPSAndLoopbackHTTP(t *testing.T) { + for _, u := range []string{"http://127.0.0.1:8080", "http://[::1]:8080", "http://localhost:8080", "http://LOCALHOST:8080", "HTTPS://app.flywp.com"} { + env := validEnv(t) + env[EnvURL] = u + if _, err := ConfigFromEnv(getenv(env)); err != nil { + t.Errorf("ConfigFromEnv(%s) error = %v, want it accepted", u, err) + } + } +} + +func TestConfigURLWithPasswordDoesNotShowIt(t *testing.T) { + for _, u := range []string{"https://user:s3cret@app.flywp.com", "http://user:s3cret@example.com"} { + env := validEnv(t) + env[EnvURL] = u + + _, err := ConfigFromEnv(getenv(env)) + if err == nil || !strings.Contains(err.Error(), "must not contain a user or a password") { + t.Errorf("ConfigFromEnv(%s) error = %v, want a refusal", u, err) + } + if err != nil && strings.Contains(err.Error(), "s3cret") { + t.Errorf("error %q shows the password", err) + } + } +} + +func TestConfigAutoUpdate(t *testing.T) { + tests := []struct { + value string + want bool + wantWarning bool + }{ + {"", true, false}, + {"on", true, false}, + {"TRUE", true, false}, + {"1", true, false}, + {"yes", true, false}, + {"off", false, false}, + {" Off ", false, false}, + {"false", false, false}, + {"0", false, false}, + {"no", false, false}, + // A typo in an optional key must not stop the metrics. + {"of", true, true}, + {"disabled", true, true}, + } + + for _, tt := range tests { + t.Run(tt.value, func(t *testing.T) { + env := validEnv(t) + env[EnvAutoUpdate] = tt.value + + cfg, err := ConfigFromEnv(getenv(env)) + if err != nil { + t.Fatalf("ConfigFromEnv() error = %v, want the agent to start", err) + } + if cfg.AutoUpdate != tt.want { + t.Errorf("AutoUpdate = %v, want %v", cfg.AutoUpdate, tt.want) + } + if got := len(cfg.Warnings) == 1 && strings.Contains(cfg.Warnings[0], EnvAutoUpdate); got != tt.wantWarning { + t.Errorf("Warnings = %q, want a warning that names %s: %v", cfg.Warnings, EnvAutoUpdate, tt.wantWarning) + } + }) + } +} diff --git a/internal/agent/fakes_test.go b/internal/agent/fakes_test.go new file mode 100644 index 0000000..1f647f7 --- /dev/null +++ b/internal/agent/fakes_test.go @@ -0,0 +1,160 @@ +package agent + +import ( + "context" + "slices" + "sync" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" +) + +// fakeCP is a control plane in memory. It keeps a copy of each request and +// the time of the request, and it answers with the reply funcs. A nil func +// accepts everything. +type fakeCP struct { + mu sync.Mutex + + metrics []wire.MetricsRequest + 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. + latency time.Duration + + // 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) { + if f.latency > 0 { + select { + case <-time.After(f.latency): + case <-ctx.Done(): + return nil, ctx.Err() + } + } + + f.mu.Lock() + defer f.mu.Unlock() + + // Copy the samples: the agent reuses the memory of its queue. + c := *req + c.Samples = slices.Clone(req.Samples) + f.metrics = append(f.metrics, c) + f.metricsAt = append(f.metricsAt, time.Now()) + + if f.metricsReply == nil { + return &wire.MetricsReply{Accepted: len(req.Samples)}, nil + } + return f.metricsReply(len(f.metrics)-1, &c) +} + +func (f *fakeCP) PostEvents(_ context.Context, req *wire.EventsRequest) (*wire.EventsReply, error) { + f.mu.Lock() + defer f.mu.Unlock() + + c := wire.EventsRequest{Events: slices.Clone(req.Events)} + f.events = append(f.events, c) + f.eventsAt = append(f.eventsAt, time.Now()) + + if f.eventsReply == nil { + return &wire.EventsReply{Accepted: len(req.Events)}, nil + } + return f.eventsReply(len(f.events)-1, &c) +} + +// sampleCounts returns the number of samples in each metrics request. +func (f *fakeCP) sampleCounts() []int { + f.mu.Lock() + defer f.mu.Unlock() + + var n []int + for _, r := range f.metrics { + n = append(n, len(r.Samples)) + } + return n +} + +// fakeCollector returns samples whose CPU value counts the samples: 1, 2, 3... +type fakeCollector struct { + mu sync.Mutex + n int + reads []time.Time + status wire.Status + err error + // readTime is the time that each reading takes. + readTime time.Duration +} + +func (c *fakeCollector) Read(now time.Time) { + c.mu.Lock() + c.reads = append(c.reads, now) + d := c.readTime + c.mu.Unlock() + time.Sleep(d) +} + +// readTimes returns the times of the readings between the ticks. +func (c *fakeCollector) readTimes() []time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return slices.Clone(c.reads) +} + +func (c *fakeCollector) Sample(time.Time) (wire.Sample, error) { + c.mu.Lock() + defer c.mu.Unlock() + + if c.err != nil { + return wire.Sample{}, c.err + } + c.n++ + return wire.Sample{CPUPercent: float64(c.n), MemoryTotalBytes: 1 << 30}, nil +} + +func (c *fakeCollector) Status(context.Context) wire.Status { + return c.status +} diff --git a/internal/agent/lock.go b/internal/agent/lock.go new file mode 100644 index 0000000..5978c20 --- /dev/null +++ b/internal/agent/lock.go @@ -0,0 +1,30 @@ +package agent + +import ( + "errors" + "fmt" + "os" + "path/filepath" + "syscall" +) + +// lock takes an exclusive lock on the lock file in dir, so that only one agent +// uses the state files. The lock ends when unlock runs or the process exits. +func lock(dir string) (unlock func(), err error) { + path := filepath.Join(dir, "lock") + f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR, 0o600) + if err != nil { + return nil, err + } + + if err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX|syscall.LOCK_NB); err != nil { + _ = f.Close() + if errors.Is(err, syscall.EWOULDBLOCK) { + return nil, fmt.Errorf("a different agent is running: %s is locked", path) + } + return nil, fmt.Errorf("locking %s: %w", path, err) + } + + // Closing the file releases the lock. + return func() { _ = f.Close() }, nil +} diff --git a/internal/agent/lock_test.go b/internal/agent/lock_test.go new file mode 100644 index 0000000..9aaac86 --- /dev/null +++ b/internal/agent/lock_test.go @@ -0,0 +1,27 @@ +package agent + +import ( + "strings" + "testing" +) + +func TestLockAllowsOneAgent(t *testing.T) { + dir := t.TempDir() + + unlock, err := lock(dir) + if err != nil { + t.Fatal(err) + } + + if _, err := lock(dir); err == nil || !strings.Contains(err.Error(), "a different agent is running") { + t.Fatalf("second lock() error = %v, want a different agent is running", err) + } + + unlock() + + unlock, err = lock(dir) + if err != nil { + t.Fatalf("lock() after unlock error = %v", err) + } + unlock() +} diff --git a/internal/agent/outbox.go b/internal/agent/outbox.go new file mode 100644 index 0000000..a7e4b27 --- /dev/null +++ b/internal/agent/outbox.go @@ -0,0 +1,99 @@ +package agent + +import ( + "errors" + "io/fs" + "log/slog" + "os" + "path/filepath" + "slices" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/statefile" +) + +const ( + // maxSamples is 24 hours of samples (contract section 4). Above it, the + // oldest sample goes first. + maxSamples = 1440 + // maxEvents limits the event queue, for example when a crash loop adds + // events while the control plane is down. + maxEvents = 1000 +) + +// outbox keeps the samples and the events that the control plane has not +// accepted yet. Each change goes to disk, so a restart or a reboot loses +// nothing. +type outbox struct { + dir string + log *slog.Logger + samples []wire.Sample + events []wire.Event +} + +func loadOutbox(dir string, log *slog.Logger) *outbox { + o := &outbox{dir: dir, log: log} + o.samples = loadQueue[wire.Sample](o, "samples.json") + o.events = loadQueue[wire.Event](o, "events.json") + return o +} + +// loadQueue reads a queue file. A file that cannot be read gives an empty +// queue: to stop would only make systemd start the agent again. +func loadQueue[T any](o *outbox, name string) []T { + var q []T + path := filepath.Join(o.dir, name) + err := statefile.Read(path, &q) + if err != nil && !errors.Is(err, fs.ErrNotExist) { + // Keep the file for a person to examine: the next save replaces it. + _ = os.Rename(path, path+".corrupt") + o.log.Error("dropping a queue that cannot be read; the file is kept as "+name+".corrupt", "file", name, "error", err) + return nil + } + + return q +} + +func (o *outbox) addSample(s wire.Sample) { + o.samples = append(o.samples, s) + if over := len(o.samples) - maxSamples; over > 0 { + o.samples = slices.Delete(o.samples, 0, over) + o.log.Warn("the sample queue is full (24 hours); dropping the oldest sample", "dropped", over) + } + o.save("samples.json", o.samples) +} + +func (o *outbox) addEvent(e wire.Event) { + // The control plane records nothing for agent.started, so only the + // last one is useful. + if e.Name == wire.EventAgentStarted { + o.events = slices.DeleteFunc(o.events, func(q wire.Event) bool { return q.Name == wire.EventAgentStarted }) + } + + o.events = append(o.events, e) + if over := len(o.events) - maxEvents; over > 0 { + o.events = slices.Delete(o.events, 0, over) + o.log.Warn("the event queue is full; dropping the oldest event", "dropped", over) + } + o.save("events.json", o.events) +} + +// dropSamples removes the n oldest samples. +func (o *outbox) dropSamples(n int) { + o.samples = slices.Delete(o.samples, 0, n) + o.save("samples.json", o.samples) +} + +// dropEvents removes the n oldest events. +func (o *outbox) dropEvents(n int) { + o.events = slices.Delete(o.events, 0, n) + o.save("events.json", o.events) +} + +// save writes a queue to disk. When the write fails, the queue stays in memory +// and goes to disk with the next change. +func (o *outbox) save(name string, q any) { + if err := statefile.Write(filepath.Join(o.dir, name), q); err != nil { + o.log.Error("saving a queue", "file", name, "error", err) + } +} diff --git a/internal/agent/outbox_test.go b/internal/agent/outbox_test.go new file mode 100644 index 0000000..12eaa6e --- /dev/null +++ b/internal/agent/outbox_test.go @@ -0,0 +1,303 @@ +package agent + +import ( + "encoding/json" + "fmt" + "log/slog" + "math" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/statefile" +) + +func TestOutboxDropsTheOldestSample(t *testing.T) { + dir := t.TempDir() + full := make([]wire.Sample, maxSamples) + for i := range full { + full[i].CPUPercent = float64(i) + } + if err := statefile.Write(filepath.Join(dir, "samples.json"), full); err != nil { + t.Fatal(err) + } + + rec := &recorder{} + o := loadOutbox(dir, slog.New(rec)) + o.addSample(wire.Sample{CPUPercent: maxSamples}) + if len(rec.times("the sample queue is full (24 hours); dropping the oldest sample")) != 1 { + t.Error("want a warning when the full queue drops a sample") + } + + if len(o.samples) != maxSamples { + t.Fatalf("queue holds %d samples, want %d", len(o.samples), maxSamples) + } + if o.samples[0].CPUPercent != 1 { + t.Errorf("oldest sample = %v, want sample 0 dropped", o.samples[0].CPUPercent) + } +} + +func TestOutboxSurvivesARestart(t *testing.T) { + dir := t.TempDir() + o := loadOutbox(dir, slog.New(&recorder{})) + o.addSample(wire.Sample{CPUPercent: 1}) + o.addSample(wire.Sample{CPUPercent: 2}) + o.addEvent(wire.Event{ID: "01JBY0000000000000000000AA", Name: wire.EventCommandCompleted, CommandID: "01JBX0000000000000000000AA"}) + o.dropSamples(1) + + again := loadOutbox(dir, slog.New(&recorder{})) + if len(again.samples) != 1 || again.samples[0].CPUPercent != 2 { + t.Errorf("samples after a restart = %+v, want the second sample only", again.samples) + } + if len(again.events) != 1 || again.events[0].ID != "01JBY0000000000000000000AA" { + t.Errorf("events after a restart = %+v, want the event with its id", again.events) + } +} + +func TestOutboxKeepsOnlyTheLastAgentStarted(t *testing.T) { + o := loadOutbox(t.TempDir(), slog.New(&recorder{})) + o.addEvent(wire.Event{ID: "a", Name: wire.EventAgentStarted}) + o.addEvent(wire.Event{ID: "b", Name: wire.EventCommandCompleted}) + o.addEvent(wire.Event{ID: "c", Name: wire.EventAgentStarted}) + + var ids []string + for _, e := range o.events { + ids = append(ids, e.ID) + } + if strings.Join(ids, ",") != "b,c" { + t.Errorf("events = %v, want b,c: an older agent.started goes", ids) + } +} + +func TestOutboxLimitsEvents(t *testing.T) { + dir := t.TempDir() + full := make([]wire.Event, maxEvents) + for i := range full { + full[i] = wire.Event{ID: fmt.Sprint(i), Name: wire.EventCommandFailed} + } + if err := statefile.Write(filepath.Join(dir, "events.json"), full); err != nil { + t.Fatal(err) + } + + o := loadOutbox(dir, slog.New(&recorder{})) + o.addEvent(wire.Event{ID: "new", Name: wire.EventCommandFailed}) + if len(o.events) != maxEvents || o.events[0].ID != "1" || o.events[maxEvents-1].ID != "new" { + t.Errorf("queue holds %d events from %s to %s, want %d with event 0 dropped", len(o.events), o.events[0].ID, o.events[len(o.events)-1].ID, maxEvents) + } +} + +func TestOutboxWithAQueueThatCannotBeRead(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "samples.json"), []byte("[{"), 0o600); err != nil { + t.Fatal(err) + } + + rec := &recorder{} + o := loadOutbox(dir, slog.New(rec)) + if len(o.samples) != 0 { + t.Errorf("samples = %d, want an empty queue", len(o.samples)) + } + if len(rec.times("dropping a queue that cannot be read; the file is kept as samples.json.corrupt")) != 1 { + t.Error("want an error log line for the queue that cannot be read") + } + if data, err := os.ReadFile(filepath.Join(dir, "samples.json.corrupt")); err != nil || string(data) != "[{" { + t.Errorf("samples.json.corrupt = %q, %v; want the file kept for a person to examine", data, err) + } +} + +func TestCleanSample(t *testing.T) { + recorded := time.Date(2026, 9, 22, 10, 0, 17, 123456789, time.FixedZone("x", 3600)) + tests := []struct { + cpu, load float64 + wantCPU, wantLoad float64 + }{ + {12.5, 0.4, 12.5, 0.4}, + {150, 2e6, 100, maxLoad}, + {-1, -3, 0, 0}, + {math.NaN(), math.NaN(), 0, 0}, + } + + for _, tt := range tests { + got := cleanSample(wire.Sample{RecordedAt: recorded, CPUPercent: tt.cpu, Load1: tt.load}) + if got.CPUPercent != tt.wantCPU || got.Load1 != tt.wantLoad { + t.Errorf("cleanSample(cpu %v, load %v) = %v, %v; want %v, %v", tt.cpu, tt.load, got.CPUPercent, got.Load1, tt.wantCPU, tt.wantLoad) + } + if want := time.Date(2026, 9, 22, 9, 0, 17, 0, time.UTC); !got.RecordedAt.Equal(want) || got.RecordedAt.Location() != time.UTC { + t.Errorf("recorded_at = %s, want %s in UTC", got.RecordedAt, want) + } + } +} + +func TestCleanTextFields(t *testing.T) { + long := strings.Repeat("é", 300) + + s := cleanStatus(wire.Status{OS: long, Kernel: long, Arch: strings.Repeat("x", 20)}) + if n := len([]rune(s.OS)); n != maxStatusTextLen { + t.Errorf("os has %d characters, want %d", n, maxStatusTextLen) + } + if n := len([]rune(s.Kernel)); n != maxStatusTextLen { + t.Errorf("kernel has %d characters, want %d", n, maxStatusTextLen) + } + if len(s.Arch) != maxArchLen { + t.Errorf("arch has %d characters, want %d", len(s.Arch), maxArchLen) + } + + e := cleanEvent(wire.Event{Name: strings.Repeat("n", 70), Data: &wire.EventData{Version: strings.Repeat("v", 40), Error: strings.Repeat("e", 2500)}}) + if len(e.Name) != maxEventNameLen || len(e.Data.Version) != maxVersionLen || len(e.Data.Error) != maxErrorLen { + t.Errorf("event lengths = %d, %d, %d; want %d, %d, %d", len(e.Name), len(e.Data.Version), len(e.Data.Error), maxEventNameLen, maxVersionLen, maxErrorLen) + } +} + +func TestCleanKeepsIntegersInTheRangeOfPHP(t *testing.T) { + s := cleanSample(wire.Sample{NetInBytes: math.MaxUint64, MemoryTotalBytes: 1 << 63, DiskUsedBytes: 42}) + if s.NetInBytes != math.MaxInt64 || s.MemoryTotalBytes != math.MaxInt64 || s.DiskUsedBytes != 42 { + t.Errorf("cleanSample() = %d, %d, %d; want %d, %d, 42", s.NetInBytes, s.MemoryTotalBytes, s.DiskUsedBytes, uint64(math.MaxInt64), uint64(math.MaxInt64)) + } + + total, security := uint64(math.MaxUint64), uint64(1<<63) + st := cleanStatus(wire.Status{UpdatesTotal: &total, UpdatesSecurity: &security, UptimeSeconds: math.MaxUint64}) + if *st.UpdatesTotal != math.MaxInt64 || *st.UpdatesSecurity != math.MaxInt64 || st.UptimeSeconds != math.MaxInt64 { + t.Errorf("cleanStatus() = %+v, want each count at most %d", st, uint64(math.MaxInt64)) + } + + // Unknown counts stay unknown, and go on the wire as null. + st = cleanStatus(wire.Status{}) + data, err := json.Marshal(st) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(data), `"updates_total":null`) || !strings.Contains(string(data), `"updates_security":null`) { + t.Errorf("status JSON = %s, want null update counts", data) + } +} + +func TestCleanPeaks(t *testing.T) { + cpu, big := 120.0, uint64(math.MaxUint64) + s := cleanSample(wire.Sample{CPUMaxPercent: &cpu, MemoryUsedMaxBytes: &big, NetInMaxBytesPerSecond: &big}) + if *s.CPUMaxPercent != 100 || *s.MemoryUsedMaxBytes != math.MaxInt64 || *s.NetInMaxBytesPerSecond != math.MaxInt64 { + t.Errorf("cleanSample() = %v, %d, %d; want 100 and the largest PHP integer", *s.CPUMaxPercent, *s.MemoryUsedMaxBytes, *s.NetInMaxBytesPerSecond) + } + if cpu != 120 { + t.Error("cleanSample() changed the value of the caller") + } + if s.SwapUsedMaxBytes != nil || s.NetOutMaxBytesPerSecond != nil { + t.Error("cleanSample() gave a value to a peak that is not known") + } +} + +// A sample that an older agent queued has no peaks. After an update, the new +// agent sends them as null: not known. +func TestQueuedSampleOfAnOlderAgentSendsNullPeaks(t *testing.T) { + dir := t.TempDir() + old := `[{"recorded_at":"2026-09-24T10:00:17Z","cpu_percent":12.5,"load_1":0.4,"memory_used_bytes":1,"memory_total_bytes":2,` + + `"swap_used_bytes":0,"swap_total_bytes":0,"disk_used_bytes":1,"disk_total_bytes":2,"net_in_bytes":5,"net_out_bytes":6,"net_counters_reset":false}]` + if err := os.WriteFile(filepath.Join(dir, "samples.json"), []byte(old), 0o600); err != nil { + t.Fatal(err) + } + + o := loadOutbox(dir, slog.New(slog.DiscardHandler)) + if len(o.samples) != 1 { + t.Fatalf("samples = %d, want the queued sample", len(o.samples)) + } + data, err := json.Marshal(o.samples[0]) + if err != nil { + t.Fatal(err) + } + for _, field := range []string{"cpu_max_percent", "memory_used_max_bytes", "swap_used_max_bytes", "net_in_max_bytes_per_second", "net_out_max_bytes_per_second"} { + if !strings.Contains(string(data), `"`+field+`":null`) { + t.Errorf("sample JSON = %s, want %s as null", data, field) + } + } +} + +func TestCleanDockerStatus(t *testing.T) { + zero, many, four := 0, 5000, 4 + status, version := "paused", strings.Repeat("9", 40) + s := cleanStatus(wire.Status{CPUCount: &zero, DockerStatus: &status, DockerVersion: &version}) + if s.CPUCount != nil || s.DockerStatus != nil || s.DockerVersion == nil || len(*s.DockerVersion) != maxVersionLen { + t.Errorf("cleanStatus() = %v, %v, %v; want null, null and %d characters", s.CPUCount, s.DockerStatus, s.DockerVersion, maxVersionLen) + } + if s := cleanStatus(wire.Status{CPUCount: &many}); s.CPUCount != nil { + t.Errorf("cpu_count = %d, want null above %d", *s.CPUCount, maxCPUCount) + } + + running := wire.DockerRunning + s = cleanStatus(wire.Status{CPUCount: &four, DockerStatus: &running}) + if s.CPUCount == nil || *s.CPUCount != 4 || s.DockerStatus == nil || *s.DockerStatus != wire.DockerRunning { + t.Errorf("cleanStatus() = %v, %v; want 4 and running", s.CPUCount, s.DockerStatus) + } + + // Not known is null on the wire. + data, err := json.Marshal(cleanStatus(wire.Status{})) + if err != nil { + t.Fatal(err) + } + for _, field := range []string{"cpu_count", "docker_status", "docker_version"} { + if !strings.Contains(string(data), `"`+field+`":null`) { + t.Errorf("status JSON = %s, want %s as null", data, field) + } + } +} + +func TestCleanSites(t *testing.T) { + if s := cleanSample(wire.Sample{}); s.Sites != nil { + t.Errorf("sites = %v, want nil to stay nil (not known)", s.Sites) + } + s := cleanSample(wire.Sample{Sites: []wire.Site{}}) + data, err := json.Marshal(s) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(data), `"sites":[]`) { + t.Errorf("sample JSON = %s, want an empty list to stay empty (no project)", data) + } + + cpu, big := 250.0, uint64(math.MaxUint64) + many := []wire.Site{{Directory: ""}, {Directory: strings.Repeat("d", 256)}, {Directory: "example.com", CPUPercent: &cpu, MemoryUsedBytes: &big}} + for i := range 1200 { + many = append(many, wire.Site{Directory: fmt.Sprintf("site%d.com", i)}) + } + s = cleanSample(wire.Sample{Sites: many}) + if len(s.Sites) != maxSites { + t.Errorf("%d sites, want at most %d", len(s.Sites), maxSites) + } + if got := s.Sites[0]; got.Directory != "example.com" || *got.CPUPercent != 100 || *got.MemoryUsedBytes != math.MaxInt64 { + t.Errorf("first site = %+v, want example.com with its values in range, after the empty and the long directory", got) + } + if cpu != 250 { + t.Error("cleanSample() changed the value of the caller") + } +} + +func TestOutboxKeepsSitesNullAndEmptyApart(t *testing.T) { + dir := t.TempDir() + o := loadOutbox(dir, slog.New(slog.DiscardHandler)) + o.addSample(cleanSample(wire.Sample{Sites: nil})) + o.addSample(cleanSample(wire.Sample{Sites: []wire.Site{}})) + + o = loadOutbox(dir, slog.New(slog.DiscardHandler)) + if len(o.samples) != 2 { + t.Fatalf("samples = %d, want 2", len(o.samples)) + } + for i, want := range []string{`"sites":null`, `"sites":[]`} { + data, err := json.Marshal(o.samples[i]) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(data), want) { + t.Errorf("sample %d JSON = %s, want %s after the queue on disk", i, data, want) + } + } +} + +func TestCleanEventDropsACommandIDThatIsNotAULID(t *testing.T) { + if e := cleanEvent(wire.Event{CommandID: "not-a-ulid"}); e.CommandID != "" { + t.Errorf("command_id = %q, want it removed", e.CommandID) + } + if e := cleanEvent(wire.Event{CommandID: "01JBX0000000000000000000AA"}); e.CommandID != "01JBX0000000000000000000AA" { + t.Errorf("command_id = %q, want the ULID kept", e.CommandID) + } +} diff --git a/internal/agent/send.go b/internal/agent/send.go new file mode 100644 index 0000000..063d4c3 --- /dev/null +++ b/internal/agent/send.go @@ -0,0 +1,225 @@ +package agent + +import ( + "context" + "errors" + "net/http" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/version" +) + +const ( + // sendBudget limits the sends of one report, so that the next tick + // comes on time. + sendBudget = 45 * time.Second + + // The contract limits each request. + maxSamplesPerRequest = 240 + maxEventsPerRequest = 100 + + // unauthorizedWait is the wait after a 401: the token is unknown or + // revoked, and the agent keeps trying (contract section 4). + unauthorizedWait = 5 * time.Minute + // throttledWait is the wait after a 429 without a Retry-After header. + throttledWait = time.Minute + // The wait after other failures doubles from 1 minute up to 10 minutes. + firstBackoff = time.Minute + maxBackoff = 10 * time.Minute +) + +// outcome tells what to do with the data of a request. +type outcome int + +const ( + // sent: drop the data, and send the next request. + sent outcome = iota + // refused: the control plane will never accept the data. Drop it, and + // send the next request. + refused + // later: keep the data, and send no more of it now. + later +) + +// backoff is the wait of one kind of request after a failure. Each kind +// waits on its own: a broken events route must not stop the metrics. +type backoff struct { + // retryAt is the earliest time of the next request, and failures is the + // number of failed requests in a row. failingSince is the start of the + // first report whose request failed and kept its data, or zero. + retryAt time.Time + failures int + failingSince time.Time +} + +// failed records a failure that keeps the data. +func (b *backoff) failed(start time.Time) { + if b.failingSince.IsZero() { + b.failingSince = start + } +} + +// waiting reports whether the request must wait at the time of the report. +func (b *backoff) waiting(at time.Time) bool { + return at.Before(b.retryAt) +} + +// send sends the events and then, with report, the samples. It returns true +// when no event is left in the queue, so that the commands can come next. +func (a *agent) send(ctx context.Context, report bool) (eventsSent bool) { + // The waits count from the start of the report, not from the end of a + // request: a 5 minute wait then ends at the tick 5 minutes later. + start := time.Now() + ctx, cancel := context.WithTimeout(ctx, sendBudget) + defer cancel() + + eventsSent = a.sendEvents(ctx, start) + if report { + a.samplesSent = a.sendSamples(ctx, start) + } + + return eventsSent +} + +// sendEvents sends the queued events, oldest first. It returns true when the +// event queue is empty. +func (a *agent) sendEvents(ctx context.Context, start time.Time) bool { + if a.eventsWait.waiting(start) { + a.log.Debug("waiting before the next events request", "until", a.eventsWait.retryAt) + return len(a.outbox.events) == 0 + } + + for len(a.outbox.events) > 0 { + n := min(len(a.outbox.events), maxEventsPerRequest) + _, err := a.cp.PostEvents(ctx, &wire.EventsRequest{Events: a.outbox.events[:n]}) + if a.outcome(ctx, &a.eventsWait, start, err, "events", n) == later { + return false + } + a.outbox.dropEvents(n) + } + + return true +} + +// sendSamples sends the queued samples, oldest first, with the status of the +// server in each request. It returns true when the sample queue is empty. +func (a *agent) sendSamples(ctx context.Context, start time.Time) bool { + if len(a.outbox.samples) == 0 { + return true + } + if a.metricsWait.waiting(start) { + a.log.Debug("waiting before the next metrics request", "until", a.metricsWait.retryAt) + return false + } + + var status *wire.Status + if a.collector != nil { + s := cleanStatus(a.collector.Status(ctx)) + status = &s + } + + for len(a.outbox.samples) > 0 { + n := min(len(a.outbox.samples), maxSamplesPerRequest) + reply, err := a.cp.PostMetrics(ctx, &wire.MetricsRequest{ + AgentVersion: truncate(version.Version, maxVersionLen), + Status: status, + Samples: a.outbox.samples[:n], + }) + + switch a.outcome(ctx, &a.metricsWait, start, err, "samples", n) { + case later: + return false + case sent: + if len(reply.Rejected) > 0 { + a.log.Warn("the control plane rejected some samples", "rejected", len(reply.Rejected), "first_reason", reply.Rejected[0].Reason) + } + a.setInterval(reply.ReportInterval) + } + a.outbox.dropSamples(n) + } + + return true +} + +// outcome applies the rules of the contract (sections 4 to 6) to the result +// of a request that carried n items of what. The waits go into b and count +// from start. +func (a *agent) outcome(ctx context.Context, b *backoff, start time.Time, err error, what string, n int) outcome { + if err == nil { + b.failures = 0 + // After the warnings of a failure, say that the data goes again. + if !b.failingSince.IsZero() { + a.log.Info("the control plane accepts the requests again", "request", what, "failing_for", start.Sub(b.failingSince).Round(time.Second)) + b.failingSince = time.Time{} + } + return sent + } + + // The time of this report ran out, or the agent stops. That is not a + // failure of the control plane: the data goes with the next report. + if ctx.Err() != nil { + if errors.Is(ctx.Err(), context.DeadlineExceeded) { + a.log.Info("the time for this report is used; the rest goes with the next report", "request", what) + } + return later + } + + var statusErr *StatusError + if errors.As(err, &statusErr) { + switch code := statusErr.StatusCode; { + case code == http.StatusBadRequest: + b.failures, b.failingSince = 0, time.Time{} + a.log.Error("the control plane refused the request as not valid; dropping its data", "request", what, "count", n, "reply", statusErr.Body) + return refused + case code == http.StatusUnprocessableEntity && what == "samples": + b.failures, b.failingSince = 0, time.Time{} + a.log.Warn("the control plane rejected every sample; dropping them", "count", n, "reply", statusErr.Body) + return refused + case code == http.StatusUnauthorized: + b.retryAt = start.Add(unauthorizedWait) + b.failed(start) + a.log.Error("the control plane does not accept the token; keeping the data", "request", what, "retry_in", unauthorizedWait) + return later + case code == http.StatusTooManyRequests: + wait := statusErr.RetryAfter + if wait <= 0 { + wait = throttledWait + } + b.retryAt = start.Add(wait) + b.failed(start) + a.log.Warn("the control plane asks the agent to wait; keeping the data", "request", what, "retry_in", wait) + return later + } + } + + // A 5xx, another status or a network error: keep the data and wait longer + // after each failure. + b.failures++ + wait := min(firstBackoff< 10*time.Second { + t.Errorf("Check() took %s, want it to stop soon after the timeout", elapsed) + } +} diff --git a/internal/dockerapi/dockerapi.go b/internal/dockerapi/dockerapi.go new file mode 100644 index 0000000..97a8ba0 --- /dev/null +++ b/internal/dockerapi/dockerapi.go @@ -0,0 +1,98 @@ +// Package dockerapi reads the Docker Engine API over its unix socket, with the +// Go standard library only. It sends two requests, GET /version and +// GET /containers/json, and never another one: access to the socket is root +// access on the server (FlyWP monitoring agent contract v0.5.0). +package dockerapi + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "time" +) + +// DefaultTimeout limits each request (contract v0.5.0). +const DefaultTimeout = 5 * time.Second + +// maxBody limits the size of an answer. +const maxBody = 16 << 20 + +// Client reads the Docker Engine API. Use New. +type Client struct { + // Timeout limits each request. Tests make it shorter. + Timeout time.Duration + + hc *http.Client +} + +// New returns a client for the socket at path, for example +// /var/run/docker.sock. +func New(path string) *Client { + var d net.Dialer + return &Client{ + Timeout: DefaultTimeout, + hc: &http.Client{ + Transport: &http.Transport{ + DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) { + return d.DialContext(ctx, "unix", path) + }, + MaxIdleConns: 1, + }, + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + }, + } +} + +// Container is a running container. +type Container struct { + ID string `json:"Id"` + Labels map[string]string `json:"Labels"` +} + +// Version returns the version of the Docker Engine, for example 29.7.1. +func (c *Client) Version(ctx context.Context) (string, error) { + var v struct { + Version string `json:"Version"` + } + if err := c.get(ctx, "/version", &v); err != nil { + return "", err + } + return v.Version, nil +} + +// Containers returns the running containers. +func (c *Client) Containers(ctx context.Context) ([]Container, error) { + var cs []Container + if err := c.get(ctx, "/containers/json", &cs); err != nil { + return nil, err + } + return cs, nil +} + +func (c *Client) get(ctx context.Context, path string, out any) error { + ctx, cancel := context.WithTimeout(ctx, c.Timeout) + defer cancel() + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, "http://docker"+path, nil) + if err != nil { + return err + } + resp, err := c.hc.Do(req) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("docker %s: %s", path, resp.Status) + } + if err := json.NewDecoder(io.LimitReader(resp.Body, maxBody)).Decode(out); err != nil { + return fmt.Errorf("docker %s: %w", path, err) + } + return nil +} diff --git a/internal/dockerapi/dockerapi_test.go b/internal/dockerapi/dockerapi_test.go new file mode 100644 index 0000000..a500470 --- /dev/null +++ b/internal/dockerapi/dockerapi_test.go @@ -0,0 +1,57 @@ +package dockerapi + +import ( + "context" + "errors" + "net/http" + "testing" + "time" + + "github.com/flywp/server-cli/internal/testutil" +) + +var engineHandler = testutil.EngineHandler("29.7.1", + `[{"Id":"abc","Names":["/x"],"Labels":{"com.docker.compose.project.working_dir":"/home/fly/example.com"}},{"Id":"def","Labels":{}}]`) + +func TestVersionAndContainers(t *testing.T) { + e := testutil.NewEngine(t, engineHandler) + c := New(e.Socket) + + v, err := c.Version(context.Background()) + if err != nil || v != "29.7.1" { + t.Errorf("Version() = %q, %v; want 29.7.1", v, err) + } + cs, err := c.Containers(context.Background()) + if err != nil || len(cs) != 2 || cs[0].ID != "abc" || cs[0].Labels["com.docker.compose.project.working_dir"] != "/home/fly/example.com" { + t.Errorf("Containers() = %+v, %v", cs, err) + } + if got := e.Requests(); len(got) != 2 || got[0] != "GET /version" || got[1] != "GET /containers/json" { + t.Errorf("requests = %v, want only GET /version and GET /containers/json", got) + } +} + +func TestErrorAnswer(t *testing.T) { + e := testutil.NewEngine(t, func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "boom", http.StatusInternalServerError) + }) + if _, err := New(e.Socket).Version(context.Background()); err == nil { + t.Error("Version() = nil error, want the 500") + } +} + +func TestTimeout(t *testing.T) { + release := make(chan struct{}) + e := testutil.NewEngine(t, func(http.ResponseWriter, *http.Request) { <-release }) + defer close(release) + + c := New(e.Socket) + c.Timeout = 50 * time.Millisecond + start := time.Now() + _, err := c.Version(context.Background()) + if !errors.Is(err, context.DeadlineExceeded) { + t.Errorf("Version() = %v, want a timeout", err) + } + if d := time.Since(start); d > 2*time.Second { + t.Errorf("Version() took %v, want the timeout", d) + } +} diff --git a/internal/metrics/disk.go b/internal/metrics/disk.go new file mode 100644 index 0000000..f5c88f8 --- /dev/null +++ b/internal/metrics/disk.go @@ -0,0 +1,99 @@ +package metrics + +import ( + "os" + "slices" + + "github.com/flywp/server-cli/internal/agent/wire" +) + +// sectorSize is the unit of the sector counts in /proc/diskstats, for each +// disk, whatever the sector size of the disk. +const sectorSize = 512 + +// readDisks returns the counters of the disks that have a hardware device: a +// name in /sys/block with a device link, for example vda, sda or nvme0n1. +// Partitions are not in /sys/block, and loop, ram, zram, device mapper and +// software RAID disks have no device: their activity is already in a hardware +// disk. It returns nil when /proc/diskstats cannot be read. +func (c *Collector) readDisks() map[string]diskCounters { + all, err := parseFile(c, "proc/diskstats", parseDiskstats) + if err != nil { + c.log.Debug("reading the disk activity", "error", err) + return nil + } + + disks := map[string]diskCounters{} + for name, d := range all { + if _, err := os.Lstat(c.file("sys/block", name, "device")); err == nil { + disks[name] = d + } + } + return disks +} + +// diskDelta returns the activity between two readings. It adds the disks +// that both readings have, so a new or a removed disk makes no spike. ok is +// false when the activity is not known: a reading without disks, a reboot, a +// reading older than maxAge, or a counter that went back. +func diskDelta(a, b reading) (d diskCounters, ok bool) { + if len(a.Disks) == 0 || len(b.Disks) == 0 || a.BootID != b.BootID { + return diskCounters{}, false + } + if age := b.At.Sub(a.At); age <= 0 || age > maxAge { + return diskCounters{}, false + } + + for name, cur := range b.Disks { + prev, found := a.Disks[name] + if !found { + continue + } + if cur.ReadOps < prev.ReadOps || cur.ReadSectors < prev.ReadSectors || + cur.WriteOps < prev.WriteOps || cur.WriteSectors < prev.WriteSectors { + return diskCounters{}, false + } + d.ReadOps += cur.ReadOps - prev.ReadOps + d.ReadSectors += cur.ReadSectors - prev.ReadSectors + d.WriteOps += cur.WriteOps - prev.WriteOps + d.WriteSectors += cur.WriteSectors - prev.WriteSectors + } + + return d, true +} + +// setDiskActivity sets the disk activity of the minute from prev to cur in s, +// and the peaks of the windows in all. The eight fields stay nil when the +// activity of the minute is not known, or when a window has a counter that +// went back: then the delta of the minute is not the real activity either. +func setDiskActivity(s *wire.Sample, prev *reading, all []reading, cur reading) { + if prev == nil { + return + } + d, ok := diskDelta(*prev, cur) + if !ok { + return + } + readBytes, writeBytes := d.ReadSectors*sectorSize, d.WriteSectors*sectorSize + + // The control plane reads the value of the minute as the value / 60. + peak := [4]uint64{readBytes / 60, writeBytes / 60, d.ReadOps / 60, d.WriteOps / 60} + // A reading without disks is left out: its windows join. + all = slices.DeleteFunc(slices.Clone(all), func(r reading) bool { return len(r.Disks) == 0 }) + for i := 1; i < len(all); i++ { + a, b := all[i-1], all[i] + w, ok := diskDelta(a, b) + if !ok { + return + } + secs := b.At.Sub(a.At).Seconds() + for j, v := range []uint64{w.ReadSectors * sectorSize, w.WriteSectors * sectorSize, w.ReadOps, w.WriteOps} { + peak[j] = max(peak[j], uint64(float64(v)/secs)) + } + } + + s.DiskReadBytes, s.DiskWriteBytes = &readBytes, &writeBytes + s.DiskReadOps, s.DiskWriteOps = &d.ReadOps, &d.WriteOps + s.DiskReadMaxBytesPerSecond, s.DiskWriteMaxBytesPerSecond = &peak[0], &peak[1] + s.DiskReadMaxOpsPerSecond, s.DiskWriteMaxOpsPerSecond = &peak[2], &peak[3] +} diff --git a/internal/metrics/disk_test.go b/internal/metrics/disk_test.go new file mode 100644 index 0000000..e155d12 --- /dev/null +++ b/internal/metrics/disk_test.go @@ -0,0 +1,334 @@ +package metrics + +import ( + "fmt" + "maps" + "os" + "path/filepath" + "slices" + "strings" + "testing" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" +) + +// A /proc/diskstats of a DigitalOcean server with Ubuntu 24.04: two disks, +// the partitions of vda, and the loop devices of snap. +const diskstats = ` 7 0 loop0 11 0 28 0 0 0 0 0 0 0 0 0 0 0 0 0 0 + 7 1 loop1 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 + 253 0 vda 33376 8239 2465787 22011 165725 77576 2598846 818732 0 202221 889331 2747 0 2299640 1263 22315 47324 + 253 1 vda1 32655 7544 2428229 21737 165698 77555 2598618 818558 0 211841 841559 2747 0 2299640 1263 0 0 + 253 14 vda14 217 0 1984 58 0 0 0 0 0 52 58 0 0 0 0 0 0 + 253 15 vda15 212 661 18510 80 2 0 2 3 0 52 83 0 0 0 0 0 0 + 259 0 vda16 192 34 13224 109 25 21 226 169 0 253 279 0 0 0 0 0 0 + 253 16 vdb 101 3 792 12 0 0 0 0 0 12 12 0 0 0 0 0 0 +` + +// blockDevice marks a disk in /sys/block as a hardware device. +func (s *server) blockDevice(name string) { + s.t.Helper() + if err := os.MkdirAll(filepath.Join(s.root, "sys/block", name, "device"), 0o755); err != nil { + s.t.Fatal(err) + } +} + +// disks writes /proc/diskstats with the counters of each disk, and a loop +// device and a partition with large counters. +func (s *server) disks(disks map[string]diskCounters) { + var b strings.Builder + fmt.Fprintf(&b, " 7 0 loop0 99999 0 99999 0 99999 0 99999 0 0 0 0 0 0 0 0 0 0\n") + for _, name := range slices.Sorted(maps.Keys(disks)) { + d := disks[name] + fmt.Fprintf(&b, " 253 0 %s %d 0 %d 0 %d 0 %d 0 0 0 0 0 0 0 0 0 0\n", name, d.ReadOps, d.ReadSectors, d.WriteOps, d.WriteSectors) + } + fmt.Fprintf(&b, " 253 1 vda1 99999 0 99999 0 99999 0 99999 0 0 0 0 0 0 0 0 0 0\n") + s.write("proc/diskstats", b.String()) +} + +func TestParseDiskstats(t *testing.T) { + all, err := parseDiskstats([]byte(diskstats)) + if err != nil { + t.Fatal(err) + } + want := diskCounters{ReadOps: 33376, ReadSectors: 2465787, WriteOps: 165725, WriteSectors: 2598846} + if all["vda"] != want { + t.Errorf("vda = %+v, want %+v", all["vda"], want) + } + if len(all) != 8 { + t.Errorf("%d lines, want 8", len(all)) + } + // A kernel before 4.18 has 14 columns: no discards and flushes. + all, err = parseDiskstats([]byte(" 8 0 sda 1 2 3 4 5 6 7 8 9 10 11\n")) + if err != nil || all["sda"] != (diskCounters{ReadOps: 1, ReadSectors: 3, WriteOps: 5, WriteSectors: 7}) { + t.Errorf("parseDiskstats(14 columns) = %+v, %v", all, err) + } + if _, err := parseDiskstats([]byte(" 8 0 sda 1 x 3 4 5 6 7 8 9 10 11\n")); err != nil { + t.Errorf("parseDiskstats() = %v, want no error: a column that is not used is not examined", err) + } + if _, err := parseDiskstats([]byte(" 8 0 sda x 2 3 4 5 6 7 8 9 10 11\n")); err == nil { + t.Error("parseDiskstats() = nil error, want an error for a bad count") + } +} + +func TestOnlyHardwareDisksAreCounted(t *testing.T) { + srv := newServer(t) + srv.write("proc/diskstats", diskstats) + srv.blockDevice("vda") + srv.blockDevice("vdb") + // A loop device is in /sys/block, but it has no device. + if err := os.MkdirAll(filepath.Join(srv.root, "sys/block/loop0"), 0o755); err != nil { + t.Fatal(err) + } + + got := slices.Sorted(maps.Keys(srv.collector(t.TempDir()).readDisks())) + if fmt.Sprint(got) != "[vda vdb]" { + t.Errorf("disks = %v, want [vda vdb]", got) + } +} + +// diskMinute takes a sample at base, five readings and the tick at base + +// 60 s, with the counters of vda at each of the seven readings. +func diskMinute(t *testing.T, srv *server, c *Collector, base time.Time, vda [7]diskCounters) wire.Sample { + t.Helper() + srv.disks(map[string]diskCounters{"vda": vda[0]}) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + for i := 1; i < 6; i++ { + srv.disks(map[string]diskCounters{"vda": vda[i]}) + c.Read(base.Add(time.Duration(i) * 10 * time.Second)) + } + srv.disks(map[string]diskCounters{"vda": vda[6]}) + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + return s +} + +func TestDiskActivity(t *testing.T) { + srv := newServer(t) + srv.blockDevice("vda") + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + // Each window reads 10 operations of 20 sectors, and writes 100 + // operations of 200 sectors. The window from 30 s to 40 s writes 5000 + // operations of 100000 sectors. + var vda [7]diskCounters + for i := 1; i < 7; i++ { + vda[i] = vda[i-1] + vda[i].ReadOps += 10 + vda[i].ReadSectors += 20 + if i == 4 { + vda[i].WriteOps += 5000 + vda[i].WriteSectors += 100000 + } else { + vda[i].WriteOps += 100 + vda[i].WriteSectors += 200 + } + } + s := diskMinute(t, srv, c, base, vda) + + for _, f := range []struct { + name string + got *uint64 + want uint64 + }{ + {"disk_read_bytes", s.DiskReadBytes, 120 * 512}, + {"disk_write_bytes", s.DiskWriteBytes, 101000 * 512}, + {"disk_read_ops", s.DiskReadOps, 60}, + {"disk_write_ops", s.DiskWriteOps, 5500}, + {"disk_read_max_bytes_per_second", s.DiskReadMaxBytesPerSecond, 20 * 512 / 10}, + {"disk_write_max_bytes_per_second", s.DiskWriteMaxBytesPerSecond, 100000 * 512 / 10}, + {"disk_read_max_ops_per_second", s.DiskReadMaxOpsPerSecond, 1}, + {"disk_write_max_ops_per_second", s.DiskWriteMaxOpsPerSecond, 500}, + } { + if f.got == nil || *f.got != f.want { + t.Errorf("%s = %v, want %d", f.name, ptr(f.got), f.want) + } + } + checkDiskOrder(t, s) +} + +func TestDiskActivityIsNotKnown(t *testing.T) { + tests := []struct { + name string + change func(*server) + after time.Duration + }{ + {"no hardware disk", func(s *server) { + if err := os.RemoveAll(filepath.Join(s.root, "sys/block")); err != nil { + s.t.Fatal(err) + } + }, time.Minute}, + {"no /proc/diskstats", func(s *server) { + if err := os.Remove(filepath.Join(s.root, "proc/diskstats")); err != nil { + s.t.Fatal(err) + } + }, time.Minute}, + {"reboot", func(s *server) { s.write("proc/sys/kernel/random/boot_id", "boot-2\n") }, time.Minute}, + {"counter went back", func(s *server) { s.disks(map[string]diskCounters{"vda": {ReadOps: 1}}) }, time.Minute}, + {"previous reading too old", func(*server) {}, 3 * time.Minute}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + srv := newServer(t) + srv.blockDevice("vda") + srv.disks(map[string]diskCounters{"vda": {ReadOps: 100, ReadSectors: 100, WriteOps: 100, WriteSectors: 100}}) + c := srv.collector(t.TempDir()) + now := time.Now().Add(time.Second) + if _, err := c.Sample(now); err != nil { + t.Fatal(err) + } + + tt.change(srv) + s, err := c.Sample(now.Add(tt.after)) + if err != nil { + t.Fatal(err) + } + checkNoDiskActivity(t, s) + }) + } +} + +func TestACounterThatWentBackInAWindowMakesTheMinuteNotKnown(t *testing.T) { + srv := newServer(t) + srv.blockDevice("vda") + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + // vda is attached again at 20 s, with the same name: its counter + // starts from 0. At the tick it is above the tick before again, but the + // delta of the minute is not the real activity. + srv.disks(map[string]diskCounters{"vda": {ReadOps: 1000}}) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + c.Read(base.Add(10 * time.Second)) + srv.disks(map[string]diskCounters{"vda": {ReadOps: 5}}) + c.Read(base.Add(20 * time.Second)) + srv.disks(map[string]diskCounters{"vda": {ReadOps: 2000}}) + + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + checkNoDiskActivity(t, s) +} + +func TestFirstSampleHasNoDiskActivity(t *testing.T) { + srv := newServer(t) + srv.blockDevice("vda") + srv.disks(map[string]diskCounters{"vda": {ReadOps: 100}}) + c := srv.collector(t.TempDir()) + srv.disks(map[string]diskCounters{"vda": {ReadOps: 200}}) + + s, err := c.Sample(time.Now().Add(30 * time.Second)) + if err != nil { + t.Fatal(err) + } + checkNoDiskActivity(t, s) +} + +func TestNewDiskMakesNoSpike(t *testing.T) { + srv := newServer(t) + srv.blockDevice("vda") + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + srv.disks(map[string]diskCounters{"vda": {ReadOps: 100}}) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + // A volume is attached at 30 s, with counters from its own start. + srv.blockDevice("sdb") + srv.disks(map[string]diskCounters{"vda": {ReadOps: 130}, "sdb": {ReadOps: 1 << 40}}) + c.Read(base.Add(30 * time.Second)) + srv.disks(map[string]diskCounters{"vda": {ReadOps: 160}, "sdb": {ReadOps: 1<<40 + 30}}) + + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.DiskReadOps == nil || *s.DiskReadOps != 60 { + t.Errorf("disk_read_ops = %v, want 60: the new sdb counts from the next sample", ptr(s.DiskReadOps)) + } + if s.DiskReadMaxOpsPerSecond == nil || *s.DiskReadMaxOpsPerSecond != 2 { + t.Errorf("disk_read_max_ops_per_second = %v, want 2: 30 of vda and 30 of sdb in 30 s", ptr(s.DiskReadMaxOpsPerSecond)) + } + checkDiskOrder(t, s) +} + +func TestAReadingWithoutDisksJoinsTheWindows(t *testing.T) { + srv := newServer(t) + srv.blockDevice("vda") + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + srv.disks(map[string]diskCounters{"vda": {}}) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + // /proc/diskstats cannot be read at 20 s. The other counters can. + if err := os.Remove(filepath.Join(srv.root, "proc/diskstats")); err != nil { + t.Fatal(err) + } + c.Read(base.Add(20 * time.Second)) + srv.disks(map[string]diskCounters{"vda": {WriteOps: 400}}) + c.Read(base.Add(40 * time.Second)) + srv.disks(map[string]diskCounters{"vda": {WriteOps: 600}}) + + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.DiskWriteMaxOpsPerSecond == nil || *s.DiskWriteMaxOpsPerSecond != 10 { + t.Errorf("disk_write_max_ops_per_second = %v, want 10: 400 in 40 s, or 200 in 20 s", ptr(s.DiskWriteMaxOpsPerSecond)) + } + + // Now the busy window is before the reading without disks: 600 in the + // joined window of 40 s is 15 each second. + if err := os.Remove(filepath.Join(srv.root, "proc/diskstats")); err != nil { + t.Fatal(err) + } + c.Read(base.Add(80 * time.Second)) + srv.disks(map[string]diskCounters{"vda": {WriteOps: 1200}}) + c.Read(base.Add(100 * time.Second)) + srv.disks(map[string]diskCounters{"vda": {WriteOps: 1200}}) + s, err = c.Sample(base.Add(2 * time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.DiskWriteMaxOpsPerSecond == nil || *s.DiskWriteMaxOpsPerSecond != 15 { + t.Errorf("disk_write_max_ops_per_second = %v, want 15 from the joined window of 60 s to 100 s", ptr(s.DiskWriteMaxOpsPerSecond)) + } +} + +func checkNoDiskActivity(t *testing.T, s wire.Sample) { + t.Helper() + for _, v := range []*uint64{s.DiskReadBytes, s.DiskWriteBytes, s.DiskReadOps, s.DiskWriteOps, + s.DiskReadMaxBytesPerSecond, s.DiskWriteMaxBytesPerSecond, s.DiskReadMaxOpsPerSecond, s.DiskWriteMaxOpsPerSecond} { + if v != nil { + t.Errorf("a disk activity field = %d, want null", *v) + } + } +} + +// checkDiskOrder checks that a value of the minute is at most 60 times its +// peak, up to rounding. +func checkDiskOrder(t *testing.T, s wire.Sample) { + t.Helper() + for _, p := range [][2]*uint64{ + {s.DiskReadBytes, s.DiskReadMaxBytesPerSecond}, + {s.DiskWriteBytes, s.DiskWriteMaxBytesPerSecond}, + {s.DiskReadOps, s.DiskReadMaxOpsPerSecond}, + {s.DiskWriteOps, s.DiskWriteMaxOpsPerSecond}, + } { + if p[0] == nil || p[1] == nil || *p[0] > 60**p[1]+59 { + t.Errorf("value %v, peak %v: want value ≤ 60 × peak", ptr(p[0]), ptr(p[1])) + } + } +} diff --git a/internal/metrics/diskusage_other.go b/internal/metrics/diskusage_other.go new file mode 100644 index 0000000..569424c --- /dev/null +++ b/internal/metrics/diskusage_other.go @@ -0,0 +1,15 @@ +//go:build !unix + +package metrics + +import ( + "context" + "errors" + "runtime" +) + +func fileID(string) uint64 { return 0 } + +func diskUsage(context.Context, string) (uint64, int, error) { + return 0, 0, errors.New("disk use is not measured on " + runtime.GOOS) +} diff --git a/internal/metrics/diskusage_unix.go b/internal/metrics/diskusage_unix.go new file mode 100644 index 0000000..583128d --- /dev/null +++ b/internal/metrics/diskusage_unix.go @@ -0,0 +1,77 @@ +//go:build unix + +package metrics + +import ( + "context" + "io/fs" + "path/filepath" + "syscall" +) + +// fileID returns the inode of a path, or 0. +func fileID(path string) uint64 { + var st syscall.Stat_t + if err := syscall.Stat(path, &st); err != nil { + return 0 + } + return uint64(st.Ino) +} + +// diskUsage returns the space that the files under root take on the disk: the +// allocated blocks, as du shows them (st_blocks × 512). It does not follow a +// symbolic link, does not go into an other file system, and counts a file +// with more than one hard link one time. It skips what it cannot read, and +// returns the number of those skips. It stops when ctx is done. +func diskUsage(ctx context.Context, root string) (used uint64, skipped int, err error) { + var dev uint64 + seen := map[[2]uint64]bool{} + + err = filepath.WalkDir(root, func(path string, d fs.DirEntry, err error) error { + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + if err != nil { + if path == root { + return err + } + // A folder that cannot be read is skipped. Its own blocks + // were counted before. + skipped++ + return nil + } + + info, err := d.Info() + if err != nil { + skipped++ + return nil + } + st, ok := info.Sys().(*syscall.Stat_t) + if !ok { + skipped++ + return nil + } + + if path == root { + dev = uint64(st.Dev) + } else if uint64(st.Dev) != dev { + // A mount point of an other file system. + if d.IsDir() { + return filepath.SkipDir + } + return nil + } + if !d.IsDir() && st.Nlink > 1 { + key := [2]uint64{uint64(st.Dev), uint64(st.Ino)} + if seen[key] { + return nil + } + seen[key] = true + } + + used += uint64(st.Blocks) * 512 + return nil + }) + + return used, skipped, err +} diff --git a/internal/metrics/diskusage_unix_test.go b/internal/metrics/diskusage_unix_test.go new file mode 100644 index 0000000..115af84 --- /dev/null +++ b/internal/metrics/diskusage_unix_test.go @@ -0,0 +1,97 @@ +//go:build unix + +package metrics + +import ( + "context" + "os" + "path/filepath" + "strings" + "syscall" + "testing" +) + +// blocks returns the allocated bytes of a path, as du counts them. +func blocks(t *testing.T, path string) uint64 { + t.Helper() + var st syscall.Stat_t + if err := syscall.Lstat(path, &st); err != nil { + t.Fatal(err) + } + return uint64(st.Blocks) * 512 +} + +func TestDiskUsage(t *testing.T) { + root := filepath.Join(t.TempDir(), "example.com") + outside := t.TempDir() + for name, size := range map[string]int{"wp-config.php": 3000, "app/index.php": 100, "app/uploads/big.jpg": 200000} { + path := filepath.Join(root, name) + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, []byte(strings.Repeat("x", size)), 0o644); err != nil { + t.Fatal(err) + } + } + // A hard link counts one time. A symbolic link to a big file outside + // the folder is not followed. + if err := os.Link(filepath.Join(root, "app/uploads/big.jpg"), filepath.Join(root, "big-link.jpg")); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(outside, "huge"), []byte(strings.Repeat("y", 1<<20)), 0o644); err != nil { + t.Fatal(err) + } + if err := os.Symlink(filepath.Join(outside, "huge"), filepath.Join(root, "huge-link")); err != nil { + t.Fatal(err) + } + + var want uint64 + for _, p := range []string{"", "app", "app/uploads", "wp-config.php", "app/index.php", "app/uploads/big.jpg", "huge-link"} { + want += blocks(t, filepath.Join(root, p)) + } + + got, skipped, err := diskUsage(context.Background(), root) + if err != nil || skipped != 0 { + t.Fatalf("diskUsage() = %d, %d skipped, %v", got, skipped, err) + } + if got != want { + t.Errorf("diskUsage() = %d, want %d: the blocks of each file one time, without the target of the symbolic link", got, want) + } +} + +func TestDiskUsageSkipsWhatItCannotRead(t *testing.T) { + if os.Getuid() == 0 { + t.Skip("root can read each folder") + } + root := t.TempDir() + locked := filepath.Join(root, "locked") + if err := os.MkdirAll(locked, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(locked, "secret"), []byte(strings.Repeat("s", 100000)), 0o644); err != nil { + t.Fatal(err) + } + if err := os.Chmod(locked, 0); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.Chmod(locked, 0o755) }) + + got, skipped, err := diskUsage(context.Background(), root) + if err != nil { + t.Fatal(err) + } + if skipped != 1 || got != blocks(t, root)+blocks(t, locked) { + t.Errorf("diskUsage() = %d with %d skipped, want the two folders and 1 skip", got, skipped) + } +} + +func TestDiskUsageStops(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, _, err := diskUsage(ctx, t.TempDir()); err == nil { + t.Error("diskUsage() = nil error, want the error of the context") + } + if _, _, err := diskUsage(context.Background(), filepath.Join(t.TempDir(), "gone")); err == nil { + t.Error("diskUsage() = nil error, want an error for a folder that does not exist") + } +} diff --git a/internal/metrics/docker.go b/internal/metrics/docker.go new file mode 100644 index 0000000..3260108 --- /dev/null +++ b/internal/metrics/docker.go @@ -0,0 +1,45 @@ +package metrics + +import ( + "context" + "errors" + "io/fs" + "os" + "syscall" + + "github.com/flywp/server-cli/internal/agent/wire" +) + +// dockerStatus returns the state and the version of Docker, from GET /version +// on the Docker socket (contract v0.5.0): +// +// - an answer: running +// - no socket, or the connection is refused: not_running when dockerd +// exists, else not_installed +// - any other result, for example "permission denied" or no answer in 5 +// seconds: nil, because the agent cannot tell +func (c *Collector) dockerStatus(ctx context.Context) (status, version *string) { + v, err := c.docker.Version(ctx) + switch { + case err == nil: + c.dockerWarned = false + s := wire.DockerRunning + if v == "" { + return &s, nil + } + return &s, &v + case errors.Is(err, fs.ErrNotExist) || errors.Is(err, syscall.ECONNREFUSED): + c.dockerWarned = false + s := wire.DockerNotInstalled + if _, err := os.Stat(c.file("usr/bin/dockerd")); err == nil { + s = wire.DockerNotRunning + } + return &s, nil + default: + if !c.dockerWarned { + c.dockerWarned = true + c.log.Warn("cannot tell whether Docker runs; sending null", "error", err) + } + return nil, nil + } +} diff --git a/internal/metrics/docker_test.go b/internal/metrics/docker_test.go new file mode 100644 index 0000000..8cc6bf0 --- /dev/null +++ b/internal/metrics/docker_test.go @@ -0,0 +1,131 @@ +package metrics + +import ( + "context" + "net" + "net/http" + "os" + "path/filepath" + "testing" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/dockerapi" + "github.com/flywp/server-cli/internal/testutil" +) + +// shortDir returns a folder with a short path, for a unix socket: its path +// has at most 104 bytes on macOS. +func shortDir(t *testing.T) string { + t.Helper() + dir, err := os.MkdirTemp("", "dk") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.RemoveAll(dir) }) + return dir +} + +func TestCPUCount(t *testing.T) { + srv := newServer(t) + srv.write("proc/stat", "cpu 1 2 3 4 5 6 7 8 0 0\ncpu0 1 2 3 4 5 6 7 8 0 0\ncpu1 1 2 3 4 5 6 7 8 0 0\ncpu2 1 2 3 4 5 6 7 8 0 0\ncpu3 1 2 3 4 5 6 7 8 0 0\nintr 1 2 3\nctxt 5\ncpufreq 1\n") + s := srv.collector(t.TempDir()).Status(context.Background()) + if s.CPUCount == nil || *s.CPUCount != 4 { + t.Errorf("cpu_count = %v, want 4", ptr(s.CPUCount)) + } +} + +func TestDockerStatus(t *testing.T) { + engine := testutil.NewEngine(t, testutil.EngineHandler("29.7.1", "[]")) + + refused := filepath.Join(shortDir(t), "docker.sock") + l, err := net.Listen("unix", refused) + if err != nil { + t.Fatal(err) + } + l.(*net.UnixListener).SetUnlinkOnClose(false) + if err := l.Close(); err != nil { + t.Fatal(err) + } + + tests := []struct { + name string + socket string + dockerd bool + wantStatus any + wantVersion any + }{ + {"running", engine.Socket, true, wire.DockerRunning, "29.7.1"}, + {"no socket, with dockerd", "/nonexistent/docker.sock", true, wire.DockerNotRunning, "null"}, + {"no socket, without dockerd", "/nonexistent/docker.sock", false, wire.DockerNotInstalled, "null"}, + {"connection refused", refused, true, wire.DockerNotRunning, "null"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + srv := newServer(t) + if tt.dockerd { + srv.write("usr/bin/dockerd", "") + } + c := srv.collector(t.TempDir()) + c.docker = dockerapi.New(tt.socket) + + s := c.Status(context.Background()) + if ptr(s.DockerStatus) != tt.wantStatus || ptr(s.DockerVersion) != tt.wantVersion { + t.Errorf("docker = %v %v, want %v %v", ptr(s.DockerStatus), ptr(s.DockerVersion), tt.wantStatus, tt.wantVersion) + } + }) + } +} + +func TestDockerStatusIsNotKnown(t *testing.T) { + t.Run("permission denied", func(t *testing.T) { + if os.Getuid() == 0 { + t.Skip("root can use any socket") + } + engine := testutil.NewEngine(t, testutil.EngineHandler("29.7.1", "[]")) + if err := os.Chmod(engine.Socket, 0); err != nil { + t.Fatal(err) + } + srv := newServer(t) + srv.write("usr/bin/dockerd", "") + c := srv.collector(t.TempDir()) + c.docker = dockerapi.New(engine.Socket) + + if s := c.Status(context.Background()); s.DockerStatus != nil || s.DockerVersion != nil { + t.Errorf("docker = %v %v, want null: the agent cannot tell", ptr(s.DockerStatus), ptr(s.DockerVersion)) + } + }) + + t.Run("no answer", func(t *testing.T) { + release := make(chan struct{}) + engine := testutil.NewEngine(t, func(http.ResponseWriter, *http.Request) { <-release }) + defer close(release) + + srv := newServer(t) + c := srv.collector(t.TempDir()) + c.docker = dockerapi.New(engine.Socket) + c.docker.Timeout = 50 * time.Millisecond + + s := c.Status(context.Background()) + if s.DockerStatus != nil || s.DockerVersion != nil { + t.Errorf("docker = %v %v, want null", ptr(s.DockerStatus), ptr(s.DockerVersion)) + } + // The other values are still there. + if s.OS == "" || s.CPUCount == nil { + t.Errorf("status = %+v, want the other values", s) + } + }) +} + +func TestDockerStatusSendsOnlyVersion(t *testing.T) { + engine := testutil.NewEngine(t, testutil.EngineHandler("29.7.1", "[]")) + srv := newServer(t) + c := srv.collector(t.TempDir()) + c.docker = dockerapi.New(engine.Socket) + + c.Status(context.Background()) + if got := engine.Requests(); len(got) != 1 || got[0] != "GET /version" { + t.Errorf("requests = %v, want only GET /version", got) + } +} diff --git a/internal/metrics/metrics.go b/internal/metrics/metrics.go new file mode 100644 index 0000000..7f48fbe --- /dev/null +++ b/internal/metrics/metrics.go @@ -0,0 +1,443 @@ +// Package metrics measures a Linux server for the monitoring agent: CPU, +// load, memory, swap, disk, network, pressure (PSI) and disk activity each +// minute, with the peaks of the minute from a reading each 10 seconds, the +// use of each site, and the status of the server (contract v0.5.0). It needs +// no root. +package metrics + +import ( + "bytes" + "context" + "errors" + "fmt" + "io/fs" + "log/slog" + "os" + "os/exec" + "path/filepath" + "runtime" + "slices" + "time" + + "github.com/flywp/server-cli/internal/agent" + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/dockerapi" + "github.com/flywp/server-cli/internal/statefile" +) + +const ( + // maxAge is the oldest previous reading that gives the traffic of one + // minute. The agent samples each 60 seconds; an older reading covers + // more than one minute. + maxAge = 90 * time.Second + + // The update counts come from apt-check, which takes some seconds. + updatesEvery = time.Hour + aptCheckPath = "/usr/lib/update-notifier/apt-check" + aptCheckTimeout = 30 * time.Second + + // minFirstMinute is the shortest first minute after a start. A shorter + // one is mostly the load of the start, for example an install: its CPU + // value would be a false spike. + minFirstMinute = 10 * time.Second + + // maxReadings limits the readings between two ticks. The agent reads + // five times between two ticks; more readings come only when ticks fail. + maxReadings = 30 +) + +// reading holds the counters at one moment. The reading of the tick is saved, +// so that the first sample after an agent restart continues from it. +type reading struct { + BootID string `json:"boot_id"` + At time.Time `json:"at"` + CPU cpuTimes `json:"cpu"` + Net map[string]netCounters `json:"net"` + // PSI is nil when the kernel has no pressure information. + PSI *psiTotals `json:"psi,omitempty"` + // Disks is nil when /proc/diskstats cannot be read, and empty when no + // disk has a hardware device. + Disks map[string]diskCounters `json:"disks,omitempty"` + // mem is not saved: only the readings of the minute give its peak. + mem memory +} + +// Collector measures the server. Use New. +type Collector struct { + root string + path string // counters.json + log *slog.Logger + + // statfs returns the size and the used space of the file system of a + // path, and release the kernel release. Tests replace them. + statfs func(path string) (total, used uint64, err error) + release func() string + aptCheck func(ctx context.Context) ([]byte, error) + docker *dockerapi.Client + // home is the home folder of the server user, where the Docker Compose + // projects of the sites are. + home string + + // prev is the reading of the last tick, and readings are the readings + // after it, oldest first. They give the windows of the next sample. + prev *reading + readings []reading + // fromStart is true while prev is the reading of the start of this + // process, not a saved reading. minFirst is the shortest first minute; + // tests set it to 0. + fromStart bool + minFirst time.Duration + // noInterface is true after the warning that no interface is counted, + // and noPSI after the warning that the kernel has no PSI. + noInterface bool + noPSI bool + // dockerWarned is true after the warning that the state of Docker is + // not known, until it is known again. + dockerWarned bool + + // prevContainers are the CPU times of the containers at the last tick, + // and walker measures the disk use of the sites. + prevContainers *containerReadings + walker *diskWalker + + updatesAt time.Time + updatesKnown bool + updatesTotal uint64 + updatesSecurity uint64 +} + +// New returns a collector that reads the files under root ("/" on a server) +// and keeps its counters in stateDir. +func New(root, stateDir string, log *slog.Logger) *Collector { + c := &Collector{ + root: root, + path: filepath.Join(stateDir, "counters.json"), + log: log, + statfs: statfs, + release: kernelRelease, + aptCheck: runAptCheck, + minFirst: minFirstMinute, + } + c.docker = dockerapi.New(c.file("var/run/docker.sock")) + c.home = homeDir() + c.walker = newDiskWalker() + + // Take a reading now, so that the first sample has a CPU value for the + // time since the start. The saved reading comes before it only when it is + // recent and from this boot: then the traffic continues without a gap. + now := time.Now() + start, startErr := c.read(now) + + var saved reading + err := statefile.Read(c.path, &saved) + if err != nil && !errors.Is(err, fs.ErrNotExist) { + log.Warn("ignoring the saved counters", "error", err) + } + + switch { + case startErr != nil: + // The first sample has no previous reading. + case err == nil && saved.BootID == start.BootID && saved.At.Before(now) && now.Sub(saved.At) <= maxAge: + c.prev = &saved + c.readings = []reading{start} + default: + // The traffic, the pressure and the disk activity since the start + // are not those of one minute: the first sample sends them as not + // known. + start.Net, start.PSI, start.Disks = nil, nil, nil + c.prev = &start + c.fromStart = true + } + + return c +} + +// Read takes a reading between two ticks, for the peaks of the minute. A +// reading that fails is left out: the windows on each side of it join into +// one (contract v0.4.0). +func (c *Collector) Read(now time.Time) { + if last := c.last(); last != nil && !now.After(last.At) { + return + } + + r, err := c.read(now) + if err != nil { + c.log.Debug("skipping a reading", "error", err) + return + } + if len(c.readings) >= maxReadings { + c.readings = slices.Delete(c.readings, 0, 1) + } + c.readings = append(c.readings, r) +} + +// last returns the newest reading, or nil. +func (c *Collector) last() *reading { + if n := len(c.readings); n > 0 { + return &c.readings[n-1] + } + return c.prev +} + +// Sample measures the minute that ends at now. At the first sample and then +// each hour, it also counts the waiting updates for Status: apt-check takes +// some seconds, and a sample comes before the sends of a report. The count +// comes after the reading of the tick, so that it does not move the reading. +func (c *Collector) Sample(now time.Time) (wire.Sample, error) { + cur, err := c.read(now) + if err != nil { + return wire.Sample{}, err + } + readAt := time.Now() + + // A tick right after the start has no minute to measure. Its reading + // starts the next minute, which then has all its values. + if c.fromStart { + c.fromStart = false + if d := cur.At.Sub(c.prev.At); d >= 0 && d < c.minFirst { + c.prev, c.readings = &cur, nil + // The CPU times of the containers start the next minute too, + // so that its sites have a CPU value. A new process has no + // result of a disk walk to lose. + c.sites(context.Background(), cur, readAt) + if err := statefile.Write(c.path, cur); err != nil { + c.log.Warn("saving the counters", "error", err) + } + return wire.Sample{}, fmt.Errorf("%w: the first minute is only %s since the start", agent.ErrNoSample, d.Round(time.Millisecond)) + } + } + + c.refreshUpdates(context.Background()) + + var s wire.Sample + if s.Load1, err = parseFile(c, "proc/loadavg", parseLoad); err != nil { + return wire.Sample{}, err + } + + mem := cur.mem + s.MemoryTotalBytes = mem.total + s.MemoryUsedBytes = mem.total - min(mem.available, mem.total) + s.SwapTotalBytes = mem.swapTotal + s.SwapUsedBytes = mem.swapTotal - min(mem.swapFree, mem.swapTotal) + + if s.DiskTotalBytes, s.DiskUsedBytes, err = c.statfs(c.file("")); err != nil { + return wire.Sample{}, err + } + + prev := c.prev + if prev != nil && prev.BootID == cur.BootID && cur.CPU.Total >= prev.CPU.Total { + s.CPUPercent = cpuPercent(prev.CPU, cur.CPU) + } + s.NetInBytes, s.NetOutBytes, s.NetCountersReset = netDelta(prev, cur) + if len(cur.Net) == 0 { + // No interface is counted, so the traffic is not known: it is not 0. + s.NetCountersReset = true + if !c.noInterface { + c.noInterface = true + c.log.Warn("no network interface to count: no interface has a hardware device, and the default route has none") + } + } + + setPeaks(&s, c.prev, c.readings, cur) + // The sites come after each step that can fail: they take the result + // of the disk walk, which must not go with a sample that is dropped. + s.Sites = c.sites(context.Background(), cur, readAt) + + c.prev = &cur + c.readings = nil + if err := statefile.Write(c.path, cur); err != nil { + c.log.Warn("saving the counters", "error", err) + } + + return s, nil +} + +// netDelta returns the traffic between two readings. It adds the interfaces +// that both readings have, so a new or a removed interface makes no spike. +// reset is true when the traffic of the minute is not known: no previous +// reading, a reboot, a reading older than maxAge, or a counter that went back. +func netDelta(prev *reading, cur reading) (in, out uint64, reset bool) { + if prev == nil || prev.Net == nil || prev.BootID != cur.BootID { + return 0, 0, true + } + if age := cur.At.Sub(prev.At); age <= 0 || age > maxAge { + return 0, 0, true + } + + for name, c := range cur.Net { + p, ok := prev.Net[name] + if !ok { + continue + } + if c.In < p.In || c.Out < p.Out { + return 0, 0, true + } + in += c.In - p.In + out += c.Out - p.Out + } + + return in, out, false +} + +// read takes the counters now. +func (c *Collector) read(now time.Time) (reading, error) { + cpu, err := parseFile(c, "proc/stat", parseCPU) + if err != nil { + return reading{}, err + } + + mem, err := parseFile(c, "proc/meminfo", parseMeminfo) + if err != nil { + return reading{}, err + } + + all, err := parseFile(c, "proc/net/dev", parseNetDev) + if err != nil { + return reading{}, err + } + + net := map[string]netCounters{} + for _, name := range c.interfaces(all) { + net[name] = all[name] + } + + bootID, err := os.ReadFile(c.file("proc/sys/kernel/random/boot_id")) + if err != nil { + return reading{}, err + } + + return reading{BootID: string(bytes.TrimSpace(bootID)), At: now, CPU: cpu, Net: net, PSI: c.readPSI(), Disks: c.readDisks(), mem: mem}, nil +} + +// interfaces returns the network interfaces that have a hardware device and +// are not a port of an other interface. Thus lo, docker0, the Docker bridges +// and the veth interfaces are left out, and container traffic is not counted +// two or three times. A port (of a bond or a bridge, or the Azure VF under +// its netvsc interface) is left out too, because its traffic is also in the +// interface above it. If no interface is left, it returns the interface of +// the default route, for example the bond or the bridge. +func (c *Collector) interfaces(all map[string]netCounters) []string { + var names []string + for name := range all { + if _, err := os.Lstat(c.file("sys/class/net", name, "device")); err != nil { + continue + } + if _, err := os.Lstat(c.file("sys/class/net", name, "master")); err == nil { + continue + } + names = append(names, name) + } + if len(names) > 0 { + return names + } + + route, err := os.ReadFile(c.file("proc/net/route")) + if err != nil { + return nil + } + if name := parseDefaultRoute(route); name != "" { + if _, ok := all[name]; ok { + return []string{name} + } + } + + return nil +} + +// Status describes the server now. A value that cannot be read stays empty, +// 0 or nil, and the problem goes to the log. The update counts come from the +// last Sample. +func (c *Collector) Status(ctx context.Context) wire.Status { + s := wire.Status{Arch: runtime.GOARCH, Kernel: c.release()} + + if _, err := os.Stat(c.file("var/run/reboot-required")); err == nil { + s.RebootRequired = true + } + + if data, err := os.ReadFile(c.file("etc/os-release")); err == nil { + s.OS = parseOSRelease(data) + } else { + c.log.Warn("reading the OS name", "error", err) + } + + if up, err := parseFile(c, "proc/uptime", parseUptime); err == nil { + s.UptimeSeconds = up + } else { + c.log.Warn("reading the uptime", "error", err) + } + + if n, err := parseFile(c, "proc/stat", parseCPUCount); err == nil { + s.CPUCount = &n + } else { + c.log.Warn("counting the CPUs", "error", err) + } + + s.DockerStatus, s.DockerVersion = c.dockerStatus(ctx) + + // Without any count, the counts are not known: null, not a false 0. + if c.updatesKnown { + total, security := c.updatesTotal, c.updatesSecurity + s.UpdatesTotal, s.UpdatesSecurity = &total, &security + } + + return s +} + +// refreshUpdates counts the waiting updates at the first call and then each +// hour. If apt-check fails, the last counts stay. Without any count, Status +// sends null (contract v0.3.1). +func (c *Collector) refreshUpdates(ctx context.Context) { + if !c.updatesAt.IsZero() && time.Since(c.updatesAt) < updatesEvery { + return + } + c.updatesAt = time.Now() + + ctx, cancel := context.WithTimeout(ctx, aptCheckTimeout) + defer cancel() + + out, err := c.aptCheck(ctx) + var total, security uint64 + if err == nil { + total, security, err = parseAptCheck(out) + } + if err != nil { + if c.updatesKnown { + c.log.Warn("cannot count the waiting updates; keeping the last counts", "error", err) + } else { + c.log.Warn("cannot count the waiting updates; sending null", "error", err) + } + return + } + + c.updatesKnown = true + c.updatesTotal, c.updatesSecurity = total, security +} + +// runAptCheck runs apt-check. It writes its result to stderr. +func runAptCheck(ctx context.Context) ([]byte, error) { + var stderr bytes.Buffer + cmd := exec.CommandContext(ctx, aptCheckPath) + cmd.Stderr = &stderr + cmd.WaitDelay = time.Second + if err := cmd.Run(); err != nil { + return nil, err + } + + return stderr.Bytes(), nil +} + +// file returns the path of a file under the root. +func (c *Collector) file(parts ...string) string { + return filepath.Join(append([]string{c.root}, parts...)...) +} + +// parseFile reads a file under the root and parses it. +func parseFile[T any](c *Collector, name string, parse func([]byte) (T, error)) (T, error) { + data, err := os.ReadFile(c.file(name)) + if err != nil { + var zero T + return zero, err + } + + return parse(data) +} diff --git a/internal/metrics/metrics_linux_test.go b/internal/metrics/metrics_linux_test.go new file mode 100644 index 0000000..8fd5bb1 --- /dev/null +++ b/internal/metrics/metrics_linux_test.go @@ -0,0 +1,36 @@ +package metrics + +import ( + "context" + "log/slog" + "runtime" + "testing" + "time" +) + +// TestRealServer measures this Linux machine: CI runs it on ubuntu-latest. +func TestRealServer(t *testing.T) { + c := New("/", t.TempDir(), slog.New(slog.DiscardHandler)) + // The sample comes right after the start: measure it anyway. + c.minFirst = 0 + + s, err := c.Sample(time.Now()) + if err != nil { + t.Fatal(err) + } + + if s.CPUPercent < 0 || s.CPUPercent > 100 { + t.Errorf("cpu_percent = %v, want 0 to 100", s.CPUPercent) + } + if s.MemoryTotalBytes == 0 || s.MemoryUsedBytes > s.MemoryTotalBytes { + t.Errorf("memory = %d of %d", s.MemoryUsedBytes, s.MemoryTotalBytes) + } + if s.DiskTotalBytes == 0 || s.DiskUsedBytes > s.DiskTotalBytes { + t.Errorf("disk = %d of %d", s.DiskUsedBytes, s.DiskTotalBytes) + } + + st := c.Status(context.Background()) + if st.Kernel == "" || st.Arch != runtime.GOARCH || st.UptimeSeconds == 0 { + t.Errorf("status = %+v, want the kernel, the arch and the uptime", st) + } +} diff --git a/internal/metrics/metrics_test.go b/internal/metrics/metrics_test.go new file mode 100644 index 0000000..41163cc --- /dev/null +++ b/internal/metrics/metrics_test.go @@ -0,0 +1,484 @@ +package metrics + +import ( + "context" + "errors" + "fmt" + "log/slog" + "os" + "path/filepath" + "runtime" + "sort" + "strings" + "testing" + "time" + + "github.com/flywp/server-cli/internal/agent" + "github.com/flywp/server-cli/internal/agent/wire" +) + +// server is a fake file system root with the files that the collector reads. +type server struct { + t *testing.T + root string +} + +func newServer(t *testing.T) *server { + t.Helper() + + s := &server{t: t, root: t.TempDir()} + s.write("proc/loadavg", "0.42 0.30 0.25 1/345 6789\n") + s.write("proc/meminfo", "MemTotal: 8000000 kB\nMemFree: 500000 kB\nMemAvailable: 6000000 kB\nSwapTotal: 2000000 kB\nSwapFree: 1500000 kB\n") + s.write("proc/uptime", "1892344.51 3700000.00\n") + s.write("proc/sys/kernel/random/boot_id", "boot-1\n") + s.write("proc/net/route", "Iface\tDestination\tGateway \tFlags\tRefCnt\tUse\tMetric\tMask\t\tMTU\tWindow\tIRTT\n"+ + "ens3\t00000000\t0101A8C0\t0003\t0\t0\t100\t00000000\t0\t0\t0\n"+ + "ens3\t0001A8C0\t00000000\t0001\t0\t0\t100\t00FFFFFF\t0\t0\t0\n") + s.write("etc/os-release", "NAME=\"Ubuntu\"\nPRETTY_NAME=\"Ubuntu 24.04.1 LTS\"\nID=ubuntu\n") + s.device("eth0") + s.device("eth1") + s.cpu(1000, 800) + s.net(map[string][2]uint64{"eth0": {1000, 500}, "eth1": {100, 50}, "lo": {9999, 9999}, "docker0": {7000, 7000}, "veth1": {7000, 7000}}) + return s +} + +func (s *server) write(name, content string) { + s.t.Helper() + path := filepath.Join(s.root, name) + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + s.t.Fatal(err) + } + if err := os.WriteFile(path, []byte(content), 0o644); err != nil { + s.t.Fatal(err) + } +} + +// device marks an interface as a hardware device. +func (s *server) device(name string) { + s.t.Helper() + if err := os.MkdirAll(filepath.Join(s.root, "sys/class/net", name, "device"), 0o755); err != nil { + s.t.Fatal(err) + } +} + +// cpu writes /proc/stat with the total and the idle ticks. +func (s *server) cpu(total, idle uint64) { + // user nice system idle iowait irq softirq steal guest guest_nice + busy := total - idle + s.write("proc/stat", fmt.Sprintf("cpu %d 0 0 %d 0 0 0 0 55 0\ncpu0 1 2 3 4 5 6 7 8 9 10\n", busy, idle)) +} + +// net writes /proc/net/dev with the received and sent bytes of each interface. +func (s *server) net(ifaces map[string][2]uint64) { + var b strings.Builder + b.WriteString("Inter-| Receive | Transmit\n") + b.WriteString(" face |bytes packets errs drop fifo frame compressed multicast|bytes packets errs drop fifo colls carrier compressed\n") + for name, c := range ifaces { + fmt.Fprintf(&b, "%6s: %d 10 0 0 0 0 0 0 %d 10 0 0 0 0 0 0\n", name, c[0], c[1]) + } + s.write("proc/net/dev", b.String()) +} + +func (s *server) collector(stateDir string) *Collector { + c := New(s.root, stateDir, slog.New(slog.DiscardHandler)) + c.statfs = func(string) (uint64, uint64, error) { return 100 << 30, 25 << 30, nil } + c.release = func() string { return "6.8.0-45-generic" } + c.aptCheck = func(context.Context) ([]byte, error) { return []byte("33;6"), nil } + // The tests take the first sample some milliseconds after the start. + c.minFirst = 0 + return c +} + +func TestSample(t *testing.T) { + srv := newServer(t) + now := time.Now() + c := srv.collector(t.TempDir()) + if _, err := c.Sample(now); err != nil { + t.Fatal(err) + } + + // One minute later: 600 more ticks with 150 idle, and some traffic. + srv.cpu(1600, 950) + srv.net(map[string][2]uint64{"eth0": {3000, 1500}, "eth1": {600, 150}, "lo": {99999, 99999}, "docker0": {70000, 70000}, "veth1": {70000, 70000}}) + + s, err := c.Sample(now.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + + if s.CPUPercent != 75 { + t.Errorf("cpu_percent = %v, want 75 (450 of 600 ticks busy)", s.CPUPercent) + } + if s.Load1 != 0.42 { + t.Errorf("load_1 = %v, want 0.42", s.Load1) + } + // MemTotal − MemAvailable, not MemTotal − MemFree. + if s.MemoryTotalBytes != 8000000*1024 || s.MemoryUsedBytes != 2000000*1024 { + t.Errorf("memory = %d of %d, want %d of %d", s.MemoryUsedBytes, s.MemoryTotalBytes, 2000000*1024, 8000000*1024) + } + if s.SwapTotalBytes != 2000000*1024 || s.SwapUsedBytes != 500000*1024 { + t.Errorf("swap = %d of %d, want %d of %d", s.SwapUsedBytes, s.SwapTotalBytes, 500000*1024, 2000000*1024) + } + if s.DiskTotalBytes != 100<<30 || s.DiskUsedBytes != 25<<30 { + t.Errorf("disk = %d of %d, want the statfs values", s.DiskUsedBytes, s.DiskTotalBytes) + } + // Only eth0 and eth1 have a device: lo, docker0 and veth1 are not counted. + if s.NetInBytes != 2000+500 || s.NetOutBytes != 1000+100 || s.NetCountersReset { + t.Errorf("net = in %d, out %d, reset %v; want in 2500, out 1100 from eth0 and eth1", s.NetInBytes, s.NetOutBytes, s.NetCountersReset) + } +} + +func TestFirstSampleHasNoTraffic(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + + s, err := c.Sample(time.Now()) + if err != nil { + t.Fatal(err) + } + if !s.NetCountersReset || s.NetInBytes != 0 || s.NetOutBytes != 0 { + t.Errorf("first sample net = %d, %d, reset %v; want 0, 0 and a reset", s.NetInBytes, s.NetOutBytes, s.NetCountersReset) + } +} + +func TestRestartContinuesFromTheSavedCounters(t *testing.T) { + srv := newServer(t) + state := t.TempDir() + now := time.Now() + if _, err := srv.collector(state).Sample(now); err != nil { + t.Fatal(err) + } + + // A new agent process starts, and one minute after the last sample it + // takes the next one. + srv.net(map[string][2]uint64{"eth0": {1500, 700}, "eth1": {100, 50}}) + s, err := srv.collector(state).Sample(now.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.NetCountersReset || s.NetInBytes != 500 || s.NetOutBytes != 200 { + t.Errorf("net after a restart = %d, %d, reset %v; want 500, 200 without a reset", s.NetInBytes, s.NetOutBytes, s.NetCountersReset) + } +} + +func TestTrafficResets(t *testing.T) { + tests := []struct { + name string + change func(*server) + after time.Duration + }{ + {"reboot", func(s *server) { s.write("proc/sys/kernel/random/boot_id", "boot-2\n") }, time.Minute}, + {"counter went back", func(s *server) { s.net(map[string][2]uint64{"eth0": {10, 10}, "eth1": {100, 50}}) }, time.Minute}, + {"previous reading too old", func(*server) {}, 3 * time.Minute}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + srv := newServer(t) + now := time.Now() + c := srv.collector(t.TempDir()) + if _, err := c.Sample(now); err != nil { + t.Fatal(err) + } + + tt.change(srv) + s, err := c.Sample(now.Add(tt.after)) + if err != nil { + t.Fatal(err) + } + if !s.NetCountersReset || s.NetInBytes != 0 || s.NetOutBytes != 0 { + t.Errorf("net = %d, %d, reset %v; want 0, 0 and a reset", s.NetInBytes, s.NetOutBytes, s.NetCountersReset) + } + }) + } +} + +func TestStaleSavedCountersAreIgnored(t *testing.T) { + srv := newServer(t) + state := t.TempDir() + if _, err := srv.collector(state).Sample(time.Now().Add(-10 * time.Minute)); err != nil { + t.Fatal(err) + } + + // The agent was stopped for 10 minutes: its saved traffic counters do not + // give the traffic of one minute. + srv.net(map[string][2]uint64{"eth0": {900000, 900000}, "eth1": {100, 50}}) + s, err := srv.collector(state).Sample(time.Now()) + if err != nil { + t.Fatal(err) + } + if !s.NetCountersReset || s.NetInBytes != 0 { + t.Errorf("net = %d, reset %v; want 0 and a reset", s.NetInBytes, s.NetCountersReset) + } +} + +func TestNewInterfaceMakesNoSpike(t *testing.T) { + srv := newServer(t) + now := time.Now() + c := srv.collector(t.TempDir()) + if _, err := c.Sample(now); err != nil { + t.Fatal(err) + } + + srv.device("eth2") + srv.net(map[string][2]uint64{"eth0": {1100, 600}, "eth1": {100, 50}, "eth2": {5 << 40, 5 << 40}}) + s, err := c.Sample(now.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.NetInBytes != 100 || s.NetOutBytes != 100 { + t.Errorf("net = %d, %d; want 100, 100 (the new eth2 counts from the next sample)", s.NetInBytes, s.NetOutBytes) + } +} + +func TestInterfacesFallBackToTheDefaultRoute(t *testing.T) { + srv := newServer(t) + if err := os.RemoveAll(filepath.Join(srv.root, "sys")); err != nil { + t.Fatal(err) + } + srv.net(map[string][2]uint64{"ens3": {1, 1}, "lo": {1, 1}, "docker0": {1, 1}}) + + all, err := parseFile(srv.collector(t.TempDir()), "proc/net/dev", parseNetDev) + if err != nil { + t.Fatal(err) + } + if got := srv.collector(t.TempDir()).interfaces(all); fmt.Sprint(got) != "[ens3]" { + t.Errorf("interfaces() = %v, want [ens3] from the default route", got) + } +} + +func TestSampleErrors(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + if err := os.Remove(filepath.Join(srv.root, "proc/meminfo")); err != nil { + t.Fatal(err) + } + if _, err := c.Sample(time.Now()); err == nil { + t.Error("Sample() = nil error, want an error without /proc/meminfo") + } + + srv = newServer(t) + c = srv.collector(t.TempDir()) + c.statfs = func(string) (uint64, uint64, error) { return 0, 0, errors.New("statfs failed") } + if _, err := c.Sample(time.Now()); err == nil { + t.Error("Sample() = nil error, want the statfs error") + } +} + +func TestStatus(t *testing.T) { + srv := newServer(t) + srv.write("var/run/reboot-required", "*** System restart required ***\n") + c := srv.collector(t.TempDir()) + + calls := 0 + c.aptCheck = func(context.Context) ([]byte, error) { + calls++ + return []byte("33;6"), nil + } + + // The count runs with the first sample, before the sends of a report. + if _, err := c.Sample(time.Now()); err != nil { + t.Fatal(err) + } + s := c.Status(context.Background()) + if !s.RebootRequired || !counts(s, 33, 6) { + t.Errorf("status = %+v, want a restart and 33 updates with 6 security updates", s) + } + if s.OS != "Ubuntu 24.04.1 LTS" || s.Kernel != "6.8.0-45-generic" || s.UptimeSeconds != 1892344 || s.Arch != runtime.GOARCH { + t.Errorf("status = %+v", s) + } + + // The counts stay for one hour: apt-check takes some seconds. + c.Status(context.Background()) + if _, err := c.Sample(time.Now()); err != nil { + t.Fatal(err) + } + if calls != 1 { + t.Errorf("apt-check ran %d times, want 1 in one hour", calls) + } + + c.updatesAt = time.Now().Add(-2 * time.Hour) + if _, err := c.Sample(time.Now()); err != nil { + t.Fatal(err) + } + if calls != 2 { + t.Errorf("apt-check ran %d times, want again after one hour", calls) + } +} + +func TestStatusWithoutAptCheck(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + c.aptCheck = func(context.Context) ([]byte, error) { return nil, os.ErrNotExist } + + if _, err := c.Sample(time.Now()); err != nil { + t.Fatal(err) + } + s := c.Status(context.Background()) + // Not known is null, not a false "no updates". + if s.UpdatesTotal != nil || s.UpdatesSecurity != nil || s.RebootRequired { + t.Errorf("status = %+v, want unknown (nil) updates and no restart", s) + } + if s.OS == "" { + t.Error("status has no OS: one missing value must not clear the others") + } +} + +func TestAptCheckWithWarningsBeforeTheResult(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + // apt-check on Ubuntu 24.04 with a source that is configured two times. + c.aptCheck = func(context.Context) ([]byte, error) { + return []byte("/usr/lib/update-notifier/apt-check:351: Warning: W:Target Packages (main/binary-amd64/Packages) is configured multiple times in /etc/apt/sources.list:1 and /etc/apt/sources.list.d/ubuntu.sources:1\n" + + " apt_pkg.init()\n29;26"), nil + } + + if _, err := c.Sample(time.Now()); err != nil { + t.Fatal(err) + } + if s := c.Status(context.Background()); !counts(s, 29, 26) { + t.Errorf("updates = %v;%v, want 29;26 from the last line", s.UpdatesTotal, s.UpdatesSecurity) + } +} + +func TestAFailedCountKeepsTheLastCounts(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + if _, err := c.Sample(time.Now()); err != nil { + t.Fatal(err) + } + + // One hour later apt-check fails, for example with a timeout. + c.aptCheck = func(context.Context) ([]byte, error) { return nil, context.DeadlineExceeded } + c.updatesAt = time.Now().Add(-2 * time.Hour) + if _, err := c.Sample(time.Now()); err != nil { + t.Fatal(err) + } + if s := c.Status(context.Background()); !counts(s, 33, 6) { + t.Errorf("updates = %v;%v, want the last counts 33;6", s.UpdatesTotal, s.UpdatesSecurity) + } +} + +// counts reports whether s has the known update counts total and security. +func counts(s wire.Status, total, security uint64) bool { + return s.UpdatesTotal != nil && s.UpdatesSecurity != nil && *s.UpdatesTotal == total && *s.UpdatesSecurity == security +} + +func TestNoCountedInterfaceIsNotZeroTraffic(t *testing.T) { + srv := newServer(t) + if err := os.RemoveAll(filepath.Join(srv.root, "sys")); err != nil { + t.Fatal(err) + } + // An IPv6-only host: no IPv4 default route. + srv.write("proc/net/route", "Iface\tDestination\tGateway \tFlags\tRefCnt\tUse\tMetric\tMask\t\tMTU\tWindow\tIRTT\n") + now := time.Now() + c := srv.collector(t.TempDir()) + if _, err := c.Sample(now); err != nil { + t.Fatal(err) + } + + s, err := c.Sample(now.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if !s.NetCountersReset { + t.Error("net_counters_reset = false, want true: the traffic is not known") + } +} + +func TestPortsOfAnOtherInterfaceAreNotCounted(t *testing.T) { + srv := newServer(t) + // Azure accelerated networking: the VF has a device, but its traffic is + // also in eth0, its master. A bond port is the same. + srv.device("enP1s1") + if err := os.Symlink("../eth0", filepath.Join(srv.root, "sys/class/net/enP1s1/master")); err != nil { + t.Fatal(err) + } + srv.net(map[string][2]uint64{"eth0": {1000, 500}, "eth1": {100, 50}, "enP1s1": {900, 400}}) + + all, err := parseFile(srv.collector(t.TempDir()), "proc/net/dev", parseNetDev) + if err != nil { + t.Fatal(err) + } + got := srv.collector(t.TempDir()).interfaces(all) + sort.Strings(got) + if fmt.Sprint(got) != "[eth0 eth1]" { + t.Errorf("interfaces() = %v, want [eth0 eth1] without the port enP1s1", got) + } +} + +func TestNewWithCountersThatCannotBeRead(t *testing.T) { + srv := newServer(t) + state := t.TempDir() + if err := os.WriteFile(filepath.Join(state, "counters.json"), []byte("{"), 0o600); err != nil { + t.Fatal(err) + } + + s, err := srv.collector(state).Sample(time.Now()) + if err != nil { + t.Fatal(err) + } + if !s.NetCountersReset { + t.Error("net_counters_reset = false, want true after counters that cannot be read") + } +} + +func TestAShortFirstMinuteHasNoSample(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + c.minFirst = minFirstMinute + + // The install ends 1 s before the tick: that second is all busy. + srv.cpu(1100, 800) + first := time.Now().Add(time.Second) + if _, err := c.Sample(first); !errors.Is(err, agent.ErrNoSample) { + t.Fatalf("Sample() 1 s after the start = %v, want ErrNoSample", err) + } + + // The next minute is a whole minute from the tick reading, with its + // traffic. + srv.cpu(1700, 1250) + srv.net(map[string][2]uint64{"eth0": {1600, 800}, "eth1": {100, 50}}) + s, err := c.Sample(first.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.CPUPercent != 25 { + t.Errorf("cpu_percent = %v, want 25 from the tick reading, not from the start", s.CPUPercent) + } + if s.NetCountersReset || s.NetInBytes != 600 || s.NetInMaxBytesPerSecond == nil { + t.Errorf("net = %d, peak %v, reset %v; want 600 and a peak: the minute has a previous reading", s.NetInBytes, ptr(s.NetInMaxBytesPerSecond), s.NetCountersReset) + } +} + +func TestALongFirstMinuteHasASample(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + c.minFirst = minFirstMinute + + srv.cpu(1600, 950) + s, err := c.Sample(time.Now().Add(20 * time.Second)) + if err != nil { + t.Fatalf("Sample() 20 s after the start = %v, want a sample", err) + } + if s.CPUPercent != 75 { + t.Errorf("cpu_percent = %v, want 75", s.CPUPercent) + } +} + +func TestARestartWithSavedCountersHasASample(t *testing.T) { + srv := newServer(t) + state := t.TempDir() + now := time.Now() + if _, err := srv.collector(state).Sample(now.Add(-59 * time.Second)); err != nil { + t.Fatal(err) + } + + // A quick restart, for example an update, 1 s before the tick: the + // minute continues from the saved reading. + c := srv.collector(state) + c.minFirst = minFirstMinute + if _, err := c.Sample(now.Add(time.Second)); err != nil { + t.Errorf("Sample() after a restart with saved counters = %v, want a sample", err) + } +} diff --git a/internal/metrics/parse.go b/internal/metrics/parse.go new file mode 100644 index 0000000..2c64d54 --- /dev/null +++ b/internal/metrics/parse.go @@ -0,0 +1,286 @@ +package metrics + +import ( + "bufio" + "bytes" + "fmt" + "strconv" + "strings" +) + +// cpuTimes are the CPU counters of /proc/stat, in clock ticks. +type cpuTimes struct { + Idle uint64 `json:"idle"` + Total uint64 `json:"total"` +} + +// parseCPU reads the first line of /proc/stat: +// +// cpu user nice system idle iowait irq softirq steal guest guest_nice +// +// Idle is idle + iowait. Total leaves out guest and guest_nice, because user +// and nice already hold them. +func parseCPU(data []byte) (cpuTimes, error) { + line, _, _ := bytes.Cut(data, []byte("\n")) + fields := strings.Fields(string(line)) + if len(fields) < 5 || fields[0] != "cpu" { + return cpuTimes{}, fmt.Errorf("/proc/stat: unexpected first line %q", line) + } + + var v [8]uint64 + for i := 0; i < len(v) && i+1 < len(fields); i++ { + n, err := strconv.ParseUint(fields[i+1], 10, 64) + if err != nil { + return cpuTimes{}, fmt.Errorf("/proc/stat: %w", err) + } + v[i] = n + } + + var t cpuTimes + for _, n := range v { + t.Total += n + } + t.Idle = v[3] + v[4] + + return t, nil +} + +// cpuPercent is the busy share of the time between two readings, 0 to 100. +func cpuPercent(prev, cur cpuTimes) float64 { + if cur.Total <= prev.Total || cur.Idle < prev.Idle { + return 0 + } + + total := float64(cur.Total - prev.Total) + idle := float64(cur.Idle - prev.Idle) + return min(max((total-idle)/total*100, 0), 100) +} + +// parseLoad reads the 1 minute load average from /proc/loadavg. +func parseLoad(data []byte) (float64, error) { + fields := strings.Fields(string(data)) + if len(fields) == 0 { + return 0, fmt.Errorf("/proc/loadavg is empty") + } + + return strconv.ParseFloat(fields[0], 64) +} + +// memory holds the values of /proc/meminfo, in bytes. +type memory struct { + total, available, swapTotal, swapFree uint64 +} + +func parseMeminfo(data []byte) (memory, error) { + values := map[string]uint64{} + s := bufio.NewScanner(bytes.NewReader(data)) + for s.Scan() { + key, rest, ok := strings.Cut(s.Text(), ":") + if !ok { + continue + } + fields := strings.Fields(rest) + if len(fields) == 0 { + continue + } + n, err := strconv.ParseUint(fields[0], 10, 64) + if err != nil { + continue + } + // The kernel shows these values in kB (KiB). + values[key] = n * 1024 + } + + for _, key := range []string{"MemTotal", "MemAvailable", "SwapTotal", "SwapFree"} { + if _, ok := values[key]; !ok { + return memory{}, fmt.Errorf("/proc/meminfo has no %s", key) + } + } + + return memory{ + total: values["MemTotal"], + available: values["MemAvailable"], + swapTotal: values["SwapTotal"], + swapFree: values["SwapFree"], + }, nil +} + +// netCounters are the received and sent bytes of one interface. +type netCounters struct { + In uint64 `json:"in"` + Out uint64 `json:"out"` +} + +// parseNetDev reads /proc/net/dev. Each interface line is +// +// name: rx_bytes rx_packets ... (8 receive fields) tx_bytes ... +func parseNetDev(data []byte) (map[string]netCounters, error) { + out := map[string]netCounters{} + s := bufio.NewScanner(bytes.NewReader(data)) + for s.Scan() { + name, rest, ok := strings.Cut(s.Text(), ":") + if !ok { + continue + } + fields := strings.Fields(rest) + if len(fields) < 9 { + continue + } + in, err1 := strconv.ParseUint(fields[0], 10, 64) + sent, err2 := strconv.ParseUint(fields[8], 10, 64) + if err1 != nil || err2 != nil { + return nil, fmt.Errorf("/proc/net/dev: bad counters for %s", strings.TrimSpace(name)) + } + out[strings.TrimSpace(name)] = netCounters{In: in, Out: sent} + } + + return out, nil +} + +// parseDefaultRoute returns the interface of the default route in +// /proc/net/route, or "". The default route has destination and mask 0: a +// VPN that routes 0.0.0.0/1 and 128.0.0.0/1 is not the default route. +// +// Iface Destination Gateway Flags RefCnt Use Metric Mask MTU Window IRTT +func parseDefaultRoute(data []byte) string { + s := bufio.NewScanner(bytes.NewReader(data)) + for s.Scan() { + fields := strings.Fields(s.Text()) + if len(fields) >= 8 && fields[1] == "00000000" && fields[7] == "00000000" { + return fields[0] + } + } + + return "" +} + +// parseOSRelease returns PRETTY_NAME from /etc/os-release. +func parseOSRelease(data []byte) string { + s := bufio.NewScanner(bytes.NewReader(data)) + for s.Scan() { + if v, ok := strings.CutPrefix(s.Text(), "PRETTY_NAME="); ok { + if unquoted, err := strconv.Unquote(v); err == nil { + return unquoted + } + return strings.Trim(v, `"'`) + } + } + + return "" +} + +// parseUptime reads the seconds since boot from /proc/uptime. +func parseUptime(data []byte) (uint64, error) { + fields := strings.Fields(string(data)) + if len(fields) == 0 { + return 0, fmt.Errorf("/proc/uptime is empty") + } + f, err := strconv.ParseFloat(fields[0], 64) + if err != nil || f < 0 { + return 0, fmt.Errorf("/proc/uptime: bad value %q", fields[0]) + } + + return uint64(f), nil +} + +// parseAptCheck reads the output of apt-check: "total;security". apt-check +// can write warnings before the result (on Ubuntu 24.04, for example for a +// source that is configured two times), so only the last line counts. +func parseAptCheck(out []byte) (total, security uint64, err error) { + lines := strings.Split(strings.TrimSpace(string(out)), "\n") + last := strings.TrimSpace(lines[len(lines)-1]) + a, b, ok := strings.Cut(last, ";") + if !ok { + return 0, 0, fmt.Errorf("apt-check: unexpected output %q", last) + } + if total, err = strconv.ParseUint(a, 10, 64); err != nil { + return 0, 0, fmt.Errorf("apt-check: %w", err) + } + if security, err = strconv.ParseUint(b, 10, 64); err != nil { + return 0, 0, fmt.Errorf("apt-check: %w", err) + } + + return total, security, nil +} + +// parsePSI reads the total of the "some" line of a /proc/pressure file: the +// microseconds in which at least one task waited for the resource. +// +// some avg10=0.00 avg60=0.00 avg300=0.00 total=12345 +// full avg10=0.00 avg60=0.00 avg300=0.00 total=0 +func parsePSI(data []byte) (uint64, error) { + s := bufio.NewScanner(bytes.NewReader(data)) + for s.Scan() { + fields := strings.Fields(s.Text()) + if len(fields) == 0 || fields[0] != "some" { + continue + } + for _, f := range fields[1:] { + if v, ok := strings.CutPrefix(f, "total="); ok { + n, err := strconv.ParseUint(v, 10, 64) + if err != nil { + return 0, fmt.Errorf("pressure: bad total %q", v) + } + return n, nil + } + } + } + + return 0, fmt.Errorf("pressure: no total on a \"some\" line") +} + +// diskCounters are the completed operations and the 512-byte sectors of one +// disk, from /proc/diskstats. +type diskCounters struct { + ReadOps uint64 `json:"read_ops"` + ReadSectors uint64 `json:"read_sectors"` + WriteOps uint64 `json:"write_ops"` + WriteSectors uint64 `json:"write_sectors"` +} + +// parseDiskstats reads /proc/diskstats. Each line is +// +// major minor name reads merged sectors ms writes merged sectors ms ... +// +// The reads and writes completed are columns 4 and 8, and the sectors read +// and written are columns 6 and 10. A sector is 512 bytes for each disk. +func parseDiskstats(data []byte) (map[string]diskCounters, error) { + out := map[string]diskCounters{} + s := bufio.NewScanner(bytes.NewReader(data)) + for s.Scan() { + fields := strings.Fields(s.Text()) + if len(fields) < 10 { + continue + } + var v [4]uint64 + for i, col := range []int{3, 5, 7, 9} { + n, err := strconv.ParseUint(fields[col], 10, 64) + if err != nil { + return nil, fmt.Errorf("/proc/diskstats: bad counters for %s", fields[2]) + } + v[i] = n + } + out[fields[2]] = diskCounters{ReadOps: v[0], ReadSectors: v[1], WriteOps: v[2], WriteSectors: v[3]} + } + + return out, nil +} + +// parseCPUCount counts the CPUs that are online: the cpuN lines of /proc/stat. +func parseCPUCount(data []byte) (int, error) { + n := 0 + s := bufio.NewScanner(bytes.NewReader(data)) + for s.Scan() { + name, _, _ := strings.Cut(s.Text(), " ") + if rest, ok := strings.CutPrefix(name, "cpu"); ok && rest != "" { + if _, err := strconv.Atoi(rest); err == nil { + n++ + } + } + } + if n == 0 { + return 0, fmt.Errorf("/proc/stat has no cpuN line") + } + + return n, nil +} diff --git a/internal/metrics/parse_test.go b/internal/metrics/parse_test.go new file mode 100644 index 0000000..ede24c7 --- /dev/null +++ b/internal/metrics/parse_test.go @@ -0,0 +1,118 @@ +package metrics + +import ( + "testing" +) + +func TestParseCPU(t *testing.T) { + // user nice system idle iowait irq softirq steal guest guest_nice + got, err := parseCPU([]byte("cpu 100 10 50 800 40 0 0 0 30 0\ncpu0 1 1 1 1 1 1 1 1 1 1\n")) + if err != nil { + t.Fatal(err) + } + // Guest time is already in user time, so it is not added again. + if got.Total != 1000 || got.Idle != 840 { + t.Errorf("parseCPU() = %+v, want total 1000 and idle 840", got) + } + + for _, bad := range []string{"", "intr 1 2 3", "cpu a b c d"} { + if _, err := parseCPU([]byte(bad)); err == nil { + t.Errorf("parseCPU(%q) = nil error, want an error", bad) + } + } +} + +func TestCPUPercent(t *testing.T) { + tests := []struct { + prev, cur cpuTimes + want float64 + }{ + {cpuTimes{Idle: 800, Total: 1000}, cpuTimes{Idle: 950, Total: 1600}, 75}, + {cpuTimes{Idle: 800, Total: 1000}, cpuTimes{Idle: 1400, Total: 1600}, 0}, + {cpuTimes{Idle: 800, Total: 1000}, cpuTimes{Idle: 800, Total: 1600}, 100}, + // No time passed, or the counters went back. + {cpuTimes{Idle: 800, Total: 1000}, cpuTimes{Idle: 800, Total: 1000}, 0}, + {cpuTimes{Idle: 800, Total: 1000}, cpuTimes{Idle: 10, Total: 20}, 0}, + } + + for _, tt := range tests { + if got := cpuPercent(tt.prev, tt.cur); got != tt.want { + t.Errorf("cpuPercent(%+v, %+v) = %v, want %v", tt.prev, tt.cur, got, tt.want) + } + } +} + +func TestParseMeminfoNeedsMemAvailable(t *testing.T) { + if _, err := parseMeminfo([]byte("MemTotal: 100 kB\nMemFree: 50 kB\nSwapTotal: 0 kB\nSwapFree: 0 kB\n")); err == nil { + t.Error("parseMeminfo() = nil error, want an error without MemAvailable") + } +} + +func TestParseNetDev(t *testing.T) { + data := "Inter-| Receive | Transmit\n face |bytes packets|bytes\n" + + " eth0: 1234 5 0 0 0 0 0 0 5678 6 0 0 0 0 0 0\n" + + " lo:1 1 0 0 0 0 0 0 2 2 0 0 0 0 0 0\n" + + got, err := parseNetDev([]byte(data)) + if err != nil { + t.Fatal(err) + } + if got["eth0"] != (netCounters{In: 1234, Out: 5678}) || got["lo"] != (netCounters{In: 1, Out: 2}) || len(got) != 2 { + t.Errorf("parseNetDev() = %+v", got) + } +} + +func TestParseDefaultRoute(t *testing.T) { + const header = "Iface\tDestination\tGateway \tFlags\tRefCnt\tUse\tMetric\tMask\t\tMTU\tWindow\tIRTT\n" + data := header + + "eth1\t0A000000\t00000000\t0001\t0\t0\t0\t000000FF\t0\t0\t0\n" + + "eth0\t00000000\t01C0A8C0\t0003\t0\t0\t100\t00000000\t0\t0\t0\n" + if got := parseDefaultRoute([]byte(data)); got != "eth0" { + t.Errorf("parseDefaultRoute() = %q, want eth0", got) + } + + // OpenVPN def1: 0.0.0.0/1 on tun0 is not the default route. + vpn := header + "tun0\t00000000\t0100080A\t0003\t0\t0\t0\t00000080\t0\t0\t0\n" + if got := parseDefaultRoute([]byte(vpn)); got != "" { + t.Errorf("parseDefaultRoute() = %q, want no default route for 0.0.0.0/1", got) + } + if got := parseDefaultRoute([]byte("Iface\tDestination\n")); got != "" { + t.Errorf("parseDefaultRoute() = %q, want no interface", got) + } +} + +func TestParseOSRelease(t *testing.T) { + tests := map[string]string{ + "NAME=\"Ubuntu\"\nPRETTY_NAME=\"Ubuntu 24.04.1 LTS\"\n": "Ubuntu 24.04.1 LTS", + "PRETTY_NAME='Debian GNU/Linux 12'\n": "Debian GNU/Linux 12", + "PRETTY_NAME=Alpine\n": "Alpine", + "NAME=x\n": "", + } + for in, want := range tests { + if got := parseOSRelease([]byte(in)); got != want { + t.Errorf("parseOSRelease(%q) = %q, want %q", in, got, want) + } + } +} + +func TestParseUptime(t *testing.T) { + if got, err := parseUptime([]byte("350735.47 234388.90\n")); err != nil || got != 350735 { + t.Errorf("parseUptime() = %d, %v; want 350735", got, err) + } + if _, err := parseUptime([]byte("")); err == nil { + t.Error("parseUptime(\"\") = nil error, want an error") + } +} + +func TestParseAptCheck(t *testing.T) { + for _, out := range []string{"33;6", "33;6\n", "Warning: W:something; else\n33;6"} { + if total, security, err := parseAptCheck([]byte(out)); err != nil || total != 33 || security != 6 { + t.Errorf("parseAptCheck(%q) = %d, %d, %v; want 33, 6", out, total, security, err) + } + } + for _, bad := range []string{"", "33", "a;b", "1;x"} { + if _, _, err := parseAptCheck([]byte(bad)); err == nil { + t.Errorf("parseAptCheck(%q) = nil error, want an error", bad) + } + } +} diff --git a/internal/metrics/pressure.go b/internal/metrics/pressure.go new file mode 100644 index 0000000..e07f15b --- /dev/null +++ b/internal/metrics/pressure.go @@ -0,0 +1,86 @@ +package metrics + +import ( + "slices" + + "github.com/flywp/server-cli/internal/agent/wire" +) + +// psiTotals are the "some" totals of /proc/pressure, in microseconds. +type psiTotals struct { + CPU uint64 `json:"cpu"` + Memory uint64 `json:"memory"` + IO uint64 `json:"io"` +} + +// readPSI reads the three pressure files. It returns nil when the kernel has +// no PSI: no /proc/pressure, or a kernel booted with psi=0, where a read +// fails. It logs this one time. +func (c *Collector) readPSI() *psiTotals { + var t psiTotals + for _, f := range []struct { + name string + v *uint64 + }{{"cpu", &t.CPU}, {"memory", &t.Memory}, {"io", &t.IO}} { + n, err := parseFile(c, "proc/pressure/"+f.name, parsePSI) + if err != nil { + if !c.noPSI { + c.noPSI = true + c.log.Warn("no pressure (PSI) to send: the kernel has no PSI, or it is off", "error", err) + } + return nil + } + *f.v = n + } + + return &t +} + +// pressure returns the share of the time between two readings in which at +// least one task waited for the CPU, the memory and the disk, 0 to 100. ok is +// false when the share is not known: a reading without PSI, a reboot, a +// reading older than maxAge, or a counter that went back. +func pressure(a, b reading) (cpu, mem, io float64, ok bool) { + if a.PSI == nil || b.PSI == nil || a.BootID != b.BootID { + return 0, 0, 0, false + } + d := b.At.Sub(a.At) + if d <= 0 || d > maxAge { + return 0, 0, 0, false + } + if b.PSI.CPU < a.PSI.CPU || b.PSI.Memory < a.PSI.Memory || b.PSI.IO < a.PSI.IO { + return 0, 0, 0, false + } + + us := float64(d.Microseconds()) + share := func(from, to uint64) float64 { + return min(float64(to-from)/us*100, 100) + } + return share(a.PSI.CPU, b.PSI.CPU), share(a.PSI.Memory, b.PSI.Memory), share(a.PSI.IO, b.PSI.IO), true +} + +// setPressure sets the pressure of the minute from prev to cur in s, and the +// peak of the windows in all. The six fields stay nil when the pressure of the +// minute is not known. +func setPressure(s *wire.Sample, prev *reading, all []reading, cur reading) { + if prev == nil { + return + } + cpu, mem, io, ok := pressure(*prev, cur) + if !ok { + return + } + + cpuMax, memMax, ioMax := cpu, mem, io + // A reading without PSI is left out: its windows join. + all = slices.DeleteFunc(slices.Clone(all), func(r reading) bool { return r.PSI == nil }) + for i := 1; i < len(all); i++ { + if c, m, o, ok := pressure(all[i-1], all[i]); ok { + cpuMax, memMax, ioMax = max(cpuMax, c), max(memMax, m), max(ioMax, o) + } + } + + s.CPUPressurePercent, s.CPUPressureMaxPercent = &cpu, &cpuMax + s.MemoryPressurePercent, s.MemoryPressureMaxPercent = &mem, &memMax + s.IOPressurePercent, s.IOPressureMaxPercent = &io, &ioMax +} diff --git a/internal/metrics/pressure_test.go b/internal/metrics/pressure_test.go new file mode 100644 index 0000000..3e698e5 --- /dev/null +++ b/internal/metrics/pressure_test.go @@ -0,0 +1,197 @@ +package metrics + +import ( + "fmt" + "os" + "path/filepath" + "testing" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" +) + +// psi writes the three /proc/pressure files with the "some" totals, in +// microseconds. +func (s *server) psi(cpu, memory, io uint64) { + for name, total := range map[string]uint64{"cpu": cpu, "memory": memory, "io": io} { + s.write("proc/pressure/"+name, fmt.Sprintf("some avg10=1.00 avg60=2.00 avg300=3.00 total=%d\nfull avg10=0.00 avg60=0.00 avg300=0.00 total=7\n", total)) + } +} + +func TestParsePSI(t *testing.T) { + n, err := parsePSI([]byte("some avg10=0.12 avg60=0.34 avg300=0.56 total=987654321\nfull avg10=0.00 avg60=0.00 avg300=0.00 total=5\n")) + if err != nil || n != 987654321 { + t.Errorf("parsePSI() = %d, %v; want the total of the some line", n, err) + } + // The CPU file of an older kernel has no full line. + if n, err := parsePSI([]byte("some avg10=0.00 avg60=0.00 avg300=0.00 total=42\n")); err != nil || n != 42 { + t.Errorf("parsePSI() = %d, %v; want 42", n, err) + } + for _, bad := range []string{"", "full avg10=0.00 total=5\n", "some avg10=0.00\n", "some total=x\n"} { + if _, err := parsePSI([]byte(bad)); err == nil { + t.Errorf("parsePSI(%q) = nil error, want an error", bad) + } + } +} + +// pressureMinute takes a sample at base, five readings and the tick at +// base + 60 s. totals are the CPU, memory and I/O totals at each of the seven +// readings. +func pressureMinute(t *testing.T, srv *server, c *Collector, base time.Time, totals [7][3]uint64) wire.Sample { + t.Helper() + srv.psi(totals[0][0], totals[0][1], totals[0][2]) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + for i := 1; i < 6; i++ { + srv.psi(totals[i][0], totals[i][1], totals[i][2]) + c.Read(base.Add(time.Duration(i) * 10 * time.Second)) + } + srv.psi(totals[6][0], totals[6][1], totals[6][2]) + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + return s +} + +func TestPressure(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + // The CPU waits 5 s in the window from 20 s to 30 s: 50% of that window, + // and 6 s of the minute: 10%. The memory never waits. The I/O waits + // 0.6 s in each window: 1%. + s := pressureMinute(t, srv, c, base, [7][3]uint64{ + {1000, 0, 0}, + {201000, 0, 100000}, + {401000, 0, 200000}, + {5401000, 0, 300000}, + {5601000, 0, 400000}, + {5801000, 0, 500000}, + {6001000, 0, 600000}, + }) + + want := map[string][2]float64{"cpu": {10, 50}, "memory": {0, 0}, "io": {1, 1}} + got := map[string][2]*float64{ + "cpu": {s.CPUPressurePercent, s.CPUPressureMaxPercent}, + "memory": {s.MemoryPressurePercent, s.MemoryPressureMaxPercent}, + "io": {s.IOPressurePercent, s.IOPressureMaxPercent}, + } + for name, w := range want { + g := got[name] + if g[0] == nil || g[1] == nil || !near(*g[0], w[0]) || !near(*g[1], w[1]) { + t.Errorf("%s pressure = %v, max %v; want %v, max %v", name, ptr(g[0]), ptr(g[1]), w[0], w[1]) + } + } +} + +func TestPressureIsNotKnown(t *testing.T) { + tests := []struct { + name string + change func(*server) + after time.Duration + }{ + {"no /proc/pressure", func(s *server) { + if err := os.RemoveAll(filepath.Join(s.root, "proc/pressure")); err != nil { + s.t.Fatal(err) + } + }, time.Minute}, + {"reboot", func(s *server) { s.write("proc/sys/kernel/random/boot_id", "boot-2\n") }, time.Minute}, + {"counter went back", func(s *server) { s.psi(10, 10, 10) }, time.Minute}, + {"previous reading too old", func(*server) {}, 3 * time.Minute}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + srv := newServer(t) + srv.psi(1000, 1000, 1000) + c := srv.collector(t.TempDir()) + now := time.Now().Add(time.Second) + if _, err := c.Sample(now); err != nil { + t.Fatal(err) + } + + tt.change(srv) + s, err := c.Sample(now.Add(tt.after)) + if err != nil { + t.Fatal(err) + } + checkNoPressure(t, s) + }) + } +} + +func TestFirstSampleHasNoPressure(t *testing.T) { + srv := newServer(t) + srv.psi(1000, 1000, 1000) + c := srv.collector(t.TempDir()) + srv.psi(2000, 2000, 2000) + + s, err := c.Sample(time.Now().Add(30 * time.Second)) + if err != nil { + t.Fatal(err) + } + checkNoPressure(t, s) +} + +func TestPressureContinuesAfterARestart(t *testing.T) { + srv := newServer(t) + srv.psi(1000, 1000, 1000) + state := t.TempDir() + now := time.Now() + if _, err := srv.collector(state).Sample(now.Add(-30 * time.Second)); err != nil { + t.Fatal(err) + } + + c := srv.collector(state) + srv.psi(601000, 1000, 1000) + s, err := c.Sample(now.Add(30 * time.Second)) + if err != nil { + t.Fatal(err) + } + if s.CPUPressurePercent == nil || !near(*s.CPUPressurePercent, 1) { + t.Errorf("cpu_pressure_percent = %v, want 1 from the saved reading", ptr(s.CPUPressurePercent)) + } +} + +func checkNoPressure(t *testing.T, s wire.Sample) { + t.Helper() + for _, v := range []*float64{s.CPUPressurePercent, s.CPUPressureMaxPercent, s.MemoryPressurePercent, s.MemoryPressureMaxPercent, s.IOPressurePercent, s.IOPressureMaxPercent} { + if v != nil { + t.Errorf("a pressure field = %v, want null", *v) + } + } +} + +// near reports whether a is b, up to the rounding of the microseconds. +func near(a, b float64) bool { + return a > b-0.001 && a < b+0.001 +} + +func TestAReadingWithoutPressureJoinsTheWindows(t *testing.T) { + srv := newServer(t) + srv.psi(0, 0, 0) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + + // The CPU waits 8 s from 0 s to 20 s, but the reading at 10 s has no + // PSI: the peak is 8 s of the joined window of 20 s. + srv.write("proc/pressure/cpu", "") + c.Read(base.Add(10 * time.Second)) + srv.psi(8000000, 0, 0) + c.Read(base.Add(20 * time.Second)) + srv.psi(9200000, 0, 0) + + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.CPUPressureMaxPercent == nil || !near(*s.CPUPressureMaxPercent, 40) { + t.Errorf("cpu_pressure_max_percent = %v, want 40 from the joined window of 0 s to 20 s", ptr(s.CPUPressureMaxPercent)) + } +} diff --git a/internal/metrics/sites.go b/internal/metrics/sites.go new file mode 100644 index 0000000..9447742 --- /dev/null +++ b/internal/metrics/sites.go @@ -0,0 +1,339 @@ +package metrics + +import ( + "bufio" + "bytes" + "context" + "fmt" + "log/slog" + "os" + "os/user" + "path/filepath" + "slices" + "strconv" + "strings" + "sync" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" +) + +const ( + // workingDirLabel is the label that Docker Compose puts on each container + // of a project: the folder of the project. + workingDirLabel = "com.docker.compose.project.working_dir" + + // The disk use of each project folder is measured at most each hour, in + // the background, and a walk of one folder stops after 5 minutes. + sitesDiskEvery = time.Hour + sitesDiskTimeout = 5 * time.Minute +) + +// containerUsage is the CPU time of one container, in microseconds, at a tick. +// cgroup identifies the cgroup folder of the container: a restart keeps the +// id of the container, but makes a new cgroup whose CPU time starts again. +type containerUsage struct { + project string + usec uint64 + cgroup uint64 +} + +// containerReadings are the CPU times of the containers at one tick. +type containerReadings struct { + bootID string + at time.Time + usage map[string]containerUsage // by container id +} + +// homeDir returns the home folder of the user that runs the agent: the +// server user. It is "" when it is not known. +func homeDir() string { + if u, err := user.Current(); err == nil && u.HomeDir != "" { + return filepath.Clean(u.HomeDir) + } + if h, err := os.UserHomeDir(); err == nil { + return filepath.Clean(h) + } + return "" +} + +// sites measures each Docker Compose project in the home folder at the tick +// cur, read at readAt (contract v0.5.0). It returns nil when the agent cannot read Docker or +// the cgroups: Docker does not run, the socket refuses the agent, or the +// server has cgroup v1. It returns an empty, non-nil slice when no project +// matches. +func (c *Collector) sites(ctx context.Context, cur reading, readAt time.Time) []wire.Site { + // Without a home folder, or with "/" as the home, no folder is a project + // of a site. + if c.home == "" || filepath.Dir(c.home) == c.home { + return nil + } + if _, err := os.Stat(c.file("sys/fs/cgroup/cgroup.controllers")); err != nil { + // cgroup v1, or no cgroups: the agent reads only cgroup v2. + return nil + } + containers, err := c.docker.Containers(ctx) + if err != nil { + c.log.Debug("listing the containers", "error", err) + return nil + } + + type project struct { + cpuUsec uint64 + cpuOK bool + mem uint64 + memOK bool + } + projects := map[string]*project{} + // The CPU times are read now, some time after the tick reading: the + // request to Docker and apt-check come before. The time of the minute + // is the time between two such reads. + at := cur.At.Add(time.Since(readAt)) + now := containerReadings{bootID: cur.BootID, at: at, usage: map[string]containerUsage{}} + prev := c.prevContainers + usable := prev != nil && prev.bootID == cur.BootID && at.After(prev.at) && at.Sub(prev.at) <= maxAge + + for _, ct := range containers { + dir := filepath.Clean(ct.Labels[workingDirLabel]) + if ct.Labels[workingDirLabel] == "" || filepath.Dir(dir) != c.home { + continue + } + name := filepath.Base(dir) + p := projects[name] + if p == nil { + p = &project{} + projects[name] = p + } + + cg := c.cgroupDir(ct.ID) + if cg == "" { + continue + } + if usec, err := parseFile(c, filepath.Join(cg, "cpu.stat"), parseCPUStat); err == nil { + id := fileID(c.file(cg)) + now.usage[ct.ID] = containerUsage{project: name, usec: usec, cgroup: id} + // Only a container that both ticks saw, with the same id and + // the same cgroup: a container that started or restarted in + // the minute is left out. + if old, ok := prev.lookup(ct.ID); usable && ok && old.project == name && old.cgroup == id && usec >= old.usec { + p.cpuUsec += usec - old.usec + p.cpuOK = true + } + } + if mem, err := c.containerMemory(cg); err == nil { + p.mem += mem + p.memOK = true + } + } + c.prevContainers = &now + + cpus, cpuErr := parseFile(c, "proc/stat", parseCPUCount) + disk := c.walker.take() + out := make([]wire.Site, 0, len(projects)) + for _, name := range slices.Sorted(func(yield func(string) bool) { + for name := range projects { + if !yield(name) { + return + } + } + }) { + p := projects[name] + site := wire.Site{Directory: name} + if p.cpuOK && cpuErr == nil { + us := float64(at.Sub(prev.at).Microseconds()) * float64(cpus) + v := min(float64(p.cpuUsec)/us*100, 100) + site.CPUPercent = &v + } + if p.memOK { + site.MemoryUsedBytes = &p.mem + } + if d, ok := disk[name]; ok { + site.DiskUsedBytes = &d + } + out = append(out, site) + } + + // Measure the disk use of the folders in the background, at most each + // hour. The result goes into the next sample. + folders := map[string]string{} + for name := range projects { + folders[name] = filepath.Join(c.home, name) + } + c.walker.start(folders, c.log) + + return out +} + +// lookup returns the reading of a container at the tick before. +func (r *containerReadings) lookup(id string) (containerUsage, bool) { + if r == nil { + return containerUsage{}, false + } + u, ok := r.usage[id] + return u, ok +} + +// cgroupDir returns the cgroup v2 folder of a container, relative to the +// root, or "". The systemd cgroup driver (the default of Ubuntu) and the +// cgroupfs driver use different folders. +func (c *Collector) cgroupDir(id string) string { + if id == "" || strings.ContainsAny(id, "/.") { + return "" + } + for _, dir := range []string{ + filepath.Join("sys/fs/cgroup/system.slice", "docker-"+id+".scope"), + filepath.Join("sys/fs/cgroup/docker", id), + } { + if _, err := os.Stat(c.file(dir, "cpu.stat")); err == nil { + return dir + } + } + return "" +} + +// containerMemory returns memory.current − inactive_file of a cgroup, as +// docker stats shows it on cgroup v2. The page cache that the kernel can drop +// is not used memory. A negative value counts as 0. +func (c *Collector) containerMemory(cg string) (uint64, error) { + current, err := parseFile(c, filepath.Join(cg, "memory.current"), func(data []byte) (uint64, error) { + return strconv.ParseUint(string(bytes.TrimSpace(data)), 10, 64) + }) + if err != nil { + return 0, err + } + inactive, err := parseFile(c, filepath.Join(cg, "memory.stat"), statValue("inactive_file")) + if err != nil { + return 0, err + } + return current - min(inactive, current), nil +} + +// parseCPUStat reads usage_usec from a cgroup v2 cpu.stat. +func parseCPUStat(data []byte) (uint64, error) { + return statValue("usage_usec")(data) +} + +// statValue returns a parser for the value of key in a cgroup file of +// "key value" lines. +func statValue(key string) func([]byte) (uint64, error) { + return func(data []byte) (uint64, error) { + s := bufio.NewScanner(bytes.NewReader(data)) + for s.Scan() { + k, v, ok := strings.Cut(s.Text(), " ") + if ok && k == key { + return strconv.ParseUint(strings.TrimSpace(v), 10, 64) + } + } + return 0, fmt.Errorf("no %s", key) + } +} + +// diskWalker measures the disk use of the project folders in a goroutine of +// its own, so that a large folder does not delay the samples. The collector +// takes each result one time. +type diskWalker struct { + mu sync.Mutex + running bool + last time.Time // the start of the last walk + results map[string]uint64 // by directory, not yet sent + // skipWarned holds the directories whose skipped files were logged as a + // warning. Only the walk goroutine uses it. + skipWarned map[string]bool + // done is closed when the running walk ends. Tests wait for it. + done chan struct{} + // walk measures one folder, and timeout limits it. Tests replace them. + walk func(ctx context.Context, root string) (uint64, int, error) + timeout time.Duration +} + +func newDiskWalker() *diskWalker { + return &diskWalker{walk: diskUsage, timeout: sitesDiskTimeout, skipWarned: map[string]bool{}} +} + +// start begins a walk of the folders, by directory, when no walk runs and the +// last walk started one hour ago or more. +func (w *diskWalker) start(folders map[string]string, log *slog.Logger) { + w.mu.Lock() + defer w.mu.Unlock() + if w.running || len(folders) == 0 || (!w.last.IsZero() && time.Since(w.last) < sitesDiskEvery) { + return + } + w.running, w.last = true, time.Now() + done := make(chan struct{}) + w.done = done + + go func() { + defer close(done) + + results := map[string]uint64{} + for _, name := range slices.Sorted(func(yield func(string) bool) { + for name := range folders { + if !yield(name) { + return + } + } + }) { + used, skipped, err := w.walkOne(folders[name]) + if err != nil { + log.Warn("cannot measure the disk use of a site; sending null", "directory", name, "error", err) + continue + } + if skipped > 0 { + // The same folders are skipped each hour, for example the + // databases in ~/.fly that belong to the container user: + // warn one time for each directory, then log at debug. + level := slog.LevelDebug + if !w.skipWarned[name] { + w.skipWarned[name] = true + level = slog.LevelWarn + } + log.Log(context.Background(), level, "some files of a site cannot be read; its disk use is lower than the real use", "directory", name, "skipped", skipped) + } + results[name] = used + } + + w.mu.Lock() + defer w.mu.Unlock() + w.results, w.running = results, false + }() +} + +// walkOne measures one folder in a goroutine of its own, and stops waiting +// for it after the timeout. A walk that hangs in a system call, for example on +// a network mount that does not answer, stays behind; the other folders and +// the next walks go on. +func (w *diskWalker) walkOne(root string) (uint64, int, error) { + ctx, cancel := context.WithTimeout(context.Background(), w.timeout) + defer cancel() + + type result struct { + used uint64 + skipped int + err error + } + ch := make(chan result, 1) + go func() { + // The walk runs at the lowest I/O priority, so that it does not + // slow the sites. The priority is a property of the thread: the + // thread stays locked, and ends with the goroutine. + lowerIOPriority() + used, skipped, err := w.walk(ctx, root) + ch <- result{used, skipped, err} + }() + + select { + case r := <-ch: + return r.used, r.skipped, r.err + case <-ctx.Done(): + return 0, 0, ctx.Err() + } +} + +// take returns the results of the last walk that ended, one time. +func (w *diskWalker) take() map[string]uint64 { + w.mu.Lock() + defer w.mu.Unlock() + r := w.results + w.results = nil + return r +} diff --git a/internal/metrics/sites_test.go b/internal/metrics/sites_test.go new file mode 100644 index 0000000..ec95d1d --- /dev/null +++ b/internal/metrics/sites_test.go @@ -0,0 +1,481 @@ +package metrics + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log/slog" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "github.com/flywp/server-cli/internal/agent" + "github.com/flywp/server-cli/internal/agent/wire" + "github.com/flywp/server-cli/internal/dockerapi" + "github.com/flywp/server-cli/internal/testutil" +) + +// fakeContainer is a running container of the fake Docker Engine. +type fakeContainer struct { + id, dir string +} + +// containerJSON is the answer of GET /containers/json for the containers. +func containerJSON(cs []fakeContainer) string { + var items []string + for _, c := range cs { + labels := "{}" + if c.dir != "" { + labels = fmt.Sprintf(`{"com.docker.compose.project.working_dir":%q,"com.docker.compose.service":"php"}`, c.dir) + } + items = append(items, fmt.Sprintf(`{"Id":%q,"Names":["/%s"],"Labels":%s}`, c.id, c.id, labels)) + } + return "[" + strings.Join(items, ",") + "]" +} + +// cgroup writes the cgroup v2 files of a container, with the systemd driver. +func (s *server) cgroup(id string, usageUsec, current, inactiveFile uint64) { + s.cgroupAt(filepath.Join("sys/fs/cgroup/system.slice", "docker-"+id+".scope"), usageUsec, current, inactiveFile) +} + +func (s *server) cgroupAt(dir string, usageUsec, current, inactiveFile uint64) { + s.write("sys/fs/cgroup/cgroup.controllers", "cpuset cpu io memory pids\n") + s.write(filepath.Join(dir, "cpu.stat"), fmt.Sprintf("usage_usec %d\nuser_usec 1\nsystem_usec 2\n", usageUsec)) + s.write(filepath.Join(dir, "memory.current"), fmt.Sprintf("%d\n", current)) + s.write(filepath.Join(dir, "memory.stat"), fmt.Sprintf("anon 1\nfile 2\nactive_file 3\ninactive_file %d\n", inactiveFile)) +} + +// sitesCollector returns a collector whose home is /home/fly under the root, +// and whose Docker Engine lists the containers. The disk walk returns 4096 +// bytes for each folder. +func sitesCollector(t *testing.T, srv *server, cs []fakeContainer) *Collector { + t.Helper() + engine := testutil.NewEngine(t, testutil.EngineHandler("29.7.1", containerJSON(cs))) + c := srv.collector(t.TempDir()) + c.docker = dockerapi.New(engine.Socket) + c.home = filepath.Join(srv.root, "home/fly") + c.walker.walk = func(context.Context, string) (uint64, int, error) { return 4096, 0, nil } + return c +} + +// home is the folder of a project in the home folder of the fake server. +func (s *server) home(name string) string { + return filepath.Join(s.root, "home/fly", name) +} + +func TestSites(t *testing.T) { + srv := newServer(t) + // Four CPUs. + srv.write("proc/stat", "cpu 1000 0 0 800 0 0 0 0 0 0\ncpu0 1 2 3 4 5 6 7 8 0 0\ncpu1 1 2 3 4 5 6 7 8 0 0\ncpu2 1 2 3 4 5 6 7 8 0 0\ncpu3 1 2 3 4 5 6 7 8 0 0\n") + cs := []fakeContainer{ + {"a1", srv.home("example.com")}, + {"a2", srv.home("example.com")}, + {"b1", srv.home(".fly")}, + {"c1", "/srv/other"}, // not in the home folder + {"d1", ""}, // not a Compose container + {"e1", srv.home("example.com/nested/deep")}, // the parent is not the home folder + } + srv.cgroup("a1", 1000000, 300<<20, 100<<20) + srv.cgroup("a2", 2000000, 50<<20, 80<<20) // more inactive file than current: 0 + srv.cgroup("b1", 3000000, 400<<20, 0) + srv.cgroup("c1", 1, 1, 0) + srv.cgroup("e1", 1, 1, 0) + c := sitesCollector(t, srv, cs) + base := time.Now().Add(time.Second) + + first, err := c.Sample(base) + if err != nil { + t.Fatal(err) + } + if got := siteNames(first.Sites); got != "[.fly example.com]" { + t.Fatalf("sites = %s, want [.fly example.com]", got) + } + for _, site := range first.Sites { + if site.CPUPercent != nil { + t.Errorf("%s cpu = %v in the first sample, want null", site.Directory, *site.CPUPercent) + } + } + if m := first.Sites[1].MemoryUsedBytes; m == nil || *m != 200<<20 { + t.Errorf("example.com memory = %v, want %d", ptr(m), 200<<20) + } + + // One minute: a1 and a2 use 2.4 s of CPU, b1 uses 4.8 s. The server has + // 240 s of CPU in one minute. + srv.cgroup("a1", 1000000+1200000, 300<<20, 100<<20) + srv.cgroup("a2", 2000000+1200000, 50<<20, 80<<20) + srv.cgroup("b1", 3000000+4800000, 400<<20, 0) + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + want := map[string]float64{".fly": 2, "example.com": 1} + for _, site := range s.Sites { + if site.CPUPercent == nil || !near(*site.CPUPercent, want[site.Directory]) { + t.Errorf("%s cpu = %v, want %v", site.Directory, ptr(site.CPUPercent), want[site.Directory]) + } + } + if m := s.Sites[0].MemoryUsedBytes; m == nil || *m != 400<<20 { + t.Errorf(".fly memory = %v, want %d", ptr(m), 400<<20) + } +} + +func TestSitesLeaveOutAContainerThatRestarted(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1000000, 1, 0) + srv.cgroup("a2", 1000000, 1, 0) + engineCS := []fakeContainer{{"a1", srv.home("example.com")}, {"a2", srv.home("example.com")}} + c := sitesCollector(t, srv, engineCS) + base := time.Now().Add(time.Second) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + + // a2 restarted: it has a new id, and its CPU time starts again from 0. + srv.cgroup("a1", 1000000+600000, 1, 0) + srv.cgroup("a3", 50000000, 1, 0) + c.docker = dockerapi.New(testutil.NewEngine(t, testutil.EngineHandler("29.7.1", + containerJSON([]fakeContainer{{"a1", srv.home("example.com")}, {"a3", srv.home("example.com")}}))).Socket) + + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if len(s.Sites) != 1 || s.Sites[0].CPUPercent == nil || !near(*s.Sites[0].CPUPercent, 1) { + t.Errorf("sites = %+v, want example.com with 1%% from a1 only", s.Sites) + } +} + +func TestSitesWithTheCgroupfsDriver(t *testing.T) { + srv := newServer(t) + srv.cgroupAt("sys/fs/cgroup/docker/a1", 1000000, 10<<20, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home("example.com")}}) + + s, err := c.Sample(time.Now().Add(time.Second)) + if err != nil { + t.Fatal(err) + } + if len(s.Sites) != 1 || s.Sites[0].MemoryUsedBytes == nil || *s.Sites[0].MemoryUsedBytes != 10<<20 { + t.Errorf("sites = %+v, want the memory of a1", s.Sites) + } +} + +func TestSitesNullAndEmpty(t *testing.T) { + t.Run("no project", func(t *testing.T) { + srv := newServer(t) + srv.cgroup("c1", 1, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"c1", "/srv/other"}}) + s, err := c.Sample(time.Now().Add(time.Second)) + if err != nil { + t.Fatal(err) + } + data, err := json.Marshal(s) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(data), `"sites":[]`) { + t.Errorf("sample JSON = %s, want \"sites\":[]: Docker runs, and no project matches", data) + } + }) + + for _, tt := range []struct { + name string + setup func(*testing.T, *server, *Collector) + }{ + {"Docker does not run", func(_ *testing.T, _ *server, c *Collector) { + c.docker = dockerapi.New("/nonexistent/docker.sock") + }}, + {"cgroup v1", func(t *testing.T, s *server, _ *Collector) { + if err := os.Remove(filepath.Join(s.root, "sys/fs/cgroup/cgroup.controllers")); err != nil { + t.Fatal(err) + } + }}, + {"no home folder", func(_ *testing.T, _ *server, c *Collector) { c.home = "" }}, + } { + t.Run(tt.name, func(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home("example.com")}}) + tt.setup(t, srv, c) + s, err := c.Sample(time.Now().Add(time.Second)) + if err != nil { + t.Fatal(err) + } + data, err := json.Marshal(s) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(data), `"sites":null`) { + t.Errorf("sample JSON = %s, want \"sites\":null", data) + } + }) + } +} + +func TestSitesDiskIsSentOneTime(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home("example.com")}}) + var walked []string + c.walker.walk = func(_ context.Context, root string) (uint64, int, error) { + walked = append(walked, root) + return 8192, 0, nil + } + base := time.Now().Add(time.Second) + + first, err := c.Sample(base) + if err != nil { + t.Fatal(err) + } + if first.Sites[0].DiskUsedBytes != nil { + t.Errorf("disk in the first sample = %d, want null: the walk runs in the background", *first.Sites[0].DiskUsedBytes) + } + waitWalk(t, c) + if len(walked) != 1 || walked[0] != srv.home("example.com") { + t.Errorf("walked %v, want the folder of example.com", walked) + } + + second, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if d := second.Sites[0].DiskUsedBytes; d == nil || *d != 8192 { + t.Errorf("disk in the second sample = %v, want 8192", ptr(d)) + } + + // One time only, and no new walk within one hour. + third, err := c.Sample(base.Add(2 * time.Minute)) + if err != nil { + t.Fatal(err) + } + waitWalk(t, c) + if third.Sites[0].DiskUsedBytes != nil || len(walked) != 1 { + t.Errorf("third sample disk = %v after %d walks, want null after 1 walk", ptr(third.Sites[0].DiskUsedBytes), len(walked)) + } +} + +func TestSitesDiskWalkThatFails(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home("example.com")}}) + c.walker.walk = func(context.Context, string) (uint64, int, error) { return 0, 0, context.DeadlineExceeded } + + if _, err := c.Sample(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + waitWalk(t, c) + s, err := c.Sample(time.Now().Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.Sites[0].DiskUsedBytes != nil { + t.Errorf("disk = %d, want null after a walk that stopped", *s.Sites[0].DiskUsedBytes) + } +} + +func TestSitesLeaveOutAContainerThatRestartedWithTheSameID(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1000000, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home("example.com")}}) + base := time.Now().Add(time.Second) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + + // docker restart keeps the id, but makes a new cgroup whose CPU time + // starts again. The new run already used more than the old one. + scope := filepath.Join(srv.root, "sys/fs/cgroup/system.slice/docker-a1.scope") + if err := os.Rename(scope, scope+".old"); err != nil { + t.Fatal(err) + } + srv.cgroup("a1", 1600000, 1, 0) + + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.Sites[0].CPUPercent != nil { + t.Errorf("cpu = %v, want null: the container restarted in the minute", *s.Sites[0].CPUPercent) + } +} + +func TestSitesWithTheRootAsHome(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", "/"}, {"a2", "/srv"}}) + c.home = "/" + s, err := c.Sample(time.Now().Add(time.Second)) + if err != nil { + t.Fatal(err) + } + if s.Sites != nil { + t.Errorf("sites = %+v, want null: with / as the home, no folder is a site", s.Sites) + } +} + +func TestSitesDiskResultSurvivesASampleThatFails(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home("example.com")}}) + base := time.Now().Add(time.Second) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + waitWalk(t, c) + + // The tick that would carry the disk use fails. + statfs := c.statfs + c.statfs = func(string) (uint64, uint64, error) { return 0, 0, errors.New("statfs failed") } + if _, err := c.Sample(base.Add(time.Minute)); err == nil { + t.Fatal("Sample() = nil error, want the statfs error") + } + c.statfs = statfs + + s, err := c.Sample(base.Add(2 * time.Minute)) + if err != nil { + t.Fatal(err) + } + if d := s.Sites[0].DiskUsedBytes; d == nil || *d != 4096 { + t.Errorf("disk = %v, want 4096 in the next sample that goes", ptr(d)) + } +} + +func TestSitesDiskWalkThatHangs(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1, 1, 0) + srv.cgroup("b1", 1, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home("a.com")}, {"b1", srv.home("b.com")}}) + // The walk of a.com hangs in a system call and does not see its + // context. The walk of b.com works. + hang := make(chan struct{}) + t.Cleanup(func() { close(hang) }) + c.walker.timeout = 50 * time.Millisecond + c.walker.walk = func(_ context.Context, root string) (uint64, int, error) { + if strings.HasSuffix(root, "a.com") { + <-hang + } + return 4096, 0, nil + } + + if _, err := c.Sample(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + waitWalk(t, c) + s, err := c.Sample(time.Now().Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.Sites[0].DiskUsedBytes != nil || s.Sites[1].DiskUsedBytes == nil { + t.Errorf("disk = %v, %v; want null for a.com and 4096 for b.com", ptr(s.Sites[0].DiskUsedBytes), ptr(s.Sites[1].DiskUsedBytes)) + } + c.walker.mu.Lock() + running := c.walker.running + c.walker.mu.Unlock() + if running { + t.Error("the walk still runs: a walk that hangs must not stop the next walks") + } +} + +func TestSitesWarnOneTimeForFilesThatCannotBeRead(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home(".fly")}}) + rec := &levels{} + c.log = slog.New(rec) + c.walker.walk = func(context.Context, string) (uint64, int, error) { return 4096, 9, nil } + + for i := range 3 { + c.walker.mu.Lock() + c.walker.last = time.Time{} // the hour passed + c.walker.mu.Unlock() + if _, err := c.Sample(time.Now().Add(time.Duration(i+1) * time.Minute)); err != nil { + t.Fatal(err) + } + waitWalk(t, c) + } + + msg := "some files of a site cannot be read; its disk use is lower than the real use" + if got := rec.count(slog.LevelWarn, msg); got != 1 { + t.Errorf("%d warnings for 3 walks, want 1", got) + } + if got := rec.count(slog.LevelDebug, msg); got != 2 { + t.Errorf("%d debug lines for 3 walks, want 2", got) + } +} + +// levels is a slog handler that keeps the level and the message of each +// record. +type levels struct { + mu sync.Mutex + records []slog.Record +} + +func (l *levels) Enabled(context.Context, slog.Level) bool { return true } +func (l *levels) WithAttrs([]slog.Attr) slog.Handler { return l } +func (l *levels) WithGroup(string) slog.Handler { return l } + +func (l *levels) Handle(_ context.Context, r slog.Record) error { + l.mu.Lock() + defer l.mu.Unlock() + l.records = append(l.records, r) + return nil +} + +func (l *levels) count(level slog.Level, msg string) int { + l.mu.Lock() + defer l.mu.Unlock() + n := 0 + for _, r := range l.records { + if r.Level == level && r.Message == msg { + n++ + } + } + return n +} + +func TestSitesHaveCPUAfterAShortFirstMinute(t *testing.T) { + srv := newServer(t) + srv.cgroup("a1", 1000000, 1, 0) + c := sitesCollector(t, srv, []fakeContainer{{"a1", srv.home("example.com")}}) + c.minFirst = minFirstMinute + + first := time.Now().Add(time.Second) + if _, err := c.Sample(first); !errors.Is(err, agent.ErrNoSample) { + t.Fatalf("Sample() 1 s after the start = %v, want ErrNoSample", err) + } + srv.cgroup("a1", 1600000, 1, 0) + s, err := c.Sample(first.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if len(s.Sites) != 1 || s.Sites[0].CPUPercent == nil || !near(*s.Sites[0].CPUPercent, 1) { + t.Errorf("sites = %+v, want example.com with 1%%: the short minute starts the CPU times too", s.Sites) + } +} + +// waitWalk waits until the disk walk that runs ends. +func waitWalk(t *testing.T, c *Collector) { + t.Helper() + c.walker.mu.Lock() + done := c.walker.done + c.walker.mu.Unlock() + if done == nil { + return + } + select { + case <-done: + case <-time.After(10 * time.Second): + t.Fatal("the disk walk did not end") + } +} + +func siteNames(sites []wire.Site) string { + var names []string + for _, s := range sites { + names = append(names, s.Directory) + } + return fmt.Sprint(names) +} diff --git a/internal/metrics/sys_linux.go b/internal/metrics/sys_linux.go new file mode 100644 index 0000000..b5dd8ec --- /dev/null +++ b/internal/metrics/sys_linux.go @@ -0,0 +1,51 @@ +package metrics + +import ( + "runtime" + + "golang.org/x/sys/unix" +) + +// statfs returns the size and the used space of the file system that holds +// path. Used space is (blocks − free blocks): the space that is reserved for +// root is used space, because the sites cannot use it. +func statfs(path string) (total, used uint64, err error) { + var st unix.Statfs_t + if err := unix.Statfs(path, &st); err != nil { + return 0, 0, err + } + + // The block counts are in units of the fragment size (f_frsize), the + // same as df uses. + size := uint64(st.Frsize) + if size == 0 { + size = uint64(st.Bsize) + } + + return st.Blocks * size, (st.Blocks - st.Bfree) * size, nil +} + +// kernelRelease returns the kernel release, as "uname -r" shows it. +func kernelRelease() string { + var u unix.Utsname + if err := unix.Uname(&u); err != nil { + return "" + } + + return unix.ByteSliceToString(u.Release[:]) +} + +// lowerIOPriority gives the calling goroutine the idle I/O class: its disk +// reads wait for all other reads. The class is a property of the thread, so +// the goroutine stays locked to its thread, and the thread ends with the +// goroutine. Without the right to set it, the priority does not change. +func lowerIOPriority() { + runtime.LockOSThread() + + const ( + whoProcess = 1 // IOPRIO_WHO_PROCESS: with id 0, the calling thread + classIdle = 3 // IOPRIO_CLASS_IDLE + classShift = 13 + ) + _, _, _ = unix.Syscall(unix.SYS_IOPRIO_SET, whoProcess, 0, classIdle< 0 && r.At.Sub(out[n-1].At) < minWindow { + continue + } + if cur.At.Sub(r.At) < minWindow { + continue + } + out = append(out, r) + } + + return append(out, cur) +} + +// setPeaks sets the peaks of the minute in s, the pressure and the disk +// activity. A peak is the highest value of the windows from one reading to +// the next (contract v0.4.0). s already holds the values of the minute, from +// prev to cur. A peak is never less than the value of its minute. +func setPeaks(s *wire.Sample, prev *reading, readings []reading, cur reading) { + all := chain(prev, readings, cur) + + // The memory peaks come from the readings of the minute: not from the + // reading of the tick before, and not from an older minute whose tick + // failed. + var memMax, swapMax uint64 + for _, r := range all { + if cur.At.Sub(r.At) >= time.Minute { + continue + } + memMax = max(memMax, r.mem.total-min(r.mem.available, r.mem.total)) + swapMax = max(swapMax, r.mem.swapTotal-min(r.mem.swapFree, r.mem.swapTotal)) + } + memMax = min(max(memMax, s.MemoryUsedBytes), s.MemoryTotalBytes) + swapMax = min(max(swapMax, s.SwapUsedBytes), s.SwapTotalBytes) + s.MemoryUsedMaxBytes, s.SwapUsedMaxBytes = &memMax, &swapMax + + if cpu, ok := cpuPeak(all); ok { + cpu = max(cpu, s.CPUPercent) + s.CPUMaxPercent = &cpu + } + + setPressure(s, prev, all, cur) + setDiskActivity(s, prev, all, cur) + + if !s.NetCountersReset { + if in, out, ok := netPeaks(all); ok { + // The control plane reads the value of the minute as bytes / 60. + in, out = max(in, s.NetInBytes/60), max(out, s.NetOutBytes/60) + s.NetInMaxBytesPerSecond, s.NetOutMaxBytesPerSecond = &in, &out + } + } +} + +// cpuPeak returns the busy share of the busiest window. A window across a +// reboot, with counters that went back, or of an older minute whose tick +// failed, is left out. +func cpuPeak(all []reading) (peak float64, ok bool) { + cur := all[len(all)-1] + for i := 1; i < len(all); i++ { + a, b := all[i-1], all[i] + if cur.At.Sub(b.At) >= time.Minute { + continue + } + if a.BootID != b.BootID || b.CPU.Total <= a.CPU.Total || b.CPU.Idle < a.CPU.Idle { + continue + } + peak, ok = max(peak, cpuPercent(a.CPU, b.CPU)), true + } + + return peak, ok +} + +// netPeaks returns the received and sent bytes each second of the busiest +// windows, rounded down. It returns false when the traffic of a window is not +// known, for example after a counter went back. +func netPeaks(all []reading) (in, out uint64, ok bool) { + for i := 1; i < len(all); i++ { + a, b := all[i-1], all[i] + dIn, dOut, reset := netDelta(&a, b) + if reset { + return 0, 0, false + } + secs := b.At.Sub(a.At).Seconds() + in = max(in, uint64(float64(dIn)/secs)) + out = max(out, uint64(float64(dOut)/secs)) + ok = true + } + + return in, out, ok +} diff --git a/internal/metrics/windows_test.go b/internal/metrics/windows_test.go new file mode 100644 index 0000000..32f9e7d --- /dev/null +++ b/internal/metrics/windows_test.go @@ -0,0 +1,341 @@ +package metrics + +import ( + "fmt" + "os" + "path/filepath" + "testing" + "time" + + "github.com/flywp/server-cli/internal/agent/wire" +) + +// mem writes /proc/meminfo with the available memory and the free swap, in +// kB. The totals are 8000000 kB and 2000000 kB. +func (s *server) mem(availableKB, swapFreeKB uint64) { + s.write("proc/meminfo", fmt.Sprintf("MemTotal: 8000000 kB\nMemFree: 500000 kB\nMemAvailable: %d kB\nSwapTotal: 2000000 kB\nSwapFree: %d kB\n", availableKB, swapFreeKB)) +} + +// step is one step of a test minute: the counters that the server shows at +// a reading. +type step struct { + total, idle uint64 // CPU ticks + in, out uint64 // eth0 bytes + available uint64 // kB + swapFree uint64 // kB +} + +func (s *server) set(st step) { + s.cpu(st.total, st.idle) + s.net(map[string][2]uint64{"eth0": {st.in, st.out}, "eth1": {100, 50}}) + s.mem(st.available, st.swapFree) +} + +// runMinute takes a sample at base, a reading each 10 seconds, and the +// sample of the next tick at base + 60 s. steps holds the counters of the +// five readings and of the tick. +func runMinute(t *testing.T, srv *server, c *Collector, base time.Time, first step, steps [6]step) wire.Sample { + t.Helper() + + srv.set(first) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + for i := range 5 { + srv.set(steps[i]) + c.Read(base.Add(time.Duration(i+1) * 10 * time.Second)) + } + srv.set(steps[5]) + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + return s +} + +func TestPeaks(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + // Each window adds 100 CPU ticks. The window from 20 s to 30 s is 90% + // busy; the others are 10% busy. eth0 receives 10000 bytes in the + // window from 40 s to 50 s, and 1000 bytes in each other window. The + // memory peaks at 30 s, and the swap at 40 s. + first := step{1000, 800, 0, 0, 6000000, 1500000} + s := runMinute(t, srv, c, base, first, [6]step{ + {1100, 890, 1000, 100, 6000000, 1500000}, + {1200, 980, 2000, 200, 5000000, 1500000}, + {1300, 990, 3000, 300, 4000000, 1500000}, + {1400, 1080, 4000, 400, 5000000, 1000000}, + {1500, 1170, 14000, 500, 6000000, 1500000}, + {1600, 1260, 15000, 600, 6000000, 1500000}, + }) + + // The minute: 600 ticks, 460 idle. + if want := float64(140) / 600 * 100; s.CPUPercent != want { + t.Errorf("cpu_percent = %v, want %v", s.CPUPercent, want) + } + if s.CPUMaxPercent == nil || *s.CPUMaxPercent != 90 { + t.Errorf("cpu_max_percent = %v, want 90", ptr(s.CPUMaxPercent)) + } + if want := uint64(4000000 * 1024); s.MemoryUsedMaxBytes == nil || *s.MemoryUsedMaxBytes != want { + t.Errorf("memory_used_max_bytes = %v, want %d", ptr(s.MemoryUsedMaxBytes), want) + } + if want := uint64(1000000 * 1024); s.SwapUsedMaxBytes == nil || *s.SwapUsedMaxBytes != want { + t.Errorf("swap_used_max_bytes = %v, want %d", ptr(s.SwapUsedMaxBytes), want) + } + if s.NetInMaxBytesPerSecond == nil || *s.NetInMaxBytesPerSecond != 1000 { + t.Errorf("net_in_max_bytes_per_second = %v, want 1000", ptr(s.NetInMaxBytesPerSecond)) + } + if s.NetOutMaxBytesPerSecond == nil || *s.NetOutMaxBytesPerSecond != 10 { + t.Errorf("net_out_max_bytes_per_second = %v, want 10", ptr(s.NetOutMaxBytesPerSecond)) + } + checkOrder(t, s) +} + +func TestAFailedReadingJoinsTwoWindows(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + srv.set(step{1000, 800, 0, 0, 6000000, 1500000}) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + + // The reading at 10 s fails. The window from 0 s to 20 s has 190 idle + // ticks of 200. + stat := filepath.Join(srv.root, "proc/stat") + if err := os.Remove(stat); err != nil { + t.Fatal(err) + } + c.Read(base.Add(10 * time.Second)) + srv.set(step{1200, 990, 0, 0, 6000000, 1500000}) + c.Read(base.Add(20 * time.Second)) + srv.set(step{1300, 1000, 0, 0, 6000000, 1500000}) + + s, err := c.Sample(base.Add(30 * time.Second)) + if err != nil { + t.Fatal(err) + } + if s.CPUMaxPercent == nil || *s.CPUMaxPercent != 90 { + t.Errorf("cpu_max_percent = %v, want 90 from the window of 20 s to 30 s", ptr(s.CPUMaxPercent)) + } + if s.NetInMaxBytesPerSecond == nil || *s.NetInMaxBytesPerSecond != 0 { + t.Errorf("net_in_max_bytes_per_second = %v, want 0", ptr(s.NetInMaxBytesPerSecond)) + } +} + +func TestFirstSamplePeaksEqualTheMinute(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + + // The agent started just now: one window, from the start to the tick. + srv.set(step{1600, 950, 5000, 5000, 5000000, 1500000}) + s, err := c.Sample(time.Now().Add(20 * time.Second)) + if err != nil { + t.Fatal(err) + } + if s.CPUMaxPercent == nil || *s.CPUMaxPercent != s.CPUPercent || s.CPUPercent != 75 { + t.Errorf("cpu = %v, max %v; want 75 and an equal peak", s.CPUPercent, ptr(s.CPUMaxPercent)) + } + if s.MemoryUsedMaxBytes == nil || *s.MemoryUsedMaxBytes != s.MemoryUsedBytes { + t.Errorf("memory_used_max_bytes = %v, want the value of the minute %d", ptr(s.MemoryUsedMaxBytes), s.MemoryUsedBytes) + } + // The traffic since the start is not known, so its peak is not known. + if !s.NetCountersReset || s.NetInMaxBytesPerSecond != nil || s.NetOutMaxBytesPerSecond != nil { + t.Errorf("net peaks = %v, %v with reset %v; want null", ptr(s.NetInMaxBytesPerSecond), ptr(s.NetOutMaxBytesPerSecond), s.NetCountersReset) + } +} + +func TestPeaksAfterARestartIncludeTheSavedReading(t *testing.T) { + srv := newServer(t) + state := t.TempDir() + now := time.Now() + + // The last tick of the old process was 30 s ago. + srv.set(step{1000, 800, 0, 0, 6000000, 1500000}) + if _, err := srv.collector(state).Sample(now.Add(-30 * time.Second)); err != nil { + t.Fatal(err) + } + + // The new process starts after 100% busy time. Then it takes the tick. + srv.set(step{1100, 800, 3000, 0, 6000000, 1500000}) + c := srv.collector(state) + srv.set(step{1600, 1250, 3000, 0, 6000000, 1500000}) + s, err := c.Sample(now.Add(30 * time.Second)) + if err != nil { + t.Fatal(err) + } + if s.CPUPercent != 25 { + t.Errorf("cpu_percent = %v, want 25 from the saved reading", s.CPUPercent) + } + if s.CPUMaxPercent == nil || *s.CPUMaxPercent != 100 { + t.Errorf("cpu_max_percent = %v, want 100 from the window of the saved reading to the start", ptr(s.CPUMaxPercent)) + } + if s.NetCountersReset || s.NetInBytes != 3000 || s.NetInMaxBytesPerSecond == nil || *s.NetInMaxBytesPerSecond < 99 { + t.Errorf("net = %d, peak %v, reset %v; want 3000 with a peak of about 100 each second (3000 bytes in 30 s)", s.NetInBytes, ptr(s.NetInMaxBytesPerSecond), s.NetCountersReset) + } + checkOrder(t, s) +} + +func TestNetPeaksAreNullAfterACounterWentBack(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + // eth0 goes back at 30 s and forward again by the tick: the minute + // looks correct, but a window is not known. + s := runMinute(t, srv, c, base, step{1000, 800, 5000, 5000, 6000000, 1500000}, [6]step{ + {1100, 890, 6000, 6000, 6000000, 1500000}, + {1200, 980, 7000, 7000, 6000000, 1500000}, + {1300, 990, 10, 10, 6000000, 1500000}, + {1400, 1080, 1000, 1000, 6000000, 1500000}, + {1500, 1170, 6000, 6000, 6000000, 1500000}, + {1600, 1260, 9000, 9000, 6000000, 1500000}, + }) + if s.NetInMaxBytesPerSecond != nil || s.NetOutMaxBytesPerSecond != nil { + t.Errorf("net peaks = %v, %v; want null after a window with a counter that went back", ptr(s.NetInMaxBytesPerSecond), ptr(s.NetOutMaxBytesPerSecond)) + } + if s.CPUMaxPercent == nil { + t.Error("cpu_max_percent = null, want a value: only the network is not known") + } +} + +func TestReadingsThatAreNotNewerAreIgnored(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + srv.set(step{1000, 800, 0, 0, 6000000, 1500000}) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + c.Read(base.Add(20 * time.Second)) + // The wall clock stepped back: the same reading comes again, and an + // older one. + c.Read(base.Add(20 * time.Second)) + c.Read(base.Add(10 * time.Second)) + if len(c.readings) != 1 { + t.Errorf("%d readings, want 1", len(c.readings)) + } + + // A reading close to the tick makes no window of some milliseconds. + c.Read(base.Add(59*time.Second + 900*time.Millisecond)) + srv.set(step{1100, 810, 0, 0, 6000000, 1500000}) + s, err := c.Sample(base.Add(time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.CPUMaxPercent == nil || *s.CPUMaxPercent != 90 { + t.Errorf("cpu_max_percent = %v, want 90 from the window of 20 s to 60 s", ptr(s.CPUMaxPercent)) + } +} + +func TestMemoryPeakAfterAFailedTickIsFromTheLastMinute(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + srv.set(step{1000, 800, 0, 0, 6000000, 1500000}) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + // A high use of memory at 30 s. The tick at 60 s fails. + srv.mem(1000000, 1500000) + c.Read(base.Add(30 * time.Second)) + srv.mem(6000000, 1500000) + c.Read(base.Add(50 * time.Second)) + stat := filepath.Join(srv.root, "proc/stat") + if err := os.Remove(stat); err != nil { + t.Fatal(err) + } + if _, err := c.Sample(base.Add(time.Minute)); err == nil { + t.Fatal("Sample() = nil error, want the error of the tick") + } + + c.Read(base.Add(70 * time.Second)) + srv.set(step{1100, 900, 0, 0, 6000000, 1500000}) + s, err := c.Sample(base.Add(2 * time.Minute)) + if err != nil { + t.Fatal(err) + } + if want := uint64(2000000 * 1024); s.MemoryUsedMaxBytes == nil || *s.MemoryUsedMaxBytes != want { + t.Errorf("memory_used_max_bytes = %v, want %d: the peak at 30 s is in the minute before", ptr(s.MemoryUsedMaxBytes), want) + } +} + +func TestCPUPeakAfterAFailedTickIsFromTheLastMinute(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + + srv.set(step{1000, 800, 0, 0, 6000000, 1500000}) + if _, err := c.Sample(base); err != nil { + t.Fatal(err) + } + // 100% busy from 0 s to 10 s. The tick at 60 s fails. + srv.set(step{1100, 800, 0, 0, 6000000, 1500000}) + c.Read(base.Add(10 * time.Second)) + if err := os.Remove(filepath.Join(srv.root, "proc/meminfo")); err != nil { + t.Fatal(err) + } + if _, err := c.Sample(base.Add(time.Minute)); err == nil { + t.Fatal("Sample() = nil error, want the error of the tick") + } + + // The next minute is idle. + srv.set(step{1100, 800, 0, 0, 6000000, 1500000}) + c.Read(base.Add(70 * time.Second)) + srv.set(step{1200, 900, 0, 0, 6000000, 1500000}) + s, err := c.Sample(base.Add(2 * time.Minute)) + if err != nil { + t.Fatal(err) + } + if s.CPUMaxPercent == nil || *s.CPUMaxPercent >= 100 { + t.Errorf("cpu_max_percent = %v, want the peak of the last minute, not the 100%% of the minute before", ptr(s.CPUMaxPercent)) + } + checkOrder(t, s) +} + +func TestReadingsAreLimited(t *testing.T) { + srv := newServer(t) + c := srv.collector(t.TempDir()) + base := time.Now().Add(time.Second) + for i := range 100 { + c.Read(base.Add(time.Duration(i) * 10 * time.Second)) + } + if len(c.readings) != maxReadings { + t.Errorf("%d readings, want at most %d", len(c.readings), maxReadings) + } +} + +// checkOrder checks the order that the contract guarantees between a value of +// the minute and its peak. +func checkOrder(t *testing.T, s wire.Sample) { + t.Helper() + if s.CPUMaxPercent != nil && s.CPUPercent > *s.CPUMaxPercent { + t.Errorf("cpu_percent %v > cpu_max_percent %v", s.CPUPercent, *s.CPUMaxPercent) + } + if m := s.MemoryUsedMaxBytes; m == nil || s.MemoryUsedBytes > *m || *m > s.MemoryTotalBytes { + t.Errorf("memory %d, max %v, total %d: want used ≤ max ≤ total", s.MemoryUsedBytes, ptr(m), s.MemoryTotalBytes) + } + if m := s.SwapUsedMaxBytes; m == nil || s.SwapUsedBytes > *m || *m > s.SwapTotalBytes { + t.Errorf("swap %d, max %v, total %d: want used ≤ max ≤ total", s.SwapUsedBytes, ptr(m), s.SwapTotalBytes) + } + if m := s.NetInMaxBytesPerSecond; m != nil && s.NetInBytes > 60**m+59 { + t.Errorf("net_in_bytes %d > 60 × %d", s.NetInBytes, *m) + } + if m := s.NetOutMaxBytesPerSecond; m != nil && s.NetOutBytes > 60**m+59 { + t.Errorf("net_out_bytes %d > 60 × %d", s.NetOutBytes, *m) + } +} + +// ptr shows a value that can be null. +func ptr[T any](v *T) any { + if v == nil { + return "null" + } + return *v +} 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..4b05a2c --- /dev/null +++ b/internal/release/download_test.go @@ -0,0 +1,227 @@ +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) + } +} + +// testRelease serves a release with the archive and a checksum file, and +// returns its description. +func testRelease(t *testing.T, archive []byte, checksums string) *GithubRelease { + t.Helper() + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/fly-linux-amd64.tar.gz": + _, _ = w.Write(archive) + case "/checksums.txt": + _, _ = w.Write([]byte(checksums)) + default: + http.NotFound(w, r) + } + })) + t.Cleanup(srv.Close) + + rel := &GithubRelease{TagName: "v0.3.0"} + for _, name := range []string{"fly-linux-amd64.tar.gz", "checksums.txt"} { + if name == "checksums.txt" && checksums == "" { + continue + } + rel.Assets = append(rel.Assets, Asset{Name: name, BrowserDownloadURL: srv.URL + "/" + name}) + } + return rel +} + +func TestSelfUpdateChecksTheArchive(t *testing.T) { + data := archive(t, map[string]string{"fly-linux-amd64": "new"}).Bytes() + other := sum([]byte("other")) + + tests := []struct { + name, checksums, want string + }{ + {"checksum agrees", sum(data) + " fly-linux-arm64.tar.gz\n" + sum(data) + " fly-linux-amd64.tar.gz\n", ""}, + {"binary mode mark", sum(data) + " *fly-linux-amd64.tar.gz\n", ""}, + {"checksum does not agree", other + " fly-linux-amd64.tar.gz\n", "the sha256 of"}, + {"no line for the archive", other + " fly-linux-arm64.tar.gz\n", "has no line for fly-linux-amd64.tar.gz"}, + {"no checksum file", "", "has no checksums.txt"}, + {"CRLF line ends", sum(data) + " fly-linux-amd64.tar.gz\r\n", ""}, + {"upper case hex", strings.ToUpper(sum(data)) + " fly-linux-amd64.tar.gz\n", ""}, + {"the same line two times", sum(data) + " fly-linux-amd64.tar.gz\n" + sum(data) + " fly-linux-amd64.tar.gz\n", ""}, + {"two different sums", sum(data) + " fly-linux-amd64.tar.gz\n" + other + " fly-linux-amd64.tar.gz\n", "two different sums"}, + {"a similar name", sum(data) + " fly-linux-amd64.tar.gz.sig\n", "has no line for fly-linux-amd64.tar.gz"}, + {"empty checksum file", "\n", "has no line for fly-linux-amd64.tar.gz"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + exe := filepath.Join(dir, "fly") + if err := os.WriteFile(exe, []byte("old"), 0o755); err != nil { + t.Fatal(err) + } + + err := selfUpdate(context.Background(), testRelease(t, data, tt.checksums), exe, "linux", "amd64") + + got, _ := os.ReadFile(exe) + if tt.want == "" { + if err != nil || string(got) != "new" { + t.Fatalf("selfUpdate() = %v, binary %q; want the new binary", err, got) + } + if entries, _ := os.ReadDir(dir); len(entries) != 1 { + t.Errorf("the directory holds %d files after the update, want only the binary", len(entries)) + } + return + } + if err == nil || !strings.Contains(err.Error(), tt.want) { + t.Fatalf("selfUpdate() error = %v, want %q", err, tt.want) + } + if string(got) != "old" { + t.Errorf("binary = %q, want the old binary unchanged", got) + } + if entries, _ := os.ReadDir(dir); len(entries) != 1 { + t.Errorf("the directory holds %d files, want only the binary", len(entries)) + } + }) + } +} diff --git a/internal/release/keys.go b/internal/release/keys.go new file mode 100644 index 0000000..f35c5ca --- /dev/null +++ b/internal/release/keys.go @@ -0,0 +1,19 @@ +package release + +import "crypto/ed25519" + +// trustedKeys are the public keys of the release signatures, in the format of +// ParseKeys. The private keys are kept outside GitHub: see "make release-key" +// and "make sign-release". Two keys can be listed while one replaces the other. +// +// Give each key a comment that names it (the COMMENT of make release-key). +// +// It is a string, not a map, so that a test build can set it with +// -ldflags "-X github.com/flywp/server-cli/internal/release.trustedKeys=...". +var trustedKeys = "21ec14c790e96d62:pyMN2C1s+lG2xXUTYeoXCgYtE1rtCh3rgYlYsBRySGk=" // server-cli release key for flywp + +// TrustedKeys returns the public keys that the agent accepts for a release +// signature. An empty map means that no release can install by itself. +func TrustedKeys() (map[string]ed25519.PublicKey, error) { + return ParseKeys(trustedKeys) +} diff --git a/internal/release/release.go b/internal/release/release.go new file mode 100644 index 0000000..6adfb0a --- /dev/null +++ b/internal/release/release.go @@ -0,0 +1,417 @@ +// Package release finds, downloads, checks and installs fly releases. "fly +// update" and the agent command agent.update use it. +package release + +import ( + "archive/tar" + "compress/gzip" + "context" + "crypto/ed25519" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "os" + "path" + "path/filepath" + "regexp" + "runtime" + "strings" + "syscall" + "time" + + "github.com/flywp/server-cli/internal/version" + "golang.org/x/mod/semver" +) + +// GithubAPI is the GitHub API URL of the latest release. +var GithubAPI = "https://api.github.com/repos/flywp/server-cli/releases/latest" + +// httpClient limits each request, including the download of the binary. +var httpClient = &http.Client{Timeout: 60 * time.Second} + +// maxBinarySize limits the size of the binary in a release archive. +const maxBinarySize = 200 << 20 + +// GithubRelease is a release in the GitHub API. +type GithubRelease struct { + TagName string `json:"tag_name"` + Assets []Asset `json:"assets"` +} + +// Asset is a file of a release. +type Asset struct { + Name string `json:"name"` + BrowserDownloadURL string `json:"browser_download_url"` +} + +// Update compares the latest release with the running version. +type Update struct { + Release *GithubRelease + // Available is true when the release is newer than the running version. + Available bool + // Comparable is false when the running version is not built from a + // release tag, for example "dev". + Comparable bool +} + +// CheckForUpdates gets the latest release from GitHub and compares it with +// the running version. +func CheckForUpdates(ctx context.Context) (*Update, error) { + release, err := LatestRelease(ctx) + if err != nil { + return nil, err + } + + available, comparable := isNewer(release.TagName, version.Version) + return &Update{Release: release, Available: available, Comparable: comparable}, nil +} + +// describeSuffix matches what git describe adds after a tag: the number of +// commits since the tag, the commit hash and "-dirty" for local changes. +var describeSuffix = regexp.MustCompile(`(-\d+-g[0-9a-f]+)?(-dirty)?$`) + +// IsLocalBuild reports whether v comes from git describe on a commit that is +// not a release tag, for example v0.2.0-3-gabcdef1 or v0.2.0-dirty. Such a +// build has code that no release has, so it must not update by itself. +func IsLocalBuild(v string) bool { + return describeSuffix.FindString(v) != "" +} + +// isNewer reports whether the release version latest is newer than current. +// comparable is false when current is not built from a release tag. +func isNewer(latest, current string) (newer, comparable bool) { + base := describeSuffix.ReplaceAllString(current, "") + if !semver.IsValid(base) { + return false, false + } + + return semver.Compare(latest, base) > 0, true +} + +// LatestRelease returns the latest release from GitHub. +func LatestRelease(ctx context.Context) (*GithubRelease, error) { + resp, err := get(ctx, GithubAPI) + if err != nil { + return nil, err + } + defer func() { _ = resp.Body.Close() }() + + // The agent reads this each day without a person: a huge reply must not + // use all of its memory. + var release GithubRelease + if err := json.NewDecoder(io.LimitReader(resp.Body, maxAssetSize)).Decode(&release); err != nil { + return nil, fmt.Errorf("reading release information: %w", err) + } + + if !semver.IsValid(release.TagName) { + return nil, fmt.Errorf("latest release has an invalid version %q", release.TagName) + } + + return &release, nil +} + +// get sends a GET request and returns the response if its status is 200 OK. +func get(ctx context.Context, url string) (*http.Response, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, err + } + req.Header.Set("Accept", "application/vnd.github+json") + req.Header.Set("User-Agent", "fly-cli/"+version.Version) + + resp, err := httpClient.Do(req) + if err != nil { + return nil, err + } + + if resp.StatusCode != http.StatusOK { + _ = resp.Body.Close() + + limited := resp.StatusCode == http.StatusForbidden || resp.StatusCode == http.StatusTooManyRequests + if limited && resp.Header.Get("X-RateLimit-Remaining") == "0" { + return nil, errors.New("the GitHub API rate limit is exceeded, try again later") + } + return nil, fmt.Errorf("unexpected response from %s: %s", url, resp.Status) + } + + return resp, nil +} + +// ChecksumsAsset is the checksum file of each release: one line for each +// archive, " " (the output of sha256sum). +const ChecksumsAsset = "checksums.txt" + +// selfUpdateTimeout limits the download of an update. +const selfUpdateTimeout = 10 * time.Minute + +// SelfUpdate replaces the running binary with the binary from release. It +// installs the archive only when its sha256 agrees with the checksum file of +// the release. +func SelfUpdate(ctx context.Context, release *GithubRelease) error { + exe, err := os.Executable() + if err != nil { + return fmt.Errorf("finding the current executable: %w", err) + } + exe, err = filepath.EvalSymlinks(exe) + if err != nil { + return fmt.Errorf("resolving symlinks: %w", err) + } + + return selfUpdate(ctx, release, exe, runtime.GOOS, runtime.GOARCH) +} + +func selfUpdate(ctx context.Context, release *GithubRelease, exe, goos, goarch string) error { + ctx, cancel := context.WithTimeout(ctx, selfUpdateTimeout) + defer cancel() + + archiveURL := assetURL(release, goos, goarch) + if archiveURL == "" { + return fmt.Errorf("no suitable binary found for this system (OS: %s, ARCH: %s)", goos, goarch) + } + + name := BinaryName(goos, goarch) + sum, err := checksum(ctx, release, name+".tar.gz") + if err != nil { + return err + } + + archive, err := Download(ctx, archiveURL, sum, filepath.Dir(exe)) + if err != nil { + return err + } + defer func() { _ = os.Remove(archive) }() + + return Install(archive, exe, name) +} + +// checksum returns the sha256 of the release file name from the checksum +// file of the release. +func checksum(ctx context.Context, release *GithubRelease, name string) (string, error) { + data, err := fetchAsset(ctx, release, ChecksumsAsset) + if err != nil { + return "", err + } + + return parseChecksums(data, release.TagName, name) +} + +// ErrNotSigned is the error of SignedChecksum for a release without a +// signature file: the maintainer did not sign it yet. +var ErrNotSigned = errors.New("the release has no signature") + +// Signed is a release archive whose sha256 comes from a signed checksum file. +type Signed struct { + URL string + SHA256 string + SignedAt time.Time +} + +// SignedChecksum returns the archive of release for goos and goarch, with its +// sha256 from the checksum file, after it checks the signature of that file +// with keys. The sum comes from the same bytes that the signature covers. +func SignedChecksum(ctx context.Context, release *GithubRelease, goos, goarch string, keys map[string]ed25519.PublicKey, now time.Time) (*Signed, error) { + archiveURL := assetURL(release, goos, goarch) + if archiveURL == "" { + return nil, fmt.Errorf("release %s has no binary for %s/%s", release.TagName, goos, goarch) + } + if asset(release, SignatureAsset) == "" { + return nil, ErrNotSigned + } + + sigFile, err := fetchAsset(ctx, release, SignatureAsset) + if err != nil { + return nil, err + } + data, err := fetchAsset(ctx, release, ChecksumsAsset) + if err != nil { + return nil, err + } + + sig, err := Verify(sigFile, data, release.TagName, keys, now) + if err != nil { + return nil, err + } + sum, err := parseChecksums(data, release.TagName, BinaryName(goos, goarch)+".tar.gz") + if err != nil { + return nil, err + } + + return &Signed{URL: archiveURL, SHA256: sum, SignedAt: sig.SignedAt}, nil +} + +// maxAssetSize limits a small release file: the checksum file and its +// signature. +const maxAssetSize = 1 << 20 + +// fetchAsset downloads the small release file name. +func fetchAsset(ctx context.Context, release *GithubRelease, name string) ([]byte, error) { + url := asset(release, name) + if url == "" { + return nil, fmt.Errorf("release %s has no %s, so its download cannot be checked", release.TagName, name) + } + + resp, err := get(ctx, url) + if err != nil { + return nil, fmt.Errorf("downloading %s: %w", name, err) + } + defer func() { _ = resp.Body.Close() }() + + data, err := io.ReadAll(io.LimitReader(resp.Body, maxAssetSize)) + if err != nil { + return nil, fmt.Errorf("downloading %s: %w", name, err) + } + + return data, nil +} + +// parseChecksums returns the sha256 of the file name from the checksum file +// data of the release tag. +func parseChecksums(data []byte, tag, name string) (string, error) { + // install.sh reads the file with the same rules. Two different sums for + // one file make the file not valid: it is not clear which one is correct. + var sum string + for line := range strings.Lines(string(data)) { + // sha256sum marks a file that it read in binary mode with "*". + fields := strings.Fields(line) + if len(fields) != 2 || strings.TrimPrefix(fields[1], "*") != name { + continue + } + if sum != "" && !strings.EqualFold(sum, fields[0]) { + return "", fmt.Errorf("%s of release %s has two different sums for %s", ChecksumsAsset, tag, name) + } + sum = fields[0] + } + if sum == "" { + return "", fmt.Errorf("%s of release %s has no line for %s", ChecksumsAsset, tag, name) + } + + return sum, nil +} + +// 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 { + return fmt.Sprintf("fly-%s-%s", goos, goarch) +} + +// assetURL returns the download URL of the release archive for goos and +// goarch, or "" if the release has none. +func assetURL(release *GithubRelease, goos, goarch string) string { + if goos != "linux" { + return "" + } + + return asset(release, BinaryName(goos, goarch)+".tar.gz") +} + +// asset returns the download URL of the release file name, or "". +func asset(release *GithubRelease, name string) string { + for _, a := range release.Assets { + if a.Name == name { + return a.BrowserDownloadURL + } + } + + return "" +} + +// replaceBinary extracts the file name from the tar.gz archive and puts it in +// place of exe. The new binary is written to a temporary file in the same +// directory and then renamed, so exe is never incomplete and the rename +// does not cross filesystems. +func replaceBinary(exe string, archive io.Reader, name string) error { + gz, err := gzip.NewReader(archive) + if err != nil { + return fmt.Errorf("reading archive: %w", err) + } + defer func() { _ = gz.Close() }() + + tr := tar.NewReader(gz) + for { + hdr, err := tr.Next() + if errors.Is(err, io.EOF) { + return fmt.Errorf("archive does not contain %s", name) + } + if err != nil { + return fmt.Errorf("reading archive: %w", err) + } + + if hdr.Typeflag != tar.TypeReg || path.Clean(hdr.Name) != name { + continue + } + if hdr.Size > maxBinarySize { + return fmt.Errorf("%s in archive is too large (%d bytes)", name, hdr.Size) + } + + return writeBinary(exe, tr) + } +} + +// writeBinary writes r to a temporary file next to exe and renames it to exe. +func writeBinary(exe string, r io.Reader) (err error) { + tmp, err := os.CreateTemp(filepath.Dir(exe), ".fly-update-*") + if err != nil { + return fmt.Errorf("creating temporary file: %w", err) + } + defer func() { + if err != nil { + _ = os.Remove(tmp.Name()) + } + }() + + if _, err = io.Copy(tmp, r); err != nil { + _ = tmp.Close() + return fmt.Errorf("writing new binary: %w", err) + } + // 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) + } + if err = keepOwner(tmp, exe); err != nil { + _ = tmp.Close() + return fmt.Errorf("keeping the owner of the binary: %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 = 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 +} + +// keepOwner gives the open file f the owner of exe. "sudo fly update" then +// keeps the binary of the agent with the server user, not with root. Only +// root can give a file to a different user; for other users the owner is +// already correct. It changes the open file, not a path: a path in a +// directory of an other user can be replaced by a link to a root file. +func keepOwner(f *os.File, exe string) error { + info, err := os.Stat(exe) + if err != nil { + return nil + } + st, ok := info.Sys().(*syscall.Stat_t) + if !ok || os.Geteuid() != 0 { + return nil + } + + return f.Chown(int(st.Uid), int(st.Gid)) +} diff --git a/internal/release/release_test.go b/internal/release/release_test.go new file mode 100644 index 0000000..a4681bb --- /dev/null +++ b/internal/release/release_test.go @@ -0,0 +1,332 @@ +package release + +import ( + "archive/tar" + "bytes" + "compress/gzip" + "context" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "syscall" + "testing" +) + +func TestIsNewer(t *testing.T) { + tests := []struct { + latest, current string + newer, wantComparable bool + }{ + {latest: "v0.1.10", current: "v0.1.9", newer: true, wantComparable: true}, + {latest: "v0.1.9", current: "v0.1.10", newer: false, wantComparable: true}, + {latest: "v0.1.1", current: "v0.1.1", newer: false, wantComparable: true}, + {latest: "v0.2.0", current: "v0.1.1", newer: true, wantComparable: true}, + {latest: "v1.0.0", current: "v0.10.0", newer: true, wantComparable: true}, + // git describe versions: built after the tag, so the tag is not newer. + {latest: "v0.1.1", current: "v0.1.1-2-g3994ef6", newer: false, wantComparable: true}, + {latest: "v0.1.2", current: "v0.1.1-2-g3994ef6", newer: true, wantComparable: true}, + {latest: "v0.1.1", current: "v0.1.1-2-g3994ef6-dirty", newer: false, wantComparable: true}, + {latest: "v0.1.1", current: "v0.1.1-dirty", newer: false, wantComparable: true}, + {latest: "v0.2.0", current: "v0.2.0-rc.1", newer: true, wantComparable: true}, + // Not built from a release tag: cannot be compared. + {latest: "v0.2.0", current: "dev", newer: false, wantComparable: false}, + {latest: "v0.2.0", current: "3994ef6", newer: false, wantComparable: false}, + } + + for _, tt := range tests { + newer, comparable := isNewer(tt.latest, tt.current) + if newer != tt.newer || comparable != tt.wantComparable { + t.Errorf("isNewer(%q, %q) = %v, %v, want %v, %v", tt.latest, tt.current, newer, comparable, tt.newer, tt.wantComparable) + } + } +} + +// serveAPI points GithubAPI at a test server that runs handler. +func serveAPI(t *testing.T, handler http.HandlerFunc) { + t.Helper() + + srv := httptest.NewServer(handler) + t.Cleanup(srv.Close) + + old := GithubAPI + GithubAPI = srv.URL + t.Cleanup(func() { GithubAPI = old }) +} + +func TestLatestRelease(t *testing.T) { + var userAgent string + serveAPI(t, func(w http.ResponseWriter, r *http.Request) { + userAgent = r.Header.Get("User-Agent") + _, _ = w.Write([]byte(`{"tag_name":"v0.2.0","assets":[{"name":"fly-linux-amd64.tar.gz","browser_download_url":"https://example.com/a"}]}`)) + }) + + release, err := LatestRelease(context.Background()) + if err != nil { + t.Fatalf("LatestRelease() error = %v", err) + } + if release.TagName != "v0.2.0" || len(release.Assets) != 1 { + t.Errorf("LatestRelease() = %+v, want v0.2.0 with one asset", release) + } + if !strings.HasPrefix(userAgent, "fly-cli/") { + t.Errorf("User-Agent = %q, want fly-cli/", userAgent) + } +} + +func TestLatestReleaseErrors(t *testing.T) { + tests := []struct { + name string + handler http.HandlerFunc + want string + }{ + { + name: "rate limit", + handler: func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-RateLimit-Remaining", "0") + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(`{"message":"API rate limit exceeded"}`)) + }, + want: "rate limit", + }, + { + name: "server error", + handler: func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + }, + want: "500", + }, + { + name: "no tag", + handler: func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"message":"Not Found"}`)) + }, + want: "invalid version", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + serveAPI(t, tt.handler) + + // An error, never an empty "latest version". + if _, err := LatestRelease(context.Background()); err == nil || !strings.Contains(err.Error(), tt.want) { + t.Errorf("LatestRelease() error = %v, want an error that contains %q", err, tt.want) + } + }) + } +} + +func TestAssetURL(t *testing.T) { + release := &GithubRelease{TagName: "v0.2.0"} + for _, name := range []string{"fly-linux-amd64.tar.gz", "fly-linux-arm64.tar.gz"} { + release.Assets = append(release.Assets, Asset{Name: name, BrowserDownloadURL: "https://example.com/" + name}) + } + + tests := []struct{ goos, goarch, want string }{ + {"linux", "amd64", "https://example.com/fly-linux-amd64.tar.gz"}, + {"linux", "arm64", "https://example.com/fly-linux-arm64.tar.gz"}, + {"linux", "386", ""}, + {"darwin", "arm64", ""}, + } + for _, tt := range tests { + if got := assetURL(release, tt.goos, tt.goarch); got != tt.want { + t.Errorf("assetURL(%s/%s) = %q, want %q", tt.goos, tt.goarch, got, tt.want) + } + } +} + +// archive returns a tar.gz archive that contains files. +func archive(t *testing.T, files map[string]string) *bytes.Buffer { + t.Helper() + + var buf bytes.Buffer + gz := gzip.NewWriter(&buf) + tw := tar.NewWriter(gz) + for name, content := range files { + hdr := &tar.Header{Name: name, Mode: 0o755, Size: int64(len(content)), Typeflag: tar.TypeReg} + if err := tw.WriteHeader(hdr); err != nil { + t.Fatal(err) + } + if _, err := tw.Write([]byte(content)); 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 +} + +func TestReplaceBinary(t *testing.T) { + dir := t.TempDir() + exe := filepath.Join(dir, "fly") + if err := os.WriteFile(exe, []byte("old"), 0o755); err != nil { + t.Fatal(err) + } + + err := replaceBinary(exe, archive(t, map[string]string{"fly-linux-amd64": "new"}), "fly-linux-amd64") + if err != nil { + t.Fatalf("replaceBinary() error = %v", err) + } + + got, err := os.ReadFile(exe) + if err != nil || string(got) != "new" { + t.Errorf("binary = %q, %v, want %q", got, err, "new") + } + if info, err := os.Stat(exe); err != nil || info.Mode().Perm() != 0o755 { + t.Errorf("binary mode = %v, %v, want 0755", info.Mode().Perm(), err) + } + assertOnlyFile(t, dir, "fly") +} + +func TestReplaceBinaryKeepsOldBinaryOnError(t *testing.T) { + tests := []struct { + name string + archive *bytes.Buffer + }{ + {name: "binary not in archive", archive: archive(t, map[string]string{"README.md": "text"})}, + {name: "not an archive", archive: bytes.NewBufferString("Not Found")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := t.TempDir() + exe := filepath.Join(dir, "fly") + if err := os.WriteFile(exe, []byte("old"), 0o755); err != nil { + t.Fatal(err) + } + + if err := replaceBinary(exe, tt.archive, "fly-linux-amd64"); err == nil { + t.Fatal("replaceBinary() = nil, want an error") + } + + if got, _ := os.ReadFile(exe); string(got) != "old" { + t.Errorf("binary = %q, want the old binary unchanged", got) + } + assertOnlyFile(t, dir, "fly") + }) + } +} + +// assertOnlyFile fails the test if dir contains anything other than name, +// for example a temporary file that was not removed. +func assertOnlyFile(t *testing.T, dir, name string) { + t.Helper() + + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + if len(entries) != 1 || entries[0].Name() != name { + var names []string + for _, e := range entries { + names = append(names, e.Name()) + } + t.Errorf("directory contains %q, want only %q", names, name) + } +} + +func TestReplaceBinaryKeepsTheOwner(t *testing.T) { + if os.Geteuid() != 0 { + t.Skip("only root can give a file to a different user") + } + + dir := t.TempDir() + exe := filepath.Join(dir, "fly") + if err := os.WriteFile(exe, []byte("old"), 0o755); err != nil { + t.Fatal(err) + } + // The binary of the agent belongs to the server user, for example 1000. + if err := os.Chown(exe, 1000, 1000); err != nil { + t.Fatal(err) + } + + if err := replaceBinary(exe, archive(t, map[string]string{"fly-linux-amd64": "new"}), "fly-linux-amd64"); err != nil { + t.Fatal(err) + } + + info, err := os.Stat(exe) + if err != nil { + t.Fatal(err) + } + if st := info.Sys().(*syscall.Stat_t); st.Uid != 1000 || st.Gid != 1000 { + t.Errorf("owner = %d:%d, want 1000:1000", st.Uid, st.Gid) + } +} + +// swapReader is an archive that, halfway, puts a link to victim in place of +// the temporary file in dir, like a hostile owner of dir. +type swapReader struct { + t *testing.T + r io.Reader + dir string + victim string + swapped bool +} + +func (s *swapReader) Read(p []byte) (int, error) { + if !s.swapped { + s.swapped = true + matches, _ := filepath.Glob(filepath.Join(s.dir, ".fly-update-*")) + for _, m := range matches { + if err := os.Rename(m, m+".moved"); err != nil { + s.t.Fatal(err) + } + if err := os.Symlink(s.victim, m); err != nil { + s.t.Fatal(err) + } + } + } + return s.r.Read(p) +} + +func TestReplaceBinaryDoesNotFollowALinkToAnOtherFile(t *testing.T) { + if os.Geteuid() != 0 { + t.Skip("the attack needs root to write the binary") + } + + dir := t.TempDir() + exe := filepath.Join(dir, "fly") + if err := os.WriteFile(exe, []byte("old"), 0o755); err != nil { + t.Fatal(err) + } + if err := os.Chown(exe, 1000, 1000); err != nil { + t.Fatal(err) + } + // A root file, for example /etc/shadow. + victim := filepath.Join(t.TempDir(), "shadow") + if err := os.WriteFile(victim, []byte("secret"), 0o640); err != nil { + t.Fatal(err) + } + + gz := &bytes.Buffer{} + tw := tar.NewWriter(gz) + _ = tw.WriteHeader(&tar.Header{Name: "fly-linux-amd64", Mode: 0o755, Size: 3, Typeflag: tar.TypeReg}) + _, _ = tw.Write([]byte("new")) + _ = tw.Close() + var zipped bytes.Buffer + zw := gzip.NewWriter(&zipped) + _, _ = zw.Write(gz.Bytes()) + _ = zw.Close() + + tr := tar.NewReader(func() io.Reader { r, _ := gzip.NewReader(&zipped); return r }()) + if _, err := tr.Next(); err != nil { + t.Fatal(err) + } + _ = writeBinary(exe, &swapReader{t: t, r: tr, dir: dir, victim: victim}) + + info, err := os.Stat(victim) + if err != nil { + t.Fatal(err) + } + st := info.Sys().(*syscall.Stat_t) + if info.Mode().Perm() != 0o640 || st.Uid != 0 { + t.Errorf("victim = %v owned by %d, want 0640 owned by root: root followed the link", info.Mode().Perm(), st.Uid) + } +} diff --git a/internal/release/signature.go b/internal/release/signature.go new file mode 100644 index 0000000..2b03ce3 --- /dev/null +++ b/internal/release/signature.go @@ -0,0 +1,151 @@ +package release + +import ( + "bytes" + "crypto/ed25519" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "errors" + "fmt" + "strings" + "time" +) + +// SignatureAsset is the signature of the checksum file of a release. The +// maintainer makes it with a key that is kept outside GitHub, after the +// release workflow publishes the release. The agent installs a release by +// itself only when this signature is valid: the checksum file alone proves +// that a download is complete, not who made the release. +const SignatureAsset = ChecksumsAsset + ".sig" + +// The lines of a signature file, in this order: +// +// fly-release-signature-v1 +// key +// tag +// signed-at +// sig +// +// The signature covers the first four lines, each with its "\n", followed by +// the exact bytes of checksums.txt. The tag stops a signature from being +// moved to an other release; signed-at is the start of the wait before +// agents install the release. +const signatureHeader = "fly-release-signature-v1" + +// maxClockSkew is how far in the future a signature time can be: the clock +// of the maintainer and of the server can differ a little. +const maxClockSkew = 10 * time.Minute + +// Signature is a verified signature file. +type Signature struct { + KeyID string + Tag string + SignedAt time.Time +} + +// KeyID returns the id of a public key: the first 8 bytes of its sha256, in +// hex. A signature names its key with it, so that a new key can replace an +// old one. +func KeyID(pub ed25519.PublicKey) string { + sum := sha256.Sum256(pub) + return hex.EncodeToString(sum[:8]) +} + +// Sign returns the signature file of checksums for the release tag. +func Sign(priv ed25519.PrivateKey, tag string, signedAt time.Time, checksums []byte) []byte { + pub, _ := priv.Public().(ed25519.PublicKey) + head := signedHead(KeyID(pub), tag, signedAt.UTC()) + sig := ed25519.Sign(priv, append([]byte(head), checksums...)) + + return []byte(head + "sig " + base64.StdEncoding.EncodeToString(sig) + "\n") +} + +// Verify checks that sigFile is a valid signature of checksums for the +// release tag, by one of the keys (key id to public key). It refuses a +// signature time more than a few minutes after now. +func Verify(sigFile, checksums []byte, tag string, keys map[string]ed25519.PublicKey, now time.Time) (*Signature, error) { + lines := strings.Split(string(sigFile), "\n") + // Five lines, each with its "\n", so the split gives an empty sixth. + if len(lines) != 6 || lines[5] != "" || lines[0] != signatureHeader { + return nil, errors.New("the signature file is not in the fly-release-signature-v1 format") + } + + keyID, ok1 := strings.CutPrefix(lines[1], "key ") + signedTag, ok2 := strings.CutPrefix(lines[2], "tag ") + at, ok3 := strings.CutPrefix(lines[3], "signed-at ") + encoded, ok4 := strings.CutPrefix(lines[4], "sig ") + if !ok1 || !ok2 || !ok3 || !ok4 { + return nil, errors.New("the signature file is not in the fly-release-signature-v1 format") + } + + pub, ok := keys[keyID] + if !ok { + return nil, fmt.Errorf("the signature uses the key %q, which this binary does not trust", keyID) + } + if signedTag != tag { + return nil, fmt.Errorf("the signature is for the release %s, not %s", signedTag, tag) + } + signedAt, err := time.Parse(time.RFC3339, at) + if err != nil { + return nil, fmt.Errorf("the signature time %q is not valid", at) + } + sig, err := base64.StdEncoding.DecodeString(encoded) + if err != nil || len(sig) != ed25519.SignatureSize { + return nil, errors.New("the signature is not a valid ed25519 signature") + } + + // The head is rebuilt from the lines as they are in the file, so that + // the check covers exactly the bytes that were signed. + head := strings.Join(lines[:4], "\n") + "\n" + if !ed25519.Verify(pub, append([]byte(head), checksums...), sig) { + return nil, fmt.Errorf("the signature of %s for %s is not correct", ChecksumsAsset, tag) + } + + // Checked after the signature: only a valid signature makes the time + // worth an error message. + if signedAt.After(now.Add(maxClockSkew)) { + return nil, fmt.Errorf("the signature time %s is in the future", at) + } + + return &Signature{KeyID: keyID, Tag: signedTag, SignedAt: signedAt}, nil +} + +func signedHead(keyID, tag string, signedAt time.Time) string { + var b bytes.Buffer + fmt.Fprintf(&b, "%s\nkey %s\ntag %s\nsigned-at %s\n", signatureHeader, keyID, tag, signedAt.Format(time.RFC3339)) + return b.String() +} + +// ParseKeys parses a list of public keys: ":", separated +// with commas. Each id must be the id of its key. +func ParseKeys(list string) (map[string]ed25519.PublicKey, error) { + keys := make(map[string]ed25519.PublicKey) + for entry := range strings.SplitSeq(list, ",") { + entry = strings.TrimSpace(entry) + if entry == "" { + continue + } + + id, encoded, ok := strings.Cut(entry, ":") + if !ok { + return nil, fmt.Errorf("the key %q is not in the : format", entry) + } + raw, err := base64.StdEncoding.DecodeString(encoded) + if err != nil || len(raw) != ed25519.PublicKeySize { + return nil, fmt.Errorf("the key %s is not a valid ed25519 public key", id) + } + pub := ed25519.PublicKey(raw) + if KeyID(pub) != id { + return nil, fmt.Errorf("the key id %s does not agree with its key (%s)", id, KeyID(pub)) + } + keys[id] = pub + } + + return keys, nil +} + +// KeyLine returns the entry of pub for the list that ParseKeys reads. +func KeyLine(pub ed25519.PublicKey) string { + return KeyID(pub) + ":" + base64.StdEncoding.EncodeToString(pub) +} diff --git a/internal/release/signature_test.go b/internal/release/signature_test.go new file mode 100644 index 0000000..5bd2094 --- /dev/null +++ b/internal/release/signature_test.go @@ -0,0 +1,257 @@ +package release + +import ( + "context" + "crypto/ed25519" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" +) + +var signTime = time.Date(2026, 10, 1, 10, 0, 0, 0, time.UTC) + +func testKey(t *testing.T) (ed25519.PrivateKey, map[string]ed25519.PublicKey) { + t.Helper() + pub, priv, err := ed25519.GenerateKey(nil) + if err != nil { + t.Fatal(err) + } + return priv, map[string]ed25519.PublicKey{KeyID(pub): pub} +} + +func TestSignThenVerify(t *testing.T) { + priv, keys := testKey(t) + checksums := []byte("abc fly-linux-amd64.tar.gz\n") + + sigFile := Sign(priv, "v0.2.1", signTime, checksums) + + sig, err := Verify(sigFile, checksums, "v0.2.1", keys, signTime.Add(time.Hour)) + if err != nil { + t.Fatalf("Verify() error = %v", err) + } + if sig.Tag != "v0.2.1" || !sig.SignedAt.Equal(signTime) { + t.Errorf("Verify() = %+v, want tag v0.2.1 signed at %v", sig, signTime) + } + if !strings.HasPrefix(string(sigFile), "fly-release-signature-v1\nkey ") { + t.Errorf("signature file = %q, want the v1 format", sigFile) + } +} + +func TestVerifyRefuses(t *testing.T) { + priv, keys := testKey(t) + otherPriv, _ := testKey(t) + checksums := []byte("abc fly-linux-amd64.tar.gz\n") + good := string(Sign(priv, "v0.2.1", signTime, checksums)) + now := signTime.Add(time.Hour) + + tests := []struct { + name string + sigFile string + checksums string + tag string + now time.Time + want string + }{ + {"an unknown key", string(Sign(otherPriv, "v0.2.1", signTime, checksums)), string(checksums), "v0.2.1", now, "does not trust"}, + {"an other tag", good, string(checksums), "v0.2.2", now, "for the release v0.2.1, not v0.2.2"}, + {"a changed checksum file", good, "def fly-linux-amd64.tar.gz\n", "v0.2.1", now, "is not correct"}, + {"a changed signature time", strings.Replace(good, "10:00:00Z", "09:00:00Z", 1), string(checksums), "v0.2.1", now, "is not correct"}, + {"a time in the future", good, string(checksums), "v0.2.1", signTime.Add(-time.Hour), "in the future"}, + {"reordered lines", reorder(good), string(checksums), "v0.2.1", now, "format"}, + {"an extra line", good + "extra\n", string(checksums), "v0.2.1", now, "format"}, + {"no end of line", strings.TrimSuffix(good, "\n"), string(checksums), "v0.2.1", now, "format"}, + {"an other format", strings.Replace(good, "-v1", "-v2", 1), string(checksums), "v0.2.1", now, "format"}, + {"a short signature", good[:strings.Index(good, "sig ")] + "sig AAAA\n", string(checksums), "v0.2.1", now, "not a valid ed25519"}, + {"an empty file", "", string(checksums), "v0.2.1", now, "format"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := Verify([]byte(tt.sigFile), []byte(tt.checksums), tt.tag, keys, tt.now) + if err == nil || !strings.Contains(err.Error(), tt.want) { + t.Errorf("Verify() error = %v, want an error that contains %q", err, tt.want) + } + }) + } +} + +func TestVerifyCoversTheTagAndTheKey(t *testing.T) { + priv, keys := testKey(t) + otherPriv, otherKeys := testKey(t) + for id, pub := range otherKeys { + keys[id] = pub + } + otherPub, _ := otherPriv.Public().(ed25519.PublicKey) + checksums := []byte("abc fly-linux-amd64.tar.gz\n") + good := string(Sign(priv, "v0.2.1", signTime, checksums)) + now := signTime.Add(time.Hour) + + // A signature moved to an other release: the tag line is part of the + // signed bytes, not only a label. + moved := strings.Replace(good, "tag v0.2.1\n", "tag v0.2.2\n", 1) + if _, err := Verify([]byte(moved), checksums, "v0.2.2", keys, now); err == nil || !strings.Contains(err.Error(), "is not correct") { + t.Errorf("Verify() of a signature moved to v0.2.2 = %v, want a signature error", err) + } + + // The key line names an other trusted key. + lines := strings.SplitN(good, "\n", 3) + swapped := lines[0] + "\nkey " + KeyID(otherPub) + "\n" + lines[2] + if _, err := Verify([]byte(swapped), checksums, "v0.2.1", keys, now); err == nil || !strings.Contains(err.Error(), "is not correct") { + t.Errorf("Verify() with an other key line = %v, want a signature error", err) + } +} + +func TestIsLocalBuild(t *testing.T) { + tests := map[string]bool{ + "v0.2.0": false, + "v0.2.0-dev.1a2b3c4": false, + "v0.2.0-rc.1": false, + "dev": false, + "v0.2.0-3-gabcdef1": true, + "v0.2.0-dirty": true, + "v0.2.0-3-gabcdef1-dirty": true, + } + for v, want := range tests { + if got := IsLocalBuild(v); got != want { + t.Errorf("IsLocalBuild(%q) = %v, want %v", v, got, want) + } + } +} + +// reorder swaps the key and the tag lines of a signature file. +func reorder(sigFile string) string { + lines := strings.Split(sigFile, "\n") + lines[1], lines[2] = lines[2], lines[1] + return strings.Join(lines, "\n") +} + +func TestVerifyWithoutKeys(t *testing.T) { + priv, _ := testKey(t) + checksums := []byte("abc fly-linux-amd64.tar.gz\n") + + // A binary without a trusted key installs no release by itself. + if _, err := Verify(Sign(priv, "v0.2.1", signTime, checksums), checksums, "v0.2.1", nil, signTime); err == nil { + t.Error("Verify() with no keys = nil, want an error") + } +} + +func TestParseKeys(t *testing.T) { + pub1, _, _ := ed25519.GenerateKey(nil) + pub2, _, _ := ed25519.GenerateKey(nil) + + keys, err := ParseKeys(KeyLine(pub1) + ", " + KeyLine(pub2)) + if err != nil { + t.Fatal(err) + } + if len(keys) != 2 || !keys[KeyID(pub1)].Equal(pub1) { + t.Errorf("ParseKeys() = %v, want both keys by id", keys) + } + + if keys, err := ParseKeys(""); err != nil || len(keys) != 0 { + t.Errorf("ParseKeys(\"\") = %v, %v, want no keys", keys, err) + } + + for _, bad := range []string{ + "no-colon", + KeyID(pub1) + ":not-base64!", + KeyID(pub1) + ":AAAA", + KeyID(pub2) + ":" + strings.SplitN(KeyLine(pub1), ":", 2)[1], + } { + if _, err := ParseKeys(bad); err == nil { + t.Errorf("ParseKeys(%q) = nil error, want an error", bad) + } + } +} + +func TestTrustedKeysParse(t *testing.T) { + // The built-in list must always parse: a bad list would stop every + // update by itself without a clear reason. An empty list would do the + // same, for every agent of the release. + keys, err := TrustedKeys() + if err != nil { + t.Fatalf("TrustedKeys() error = %v", err) + } + if len(keys) == 0 { + t.Fatal("TrustedKeys() is empty: add the public key line of the release key to keys.go") + } +} + +// signedRelease serves a release with an archive, a checksum file and, when +// sigFile is not nil, a signature file. +func signedRelease(t *testing.T, checksums string, sigFile []byte) *GithubRelease { + t.Helper() + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/checksums.txt": + _, _ = w.Write([]byte(checksums)) + case "/checksums.txt.sig": + _, _ = w.Write(sigFile) + default: + http.NotFound(w, r) + } + })) + t.Cleanup(srv.Close) + + rel := &GithubRelease{TagName: "v0.2.1"} + names := []string{"fly-linux-amd64.tar.gz", "checksums.txt"} + if sigFile != nil { + names = append(names, "checksums.txt.sig") + } + for _, name := range names { + rel.Assets = append(rel.Assets, Asset{Name: name, BrowserDownloadURL: srv.URL + "/" + name}) + } + return rel +} + +func TestSignedChecksum(t *testing.T) { + priv, keys := testKey(t) + archiveSum := sum([]byte("archive")) + checksums := archiveSum + " fly-linux-amd64.tar.gz\n" + rel := signedRelease(t, checksums, Sign(priv, "v0.2.1", signTime, []byte(checksums))) + + got, err := SignedChecksum(context.Background(), rel, "linux", "amd64", keys, signTime) + if err != nil { + t.Fatalf("SignedChecksum() error = %v", err) + } + if got.SHA256 != archiveSum || !got.SignedAt.Equal(signTime) || !strings.HasSuffix(got.URL, "/fly-linux-amd64.tar.gz") { + t.Errorf("SignedChecksum() = %+v, want the signed sum, time and archive URL", got) + } +} + +func TestSignedChecksumRefuses(t *testing.T) { + priv, keys := testKey(t) + checksums := sum([]byte("archive")) + " fly-linux-amd64.tar.gz\n" + + t.Run("no signature", func(t *testing.T) { + rel := signedRelease(t, checksums, nil) + if _, err := SignedChecksum(context.Background(), rel, "linux", "amd64", keys, signTime); !errors.Is(err, ErrNotSigned) { + t.Errorf("SignedChecksum() error = %v, want ErrNotSigned", err) + } + }) + + t.Run("a signature of other checksums", func(t *testing.T) { + rel := signedRelease(t, checksums, Sign(priv, "v0.2.1", signTime, []byte("other\n"))) + if _, err := SignedChecksum(context.Background(), rel, "linux", "amd64", keys, signTime); err == nil || !strings.Contains(err.Error(), "is not correct") { + t.Errorf("SignedChecksum() error = %v, want a signature error", err) + } + }) + + t.Run("no line for this arch", func(t *testing.T) { + rel := signedRelease(t, checksums, Sign(priv, "v0.2.1", signTime, []byte(checksums))) + rel.Assets = append(rel.Assets, Asset{Name: "fly-linux-arm64.tar.gz", BrowserDownloadURL: rel.Assets[0].BrowserDownloadURL}) + if _, err := SignedChecksum(context.Background(), rel, "linux", "arm64", keys, signTime); err == nil || !strings.Contains(err.Error(), "has no line for fly-linux-arm64.tar.gz") { + t.Errorf("SignedChecksum() error = %v, want no line for arm64", err) + } + }) + + t.Run("no archive for this system", func(t *testing.T) { + rel := signedRelease(t, checksums, Sign(priv, "v0.2.1", signTime, []byte(checksums))) + if _, err := SignedChecksum(context.Background(), rel, "darwin", "arm64", keys, signTime); err == nil || !strings.Contains(err.Error(), "has no binary") { + t.Errorf("SignedChecksum() error = %v, want no binary for darwin", err) + } + }) +} diff --git a/internal/service/service.go b/internal/service/service.go new file mode 100644 index 0000000..22f8b59 --- /dev/null +++ b/internal/service/service.go @@ -0,0 +1,85 @@ +// Package service controls the systemd service of the monitoring agent, +// fly-agent.service, for "fly update". +package service + +import ( + "context" + "errors" + "fmt" + "io/fs" + "os" + "os/exec" + "path/filepath" + "strconv" + "strings" + "time" +) + +const unit = "fly-agent" + +// UnitPath is the unit file that the FlyWP installer writes. procRoot is the +// proc file system. Tests replace them. +var ( + UnitPath = "/etc/systemd/system/fly-agent.service" + procRoot = "/proc" +) + +// Installed reports whether the server has the monitoring agent. +func Installed() bool { + _, err := os.Stat(UnitPath) + return err == nil +} + +// Stale reports whether the agent runs a binary that was replaced: after a +// rename over the binary, the kernel shows the old one as "(deleted)". +func Stale(ctx context.Context) (bool, error) { + out, err := systemctl(ctx, "show", "--property=MainPID", "--value", unit) + if err != nil { + return false, err + } + + pid, err := strconv.Atoi(strings.TrimSpace(out)) + if err != nil { + return false, fmt.Errorf("reading the process id of %s: %q", unit, out) + } + if pid == 0 { + // The agent does not run: systemd starts the binary on the disk. + return false, nil + } + + exe, err := os.Readlink(filepath.Join(procRoot, strconv.Itoa(pid), "exe")) + if errors.Is(err, fs.ErrNotExist) { + // The process ended after systemctl showed it, for example for its + // own update or restart. systemd starts the binary on the disk. + return false, nil + } + if err != nil { + return false, fmt.Errorf("reading the binary of the agent: %w", err) + } + + return strings.HasSuffix(exe, " (deleted)"), nil +} + +// Restart restarts the agent if it runs, so that it runs the binary on the +// disk. An agent that an administrator stopped stays stopped. +func Restart(ctx context.Context) error { + _, err := systemctl(ctx, "try-restart", unit) + return err +} + +// systemctlTimeout is longer than the default stop timeout of systemd (90 s). +const systemctlTimeout = 2 * time.Minute + +func systemctl(ctx context.Context, args ...string) (string, error) { + ctx, cancel := context.WithTimeout(ctx, systemctlTimeout) + defer cancel() + + out, err := exec.CommandContext(ctx, "systemctl", args...).CombinedOutput() + if err != nil { + // %v, not %w: an *exec.ExitError would tell fly that the child + // already showed its error, and the output here would be lost. + return "", fmt.Errorf("systemctl %s: %v: %s", strings.Join(args, " "), err, strings.TrimSpace(string(out))) + } + + return string(out), nil +} diff --git a/internal/service/service_test.go b/internal/service/service_test.go new file mode 100644 index 0000000..7b22e6d --- /dev/null +++ b/internal/service/service_test.go @@ -0,0 +1,141 @@ +package service + +import ( + "context" + "os" + "path/filepath" + "strconv" + "strings" + "testing" +) + +// fakeSystemctl puts a systemctl command first in PATH. It logs its arguments +// and prints mainPID for "show". +func fakeSystemctl(t *testing.T, mainPID string, exit int) (log string) { + t.Helper() + + dir := t.TempDir() + log = filepath.Join(t.TempDir(), "calls") + script := "#!/bin/sh\nprintf '%s\\n' \"$*\" >> " + log + "\n" + + "[ \"$1\" = show ] && echo " + mainPID + "\n" + + "exit " + strconv.Itoa(exit) + "\n" + if err := os.WriteFile(filepath.Join(dir, "systemctl"), []byte(script), 0o755); err != nil { + t.Fatal(err) + } + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + + return log +} + +// fakeProc makes a proc directory in which process 4242 runs exe. +func fakeProc(t *testing.T, exe string) { + t.Helper() + + dir := t.TempDir() + if err := os.MkdirAll(filepath.Join(dir, "4242"), 0o755); err != nil { + t.Fatal(err) + } + if err := os.Symlink(exe, filepath.Join(dir, "4242", "exe")); err != nil { + t.Fatal(err) + } + + old := procRoot + procRoot = dir + t.Cleanup(func() { procRoot = old }) +} + +func calls(t *testing.T, log string) string { + t.Helper() + data, err := os.ReadFile(log) + if err != nil && !os.IsNotExist(err) { + t.Fatal(err) + } + return strings.TrimSpace(string(data)) +} + +func TestInstalled(t *testing.T) { + old := UnitPath + t.Cleanup(func() { UnitPath = old }) + + UnitPath = filepath.Join(t.TempDir(), "fly-agent.service") + if Installed() { + t.Error("Installed() = true without the unit file") + } + if err := os.WriteFile(UnitPath, []byte("[Unit]\n"), 0o644); err != nil { + t.Fatal(err) + } + if !Installed() { + t.Error("Installed() = false with the unit file") + } +} + +func TestStale(t *testing.T) { + tests := []struct { + name, mainPID, exe string + want bool + }{ + {"replaced binary", "4242", "/home/fly/.fly/bin/fly (deleted)", true}, + {"current binary", "4242", "/home/fly/.fly/bin/fly", false}, + {"agent not running", "0", "", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + log := fakeSystemctl(t, tt.mainPID, 0) + if tt.exe != "" { + fakeProc(t, tt.exe) + } + + got, err := Stale(context.Background()) + if err != nil { + t.Fatal(err) + } + if got != tt.want { + t.Errorf("Stale() = %v, want %v", got, tt.want) + } + if c := calls(t, log); c != "show --property=MainPID --value fly-agent" { + t.Errorf("systemctl calls = %q", c) + } + }) + } +} + +func TestStaleWhenTheProcessJustEnded(t *testing.T) { + fakeSystemctl(t, "4242", 0) + // No /proc/4242: the process ended after systemctl showed it. + old := procRoot + procRoot = t.TempDir() + t.Cleanup(func() { procRoot = old }) + + stale, err := Stale(context.Background()) + if err != nil || stale { + t.Errorf("Stale() = %v, %v; want false without an error", stale, err) + } +} + +func TestStaleErrors(t *testing.T) { + fakeSystemctl(t, "not-a-number", 0) + if _, err := Stale(context.Background()); err == nil { + t.Error("Stale() = nil error, want an error for a bad process id") + } + + fakeSystemctl(t, "4242", 1) + if _, err := Stale(context.Background()); err == nil || !strings.Contains(err.Error(), "systemctl show") { + t.Errorf("Stale() error = %v, want the systemctl error", err) + } +} + +func TestRestart(t *testing.T) { + log := fakeSystemctl(t, "0", 0) + if err := Restart(context.Background()); err != nil { + t.Fatal(err) + } + if c := calls(t, log); c != "try-restart fly-agent" { + t.Errorf("systemctl calls = %q, want try-restart fly-agent: a stopped agent stays stopped", c) + } + + fakeSystemctl(t, "0", 1) + if err := Restart(context.Background()); err == nil { + t.Error("Restart() = nil error, want the systemctl error") + } +} diff --git a/internal/statefile/statefile.go b/internal/statefile/statefile.go new file mode 100644 index 0000000..d3ccdac --- /dev/null +++ b/internal/statefile/statefile.go @@ -0,0 +1,79 @@ +// Package statefile keeps small JSON files that must survive a crash. +package statefile + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" +) + +// tempSuffix marks the temporary files of Write. +const tempSuffix = ".tmp-" + +// Write stores v as JSON in path. It writes a temporary file in the same +// directory, syncs it and renames it over path, so path always holds a +// complete file, also after a crash or a power loss. +func Write(path string, v any) (err error) { + data, err := json.Marshal(v) + if err != nil { + return fmt.Errorf("encoding %s: %w", filepath.Base(path), err) + } + + dir := filepath.Dir(path) + tmp, err := os.CreateTemp(dir, "."+filepath.Base(path)+tempSuffix+"*") + if err != nil { + return err + } + defer func() { + if err != nil { + _ = os.Remove(tmp.Name()) + } + }() + + if _, err = tmp.Write(data); err != nil { + _ = tmp.Close() + return err + } + if err = tmp.Sync(); err != nil { + _ = tmp.Close() + return err + } + if err = tmp.Close(); err != nil { + return err + } + if err = os.Rename(tmp.Name(), path); err != nil { + return err + } + + // Sync the directory too, so that the rename itself survives a power loss. + d, err := os.Open(dir) + if err != nil { + return err + } + defer func() { _ = d.Close() }() + return d.Sync() +} + +// Read decodes the JSON in path into v. When path does not exist, the error +// matches fs.ErrNotExist and v is not changed. +func Read(path string, v any) error { + data, err := os.ReadFile(path) + if err != nil { + return err + } + if err := json.Unmarshal(data, v); err != nil { + return fmt.Errorf("reading %s: %w", filepath.Base(path), err) + } + + return nil +} + +// RemoveTemp removes the temporary files that Write leaves in dir after a +// crash. Call it before any Write in dir starts. +func RemoveTemp(dir string) { + matches, _ := filepath.Glob(filepath.Join(dir, ".*"+tempSuffix+"*")) + for _, m := range matches { + _ = os.Remove(m) + } +} diff --git a/internal/statefile/statefile_test.go b/internal/statefile/statefile_test.go new file mode 100644 index 0000000..049c969 --- /dev/null +++ b/internal/statefile/statefile_test.go @@ -0,0 +1,129 @@ +package statefile + +import ( + "errors" + "io/fs" + "os" + "path/filepath" + "testing" +) + +type state struct { + Interval int `json:"interval"` + Names []string `json:"names"` +} + +func TestWriteThenRead(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + + want := state{Interval: 5, Names: []string{"a", "b"}} + if err := Write(path, want); err != nil { + t.Fatal(err) + } + + var got state + if err := Read(path, &got); err != nil { + t.Fatal(err) + } + if got.Interval != want.Interval || len(got.Names) != 2 { + t.Errorf("Read() = %+v, want %+v", got, want) + } + + // Only the file itself is left: no temporary files. + entries, err := os.ReadDir(filepath.Dir(path)) + if err != nil { + t.Fatal(err) + } + if len(entries) != 1 { + t.Errorf("directory holds %d entries, want only state.json", len(entries)) + } +} + +func TestWriteReplaces(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + + for i := 1; i <= 3; i++ { + if err := Write(path, state{Interval: i}); err != nil { + t.Fatal(err) + } + } + + var got state + if err := Read(path, &got); err != nil { + t.Fatal(err) + } + if got.Interval != 3 { + t.Errorf("Interval = %d, want 3", got.Interval) + } +} + +func TestReadMissing(t *testing.T) { + got := state{Interval: 7} + err := Read(filepath.Join(t.TempDir(), "missing.json"), &got) + if !errors.Is(err, fs.ErrNotExist) { + t.Fatalf("Read() error = %v, want fs.ErrNotExist", err) + } + if got.Interval != 7 { + t.Errorf("Read() changed v to %+v", got) + } +} + +func TestReadIgnoresTornTemporaryFile(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "state.json") + if err := Write(path, state{Interval: 2}); err != nil { + t.Fatal(err) + } + + // A crash during a write leaves a partial temporary file next to the + // complete one. Read must still see the complete file. + if err := os.WriteFile(filepath.Join(dir, ".state.json.tmp-123"), []byte(`{"interv`), 0o600); err != nil { + t.Fatal(err) + } + + var got state + if err := Read(path, &got); err != nil { + t.Fatal(err) + } + if got.Interval != 2 { + t.Errorf("Interval = %d, want 2", got.Interval) + } +} + +func TestReadCorrupt(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + if err := os.WriteFile(path, []byte("{"), 0o600); err != nil { + t.Fatal(err) + } + + var got state + if err := Read(path, &got); err == nil { + t.Fatal("Read() = nil, want an error for invalid JSON") + } +} + +func TestRemoveTemp(t *testing.T) { + dir := t.TempDir() + if err := Write(filepath.Join(dir, "state.json"), state{Interval: 1}); err != nil { + t.Fatal(err) + } + for _, name := range []string{".state.json.tmp-123", ".samples.json.tmp-9"} { + if err := os.WriteFile(filepath.Join(dir, name), []byte("{"), 0o600); err != nil { + t.Fatal(err) + } + } + + RemoveTemp(dir) + + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + if len(entries) != 1 || entries[0].Name() != "state.json" { + var names []string + for _, e := range entries { + names = append(names, e.Name()) + } + t.Errorf("directory holds %v, want only state.json", names) + } +} diff --git a/internal/testutil/fakedocker.go b/internal/testutil/fakedocker.go new file mode 100644 index 0000000..c7ec06f --- /dev/null +++ b/internal/testutil/fakedocker.go @@ -0,0 +1,153 @@ +// Package testutil provides helpers for tests that run fly against a fake +// docker command or a fake Docker Engine API instead of a real Docker +// installation. +package testutil + +import ( + "fmt" + "os" + "path/filepath" + "strings" + "testing" +) + +// Environment variables that control the fake docker command. +const ( + // EnvLog is the file that receives one line per docker call. + EnvLog = "FAKE_DOCKER_LOG" + // EnvExit is the exit status of "docker compose" calls (default 0). + EnvExit = "FAKE_DOCKER_EXIT" + // EnvMode selects a failure mode: "no-compose", "daemon-down" or + // "daemon-hang". + EnvMode = "FAKE_DOCKER_MODE" +) + +const fakeDockerScript = `#!/bin/sh +printf '%s\n' "$*" >> "$FAKE_DOCKER_LOG" +case "$FAKE_DOCKER_MODE" in +no-compose) + if [ "$1" = compose ]; then + echo "docker: 'compose' is not a docker command." >&2 + exit 1 + fi + ;; +daemon-hang) + if [ "$1" = version ]; then + exec sleep 30 + fi + ;; +daemon-down) + if [ "$1 $2" != "compose version" ]; then + echo "Cannot connect to the Docker daemon at unix:///var/run/docker.sock. Is the docker daemon running?" >&2 + exit 1 + fi + ;; +esac +case "$1 $2" in +"compose version") echo "2.40.0"; exit 0 ;; +"version "*) echo "29.0.0"; exit 0 ;; +esac +exit "${FAKE_DOCKER_EXIT:-0}" +` + +// FakeDocker is a fake docker command in its own directory. +type FakeDocker struct { + // Dir holds the docker script. Put it first in PATH. + Dir string + // Log receives the arguments of each docker call, one call per line. + Log string +} + +// NewFakeDocker writes a fake docker command to a temporary directory. +func NewFakeDocker(t *testing.T) *FakeDocker { + t.Helper() + + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "docker"), []byte(fakeDockerScript), 0o755); err != nil { + t.Fatal(err) + } + + return &FakeDocker{Dir: dir, Log: filepath.Join(t.TempDir(), "docker.log")} +} + +// Install puts the fake docker first in PATH and points it at its log for +// the duration of the test. +func (f *FakeDocker) Install(t *testing.T) { + t.Helper() + t.Setenv("PATH", f.Dir+string(os.PathListSeparator)+os.Getenv("PATH")) + t.Setenv(EnvLog, f.Log) +} + +// Env returns the environment for a child process that uses the fake docker. +// The PATH holds only the fake docker and the system directories. +func (f *FakeDocker) Env() []string { + return []string{ + "PATH=" + f.Dir + string(os.PathListSeparator) + "/usr/bin:/bin", + EnvLog + "=" + f.Log, + } +} + +// Calls returns the arguments of each docker call so far. +func (f *FakeDocker) Calls(t *testing.T) []string { + t.Helper() + + data, err := os.ReadFile(f.Log) + if os.IsNotExist(err) { + return nil + } + if err != nil { + t.Fatal(err) + } + + return strings.Split(strings.TrimSuffix(string(data), "\n"), "\n") +} + +// ComposeCalls returns the "docker compose -f" calls so far, without the +// calls that check whether Docker is available. +func (f *FakeDocker) ComposeCalls(t *testing.T) []string { + t.Helper() + + var calls []string + for _, c := range f.Calls(t) { + if strings.HasPrefix(c, "compose -f ") { + calls = append(calls, c) + } + } + + return calls +} + +// WriteSite creates dir with a docker-compose.yml that defines services and +// returns the path of the compose file. +func WriteSite(t *testing.T, dir string, services ...string) string { + t.Helper() + + var b strings.Builder + b.WriteString("services:\n") + for _, s := range services { + fmt.Fprintf(&b, " %s:\n image: example\n", s) + } + + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatal(err) + } + composePath := filepath.Join(dir, "docker-compose.yml") + if err := os.WriteFile(composePath, []byte(b.String()), 0o644); err != nil { + t.Fatal(err) + } + + return composePath +} + +// TempDir returns a new temporary directory with symlinks resolved, so that +// it compares equal to the working directory a child process sees. +func TempDir(t *testing.T) string { + t.Helper() + + dir, err := filepath.EvalSymlinks(t.TempDir()) + if err != nil { + t.Fatal(err) + } + + return dir +} diff --git a/internal/testutil/fakeengine.go b/internal/testutil/fakeengine.go new file mode 100644 index 0000000..7df39a6 --- /dev/null +++ b/internal/testutil/fakeengine.go @@ -0,0 +1,69 @@ +package testutil + +import ( + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + "testing" +) + +// Engine is a fake Docker Engine API on a unix socket. It records the +// requests. +type Engine struct { + Socket string + + mu sync.Mutex + requests []string +} + +// NewEngine starts a fake engine that answers with handler. The socket path is +// short: a unix socket path has at most 104 bytes on macOS. +func NewEngine(t *testing.T, handler http.HandlerFunc) *Engine { + t.Helper() + dir, err := os.MkdirTemp("", "dk") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.RemoveAll(dir) }) + + e := &Engine{Socket: filepath.Join(dir, "docker.sock")} + l, err := net.Listen("unix", e.Socket) + if err != nil { + t.Fatal(err) + } + srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + e.mu.Lock() + e.requests = append(e.requests, r.Method+" "+r.URL.Path) + e.mu.Unlock() + handler(w, r) + })) + srv.Listener = l + srv.Start() + t.Cleanup(srv.Close) + return e +} + +// Requests returns the method and the path of each request. +func (e *Engine) Requests() []string { + e.mu.Lock() + defer e.mu.Unlock() + return append([]string(nil), e.requests...) +} + +// EngineHandler answers GET /version with version, and GET /containers/json +// with containers (JSON). +func EngineHandler(version, containers string) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/version": + _, _ = w.Write([]byte(`{"Platform":{"Name":"Docker Engine - Community"},"Version":"` + version + `","ApiVersion":"1.52"}`)) + case "/containers/json": + _, _ = w.Write([]byte(containers)) + default: + http.NotFound(w, r) + } + } +} diff --git a/internal/utils/error.go b/internal/utils/error.go deleted file mode 100644 index 2f7bc5e..0000000 --- a/internal/utils/error.go +++ /dev/null @@ -1,14 +0,0 @@ -package utils - -import "github.com/fatih/color" - -func ShowNoComposeError() { - color.Red("No docker-compose.yml file found!") - - color.Yellow("You are not inside a site directory.") - color.Yellow("Please run this command from inside a site directory, e.g:") - color.Yellow(" cd ~/example.com") - color.Yellow(" fly start") - color.Yellow("\nOr specify the domain name:") - color.Yellow(" fly start --domain example.com") -} diff --git a/internal/utils/finder.go b/internal/utils/finder.go index 954bd37..74e6026 100644 --- a/internal/utils/finder.go +++ b/internal/utils/finder.go @@ -1,36 +1,89 @@ package utils import ( + "errors" + "fmt" + "io/fs" "os" "path/filepath" + "regexp" + "strings" ) -func FindComposeFile(domain string) string { - homeDir, _ := os.UserHomeDir() +// ErrComposeNotFound means that the site has no docker-compose.yml file. +var ErrComposeNotFound = errors.New("no docker-compose.yml file found") + +// hostnameLabel matches one DNS label: letters, digits and inner hyphens. +var hostnameLabel = regexp.MustCompile(`^[A-Za-z0-9]([A-Za-z0-9-]{0,61}[A-Za-z0-9])?$`) + +// ValidateDomain returns an error if domain is not a valid hostname. A valid +// hostname has no path separators and no "..", so it can only name a +// directory directly inside the sites directory. +func ValidateDomain(domain string) error { + if len(domain) == 0 || len(domain) > 253 { + return fmt.Errorf("invalid domain %q", domain) + } + + for _, label := range strings.Split(domain, ".") { + if !hostnameLabel.MatchString(label) { + return fmt.Errorf("invalid domain %q", domain) + } + } + + return nil +} + +// FindComposeFile returns the docker-compose.yml file of a site. With a +// domain, it looks in ~/. Without one, it searches from the current +// directory up to the home directory. It returns ErrComposeNotFound if the +// site has no compose file. +func FindComposeFile(domain string) (string, error) { + homeDir, err := os.UserHomeDir() + if err != nil { + return "", err + } + homeDir = filepath.Clean(homeDir) // If domain is provided, check in ~/domain/docker-compose.yml if domain != "" { - composePath := filepath.Join(homeDir, domain, "docker-compose.yml") - if _, err := os.Stat(composePath); err == nil { - return composePath + if err := ValidateDomain(domain); err != nil { + return "", err } - return "" + + siteDir := filepath.Join(homeDir, domain) + if filepath.Dir(siteDir) != homeDir { + return "", fmt.Errorf("invalid domain %q", domain) + } + + composePath := filepath.Join(siteDir, "docker-compose.yml") + if _, err := os.Stat(composePath); errors.Is(err, fs.ErrNotExist) { + return "", fmt.Errorf("%w for domain %q", ErrComposeNotFound, domain) + } else if err != nil { + return "", err + } + + return composePath, nil } // Otherwise search from current directory up to home directory - dir, _ := os.Getwd() + dir, err := os.Getwd() + if err != nil { + return "", err + } for { composePath := filepath.Join(dir, "docker-compose.yml") - if _, err := os.Stat(composePath); err == nil { - return composePath + return composePath, nil } - if dir == homeDir { - return "" + // Stop at the home directory, or at the filesystem root when the + // current directory is outside the home directory. + parent := filepath.Dir(dir) + if dir == homeDir || parent == dir { + return "", ErrComposeNotFound } - dir = filepath.Dir(dir) + dir = parent } } diff --git a/internal/utils/finder_test.go b/internal/utils/finder_test.go new file mode 100644 index 0000000..eeca70e --- /dev/null +++ b/internal/utils/finder_test.go @@ -0,0 +1,142 @@ +package utils + +import ( + "errors" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/flywp/server-cli/internal/testutil" +) + +// findWithTimeout fails the test if FindComposeFile does not return quickly. +func findWithTimeout(t *testing.T, domain string) (string, error) { + t.Helper() + + type result struct { + path string + err error + } + done := make(chan result, 1) + go func() { + path, err := FindComposeFile(domain) + done <- result{path, err} + }() + + select { + case r := <-done: + return r.path, r.err + case <-time.After(2 * time.Second): + t.Fatal("FindComposeFile did not return within 2s") + return "", nil + } +} + +func TestFindComposeFileOutsideHomeStops(t *testing.T) { + root := testutil.TempDir(t) + t.Setenv("HOME", filepath.Join(root, "home")) + + outside := filepath.Join(root, "var", "www") + if err := os.MkdirAll(outside, 0o755); err != nil { + t.Fatal(err) + } + t.Chdir(outside) + + if _, err := findWithTimeout(t, ""); !errors.Is(err, ErrComposeNotFound) { + t.Errorf("FindComposeFile() error = %v, want ErrComposeNotFound", err) + } +} + +func TestFindComposeFileFromNestedDirectory(t *testing.T) { + home := testutil.TempDir(t) + t.Setenv("HOME", home) + + want := testutil.WriteSite(t, filepath.Join(home, "example.com"), "php") + nested := filepath.Join(home, "example.com", "app", "public", "wp-content") + if err := os.MkdirAll(nested, 0o755); err != nil { + t.Fatal(err) + } + t.Chdir(nested) + + got, err := findWithTimeout(t, "") + if err != nil || got != want { + t.Errorf("FindComposeFile() = %q, %v, want %q, nil", got, err, want) + } +} + +func TestFindComposeFileStopsAtHome(t *testing.T) { + root := testutil.TempDir(t) + testutil.WriteSite(t, root, "php") // above the home directory: not a site + home := filepath.Join(root, "home") + t.Setenv("HOME", home) + + dir := filepath.Join(home, "notes") + if err := os.MkdirAll(dir, 0o755); err != nil { + t.Fatal(err) + } + t.Chdir(dir) + + if got, err := findWithTimeout(t, ""); !errors.Is(err, ErrComposeNotFound) { + t.Errorf("FindComposeFile() = %q, %v, want ErrComposeNotFound", got, err) + } +} + +func TestFindComposeFileWithDomain(t *testing.T) { + home := testutil.TempDir(t) + t.Setenv("HOME", home) + want := testutil.WriteSite(t, filepath.Join(home, "example.com"), "php") + + if got, err := findWithTimeout(t, "example.com"); err != nil || got != want { + t.Errorf("FindComposeFile(example.com) = %q, %v, want %q, nil", got, err, want) + } + + if _, err := findWithTimeout(t, "other.com"); !errors.Is(err, ErrComposeNotFound) { + t.Errorf("FindComposeFile(other.com) error = %v, want ErrComposeNotFound", err) + } + + if _, err := findWithTimeout(t, "../"); err == nil || errors.Is(err, ErrComposeNotFound) { + t.Errorf("FindComposeFile(../) error = %v, want an invalid domain error", err) + } +} + +func TestValidateDomain(t *testing.T) { + valid := []string{ + "example.com", + "www.example.co.uk", + "my-site.example.com", + "xn--mnchen-3ya.de", + "localhost", + strings.Repeat("a", 63) + ".com", + } + for _, d := range valid { + if err := ValidateDomain(d); err != nil { + t.Errorf("ValidateDomain(%q) = %v, want nil", d, err) + } + } + + invalid := []string{ + "", + ".", + "..", + "../", + "../etc", + "a/b", + "/etc", + ".example.com", + "example..com", + "example.com.", + "-example.com", + "example-.com", + "exa mple.com", + "exa_mple.com", + strings.Repeat("a", 64) + ".com", + strings.Repeat("a.", 127) + "com", + } + for _, d := range invalid { + if err := ValidateDomain(d); err == nil { + t.Errorf("ValidateDomain(%q) = nil, want an error", d) + } + } +} diff --git a/internal/utils/version.go b/internal/utils/version.go deleted file mode 100644 index b5fefa5..0000000 --- a/internal/utils/version.go +++ /dev/null @@ -1,158 +0,0 @@ -package utils - -import ( - "encoding/json" - "fmt" - "io" - "net/http" - "os" - "os/exec" - "path/filepath" - "runtime" - - "github.com/flywp/server-cli/internal/version" -) - -const GithubAPI = "https://api.github.com/repos/flywp/server-cli/releases/latest" - -type GithubRelease struct { - TagName string `json:"tag_name"` - Assets []struct { - Name string `json:"name"` - BrowserDownloadURL string `json:"browser_download_url"` - } `json:"assets"` -} - -func CheckForUpdates() (string, bool, error) { - resp, err := http.Get(GithubAPI) - if err != nil { - return "", false, err - } - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return "", false, err - } - - var release GithubRelease - if err := json.Unmarshal(body, &release); err != nil { - return "", false, err - } - - return release.TagName, release.TagName > version.Version, nil -} - -func SelfUpdate() error { - if os.Geteuid() != 0 { - return fmt.Errorf("the update command must be run as root") - } - - release, err := getLatestRelease() - if err != nil { - return fmt.Errorf("failed to get latest release: %w", err) - } - - assetURL := getAssetURL(release) - if assetURL == "" { - return fmt.Errorf("no suitable binary found for this system (OS: %s, ARCH: %s)", runtime.GOOS, runtime.GOARCH) - } - - // Create a temporary directory - tmpDir, err := os.MkdirTemp("", "fly-cli-update") - if err != nil { - return fmt.Errorf("failed to create temp directory: %w", err) - } - defer os.RemoveAll(tmpDir) - - // Download the archive - resp, err := http.Get(assetURL) - if err != nil { - return fmt.Errorf("failed to download update: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("failed to download update: HTTP %d", resp.StatusCode) - } - - // Create the archive file - archivePath := filepath.Join(tmpDir, "update.tar.gz") - out, err := os.Create(archivePath) - if err != nil { - return fmt.Errorf("failed to create archive file: %w", err) - } - - // Write the body to file - _, err = io.Copy(out, resp.Body) - out.Close() - if err != nil { - return fmt.Errorf("failed to write archive file: %w", err) - } - - // Extract the archive - binaryName := fmt.Sprintf("fly-%s-%s", runtime.GOOS, runtime.GOARCH) - cmd := exec.Command("tar", "-xzf", archivePath, "-C", tmpDir) - if err := cmd.Run(); err != nil { - return fmt.Errorf("failed to extract archive: %w", err) - } - - // Get the current executable path - exe, err := os.Executable() - if err != nil { - return fmt.Errorf("failed to get current executable path: %w", err) - } - exe, err = filepath.EvalSymlinks(exe) - if err != nil { - return fmt.Errorf("failed to resolve symlinks: %w", err) - } - - // Make the new binary executable - extractedBinary := filepath.Join(tmpDir, binaryName) - if err := os.Chmod(extractedBinary, 0755); err != nil { - return fmt.Errorf("failed to make binary executable: %w", err) - } - - // Rename the temporary file to the executable name - if err := os.Rename(extractedBinary, exe); err != nil { - return fmt.Errorf("failed to replace old binary: %w", err) - } - - return nil -} - -func getLatestRelease() (*GithubRelease, error) { - resp, err := http.Get(GithubAPI) - if err != nil { - return nil, err - } - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return nil, err - } - - var release GithubRelease - if err := json.Unmarshal(body, &release); err != nil { - return nil, err - } - - return &release, nil -} - -func getAssetURL(release *GithubRelease) string { - arch := runtime.GOARCH - if runtime.GOOS != "linux" { - return "" - } - - expectedName := fmt.Sprintf("fly-linux-%s.tar.gz", arch) - for _, asset := range release.Assets { - if asset.Name == expectedName { - return asset.BrowserDownloadURL - } - } - - return "" -} diff --git a/main_test.go b/main_test.go new file mode 100644 index 0000000..e9de3ab --- /dev/null +++ b/main_test.go @@ -0,0 +1,396 @@ +package main + +// End-to-end tests: build the fly binary once and run it against a fake +// docker command, so that exit codes and output are the ones users see. + +import ( + "bytes" + "context" + "errors" + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/flywp/server-cli/internal/testutil" +) + +var flyBin string + +func TestMain(m *testing.M) { + dir, err := os.MkdirTemp("", "fly-e2e") + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + + flyBin = filepath.Join(dir, "fly") + if out, err := exec.Command("go", "build", "-o", flyBin, ".").CombinedOutput(); err != nil { + fmt.Fprintf(os.Stderr, "building fly: %v\n%s", err, out) + os.Exit(1) + } + + code := m.Run() + _ = os.RemoveAll(dir) + os.Exit(code) +} + +type result struct { + code int + stdout, stderr string +} + +// env is a test environment: a home directory with one site in it and a +// fake docker command. +type env struct { + home, site string + docker *testutil.FakeDocker + vars []string +} + +func newEnv(t *testing.T, services ...string) *env { + t.Helper() + + if os.Geteuid() == 0 { + t.Skip("fly refuses to run as root") + } + if len(services) == 0 { + services = []string{"php", "nginx"} + } + + home := testutil.TempDir(t) + site := filepath.Join(home, "example.com") + testutil.WriteSite(t, site, services...) + + return &env{home: home, site: site, docker: testutil.NewFakeDocker(t)} +} + +// run executes fly in dir with the test environment. It kills fly and fails +// the test if fly does not finish within 10 seconds. +func (e *env) run(t *testing.T, dir string, args ...string) result { + t.Helper() + return runFly(t, dir, append(append(e.docker.Env(), "HOME="+e.home), e.vars...), args...) +} + +// runFly executes fly in dir with the environment environ. It kills fly and +// fails the test if fly does not finish within 10 seconds. +func runFly(t *testing.T, dir string, environ []string, args ...string) result { + t.Helper() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + cmd := exec.CommandContext(ctx, flyBin, args...) + cmd.Dir = dir + cmd.Env = environ + + var stdout, stderr bytes.Buffer + cmd.Stdout, cmd.Stderr = &stdout, &stderr + + err := cmd.Run() + res := result{stdout: stdout.String(), stderr: stderr.String()} + + var exitErr *exec.ExitError + switch { + case ctx.Err() != nil: + t.Fatalf("fly %s did not finish within 10s", strings.Join(args, " ")) + case errors.As(err, &exitErr): + res.code = exitErr.ExitCode() + case err != nil: + t.Fatalf("running fly %s: %v", strings.Join(args, " "), err) + } + + return res +} + +func TestChildExitStatusPassesThrough(t *testing.T) { + e := newEnv(t) + e.vars = append(e.vars, testutil.EnvExit+"=7") + + res := e.run(t, e.site, "exec", "--", "php", "sh", "-c", "exit 7") + if res.code != 7 { + t.Errorf("exit code = %d, want 7 (stderr %q)", res.code, res.stderr) + } + if res.stderr != "" { + t.Errorf("stderr = %q, want nothing: the child reports its own error", res.stderr) + } +} + +func TestDockerFailureExitsNonZero(t *testing.T) { + commands := [][]string{ + {"start"}, + {"stop"}, + {"restart"}, + {"restart", "php"}, + {"wp", "--", "plugin", "list"}, + {"exec", "--", "php", "ls"}, + {"logs"}, + {"base", "start"}, + {"base", "stop"}, + {"base", "restart"}, + } + + for _, args := range commands { + t.Run(strings.Join(args, " "), func(t *testing.T) { + e := newEnv(t) + e.vars = append(e.vars, testutil.EnvExit+"=3") + + res := e.run(t, e.site, args...) + if res.code != 3 { + t.Errorf("exit code = %d, want 3 (stderr %q)", res.code, res.stderr) + } + if strings.Contains(res.stdout, "successfully") { + t.Errorf("stdout = %q, want no success message after a failure", res.stdout) + } + if strings.Contains(res.stdout+res.stderr, "%!") { + t.Errorf("output has a format error: stdout %q, stderr %q", res.stdout, res.stderr) + } + }) + } +} + +func TestSuccessExitsZero(t *testing.T) { + e := newEnv(t) + + res := e.run(t, e.site, "start") + if res.code != 0 { + t.Fatalf("exit code = %d, want 0 (stderr %q)", res.code, res.stderr) + } + + want := "compose -f " + filepath.Join(e.site, "docker-compose.yml") + " up -d" + if calls := e.docker.ComposeCalls(t); len(calls) != 1 || calls[0] != want { + t.Errorf("docker calls = %q, want [%q]", calls, want) + } +} + +func TestNoSiteIsAnError(t *testing.T) { + e := newEnv(t) + dir := filepath.Join(e.home, "not-a-site") + if err := os.Mkdir(dir, 0o755); err != nil { + t.Fatal(err) + } + + res := e.run(t, dir, "start") + if res.code != 1 { + t.Errorf("exit code = %d, want 1", res.code) + } + if !strings.Contains(res.stderr, "no docker-compose.yml file found") { + t.Errorf("stderr = %q, want the no-site error", res.stderr) + } + if res.stdout != "" { + t.Errorf("stdout = %q, want errors on stderr only", res.stdout) + } +} + +func TestOutsideHomeDoesNotHang(t *testing.T) { + e := newEnv(t) + outside := testutil.TempDir(t) // not inside e.home + + // run fails the test if fly does not stop. + res := e.run(t, outside, "start") + if res.code != 1 || !strings.Contains(res.stderr, "no docker-compose.yml file found") { + t.Errorf("got exit %d, stderr %q, want exit 1 and the no-site error", res.code, res.stderr) + } +} + +func TestDomainFlag(t *testing.T) { + e := newEnv(t) + outside := testutil.TempDir(t) + + res := e.run(t, outside, "start", "--domain", "example.com") + if res.code != 0 { + t.Fatalf("--domain example.com: exit code = %d, want 0 (stderr %q)", res.code, res.stderr) + } + + for _, bad := range []string{"../", "..", "example.com/../..", "a/b"} { + res := e.run(t, outside, "start", "--domain", bad) + if res.code != 1 || !strings.Contains(res.stderr, "invalid domain") { + t.Errorf("--domain %q: exit %d, stderr %q, want exit 1 and an invalid domain error", bad, res.code, res.stderr) + } + } + + if calls := e.docker.ComposeCalls(t); len(calls) != 1 { + t.Errorf("docker calls = %q, want only the call for example.com", calls) + } +} + +func TestFlagsPassThrough(t *testing.T) { + tests := []struct { + name string + services []string + outside bool // run outside the site directory + args []string + want string // docker arguments after "compose -f " + }{ + { + name: "wp-cli flags", + args: []string{"wp", "plugin", "list", "--format=json"}, + want: "exec -T php wp plugin list --format=json", + }, + { + name: "domain before the wp-cli command", + outside: true, + args: []string{"--domain", "example.com", "wp", "plugin", "list", "--format=json"}, + want: "exec -T php wp plugin list --format=json", + }, + { + name: "exec flags", + args: []string{"exec", "php", "ls", "-la"}, + want: "exec -T php ls -la", + }, + { + name: "exec in the default service", + args: []string{"exec", "ls", "-la"}, + want: "exec -T php ls -la", + }, + { + name: "exec on an OpenLiteSpeed site", + services: []string{"openlitespeed"}, + args: []string{"exec", "ls"}, + want: "exec -T openlitespeed ls", + }, + { + name: "wp on an OpenLiteSpeed site", + services: []string{"openlitespeed"}, + args: []string{"wp", "plugin", "list"}, + want: "exec -T --user www-data openlitespeed wp plugin list", + }, + { + name: "follow logs of one service", + args: []string{"logs", "-f", "php"}, + want: "logs --follow php", + }, + { + name: "tail logs", + args: []string{"logs", "--tail", "50"}, + want: "logs --tail 50", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + e := newEnv(t, tt.services...) + dir := e.site + if tt.outside { + dir = testutil.TempDir(t) + } + + res := e.run(t, dir, tt.args...) + if res.code != 0 { + t.Fatalf("exit code = %d, want 0 (stderr %q)", res.code, res.stderr) + } + + want := "compose -f " + filepath.Join(e.site, "docker-compose.yml") + " " + tt.want + if calls := e.docker.ComposeCalls(t); len(calls) != 1 || calls[0] != want { + t.Errorf("docker calls = %q, want [%q]", calls, want) + } + }) + } +} + +func TestExecNeedsACommand(t *testing.T) { + e := newEnv(t) + + res := e.run(t, e.site, "exec", "php") + if res.code != 1 || !strings.Contains(res.stderr, "no command given") { + t.Errorf("exit %d, stderr %q, want exit 1 and a missing-command error", res.code, res.stderr) + } +} + +func TestDockerUnavailable(t *testing.T) { + modes := []struct { + name string + mode string // fake docker failure mode + noDocker bool // no docker command in PATH at all + want string + }{ + {name: "no docker CLI", noDocker: true, want: "Warning: Docker CLI is not available"}, + {name: "no compose plugin", mode: "no-compose", want: "Warning: Docker Compose plugin is not available"}, + {name: "daemon down", mode: "daemon-down", want: "Warning: Docker daemon is not available"}, + } + commands := [][]string{ + {"start"}, + {"wp", "plugin", "list"}, + {"exec", "php", "ls"}, + {"base", "stop"}, + {"sites", "start"}, + } + + for _, m := range modes { + for _, args := range commands { + t.Run(m.name+"/"+strings.Join(args, " "), func(t *testing.T) { + e := newEnv(t) + if m.noDocker { + e.vars = append(e.vars, "PATH="+testutil.TempDir(t)) + } else { + e.vars = append(e.vars, testutil.EnvMode+"="+m.mode) + } + + res := e.run(t, e.site, args...) + if res.code != 69 { + t.Errorf("exit code = %d, want 69", res.code) + } + lines := strings.Split(strings.TrimSpace(res.stderr), "\n") + if len(lines) != 1 || !strings.HasPrefix(lines[0], m.want) { + t.Errorf("stderr = %q, want one line that starts with %q", res.stderr, m.want) + } + if strings.Contains(res.stderr, "exec:") || strings.Contains(res.stdout, "successfully") { + t.Errorf("output has a raw error or a false success: stdout %q, stderr %q", res.stdout, res.stderr) + } + if calls := e.docker.ComposeCalls(t); len(calls) != 0 { + t.Errorf("compose calls = %q, want none", calls) + } + }) + } + + t.Run(m.name+"/status", func(t *testing.T) { + e := newEnv(t) + if m.noDocker { + e.vars = append(e.vars, "PATH="+testutil.TempDir(t)) + } else { + e.vars = append(e.vars, testutil.EnvMode+"="+m.mode) + } + + res := e.run(t, e.site, "status") + if res.code != 0 { + t.Errorf("exit code = %d, want 0", res.code) + } + for _, part := range []string{"Docker CLI is", "Docker Compose plugin is", "Docker daemon is"} { + if !strings.Contains(res.stdout, part) { + t.Errorf("stdout = %q, want a line for %q", res.stdout, part) + } + } + if !strings.Contains(res.stdout, strings.TrimPrefix(m.want, "Warning: ")) { + t.Errorf("stdout = %q, want %q", res.stdout, strings.TrimPrefix(m.want, "Warning: ")) + } + }) + } +} + +func TestUsageErrors(t *testing.T) { + tests := []struct { + args []string + want string + }{ + {args: []string{"start", "--bogus"}, want: "Run 'fly start --help' for usage"}, + {args: []string{"exec"}, want: "requires at least 1 arg"}, + {args: []string{"update"}, want: "sudo fly update"}, + } + + for _, tt := range tests { + t.Run(strings.Join(tt.args, " "), func(t *testing.T) { + e := newEnv(t) + + res := e.run(t, e.site, tt.args...) + if res.code != 1 { + t.Errorf("exit code = %d, want 1", res.code) + } + if !strings.Contains(res.stderr, tt.want) { + t.Errorf("stderr = %q, want it to contain %q", res.stderr, tt.want) + } + }) + } +} diff --git a/tools/releasesign/main.go b/tools/releasesign/main.go new file mode 100644 index 0000000..dc18fcd --- /dev/null +++ b/tools/releasesign/main.go @@ -0,0 +1,212 @@ +// Command releasesign makes the release signing key and signs the checksum +// file of a release. The agent installs a release by itself only when the +// signature is valid. It is not part of the fly binary. +// +// go run ./tools/releasesign keygen -out [-comment "server-cli release key for flywp"] +// go run ./tools/releasesign sign -key -tag v0.2.1 checksums.txt > checksums.txt.sig +// go run ./tools/releasesign verify -tag v0.2.1 checksums.txt checksums.txt.sig +// +// Keep the private key outside GitHub, for example in a password manager. +// Never commit it, and never put it in a GitHub secret. +package main + +import ( + "crypto/ed25519" + "crypto/x509" + "encoding/pem" + "errors" + "flag" + "fmt" + "io" + "os" + "path/filepath" + "strings" + "time" + + "github.com/flywp/server-cli/internal/release" +) + +func main() { + if err := run(os.Args[1:], os.Stdin, os.Stdout); err != nil { + fmt.Fprintln(os.Stderr, "releasesign:", err) + os.Exit(1) + } +} + +func run(args []string, stdin io.Reader, stdout io.Writer) error { + if len(args) == 0 { + return errors.New("usage: releasesign keygen|sign|verify [flags]") + } + + switch args[0] { + case "keygen": + return keygen(args[1:], stdout) + case "sign": + return sign(args[1:], stdin, stdout) + case "verify": + return verify(args[1:], stdout) + default: + return fmt.Errorf("unknown command %q: use keygen, sign or verify", args[0]) + } +} + +// keygen writes a new private key to a file that must not exist, and prints +// the public key line for internal/release/keys.go. It never prints the +// private key. The comment names the key: it goes in the key file as a PEM +// header, and next to the public key line. +func keygen(args []string, stdout io.Writer) error { + fs := flag.NewFlagSet("keygen", flag.ContinueOnError) + out := fs.String("out", "", "the file for the private key (it must not exist)") + comment := fs.String("comment", "", "a name for the key, for example \"server-cli release key for flywp\"") + if err := fs.Parse(args); err != nil { + return err + } + if *out == "" { + return errors.New("keygen needs -out ") + } + path, err := expandHome(*out) + if err != nil { + return err + } + if strings.ContainsAny(*comment, "\r\n") { + return errors.New("the comment must be one line") + } + + pub, priv, err := ed25519.GenerateKey(nil) + if err != nil { + return err + } + der, err := x509.MarshalPKCS8PrivateKey(priv) + if err != nil { + return err + } + + f, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o600) + if err != nil { + return err + } + block := &pem.Block{Type: "PRIVATE KEY", Bytes: der} + if *comment != "" { + block.Headers = map[string]string{"Comment": *comment} + } + if err := pem.Encode(f, block); err != nil { + _ = f.Close() + return err + } + if err := f.Close(); err != nil { + return err + } + + msg := fmt.Sprintf("Private key: %s (keep it outside GitHub, with a backup)\n", path) + if *comment != "" { + msg += fmt.Sprintf("Comment: %s\n", *comment) + } + msg += fmt.Sprintf("Public key line for internal/release/keys.go:\n%s\n", release.KeyLine(pub)) + _, err = io.WriteString(stdout, msg) + return err +} + +// sign prints the signature file of a checksum file. +func sign(args []string, stdin io.Reader, stdout io.Writer) error { + fs := flag.NewFlagSet("sign", flag.ContinueOnError) + keyPath := fs.String("key", "", "the private key file, or - to read it from stdin") + tag := fs.String("tag", "", "the release tag, for example v0.2.1") + if err := fs.Parse(args); err != nil { + return err + } + if *keyPath == "" || *tag == "" || fs.NArg() != 1 { + return errors.New("usage: releasesign sign -key -tag checksums.txt") + } + + priv, err := readKey(*keyPath, stdin) + if err != nil { + return err + } + checksums, err := os.ReadFile(fs.Arg(0)) + if err != nil { + return err + } + + _, err = stdout.Write(release.Sign(priv, *tag, time.Now().UTC().Truncate(time.Second), checksums)) + return err +} + +// verify checks a signature file with the keys that this build trusts: the +// keys of the agents that are built from the same commit. +func verify(args []string, stdout io.Writer) error { + fs := flag.NewFlagSet("verify", flag.ContinueOnError) + tag := fs.String("tag", "", "the release tag, for example v0.2.1") + if err := fs.Parse(args); err != nil { + return err + } + if *tag == "" || fs.NArg() != 2 { + return errors.New("usage: releasesign verify -tag checksums.txt checksums.txt.sig") + } + + keys, err := release.TrustedKeys() + if err != nil { + return err + } + if len(keys) == 0 { + return errors.New("this build trusts no key: add the public key line to internal/release/keys.go") + } + checksums, err := os.ReadFile(fs.Arg(0)) + if err != nil { + return err + } + sigFile, err := os.ReadFile(fs.Arg(1)) + if err != nil { + return err + } + + sig, err := release.Verify(sigFile, checksums, *tag, keys, time.Now()) + if err != nil { + return err + } + + _, err = fmt.Fprintf(stdout, "The signature of %s is valid (key %s, signed at %s).\n", *tag, sig.KeyID, sig.SignedAt.Format(time.RFC3339)) + return err +} + +// expandHome replaces a leading "~/" with the home directory. A shell does +// not do this in "make release-key KEY=~/key": the "~" is not at the start of +// a word. +func expandHome(path string) (string, error) { + rest, ok := strings.CutPrefix(path, "~/") + if !ok { + return path, nil + } + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + return filepath.Join(home, rest), nil +} + +func readKey(path string, stdin io.Reader) (ed25519.PrivateKey, error) { + var data []byte + var err error + if path == "-" { + data, err = io.ReadAll(io.LimitReader(stdin, 1<<16)) + } else if path, err = expandHome(path); err == nil { + data, err = os.ReadFile(path) + } + if err != nil { + return nil, err + } + + block, _ := pem.Decode(data) + if block == nil || block.Type != "PRIVATE KEY" { + return nil, errors.New("the key is not a PEM PRIVATE KEY") + } + key, err := x509.ParsePKCS8PrivateKey(block.Bytes) + if err != nil { + return nil, fmt.Errorf("reading the key: %w", err) + } + priv, ok := key.(ed25519.PrivateKey) + if !ok { + return nil, errors.New("the key is not an ed25519 key") + } + + return priv, nil +} diff --git a/tools/releasesign/main_test.go b/tools/releasesign/main_test.go new file mode 100644 index 0000000..08584b2 --- /dev/null +++ b/tools/releasesign/main_test.go @@ -0,0 +1,107 @@ +package main + +import ( + "bytes" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/flywp/server-cli/internal/release" +) + +func TestKeygenThenSign(t *testing.T) { + dir := t.TempDir() + keyPath := filepath.Join(dir, "release.key") + + var out bytes.Buffer + if err := run([]string{"keygen", "-out", keyPath, "-comment", "server-cli release key for flywp"}, nil, &out); err != nil { + t.Fatal(err) + } + + info, err := os.Stat(keyPath) + if err != nil { + t.Fatal(err) + } + if info.Mode().Perm() != 0o600 { + t.Errorf("key file mode = %v, want 0600", info.Mode().Perm()) + } + key, _ := os.ReadFile(keyPath) + if strings.Contains(out.String(), "PRIVATE KEY") || strings.Contains(out.String(), string(key)) { + t.Fatal("keygen printed the private key") + } + // The comment names the key in the key file and in the output. + if !strings.Contains(string(key), "Comment: server-cli release key for flywp") || !strings.Contains(out.String(), "Comment: server-cli release key for flywp") { + t.Errorf("the comment is not in the key file and the output:\n%s", out.String()) + } + + // The printed line parses, and it is the key of the private key file. + lines := strings.Split(strings.TrimSpace(out.String()), "\n") + keys, err := release.ParseKeys(lines[len(lines)-1]) + if err != nil || len(keys) != 1 { + t.Fatalf("public key line %q: %v, %v", lines[len(lines)-1], keys, err) + } + + checksums := filepath.Join(dir, "checksums.txt") + if err := os.WriteFile(checksums, []byte("abc fly-linux-amd64.tar.gz\n"), 0o600); err != nil { + t.Fatal(err) + } + + // The key comes from stdin, for example from a password manager. + var sigFile bytes.Buffer + if err := run([]string{"sign", "-key", "-", "-tag", "v0.2.1", checksums}, bytes.NewReader(key), &sigFile); err != nil { + t.Fatal(err) + } + data, _ := os.ReadFile(checksums) + if _, err := release.Verify(sigFile.Bytes(), data, "v0.2.1", keys, time.Now()); err != nil { + t.Errorf("Verify() of the signed file = %v", err) + } +} + +func TestKeygenRefusesAMultiLineComment(t *testing.T) { + keyPath := filepath.Join(t.TempDir(), "release.key") + if err := run([]string{"keygen", "-out", keyPath, "-comment", "a\nProc-Type: 4,ENCRYPTED"}, nil, &bytes.Buffer{}); err == nil { + t.Fatal("keygen with a multi-line comment = nil error, want an error") + } + if _, err := os.Stat(keyPath); err == nil { + t.Error("keygen wrote a key file for a bad comment") + } +} + +func TestKeygenDoesNotOverwrite(t *testing.T) { + keyPath := filepath.Join(t.TempDir(), "release.key") + if err := os.WriteFile(keyPath, []byte("old key"), 0o600); err != nil { + t.Fatal(err) + } + + if err := run([]string{"keygen", "-out", keyPath}, nil, &bytes.Buffer{}); err == nil { + t.Fatal("keygen over an existing file = nil error, want an error") + } + if data, _ := os.ReadFile(keyPath); string(data) != "old key" { + t.Error("keygen changed an existing key file") + } +} + +func TestKeyPathsExpandTheHomeDirectory(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + + // make passes "~/release.key" as it is: the shell does not expand it. + var out bytes.Buffer + if err := run([]string{"keygen", "-out", "~/release.key"}, nil, &out); err != nil { + t.Fatalf("keygen -out ~/release.key: %v", err) + } + if _, err := os.Stat(filepath.Join(home, "release.key")); err != nil { + t.Fatalf("the key is not in the home directory: %v", err) + } + if _, err := readKey("~/release.key", nil); err != nil { + t.Errorf("readKey(~/release.key) = %v", err) + } +} + +func TestReadKeyRefusesText(t *testing.T) { + if _, err := readKey("-", strings.NewReader("not a key")); err == nil { + t.Error("readKey() of text = nil error, want an error") + } +} diff --git a/tools/sign-release.sh b/tools/sign-release.sh new file mode 100755 index 0000000..06993c1 --- /dev/null +++ b/tools/sign-release.sh @@ -0,0 +1,103 @@ +#!/usr/bin/env bash +# +# Sign a published release: make sign-release VERSION=v0.2.1 KEY= +# +# The signature says "this release is the code of my tag". So the script +# signs only what it can build again: it builds the release from your local +# tag, with the Go version of the CI build, and compares the binaries with +# the binaries that GitHub serves. A swapped archive or checksums.txt on +# GitHub, or a tag that moved on GitHub, stops the script before it signs. +# +# The binary holds the commit (vcs.revision), so equal binaries also prove +# that CI built the commit of your local tag. Review that commit before you +# sign: the signature is the approval. +# +# UPLOAD=0 does all checks and writes checksums.txt.sig, but does not upload it. +# SIGN_YES=1 skips the typed confirmation (tests only). + +set -euo pipefail + +version=${1:?usage: sign-release.sh } +key=${2:?usage: sign-release.sh } + +root=$(git rev-parse --show-toplevel) +dir="$root/build/sign/$version" + +# make passes "~/key" as it is: the shell does not expand a "~" that is not +# at the start of a word. +case "$key" in "~/"*) key="$HOME/${key#\~/}" ;; esac +if [ "$key" != "-" ]; then + [ -f "$key" ] || { echo "No key file: $key" >&2; exit 1; } + key="$(cd "$(dirname "$key")" && pwd)/$(basename "$key")" +fi + +# The source of trust is your local tag, not the tag on GitHub. +if ! commit=$(git -C "$root" rev-parse --verify --quiet "refs/tags/$version^{commit}"); then + echo "There is no local tag $version. Sign only a tag that you made or reviewed." >&2 + exit 1 +fi + +echo "Release $version is commit $(git -C "$root" log -1 --format='%h %s' "$commit")" +if [ "${SIGN_YES:-}" != "1" ]; then + printf 'Type the tag to sign it: ' > /dev/tty + read -r answer < /dev/tty + [ "$answer" = "$version" ] || { echo "Not signed." >&2; exit 1; } +fi + +rm -rf "$dir" +mkdir -p "$dir/ci" "$dir/ci-bin" +cleanup() { git -C "$root" worktree remove --force "$dir/src" >/dev/null 2>&1 || true; } +trap cleanup EXIT + +echo "Downloading the release from GitHub..." +gh release download "$version" --pattern 'fly-linux-*.tar.gz' --pattern checksums.txt --dir "$dir/ci" + +# The archives agree with checksums.txt. +if command -v sha256sum >/dev/null 2>&1; then + (cd "$dir/ci" && sha256sum -c checksums.txt) +else + (cd "$dir/ci" && shasum -a 256 -c checksums.txt) +fi + +# Build the release again from the local tag, with the Go version of CI. +git -C "$root" worktree add --detach --quiet "$dir/src" "$commit" +for archive in "$dir"/ci/fly-linux-*.tar.gz; do + tar -xzf "$archive" -C "$dir/ci-bin" +done +goversions=$(for bin in "$dir"/ci-bin/*; do go version "$bin" | awk '{ print $NF }'; done | sort -u) +if [ "$(printf '%s\n' "$goversions" | wc -l)" -ne 1 ]; then + echo "The CI binaries were built with different Go versions: $goversions" >&2 + exit 1 +fi +echo "Building $version again with $goversions..." +(cd "$dir/src" && GOTOOLCHAIN="$goversions" make release VERSION="$version" >/dev/null) + +# checksums.txt must name exactly the archives of the release build. +names() { sed 's/\r$//' "$1" | awk '{ sub(/^\*/, "", $2); print $2 }' | sort; } +if [ "$(names "$dir/ci/checksums.txt")" != "$(names "$dir/src/build/checksums.txt")" ]; then + echo "checksums.txt on GitHub does not name the archives of the release build." >&2 + exit 1 +fi + +# The binaries are the same, byte for byte. +for bin in "$dir"/src/build/fly-linux-*; do + case "$bin" in *.tar.gz) continue ;; esac + name=$(basename "$bin") + if ! cmp -s "$bin" "$dir/ci-bin/$name"; then + echo "$name on GitHub is not the build of $version. Do not sign it." >&2 + exit 1 + fi + echo "$name: the same as the build of the local tag" +done + +# Sign and verify with the tool and the trusted keys of the tag: the keys +# that the agents of this release have. +(cd "$dir/src" && go run ./tools/releasesign sign -key "$key" -tag "$version" "$dir/ci/checksums.txt") > "$dir/checksums.txt.sig" +(cd "$dir/src" && go run ./tools/releasesign verify -tag "$version" "$dir/ci/checksums.txt" "$dir/checksums.txt.sig") + +if [ "${UPLOAD:-1}" = "0" ]; then + echo "Not uploaded (UPLOAD=0): $dir/checksums.txt.sig" + exit 0 +fi +gh release upload "$version" "$dir/checksums.txt.sig" --clobber +echo "Signed $version. Agents install it 24 hours after the signature."