Files
Felis/internal/dbbackup/restore.go
T

355 lines
11 KiB
Go

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 := listArchive(ctx, c, o.Tools, 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(t.command(ctx, c, t.psql(), "-X", "-q", "-t", "-A", "-w", "-d", t.dsn(c), "-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.
//
// pg_restore reads the archive on stdin, sequentially, which is all a full
// restore needs and the only way in under Tools.Exec.
func replay(ctx context.Context, c conn, t Tools, dump string) error {
ctx, cancel := context.WithCancel(ctx)
defer cancel()
in, err := os.Open(dump)
if err != nil {
return err
}
defer in.Close()
r, w, err := os.Pipe()
if err != nil {
return err
}
var restoreErr, psqlErr bytes.Buffer
gen := t.command(ctx, c, t.pgRestore(), "--no-owner", "--no-privileges", "--file=-")
gen.Stdin, gen.Stdout, gen.Stderr = in, 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.Cmd}
apply := t.command(ctx, c, t.psql(), "-X", "-q", "-w", "-v", "ON_ERROR_STOP=1", "-d", t.dsn(c))
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
}