Compare commits
1
Commits
next
..
68f4ccf5b9
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
68f4ccf5b9 |
+5
-4
@@ -63,7 +63,7 @@ A content-addressed unit of data. Files are split into variable-size chunks usin
|
|||||||
- `ChunkHash`: SHA256 hash of chunk content (primary key)
|
- `ChunkHash`: SHA256 hash of chunk content (primary key)
|
||||||
- `Size`: Chunk size in bytes
|
- `Size`: Chunk size in bytes
|
||||||
|
|
||||||
Chunk sizes vary between `avgChunkSize/4` and `avgChunkSize*4` (2.5MB-40MB for the 10MB default average).
|
Chunk sizes vary between `avgChunkSize/4` and `avgChunkSize*4` (typically 16KB-256KB for 64KB average).
|
||||||
|
|
||||||
#### FileChunk (`database.FileChunk`)
|
#### FileChunk (`database.FileChunk`)
|
||||||
Maps files to their constituent chunks:
|
Maps files to their constituent chunks:
|
||||||
@@ -120,7 +120,7 @@ The CLI uses fx for dependency injection. Here's the instantiation order:
|
|||||||
```go
|
```go
|
||||||
// cli/app.go: NewApp()
|
// cli/app.go: NewApp()
|
||||||
fx.New(
|
fx.New(
|
||||||
fx.Supply(config.Path(opts.ConfigPath)), // 1. Config path
|
fx.Supply(config.ConfigPath(opts.ConfigPath)), // 1. Config path
|
||||||
fx.Supply(opts.LogOptions), // 2. Log options
|
fx.Supply(opts.LogOptions), // 2. Log options
|
||||||
fx.Provide(globals.New), // 3. Globals
|
fx.Provide(globals.New), // 3. Globals
|
||||||
fx.Provide(log.New), // 4. Logger config
|
fx.Provide(log.New), // 4. Logger config
|
||||||
@@ -193,7 +193,7 @@ scanner := v.ScannerFactory(snapshot.ScannerParams{
|
|||||||
- **Created by**: `chunker.NewChunker(avgChunkSize)`
|
- **Created by**: `chunker.NewChunker(avgChunkSize)`
|
||||||
- **When**: Inside `snapshot.NewScanner()`
|
- **When**: Inside `snapshot.NewScanner()`
|
||||||
- **Configuration**:
|
- **Configuration**:
|
||||||
- `avgChunkSize`: From config (default 10MB)
|
- `avgChunkSize`: From config (typically 64KB)
|
||||||
- `minChunkSize`: avgChunkSize / 4
|
- `minChunkSize`: avgChunkSize / 4
|
||||||
- `maxChunkSize`: avgChunkSize * 4
|
- `maxChunkSize`: avgChunkSize * 4
|
||||||
|
|
||||||
@@ -286,6 +286,7 @@ Key methods:
|
|||||||
- `CreateSnapshot(ctx, hostname, version, commit)` → Create snapshot record
|
- `CreateSnapshot(ctx, hostname, version, commit)` → Create snapshot record
|
||||||
- `CompleteSnapshot(ctx, snapshotID)` → Mark snapshot complete
|
- `CompleteSnapshot(ctx, snapshotID)` → Mark snapshot complete
|
||||||
- `ExportSnapshotMetadata(ctx, dbPath, snapshotID)` → Export to S3
|
- `ExportSnapshotMetadata(ctx, dbPath, snapshotID)` → Export to S3
|
||||||
|
- `CleanupIncompleteSnapshots(ctx, hostname)` → Remove failed snapshots
|
||||||
|
|
||||||
### `internal/database`
|
### `internal/database`
|
||||||
SQLite database for local index. Single-writer mode for thread safety.
|
SQLite database for local index. Single-writer mode for thread safety.
|
||||||
@@ -306,7 +307,7 @@ Repository interfaces:
|
|||||||
```
|
```
|
||||||
CreateSnapshot(opts)
|
CreateSnapshot(opts)
|
||||||
│
|
│
|
||||||
├─► PruneDatabase() // Critical: avoid dedup errors
|
├─► CleanupIncompleteSnapshots() // Critical: avoid dedup errors
|
||||||
│
|
│
|
||||||
├─► SnapshotManager.CreateSnapshot() // Create DB record
|
├─► SnapshotManager.CreateSnapshot() // Create DB record
|
||||||
│
|
│
|
||||||
|
|||||||
+3
-24
@@ -20,6 +20,8 @@
|
|||||||
# golang:1.26.1-alpine, 2026-03-17
|
# golang:1.26.1-alpine, 2026-03-17
|
||||||
FROM golang:1.26.1-alpine@sha256:2389ebfa5b7f43eeafbd6be0c3700cc46690ef842ad962f6c5bd6be49ed82039 AS builder
|
FROM golang:1.26.1-alpine@sha256:2389ebfa5b7f43eeafbd6be0c3700cc46690ef842ad962f6c5bd6be49ed82039 AS builder
|
||||||
|
|
||||||
|
ARG VERSION=dev
|
||||||
|
|
||||||
# Build tooling: make, plus a C toolchain because `go test -race` needs cgo.
|
# Build tooling: make, plus a C toolchain because `go test -race` needs cgo.
|
||||||
# The sqlite driver is pure Go (modernc.org/sqlite), so no sqlite library or
|
# The sqlite driver is pure Go (modernc.org/sqlite), so no sqlite library or
|
||||||
# CLI is required.
|
# CLI is required.
|
||||||
@@ -64,31 +66,8 @@ RUN [ -n "$CHECK_EPOCH" ] || exit 1
|
|||||||
RUN echo "check epoch: ${CHECK_EPOCH}" && make fmt-check
|
RUN echo "check epoch: ${CHECK_EPOCH}" && make fmt-check
|
||||||
RUN echo "check epoch: ${CHECK_EPOCH}" && make test
|
RUN echo "check epoch: ${CHECK_EPOCH}" && make test
|
||||||
|
|
||||||
# Version, commit and build date are computed on the host by
|
|
||||||
# script/docker and script/cibuild (where .git exists) and passed in as
|
|
||||||
# build args. The build context excludes .git (see .dockerignore), so
|
|
||||||
# the build cannot derive them itself: it used to try, with `git
|
|
||||||
# rev-parse` inside this stage, and always got "unknown". VERSION comes
|
|
||||||
# from script/version, the source of truth shared with the Makefile, so
|
|
||||||
# it carries the same tag / dev-<sha> / -dirty rules and a Docker image
|
|
||||||
# reports the same string a local build of the same tree would.
|
|
||||||
#
|
|
||||||
# The defaults are the fallback for a bare `docker build .` that passes
|
|
||||||
# none of them: an unset arg would otherwise stamp an empty string and
|
|
||||||
# produce an image that cannot report its own version, commit or date.
|
|
||||||
# They match what an out-of-git build reports elsewhere.
|
|
||||||
#
|
|
||||||
# These ARGs sit here, after the checks, rather than at the top of the
|
|
||||||
# stage: every commit changes their values, and a value change
|
|
||||||
# invalidates all layers below the ARG. Declared up top they would bust
|
|
||||||
# `go mod download`; here they only rekey this build layer, which the
|
|
||||||
# COPY of the sources above already rebuilds on any change anyway.
|
|
||||||
ARG VERSION=dev
|
|
||||||
ARG COMMIT=unknown
|
|
||||||
ARG COMMIT_DATE=unknown
|
|
||||||
|
|
||||||
# Build (pure Go, no CGO required since we use modernc.org/sqlite)
|
# Build (pure Go, no CGO required since we use modernc.org/sqlite)
|
||||||
RUN CGO_ENABLED=0 go build -ldflags "-X 'sneak.berlin/go/vaultik/internal/globals.Version=${VERSION}' -X 'sneak.berlin/go/vaultik/internal/globals.Commit=${COMMIT}' -X 'sneak.berlin/go/vaultik/internal/globals.CommitDate=${COMMIT_DATE}'" -o /vaultik ./cmd/vaultik
|
RUN CGO_ENABLED=0 go build -ldflags "-X 'sneak.berlin/go/vaultik/internal/globals.Version=${VERSION}' -X 'sneak.berlin/go/vaultik/internal/globals.Commit=$(git rev-parse HEAD 2>/dev/null || echo unknown)' -X 'sneak.berlin/go/vaultik/internal/globals.CommitDate=$(git show -s --format=%cs HEAD 2>/dev/null || echo unknown)'" -o /vaultik ./cmd/vaultik
|
||||||
|
|
||||||
# Runtime stage
|
# Runtime stage
|
||||||
# alpine:3.21, 2026-02-25
|
# alpine:3.21, 2026-02-25
|
||||||
|
|||||||
@@ -71,19 +71,14 @@ Requirements that no existing tool meets:
|
|||||||
## daily use
|
## daily use
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
# verify a snapshot (shallow: checks all blobs are present with the listed size)
|
# verify a snapshot (shallow: checks all blobs exist)
|
||||||
vaultik snapshot verify <snapshot-id>
|
vaultik snapshot verify <snapshot-id>
|
||||||
|
|
||||||
# put the private key file in the environment (reading it from the file
|
|
||||||
# keeps the key out of your shell history); the whole age-keygen file,
|
|
||||||
# with one or more identities, is accepted
|
|
||||||
export VAULTIK_AGE_SECRET_KEY="$(cat vaultik_backup_private_key.txt)"
|
|
||||||
|
|
||||||
# deep verify (downloads and cryptographically verifies every blob)
|
# deep verify (downloads and cryptographically verifies every blob)
|
||||||
vaultik snapshot verify --deep <snapshot-id>
|
VAULTIK_AGE_SECRET_KEY='AGE-SECRET-KEY-...' vaultik snapshot verify --deep <snapshot-id>
|
||||||
|
|
||||||
# restore (requires the private key)
|
# restore (requires the private key)
|
||||||
vaultik snapshot restore <snapshot-id> /tmp/restored
|
VAULTIK_AGE_SECRET_KEY='AGE-SECRET-KEY-...' vaultik snapshot restore <snapshot-id> /tmp/restored
|
||||||
|
|
||||||
# daily cron job: back up, keep a 4-week rolling window of snapshots
|
# daily cron job: back up, keep a 4-week rolling window of snapshots
|
||||||
# 0 3 * * * vaultik snapshot create --cron --prune --keep-newer-than 4w
|
# 0 3 * * * vaultik snapshot create --cron --prune --keep-newer-than 4w
|
||||||
@@ -124,17 +119,15 @@ Use that remote key — the hex printed inside `<remote only:...>`, or the
|
|||||||
full `remote_key` from `snapshot list --json` — to restore and verify:
|
full `remote_key` from `snapshot list --json` — to restore and verify:
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
# put the private key file in the environment (reading it from the file
|
|
||||||
# keeps the key out of your shell history)
|
|
||||||
export VAULTIK_AGE_SECRET_KEY="$(cat vaultik_backup_private_key.txt)"
|
|
||||||
|
|
||||||
# restore everything to /tmp/restored, then check every restored file's
|
# restore everything to /tmp/restored, then check every restored file's
|
||||||
# chunk hashes
|
# chunk hashes
|
||||||
vaultik snapshot restore --verify <remote-key> /tmp/restored
|
VAULTIK_AGE_SECRET_KEY='AGE-SECRET-KEY-...' \
|
||||||
|
vaultik snapshot restore --verify <remote-key> /tmp/restored
|
||||||
|
|
||||||
# optionally, deep-verify the snapshot against the store (downloads and
|
# optionally, deep-verify the snapshot against the store (downloads and
|
||||||
# cryptographically checks every blob)
|
# cryptographically checks every blob)
|
||||||
vaultik snapshot verify --deep <remote-key>
|
VAULTIK_AGE_SECRET_KEY='AGE-SECRET-KEY-...' \
|
||||||
|
vaultik snapshot verify --deep <remote-key>
|
||||||
```
|
```
|
||||||
|
|
||||||
`age_recipients` (the public key) is not needed to restore — only the
|
`age_recipients` (the public key) is not needed to restore — only the
|
||||||
@@ -154,10 +147,10 @@ vaultik [--config <path>] config edit
|
|||||||
vaultik [--config <path>] config get <key>
|
vaultik [--config <path>] config get <key>
|
||||||
vaultik [--config <path>] config set <key> <value>
|
vaultik [--config <path>] config set <key> <value>
|
||||||
vaultik [--config <path>] snapshot create [snapshot-names...] [--cron] [--prune] [--keep-newer-than <duration>]
|
vaultik [--config <path>] snapshot create [snapshot-names...] [--cron] [--prune] [--keep-newer-than <duration>]
|
||||||
vaultik [--config <path>] snapshot list [--json] # alias: ls
|
vaultik [--config <path>] snapshot list [--json]
|
||||||
vaultik [--config <path>] snapshot verify <snapshot-id> [--deep] [--json]
|
vaultik [--config <path>] snapshot verify <snapshot-id> [--deep] [--json]
|
||||||
vaultik [--config <path>] snapshot purge [--keep-latest | --older-than <duration>] [--snapshot <name>...] [--force]
|
vaultik [--config <path>] snapshot purge [--keep-latest | --older-than <duration>] [--snapshot <name>...] [--force]
|
||||||
vaultik [--config <path>] snapshot remove <snapshot-id> [--dry-run] [--force] [--local-only] [--json] # alias: rm
|
vaultik [--config <path>] snapshot remove <snapshot-id> [--dry-run] [--force] [--local-only] [--json]
|
||||||
vaultik [--config <path>] snapshot restore <snapshot-id> <target-dir> [paths...] [--verify]
|
vaultik [--config <path>] snapshot restore <snapshot-id> <target-dir> [paths...] [--verify]
|
||||||
vaultik [--config <path>] prune [--force] [--json]
|
vaultik [--config <path>] prune [--force] [--json]
|
||||||
vaultik [--config <path>] info
|
vaultik [--config <path>] info
|
||||||
@@ -174,24 +167,7 @@ vaultik version
|
|||||||
* `--verbose`, `-v`: Enable verbose output (on stderr — see below)
|
* `--verbose`, `-v`: Enable verbose output (on stderr — see below)
|
||||||
* `--debug`: Enable debug output (on stderr — see below)
|
* `--debug`: Enable debug output (on stderr — see below)
|
||||||
* `--quiet`, `-q`: Suppress non-error output (also suppresses startup banner)
|
* `--quiet`, `-q`: Suppress non-error output (also suppresses startup banner)
|
||||||
* `--skip-errors`: Skip files that cannot be read when creating a snapshot, or that cannot be restored when restoring, instead of aborting. Packing and storage errors (which would leave a chunk recorded but not stored) still abort the run.
|
* `--skip-errors`: Continue past per-file errors instead of aborting (applies to `snapshot create` and `restore`)
|
||||||
|
|
||||||
### locking
|
|
||||||
|
|
||||||
Commands that write persistent state — `snapshot create`, `snapshot
|
|
||||||
purge`, `snapshot remove`, `prune`, and `remote nuke` — take a
|
|
||||||
process-wide lock at `$XDG_DATA_HOME/vaultik/vaultik.pid`
|
|
||||||
(`~/.local/share/vaultik/vaultik.pid` on Linux) for the whole run. Only
|
|
||||||
one of them runs at a time: a second one exits immediately with an
|
|
||||||
"already running" error rather than waiting, so two writers can never
|
|
||||||
corrupt the local index or the destination store.
|
|
||||||
|
|
||||||
Read-only commands — `info`, `snapshot list`, `snapshot verify`, and
|
|
||||||
`remote info` — do not take the lock and are never blocked, so they run
|
|
||||||
even while a backup is in progress. `snapshot restore` does not take the
|
|
||||||
lock either: it writes only to the target directory you name, not the
|
|
||||||
local index or the destination store. `config`, `database delete`,
|
|
||||||
`completion`, and `version` do not take the lock.
|
|
||||||
|
|
||||||
### stdout and stderr
|
### stdout and stderr
|
||||||
|
|
||||||
@@ -224,11 +200,9 @@ and `vaultik prune --json | jq .` both work as written.
|
|||||||
|
|
||||||
### environment variables
|
### environment variables
|
||||||
|
|
||||||
* `VAULTIK_AGE_SECRET_KEY`: Age private key for decryption (required for `snapshot restore` and `snapshot verify --deep`). May hold the whole `age-keygen` file — comments and every identity in it are accepted. Set it from the file, e.g. `export VAULTIK_AGE_SECRET_KEY="$(cat vaultik_backup_private_key.txt)"`, so the key is not typed into your shell history.
|
* `VAULTIK_AGE_SECRET_KEY`: Age private key for decryption (required for `snapshot restore` and `snapshot verify --deep`)
|
||||||
* `VAULTIK_CONFIG`: Path to config file (overridden by `--config`)
|
* `VAULTIK_CONFIG`: Path to config file (overridden by `--config`)
|
||||||
* `VAULTIK_INDEX_PATH`: Override local SQLite index path
|
* `VAULTIK_INDEX_PATH`: Override local SQLite index path
|
||||||
* `VAULTIK_CPUPROFILE`: Write a CPU profile to this path for the duration of the run (development/debugging)
|
|
||||||
* `VAULTIK_MEMPROFILE`: Write a heap profile to this path when the run exits (development/debugging)
|
|
||||||
|
|
||||||
### shell completion
|
### shell completion
|
||||||
|
|
||||||
@@ -319,9 +293,7 @@ local index alone, and still exits zero.
|
|||||||
logger, so stdout stays a single parseable document.
|
logger, so stdout stays a single parseable document.
|
||||||
|
|
||||||
**`snapshot verify`**: Verify snapshot integrity.
|
**`snapshot verify`**: Verify snapshot integrity.
|
||||||
* Default (shallow): checks that every blob the manifest lists is present in
|
* Default (shallow): checks that all blobs referenced in the manifest exist in storage
|
||||||
storage with the size the manifest records, and that the encrypted database is
|
|
||||||
present. It does not read blob contents.
|
|
||||||
* `--deep`: Downloads and decrypts each blob, verifies chunk hashes against the
|
* `--deep`: Downloads and decrypts each blob, verifies chunk hashes against the
|
||||||
encrypted metadata database
|
encrypted metadata database
|
||||||
* Accepts the same identifiers as `snapshot restore`: a snapshot ID, or a
|
* Accepts the same identifiers as `snapshot restore`: a snapshot ID, or a
|
||||||
@@ -423,10 +395,6 @@ both are set.
|
|||||||
|
|
||||||
## architecture
|
## architecture
|
||||||
|
|
||||||
For an implementation-level view of the internals — the data model, the
|
|
||||||
`fx` dependency-injection wiring, and the scanner — see
|
|
||||||
[`ARCHITECTURE.md`](ARCHITECTURE.md).
|
|
||||||
|
|
||||||
### remote storage layout
|
### remote storage layout
|
||||||
|
|
||||||
```
|
```
|
||||||
@@ -504,30 +472,25 @@ derivation.
|
|||||||
|
|
||||||
### compression
|
### compression
|
||||||
|
|
||||||
* zstd compression at configurable level (1-19, default 3). The level is
|
* zstd compression at configurable level (1-19, default 3)
|
||||||
accepted as 1-19 but maps onto zstd's four internal speed presets:
|
|
||||||
1-2 fastest, 3-5 default, 6-9 better, 10-19 best. Levels within the
|
|
||||||
same band compress identically.
|
|
||||||
* Applied before encryption at the blob level
|
* Applied before encryption at the blob level
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## configuration reference
|
## configuration reference
|
||||||
|
|
||||||
Run `vaultik config init` to generate a fully commented config file; a
|
Run `vaultik config init` to generate a fully commented config file.
|
||||||
complete annotated example also lives in
|
Key fields:
|
||||||
[`config.example.yml`](config.example.yml). Key fields:
|
|
||||||
|
|
||||||
| Field | Default | Description |
|
| Field | Default | Description |
|
||||||
|-------|---------|-------------|
|
|-------|---------|-------------|
|
||||||
| `age_recipients` | (required) | Age public keys for encryption |
|
| `age_recipients` | (required) | Age public keys for encryption |
|
||||||
| `age_secret_key` | (unset) | Age private key for decryption (`snapshot restore`, `snapshot verify --deep`). Setting it in the config file places the private key on the backed-up host, defeating the public-key-only design (see "why" above). Prefer the `VAULTIK_AGE_SECRET_KEY` environment variable, supplied only on the machine you restore from. |
|
|
||||||
| `snapshots` | (required) | Named snapshot definitions with paths and excludes |
|
| `snapshots` | (required) | Named snapshot definitions with paths and excludes |
|
||||||
| `storage_url` | | Storage backend URL (`s3://`, `file://`, `rclone://`) |
|
| `storage_url` | | Storage backend URL (`s3://`, `file://`, `rclone://`) |
|
||||||
| `s3.*` | | Legacy S3 configuration (endpoint, bucket, credentials) |
|
| `s3.*` | | Legacy S3 configuration (endpoint, bucket, credentials) |
|
||||||
| `exclude` | | Global exclude patterns (applied to all snapshots) |
|
| `exclude` | | Global exclude patterns (applied to all snapshots) |
|
||||||
| `chunk_size` | `10MB` | Average chunk size for content-defined chunking |
|
| `chunk_size` | `10MB` | Average chunk size for content-defined chunking |
|
||||||
| `blob_size_limit` | `10GB` | Maximum blob size before splitting. Must be at least four times `chunk_size` (the largest chunk the chunker can emit), otherwise a single-chunk blob could exceed the limit |
|
| `blob_size_limit` | `10GB` | Maximum blob size before splitting |
|
||||||
| `compression_level` | `3` | zstd compression level (1-19) |
|
| `compression_level` | `3` | zstd compression level (1-19) |
|
||||||
| `hostname` | system hostname | Hostname used in snapshot IDs |
|
| `hostname` | system hostname | Hostname used in snapshot IDs |
|
||||||
| `index_path` | platform data dir | Local SQLite index path |
|
| `index_path` | platform data dir | Local SQLite index path |
|
||||||
@@ -637,17 +600,9 @@ priority.
|
|||||||
|
|
||||||
## output style
|
## output style
|
||||||
|
|
||||||
The operational narration of the long-running commands — the Begin,
|
All user-facing output goes through helpers in `internal/ui` and conforms
|
||||||
Complete, Progress, and status lines of `snapshot create`, `prune`,
|
to a uniform style. Color is enabled when stdout is a TTY and the
|
||||||
`snapshot restore`, and the like — goes through helpers in `internal/ui`
|
`NO_COLOR` environment variable is unset (https://no-color.org/).
|
||||||
and conforms to the uniform style below. Some commands instead write
|
|
||||||
plain text straight to stdout (`version`, `info`, `config`, the
|
|
||||||
`database delete` prompt, and the `snapshot list` table); that output is
|
|
||||||
unstyled and does not honor `--quiet`. Routing it through `internal/ui`
|
|
||||||
is tracked in
|
|
||||||
[issue #149](https://git.eeqj.de/sneak/vaultik/issues/149). Color is
|
|
||||||
enabled when stdout is a TTY and the `NO_COLOR` environment variable is
|
|
||||||
unset (https://no-color.org/).
|
|
||||||
|
|
||||||
`internal/ui` writes to stdout; it is the output the user asked for.
|
`internal/ui` writes to stdout; it is the output the user asked for.
|
||||||
Structured log records are a different thing and go through
|
Structured log records are a different thing and go through
|
||||||
|
|||||||
@@ -25,43 +25,6 @@ release" is exactly the contradiction
|
|||||||
|
|
||||||
# Completed Steps
|
# Completed Steps
|
||||||
|
|
||||||
- 2026-09-21: Stopped an interrupted blob upload from making a later
|
|
||||||
backup deduplicate against data that was never stored
|
|
||||||
([issue #148](https://git.eeqj.de/sneak/vaultik/issues/148)). The
|
|
||||||
packer commits a blob's `chunks`, `blob_chunks`, and `blobs` rows
|
|
||||||
before the upload is attempted, so a failed upload left chunk rows
|
|
||||||
behind and the next run skipped re-uploading them, producing a
|
|
||||||
snapshot that reported success but could not be restored. A run now
|
|
||||||
deduplicates only against chunks held by a blob whose `uploaded_ts` is
|
|
||||||
set, and at startup drops any un-uploaded blob rows (and the chunks
|
|
||||||
they orphan) so the affected data is re-chunked and re-uploaded. Blobs
|
|
||||||
recorded with no remote backend are marked uploaded so this invariant
|
|
||||||
holds uniformly.
|
|
||||||
|
|
||||||
- 2026-09-22: Made restore refuse any snapshot path that would write
|
|
||||||
outside the target directory
|
|
||||||
([issue #154](https://git.eeqj.de/sneak/vaultik/issues/154)).
|
|
||||||
`restoreFile` and `verifyRestoredFiles` joined the stored path onto the
|
|
||||||
target with no containment check, so a `..` segment or an absolute path
|
|
||||||
escaped the target and a restored symlink could redirect a later child
|
|
||||||
write anywhere on disk. Every stored path is now rejected unless
|
|
||||||
`filepath.IsLocal` accepts it with the leading separator removed, and
|
|
||||||
each existing ancestor directory below the target is `Lstat`ed to refuse
|
|
||||||
descending through a symlink; honest symlinks pointing outside the tree
|
|
||||||
are still written verbatim. age decryption proves a snapshot is
|
|
||||||
readable, not honest, and restore usually runs as root.
|
|
||||||
|
|
||||||
- 2026-09-21: Stopped `--json` from silencing stderr diagnostics
|
|
||||||
([issue #112](https://git.eeqj.de/sneak/vaultik/issues/112)). `--json`
|
|
||||||
used to be folded into `Quiet`, which pinned the log level to `WARN`,
|
|
||||||
so `prune --json` gave a machine consumer no record of the local index
|
|
||||||
rows it deleted even under `--verbose`. `--json` now quiets only the
|
|
||||||
stdout UI (the JSON document must stay clean, per
|
|
||||||
[issue #108](https://git.eeqj.de/sneak/vaultik/issues/108)); the stderr
|
|
||||||
log level follows `--verbose`/`--debug` again. The coupling was
|
|
||||||
removed the same way for `snapshot verify`, `snapshot remove`, and
|
|
||||||
`remote info`, which carried it for the same outdated reason.
|
|
||||||
|
|
||||||
- 2026-09-21: Stopped `prune` from reporting a failed row count as 0
|
- 2026-09-21: Stopped `prune` from reporting a failed row count as 0
|
||||||
([issue #96](https://git.eeqj.de/sneak/vaultik/issues/96)). The seven
|
([issue #96](https://git.eeqj.de/sneak/vaultik/issues/96)). The seven
|
||||||
`getTableCount` reads in `PruneDatabase` discarded their error, so a
|
`getTableCount` reads in `PruneDatabase` discarded their error, so a
|
||||||
|
|||||||
@@ -1,102 +0,0 @@
|
|||||||
package main_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// This file guards the version stamping of the product image (issue
|
|
||||||
// #75). The failure it protects against is silent: the image still
|
|
||||||
// builds and runs, but `vaultik version` inside it reports "commit:
|
|
||||||
// unknown", so an operator cannot tell which source produced a given
|
|
||||||
// backup. .dockerignore excludes .git, so the build cannot derive the
|
|
||||||
// commit itself; the values must be computed on the host and passed in.
|
|
||||||
//
|
|
||||||
// These are parses of the committed files, for the same reason the lint
|
|
||||||
// guards next door are: shelling out to docker would nest a build
|
|
||||||
// inside `make test`. That `vaultik version` in the built image really
|
|
||||||
// prints the host's version is verified by hand and recorded on the
|
|
||||||
// pull request.
|
|
||||||
|
|
||||||
// dockerScript is script/docker, relative to the repository root.
|
|
||||||
const dockerScript = "script/docker"
|
|
||||||
|
|
||||||
// versionArgs are the ldflag targets the build stamps and, matching
|
|
||||||
// them, the build args the host must supply. The names line up so the
|
|
||||||
// same list checks both files.
|
|
||||||
func versionArgs() []string {
|
|
||||||
return []string{"VERSION", "COMMIT", "COMMIT_DATE"}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProductDockerfileTakesVersionAsBuildArgs fails unless the build
|
|
||||||
// declares each version arg and stamps it into the binary by ldflag
|
|
||||||
// reference, rather than computing it in the container.
|
|
||||||
func TestProductDockerfileTakesVersionAsBuildArgs(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
found := instructions(t, productDockerfile)
|
|
||||||
|
|
||||||
for _, arg := range versionArgs() {
|
|
||||||
require.GreaterOrEqual(t, indexOf(found, "ARG "+arg), 0,
|
|
||||||
"%s must declare `ARG %s` so the host can pass it in",
|
|
||||||
productDockerfile, arg)
|
|
||||||
|
|
||||||
assertLdflagReferences(t, found, arg)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProductDockerfileDoesNotDeriveVersionItself is the anti-regression
|
|
||||||
// for the original defect: the container ran `git rev-parse`, but .git
|
|
||||||
// is not in the build context, so it always resolved to "unknown". No
|
|
||||||
// git command may reach into a build that cannot see the history.
|
|
||||||
func TestProductDockerfileDoesNotDeriveVersionItself(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
text := instructionText(readRepoFile(t, productDockerfile))
|
|
||||||
|
|
||||||
assert.NotContains(t, text, "git ",
|
|
||||||
"%s must not run git: .git is excluded from the build context, so"+
|
|
||||||
" any value it derives is wrong. Pass version, commit and date"+
|
|
||||||
" in as build args instead.", productDockerfile)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestDockerScriptComputesVersionOnTheHost fails unless script/docker
|
|
||||||
// derives each value where .git exists and passes it as a build arg,
|
|
||||||
// with VERSION coming from script/version so a Docker build reports the
|
|
||||||
// same string a local build of the same tree would.
|
|
||||||
func TestDockerScriptComputesVersionOnTheHost(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
script := readRepoFile(t, dockerScript)
|
|
||||||
|
|
||||||
for _, arg := range versionArgs() {
|
|
||||||
assert.Contains(t, script, "--build-arg "+arg+"=",
|
|
||||||
"%s must pass --build-arg %s to the build", dockerScript, arg)
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.Contains(t, script, "/version",
|
|
||||||
"%s must take VERSION from script/version, the source of truth"+
|
|
||||||
" shared with the Makefile", dockerScript)
|
|
||||||
}
|
|
||||||
|
|
||||||
// assertLdflagReferences fails unless some build instruction stamps the
|
|
||||||
// named variable from the ARG (a ${arg} reference), not from a value
|
|
||||||
// computed inside the container.
|
|
||||||
func assertLdflagReferences(t *testing.T, found []string, arg string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
for _, instruction := range found {
|
|
||||||
if strings.HasPrefix(instruction, "RUN ") &&
|
|
||||||
strings.Contains(instruction, "go build") &&
|
|
||||||
strings.Contains(instruction, "${"+arg+"}") {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.Fail(t, "version arg is declared but never stamped",
|
|
||||||
"the go build in %s must reference ${%s} in its ldflags, or the"+
|
|
||||||
" arg is passed and discarded", productDockerfile, arg)
|
|
||||||
}
|
|
||||||
@@ -304,14 +304,10 @@ func instructionText(contents string) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// indexOf returns the position of the first instruction equal to, or
|
// indexOf returns the position of the first instruction equal to, or
|
||||||
// beginning with, want; -1 if there is none. An `ARG NAME=default`
|
// beginning with, want; -1 if there is none.
|
||||||
// counts as beginning with `ARG NAME`, so a declared arg is found
|
|
||||||
// whether or not it carries a default.
|
|
||||||
func indexOf(found []string, want string) int {
|
func indexOf(found []string, want string) int {
|
||||||
for i, instruction := range found {
|
for i, instruction := range found {
|
||||||
if instruction == want ||
|
if instruction == want || strings.HasPrefix(instruction, want+" ") {
|
||||||
strings.HasPrefix(instruction, want+" ") ||
|
|
||||||
strings.HasPrefix(instruction, want+"=") {
|
|
||||||
return i
|
return i
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-11
@@ -10,16 +10,6 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func main() {
|
func main() {
|
||||||
os.Exit(run())
|
|
||||||
}
|
|
||||||
|
|
||||||
// run sets up optional profiling, runs the CLI, and returns the process
|
|
||||||
// exit code. os.Exit lives in main so it fires only after run's deferred
|
|
||||||
// profile writers have flushed. cli.Entry returns a status code rather
|
|
||||||
// than calling os.Exit itself: an os.Exit from inside it would skip
|
|
||||||
// these defers and truncate the profile of a failing command -- exactly
|
|
||||||
// the command one most often wants to profile.
|
|
||||||
func run() int {
|
|
||||||
// CPU profiling: set VAULTIK_CPUPROFILE=/path/to/cpu.prof
|
// CPU profiling: set VAULTIK_CPUPROFILE=/path/to/cpu.prof
|
||||||
if cpuProfile := os.Getenv("VAULTIK_CPUPROFILE"); cpuProfile != "" {
|
if cpuProfile := os.Getenv("VAULTIK_CPUPROFILE"); cpuProfile != "" {
|
||||||
f, err := os.Create(cpuProfile) //nolint:gosec // G304: operator-set path
|
f, err := os.Create(cpuProfile) //nolint:gosec // G304: operator-set path
|
||||||
@@ -56,5 +46,5 @@ func run() int {
|
|||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
return cli.Entry()
|
cli.Entry()
|
||||||
}
|
}
|
||||||
|
|||||||
+5
-7
@@ -257,16 +257,16 @@ exclude:
|
|||||||
|
|
||||||
# Storage URL - use either this OR the s3 section below
|
# Storage URL - use either this OR the s3 section below
|
||||||
# Supports: s3://bucket/prefix, file:///path, rclone://remote/path
|
# Supports: s3://bucket/prefix, file:///path, rclone://remote/path
|
||||||
storage_url: "rclone://myremote/path/to/backups"
|
storage_url: "rclone://las1stor1//srv/pool.2024.04/backups/heraklion"
|
||||||
|
|
||||||
# S3-compatible storage configuration
|
# S3-compatible storage configuration
|
||||||
#s3:
|
#s3:
|
||||||
# # S3-compatible endpoint URL
|
# # S3-compatible endpoint URL
|
||||||
# # Examples: https://s3.amazonaws.com, https://storage.googleapis.com
|
# # Examples: https://s3.amazonaws.com, https://storage.googleapis.com
|
||||||
# endpoint: https://s3.example.com
|
# endpoint: http://10.100.205.122:8333
|
||||||
#
|
#
|
||||||
# # Bucket name where backups will be stored
|
# # Bucket name where backups will be stored
|
||||||
# bucket: mybucket
|
# bucket: testbucket
|
||||||
#
|
#
|
||||||
# # Prefix (folder) within the bucket for this host's backups
|
# # Prefix (folder) within the bucket for this host's backups
|
||||||
# # Useful for organizing backups from multiple hosts
|
# # Useful for organizing backups from multiple hosts
|
||||||
@@ -274,8 +274,8 @@ storage_url: "rclone://myremote/path/to/backups"
|
|||||||
# #prefix: "hosts/myserver/"
|
# #prefix: "hosts/myserver/"
|
||||||
#
|
#
|
||||||
# # S3 access credentials
|
# # S3 access credentials
|
||||||
# access_key_id: YOUR_ACCESS_KEY
|
# access_key_id: Z9GT22M9YFU08WRMC5D4
|
||||||
# secret_access_key: YOUR_SECRET_KEY
|
# secret_access_key: Pi0tPKjFbN4rZlRhcA4zBtEkib04yy2WcIzI+AXk
|
||||||
#
|
#
|
||||||
# # S3 region
|
# # S3 region
|
||||||
# # Default: us-east-1
|
# # Default: us-east-1
|
||||||
@@ -304,8 +304,6 @@ storage_url: "rclone://myremote/path/to/backups"
|
|||||||
|
|
||||||
# Maximum blob size
|
# Maximum blob size
|
||||||
# Multiple chunks are packed into blobs up to this size
|
# Multiple chunks are packed into blobs up to this size
|
||||||
# Must be at least four times chunk_size (the largest chunk the chunker can
|
|
||||||
# emit); a smaller limit would let a single-chunk blob exceed it.
|
|
||||||
# Supports: 1GB, 10G, 500MB, 1GiB, etc.
|
# Supports: 1GB, 10G, 500MB, 1GiB, etc.
|
||||||
# Default: 10GB
|
# Default: 10GB
|
||||||
#blob_size_limit: 10GB
|
#blob_size_limit: 10GB
|
||||||
|
|||||||
@@ -145,10 +145,10 @@ An observer cannot determine:
|
|||||||
## Pruning Safety
|
## Pruning Safety
|
||||||
|
|
||||||
The prune operation is safe because:
|
The prune operation is safe because:
|
||||||
1. It keeps every blob listed in any snapshot's manifest and deletes only blobs that no manifest references
|
1. It only deletes blobs not referenced in any manifest
|
||||||
2. Manifests are unencrypted and can be read without keys
|
2. Manifests are unencrypted and can be read without keys
|
||||||
3. If any manifest cannot be downloaded or decoded, prune deletes nothing and exits with an error, rather than treating that snapshot's blobs as unreferenced
|
3. The operation compares the latest local DB snapshot with the latest S3 snapshot to ensure consistency
|
||||||
4. Prune requires exclusive access to the destination: running it during a concurrent backup can race a snapshot whose manifest is not yet written, so do not prune while a backup is in progress
|
4. Pruning will fail if these don't match, preventing accidental deletion of needed blobs
|
||||||
|
|
||||||
## Restoration Requirements
|
## Restoration Requirements
|
||||||
|
|
||||||
|
|||||||
@@ -487,7 +487,7 @@ func (p *Packer) closeBlobWriter() (string, int64, error) {
|
|||||||
return "", 0, fmt.Errorf("seeking to start: %w", err)
|
return "", 0, fmt.Errorf("seeking to start: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
finalHash := p.currentBlob.writer.ContentID()
|
finalHash := p.currentBlob.writer.Sum256()
|
||||||
|
|
||||||
return hex.EncodeToString(finalHash), finalSize, nil
|
return hex.EncodeToString(finalHash), finalSize, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,89 @@
|
|||||||
|
// Package blobgen implements the blob data pipeline: streaming zstd
|
||||||
|
// compression, age encryption, and SHA256 content hashing for blob
|
||||||
|
// creation, plus the matching decrypt/decompress/verify reader.
|
||||||
|
package blobgen
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/hex"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CompressResult contains the results of compression
|
||||||
|
type CompressResult struct {
|
||||||
|
Data []byte
|
||||||
|
UncompressedSize int64
|
||||||
|
CompressedSize int64
|
||||||
|
SHA256 string
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompressData compresses and encrypts data, returning the result with hash
|
||||||
|
func CompressData(
|
||||||
|
data []byte, compressionLevel int, recipients []string,
|
||||||
|
) (*CompressResult, error) {
|
||||||
|
var buf bytes.Buffer
|
||||||
|
|
||||||
|
// Create writer
|
||||||
|
w, err := NewWriter(&buf, compressionLevel, recipients)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("creating writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write data
|
||||||
|
_, err = w.Write(data)
|
||||||
|
if err != nil {
|
||||||
|
_ = w.Close()
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("writing data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close to flush
|
||||||
|
err = w.Close()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("closing writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &CompressResult{
|
||||||
|
Data: buf.Bytes(),
|
||||||
|
UncompressedSize: int64(len(data)),
|
||||||
|
CompressedSize: int64(buf.Len()),
|
||||||
|
SHA256: hex.EncodeToString(w.Sum256()),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompressStream compresses and encrypts from reader to writer, returning
|
||||||
|
// the number of uncompressed bytes written and the content hash.
|
||||||
|
func CompressStream(
|
||||||
|
dst io.Writer, src io.Reader, compressionLevel int, recipients []string,
|
||||||
|
) (int64, string, error) {
|
||||||
|
// Create writer
|
||||||
|
w, err := NewWriter(dst, compressionLevel, recipients)
|
||||||
|
if err != nil {
|
||||||
|
return 0, "", fmt.Errorf("creating writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
closed := false
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
if !closed {
|
||||||
|
_ = w.Close()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Copy data
|
||||||
|
_, err = io.Copy(w, src)
|
||||||
|
if err != nil {
|
||||||
|
return 0, "", fmt.Errorf("copying data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close to flush
|
||||||
|
err = w.Close()
|
||||||
|
if err != nil {
|
||||||
|
return 0, "", fmt.Errorf("closing writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
closed = true
|
||||||
|
|
||||||
|
return w.BytesWritten(), hex.EncodeToString(w.Sum256()), nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,80 @@
|
|||||||
|
package blobgen_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"crypto/rand"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testRecipient is a static age recipient for tests.
|
||||||
|
const testRecipient = "age1cplgrwj77ta54dnmydvvmzn64ltk83ankxl5sww04mrtmu62kv3s89gmvv"
|
||||||
|
|
||||||
|
// TestCompressStreamNoDoubleClose is a regression test for issue #28.
|
||||||
|
// It verifies that CompressStream does not panic or return an error due to
|
||||||
|
// double-closing the underlying blobgen.Writer. Before the fix in PR #33,
|
||||||
|
// the explicit Close() on the happy path combined with defer Close() would
|
||||||
|
// cause a double close.
|
||||||
|
func TestCompressStreamNoDoubleClose(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
input := []byte("regression test data for issue #28 double-close fix")
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
|
||||||
|
written, hash, err := blobgen.CompressStream(
|
||||||
|
&buf, bytes.NewReader(input), 3, []string{testRecipient})
|
||||||
|
require.NoError(t, err, "CompressStream should not return an error")
|
||||||
|
assert.Positive(t, written, "expected bytes written > 0")
|
||||||
|
assert.NotEmpty(t, hash, "expected non-empty hash")
|
||||||
|
assert.Positive(t, buf.Len(), "expected non-empty output")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCompressStreamLargeInput exercises CompressStream with a larger payload
|
||||||
|
// to ensure no double-close issues surface under heavier I/O.
|
||||||
|
func TestCompressStreamLargeInput(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
data := make([]byte, 512*1024) // 512 KB
|
||||||
|
_, err := rand.Read(data)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
|
||||||
|
written, hash, err := blobgen.CompressStream(
|
||||||
|
&buf, bytes.NewReader(data), 3, []string{testRecipient})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Positive(t, written)
|
||||||
|
assert.NotEmpty(t, hash)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCompressStreamEmptyInput verifies CompressStream handles empty input
|
||||||
|
// without double-close issues.
|
||||||
|
func TestCompressStreamEmptyInput(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
|
||||||
|
_, hash, err := blobgen.CompressStream(
|
||||||
|
&buf, strings.NewReader(""), 3, []string{testRecipient})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NotEmpty(t, hash)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCompressDataNoDoubleClose mirrors the stream test for CompressData,
|
||||||
|
// ensuring the explicit Close + error-path Close pattern is also safe.
|
||||||
|
func TestCompressDataNoDoubleClose(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
input := []byte("CompressData regression test for double-close")
|
||||||
|
|
||||||
|
result, err := blobgen.CompressData(input, 3, []string{testRecipient})
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Positive(t, result.CompressedSize)
|
||||||
|
assert.Equal(t, result.UncompressedSize, int64(len(input)))
|
||||||
|
assert.NotEmpty(t, result.SHA256)
|
||||||
|
}
|
||||||
@@ -1,119 +0,0 @@
|
|||||||
package blobgen_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"crypto/rand"
|
|
||||||
"errors"
|
|
||||||
"io"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"filippo.io/age"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ageChunkSize is age's STREAM plaintext chunk size (64 KiB); each encrypted
|
|
||||||
// chunk adds a 16-byte ChaCha20-Poly1305 tag.
|
|
||||||
const (
|
|
||||||
ageChunkSize = 64 * 1024
|
|
||||||
ageChunkTagSize = 16
|
|
||||||
ageSegmentSize = ageChunkSize + ageChunkTagSize
|
|
||||||
ageNonceSize = 16
|
|
||||||
)
|
|
||||||
|
|
||||||
// makeIdentity returns a fresh X25519 identity and its recipient string.
|
|
||||||
func makeIdentity(t *testing.T) (*age.X25519Identity, string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
id, err := age.GenerateX25519Identity()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
return id, id.Recipient().String()
|
|
||||||
}
|
|
||||||
|
|
||||||
// randomBytes returns n cryptographically random bytes, which do not compress
|
|
||||||
// so the encrypted payload spans multiple age segments.
|
|
||||||
func randomBytes(t *testing.T, n int) []byte {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
b := make([]byte, n)
|
|
||||||
_, err := rand.Read(b)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
// compressibleBytes returns n bytes of a repeating pattern, which zstd packs
|
|
||||||
// down to a small payload.
|
|
||||||
func compressibleBytes(n int) []byte {
|
|
||||||
pattern := bytes.Repeat([]byte("compressible-"), n/13+1)
|
|
||||||
|
|
||||||
return pattern[:n]
|
|
||||||
}
|
|
||||||
|
|
||||||
// encryptBlob compresses, encrypts and returns a blob for plaintext at
|
|
||||||
// compression level 1.
|
|
||||||
func encryptBlob(t *testing.T, plaintext []byte, recipients ...string) []byte {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var buf bytes.Buffer
|
|
||||||
|
|
||||||
w, err := blobgen.NewWriter(&buf, 1, recipients)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = w.Write(plaintext)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, w.Close())
|
|
||||||
|
|
||||||
return buf.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
// ageHeaderLen returns the byte length of blob's age header, i.e. the offset
|
|
||||||
// of the 16-byte payload nonce that follows it. The header ends with a MAC
|
|
||||||
// line "--- <mac>\n"; the nonce begins right after that newline.
|
|
||||||
func ageHeaderLen(t *testing.T, blob []byte) int {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
i := bytes.Index(blob, []byte("\n--- "))
|
|
||||||
require.GreaterOrEqual(t, i, 0, "age MAC footer line not found")
|
|
||||||
|
|
||||||
nl := bytes.IndexByte(blob[i+1:], '\n')
|
|
||||||
require.GreaterOrEqual(t, nl, 0, "newline ending MAC line not found")
|
|
||||||
|
|
||||||
return i + 1 + nl + 1
|
|
||||||
}
|
|
||||||
|
|
||||||
// requireBlobUnreadable asserts that data never decrypts to a plaintext with a
|
|
||||||
// nil error: either NewReader fails, or reading it does.
|
|
||||||
func requireBlobUnreadable(t *testing.T, data []byte, id age.Identity) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
r, err := blobgen.NewReader(bytes.NewReader(data), id)
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err = io.ReadAll(r)
|
|
||||||
_ = r.Close()
|
|
||||||
|
|
||||||
require.Error(t, err, "reading a damaged blob must fail")
|
|
||||||
}
|
|
||||||
|
|
||||||
// errFailWriter is returned by failAfterWriter once its byte limit is passed.
|
|
||||||
var errFailWriter = errors.New("destination write failed")
|
|
||||||
|
|
||||||
// failAfterWriter accepts writes until more than limit bytes have been sent,
|
|
||||||
// then fails every write. It models a destination that dies mid-blob.
|
|
||||||
type failAfterWriter struct {
|
|
||||||
limit int
|
|
||||||
written int
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *failAfterWriter) Write(p []byte) (int, error) {
|
|
||||||
f.written += len(p)
|
|
||||||
if f.written > f.limit {
|
|
||||||
return 0, errFailWriter
|
|
||||||
}
|
|
||||||
|
|
||||||
return len(p), nil
|
|
||||||
}
|
|
||||||
@@ -1,190 +0,0 @@
|
|||||||
package blobgen_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"fmt"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"filippo.io/age"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestNewReaderWrongIdentity covers issue case 4: opening a blob with an
|
|
||||||
// identity other than the recipient reports no matching identity.
|
|
||||||
func TestNewReaderWrongIdentity(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
_, recipient := makeIdentity(t)
|
|
||||||
other, _ := makeIdentity(t)
|
|
||||||
|
|
||||||
blob := encryptBlob(t, []byte("secret payload"), recipient)
|
|
||||||
|
|
||||||
_, err := blobgen.NewReader(bytes.NewReader(blob), other)
|
|
||||||
require.Error(t, err)
|
|
||||||
|
|
||||||
var noMatch *age.NoIdentityMatchError
|
|
||||||
assert.ErrorAs(t, err, &noMatch)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewReaderTruncated covers issue case 6: a multi-segment blob cut at
|
|
||||||
// several points must never read back as valid data. The point immediately
|
|
||||||
// after the header and nonce is intentionally excluded: it reads as a valid
|
|
||||||
// empty blob today and is the regression case for
|
|
||||||
// https://git.eeqj.de/sneak/vaultik/issues/152.
|
|
||||||
func TestNewReaderTruncated(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
id, recipient := makeIdentity(t)
|
|
||||||
blob := encryptBlob(t, randomBytes(t, 4*65536+123), recipient)
|
|
||||||
h := ageHeaderLen(t, blob)
|
|
||||||
|
|
||||||
require.Greater(t, len(blob), h+ageNonceSize+ageSegmentSize,
|
|
||||||
"test needs a blob of at least two age segments")
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
size int
|
|
||||||
}{
|
|
||||||
{"inside header", h / 2},
|
|
||||||
{"inside nonce", h + 8},
|
|
||||||
{"inside first segment", h + ageNonceSize + 100},
|
|
||||||
{"end of first full segment", h + ageNonceSize + ageSegmentSize},
|
|
||||||
{"last byte removed", len(blob) - 1},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range cases {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
requireBlobUnreadable(t, blob[:tc.size], id)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewReaderCorrupted covers issue case 7: one flipped byte in each region
|
|
||||||
// of a multi-segment blob makes it unreadable.
|
|
||||||
func TestNewReaderCorrupted(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
id, recipient := makeIdentity(t)
|
|
||||||
blob := encryptBlob(t, randomBytes(t, 4*65536+123), recipient)
|
|
||||||
h := ageHeaderLen(t, blob)
|
|
||||||
|
|
||||||
firstNL := bytes.IndexByte(blob, '\n')
|
|
||||||
require.Positive(t, firstNL, "header must have a version line")
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
pos int
|
|
||||||
}{
|
|
||||||
{"header stanza", firstNL + 5},
|
|
||||||
{"header MAC line", h - 2},
|
|
||||||
{"nonce", h + 4},
|
|
||||||
{"body segment", h + ageNonceSize + 50},
|
|
||||||
{"final tag", len(blob) - 1},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range cases {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
corrupt := append([]byte(nil), blob...)
|
|
||||||
corrupt[tc.pos] ^= 0xff
|
|
||||||
requireBlobUnreadable(t, corrupt, id)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewReaderTrailingAndGarbage covers issue case 8: bytes appended after a
|
|
||||||
// valid blob, empty input, and random garbage each fail to read.
|
|
||||||
func TestNewReaderTrailingAndGarbage(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
id, recipient := makeIdentity(t)
|
|
||||||
|
|
||||||
valid := encryptBlob(t, []byte("small payload"), recipient)
|
|
||||||
appended := append(append([]byte(nil), valid...), []byte("trailing junk")...)
|
|
||||||
|
|
||||||
t.Run("appended bytes", func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
requireBlobUnreadable(t, appended, id)
|
|
||||||
})
|
|
||||||
t.Run("empty input", func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
requireBlobUnreadable(t, []byte{}, id)
|
|
||||||
})
|
|
||||||
t.Run("random garbage", func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
requireBlobUnreadable(t, randomBytes(t, 512), id)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewWriterInvalidLevel covers the rejected end of issue case 9: an
|
|
||||||
// out-of-range compression level errors and writes nothing to the destination.
|
|
||||||
func TestNewWriterInvalidLevel(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
_, recipient := makeIdentity(t)
|
|
||||||
|
|
||||||
for _, level := range []int{0, -1, 20} {
|
|
||||||
t.Run(fmt.Sprintf("level%d", level), func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
var buf bytes.Buffer
|
|
||||||
|
|
||||||
w, err := blobgen.NewWriter(&buf, level, []string{recipient})
|
|
||||||
require.ErrorIs(t, err, blobgen.ErrInvalidCompressionLevel)
|
|
||||||
assert.Nil(t, w)
|
|
||||||
assert.Zero(t, buf.Len(), "nothing written on an invalid level")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewWriterInvalidRecipients covers issue case 10: nil and empty recipient
|
|
||||||
// lists and an unparsable recipient string each error.
|
|
||||||
func TestNewWriterInvalidRecipients(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
recipients []string
|
|
||||||
}{
|
|
||||||
{"nil list", nil},
|
|
||||||
{"empty list", []string{}},
|
|
||||||
{"invalid recipient string", []string{"not-a-recipient"}},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range cases {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
var buf bytes.Buffer
|
|
||||||
|
|
||||||
w, err := blobgen.NewWriter(&buf, 1, tc.recipients)
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.Nil(t, w)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewWriterFailingDestination covers issue case 11: a destination that
|
|
||||||
// fails mid-blob surfaces its error from Write or Close.
|
|
||||||
func TestNewWriterFailingDestination(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
_, recipient := makeIdentity(t)
|
|
||||||
|
|
||||||
// The limit clears the age header and nonce so NewWriter succeeds, then
|
|
||||||
// trips once the compressed body starts flowing.
|
|
||||||
dst := &failAfterWriter{limit: 512}
|
|
||||||
|
|
||||||
w, err := blobgen.NewWriter(dst, 1, []string{recipient})
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, writeErr := w.Write(randomBytes(t, 256*1024))
|
|
||||||
closeErr := w.Close()
|
|
||||||
|
|
||||||
assert.True(t, writeErr != nil || closeErr != nil,
|
|
||||||
"destination failure must surface from Write or Close")
|
|
||||||
}
|
|
||||||
@@ -20,12 +20,10 @@ type Reader struct {
|
|||||||
bytesRead int64
|
bytesRead int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewReader creates a new Reader that decrypts, decompresses, and verifies
|
// NewReader creates a new Reader that decrypts, decompresses, and verifies data
|
||||||
// data. Every supplied identity is offered to age.Decrypt, so a blob
|
func NewReader(r io.Reader, identity age.Identity) (*Reader, error) {
|
||||||
// encrypted to any one of them can be read.
|
|
||||||
func NewReader(r io.Reader, identities ...age.Identity) (*Reader, error) {
|
|
||||||
// Create decryption reader
|
// Create decryption reader
|
||||||
decReader, err := age.Decrypt(r, identities...)
|
decReader, err := age.Decrypt(r, identity)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("creating decryption reader: %w", err)
|
return nil, fmt.Errorf("creating decryption reader: %w", err)
|
||||||
}
|
}
|
||||||
@@ -66,9 +64,7 @@ func (r *Reader) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Sum256 returns the single SHA-256 of the plaintext read so far. This is the
|
// Sum256 returns the SHA256 hash of all data read
|
||||||
// first hash only; the stored object name is its double hash, which callers
|
|
||||||
// obtain by passing this digest to DoubleSHA256.
|
|
||||||
func (r *Reader) Sum256() []byte {
|
func (r *Reader) Sum256() []byte {
|
||||||
return r.hasher.Sum(nil)
|
return r.hasher.Sum(nil)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,54 +0,0 @@
|
|||||||
package blobgen_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"io"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"filippo.io/age"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestMultipleRecipients verifies that data written for several recipients can
|
|
||||||
// be read back by each recipient's identity. Moved from internal/crypto, which
|
|
||||||
// held the only multi-recipient test; blobgen is now the sole encryption path.
|
|
||||||
func TestMultipleRecipients(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
identities := make([]*age.X25519Identity, 3)
|
|
||||||
recipients := make([]string, 3)
|
|
||||||
|
|
||||||
for i := range identities {
|
|
||||||
identity, err := age.GenerateX25519Identity()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
identities[i] = identity
|
|
||||||
recipients[i] = identity.Recipient().String()
|
|
||||||
}
|
|
||||||
|
|
||||||
plaintext := []byte("Secret message for multiple recipients")
|
|
||||||
|
|
||||||
var encrypted bytes.Buffer
|
|
||||||
|
|
||||||
writer, err := blobgen.NewWriter(&encrypted, 3, recipients)
|
|
||||||
require.NoError(t, err)
|
|
||||||
_, err = writer.Write(plaintext)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, writer.Close())
|
|
||||||
|
|
||||||
// Every recipient's identity must recover the original plaintext.
|
|
||||||
for i, identity := range identities {
|
|
||||||
reader, err := blobgen.NewReader(
|
|
||||||
bytes.NewReader(encrypted.Bytes()), identity)
|
|
||||||
require.NoError(t, err, "recipient %d should open the reader", i+1)
|
|
||||||
|
|
||||||
got, err := io.ReadAll(reader)
|
|
||||||
require.NoError(t, err, "recipient %d should read the plaintext", i+1)
|
|
||||||
require.NoError(t, reader.Close())
|
|
||||||
|
|
||||||
assert.Equal(t, plaintext, got,
|
|
||||||
"recipient %d should recover the original plaintext", i+1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,132 +0,0 @@
|
|||||||
package blobgen_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"crypto/sha256"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"filippo.io/age"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
|
||||||
)
|
|
||||||
|
|
||||||
// checkRoundTrip writes input through a Writer, reads it back through a Reader,
|
|
||||||
// and verifies the plaintext, the byte counts, and the content hashes.
|
|
||||||
func checkRoundTrip(
|
|
||||||
t *testing.T, id *age.X25519Identity, recipient string,
|
|
||||||
level int, input []byte,
|
|
||||||
) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var buf bytes.Buffer
|
|
||||||
|
|
||||||
w, err := blobgen.NewWriter(&buf, level, []string{recipient})
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
n, err := w.Write(input)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, len(input), n)
|
|
||||||
require.NoError(t, w.Close())
|
|
||||||
require.Equal(t, int64(len(input)), w.BytesWritten())
|
|
||||||
|
|
||||||
r, err := blobgen.NewReader(bytes.NewReader(buf.Bytes()), id)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
got, err := io.ReadAll(r)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, r.Close())
|
|
||||||
|
|
||||||
assert.Equal(t, input, got, "decrypted output must equal input")
|
|
||||||
require.Equal(t, int64(len(input)), r.BytesRead())
|
|
||||||
|
|
||||||
// The hash values are checked by decrypting: the reader's single SHA-256
|
|
||||||
// is the hash of the plaintext, and hashing it once more (DoubleSHA256)
|
|
||||||
// gives the writer's ContentID.
|
|
||||||
single := sha256.Sum256(got)
|
|
||||||
assert.Equal(t, single[:], r.Sum256())
|
|
||||||
assert.Equal(t, blobgen.DoubleSHA256(r.Sum256()), w.ContentID())
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWriterReaderRoundTrip covers issue cases 1 and 2: every size round trips
|
|
||||||
// for both random and compressible data, and the reader hash, its double hash
|
|
||||||
// and the byte counts all agree.
|
|
||||||
func TestWriterReaderRoundTrip(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
id, recipient := makeIdentity(t)
|
|
||||||
|
|
||||||
// Sizes exercise the age segment boundary (64 KiB) from just below to a
|
|
||||||
// few segments above it, plus the empty and single-byte edges.
|
|
||||||
sizes := []int{0, 1, 65535, 65536, 65537, 4*65536 + 123}
|
|
||||||
|
|
||||||
kinds := []struct {
|
|
||||||
name string
|
|
||||||
fill func(*testing.T, int) []byte
|
|
||||||
}{
|
|
||||||
{"random", randomBytes},
|
|
||||||
{"compressible", func(_ *testing.T, n int) []byte {
|
|
||||||
return compressibleBytes(n)
|
|
||||||
}},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, k := range kinds {
|
|
||||||
for _, size := range sizes {
|
|
||||||
name := fmt.Sprintf("%s/%d", k.name, size)
|
|
||||||
t.Run(name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
checkRoundTrip(t, id, recipient, 1, k.fill(t, size))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestZeroLengthNoWrite covers issue case 3: a Writer closed with no Write at
|
|
||||||
// all produces the double hash of the empty input, and the blob reads back as
|
|
||||||
// empty with no error.
|
|
||||||
func TestZeroLengthNoWrite(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
id, recipient := makeIdentity(t)
|
|
||||||
|
|
||||||
var buf bytes.Buffer
|
|
||||||
|
|
||||||
w, err := blobgen.NewWriter(&buf, 1, []string{recipient})
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, w.Close())
|
|
||||||
assert.Equal(t, int64(0), w.BytesWritten())
|
|
||||||
|
|
||||||
empty := sha256.Sum256(nil)
|
|
||||||
doubled := sha256.Sum256(empty[:])
|
|
||||||
assert.Equal(t, doubled[:], w.ContentID(),
|
|
||||||
"ContentID of empty input is SHA256(SHA256(\"\"))")
|
|
||||||
|
|
||||||
r, err := blobgen.NewReader(bytes.NewReader(buf.Bytes()), id)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
got, err := io.ReadAll(r)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, r.Close())
|
|
||||||
|
|
||||||
assert.Empty(t, got, "empty blob decrypts to empty output")
|
|
||||||
assert.Equal(t, int64(0), r.BytesRead())
|
|
||||||
assert.Equal(t, empty[:], r.Sum256())
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewWriterValidLevelsRoundTrip covers the accepted end of issue case 9:
|
|
||||||
// the boundary compression levels 1 and 19 both round trip.
|
|
||||||
func TestNewWriterValidLevelsRoundTrip(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
id, recipient := makeIdentity(t)
|
|
||||||
input := randomBytes(t, 4096)
|
|
||||||
|
|
||||||
for _, level := range []int{1, 19} {
|
|
||||||
t.Run(fmt.Sprintf("level%d", level), func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
checkRoundTrip(t, id, recipient, level, input)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+13
-30
@@ -1,6 +1,3 @@
|
|||||||
// Package blobgen implements the blob data pipeline: streaming zstd
|
|
||||||
// compression, age encryption, and SHA256 content hashing for blob
|
|
||||||
// creation, plus the matching decrypt/decompress/verify reader.
|
|
||||||
package blobgen
|
package blobgen
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -15,18 +12,6 @@ import (
|
|||||||
"github.com/klauspost/compress/zstd"
|
"github.com/klauspost/compress/zstd"
|
||||||
)
|
)
|
||||||
|
|
||||||
// DoubleSHA256 returns the double SHA-256 of content whose single SHA-256
|
|
||||||
// digest is sum: it hashes that digest once more. Stored objects are named by
|
|
||||||
// this second hash so that a name never reveals whether known content is
|
|
||||||
// present — an attacker who knows a plaintext, and thus its SHA-256, still
|
|
||||||
// cannot derive the stored name without hashing the digest again. Both a blob
|
|
||||||
// and the metadata database export are named this way.
|
|
||||||
func DoubleSHA256(sum []byte) []byte {
|
|
||||||
h := sha256.Sum256(sum)
|
|
||||||
|
|
||||||
return h[:]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Zstd compression level bounds accepted by NewWriter.
|
// Zstd compression level bounds accepted by NewWriter.
|
||||||
const (
|
const (
|
||||||
minCompressionLevel = 1
|
minCompressionLevel = 1
|
||||||
@@ -42,11 +27,6 @@ const reservedCompressionCPUs = 2
|
|||||||
var ErrInvalidCompressionLevel = errors.New(
|
var ErrInvalidCompressionLevel = errors.New(
|
||||||
"invalid compression level: must be between 1 and 19")
|
"invalid compression level: must be between 1 and 19")
|
||||||
|
|
||||||
// errInvalidRecipient is returned when a recipient string does not parse as
|
|
||||||
// an X25519 age1... public key. It omits the value, which can be sensitive.
|
|
||||||
var errInvalidRecipient = errors.New(
|
|
||||||
"not a valid X25519 age1... recipient")
|
|
||||||
|
|
||||||
// Writer wraps compression and encryption with SHA256 hashing.
|
// Writer wraps compression and encryption with SHA256 hashing.
|
||||||
// Data flows: input -> tee(hasher, compressor -> encryptor -> destination)
|
// Data flows: input -> tee(hasher, compressor -> encryptor -> destination)
|
||||||
// The hash is computed on the uncompressed input for deterministic content-addressing.
|
// The hash is computed on the uncompressed input for deterministic content-addressing.
|
||||||
@@ -77,12 +57,10 @@ func NewWriter(
|
|||||||
// Parse recipients
|
// Parse recipients
|
||||||
var ageRecipients []age.Recipient
|
var ageRecipients []age.Recipient
|
||||||
|
|
||||||
for i, recipient := range recipients {
|
for _, recipient := range recipients {
|
||||||
// The recipient string can be sensitive (e.g. a secret key pasted by
|
|
||||||
// mistake), so the error names its position, never its value.
|
|
||||||
r, err := age.ParseX25519Recipient(recipient)
|
r, err := age.ParseX25519Recipient(recipient)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("%w: recipient %d", errInvalidRecipient, i)
|
return nil, fmt.Errorf("parsing recipient %s: %w", recipient, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ageRecipients = append(ageRecipients, r)
|
ageRecipients = append(ageRecipients, r)
|
||||||
@@ -145,12 +123,17 @@ func (w *Writer) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ContentID returns the double SHA-256 of the uncompressed input data: the
|
// Sum256 returns the double SHA256 hash of the uncompressed input data.
|
||||||
// name under which this content is stored. It is the second hash of the
|
// Double hashing (SHA256(SHA256(data))) prevents information leakage about
|
||||||
// running SHA-256, via DoubleSHA256; see that function for why content is
|
// the plaintext - an attacker cannot confirm existence of known content
|
||||||
// named this way rather than by its plain SHA-256.
|
// by computing its hash and checking for a matching blob filename.
|
||||||
func (w *Writer) ContentID() []byte {
|
func (w *Writer) Sum256() []byte {
|
||||||
return DoubleSHA256(w.hasher.Sum(nil))
|
// First hash: SHA256(plaintext)
|
||||||
|
firstHash := w.hasher.Sum(nil)
|
||||||
|
// Second hash: SHA256(firstHash) - this is the blob ID
|
||||||
|
secondHash := sha256.Sum256(firstHash)
|
||||||
|
|
||||||
|
return secondHash[:]
|
||||||
}
|
}
|
||||||
|
|
||||||
// BytesWritten returns the number of uncompressed bytes written
|
// BytesWritten returns the number of uncompressed bytes written
|
||||||
|
|||||||
@@ -12,10 +12,9 @@ import (
|
|||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestWriterHashIsDoubleHash verifies that Writer.ContentID() returns
|
// TestWriterHashIsDoubleHash verifies that Writer.Sum256() returns
|
||||||
// SHA256(SHA256(plaintext)). Stored objects are named by this second hash so a
|
// the double hash SHA256(SHA256(plaintext)) for security.
|
||||||
// name is not the plaintext's own SHA-256; this does not stop someone who
|
// Double hashing prevents attackers from confirming existence of known content.
|
||||||
// already holds the plaintext from confirming it.
|
|
||||||
func TestWriterHashIsDoubleHash(t *testing.T) {
|
func TestWriterHashIsDoubleHash(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -44,7 +43,7 @@ func TestWriterHashIsDoubleHash(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
// Get the hash from the writer
|
// Get the hash from the writer
|
||||||
writerHash := hex.EncodeToString(writer.ContentID())
|
writerHash := hex.EncodeToString(writer.Sum256())
|
||||||
|
|
||||||
// Calculate the expected double hash: SHA256(SHA256(plaintext))
|
// Calculate the expected double hash: SHA256(SHA256(plaintext))
|
||||||
firstHash := sha256.Sum256(testData)
|
firstHash := sha256.Sum256(testData)
|
||||||
@@ -61,11 +60,11 @@ func TestWriterHashIsDoubleHash(t *testing.T) {
|
|||||||
|
|
||||||
// The writer hash should match the double hash
|
// The writer hash should match the double hash
|
||||||
assert.Equal(t, expectedDoubleHash, writerHash,
|
assert.Equal(t, expectedDoubleHash, writerHash,
|
||||||
"Writer.ContentID() must be SHA256(SHA256(plaintext))")
|
"Writer.Sum256() should return SHA256(SHA256(plaintext)) for security")
|
||||||
|
|
||||||
// It must be the second hash, not the plaintext's own SHA-256.
|
// Verify it's NOT the single hash (would leak information)
|
||||||
assert.NotEqual(t, singleHashStr, writerHash,
|
assert.NotEqual(t, singleHashStr, writerHash,
|
||||||
"Writer hash must be the double hash, not the single SHA-256")
|
"Writer hash should not be single hash (would allow content confirmation attacks)")
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestWriterDeterministicHash verifies that the same input always produces
|
// TestWriterDeterministicHash verifies that the same input always produces
|
||||||
@@ -94,8 +93,8 @@ func TestWriterDeterministicHash(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.NoError(t, writer2.Close())
|
require.NoError(t, writer2.Close())
|
||||||
|
|
||||||
hash1 := hex.EncodeToString(writer1.ContentID())
|
hash1 := hex.EncodeToString(writer1.Sum256())
|
||||||
hash2 := hex.EncodeToString(writer2.ContentID())
|
hash2 := hex.EncodeToString(writer2.Sum256())
|
||||||
|
|
||||||
// Hashes should be identical (deterministic)
|
// Hashes should be identical (deterministic)
|
||||||
assert.Equal(t, hash1, hash2, "Same input should produce same hash")
|
assert.Equal(t, hash1, hash2, "Same input should produce same hash")
|
||||||
@@ -109,20 +108,3 @@ func TestWriterDeterministicHash(t *testing.T) {
|
|||||||
t.Logf("Encrypted size 1: %d bytes", buf1.Len())
|
t.Logf("Encrypted size 1: %d bytes", buf1.Len())
|
||||||
t.Logf("Encrypted size 2: %d bytes", buf2.Len())
|
t.Logf("Encrypted size 2: %d bytes", buf2.Len())
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestNewWriterSecretKeyNotEchoed verifies that a secret key mistakenly passed
|
|
||||||
// as a recipient does not appear in the returned error. A recipient string can
|
|
||||||
// be sensitive, so the error must name only the position, not the value.
|
|
||||||
func TestNewWriterSecretKeyNotEchoed(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
secretKey := "AGE-SECRET-KEY-19CR5YSFW59HM4TLD6GX" +
|
|
||||||
"VEDMZFTVVF7PPHKUT68TXSFPK7APHXA2QS2NJA5"
|
|
||||||
|
|
||||||
var buf bytes.Buffer
|
|
||||||
|
|
||||||
_, err := blobgen.NewWriter(&buf, 3, []string{secretKey})
|
|
||||||
require.Error(t, err)
|
|
||||||
assert.NotContains(t, err.Error(), secretKey,
|
|
||||||
"error must not echo the recipient value")
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -33,10 +33,9 @@ type Chunker struct {
|
|||||||
maxChunkSize int
|
maxChunkSize int
|
||||||
}
|
}
|
||||||
|
|
||||||
// ChunkSizeSpread is the FastCDC-recommended factor between the average
|
// chunkSizeSpread is the FastCDC-recommended factor between the average
|
||||||
// chunk size and the minimum (avg/spread) and maximum (avg*spread) sizes.
|
// chunk size and the minimum (avg/spread) and maximum (avg*spread) sizes.
|
||||||
// The largest chunk the chunker can emit is therefore avg*ChunkSizeSpread.
|
const chunkSizeSpread = 4
|
||||||
const ChunkSizeSpread = 4
|
|
||||||
|
|
||||||
// NewChunker creates a new chunker with the specified average chunk size.
|
// NewChunker creates a new chunker with the specified average chunk size.
|
||||||
// The actual chunk sizes will vary between avgChunkSize/4 and avgChunkSize*4
|
// The actual chunk sizes will vary between avgChunkSize/4 and avgChunkSize*4
|
||||||
@@ -46,8 +45,8 @@ func NewChunker(avgChunkSize int64) *Chunker {
|
|||||||
// FastCDC recommends min = avg/4 and max = avg*4
|
// FastCDC recommends min = avg/4 and max = avg*4
|
||||||
return &Chunker{
|
return &Chunker{
|
||||||
avgChunkSize: int(avgChunkSize),
|
avgChunkSize: int(avgChunkSize),
|
||||||
minChunkSize: int(avgChunkSize / ChunkSizeSpread),
|
minChunkSize: int(avgChunkSize / chunkSizeSpread),
|
||||||
maxChunkSize: int(avgChunkSize * ChunkSizeSpread),
|
maxChunkSize: int(avgChunkSize * chunkSizeSpread),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+104
-192
@@ -7,9 +7,11 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"os/signal"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/adrg/xdg"
|
"github.com/adrg/xdg"
|
||||||
@@ -30,33 +32,14 @@ import (
|
|||||||
// may take before we give up.
|
// may take before we give up.
|
||||||
const shutdownTimeout = 30 * time.Second
|
const shutdownTimeout = 30 * time.Second
|
||||||
|
|
||||||
// lockMode says whether a command mutates persistent state — the local
|
// AppOptions contains common options for creating the fx application.
|
||||||
// index database or the remote store — and so must hold the process-wide
|
// It includes the configuration file path, logging options, and additional
|
||||||
// PID lock, or only reads that state and may run alongside a mutator.
|
// fx modules and invocations that should be included in the application.
|
||||||
type lockMode int
|
|
||||||
|
|
||||||
const (
|
|
||||||
// mutating commands (snapshot create, snapshot purge, snapshot remove,
|
|
||||||
// prune, remote nuke) write the local index or the remote store. They
|
|
||||||
// hold the PID lock so that at most one runs at a time.
|
|
||||||
mutating lockMode = iota
|
|
||||||
// readOnly commands (info, snapshot list, snapshot verify, remote info,
|
|
||||||
// snapshot restore) do not write the local index or the remote store,
|
|
||||||
// so they run without the lock and are never blocked by a running
|
|
||||||
// mutator. restore writes only to the target directory it is given.
|
|
||||||
readOnly
|
|
||||||
)
|
|
||||||
|
|
||||||
// AppOptions contains common options for creating and running the fx
|
|
||||||
// application: the configuration file path, logging options, additional fx
|
|
||||||
// modules and invocations, and whether the command mutates persistent
|
|
||||||
// state (which decides whether it takes the PID lock).
|
|
||||||
type AppOptions struct {
|
type AppOptions struct {
|
||||||
ConfigPath string
|
ConfigPath string
|
||||||
LogOptions log.Options
|
LogOptions log.Options
|
||||||
Modules []fx.Option
|
Modules []fx.Option
|
||||||
Invokes []fx.Option
|
Invokes []fx.Option
|
||||||
Mode lockMode
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// setupGlobals records the startup time and, when an output-suppression
|
// setupGlobals records the startup time and, when an output-suppression
|
||||||
@@ -65,11 +48,6 @@ type AppOptions struct {
|
|||||||
// silenced — per the documented convention that --quiet suppresses
|
// silenced — per the documented convention that --quiet suppresses
|
||||||
// non-error output only. The startup banner is printed by Entry
|
// non-error output only. The startup banner is printed by Entry
|
||||||
// before cobra parses arguments, gated by the same arg-level check.
|
// before cobra parses arguments, gated by the same arg-level check.
|
||||||
//
|
|
||||||
// --json quiets the UI here too, because stdout then carries a JSON
|
|
||||||
// document and human narration would corrupt it. Unlike Quiet it does
|
|
||||||
// not lower the stderr log level (issue #112), so --verbose/--debug
|
|
||||||
// still surface diagnostics alongside the document.
|
|
||||||
func setupGlobals(
|
func setupGlobals(
|
||||||
lc fx.Lifecycle, g *globals.Globals, v *vaultik.Vaultik, opts log.Options,
|
lc fx.Lifecycle, g *globals.Globals, v *vaultik.Vaultik, opts log.Options,
|
||||||
) {
|
) {
|
||||||
@@ -77,7 +55,7 @@ func setupGlobals(
|
|||||||
OnStart: func(_ context.Context) error {
|
OnStart: func(_ context.Context) error {
|
||||||
g.StartTime = time.Now().UTC()
|
g.StartTime = time.Now().UTC()
|
||||||
|
|
||||||
if opts.Cron || opts.Quiet || opts.JSON {
|
if opts.Cron || opts.Quiet {
|
||||||
v.UI.SetQuiet(true)
|
v.UI.SetQuiet(true)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -158,148 +136,75 @@ func cleanStartupError(err error) error {
|
|||||||
return &startupError{msg: msg}
|
return &startupError{msg: msg}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RunApp starts the fx application, blocks until it is asked to stop, and
|
// RunApp starts and stops the fx application within the given context.
|
||||||
// then stops it. The app is asked to stop either by an OS interrupt
|
// It handles graceful shutdown on interrupt signals (SIGINT, SIGTERM) and
|
||||||
// (SIGINT/SIGTERM — fx installs its own handler when app.Wait is called) or,
|
// ensures the application stops cleanly. The function blocks until the
|
||||||
// on normal completion, by the finished operation calling
|
// application completes or is interrupted. Returns an error if startup fails.
|
||||||
// Shutdowner.Shutdown(); both arrive on the app.Wait channel.
|
|
||||||
//
|
|
||||||
// Stopping runs the fx OnStop hooks, and RunApp does not return until Stop
|
|
||||||
// returns. On an interrupt the operation's OnStop hook cancels the running
|
|
||||||
// command and waits for it to unwind — removing its decrypted scratch files —
|
|
||||||
// so the process cannot proceed to exit mid-cleanup (issue #159). Waiting for
|
|
||||||
// Stop before returning is what makes that hook effective: routing the
|
|
||||||
// interrupt through app.Stop and not returning until it completes is required,
|
|
||||||
// because fx also fires the app.Wait channel on the signal, and an earlier
|
|
||||||
// version returned on that alone — unwinding to os.Exit while the concurrent
|
|
||||||
// cleanup still ran. The stop is bounded by shutdownTimeout. Returns an error
|
|
||||||
// if startup fails.
|
|
||||||
func RunApp(ctx context.Context, app *fx.App) error {
|
func RunApp(ctx context.Context, app *fx.App) error {
|
||||||
|
// Set up signal handling for graceful shutdown
|
||||||
|
sigChan := make(chan os.Signal, 1)
|
||||||
|
signal.Notify(sigChan, os.Interrupt, syscall.SIGTERM)
|
||||||
|
|
||||||
|
// Create a context that will be cancelled on signal
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
// Start the app
|
||||||
err := app.Start(ctx)
|
err := app.Start(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return cleanStartupError(err)
|
return cleanStartupError(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Block until an interrupt or the finished operation's
|
// Handle shutdown
|
||||||
// Shutdowner.Shutdown() arrives, then stop the app in this goroutine so we
|
shutdownComplete := make(chan struct{})
|
||||||
// return only after its OnStop hooks — including the operation's cleanup
|
|
||||||
// wait — have run. Detach the stop from ctx's cancellation but keep its
|
|
||||||
// values, and bound it by shutdownTimeout.
|
|
||||||
<-app.Wait()
|
|
||||||
|
|
||||||
shutdownCtx, cancel := context.WithTimeout(
|
go func() {
|
||||||
|
defer close(shutdownComplete)
|
||||||
|
|
||||||
|
<-sigChan
|
||||||
|
log.Notice("Received interrupt signal, shutting down gracefully...")
|
||||||
|
|
||||||
|
// Create a timeout context for shutdown. The parent ctx is being
|
||||||
|
// cancelled, so detach from its cancellation but keep its values.
|
||||||
|
shutdownCtx, shutdownCancel := context.WithTimeout(
|
||||||
context.WithoutCancel(ctx), shutdownTimeout)
|
context.WithoutCancel(ctx), shutdownTimeout)
|
||||||
defer cancel()
|
defer shutdownCancel()
|
||||||
|
|
||||||
err = app.Stop(shutdownCtx)
|
err := app.Stop(shutdownCtx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Error("Error during shutdown", "error", err)
|
log.Error("Error during shutdown", "error", err)
|
||||||
}
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Wait for the signal handler to complete shutdown or the app to
|
||||||
|
// request shutdown.
|
||||||
|
select {
|
||||||
|
case <-shutdownComplete:
|
||||||
|
// Shutdown completed via signal
|
||||||
return nil
|
return nil
|
||||||
}
|
case <-ctx.Done():
|
||||||
|
// Context cancelled (shouldn't happen in normal operation)
|
||||||
// errReported marks a failure the operation has already shown the user
|
err := app.Stop(context.WithoutCancel(ctx))
|
||||||
// (and deliberately withheld under --json). Entry turns it into a
|
|
||||||
// non-zero exit status without printing anything further, so the error
|
|
||||||
// line is not doubled. It flows up from RunOperation through cobra to
|
|
||||||
// Entry.
|
|
||||||
var errReported = errors.New("operation failed")
|
|
||||||
|
|
||||||
// RunOperation runs op against the Vaultik instance inside the fx app
|
|
||||||
// and turns a failure into a returned error rather than an os.Exit from
|
|
||||||
// within the goroutine. An os.Exit there skipped main's deferred
|
|
||||||
// profile writers -- so profiling a failing command yielded a truncated
|
|
||||||
// profile (issue #75) -- and RunWithApp's PID-lock release, and denied
|
|
||||||
// the app any graceful shutdown; returning the error to the top runs
|
|
||||||
// all three.
|
|
||||||
//
|
|
||||||
// op runs in a goroutine so OnStart returns promptly and an interrupt
|
|
||||||
// can still cancel through OnStop; when it finishes, success or failure,
|
|
||||||
// it triggers shutdown, which is what lets RunWithApp return. On an
|
|
||||||
// interrupt OnStop cancels op and waits for the goroutine to return, so
|
|
||||||
// op's cleanup (removing decrypted scratch files) runs before the
|
|
||||||
// process exits; the wait is bounded by shutdownTimeout. report is
|
|
||||||
// called with a non-canceled failure so the caller can log it (and
|
|
||||||
// suppress it under --json) before it becomes errReported. A context
|
|
||||||
// cancellation is the interrupt path, not a failure: it is neither
|
|
||||||
// reported nor counted as one.
|
|
||||||
func RunOperation(
|
|
||||||
ctx context.Context, opts AppOptions,
|
|
||||||
op func(v *vaultik.Vaultik) error, report func(err error),
|
|
||||||
) error {
|
|
||||||
var (
|
|
||||||
mu sync.Mutex
|
|
||||||
failed bool
|
|
||||||
)
|
|
||||||
|
|
||||||
opts.Invokes = append(opts.Invokes,
|
|
||||||
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
|
||||||
var stop func(context.Context) bool
|
|
||||||
|
|
||||||
lc.Append(fx.Hook{
|
|
||||||
OnStart: func(_ context.Context) error {
|
|
||||||
stop = v.StartOperation(func() {
|
|
||||||
err := op(v)
|
|
||||||
if err != nil && !errors.Is(err, context.Canceled) {
|
|
||||||
report(err)
|
|
||||||
|
|
||||||
mu.Lock()
|
|
||||||
failed = true
|
|
||||||
mu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
stopErr := v.Shutdowner.Shutdown()
|
|
||||||
if stopErr != nil {
|
|
||||||
log.Error("Failed to shutdown", "error", stopErr)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
return nil
|
|
||||||
},
|
|
||||||
// On an interrupt, cancel the operation and wait for it to
|
|
||||||
// unwind so its cleanup defers (which remove decrypted
|
|
||||||
// scratch files from the temp directory) run before the
|
|
||||||
// process exits. The wait is bounded by ctx, the existing
|
|
||||||
// shutdownTimeout.
|
|
||||||
OnStop: func(ctx context.Context) error {
|
|
||||||
if !stop(ctx) {
|
|
||||||
log.Warn("Shutdown timed out before the operation " +
|
|
||||||
"finished; decrypted temporary files may remain")
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}))
|
|
||||||
|
|
||||||
err := RunWithApp(ctx, opts)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
log.Error("Error stopping app", "error", err)
|
||||||
}
|
|
||||||
|
|
||||||
// The goroutine sets failed before triggering the shutdown that lets
|
|
||||||
// RunWithApp return, so the write is in place by the time we read it.
|
|
||||||
mu.Lock()
|
|
||||||
defer mu.Unlock()
|
|
||||||
|
|
||||||
if failed {
|
|
||||||
return errReported
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return ctx.Err()
|
||||||
|
case <-app.Done():
|
||||||
|
// App finished running (e.g., backup completed)
|
||||||
return nil
|
return nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// runVaultikApp runs the standard single-operation command lifecycle
|
// runVaultikApp runs the standard single-operation command lifecycle
|
||||||
// shared by the snapshot list/purge/remove and remote nuke subcommands:
|
// shared by the list/purge/verify/remove/remote-info subcommands:
|
||||||
// resolve the config, then run op against the Vaultik instance through
|
// resolve the config, start the fx app, run op against the Vaultik
|
||||||
// RunOperation, reporting a failure prefixed with failMsg (suppressed
|
// instance in a goroutine, report a failure prefixed with failMsg
|
||||||
// while suppressErrors is true, e.g. under --json). mode says whether the
|
// (suppressed while suppressErrors is true, e.g. under --json), then
|
||||||
// command takes the PID lock. jsonOutput marks a command whose stdout is a
|
// trigger shutdown. The operation is cancelled when the app stops.
|
||||||
// JSON document: it quiets the UI but, unlike Quiet, leaves the stderr log
|
// extraQuiet is OR-ed into LogOptions.Quiet (e.g. --json output modes).
|
||||||
// level alone.
|
|
||||||
func runVaultikApp(
|
func runVaultikApp(
|
||||||
cmd *cobra.Command, mode lockMode, jsonOutput, suppressErrors bool,
|
cmd *cobra.Command, extraQuiet, suppressErrors bool,
|
||||||
failMsg string, op func(v *vaultik.Vaultik) error,
|
failMsg string, op func(v *vaultik.Vaultik) error,
|
||||||
) error {
|
) error {
|
||||||
configPath, err := ResolveConfigPath()
|
configPath, err := ResolveConfigPath()
|
||||||
@@ -309,68 +214,75 @@ func runVaultikApp(
|
|||||||
|
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
return RunOperation(cmd.Context(), AppOptions{
|
return RunWithApp(cmd.Context(), AppOptions{
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogOptions: log.Options{
|
LogOptions: log.Options{
|
||||||
Verbose: rootFlags.Verbose,
|
Verbose: rootFlags.Verbose,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet || extraQuiet,
|
||||||
JSON: jsonOutput,
|
|
||||||
},
|
},
|
||||||
Mode: mode,
|
Modules: []fx.Option{},
|
||||||
}, op, func(err error) {
|
Invokes: []fx.Option{
|
||||||
if suppressErrors {
|
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
||||||
return
|
lc.Append(fx.Hook{
|
||||||
}
|
OnStart: func(_ context.Context) error {
|
||||||
|
go func() {
|
||||||
|
err := op(v)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, context.Canceled) {
|
||||||
|
if !suppressErrors {
|
||||||
log.Error(failMsg, "error", err)
|
log.Error(failMsg, "error", err)
|
||||||
ReportErrorf("%s: %v", failMsg, err)
|
ReportErrorf("%s: %v", failMsg, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = v.Shutdowner.Shutdown()
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to shutdown", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
OnStop: func(_ context.Context) error {
|
||||||
|
v.Cancel()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}),
|
||||||
|
},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// RunWithApp is a helper that creates and runs an fx app with the given options.
|
// RunWithApp is a helper that creates and runs an fx app with the given options.
|
||||||
// It combines NewApp and RunApp into a single convenient function. This is the
|
// It combines NewApp and RunApp into a single convenient function. This is the
|
||||||
// preferred way to run CLI commands that need the full application context.
|
// preferred way to run CLI commands that need the full application context.
|
||||||
// A mutating command takes the process-wide PID lock before starting so that
|
// It acquires a PID lock before starting to prevent concurrent instances.
|
||||||
// only one runs at a time; a read-only command runs without it and is not
|
|
||||||
// blocked while a mutator holds the lock (opts.Mode).
|
|
||||||
func RunWithApp(ctx context.Context, opts AppOptions) error {
|
func RunWithApp(ctx context.Context, opts AppOptions) error {
|
||||||
release, err := acquireLockIfMutating(opts.Mode,
|
// Acquire PID lock to prevent concurrent instances
|
||||||
filepath.Join(xdg.DataHome, "vaultik"))
|
lockDir := filepath.Join(xdg.DataHome, "vaultik")
|
||||||
|
|
||||||
|
lock, err := pidlock.Acquire(lockDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
if errors.Is(err, pidlock.ErrAlreadyRunning) {
|
||||||
|
return fmt.Errorf("cannot start: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
defer release()
|
return fmt.Errorf("failed to acquire lock: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
err := lock.Release()
|
||||||
|
if err != nil {
|
||||||
|
log.Warn("Failed to release PID lock", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
app := NewApp(opts)
|
app := NewApp(opts)
|
||||||
|
|
||||||
return RunApp(ctx, app)
|
return RunApp(ctx, app)
|
||||||
}
|
}
|
||||||
|
|
||||||
// acquireLockIfMutating takes the process-wide PID lock in lockDir for a
|
|
||||||
// mutating command and returns a function that releases it. A read-only
|
|
||||||
// command takes no lock, so it returns a no-op release and is never blocked
|
|
||||||
// while a mutator holds the lock. ErrAlreadyRunning (another mutator holds
|
|
||||||
// the lock) is surfaced as a "cannot start" error.
|
|
||||||
func acquireLockIfMutating(mode lockMode, lockDir string) (func(), error) {
|
|
||||||
if mode != mutating {
|
|
||||||
return func() {}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
lock, err := pidlock.Acquire(lockDir)
|
|
||||||
if err != nil {
|
|
||||||
if errors.Is(err, pidlock.ErrAlreadyRunning) {
|
|
||||||
return nil, fmt.Errorf("cannot start: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil, fmt.Errorf("failed to acquire lock: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return func() {
|
|
||||||
err := lock.Release()
|
|
||||||
if err != nil {
|
|
||||||
log.Warn("Failed to release PID lock", "error", err)
|
|
||||||
}
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -2,10 +2,7 @@ package cli //nolint:testpackage // needs access to unexported cleanStartupError
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/pidlock"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestCleanStartupError(t *testing.T) {
|
func TestCleanStartupError(t *testing.T) {
|
||||||
@@ -56,42 +53,3 @@ func TestCleanStartupError(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestLockScopedToMutatingCommands proves the partition the PID lock now
|
|
||||||
// enforces: a read-only command runs while a mutator holds the lock, and
|
|
||||||
// two mutating commands still mutually exclude.
|
|
||||||
func TestLockScopedToMutatingCommands(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
lockDir := filepath.Join(t.TempDir(), "vaultik")
|
|
||||||
|
|
||||||
// A mutating command takes the process-wide lock.
|
|
||||||
releaseMutator, err := acquireLockIfMutating(mutating, lockDir)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("mutating command could not acquire lock: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// A read-only command runs to completion even while the lock is held.
|
|
||||||
releaseReader, err := acquireLockIfMutating(readOnly, lockDir)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read-only command was blocked by held lock: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
releaseReader()
|
|
||||||
|
|
||||||
// A second mutating command is refused while the first holds the lock.
|
|
||||||
_, err = acquireLockIfMutating(mutating, lockDir)
|
|
||||||
if !errors.Is(err, pidlock.ErrAlreadyRunning) {
|
|
||||||
t.Fatalf("second mutating command was not excluded, got: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Once the first mutator releases, another mutating command may run.
|
|
||||||
releaseMutator()
|
|
||||||
|
|
||||||
release, err := acquireLockIfMutating(mutating, lockDir)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("mutating command could not acquire released lock: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
release()
|
|
||||||
}
|
|
||||||
|
|||||||
+13
-29
@@ -4,7 +4,6 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -193,8 +192,8 @@ storage_url: ""
|
|||||||
# access_key_id: YOUR_ACCESS_KEY
|
# access_key_id: YOUR_ACCESS_KEY
|
||||||
# secret_access_key: YOUR_SECRET_KEY
|
# secret_access_key: YOUR_SECRET_KEY
|
||||||
# # region: us-east-1 # Default: us-east-1
|
# # region: us-east-1 # Default: us-east-1
|
||||||
|
# # use_ssl: true # Default: true
|
||||||
# # part_size: 5MB # Multipart upload part size. Default: 5MB
|
# # part_size: 5MB # Multipart upload part size. Default: 5MB
|
||||||
# # For the s3:// form, disable TLS with ?ssl=false in the URL, not use_ssl.
|
|
||||||
|
|
||||||
# ─── OPTIONAL ────────────────────────────────────────────────────────────────
|
# ─── OPTIONAL ────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
@@ -213,8 +212,6 @@ storage_url: ""
|
|||||||
# chunk_size: 10MB
|
# chunk_size: 10MB
|
||||||
|
|
||||||
# Maximum blob size before splitting into a new blob.
|
# Maximum blob size before splitting into a new blob.
|
||||||
# Must be at least four times chunk_size (the largest chunk the chunker can
|
|
||||||
# emit); a smaller limit would let a single-chunk blob exceed it.
|
|
||||||
# Accepts: 1GB, 10G, 500MB, etc.
|
# Accepts: 1GB, 10G, 500MB, etc.
|
||||||
# Default: 10GB
|
# Default: 10GB
|
||||||
# blob_size_limit: 10GB
|
# blob_size_limit: 10GB
|
||||||
@@ -380,23 +377,12 @@ Examples:
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return writeConfigSet(os.Stdout, path, args[0], args[1])
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeConfigSet applies key=value to the config at path, writes it back
|
|
||||||
// owner-only, and confirms the write by printing just the key name to w.
|
|
||||||
// The value is never echoed: it may be a secret such as
|
|
||||||
// s3.secret_access_key, and captured stdout or a pasted terminal would
|
|
||||||
// then leak it.
|
|
||||||
func writeConfigSet(w io.Writer, path, key, value string) error {
|
|
||||||
root, err := loadYAMLFile(path)
|
root, err := loadYAMLFile(path)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
err = yamlPathSet(root, strings.Split(key, "."), value)
|
err = yamlPathSet(root, strings.Split(args[0], "."), args[1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -406,25 +392,23 @@ func writeConfigSet(w io.Writer, path, key, value string) error {
|
|||||||
return fmt.Errorf("marshaling config: %w", err)
|
return fmt.Errorf("marshaling config: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = os.WriteFile(path, out, configFileMode)
|
mode := os.FileMode(configFileMode)
|
||||||
|
|
||||||
|
info, statErr := os.Stat(path)
|
||||||
|
if statErr == nil {
|
||||||
|
mode = info.Mode().Perm()
|
||||||
|
}
|
||||||
|
|
||||||
|
err = os.WriteFile(path, out, mode)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("writing config file: %w", err)
|
return fmt.Errorf("writing config file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// os.WriteFile does not change the mode of a file that already exists,
|
_, _ = fmt.Fprintf(os.Stdout, "%s = %s\n", args[0], args[1])
|
||||||
// so a config that was group- or world-readable stays that way. As it
|
|
||||||
// may hold S3 credentials, tighten it to owner-only after writing.
|
|
||||||
info, statErr := os.Stat(path)
|
|
||||||
if statErr == nil && info.Mode().Perm()&0o044 != 0 {
|
|
||||||
err = os.Chmod(path, configFileMode)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("tightening config file permissions: %w", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
_, _ = fmt.Fprintln(w, key)
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// marshalConfigYAML renders a config document tree with 2-space indentation,
|
// marshalConfigYAML renders a config document tree with 2-space indentation,
|
||||||
|
|||||||
@@ -1,9 +1,6 @@
|
|||||||
package cli //nolint:testpackage // exercises unexported yamlPathGet/yamlPathSet
|
package cli //nolint:testpackage // exercises unexported yamlPathGet/yamlPathSet
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -232,68 +229,6 @@ func TestConfigSetPreservesFormatting(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestWriteConfigSetHidesSecret checks that setting a secret key prints
|
|
||||||
// only the key name, never the value, to the confirmation output.
|
|
||||||
func TestWriteConfigSetHidesSecret(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const secret = "SUPERSECRETVALUE"
|
|
||||||
|
|
||||||
path := filepath.Join(t.TempDir(), "config.yaml")
|
|
||||||
|
|
||||||
err := os.WriteFile(path, []byte("version: 1\n"), 0o600)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("seed config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var out bytes.Buffer
|
|
||||||
|
|
||||||
err = writeConfigSet(&out, path, "s3.secret_access_key", secret)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("writeConfigSet: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if strings.Contains(out.String(), secret) {
|
|
||||||
t.Errorf("output echoed the secret value: %q", out.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
if !strings.Contains(out.String(), "s3.secret_access_key") {
|
|
||||||
t.Errorf("output did not confirm the key name: %q", out.String())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWriteConfigSetTightensMode checks that a pre-existing group- or
|
|
||||||
// world-readable config is tightened to owner-only after a set, since
|
|
||||||
// os.WriteFile leaves an existing file's mode untouched.
|
|
||||||
func TestWriteConfigSetTightensMode(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
path := filepath.Join(t.TempDir(), "config.yaml")
|
|
||||||
|
|
||||||
// Seed a world-readable config; the loose mode is the condition under
|
|
||||||
// test, so gosec's G306 is expected here.
|
|
||||||
err := os.WriteFile(path, []byte("version: 1\n"), 0o644) //nolint:gosec // G306
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("seed config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var out bytes.Buffer
|
|
||||||
|
|
||||||
err = writeConfigSet(&out, path, "compression_level", "9")
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("writeConfigSet: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
info, err := os.Stat(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("stat config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if info.Mode().Perm() != 0o600 {
|
|
||||||
t.Errorf("config mode = %04o, want 0600", info.Mode().Perm())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func splitPath(s string) []string {
|
func splitPath(s string) []string {
|
||||||
return strings.Split(s, ".")
|
return strings.Split(s, ".")
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-17
@@ -1,7 +1,6 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -20,11 +19,7 @@ const shortCommitLen = 12
|
|||||||
// flag is present in os.Args — see bannerSuppressedInArgs), executes the
|
// flag is present in os.Args — see bannerSuppressedInArgs), executes the
|
||||||
// root cobra command, and routes any returned error through the
|
// root cobra command, and routes any returned error through the
|
||||||
// ui.Writer so the user sees a properly formatted "🛑 ERROR:" line.
|
// ui.Writer so the user sees a properly formatted "🛑 ERROR:" line.
|
||||||
//
|
func Entry() {
|
||||||
// It returns the process exit code (0 on success, 1 on error) rather
|
|
||||||
// than calling os.Exit, so that main's deferred profile writers run
|
|
||||||
// before the process ends. See run in cmd/vaultik/main.go.
|
|
||||||
func Entry() int {
|
|
||||||
emitStartupBanner(os.Args[1:], os.Stdout)
|
emitStartupBanner(os.Args[1:], os.Stdout)
|
||||||
|
|
||||||
rootCmd := NewRootCommand()
|
rootCmd := NewRootCommand()
|
||||||
@@ -32,19 +27,9 @@ func Entry() int {
|
|||||||
|
|
||||||
err := rootCmd.Execute()
|
err := rootCmd.Execute()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// An operation that ran inside the fx app has already reported
|
|
||||||
// its own failure (and suppressed it under --json); errReported
|
|
||||||
// says so. Printing it again here would double the error line.
|
|
||||||
// Every other error — bad arguments, a config that would not
|
|
||||||
// load — reaches Entry unreported, so it is shown here.
|
|
||||||
if !errors.Is(err, errReported) {
|
|
||||||
ReportErrorf("%s", err.Error())
|
ReportErrorf("%s", err.Error())
|
||||||
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
|
|
||||||
return 0
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// emitStartupBanner writes the startup banner to w unless args (the
|
// emitStartupBanner writes the startup banner to w unless args (the
|
||||||
|
|||||||
@@ -230,7 +230,7 @@ func TestEntryJSONStdoutIsExactlyOneDocument(t *testing.T) {
|
|||||||
programName, flagConfig, configPath, cmdSnapshot, cmdList, flagJSON,
|
programName, flagConfig, configPath, cmdSnapshot, cmdList, flagJSON,
|
||||||
}
|
}
|
||||||
|
|
||||||
stdout := captureProcessStdout(t, func() { _ = Entry() })
|
stdout := captureProcessStdout(t, Entry)
|
||||||
|
|
||||||
requireExactlyOneJSONDocument(t, stdout)
|
requireExactlyOneJSONDocument(t, stdout)
|
||||||
|
|
||||||
|
|||||||
@@ -1,140 +0,0 @@
|
|||||||
package cli //nolint:testpackage // shares the prune fixtures and capture helpers
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
// staleRecordLogMessage is the local-cleanup audit line CleanupLocalSnapshots
|
|
||||||
// logs for each stale record. It is exactly the signal issue #112 says a
|
|
||||||
// machine consumer lost under --json: gated off stdout, and pinned below
|
|
||||||
// the log level on stderr because --json used to force Quiet.
|
|
||||||
const staleRecordLogMessage = "Removing stale local snapshot record"
|
|
||||||
|
|
||||||
// TestEntryPruneJSONStderrHonoursVerbosity is the end-to-end regression
|
|
||||||
// guard for issue #112. Under --json the log level must still follow
|
|
||||||
// --verbose/--debug rather than being pinned to WARN, so the
|
|
||||||
// local-cleanup records reach stderr under --verbose while stdout stays
|
|
||||||
// exactly one JSON document; without --verbose they stay below the
|
|
||||||
// level, as they do without --json.
|
|
||||||
//
|
|
||||||
// Both halves are asserted together on the same run, because the fix has
|
|
||||||
// to keep the document clean (issue #108) while freeing stderr.
|
|
||||||
//
|
|
||||||
// Not parallel: it replaces os.Args, os.Stdout, os.Stderr and the xdg
|
|
||||||
// globals.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // replaces os.Args, os.Stdout, os.Stderr and the xdg globals
|
|
||||||
func TestEntryPruneJSONStderrHonoursVerbosity(t *testing.T) {
|
|
||||||
for _, testCase := range []struct {
|
|
||||||
name string
|
|
||||||
verbose bool
|
|
||||||
wantOnStderr bool
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "verbose json surfaces the cleanup record on stderr",
|
|
||||||
verbose: true,
|
|
||||||
wantOnStderr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "json alone keeps the cleanup record below the level",
|
|
||||||
verbose: false,
|
|
||||||
wantOnStderr: false,
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(testCase.name, func(t *testing.T) {
|
|
||||||
configPath := writeHermeticPruneConfig(t, true)
|
|
||||||
|
|
||||||
previousArgs := os.Args
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
os.Args = previousArgs
|
|
||||||
rootFlags = RootFlags{}
|
|
||||||
})
|
|
||||||
|
|
||||||
args := []string{
|
|
||||||
programName, flagConfig, configPath, cmdPrune, flagJSON,
|
|
||||||
}
|
|
||||||
if testCase.verbose {
|
|
||||||
args = append(args, "--verbose")
|
|
||||||
}
|
|
||||||
|
|
||||||
os.Args = args
|
|
||||||
|
|
||||||
stdout, stderr := captureProcessStdoutAndStderr(t,
|
|
||||||
func() { _ = Entry() })
|
|
||||||
|
|
||||||
// The document stays clean in both cases: freeing stderr must
|
|
||||||
// not regress issue #108.
|
|
||||||
requireExactlyOneJSONDocument(t, stdout)
|
|
||||||
|
|
||||||
if testCase.wantOnStderr {
|
|
||||||
assert.Contains(t, stderr, staleRecordLogMessage,
|
|
||||||
"--verbose --json must emit the cleanup record on stderr")
|
|
||||||
assert.Contains(t, stderr, stalePruneSnapshotID,
|
|
||||||
"the record must name the snapshot it removed")
|
|
||||||
} else {
|
|
||||||
assert.NotContains(t, stderr, staleRecordLogMessage,
|
|
||||||
"without --verbose the record stays below the log level")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// captureProcessStdoutAndStderr redirects both of the process's own
|
|
||||||
// standard streams to pipes for the duration of fn and returns what was
|
|
||||||
// written to each. The redirection is at the file-descriptor level
|
|
||||||
// because the logger binds os.Stderr when it initializes inside fn, and
|
|
||||||
// the JSON document reaches os.Stdout independently; the point is to see
|
|
||||||
// where each actually lands.
|
|
||||||
//
|
|
||||||
// Not parallel-safe: os.Stdout and os.Stderr are process-global.
|
|
||||||
func captureProcessStdoutAndStderr(t *testing.T, fn func()) (string, string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
outReader, outWriter, err := os.Pipe()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
errReader, errWriter, err := os.Pipe()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
previousOut, previousErr := os.Stdout, os.Stderr
|
|
||||||
os.Stdout, os.Stderr = outWriter, errWriter
|
|
||||||
|
|
||||||
capturedOut := drain(outReader)
|
|
||||||
capturedErr := drain(errReader)
|
|
||||||
|
|
||||||
fn()
|
|
||||||
|
|
||||||
os.Stdout, os.Stderr = previousOut, previousErr
|
|
||||||
|
|
||||||
require.NoError(t, outWriter.Close())
|
|
||||||
require.NoError(t, errWriter.Close())
|
|
||||||
|
|
||||||
out, errOut := <-capturedOut, <-capturedErr
|
|
||||||
|
|
||||||
require.NoError(t, outReader.Close())
|
|
||||||
require.NoError(t, errReader.Close())
|
|
||||||
|
|
||||||
return out, errOut
|
|
||||||
}
|
|
||||||
|
|
||||||
// drain copies a reader to a string on a goroutine and delivers the
|
|
||||||
// result once the writer end is closed.
|
|
||||||
func drain(reader io.Reader) <-chan string {
|
|
||||||
captured := make(chan string, 1)
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
var buf bytes.Buffer
|
|
||||||
|
|
||||||
_, _ = io.Copy(&buf, reader)
|
|
||||||
captured <- buf.String()
|
|
||||||
}()
|
|
||||||
|
|
||||||
return captured
|
|
||||||
}
|
|
||||||
@@ -81,7 +81,7 @@ func TestEntryPruneJSONStdoutIsExactlyOneDocument(t *testing.T) {
|
|||||||
programName, flagConfig, configPath, cmdPrune, flagJSON,
|
programName, flagConfig, configPath, cmdPrune, flagJSON,
|
||||||
}
|
}
|
||||||
|
|
||||||
stdout := captureProcessStdout(t, func() { _ = Entry() })
|
stdout := captureProcessStdout(t, Entry)
|
||||||
|
|
||||||
requireExactlyOneJSONDocument(t, stdout)
|
requireExactlyOneJSONDocument(t, stdout)
|
||||||
|
|
||||||
|
|||||||
@@ -1,58 +0,0 @@
|
|||||||
package cli //nolint:testpackage // shares programName and the capture helpers
|
|
||||||
|
|
||||||
import (
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestEntryReturnsStatusCode pins the contract main() relies on for
|
|
||||||
// issue #75: Entry reports success or failure through its return value
|
|
||||||
// and never calls os.Exit. An os.Exit from inside Entry would skip
|
|
||||||
// main's deferred profile writers and truncate the profile of a failing
|
|
||||||
// command. main turns this code into os.Exit only after those defers
|
|
||||||
// run, so a failing command must come back with a non-zero code rather
|
|
||||||
// than ending the process here.
|
|
||||||
//
|
|
||||||
// Stdout is captured only to keep the banner and command output off the
|
|
||||||
// test log; the assertion is on the returned code.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // replaces os.Args and rootFlags
|
|
||||||
func TestEntryReturnsStatusCode(t *testing.T) {
|
|
||||||
for _, testCase := range []struct {
|
|
||||||
name string
|
|
||||||
args []string
|
|
||||||
want int
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
// version is self-contained: it needs no config and no
|
|
||||||
// destination store, so it exercises the success path.
|
|
||||||
name: "successful command returns zero",
|
|
||||||
args: []string{programName, "version"},
|
|
||||||
want: 0,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "unknown command returns one",
|
|
||||||
args: []string{programName, "no-such-command"},
|
|
||||||
want: 1,
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(testCase.name, func(t *testing.T) {
|
|
||||||
previousArgs := os.Args
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
os.Args = previousArgs
|
|
||||||
rootFlags = RootFlags{}
|
|
||||||
})
|
|
||||||
|
|
||||||
os.Args = testCase.args
|
|
||||||
|
|
||||||
var code int
|
|
||||||
|
|
||||||
_ = captureProcessStdout(t, func() { code = Entry() })
|
|
||||||
|
|
||||||
assert.Equal(t, testCase.want, code)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+35
-5
@@ -1,7 +1,12 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
"go.uber.org/fx"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
"sneak.berlin/go/vaultik/internal/vaultik"
|
||||||
)
|
)
|
||||||
@@ -28,19 +33,44 @@ func NewInfoCommand() *cobra.Command {
|
|||||||
// Use the app framework
|
// Use the app framework
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
return RunOperation(cmd.Context(), AppOptions{
|
return RunWithApp(cmd.Context(), AppOptions{
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogOptions: log.Options{
|
LogOptions: log.Options{
|
||||||
Verbose: rootFlags.Verbose,
|
Verbose: rootFlags.Verbose,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet,
|
||||||
},
|
},
|
||||||
Mode: readOnly,
|
Modules: []fx.Option{},
|
||||||
}, func(v *vaultik.Vaultik) error {
|
Invokes: []fx.Option{
|
||||||
return v.ShowInfo()
|
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
||||||
}, func(err error) {
|
lc.Append(fx.Hook{
|
||||||
|
OnStart: func(_ context.Context) error {
|
||||||
|
go func() {
|
||||||
|
err := v.ShowInfo()
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, context.Canceled) {
|
||||||
log.Error("Failed to show info", "error", err)
|
log.Error("Failed to show info", "error", err)
|
||||||
ReportErrorf("Failed to show info: %v", err)
|
ReportErrorf("Failed to show info: %v", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = v.Shutdowner.Shutdown()
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to shutdown", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
OnStop: func(_ context.Context) error {
|
||||||
|
v.Cancel()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}),
|
||||||
|
},
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
+43
-11
@@ -1,7 +1,12 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
"go.uber.org/fx"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
"sneak.berlin/go/vaultik/internal/vaultik"
|
||||||
)
|
)
|
||||||
@@ -36,24 +41,51 @@ work (e.g. after a crashed backup or to reclaim storage).`,
|
|||||||
// Use the app framework like other commands
|
// Use the app framework like other commands
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
return RunOperation(cmd.Context(), AppOptions{
|
return RunWithApp(cmd.Context(), AppOptions{
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogOptions: log.Options{
|
LogOptions: log.Options{
|
||||||
Verbose: rootFlags.Verbose,
|
Verbose: rootFlags.Verbose,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet || opts.JSON,
|
||||||
JSON: opts.JSON,
|
|
||||||
},
|
},
|
||||||
Mode: mutating,
|
Modules: []fx.Option{},
|
||||||
}, func(v *vaultik.Vaultik) error {
|
Invokes: []fx.Option{
|
||||||
return v.Prune(opts)
|
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
||||||
}, func(err error) {
|
lc.Append(fx.Hook{
|
||||||
if opts.JSON {
|
OnStart: func(_ context.Context) error {
|
||||||
return
|
// Start the prune operation in a goroutine
|
||||||
}
|
go func() {
|
||||||
|
// Run the prune operation
|
||||||
|
err := v.Prune(opts)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, context.Canceled) {
|
||||||
|
if !opts.JSON {
|
||||||
log.Error("Prune operation failed", "error", err)
|
log.Error("Prune operation failed", "error", err)
|
||||||
ReportErrorf("Prune failed: %v", err)
|
ReportErrorf("Prune failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Shutdown the app when prune completes
|
||||||
|
err = v.Shutdowner.Shutdown()
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to shutdown", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
OnStop: func(_ context.Context) error {
|
||||||
|
log.Debug("Stopping prune operation")
|
||||||
|
v.Cancel()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}),
|
||||||
|
},
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
+38
-12
@@ -1,9 +1,12 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
"os"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
"go.uber.org/fx"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
"sneak.berlin/go/vaultik/internal/vaultik"
|
||||||
)
|
)
|
||||||
@@ -45,7 +48,7 @@ This is destructive and irreversible. Requires --force.`,
|
|||||||
return errNukeNeedsForce
|
return errNukeNeedsForce
|
||||||
}
|
}
|
||||||
|
|
||||||
return runVaultikApp(cmd, mutating, false, false, "Remote nuke failed",
|
return runVaultikApp(cmd, false, false, "Remote nuke failed",
|
||||||
func(v *vaultik.Vaultik) error {
|
func(v *vaultik.Vaultik) error {
|
||||||
return v.NukeRemote(true)
|
return v.NukeRemote(true)
|
||||||
})
|
})
|
||||||
@@ -80,24 +83,47 @@ func newRemoteInfoCommand() *cobra.Command {
|
|||||||
|
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
return RunOperation(cmd.Context(), AppOptions{
|
return RunWithApp(cmd.Context(), AppOptions{
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogOptions: log.Options{
|
LogOptions: log.Options{
|
||||||
Verbose: rootFlags.Verbose,
|
Verbose: rootFlags.Verbose,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet || jsonOutput,
|
||||||
JSON: jsonOutput,
|
|
||||||
},
|
},
|
||||||
Mode: readOnly,
|
Modules: []fx.Option{},
|
||||||
}, func(v *vaultik.Vaultik) error {
|
Invokes: []fx.Option{
|
||||||
return v.RemoteInfo(jsonOutput)
|
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
||||||
}, func(err error) {
|
lc.Append(fx.Hook{
|
||||||
if jsonOutput {
|
OnStart: func(_ context.Context) error {
|
||||||
return
|
go func() {
|
||||||
}
|
err := v.RemoteInfo(jsonOutput)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, context.Canceled) {
|
||||||
|
if !jsonOutput {
|
||||||
log.Error("Failed to get remote info", "error", err)
|
log.Error("Failed to get remote info", "error", err)
|
||||||
ReportErrorf("Failed to get remote info: %v", err)
|
ReportErrorf("Failed to get remote info: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = v.Shutdowner.Shutdown()
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to shutdown", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
OnStop: func(_ context.Context) error {
|
||||||
|
v.Cancel()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}),
|
||||||
|
},
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -57,9 +57,8 @@ on the source system.`,
|
|||||||
cmd.PersistentFlags().BoolVarP(&rootFlags.Quiet, "quiet", "q", false,
|
cmd.PersistentFlags().BoolVarP(&rootFlags.Quiet, "quiet", "q", false,
|
||||||
"Suppress non-error output")
|
"Suppress non-error output")
|
||||||
cmd.PersistentFlags().BoolVar(&rootFlags.SkipErrors, "skip-errors", false,
|
cmd.PersistentFlags().BoolVar(&rootFlags.SkipErrors, "skip-errors", false,
|
||||||
"Skip files that cannot be read when creating a snapshot, or "+
|
"Continue past per-file errors instead of aborting "+
|
||||||
"that cannot be restored when restoring, instead of aborting "+
|
"(applies to snapshot create and restore)")
|
||||||
"(packing and storage errors still abort)")
|
|
||||||
|
|
||||||
// Add subcommands
|
// Add subcommands
|
||||||
cmd.AddCommand(
|
cmd.AddCommand(
|
||||||
|
|||||||
@@ -1,97 +0,0 @@
|
|||||||
package cli_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"go.uber.org/fx"
|
|
||||||
"sneak.berlin/go/vaultik/internal/cli"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestRunAppWaitsForOperationCleanupOnShutdown drives RunApp with an fx app
|
|
||||||
// wired the way RunOperation wires a command: a single lifecycle hook whose
|
|
||||||
// OnStart launches the operation in its own goroutine and whose OnStop cancels
|
|
||||||
// it and blocks until that goroutine returns. The operation stands in for a
|
|
||||||
// restore blocked mid-download — it holds a decrypted "scratch" file and only
|
|
||||||
// removes it as it unwinds on cancellation.
|
|
||||||
//
|
|
||||||
// The app is asked to stop once the operation is running (standing in for an
|
|
||||||
// OS interrupt; fx delivers a real signal and Shutdowner.Shutdown() on the
|
|
||||||
// same app.Wait channel, so both drive the identical shutdown path). RunApp
|
|
||||||
// must not return until app.Stop has run the OnStop hook, so the scratch file
|
|
||||||
// must be gone by the time RunApp returns. Before the fix RunApp returned as
|
|
||||||
// soon as the app.Wait/Done channel fired, without running app.Stop, so the
|
|
||||||
// cleanup never ran and this file would still be on disk (issue #159).
|
|
||||||
func TestRunAppWaitsForOperationCleanupOnShutdown(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
scratch := filepath.Join(t.TempDir(), "decrypted-scratch")
|
|
||||||
require.NoError(t, os.WriteFile(scratch, []byte("secret"), 0o600))
|
|
||||||
|
|
||||||
// Cancel and reap the operation even if RunApp returns without doing so
|
|
||||||
// (the buggy path), so the goroutine cannot leak past the test.
|
|
||||||
opCtx, opCancel := context.WithCancel(context.Background())
|
|
||||||
t.Cleanup(opCancel)
|
|
||||||
|
|
||||||
var stop func(context.Context) bool
|
|
||||||
|
|
||||||
app := fx.New(
|
|
||||||
fx.NopLogger,
|
|
||||||
fx.Invoke(func(lc fx.Lifecycle, sh fx.Shutdowner) {
|
|
||||||
lc.Append(fx.Hook{
|
|
||||||
OnStart: func(_ context.Context) error {
|
|
||||||
done := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
defer close(done)
|
|
||||||
|
|
||||||
// Blocked mid-operation until cancelled, then run the
|
|
||||||
// cleanup an interrupted restore would run.
|
|
||||||
<-opCtx.Done()
|
|
||||||
|
|
||||||
_ = os.Remove(scratch)
|
|
||||||
}()
|
|
||||||
|
|
||||||
stop = func(ctx context.Context) bool {
|
|
||||||
opCancel()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
return true
|
|
||||||
case <-ctx.Done():
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Ask the app to stop now that the operation is running.
|
|
||||||
go func() { _ = sh.Shutdown() }()
|
|
||||||
|
|
||||||
return nil
|
|
||||||
},
|
|
||||||
OnStop: func(ctx context.Context) error {
|
|
||||||
stop(ctx)
|
|
||||||
|
|
||||||
return nil
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}),
|
|
||||||
)
|
|
||||||
|
|
||||||
done := make(chan error, 1)
|
|
||||||
go func() { done <- cli.RunApp(context.Background(), app) }()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case err := <-done:
|
|
||||||
require.NoError(t, err)
|
|
||||||
case <-time.After(30 * time.Second):
|
|
||||||
t.Fatal("RunApp did not return after shutdown was requested")
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err := os.Stat(scratch)
|
|
||||||
require.True(t, os.IsNotExist(err),
|
|
||||||
"RunApp returned before the operation removed its decrypted scratch file")
|
|
||||||
}
|
|
||||||
+77
-25
@@ -1,10 +1,13 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"os"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
"go.uber.org/fx"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
"sneak.berlin/go/vaultik/internal/vaultik"
|
||||||
)
|
)
|
||||||
@@ -83,8 +86,7 @@ specifying a path using --config or by setting VAULTIK_CONFIG to a path.`,
|
|||||||
// Use the backup functionality from cli package
|
// Use the backup functionality from cli package
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
// --cron suppression is wired through v.UI by setupGlobals.
|
return RunWithApp(cmd.Context(), AppOptions{
|
||||||
return RunOperation(cmd.Context(), AppOptions{
|
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogOptions: log.Options{
|
LogOptions: log.Options{
|
||||||
Verbose: rootFlags.Verbose,
|
Verbose: rootFlags.Verbose,
|
||||||
@@ -92,12 +94,42 @@ specifying a path using --config or by setting VAULTIK_CONFIG to a path.`,
|
|||||||
Cron: opts.Cron,
|
Cron: opts.Cron,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet,
|
||||||
},
|
},
|
||||||
Mode: mutating,
|
Modules: []fx.Option{},
|
||||||
}, func(v *vaultik.Vaultik) error {
|
Invokes: []fx.Option{
|
||||||
return v.CreateSnapshot(opts)
|
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
||||||
}, func(err error) {
|
lc.Append(fx.Hook{
|
||||||
|
OnStart: func(_ context.Context) error {
|
||||||
|
// Start the snapshot creation in a goroutine
|
||||||
|
go func() {
|
||||||
|
// --cron suppression is wired through v.UI by setupGlobals.
|
||||||
|
err := v.CreateSnapshot(opts)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, context.Canceled) {
|
||||||
log.Error("Snapshot creation failed", "error", err)
|
log.Error("Snapshot creation failed", "error", err)
|
||||||
ReportErrorf("Snapshot creation failed: %v", err)
|
ReportErrorf("Snapshot creation failed: %v", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Shutdown the app when snapshot completes
|
||||||
|
err = v.Shutdowner.Shutdown()
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to shutdown", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
OnStop: func(_ context.Context) error {
|
||||||
|
log.Debug("Stopping snapshot creation")
|
||||||
|
// Cancel the Vaultik context
|
||||||
|
v.Cancel()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}),
|
||||||
|
},
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -126,7 +158,7 @@ func newSnapshotListCommand() *cobra.Command {
|
|||||||
Long: "Lists all snapshots with their ID, timestamp, and compressed size",
|
Long: "Lists all snapshots with their ID, timestamp, and compressed size",
|
||||||
Args: cobra.NoArgs,
|
Args: cobra.NoArgs,
|
||||||
RunE: func(cmd *cobra.Command, _ []string) error {
|
RunE: func(cmd *cobra.Command, _ []string) error {
|
||||||
return runVaultikApp(cmd, readOnly, false, false,
|
return runVaultikApp(cmd, false, false,
|
||||||
"Failed to list snapshots",
|
"Failed to list snapshots",
|
||||||
func(v *vaultik.Vaultik) error {
|
func(v *vaultik.Vaultik) error {
|
||||||
return v.ListSnapshots(jsonOutput)
|
return v.ListSnapshots(jsonOutput)
|
||||||
@@ -162,7 +194,7 @@ restrict the operation to specific snapshot names.`,
|
|||||||
return errPurgeCriteriaBoth
|
return errPurgeCriteriaBoth
|
||||||
}
|
}
|
||||||
|
|
||||||
return runVaultikApp(cmd, mutating, false, false,
|
return runVaultikApp(cmd, false, false,
|
||||||
"Failed to purge snapshots",
|
"Failed to purge snapshots",
|
||||||
func(v *vaultik.Vaultik) error {
|
func(v *vaultik.Vaultik) error {
|
||||||
return v.PurgeSnapshotsWithOptions(opts)
|
return v.PurgeSnapshotsWithOptions(opts)
|
||||||
@@ -188,11 +220,8 @@ func newSnapshotVerifyCommand() *cobra.Command {
|
|||||||
|
|
||||||
cmd := &cobra.Command{
|
cmd := &cobra.Command{
|
||||||
Use: "verify <snapshot-id>",
|
Use: "verify <snapshot-id>",
|
||||||
Short: "Check a snapshot's blobs are present with the listed size",
|
Short: "Verify snapshot integrity",
|
||||||
Long: "Checks that every blob the snapshot's manifest lists is present\n" +
|
Long: "Verifies that all blobs referenced in a snapshot exist.\n\n" +
|
||||||
"in storage with the size the manifest records, and that the\n" +
|
|
||||||
"snapshot's encrypted database is present. It does not read blob\n" +
|
|
||||||
"contents; use --deep to download and cryptographically verify them.\n\n" +
|
|
||||||
"The snapshot may be named by its ID or, on a host with no local\n" +
|
"The snapshot may be named by its ID or, on a host with no local\n" +
|
||||||
"index, by the remote key that 'snapshot list' prints for a\n" +
|
"index, by the remote key that 'snapshot list' prints for a\n" +
|
||||||
"remote-only snapshot (an unambiguous leading part is enough).",
|
"remote-only snapshot (an unambiguous leading part is enough).",
|
||||||
@@ -208,24 +237,47 @@ func newSnapshotVerifyCommand() *cobra.Command {
|
|||||||
|
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
return RunOperation(cmd.Context(), AppOptions{
|
return RunWithApp(cmd.Context(), AppOptions{
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogOptions: log.Options{
|
LogOptions: log.Options{
|
||||||
Verbose: rootFlags.Verbose,
|
Verbose: rootFlags.Verbose,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet || opts.JSON,
|
||||||
JSON: opts.JSON,
|
|
||||||
},
|
},
|
||||||
Mode: readOnly,
|
Modules: []fx.Option{},
|
||||||
}, func(v *vaultik.Vaultik) error {
|
Invokes: []fx.Option{
|
||||||
return v.VerifySnapshotWithOptions(snapshotID, opts)
|
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
||||||
}, func(err error) {
|
lc.Append(fx.Hook{
|
||||||
if opts.JSON {
|
OnStart: func(_ context.Context) error {
|
||||||
return
|
go func() {
|
||||||
}
|
err := v.VerifySnapshotWithOptions(snapshotID, opts)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, context.Canceled) {
|
||||||
|
if !opts.JSON {
|
||||||
log.Error("Verification failed", "error", err)
|
log.Error("Verification failed", "error", err)
|
||||||
ReportErrorf("Verification failed: %v", err)
|
ReportErrorf("Verification failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err = v.Shutdowner.Shutdown()
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to shutdown", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
OnStop: func(_ context.Context) error {
|
||||||
|
v.Cancel()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}),
|
||||||
|
},
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -264,7 +316,7 @@ To wipe the entire destination store and start over, use 'vaultik remote
|
|||||||
nuke --force' — it is the single supported entry point for that.`,
|
nuke --force' — it is the single supported entry point for that.`,
|
||||||
Args: requireSnapshotIDArg,
|
Args: requireSnapshotIDArg,
|
||||||
RunE: func(cmd *cobra.Command, args []string) error {
|
RunE: func(cmd *cobra.Command, args []string) error {
|
||||||
return runVaultikApp(cmd, mutating, opts.JSON, opts.JSON,
|
return runVaultikApp(cmd, opts.JSON, opts.JSON,
|
||||||
"Failed to remove snapshot",
|
"Failed to remove snapshot",
|
||||||
func(v *vaultik.Vaultik) error {
|
func(v *vaultik.Vaultik) error {
|
||||||
_, err := v.RemoveSnapshot(args[0], opts)
|
_, err := v.RemoveSnapshot(args[0], opts)
|
||||||
|
|||||||
@@ -1,8 +1,16 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
"go.uber.org/fx"
|
||||||
|
"sneak.berlin/go/vaultik/internal/config"
|
||||||
|
"sneak.berlin/go/vaultik/internal/globals"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
|
"sneak.berlin/go/vaultik/internal/storage"
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
"sneak.berlin/go/vaultik/internal/vaultik"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -17,6 +25,15 @@ type RestoreOptions struct {
|
|||||||
Verify bool // Verify restored files after restore
|
Verify bool // Verify restored files after restore
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RestoreApp contains all dependencies needed for restore
|
||||||
|
type RestoreApp struct {
|
||||||
|
Globals *globals.Globals
|
||||||
|
Config *config.Config
|
||||||
|
Storage storage.Storer
|
||||||
|
Vaultik *vaultik.Vaultik
|
||||||
|
Shutdowner fx.Shutdowner
|
||||||
|
}
|
||||||
|
|
||||||
// newSnapshotRestoreCommand creates the 'snapshot restore' subcommand
|
// newSnapshotRestoreCommand creates the 'snapshot restore' subcommand
|
||||||
func newSnapshotRestoreCommand() *cobra.Command {
|
func newSnapshotRestoreCommand() *cobra.Command {
|
||||||
opts := &RestoreOptions{}
|
opts := &RestoreOptions{}
|
||||||
@@ -35,12 +52,8 @@ The snapshot may be named by its ID or, when restoring on a host with no
|
|||||||
local index, by the remote key that 'snapshot list' prints for a
|
local index, by the remote key that 'snapshot list' prints for a
|
||||||
remote-only snapshot (an unambiguous leading part is enough).
|
remote-only snapshot (an unambiguous leading part is enough).
|
||||||
|
|
||||||
Requires the age private key in the VAULTIK_AGE_SECRET_KEY environment
|
Requires the VAULTIK_AGE_SECRET_KEY environment variable to be set with
|
||||||
variable. The variable may hold the whole age-keygen file (comments and
|
the age private key.
|
||||||
all of its identities are accepted); read it from the file rather than
|
|
||||||
typing the key, so it does not land in your shell history:
|
|
||||||
|
|
||||||
export VAULTIK_AGE_SECRET_KEY="$(cat vaultik_backup_private_key.txt)"
|
|
||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Restore entire snapshot
|
# Restore entire snapshot
|
||||||
@@ -68,8 +81,7 @@ Examples:
|
|||||||
return cmd
|
return cmd
|
||||||
}
|
}
|
||||||
|
|
||||||
// runRestore parses arguments and runs the restore operation through the
|
// runRestore parses arguments and runs the restore operation through the app framework
|
||||||
// app framework.
|
|
||||||
func runRestore(cmd *cobra.Command, args []string, opts *RestoreOptions) error {
|
func runRestore(cmd *cobra.Command, args []string, opts *RestoreOptions) error {
|
||||||
snapshotID := args[0]
|
snapshotID := args[0]
|
||||||
|
|
||||||
@@ -78,31 +90,87 @@ func runRestore(cmd *cobra.Command, args []string, opts *RestoreOptions) error {
|
|||||||
opts.Paths = args[restoreMinArgs:]
|
opts.Paths = args[restoreMinArgs:]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Use unified config resolution
|
||||||
configPath, err := ResolveConfigPath()
|
configPath, err := ResolveConfigPath()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Use the app framework like other commands
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
return RunOperation(cmd.Context(), AppOptions{
|
return RunWithApp(cmd.Context(), AppOptions{
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogOptions: log.Options{
|
LogOptions: log.Options{
|
||||||
Verbose: rootFlags.Verbose,
|
Verbose: rootFlags.Verbose,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet,
|
||||||
},
|
},
|
||||||
Mode: readOnly,
|
Modules: buildRestoreModules(),
|
||||||
}, func(v *vaultik.Vaultik) error {
|
Invokes: buildRestoreInvokes(snapshotID, opts),
|
||||||
return v.Restore(&vaultik.RestoreOptions{
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildRestoreModules returns the fx.Options for dependency injection in restore
|
||||||
|
func buildRestoreModules() []fx.Option {
|
||||||
|
return []fx.Option{
|
||||||
|
fx.Provide(fx.Annotate(
|
||||||
|
func(g *globals.Globals, cfg *config.Config,
|
||||||
|
storer storage.Storer, v *vaultik.Vaultik, shutdowner fx.Shutdowner) *RestoreApp {
|
||||||
|
return &RestoreApp{
|
||||||
|
Globals: g,
|
||||||
|
Config: cfg,
|
||||||
|
Storage: storer,
|
||||||
|
Vaultik: v,
|
||||||
|
Shutdowner: shutdowner,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildRestoreInvokes returns the fx.Options that wire up the restore lifecycle
|
||||||
|
func buildRestoreInvokes(snapshotID string, opts *RestoreOptions) []fx.Option {
|
||||||
|
return []fx.Option{
|
||||||
|
fx.Invoke(func(app *RestoreApp, lc fx.Lifecycle) {
|
||||||
|
lc.Append(fx.Hook{
|
||||||
|
OnStart: func(_ context.Context) error {
|
||||||
|
// Start the restore operation in a goroutine
|
||||||
|
go func() {
|
||||||
|
// Run the restore operation
|
||||||
|
restoreOpts := &vaultik.RestoreOptions{
|
||||||
SnapshotID: snapshotID,
|
SnapshotID: snapshotID,
|
||||||
TargetDir: opts.TargetDir,
|
TargetDir: opts.TargetDir,
|
||||||
Paths: opts.Paths,
|
Paths: opts.Paths,
|
||||||
Verify: opts.Verify,
|
Verify: opts.Verify,
|
||||||
SkipErrors: rootFlags.SkipErrors,
|
SkipErrors: GetRootFlags().SkipErrors,
|
||||||
})
|
}
|
||||||
}, func(err error) {
|
|
||||||
|
err := app.Vaultik.Restore(restoreOpts)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, context.Canceled) {
|
||||||
log.Error("Restore operation failed", "error", err)
|
log.Error("Restore operation failed", "error", err)
|
||||||
ReportErrorf("Restore failed: %v", err)
|
ReportErrorf("Restore failed: %v", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Shutdown the app when restore completes
|
||||||
|
err = app.Shutdowner.Shutdown()
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to shutdown", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
OnStop: func(_ context.Context) error {
|
||||||
|
log.Debug("Stopping restore operation")
|
||||||
|
app.Vaultik.Cancel()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
},
|
||||||
})
|
})
|
||||||
|
}),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,37 +0,0 @@
|
|||||||
package cli //nolint:testpackage // exercises the unexported command constructor
|
|
||||||
|
|
||||||
import (
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/spf13/pflag"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestRestoreCommandDoesNotTakeKeyAsArgument guards the fix for the age
|
|
||||||
// key being echoed on the command line: restore must take the key only
|
|
||||||
// from the environment, never as a flag value, and its help must show the
|
|
||||||
// file-based form rather than a literal key that would land in shell
|
|
||||||
// history.
|
|
||||||
func TestRestoreCommandDoesNotTakeKeyAsArgument(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
cmd := newSnapshotRestoreCommand()
|
|
||||||
|
|
||||||
cmd.Flags().VisitAll(func(f *pflag.Flag) {
|
|
||||||
lower := strings.ToLower(f.Name)
|
|
||||||
for _, banned := range []string{"key", "secret", "age", "identity"} {
|
|
||||||
if strings.Contains(lower, banned) {
|
|
||||||
t.Errorf("restore must not accept the key as a flag; found --%s", f.Name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
help := cmd.Long
|
|
||||||
if strings.Contains(help, "AGE-SECRET-KEY-") {
|
|
||||||
t.Error("restore help must not show a literal age private key to type")
|
|
||||||
}
|
|
||||||
|
|
||||||
if !strings.Contains(help, "$(cat ") {
|
|
||||||
t.Error("restore help should read the key from a file, e.g. $(cat ...)")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+26
-97
@@ -16,19 +16,11 @@ import (
|
|||||||
"github.com/adrg/xdg"
|
"github.com/adrg/xdg"
|
||||||
"go.uber.org/fx"
|
"go.uber.org/fx"
|
||||||
"gopkg.in/yaml.v3"
|
"gopkg.in/yaml.v3"
|
||||||
"sneak.berlin/go/vaultik/internal/chunker"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
)
|
)
|
||||||
|
|
||||||
const appName = "vaultik"
|
const appName = "vaultik"
|
||||||
|
|
||||||
// secretKeyPrefix marks an age secret (private) key. It is compared
|
|
||||||
// case-insensitively so a recipient entry that is actually a private key is
|
|
||||||
// caught and never passed to age or echoed back.
|
|
||||||
//
|
|
||||||
//nolint:gosec // G101: marker for detecting a pasted secret key, not a credential
|
|
||||||
const secretKeyPrefix = "AGE-SECRET-KEY-"
|
|
||||||
|
|
||||||
// Defaults and validation bounds for tunable settings.
|
// Defaults and validation bounds for tunable settings.
|
||||||
const (
|
const (
|
||||||
defaultBlobSizeLimit = Size(10 * 1024 * 1024 * 1024) // 10GB
|
defaultBlobSizeLimit = Size(10 * 1024 * 1024 * 1024) // 10GB
|
||||||
@@ -45,17 +37,11 @@ var (
|
|||||||
errNoConfigPath = errors.New("config path not provided")
|
errNoConfigPath = errors.New("config path not provided")
|
||||||
errNoAgeRecipients = errors.New(
|
errNoAgeRecipients = errors.New(
|
||||||
"at least one age_recipient is required (generate with: age-keygen)")
|
"at least one age_recipient is required (generate with: age-keygen)")
|
||||||
errRecipientIsSecretKey = errors.New(
|
|
||||||
"an age secret key was given where a public key (age1...) belongs")
|
|
||||||
errRecipientNotX25519 = errors.New(
|
|
||||||
"not a valid recipient; only X25519 age1... public keys are supported")
|
|
||||||
errNoSnapshots = errors.New(
|
errNoSnapshots = errors.New(
|
||||||
"at least one snapshot must be configured (see config.example.yml)")
|
"at least one snapshot must be configured (see config.example.yml)")
|
||||||
errSnapshotNoPaths = errors.New("snapshot must have at least one path")
|
errSnapshotNoPaths = errors.New("snapshot must have at least one path")
|
||||||
errChunkSizeTooSmall = errors.New("chunk_size must be at least 1MB")
|
errChunkSizeTooSmall = errors.New("chunk_size must be at least 1MB")
|
||||||
errBlobSizeTooSmall = errors.New(
|
errBlobSizeTooSmall = errors.New("blob_size_limit must be at least chunk_size")
|
||||||
"blob_size_limit must be at least the largest chunk the chunker can " +
|
|
||||||
"emit (chunk_size times the FastCDC size spread)")
|
|
||||||
errBadCompression = errors.New("compression_level must be between 1 and 19")
|
errBadCompression = errors.New("compression_level must be between 1 and 19")
|
||||||
errBadStorageScheme = errors.New(
|
errBadStorageScheme = errors.New(
|
||||||
"storage_url must start with s3://, file://, or rclone://")
|
"storage_url must start with s3://, file://, or rclone://")
|
||||||
@@ -135,26 +121,6 @@ func (c *Config) SnapshotNames() []string {
|
|||||||
return names
|
return names
|
||||||
}
|
}
|
||||||
|
|
||||||
// Names of the two places the age secret key can be configured, used by
|
|
||||||
// AgeSecretKeySourceName for error messages that must not echo the value.
|
|
||||||
//
|
|
||||||
//nolint:gosec // G101: these are the names of the config sources, not a key
|
|
||||||
const (
|
|
||||||
ageSecretKeySourceEnv = "VAULTIK_AGE_SECRET_KEY"
|
|
||||||
ageSecretKeySourceConfig = "age_secret_key"
|
|
||||||
)
|
|
||||||
|
|
||||||
// AgeSecretKeySourceName returns the human name of where AgeSecretKey was
|
|
||||||
// configured. A Config built directly (as in tests) has no recorded
|
|
||||||
// source, so it reports the config-file field name.
|
|
||||||
func (c *Config) AgeSecretKeySourceName() string {
|
|
||||||
if c.AgeSecretKeySource != "" {
|
|
||||||
return c.AgeSecretKeySource
|
|
||||||
}
|
|
||||||
|
|
||||||
return ageSecretKeySourceConfig
|
|
||||||
}
|
|
||||||
|
|
||||||
// Config represents the application configuration for Vaultik.
|
// Config represents the application configuration for Vaultik.
|
||||||
// It defines all settings for backup operations, including source directories,
|
// It defines all settings for backup operations, including source directories,
|
||||||
// encryption recipients, storage configuration, and performance tuning parameters.
|
// encryption recipients, storage configuration, and performance tuning parameters.
|
||||||
@@ -164,11 +130,6 @@ func (c *Config) AgeSecretKeySourceName() string {
|
|||||||
type Config struct {
|
type Config struct {
|
||||||
AgeRecipients []string `yaml:"age_recipients"`
|
AgeRecipients []string `yaml:"age_recipients"`
|
||||||
AgeSecretKey string `yaml:"age_secret_key"`
|
AgeSecretKey string `yaml:"age_secret_key"`
|
||||||
// AgeSecretKeySource names where AgeSecretKey was configured
|
|
||||||
// ("VAULTIK_AGE_SECRET_KEY" or "age_secret_key") so a later parse
|
|
||||||
// failure can name the source without echoing the secret value. It is
|
|
||||||
// set by Load and never read from or written to the config file.
|
|
||||||
AgeSecretKeySource string `yaml:"-"`
|
|
||||||
BlobSizeLimit Size `yaml:"blob_size_limit"`
|
BlobSizeLimit Size `yaml:"blob_size_limit"`
|
||||||
ChunkSize Size `yaml:"chunk_size"`
|
ChunkSize Size `yaml:"chunk_size"`
|
||||||
// Exclude holds global excludes applied to all snapshots.
|
// Exclude holds global excludes applied to all snapshots.
|
||||||
@@ -201,9 +162,7 @@ type S3Config struct {
|
|||||||
AccessKeyID string `yaml:"access_key_id"`
|
AccessKeyID string `yaml:"access_key_id"`
|
||||||
SecretAccessKey string `yaml:"secret_access_key"`
|
SecretAccessKey string `yaml:"secret_access_key"`
|
||||||
Region string `yaml:"region"`
|
Region string `yaml:"region"`
|
||||||
// UseSSL selects HTTPS for a scheme-less endpoint. Omitted (nil) means
|
UseSSL bool `yaml:"use_ssl"`
|
||||||
// the default, TLS; set it to false only to force plain HTTP.
|
|
||||||
UseSSL *bool `yaml:"use_ssl"`
|
|
||||||
PartSize Size `yaml:"part_size"`
|
PartSize Size `yaml:"part_size"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -279,7 +238,10 @@ func Load(path string) (*Config, error) {
|
|||||||
cfg.IndexPath = expandTilde(envIndexPath)
|
cfg.IndexPath = expandTilde(envIndexPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg.setAgeSecretKey()
|
// Check for environment variable override for AgeSecretKey
|
||||||
|
if envAgeSecretKey := os.Getenv("VAULTIK_AGE_SECRET_KEY"); envAgeSecretKey != "" {
|
||||||
|
cfg.AgeSecretKey = extractAgeSecretKey(envAgeSecretKey)
|
||||||
|
}
|
||||||
|
|
||||||
// Get hostname if not set
|
// Get hostname if not set
|
||||||
if cfg.Hostname == "" {
|
if cfg.Hostname == "" {
|
||||||
@@ -323,30 +285,18 @@ func Load(path string) (*Config, error) {
|
|||||||
|
|
||||||
// Validate checks if the configuration is valid and complete.
|
// Validate checks if the configuration is valid and complete.
|
||||||
// It ensures all required fields are present and have valid values:
|
// It ensures all required fields are present and have valid values:
|
||||||
// - At least one age recipient must be specified, and every recipient must
|
// - At least one age recipient must be specified
|
||||||
// parse as an X25519 age1... public key (so a bad entry fails at load, not
|
|
||||||
// mid-backup); errors name the position, never the value
|
|
||||||
// - At least one snapshot must be configured with at least one path
|
// - At least one snapshot must be configured with at least one path
|
||||||
// - Storage must be configured (either storage_url or s3.* fields)
|
// - Storage must be configured (either storage_url or s3.* fields)
|
||||||
// - Chunk size must be at least 1MB
|
// - Chunk size must be at least 1MB
|
||||||
// - Blob size limit must be at least the largest chunk the chunker can emit
|
// - Blob size limit must be at least the chunk size
|
||||||
// (chunk_size times chunker.ChunkSizeSpread), so a single-chunk blob never
|
|
||||||
// exceeds the configured limit
|
|
||||||
// - Compression level must be between 1 and 19
|
// - Compression level must be between 1 and 19
|
||||||
//
|
|
||||||
// Returns an error describing the first validation failure encountered.
|
// Returns an error describing the first validation failure encountered.
|
||||||
func (c *Config) Validate() error {
|
func (c *Config) Validate() error {
|
||||||
if len(c.AgeRecipients) == 0 {
|
if len(c.AgeRecipients) == 0 {
|
||||||
return errNoAgeRecipients
|
return errNoAgeRecipients
|
||||||
}
|
}
|
||||||
|
|
||||||
for i, recipient := range c.AgeRecipients {
|
|
||||||
err := validateAgeRecipient(recipient)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("age_recipients[%d]: %w", i, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(c.Snapshots) == 0 {
|
if len(c.Snapshots) == 0 {
|
||||||
return errNoSnapshots
|
return errNoSnapshots
|
||||||
}
|
}
|
||||||
@@ -367,13 +317,8 @@ func (c *Config) Validate() error {
|
|||||||
return errChunkSizeTooSmall
|
return errChunkSizeTooSmall
|
||||||
}
|
}
|
||||||
|
|
||||||
// The chunker can emit chunks up to chunk_size * ChunkSizeSpread, and the
|
if c.BlobSizeLimit.Int64() < c.ChunkSize.Int64() {
|
||||||
// packer places a single such chunk into an otherwise empty blob. A limit
|
return errBlobSizeTooSmall
|
||||||
// below that bound would let a blob exceed it, so reject it.
|
|
||||||
largestChunk := c.ChunkSize.Int64() * chunker.ChunkSizeSpread
|
|
||||||
if c.BlobSizeLimit.Int64() < largestChunk {
|
|
||||||
return fmt.Errorf("%w: need at least %d bytes",
|
|
||||||
errBlobSizeTooSmall, largestChunk)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if c.CompressionLevel < minCompressionLevel ||
|
if c.CompressionLevel < minCompressionLevel ||
|
||||||
@@ -384,38 +329,6 @@ func (c *Config) Validate() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// validateAgeRecipient parses one age_recipients entry with the age library
|
|
||||||
// and returns a value-free error on failure. A recipient string can be
|
|
||||||
// sensitive (an operator may paste a secret key by mistake), so neither the
|
|
||||||
// entry nor age's own error (which quotes its input) is ever included.
|
|
||||||
func validateAgeRecipient(recipient string) error {
|
|
||||||
if strings.HasPrefix(strings.ToUpper(recipient), secretKeyPrefix) {
|
|
||||||
return errRecipientIsSecretKey
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err := age.ParseX25519Recipient(recipient)
|
|
||||||
if err != nil {
|
|
||||||
return errRecipientNotX25519
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// setAgeSecretKey records the age secret key and where it came from. The
|
|
||||||
// value is stored raw and parsed only where decryption happens
|
|
||||||
// (internal/vaultik), so backup, list and prune keep working whatever the
|
|
||||||
// field holds. The environment variable overrides the config-file field.
|
|
||||||
func (c *Config) setAgeSecretKey() {
|
|
||||||
if c.AgeSecretKey != "" {
|
|
||||||
c.AgeSecretKeySource = ageSecretKeySourceConfig
|
|
||||||
}
|
|
||||||
|
|
||||||
if env := os.Getenv("VAULTIK_AGE_SECRET_KEY"); env != "" {
|
|
||||||
c.AgeSecretKey = env
|
|
||||||
c.AgeSecretKeySource = ageSecretKeySourceEnv
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// validateStorage validates storage configuration.
|
// validateStorage validates storage configuration.
|
||||||
// If StorageURL is set, it takes precedence. S3 URLs require credentials.
|
// If StorageURL is set, it takes precedence. S3 URLs require credentials.
|
||||||
// File URLs don't require any S3 configuration.
|
// File URLs don't require any S3 configuration.
|
||||||
@@ -472,6 +385,22 @@ func (c *Config) validateStorageURL() error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// extractAgeSecretKey extracts the AGE-SECRET-KEY from the input using
|
||||||
|
// the age library's parser, which handles comments and whitespace.
|
||||||
|
func extractAgeSecretKey(input string) string {
|
||||||
|
identities, err := age.ParseIdentities(strings.NewReader(input))
|
||||||
|
if err != nil || len(identities) == 0 {
|
||||||
|
// Fall back to trimmed input if parsing fails
|
||||||
|
return strings.TrimSpace(input)
|
||||||
|
}
|
||||||
|
// Return the string representation of the first identity
|
||||||
|
if id, ok := identities[0].(*age.X25519Identity); ok {
|
||||||
|
return id.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.TrimSpace(input)
|
||||||
|
}
|
||||||
|
|
||||||
// Module exports the config module for fx dependency injection.
|
// Module exports the config module for fx dependency injection.
|
||||||
// It provides the Config type to other modules in the application.
|
// It provides the Config type to other modules in the application.
|
||||||
//
|
//
|
||||||
|
|||||||
+38
-212
@@ -1,13 +1,9 @@
|
|||||||
package config //nolint:testpackage // exercises unexported source constants
|
package config //nolint:testpackage // exercises unexported extractAgeSecretKey
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/chunker"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -87,48 +83,6 @@ func TestConfigLoad(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestExampleConfigIsScrubbedAndLoads checks that the shipped
|
|
||||||
// config.example.yml carries only neutral placeholders (no real credentials,
|
|
||||||
// private addresses, or internal host names) and still parses.
|
|
||||||
func TestExampleConfigIsScrubbedAndLoads(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
examplePath := filepath.Join("..", "..", "config.example.yml")
|
|
||||||
|
|
||||||
cfg, err := Load(examplePath)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to load config.example.yml: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if cfg.StorageURL != "rclone://myremote/path/to/backups" {
|
|
||||||
t.Errorf("Expected neutral storage_url, got '%s'", cfg.StorageURL)
|
|
||||||
}
|
|
||||||
|
|
||||||
//nolint:gosec // G304: examplePath is a fixed in-repo path, not user input
|
|
||||||
raw, err := os.ReadFile(examplePath)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Failed to read config.example.yml: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
text := string(raw)
|
|
||||||
|
|
||||||
wantSubstrings := []string{
|
|
||||||
"YOUR_ACCESS_KEY",
|
|
||||||
"YOUR_SECRET_KEY",
|
|
||||||
"endpoint: https://",
|
|
||||||
}
|
|
||||||
for _, want := range wantSubstrings {
|
|
||||||
if !strings.Contains(text, want) {
|
|
||||||
t.Errorf("Expected config.example.yml to contain %q", want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// A raw "http://" scheme would mean a plaintext, likely private endpoint.
|
|
||||||
if strings.Contains(text, "http://") {
|
|
||||||
t.Error("config.example.yml should not contain an http:// endpoint")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestConfigFromEnv tests loading config path from environment variable
|
// TestConfigFromEnv tests loading config path from environment variable
|
||||||
func TestConfigFromEnv(t *testing.T) {
|
func TestConfigFromEnv(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
@@ -147,57 +101,53 @@ func TestConfigFromEnv(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestValidateBlobSizeLimit checks the blob_size_limit boundary: it must be at
|
// TestExtractAgeSecretKey tests extraction of AGE-SECRET-KEY from various inputs
|
||||||
// least the largest chunk the chunker can emit (chunk_size times
|
func TestExtractAgeSecretKey(t *testing.T) {
|
||||||
// chunker.ChunkSizeSpread), because the packer places a single such chunk into
|
|
||||||
// an otherwise empty blob. A limit between chunk_size and that bound is rejected.
|
|
||||||
func TestValidateBlobSizeLimit(t *testing.T) {
|
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
const chunkSize = Size(10 * 1024 * 1024) // 10MB
|
|
||||||
|
|
||||||
largestChunk := chunkSize.Int64() * chunker.ChunkSizeSpread
|
|
||||||
|
|
||||||
newConfig := func(blobLimit Size) *Config {
|
|
||||||
return &Config{
|
|
||||||
AgeRecipients: []string{testSneakAgePublicKey},
|
|
||||||
Snapshots: map[string]SnapshotConfig{"test": {Paths: []string{"/tmp/src"}}},
|
|
||||||
StorageURL: "file:///tmp/vaultik-test-store",
|
|
||||||
ChunkSize: chunkSize,
|
|
||||||
BlobSizeLimit: blobLimit,
|
|
||||||
CompressionLevel: 3,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
blobLimit Size
|
input string
|
||||||
wantErr bool
|
expected string
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "at chunk_size but below largest chunk is rejected",
|
name: "plain key",
|
||||||
blobLimit: chunkSize,
|
input: testIntegrationAgePrivateKey,
|
||||||
wantErr: true,
|
expected: testIntegrationAgePrivateKey,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "between chunk_size and largest chunk is rejected",
|
name: "key with trailing newline",
|
||||||
blobLimit: Size(chunkSize.Int64() * 2),
|
input: testIntegrationAgePrivateKey + "\n",
|
||||||
wantErr: true,
|
expected: testIntegrationAgePrivateKey,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "one byte below largest chunk is rejected",
|
name: "full age-keygen output",
|
||||||
blobLimit: Size(largestChunk - 1),
|
input: "# created: 2025-01-14T12:00:00Z\n" +
|
||||||
wantErr: true,
|
"# public key: " + testIntegrationAgePublicKey + "\n" +
|
||||||
|
testIntegrationAgePrivateKey + "\n",
|
||||||
|
expected: testIntegrationAgePrivateKey,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "exactly at largest chunk is accepted",
|
name: "age-keygen output with extra blank lines",
|
||||||
blobLimit: Size(largestChunk),
|
input: "# created: 2025-01-14T12:00:00Z\n" +
|
||||||
wantErr: false,
|
"# public key: " + testIntegrationAgePublicKey + "\n\n" +
|
||||||
|
testIntegrationAgePrivateKey + "\n\n",
|
||||||
|
expected: testIntegrationAgePrivateKey,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "above largest chunk is accepted",
|
name: "key with leading whitespace",
|
||||||
blobLimit: Size(largestChunk * 100),
|
input: " " + testIntegrationAgePrivateKey + " ",
|
||||||
wantErr: false,
|
expected: testIntegrationAgePrivateKey,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty input",
|
||||||
|
input: "",
|
||||||
|
expected: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "only comments",
|
||||||
|
input: "# this is a comment\n# another comment",
|
||||||
|
expected: "# this is a comment\n# another comment",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -205,134 +155,10 @@ func TestValidateBlobSizeLimit(t *testing.T) {
|
|||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
err := newConfig(tt.blobLimit).Validate()
|
result := extractAgeSecretKey(tt.input)
|
||||||
if tt.wantErr {
|
if result != tt.expected {
|
||||||
if !errors.Is(err, errBlobSizeTooSmall) {
|
t.Errorf("extractAgeSecretKey(%q) = %q, want %q",
|
||||||
t.Fatalf("Validate() error = %v, want errBlobSizeTooSmall", err)
|
tt.input, result, tt.expected)
|
||||||
}
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Validate() unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestValidateAgeRecipients checks that recipients are parsed at config load
|
|
||||||
// (a bad entry fails immediately, not mid-backup) and that no invalid entry —
|
|
||||||
// least of all a pasted secret key — is echoed in the error.
|
|
||||||
func TestValidateAgeRecipients(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
baseConfig := func(recipients []string) *Config {
|
|
||||||
return &Config{
|
|
||||||
AgeRecipients: recipients,
|
|
||||||
Snapshots: map[string]SnapshotConfig{"test": {Paths: []string{"/tmp/src"}}},
|
|
||||||
StorageURL: "file:///tmp/vaultik-test-store",
|
|
||||||
ChunkSize: Size(10 * 1024 * 1024),
|
|
||||||
BlobSizeLimit: Size(10 * 1024 * 1024 * 1024),
|
|
||||||
CompressionLevel: 3,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
recipients []string
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "config init placeholder is rejected",
|
|
||||||
recipients: []string{"age1REPLACE_WITH_YOUR_PUBLIC_KEY"},
|
|
||||||
wantErr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "ssh-ed25519 recipient is rejected",
|
|
||||||
recipients: []string{"ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIexamplekeydata"},
|
|
||||||
wantErr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "truncated age1 string is rejected",
|
|
||||||
recipients: []string{"age1short"},
|
|
||||||
wantErr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "secret key passed as recipient is rejected",
|
|
||||||
recipients: []string{testIntegrationAgePrivateKey},
|
|
||||||
wantErr: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "two valid recipients are accepted",
|
|
||||||
recipients: []string{testSneakAgePublicKey, testIntegrationAgePublicKey},
|
|
||||||
wantErr: false,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
err := baseConfig(tt.recipients).Validate()
|
|
||||||
if !tt.wantErr {
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Validate() unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("Validate() returned nil, want error")
|
|
||||||
}
|
|
||||||
|
|
||||||
// The entry itself must never appear in the error, since a
|
|
||||||
// recipient string can be a secret key.
|
|
||||||
for _, recipient := range tt.recipients {
|
|
||||||
if strings.Contains(err.Error(), recipient) {
|
|
||||||
t.Fatalf("Validate() error echoed the recipient value: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestAgeSecretKeySourceName checks the name reported for the configured
|
|
||||||
// age secret key: the recorded source when Load set one, and the
|
|
||||||
// config-file field name for a Config built directly (as in tests).
|
|
||||||
func TestAgeSecretKeySourceName(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
source string
|
|
||||||
want string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "unset defaults to config field",
|
|
||||||
source: "",
|
|
||||||
want: ageSecretKeySourceConfig,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "environment source",
|
|
||||||
source: ageSecretKeySourceEnv,
|
|
||||||
want: ageSecretKeySourceEnv,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "config-file source",
|
|
||||||
source: ageSecretKeySourceConfig,
|
|
||||||
want: ageSecretKeySourceConfig,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
cfg := &Config{AgeSecretKeySource: tt.source}
|
|
||||||
if got := cfg.AgeSecretKeySourceName(); got != tt.want {
|
|
||||||
t.Errorf("AgeSecretKeySourceName() = %q, want %q", got, tt.want)
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,224 @@
|
|||||||
|
// Package crypto provides thread-safe age encryption and decryption
|
||||||
|
// helpers used to protect blob and metadata content.
|
||||||
|
package crypto //nolint:revive,nolintlint // stdlib crypto unused; see #76
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"filippo.io/age"
|
||||||
|
"go.uber.org/fx"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ErrNoRecipients is returned when an encryptor is created or updated
|
||||||
|
// without any recipient public keys.
|
||||||
|
var ErrNoRecipients = errors.New("at least one recipient is required")
|
||||||
|
|
||||||
|
// Encryptor provides thread-safe encryption using the age encryption library.
|
||||||
|
// It supports encrypting data for multiple recipients simultaneously, allowing
|
||||||
|
// any of the corresponding private keys to decrypt the data. This is useful
|
||||||
|
// for backup scenarios where multiple parties should be able to decrypt the data.
|
||||||
|
type Encryptor struct {
|
||||||
|
recipients []age.Recipient
|
||||||
|
mu sync.RWMutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewEncryptor creates a new encryptor with the given age public keys.
|
||||||
|
// Each public key should be a valid age X25519 recipient string (e.g., "age1...")
|
||||||
|
// At least one recipient must be provided. Returns an error if any of the
|
||||||
|
// public keys are invalid or if no recipients are specified.
|
||||||
|
func NewEncryptor(publicKeys []string) (*Encryptor, error) {
|
||||||
|
if len(publicKeys) == 0 {
|
||||||
|
return nil, ErrNoRecipients
|
||||||
|
}
|
||||||
|
|
||||||
|
recipients := make([]age.Recipient, 0, len(publicKeys))
|
||||||
|
for _, key := range publicKeys {
|
||||||
|
recipient, err := age.ParseX25519Recipient(key)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("parsing age recipient %s: %w", key, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
recipients = append(recipients, recipient)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Encryptor{
|
||||||
|
recipients: recipients,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Encrypt encrypts data using age encryption for all configured recipients.
|
||||||
|
// The encrypted data can be decrypted by any of the corresponding private keys.
|
||||||
|
// This method is suitable for small to medium amounts of data that fit in memory.
|
||||||
|
// For large data streams, use EncryptStream or EncryptWriter instead.
|
||||||
|
func (e *Encryptor) Encrypt(data []byte) ([]byte, error) {
|
||||||
|
e.mu.RLock()
|
||||||
|
recipients := e.recipients
|
||||||
|
e.mu.RUnlock()
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
|
||||||
|
// Create encrypted writer for all recipients
|
||||||
|
w, err := age.Encrypt(&buf, recipients...)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("creating encrypted writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write data
|
||||||
|
_, err = w.Write(data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("writing encrypted data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close to flush
|
||||||
|
err = w.Close()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("closing encrypted writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return buf.Bytes(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// EncryptStream encrypts data from reader to writer using age encryption.
|
||||||
|
// This method is suitable for encrypting large files or streams as it processes
|
||||||
|
// data in a streaming fashion without loading everything into memory.
|
||||||
|
// The encrypted data is written directly to the destination writer.
|
||||||
|
func (e *Encryptor) EncryptStream(dst io.Writer, src io.Reader) error {
|
||||||
|
e.mu.RLock()
|
||||||
|
recipients := e.recipients
|
||||||
|
e.mu.RUnlock()
|
||||||
|
|
||||||
|
// Create encrypted writer for all recipients
|
||||||
|
w, err := age.Encrypt(dst, recipients...)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("creating encrypted writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Copy data
|
||||||
|
_, err = io.Copy(w, src)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("copying encrypted data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close to flush
|
||||||
|
err = w.Close()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("closing encrypted writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// EncryptWriter creates a writer that encrypts data written to it.
|
||||||
|
// All data written to the returned WriteCloser will be encrypted and written
|
||||||
|
// to the destination writer. The caller must call Close() on the returned
|
||||||
|
// writer to ensure all encrypted data is properly flushed and finalized.
|
||||||
|
// This is useful for integrating encryption into existing writer-based pipelines.
|
||||||
|
func (e *Encryptor) EncryptWriter(dst io.Writer) (io.WriteCloser, error) {
|
||||||
|
e.mu.RLock()
|
||||||
|
recipients := e.recipients
|
||||||
|
e.mu.RUnlock()
|
||||||
|
|
||||||
|
// Create encrypted writer for all recipients
|
||||||
|
w, err := age.Encrypt(dst, recipients...)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("creating encrypted writer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return w, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateRecipients updates the recipients for future encryption operations.
|
||||||
|
// This method is thread-safe and can be called while other encryption operations
|
||||||
|
// are in progress. Existing encryption operations will continue with the old
|
||||||
|
// recipients. At least one recipient must be provided. Returns an error if any
|
||||||
|
// of the public keys are invalid or if no recipients are specified.
|
||||||
|
func (e *Encryptor) UpdateRecipients(publicKeys []string) error {
|
||||||
|
if len(publicKeys) == 0 {
|
||||||
|
return ErrNoRecipients
|
||||||
|
}
|
||||||
|
|
||||||
|
recipients := make([]age.Recipient, 0, len(publicKeys))
|
||||||
|
for _, key := range publicKeys {
|
||||||
|
recipient, err := age.ParseX25519Recipient(key)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("parsing age recipient %s: %w", key, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
recipients = append(recipients, recipient)
|
||||||
|
}
|
||||||
|
|
||||||
|
e.mu.Lock()
|
||||||
|
e.recipients = recipients
|
||||||
|
e.mu.Unlock()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Decryptor provides thread-safe decryption using the age encryption library.
|
||||||
|
// It uses a private key to decrypt data that was encrypted for the corresponding
|
||||||
|
// public key.
|
||||||
|
type Decryptor struct {
|
||||||
|
identity age.Identity
|
||||||
|
mu sync.RWMutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDecryptor creates a new decryptor with the given age private key.
|
||||||
|
// The private key should be a valid age X25519 identity string.
|
||||||
|
// Returns an error if the private key is invalid.
|
||||||
|
func NewDecryptor(privateKey string) (*Decryptor, error) {
|
||||||
|
identity, err := age.ParseX25519Identity(privateKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("parsing age identity: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Decryptor{
|
||||||
|
identity: identity,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Decrypt decrypts data using age decryption.
|
||||||
|
// This method is suitable for small to medium amounts of data that fit in memory.
|
||||||
|
// For large data streams, use DecryptStream instead.
|
||||||
|
func (d *Decryptor) Decrypt(data []byte) ([]byte, error) {
|
||||||
|
d.mu.RLock()
|
||||||
|
identity := d.identity
|
||||||
|
d.mu.RUnlock()
|
||||||
|
|
||||||
|
r, err := age.Decrypt(bytes.NewReader(data), identity)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("creating decrypted reader: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
decrypted, err := io.ReadAll(r)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("reading decrypted data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return decrypted, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DecryptStream returns a reader that decrypts data from the provided reader.
|
||||||
|
// This method is suitable for decrypting large files or streams as it processes
|
||||||
|
// data in a streaming fashion without loading everything into memory.
|
||||||
|
// The caller should close the input reader when done.
|
||||||
|
func (d *Decryptor) DecryptStream(src io.Reader) (io.Reader, error) {
|
||||||
|
d.mu.RLock()
|
||||||
|
identity := d.identity
|
||||||
|
d.mu.RUnlock()
|
||||||
|
|
||||||
|
r, err := age.Decrypt(src, identity)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("creating decrypted reader: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return r, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Module exports the crypto module for fx dependency injection.
|
||||||
|
//
|
||||||
|
//nolint:gochecknoglobals // fx module definitions are package globals
|
||||||
|
var Module = fx.Module("crypto")
|
||||||
@@ -0,0 +1,178 @@
|
|||||||
|
package crypto_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"filippo.io/age"
|
||||||
|
"sneak.berlin/go/vaultik/internal/crypto"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEncryptor(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Generate a test key pair
|
||||||
|
identity, err := age.GenerateX25519Identity()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to generate identity: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
publicKey := identity.Recipient().String()
|
||||||
|
|
||||||
|
// Create encryptor
|
||||||
|
enc, err := crypto.NewEncryptor([]string{publicKey})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to create encryptor: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test data
|
||||||
|
plaintext := []byte("Hello, World! This is a test message.")
|
||||||
|
|
||||||
|
// Encrypt
|
||||||
|
ciphertext, err := enc.Encrypt(plaintext)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to encrypt: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify it's actually encrypted (should be larger and different)
|
||||||
|
if bytes.Equal(plaintext, ciphertext) {
|
||||||
|
t.Error("ciphertext equals plaintext")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Decrypt to verify
|
||||||
|
r, err := age.Decrypt(bytes.NewReader(ciphertext), identity)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to decrypt: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var decrypted bytes.Buffer
|
||||||
|
|
||||||
|
_, err = decrypted.ReadFrom(r)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to read decrypted data: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !bytes.Equal(plaintext, decrypted.Bytes()) {
|
||||||
|
t.Error("decrypted data doesn't match original")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEncryptorMultipleRecipients(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Generate three test key pairs
|
||||||
|
identity1, err := age.GenerateX25519Identity()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to generate identity1: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
identity2, err := age.GenerateX25519Identity()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to generate identity2: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
identity3, err := age.GenerateX25519Identity()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to generate identity3: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
publicKeys := []string{
|
||||||
|
identity1.Recipient().String(),
|
||||||
|
identity2.Recipient().String(),
|
||||||
|
identity3.Recipient().String(),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create encryptor with multiple recipients
|
||||||
|
enc, err := crypto.NewEncryptor(publicKeys)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to create encryptor: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test data
|
||||||
|
plaintext := []byte("Secret message for multiple recipients")
|
||||||
|
|
||||||
|
// Encrypt
|
||||||
|
ciphertext, err := enc.Encrypt(plaintext)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to encrypt: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify each recipient can decrypt
|
||||||
|
identities := []age.Identity{identity1, identity2, identity3}
|
||||||
|
for i, identity := range identities {
|
||||||
|
r, err := age.Decrypt(bytes.NewReader(ciphertext), identity)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("recipient %d failed to decrypt: %v", i+1, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var decrypted bytes.Buffer
|
||||||
|
|
||||||
|
_, err = decrypted.ReadFrom(r)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("recipient %d failed to read decrypted data: %v", i+1, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !bytes.Equal(plaintext, decrypted.Bytes()) {
|
||||||
|
t.Errorf("recipient %d: decrypted data doesn't match original", i+1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEncryptorUpdateRecipients(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Generate two identities
|
||||||
|
identity1, _ := age.GenerateX25519Identity()
|
||||||
|
identity2, _ := age.GenerateX25519Identity()
|
||||||
|
|
||||||
|
publicKey1 := identity1.Recipient().String()
|
||||||
|
publicKey2 := identity2.Recipient().String()
|
||||||
|
|
||||||
|
// Create encryptor with first key
|
||||||
|
enc, err := crypto.NewEncryptor([]string{publicKey1})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to create encryptor: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Encrypt with first key
|
||||||
|
plaintext := []byte("test data")
|
||||||
|
|
||||||
|
ciphertext1, err := enc.Encrypt(plaintext)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to encrypt: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update to second key
|
||||||
|
err = enc.UpdateRecipients([]string{publicKey2})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to update recipients: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Encrypt with second key
|
||||||
|
ciphertext2, err := enc.Encrypt(plaintext)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to encrypt: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// First ciphertext should only decrypt with first identity
|
||||||
|
_, err = age.Decrypt(bytes.NewReader(ciphertext1), identity1)
|
||||||
|
if err != nil {
|
||||||
|
t.Error("failed to decrypt with identity1")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = age.Decrypt(bytes.NewReader(ciphertext1), identity2)
|
||||||
|
if err == nil {
|
||||||
|
t.Error("should not decrypt with identity2")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Second ciphertext should only decrypt with second identity
|
||||||
|
_, err = age.Decrypt(bytes.NewReader(ciphertext2), identity2)
|
||||||
|
if err != nil {
|
||||||
|
t.Error("failed to decrypt with identity2")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = age.Decrypt(bytes.NewReader(ciphertext2), identity1)
|
||||||
|
if err == nil {
|
||||||
|
t.Error("should not decrypt with identity1")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -208,30 +208,6 @@ func (r *BlobRepository) DeleteOrphaned(ctx context.Context) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteUnuploaded deletes blob rows whose upload never completed
|
|
||||||
// (uploaded_ts IS NULL) and returns how many were removed. Their
|
|
||||||
// blob_chunks rows are removed by the ON DELETE CASCADE foreign key.
|
|
||||||
// A blob is only ever attached to a snapshot once its upload has been
|
|
||||||
// recorded, so an un-uploaded blob is never referenced by a completed
|
|
||||||
// snapshot: dropping it discards chunk rows that point at data which
|
|
||||||
// was never stored remotely, so the affected content is re-chunked and
|
|
||||||
// re-uploaded on the next run.
|
|
||||||
func (r *BlobRepository) DeleteUnuploaded(ctx context.Context) (int64, error) {
|
|
||||||
query := `DELETE FROM blobs WHERE uploaded_ts IS NULL`
|
|
||||||
|
|
||||||
result, err := r.db.ExecWithLog(ctx, query)
|
|
||||||
if err != nil {
|
|
||||||
return 0, fmt.Errorf("deleting un-uploaded blobs: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
rowsAffected, _ := result.RowsAffected()
|
|
||||||
if rowsAffected > 0 {
|
|
||||||
log.Debug("Deleted un-uploaded blobs", "count", rowsAffected)
|
|
||||||
}
|
|
||||||
|
|
||||||
return rowsAffected, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// getOne fetches a single blob row matched on the given column, or
|
// getOne fetches a single blob row matched on the given column, or
|
||||||
// (nil, nil) when no row matches.
|
// (nil, nil) when no row matches.
|
||||||
func (r *BlobRepository) getOne(
|
func (r *BlobRepository) getOne(
|
||||||
|
|||||||
@@ -7,32 +7,12 @@ import (
|
|||||||
|
|
||||||
// List returns every chunk in the index, ordered by chunk hash.
|
// List returns every chunk in the index, ordered by chunk hash.
|
||||||
func (r *ChunkRepository) List(ctx context.Context) ([]*Chunk, error) {
|
func (r *ChunkRepository) List(ctx context.Context) ([]*Chunk, error) {
|
||||||
return r.list(ctx, `
|
query := `
|
||||||
SELECT chunk_hash, size
|
SELECT chunk_hash, size
|
||||||
FROM chunks
|
FROM chunks
|
||||||
ORDER BY chunk_hash
|
ORDER BY chunk_hash
|
||||||
`)
|
`
|
||||||
}
|
|
||||||
|
|
||||||
// ListInUploadedBlobs returns the chunks that are stored in a blob whose
|
|
||||||
// upload has completed (uploaded_ts set), ordered by chunk hash. These
|
|
||||||
// are the only chunks a backup may safely deduplicate against: a chunk
|
|
||||||
// recorded solely in a blob that was never uploaded refers to data that
|
|
||||||
// is not in remote storage, so trusting it would silently drop that data
|
|
||||||
// from later snapshots.
|
|
||||||
func (r *ChunkRepository) ListInUploadedBlobs(ctx context.Context) ([]*Chunk, error) {
|
|
||||||
return r.list(ctx, `
|
|
||||||
SELECT DISTINCT c.chunk_hash, c.size
|
|
||||||
FROM chunks c
|
|
||||||
JOIN blob_chunks bc ON c.chunk_hash = bc.chunk_hash
|
|
||||||
JOIN blobs b ON bc.blob_id = b.id
|
|
||||||
WHERE b.uploaded_ts IS NOT NULL
|
|
||||||
ORDER BY c.chunk_hash
|
|
||||||
`)
|
|
||||||
}
|
|
||||||
|
|
||||||
// list runs a chunk-selecting query and scans the (chunk_hash, size) rows.
|
|
||||||
func (r *ChunkRepository) list(ctx context.Context, query string) ([]*Chunk, error) {
|
|
||||||
rows, err := r.db.conn.QueryContext(ctx, query)
|
rows, err := r.db.conn.QueryContext(ctx, query)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("querying chunks: %w", err)
|
return nil, fmt.Errorf("querying chunks: %w", err)
|
||||||
|
|||||||
@@ -17,7 +17,6 @@ import (
|
|||||||
"embed"
|
"embed"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/url"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"sort"
|
"sort"
|
||||||
@@ -220,135 +219,6 @@ func openWithRecovery(ctx context.Context, path string) (*DB, error) {
|
|||||||
return db, nil
|
return db, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// errUntrustedSnapshotSchema is returned when a downloaded snapshot
|
|
||||||
// database carries schema objects the real schema never defines, or is
|
|
||||||
// missing a table the restore and deep-verify queries read.
|
|
||||||
var errUntrustedSnapshotSchema = errors.New(
|
|
||||||
"downloaded snapshot database has an untrusted schema")
|
|
||||||
|
|
||||||
// snapshotReadOnlyDSN builds the driver DSN that opens a materialized
|
|
||||||
// snapshot database file read-only. mode=ro opens the file read-only at
|
|
||||||
// the OS level, query_only rejects any write the engine is asked to make,
|
|
||||||
// and trusted_schema=OFF refuses to run application code named in the
|
|
||||||
// schema. The file: URI form is required for the driver to honour the
|
|
||||||
// mode parameter.
|
|
||||||
func snapshotReadOnlyDSN(path string) string {
|
|
||||||
u := url.URL{
|
|
||||||
Scheme: "file",
|
|
||||||
Path: path,
|
|
||||||
RawQuery: "mode=ro&_pragma=query_only(true)&_pragma=trusted_schema(false)",
|
|
||||||
}
|
|
||||||
|
|
||||||
return u.String()
|
|
||||||
}
|
|
||||||
|
|
||||||
// OpenReadOnly opens an already-materialized SQLite file for read-only
|
|
||||||
// querying of a snapshot database downloaded from the store, used by
|
|
||||||
// restore and deep verify. Unlike New it never applies schema migrations
|
|
||||||
// and never writes: the connection is opened read-only with query_only
|
|
||||||
// and trusted_schema=OFF. It refuses any file whose schema carries a
|
|
||||||
// trigger, view or virtual table, or lacks an expected table, so a forged
|
|
||||||
// file cannot redefine what the restore queries return. The caller owns
|
|
||||||
// the file and must remove it.
|
|
||||||
func OpenReadOnly(ctx context.Context, path string) (*DB, error) {
|
|
||||||
conn, err := sql.Open("sqlite", snapshotReadOnlyDSN(path))
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("opening read-only database: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
configureConnPool(conn)
|
|
||||||
|
|
||||||
err = conn.PingContext(ctx)
|
|
||||||
if err != nil {
|
|
||||||
_ = conn.Close()
|
|
||||||
|
|
||||||
return nil, fmt.Errorf("opening read-only database: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = verifySnapshotSchema(ctx, conn)
|
|
||||||
if err != nil {
|
|
||||||
_ = conn.Close()
|
|
||||||
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return &DB{conn: conn, path: path}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// verifySnapshotSchema rejects a downloaded database whose schema is not
|
|
||||||
// the plain table set the real schema defines. Any trigger, view or
|
|
||||||
// virtual table, or a missing expected table, fails the open.
|
|
||||||
func verifySnapshotSchema(ctx context.Context, conn *sql.DB) error {
|
|
||||||
// expectedSnapshotTables are the tables the restore and deep-verify
|
|
||||||
// queries read. A downloaded database missing any of them is not a
|
|
||||||
// genuine snapshot database and is refused.
|
|
||||||
expectedSnapshotTables := []string{
|
|
||||||
"blob_chunks",
|
|
||||||
"blobs",
|
|
||||||
"chunks",
|
|
||||||
"file_chunks",
|
|
||||||
"files",
|
|
||||||
}
|
|
||||||
|
|
||||||
rows, err := conn.QueryContext(
|
|
||||||
ctx, "SELECT type, name, sql FROM sqlite_master")
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("reading snapshot schema: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() { _ = rows.Close() }()
|
|
||||||
|
|
||||||
present := make(map[string]struct{})
|
|
||||||
|
|
||||||
for rows.Next() {
|
|
||||||
var objType, name string
|
|
||||||
|
|
||||||
var objSQL sql.NullString
|
|
||||||
|
|
||||||
err = rows.Scan(&objType, &name, &objSQL)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("reading snapshot schema: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
switch objType {
|
|
||||||
case "trigger", "view":
|
|
||||||
return fmt.Errorf(
|
|
||||||
"%w: unexpected %s %q", errUntrustedSnapshotSchema, objType, name)
|
|
||||||
case "table":
|
|
||||||
if isVirtualTableSQL(objSQL.String) {
|
|
||||||
return fmt.Errorf(
|
|
||||||
"%w: unexpected virtual table %q",
|
|
||||||
errUntrustedSnapshotSchema, name)
|
|
||||||
}
|
|
||||||
|
|
||||||
present[name] = struct{}{}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
err = rows.Err()
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("reading snapshot schema: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, table := range expectedSnapshotTables {
|
|
||||||
if _, ok := present[table]; !ok {
|
|
||||||
return fmt.Errorf(
|
|
||||||
"%w: missing table %q", errUntrustedSnapshotSchema, table)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// isVirtualTableSQL reports whether a sqlite_master row's SQL defines a
|
|
||||||
// virtual table. Virtual tables are recorded with type 'table' but a
|
|
||||||
// "CREATE VIRTUAL TABLE" definition and can run module code, so they are
|
|
||||||
// refused alongside triggers and views.
|
|
||||||
func isVirtualTableSQL(createSQL string) bool {
|
|
||||||
return strings.HasPrefix(
|
|
||||||
strings.ToUpper(strings.TrimSpace(createSQL)), "CREATE VIRTUAL TABLE")
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewTestDB creates an in-memory SQLite database for testing purposes.
|
// NewTestDB creates an in-memory SQLite database for testing purposes.
|
||||||
// The database is automatically initialized with the schema and is ready
|
// The database is automatically initialized with the schema and is ready
|
||||||
// for use. Each call creates a new independent database instance.
|
// for use. Each call creates a new independent database instance.
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ package database
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -16,44 +15,6 @@ import (
|
|||||||
// the index describes the backed-up file tree and must stay private.
|
// the index describes the backed-up file tree and must stay private.
|
||||||
const indexDirPerm = 0o700
|
const indexDirPerm = 0o700
|
||||||
|
|
||||||
// indexFilePerm restricts the index file to the owning user; it lists every
|
|
||||||
// backed-up path and chunk hash and must stay private.
|
|
||||||
const indexFilePerm = 0o600
|
|
||||||
|
|
||||||
// ensureIndexFileMode makes the index file owner-only before the SQLite
|
|
||||||
// driver opens it: it creates the file 0600 if absent, or chmods an existing
|
|
||||||
// one to 0600. Doing this first matters because SQLite creates its -wal and
|
|
||||||
// -shm side files with the mode of the main database file, so a private main
|
|
||||||
// file yields private side files. The driver treats a zero-byte file as an
|
|
||||||
// empty database, so pre-creating it here is safe.
|
|
||||||
func ensureIndexFileMode(path string) error {
|
|
||||||
info, err := os.Stat(path)
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case err == nil:
|
|
||||||
if info.Mode().Perm() == indexFilePerm {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
err = os.Chmod(path, indexFilePerm)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("restricting index file permissions: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
case errors.Is(err, os.ErrNotExist):
|
|
||||||
//nolint:gosec // G304: the index path is operator-configured by design
|
|
||||||
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY, indexFilePerm)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("creating index file: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return f.Close()
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("checking index file: %w", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Module provides database dependencies
|
// Module provides database dependencies
|
||||||
//
|
//
|
||||||
//nolint:gochecknoglobals // fx module definitions are package globals by convention
|
//nolint:gochecknoglobals // fx module definitions are package globals by convention
|
||||||
@@ -73,11 +34,6 @@ func provideDatabase(lc fx.Lifecycle, cfg *config.Config) (*DB, error) {
|
|||||||
return nil, fmt.Errorf("creating index directory: %w", err)
|
return nil, fmt.Errorf("creating index directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = ensureIndexFileMode(cfg.IndexPath)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
db, err := New(context.Background(), cfg.IndexPath)
|
db, err := New(context.Background(), cfg.IndexPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("opening database: %w", err)
|
return nil, fmt.Errorf("opening database: %w", err)
|
||||||
|
|||||||
@@ -1,100 +0,0 @@
|
|||||||
package database
|
|
||||||
|
|
||||||
import (
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"syscall"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"go.uber.org/fx/fxtest"
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestProvideDatabaseFreshIndexMode verifies that provideDatabase creates a
|
|
||||||
// missing index file owner-only (0600), even under a lenient 022 umask that
|
|
||||||
// would otherwise leave a freshly created file world-readable.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // syscall.Umask is process-global; parallel tests would clash
|
|
||||||
func TestProvideDatabaseFreshIndexMode(t *testing.T) {
|
|
||||||
restore := syscall.Umask(0o022)
|
|
||||||
defer syscall.Umask(restore)
|
|
||||||
|
|
||||||
indexPath := filepath.Join(t.TempDir(), "index.sqlite")
|
|
||||||
|
|
||||||
openIndex(t, indexPath)
|
|
||||||
assertPerm(t, indexPath, 0o600)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestProvideDatabaseExistingIndexMode verifies that provideDatabase tightens
|
|
||||||
// an existing world-readable index (0644) in a group/other-readable directory
|
|
||||||
// down to owner-only (0600).
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // syscall.Umask is process-global; parallel tests would clash
|
|
||||||
func TestProvideDatabaseExistingIndexMode(t *testing.T) {
|
|
||||||
restore := syscall.Umask(0o022)
|
|
||||||
defer syscall.Umask(restore)
|
|
||||||
|
|
||||||
dir := filepath.Join(t.TempDir(), "data")
|
|
||||||
|
|
||||||
//nolint:gosec // G301: the test intentionally uses a 0755 directory
|
|
||||||
err := os.MkdirAll(dir, 0o755)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("creating index directory: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
//nolint:gosec // G302: the test intentionally uses a 0755 directory
|
|
||||||
err = os.Chmod(dir, 0o755)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("relaxing index directory permissions: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
indexPath := filepath.Join(dir, "index.sqlite")
|
|
||||||
|
|
||||||
//nolint:gosec // G306: the test intentionally starts from a 0644 index
|
|
||||||
err = os.WriteFile(indexPath, nil, 0o644)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("creating pre-existing index: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
//nolint:gosec // G302: the test intentionally starts from a 0644 index
|
|
||||||
err = os.Chmod(indexPath, 0o644)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("relaxing pre-existing index permissions: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
openIndex(t, indexPath)
|
|
||||||
assertPerm(t, indexPath, 0o600)
|
|
||||||
}
|
|
||||||
|
|
||||||
// openIndex runs provideDatabase against indexPath and closes the resulting
|
|
||||||
// database before returning.
|
|
||||||
func openIndex(t *testing.T, indexPath string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
cfg := &config.Config{IndexPath: indexPath}
|
|
||||||
|
|
||||||
db, err := provideDatabase(fxtest.NewLifecycle(t), cfg)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("provideDatabase: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = db.Close()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("closing database: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// assertPerm fails the test unless path has exactly the given permission bits.
|
|
||||||
func assertPerm(t *testing.T, path string, want os.FileMode) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
info, err := os.Stat(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("stat %s: %v", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
got := info.Mode().Perm()
|
|
||||||
if got != want {
|
|
||||||
t.Fatalf("permissions of %s = %#o, want %#o", path, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,145 +0,0 @@
|
|||||||
//nolint:testpackage // exercises unexported read-only open internals
|
|
||||||
package database
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"database/sql"
|
|
||||||
"errors"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// genuineSnapshotDB writes a real snapshot database (the full schema
|
|
||||||
// applied) to a fresh file and returns its path.
|
|
||||||
func genuineSnapshotDB(t *testing.T) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
path := filepath.Join(t.TempDir(), "snapshot.db")
|
|
||||||
|
|
||||||
db, err := New(context.Background(), path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("creating snapshot database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = db.Close()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("closing snapshot database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return path
|
|
||||||
}
|
|
||||||
|
|
||||||
// forgedDB creates an empty database file and runs the given statements
|
|
||||||
// against it read-write, so a test can plant schema objects the real
|
|
||||||
// schema never defines.
|
|
||||||
func forgedDB(t *testing.T, stmts ...string) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
path := filepath.Join(t.TempDir(), "forged.db")
|
|
||||||
|
|
||||||
db, err := sql.Open("sqlite", path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("opening forged database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, stmt := range stmts {
|
|
||||||
_, err = db.ExecContext(context.Background(), stmt)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("executing %q: %v", stmt, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
err = db.Close()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("closing forged database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return path
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenReadOnlyAcceptsGenuineSnapshot(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
db, err := OpenReadOnly(context.Background(), genuineSnapshotDB(t))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("OpenReadOnly refused a genuine snapshot database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = db.Close() })
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenReadOnlyRefusesWrites(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
db, err := OpenReadOnly(context.Background(), genuineSnapshotDB(t))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("OpenReadOnly: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = db.Close() })
|
|
||||||
|
|
||||||
// A schema write depends on no table columns, so the only reason it
|
|
||||||
// can fail is that the database is open read-only.
|
|
||||||
_, err = db.Conn().ExecContext(context.Background(),
|
|
||||||
"CREATE TABLE probe_readonly (x)")
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("expected a write to a read-only snapshot database to fail")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenReadOnlyRejectsView(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
path := forgedDB(t, "CREATE VIEW files AS SELECT 1 AS path")
|
|
||||||
|
|
||||||
_, err := OpenReadOnly(context.Background(), path)
|
|
||||||
if !errors.Is(err, errUntrustedSnapshotSchema) {
|
|
||||||
t.Fatalf("expected a view named files to be refused, got %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenReadOnlyRejectsTrigger(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
path := forgedDB(t,
|
|
||||||
"CREATE TABLE files (path TEXT)",
|
|
||||||
"CREATE TRIGGER t AFTER INSERT ON files BEGIN SELECT 1; END")
|
|
||||||
|
|
||||||
_, err := OpenReadOnly(context.Background(), path)
|
|
||||||
if !errors.Is(err, errUntrustedSnapshotSchema) {
|
|
||||||
t.Fatalf("expected a trigger to be refused, got %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOpenReadOnlyRejectsMissingTable(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// Only one of the expected tables is present.
|
|
||||||
path := forgedDB(t, "CREATE TABLE files (path TEXT)")
|
|
||||||
|
|
||||||
_, err := OpenReadOnly(context.Background(), path)
|
|
||||||
if !errors.Is(err, errUntrustedSnapshotSchema) {
|
|
||||||
t.Fatalf("expected a missing expected table to be refused, got %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIsVirtualTableSQL(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
sql string
|
|
||||||
want bool
|
|
||||||
}{
|
|
||||||
{"CREATE VIRTUAL TABLE t USING fts5(x)", true},
|
|
||||||
{" create virtual table t using fts5(x)", true},
|
|
||||||
{"CREATE TABLE t (x)", false},
|
|
||||||
{"CREATE VIEW t AS SELECT 1", false},
|
|
||||||
{"", false},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, c := range cases {
|
|
||||||
if got := isVirtualTableSQL(c.sql); got != c.want {
|
|
||||||
t.Errorf("isVirtualTableSQL(%q) = %v, want %v", c.sql, got, c.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -11,18 +11,6 @@ import (
|
|||||||
"sneak.berlin/go/vaultik/internal/types"
|
"sneak.berlin/go/vaultik/internal/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Sentinel errors for the single-snapshot invariant that an exported
|
|
||||||
// per-snapshot metadata database must satisfy.
|
|
||||||
var (
|
|
||||||
// ErrNoSnapshotInDatabase means the metadata database has no snapshot
|
|
||||||
// row at all.
|
|
||||||
ErrNoSnapshotInDatabase = errors.New("database contains no snapshot")
|
|
||||||
// ErrMultipleSnapshotsInDatabase means the metadata database holds
|
|
||||||
// more than the single snapshot an export is supposed to contain.
|
|
||||||
ErrMultipleSnapshotsInDatabase = errors.New(
|
|
||||||
"database contains more than one snapshot")
|
|
||||||
)
|
|
||||||
|
|
||||||
// SnapshotRepository provides access to the snapshots table and its
|
// SnapshotRepository provides access to the snapshots table and its
|
||||||
// snapshot_files / snapshot_blobs association tables.
|
// snapshot_files / snapshot_blobs association tables.
|
||||||
type SnapshotRepository struct {
|
type SnapshotRepository struct {
|
||||||
@@ -218,48 +206,6 @@ func (r *SnapshotRepository) GetByID(
|
|||||||
return &snapshot, nil
|
return &snapshot, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetOnlySnapshot returns the sole snapshot in an exported per-snapshot
|
|
||||||
// metadata database. The backup path writes each snapshot's database with
|
|
||||||
// exactly one snapshot row (see cleanSnapshotDB), so restore and deep
|
|
||||||
// verify expect exactly one. Zero rows return ErrNoSnapshotInDatabase and
|
|
||||||
// more than one returns ErrMultipleSnapshotsInDatabase; callers treat
|
|
||||||
// either as a failed identity check on the downloaded database.
|
|
||||||
func (r *SnapshotRepository) GetOnlySnapshot(ctx context.Context) (*Snapshot, error) {
|
|
||||||
query := `
|
|
||||||
SELECT id, hostname, vaultik_version, vaultik_git_revision,
|
|
||||||
started_at, completed_at, file_count, chunk_count, blob_count,
|
|
||||||
total_size, blob_size, compression_ratio
|
|
||||||
FROM snapshots
|
|
||||||
LIMIT 2
|
|
||||||
`
|
|
||||||
|
|
||||||
rows, err := r.db.conn.QueryContext(ctx, query)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("querying snapshots: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() {
|
|
||||||
err := rows.Close()
|
|
||||||
if err != nil {
|
|
||||||
Fatalf("failed to close rows: %v", err)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
snapshots, err := r.scanSnapshotRows(rows)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
switch len(snapshots) {
|
|
||||||
case 1:
|
|
||||||
return snapshots[0], nil
|
|
||||||
case 0:
|
|
||||||
return nil, ErrNoSnapshotInDatabase
|
|
||||||
default:
|
|
||||||
return nil, ErrMultipleSnapshotsInDatabase
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ListRecent returns up to limit snapshots, most recently started first.
|
// ListRecent returns up to limit snapshots, most recently started first.
|
||||||
func (r *SnapshotRepository) ListRecent(
|
func (r *SnapshotRepository) ListRecent(
|
||||||
ctx context.Context, limit int,
|
ctx context.Context, limit int,
|
||||||
|
|||||||
+1
-15
@@ -14,18 +14,8 @@ var Module = fx.Module("log",
|
|||||||
)
|
)
|
||||||
|
|
||||||
// New creates a new logger configuration from provided options.
|
// New creates a new logger configuration from provided options.
|
||||||
//
|
|
||||||
// JSON is intentionally not carried into Config: a command emitting a
|
|
||||||
// JSON document on stdout must keep its stderr log level under
|
|
||||||
// --verbose/--debug, so --json must not lower it (issue #112). JSON
|
|
||||||
// silences the stdout UI in setupGlobals instead.
|
|
||||||
func New(opts Options) Config {
|
func New(opts Options) Config {
|
||||||
return Config{
|
return Config(opts)
|
||||||
Verbose: opts.Verbose,
|
|
||||||
Debug: opts.Debug,
|
|
||||||
Cron: opts.Cron,
|
|
||||||
Quiet: opts.Quiet,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Options are provided by the CLI.
|
// Options are provided by the CLI.
|
||||||
@@ -34,8 +24,4 @@ type Options struct {
|
|||||||
Debug bool
|
Debug bool
|
||||||
Cron bool
|
Cron bool
|
||||||
Quiet bool
|
Quiet bool
|
||||||
// JSON marks a command whose stdout carries a machine-readable
|
|
||||||
// document. It silences the human UI on stdout (see setupGlobals),
|
|
||||||
// but unlike Quiet it leaves the stderr log level alone.
|
|
||||||
JSON bool
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,49 +0,0 @@
|
|||||||
//nolint:testpackage // exercises the unexported copyFile helper
|
|
||||||
package snapshot
|
|
||||||
|
|
||||||
import (
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"syscall"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestCopyFileExportCopyMode verifies that the exported snapshot database
|
|
||||||
// copy is created owner-only (0600), even under a lenient 022 umask that
|
|
||||||
// would otherwise leave a fresh file world-readable.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // syscall.Umask is process-global; parallel tests would clash
|
|
||||||
func TestCopyFileExportCopyMode(t *testing.T) {
|
|
||||||
restore := syscall.Umask(0o022)
|
|
||||||
defer syscall.Umask(restore)
|
|
||||||
|
|
||||||
dir := t.TempDir()
|
|
||||||
|
|
||||||
src := filepath.Join(dir, "index.sqlite")
|
|
||||||
|
|
||||||
err := os.WriteFile(src, []byte("index data"), 0o600)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("creating source index: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
dst := filepath.Join(dir, "snapshot.db")
|
|
||||||
|
|
||||||
sm := &SnapshotManager{fs: afero.NewOsFs()}
|
|
||||||
|
|
||||||
err = sm.copyFile(src, dst)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("copyFile: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
info, err := os.Stat(dst)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("stat export copy: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
got := info.Mode().Perm()
|
|
||||||
if got != 0o600 {
|
|
||||||
t.Fatalf("export copy permissions = %#o, want %#o", got, 0o600)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,64 +0,0 @@
|
|||||||
//nolint:testpackage // exercises the unexported generateBlobManifest
|
|
||||||
package snapshot
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/types"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestGenerateBlobManifest_MissingBlobFails is the regression guard for
|
|
||||||
// issue #157: a blob the snapshot references but that is absent from the
|
|
||||||
// blobs table used to be logged and skipped, yielding a manifest with
|
|
||||||
// fewer blobs than the snapshot needs. Since prune trusts the manifest
|
|
||||||
// alone, that omitted blob would be deleted at the next prune. Manifest
|
|
||||||
// generation must fail instead.
|
|
||||||
func TestGenerateBlobManifest_MissingBlobFails(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
dbPath := filepath.Join(t.TempDir(), "snapshot.db")
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
// A real blob row satisfies the snapshot_blobs foreign key on
|
|
||||||
// blob_id; the snapshot then references a different, absent hash.
|
|
||||||
presentBlob := &database.Blob{
|
|
||||||
ID: types.NewBlobID(),
|
|
||||||
Hash: types.BlobHash("present-blob-hash"),
|
|
||||||
CreatedTS: time.Now().Truncate(time.Second),
|
|
||||||
}
|
|
||||||
require.NoError(t, repos.Blobs.Create(ctx, nil, presentBlob))
|
|
||||||
|
|
||||||
snap := &database.Snapshot{
|
|
||||||
ID: "testhost_home_2026-05-01T00:00:00Z",
|
|
||||||
Hostname: "testhost",
|
|
||||||
}
|
|
||||||
require.NoError(t, repos.Snapshots.Create(ctx, nil, snap))
|
|
||||||
require.NoError(t, repos.Snapshots.AddBlob(ctx, nil,
|
|
||||||
snap.ID.String(), presentBlob.ID, types.BlobHash("absent-blob-hash")))
|
|
||||||
|
|
||||||
require.NoError(t, db.Close())
|
|
||||||
|
|
||||||
sm := &SnapshotManager{
|
|
||||||
config: &config.Config{CompressionLevel: 3},
|
|
||||||
fs: afero.NewOsFs(),
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err = sm.generateBlobManifest(ctx, dbPath, snap.ID.String())
|
|
||||||
require.Error(t, err, "manifest generation must fail on a missing blob")
|
|
||||||
assert.Contains(t, err.Error(), "absent-blob-hash")
|
|
||||||
}
|
|
||||||
@@ -63,9 +63,7 @@ type Scanner struct {
|
|||||||
exclude []string // Glob patterns for files/directories to exclude
|
exclude []string // Glob patterns for files/directories to exclude
|
||||||
compiledExclude []compiledPattern // Compiled glob patterns
|
compiledExclude []compiledPattern // Compiled glob patterns
|
||||||
progress *ProgressReporter
|
progress *ProgressReporter
|
||||||
// skipErrors skips files that cannot be opened or read (logged loudly);
|
skipErrors bool // Skip file read errors (log loudly but continue)
|
||||||
// packer, database, encryption, and upload errors still abort the run.
|
|
||||||
skipErrors bool
|
|
||||||
// ui is the user-facing output; never nil (defaults to a discarding writer).
|
// ui is the user-facing output; never nil (defaults to a discarding writer).
|
||||||
ui *ui.Writer
|
ui *ui.Writer
|
||||||
|
|
||||||
@@ -123,9 +121,7 @@ type ScannerConfig struct {
|
|||||||
EnableProgress bool // Enable the live progress reporter (ETAs, throughput)
|
EnableProgress bool // Enable the live progress reporter (ETAs, throughput)
|
||||||
UI *ui.Writer // Where user-facing scanner messages go; nil = discard
|
UI *ui.Writer // Where user-facing scanner messages go; nil = discard
|
||||||
Exclude []string // Glob patterns for files/directories to exclude
|
Exclude []string // Glob patterns for files/directories to exclude
|
||||||
// SkipErrors skips files that cannot be opened or read (log loudly but
|
SkipErrors bool // Skip file read errors (log loudly but continue)
|
||||||
// continue); packer, database, encryption, and upload errors still abort.
|
|
||||||
SkipErrors bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// ScanResult contains the results of a scan operation
|
// ScanResult contains the results of a scan operation
|
||||||
@@ -224,14 +220,7 @@ func (s *Scanner) Scan(
|
|||||||
defer s.progress.Stop()
|
defer s.progress.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Phase 0: Repair any state left by an interrupted previous run, then
|
// Phase 0: Load known files and chunks from database into memory for fast lookup
|
||||||
// load known files and chunks from the database into memory for fast
|
|
||||||
// lookup.
|
|
||||||
err := s.repairInterruptedBlobs(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
knownFiles, err := s.loadDatabaseState(ctx, path)
|
knownFiles, err := s.loadDatabaseState(ctx, path)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -328,38 +317,6 @@ func (s *Scanner) loadDatabaseState(
|
|||||||
return knownFiles, nil
|
return knownFiles, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// repairInterruptedBlobs discards blob rows left by a previous run whose
|
|
||||||
// upload never completed. Such a blob has its chunks, blob_chunks, and
|
|
||||||
// blobs rows committed to the local index before the upload is attempted,
|
|
||||||
// so a crash or dropped connection mid-upload leaves them behind while the
|
|
||||||
// data never reaches remote storage. Deduplicating against those chunks on
|
|
||||||
// a later run would produce a snapshot that reports success but cannot be
|
|
||||||
// restored. Dropping the un-uploaded blobs (their blob_chunks cascade) and
|
|
||||||
// then any chunks left unreferenced forces the affected data to be
|
|
||||||
// re-chunked and re-uploaded this run. A blob is attached to a snapshot
|
|
||||||
// only once its upload is recorded, so this never touches a completed
|
|
||||||
// snapshot's data.
|
|
||||||
func (s *Scanner) repairInterruptedBlobs(ctx context.Context) error {
|
|
||||||
removed, err := s.repos.Blobs.DeleteUnuploaded(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("removing un-uploaded blob records: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if removed == 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Warn("Discarded blob records from an interrupted previous run; "+
|
|
||||||
"their data will be re-uploaded", "blobs", removed)
|
|
||||||
|
|
||||||
err = s.repos.Chunks.DeleteOrphaned(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("removing orphaned chunks: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// summarizeScanPhase calculates total size to process, updates progress tracking,
|
// summarizeScanPhase calculates total size to process, updates progress tracking,
|
||||||
// and prints the scan phase summary with file counts and sizes
|
// and prints the scan phase summary with file counts and sizes
|
||||||
func (s *Scanner) summarizeScanPhase(
|
func (s *Scanner) summarizeScanPhase(
|
||||||
@@ -435,14 +392,11 @@ func (s *Scanner) loadKnownFiles(
|
|||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// loadKnownChunks loads the chunk hashes safe to deduplicate against into
|
// loadKnownChunks loads all known chunk hashes from the database into a
|
||||||
// an in-memory map for fast lookup, avoiding per-chunk database queries
|
// map for fast lookup. This avoids per-chunk database queries during file
|
||||||
// during file processing. Only chunks held by a blob whose upload
|
// processing.
|
||||||
// completed are loaded: a chunk left behind by an interrupted upload
|
|
||||||
// refers to data that never reached remote storage, and deduplicating
|
|
||||||
// against it would silently produce an unrestorable snapshot.
|
|
||||||
func (s *Scanner) loadKnownChunks(ctx context.Context) error {
|
func (s *Scanner) loadKnownChunks(ctx context.Context) error {
|
||||||
chunks, err := s.repos.Chunks.ListInUploadedBlobs(ctx)
|
chunks, err := s.repos.Chunks.List(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("listing chunks: %w", err)
|
return fmt.Errorf("listing chunks: %w", err)
|
||||||
}
|
}
|
||||||
@@ -1340,15 +1294,6 @@ func (s *Scanner) processFileWithErrorHandling(
|
|||||||
) (bool, error) {
|
) (bool, error) {
|
||||||
err := s.processFileStreaming(ctx, fileToProcess, result)
|
err := s.processFileStreaming(ctx, fileToProcess, result)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// A packer/database/encryption/upload failure means the chunk's data
|
|
||||||
// may not have been stored. Skipping the file would let the snapshot
|
|
||||||
// record a file whose chunk is in no blob and cannot be restored, so
|
|
||||||
// abort the run even under --skip-errors. Only open and read errors
|
|
||||||
// are skipped below.
|
|
||||||
var pErr *packerError
|
|
||||||
if errors.As(err, &pErr) {
|
|
||||||
return false, fmt.Errorf("processing file %s: %w", fileToProcess.Path, err)
|
|
||||||
}
|
|
||||||
// Handle files that were deleted between scan and process phases
|
// Handle files that were deleted between scan and process phases
|
||||||
if errors.Is(err, os.ErrNotExist) {
|
if errors.Is(err, os.ErrNotExist) {
|
||||||
log.Warn("File was deleted during backup, skipping",
|
log.Warn("File was deleted during backup, skipping",
|
||||||
@@ -1358,7 +1303,7 @@ func (s *Scanner) processFileWithErrorHandling(
|
|||||||
|
|
||||||
return true, nil
|
return true, nil
|
||||||
}
|
}
|
||||||
// Skip open/read errors if --skip-errors is enabled
|
// Skip file read errors if --skip-errors is enabled
|
||||||
if s.skipErrors {
|
if s.skipErrors {
|
||||||
log.Error("Failed to process file (skipping due to --skip-errors)",
|
log.Error("Failed to process file (skipping due to --skip-errors)",
|
||||||
"path", fileToProcess.Path, "error", err)
|
"path", fileToProcess.Path, "error", err)
|
||||||
@@ -1456,17 +1401,7 @@ func (s *Scanner) finalizeProcessPhase(ctx context.Context, result *ScanResult)
|
|||||||
return fmt.Errorf("parsing blob ID: %w", err)
|
return fmt.Errorf("parsing blob ID: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// With no remote backend the blob's lifecycle ends here, so
|
|
||||||
// mark it uploaded in the same transaction that attaches it to
|
|
||||||
// the snapshot. This keeps the invariant that any blob a
|
|
||||||
// snapshot references has uploaded_ts set, so deduplication and
|
|
||||||
// interrupted-run repair treat these blobs as trustworthy.
|
|
||||||
err = s.repos.WithTx(ctx, func(ctx context.Context, tx *sql.Tx) error {
|
err = s.repos.WithTx(ctx, func(ctx context.Context, tx *sql.Tx) error {
|
||||||
err := s.repos.Blobs.UpdateUploaded(ctx, tx, b.ID)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("marking blob uploaded: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return s.repos.Snapshots.AddBlob(ctx, tx, s.snapshotID, blobID,
|
return s.repos.Snapshots.AddBlob(ctx, tx, s.snapshotID, blobID,
|
||||||
types.BlobHash(b.Hash))
|
types.BlobHash(b.Hash))
|
||||||
})
|
})
|
||||||
@@ -1725,20 +1660,6 @@ type streamingChunkInfo struct {
|
|||||||
size int64
|
size int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// packerError marks an error that came from adding a chunk to the packer
|
|
||||||
// (packing, database, encryption, or upload). Such an error means the chunk's
|
|
||||||
// data may not have been stored, so the run must abort even under --skip-errors:
|
|
||||||
// skipping the file would leave the chunk recorded as backed up while it lives
|
|
||||||
// in no blob, and a later snapshot could record a file that cannot be restored.
|
|
||||||
// Only open and read errors are safe to skip.
|
|
||||||
type packerError struct {
|
|
||||||
err error
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *packerError) Error() string { return e.err.Error() }
|
|
||||||
|
|
||||||
func (e *packerError) Unwrap() error { return e.err }
|
|
||||||
|
|
||||||
// processFileStreaming processes a file by streaming chunks directly to the packer
|
// processFileStreaming processes a file by streaming chunks directly to the packer
|
||||||
func (s *Scanner) processFileStreaming(
|
func (s *Scanner) processFileStreaming(
|
||||||
ctx context.Context, fileToProcess *FileToProcess, result *ScanResult,
|
ctx context.Context, fileToProcess *FileToProcess, result *ScanResult,
|
||||||
@@ -1789,11 +1710,7 @@ func (s *Scanner) processFileStreaming(
|
|||||||
if !chunkExists {
|
if !chunkExists {
|
||||||
err := s.addChunkToPacker(ctx, chunk)
|
err := s.addChunkToPacker(ctx, chunk)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Mark as a packer error so --skip-errors cannot swallow it:
|
return err
|
||||||
// the chunk was registered as pending before packing, so a
|
|
||||||
// skipped file here would be recorded as backed up while its
|
|
||||||
// data was never stored.
|
|
||||||
return &packerError{err: err}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,216 +0,0 @@
|
|||||||
package snapshot_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
)
|
|
||||||
|
|
||||||
// errSimTempFail is the one-time temp-file creation failure blobTempFailFs
|
|
||||||
// injects, mirroring a full temp filesystem.
|
|
||||||
var errSimTempFail = errors.New("simulated temp-file creation failure")
|
|
||||||
|
|
||||||
// errSimRead is the read failure readFailFile injects for a file that opens
|
|
||||||
// but cannot be read.
|
|
||||||
var errSimRead = errors.New("simulated read failure")
|
|
||||||
|
|
||||||
// blobTempFailFs fails the first temp-file creation for a packer blob, then
|
|
||||||
// behaves normally, simulating a one-time failure to start a new blob.
|
|
||||||
type blobTempFailFs struct {
|
|
||||||
afero.Fs
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
failed bool
|
|
||||||
}
|
|
||||||
|
|
||||||
//nolint:ireturn // afero.Fs.OpenFile is defined to return the interface.
|
|
||||||
func (f *blobTempFailFs) OpenFile(
|
|
||||||
name string, flag int, perm os.FileMode,
|
|
||||||
) (afero.File, error) {
|
|
||||||
if strings.Contains(name, "vaultik-blob-") {
|
|
||||||
f.mu.Lock()
|
|
||||||
firstTime := !f.failed
|
|
||||||
f.failed = true
|
|
||||||
f.mu.Unlock()
|
|
||||||
|
|
||||||
if firstTime {
|
|
||||||
return nil, errSimTempFail
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return f.Fs.OpenFile(name, flag, perm)
|
|
||||||
}
|
|
||||||
|
|
||||||
// readFailFile wraps an afero.File whose Read always fails.
|
|
||||||
type readFailFile struct {
|
|
||||||
afero.File
|
|
||||||
}
|
|
||||||
|
|
||||||
func (readFailFile) Read([]byte) (int, error) {
|
|
||||||
return 0, errSimRead
|
|
||||||
}
|
|
||||||
|
|
||||||
// readFailFs fails reads of one target path after a successful open.
|
|
||||||
type readFailFs struct {
|
|
||||||
afero.Fs
|
|
||||||
|
|
||||||
target string
|
|
||||||
}
|
|
||||||
|
|
||||||
//nolint:ireturn // afero.Fs.Open is defined to return the interface.
|
|
||||||
func (f *readFailFs) Open(name string) (afero.File, error) {
|
|
||||||
file, err := f.Fs.Open(name)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if name == f.target {
|
|
||||||
return readFailFile{File: file}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return file, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeSkipErrorTestFile writes one file into fs with a fixed mtime.
|
|
||||||
func writeSkipErrorTestFile(t *testing.T, fs afero.Fs, path, content string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
err := fs.MkdirAll(filepath.Dir(path), 0755)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("mkdir: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = afero.WriteFile(fs, path, []byte(content), 0644)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write %s: %v", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
when := time.Date(2024, 1, 1, 12, 0, 0, 0, time.UTC)
|
|
||||||
|
|
||||||
err = fs.Chtimes(path, when, when)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("chtimes %s: %v", path, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// runSkipErrorScan scans /source on fs with the given skip-errors setting and
|
|
||||||
// returns the repositories (for inspection) and the scan error.
|
|
||||||
func runSkipErrorScan(
|
|
||||||
t *testing.T, fs afero.Fs, skipErrors bool,
|
|
||||||
) (*database.Repositories, error) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
db, err := database.NewTestDB()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("create test db: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
cerr := db.Close()
|
|
||||||
if cerr != nil {
|
|
||||||
t.Errorf("close db: %v", cerr)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
scanner := snapshot.NewScanner(snapshot.ScannerConfig{
|
|
||||||
FS: fs,
|
|
||||||
ChunkSize: int64(1024 * 16),
|
|
||||||
Repositories: repos,
|
|
||||||
MaxBlobSize: int64(1024 * 1024),
|
|
||||||
CompressionLevel: 3,
|
|
||||||
AgeRecipients: []string{testAgePublicKey},
|
|
||||||
SkipErrors: skipErrors,
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
snapshotID := "test-snapshot-skip-errors"
|
|
||||||
createTestSnapshotRecord(ctx, t, repos, snapshotID)
|
|
||||||
|
|
||||||
_, err = scanner.Scan(ctx, "/source", snapshotID)
|
|
||||||
|
|
||||||
return repos, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestScannerPackingFailureAbortsUnderSkipErrors checks that a failure to start
|
|
||||||
// a new blob aborts the run even with --skip-errors. Otherwise the file would
|
|
||||||
// be skipped while its chunk had already been registered as pending, letting a
|
|
||||||
// later blob record that chunk in the chunks table with no blob to back it —
|
|
||||||
// a snapshot that completes with a file that cannot be restored.
|
|
||||||
func TestScannerPackingFailureAbortsUnderSkipErrors(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// Two files with distinct content so each yields a distinct chunk: the
|
|
||||||
// first fails to start a blob, and without the fix the second's blob would
|
|
||||||
// commit the first's orphaned chunk row.
|
|
||||||
fs := &blobTempFailFs{Fs: afero.NewMemMapFs()}
|
|
||||||
writeSkipErrorTestFile(t, fs, "/source/file1.txt", "first file content")
|
|
||||||
writeSkipErrorTestFile(t, fs, "/source/file2.txt", "second file content")
|
|
||||||
|
|
||||||
repos, err := runSkipErrorScan(t, fs, true)
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("expected scan to abort on the packer error, got nil")
|
|
||||||
}
|
|
||||||
|
|
||||||
// ListUnpacked returns chunks recorded with no blob_chunks row: exactly the
|
|
||||||
// unrestorable state this fix prevents.
|
|
||||||
unpacked, err := repos.Chunks.ListUnpacked(context.Background(), 10)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("listing unpacked chunks: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(unpacked) != 0 {
|
|
||||||
t.Fatalf("expected no chunk recorded without a blob, got %d", len(unpacked))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestScannerReadErrorAbortsWithoutSkipErrors checks that a file read error
|
|
||||||
// aborts the run when --skip-errors is not set.
|
|
||||||
func TestScannerReadErrorAbortsWithoutSkipErrors(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const target = "/source/unreadable.txt"
|
|
||||||
|
|
||||||
fs := &readFailFs{Fs: afero.NewMemMapFs(), target: target}
|
|
||||||
writeSkipErrorTestFile(t, fs, target, "content that cannot be read")
|
|
||||||
|
|
||||||
_, err := runSkipErrorScan(t, fs, false)
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("expected scan to fail on the read error, got nil")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestScannerReadErrorSkippedWithSkipErrors checks that a file read error is
|
|
||||||
// skipped and the run completes when --skip-errors is set.
|
|
||||||
func TestScannerReadErrorSkippedWithSkipErrors(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const target = "/source/unreadable.txt"
|
|
||||||
|
|
||||||
fs := &readFailFs{Fs: afero.NewMemMapFs(), target: target}
|
|
||||||
writeSkipErrorTestFile(t, fs, target, "content that cannot be read")
|
|
||||||
|
|
||||||
repos, err := runSkipErrorScan(t, fs, true)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("expected scan to complete with --skip-errors, got %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
chunks, err := repos.FileChunks.GetByFile(context.Background(), target)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("getting file chunks: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(chunks) != 0 {
|
|
||||||
t.Fatalf("expected unreadable file skipped, got %d chunks", len(chunks))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+108
-27
@@ -44,7 +44,6 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -295,6 +294,68 @@ func (sm *SnapshotManager) ExportSnapshotMetadata(
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// CleanupIncompleteSnapshots removes incomplete snapshots that don't have
|
||||||
|
// metadata in S3. This is critical for data safety: incomplete snapshots
|
||||||
|
// can cause deduplication to skip files that were never successfully
|
||||||
|
// backed up, resulting in data loss.
|
||||||
|
func (sm *SnapshotManager) CleanupIncompleteSnapshots(
|
||||||
|
ctx context.Context, hostname string,
|
||||||
|
) error {
|
||||||
|
log.Info("Checking for incomplete snapshots", "hostname", hostname)
|
||||||
|
|
||||||
|
// Get all incomplete snapshots for this hostname
|
||||||
|
incompleteSnapshots, err := sm.repos.Snapshots.GetIncompleteByHostname(ctx, hostname)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("getting incomplete snapshots: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(incompleteSnapshots) == 0 {
|
||||||
|
log.Debug("No incomplete snapshots found")
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Info("Found incomplete snapshots", "count", len(incompleteSnapshots))
|
||||||
|
|
||||||
|
// Check each incomplete snapshot for metadata in storage
|
||||||
|
for _, snapshot := range incompleteSnapshots {
|
||||||
|
// Check if metadata exists in storage (paths use the hashed
|
||||||
|
// remote key so we don't leak host info to the listing).
|
||||||
|
metadataKey := fmt.Sprintf("metadata/%s/db.zst",
|
||||||
|
RemoteSnapshotKey(snapshot.ID.String()))
|
||||||
|
|
||||||
|
_, err := sm.storage.Stat(ctx, metadataKey)
|
||||||
|
if err != nil {
|
||||||
|
// Metadata doesn't exist in S3 - this is an incomplete snapshot
|
||||||
|
log.Info("Cleaning up incomplete snapshot record",
|
||||||
|
"snapshot_id", snapshot.ID, "started_at", snapshot.StartedAt)
|
||||||
|
|
||||||
|
// Delete the snapshot and all its associations
|
||||||
|
err := sm.deleteSnapshot(ctx, snapshot.ID.String())
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting incomplete snapshot %s: %w",
|
||||||
|
snapshot.ID, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Info("Deleted incomplete snapshot record and associated data",
|
||||||
|
"snapshot_id", snapshot.ID)
|
||||||
|
} else {
|
||||||
|
// Metadata exists - this snapshot was completed but database wasn't updated
|
||||||
|
// This shouldn't happen in normal operation, but mark it complete
|
||||||
|
log.Warn("Found snapshot with remote metadata but incomplete in database",
|
||||||
|
"snapshot_id", snapshot.ID)
|
||||||
|
|
||||||
|
err := sm.repos.Snapshots.MarkComplete(ctx, nil, snapshot.ID.String())
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to mark snapshot as complete in database",
|
||||||
|
"snapshot_id", snapshot.ID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// CleanupOrphanedData removes files, chunks, and blobs that are no longer
|
// CleanupOrphanedData removes files, chunks, and blobs that are no longer
|
||||||
// referenced by any snapshot. This should be called periodically to clean
|
// referenced by any snapshot. This should be called periodically to clean
|
||||||
// up data from deleted or incomplete snapshots.
|
// up data from deleted or incomplete snapshots.
|
||||||
@@ -697,18 +758,12 @@ func (sm *SnapshotManager) compressFile(inputPath, outputPath string) error {
|
|||||||
|
|
||||||
writerClosed = true
|
writerClosed = true
|
||||||
|
|
||||||
log.Debug("Compression complete", "hash", hex.EncodeToString(writer.ContentID()))
|
log.Debug("Compression complete", "hash", hex.EncodeToString(writer.Sum256()))
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// exportCopyPerm restricts the exported snapshot database copy to the owning
|
// copyFile copies a file from src to dst
|
||||||
// user; it holds the same private index data as the local index file.
|
|
||||||
const exportCopyPerm = 0o600
|
|
||||||
|
|
||||||
// copyFile copies a file from src to dst. The destination is the exported
|
|
||||||
// snapshot database, so it is created owner-only rather than with the
|
|
||||||
// umask-dependent default.
|
|
||||||
func (sm *SnapshotManager) copyFile(src, dst string) error {
|
func (sm *SnapshotManager) copyFile(src, dst string) error {
|
||||||
log.Debug("Opening source file for copy", "path", src)
|
log.Debug("Opening source file for copy", "path", src)
|
||||||
|
|
||||||
@@ -728,9 +783,7 @@ func (sm *SnapshotManager) copyFile(src, dst string) error {
|
|||||||
|
|
||||||
log.Debug("Creating destination file", "path", dst)
|
log.Debug("Creating destination file", "path", dst)
|
||||||
|
|
||||||
destFile, err := sm.fs.OpenFile(
|
destFile, err := sm.fs.Create(dst)
|
||||||
dst, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, exportCopyPerm,
|
|
||||||
)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -756,11 +809,6 @@ func (sm *SnapshotManager) copyFile(src, dst string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// errBlobMissingFromDatabase means a snapshot references a blob that is
|
|
||||||
// absent from the blobs table, so a complete manifest cannot be built.
|
|
||||||
var errBlobMissingFromDatabase = errors.New(
|
|
||||||
"blob referenced by snapshot is not in the database")
|
|
||||||
|
|
||||||
// generateBlobManifest creates a compressed JSON list of all blobs in the snapshot
|
// generateBlobManifest creates a compressed JSON list of all blobs in the snapshot
|
||||||
func (sm *SnapshotManager) generateBlobManifest(
|
func (sm *SnapshotManager) generateBlobManifest(
|
||||||
ctx context.Context, dbPath string, snapshotID string,
|
ctx context.Context, dbPath string, snapshotID string,
|
||||||
@@ -791,27 +839,21 @@ func (sm *SnapshotManager) generateBlobManifest(
|
|||||||
totalCompressedSize := int64(0)
|
totalCompressedSize := int64(0)
|
||||||
|
|
||||||
for _, hash := range blobHashes {
|
for _, hash := range blobHashes {
|
||||||
// Every blob the snapshot references must appear in the manifest.
|
|
||||||
// Prune consults only the manifest to decide what is still in use,
|
|
||||||
// so silently dropping a blob here would let a later prune delete
|
|
||||||
// it while this snapshot still needs it. A lookup failure or a
|
|
||||||
// missing blob row therefore fails manifest generation.
|
|
||||||
blob, err := repos.Blobs.GetByHash(ctx, hash)
|
blob, err := repos.Blobs.GetByHash(ctx, hash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("getting blob details for %s: %w", hash, err)
|
log.Warn("Failed to get blob details", "hash", hash, "error", err)
|
||||||
}
|
|
||||||
|
|
||||||
if blob == nil {
|
continue
|
||||||
return nil, fmt.Errorf("%w: blob %s, snapshot %s",
|
|
||||||
errBlobMissingFromDatabase, hash, snapshotID)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if blob != nil {
|
||||||
blobs = append(blobs, BlobInfo{
|
blobs = append(blobs, BlobInfo{
|
||||||
Hash: hash,
|
Hash: hash,
|
||||||
CompressedSize: blob.CompressedSize,
|
CompressedSize: blob.CompressedSize,
|
||||||
})
|
})
|
||||||
totalCompressedSize += blob.CompressedSize
|
totalCompressedSize += blob.CompressedSize
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Create manifest. SnapshotID in the unencrypted manifest is the
|
// Create manifest. SnapshotID in the unencrypted manifest is the
|
||||||
// double-SHA256 remote key (see RemoteSnapshotKey), not the human ID,
|
// double-SHA256 remote key (see RemoteSnapshotKey), not the human ID,
|
||||||
@@ -871,6 +913,45 @@ type ExtendedBackupStats struct {
|
|||||||
UploadDurationMs int64 // Total milliseconds spent uploading to S3
|
UploadDurationMs int64 // Total milliseconds spent uploading to S3
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// deleteSnapshot removes a snapshot and all its associations from the database
|
||||||
|
func (sm *SnapshotManager) deleteSnapshot(
|
||||||
|
ctx context.Context, snapshotID string,
|
||||||
|
) error {
|
||||||
|
// Delete snapshot_files entries
|
||||||
|
err := sm.repos.Snapshots.DeleteSnapshotFiles(ctx, snapshotID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting snapshot files: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete snapshot_blobs entries
|
||||||
|
err = sm.repos.Snapshots.DeleteSnapshotBlobs(ctx, snapshotID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting snapshot blobs: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete uploads entries (has foreign key to snapshots without CASCADE)
|
||||||
|
err = sm.repos.Snapshots.DeleteSnapshotUploads(ctx, snapshotID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting snapshot uploads: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete the snapshot itself
|
||||||
|
err = sm.repos.Snapshots.Delete(ctx, snapshotID)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("deleting snapshot: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clean up orphaned data
|
||||||
|
log.Debug("Cleaning up orphaned records in main database")
|
||||||
|
|
||||||
|
err = sm.CleanupOrphanedData(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("cleaning up orphaned data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// deleteOtherSnapshots deletes all snapshots except the current one
|
// deleteOtherSnapshots deletes all snapshots except the current one
|
||||||
func (sm *SnapshotManager) deleteOtherSnapshots(
|
func (sm *SnapshotManager) deleteOtherSnapshots(
|
||||||
ctx context.Context, tx *sql.Tx, currentSnapshotID string,
|
ctx context.Context, tx *sql.Tx, currentSnapshotID string,
|
||||||
|
|||||||
@@ -1,198 +0,0 @@
|
|||||||
package storage_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"io"
|
|
||||||
"reflect"
|
|
||||||
"sort"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
)
|
|
||||||
|
|
||||||
// runStorerConformance is the shared Storer contract. Every backend that
|
|
||||||
// can run in-process is expected to pass it: TestFileStorer runs it against
|
|
||||||
// file://, TestS3Storer against s3://. A new backend inherits this coverage
|
|
||||||
// by passing its own constructor, so the contract is defined once.
|
|
||||||
//
|
|
||||||
// It exercises the public Storer interface: round-trip, stat, list with
|
|
||||||
// prefix filtering, overwrite, delete, delete-of-missing, and not-found on
|
|
||||||
// Get and Stat. Each section takes its own fresh backend instance, so the
|
|
||||||
// order of sections never matters and no section sees another's objects.
|
|
||||||
func runStorerConformance(t *testing.T, newStorer func(*testing.T) storage.Storer) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
conformanceRoundTrip(t, newStorer(t))
|
|
||||||
conformanceOverwrite(t, newStorer(t))
|
|
||||||
conformanceList(t, newStorer(t))
|
|
||||||
conformanceDelete(t, newStorer(t))
|
|
||||||
conformanceNotFound(t, newStorer(t))
|
|
||||||
}
|
|
||||||
|
|
||||||
// conformanceRoundTrip stores a nested key, then reads it back and stats it.
|
|
||||||
func conformanceRoundTrip(t *testing.T, s storage.Storer) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
key := "blobs/aa/bb/object.bin"
|
|
||||||
want := []byte("round-trip payload")
|
|
||||||
|
|
||||||
err := s.Put(ctx, key, bytes.NewReader(want))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
got := getBytes(t, s, key)
|
|
||||||
if !bytes.Equal(got, want) {
|
|
||||||
t.Errorf("Get returned %q, want %q", got, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
info, err := s.Stat(ctx, key)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Stat: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if info.Key != key {
|
|
||||||
t.Errorf("Stat key = %q, want %q", info.Key, key)
|
|
||||||
}
|
|
||||||
|
|
||||||
if info.Size != int64(len(want)) {
|
|
||||||
t.Errorf("Stat size = %d, want %d", info.Size, len(want))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// conformanceOverwrite checks that a second Put replaces the first.
|
|
||||||
func conformanceOverwrite(t *testing.T, s storage.Storer) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
key := "meta/snapshot.json"
|
|
||||||
|
|
||||||
err := s.Put(ctx, key, bytes.NewReader([]byte("first")))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("first Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
want := []byte("second and longer payload")
|
|
||||||
|
|
||||||
err = s.Put(ctx, key, bytes.NewReader(want))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("second Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
got := getBytes(t, s, key)
|
|
||||||
if !bytes.Equal(got, want) {
|
|
||||||
t.Errorf("after overwrite Get returned %q, want %q", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// conformanceList checks prefix filtering and the empty result for a
|
|
||||||
// prefix that matches nothing.
|
|
||||||
func conformanceList(t *testing.T, s storage.Storer) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
keys := []string{"blobs/aa/one", "blobs/bb/two", "meta/three"}
|
|
||||||
|
|
||||||
for _, k := range keys {
|
|
||||||
err := s.Put(ctx, k, bytes.NewReader([]byte("data")))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Put %q: %v", k, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if got := listSorted(t, s, ""); !reflect.DeepEqual(got, keys) {
|
|
||||||
t.Errorf("List(\"\") = %v, want %v", got, keys)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantBlobs := []string{"blobs/aa/one", "blobs/bb/two"}
|
|
||||||
if got := listSorted(t, s, "blobs/"); !reflect.DeepEqual(got, wantBlobs) {
|
|
||||||
t.Errorf("List(\"blobs/\") = %v, want %v", got, wantBlobs)
|
|
||||||
}
|
|
||||||
|
|
||||||
if got := listSorted(t, s, "absent/"); len(got) != 0 {
|
|
||||||
t.Errorf("List(\"absent/\") = %v, want empty", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// conformanceDelete checks that Delete removes an object and that deleting
|
|
||||||
// a missing key is not an error.
|
|
||||||
func conformanceDelete(t *testing.T, s storage.Storer) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
key := "blobs/cc/gone.bin"
|
|
||||||
|
|
||||||
err := s.Put(ctx, key, bytes.NewReader([]byte("temporary")))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Put: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = s.Delete(ctx, key)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Delete: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err = s.Get(ctx, key)
|
|
||||||
if !errors.Is(err, storage.ErrNotFound) {
|
|
||||||
t.Errorf("Get after Delete error = %v, want ErrNotFound", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = s.Delete(ctx, key)
|
|
||||||
if err != nil {
|
|
||||||
t.Errorf("Delete of missing key = %v, want nil", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// conformanceNotFound checks Get and Stat on an absent key.
|
|
||||||
func conformanceNotFound(t *testing.T, s storage.Storer) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
key := "never/written"
|
|
||||||
|
|
||||||
_, err := s.Get(ctx, key)
|
|
||||||
if !errors.Is(err, storage.ErrNotFound) {
|
|
||||||
t.Errorf("Get error = %v, want ErrNotFound", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err = s.Stat(ctx, key)
|
|
||||||
if !errors.Is(err, storage.ErrNotFound) {
|
|
||||||
t.Errorf("Stat error = %v, want ErrNotFound", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// getBytes reads a key fully and closes the reader.
|
|
||||||
func getBytes(t *testing.T, s storage.Storer, key string) []byte {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
rc, err := s.Get(context.Background(), key)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Get %q: %v", key, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() { _ = rc.Close() }()
|
|
||||||
|
|
||||||
data, err := io.ReadAll(rc)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read %q: %v", key, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return data
|
|
||||||
}
|
|
||||||
|
|
||||||
// listSorted returns the keys under a prefix in a stable order.
|
|
||||||
func listSorted(t *testing.T, s storage.Storer, prefix string) []string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
keys, err := s.List(context.Background(), prefix)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("List %q: %v", prefix, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
sort.Strings(keys)
|
|
||||||
|
|
||||||
return keys
|
|
||||||
}
|
|
||||||
@@ -1,209 +0,0 @@
|
|||||||
// Package faultstore provides a storage.Storer wrapper that injects
|
|
||||||
// faults on demand, so tests can reproduce the failure modes a real
|
|
||||||
// backend exhibits: an upload that fails partway, a backend that reports
|
|
||||||
// success while storing nothing, and reads that return corrupt or
|
|
||||||
// truncated bytes. It is the seam called for by the fault-injection
|
|
||||||
// tests (sneak/vaultik issue 72) and is meant to be reused by future
|
|
||||||
// tests rather than re-implemented per case.
|
|
||||||
//
|
|
||||||
// The wrapper delegates every method to the inner Storer. Two hooks
|
|
||||||
// change that: OnPut decides the fate of each write, and OnGet decides
|
|
||||||
// how each read's bytes are returned. Both are keyed by the object key,
|
|
||||||
// so a test can fault only blobs, only metadata, or a single object.
|
|
||||||
package faultstore
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ErrInjectedUpload is returned by a Put the OnPut hook chose to fail.
|
|
||||||
var ErrInjectedUpload = errors.New("faultstore: injected upload failure")
|
|
||||||
|
|
||||||
// PutAction is the disposition OnPut assigns to a write.
|
|
||||||
type PutAction int
|
|
||||||
|
|
||||||
const (
|
|
||||||
// PutNormal writes through to the inner Storer.
|
|
||||||
PutNormal PutAction = iota
|
|
||||||
// PutFail reads part of the stream, then fails without storing the
|
|
||||||
// object — a network upload that dies partway through.
|
|
||||||
PutFail
|
|
||||||
// PutSwallow reports success but stores nothing — a backend that
|
|
||||||
// lies about durability.
|
|
||||||
PutSwallow
|
|
||||||
)
|
|
||||||
|
|
||||||
// GetFault is how OnGet chooses to damage a read.
|
|
||||||
type GetFault int
|
|
||||||
|
|
||||||
const (
|
|
||||||
// GetNormal returns the stored bytes unchanged.
|
|
||||||
GetNormal GetFault = iota
|
|
||||||
// GetCorrupt flips a byte so the returned object no longer matches
|
|
||||||
// what was stored.
|
|
||||||
GetCorrupt
|
|
||||||
// GetTruncate returns a short read: the object's bytes cut off
|
|
||||||
// before the end.
|
|
||||||
GetTruncate
|
|
||||||
)
|
|
||||||
|
|
||||||
// Storer wraps an inner storage.Storer with fault-injection hooks. A
|
|
||||||
// zero-valued hook means "no fault": construct with New and set only the
|
|
||||||
// hook a test needs.
|
|
||||||
type Storer struct {
|
|
||||||
inner storage.Storer
|
|
||||||
|
|
||||||
// OnPut, when set, is consulted before every Put and
|
|
||||||
// PutWithProgress with the object key.
|
|
||||||
OnPut func(key string) PutAction
|
|
||||||
|
|
||||||
// OnGet, when set, is consulted for every Get with the object key
|
|
||||||
// and damages the returned bytes accordingly.
|
|
||||||
OnGet func(key string) GetFault
|
|
||||||
}
|
|
||||||
|
|
||||||
// New wraps inner. inner must be non-nil.
|
|
||||||
func New(inner storage.Storer) *Storer {
|
|
||||||
return &Storer{inner: inner}
|
|
||||||
}
|
|
||||||
|
|
||||||
// midStreamBytes is how far a PutFail reads before failing, enough to be
|
|
||||||
// past the start of any real blob without depending on the blob's size.
|
|
||||||
const midStreamBytes = 512
|
|
||||||
|
|
||||||
// Put stores data unless OnPut faults the write.
|
|
||||||
func (f *Storer) Put(ctx context.Context, key string, data io.Reader) error {
|
|
||||||
handled, err := f.injectPut(key, data)
|
|
||||||
if handled {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
return f.inner.Put(ctx, key, data)
|
|
||||||
}
|
|
||||||
|
|
||||||
// PutWithProgress stores data unless OnPut faults the write.
|
|
||||||
func (f *Storer) PutWithProgress(
|
|
||||||
ctx context.Context, key string, data io.Reader,
|
|
||||||
size int64, progress storage.ProgressCallback,
|
|
||||||
) error {
|
|
||||||
handled, err := f.injectPut(key, data)
|
|
||||||
if handled {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
return f.inner.PutWithProgress(ctx, key, data, size, progress)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Get retrieves data, damaging it if OnGet faults the read.
|
|
||||||
func (f *Storer) Get(ctx context.Context, key string) (io.ReadCloser, error) {
|
|
||||||
rc, err := f.inner.Get(ctx, key)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
fault := GetNormal
|
|
||||||
if f.OnGet != nil {
|
|
||||||
fault = f.OnGet(key)
|
|
||||||
}
|
|
||||||
|
|
||||||
if fault == GetNormal {
|
|
||||||
return rc, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
data, err := io.ReadAll(rc)
|
|
||||||
_ = rc.Close()
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return io.NopCloser(bytes.NewReader(damage(fault, data))), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// damage returns a faulted copy of the stored bytes. GetCorrupt flips a
|
|
||||||
// byte in the middle so decryption authentication fails; GetTruncate
|
|
||||||
// drops the final byte so the read ends short. Both are no-ops on empty
|
|
||||||
// input, which cannot be damaged into something distinguishable.
|
|
||||||
func damage(fault GetFault, data []byte) []byte {
|
|
||||||
out := make([]byte, len(data))
|
|
||||||
copy(out, data)
|
|
||||||
|
|
||||||
if len(out) == 0 {
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
switch fault {
|
|
||||||
case GetCorrupt:
|
|
||||||
out[len(out)/2] ^= 0xff
|
|
||||||
case GetTruncate:
|
|
||||||
out = out[:len(out)-1]
|
|
||||||
case GetNormal:
|
|
||||||
}
|
|
||||||
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// Stat delegates unchanged.
|
|
||||||
func (f *Storer) Stat(ctx context.Context, key string) (*storage.ObjectInfo, error) {
|
|
||||||
return f.inner.Stat(ctx, key)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Delete delegates unchanged.
|
|
||||||
func (f *Storer) Delete(ctx context.Context, key string) error {
|
|
||||||
return f.inner.Delete(ctx, key)
|
|
||||||
}
|
|
||||||
|
|
||||||
// List delegates unchanged.
|
|
||||||
func (f *Storer) List(ctx context.Context, prefix string) ([]string, error) {
|
|
||||||
return f.inner.List(ctx, prefix)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ListStream delegates unchanged.
|
|
||||||
func (f *Storer) ListStream(
|
|
||||||
ctx context.Context, prefix string,
|
|
||||||
) <-chan storage.ObjectInfo {
|
|
||||||
return f.inner.ListStream(ctx, prefix)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Info delegates unchanged.
|
|
||||||
func (f *Storer) Info() storage.Info {
|
|
||||||
return f.inner.Info()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Storer) putAction(key string) PutAction {
|
|
||||||
if f.OnPut == nil {
|
|
||||||
return PutNormal
|
|
||||||
}
|
|
||||||
|
|
||||||
return f.OnPut(key)
|
|
||||||
}
|
|
||||||
|
|
||||||
// injectPut handles the non-normal write dispositions. It reports
|
|
||||||
// whether it handled the write and, if so, with what error.
|
|
||||||
func (f *Storer) injectPut(key string, data io.Reader) (bool, error) {
|
|
||||||
switch f.putAction(key) {
|
|
||||||
case PutFail:
|
|
||||||
// Consume part of the stream so the failure lands mid-transfer,
|
|
||||||
// the way a dropped connection would, then error without
|
|
||||||
// storing anything.
|
|
||||||
_, _ = io.CopyN(io.Discard, data, midStreamBytes)
|
|
||||||
|
|
||||||
return true, fmt.Errorf("%w for %q", ErrInjectedUpload, key)
|
|
||||||
case PutSwallow:
|
|
||||||
// A lying backend still drains the request body, then keeps
|
|
||||||
// nothing.
|
|
||||||
_, _ = io.Copy(io.Discard, data)
|
|
||||||
|
|
||||||
return true, nil
|
|
||||||
case PutNormal:
|
|
||||||
return false, nil
|
|
||||||
default:
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,27 +0,0 @@
|
|||||||
package storage_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
)
|
|
||||||
|
|
||||||
// newFileStorer builds a file:// backend rooted at a fresh temp directory.
|
|
||||||
//
|
|
||||||
//nolint:ireturn // conformance runs against the Storer interface by design
|
|
||||||
func newFileStorer(t *testing.T) storage.Storer {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
s, err := storage.NewFileStorer(t.TempDir())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("NewFileStorer: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestFileStorer runs the shared Storer contract against the file:// backend.
|
|
||||||
func TestFileStorer(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
runStorerConformance(t, newFileStorer)
|
|
||||||
}
|
|
||||||
@@ -111,11 +111,10 @@ func storerFromParsedS3URL(parsed *URL, cfg *config.Config) (Storer, error) {
|
|||||||
func storerFromLegacyS3Config(cfg *config.Config) (Storer, error) {
|
func storerFromLegacyS3Config(cfg *config.Config) (Storer, error) {
|
||||||
endpoint := cfg.S3.Endpoint
|
endpoint := cfg.S3.Endpoint
|
||||||
|
|
||||||
// Ensure protocol is present. Absent an explicit use_ssl, default to TLS;
|
// Ensure protocol is present
|
||||||
// plain HTTP only when use_ssl is written as false.
|
|
||||||
if !strings.HasPrefix(endpoint, "http://") &&
|
if !strings.HasPrefix(endpoint, "http://") &&
|
||||||
!strings.HasPrefix(endpoint, "https://") {
|
!strings.HasPrefix(endpoint, "https://") {
|
||||||
if cfg.S3.UseSSL == nil || *cfg.S3.UseSSL {
|
if cfg.S3.UseSSL {
|
||||||
endpoint = "https://" + endpoint
|
endpoint = "https://" + endpoint
|
||||||
} else {
|
} else {
|
||||||
endpoint = "http://" + endpoint
|
endpoint = "http://" + endpoint
|
||||||
|
|||||||
@@ -1,61 +0,0 @@
|
|||||||
package storage_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
)
|
|
||||||
|
|
||||||
// legacyS3Config returns a minimal s3.* (no storage_url) configuration with a
|
|
||||||
// scheme-less endpoint. useSSL mirrors the config file: nil means the key is
|
|
||||||
// omitted, a pointer means it was written explicitly.
|
|
||||||
func legacyS3Config(useSSL *bool) *config.Config {
|
|
||||||
return &config.Config{
|
|
||||||
S3: config.S3Config{
|
|
||||||
Endpoint: "s3.example.com",
|
|
||||||
Bucket: "bucket",
|
|
||||||
AccessKeyID: "key",
|
|
||||||
SecretAccessKey: "secret",
|
|
||||||
Region: "us-east-1",
|
|
||||||
UseSSL: useSSL,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// endpointScheme builds the storer from cfg and returns the scheme its
|
|
||||||
// resolved endpoint carries (Info().Location is "endpoint/bucket").
|
|
||||||
func endpointScheme(t *testing.T, cfg *config.Config) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
storer, err := storage.NewStorer(cfg)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("NewStorer: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
location := storer.Info().Location
|
|
||||||
switch {
|
|
||||||
case strings.HasPrefix(location, "https://"):
|
|
||||||
return "https"
|
|
||||||
case strings.HasPrefix(location, "http://"):
|
|
||||||
return "http"
|
|
||||||
default:
|
|
||||||
t.Fatalf("endpoint has no http(s) scheme: %q", location)
|
|
||||||
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLegacyS3SchemelessEndpointDefaultsToTLS(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
if got := endpointScheme(t, legacyS3Config(nil)); got != "https" {
|
|
||||||
t.Errorf("use_ssl omitted: got %q scheme, want https", got)
|
|
||||||
}
|
|
||||||
|
|
||||||
no := false
|
|
||||||
if got := endpointScheme(t, legacyS3Config(&no)); got != "http" {
|
|
||||||
t.Errorf("use_ssl: false: got %q scheme, want http", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,58 +0,0 @@
|
|||||||
package storage_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The rclone backend is a thin adapter over the rclone library: it turns a
|
|
||||||
// (remote, path) pair into rclone's "remote:path" string, hands it to
|
|
||||||
// rclone, and maps rclone's own results back to the Storer interface. What
|
|
||||||
// can be tested in-process, without a configured remote or network, is that
|
|
||||||
// adapter layer — how the arguments are shaped and how construction errors
|
|
||||||
// are reported. The data-plane operations (Put/Get/List/Delete) are rclone's
|
|
||||||
// own, exercised against a real provider (drive, s3-via-rclone, ...), which
|
|
||||||
// needs a configured remote with credentials and network access and so is
|
|
||||||
// out of reach of a unit test. The shared Storer conformance suite therefore
|
|
||||||
// runs against the in-process file and s3 backends; the rclone backend
|
|
||||||
// inherits that contract once a remote is configured.
|
|
||||||
//
|
|
||||||
// These tests use rclone's ":local:" on-the-fly backend, which addresses the
|
|
||||||
// local filesystem directly without any configured remote, so construction
|
|
||||||
// runs entirely in-process.
|
|
||||||
|
|
||||||
// TestNewRcloneStorerConstruction checks that a valid remote constructs a
|
|
||||||
// backend and that Info() reports the shaped "remote:path" location.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // NewRcloneStorer installs the process-global rclone config
|
|
||||||
func TestNewRcloneStorerConstruction(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
|
|
||||||
s, err := storage.NewRcloneStorer(context.Background(), ":local", dir)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("NewRcloneStorer: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Info().Location is the "remote:path" string the adapter builds from
|
|
||||||
// its two arguments, so asserting it confirms the argument shaping.
|
|
||||||
want := ":local:" + dir
|
|
||||||
if got := s.Info().Location; got != want {
|
|
||||||
t.Errorf("Info().Location = %q, want %q", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestNewRcloneStorerUnknownRemote checks that a remote that is not in the
|
|
||||||
// rclone config fails construction with the ErrRemoteNotFound sentinel,
|
|
||||||
// rather than silently returning a backend pointed nowhere.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // NewRcloneStorer installs the process-global rclone config
|
|
||||||
func TestNewRcloneStorerUnknownRemote(t *testing.T) {
|
|
||||||
_, err := storage.NewRcloneStorer(
|
|
||||||
context.Background(), "vaultik-no-such-remote", "path")
|
|
||||||
if !errors.Is(err, storage.ErrRemoteNotFound) {
|
|
||||||
t.Errorf("NewRcloneStorer error = %v, want ErrRemoteNotFound", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+14
-36
@@ -13,23 +13,18 @@ import (
|
|||||||
"sneak.berlin/go/vaultik/internal/storage"
|
"sneak.berlin/go/vaultik/internal/storage"
|
||||||
)
|
)
|
||||||
|
|
||||||
// s3TestBucket is the bucket created for each in-process S3 server.
|
// TestS3StorerMissingKeyMapsToErrNotFound verifies that the s3 backend reports
|
||||||
const s3TestBucket = "test-bucket"
|
// a missing object as storage.ErrNotFound, matching the file and rclone
|
||||||
|
// backends and the Storer contract. Without the mapping, Get and Stat leak the
|
||||||
// newS3Storer builds an s3:// backend backed by a fresh in-process
|
// raw SDK error and errors.Is(err, storage.ErrNotFound) is false.
|
||||||
// S3 server. It reuses the same in-memory S3 harness (gofakes3 + s3mem
|
|
||||||
// over httptest) that internal/s3 and the not-found regression test use,
|
|
||||||
// so no new mock or dependency is introduced. Each call gets its own
|
|
||||||
// server, bucket, and client, so the conformance suite's per-section
|
|
||||||
// instances stay isolated.
|
|
||||||
//
|
//
|
||||||
//nolint:ireturn // conformance runs against the Storer interface by design
|
//nolint:paralleltest // shares an in-process S3 server via t.Cleanup
|
||||||
func newS3Storer(t *testing.T) storage.Storer {
|
func TestS3StorerMissingKeyMapsToErrNotFound(t *testing.T) {
|
||||||
t.Helper()
|
const bucket = "test-bucket"
|
||||||
|
|
||||||
backend := s3mem.New()
|
backend := s3mem.New()
|
||||||
|
|
||||||
err := backend.CreateBucket(s3TestBucket)
|
err := backend.CreateBucket(bucket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("create bucket: %v", err)
|
t.Fatalf("create bucket: %v", err)
|
||||||
}
|
}
|
||||||
@@ -37,9 +32,11 @@ func newS3Storer(t *testing.T) storage.Storer {
|
|||||||
srv := httptest.NewServer(gofakes3.New(backend).Server())
|
srv := httptest.NewServer(gofakes3.New(backend).Server())
|
||||||
t.Cleanup(srv.Close)
|
t.Cleanup(srv.Close)
|
||||||
|
|
||||||
client, err := s3.NewClient(context.Background(), s3.Config{
|
ctx := context.Background()
|
||||||
|
|
||||||
|
client, err := s3.NewClient(ctx, s3.Config{
|
||||||
Endpoint: srv.URL,
|
Endpoint: srv.URL,
|
||||||
Bucket: s3TestBucket,
|
Bucket: bucket,
|
||||||
AccessKeyID: "test",
|
AccessKeyID: "test",
|
||||||
SecretAccessKey: "test",
|
SecretAccessKey: "test",
|
||||||
Region: "us-east-1",
|
Region: "us-east-1",
|
||||||
@@ -48,28 +45,9 @@ func newS3Storer(t *testing.T) storage.Storer {
|
|||||||
t.Fatalf("new client: %v", err)
|
t.Fatalf("new client: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return storage.NewS3Storer(client)
|
storer := storage.NewS3Storer(client)
|
||||||
}
|
|
||||||
|
|
||||||
// TestS3Storer runs the shared Storer contract against the s3:// backend,
|
_, err = storer.Get(ctx, "does-not-exist")
|
||||||
// so it is held to the same round-trip, list, delete, and not-found
|
|
||||||
// behaviour as the file:// backend.
|
|
||||||
func TestS3Storer(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
runStorerConformance(t, newS3Storer)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestS3StorerMissingKeyMapsToErrNotFound pins the specific contract that a
|
|
||||||
// missing object surfaces as storage.ErrNotFound rather than the raw AWS SDK
|
|
||||||
// error. Without the mapping, errors.Is(err, storage.ErrNotFound) is false on
|
|
||||||
// s3 and callers would branch differently per backend.
|
|
||||||
func TestS3StorerMissingKeyMapsToErrNotFound(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
storer := newS3Storer(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
_, err := storer.Get(ctx, "does-not-exist")
|
|
||||||
if !errors.Is(err, storage.ErrNotFound) {
|
if !errors.Is(err, storage.ErrNotFound) {
|
||||||
t.Errorf("Get on missing key: got %v, want ErrNotFound", err)
|
t.Errorf("Get on missing key: got %v, want ErrNotFound", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+17
-72
@@ -4,7 +4,6 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/url"
|
"net/url"
|
||||||
"slices"
|
|
||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -24,10 +23,6 @@ var (
|
|||||||
ErrUnsupportedScheme = errors.New(
|
ErrUnsupportedScheme = errors.New(
|
||||||
"unsupported URL scheme: must start with s3://, file://, or rclone://")
|
"unsupported URL scheme: must start with s3://, file://, or rclone://")
|
||||||
ErrUnsupportedStorage = errors.New("unsupported storage scheme")
|
ErrUnsupportedStorage = errors.New("unsupported storage scheme")
|
||||||
ErrURLCredentials = errors.New(
|
|
||||||
"storage URL must not carry credentials; " +
|
|
||||||
"set s3.access_key_id and s3.secret_access_key in the config instead")
|
|
||||||
ErrURLUnknownParam = errors.New("unknown query parameter in storage URL")
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// URL represents a parsed storage URL.
|
// URL represents a parsed storage URL.
|
||||||
@@ -64,28 +59,11 @@ func ParseStorageURL(rawURL string) (*URL, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Handle s3:// URLs
|
||||||
if strings.HasPrefix(rawURL, "s3://") {
|
if strings.HasPrefix(rawURL, "s3://") {
|
||||||
return parseS3URL(rawURL)
|
|
||||||
}
|
|
||||||
|
|
||||||
if strings.HasPrefix(rawURL, "rclone://") {
|
|
||||||
return parseRcloneURL(rawURL)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil, ErrUnsupportedScheme
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseS3URL parses an s3://bucket/prefix URL. It rejects credentials in
|
|
||||||
// the userinfo and any query parameter other than endpoint, region and
|
|
||||||
// ssl, so a credential-bearing URL is never stored or echoed.
|
|
||||||
func parseS3URL(rawURL string) (*URL, error) {
|
|
||||||
u, err := url.Parse(rawURL)
|
u, err := url.Parse(rawURL)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, wrapParseError(err)
|
return nil, fmt.Errorf("invalid URL: %w", err)
|
||||||
}
|
|
||||||
|
|
||||||
if u.User != nil {
|
|
||||||
return nil, ErrURLCredentials
|
|
||||||
}
|
}
|
||||||
|
|
||||||
bucket := u.Host
|
bucket := u.Host
|
||||||
@@ -93,34 +71,30 @@ func parseS3URL(rawURL string) (*URL, error) {
|
|||||||
return nil, ErrMissingBucket
|
return nil, ErrMissingBucket
|
||||||
}
|
}
|
||||||
|
|
||||||
|
prefix := strings.TrimPrefix(u.Path, "/")
|
||||||
|
|
||||||
query := u.Query()
|
query := u.Query()
|
||||||
|
|
||||||
err = rejectUnknownParams(query, "endpoint", "region", "ssl")
|
useSSL := true
|
||||||
if err != nil {
|
if query.Get("ssl") == "false" {
|
||||||
return nil, err
|
useSSL = false
|
||||||
}
|
}
|
||||||
|
|
||||||
return &URL{
|
return &URL{
|
||||||
Scheme: schemeS3,
|
Scheme: schemeS3,
|
||||||
Bucket: bucket,
|
Bucket: bucket,
|
||||||
Prefix: strings.TrimPrefix(u.Path, "/"),
|
Prefix: prefix,
|
||||||
Endpoint: query.Get("endpoint"),
|
Endpoint: query.Get("endpoint"),
|
||||||
Region: query.Get("region"),
|
Region: query.Get("region"),
|
||||||
UseSSL: query.Get("ssl") != "false",
|
UseSSL: useSSL,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
|
||||||
|
|
||||||
// parseRcloneURL parses an rclone://remote/path URL. rclone:// takes no
|
|
||||||
// query parameters, so credentials in the userinfo and any parameter at
|
|
||||||
// all are rejected rather than silently ignored.
|
|
||||||
func parseRcloneURL(rawURL string) (*URL, error) {
|
|
||||||
u, err := url.Parse(rawURL)
|
|
||||||
if err != nil {
|
|
||||||
return nil, wrapParseError(err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if u.User != nil {
|
// Handle rclone:// URLs
|
||||||
return nil, ErrURLCredentials
|
if strings.HasPrefix(rawURL, "rclone://") {
|
||||||
|
u, err := url.Parse(rawURL)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid URL: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
remote := u.Host
|
remote := u.Host
|
||||||
@@ -128,45 +102,16 @@ func parseRcloneURL(rawURL string) (*URL, error) {
|
|||||||
return nil, ErrMissingRemote
|
return nil, ErrMissingRemote
|
||||||
}
|
}
|
||||||
|
|
||||||
err = rejectUnknownParams(u.Query())
|
path := strings.TrimPrefix(u.Path, "/")
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return &URL{
|
return &URL{
|
||||||
Scheme: schemeRclone,
|
Scheme: schemeRclone,
|
||||||
Prefix: strings.TrimPrefix(u.Path, "/"),
|
Prefix: path,
|
||||||
RcloneRemote: remote,
|
RcloneRemote: remote,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
|
||||||
|
|
||||||
// rejectUnknownParams returns an error naming the first query parameter
|
|
||||||
// not in allowed. The parameter's name is included (so a misspelt
|
|
||||||
// endpoint= is caught), but never its value, which could be a secret,
|
|
||||||
// and never the whole URL.
|
|
||||||
func rejectUnknownParams(query url.Values, allowed ...string) error {
|
|
||||||
for name := range query {
|
|
||||||
if !slices.Contains(allowed, name) {
|
|
||||||
return fmt.Errorf(
|
|
||||||
"%w: %q; put credentials in s3.access_key_id and "+
|
|
||||||
"s3.secret_access_key, not the URL",
|
|
||||||
ErrURLUnknownParam, name)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil, ErrUnsupportedScheme
|
||||||
}
|
|
||||||
|
|
||||||
// wrapParseError wraps only the inner cause of a url.Parse failure. The
|
|
||||||
// *url.Error that url.Parse returns embeds the raw URL in its message, so
|
|
||||||
// wrapping it directly would echo a credential-bearing URL into logs.
|
|
||||||
func wrapParseError(err error) error {
|
|
||||||
var uerr *url.Error
|
|
||||||
if errors.As(err, &uerr) {
|
|
||||||
return fmt.Errorf("invalid URL: %w", uerr.Err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return fmt.Errorf("invalid URL: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// String returns a human-readable representation of the storage URL.
|
// String returns a human-readable representation of the storage URL.
|
||||||
|
|||||||
@@ -1,208 +0,0 @@
|
|||||||
package storage_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"reflect"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestParseStorageURLValid checks that each supported scheme parses into
|
|
||||||
// the expected fields, since those fields decide which backend is built.
|
|
||||||
func TestParseStorageURLValid(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const bucket = "mybucket"
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
raw string
|
|
||||||
want *storage.URL
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "file absolute path",
|
|
||||||
raw: "file:///var/backups/vaultik",
|
|
||||||
want: &storage.URL{Scheme: "file", Prefix: "/var/backups/vaultik"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "s3 bucket and prefix, ssl defaults on",
|
|
||||||
raw: "s3://mybucket/backups/host",
|
|
||||||
want: &storage.URL{
|
|
||||||
Scheme: "s3", Bucket: bucket,
|
|
||||||
Prefix: "backups/host", UseSSL: true,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "s3 bucket only",
|
|
||||||
raw: "s3://mybucket",
|
|
||||||
want: &storage.URL{Scheme: "s3", Bucket: bucket, UseSSL: true},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "s3 with endpoint, region, ssl off",
|
|
||||||
raw: "s3://mybucket?endpoint=minio.example.com®ion=us-west-2&ssl=false",
|
|
||||||
want: &storage.URL{
|
|
||||||
Scheme: "s3", Bucket: bucket,
|
|
||||||
Endpoint: "minio.example.com", Region: "us-west-2", UseSSL: false,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "rclone remote and path",
|
|
||||||
raw: "rclone://gdrive/backups/host",
|
|
||||||
want: &storage.URL{
|
|
||||||
Scheme: "rclone", RcloneRemote: "gdrive", Prefix: "backups/host",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "rclone remote only",
|
|
||||||
raw: "rclone://gdrive",
|
|
||||||
want: &storage.URL{Scheme: "rclone", RcloneRemote: "gdrive"},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range cases {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
got, err := storage.ParseStorageURL(tc.raw)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("ParseStorageURL(%q) returned error: %v", tc.raw, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !reflect.DeepEqual(got, tc.want) {
|
|
||||||
t.Errorf("ParseStorageURL(%q) = %+v, want %+v", tc.raw, got, tc.want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestParseStorageURLErrors checks that empty, missing, and unknown-scheme
|
|
||||||
// inputs fail with the documented sentinel errors instead of parsing to a
|
|
||||||
// wrong destination.
|
|
||||||
func TestParseStorageURLErrors(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
raw string
|
|
||||||
wantErr error
|
|
||||||
}{
|
|
||||||
{"empty url", "", storage.ErrEmptyStorageURL},
|
|
||||||
{"file empty path", "file://", storage.ErrEmptyFilePath},
|
|
||||||
{"s3 missing bucket", "s3://", storage.ErrMissingBucket},
|
|
||||||
{"s3 missing bucket with path", "s3:///justprefix", storage.ErrMissingBucket},
|
|
||||||
{"rclone missing remote", "rclone://", storage.ErrMissingRemote},
|
|
||||||
{"unknown scheme", "gs://bucket/x", storage.ErrUnsupportedScheme},
|
|
||||||
{"no scheme", "/local/path", storage.ErrUnsupportedScheme},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range cases {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
_, err := storage.ParseStorageURL(tc.raw)
|
|
||||||
if !errors.Is(err, tc.wantErr) {
|
|
||||||
t.Errorf("ParseStorageURL(%q) error = %v, want %v",
|
|
||||||
tc.raw, err, tc.wantErr)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestParseStorageURLRejectsCredentials checks that a URL carrying
|
|
||||||
// credentials in its userinfo or in an unknown query parameter is
|
|
||||||
// rejected, and that the error never echoes the secret-bearing URL back
|
|
||||||
// into logs or output.
|
|
||||||
func TestParseStorageURLRejectsCredentials(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// Split so the literals never form a "user:pass@" URL pattern that
|
|
||||||
// tooling would flag as a real hardcoded credential.
|
|
||||||
const (
|
|
||||||
key = "AKIAKEY"
|
|
||||||
secret = "topsecret"
|
|
||||||
)
|
|
||||||
|
|
||||||
cases := []struct {
|
|
||||||
name string
|
|
||||||
raw string
|
|
||||||
wantErr error
|
|
||||||
secrets []string // must not appear in the error message
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "s3 userinfo",
|
|
||||||
raw: "s3://" + key + ":" + secret + "@mybucket/prefix",
|
|
||||||
wantErr: storage.ErrURLCredentials,
|
|
||||||
secrets: []string{key, secret, "mybucket"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "s3 unknown query param",
|
|
||||||
raw: "s3://mybucket?access_key=" + key + "&secret=" + secret,
|
|
||||||
wantErr: storage.ErrURLUnknownParam,
|
|
||||||
secrets: []string{key, secret},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "s3 misspelt endpoint",
|
|
||||||
raw: "s3://mybucket?endpiont=minio.example.com",
|
|
||||||
wantErr: storage.ErrURLUnknownParam,
|
|
||||||
secrets: nil,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "rclone userinfo",
|
|
||||||
raw: "rclone://user:" + secret + "@gdrive/backups",
|
|
||||||
wantErr: storage.ErrURLCredentials,
|
|
||||||
secrets: []string{secret},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "rclone query param",
|
|
||||||
raw: "rclone://gdrive/backups?token=" + secret,
|
|
||||||
wantErr: storage.ErrURLUnknownParam,
|
|
||||||
secrets: []string{secret},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range cases {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
_, err := storage.ParseStorageURL(tc.raw)
|
|
||||||
if !errors.Is(err, tc.wantErr) {
|
|
||||||
t.Fatalf("ParseStorageURL(%q) error = %v, want %v",
|
|
||||||
tc.raw, err, tc.wantErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The rejection must name the proper config keys so the
|
|
||||||
// operator knows where credentials belong.
|
|
||||||
for _, key := range []string{"s3.access_key_id", "s3.secret_access_key"} {
|
|
||||||
if !strings.Contains(err.Error(), key) {
|
|
||||||
t.Errorf("error %q does not name %q", err.Error(), key)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, secret := range tc.secrets {
|
|
||||||
if strings.Contains(err.Error(), secret) {
|
|
||||||
t.Errorf("error message leaked %q: %v", secret, err.Error())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestParseStorageURLParseFailureHidesURL checks that when url.Parse
|
|
||||||
// itself fails, the wrapped error carries only the inner cause, not the
|
|
||||||
// *url.Error whose text embeds the raw (possibly credential-bearing) URL.
|
|
||||||
func TestParseStorageURLParseFailureHidesURL(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const raw = "s3://mybucket/%zz"
|
|
||||||
|
|
||||||
_, err := storage.ParseStorageURL(raw)
|
|
||||||
if err == nil {
|
|
||||||
t.Fatalf("ParseStorageURL(%q) returned no error", raw)
|
|
||||||
}
|
|
||||||
|
|
||||||
if strings.Contains(err.Error(), "mybucket") {
|
|
||||||
t.Errorf("error message echoed the raw URL: %v", err.Error())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+48
-3
@@ -1,7 +1,7 @@
|
|||||||
// Package types provides custom types for better type safety across the
|
// Package types provides custom types for better type safety across the
|
||||||
// vaultik codebase. Using distinct types for IDs, hashes, and paths prevents
|
// vaultik codebase. Using distinct types for IDs, hashes, paths, and
|
||||||
// accidental mixing of semantically different values that happen to share the
|
// credentials prevents accidental mixing of semantically different values
|
||||||
// same underlying type.
|
// that happen to share the same underlying type.
|
||||||
package types //nolint:revive,nolintlint // rename decision tracked in #76
|
package types //nolint:revive,nolintlint // rename decision tracked in #76
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -157,6 +157,34 @@ type FilePath string
|
|||||||
// Used during restore to strip the source prefix from paths.
|
// Used during restore to strip the source prefix from paths.
|
||||||
type SourcePath string
|
type SourcePath string
|
||||||
|
|
||||||
|
// AgeRecipient is an age public key used for encryption.
|
||||||
|
// Format: age1... (Bech32-encoded X25519 public key)
|
||||||
|
type AgeRecipient string
|
||||||
|
|
||||||
|
// AgeSecretKey is an age private key used for decryption.
|
||||||
|
// Format: AGE-SECRET-KEY-... (Bech32-encoded X25519 private key)
|
||||||
|
// This type should never be logged or serialized in plaintext.
|
||||||
|
type AgeSecretKey string
|
||||||
|
|
||||||
|
// S3Endpoint is the URL of an S3-compatible storage endpoint.
|
||||||
|
type S3Endpoint string
|
||||||
|
|
||||||
|
// BucketName is the name of an S3 bucket.
|
||||||
|
type BucketName string
|
||||||
|
|
||||||
|
// S3Prefix is the path prefix within an S3 bucket.
|
||||||
|
type S3Prefix string
|
||||||
|
|
||||||
|
// AWSRegion is an AWS region identifier (e.g., "us-east-1").
|
||||||
|
type AWSRegion string
|
||||||
|
|
||||||
|
// AWSAccessKeyID is an AWS access key ID for authentication.
|
||||||
|
type AWSAccessKeyID string
|
||||||
|
|
||||||
|
// AWSSecretAccessKey is an AWS secret access key for authentication.
|
||||||
|
// This type should never be logged or serialized in plaintext.
|
||||||
|
type AWSSecretAccessKey string
|
||||||
|
|
||||||
// Hostname identifies a host machine.
|
// Hostname identifies a host machine.
|
||||||
type Hostname string
|
type Hostname string
|
||||||
|
|
||||||
@@ -178,7 +206,24 @@ func (h ChunkHash) String() string { return string(h) }
|
|||||||
func (h BlobHash) String() string { return string(h) }
|
func (h BlobHash) String() string { return string(h) }
|
||||||
func (p FilePath) String() string { return string(p) }
|
func (p FilePath) String() string { return string(p) }
|
||||||
func (p SourcePath) String() string { return string(p) }
|
func (p SourcePath) String() string { return string(p) }
|
||||||
|
func (r AgeRecipient) String() string { return string(r) }
|
||||||
|
func (e S3Endpoint) String() string { return string(e) }
|
||||||
|
func (b BucketName) String() string { return string(b) }
|
||||||
|
func (p S3Prefix) String() string { return string(p) }
|
||||||
|
func (r AWSRegion) String() string { return string(r) }
|
||||||
|
func (k AWSAccessKeyID) String() string { return string(k) }
|
||||||
func (h Hostname) String() string { return string(h) }
|
func (h Hostname) String() string { return string(h) }
|
||||||
func (v Version) String() string { return string(v) }
|
func (v Version) String() string { return string(v) }
|
||||||
func (r GitRevision) String() string { return string(r) }
|
func (r GitRevision) String() string { return string(r) }
|
||||||
func (p GlobPattern) String() string { return string(p) }
|
func (p GlobPattern) String() string { return string(p) }
|
||||||
|
|
||||||
|
// Redacted String methods for sensitive types - prevents accidental logging
|
||||||
|
|
||||||
|
func (k AgeSecretKey) String() string { return "[REDACTED]" }
|
||||||
|
func (k AWSSecretAccessKey) String() string { return "[REDACTED]" }
|
||||||
|
|
||||||
|
// Raw returns the actual value for sensitive types when explicitly needed.
|
||||||
|
func (k AgeSecretKey) Raw() string { return string(k) }
|
||||||
|
|
||||||
|
// Raw returns the actual value for sensitive types when explicitly needed.
|
||||||
|
func (k AWSSecretAccessKey) Raw() string { return string(k) }
|
||||||
|
|||||||
@@ -1,136 +0,0 @@
|
|||||||
package types_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"database/sql"
|
|
||||||
"database/sql/driver"
|
|
||||||
"fmt"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/types"
|
|
||||||
)
|
|
||||||
|
|
||||||
// scannableID is the shared behaviour of the UUID-backed id types. A pointer
|
|
||||||
// to FileID or BlobID satisfies it, so both are tested through one set of
|
|
||||||
// cases.
|
|
||||||
type scannableID interface {
|
|
||||||
driver.Valuer
|
|
||||||
sql.Scanner
|
|
||||||
fmt.Stringer
|
|
||||||
IsZero() bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// idKind adapts one id type to the generic tests below.
|
|
||||||
type idKind struct {
|
|
||||||
name string
|
|
||||||
newZero func() scannableID
|
|
||||||
newRandom func() scannableID
|
|
||||||
parse func(string) (scannableID, error)
|
|
||||||
}
|
|
||||||
|
|
||||||
func idKinds() []idKind {
|
|
||||||
return []idKind{
|
|
||||||
{
|
|
||||||
name: "FileID",
|
|
||||||
newZero: func() scannableID { return &types.FileID{} },
|
|
||||||
newRandom: func() scannableID {
|
|
||||||
id := types.NewFileID()
|
|
||||||
|
|
||||||
return &id
|
|
||||||
},
|
|
||||||
parse: func(s string) (scannableID, error) {
|
|
||||||
id, err := types.ParseFileID(s)
|
|
||||||
|
|
||||||
return &id, err
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "BlobID",
|
|
||||||
newZero: func() scannableID { return &types.BlobID{} },
|
|
||||||
newRandom: func() scannableID {
|
|
||||||
id := types.NewBlobID()
|
|
||||||
|
|
||||||
return &id
|
|
||||||
},
|
|
||||||
parse: func(s string) (scannableID, error) {
|
|
||||||
id, err := types.ParseBlobID(s)
|
|
||||||
|
|
||||||
return &id, err
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIDValueScan checks that Value then Scan round trips from both a string
|
|
||||||
// and a []byte, that a NULL scans to the zero id, and that a non-string type
|
|
||||||
// and malformed text are rejected.
|
|
||||||
func TestIDValueScan(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
for _, k := range idKinds() {
|
|
||||||
t.Run(k.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
orig := k.newRandom()
|
|
||||||
v, err := orig.Value()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
s, ok := v.(string)
|
|
||||||
require.True(t, ok, "Value must yield a string")
|
|
||||||
|
|
||||||
fromString := k.newZero()
|
|
||||||
require.NoError(t, fromString.Scan(s))
|
|
||||||
assert.Equal(t, orig.String(), fromString.String())
|
|
||||||
assert.False(t, fromString.IsZero())
|
|
||||||
|
|
||||||
fromBytes := k.newZero()
|
|
||||||
require.NoError(t, fromBytes.Scan([]byte(s)))
|
|
||||||
assert.Equal(t, orig.String(), fromBytes.String())
|
|
||||||
|
|
||||||
nulled := k.newRandom()
|
|
||||||
require.NoError(t, nulled.Scan(nil))
|
|
||||||
assert.True(t, nulled.IsZero(), "NULL scans to the zero id")
|
|
||||||
|
|
||||||
require.Error(t, k.newZero().Scan(42),
|
|
||||||
"a non-string type must be rejected")
|
|
||||||
assert.Error(t, k.newZero().Scan("not-a-uuid"),
|
|
||||||
"malformed text must be rejected")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIDParse checks that the Parse function accepts a canonical id and
|
|
||||||
// rejects malformed text.
|
|
||||||
func TestIDParse(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
for _, k := range idKinds() {
|
|
||||||
t.Run(k.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
canonical := k.newRandom().String()
|
|
||||||
|
|
||||||
parsed, err := k.parse(canonical)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, canonical, parsed.String())
|
|
||||||
|
|
||||||
_, err = k.parse("not-a-uuid")
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestIDIsZero checks IsZero on the zero and on a freshly generated id.
|
|
||||||
func TestIDIsZero(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
for _, k := range idKinds() {
|
|
||||||
t.Run(k.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
assert.True(t, k.newZero().IsZero())
|
|
||||||
assert.False(t, k.newRandom().IsZero())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -2,6 +2,7 @@ package vaultik
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
@@ -17,13 +18,6 @@ import (
|
|||||||
// not match the expected double-SHA-256 hash.
|
// not match the expected double-SHA-256 hash.
|
||||||
var errBlobHashMismatch = errors.New("blob hash mismatch")
|
var errBlobHashMismatch = errors.New("blob hash mismatch")
|
||||||
|
|
||||||
// errBlobNotFullyRead is returned when the verifying reader is closed
|
|
||||||
// before its plaintext reached EOF. The hash can only be checked once
|
|
||||||
// the whole stream has been read, so an early or short-read close must
|
|
||||||
// fail rather than silently skip verification.
|
|
||||||
var errBlobNotFullyRead = errors.New(
|
|
||||||
"blob closed before fully read; hash not verified")
|
|
||||||
|
|
||||||
// hashVerifyReader wraps a blobgen.Reader and verifies the double-SHA-256 hash
|
// hashVerifyReader wraps a blobgen.Reader and verifies the double-SHA-256 hash
|
||||||
// of decrypted plaintext when Close is called. It reuses the hash that
|
// of decrypted plaintext when Close is called. It reuses the hash that
|
||||||
// blobgen.Reader already computes internally via its TeeReader, avoiding
|
// blobgen.Reader already computes internally via its TeeReader, avoiding
|
||||||
@@ -44,23 +38,22 @@ func (h *hashVerifyReader) Read(p []byte) (int, error) {
|
|||||||
return n, err
|
return n, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close closes the underlying readers and verifies the blob hash. The
|
// Close verifies the hash (if the stream was fully read) and closes underlying readers.
|
||||||
// hash check cannot be skipped: closing before the plaintext reached
|
|
||||||
// EOF (a short read or an early close) is an error, so a caller can
|
|
||||||
// never obtain unverified blob bytes.
|
|
||||||
func (h *hashVerifyReader) Close() error {
|
func (h *hashVerifyReader) Close() error {
|
||||||
readerErr := h.reader.Close()
|
readerErr := h.reader.Close()
|
||||||
fetcherErr := h.fetcher.Close()
|
fetcherErr := h.fetcher.Close()
|
||||||
|
|
||||||
if !h.done {
|
if h.done {
|
||||||
return errBlobNotFullyRead
|
firstHash := h.reader.Sum256()
|
||||||
}
|
secondHasher := sha256.New()
|
||||||
|
secondHasher.Write(firstHash)
|
||||||
|
|
||||||
actualHashHex := hex.EncodeToString(blobgen.DoubleSHA256(h.reader.Sum256()))
|
actualHashHex := hex.EncodeToString(secondHasher.Sum(nil))
|
||||||
if actualHashHex != h.blobHash {
|
if actualHashHex != h.blobHash {
|
||||||
return fmt.Errorf("%w: expected %s, got %s",
|
return fmt.Errorf("%w: expected %s, got %s",
|
||||||
errBlobHashMismatch, h.blobHash[:16], actualHashHex[:16])
|
errBlobHashMismatch, h.blobHash[:16], actualHashHex[:16])
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if readerErr != nil {
|
if readerErr != nil {
|
||||||
return readerErr
|
return readerErr
|
||||||
@@ -74,15 +67,14 @@ func (h *hashVerifyReader) Close() error {
|
|||||||
// The hash is verified when the returned reader is closed (after fully reading).
|
// The hash is verified when the returned reader is closed (after fully reading).
|
||||||
// This avoids buffering the entire blob in memory.
|
// This avoids buffering the entire blob in memory.
|
||||||
func (v *Vaultik) FetchAndDecryptBlob(
|
func (v *Vaultik) FetchAndDecryptBlob(
|
||||||
ctx context.Context, blobHash string, expectedSize int64,
|
ctx context.Context, blobHash string, expectedSize int64, identity age.Identity,
|
||||||
identities ...age.Identity,
|
|
||||||
) (io.ReadCloser, error) {
|
) (io.ReadCloser, error) {
|
||||||
rc, _, err := v.FetchBlob(ctx, blobHash, expectedSize)
|
rc, _, err := v.FetchBlob(ctx, blobHash, expectedSize)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
reader, err := blobgen.NewReader(rc, identities...)
|
reader, err := blobgen.NewReader(rc, identity)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = rc.Close()
|
_ = rc.Close()
|
||||||
|
|
||||||
|
|||||||
@@ -40,13 +40,13 @@ func buildHashTestBlob(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Compute the double-SHA-256 hash of the plaintext (matches
|
// Compute the double-SHA-256 hash of the plaintext (matches
|
||||||
// blobgen.Writer.ContentID).
|
// blobgen.Writer.Sum256).
|
||||||
firstHash := sha256.Sum256(plaintext)
|
firstHash := sha256.Sum256(plaintext)
|
||||||
secondHash := sha256.Sum256(firstHash[:])
|
secondHash := sha256.Sum256(firstHash[:])
|
||||||
correctHash := hex.EncodeToString(secondHash[:])
|
correctHash := hex.EncodeToString(secondHash[:])
|
||||||
|
|
||||||
// Verify our hash matches what blobgen.Writer produces
|
// Verify our hash matches what blobgen.Writer produces
|
||||||
writerHash := hex.EncodeToString(writer.ContentID())
|
writerHash := hex.EncodeToString(writer.Sum256())
|
||||||
if correctHash != writerHash {
|
if correctHash != writerHash {
|
||||||
t.Fatalf("hash computation mismatch: manual=%s, writer=%s",
|
t.Fatalf("hash computation mismatch: manual=%s, writer=%s",
|
||||||
correctHash, writerHash)
|
correctHash, writerHash)
|
||||||
@@ -133,51 +133,3 @@ func TestFetchAndDecryptBlobVerifiesHash(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestFetchAndDecryptBlobCloseBeforeEOFFails verifies the hash check
|
|
||||||
// cannot be skipped: a caller that reads only part of the blob and then
|
|
||||||
// closes gets an error rather than silently unverified bytes.
|
|
||||||
func TestFetchAndDecryptBlobCloseBeforeEOFFails(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
identity, err := age.GenerateX25519Identity()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("generating identity: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
plaintext := []byte("hello world test data for blob hash verification")
|
|
||||||
encryptedData, correctHash := buildHashTestBlob(t, identity, plaintext)
|
|
||||||
|
|
||||||
mockStorage := NewMockStorer()
|
|
||||||
blobPath := "blobs/" + correctHash[:2] + "/" +
|
|
||||||
correctHash[2:4] + "/" + correctHash
|
|
||||||
|
|
||||||
mockStorage.mu.Lock()
|
|
||||||
mockStorage.data[blobPath] = encryptedData
|
|
||||||
mockStorage.mu.Unlock()
|
|
||||||
|
|
||||||
tv := vaultik.NewForTesting(mockStorage)
|
|
||||||
|
|
||||||
rc, err := tv.FetchAndDecryptBlob(
|
|
||||||
context.Background(), correctHash, int64(len(encryptedData)), identity)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("unexpected error opening stream: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read one byte, far short of the plaintext length, then close.
|
|
||||||
buf := make([]byte, 1)
|
|
||||||
|
|
||||||
_, err = rc.Read(buf)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("reading first byte: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = rc.Close()
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("expected error closing before EOF, got nil")
|
|
||||||
}
|
|
||||||
|
|
||||||
if !strings.Contains(err.Error(), "hash not verified") {
|
|
||||||
t.Fatalf("expected not-verified error, got: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,647 +0,0 @@
|
|||||||
package vaultik_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage/faultstore"
|
|
||||||
"sneak.berlin/go/vaultik/internal/ui"
|
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
|
||||||
)
|
|
||||||
|
|
||||||
// These tests cover the failure modes a backup tool must survive:
|
|
||||||
// interrupted uploads, an interrupted metadata export, corrupt and
|
|
||||||
// truncated reads, a full restore disk, and a backend that reports
|
|
||||||
// success while storing nothing. Faults are injected through the
|
|
||||||
// storage.Storer seam (internal/storage/faultstore), never by patching
|
|
||||||
// production code. Each test asserts on the observable end state — what
|
|
||||||
// is in the index, what is at the destination, what the user is told —
|
|
||||||
// not merely that an error was returned. See
|
|
||||||
// https://git.eeqj.de/sneak/vaultik/issues/72.
|
|
||||||
//
|
|
||||||
// Object-level write atomicity (no partial blob object left behind) is
|
|
||||||
// covered by the file:// backend's atomic-write work
|
|
||||||
// (https://git.eeqj.de/sneak/vaultik/issues/130) and is not re-tested
|
|
||||||
// here; these tests target the layers above the backend.
|
|
||||||
//
|
|
||||||
// The tests run serially, not with t.Parallel: each calls
|
|
||||||
// log.Initialize, which replaces the package-global logger, and a
|
|
||||||
// backup or restore running concurrently reads that same logger. Under
|
|
||||||
// -race the two collide. Running one at a time is the same choice
|
|
||||||
// prune_count_test.go already makes for the same reason.
|
|
||||||
|
|
||||||
const (
|
|
||||||
faultChunkSize = int64(64 * 1024)
|
|
||||||
faultMaxBlobSize = int64(256 * 1024)
|
|
||||||
)
|
|
||||||
|
|
||||||
// faultTestConfig returns the config shared by the fault-injection
|
|
||||||
// tests: a real recipient/secret keypair so blobs are genuinely
|
|
||||||
// encrypted, and a blob size limit the restore sweeper can divide.
|
|
||||||
func faultTestConfig() *config.Config {
|
|
||||||
return &config.Config{
|
|
||||||
AgeRecipients: []string{testAgePublicKey},
|
|
||||||
AgeSecretKey: testAgeSecretKey,
|
|
||||||
CompressionLevel: 3,
|
|
||||||
Hostname: testHostname,
|
|
||||||
BlobSizeLimit: config.Size(faultMaxBlobSize),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeFaultSourceTree writes a spread of file sizes that forces several
|
|
||||||
// chunks across more than one blob, so a fault landing on a single blob
|
|
||||||
// still leaves other data intact. Returns the expected content by path.
|
|
||||||
func writeFaultSourceTree(
|
|
||||||
t *testing.T, fs afero.Fs, dataDir string,
|
|
||||||
) map[string][]byte {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
files := map[string][]byte{
|
|
||||||
filepath.Join(dataDir, "small.txt"): []byte("hello vaultik"),
|
|
||||||
filepath.Join(dataDir, "a.bin"): bytesPattern("a-", int(faultChunkSize*3)),
|
|
||||||
filepath.Join(dataDir, "sub", "b.bin"): bytesPattern("b-", int(faultChunkSize*3)),
|
|
||||||
filepath.Join(dataDir, "sub", "c.bin"): bytesPattern("c-", int(faultChunkSize*2)),
|
|
||||||
}
|
|
||||||
|
|
||||||
for path, content := range files {
|
|
||||||
require.NoError(t, fs.MkdirAll(filepath.Dir(path), 0o755))
|
|
||||||
require.NoError(t, afero.WriteFile(fs, path, content, 0o644))
|
|
||||||
}
|
|
||||||
|
|
||||||
return files
|
|
||||||
}
|
|
||||||
|
|
||||||
// newFaultScanner builds a scanner writing through the given storer.
|
|
||||||
func newFaultScanner(
|
|
||||||
fs afero.Fs, storer storage.Storer,
|
|
||||||
cfg *config.Config, repos *database.Repositories,
|
|
||||||
) *snapshot.Scanner {
|
|
||||||
return snapshot.NewScanner(snapshot.ScannerConfig{
|
|
||||||
FS: fs,
|
|
||||||
Storage: storer,
|
|
||||||
ChunkSize: faultChunkSize,
|
|
||||||
MaxBlobSize: faultMaxBlobSize,
|
|
||||||
CompressionLevel: cfg.CompressionLevel,
|
|
||||||
AgeRecipients: cfg.AgeRecipients,
|
|
||||||
Repositories: repos,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// newFaultSnapshotManager builds a snapshot manager writing through the
|
|
||||||
// given storer.
|
|
||||||
func newFaultSnapshotManager(
|
|
||||||
fs afero.Fs, storer storage.Storer,
|
|
||||||
cfg *config.Config, repos *database.Repositories,
|
|
||||||
) *snapshot.SnapshotManager {
|
|
||||||
sm := snapshot.NewSnapshotManager(snapshot.SnapshotManagerParams{
|
|
||||||
Repos: repos,
|
|
||||||
Storage: storer,
|
|
||||||
Config: cfg,
|
|
||||||
})
|
|
||||||
sm.SetFilesystem(fs)
|
|
||||||
|
|
||||||
return sm
|
|
||||||
}
|
|
||||||
|
|
||||||
// fullFaultBackup runs a complete backup (create, scan, complete,
|
|
||||||
// export) through storer and returns the snapshot ID.
|
|
||||||
func fullFaultBackup(
|
|
||||||
ctx context.Context, t *testing.T, fs afero.Fs, storer storage.Storer,
|
|
||||||
cfg *config.Config, repos *database.Repositories,
|
|
||||||
dataDir, dbPath, name string,
|
|
||||||
) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
sm := newFaultSnapshotManager(fs, storer, cfg, repos)
|
|
||||||
scanner := newFaultScanner(fs, storer, cfg, repos)
|
|
||||||
|
|
||||||
id, err := sm.CreateSnapshotWithName(ctx, cfg.Hostname, name, "v", "g")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = scanner.Scan(ctx, dataDir, id)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
require.NoError(t, sm.CompleteSnapshot(ctx, id))
|
|
||||||
require.NoError(t, sm.ExportSnapshotMetadata(ctx, dbPath, id))
|
|
||||||
|
|
||||||
return id
|
|
||||||
}
|
|
||||||
|
|
||||||
// newReaderVaultik builds a Vaultik that reads (restore/verify) through
|
|
||||||
// storer, with the given repositories (nil is fine for restore/verify,
|
|
||||||
// which read metadata from storage).
|
|
||||||
func newReaderVaultik(
|
|
||||||
ctx context.Context, cfg *config.Config, storer storage.Storer,
|
|
||||||
repos *database.Repositories, fs afero.Fs,
|
|
||||||
) *vaultik.Vaultik {
|
|
||||||
v := &vaultik.Vaultik{
|
|
||||||
Config: cfg,
|
|
||||||
Storage: storer,
|
|
||||||
Repositories: repos,
|
|
||||||
Fs: fs,
|
|
||||||
Stdout: io.Discard,
|
|
||||||
Stderr: io.Discard,
|
|
||||||
UI: ui.NewWithColor(io.Discard, false),
|
|
||||||
}
|
|
||||||
v.SetContext(ctx)
|
|
||||||
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
|
|
||||||
// Scenario 3: a stored blob's bytes are flipped before restore reads
|
|
||||||
// them. Restore must fail loudly, and no file must be left on the
|
|
||||||
// restore target holding corrupt content.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // installs the global logger via log.Initialize
|
|
||||||
func TestRestoreRejectsCorruptBlob(t *testing.T) {
|
|
||||||
assertRestoreRejectsDamagedBlob(t, faultstore.GetCorrupt, "corrupt")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Scenario 4: a stored blob is truncated before restore reads it. Same
|
|
||||||
// contract as the corrupt case.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // installs the global logger via log.Initialize
|
|
||||||
func TestRestoreRejectsTruncatedBlob(t *testing.T) {
|
|
||||||
assertRestoreRejectsDamagedBlob(t, faultstore.GetTruncate, "truncated")
|
|
||||||
}
|
|
||||||
|
|
||||||
// assertRestoreRejectsDamagedBlob backs up the source tree, then restores
|
|
||||||
// through a store that damages every blob read with the given fault, and
|
|
||||||
// asserts restore fails naming a blob and leaves no file on the target
|
|
||||||
// holding wrong bytes. Metadata reads are returned intact so the failure
|
|
||||||
// is isolated to the blob.
|
|
||||||
func assertRestoreRejectsDamagedBlob(
|
|
||||||
t *testing.T, fault faultstore.GetFault, name string,
|
|
||||||
) {
|
|
||||||
t.Helper()
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
dataDir := filepath.Join(tempDir, "src")
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
restoreDir := filepath.Join(tempDir, "restored")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
cfg := faultTestConfig()
|
|
||||||
testFiles := writeFaultSourceTree(t, fs, dataDir)
|
|
||||||
|
|
||||||
inner, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
id := fullFaultBackup(ctx, t, fs, inner, cfg, repos, dataDir, dbPath, name)
|
|
||||||
require.NoError(t, db.Close())
|
|
||||||
|
|
||||||
faultStore := faultstore.New(inner)
|
|
||||||
faultStore.OnGet = func(key string) faultstore.GetFault {
|
|
||||||
if strings.HasPrefix(key, "blobs/") {
|
|
||||||
return fault
|
|
||||||
}
|
|
||||||
|
|
||||||
return faultstore.GetNormal
|
|
||||||
}
|
|
||||||
|
|
||||||
v := newReaderVaultik(ctx, cfg, faultStore, nil, fs)
|
|
||||||
err = v.Restore(&vaultik.RestoreOptions{SnapshotID: id, TargetDir: restoreDir})
|
|
||||||
|
|
||||||
require.Error(t, err, "restore must fail on a damaged blob")
|
|
||||||
assert.Contains(t, err.Error(), "blob",
|
|
||||||
"error should name the blob that failed")
|
|
||||||
assertNoCorruptFiles(t, fs, restoreDir, testFiles)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Scenario 6: the backend accepts blob uploads and reports success but
|
|
||||||
// stores nothing. verify --deep must catch it.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // installs the global logger via log.Initialize
|
|
||||||
func TestDeepVerifyCatchesLyingBackend(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
dataDir := filepath.Join(tempDir, "src")
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
cfg := faultTestConfig()
|
|
||||||
|
|
||||||
writeFaultSourceTree(t, fs, dataDir)
|
|
||||||
|
|
||||||
inner, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Blob uploads are swallowed; metadata uploads land, so verify can
|
|
||||||
// download the manifest and database and then discover the blobs are
|
|
||||||
// absent.
|
|
||||||
lying := faultstore.New(inner)
|
|
||||||
lying.OnPut = func(key string) faultstore.PutAction {
|
|
||||||
if strings.HasPrefix(key, "blobs/") {
|
|
||||||
return faultstore.PutSwallow
|
|
||||||
}
|
|
||||||
|
|
||||||
return faultstore.PutNormal
|
|
||||||
}
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
id := fullFaultBackup(ctx, t, fs, lying, cfg, repos, dataDir, dbPath, "lying")
|
|
||||||
require.NoError(t, db.Close())
|
|
||||||
|
|
||||||
// No blob objects were actually written.
|
|
||||||
blobKeys, err := inner.List(ctx, "blobs/")
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, blobKeys, "lying backend should have stored no blobs")
|
|
||||||
|
|
||||||
// Read back through the honest underlying store.
|
|
||||||
v := newReaderVaultik(ctx, cfg, inner, nil, fs)
|
|
||||||
err = v.VerifySnapshotWithOptions(id, &vaultik.VerifyOptions{Deep: true})
|
|
||||||
require.Error(t, err, "deep verify must catch a backend that stored nothing")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Scenario 1a: a blob upload fails partway through. The interrupted run
|
|
||||||
// must not record the blob as uploaded, must not reference it from the
|
|
||||||
// snapshot, and must leave no blob object at the destination.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // installs the global logger via log.Initialize
|
|
||||||
func TestInterruptedBlobUploadRecordsNoUploadedBlob(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
dataDir := filepath.Join(tempDir, "src")
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
cfg := faultTestConfig()
|
|
||||||
|
|
||||||
writeFaultSourceTree(t, fs, dataDir)
|
|
||||||
|
|
||||||
inner, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
defer func() { _ = db.Close() }()
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
// Every blob upload fails partway through. The scan must surface it.
|
|
||||||
fault := faultstore.New(inner)
|
|
||||||
fault.OnPut = func(key string) faultstore.PutAction {
|
|
||||||
if strings.HasPrefix(key, "blobs/") {
|
|
||||||
return faultstore.PutFail
|
|
||||||
}
|
|
||||||
|
|
||||||
return faultstore.PutNormal
|
|
||||||
}
|
|
||||||
|
|
||||||
sm := newFaultSnapshotManager(fs, fault, cfg, repos)
|
|
||||||
scanner := newFaultScanner(fs, fault, cfg, repos)
|
|
||||||
|
|
||||||
id, err := sm.CreateSnapshotWithName(ctx, cfg.Hostname, "interrupted", "v", "g")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = scanner.Scan(ctx, dataDir, id)
|
|
||||||
require.Error(t, err, "scan must fail when a blob upload fails")
|
|
||||||
|
|
||||||
// No blob may claim to be uploaded.
|
|
||||||
blobs, err := repos.Blobs.GetAll(ctx)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
for _, b := range blobs {
|
|
||||||
assert.Nilf(t, b.UploadedTS,
|
|
||||||
"blob %s marked uploaded after a failed upload", b.Hash)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The snapshot may reference no blobs, and the destination holds none.
|
|
||||||
hashes, err := repos.Snapshots.GetBlobHashes(ctx, id)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, hashes, "interrupted snapshot must reference no blobs")
|
|
||||||
|
|
||||||
blobKeys, err := inner.List(ctx, "blobs/")
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Empty(t, blobKeys, "no blob object may survive at the destination")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Scenario 1b: after an interrupted upload, a retry on the same local
|
|
||||||
// index must produce a restorable snapshot. The interrupted run leaves
|
|
||||||
// the blob's chunk rows in the index; the fix for
|
|
||||||
// https://git.eeqj.de/sneak/vaultik/issues/148 discards those un-uploaded
|
|
||||||
// blob rows at the start of the next scan and deduplicates only against
|
|
||||||
// chunks in a blob that was actually uploaded, so the retry re-chunks and
|
|
||||||
// re-uploads the affected data instead of silently referencing data that
|
|
||||||
// never reached storage.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // installs the global logger via log.Initialize
|
|
||||||
func TestBackupRetryAfterInterruptedUploadIsRestorable(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
dataDir := filepath.Join(tempDir, "src")
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
restoreDir := filepath.Join(tempDir, "restored")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
cfg := faultTestConfig()
|
|
||||||
testFiles := writeFaultSourceTree(t, fs, dataDir)
|
|
||||||
|
|
||||||
inner, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
// Attempt 1: every blob upload fails.
|
|
||||||
fault := faultstore.New(inner)
|
|
||||||
fault.OnPut = func(key string) faultstore.PutAction {
|
|
||||||
if strings.HasPrefix(key, "blobs/") {
|
|
||||||
return faultstore.PutFail
|
|
||||||
}
|
|
||||||
|
|
||||||
return faultstore.PutNormal
|
|
||||||
}
|
|
||||||
|
|
||||||
sm := newFaultSnapshotManager(fs, fault, cfg, repos)
|
|
||||||
scanner := newFaultScanner(fs, fault, cfg, repos)
|
|
||||||
|
|
||||||
id1, err := sm.CreateSnapshotWithName(ctx, cfg.Hostname, "interrupted", "v", "g")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = scanner.Scan(ctx, dataDir, id1)
|
|
||||||
require.Error(t, err)
|
|
||||||
|
|
||||||
// Retry on the same local index with a working backend.
|
|
||||||
id2 := fullFaultBackup(ctx, t, fs, inner, cfg, repos, dataDir, dbPath, "retry")
|
|
||||||
require.NoError(t, db.Close())
|
|
||||||
|
|
||||||
v := newReaderVaultik(ctx, cfg, inner, nil, fs)
|
|
||||||
require.NoError(t, v.Restore(&vaultik.RestoreOptions{
|
|
||||||
SnapshotID: id2,
|
|
||||||
TargetDir: restoreDir,
|
|
||||||
Verify: true,
|
|
||||||
}), "retry after an interrupted upload must produce a restorable snapshot")
|
|
||||||
|
|
||||||
assertRestoredTree(t, fs, restoreDir, testFiles)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Scenario 2: the process dies during the metadata export, after the
|
|
||||||
// database is uploaded but before the manifest. The destination is left
|
|
||||||
// with blobs and a database but no manifest. verify and snapshot list
|
|
||||||
// must report the damage honestly rather than crashing or passing.
|
|
||||||
// Automatic detection and repair of this partial state on the next run
|
|
||||||
// is tracked in https://git.eeqj.de/sneak/vaultik/issues/177 and is not
|
|
||||||
// asserted here.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // installs the global logger via log.Initialize
|
|
||||||
func TestBackupSurvivesMetadataExportInterruption(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
dataDir := filepath.Join(tempDir, "src")
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
cfg := faultTestConfig()
|
|
||||||
|
|
||||||
writeFaultSourceTree(t, fs, dataDir)
|
|
||||||
|
|
||||||
inner, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
// Back up and complete with a working backend.
|
|
||||||
sm := newFaultSnapshotManager(fs, inner, cfg, repos)
|
|
||||||
scanner := newFaultScanner(fs, inner, cfg, repos)
|
|
||||||
|
|
||||||
id, err := sm.CreateSnapshotWithName(ctx, cfg.Hostname, "export", "v", "g")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = scanner.Scan(ctx, dataDir, id)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, sm.CompleteSnapshot(ctx, id))
|
|
||||||
|
|
||||||
// Export through a backend that fails only the manifest upload. The
|
|
||||||
// database uploads first and lands; the manifest does not.
|
|
||||||
fault := faultstore.New(inner)
|
|
||||||
fault.OnPut = func(key string) faultstore.PutAction {
|
|
||||||
if strings.HasSuffix(key, "manifest.json.zst") {
|
|
||||||
return faultstore.PutFail
|
|
||||||
}
|
|
||||||
|
|
||||||
return faultstore.PutNormal
|
|
||||||
}
|
|
||||||
|
|
||||||
smFault := newFaultSnapshotManager(fs, fault, cfg, repos)
|
|
||||||
|
|
||||||
err = smFault.ExportSnapshotMetadata(ctx, dbPath, id)
|
|
||||||
require.Error(t, err, "export must fail when the manifest upload fails")
|
|
||||||
|
|
||||||
// The destination is in the partial state the scenario describes.
|
|
||||||
key := snapshot.RemoteSnapshotKey(id)
|
|
||||||
|
|
||||||
_, err = inner.Stat(ctx, "metadata/"+key+"/db.zst.age")
|
|
||||||
require.NoError(t, err, "database should have been uploaded before the manifest")
|
|
||||||
|
|
||||||
_, err = inner.Stat(ctx, "metadata/"+key+"/manifest.json.zst")
|
|
||||||
require.ErrorIs(t, err, storage.ErrNotFound, "manifest upload should not have landed")
|
|
||||||
|
|
||||||
// verify must fail loudly for this snapshot, in both modes.
|
|
||||||
reader := newReaderVaultik(ctx, cfg, inner, repos, fs)
|
|
||||||
|
|
||||||
deepOpts := &vaultik.VerifyOptions{Deep: true}
|
|
||||||
require.Error(t, reader.VerifySnapshotWithOptions(id, deepOpts),
|
|
||||||
"deep verify must report the missing manifest")
|
|
||||||
|
|
||||||
shallowOpts := &vaultik.VerifyOptions{Deep: false}
|
|
||||||
require.Error(t, reader.VerifySnapshotWithOptions(id, shallowOpts),
|
|
||||||
"shallow verify must report the missing manifest")
|
|
||||||
|
|
||||||
// snapshot list must not crash on the partial snapshot.
|
|
||||||
require.NoError(t, reader.ListSnapshots(false),
|
|
||||||
"snapshot list must tolerate a partially-exported snapshot")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Scenario 5: the restore target runs out of space mid-file. Restore
|
|
||||||
// must fail with an out-of-space error, and must not leave a truncated
|
|
||||||
// file at the target path presenting as a complete restore. Restore
|
|
||||||
// today writes each file straight to its final path and does not remove
|
|
||||||
// it when a write fails, so the truncated file survives; deleting it is
|
|
||||||
// tracked by https://git.eeqj.de/sneak/vaultik/issues/163. Skipped until
|
|
||||||
// that lands, so the destination assertion below is recorded rather than
|
|
||||||
// dropped.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // installs the global logger via log.Initialize
|
|
||||||
func TestRestoreReportsDiskFull(t *testing.T) {
|
|
||||||
t.Skip("blocked on https://git.eeqj.de/sneak/vaultik/issues/163: " +
|
|
||||||
"a disk-full write leaves a truncated file at the target path " +
|
|
||||||
"instead of removing it")
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
|
|
||||||
osFS := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
dataDir := filepath.Join(tempDir, "src")
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
restoreDir := filepath.Join(tempDir, "restored")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
cfg := faultTestConfig()
|
|
||||||
|
|
||||||
testFiles := writeFaultSourceTree(t, osFS, dataDir)
|
|
||||||
|
|
||||||
inner, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
id := fullFaultBackup(ctx, t, osFS, inner, cfg, repos, dataDir, dbPath, "diskfull")
|
|
||||||
require.NoError(t, db.Close())
|
|
||||||
|
|
||||||
// Restore onto a filesystem that allows only a few bytes of file
|
|
||||||
// content: enough to create files, far too little to hold them.
|
|
||||||
budget := int64(8)
|
|
||||||
quota := "aFS{Fs: osFS, remaining: &budget}
|
|
||||||
|
|
||||||
v := newReaderVaultik(ctx, cfg, inner, nil, quota)
|
|
||||||
err = v.Restore(&vaultik.RestoreOptions{SnapshotID: id, TargetDir: restoreDir})
|
|
||||||
|
|
||||||
require.Error(t, err, "restore must fail when the target disk is full")
|
|
||||||
assert.Contains(t, err.Error(), errNoSpace.Error(),
|
|
||||||
"restore error should surface the out-of-space cause")
|
|
||||||
|
|
||||||
// The failure must not leave a truncated file behind presenting as a
|
|
||||||
// complete restore: any file at the target must hold the original
|
|
||||||
// bytes, or be absent.
|
|
||||||
assertNoCorruptFiles(t, osFS, restoreDir, testFiles)
|
|
||||||
}
|
|
||||||
|
|
||||||
// assertRestoredTree byte-compares every restored file against the
|
|
||||||
// original.
|
|
||||||
func assertRestoredTree(
|
|
||||||
t *testing.T, fs afero.Fs, restoreDir string, testFiles map[string][]byte,
|
|
||||||
) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
for origPath, expected := range testFiles {
|
|
||||||
restoredPath := filepath.Join(restoreDir, origPath)
|
|
||||||
got, err := afero.ReadFile(fs, restoredPath)
|
|
||||||
require.NoErrorf(t, err, "restored file missing: %s", origPath)
|
|
||||||
require.Equalf(t, expected, got, "restored content mismatch for %s", origPath)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// errNoSpace is the out-of-space error quotaFS returns once its byte
|
|
||||||
// budget is exhausted, mirroring a real ENOSPC.
|
|
||||||
var errNoSpace = errors.New("no space left on device")
|
|
||||||
|
|
||||||
// quotaFS is an afero.Fs whose files may write only a fixed total number
|
|
||||||
// of content bytes before failing, simulating a full restore target. It
|
|
||||||
// wraps the interface so every method except Create delegates to the
|
|
||||||
// real filesystem; only file writes are capped.
|
|
||||||
type quotaFS struct {
|
|
||||||
afero.Fs
|
|
||||||
|
|
||||||
remaining *int64
|
|
||||||
}
|
|
||||||
|
|
||||||
//nolint:ireturn // afero.Fs.Create's signature requires returning afero.File.
|
|
||||||
func (q *quotaFS) Create(name string) (afero.File, error) {
|
|
||||||
f, err := q.Fs.Create(name)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return "aFile{File: f, remaining: q.remaining}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// quotaFile fails writes once the shared byte budget is exhausted.
|
|
||||||
type quotaFile struct {
|
|
||||||
afero.File
|
|
||||||
|
|
||||||
remaining *int64
|
|
||||||
}
|
|
||||||
|
|
||||||
func (q *quotaFile) Write(p []byte) (int, error) {
|
|
||||||
if *q.remaining <= 0 {
|
|
||||||
return 0, errNoSpace
|
|
||||||
}
|
|
||||||
|
|
||||||
allowed := min(int64(len(p)), *q.remaining)
|
|
||||||
|
|
||||||
n, err := q.File.Write(p[:allowed])
|
|
||||||
*q.remaining -= int64(n)
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
return n, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if int64(n) < int64(len(p)) {
|
|
||||||
return n, errNoSpace
|
|
||||||
}
|
|
||||||
|
|
||||||
return n, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// assertNoCorruptFiles fails if any file that made it to the restore
|
|
||||||
// target holds content that differs from the original: a failed restore
|
|
||||||
// may leave a file absent, but must never leave wrong bytes presenting
|
|
||||||
// as the real file.
|
|
||||||
func assertNoCorruptFiles(
|
|
||||||
t *testing.T, fs afero.Fs, restoreDir string, testFiles map[string][]byte,
|
|
||||||
) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
for origPath, expected := range testFiles {
|
|
||||||
restoredPath := filepath.Join(restoreDir, origPath)
|
|
||||||
|
|
||||||
got, err := afero.ReadFile(fs, restoredPath)
|
|
||||||
if err != nil {
|
|
||||||
if os.IsNotExist(err) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
assert.Equalf(t, expected, got,
|
|
||||||
"restored file %s holds corrupt content", origPath)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -167,12 +167,6 @@ func (v *Vaultik) PruneBlobs(opts *PruneOptions) error {
|
|||||||
|
|
||||||
// collectReferencedBlobs downloads all manifests and returns the set of
|
// collectReferencedBlobs downloads all manifests and returns the set of
|
||||||
// referenced blob hashes.
|
// referenced blob hashes.
|
||||||
//
|
|
||||||
// Every manifest must be read successfully. A manifest that cannot be
|
|
||||||
// downloaded or decoded means its snapshot's blobs are unknown, so
|
|
||||||
// treating them as unreferenced would let prune delete data a snapshot
|
|
||||||
// still needs. Rather than risk that silent loss, any failure returns an
|
|
||||||
// error naming the remote key and prune deletes nothing.
|
|
||||||
func (v *Vaultik) collectReferencedBlobs() (map[string]bool, error) {
|
func (v *Vaultik) collectReferencedBlobs() (map[string]bool, error) {
|
||||||
log.Info("Listing remote snapshots")
|
log.Info("Listing remote snapshots")
|
||||||
// IDs returned by listUniqueSnapshotIDs are remote keys (hashed
|
// IDs returned by listUniqueSnapshotIDs are remote keys (hashed
|
||||||
@@ -185,22 +179,27 @@ func (v *Vaultik) collectReferencedBlobs() (map[string]bool, error) {
|
|||||||
log.Info("Found manifests in remote storage", "count", len(remoteKeys))
|
log.Info("Found manifests in remote storage", "count", len(remoteKeys))
|
||||||
|
|
||||||
allBlobsReferenced := make(map[string]bool)
|
allBlobsReferenced := make(map[string]bool)
|
||||||
|
manifestCount := 0
|
||||||
|
|
||||||
for _, remoteKey := range remoteKeys {
|
for _, remoteKey := range remoteKeys {
|
||||||
log.Debug("Processing manifest", "remote_key", remoteKey)
|
log.Debug("Processing manifest", "remote_key", remoteKey)
|
||||||
|
|
||||||
manifest, err := v.downloadManifestByKey(remoteKey)
|
manifest, err := v.downloadManifestByKey(remoteKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("reading manifest %s: %w", remoteKey, err)
|
log.Error("Failed to download manifest", "remote_key", remoteKey, "error", err)
|
||||||
|
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, blob := range manifest.Blobs {
|
for _, blob := range manifest.Blobs {
|
||||||
allBlobsReferenced[blob.Hash] = true
|
allBlobsReferenced[blob.Hash] = true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
manifestCount++
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Info("Processed manifests",
|
log.Info("Processed manifests",
|
||||||
"count", len(remoteKeys), "unique_blobs_referenced", len(allBlobsReferenced))
|
"count", manifestCount, "unique_blobs_referenced", len(allBlobsReferenced))
|
||||||
|
|
||||||
return allBlobsReferenced, nil
|
return allBlobsReferenced, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,47 +0,0 @@
|
|||||||
package vaultik_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestPruneBlobs_UnreadableManifestDeletesNothing is the regression guard
|
|
||||||
// for issue #157: prune identifies referenced blobs by reading every
|
|
||||||
// snapshot's manifest, and a manifest it cannot decode used to be logged
|
|
||||||
// and skipped. Blobs referenced only by that snapshot then looked
|
|
||||||
// unreferenced and were deleted, with a zero exit — silent backup loss,
|
|
||||||
// made worse by `snapshot create --prune` running unattended with force.
|
|
||||||
//
|
|
||||||
// The single blob here is referenced only by the snapshot whose manifest
|
|
||||||
// is corrupt, so the old behaviour would delete it and succeed. Prune
|
|
||||||
// must instead delete nothing and return an error.
|
|
||||||
func TestPruneBlobs_UnreadableManifestDeletesNothing(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
env := newListEnv(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
blobKey := "blobs/" + testBlobHashA[:2] + "/" + testBlobHashA[2:4] +
|
|
||||||
"/" + testBlobHashA
|
|
||||||
require.NoError(t, env.store.Put(ctx, blobKey,
|
|
||||||
bytes.NewReader([]byte("blob-bytes"))))
|
|
||||||
|
|
||||||
// A manifest at the path prune reads, but with contents it cannot
|
|
||||||
// decode.
|
|
||||||
require.NoError(t, env.store.Put(ctx,
|
|
||||||
"metadata/corruptkey/manifest.json.zst",
|
|
||||||
bytes.NewReader([]byte("not a valid manifest"))))
|
|
||||||
|
|
||||||
err := env.v.PruneBlobs(&vaultik.PruneOptions{Force: true})
|
|
||||||
|
|
||||||
require.Error(t, err, "prune must fail when a manifest cannot be read")
|
|
||||||
assert.True(t, env.store.hasKey(blobKey),
|
|
||||||
"no blob may be deleted when a manifest is unreadable")
|
|
||||||
}
|
|
||||||
@@ -1,146 +0,0 @@
|
|||||||
package vaultik_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"database/sql"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
"sneak.berlin/go/vaultik/internal/types"
|
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
|
||||||
)
|
|
||||||
|
|
||||||
// setupConsistencyTest builds a Vaultik whose local database and mock
|
|
||||||
// remote both hold the given snapshots. Remote metadata is stored under
|
|
||||||
// the production layout, metadata/<RemoteSnapshotKey(id)>/manifest.json.zst.
|
|
||||||
// It returns the instance and the mock so a test can inspect the remote.
|
|
||||||
func setupConsistencyTest(
|
|
||||||
t *testing.T, snapshotIDs []string,
|
|
||||||
) (*vaultik.Vaultik, *MockStorer) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
db, err := database.New(ctx, ":memory:")
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = db.Close() })
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
mockStorage := NewMockStorer()
|
|
||||||
|
|
||||||
for _, id := range snapshotIDs {
|
|
||||||
parts := strings.Split(id, "_")
|
|
||||||
startedAt, err := time.Parse(time.RFC3339, parts[len(parts)-1])
|
|
||||||
require.NoError(t, err, "parsing timestamp from snapshot ID %q", id)
|
|
||||||
|
|
||||||
completedAt := startedAt.Add(5 * time.Minute)
|
|
||||||
snap := &database.Snapshot{
|
|
||||||
ID: types.SnapshotID(id),
|
|
||||||
Hostname: testHostname,
|
|
||||||
VaultikVersion: testLabel,
|
|
||||||
StartedAt: startedAt,
|
|
||||||
CompletedAt: &completedAt,
|
|
||||||
}
|
|
||||||
err = repos.WithTx(ctx, func(ctx context.Context, tx *sql.Tx) error {
|
|
||||||
return repos.Snapshots.Create(ctx, tx, snap)
|
|
||||||
})
|
|
||||||
require.NoError(t, err, "creating snapshot %s", id)
|
|
||||||
|
|
||||||
metadataKey := "metadata/" + snapshot.RemoteSnapshotKey(id) +
|
|
||||||
"/manifest.json.zst"
|
|
||||||
err = mockStorage.Put(ctx, metadataKey, strings.NewReader("stub"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
v := &vaultik.Vaultik{
|
|
||||||
Storage: mockStorage,
|
|
||||||
Repositories: repos,
|
|
||||||
DB: db,
|
|
||||||
Stdout: &bytes.Buffer{},
|
|
||||||
Stderr: &bytes.Buffer{},
|
|
||||||
Stdin: &bytes.Buffer{},
|
|
||||||
}
|
|
||||||
v.SetContext(ctx)
|
|
||||||
|
|
||||||
return v, mockStorage
|
|
||||||
}
|
|
||||||
|
|
||||||
func remoteHasSnapshot(t *testing.T, m *MockStorer, id string) bool {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
prefix := "metadata/" + snapshot.RemoteSnapshotKey(id) + "/"
|
|
||||||
keys, err := m.List(context.Background(), prefix)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
return len(keys) > 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestPurgeKeepsRemotelyBackedLocalRows guards against issue #160
|
|
||||||
// (https://git.eeqj.de/sneak/vaultik/issues/160): purge reconciles local
|
|
||||||
// rows against the remote first, and that step compared human snapshot IDs
|
|
||||||
// against the hashed remote directory names, which never match — so it
|
|
||||||
// deleted every local record and the purge itself then removed nothing.
|
|
||||||
//
|
|
||||||
// With every snapshot still present remotely and nothing old enough to
|
|
||||||
// purge, all local rows must survive the reconcile untouched.
|
|
||||||
func TestPurgeKeepsRemotelyBackedLocalRows(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ids := []string{snapHomeT0, snapHomeT1, snapSystemT0}
|
|
||||||
|
|
||||||
v, _ := setupConsistencyTest(t, ids)
|
|
||||||
|
|
||||||
err := v.PurgeSnapshotsWithOptions(&vaultik.SnapshotPurgeOptions{
|
|
||||||
// 100 years: nothing is old enough to delete, so the reconcile
|
|
||||||
// is the only thing that touches the rows.
|
|
||||||
OlderThan: "36500d",
|
|
||||||
Force: true,
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
remaining := listRemainingSnapshots(t, v)
|
|
||||||
assert.Len(t, remaining, len(ids),
|
|
||||||
"remotely-backed local rows must survive the reconcile")
|
|
||||||
assert.Contains(t, remaining, snapHomeT0)
|
|
||||||
assert.Contains(t, remaining, snapHomeT1)
|
|
||||||
assert.Contains(t, remaining, snapSystemT0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestPurgeRemovesLocalAndRemoteTogether proves the two halves stay
|
|
||||||
// consistent: a purged snapshot is gone both locally and remotely, while a
|
|
||||||
// retained one keeps both. Before the fix, the reconcile dropped every
|
|
||||||
// local row yet the remote metadata was left in place.
|
|
||||||
func TestPurgeRemovesLocalAndRemoteTogether(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ids := []string{snapHomeT0, snapHomeT1, snapSystemT0}
|
|
||||||
|
|
||||||
v, mock := setupConsistencyTest(t, ids)
|
|
||||||
|
|
||||||
err := v.PurgeSnapshotsWithOptions(&vaultik.SnapshotPurgeOptions{
|
|
||||||
KeepLatest: true,
|
|
||||||
Force: true,
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// Keep latest per name: newest home and the lone system are kept.
|
|
||||||
remaining := listRemainingSnapshots(t, v)
|
|
||||||
assert.ElementsMatch(t, []string{snapHomeT1, snapSystemT0}, remaining)
|
|
||||||
|
|
||||||
// Local and remote agree: the older home snapshot is gone from both,
|
|
||||||
// the retained ones are present in both.
|
|
||||||
assert.False(t, remoteHasSnapshot(t, mock, snapHomeT0),
|
|
||||||
"purged snapshot must also be removed remotely")
|
|
||||||
assert.True(t, remoteHasSnapshot(t, mock, snapHomeT1),
|
|
||||||
"retained snapshot must remain remotely")
|
|
||||||
assert.True(t, remoteHasSnapshot(t, mock, snapSystemT0),
|
|
||||||
"retained snapshot must remain remotely")
|
|
||||||
}
|
|
||||||
@@ -12,7 +12,6 @@ import (
|
|||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
"sneak.berlin/go/vaultik/internal/database"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
"sneak.berlin/go/vaultik/internal/types"
|
"sneak.berlin/go/vaultik/internal/types"
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
"sneak.berlin/go/vaultik/internal/vaultik"
|
||||||
)
|
)
|
||||||
@@ -61,11 +60,8 @@ func setupPurgeTest(t *testing.T, snapshotIDs []string) *vaultik.Vaultik {
|
|||||||
})
|
})
|
||||||
require.NoError(t, err, "creating snapshot %s", id)
|
require.NoError(t, err, "creating snapshot %s", id)
|
||||||
|
|
||||||
// Create the remote metadata stub under the production layout so
|
// Create remote metadata stub so syncWithRemote keeps it
|
||||||
// syncWithRemote keeps the local row. Production stores metadata
|
metadataKey := "metadata/" + id + "/manifest.json.zst"
|
||||||
// under the hashed remote key, not the human snapshot ID.
|
|
||||||
metadataKey := "metadata/" + snapshot.RemoteSnapshotKey(id) +
|
|
||||||
"/manifest.json.zst"
|
|
||||||
err = mockStorage.Put(ctx, metadataKey, strings.NewReader("stub"))
|
err = mockStorage.Put(ctx, metadataKey, strings.NewReader("stub"))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|||||||
+69
-317
@@ -11,7 +11,6 @@ import (
|
|||||||
"math"
|
"math"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"filippo.io/age"
|
"filippo.io/age"
|
||||||
@@ -19,7 +18,6 @@ import (
|
|||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
"sneak.berlin/go/vaultik/internal/database"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
"sneak.berlin/go/vaultik/internal/types"
|
"sneak.berlin/go/vaultik/internal/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -30,43 +28,19 @@ var (
|
|||||||
errDecryptionKeyRequired = errors.New(
|
errDecryptionKeyRequired = errors.New(
|
||||||
"decryption key required for restore\n\n" +
|
"decryption key required for restore\n\n" +
|
||||||
"Set the VAULTIK_AGE_SECRET_KEY environment variable to your " +
|
"Set the VAULTIK_AGE_SECRET_KEY environment variable to your " +
|
||||||
"age private key file:\n" +
|
"age private key:\n" +
|
||||||
" export VAULTIK_AGE_SECRET_KEY=\"$(cat vaultik_backup_private_key.txt)\"")
|
" export VAULTIK_AGE_SECRET_KEY='AGE-SECRET-KEY-...'")
|
||||||
// errInvalidAgeSecretKey is returned when the configured key does not
|
|
||||||
// parse as any age identity. It names the source but never the value,
|
|
||||||
// which is secret, so the message is safe to print and log.
|
|
||||||
errInvalidAgeSecretKey = errors.New(
|
|
||||||
"configured age secret key holds no usable age identity")
|
|
||||||
errBlobMissingFromIndex = errors.New("blob hash missing from blob index")
|
errBlobMissingFromIndex = errors.New("blob hash missing from blob index")
|
||||||
errChunkNotInAnyBlob = errors.New("chunk not found in any blob")
|
errChunkNotInAnyBlob = errors.New("chunk not found in any blob")
|
||||||
errBlobIDNotInHashIndex = errors.New("blob id missing from hash index")
|
errBlobIDNotInHashIndex = errors.New("blob id missing from hash index")
|
||||||
errShortChunkRead = errors.New("short read")
|
errShortChunkRead = errors.New("short read")
|
||||||
errRestorePathEscapesTarget = errors.New(
|
|
||||||
"refusing to restore path outside the target directory")
|
|
||||||
errTrailingRestoreData = errors.New(
|
|
||||||
"restored file has trailing data after its last chunk")
|
|
||||||
errRestoreIncomplete = errors.New(
|
|
||||||
"restore loop ended with files still pending")
|
|
||||||
errSnapshotDBMismatch = errors.New(
|
|
||||||
"decrypted database is not the requested snapshot")
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// snapshotDBFilename is the name the decrypted snapshot database is
|
|
||||||
// written under inside its private temp directory.
|
|
||||||
const snapshotDBFilename = "snapshot.db"
|
|
||||||
|
|
||||||
// restoreDirMode is the permission mode for directories created while
|
// restoreDirMode is the permission mode for directories created while
|
||||||
// restoring (parent directories and the target root; restored
|
// restoring (parent directories and the target root; restored
|
||||||
// directories themselves get their stored mode).
|
// directories themselves get their stored mode).
|
||||||
const restoreDirMode = 0o755
|
const restoreDirMode = 0o755
|
||||||
|
|
||||||
// restoreFileMode is the restrictive mode a regular file is created with
|
|
||||||
// during restore. Content is written while the file holds this mode; the
|
|
||||||
// stored mode is applied only after the file is fully written and closed,
|
|
||||||
// so a file whose stored mode is restrictive is never briefly readable by
|
|
||||||
// other local users while its content is being written.
|
|
||||||
const restoreFileMode = 0o600
|
|
||||||
|
|
||||||
// sweepIntervalDivisor sets the sweeper threshold to one N-th of the
|
// sweepIntervalDivisor sets the sweeper threshold to one N-th of the
|
||||||
// configured blob size limit.
|
// configured blob size limit.
|
||||||
const sweepIntervalDivisor = 100
|
const sweepIntervalDivisor = 100
|
||||||
@@ -102,7 +76,7 @@ type RestoreResult struct {
|
|||||||
func (v *Vaultik) Restore(opts *RestoreOptions) error {
|
func (v *Vaultik) Restore(opts *RestoreOptions) error {
|
||||||
startTime := time.Now()
|
startTime := time.Now()
|
||||||
|
|
||||||
identities, err := v.restoreIdentities()
|
identity, err := v.prepareRestoreIdentity()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -116,7 +90,7 @@ func (v *Vaultik) Restore(opts *RestoreOptions) error {
|
|||||||
// Step 1: Download and decrypt the snapshot metadata database
|
// Step 1: Download and decrypt the snapshot metadata database
|
||||||
log.Info("Downloading snapshot metadata...")
|
log.Info("Downloading snapshot metadata...")
|
||||||
|
|
||||||
tempDB, tempDir, err := v.downloadSnapshotDB(opts.SnapshotID, identities)
|
tempDB, err := v.downloadSnapshotDB(opts.SnapshotID, identity)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("downloading snapshot database: %w", err)
|
return fmt.Errorf("downloading snapshot database: %w", err)
|
||||||
}
|
}
|
||||||
@@ -126,11 +100,10 @@ func (v *Vaultik) Restore(opts *RestoreOptions) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
log.Debug("Failed to close temp database", "error", err)
|
log.Debug("Failed to close temp database", "error", err)
|
||||||
}
|
}
|
||||||
// Remove the whole private directory, so the decrypted database
|
// Clean up temp file
|
||||||
// and any SQLite side files it produced are gone on every path.
|
err = v.Fs.Remove(tempDB.Path())
|
||||||
err = v.Fs.RemoveAll(tempDir)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Debug("Failed to remove temp database directory", "error", err)
|
log.Debug("Failed to remove temp database", "error", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
@@ -165,7 +138,7 @@ func (v *Vaultik) Restore(opts *RestoreOptions) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Step 5: Restore files
|
// Step 5: Restore files
|
||||||
result, err := v.restoreAllFiles(files, repos, opts, identities, chunkToBlobMap)
|
result, err := v.restoreAllFiles(files, repos, opts, identity, chunkToBlobMap)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -226,28 +199,21 @@ func (v *Vaultik) finishRestore(
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// restoreIdentities parses the configured age secret key once into every
|
// prepareRestoreIdentity validates that an age secret key is configured
|
||||||
// identity it contains. The value may be a single key line or a whole
|
// and parses it.
|
||||||
// age-keygen file with several identities; all of them are returned so
|
//
|
||||||
// blobgen (via age.Decrypt) can read a blob encrypted to any of their
|
//nolint:ireturn // age.Identity is the decryption abstraction by design
|
||||||
// recipients. This is the first step of both restore and deep verify, so
|
func (v *Vaultik) prepareRestoreIdentity() (age.Identity, error) {
|
||||||
// a missing or unparseable key fails before anything is downloaded. The
|
|
||||||
// error names the configuration source but never the key value.
|
|
||||||
func (v *Vaultik) restoreIdentities() ([]age.Identity, error) {
|
|
||||||
if v.Config.AgeSecretKey == "" {
|
if v.Config.AgeSecretKey == "" {
|
||||||
return nil, errDecryptionKeyRequired
|
return nil, errDecryptionKeyRequired
|
||||||
}
|
}
|
||||||
|
|
||||||
// age.ParseIdentities skips comment and blank lines and rejects a
|
identity, err := age.ParseX25519Identity(v.Config.AgeSecretKey)
|
||||||
// malformed key. Its error can quote the offending line, so it is not
|
|
||||||
// wrapped here — that would leak the secret into the message.
|
|
||||||
identities, err := age.ParseIdentities(strings.NewReader(v.Config.AgeSecretKey))
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("%w (source: %s)",
|
return nil, fmt.Errorf("parsing age secret key: %w", err)
|
||||||
errInvalidAgeSecretKey, v.Config.AgeSecretKeySourceName())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return identities, nil
|
return identity, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// restoreAllFiles processes files in blob-locality order: drain every
|
// restoreAllFiles processes files in blob-locality order: drain every
|
||||||
@@ -260,7 +226,7 @@ func (v *Vaultik) restoreAllFiles(
|
|||||||
files []*database.File,
|
files []*database.File,
|
||||||
repos *database.Repositories,
|
repos *database.Repositories,
|
||||||
opts *RestoreOptions,
|
opts *RestoreOptions,
|
||||||
identities []age.Identity,
|
identity age.Identity,
|
||||||
chunkToBlobMap map[string]*database.BlobChunk,
|
chunkToBlobMap map[string]*database.BlobChunk,
|
||||||
) (*RestoreResult, error) {
|
) (*RestoreResult, error) {
|
||||||
result := &RestoreResult{}
|
result := &RestoreResult{}
|
||||||
@@ -314,7 +280,7 @@ func (v *Vaultik) restoreAllFiles(
|
|||||||
ctx: v.ctx,
|
ctx: v.ctx,
|
||||||
repos: repos,
|
repos: repos,
|
||||||
opts: opts,
|
opts: opts,
|
||||||
identities: identities,
|
identity: identity,
|
||||||
chunkToBlobMap: chunkToBlobMap,
|
chunkToBlobMap: chunkToBlobMap,
|
||||||
blobByHash: blobByHash,
|
blobByHash: blobByHash,
|
||||||
blobIDToHash: blobIDToHash,
|
blobIDToHash: blobIDToHash,
|
||||||
@@ -390,13 +356,6 @@ func (v *Vaultik) runRestoreLoop(
|
|||||||
totalBytesExpected, startTime, &lastStatusTime)
|
totalBytesExpected, startTime, &lastStatusTime)
|
||||||
}
|
}
|
||||||
|
|
||||||
// The loop above stops as soon as nothing is ready and nothing more
|
|
||||||
// can be downloaded. If files still remain, they were abandoned
|
|
||||||
// rather than restored; fail loudly instead of reporting success.
|
|
||||||
if plan.hasPending() {
|
|
||||||
return errRestoreIncomplete
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -411,18 +370,12 @@ func (v *Vaultik) runRestoreLoop(
|
|||||||
func (s *restoreSession) downloadNextBlobSet(plan *restorePlan) (bool, error) {
|
func (s *restoreSession) downloadNextBlobSet(plan *restorePlan) (bool, error) {
|
||||||
s.sweeper.sweep()
|
s.sweeper.sweep()
|
||||||
|
|
||||||
next, ok := plan.pickNextDownload()
|
next := plan.pickNextDownload()
|
||||||
if !ok {
|
if next.IsZero() {
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, hash := range plan.blobsNeeded(next) {
|
for _, hash := range plan.blobsNeeded(next) {
|
||||||
// Stop between blobs on cancel so an interrupt ends the download
|
|
||||||
// phase promptly rather than fetching the rest of the set.
|
|
||||||
if s.ctx.Err() != nil {
|
|
||||||
return false, s.ctx.Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
blob, ok := s.blobByHash[hash]
|
blob, ok := s.blobByHash[hash]
|
||||||
if !ok {
|
if !ok {
|
||||||
return false, fmt.Errorf("%w: %s", errBlobMissingFromIndex, hash[:16])
|
return false, fmt.Errorf("%w: %s", errBlobMissingFromIndex, hash[:16])
|
||||||
@@ -628,11 +581,11 @@ func (v *Vaultik) handleRestoreVerification(
|
|||||||
// for a remote-only snapshot) is used as-is, so a host with no local
|
// for a remote-only snapshot) is used as-is, so a host with no local
|
||||||
// index can restore the snapshots it can only see on the store.
|
// index can restore the snapshots it can only see on the store.
|
||||||
func (v *Vaultik) downloadSnapshotDB(
|
func (v *Vaultik) downloadSnapshotDB(
|
||||||
snapshotID string, identities []age.Identity,
|
snapshotID string, identity age.Identity,
|
||||||
) (*database.DB, string, error) {
|
) (*database.DB, error) {
|
||||||
remoteKey, err := v.resolveSnapshotRemoteKey(snapshotID)
|
remoteKey, err := v.resolveSnapshotRemoteKey(snapshotID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Download encrypted database from storage
|
// Download encrypted database from storage
|
||||||
@@ -640,7 +593,7 @@ func (v *Vaultik) downloadSnapshotDB(
|
|||||||
|
|
||||||
reader, err := v.Storage.Get(v.ctx, dbKey)
|
reader, err := v.Storage.Get(v.ctx, dbKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", fmt.Errorf("downloading %s: %w", dbKey, err)
|
return nil, fmt.Errorf("downloading %s: %w", dbKey, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
defer func() { _ = reader.Close() }()
|
defer func() { _ = reader.Close() }()
|
||||||
@@ -648,16 +601,16 @@ func (v *Vaultik) downloadSnapshotDB(
|
|||||||
// Read all data
|
// Read all data
|
||||||
encryptedData, err := io.ReadAll(reader)
|
encryptedData, err := io.ReadAll(reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", fmt.Errorf("reading encrypted data: %w", err)
|
return nil, fmt.Errorf("reading encrypted data: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Debug("Downloaded encrypted database",
|
log.Debug("Downloaded encrypted database",
|
||||||
"size", ubytes(int64(len(encryptedData))))
|
"size", ubytes(int64(len(encryptedData))))
|
||||||
|
|
||||||
// Decrypt and decompress using blobgen.Reader
|
// Decrypt and decompress using blobgen.Reader
|
||||||
blobReader, err := blobgen.NewReader(bytes.NewReader(encryptedData), identities...)
|
blobReader, err := blobgen.NewReader(bytes.NewReader(encryptedData), identity)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", fmt.Errorf("creating decryption reader: %w", err)
|
return nil, fmt.Errorf("creating decryption reader: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
defer func() { _ = blobReader.Close() }()
|
defer func() { _ = blobReader.Close() }()
|
||||||
@@ -665,98 +618,44 @@ func (v *Vaultik) downloadSnapshotDB(
|
|||||||
// Read the binary SQLite database
|
// Read the binary SQLite database
|
||||||
dbData, err := io.ReadAll(blobReader)
|
dbData, err := io.ReadAll(blobReader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", fmt.Errorf("decrypting and decompressing: %w", err)
|
return nil, fmt.Errorf("decrypting and decompressing: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Debug("Decrypted database", "size", ubytes(int64(len(dbData))))
|
log.Debug("Decrypted database", "size", ubytes(int64(len(dbData))))
|
||||||
|
|
||||||
db, tempDir, err := v.materializeSnapshotDB(dbData)
|
// Create a temporary database file and write the binary SQLite data directly
|
||||||
|
tempFile, err := afero.TempFile(v.Fs, "", "vaultik-restore-*.db")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", err
|
return nil, fmt.Errorf("creating temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Confirm the decrypted database really is the snapshot named by
|
tempPath := tempFile.Name()
|
||||||
// remoteKey before any files are read from it. On mismatch, close the
|
|
||||||
// database and remove its private directory so nothing is left behind.
|
// Write the binary SQLite database directly
|
||||||
err = v.verifySnapshotDBIdentity(db, snapshotID, remoteKey)
|
_, err = tempFile.Write(dbData)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = db.Close()
|
_ = tempFile.Close()
|
||||||
_ = v.Fs.RemoveAll(tempDir)
|
_ = v.Fs.Remove(tempPath)
|
||||||
|
|
||||||
return nil, "", err
|
return nil, fmt.Errorf("writing database file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return db, tempDir, nil
|
err = tempFile.Close()
|
||||||
}
|
|
||||||
|
|
||||||
// verifySnapshotDBIdentity confirms the decrypted metadata database really
|
|
||||||
// is the snapshot named by remoteKey. age decryption proves the database
|
|
||||||
// is readable, not that the object served at
|
|
||||||
// metadata/<remoteKey>/db.zst.age is the snapshot that was requested: an
|
|
||||||
// attacker who swaps in another valid db.zst.age (which needs no key
|
|
||||||
// material) would otherwise redirect restore and deep verify to a
|
|
||||||
// different snapshot's contents. The exported per-snapshot database holds
|
|
||||||
// exactly one snapshot row, and a snapshot's remote key is derived from
|
|
||||||
// that row's ID, so the database is the requested one exactly when its
|
|
||||||
// sole snapshot hashes back to remoteKey. Comparing the requested
|
|
||||||
// identifier directly would not do: it may be a remote-key prefix a
|
|
||||||
// recovery host uses in place of a human snapshot ID it cannot know.
|
|
||||||
func (v *Vaultik) verifySnapshotDBIdentity(
|
|
||||||
db *database.DB, requested, remoteKey string,
|
|
||||||
) error {
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
snap, err := repos.Snapshots.GetOnlySnapshot(v.ctx)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("checking identity of database for %s: %w", requested, err)
|
_ = v.Fs.Remove(tempPath)
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("closing temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if snapshot.RemoteSnapshotKey(snap.ID.String()) != remoteKey {
|
log.Debug("Created restore database", "path", tempPath)
|
||||||
return fmt.Errorf("%w: requested %s but the database is snapshot %s",
|
|
||||||
errSnapshotDBMismatch, requested, snap.ID)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
// Open the database
|
||||||
}
|
db, err := database.New(v.ctx, tempPath)
|
||||||
|
|
||||||
// materializeSnapshotDB writes the decrypted snapshot database bytes into
|
|
||||||
// a fresh private (0700) temp directory and opens the file read-only. On
|
|
||||||
// any failure it removes the directory before returning, so no decrypted
|
|
||||||
// metadata is left on disk when the open is interrupted or the payload is
|
|
||||||
// damaged. On success the returned directory is the caller's to remove.
|
|
||||||
func (v *Vaultik) materializeSnapshotDB(
|
|
||||||
dbData []byte,
|
|
||||||
) (*database.DB, string, error) {
|
|
||||||
tempDir, err := afero.TempDir(v.Fs, "", "vaultik-restore-")
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, "", fmt.Errorf("creating temp directory: %w", err)
|
return nil, fmt.Errorf("opening restore database: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
success := false
|
return db, nil
|
||||||
|
|
||||||
defer func() {
|
|
||||||
if !success {
|
|
||||||
_ = v.Fs.RemoveAll(tempDir)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
dbPath := filepath.Join(tempDir, snapshotDBFilename)
|
|
||||||
|
|
||||||
err = afero.WriteFile(v.Fs, dbPath, dbData, restoreFileMode)
|
|
||||||
if err != nil {
|
|
||||||
return nil, "", fmt.Errorf("writing database file: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Debug("Created restore database", "path", dbPath)
|
|
||||||
|
|
||||||
db, err := database.OpenReadOnly(v.ctx, dbPath)
|
|
||||||
if err != nil {
|
|
||||||
return nil, "", fmt.Errorf("opening restore database: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
success = true
|
|
||||||
|
|
||||||
return db, tempDir, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// getFilesToRestore returns the list of files to restore based on path filters
|
// getFilesToRestore returns the list of files to restore based on path filters
|
||||||
@@ -845,7 +744,7 @@ type restoreSession struct {
|
|||||||
ctx context.Context //nolint:containedctx // per-restore state by design
|
ctx context.Context //nolint:containedctx // per-restore state by design
|
||||||
repos *database.Repositories
|
repos *database.Repositories
|
||||||
opts *RestoreOptions
|
opts *RestoreOptions
|
||||||
identities []age.Identity
|
identity age.Identity
|
||||||
chunkToBlobMap map[string]*database.BlobChunk
|
chunkToBlobMap map[string]*database.BlobChunk
|
||||||
blobByHash map[string]*database.Blob
|
blobByHash map[string]*database.Blob
|
||||||
blobIDToHash map[string]string
|
blobIDToHash map[string]string
|
||||||
@@ -861,85 +760,13 @@ type restoreSession struct {
|
|||||||
runningAsRoot bool
|
runningAsRoot bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// containedRestorePath resolves rel — a path read from the snapshot
|
|
||||||
// database — to its location under targetDir and confirms the write will
|
|
||||||
// stay inside the target.
|
|
||||||
//
|
|
||||||
// age decryption proves a snapshot is readable, not that it is honest, so
|
|
||||||
// every stored path is treated as hostile. rel is rejected unless
|
|
||||||
// filepath.IsLocal accepts it once the leading separator is stripped:
|
|
||||||
// stored paths are absolute and the join to targetDir drops that
|
|
||||||
// separator, so "/etc/passwd" is judged as the relative "etc/passwd" it
|
|
||||||
// becomes on disk. This bars "..", absolute, and empty paths.
|
|
||||||
//
|
|
||||||
// A stored symlink whose target points outside the tree is still honest
|
|
||||||
// (and restored verbatim), but a later entry must not be written through
|
|
||||||
// it. Each existing ancestor directory below the target is therefore
|
|
||||||
// Lstat'ed and a symlink among them is refused. The leaf itself is not
|
|
||||||
// traversed: honest snapshots restore symlinks at leaf positions, and the
|
|
||||||
// unique-path constraint keeps a leaf from being both a symlink and a
|
|
||||||
// regular file. The target directory itself may be a symlink; only
|
|
||||||
// components below it are checked.
|
|
||||||
func containedRestorePath(fs afero.Fs, targetDir, rel string) (string, error) {
|
|
||||||
local := strings.TrimPrefix(rel, string(filepath.Separator))
|
|
||||||
if !filepath.IsLocal(local) {
|
|
||||||
return "", fmt.Errorf("%w: %s", errRestorePathEscapesTarget, rel)
|
|
||||||
}
|
|
||||||
|
|
||||||
local = filepath.Clean(local)
|
|
||||||
targetPath := filepath.Join(targetDir, local)
|
|
||||||
|
|
||||||
relDir := filepath.Dir(local)
|
|
||||||
if relDir == "." {
|
|
||||||
return targetPath, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
current := targetDir
|
|
||||||
for component := range strings.SplitSeq(relDir, string(filepath.Separator)) {
|
|
||||||
current = filepath.Join(current, component)
|
|
||||||
|
|
||||||
info, err := lstatIfPossible(fs, current)
|
|
||||||
if err != nil {
|
|
||||||
if os.IsNotExist(err) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
return "", fmt.Errorf("checking restore path %s: %w", current, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if info.Mode()&os.ModeSymlink != 0 {
|
|
||||||
return "", fmt.Errorf("%w: %s descends through symlink %s",
|
|
||||||
errRestorePathEscapesTarget, rel, current)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return targetPath, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// lstatIfPossible performs a symlink-aware stat when the filesystem
|
|
||||||
// supports it. afero.OsFs does; MemMapFs, which has no symlinks, reports
|
|
||||||
// that Lstat was not used and its result never carries ModeSymlink.
|
|
||||||
func lstatIfPossible(fs afero.Fs, name string) (os.FileInfo, error) {
|
|
||||||
if lstater, ok := fs.(afero.Lstater); ok {
|
|
||||||
info, _, err := lstater.LstatIfPossible(name)
|
|
||||||
|
|
||||||
return info, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return fs.Stat(name)
|
|
||||||
}
|
|
||||||
|
|
||||||
// restoreFile dispatches to the right per-kind restorer.
|
// restoreFile dispatches to the right per-kind restorer.
|
||||||
func (s *restoreSession) restoreFile(file *database.File) error {
|
func (s *restoreSession) restoreFile(file *database.File) error {
|
||||||
targetPath, err := containedRestorePath(
|
targetPath := filepath.Join(s.opts.TargetDir, file.Path.String())
|
||||||
s.v.Fs, s.opts.TargetDir, file.Path.String())
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
parentDir := filepath.Dir(targetPath)
|
parentDir := filepath.Dir(targetPath)
|
||||||
|
|
||||||
err = s.v.Fs.MkdirAll(parentDir, restoreDirMode)
|
err := s.v.Fs.MkdirAll(parentDir, restoreDirMode)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("creating parent directory: %w", err)
|
return fmt.Errorf("creating parent directory: %w", err)
|
||||||
}
|
}
|
||||||
@@ -987,13 +814,6 @@ func (s *restoreSession) restoreDirectory(
|
|||||||
return fmt.Errorf("creating directory: %w", err)
|
return fmt.Errorf("creating directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// MkdirAll applies the process umask, so chmod to the exact stored
|
|
||||||
// mode. A failure here is non-fatal.
|
|
||||||
err = s.v.Fs.Chmod(targetPath, os.FileMode(file.Mode))
|
|
||||||
if err != nil {
|
|
||||||
log.Debug("Failed to set permissions", "path", targetPath, "error", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
s.applyFileMetadata(file, targetPath)
|
s.applyFileMetadata(file, targetPath)
|
||||||
|
|
||||||
s.result.FilesRestored++
|
s.result.FilesRestored++
|
||||||
@@ -1001,22 +821,25 @@ func (s *restoreSession) restoreDirectory(
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// applyFileMetadata applies ownership (when running as root on a real
|
// applyFileMetadata applies stored permissions, ownership (when running
|
||||||
// filesystem) and mtime to a restored path. Permission mode is applied
|
// as root on a real filesystem), and mtime to a restored path. Failures
|
||||||
// separately by each caller, with different failure handling, so it is
|
// are logged at debug level and do not abort the restore.
|
||||||
// not touched here. Failures are logged at debug level and do not abort
|
|
||||||
// the restore.
|
|
||||||
func (s *restoreSession) applyFileMetadata(file *database.File, targetPath string) {
|
func (s *restoreSession) applyFileMetadata(file *database.File, targetPath string) {
|
||||||
|
err := s.v.Fs.Chmod(targetPath, os.FileMode(file.Mode))
|
||||||
|
if err != nil {
|
||||||
|
log.Debug("Failed to set permissions", "path", targetPath, "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
if s.runningAsRoot {
|
if s.runningAsRoot {
|
||||||
if _, ok := s.v.Fs.(*afero.OsFs); ok {
|
if _, ok := s.v.Fs.(*afero.OsFs); ok {
|
||||||
err := os.Chown(targetPath, int(file.UID), int(file.GID))
|
err = os.Chown(targetPath, int(file.UID), int(file.GID))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Debug("Failed to set ownership", "path", targetPath, "error", err)
|
log.Debug("Failed to set ownership", "path", targetPath, "error", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
err := s.v.Fs.Chtimes(targetPath, file.MTime, file.MTime)
|
err = s.v.Fs.Chtimes(targetPath, file.MTime, file.MTime)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Debug("Failed to set mtime", "path", targetPath, "error", err)
|
log.Debug("Failed to set mtime", "path", targetPath, "error", err)
|
||||||
}
|
}
|
||||||
@@ -1050,30 +873,17 @@ func (s *restoreSession) restoreRegularFile(
|
|||||||
|
|
||||||
t0 = time.Now()
|
t0 = time.Now()
|
||||||
|
|
||||||
// Remove any existing entry, then create the file with a restrictive
|
outFile, err := s.v.Fs.Create(targetPath)
|
||||||
// mode via O_EXCL. The stored mode is applied only after the content
|
|
||||||
// is written and the file closed, so a file whose stored mode is
|
|
||||||
// restrictive is never briefly readable by other local users while
|
|
||||||
// its content is written. Removing first (rather than failing on a
|
|
||||||
// leftover file) matches the documented behaviour that re-running
|
|
||||||
// restore overwrites partial output.
|
|
||||||
_ = s.v.Fs.Remove(targetPath)
|
|
||||||
|
|
||||||
outFile, err := s.v.Fs.OpenFile(
|
|
||||||
targetPath, os.O_CREATE|os.O_EXCL|os.O_WRONLY, restoreFileMode)
|
|
||||||
createDur := time.Since(t0)
|
createDur := time.Since(t0)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("creating output file: %w", err)
|
return fmt.Errorf("creating output file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
defer func() { _ = outFile.Close() }()
|
||||||
|
|
||||||
bytesWritten, timings, err := s.writeFileChunks(outFile, fileChunks)
|
bytesWritten, timings, err := s.writeFileChunks(outFile, fileChunks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Do not leave a partial file behind.
|
|
||||||
_ = outFile.Close()
|
|
||||||
|
|
||||||
s.removePartialRestore(targetPath)
|
|
||||||
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1091,12 +901,9 @@ func (s *restoreSession) restoreRegularFile(
|
|||||||
|
|
||||||
err = outFile.Close()
|
err = outFile.Close()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.removePartialRestore(targetPath)
|
|
||||||
|
|
||||||
return fmt.Errorf("closing output file: %w", err)
|
return fmt.Errorf("closing output file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
s.applyRestoredFileMode(file, targetPath)
|
|
||||||
s.applyFileMetadata(file, targetPath)
|
s.applyFileMetadata(file, targetPath)
|
||||||
|
|
||||||
s.result.FilesRestored++
|
s.result.FilesRestored++
|
||||||
@@ -1107,31 +914,6 @@ func (s *restoreSession) restoreRegularFile(
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// applyRestoredFileMode applies the stored permission bits to a
|
|
||||||
// just-written regular file (created with restoreFileMode). A failure is
|
|
||||||
// a user-visible warning, not a fatal error: the file's content is
|
|
||||||
// intact and it remains at the restrictive create-time mode, so the
|
|
||||||
// restore is not aborted or discarded over it.
|
|
||||||
func (s *restoreSession) applyRestoredFileMode(
|
|
||||||
file *database.File, targetPath string,
|
|
||||||
) {
|
|
||||||
err := s.v.Fs.Chmod(targetPath, os.FileMode(file.Mode))
|
|
||||||
if err != nil {
|
|
||||||
s.v.UI.Warningf("Failed to set mode %s on %s: %v",
|
|
||||||
os.FileMode(file.Mode).Perm(), s.v.UI.Path(targetPath), err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// removePartialRestore deletes a restore output file whose write did not
|
|
||||||
// complete, so a failed restore never leaves a partial file behind.
|
|
||||||
func (s *restoreSession) removePartialRestore(targetPath string) {
|
|
||||||
err := s.v.Fs.Remove(targetPath)
|
|
||||||
if err != nil {
|
|
||||||
log.Debug("Failed to remove partial restore file",
|
|
||||||
"path", targetPath, "error", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeFileChunks streams each of the file's chunks from the blob disk
|
// writeFileChunks streams each of the file's chunks from the blob disk
|
||||||
// cache into outFile, crediting restored bytes to the sweeper as it
|
// cache into outFile, crediting restored bytes to the sweeper as it
|
||||||
// goes. Returns the bytes written plus per-phase timing accumulators.
|
// goes. Returns the bytes written plus per-phase timing accumulators.
|
||||||
@@ -1144,12 +926,6 @@ func (s *restoreSession) writeFileChunks(
|
|||||||
)
|
)
|
||||||
|
|
||||||
for _, fc := range fileChunks {
|
for _, fc := range fileChunks {
|
||||||
// Stop between chunks on cancel so an interrupt does not keep
|
|
||||||
// writing a large file after the operation has been told to stop.
|
|
||||||
if s.ctx.Err() != nil {
|
|
||||||
return bytesWritten, timings, s.ctx.Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
chunkHashStr := fc.ChunkHash.String()
|
chunkHashStr := fc.ChunkHash.String()
|
||||||
|
|
||||||
blobChunk, ok := s.chunkToBlobMap[chunkHashStr]
|
blobChunk, ok := s.chunkToBlobMap[chunkHashStr]
|
||||||
@@ -1207,7 +983,7 @@ func (s *restoreSession) downloadBlobToCache(
|
|||||||
start := time.Now()
|
start := time.Now()
|
||||||
|
|
||||||
t0 := time.Now()
|
t0 := time.Now()
|
||||||
rc, err := s.v.FetchAndDecryptBlob(s.ctx, blobHash, expectedSize, s.identities...)
|
rc, err := s.v.FetchAndDecryptBlob(s.ctx, blobHash, expectedSize, s.identity)
|
||||||
fetchSetupDur := time.Since(t0)
|
fetchSetupDur := time.Since(t0)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -1219,19 +995,11 @@ func (s *restoreSession) downloadBlobToCache(
|
|||||||
streamDur := time.Since(t0)
|
streamDur := time.Since(t0)
|
||||||
closeErr := rc.Close()
|
closeErr := rc.Close()
|
||||||
|
|
||||||
// closeErr carries the blob's hash-verification result (a mismatch,
|
|
||||||
// or the stream not being fully read). On any failure, drop the
|
|
||||||
// cache entry so a blob that failed verification is never read back
|
|
||||||
// as if it were valid.
|
|
||||||
if copyErr != nil {
|
if copyErr != nil {
|
||||||
s.blobCache.Delete(blobHash)
|
|
||||||
|
|
||||||
return copyErr
|
return copyErr
|
||||||
}
|
}
|
||||||
|
|
||||||
if closeErr != nil {
|
if closeErr != nil {
|
||||||
s.blobCache.Delete(blobHash)
|
|
||||||
|
|
||||||
return closeErr
|
return closeErr
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1294,22 +1062,17 @@ func (v *Vaultik) verifyRestoredFiles(
|
|||||||
return ctx.Err()
|
return ctx.Err()
|
||||||
}
|
}
|
||||||
|
|
||||||
targetPath, err := containedRestorePath(v.Fs, targetDir, file.Path.String())
|
targetPath := filepath.Join(targetDir, file.Path.String())
|
||||||
if err == nil {
|
|
||||||
var bytesVerified int64
|
|
||||||
|
|
||||||
bytesVerified, err = v.verifyFile(ctx, repos, file, targetPath)
|
|
||||||
if err == nil {
|
|
||||||
result.FilesVerified++
|
|
||||||
result.BytesVerified += bytesVerified
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
|
bytesVerified, err := v.verifyFile(ctx, repos, file, targetPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Error("File verification failed", "path", file.Path, "error", err)
|
log.Error("File verification failed", "path", file.Path, "error", err)
|
||||||
|
|
||||||
result.FilesFailed++
|
result.FilesFailed++
|
||||||
result.FailedFiles = append(result.FailedFiles, file.Path.String())
|
result.FailedFiles = append(result.FailedFiles, file.Path.String())
|
||||||
|
} else {
|
||||||
|
result.FilesVerified++
|
||||||
|
result.BytesVerified += bytesVerified
|
||||||
}
|
}
|
||||||
|
|
||||||
bytesProcessed += file.Size
|
bytesProcessed += file.Size
|
||||||
@@ -1399,17 +1162,6 @@ func (v *Vaultik) verifyFile(
|
|||||||
bytesVerified += int64(n)
|
bytesVerified += int64(n)
|
||||||
}
|
}
|
||||||
|
|
||||||
// The stored chunks account for the whole file, so the reader must
|
|
||||||
// be at EOF now. Trailing bytes past the last chunk are corruption
|
|
||||||
// the per-chunk loop cannot see.
|
|
||||||
extra := make([]byte, 1)
|
|
||||||
|
|
||||||
n, err := f.Read(extra)
|
|
||||||
if n != 0 || !errors.Is(err, io.EOF) {
|
|
||||||
return bytesVerified, fmt.Errorf("%w: file longer than its %d chunk(s)",
|
|
||||||
errTrailingRestoreData, len(fileChunks))
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Debug("File verified",
|
log.Debug("File verified",
|
||||||
"path", file.Path, "bytes", bytesVerified, "chunks", len(fileChunks))
|
"path", file.Path, "bytes", bytesVerified, "chunks", len(fileChunks))
|
||||||
|
|
||||||
|
|||||||
@@ -1,175 +0,0 @@
|
|||||||
package vaultik //nolint:testpackage // drives unexported restore internals
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/types"
|
|
||||||
"sneak.berlin/go/vaultik/internal/ui"
|
|
||||||
)
|
|
||||||
|
|
||||||
// These tests exercise the path-containment guard that keeps restore from
|
|
||||||
// writing outside its target directory. age decryption proves only that a
|
|
||||||
// snapshot is readable, not that its recorded paths are honest, so restore
|
|
||||||
// treats every stored path as hostile: a compromised backed-up host could
|
|
||||||
// forge a snapshot that decrypts cleanly, and restore usually runs as root.
|
|
||||||
//
|
|
||||||
// They drive restoreAllFiles directly (rather than the full Restore, which
|
|
||||||
// downloads and decrypts the metadata database from storage) so a snapshot
|
|
||||||
// database with adversarial rows can be handed to the restore loop without
|
|
||||||
// the surrounding blob/storage machinery. Directory and symlink entries
|
|
||||||
// carry no chunks, so no blobs are needed.
|
|
||||||
|
|
||||||
// containmentDirMode marks a File row as a directory for the restore loop.
|
|
||||||
const containmentDirMode = uint32(os.ModeDir | 0o755)
|
|
||||||
|
|
||||||
// newContainmentVaultik builds the minimal Vaultik needed to run
|
|
||||||
// restoreAllFiles against fs.
|
|
||||||
func newContainmentVaultik(ctx context.Context, fs afero.Fs) *Vaultik {
|
|
||||||
v := &Vaultik{
|
|
||||||
Config: &config.Config{
|
|
||||||
BlobSizeLimit: config.Size(10 * 1024 * 1024),
|
|
||||||
},
|
|
||||||
Fs: fs,
|
|
||||||
Stdout: io.Discard,
|
|
||||||
Stderr: io.Discard,
|
|
||||||
UI: ui.NewWithColor(io.Discard, false),
|
|
||||||
}
|
|
||||||
v.SetContext(ctx)
|
|
||||||
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
|
|
||||||
// makeFiles inserts the given rows into a fresh in-memory snapshot database
|
|
||||||
// and returns them (with IDs assigned) plus the repositories.
|
|
||||||
func makeFiles(
|
|
||||||
ctx context.Context, t *testing.T, rows []*database.File,
|
|
||||||
) ([]*database.File, *database.Repositories) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
db, err := database.New(ctx, filepath.Join(t.TempDir(), "index.sqlite"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = db.Close() })
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
for _, f := range rows {
|
|
||||||
require.NoError(t, repos.Files.Create(ctx, nil, f))
|
|
||||||
}
|
|
||||||
|
|
||||||
return rows, repos
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRestoreRejectsPathTraversal(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
// rows are inserted in order; the escape entry is restored after
|
|
||||||
// any entry it depends on (the symlink case needs its link first).
|
|
||||||
rows func(outsideDir string) []*database.File
|
|
||||||
// escaped is the path, outside the target, that must not appear.
|
|
||||||
escaped func(tempDir, outsideDir string) string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "relative dotdot",
|
|
||||||
rows: func(_ string) []*database.File {
|
|
||||||
return []*database.File{{
|
|
||||||
Path: "../escaped-relative",
|
|
||||||
Mode: containmentDirMode,
|
|
||||||
}}
|
|
||||||
},
|
|
||||||
escaped: func(tempDir, _ string) string {
|
|
||||||
return filepath.Join(tempDir, "escaped-relative")
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "absolute with dotdot",
|
|
||||||
rows: func(_ string) []*database.File {
|
|
||||||
return []*database.File{{
|
|
||||||
Path: "/a/../../escaped-absolute",
|
|
||||||
Mode: containmentDirMode,
|
|
||||||
}}
|
|
||||||
},
|
|
||||||
escaped: func(tempDir, _ string) string {
|
|
||||||
return filepath.Join(tempDir, "escaped-absolute")
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "child through symlink",
|
|
||||||
rows: func(outsideDir string) []*database.File {
|
|
||||||
return []*database.File{
|
|
||||||
// Restored first: an in-target symlink pointing out.
|
|
||||||
{Path: "linkdir", LinkTarget: types.FilePath(outsideDir)},
|
|
||||||
// Restored second: a child written through that link.
|
|
||||||
{Path: "linkdir/child", Mode: containmentDirMode},
|
|
||||||
}
|
|
||||||
},
|
|
||||||
escaped: func(_, outsideDir string) string {
|
|
||||||
return filepath.Join(outsideDir, "child")
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tc := range tests {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
targetDir := filepath.Join(tempDir, "target")
|
|
||||||
outsideDir := filepath.Join(tempDir, "outside")
|
|
||||||
require.NoError(t, fs.MkdirAll(outsideDir, 0o755))
|
|
||||||
|
|
||||||
rows, repos := makeFiles(ctx, t, tc.rows(outsideDir))
|
|
||||||
v := newContainmentVaultik(ctx, fs)
|
|
||||||
|
|
||||||
_, err := v.restoreAllFiles(rows, repos,
|
|
||||||
&RestoreOptions{TargetDir: targetDir}, nil, nil)
|
|
||||||
|
|
||||||
require.ErrorIs(t, err, errRestorePathEscapesTarget)
|
|
||||||
|
|
||||||
escaped := tc.escaped(tempDir, outsideDir)
|
|
||||||
_, statErr := os.Lstat(escaped)
|
|
||||||
require.Truef(t, os.IsNotExist(statErr),
|
|
||||||
"restore wrote outside the target at %s", escaped)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestRestoreAllowsSymlinkPointingOutsideTree confirms the guard does not
|
|
||||||
// over-block: an honest snapshot may contain a symlink whose target lies
|
|
||||||
// outside the restored tree, and it must still be restored verbatim.
|
|
||||||
func TestRestoreAllowsSymlinkPointingOutsideTree(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
targetDir := filepath.Join(tempDir, "target")
|
|
||||||
linkTarget := filepath.Join(tempDir, "outside", "data")
|
|
||||||
|
|
||||||
rows, repos := makeFiles(ctx, t, []*database.File{
|
|
||||||
{Path: "goodlink", LinkTarget: types.FilePath(linkTarget), MTime: time.Unix(0, 0)},
|
|
||||||
})
|
|
||||||
v := newContainmentVaultik(ctx, fs)
|
|
||||||
|
|
||||||
_, err := v.restoreAllFiles(rows, repos,
|
|
||||||
&RestoreOptions{TargetDir: targetDir}, nil, nil)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
got, err := os.Readlink(filepath.Join(targetDir, "goodlink"))
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, linkTarget, got)
|
|
||||||
}
|
|
||||||
@@ -1,108 +0,0 @@
|
|||||||
package vaultik //nolint:testpackage // exercises unexported restoreIdentities
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"io"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"filippo.io/age"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// encryptBlobTo returns a blobgen blob of plaintext encrypted to exactly
|
|
||||||
// one recipient, so a decryptor succeeds only if it holds that recipient's
|
|
||||||
// identity.
|
|
||||||
func encryptBlobTo(t *testing.T, recipient string, plaintext []byte) []byte {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var buf bytes.Buffer
|
|
||||||
|
|
||||||
writer, err := blobgen.NewWriter(&buf, 1, []string{recipient})
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = writer.Write(plaintext)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, writer.Close())
|
|
||||||
|
|
||||||
return buf.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
// decryptBlobWith reads a blob back through the identities and returns its
|
|
||||||
// plaintext.
|
|
||||||
func decryptBlobWith(t *testing.T, blob []byte, identities []age.Identity) []byte {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
reader, err := blobgen.NewReader(bytes.NewReader(blob), identities...)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
plaintext, err := io.ReadAll(reader)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, reader.Close())
|
|
||||||
|
|
||||||
return plaintext
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestRestoreIdentitiesAcceptsEveryIdentity proves a key file holding two
|
|
||||||
// identities yields both, so a blob encrypted only to the second
|
|
||||||
// recipient — the one the previous single-identity parse dropped — still
|
|
||||||
// decrypts.
|
|
||||||
func TestRestoreIdentitiesAcceptsEveryIdentity(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
first, err := age.GenerateX25519Identity()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
second, err := age.GenerateX25519Identity()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
// A whole age-keygen-style file: comment lines plus two identity lines.
|
|
||||||
keyFile := "# public key: " + first.Recipient().String() + "\n" +
|
|
||||||
first.String() + "\n" +
|
|
||||||
"# public key: " + second.Recipient().String() + "\n" +
|
|
||||||
second.String() + "\n"
|
|
||||||
|
|
||||||
v := &Vaultik{Config: &config.Config{AgeSecretKey: keyFile}}
|
|
||||||
|
|
||||||
identities, err := v.restoreIdentities()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Len(t, identities, 2)
|
|
||||||
|
|
||||||
plaintext := []byte("payload encrypted only to the second identity")
|
|
||||||
blob := encryptBlobTo(t, second.Recipient().String(), plaintext)
|
|
||||||
|
|
||||||
require.Equal(t, plaintext, decryptBlobWith(t, blob, identities))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestRestoreIdentitiesAcceptsTrailingNewline mirrors a YAML
|
|
||||||
// age_secret_key value that carries a trailing newline: it must still
|
|
||||||
// parse to its one identity and decrypt a blob encrypted to it.
|
|
||||||
func TestRestoreIdentitiesAcceptsTrailingNewline(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
id, err := age.GenerateX25519Identity()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
v := &Vaultik{Config: &config.Config{AgeSecretKey: id.String() + "\n"}}
|
|
||||||
|
|
||||||
identities, err := v.restoreIdentities()
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Len(t, identities, 1)
|
|
||||||
|
|
||||||
plaintext := []byte("value with a trailing newline")
|
|
||||||
blob := encryptBlobTo(t, id.Recipient().String(), plaintext)
|
|
||||||
|
|
||||||
require.Equal(t, plaintext, decryptBlobWith(t, blob, identities))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestRestoreIdentitiesMissingKey reports the dedicated missing-key error
|
|
||||||
// rather than a parse failure, so the user is told to set the key.
|
|
||||||
func TestRestoreIdentitiesMissingKey(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
v := &Vaultik{Config: &config.Config{}}
|
|
||||||
|
|
||||||
_, err := v.restoreIdentities()
|
|
||||||
require.ErrorIs(t, err, errDecryptionKeyRequired)
|
|
||||||
}
|
|
||||||
@@ -1,159 +0,0 @@
|
|||||||
package vaultik //nolint:testpackage // sets ctx/cancel and inspects scratch files
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
"sneak.berlin/go/vaultik/internal/ui"
|
|
||||||
)
|
|
||||||
|
|
||||||
// blockingBlobStorer wraps a Storer and blocks the first blob download
|
|
||||||
// until its context is cancelled, so a test can catch a restore while it
|
|
||||||
// is mid-download. Metadata reads pass straight through, so the restore
|
|
||||||
// reaches the blob-download phase — having already written its decrypted
|
|
||||||
// scratch files — before it blocks.
|
|
||||||
type blockingBlobStorer struct {
|
|
||||||
storage.Storer
|
|
||||||
|
|
||||||
once sync.Once
|
|
||||||
entered chan struct{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newBlockingBlobStorer(inner storage.Storer) *blockingBlobStorer {
|
|
||||||
return &blockingBlobStorer{Storer: inner, entered: make(chan struct{})}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (b *blockingBlobStorer) Get(
|
|
||||||
ctx context.Context, key string,
|
|
||||||
) (io.ReadCloser, error) {
|
|
||||||
if strings.HasPrefix(key, "blobs/") {
|
|
||||||
b.once.Do(func() { close(b.entered) })
|
|
||||||
<-ctx.Done()
|
|
||||||
|
|
||||||
return nil, ctx.Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
return b.Storer.Get(ctx, key)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestRestoreCleansTempDirOnInterrupt drives a restore through the stop
|
|
||||||
// path (v.StartOperation, which is what the fx OnStop hook uses) instead
|
|
||||||
// of calling Restore directly, catches it mid-download, and asserts that
|
|
||||||
// stopping waits for the operation to unwind and removes its decrypted
|
|
||||||
// scratch files — the blob cache and the temporary snapshot database —
|
|
||||||
// from the temp directory. Without the wait a SIGINT exits the process
|
|
||||||
// before those defers run, leaving decrypted data on disk (issue #159).
|
|
||||||
//
|
|
||||||
// Not parallel: it points TMPDIR at a private directory (via t.Setenv)
|
|
||||||
// so it can assert on exactly the scratch files this restore created.
|
|
||||||
func TestRestoreCleansTempDirOnInterrupt(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
root := t.TempDir()
|
|
||||||
|
|
||||||
dataDir := filepath.Join(root, "source")
|
|
||||||
storeDir := filepath.Join(root, "remote")
|
|
||||||
restoreDir := filepath.Join(root, "restored")
|
|
||||||
dbPath := filepath.Join(root, "index.sqlite")
|
|
||||||
|
|
||||||
require.NoError(t, fs.MkdirAll(dataDir, 0o755))
|
|
||||||
|
|
||||||
buildLocalityFixture(t, fs, dataDir)
|
|
||||||
|
|
||||||
cfg, storer, snapshotID := setupLocalityBackup(
|
|
||||||
context.Background(), t, fs, dataDir, storeDir, dbPath)
|
|
||||||
|
|
||||||
// Point the "" temp paths (the blob cache directory and the
|
|
||||||
// snapshot-database directory) at a private directory so the test can
|
|
||||||
// assert on exactly the scratch this restore creates.
|
|
||||||
scratch := filepath.Join(root, "scratch")
|
|
||||||
require.NoError(t, fs.MkdirAll(scratch, 0o755))
|
|
||||||
t.Setenv("TMPDIR", scratch)
|
|
||||||
|
|
||||||
gate := newBlockingBlobStorer(storer)
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
v := &Vaultik{
|
|
||||||
Config: cfg,
|
|
||||||
Storage: gate,
|
|
||||||
Fs: fs,
|
|
||||||
Stdout: io.Discard,
|
|
||||||
Stderr: io.Discard,
|
|
||||||
UI: ui.NewWithColor(io.Discard, false),
|
|
||||||
ctx: ctx,
|
|
||||||
cancel: cancel,
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
|
||||||
opReturned atomic.Bool
|
|
||||||
restoreErr error
|
|
||||||
)
|
|
||||||
|
|
||||||
stop := v.StartOperation(func() {
|
|
||||||
defer opReturned.Store(true)
|
|
||||||
|
|
||||||
restoreErr = v.Restore(&RestoreOptions{
|
|
||||||
SnapshotID: snapshotID,
|
|
||||||
TargetDir: restoreDir,
|
|
||||||
})
|
|
||||||
})
|
|
||||||
|
|
||||||
// Wait until the restore is blocked mid-download; its decrypted
|
|
||||||
// scratch files exist by now.
|
|
||||||
select {
|
|
||||||
case <-gate.entered:
|
|
||||||
case <-time.After(30 * time.Second):
|
|
||||||
t.Fatal("restore never reached the blob-download phase")
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NotEmpty(t, scratchEntries(t, scratch),
|
|
||||||
"expected decrypted scratch files to exist mid-restore")
|
|
||||||
|
|
||||||
// Stop the operation the way the fx OnStop hook does.
|
|
||||||
stopCtx, stopCancel := context.WithTimeout(
|
|
||||||
context.Background(), 30*time.Second)
|
|
||||||
defer stopCancel()
|
|
||||||
|
|
||||||
require.True(t, stop(stopCtx),
|
|
||||||
"stop timed out; the operation goroutine did not return")
|
|
||||||
|
|
||||||
// stop returns only once the operation goroutine has returned, so its
|
|
||||||
// cleanup defers have run by the time we read these.
|
|
||||||
require.True(t, opReturned.Load(),
|
|
||||||
"stop returned before the operation goroutine finished")
|
|
||||||
require.ErrorIs(t, restoreErr, context.Canceled)
|
|
||||||
require.Empty(t, scratchEntries(t, scratch),
|
|
||||||
"decrypted scratch files remained after the interrupt")
|
|
||||||
}
|
|
||||||
|
|
||||||
// scratchEntries returns the vaultik blob-cache and snapshot-database
|
|
||||||
// scratch entries currently present in dir.
|
|
||||||
func scratchEntries(t *testing.T, dir string) []string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var matches []string
|
|
||||||
|
|
||||||
for _, pattern := range []string{
|
|
||||||
"vaultik-blobcache-*", "vaultik-restore-*",
|
|
||||||
} {
|
|
||||||
found, err := filepath.Glob(filepath.Join(dir, pattern))
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
matches = append(matches, found...)
|
|
||||||
}
|
|
||||||
|
|
||||||
return matches
|
|
||||||
}
|
|
||||||
@@ -1,44 +0,0 @@
|
|||||||
package vaultik_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
"sneak.berlin/go/vaultik/internal/ui"
|
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestRestoreRejectsMalformedKeyBeforeDownload verifies that a malformed
|
|
||||||
// age secret key stops restore at the parse step: nothing is fetched from
|
|
||||||
// the store, and the error does not echo the key value (which is secret).
|
|
||||||
func TestRestoreRejectsMalformedKeyBeforeDownload(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const malformed = "this-is-not-a-valid-age-key"
|
|
||||||
|
|
||||||
mock := NewMockStorer()
|
|
||||||
|
|
||||||
v := &vaultik.Vaultik{
|
|
||||||
Config: &config.Config{AgeSecretKey: malformed},
|
|
||||||
Storage: mock,
|
|
||||||
Stdout: io.Discard,
|
|
||||||
Stderr: io.Discard,
|
|
||||||
UI: ui.NewWithColor(io.Discard, false),
|
|
||||||
}
|
|
||||||
v.SetContext(context.Background())
|
|
||||||
|
|
||||||
err := v.Restore(&vaultik.RestoreOptions{
|
|
||||||
SnapshotID: "any-snapshot",
|
|
||||||
TargetDir: t.TempDir(),
|
|
||||||
})
|
|
||||||
require.Error(t, err)
|
|
||||||
require.NotContains(t, err.Error(), malformed,
|
|
||||||
"error must not echo the key value")
|
|
||||||
require.Contains(t, err.Error(), "age_secret_key",
|
|
||||||
"error should name the configuration source")
|
|
||||||
require.Empty(t, mock.GetCalls(),
|
|
||||||
"a malformed key must fail before anything is fetched")
|
|
||||||
}
|
|
||||||
@@ -1,306 +0,0 @@
|
|||||||
package vaultik //nolint:testpackage // drives restore through unexported session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"syscall"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
"sneak.berlin/go/vaultik/internal/ui"
|
|
||||||
)
|
|
||||||
|
|
||||||
// errSpyWrite is the injected write failure used to exercise the
|
|
||||||
// partial-file cleanup path.
|
|
||||||
var errSpyWrite = errors.New("injected write failure")
|
|
||||||
|
|
||||||
// modeSpyFs wraps a real filesystem so restore tests can observe and
|
|
||||||
// perturb the single output file whose path contains watch. It records
|
|
||||||
// the on-disk permission bits seen at the moment content is first
|
|
||||||
// written (the window during which another local user could read it),
|
|
||||||
// and can inject a write failure or append trailing bytes on close.
|
|
||||||
type modeSpyFs struct {
|
|
||||||
afero.Fs
|
|
||||||
|
|
||||||
watch string
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
writeModes []os.FileMode
|
|
||||||
failWrite bool
|
|
||||||
trailing int
|
|
||||||
}
|
|
||||||
|
|
||||||
//nolint:ireturn // afero.Fs.OpenFile is defined to return the interface
|
|
||||||
func (m *modeSpyFs) OpenFile(
|
|
||||||
name string, flag int, perm os.FileMode,
|
|
||||||
) (afero.File, error) {
|
|
||||||
f, err := m.Fs.OpenFile(name, flag, perm)
|
|
||||||
if err != nil || !strings.Contains(name, m.watch) {
|
|
||||||
return f, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return &modeSpyFile{File: f, fs: m, path: name}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type modeSpyFile struct {
|
|
||||||
afero.File
|
|
||||||
|
|
||||||
fs *modeSpyFs
|
|
||||||
path string
|
|
||||||
written bool
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *modeSpyFile) Write(p []byte) (int, error) {
|
|
||||||
if !f.written {
|
|
||||||
f.written = true
|
|
||||||
|
|
||||||
info, err := f.fs.Stat(f.path)
|
|
||||||
if err == nil {
|
|
||||||
f.fs.mu.Lock()
|
|
||||||
f.fs.writeModes = append(f.fs.writeModes, info.Mode().Perm())
|
|
||||||
f.fs.mu.Unlock()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if f.fs.failWrite {
|
|
||||||
return 0, errSpyWrite
|
|
||||||
}
|
|
||||||
|
|
||||||
return f.File.Write(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *modeSpyFile) Close() error {
|
|
||||||
if f.fs.trailing > 0 {
|
|
||||||
_, _ = f.File.Write(bytes.Repeat([]byte{'x'}, f.fs.trailing))
|
|
||||||
}
|
|
||||||
|
|
||||||
return f.File.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
// backupOneFile writes a single source file with the given mode and
|
|
||||||
// backs it up into a fresh file storer, returning everything a restore
|
|
||||||
// needs. The index database is closed before returning so the restore
|
|
||||||
// half runs from the exported metadata and remote bytes only.
|
|
||||||
func backupOneFile(
|
|
||||||
ctx context.Context, t *testing.T, fs afero.Fs, tempDir, name string,
|
|
||||||
content []byte, mode os.FileMode,
|
|
||||||
) (*config.Config, *storage.FileStorer, string, string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
dataDir := filepath.Join(tempDir, "src")
|
|
||||||
require.NoError(t, fs.MkdirAll(dataDir, 0o755))
|
|
||||||
|
|
||||||
srcPath := filepath.Join(dataDir, name)
|
|
||||||
require.NoError(t, afero.WriteFile(fs, srcPath, content, mode))
|
|
||||||
require.NoError(t, fs.Chmod(srcPath, mode))
|
|
||||||
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
storer, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
cfg := &config.Config{
|
|
||||||
AgeRecipients: []string{
|
|
||||||
"age1ezrjmfpwsc95svdg0y54mums3zevgzu0x0ecq2f7tp8a05gl0sjq9q9wjg",
|
|
||||||
},
|
|
||||||
AgeSecretKey: "AGE-SECRET-KEY-19CR5YSFW59HM4TLD6GXVEDMZFTVVF7PPHKU" +
|
|
||||||
"T68TXSFPK7APHXA2QS2NJA5",
|
|
||||||
CompressionLevel: 3,
|
|
||||||
Hostname: "test-host",
|
|
||||||
BlobSizeLimit: config.Size(5 * 1024 * 1024),
|
|
||||||
}
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
sm := snapshot.NewSnapshotManager(snapshot.SnapshotManagerParams{
|
|
||||||
Repos: repos,
|
|
||||||
Storage: storer,
|
|
||||||
Config: cfg,
|
|
||||||
})
|
|
||||||
sm.SetFilesystem(fs)
|
|
||||||
|
|
||||||
scanner := snapshot.NewScanner(snapshot.ScannerConfig{
|
|
||||||
FS: fs,
|
|
||||||
Storage: storer,
|
|
||||||
ChunkSize: 4 * 1024 * 1024,
|
|
||||||
MaxBlobSize: 5 * 1024 * 1024,
|
|
||||||
CompressionLevel: cfg.CompressionLevel,
|
|
||||||
AgeRecipients: cfg.AgeRecipients,
|
|
||||||
Repositories: repos,
|
|
||||||
})
|
|
||||||
|
|
||||||
snapshotID, err := sm.CreateSnapshotWithName(
|
|
||||||
ctx, cfg.Hostname, "perms", "test-version", "test-git")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = scanner.Scan(ctx, dataDir, snapshotID)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
require.NoError(t, sm.CompleteSnapshot(ctx, snapshotID))
|
|
||||||
require.NoError(t, sm.ExportSnapshotMetadata(ctx, dbPath, snapshotID))
|
|
||||||
require.NoError(t, db.Close())
|
|
||||||
|
|
||||||
return cfg, storer, snapshotID, srcPath
|
|
||||||
}
|
|
||||||
|
|
||||||
// restoredPathFor returns where backupOneFile's source lands under a
|
|
||||||
// restore target: restore recreates each file at its original absolute
|
|
||||||
// path beneath TargetDir.
|
|
||||||
func restoredPathFor(restoreDir, srcPath string) string {
|
|
||||||
return filepath.Join(restoreDir, srcPath)
|
|
||||||
}
|
|
||||||
|
|
||||||
// withUmask022 forces the process umask to 022 for the duration of a
|
|
||||||
// test, so the difference between a 0600 create and a default create is
|
|
||||||
// observable. Restored serially (no t.Parallel) so it does not race
|
|
||||||
// other tests.
|
|
||||||
func withUmask022(t *testing.T) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
old := syscall.Umask(0o022)
|
|
||||||
|
|
||||||
t.Cleanup(func() { syscall.Umask(old) })
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestRestoreCreatesFileNeverWiderThanStoredMode checks that a file with
|
|
||||||
// a restrictive stored mode (0600) is never observable with a wider mode
|
|
||||||
// while its content is being written, and ends at its stored mode.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // sets the process umask; must run serially
|
|
||||||
func TestRestoreCreatesFileNeverWiderThanStoredMode(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
withUmask022(t)
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
content := randomBytes(t, 4096)
|
|
||||||
|
|
||||||
cfg, storer, snapshotID, srcPath := backupOneFile(
|
|
||||||
ctx, t, fs, tempDir, "secret.bin", content, 0o600)
|
|
||||||
|
|
||||||
restoreDir := filepath.Join(tempDir, "restored")
|
|
||||||
spy := &modeSpyFs{Fs: fs, watch: "secret.bin"}
|
|
||||||
|
|
||||||
v := newRestoreVaultik(ctx, cfg, storer, spy)
|
|
||||||
require.NoError(t, v.Restore(&RestoreOptions{
|
|
||||||
SnapshotID: snapshotID,
|
|
||||||
TargetDir: restoreDir,
|
|
||||||
}))
|
|
||||||
|
|
||||||
spy.mu.Lock()
|
|
||||||
observed := append([]os.FileMode(nil), spy.writeModes...)
|
|
||||||
spy.mu.Unlock()
|
|
||||||
|
|
||||||
require.NotEmpty(t, observed,
|
|
||||||
"spy never saw the output file being written")
|
|
||||||
|
|
||||||
for _, m := range observed {
|
|
||||||
assert.Equalf(t, os.FileMode(0o600), m,
|
|
||||||
"file was observable at mode %o during write; must be 0600", m)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The stored mode is applied after the content is written.
|
|
||||||
info, err := fs.Stat(restoredPathFor(restoreDir, srcPath))
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, os.FileMode(0o600), info.Mode().Perm())
|
|
||||||
|
|
||||||
got, err := afero.ReadFile(fs, restoredPathFor(restoreDir, srcPath))
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.True(t, bytes.Equal(got, content))
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestRestoreRemovesPartialFileOnWriteFailure checks that a file whose
|
|
||||||
// content write fails is not left behind.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // sets the process umask; must run serially
|
|
||||||
func TestRestoreRemovesPartialFileOnWriteFailure(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
withUmask022(t)
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
cfg, storer, snapshotID, srcPath := backupOneFile(
|
|
||||||
ctx, t, fs, tempDir, "doomed.bin", randomBytes(t, 4096), 0o600)
|
|
||||||
|
|
||||||
restoreDir := filepath.Join(tempDir, "restored")
|
|
||||||
spy := &modeSpyFs{Fs: fs, watch: "doomed.bin", failWrite: true}
|
|
||||||
|
|
||||||
v := newRestoreVaultik(ctx, cfg, storer, spy)
|
|
||||||
err := v.Restore(&RestoreOptions{
|
|
||||||
SnapshotID: snapshotID,
|
|
||||||
TargetDir: restoreDir,
|
|
||||||
})
|
|
||||||
require.Error(t, err, "restore should fail when the write fails")
|
|
||||||
|
|
||||||
exists, err := afero.Exists(fs, restoredPathFor(restoreDir, srcPath))
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.False(t, exists, "partial file must be removed after a failed write")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestVerifyRejectsTrailingBytes checks that --verify fails a restored
|
|
||||||
// file that has bytes past its last chunk.
|
|
||||||
//
|
|
||||||
//nolint:paralleltest // sets the process umask; must run serially
|
|
||||||
func TestVerifyRejectsTrailingBytes(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
withUmask022(t)
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
cfg, storer, snapshotID, _ := backupOneFile(
|
|
||||||
ctx, t, fs, tempDir, "padded.bin", randomBytes(t, 4096), 0o600)
|
|
||||||
|
|
||||||
restoreDir := filepath.Join(tempDir, "restored")
|
|
||||||
// Append one byte to the file as it is written, so its content still
|
|
||||||
// matches the stored chunks but it is one byte too long.
|
|
||||||
spy := &modeSpyFs{Fs: fs, watch: "padded.bin", trailing: 1}
|
|
||||||
|
|
||||||
v := newRestoreVaultik(ctx, cfg, storer, spy)
|
|
||||||
err := v.Restore(&RestoreOptions{
|
|
||||||
SnapshotID: snapshotID,
|
|
||||||
TargetDir: restoreDir,
|
|
||||||
Verify: true,
|
|
||||||
})
|
|
||||||
require.Error(t, err, "verify should fail on a file with trailing bytes")
|
|
||||||
assert.ErrorIs(t, err, errFilesFailedVerify)
|
|
||||||
}
|
|
||||||
|
|
||||||
// newRestoreVaultik builds a Vaultik wired for a restore-only test.
|
|
||||||
func newRestoreVaultik(
|
|
||||||
ctx context.Context, cfg *config.Config, storer storage.Storer, fs afero.Fs,
|
|
||||||
) *Vaultik {
|
|
||||||
v := &Vaultik{
|
|
||||||
Config: cfg,
|
|
||||||
Storage: storer,
|
|
||||||
Fs: fs,
|
|
||||||
Stdout: io.Discard,
|
|
||||||
Stderr: io.Discard,
|
|
||||||
UI: ui.NewWithColor(io.Discard, false),
|
|
||||||
}
|
|
||||||
v.SetContext(ctx)
|
|
||||||
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
@@ -171,13 +171,10 @@ func (p *restorePlan) finishFile(fileID types.FileID) {
|
|||||||
// downloaded next, after which it — together with any other pending
|
// downloaded next, after which it — together with any other pending
|
||||||
// files whose blob sets become empty — moves to the ready queue.
|
// files whose blob sets become empty — moves to the ready queue.
|
||||||
//
|
//
|
||||||
// The second return value is false when no file needs a download, so a
|
// The zero FileID return means nothing is pending.
|
||||||
// genuine file carrying the nil UUID is picked rather than mistaken for
|
func (p *restorePlan) pickNextDownload() types.FileID {
|
||||||
// "nothing left".
|
|
||||||
func (p *restorePlan) pickNextDownload() (types.FileID, bool) {
|
|
||||||
var best types.FileID
|
var best types.FileID
|
||||||
|
|
||||||
found := false
|
|
||||||
bestCount := math.MaxInt
|
bestCount := math.MaxInt
|
||||||
|
|
||||||
var bestID string
|
var bestID string
|
||||||
@@ -191,15 +188,14 @@ func (p *restorePlan) pickNextDownload() (types.FileID, bool) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
idStr := id.String()
|
idStr := id.String()
|
||||||
if !found || n < bestCount || (n == bestCount && idStr < bestID) {
|
if n < bestCount || (n == bestCount && (best.IsZero() || idStr < bestID)) {
|
||||||
best = id
|
best = id
|
||||||
found = true
|
|
||||||
bestCount = n
|
bestCount = n
|
||||||
bestID = idStr
|
bestID = idStr
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return best, found
|
return best
|
||||||
}
|
}
|
||||||
|
|
||||||
// blobsNeeded returns the uncached blob hashes for fileID in any order.
|
// blobsNeeded returns the uncached blob hashes for fileID in any order.
|
||||||
|
|||||||
@@ -1,88 +0,0 @@
|
|||||||
package vaultik //nolint:testpackage // inspects unexported restore plan internals
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"math"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
"sneak.berlin/go/vaultik/internal/types"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestPickNextDownloadReturnsNilUUIDFile proves a genuine pending file
|
|
||||||
// carrying the nil UUID is picked for download rather than mistaken for
|
|
||||||
// "nothing left" — the bug that could abandon every remaining file.
|
|
||||||
func TestPickNextDownloadReturnsNilUUIDFile(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
var nilID types.FileID // zero value is the nil UUID
|
|
||||||
|
|
||||||
plan := &restorePlan{
|
|
||||||
fileBlobs: map[types.FileID]map[string]struct{}{
|
|
||||||
nilID: {"blobhash": {}},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
id, ok := plan.pickNextDownload()
|
|
||||||
require.True(t, ok,
|
|
||||||
"pickNextDownload treated a pending nil-UUID file as nothing to do")
|
|
||||||
require.True(t, id.IsZero(), "expected the nil-UUID file to be picked")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestPickNextDownloadEmptyPlan confirms the second return value is false
|
|
||||||
// only when no file needs a download.
|
|
||||||
func TestPickNextDownloadEmptyPlan(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
plan := &restorePlan{
|
|
||||||
fileBlobs: map[types.FileID]map[string]struct{}{},
|
|
||||||
}
|
|
||||||
|
|
||||||
_, ok := plan.pickNextDownload()
|
|
||||||
require.False(t, ok, "pickNextDownload reported work on an empty plan")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestRunRestoreLoopFailsOnAbandonedFiles proves the loop returns an
|
|
||||||
// error rather than silent success when files remain pending after it
|
|
||||||
// can make no further progress.
|
|
||||||
func TestRunRestoreLoopFailsOnAbandonedFiles(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
db, err := database.NewTestDB()
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = db.Close() })
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
cache, err := newBlobDiskCache(math.MaxInt64)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = cache.Close() })
|
|
||||||
|
|
||||||
v := &Vaultik{ctx: ctx}
|
|
||||||
session := &restoreSession{
|
|
||||||
v: v,
|
|
||||||
ctx: ctx,
|
|
||||||
repos: repos,
|
|
||||||
sweeper: newRestoreSweeper(ctx, repos, cache, 1),
|
|
||||||
result: &RestoreResult{},
|
|
||||||
}
|
|
||||||
|
|
||||||
// A file that is still pending but whose uncached-blob set is empty
|
|
||||||
// and which was never queued as ready: the loop can neither restore
|
|
||||||
// nor download it. This is the abandonment the guard must catch.
|
|
||||||
var stuck types.FileID
|
|
||||||
|
|
||||||
plan := &restorePlan{
|
|
||||||
fileBlobs: map[types.FileID]map[string]struct{}{stuck: {}},
|
|
||||||
blobFiles: map[string]map[types.FileID]struct{}{},
|
|
||||||
cached: map[string]struct{}{},
|
|
||||||
}
|
|
||||||
|
|
||||||
err = v.runRestoreLoop(session, plan, map[types.FileID]*database.File{}, 0)
|
|
||||||
require.ErrorIs(t, err, errRestoreIncomplete)
|
|
||||||
}
|
|
||||||
@@ -1,251 +0,0 @@
|
|||||||
package vaultik_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
"sneak.berlin/go/vaultik/internal/storage"
|
|
||||||
"sneak.berlin/go/vaultik/internal/ui"
|
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestRestoreAndDeepVerifyRejectSwappedDatabase proves that swapping two
|
|
||||||
// snapshots' encrypted databases on the store is caught. age decryption
|
|
||||||
// alone proves only that a database is readable; without an identity check
|
|
||||||
// restore would happily write the wrong snapshot's files and deep verify
|
|
||||||
// would report success. After the swap, restore and deep verify of A both
|
|
||||||
// fail, and the error names the snapshot the database actually holds (B).
|
|
||||||
func TestRestoreAndDeepVerifyRejectSwappedDatabase(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
|
|
||||||
chunkSize := int64(64 * 1024)
|
|
||||||
|
|
||||||
storer, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
// Two snapshots with different content, backed up into one shared
|
|
||||||
// store. Different names give them different remote keys, so their
|
|
||||||
// metadata directories are distinct and can be tampered with alone.
|
|
||||||
dataA := filepath.Join(tempDir, "srcA")
|
|
||||||
require.NoError(t, fs.MkdirAll(dataA, 0o755))
|
|
||||||
require.NoError(t, afero.WriteFile(fs, filepath.Join(dataA, "a.bin"),
|
|
||||||
bytesPattern("alpha-", int(chunkSize*2)), 0o644))
|
|
||||||
|
|
||||||
dataB := filepath.Join(tempDir, "srcB")
|
|
||||||
require.NoError(t, fs.MkdirAll(dataB, 0o755))
|
|
||||||
require.NoError(t, afero.WriteFile(fs, filepath.Join(dataB, "b.bin"),
|
|
||||||
bytesPattern("beta-", int(chunkSize*2)), 0o644))
|
|
||||||
|
|
||||||
idA := backupNamedSnapshotToStore(ctx, t, fs, dataA, storer,
|
|
||||||
filepath.Join(tempDir, "idxA.sqlite"), "alpha")
|
|
||||||
idB := backupNamedSnapshotToStore(ctx, t, fs, dataB, storer,
|
|
||||||
filepath.Join(tempDir, "idxB.sqlite"), "beta")
|
|
||||||
require.NotEqual(t, idA, idB)
|
|
||||||
|
|
||||||
keyA := snapshot.RemoteSnapshotKey(idA)
|
|
||||||
keyB := snapshot.RemoteSnapshotKey(idB)
|
|
||||||
require.NotEqual(t, keyA, keyB)
|
|
||||||
|
|
||||||
// Baseline: each snapshot verifies against its own intact metadata.
|
|
||||||
require.NoError(t, newStoreClient(ctx, t, fs, storer).RunDeepVerify(
|
|
||||||
idA, &vaultik.VerifyOptions{Deep: true}))
|
|
||||||
require.NoError(t, newStoreClient(ctx, t, fs, storer).RunDeepVerify(
|
|
||||||
idB, &vaultik.VerifyOptions{Deep: true}))
|
|
||||||
|
|
||||||
// Swap the two snapshots' encrypted databases on the store.
|
|
||||||
swapStoreFiles(t, fs,
|
|
||||||
filepath.Join(storeDir, "metadata", keyA, "db.zst.age"),
|
|
||||||
filepath.Join(storeDir, "metadata", keyB, "db.zst.age"))
|
|
||||||
|
|
||||||
// Restore of A now decrypts B's database; the identity check must
|
|
||||||
// reject it and name the snapshot it actually found.
|
|
||||||
restoreErr := newStoreClient(ctx, t, fs, storer).Restore(&vaultik.RestoreOptions{
|
|
||||||
SnapshotID: idA,
|
|
||||||
TargetDir: filepath.Join(tempDir, "restoreA"),
|
|
||||||
})
|
|
||||||
require.Error(t, restoreErr)
|
|
||||||
require.ErrorContains(t, restoreErr, idB)
|
|
||||||
|
|
||||||
// Deep verify of A must reject the swapped database for the same reason.
|
|
||||||
verifyErr := newStoreClient(ctx, t, fs, storer).RunDeepVerify(
|
|
||||||
idA, &vaultik.VerifyOptions{Deep: true})
|
|
||||||
require.Error(t, verifyErr)
|
|
||||||
require.ErrorContains(t, verifyErr, idB)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestDeepVerifyRejectsSwappedDatabaseWithEmptyManifest covers the case the
|
|
||||||
// issue calls out: swapping in a database whose blob set is empty and
|
|
||||||
// pairing it with an equally empty manifest. The manifest then agrees with
|
|
||||||
// the database, so every blob-level check passes and deep verify used to
|
|
||||||
// report success with zero blobs verified. The identity check rejects it
|
|
||||||
// before any blob check runs.
|
|
||||||
func TestDeepVerifyRejectsSwappedDatabaseWithEmptyManifest(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
|
|
||||||
chunkSize := int64(64 * 1024)
|
|
||||||
|
|
||||||
storer, err := storage.NewFileStorer(storeDir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
// Snapshot A: real content, so its manifest lists blobs.
|
|
||||||
dataA := filepath.Join(tempDir, "srcA")
|
|
||||||
require.NoError(t, fs.MkdirAll(dataA, 0o755))
|
|
||||||
require.NoError(t, afero.WriteFile(fs, filepath.Join(dataA, "a.bin"),
|
|
||||||
bytesPattern("alpha-", int(chunkSize*2)), 0o644))
|
|
||||||
idA := backupNamedSnapshotToStore(ctx, t, fs, dataA, storer,
|
|
||||||
filepath.Join(tempDir, "idxA.sqlite"), "alpha")
|
|
||||||
|
|
||||||
// Snapshot C: a single empty file, so it references no blobs and its
|
|
||||||
// manifest is empty. This is the database/manifest pair an attacker
|
|
||||||
// would swap in to make the blob checks vacuously pass.
|
|
||||||
dataC := filepath.Join(tempDir, "srcC")
|
|
||||||
require.NoError(t, fs.MkdirAll(dataC, 0o755))
|
|
||||||
require.NoError(t, afero.WriteFile(fs,
|
|
||||||
filepath.Join(dataC, "empty.bin"), []byte{}, 0o644))
|
|
||||||
idC := backupNamedSnapshotToStore(ctx, t, fs, dataC, storer,
|
|
||||||
filepath.Join(tempDir, "idxC.sqlite"), "charlie")
|
|
||||||
|
|
||||||
keyA := snapshot.RemoteSnapshotKey(idA)
|
|
||||||
keyC := snapshot.RemoteSnapshotKey(idC)
|
|
||||||
|
|
||||||
// Replace A's database and manifest with C's empty pair.
|
|
||||||
copyStoreFile(t, fs,
|
|
||||||
filepath.Join(storeDir, "metadata", keyC, "db.zst.age"),
|
|
||||||
filepath.Join(storeDir, "metadata", keyA, "db.zst.age"))
|
|
||||||
copyStoreFile(t, fs,
|
|
||||||
filepath.Join(storeDir, "metadata", keyC, "manifest.json.zst"),
|
|
||||||
filepath.Join(storeDir, "metadata", keyA, "manifest.json.zst"))
|
|
||||||
|
|
||||||
verifyErr := newStoreClient(ctx, t, fs, storer).RunDeepVerify(
|
|
||||||
idA, &vaultik.VerifyOptions{Deep: true})
|
|
||||||
require.Error(t, verifyErr)
|
|
||||||
require.ErrorContains(t, verifyErr, idC)
|
|
||||||
}
|
|
||||||
|
|
||||||
// backupNamedSnapshotToStore backs up dataDir into the shared storer under
|
|
||||||
// the given snapshot name and returns the human snapshot ID. Two snapshots
|
|
||||||
// backed up under different names get different remote keys, so their
|
|
||||||
// metadata directories on the store are distinct.
|
|
||||||
func backupNamedSnapshotToStore(
|
|
||||||
ctx context.Context, t *testing.T, fs afero.Fs,
|
|
||||||
dataDir string, storer storage.Storer, dbPath, name string,
|
|
||||||
) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
const (
|
|
||||||
chunkSize = int64(64 * 1024)
|
|
||||||
maxBlobSize = int64(512 * 1024)
|
|
||||||
)
|
|
||||||
|
|
||||||
cfg := &config.Config{
|
|
||||||
AgeRecipients: []string{testAgePublicKey},
|
|
||||||
AgeSecretKey: testAgeSecretKey,
|
|
||||||
CompressionLevel: 3,
|
|
||||||
Hostname: testHostname,
|
|
||||||
}
|
|
||||||
|
|
||||||
db, err := database.New(ctx, dbPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
repos := database.NewRepositories(db)
|
|
||||||
|
|
||||||
sm := snapshot.NewSnapshotManager(snapshot.SnapshotManagerParams{
|
|
||||||
Repos: repos,
|
|
||||||
Storage: storer,
|
|
||||||
Config: cfg,
|
|
||||||
})
|
|
||||||
sm.SetFilesystem(fs)
|
|
||||||
|
|
||||||
scanner := snapshot.NewScanner(snapshot.ScannerConfig{
|
|
||||||
FS: fs,
|
|
||||||
Storage: storer,
|
|
||||||
ChunkSize: chunkSize,
|
|
||||||
MaxBlobSize: maxBlobSize,
|
|
||||||
CompressionLevel: cfg.CompressionLevel,
|
|
||||||
AgeRecipients: cfg.AgeRecipients,
|
|
||||||
Repositories: repos,
|
|
||||||
})
|
|
||||||
|
|
||||||
snapshotID, err := sm.CreateSnapshotWithName(
|
|
||||||
ctx, cfg.Hostname, name, "test-version", "test-git")
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
_, err = scanner.Scan(ctx, dataDir, snapshotID)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
require.NoError(t, sm.CompleteSnapshot(ctx, snapshotID))
|
|
||||||
require.NoError(t, sm.ExportSnapshotMetadata(ctx, dbPath, snapshotID))
|
|
||||||
require.NoError(t, db.Close())
|
|
||||||
|
|
||||||
return snapshotID
|
|
||||||
}
|
|
||||||
|
|
||||||
// newStoreClient builds a Vaultik that reads only from the store: the
|
|
||||||
// secret key, the storer, and a filesystem, with no local index. This is
|
|
||||||
// what restore and deep verify need.
|
|
||||||
func newStoreClient(
|
|
||||||
ctx context.Context, t *testing.T, fs afero.Fs, storer storage.Storer,
|
|
||||||
) *vaultik.Vaultik {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
v := &vaultik.Vaultik{
|
|
||||||
Config: &config.Config{
|
|
||||||
AgeSecretKey: testAgeSecretKey,
|
|
||||||
Hostname: testHostname,
|
|
||||||
},
|
|
||||||
Storage: storer,
|
|
||||||
Fs: fs,
|
|
||||||
Stdout: io.Discard,
|
|
||||||
Stderr: io.Discard,
|
|
||||||
UI: ui.NewWithColor(io.Discard, false),
|
|
||||||
}
|
|
||||||
v.SetContext(ctx)
|
|
||||||
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
|
|
||||||
// swapStoreFiles exchanges the contents of two files on the store.
|
|
||||||
func swapStoreFiles(t *testing.T, fs afero.Fs, a, b string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
dataA, err := afero.ReadFile(fs, a)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
dataB, err := afero.ReadFile(fs, b)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
require.NoError(t, afero.WriteFile(fs, a, dataB, 0o644))
|
|
||||||
require.NoError(t, afero.WriteFile(fs, b, dataA, 0o644))
|
|
||||||
}
|
|
||||||
|
|
||||||
// copyStoreFile overwrites dst with the contents of src on the store.
|
|
||||||
func copyStoreFile(t *testing.T, fs afero.Fs, src, dst string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
data, err := afero.ReadFile(fs, src)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
require.NoError(t, afero.WriteFile(fs, dst, data, 0o644))
|
|
||||||
}
|
|
||||||
@@ -1,73 +0,0 @@
|
|||||||
package vaultik //nolint:testpackage // inspects unexported snapshot-db materialization
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
|
||||||
)
|
|
||||||
|
|
||||||
// genuineSnapshotDBBytes returns the on-disk bytes of a real snapshot
|
|
||||||
// database (the full schema applied).
|
|
||||||
func genuineSnapshotDBBytes(t *testing.T) []byte {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
path := filepath.Join(t.TempDir(), "snapshot.db")
|
|
||||||
|
|
||||||
db, err := database.New(context.Background(), path)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NoError(t, db.Close())
|
|
||||||
|
|
||||||
data, err := os.ReadFile(path) //nolint:gosec // G304: test-controlled temp path
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
return data
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMaterializeSnapshotDBPrivateDir proves the decrypted database lands
|
|
||||||
// in a private (0700) directory and opens read-only.
|
|
||||||
func TestMaterializeSnapshotDBPrivateDir(t *testing.T) {
|
|
||||||
dbData := genuineSnapshotDBBytes(t)
|
|
||||||
|
|
||||||
t.Setenv("TMPDIR", t.TempDir())
|
|
||||||
|
|
||||||
v := &Vaultik{ctx: context.Background(), Fs: afero.NewOsFs()}
|
|
||||||
|
|
||||||
db, dir, err := v.materializeSnapshotDB(dbData)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = db.Close()
|
|
||||||
_ = os.RemoveAll(dir)
|
|
||||||
})
|
|
||||||
|
|
||||||
info, err := os.Stat(dir)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.Equal(t, os.FileMode(0o700), info.Mode().Perm(),
|
|
||||||
"snapshot database directory must not be world-readable")
|
|
||||||
|
|
||||||
_, err = db.Conn().ExecContext(context.Background(),
|
|
||||||
"CREATE TABLE probe_readonly (x)")
|
|
||||||
require.Error(t, err, "materialized snapshot database must be read-only")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMaterializeSnapshotDBRemovesDirOnOpenFailure proves a failed open
|
|
||||||
// leaves no temp directory behind.
|
|
||||||
func TestMaterializeSnapshotDBRemovesDirOnOpenFailure(t *testing.T) {
|
|
||||||
base := t.TempDir()
|
|
||||||
|
|
||||||
t.Setenv("TMPDIR", base)
|
|
||||||
|
|
||||||
v := &Vaultik{ctx: context.Background(), Fs: afero.NewOsFs()}
|
|
||||||
|
|
||||||
_, _, err := v.materializeSnapshotDB([]byte("this is not a sqlite database"))
|
|
||||||
require.Error(t, err)
|
|
||||||
|
|
||||||
entries, rerr := os.ReadDir(base)
|
|
||||||
require.NoError(t, rerr)
|
|
||||||
require.Empty(t, entries, "temp directory left behind after open failure")
|
|
||||||
}
|
|
||||||
@@ -1,191 +0,0 @@
|
|||||||
package vaultik_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/spf13/afero"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
"sneak.berlin/go/vaultik/internal/ui"
|
|
||||||
"sneak.berlin/go/vaultik/internal/vaultik"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestShallowVerifyDetectsWrongBlobSize backs up a real snapshot, runs
|
|
||||||
// shallow verify (which passes and reports exactly the blobs it checked),
|
|
||||||
// then grows one stored blob so its size no longer matches the manifest.
|
|
||||||
// Shallow verify must then fail, count the grown blob as a size mismatch,
|
|
||||||
// and drop it from the verified count rather than continuing to report it
|
|
||||||
// as checked.
|
|
||||||
func TestShallowVerifyDetectsWrongBlobSize(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
|
|
||||||
dataDir := filepath.Join(tempDir, "source")
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
chunkSize := int64(32 * 1024)
|
|
||||||
maxBlobSize := int64(128 * 1024)
|
|
||||||
|
|
||||||
// Enough data to span several blobs, so the mismatch count and the
|
|
||||||
// dropped verified count are both meaningful.
|
|
||||||
require.NoError(t, fs.MkdirAll(dataDir, 0o755))
|
|
||||||
require.NoError(t, afero.WriteFile(fs,
|
|
||||||
filepath.Join(dataDir, "data.bin"),
|
|
||||||
bytesPattern("shallow-", int(maxBlobSize*4)), 0o644))
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
cfg, storer, snapshotID := runFileStorageBackup(
|
|
||||||
ctx, t, fs, dataDir, storeDir, dbPath, chunkSize, maxBlobSize)
|
|
||||||
|
|
||||||
newVerifier := func(out io.Writer) *vaultik.Vaultik {
|
|
||||||
v := &vaultik.Vaultik{
|
|
||||||
Config: cfg,
|
|
||||||
Storage: storer,
|
|
||||||
Fs: fs,
|
|
||||||
Stdout: out,
|
|
||||||
Stderr: io.Discard,
|
|
||||||
UI: ui.NewWithColor(io.Discard, false),
|
|
||||||
}
|
|
||||||
v.SetContext(ctx)
|
|
||||||
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
|
|
||||||
var out bytes.Buffer
|
|
||||||
|
|
||||||
require.NoError(t,
|
|
||||||
newVerifier(&out).VerifySnapshotWithOptions(
|
|
||||||
snapshotID, &vaultik.VerifyOptions{JSON: true}),
|
|
||||||
"shallow verify should pass on a healthy snapshot")
|
|
||||||
|
|
||||||
healthy := decodeVerifyResult(t, out.Bytes())
|
|
||||||
require.Equal(t, "ok", healthy.Status)
|
|
||||||
require.Positive(t, healthy.BlobCount)
|
|
||||||
require.Equal(t, healthy.BlobCount, healthy.Verified,
|
|
||||||
"shallow verify must report exactly the blobs it checked")
|
|
||||||
require.Zero(t, healthy.Mismatched)
|
|
||||||
|
|
||||||
// Grow one stored blob so its size no longer matches the manifest.
|
|
||||||
growOneBlob(t, fs, filepath.Join(storeDir, "blobs"))
|
|
||||||
|
|
||||||
out.Reset()
|
|
||||||
err := newVerifier(&out).VerifySnapshotWithOptions(
|
|
||||||
snapshotID, &vaultik.VerifyOptions{JSON: true})
|
|
||||||
require.Error(t, err,
|
|
||||||
"shallow verify must fail when a blob's stored size differs from the manifest")
|
|
||||||
|
|
||||||
bad := decodeVerifyResult(t, out.Bytes())
|
|
||||||
require.Equal(t, "failed", bad.Status)
|
|
||||||
require.Equal(t, 1, bad.Mismatched)
|
|
||||||
require.Equal(t, healthy.BlobCount-1, bad.Verified,
|
|
||||||
"the wrong-sized blob must not be counted as verified")
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestShallowVerifyDetectsMissingDatabase backs up a real snapshot,
|
|
||||||
// confirms shallow verify passes, then deletes the snapshot's encrypted
|
|
||||||
// database. Shallow verify must fail: a snapshot without its database is
|
|
||||||
// not restorable, even when every blob is present.
|
|
||||||
func TestShallowVerifyDetectsMissingDatabase(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
fs := afero.NewOsFs()
|
|
||||||
tempDir := t.TempDir()
|
|
||||||
|
|
||||||
dataDir := filepath.Join(tempDir, "source")
|
|
||||||
storeDir := filepath.Join(tempDir, "remote")
|
|
||||||
dbPath := filepath.Join(tempDir, "index.sqlite")
|
|
||||||
|
|
||||||
chunkSize := int64(32 * 1024)
|
|
||||||
maxBlobSize := int64(128 * 1024)
|
|
||||||
|
|
||||||
require.NoError(t, fs.MkdirAll(dataDir, 0o755))
|
|
||||||
require.NoError(t, afero.WriteFile(fs,
|
|
||||||
filepath.Join(dataDir, "data.bin"),
|
|
||||||
bytesPattern("shallow-db-", int(maxBlobSize*2)), 0o644))
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
cfg, storer, snapshotID := runFileStorageBackup(
|
|
||||||
ctx, t, fs, dataDir, storeDir, dbPath, chunkSize, maxBlobSize)
|
|
||||||
|
|
||||||
newVerifier := func() *vaultik.Vaultik {
|
|
||||||
v := &vaultik.Vaultik{
|
|
||||||
Config: cfg,
|
|
||||||
Storage: storer,
|
|
||||||
Fs: fs,
|
|
||||||
Stdout: io.Discard,
|
|
||||||
Stderr: io.Discard,
|
|
||||||
UI: ui.NewWithColor(io.Discard, false),
|
|
||||||
}
|
|
||||||
v.SetContext(ctx)
|
|
||||||
|
|
||||||
return v
|
|
||||||
}
|
|
||||||
|
|
||||||
require.NoError(t,
|
|
||||||
newVerifier().VerifySnapshotWithOptions(
|
|
||||||
snapshotID, &vaultik.VerifyOptions{}),
|
|
||||||
"shallow verify should pass on a healthy snapshot")
|
|
||||||
|
|
||||||
// The database lives under the hashed remote key, not the human ID.
|
|
||||||
dbObject := filepath.Join(storeDir, "metadata",
|
|
||||||
snapshot.RemoteSnapshotKey(snapshotID), "db.zst.age")
|
|
||||||
require.NoError(t, os.Remove(dbObject))
|
|
||||||
|
|
||||||
require.Error(t,
|
|
||||||
newVerifier().VerifySnapshotWithOptions(
|
|
||||||
snapshotID, &vaultik.VerifyOptions{}),
|
|
||||||
"shallow verify must fail when db.zst.age is absent")
|
|
||||||
}
|
|
||||||
|
|
||||||
// growOneBlob appends bytes to the first blob file found under blobsDir,
|
|
||||||
// changing its on-disk size so it no longer matches the manifest.
|
|
||||||
func growOneBlob(t *testing.T, fs afero.Fs, blobsDir string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var blobPath string
|
|
||||||
|
|
||||||
err := afero.Walk(fs, blobsDir,
|
|
||||||
func(path string, info os.FileInfo, err error) error {
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
if blobPath == "" && !info.IsDir() {
|
|
||||||
blobPath = path
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.NotEmpty(t, blobPath, "expected at least one blob on disk")
|
|
||||||
|
|
||||||
data, err := afero.ReadFile(fs, blobPath)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
data = append(data, []byte("extra")...)
|
|
||||||
require.NoError(t, afero.WriteFile(fs, blobPath, data, 0o644))
|
|
||||||
}
|
|
||||||
|
|
||||||
func decodeVerifyResult(t *testing.T, b []byte) vaultik.VerifyResult {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var result vaultik.VerifyResult
|
|
||||||
|
|
||||||
require.NoError(t, json.Unmarshal(b, &result))
|
|
||||||
|
|
||||||
return result
|
|
||||||
}
|
|
||||||
+57
-112
@@ -20,6 +20,7 @@ import (
|
|||||||
var (
|
var (
|
||||||
errSnapshotNotInConfig = errors.New("snapshot not found in config")
|
errSnapshotNotInConfig = errors.New("snapshot not found in config")
|
||||||
errNoSnapshotsInConfig = errors.New("no snapshots configured")
|
errNoSnapshotsInConfig = errors.New("no snapshots configured")
|
||||||
|
errBlobsMissing = errors.New("blobs are missing")
|
||||||
errSnapshotVerifyFailed = errors.New("verification failed")
|
errSnapshotVerifyFailed = errors.New("verification failed")
|
||||||
errRemoveAllNeedsForce = errors.New("--all requires --force")
|
errRemoveAllNeedsForce = errors.New("--all requires --force")
|
||||||
errInvalidTableName = errors.New("invalid table name")
|
errInvalidTableName = errors.New("invalid table name")
|
||||||
@@ -55,8 +56,8 @@ func (v *Vaultik) CreateSnapshot(opts *SnapshotCreateOptions) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Clean up incomplete snapshots FIRST, before any scanning.
|
// Clean up incomplete snapshots FIRST, before any scanning
|
||||||
// This is critical for data safety; PruneDatabase below does it.
|
// This is critical for data safety - see CleanupIncompleteSnapshots for details
|
||||||
hostname := v.Config.Hostname
|
hostname := v.Config.Hostname
|
||||||
if hostname == "" {
|
if hostname == "" {
|
||||||
hostname, _ = os.Hostname()
|
hostname, _ = os.Hostname()
|
||||||
@@ -669,24 +670,11 @@ func (v *Vaultik) VerifySnapshotWithOptions(
|
|||||||
|
|
||||||
v.printVerifyHeader(snapshotID, opts)
|
v.printVerifyHeader(snapshotID, opts)
|
||||||
|
|
||||||
// Resolve the identifier to the snapshot's remote key. A human ID is
|
// Resolve the identifier to the snapshot's remote key and download the
|
||||||
// hashed; a remote key (or its abbreviation, as printed for a
|
// manifest. A human ID is hashed; a remote key (or its abbreviation,
|
||||||
// remote-only snapshot) is used as-is, so a host with no local index
|
// as printed for a remote-only snapshot) is used as-is, so a host with
|
||||||
// can verify a snapshot it can only see on the store. The key is kept
|
// no local index can verify a snapshot it can only see on the store.
|
||||||
// so we can also check for the snapshot's encrypted database below.
|
manifest, err := v.resolveAndDownloadManifest(snapshotID)
|
||||||
remoteKey, err := v.resolveSnapshotRemoteKey(snapshotID)
|
|
||||||
if err != nil {
|
|
||||||
if opts.JSON {
|
|
||||||
result.Status = verifyStatusFailed
|
|
||||||
result.ErrorMessage = fmt.Sprintf("resolving snapshot identifier: %v", err)
|
|
||||||
|
|
||||||
return v.outputVerifyJSON(result)
|
|
||||||
}
|
|
||||||
|
|
||||||
return fmt.Errorf("resolving snapshot identifier: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
manifest, err := v.downloadManifestByKey(remoteKey)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if opts.JSON {
|
if opts.JSON {
|
||||||
result.Status = verifyStatusFailed
|
result.Status = verifyStatusFailed
|
||||||
@@ -716,24 +704,14 @@ func (v *Vaultik) VerifySnapshotWithOptions(
|
|||||||
|
|
||||||
v.printlnStdout()
|
v.printlnStdout()
|
||||||
|
|
||||||
// Check each blob is present with the size the manifest records.
|
// Check each blob exists
|
||||||
v.stdoutf("Checking blob presence and sizes...\n")
|
v.stdoutf("Checking blob existence...\n")
|
||||||
}
|
}
|
||||||
|
|
||||||
// A snapshot is only restorable if its encrypted database is present
|
result.Verified, result.Missing, result.MissingSize =
|
||||||
// alongside the blobs. Shallow verify checks that the object exists; it
|
v.verifyManifestBlobsExist(manifest, opts)
|
||||||
// does not decrypt it (that is deep verify's job).
|
|
||||||
dbPath := fmt.Sprintf("metadata/%s/db.zst.age", remoteKey)
|
|
||||||
|
|
||||||
_, dbErr := v.Storage.Stat(v.ctx, dbPath)
|
return v.formatVerifyResult(result, manifest, opts)
|
||||||
if dbErr != nil {
|
|
||||||
result.DatabaseMissing = true
|
|
||||||
}
|
|
||||||
|
|
||||||
result.Verified, result.Missing, result.Mismatched, result.MissingSize =
|
|
||||||
v.verifyManifestBlobs(manifest, opts)
|
|
||||||
|
|
||||||
return v.formatVerifyResult(result, opts)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// printVerifyHeader prints the snapshot ID and parsed timestamp for
|
// printVerifyHeader prints the snapshot ID and parsed timestamp for
|
||||||
@@ -758,17 +736,14 @@ func (v *Vaultik) printVerifyHeader(snapshotID string, opts *VerifyOptions) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// verifyManifestBlobs checks that each blob in the manifest is present in
|
// verifyManifestBlobsExist checks that each blob in the manifest exists
|
||||||
// storage with the size the manifest records, returning the counts of
|
// in storage, returning the verified count, missing count, and total
|
||||||
// blobs that were present with the right size, absent, and present but the
|
// missing bytes.
|
||||||
// wrong size, plus the total bytes of the absent blobs. It does not read
|
func (v *Vaultik) verifyManifestBlobsExist(
|
||||||
// blob contents; deep verification (RunDeepVerify) does that. The size
|
|
||||||
// comparison matches the deep path (see verifyBlobExistenceFromDB).
|
|
||||||
func (v *Vaultik) verifyManifestBlobs(
|
|
||||||
manifest *snapshot.Manifest, opts *VerifyOptions,
|
manifest *snapshot.Manifest, opts *VerifyOptions,
|
||||||
) (int, int, int, int64) {
|
) (int, int, int64) {
|
||||||
var (
|
var (
|
||||||
verified, missing, mismatched int
|
verified, missing int
|
||||||
missingSize int64
|
missingSize int64
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -776,9 +751,10 @@ func (v *Vaultik) verifyManifestBlobs(
|
|||||||
blobPath := fmt.Sprintf("blobs/%s/%s/%s",
|
blobPath := fmt.Sprintf("blobs/%s/%s/%s",
|
||||||
blob.Hash[:2], blob.Hash[2:4], blob.Hash)
|
blob.Hash[:2], blob.Hash[2:4], blob.Hash)
|
||||||
|
|
||||||
stat, err := v.Storage.Stat(v.ctx, blobPath)
|
// Shallow: check existence only (deep verification is handled
|
||||||
switch {
|
// by RunDeepVerify).
|
||||||
case err != nil:
|
_, err := v.Storage.Stat(v.ctx, blobPath)
|
||||||
|
if err != nil {
|
||||||
if !opts.JSON {
|
if !opts.JSON {
|
||||||
v.stdoutf(" Missing: %s (%s)\n",
|
v.stdoutf(" Missing: %s (%s)\n",
|
||||||
blob.Hash, ubytes(blob.CompressedSize))
|
blob.Hash, ubytes(blob.CompressedSize))
|
||||||
@@ -786,32 +762,23 @@ func (v *Vaultik) verifyManifestBlobs(
|
|||||||
|
|
||||||
missing++
|
missing++
|
||||||
missingSize += blob.CompressedSize
|
missingSize += blob.CompressedSize
|
||||||
case stat.Size != blob.CompressedSize:
|
} else {
|
||||||
if !opts.JSON {
|
|
||||||
v.stdoutf(" Wrong size: %s (store has %s, manifest lists %s)\n",
|
|
||||||
blob.Hash, ubytes(stat.Size), ubytes(blob.CompressedSize))
|
|
||||||
}
|
|
||||||
|
|
||||||
mismatched++
|
|
||||||
default:
|
|
||||||
verified++
|
verified++
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return verified, missing, mismatched, missingSize
|
return verified, missing, missingSize
|
||||||
}
|
}
|
||||||
|
|
||||||
// formatVerifyResult outputs the final verification results as JSON or
|
// formatVerifyResult outputs the final verification results as JSON or
|
||||||
// human-readable text.
|
// human-readable text.
|
||||||
func (v *Vaultik) formatVerifyResult(
|
func (v *Vaultik) formatVerifyResult(
|
||||||
result *VerifyResult, opts *VerifyOptions,
|
result *VerifyResult, manifest *snapshot.Manifest, opts *VerifyOptions,
|
||||||
) error {
|
) error {
|
||||||
failure := shallowVerifyFailure(result)
|
|
||||||
|
|
||||||
if opts.JSON {
|
if opts.JSON {
|
||||||
if failure != "" {
|
if result.Missing > 0 {
|
||||||
result.Status = verifyStatusFailed
|
result.Status = verifyStatusFailed
|
||||||
result.ErrorMessage = failure
|
result.ErrorMessage = fmt.Sprintf("%d blobs are missing", result.Missing)
|
||||||
} else {
|
} else {
|
||||||
result.Status = "ok"
|
result.Status = "ok"
|
||||||
}
|
}
|
||||||
@@ -820,57 +787,29 @@ func (v *Vaultik) formatVerifyResult(
|
|||||||
}
|
}
|
||||||
|
|
||||||
v.stdoutf("\nVerification complete:\n")
|
v.stdoutf("\nVerification complete:\n")
|
||||||
v.stdoutf(" Present with listed size: %d blobs\n", result.Verified)
|
v.stdoutf(" Verified: %d blobs (%s)\n", result.Verified,
|
||||||
|
ubytes(manifest.TotalCompressedSize-result.MissingSize))
|
||||||
|
|
||||||
if result.Missing > 0 {
|
if result.Missing > 0 {
|
||||||
v.stdoutf(" Missing: %d blobs (%s)\n",
|
v.stdoutf(" Missing: %d blobs (%s)\n",
|
||||||
result.Missing, ubytes(result.MissingSize))
|
result.Missing, ubytes(result.MissingSize))
|
||||||
}
|
} else {
|
||||||
|
v.stdoutf(" Missing: 0 blobs\n")
|
||||||
if result.Mismatched > 0 {
|
|
||||||
v.stdoutf(" Wrong size: %d blobs\n", result.Mismatched)
|
|
||||||
}
|
|
||||||
|
|
||||||
if result.DatabaseMissing {
|
|
||||||
v.stdoutf(" Encrypted database: missing\n")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
v.stdoutf(" Status: ")
|
v.stdoutf(" Status: ")
|
||||||
|
|
||||||
if failure != "" {
|
if result.Missing > 0 {
|
||||||
v.stdoutf("FAILED - %s\n", failure)
|
v.stdoutf("FAILED - %d blobs are missing\n", result.Missing)
|
||||||
|
|
||||||
return fmt.Errorf("%w: %s", errSnapshotVerifyFailed, failure)
|
return fmt.Errorf("%d %w", result.Missing, errBlobsMissing)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Report only what was actually checked: presence and size, not contents.
|
v.stdoutf("OK - All blobs verified\n")
|
||||||
v.stdoutf("OK - all %d blobs listed in the manifest are present with the "+
|
|
||||||
"listed size; contents not checked (use --deep)\n", result.Verified)
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// shallowVerifyFailure returns a human-readable description of everything
|
|
||||||
// that failed shallow verification, or the empty string if it passed.
|
|
||||||
func shallowVerifyFailure(result *VerifyResult) string {
|
|
||||||
var parts []string
|
|
||||||
|
|
||||||
if result.Missing > 0 {
|
|
||||||
parts = append(parts, fmt.Sprintf("%d blobs are missing", result.Missing))
|
|
||||||
}
|
|
||||||
|
|
||||||
if result.Mismatched > 0 {
|
|
||||||
parts = append(parts,
|
|
||||||
fmt.Sprintf("%d blobs have the wrong size", result.Mismatched))
|
|
||||||
}
|
|
||||||
|
|
||||||
if result.DatabaseMissing {
|
|
||||||
parts = append(parts, "the encrypted database is missing")
|
|
||||||
}
|
|
||||||
|
|
||||||
return strings.Join(parts, "; ")
|
|
||||||
}
|
|
||||||
|
|
||||||
// outputVerifyJSON outputs the verification result as JSON
|
// outputVerifyJSON outputs the verification result as JSON
|
||||||
func (v *Vaultik) outputVerifyJSON(result *VerifyResult) error {
|
func (v *Vaultik) outputVerifyJSON(result *VerifyResult) error {
|
||||||
encoder := json.NewEncoder(v.Stdout)
|
encoder := json.NewEncoder(v.Stdout)
|
||||||
@@ -996,23 +935,29 @@ func (v *Vaultik) downloadManifestByKey(remoteKey string) (*snapshot.Manifest, e
|
|||||||
func (v *Vaultik) syncWithRemote() error {
|
func (v *Vaultik) syncWithRemote() error {
|
||||||
log.Info("Syncing with remote snapshots")
|
log.Info("Syncing with remote snapshots")
|
||||||
|
|
||||||
// Remote metadata lives under metadata/<remote-key>/, where the
|
// Get all remote snapshot IDs
|
||||||
// directory name is snapshot.RemoteSnapshotKey(id), not the human
|
remoteSnapshots := make(map[string]bool)
|
||||||
// snapshot ID. Compare each local row's hashed key against that set
|
objectCh := v.Storage.ListStream(v.ctx, "metadata/")
|
||||||
// so a row still backed by remote metadata is kept. Comparing human
|
|
||||||
// IDs against the hashed directory names matches nothing and deletes
|
for object := range objectCh {
|
||||||
// every local snapshot record (issue #160).
|
if object.Err != nil {
|
||||||
remoteKeys, err := v.listAllRemoteSnapshotKeys()
|
return fmt.Errorf("listing remote snapshots: %w", object.Err)
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("listing remote snapshots: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
remoteKeySet := make(map[string]bool, len(remoteKeys))
|
// Extract snapshot ID from paths like metadata/hostname-20240115-143052Z/
|
||||||
for _, k := range remoteKeys {
|
parts := strings.Split(object.Key, "/")
|
||||||
remoteKeySet[k] = true
|
if len(parts) >= minSnapshotIDParts &&
|
||||||
|
parts[0] == metadataDirName && parts[1] != "" {
|
||||||
|
// Skip macOS resource fork files (._*) and other hidden files
|
||||||
|
if strings.HasPrefix(parts[1], ".") {
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Debug("Found remote snapshots", "count", len(remoteKeySet))
|
remoteSnapshots[parts[1]] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debug("Found remote snapshots", "count", len(remoteSnapshots))
|
||||||
|
|
||||||
// Get all local snapshots (use a high limit to get all)
|
// Get all local snapshots (use a high limit to get all)
|
||||||
localSnapshots, err := v.Repositories.Snapshots.ListRecent(v.ctx, listRecentLimit)
|
localSnapshots, err := v.Repositories.Snapshots.ListRecent(v.ctx, listRecentLimit)
|
||||||
@@ -1020,12 +965,12 @@ func (v *Vaultik) syncWithRemote() error {
|
|||||||
return fmt.Errorf("listing local snapshots: %w", err)
|
return fmt.Errorf("listing local snapshots: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Remove local snapshots whose metadata is absent from the remote.
|
// Remove local snapshots that don't exist remotely
|
||||||
removedCount := 0
|
removedCount := 0
|
||||||
|
|
||||||
for _, snap := range localSnapshots {
|
for _, snap := range localSnapshots {
|
||||||
snapshotIDStr := snap.ID.String()
|
snapshotIDStr := snap.ID.String()
|
||||||
if !remoteKeySet[snapshot.RemoteSnapshotKey(snapshotIDStr)] {
|
if !remoteSnapshots[snapshotIDStr] {
|
||||||
log.Info("Removing local snapshot not found in remote",
|
log.Info("Removing local snapshot not found in remote",
|
||||||
"snapshot_id", snap.ID)
|
"snapshot_id", snap.ID)
|
||||||
|
|
||||||
|
|||||||
@@ -69,6 +69,20 @@ func (v *Vaultik) resolveSnapshotRemoteKey(identifier string) (string, error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// resolveAndDownloadManifest resolves a snapshot identifier to its remote
|
||||||
|
// key (see resolveSnapshotRemoteKey) and downloads that snapshot's
|
||||||
|
// manifest.
|
||||||
|
func (v *Vaultik) resolveAndDownloadManifest(
|
||||||
|
identifier string,
|
||||||
|
) (*snapshot.Manifest, error) {
|
||||||
|
remoteKey, err := v.resolveSnapshotRemoteKey(identifier)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return v.downloadManifestByKey(remoteKey)
|
||||||
|
}
|
||||||
|
|
||||||
// isRemoteKeyOrPrefix reports whether s is a full remote key or the
|
// isRemoteKeyOrPrefix reports whether s is a full remote key or the
|
||||||
// leading part of one: 1 to 64 lowercase hex characters. A human snapshot
|
// leading part of one: 1 to 64 lowercase hex characters. A human snapshot
|
||||||
// ID is never all hex, so this shape test is enough to tell the two apart.
|
// ID is never all hex, so this shape test is enough to tell the two apart.
|
||||||
|
|||||||
+28
-34
@@ -5,6 +5,7 @@ package vaultik
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
@@ -12,6 +13,7 @@ import (
|
|||||||
"github.com/spf13/afero"
|
"github.com/spf13/afero"
|
||||||
"go.uber.org/fx"
|
"go.uber.org/fx"
|
||||||
"sneak.berlin/go/vaultik/internal/config"
|
"sneak.berlin/go/vaultik/internal/config"
|
||||||
|
"sneak.berlin/go/vaultik/internal/crypto"
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
"sneak.berlin/go/vaultik/internal/database"
|
||||||
"sneak.berlin/go/vaultik/internal/globals"
|
"sneak.berlin/go/vaultik/internal/globals"
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
"sneak.berlin/go/vaultik/internal/snapshot"
|
||||||
@@ -19,6 +21,12 @@ import (
|
|||||||
"sneak.berlin/go/vaultik/internal/ui"
|
"sneak.berlin/go/vaultik/internal/ui"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Sentinel errors for misconfigured encryption settings.
|
||||||
|
var (
|
||||||
|
errNoAgeRecipients = errors.New("no age recipients configured")
|
||||||
|
errNoAgeSecretKey = errors.New("no age secret key configured")
|
||||||
|
)
|
||||||
|
|
||||||
// Vaultik contains all dependencies needed for vaultik operations
|
// Vaultik contains all dependencies needed for vaultik operations
|
||||||
type Vaultik struct {
|
type Vaultik struct {
|
||||||
Globals *globals.Globals
|
Globals *globals.Globals
|
||||||
@@ -128,45 +136,31 @@ func (v *Vaultik) Cancel() {
|
|||||||
v.cancel()
|
v.cancel()
|
||||||
}
|
}
|
||||||
|
|
||||||
// StartOperation runs fn in its own goroutine and returns a stop
|
|
||||||
// function. fn is the command being run (a restore, verify, prune, and
|
|
||||||
// so on); it observes cancellation through the Vaultik context and
|
|
||||||
// removes its decrypted scratch files (the blob cache and the temporary
|
|
||||||
// snapshot database) from the temp directory as it unwinds.
|
|
||||||
//
|
|
||||||
// Calling stop cancels the Vaultik context and then blocks until fn has
|
|
||||||
// returned — so that unwinding, and the cleanup it does, completes
|
|
||||||
// before the caller proceeds — or until the passed context is done,
|
|
||||||
// whichever comes first. It reports whether fn returned before that
|
|
||||||
// deadline. A signal-driven shutdown must call stop before the process
|
|
||||||
// exits; otherwise the process can exit mid-operation and leave
|
|
||||||
// decrypted data behind.
|
|
||||||
func (v *Vaultik) StartOperation(fn func()) func(context.Context) bool {
|
|
||||||
done := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
defer close(done)
|
|
||||||
|
|
||||||
fn()
|
|
||||||
}()
|
|
||||||
|
|
||||||
return func(ctx context.Context) bool {
|
|
||||||
v.Cancel()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
return true
|
|
||||||
case <-ctx.Done():
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// CanDecrypt returns true if this Vaultik instance has decryption capabilities
|
// CanDecrypt returns true if this Vaultik instance has decryption capabilities
|
||||||
func (v *Vaultik) CanDecrypt() bool {
|
func (v *Vaultik) CanDecrypt() bool {
|
||||||
return v.Config.AgeSecretKey != ""
|
return v.Config.AgeSecretKey != ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetEncryptor creates a new Encryptor instance based on the configured age recipients
|
||||||
|
// Returns an error if no recipients are configured
|
||||||
|
func (v *Vaultik) GetEncryptor() (*crypto.Encryptor, error) {
|
||||||
|
if len(v.Config.AgeRecipients) == 0 {
|
||||||
|
return nil, errNoAgeRecipients
|
||||||
|
}
|
||||||
|
|
||||||
|
return crypto.NewEncryptor(v.Config.AgeRecipients)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetDecryptor creates a new Decryptor instance based on the configured age secret key
|
||||||
|
// Returns an error if no secret key is configured
|
||||||
|
func (v *Vaultik) GetDecryptor() (*crypto.Decryptor, error) {
|
||||||
|
if v.Config.AgeSecretKey == "" {
|
||||||
|
return nil, errNoAgeSecretKey
|
||||||
|
}
|
||||||
|
|
||||||
|
return crypto.NewDecryptor(v.Config.AgeSecretKey)
|
||||||
|
}
|
||||||
|
|
||||||
// GetFilesystem returns the filesystem instance used by Vaultik
|
// GetFilesystem returns the filesystem instance used by Vaultik
|
||||||
//
|
//
|
||||||
//nolint:ireturn // afero.Fs is the filesystem abstraction by design
|
//nolint:ireturn // afero.Fs is the filesystem abstraction by design
|
||||||
|
|||||||
+102
-143
@@ -6,15 +6,15 @@ import (
|
|||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"hash"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"filippo.io/age"
|
"github.com/klauspost/compress/zstd"
|
||||||
|
|
||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
// Blank import registers the pure-Go sqlite driver for database/sql.
|
||||||
"sneak.berlin/go/vaultik/internal/database"
|
_ "modernc.org/sqlite"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
"sneak.berlin/go/vaultik/internal/snapshot"
|
||||||
)
|
)
|
||||||
@@ -29,8 +29,6 @@ var (
|
|||||||
errTrailingBlobData = errors.New(
|
errTrailingBlobData = errors.New(
|
||||||
"blob has unexpected trailing bytes not covered by chunk list")
|
"blob has unexpected trailing bytes not covered by chunk list")
|
||||||
errManifestExtraBlob = errors.New("manifest contains blob not in database")
|
errManifestExtraBlob = errors.New("manifest contains blob not in database")
|
||||||
errManifestMissingBlob = errors.New(
|
|
||||||
"manifest omits blob present in database")
|
|
||||||
errBlobSizeMismatch = errors.New("blob size mismatch")
|
errBlobSizeMismatch = errors.New("blob size mismatch")
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -55,11 +53,6 @@ type VerifyResult struct {
|
|||||||
Verified int `json:"verified"`
|
Verified int `json:"verified"`
|
||||||
Missing int `json:"missing"`
|
Missing int `json:"missing"`
|
||||||
MissingSize int64 `json:"missing_size,omitempty"`
|
MissingSize int64 `json:"missing_size,omitempty"`
|
||||||
Mismatched int `json:"mismatched,omitempty"`
|
|
||||||
// DatabaseMissing is set by shallow verify when the snapshot's
|
|
||||||
// encrypted database (metadata/<key>/db.zst.age) is absent, which
|
|
||||||
// makes the snapshot unrestorable regardless of the blobs.
|
|
||||||
DatabaseMissing bool `json:"database_missing,omitempty"`
|
|
||||||
ErrorMessage string `json:"error,omitempty"`
|
ErrorMessage string `json:"error,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -93,21 +86,13 @@ func (v *Vaultik) RunDeepVerify(snapshotID string, opts *VerifyOptions) error {
|
|||||||
errSecretKeyRequired.Error(), errSecretKeyRequired)
|
errSecretKeyRequired.Error(), errSecretKeyRequired)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse the age secret key once, the same way restore does, and reuse
|
|
||||||
// the identities for the database and every blob.
|
|
||||||
identities, err := v.restoreIdentities()
|
|
||||||
if err != nil {
|
|
||||||
return v.deepVerifyFailure(result, opts, err.Error(), err)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Info("Starting snapshot verification", "snapshot_id", snapshotID, "mode", "deep")
|
log.Info("Starting snapshot verification", "snapshot_id", snapshotID, "mode", "deep")
|
||||||
|
|
||||||
if !opts.JSON {
|
if !opts.JSON {
|
||||||
v.stdoutf("Deep verification of snapshot: %s\n\n", snapshotID)
|
v.stdoutf("Deep verification of snapshot: %s\n\n", snapshotID)
|
||||||
}
|
}
|
||||||
|
|
||||||
manifest, tempDB, dbBlobs, err := v.loadVerificationData(
|
manifest, tempDB, dbBlobs, err := v.loadVerificationData(snapshotID, opts, result)
|
||||||
snapshotID, opts, result, identities)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -127,8 +112,7 @@ func (v *Vaultik) RunDeepVerify(snapshotID string, opts *VerifyOptions) error {
|
|||||||
|
|
||||||
result.TotalSize = totalSize
|
result.TotalSize = totalSize
|
||||||
|
|
||||||
err = v.runVerificationSteps(
|
err = v.runVerificationSteps(manifest, dbBlobs, tempDB, opts, result, totalSize)
|
||||||
manifest, dbBlobs, tempDB, opts, result, totalSize, identities)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -153,7 +137,6 @@ func (v *Vaultik) RunDeepVerify(snapshotID string, opts *VerifyOptions) error {
|
|||||||
// loadVerificationData downloads manifest, database, and blob list for verification
|
// loadVerificationData downloads manifest, database, and blob list for verification
|
||||||
func (v *Vaultik) loadVerificationData(
|
func (v *Vaultik) loadVerificationData(
|
||||||
snapshotID string, opts *VerifyOptions, result *VerifyResult,
|
snapshotID string, opts *VerifyOptions, result *VerifyResult,
|
||||||
identities []age.Identity,
|
|
||||||
) (*snapshot.Manifest, *tempDB, []snapshot.BlobInfo, error) {
|
) (*snapshot.Manifest, *tempDB, []snapshot.BlobInfo, error) {
|
||||||
// Resolve the identifier to the snapshot's remote key. A human ID is
|
// Resolve the identifier to the snapshot's remote key. A human ID is
|
||||||
// hashed; a remote key (or its abbreviation, as printed for a
|
// hashed; a remote key (or its abbreviation, as printed for a
|
||||||
@@ -190,13 +173,27 @@ func (v *Vaultik) loadVerificationData(
|
|||||||
v.stdoutf("Downloading and decrypting database...\n")
|
v.stdoutf("Downloading and decrypting database...\n")
|
||||||
}
|
}
|
||||||
|
|
||||||
tdb, err := v.downloadVerifiedSnapshotDB(
|
// Download and decrypt database
|
||||||
snapshotID, remoteKey, opts, result, identities)
|
dbPath := fmt.Sprintf("metadata/%s/db.zst.age", remoteKey)
|
||||||
|
log.Info("Downloading encrypted database", "path", dbPath)
|
||||||
|
|
||||||
|
dbReader, err := v.Storage.Get(v.ctx, dbPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, nil, nil, err
|
return nil, nil, nil, v.deepVerifyFailure(result, opts,
|
||||||
|
fmt.Sprintf("failed to download database: %v", err),
|
||||||
|
fmt.Errorf("failed to download database: %w", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
dbBlobs, err := v.getBlobsFromDatabase(tdb.db.Conn())
|
defer func() { _ = dbReader.Close() }()
|
||||||
|
|
||||||
|
tdb, err := v.decryptAndLoadDatabase(dbReader)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, nil, v.deepVerifyFailure(result, opts,
|
||||||
|
fmt.Sprintf("failed to decrypt database: %v", err),
|
||||||
|
fmt.Errorf("failed to decrypt database: %w", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
dbBlobs, err := v.getBlobsFromDatabase(tdb.DB)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = tdb.Close()
|
_ = tdb.Close()
|
||||||
|
|
||||||
@@ -222,45 +219,6 @@ func (v *Vaultik) loadVerificationData(
|
|||||||
return manifest, tdb, dbBlobs, nil
|
return manifest, tdb, dbBlobs, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// downloadVerifiedSnapshotDB downloads and decrypts the snapshot metadata
|
|
||||||
// database and confirms it really is the snapshot named by remoteKey
|
|
||||||
// before any of its rows are trusted (see verifySnapshotDBIdentity). On
|
|
||||||
// any failure it records the failure in result and returns the error the
|
|
||||||
// caller should propagate; the temp database is closed on a rejected
|
|
||||||
// identity so nothing is left on disk.
|
|
||||||
func (v *Vaultik) downloadVerifiedSnapshotDB(
|
|
||||||
snapshotID, remoteKey string, opts *VerifyOptions, result *VerifyResult,
|
|
||||||
identities []age.Identity,
|
|
||||||
) (*tempDB, error) {
|
|
||||||
dbPath := fmt.Sprintf("metadata/%s/db.zst.age", remoteKey)
|
|
||||||
log.Info("Downloading encrypted database", "path", dbPath)
|
|
||||||
|
|
||||||
dbReader, err := v.Storage.Get(v.ctx, dbPath)
|
|
||||||
if err != nil {
|
|
||||||
return nil, v.deepVerifyFailure(result, opts,
|
|
||||||
fmt.Sprintf("failed to download database: %v", err),
|
|
||||||
fmt.Errorf("failed to download database: %w", err))
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() { _ = dbReader.Close() }()
|
|
||||||
|
|
||||||
tdb, err := v.decryptAndLoadDatabase(dbReader, identities)
|
|
||||||
if err != nil {
|
|
||||||
return nil, v.deepVerifyFailure(result, opts,
|
|
||||||
fmt.Sprintf("failed to decrypt database: %v", err),
|
|
||||||
fmt.Errorf("failed to decrypt database: %w", err))
|
|
||||||
}
|
|
||||||
|
|
||||||
err = v.verifySnapshotDBIdentity(tdb.db, snapshotID, remoteKey)
|
|
||||||
if err != nil {
|
|
||||||
_ = tdb.Close()
|
|
||||||
|
|
||||||
return nil, v.deepVerifyFailure(result, opts, err.Error(), err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return tdb, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// runVerificationSteps executes manifest verification, blob existence
|
// runVerificationSteps executes manifest verification, blob existence
|
||||||
// check, and deep content verification.
|
// check, and deep content verification.
|
||||||
func (v *Vaultik) runVerificationSteps(
|
func (v *Vaultik) runVerificationSteps(
|
||||||
@@ -270,7 +228,6 @@ func (v *Vaultik) runVerificationSteps(
|
|||||||
opts *VerifyOptions,
|
opts *VerifyOptions,
|
||||||
result *VerifyResult,
|
result *VerifyResult,
|
||||||
totalSize int64,
|
totalSize int64,
|
||||||
identities []age.Identity,
|
|
||||||
) error {
|
) error {
|
||||||
if !opts.JSON {
|
if !opts.JSON {
|
||||||
v.stdoutf("Verifying manifest against database...\n")
|
v.stdoutf("Verifying manifest against database...\n")
|
||||||
@@ -297,7 +254,7 @@ func (v *Vaultik) runVerificationSteps(
|
|||||||
len(dbBlobs), ubytes(totalSize))
|
len(dbBlobs), ubytes(totalSize))
|
||||||
}
|
}
|
||||||
|
|
||||||
err = v.performDeepVerificationFromDB(dbBlobs, tdb.db.Conn(), opts, identities)
|
err = v.performDeepVerificationFromDB(dbBlobs, tdb.DB, opts)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return v.deepVerifyFailure(result, opts, err.Error(), err)
|
return v.deepVerifyFailure(result, opts, err.Error(), err)
|
||||||
}
|
}
|
||||||
@@ -305,92 +262,81 @@ func (v *Vaultik) runVerificationSteps(
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// tempDB is the downloaded snapshot database opened read-only for deep
|
// tempDB wraps sql.DB with cleanup
|
||||||
// verify, held in a private temp directory removed in full on Close.
|
|
||||||
type tempDB struct {
|
type tempDB struct {
|
||||||
db *database.DB
|
*sql.DB
|
||||||
tempDir string
|
|
||||||
|
tempPath string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tempDB) Close() error {
|
func (t *tempDB) Close() error {
|
||||||
err := t.db.Close()
|
err := t.DB.Close()
|
||||||
// Remove the whole private directory so the decrypted database and
|
_ = os.Remove(t.tempPath)
|
||||||
// any SQLite side files are gone on every path.
|
|
||||||
_ = os.RemoveAll(t.tempDir)
|
|
||||||
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// decryptAndLoadDatabase decrypts and loads the binary SQLite database
|
// decryptAndLoadDatabase decrypts and loads the binary SQLite database
|
||||||
// from the encrypted stream. It reads through the same blobgen reader restore
|
// from the encrypted stream.
|
||||||
// uses, streaming the decrypted, decompressed database to a temp file.
|
func (v *Vaultik) decryptAndLoadDatabase(reader io.ReadCloser) (*tempDB, error) {
|
||||||
func (v *Vaultik) decryptAndLoadDatabase(
|
// Get decryptor
|
||||||
reader io.ReadCloser, identities []age.Identity,
|
decryptor, err := v.GetDecryptor()
|
||||||
) (*tempDB, error) {
|
|
||||||
// Decrypt and decompress through the shared blobgen reader.
|
|
||||||
blobReader, err := blobgen.NewReader(reader, identities...)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create decryption reader: %w", err)
|
return nil, fmt.Errorf("failed to get decryptor: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
defer func() { _ = blobReader.Close() }()
|
// Decrypt the stream
|
||||||
|
decryptedReader, err := decryptor.DecryptStream(reader)
|
||||||
// Materialize the decrypted database inside a private (0700) temp
|
|
||||||
// directory so it is never world-readable, and remove the whole
|
|
||||||
// directory on any failure below.
|
|
||||||
tempDir, err := os.MkdirTemp("", "vaultik-verify-")
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create temp directory: %w", err)
|
return nil, fmt.Errorf("failed to decrypt database: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
success := false
|
// Decompress the binary database
|
||||||
|
decompressor, err := zstd.NewReader(decryptedReader)
|
||||||
defer func() {
|
if err != nil {
|
||||||
if !success {
|
return nil, fmt.Errorf("failed to create decompressor: %w", err)
|
||||||
_ = os.RemoveAll(tempDir)
|
|
||||||
}
|
}
|
||||||
}()
|
defer decompressor.Close()
|
||||||
|
|
||||||
dbPath := filepath.Join(tempDir, snapshotDBFilename)
|
// Create temporary file for the database
|
||||||
|
tempFile, err := os.CreateTemp("", "vaultik-verify-*.db")
|
||||||
//nolint:gosec // G304: dbPath is our MkdirTemp dir plus a constant filename
|
|
||||||
tempFile, err := os.OpenFile(
|
|
||||||
dbPath, os.O_CREATE|os.O_EXCL|os.O_WRONLY, restoreFileMode)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create temp file: %w", err)
|
return nil, fmt.Errorf("failed to create temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
tempPath := tempFile.Name()
|
||||||
|
|
||||||
// Stream decompress directly to file
|
// Stream decompress directly to file
|
||||||
log.Info("Decompressing database...")
|
log.Info("Decompressing database...")
|
||||||
|
|
||||||
written, err := io.Copy(tempFile, blobReader)
|
written, err := io.Copy(tempFile, decompressor)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = tempFile.Close()
|
_ = tempFile.Close()
|
||||||
|
_ = os.Remove(tempPath)
|
||||||
|
|
||||||
return nil, fmt.Errorf("failed to decompress database: %w", err)
|
return nil, fmt.Errorf("failed to decompress database: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = tempFile.Close()
|
_ = tempFile.Close()
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to close temp database file: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Info("Database decompressed", "size", ubytes(written))
|
log.Info("Database decompressed", "size", ubytes(written))
|
||||||
|
|
||||||
db, err := database.OpenReadOnly(v.ctx, dbPath)
|
// Open the database
|
||||||
|
db, err := sql.Open("sqlite", tempPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
_ = os.Remove(tempPath)
|
||||||
|
|
||||||
return nil, fmt.Errorf("failed to open database: %w", err)
|
return nil, fmt.Errorf("failed to open database: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
success = true
|
return &tempDB{
|
||||||
|
DB: db,
|
||||||
return &tempDB{db: db, tempDir: tempDir}, nil
|
tempPath: tempPath,
|
||||||
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// verifyBlob downloads and verifies a single blob
|
// verifyBlob downloads and verifies a single blob
|
||||||
func (v *Vaultik) verifyBlob(
|
func (v *Vaultik) verifyBlob(blobInfo snapshot.BlobInfo, db *sql.DB) error {
|
||||||
blobInfo snapshot.BlobInfo, db *sql.DB, identities []age.Identity,
|
|
||||||
) error {
|
|
||||||
// Download blob using shared fetch method
|
// Download blob using shared fetch method
|
||||||
reader, _, err := v.FetchBlob(v.ctx, blobInfo.Hash, blobInfo.CompressedSize)
|
reader, _, err := v.FetchBlob(v.ctx, blobInfo.Hash, blobInfo.CompressedSize)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -399,23 +345,38 @@ func (v *Vaultik) verifyBlob(
|
|||||||
|
|
||||||
defer func() { _ = reader.Close() }()
|
defer func() { _ = reader.Close() }()
|
||||||
|
|
||||||
// Decrypt and decompress through the shared blobgen reader, which hashes
|
// Get decryptor
|
||||||
// the plaintext as it is read. A blob's hash — its remote name — is the
|
decryptor, err := v.GetDecryptor()
|
||||||
// double SHA-256 of that plaintext (see blobgen.DoubleSHA256), not of the
|
|
||||||
// encrypted bytes.
|
|
||||||
blobReader, err := blobgen.NewReader(reader, identities...)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to create blob reader: %w", err)
|
return fmt.Errorf("failed to get decryptor: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
defer func() { _ = blobReader.Close() }()
|
// Decrypt blob
|
||||||
|
decryptedReader, err := decryptor.DecryptStream(reader)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to decrypt: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
chunkCount, err := v.verifyBlobChunks(db, blobInfo.Hash, blobReader)
|
// Decompress blob
|
||||||
|
decompressor, err := zstd.NewReader(decryptedReader)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to decompress: %w", err)
|
||||||
|
}
|
||||||
|
defer decompressor.Close()
|
||||||
|
|
||||||
|
// A blob's hash — its remote name — is the double SHA256 of its
|
||||||
|
// decompressed plaintext (see blobgen.Writer.Sum256), not of the
|
||||||
|
// encrypted bytes. Hash the plaintext as chunk verification streams
|
||||||
|
// it, then compare on completion.
|
||||||
|
plaintextHasher := sha256.New()
|
||||||
|
hashedStream := io.TeeReader(decompressor, plaintextHasher)
|
||||||
|
|
||||||
|
chunkCount, err := v.verifyBlobChunks(db, blobInfo.Hash, hashedStream)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
err = v.verifyBlobFinalIntegrity(blobReader, blobInfo.Hash)
|
err = v.verifyBlobFinalIntegrity(hashedStream, plaintextHasher, blobInfo.Hash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -521,12 +482,11 @@ func (v *Vaultik) verifyBlobChunks(
|
|||||||
// verifyBlobFinalIntegrity checks that no trailing data exists in the
|
// verifyBlobFinalIntegrity checks that no trailing data exists in the
|
||||||
// decompressed stream and that the blob hash matches the expected value.
|
// decompressed stream and that the blob hash matches the expected value.
|
||||||
func (v *Vaultik) verifyBlobFinalIntegrity(
|
func (v *Vaultik) verifyBlobFinalIntegrity(
|
||||||
blobReader *blobgen.Reader, expectedHash string,
|
plaintext io.Reader, plaintextHasher hash.Hash, expectedHash string,
|
||||||
) error {
|
) error {
|
||||||
// Verify no remaining data in blob - if the chunk list is accurate,
|
// Verify no remaining data in blob - if the chunk list is accurate,
|
||||||
// the blob should be fully consumed. Draining to EOF also completes the
|
// the blob should be fully consumed.
|
||||||
// reader's plaintext hash.
|
remaining, err := io.Copy(io.Discard, plaintext)
|
||||||
remaining, err := io.Copy(io.Discard, blobReader)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to check for remaining blob data: %w", err)
|
return fmt.Errorf("failed to check for remaining blob data: %w", err)
|
||||||
}
|
}
|
||||||
@@ -535,9 +495,10 @@ func (v *Vaultik) verifyBlobFinalIntegrity(
|
|||||||
return fmt.Errorf("%w: %d bytes", errTrailingBlobData, remaining)
|
return fmt.Errorf("%w: %d bytes", errTrailingBlobData, remaining)
|
||||||
}
|
}
|
||||||
|
|
||||||
// The blob hash is the double SHA-256 of its plaintext content.
|
// The blob hash is the double SHA256 of its plaintext content.
|
||||||
calculatedBlobHash := hex.EncodeToString(
|
firstHash := plaintextHasher.Sum(nil)
|
||||||
blobgen.DoubleSHA256(blobReader.Sum256()))
|
secondHash := sha256.Sum256(firstHash)
|
||||||
|
calculatedBlobHash := hex.EncodeToString(secondHash[:])
|
||||||
|
|
||||||
if calculatedBlobHash != expectedHash {
|
if calculatedBlobHash != expectedHash {
|
||||||
return fmt.Errorf("%w: calculated %s, expected %s",
|
return fmt.Errorf("%w: calculated %s, expected %s",
|
||||||
@@ -614,11 +575,16 @@ func (v *Vaultik) verifyManifestAgainstDatabase(
|
|||||||
manifestBlobMap[blob.Hash] = blob.CompressedSize
|
manifestBlobMap[blob.Hash] = blob.CompressedSize
|
||||||
}
|
}
|
||||||
|
|
||||||
// The manifest is the only blob list prune consults, so it must match
|
// Check counts match
|
||||||
// the database exactly. A blob in the manifest but not the database
|
if len(dbBlobMap) != len(manifestBlobMap) {
|
||||||
// points at a corrupt manifest; a blob in the database but omitted
|
log.Warn("Manifest blob count mismatch",
|
||||||
// from the manifest would be pruned away while this snapshot still
|
"database_blobs", len(dbBlobMap),
|
||||||
// needs it. Either divergence fails verification.
|
"manifest_blobs", len(manifestBlobMap),
|
||||||
|
)
|
||||||
|
// This is a warning, not an error - database is authoritative
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check each manifest blob exists in database with correct size
|
||||||
for hash, manifestSize := range manifestBlobMap {
|
for hash, manifestSize := range manifestBlobMap {
|
||||||
dbSize, exists := dbBlobMap[hash]
|
dbSize, exists := dbBlobMap[hash]
|
||||||
if !exists {
|
if !exists {
|
||||||
@@ -632,12 +598,6 @@ func (v *Vaultik) verifyManifestAgainstDatabase(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
for hash := range dbBlobMap {
|
|
||||||
if _, exists := manifestBlobMap[hash]; !exists {
|
|
||||||
return fmt.Errorf("%w: %s", errManifestMissingBlob, hash)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Info("✓ Manifest verified against database",
|
log.Info("✓ Manifest verified against database",
|
||||||
"manifest_blobs", len(manifestBlobMap),
|
"manifest_blobs", len(manifestBlobMap),
|
||||||
"database_blobs", len(dbBlobMap),
|
"database_blobs", len(dbBlobMap),
|
||||||
@@ -687,7 +647,6 @@ func (v *Vaultik) verifyBlobExistenceFromDB(blobs []snapshot.BlobInfo) error {
|
|||||||
// each blob using the database as source.
|
// each blob using the database as source.
|
||||||
func (v *Vaultik) performDeepVerificationFromDB(
|
func (v *Vaultik) performDeepVerificationFromDB(
|
||||||
blobs []snapshot.BlobInfo, db *sql.DB, opts *VerifyOptions,
|
blobs []snapshot.BlobInfo, db *sql.DB, opts *VerifyOptions,
|
||||||
identities []age.Identity,
|
|
||||||
) error {
|
) error {
|
||||||
// Calculate total bytes for ETA
|
// Calculate total bytes for ETA
|
||||||
var totalBytesExpected int64
|
var totalBytesExpected int64
|
||||||
@@ -705,7 +664,7 @@ func (v *Vaultik) performDeepVerificationFromDB(
|
|||||||
|
|
||||||
for i, blobInfo := range blobs {
|
for i, blobInfo := range blobs {
|
||||||
// Verify individual blob
|
// Verify individual blob
|
||||||
err := v.verifyBlob(blobInfo, db, identities)
|
err := v.verifyBlob(blobInfo, db)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("blob %s verification failed: %w", blobInfo.Hash, err)
|
return fmt.Errorf("blob %s verification failed: %w", blobInfo.Hash, err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,61 +0,0 @@
|
|||||||
package vaultik //nolint:testpackage // calls unexported verifyManifestAgainstDatabase
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
"sneak.berlin/go/vaultik/internal/snapshot"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Blob hashes shared by the manifest-verification tests below.
|
|
||||||
const (
|
|
||||||
manifestTestBlobA = "blob-a"
|
|
||||||
manifestTestBlobB = "blob-b"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestVerifyManifestAgainstDatabase_MissingBlobFails is the regression
|
|
||||||
// guard for issue #157: deep verify must fail when the manifest omits a
|
|
||||||
// blob the database records. The divergence used to be logged as a
|
|
||||||
// warning while verification still returned ok, so an incomplete
|
|
||||||
// manifest — the exact defect that lets prune later delete a needed blob
|
|
||||||
// — passed unnoticed.
|
|
||||||
func TestVerifyManifestAgainstDatabase_MissingBlobFails(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
v := &Vaultik{}
|
|
||||||
|
|
||||||
dbBlobs := []snapshot.BlobInfo{
|
|
||||||
{Hash: manifestTestBlobA, CompressedSize: 10},
|
|
||||||
{Hash: manifestTestBlobB, CompressedSize: 20},
|
|
||||||
}
|
|
||||||
manifest := &snapshot.Manifest{
|
|
||||||
Blobs: []snapshot.BlobInfo{
|
|
||||||
{Hash: manifestTestBlobA, CompressedSize: 10},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
err := v.verifyManifestAgainstDatabase(manifest, dbBlobs)
|
|
||||||
require.Error(t, err, "verify must fail when the manifest omits a database blob")
|
|
||||||
assert.Contains(t, err.Error(), manifestTestBlobB)
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestVerifyManifestAgainstDatabase_MatchingSetsPass keeps the other half
|
|
||||||
// honest: identical blob sets still verify, so the check above cannot be
|
|
||||||
// satisfied by failing everything.
|
|
||||||
func TestVerifyManifestAgainstDatabase_MatchingSetsPass(t *testing.T) {
|
|
||||||
log.Initialize(log.Config{})
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
v := &Vaultik{}
|
|
||||||
|
|
||||||
blobs := []snapshot.BlobInfo{
|
|
||||||
{Hash: manifestTestBlobA, CompressedSize: 10},
|
|
||||||
{Hash: manifestTestBlobB, CompressedSize: 20},
|
|
||||||
}
|
|
||||||
manifest := &snapshot.Manifest{Blobs: blobs}
|
|
||||||
|
|
||||||
require.NoError(t, v.verifyManifestAgainstDatabase(manifest, blobs))
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,100 @@
|
|||||||
|
package vaultik_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"io"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/klauspost/compress/zstd"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"sneak.berlin/go/vaultik/internal/crypto"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestTeeReaderWithDecryption tests that TeeReader correctly hashes all encrypted
|
||||||
|
// bytes when streaming through age decryption and zstd decompression.
|
||||||
|
// This validates the verification path: hash encrypted blob -> decrypt -> decompress.
|
||||||
|
func TestTeeReaderWithDecryption(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Test data - use random data that doesn't compress well (5MB)
|
||||||
|
testData := make([]byte, 5*1024*1024)
|
||||||
|
_, err := rand.Read(testData)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Compress the data
|
||||||
|
var compressedBuf bytes.Buffer
|
||||||
|
|
||||||
|
compressor, err := zstd.NewWriter(&compressedBuf,
|
||||||
|
zstd.WithEncoderLevel(zstd.SpeedDefault))
|
||||||
|
require.NoError(t, err)
|
||||||
|
_, err = compressor.Write(testData)
|
||||||
|
require.NoError(t, err)
|
||||||
|
err = compressor.Close()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Encrypt the compressed data
|
||||||
|
testRecipient := "age1cplgrwj77ta54dnmydvvmzn64ltk83ankxl5sww04mrt" +
|
||||||
|
"mu62kv3s89gmvv"
|
||||||
|
testSecretKey := "AGE-SECRET-KEY-1C77PYNTHXSHNNC6EYR2W52UWYXACXA5J" +
|
||||||
|
"T00J9CCW9986M3XY87PSGP89AQ"
|
||||||
|
|
||||||
|
encryptor, err := crypto.NewEncryptor([]string{testRecipient})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var encryptedBuf bytes.Buffer
|
||||||
|
|
||||||
|
err = encryptor.EncryptStream(&encryptedBuf, bytes.NewReader(compressedBuf.Bytes()))
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
encryptedData := encryptedBuf.Bytes()
|
||||||
|
|
||||||
|
// Calculate the expected hash of the encrypted data directly
|
||||||
|
expectedHash := sha256.Sum256(encryptedData)
|
||||||
|
expectedHashStr := hex.EncodeToString(expectedHash[:])
|
||||||
|
|
||||||
|
t.Logf("Encrypted data size: %d bytes", len(encryptedData))
|
||||||
|
t.Logf("Expected hash: %s", expectedHashStr)
|
||||||
|
|
||||||
|
// Now simulate what verifyBlob does: use TeeReader to hash while decrypting
|
||||||
|
decryptor, err := crypto.NewDecryptor(testSecretKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Create hasher and tee reader
|
||||||
|
hasher := sha256.New()
|
||||||
|
reader := bytes.NewReader(encryptedData)
|
||||||
|
teeReader := io.TeeReader(reader, hasher)
|
||||||
|
|
||||||
|
// Decrypt through the tee reader
|
||||||
|
decryptedReader, err := decryptor.DecryptStream(teeReader)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Decompress
|
||||||
|
decompressor, err := zstd.NewReader(decryptedReader)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
defer decompressor.Close()
|
||||||
|
|
||||||
|
// Read all decompressed data (simulating chunk verification)
|
||||||
|
decompressedData, err := io.ReadAll(decompressor)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Verify we got the original data back
|
||||||
|
assert.Equal(t, testData, decompressedData, "Decompressed data should match original")
|
||||||
|
|
||||||
|
// Drain remaining decompressed data (should be 0)
|
||||||
|
remaining, err := io.Copy(io.Discard, decompressor)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, int64(0), remaining, "No remaining decompressed data")
|
||||||
|
|
||||||
|
// Calculate hash from tee reader
|
||||||
|
calculatedHashStr := hex.EncodeToString(hasher.Sum(nil))
|
||||||
|
t.Logf("Calculated hash (before drain): %s", calculatedHashStr)
|
||||||
|
|
||||||
|
// Verify the hash matches the direct hash of encrypted data
|
||||||
|
assert.Equal(t, expectedHashStr, calculatedHashStr,
|
||||||
|
"Hash calculated via TeeReader should match direct hash of encrypted data")
|
||||||
|
}
|
||||||
+1
-16
@@ -56,23 +56,8 @@ main() {
|
|||||||
docker build --output=type=cacheonly \
|
docker build --output=type=cacheonly \
|
||||||
--build-arg CHECK_EPOCH="$epoch" -f Dockerfile.lint .
|
--build-arg CHECK_EPOCH="$epoch" -f Dockerfile.lint .
|
||||||
|
|
||||||
# Version, commit and build date are computed here on the host, the
|
|
||||||
# same way script/docker does, and passed into the product build so
|
|
||||||
# the CI-built image reports its real source. The build context
|
|
||||||
# excludes .git (see .dockerignore), so the build cannot derive them
|
|
||||||
# itself; without these it would stamp the Dockerfile's dev/unknown
|
|
||||||
# fallbacks. VERSION comes from script/version, the source of truth
|
|
||||||
# shared with the Makefile.
|
|
||||||
version="$("$ROOT/script/version")"
|
|
||||||
commit="$(git rev-parse HEAD 2>/dev/null || echo unknown)"
|
|
||||||
commit_date="$(git show -s --format=%cs HEAD 2>/dev/null || echo unknown)"
|
|
||||||
|
|
||||||
epoch="$(date +%s%N)$$"
|
epoch="$(date +%s%N)$$"
|
||||||
docker build --build-arg CHECK_EPOCH="$epoch" \
|
docker build --build-arg CHECK_EPOCH="$epoch" .
|
||||||
--build-arg VERSION="$version" \
|
|
||||||
--build-arg COMMIT="$commit" \
|
|
||||||
--build-arg COMMIT_DATE="$commit_date" \
|
|
||||||
.
|
|
||||||
}
|
}
|
||||||
|
|
||||||
main "$@"
|
main "$@"
|
||||||
|
|||||||
@@ -24,24 +24,7 @@ main() {
|
|||||||
# whether the tree is clean. The Dockerfile now refuses to build
|
# whether the tree is clean. The Dockerfile now refuses to build
|
||||||
# without a non-empty value, so this is required, not optional.
|
# without a non-empty value, so this is required, not optional.
|
||||||
epoch="$(date +%s%N)$$"
|
epoch="$(date +%s%N)$$"
|
||||||
|
|
||||||
# Version, commit and build date are computed here on the host,
|
|
||||||
# where .git exists, and passed into the build. The build context
|
|
||||||
# excludes .git (see .dockerignore), so the container cannot derive
|
|
||||||
# them itself -- it used to try and always got "unknown", giving
|
|
||||||
# every image a "commit: unknown" it could not be traced from.
|
|
||||||
# VERSION comes from script/version, the source of truth shared with
|
|
||||||
# the Makefile, so a Docker build reports the same string (tag,
|
|
||||||
# dev-<sha>, or a -dirty variant) that a local build of the same
|
|
||||||
# tree would.
|
|
||||||
version="$("$SCRIPT_DIR/version")"
|
|
||||||
commit="$(git rev-parse HEAD 2>/dev/null || echo unknown)"
|
|
||||||
commit_date="$(git show -s --format=%cs HEAD 2>/dev/null || echo unknown)"
|
|
||||||
|
|
||||||
docker build --build-arg CHECK_EPOCH="$epoch" \
|
docker build --build-arg CHECK_EPOCH="$epoch" \
|
||||||
--build-arg VERSION="$version" \
|
|
||||||
--build-arg COMMIT="$commit" \
|
|
||||||
--build-arg COMMIT_DATE="$commit_date" \
|
|
||||||
-t "$("$SCRIPT_DIR/projectname")" .
|
-t "$("$SCRIPT_DIR/projectname")" .
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -1,6 +1,6 @@
|
|||||||
age_recipients:
|
age_recipients:
|
||||||
- age1278m9q7dp3chsh2dcy82qk27v047zywyvtxwnj4cvt0z65jw6a7q5dqhfj # sneak's long term age key
|
- age1278m9q7dp3chsh2dcy82qk27v047zywyvtxwnj4cvt0z65jw6a7q5dqhfj # sneak's long term age key
|
||||||
- age1ezrjmfpwsc95svdg0y54mums3zevgzu0x0ecq2f7tp8a05gl0sjq9q9wjg # add additional recipients as needed
|
- age1otherpubkey... # add additional recipients as needed
|
||||||
snapshots:
|
snapshots:
|
||||||
test:
|
test:
|
||||||
paths:
|
paths:
|
||||||
|
|||||||
Reference in New Issue
Block a user