feat(db): 控制面 PG 定时备份、迁移前快照与原子恢复

This commit is contained in:
Lemon-miaow committed 2026-09-24 15:19:42 +08:00
1 parent abfe60d62d
commit c7db7d4126
31 files changed
+3217 -17

No files matched your search

+707
View File
@@ -0,0 +1,707 @@
// Package dbbackup takes and restores logical backups of the control-plane
// PostgreSQL: users, passkeys, account links, server ownership, quotas, audit
// logs and the world_backups index that maps a world archive back to its owner.
// World archives live on their own volume (internal/archive); without this
// database they are files nobody can be matched to.
//
// A backup is one bundle, felis-db-<UTC stamp>-<label>.tar, holding
//
// MANIFEST.json what is in the bundle and each member's sha256
// db.dump pg_dump --format=custom of the felis database
// state/etc/felis/... the host state a rebuild needs: secrets.env
// (the DB password, session and forwarding
// secrets, registry tokens), felis.{host,pod}.toml,
// the panel TLS pair
// k8s/minecraftservers.json the MinecraftServer objects, when the cluster
// answered (best effort)
//
// plus a felis-db-....tar.sha256 sidecar in sha256sum format, so a copy shipped
// off the host can be checked with `sha256sum -c` before anyone relies on it.
//
// Bundles are written as a hidden .partial and renamed into place after an
// fsync, so a crash or a full disk leaves either a complete bundle or none.
// The dump is proven readable (pg_restore --list) before the bundle counts.
//
// Restore is all-or-nothing: the dump is replayed through psql in a single
// transaction that first drops everything the felis role owns, so a failure
// anywhere leaves the database exactly as it was, and objects a newer schema
// added do not survive to collide with the next `felis migrate up`.
package dbbackup
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"io/fs"
"net/url"
"os"
"os/exec"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"syscall"
"time"
)
const (
// DefaultDir is where the host keeps its bundles. It is deliberately off the
// k3s storage tree: `rm -rf /var/lib/rancher` (a k3s reinstall) must not take
// the database backups with it.
DefaultDir = "/var/lib/felis/db-backups"
// DefaultStateDir is the host state directory bootstrap writes.
DefaultStateDir = "/etc/felis"
LabelDaily = "daily"
LabelPreMigrate = "pre-migrate"
LabelPreRestore = "pre-restore"
LabelManual = "manual"
manifestEntry = "MANIFEST.json"
dumpEntry = "db.dump"
stateEntry = "state"
serversEntry = "k8s/minecraftservers.json"
bundlePrefix = "felis-db-"
bundleExt = ".tar"
sumExt = ".sha256"
stampLayout = "20060102T150405Z"
formatV1 = 1
)
// stateSkip are files in the state directory a bundle leaves out:
// bootstrap.done marks THIS host as installed, and carrying it to a fresh host
// would make the installer treat a first install as an upgrade.
var stateSkip = map[string]bool{"bootstrap.done": true}
var labelRe = regexp.MustCompile(`^[a-z][a-z0-9-]{0,31}$`)
// Tools names the PostgreSQL client binaries. Empty fields take the names on
// PATH; tests point them at fakes.
type Tools struct {
PGDump, PGRestore, PSQL string
}
func (t Tools) pgDump() string { return orDefault(t.PGDump, "pg_dump") }
func (t Tools) pgRestore() string { return orDefault(t.PGRestore, "pg_restore") }
func (t Tools) psql() string { return orDefault(t.PSQL, "psql") }
func orDefault(s, d string) string {
if s == "" {
return d
}
return s
}
// BackupOptions configures one backup.
type BackupOptions struct {
DatabaseURL string
Dir string // bundle directory; created 0700
Label string // daily | pre-migrate | pre-restore | manual | any [a-z0-9-]
// Keep is how many bundles of this label survive the post-backup prune;
// zero or less prunes nothing.
Keep int
StateDir string // host state to bundle; "" bundles none
Version string // felis build stamp, recorded in the manifest
Tools Tools
// ExportServers returns the cluster's MinecraftServer objects as JSON. A
// failure is recorded in the manifest and does not fail the backup: the
// database is what must not be lost, and a nightly run cannot hang on a
// cluster that happens to be down.
ExportServers func(ctx context.Context) ([]byte, error)
// MetricsFile, when set, is rewritten after a successful backup with
// node-exporter textfile metrics (felis_db_backup_last_success_timestamp_seconds
// and felis_db_backup_last_size_bytes), which FelisDBBackupStale alerts on.
MetricsFile string
// Record stores a summary of the backup in platform_settings under
// StatusKey, which the admin panel reads to show how fresh the newest
// backup is. A failure to record is logged, not fatal.
Record bool
Now func() time.Time
Log io.Writer
}
// StatusKey is the platform_settings key Record writes; internal/api reads it.
const StatusKey = "db_backup_last"
// StaleAfter is how old the newest backup may get before it counts as missed:
// a day plus the timer's randomized delay and a slow dump. `felis db check`,
// the admin panel and the FelisDBBackupStale alert (deploy/alerts) share it.
const StaleAfter = 26 * time.Hour
// Status is the value stored under StatusKey.
type Status struct {
At time.Time `json:"at"`
Name string `json:"name"`
Label string `json:"label"`
SizeBytes int64 `json:"size_bytes"`
FelisVersion string `json:"felis_version,omitempty"`
SchemaVersion int `json:"schema_version,omitempty"`
Dir string `json:"dir"`
}
// Manifest describes a bundle.
type Manifest struct {
Format int `json:"format"`
CreatedAt time.Time `json:"created_at"`
Label string `json:"label"`
FelisVersion string `json:"felis_version,omitempty"`
Database DatabaseInfo `json:"database"`
SchemaVersion int `json:"schema_version,omitempty"`
PGDumpVersion string `json:"pg_dump_version,omitempty"`
Files []ManifestEntry `json:"files"`
// ServersError is why k8s/minecraftservers.json is absent, when it is.
ServersError string `json:"servers_error,omitempty"`
}
// DatabaseInfo is the connection a bundle was taken from, password excluded.
type DatabaseInfo struct {
Host string `json:"host"`
Port string `json:"port,omitempty"`
Name string `json:"name"`
User string `json:"user"`
}
// ManifestEntry is one bundle member.
type ManifestEntry struct {
Name string `json:"name"`
Size int64 `json:"size"`
SHA256 string `json:"sha256,omitempty"` // absent for symlinks
Mode uint32 `json:"mode"`
Link string `json:"link,omitempty"`
}
// Bundle is one bundle on disk.
type Bundle struct {
Name string
Path string
Label string
Created time.Time
Size int64
}
// conn splits a postgres:// URL into the URL libpq should see (password
// removed) and the password, which goes to the child through PGPASSWORD so it
// never shows up in ps.
type conn struct {
uri string
password string
info DatabaseInfo
}
func parseConn(raw string) (conn, error) {
u, err := url.Parse(raw)
if err != nil || (u.Scheme != "postgres" && u.Scheme != "postgresql") {
return conn{}, errors.New("database url must be a postgres:// URL")
}
c := conn{info: DatabaseInfo{Host: u.Hostname(), Port: u.Port(), Name: strings.TrimPrefix(u.Path, "/")}}
if u.User != nil {
c.info.User = u.User.Username()
c.password, _ = u.User.Password()
u.User = url.User(c.info.User)
}
if c.info.Name == "" {
return conn{}, errors.New("database url names no database")
}
c.uri = u.String()
return c, nil
}
func (c conn) env() []string {
env := os.Environ()
if c.password != "" {
env = append(env, "PGPASSWORD="+c.password)
}
// Never prompt: a timer-driven run with a wrong password must fail, not hang.
return append(env, "PGCONNECT_TIMEOUT=15")
}
func (c conn) command(ctx context.Context, bin string, args ...string) *exec.Cmd {
cmd := exec.CommandContext(ctx, bin, args...)
cmd.Env = c.env()
return cmd
}
// run executes cmd and folds its stderr into the error.
func run(cmd *exec.Cmd) ([]byte, error) {
var stderr bytes.Buffer
cmd.Stderr = &stderr
out, err := cmd.Output()
if err != nil {
msg := strings.TrimSpace(stderr.String())
if msg == "" {
return out, fmt.Errorf("%s: %w", filepath.Base(cmd.Path), err)
}
return out, fmt.Errorf("%s: %w: %s", filepath.Base(cmd.Path), err, msg)
}
return out, nil
}
// foreignObjectRe finds the object pg_dump could not read, and its kind.
var foreignObjectRe = regexp.MustCompile(`permission denied for (table|sequence|schema|view|materialized view) ([^\s]+)`)
// dumpHint explains the one pg_dump failure an operator causes without noticing:
// an object created in the felis database by another role (typically postgres,
// from a manual psql session). The dump runs as the felis role and must read
// everything; leaving the object out would make the bundle an incomplete restore.
func dumpHint(err error, db DatabaseInfo) string {
m := foreignObjectRe.FindStringSubmatch(err.Error())
if m == nil {
return ""
}
kind, name := strings.ToUpper(m[1]), m[2]
return fmt.Sprintf("\n %s %s belongs to another role, so %s cannot dump it. Hand it over with\n"+
" sudo -u postgres psql -d %s -c 'ALTER %s %s OWNER TO %s'\n"+
" or drop it if it is a leftover.", strings.ToLower(kind), name, db.User, db.Name, kind, name, db.User)
}
// BundleName is the file name of a bundle taken at t with label.
func BundleName(t time.Time, label string) string {
return bundlePrefix + t.UTC().Format(stampLayout) + "-" + label + bundleExt
}
// parseBundleName reverses BundleName.
func parseBundleName(name string) (time.Time, string, bool) {
if !strings.HasPrefix(name, bundlePrefix) || !strings.HasSuffix(name, bundleExt) {
return time.Time{}, "", false
}
rest := strings.TrimSuffix(strings.TrimPrefix(name, bundlePrefix), bundleExt)
if len(rest) < len(stampLayout)+2 || rest[len(stampLayout)] != '-' {
return time.Time{}, "", false
}
t, err := time.Parse(stampLayout, rest[:len(stampLayout)])
label := rest[len(stampLayout)+1:]
if err != nil || !labelRe.MatchString(label) {
return time.Time{}, "", false
}
return t, label, true
}
// List returns the bundles in dir, newest first. A missing dir is no bundles.
func List(dir string) ([]Bundle, error) {
entries, err := os.ReadDir(dir)
if errors.Is(err, os.ErrNotExist) {
return nil, nil
}
if err != nil {
return nil, err
}
var out []Bundle
for _, e := range entries {
if !e.Type().IsRegular() {
continue
}
t, label, ok := parseBundleName(e.Name())
if !ok {
continue
}
info, err := e.Info()
if err != nil {
continue
}
out = append(out, Bundle{Name: e.Name(), Path: filepath.Join(dir, e.Name()), Label: label, Created: t, Size: info.Size()})
}
sort.Slice(out, func(i, j int) bool {
if !out[i].Created.Equal(out[j].Created) {
return out[i].Created.After(out[j].Created)
}
return out[i].Name > out[j].Name
})
return out, nil
}
// Prune deletes all but the newest keep bundles of label (and their sidecars)
// and returns what it removed. keep <= 0 removes nothing.
func Prune(dir, label string, keep int) ([]string, error) {
if keep <= 0 {
return nil, nil
}
all, err := List(dir)
if err != nil {
return nil, err
}
var removed []string
n := 0
for _, b := range all {
if b.Label != label {
continue
}
if n++; n <= keep {
continue
}
if err := os.Remove(b.Path); err != nil && !errors.Is(err, os.ErrNotExist) {
return removed, err
}
_ = os.Remove(b.Path + sumExt)
removed = append(removed, b.Name)
}
return removed, nil
}
// lockDir serializes bundle writers (the nightly timer, a pre-migrate snapshot
// and an operator's manual run) on dir/.lock.
func lockDir(dir string) (func(), error) {
f, err := os.OpenFile(filepath.Join(dir, ".lock"), os.O_CREATE|os.O_RDWR, 0o600)
if err != nil {
return nil, err
}
if err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX); err != nil {
f.Close()
return nil, fmt.Errorf("lock %s: %w", dir, err)
}
return func() {
_ = syscall.Flock(int(f.Fd()), syscall.LOCK_UN)
f.Close()
}, nil
}
// removeStalePartials drops leftovers of a writer that died mid-bundle. Only
// called under the directory lock, so no live writer owns them.
func removeStalePartials(dir string) {
entries, _ := os.ReadDir(dir)
for _, e := range entries {
if strings.HasPrefix(e.Name(), "."+bundlePrefix) && strings.HasSuffix(e.Name(), ".partial") {
_ = os.Remove(filepath.Join(dir, e.Name()))
}
}
}
// Backup writes one bundle and returns its path.
func Backup(ctx context.Context, o BackupOptions) (string, error) {
if !labelRe.MatchString(o.Label) {
return "", fmt.Errorf("invalid label %q (want [a-z0-9-], e.g. daily or manual)", o.Label)
}
c, err := parseConn(o.DatabaseURL)
if err != nil {
return "", err
}
now := time.Now
if o.Now != nil {
now = o.Now
}
logw := o.Log
if logw == nil {
logw = io.Discard
}
if err := os.MkdirAll(o.Dir, 0o700); err != nil {
return "", err
}
if err := os.Chmod(o.Dir, 0o700); err != nil {
return "", err
}
unlock, err := lockDir(o.Dir)
if err != nil {
return "", err
}
defer unlock()
removeStalePartials(o.Dir)
// Names have one-second resolution. A second bundle of the same label within
// that second (a restore retried right after a failed one takes two
// pre-restore snapshots) moves to the next free second; the lock makes the
// probe race-free.
created := now().UTC().Truncate(time.Second)
name := BundleName(created, o.Label)
for i := 0; ; i++ {
if _, err := os.Lstat(filepath.Join(o.Dir, name)); errors.Is(err, fs.ErrNotExist) {
break
} else if err != nil {
return "", err
}
if i == 60 {
return "", fmt.Errorf("no free bundle name after %s", name)
}
created = created.Add(time.Second)
name = BundleName(created, o.Label)
}
final := filepath.Join(o.Dir, name)
dump := filepath.Join(o.Dir, "."+name+".dump.partial")
defer os.Remove(dump)
if _, err := run(c.command(ctx, o.Tools.pgDump(), "--format=custom", "--no-password", "--file="+dump, "--dbname="+c.uri)); err != nil {
return "", fmt.Errorf("dump the database: %w%s", err, dumpHint(err, c.info))
}
if err := os.Chmod(dump, 0o600); err != nil {
return "", err
}
// A dump pg_restore cannot read is not a backup; find out now, not on the
// day it is needed.
if _, err := run(exec.CommandContext(ctx, o.Tools.pgRestore(), "--list", dump)); err != nil {
return "", fmt.Errorf("the dump does not read back: %w", err)
}
m := Manifest{Format: formatV1, CreatedAt: created, Label: o.Label, FelisVersion: o.Version, Database: c.info}
if out, err := run(exec.CommandContext(ctx, o.Tools.pgDump(), "--version")); err == nil {
m.PGDumpVersion = strings.TrimSpace(string(out))
}
m.SchemaVersion = schemaVersion(ctx, c, o.Tools)
var members []member
dm, err := fileMember(dumpEntry, dump)
if err != nil {
return "", err
}
members = append(members, dm)
if o.StateDir != "" {
sm, err := stateMembers(o.StateDir)
if err != nil {
return "", fmt.Errorf("read host state: %w", err)
}
members = append(members, sm...)
}
if o.ExportServers != nil {
if data, err := o.ExportServers(ctx); err != nil {
m.ServersError = err.Error()
fmt.Fprintf(logw, "felis db backup: MinecraftServer objects not included: %v\n", err)
} else {
members = append(members, bytesMember(serversEntry, data))
}
}
for _, mb := range members {
m.Files = append(m.Files, mb.entry)
}
sum, err := writeBundle(o.Dir, final, m, members)
if err != nil {
return "", err
}
if err := writeFileAtomic(final+sumExt, []byte(sum+" "+name+"\n")); err != nil {
return "", fmt.Errorf("write checksum: %w", err)
}
if info, err := os.Stat(final); err == nil {
st := Status{At: created, Name: name, Label: o.Label, SizeBytes: info.Size(),
FelisVersion: o.Version, SchemaVersion: m.SchemaVersion, Dir: o.Dir}
if o.Record {
if err := record(ctx, c, o.Tools, st); err != nil {
fmt.Fprintf(logw, "felis db backup: record the backup for the panel: %v\n", err)
}
}
if o.MetricsFile != "" {
if err := writeMetrics(o.MetricsFile, st); err != nil {
fmt.Fprintf(logw, "felis db backup: write %s: %v\n", o.MetricsFile, err)
}
}
}
if removed, err := Prune(o.Dir, o.Label, o.Keep); err != nil {
fmt.Fprintf(logw, "felis db backup: prune old %s bundles: %v\n", o.Label, err)
} else if len(removed) > 0 {
fmt.Fprintf(logw, "felis db backup: pruned %d old %s bundle(s)\n", len(removed), o.Label)
}
return final, nil
}
// record upserts st into platform_settings. The JSON travels as a psql
// variable, quoted by psql itself, over stdin (-c does not interpolate).
func record(ctx context.Context, c conn, t Tools, st Status) error {
v, err := json.Marshal(st)
if err != nil {
return err
}
cmd := c.command(ctx, t.psql(), "-X", "-q", "-w", "-v", "ON_ERROR_STOP=1", "-v", "v="+string(v), "-d", c.uri)
cmd.Stdin = strings.NewReader("INSERT INTO platform_settings (key, value) VALUES ('" + StatusKey + "', :'v'::jsonb)\n" +
"ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value, updated_at = now();\n")
_, err = run(cmd)
return err
}
// writeMetrics rewrites a node-exporter textfile-collector file for st.
func writeMetrics(path string, st Status) error {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return err
}
body := fmt.Sprintf(`# HELP felis_db_backup_last_success_timestamp_seconds Unix time of the newest successful control-plane database backup.
# TYPE felis_db_backup_last_success_timestamp_seconds gauge
felis_db_backup_last_success_timestamp_seconds{label=%q} %d
# HELP felis_db_backup_last_size_bytes Size of the newest control-plane database backup bundle.
# TYPE felis_db_backup_last_size_bytes gauge
felis_db_backup_last_size_bytes{label=%q} %d
`, st.Label, st.At.Unix(), st.Label, st.SizeBytes)
if err := writeFileAtomic(path, []byte(body)); err != nil {
return err
}
// Read by node-exporter, which usually runs unprivileged.
return os.Chmod(path, 0o644)
}
// schemaVersion reads the newest applied migration, or 0 when it cannot.
func schemaVersion(ctx context.Context, c conn, t Tools) int {
out, err := run(c.command(ctx, t.psql(), "-X", "-q", "-t", "-A", "-w", "-d", c.uri,
"-c", "SELECT coalesce(max(version), 0) FROM schema_migrations"))
if err != nil {
return 0
}
v, _ := strconv.Atoi(strings.TrimSpace(string(out)))
return v
}
// member is one bundle entry: a file on disk, bytes, or a symlink.
type member struct {
entry ManifestEntry
path string
data []byte
}
func fileMember(name, path string) (member, error) {
f, err := os.Open(path)
if err != nil {
return member{}, err
}
defer f.Close()
info, err := f.Stat()
if err != nil {
return member{}, err
}
h := sha256.New()
n, err := io.Copy(h, f)
if err != nil {
return member{}, err
}
return member{entry: ManifestEntry{Name: name, Size: n, SHA256: hex.EncodeToString(h.Sum(nil)), Mode: uint32(info.Mode().Perm())}, path: path}, nil
}
func bytesMember(name string, data []byte) member {
s := sha256.Sum256(data)
return member{entry: ManifestEntry{Name: name, Size: int64(len(data)), SHA256: hex.EncodeToString(s[:]), Mode: 0o600}, data: data}
}
// stateMembers bundles the regular files and symlinks directly in dir, under
// state/<absolute dir>/. Subdirectories are not state bootstrap writes.
func stateMembers(dir string) ([]member, error) {
abs, err := filepath.Abs(dir)
if err != nil {
return nil, err
}
entries, err := os.ReadDir(abs)
if err != nil {
return nil, err
}
prefix := stateEntry + filepath.ToSlash(abs) + "/"
var out []member
for _, e := range entries {
if stateSkip[e.Name()] {
continue
}
p := filepath.Join(abs, e.Name())
switch {
case e.Type()&os.ModeSymlink != 0:
target, err := os.Readlink(p)
if err != nil {
return nil, err
}
out = append(out, member{entry: ManifestEntry{Name: prefix + e.Name(), Mode: 0o777, Link: target}})
case e.Type().IsRegular():
m, err := fileMember(prefix+e.Name(), p)
if err != nil {
return nil, err
}
out = append(out, m)
}
}
return out, nil
}
// writeBundle writes MANIFEST.json and the members to final via a .partial and
// returns the bundle's sha256.
func writeBundle(dir, final string, m Manifest, members []member) (string, error) {
manifest, err := json.MarshalIndent(m, "", " ")
if err != nil {
return "", err
}
partial := filepath.Join(dir, "."+filepath.Base(final)+".partial")
f, err := os.OpenFile(partial, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600)
if err != nil {
return "", err
}
defer os.Remove(partial)
h := sha256.New()
if err := writeTar(io.MultiWriter(f, h), m.CreatedAt, manifest, members); err != nil {
f.Close()
return "", fmt.Errorf("write bundle: %w", err)
}
if err := f.Sync(); err != nil {
f.Close()
return "", err
}
if err := f.Close(); err != nil {
return "", err
}
if err := os.Rename(partial, final); err != nil {
return "", err
}
if err := syncDir(dir); err != nil {
return "", err
}
return hex.EncodeToString(h.Sum(nil)), nil
}
func writeFileAtomic(path string, data []byte) error {
dir := filepath.Dir(path)
partial := filepath.Join(dir, "."+filepath.Base(path)+".partial")
defer os.Remove(partial)
f, err := os.OpenFile(partial, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600)
if err != nil {
return err
}
if _, err := f.Write(data); err != nil {
f.Close()
return err
}
if err := f.Sync(); err != nil {
f.Close()
return err
}
if err := f.Close(); err != nil {
return err
}
if err := os.Rename(partial, path); err != nil {
return err
}
return syncDir(dir)
}
func syncDir(dir string) error {
d, err := os.Open(dir)
if err != nil {
return err
}
defer d.Close()
return d.Sync()
}
// Age renders a bundle's age for a human: seconds under a minute, then
// minutes, then days and hours past two days.
func Age(d time.Duration) string {
switch {
case d < 0:
return "0s"
case d < time.Minute:
return d.Truncate(time.Second).String()
case d < 48*time.Hour:
return strings.TrimSuffix(d.Truncate(time.Minute).String(), "0s")
default:
return fmt.Sprintf("%dd%dh", d/(24*time.Hour), d%(24*time.Hour)/time.Hour)
}
}
// Check reports the newest bundle in dir and an error when there is none or it
// is older than maxAge.
func Check(dir string, maxAge time.Duration, now time.Time) (*Bundle, error) {
all, err := List(dir)
if err != nil {
return nil, err
}
if len(all) == 0 {
return nil, fmt.Errorf("no database backup in %s", dir)
}
newest := all[0]
if age := now.Sub(newest.Created); age > maxAge {
return &newest, fmt.Errorf("newest database backup %s is %s old (limit %s)", newest.Name, Age(age), maxAge)
}
return &newest, nil
}
+512
View File
@@ -0,0 +1,512 @@
package dbbackup
import (
"archive/tar"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"slices"
"strings"
"testing"
"time"
)
// The PostgreSQL client tools are faked with shell scripts over a one-file
// "database" ($FAKE_DIR/db). The fake psql only writes the replayed rows back
// when its input ends in COMMIT, which is how a real server treats an open
// transaction at disconnect, so rollback-on-failure is observable.
const fakePGDump = `#!/bin/sh
D="$FAKE_DIR"
if [ "$1" = "--version" ]; then echo "pg_dump (PostgreSQL) 13.23"; exit 0; fi
printf '%s\n' "$*" > "$D/pg_dump.args"
printf '%s' "$PGPASSWORD" > "$D/pg_dump.password"
[ -f "$D/dump_fail" ] && { echo "pg_dump: error: connection refused" >&2; exit 1; }
[ -f "$D/dump_denied" ] && { printf 'pg_dump: error: query failed: ERROR: permission denied for table servers_preserve\npg_dump: error: query was: LOCK TABLE public.servers_preserve IN ACCESS SHARE MODE\n' >&2; exit 1; }
for a in "$@"; do case "$a" in --file=*) out="${a#--file=}";; esac; done
if [ -f "$D/dump_garbage" ]; then echo garbage > "$out"; exit 0; fi
{ printf 'PGDMP\n'; cat "$D/db"; } > "$out"
`
const fakePGRestore = `#!/bin/sh
D="$FAKE_DIR"
list=0
for a in "$@"; do case "$a" in --list) list=1;; esac; last="$a"; done
head -n 1 "$last" | grep -q '^PGDMP$' || { echo "pg_restore: error: input file does not appear to be a valid archive" >&2; exit 1; }
[ $list = 1 ] && { echo "; Archive created"; exit 0; }
echo "-- restore script"
tail -n +2 "$last" | sed 's/^/DATA /'
[ -f "$D/restore_fail" ] && { echo "pg_restore: error: could not read input" >&2; exit 1; }
exit 0
`
const fakePSQL = `#!/bin/sh
D="$FAKE_DIR"
printf '%s\n' "$@" > "$D/psql.args"
q=""
while [ $# -gt 0 ]; do case "$1" in -c) q="$2"; shift;; esac; shift; done
if [ -n "$q" ]; then
case "$q" in
*schema_migrations*) echo 21;;
*pg_stat_activity*) cat "$D/clients" 2>/dev/null || echo 0;;
esac
exit 0
fi
cat > "$D/psql.in"
# The freshness record is its own psql call; keep it apart from the replay.
if grep -q platform_settings "$D/psql.in"; then
mv "$D/psql.in" "$D/record.stdin"; cp "$D/psql.args" "$D/record.args"; exit 0
fi
mv "$D/psql.in" "$D/psql.stdin"
[ -f "$D/psql_fail" ] && { echo 'ERROR: relation "x" already exists' >&2; exit 3; }
tail -n 1 "$D/psql.stdin" | grep -q '^COMMIT;$' || exit 0
grep '^DATA ' "$D/psql.stdin" | sed 's/^DATA //' > "$D/db"
`
const testURL = "postgres://felis:[email protected]:5432/felis?sslmode=disable"
type fakePG struct {
dir string
tools Tools
}
func newFakePG(t *testing.T, db string) *fakePG {
t.Helper()
dir := t.TempDir()
write := func(name, body string) string {
p := filepath.Join(dir, name)
if err := os.WriteFile(p, []byte(body), 0o755); err != nil {
t.Fatal(err)
}
return p
}
f := &fakePG{dir: dir, tools: Tools{
PGDump: write("pg_dump", fakePGDump), PGRestore: write("pg_restore", fakePGRestore), PSQL: write("psql", fakePSQL),
}}
f.setDB(t, db)
t.Setenv("FAKE_DIR", dir)
return f
}
func (f *fakePG) setDB(t *testing.T, s string) {
t.Helper()
if err := os.WriteFile(filepath.Join(f.dir, "db"), []byte(s), 0o600); err != nil {
t.Fatal(err)
}
}
func (f *fakePG) db(t *testing.T) string {
t.Helper()
b, err := os.ReadFile(filepath.Join(f.dir, "db"))
if err != nil {
t.Fatal(err)
}
return string(b)
}
func (f *fakePG) flag(t *testing.T, name, content string) {
t.Helper()
if err := os.WriteFile(filepath.Join(f.dir, name), []byte(content), 0o600); err != nil {
t.Fatal(err)
}
}
func stateDir(t *testing.T) string {
t.Helper()
d := t.TempDir()
for name, body := range map[string]string{
"secrets.env": "DB_PASSWORD=s3cret-pw\n",
"felis.host.toml": "[database]\n",
"bootstrap.done": "2026-09-24T00:00:00Z\n",
} {
if err := os.WriteFile(filepath.Join(d, name), []byte(body), 0o600); err != nil {
t.Fatal(err)
}
}
if err := os.Symlink(filepath.Join(d, "felis.host.toml"), filepath.Join(d, "felis.toml")); err != nil {
t.Fatal(err)
}
return d
}
var t0 = time.Date(2026, 9, 24, 3, 30, 0, 0, time.UTC)
// recorded is the Status of the last freshness record, or the zero Status.
func (pg *fakePG) recorded(t *testing.T) Status {
t.Helper()
args, _ := os.ReadFile(filepath.Join(pg.dir, "record.args"))
var st Status
for _, a := range strings.Split(string(args), "\n") {
if v, ok := strings.CutPrefix(a, "v="); ok {
if err := json.Unmarshal([]byte(v), &st); err != nil {
t.Fatal(err)
}
}
}
return st
}
func at(t time.Time) func() time.Time { return func() time.Time { return t } }
func TestBackupWritesAVerifiableBundle(t *testing.T) {
pg := newFakePG(t, "users: alice\n")
dir := filepath.Join(t.TempDir(), "db-backups")
state := stateDir(t)
path, err := Backup(context.Background(), BackupOptions{
DatabaseURL: testURL, Dir: dir, Label: LabelDaily, StateDir: state, Version: "v1.2.3",
Tools: pg.tools, Now: at(t0),
ExportServers: func(context.Context) ([]byte, error) { return []byte(`{"kind":"List","items":[]}`), nil },
})
if err != nil {
t.Fatalf("Backup: %v", err)
}
if want := filepath.Join(dir, "felis-db-20260924T033000Z-daily.tar"); path != want {
t.Fatalf("path = %s, want %s", path, want)
}
for p, mode := range map[string]os.FileMode{dir: 0o700, path: 0o600, path + ".sha256": 0o600} {
info, err := os.Stat(p)
if err != nil {
t.Fatal(err)
}
if info.Mode().Perm() != mode {
t.Errorf("%s mode = %v, want %v", p, info.Mode().Perm(), mode)
}
}
// The password reaches pg_dump through the environment, never argv.
args, _ := os.ReadFile(filepath.Join(pg.dir, "pg_dump.args"))
if strings.Contains(string(args), "s3cret-pw") {
t.Errorf("password on the pg_dump command line: %s", args)
}
if pw, _ := os.ReadFile(filepath.Join(pg.dir, "pg_dump.password")); string(pw) != "s3cret-pw" {
t.Errorf("PGPASSWORD = %q", pw)
}
m, err := Verify(path)
if err != nil {
t.Fatalf("Verify: %v", err)
}
if m.Label != LabelDaily || m.FelisVersion != "v1.2.3" || m.SchemaVersion != 21 || !m.CreatedAt.Equal(t0) {
t.Errorf("manifest = %+v", m)
}
if m.Database != (DatabaseInfo{Host: "127.0.0.1", Port: "5432", Name: "felis", User: "felis"}) {
t.Errorf("database = %+v", m.Database)
}
if !strings.Contains(m.PGDumpVersion, "13.23") {
t.Errorf("pg_dump version = %q", m.PGDumpVersion)
}
var names []string
for _, f := range m.Files {
names = append(names, f.Name)
}
abs, _ := filepath.Abs(state)
prefix := "state" + filepath.ToSlash(abs) + "/"
want := []string{"db.dump", prefix + "felis.host.toml", prefix + "felis.toml", prefix + "secrets.env", "k8s/minecraftservers.json"}
if !slices.Equal(names, want) {
t.Errorf("members = %v, want %v (bootstrap.done left out)", names, want)
}
for _, f := range m.Files {
if f.Name == prefix+"felis.toml" && f.Link != filepath.Join(state, "felis.host.toml") {
t.Errorf("symlink recorded as %+v", f)
}
}
// No scratch or partial files survive a successful run.
entries, _ := os.ReadDir(dir)
for _, e := range entries {
if strings.HasSuffix(e.Name(), ".partial") || strings.HasSuffix(e.Name(), ".dump") {
t.Errorf("leftover %s", e.Name())
}
}
}
func TestBackupRecordsFreshness(t *testing.T) {
pg := newFakePG(t, "x\n")
dir := t.TempDir()
metrics := filepath.Join(t.TempDir(), "textfile", "felis_db_backup.prom")
path, err := Backup(context.Background(), BackupOptions{
DatabaseURL: testURL, Dir: dir, Label: LabelDaily, Version: "v9", Tools: pg.tools, Now: at(t0),
Record: true, MetricsFile: metrics,
})
if err != nil {
t.Fatal(err)
}
stdin, _ := os.ReadFile(filepath.Join(pg.dir, "record.stdin"))
if !strings.Contains(string(stdin), "INSERT INTO platform_settings") || !strings.Contains(string(stdin), ":'v'::jsonb") {
t.Fatalf("record SQL = %q", stdin)
}
st := pg.recorded(t)
info, _ := os.Stat(path)
if st.Name != filepath.Base(path) || st.Label != LabelDaily || !st.At.Equal(t0) || st.SizeBytes != info.Size() || st.SchemaVersion != 21 {
t.Fatalf("recorded %+v", st)
}
prom, err := os.ReadFile(metrics)
if err != nil {
t.Fatal(err)
}
want := fmt.Sprintf("felis_db_backup_last_success_timestamp_seconds{label=\"daily\"} %d\n", t0.Unix())
if !strings.Contains(string(prom), want) {
t.Fatalf("metrics = %s, want %s", prom, want)
}
if fi, _ := os.Stat(metrics); fi.Mode().Perm() != 0o644 {
t.Errorf("metrics mode %v", fi.Mode().Perm())
}
}
func TestBackupRecordsAClusterThatDidNotAnswer(t *testing.T) {
pg := newFakePG(t, "x\n")
dir := t.TempDir()
path, err := Backup(context.Background(), BackupOptions{
DatabaseURL: testURL, Dir: dir, Label: LabelDaily, Tools: pg.tools, Now: at(t0),
ExportServers: func(context.Context) ([]byte, error) { return nil, errors.New("connection refused") },
})
if err != nil {
t.Fatalf("a cluster outage must not fail the database backup: %v", err)
}
m, err := Verify(path)
if err != nil {
t.Fatal(err)
}
if m.ServersError != "connection refused" || len(m.Files) != 1 {
t.Errorf("manifest = %+v", m)
}
}
func TestBackupsWithinOneSecondGetDistinctNames(t *testing.T) {
pg := newFakePG(t, "x\n")
dir := t.TempDir()
o := BackupOptions{DatabaseURL: testURL, Dir: dir, Label: LabelPreRestore, Tools: pg.tools, Now: at(t0)}
first, err := Backup(context.Background(), o)
if err != nil {
t.Fatal(err)
}
second, err := Backup(context.Background(), o)
if err != nil {
t.Fatalf("second backup in the same second: %v", err)
}
if first == second {
t.Fatalf("both backups wrote %s", first)
}
all, err := List(dir)
if err != nil || len(all) != 2 || !all[0].Created.After(all[1].Created) {
t.Fatalf("list = %+v, %v", all, err)
}
if _, err := Verify(second); err != nil {
t.Fatal(err)
}
}
func TestBackupFailures(t *testing.T) {
for _, tc := range []struct {
flag, want string
}{
{"dump_fail", "connection refused"},
{"dump_garbage", "does not read back"},
// An object another role created in the database: the error names the fix.
{"dump_denied", "sudo -u postgres psql -d felis -c 'ALTER TABLE servers_preserve OWNER TO felis'"},
} {
t.Run(tc.flag, func(t *testing.T) {
pg := newFakePG(t, "x\n")
pg.flag(t, tc.flag, "")
dir := t.TempDir()
_, err := Backup(context.Background(), BackupOptions{DatabaseURL: testURL, Dir: dir, Label: LabelDaily, Tools: pg.tools, Now: at(t0)})
if err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("err = %v, want %q", err, tc.want)
}
entries, _ := os.ReadDir(dir)
for _, e := range entries {
if e.Name() != ".lock" {
t.Errorf("a failed backup left %s behind", e.Name())
}
}
})
}
t.Run("bad label and url", func(t *testing.T) {
pg := newFakePG(t, "x\n")
if _, err := Backup(context.Background(), BackupOptions{DatabaseURL: testURL, Dir: t.TempDir(), Label: "../x", Tools: pg.tools}); err == nil {
t.Error("label with a path separator accepted")
}
if _, err := Backup(context.Background(), BackupOptions{DatabaseURL: "host=x dbname=y", Dir: t.TempDir(), Label: "daily", Tools: pg.tools}); err == nil {
t.Error("non-URL database accepted")
}
})
t.Run("stale partials from a crashed run are cleared", func(t *testing.T) {
pg := newFakePG(t, "x\n")
dir := t.TempDir()
stale := filepath.Join(dir, ".felis-db-20260101T000000Z-daily.tar.partial")
if err := os.WriteFile(stale, []byte("half"), 0o600); err != nil {
t.Fatal(err)
}
if _, err := Backup(context.Background(), BackupOptions{DatabaseURL: testURL, Dir: dir, Label: LabelDaily, Tools: pg.tools, Now: at(t0)}); err != nil {
t.Fatal(err)
}
if _, err := os.Stat(stale); !errors.Is(err, os.ErrNotExist) {
t.Error("stale partial survived")
}
})
}
func TestVerifyCatchesCorruption(t *testing.T) {
pg := newFakePG(t, strings.Repeat("row\n", 64))
dir := t.TempDir()
path, err := Backup(context.Background(), BackupOptions{DatabaseURL: testURL, Dir: dir, Label: LabelManual, Tools: pg.tools, Now: at(t0)})
if err != nil {
t.Fatal(err)
}
raw, _ := os.ReadFile(path)
i := strings.Index(string(raw), "PGDMP")
raw[i+10] ^= 0x20
if err := os.WriteFile(path, raw, 0o600); err != nil {
t.Fatal(err)
}
if _, err := Verify(path); err == nil || !strings.Contains(err.Error(), "corrupt") {
t.Fatalf("flipped dump byte: err = %v, want member corrupt", err)
}
// A bundle intact inside but not the one the sidecar vouches for.
raw[i+10] ^= 0x20
raw = append(raw, make([]byte, 512)...)
if err := os.WriteFile(path, raw, 0o600); err != nil {
t.Fatal(err)
}
if _, err := Verify(path); err == nil || !strings.Contains(err.Error(), ".sha256") {
t.Fatalf("sidecar mismatch: err = %v", err)
}
junk := filepath.Join(dir, "junk.tar")
if err := os.WriteFile(junk, []byte("not a tar"), 0o600); err != nil {
t.Fatal(err)
}
if _, err := Verify(junk); err == nil {
t.Fatal("junk verified")
}
}
func TestVerifyRejectsUnlistedMember(t *testing.T) {
pg := newFakePG(t, "x\n")
dir := t.TempDir()
path, err := Backup(context.Background(), BackupOptions{DatabaseURL: testURL, Dir: dir, Label: LabelManual, Tools: pg.tools, Now: at(t0)})
if err != nil {
t.Fatal(err)
}
// Re-pack with an extra member the manifest does not list.
src, _ := os.Open(path)
defer src.Close()
out := filepath.Join(dir, "felis-db-20260924T040000Z-manual.tar")
dst, _ := os.Create(out)
tr, tw := tar.NewReader(src), tar.NewWriter(dst)
for {
h, err := tr.Next()
if err == io.EOF {
break
}
if err != nil {
t.Fatal(err)
}
_ = tw.WriteHeader(h)
_, _ = io.Copy(tw, tr)
}
_ = tw.WriteHeader(&tar.Header{Name: "state/etc/cron.d/evil", Mode: 0o644, Size: 1, Typeflag: tar.TypeReg})
_, _ = tw.Write([]byte("x"))
_ = tw.Close()
_ = dst.Close()
if _, err := Verify(out); err == nil || !strings.Contains(err.Error(), "not in the manifest") {
t.Fatalf("err = %v", err)
}
}
func TestListPruneCheck(t *testing.T) {
pg := newFakePG(t, "x\n")
dir := t.TempDir()
backup := func(label string, when time.Time, keep int) {
t.Helper()
if _, err := Backup(context.Background(), BackupOptions{DatabaseURL: testURL, Dir: dir, Label: label, Keep: keep, Tools: pg.tools, Now: at(when)}); err != nil {
t.Fatal(err)
}
}
backup(LabelManual, t0.Add(-100*time.Hour), 0)
for i := range 5 {
backup(LabelDaily, t0.Add(time.Duration(i-5)*24*time.Hour), 3)
}
backup(LabelPreMigrate, t0.Add(-time.Hour), 3)
all, err := List(dir)
if err != nil {
t.Fatal(err)
}
var got []string
for _, b := range all {
got = append(got, b.Label+"@"+b.Created.Format("0102T15"))
}
want := []string{"pre-migrate@0924T02", "daily@0923T03", "daily@0922T03", "daily@0921T03", "manual@0919T23"}
if !slices.Equal(got, want) {
t.Fatalf("List = %v, want %v (dailies pruned to 3, others untouched)", got, want)
}
if _, err := os.Stat(filepath.Join(dir, BundleName(t0.Add(-5*24*time.Hour), LabelDaily)+".sha256")); !errors.Is(err, os.ErrNotExist) {
t.Error("a pruned bundle's sidecar survived")
}
if b, err := Check(dir, 26*time.Hour, t0); err != nil || b.Label != LabelPreMigrate {
t.Errorf("Check fresh = %v, %v", b, err)
}
if _, err := Check(dir, 26*time.Hour, t0.Add(30*time.Hour)); err == nil {
t.Error("Check accepted a 31h-old newest bundle")
}
if _, err := Check(filepath.Join(dir, "none"), time.Hour, t0); err == nil {
t.Error("Check accepted an empty directory")
}
}
func TestParseBundleName(t *testing.T) {
for name, ok := range map[string]bool{
"felis-db-20260924T033000Z-daily.tar": true,
"felis-db-20260924T033000Z-pre-migrate.tar": true,
"felis-db-20260924T033000Z-.tar": false,
"felis-db-20260924T033000Z-Daily.tar": false,
"felis-db-2026-daily.tar": false,
"felis-db-20260924T033000Z-daily.tar.sha256": false,
"other.tar": false,
} {
if _, _, got := parseBundleName(name); got != ok {
t.Errorf("parseBundleName(%q) ok = %v, want %v", name, got, ok)
}
}
}
func TestParseConnStripsPassword(t *testing.T) {
c, err := parseConn(testURL)
if err != nil {
t.Fatal(err)
}
if strings.Contains(c.uri, "s3cret") || c.password != "s3cret-pw" {
t.Fatalf("conn = %+v", c)
}
if !strings.Contains(c.uri, "sslmode=disable") || !strings.HasPrefix(c.uri, "postgres://[email protected]:5432/felis") {
t.Fatalf("uri = %s", c.uri)
}
if _, err := parseConn("postgres://127.0.0.1/"); err == nil {
t.Fatal("URL without a database accepted")
}
}
func TestAge(t *testing.T) {
for d, want := range map[time.Duration]string{
-time.Second: "0s",
44 * time.Second: "44s",
90 * time.Second: "1m",
26*time.Hour + 5*time.Minute: "26h5m",
3*24*time.Hour + 4*time.Hour: "3d4h",
2*time.Hour + 30*time.Second: "2h0m",
} {
if got := Age(d); got != want {
t.Errorf("Age(%s) = %q, want %q", d, got, want)
}
}
}
+346
View File
@@ -0,0 +1,346 @@
package dbbackup
import (
"archive/tar"
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"strconv"
"strings"
"time"
)
// writeTar streams MANIFEST.json and then every member.
func writeTar(w io.Writer, mtime time.Time, manifest []byte, members []member) error {
tw := tar.NewWriter(w)
hdr := func(name string, mode uint32, size int64) *tar.Header {
return &tar.Header{Name: name, Mode: int64(mode), Size: size, ModTime: mtime, Typeflag: tar.TypeReg, Format: tar.FormatPAX}
}
if err := tw.WriteHeader(hdr(manifestEntry, 0o600, int64(len(manifest)))); err != nil {
return err
}
if _, err := tw.Write(manifest); err != nil {
return err
}
for _, m := range members {
if m.entry.Link != "" {
if err := tw.WriteHeader(&tar.Header{Name: m.entry.Name, Linkname: m.entry.Link, Mode: 0o777,
ModTime: mtime, Typeflag: tar.TypeSymlink, Format: tar.FormatPAX}); err != nil {
return err
}
continue
}
if err := tw.WriteHeader(hdr(m.entry.Name, m.entry.Mode, m.entry.Size)); err != nil {
return err
}
if m.data != nil {
if _, err := tw.Write(m.data); err != nil {
return err
}
continue
}
f, err := os.Open(m.path)
if err != nil {
return err
}
// CopyN: the size went into the header from the hash pass; a file that
// changed since is an error here, not a silently short member.
_, err = io.CopyN(tw, f, m.entry.Size)
f.Close()
if err != nil {
return fmt.Errorf("%s: %w", m.entry.Name, err)
}
}
return tw.Close()
}
// Verify reads a whole bundle, checks it against its sidecar (when present)
// and every member against the manifest, and returns the manifest.
func Verify(path string) (Manifest, error) {
return readBundle(path, nil)
}
// readBundle is Verify that also copies db.dump to dumpTo when non-nil.
func readBundle(path string, dumpTo io.Writer) (Manifest, error) {
var m Manifest
f, err := os.Open(path)
if err != nil {
return m, err
}
defer f.Close()
whole := sha256.New()
tr := tar.NewReader(io.TeeReader(f, whole))
first, err := tr.Next()
if err != nil || first.Name != manifestEntry {
return m, fmt.Errorf("%s is not a felis database bundle (no %s)", filepath.Base(path), manifestEntry)
}
raw, err := io.ReadAll(io.LimitReader(tr, 1<<20))
if err != nil {
return m, err
}
if err := json.Unmarshal(raw, &m); err != nil {
return m, fmt.Errorf("read %s: %w", manifestEntry, err)
}
if m.Format != formatV1 {
return m, fmt.Errorf("bundle format %d is not one this felis reads (want %d)", m.Format, formatV1)
}
want := map[string]ManifestEntry{}
for _, e := range m.Files {
want[e.Name] = e
}
seen := map[string]bool{}
for {
h, err := tr.Next()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return m, fmt.Errorf("read bundle: %w", err)
}
e, ok := want[h.Name]
if !ok {
return m, fmt.Errorf("bundle member %s is not in the manifest", h.Name)
}
seen[h.Name] = true
if h.Typeflag == tar.TypeSymlink {
if h.Linkname != e.Link {
return m, fmt.Errorf("bundle member %s links to %q, manifest says %q", h.Name, h.Linkname, e.Link)
}
continue
}
dst := io.Discard
if h.Name == dumpEntry && dumpTo != nil {
dst = dumpTo
}
sum := sha256.New()
n, err := io.Copy(io.MultiWriter(dst, sum), tr)
if err != nil {
return m, fmt.Errorf("read %s: %w", h.Name, err)
}
if n != e.Size || hex.EncodeToString(sum.Sum(nil)) != e.SHA256 {
return m, fmt.Errorf("bundle member %s is corrupt (size or sha256 differs from the manifest)", h.Name)
}
}
for name := range want {
if !seen[name] {
return m, fmt.Errorf("bundle is missing %s", name)
}
}
if _, ok := want[dumpEntry]; !ok {
return m, fmt.Errorf("bundle holds no %s", dumpEntry)
}
// The tar end marker is not the end of the file; the sidecar covers every byte.
if _, err := io.Copy(io.Discard, io.TeeReader(f, whole)); err != nil {
return m, err
}
if sidecar, err := os.ReadFile(path + sumExt); err == nil {
fields := strings.Fields(string(sidecar))
if len(fields) == 0 || fields[0] != hex.EncodeToString(whole.Sum(nil)) {
return m, fmt.Errorf("%s does not match %s%s", filepath.Base(path), filepath.Base(path), sumExt)
}
}
return m, nil
}
// RestoreOptions configures a restore.
type RestoreOptions struct {
DatabaseURL string
Bundle string
// Dir holds the scratch copy of the dump and the pre-restore safety bundle.
Dir string
// Force restores even while other clients are connected to the database.
Force bool
// SkipSafetyBackup skips the bundle of the current database taken before it
// is replaced.
SkipSafetyBackup bool
// Safety configures that bundle; its DatabaseURL, Dir and Label are set here.
Safety BackupOptions
Tools Tools
Log io.Writer
}
// ErrClientsConnected refuses a restore under live clients: the replay needs
// exclusive locks on every table, and felis-api would be serving from a
// database that is about to change under it.
var ErrClientsConnected = errors.New("other clients are connected to the database")
// Restore replaces the database's contents with the bundle's dump, atomically.
// It returns the bundle's manifest and, unless skipped, the path of the safety
// bundle of what was there before.
func Restore(ctx context.Context, o RestoreOptions) (Manifest, string, error) {
c, err := parseConn(o.DatabaseURL)
if err != nil {
return Manifest{}, "", err
}
logw := o.Log
if logw == nil {
logw = io.Discard
}
if err := os.MkdirAll(o.Dir, 0o700); err != nil {
return Manifest{}, "", err
}
scratch, err := os.CreateTemp(o.Dir, ".restore-*.dump")
if err != nil {
return Manifest{}, "", err
}
defer os.Remove(scratch.Name())
m, err := readBundle(o.Bundle, scratch)
if cerr := scratch.Close(); err == nil {
err = cerr
}
if err != nil {
return m, "", err
}
if _, err := run(exec.CommandContext(ctx, o.Tools.pgRestore(), "--list", scratch.Name())); err != nil {
return m, "", fmt.Errorf("the bundle's dump does not read: %w", err)
}
if !o.Force {
n, err := otherClients(ctx, c, o.Tools)
if err != nil {
return m, "", fmt.Errorf("count connected clients: %w", err)
}
if n > 0 {
return m, "", fmt.Errorf("%w (%d); scale felis-api and felis-operator to 0 first, or pass -force", ErrClientsConnected, n)
}
}
var safety string
if !o.SkipSafetyBackup {
so := o.Safety
so.DatabaseURL, so.Dir, so.Label, so.Tools = o.DatabaseURL, o.Dir, LabelPreRestore, o.Tools
if so.Log == nil {
so.Log = logw
}
if safety, err = Backup(ctx, so); err != nil {
return m, "", fmt.Errorf("safety backup of the current database: %w (pass -no-safety-backup to restore without one)", err)
}
fmt.Fprintf(logw, "felis db restore: current database saved to %s\n", safety)
}
if err := replay(ctx, c, o.Tools, scratch.Name()); err != nil {
return m, safety, err
}
// The dump carried its own freshness record, older than the bundle it is in
// (the record is written after the bundle). Point it at the newest bundle on
// disk, or the panel reports a missing backup right after a restore.
if err := recordNewest(ctx, c, o.Tools, o.Dir); err != nil {
fmt.Fprintf(logw, "felis db restore: record the newest backup for the panel: %v\n", err)
}
return m, safety, nil
}
// recordNewest records the newest bundle in dir as the latest backup.
func recordNewest(ctx context.Context, c conn, t Tools, dir string) error {
all, err := List(dir)
if err != nil || len(all) == 0 {
return err
}
b := all[0]
st := Status{At: b.Created, Name: b.Name, Label: b.Label, SizeBytes: b.Size, Dir: dir}
if m, err := Verify(b.Path); err == nil {
st.FelisVersion, st.SchemaVersion = m.FelisVersion, m.SchemaVersion
}
return record(ctx, c, t, st)
}
// otherClients counts client sessions on the database other than this one.
func otherClients(ctx context.Context, c conn, t Tools) (int, error) {
out, err := run(c.command(ctx, t.psql(), "-X", "-q", "-t", "-A", "-w", "-d", c.uri, "-c",
"SELECT count(*) FROM pg_stat_activity WHERE datname = current_database() AND pid <> pg_backend_pid() AND backend_type = 'client backend'"))
if err != nil {
return 0, err
}
return strconv.Atoi(strings.TrimSpace(string(out)))
}
// dropOwned clears everything the connecting role owns in the database, the
// first statement of the restore transaction.
const dropOwned = "DROP OWNED BY CURRENT_USER;\n"
// replay pipes `pg_restore --file=-` into one psql transaction that starts by
// dropping what the role owns. The COMMIT is only written once pg_restore has
// exited cleanly: a generator that dies mid-stream leaves psql at EOF inside an
// open transaction, which the server rolls back when psql disconnects. psql's
// own --single-transaction would commit whatever arrived before that EOF.
func replay(ctx context.Context, c conn, t Tools, dump string) error {
ctx, cancel := context.WithCancel(ctx)
defer cancel()
r, w, err := os.Pipe()
if err != nil {
return err
}
var restoreErr, psqlErr bytes.Buffer
gen := exec.CommandContext(ctx, t.pgRestore(), "--no-owner", "--no-privileges", "--file=-", dump)
gen.Stdout, gen.Stderr = w, &restoreErr
if err := gen.Start(); err != nil {
r.Close()
w.Close()
return fmt.Errorf("pg_restore: %w", err)
}
w.Close()
tail := &commitAfter{gen: gen}
apply := c.command(ctx, t.psql(), "-X", "-q", "-w", "-v", "ON_ERROR_STOP=1", "-d", c.uri)
apply.Stdin = io.MultiReader(strings.NewReader("BEGIN;\n"+dropOwned), r, tail)
apply.Stdout, apply.Stderr = io.Discard, &psqlErr
aerr := apply.Run()
// Closing the read end makes a pg_restore still writing (psql stopped early)
// die on the broken pipe instead of blocking, and it is reaped exactly once.
r.Close()
gerr := tail.wait()
if aerr == nil {
// psql read to the end, and the end is a COMMIT only pg_restore's clean
// exit releases.
return nil
}
msg := "replay the dump (rolled back, the database is unchanged)"
if s := strings.TrimSpace(psqlErr.String()); s != "" {
msg += ": psql: " + s
}
if s := strings.TrimSpace(restoreErr.String()); gerr != nil && s != "" {
msg += ": pg_restore: " + s
}
return fmt.Errorf("%s: %w", msg, aerr)
}
// commitAfter yields "COMMIT;" once, and only if the pg_restore feeding the
// pipe exited cleanly; otherwise it fails the stream so psql never sees one.
type commitAfter struct {
gen *exec.Cmd
waited bool
err error
committed bool
rest []byte
}
func (c *commitAfter) wait() error {
if !c.waited {
c.waited, c.err = true, c.gen.Wait()
}
return c.err
}
func (c *commitAfter) Read(p []byte) (int, error) {
if !c.committed {
if err := c.wait(); err != nil {
return 0, err
}
c.committed, c.rest = true, []byte("COMMIT;\n")
}
if len(c.rest) == 0 {
return 0, io.EOF
}
n := copy(p, c.rest)
c.rest = c.rest[n:]
return n, nil
}
+139
View File
@@ -0,0 +1,139 @@
package dbbackup
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
"time"
)
func takeBackup(t *testing.T, pg *fakePG, dir string, when time.Time) string {
t.Helper()
path, err := Backup(context.Background(), BackupOptions{DatabaseURL: testURL, Dir: dir, Label: LabelManual, Tools: pg.tools, Now: at(when)})
if err != nil {
t.Fatal(err)
}
return path
}
func restore(pg *fakePG, dir, bundle string, mut func(*RestoreOptions)) (string, error) {
o := RestoreOptions{DatabaseURL: testURL, Bundle: bundle, Dir: dir, Tools: pg.tools,
Safety: BackupOptions{Now: at(t0.Add(time.Hour))}}
if mut != nil {
mut(&o)
}
_, safety, err := Restore(context.Background(), o)
return safety, err
}
func TestRestoreReplacesTheDatabase(t *testing.T) {
pg := newFakePG(t, "alice\n")
dir := t.TempDir()
bundle := takeBackup(t, pg, dir, t0)
pg.setDB(t, "alice\nbob\n")
safety, err := restore(pg, dir, bundle, nil)
if err != nil {
t.Fatalf("Restore: %v", err)
}
if got := pg.db(t); got != "alice\n" {
t.Fatalf("db after restore = %q", got)
}
// The replay dropped what the role owns inside the same transaction.
stdin, _ := os.ReadFile(filepath.Join(pg.dir, "psql.stdin"))
if !strings.HasPrefix(string(stdin), "BEGIN;\n"+dropOwned) || !strings.HasSuffix(string(stdin), "COMMIT;\n") {
t.Fatalf("psql input = %q", stdin)
}
// The safety bundle holds what was replaced, and restores it.
if filepath.Base(safety) != "felis-db-20260924T043000Z-pre-restore.tar" {
t.Fatalf("safety = %s", safety)
}
// The dump brought back its own, older freshness record; the restore
// points it at the newest bundle on disk again.
if st := pg.recorded(t); st.Name != filepath.Base(safety) || st.Label != LabelPreRestore || st.SchemaVersion != 21 || st.Dir != dir {
t.Fatalf("recorded after restore = %+v", st)
}
if _, err := restore(pg, dir, safety, func(o *RestoreOptions) { o.SkipSafetyBackup = true }); err != nil {
t.Fatal(err)
}
if got := pg.db(t); got != "alice\nbob\n" {
t.Fatalf("db after undo = %q", got)
}
}
func TestRestoreRefusesLiveClients(t *testing.T) {
pg := newFakePG(t, "alice\n")
dir := t.TempDir()
bundle := takeBackup(t, pg, dir, t0)
pg.setDB(t, "changed\n")
pg.flag(t, "clients", "2\n")
if _, err := restore(pg, dir, bundle, nil); !errors.Is(err, ErrClientsConnected) {
t.Fatalf("err = %v, want ErrClientsConnected", err)
}
if pg.db(t) != "changed\n" {
t.Fatal("database touched despite the refusal")
}
if _, err := restore(pg, dir, bundle, func(o *RestoreOptions) { o.Force = true }); err != nil {
t.Fatalf("forced: %v", err)
}
if pg.db(t) != "alice\n" {
t.Fatal("forced restore did not apply")
}
}
func TestRestoreFailureLeavesTheDatabaseAlone(t *testing.T) {
for _, tc := range []struct {
flag, want string
}{
// The replay fails in the server: psql stops, the transaction dies with it.
{"psql_fail", "already exists"},
// The dump reader dies after streaming everything: psql never gets the
// COMMIT that only a clean pg_restore exit releases.
{"restore_fail", "could not read input"},
} {
t.Run(tc.flag, func(t *testing.T) {
pg := newFakePG(t, "alice\n")
dir := t.TempDir()
bundle := takeBackup(t, pg, dir, t0)
pg.setDB(t, "current\n")
pg.flag(t, tc.flag, "")
_, err := restore(pg, dir, bundle, func(o *RestoreOptions) { o.SkipSafetyBackup = true })
if err == nil || !strings.Contains(err.Error(), "rolled back") || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("err = %v, want rolled back + %q", err, tc.want)
}
if got := pg.db(t); got != "current\n" {
t.Fatalf("db = %q, want it untouched", got)
}
})
}
}
func TestRestoreRefusesACorruptBundle(t *testing.T) {
pg := newFakePG(t, "alice\n")
dir := t.TempDir()
bundle := takeBackup(t, pg, dir, t0)
if err := os.WriteFile(bundle+".sha256", []byte("0000 x\n"), 0o600); err != nil {
t.Fatal(err)
}
pg.setDB(t, "current\n")
if _, err := restore(pg, dir, bundle, nil); err == nil {
t.Fatal("restored a bundle its checksum disowns")
}
if pg.db(t) != "current\n" {
t.Fatal("database touched")
}
entries, _ := os.ReadDir(dir)
for _, e := range entries {
if strings.HasPrefix(e.Name(), ".restore-") {
t.Errorf("scratch dump %s left behind", e.Name())
}
if strings.Contains(e.Name(), LabelPreRestore) {
t.Errorf("safety bundle %s taken before the bundle was verified", e.Name())
}
}
}