Compare commits
30
Commits
b76437aa7b
..
next
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c3bec7d3aa | ||
|
|
1244c9e48d | ||
|
|
82c51a5337 | ||
|
|
d88ed64489 | ||
|
|
bd9656dbd4 | ||
|
|
7e611b95db | ||
|
|
4f27608560 | ||
|
|
ae6aaaa388 | ||
|
|
f788668287 | ||
|
|
238ce3985f | ||
|
|
548a7ae156 | ||
|
|
3a58377127 | ||
|
|
a6434de57f | ||
|
|
b4654f8e52 | ||
|
|
39aef1c47c | ||
|
|
96ebcd40d7 | ||
|
|
d9f0220f94 | ||
|
|
4c83e82543 | ||
|
|
86361c8b50 | ||
|
|
d77663d039 | ||
|
|
3abe9cbd9e | ||
|
|
76a6917a35 | ||
|
|
38ebfd843a | ||
|
|
6b7517a4dc | ||
|
|
994e5de613 | ||
|
|
42f4e648d7 | ||
|
|
343129f891 | ||
|
|
a50e3fa038 | ||
|
|
6fcd8e1668 | ||
|
|
aab6a87f8c |
@@ -104,7 +104,12 @@ Version: 2025-06-08
|
|||||||
|
|
||||||
13. Pre-1.0: NEVER write database migrations. There are no live databases
|
13. Pre-1.0: NEVER write database migrations. There are no live databases
|
||||||
anywhere — every user's local index can be rebuilt from a fresh full
|
anywhere — every user's local index can be rebuilt from a fresh full
|
||||||
backup. When the schema changes, just change `schema.sql` (and any code
|
backup. To change the schema, edit `internal/database/schema/001.sql`
|
||||||
that touches the affected tables). The local index is disposable until
|
(and any code that touches the affected tables) directly; do not add new
|
||||||
1.0 ships and is tagged.
|
numbered schema files. Those numbered files and the `schema_migrations`
|
||||||
|
table they populate only bootstrap a fresh database — they are not an
|
||||||
|
upgrade path. The local index is disposable until 1.0 ships and is
|
||||||
|
tagged; once 1.0 is tagged that clause expires and the question of
|
||||||
|
upgrading existing indexes returns. See [`docs/DATAMODEL.md`](docs/DATAMODEL.md)
|
||||||
|
for the full explanation.
|
||||||
|
|
||||||
|
|||||||
+4
-5
@@ -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` (typically 16KB-256KB for 64KB average).
|
Chunk sizes vary between `avgChunkSize/4` and `avgChunkSize*4` (2.5MB-40MB for the 10MB default 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.ConfigPath(opts.ConfigPath)), // 1. Config path
|
fx.Supply(config.Path(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 (typically 64KB)
|
- `avgChunkSize`: From config (default 10MB)
|
||||||
- `minChunkSize`: avgChunkSize / 4
|
- `minChunkSize`: avgChunkSize / 4
|
||||||
- `maxChunkSize`: avgChunkSize * 4
|
- `maxChunkSize`: avgChunkSize * 4
|
||||||
|
|
||||||
@@ -286,7 +286,6 @@ 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.
|
||||||
@@ -307,7 +306,7 @@ Repository interfaces:
|
|||||||
```
|
```
|
||||||
CreateSnapshot(opts)
|
CreateSnapshot(opts)
|
||||||
│
|
│
|
||||||
├─► CleanupIncompleteSnapshots() // Critical: avoid dedup errors
|
├─► PruneDatabase() // Critical: avoid dedup errors
|
||||||
│
|
│
|
||||||
├─► SnapshotManager.CreateSnapshot() // Create DB record
|
├─► SnapshotManager.CreateSnapshot() // Create DB record
|
||||||
│
|
│
|
||||||
|
|||||||
+24
-3
@@ -20,8 +20,6 @@
|
|||||||
# 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.
|
||||||
@@ -66,8 +64,31 @@ 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=$(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
|
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
|
||||||
|
|
||||||
# Runtime stage
|
# Runtime stage
|
||||||
# alpine:3.21, 2026-02-25
|
# alpine:3.21, 2026-02-25
|
||||||
|
|||||||
@@ -71,14 +71,19 @@ Requirements that no existing tool meets:
|
|||||||
## daily use
|
## daily use
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
# verify a snapshot (shallow: checks all blobs exist)
|
# verify a snapshot (shallow: checks all blobs are present with the listed size)
|
||||||
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_AGE_SECRET_KEY='AGE-SECRET-KEY-...' vaultik snapshot verify --deep <snapshot-id>
|
vaultik snapshot verify --deep <snapshot-id>
|
||||||
|
|
||||||
# restore (requires the private key)
|
# restore (requires the private key)
|
||||||
VAULTIK_AGE_SECRET_KEY='AGE-SECRET-KEY-...' vaultik snapshot restore <snapshot-id> /tmp/restored
|
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
|
||||||
@@ -119,15 +124,17 @@ 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_AGE_SECRET_KEY='AGE-SECRET-KEY-...' \
|
vaultik snapshot restore --verify <remote-key> /tmp/restored
|
||||||
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_AGE_SECRET_KEY='AGE-SECRET-KEY-...' \
|
vaultik snapshot verify --deep <remote-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
|
||||||
@@ -147,10 +154,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]
|
vaultik [--config <path>] snapshot list [--json] # alias: ls
|
||||||
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]
|
vaultik [--config <path>] snapshot remove <snapshot-id> [--dry-run] [--force] [--local-only] [--json] # alias: rm
|
||||||
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
|
||||||
@@ -167,7 +174,24 @@ 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`: Continue past per-file errors instead of aborting (applies to `snapshot create` and `restore`)
|
* `--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.
|
||||||
|
|
||||||
|
### 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
|
||||||
|
|
||||||
@@ -200,9 +224,11 @@ 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`)
|
* `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_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
|
||||||
|
|
||||||
@@ -293,7 +319,9 @@ 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 all blobs referenced in the manifest exist in storage
|
* Default (shallow): checks that every blob the manifest lists is present in
|
||||||
|
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
|
||||||
@@ -395,6 +423,10 @@ 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
|
||||||
|
|
||||||
```
|
```
|
||||||
@@ -472,25 +504,30 @@ derivation.
|
|||||||
|
|
||||||
### compression
|
### compression
|
||||||
|
|
||||||
* zstd compression at configurable level (1-19, default 3)
|
* zstd compression at configurable level (1-19, default 3). The level is
|
||||||
|
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.
|
Run `vaultik config init` to generate a fully commented config file; a
|
||||||
Key fields:
|
complete annotated example also lives in
|
||||||
|
[`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 |
|
| `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 |
|
||||||
| `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 |
|
||||||
@@ -514,9 +551,13 @@ Key fields:
|
|||||||
sequentially. Restore speed is bound by single-stream throughput.
|
sequentially. Restore speed is bound by single-stream throughput.
|
||||||
* **Device nodes, named pipes, and sockets are silently skipped.** Only
|
* **Device nodes, named pipes, and sockets are silently skipped.** Only
|
||||||
regular files, directories, and symlinks are backed up.
|
regular files, directories, and symlinks are backed up.
|
||||||
* **No database migrations.** If the local SQLite schema changes between
|
* **No upgrade path between versions.** There is no supported way to carry
|
||||||
versions, delete the local database (`vaultik database delete`) and run
|
an existing local index across a schema change; if the local SQLite
|
||||||
a full backup. Remote storage is unaffected.
|
schema changes between versions, delete the local database (`vaultik
|
||||||
|
database delete`) and run a full backup. Remote storage is unaffected.
|
||||||
|
(The binary does embed numbered schema files and a `schema_migrations`
|
||||||
|
table to bootstrap a fresh database — see [`docs/DATAMODEL.md`](docs/DATAMODEL.md)
|
||||||
|
— but that is not an upgrade path.)
|
||||||
* **Files that change during backup may be inconsistent.** There is no
|
* **Files that change during backup may be inconsistent.** There is no
|
||||||
filesystem snapshot or freeze. If a file is modified between the scan
|
filesystem snapshot or freeze. If a file is modified between the scan
|
||||||
and chunk phases, the backed-up copy may reflect a partial write.
|
and chunk phases, the backed-up copy may reflect a partial write.
|
||||||
@@ -582,10 +623,12 @@ priority.
|
|||||||
|
|
||||||
### infrastructure
|
### infrastructure
|
||||||
|
|
||||||
* **Schema migrations.** Currently nonexistent — pre-1.0 schema
|
* **Cross-version schema upgrades.** There is no upgrade path between
|
||||||
changes are handled by `vaultik database delete` plus a full
|
released versions — pre-1.0 schema changes are handled by `vaultik
|
||||||
re-scan. Post-1.0 we'll need a migration story to keep existing
|
database delete` plus a full re-scan (see
|
||||||
index databases usable across upgrades.
|
[`docs/DATAMODEL.md`](docs/DATAMODEL.md)). Post-1.0 we'll need a
|
||||||
|
migration story to keep existing index databases usable across
|
||||||
|
upgrades.
|
||||||
* **Storage backend coverage tests.** S3, file://, and rclone://
|
* **Storage backend coverage tests.** S3, file://, and rclone://
|
||||||
all share the Storer interface but the rclone path is the least
|
all share the Storer interface but the rclone path is the least
|
||||||
exercised in CI.
|
exercised in CI.
|
||||||
@@ -594,9 +637,17 @@ priority.
|
|||||||
|
|
||||||
## output style
|
## output style
|
||||||
|
|
||||||
All user-facing output goes through helpers in `internal/ui` and conforms
|
The operational narration of the long-running commands — the Begin,
|
||||||
to a uniform style. Color is enabled when stdout is a TTY and the
|
Complete, Progress, and status lines of `snapshot create`, `prune`,
|
||||||
`NO_COLOR` environment variable is unset (https://no-color.org/).
|
`snapshot restore`, and the like — goes through helpers in `internal/ui`
|
||||||
|
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,6 +25,66 @@ release" is exactly the contradiction
|
|||||||
|
|
||||||
# Completed Steps
|
# Completed Steps
|
||||||
|
|
||||||
|
- 2026-09-22: Validated blob hashes, offsets and lengths read back from
|
||||||
|
the destination before using them
|
||||||
|
([issue #155](https://git.eeqj.de/sneak/vaultik/issues/155)). A blob
|
||||||
|
hash taken from the downloaded database or the store listing was
|
||||||
|
trusted unchecked, so a hostile remote could drive a decrypted blob to
|
||||||
|
be written outside the cache directory (a hash like `aa/../../etc`) or
|
||||||
|
crash a command with a short or negative value. `blobDiskCache.path`
|
||||||
|
now refuses any key containing a path separator, and `ReadAt` rejects a
|
||||||
|
negative offset or length and bounds with `length > size-offset` so a
|
||||||
|
sum cannot overflow past the check. A new `isBlobHash` helper (a plain
|
||||||
|
function, not a method — the packer stores `temp-placeholder-{uuid}` as
|
||||||
|
a hash) gates `FetchBlob`, shallow and deep verify; the `blobs/` and
|
||||||
|
`metadata/` listings skip a non-conforming name with a warning; and
|
||||||
|
short-hash prefixes in log and error text go through a `shortHash`
|
||||||
|
helper that cannot panic. `verify`'s chunk reader also rejects a
|
||||||
|
negative `blob_chunks` length and streams the chunk instead of
|
||||||
|
allocating a database-supplied size. `restore.go` and
|
||||||
|
`internal/database` were left untouched to avoid colliding with the
|
||||||
|
in-flight [issue #156](https://git.eeqj.de/sneak/vaultik/issues/156)
|
||||||
|
work; the cache-path and `FetchBlob` guards already stop the unsafe
|
||||||
|
write and fetch, so restore's own `buildBlobIndexes` early check is
|
||||||
|
deferred as fail-fast defense in depth.
|
||||||
|
|
||||||
|
- 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
|
||||||
|
|||||||
@@ -0,0 +1,102 @@
|
|||||||
|
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,10 +304,14 @@ 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.
|
// beginning with, want; -1 if there is none. An `ARG NAME=default`
|
||||||
|
// 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 || strings.HasPrefix(instruction, want+" ") {
|
if instruction == want ||
|
||||||
|
strings.HasPrefix(instruction, want+" ") ||
|
||||||
|
strings.HasPrefix(instruction, want+"=") {
|
||||||
return i
|
return i
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+11
-1
@@ -10,6 +10,16 @@ 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
|
||||||
@@ -46,5 +56,5 @@ func main() {
|
|||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
cli.Entry()
|
return cli.Entry()
|
||||||
}
|
}
|
||||||
|
|||||||
+7
-5
@@ -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://las1stor1//srv/pool.2024.04/backups/heraklion"
|
storage_url: "rclone://myremote/path/to/backups"
|
||||||
|
|
||||||
# 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: http://10.100.205.122:8333
|
# endpoint: https://s3.example.com
|
||||||
#
|
#
|
||||||
# # Bucket name where backups will be stored
|
# # Bucket name where backups will be stored
|
||||||
# bucket: testbucket
|
# bucket: mybucket
|
||||||
#
|
#
|
||||||
# # 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://las1stor1//srv/pool.2024.04/backups/heraklion"
|
|||||||
# #prefix: "hosts/myserver/"
|
# #prefix: "hosts/myserver/"
|
||||||
#
|
#
|
||||||
# # S3 access credentials
|
# # S3 access credentials
|
||||||
# access_key_id: Z9GT22M9YFU08WRMC5D4
|
# access_key_id: YOUR_ACCESS_KEY
|
||||||
# secret_access_key: Pi0tPKjFbN4rZlRhcA4zBtEkib04yy2WcIzI+AXk
|
# secret_access_key: YOUR_SECRET_KEY
|
||||||
#
|
#
|
||||||
# # S3 region
|
# # S3 region
|
||||||
# # Default: us-east-1
|
# # Default: us-east-1
|
||||||
@@ -304,6 +304,8 @@ storage_url: "rclone://las1stor1//srv/pool.2024.04/backups/heraklion"
|
|||||||
|
|
||||||
# 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
|
||||||
|
|||||||
+24
-5
@@ -5,11 +5,30 @@
|
|||||||
Vaultik uses a local SQLite database to track file metadata, chunk mappings, and blob associations during the backup process. This database serves as an index for incremental backups and enables efficient deduplication.
|
Vaultik uses a local SQLite database to track file metadata, chunk mappings, and blob associations during the backup process. This database serves as an index for incremental backups and enables efficient deduplication.
|
||||||
|
|
||||||
**Important Notes:**
|
**Important Notes:**
|
||||||
- **No Migration Support (pre-1.0)**: Vaultik does not support database schema
|
|
||||||
migrations. The local index is treated as disposable — if the schema changes,
|
This section is the authoritative explanation of the schema/migration story;
|
||||||
delete the local SQLite database (`vaultik database delete`) and run a full
|
other documents (the README and `AGENTS.md`) link here.
|
||||||
backup. The remote storage is unaffected; the new index will re-deduplicate
|
|
||||||
against existing remote blobs.
|
- **No upgrade path between versions (pre-1.0)**: Vaultik has no supported way to
|
||||||
|
carry an existing local index across a schema change. The index is disposable
|
||||||
|
— if the on-disk schema changes between versions, delete the local SQLite
|
||||||
|
database (`vaultik database delete`) and run a full backup. Remote storage is
|
||||||
|
unaffected; the new index re-deduplicates against existing remote blobs. This
|
||||||
|
is the standing project policy, and it is separate from the schema bootstrap
|
||||||
|
described next.
|
||||||
|
- **Schema bootstrap**: a fresh database is populated from numbered SQL files
|
||||||
|
embedded in the binary under `internal/database/schema/`. `000.sql` creates the
|
||||||
|
`schema_migrations` table; `001.sql` creates the application tables. On opening
|
||||||
|
a database the code applies each numbered file that has not yet run and records
|
||||||
|
its version in `schema_migrations`. This bootstraps a new database; it does not
|
||||||
|
upgrade an existing one between released versions.
|
||||||
|
- **Changing the schema (pre-1.0)**: edit `internal/database/schema/001.sql` (and
|
||||||
|
the code that touches the affected tables) directly. Do not add new numbered
|
||||||
|
files — there is no installed base to migrate.
|
||||||
|
- **Disposability expires at 1.0**: the index is treated as disposable only until
|
||||||
|
1.0 ships and is tagged. Once 1.0 is tagged that clause expires and the
|
||||||
|
question of upgrading existing indexes returns. It is deliberately left open
|
||||||
|
here.
|
||||||
- **Version Compatibility**: In rare cases, you may need to use the same version
|
- **Version Compatibility**: In rare cases, you may need to use the same version
|
||||||
of Vaultik to restore a backup as was used to create it. This ensures
|
of Vaultik to restore a backup as was used to create it. This ensures
|
||||||
compatibility with the metadata format stored in S3.
|
compatibility with the metadata format stored in S3.
|
||||||
|
|||||||
@@ -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 only deletes blobs not referenced in any manifest
|
1. It keeps every blob listed in any snapshot's manifest and deletes only blobs that no manifest references
|
||||||
2. Manifests are unencrypted and can be read without keys
|
2. Manifests are unencrypted and can be read without keys
|
||||||
3. The operation compares the latest local DB snapshot with the latest S3 snapshot to ensure consistency
|
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
|
||||||
4. Pruning will fail if these don't match, preventing accidental deletion of needed blobs
|
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
|
||||||
|
|
||||||
## 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.Sum256()
|
finalHash := p.currentBlob.writer.ContentID()
|
||||||
|
|
||||||
return hex.EncodeToString(finalHash), finalSize, nil
|
return hex.EncodeToString(finalHash), finalSize, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,89 +0,0 @@
|
|||||||
// 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
|
|
||||||
}
|
|
||||||
@@ -1,80 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,119 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -0,0 +1,49 @@
|
|||||||
|
package blobgen
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ErrOutputTooLarge is returned by a reader from LimitReader once it has
|
||||||
|
// been asked for more than its limit. It bounds how far an untrusted
|
||||||
|
// compressed stream may expand, so a small, highly compressible object
|
||||||
|
// from the store cannot decompress without limit.
|
||||||
|
var ErrOutputTooLarge = errors.New("output exceeds size limit")
|
||||||
|
|
||||||
|
// LimitReader returns a reader that yields at most limit bytes from r and
|
||||||
|
// then fails with ErrOutputTooLarge. Unlike io.LimitReader, which reports
|
||||||
|
// a silent io.EOF at the limit (indistinguishable from a stream that
|
||||||
|
// simply ended), this fails, so a caller decoding or copying the stream
|
||||||
|
// sees an error rather than a truncated value. A stream of exactly limit
|
||||||
|
// bytes reads back cleanly to EOF; the first byte beyond it is the error.
|
||||||
|
func LimitReader(r io.Reader, limit int64) io.Reader {
|
||||||
|
// remaining counts down from limit+1: the extra byte is the one that,
|
||||||
|
// if it ever arrives, proves the stream is longer than the limit.
|
||||||
|
return &limitReader{r: r, remaining: limit + 1}
|
||||||
|
}
|
||||||
|
|
||||||
|
type limitReader struct {
|
||||||
|
r io.Reader
|
||||||
|
remaining int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *limitReader) Read(p []byte) (int, error) {
|
||||||
|
if l.remaining <= 0 {
|
||||||
|
return 0, ErrOutputTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
if int64(len(p)) > l.remaining {
|
||||||
|
p = p[:l.remaining]
|
||||||
|
}
|
||||||
|
|
||||||
|
n, err := l.r.Read(p)
|
||||||
|
l.remaining -= int64(n)
|
||||||
|
|
||||||
|
if l.remaining <= 0 {
|
||||||
|
// The (limit+1)th byte was just read: the stream is too long.
|
||||||
|
return n, ErrOutputTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
package blobgen_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"io"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestLimitReaderPassesExactSize checks that a stream of exactly the limit
|
||||||
|
// reads back cleanly to EOF: the bound must not reject a legitimate blob
|
||||||
|
// whose plaintext equals its recorded size.
|
||||||
|
func TestLimitReaderPassesExactSize(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const n = 1000
|
||||||
|
|
||||||
|
r := blobgen.LimitReader(bytes.NewReader(bytes.Repeat([]byte("a"), n)), n)
|
||||||
|
|
||||||
|
got, err := io.ReadAll(r)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, got, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLimitReaderFailsPastLimit feeds a large, highly compressible run of
|
||||||
|
// zeros — the decompressed output a zip bomb would produce — through a
|
||||||
|
// small limit and checks it fails within the bound rather than passing
|
||||||
|
// the whole stream through.
|
||||||
|
func TestLimitReaderFailsPastLimit(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const limit = 1000
|
||||||
|
|
||||||
|
r := blobgen.LimitReader(
|
||||||
|
bytes.NewReader(bytes.Repeat([]byte{0}, limit*1000)), limit)
|
||||||
|
|
||||||
|
n, err := io.Copy(io.Discard, r)
|
||||||
|
require.ErrorIs(t, err, blobgen.ErrOutputTooLarge)
|
||||||
|
require.LessOrEqual(t, n, int64(limit)+1,
|
||||||
|
"reader must stop within one byte of the limit")
|
||||||
|
}
|
||||||
@@ -0,0 +1,190 @@
|
|||||||
|
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")
|
||||||
|
}
|
||||||
@@ -2,6 +2,7 @@ package blobgen
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"hash"
|
"hash"
|
||||||
"io"
|
"io"
|
||||||
@@ -20,10 +21,12 @@ type Reader struct {
|
|||||||
bytesRead int64
|
bytesRead int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewReader creates a new Reader that decrypts, decompresses, and verifies data
|
// NewReader creates a new Reader that decrypts, decompresses, and verifies
|
||||||
func NewReader(r io.Reader, identity age.Identity) (*Reader, error) {
|
// data. Every supplied identity is offered to age.Decrypt, so a blob
|
||||||
|
// 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, identity)
|
decReader, err := age.Decrypt(r, identities...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("creating decryption reader: %w", err)
|
return nil, fmt.Errorf("creating decryption reader: %w", err)
|
||||||
}
|
}
|
||||||
@@ -54,6 +57,22 @@ func (r *Reader) Read(p []byte) (int, error) {
|
|||||||
n, err := r.teeReader.Read(p)
|
n, err := r.teeReader.Read(p)
|
||||||
r.bytesRead += int64(n)
|
r.bytesRead += int64(n)
|
||||||
|
|
||||||
|
// When the ciphertext is cut right after the age header plus its
|
||||||
|
// 16-byte nonce, the age reader's first read fails with
|
||||||
|
// io.ErrUnexpectedEOF, and the zstd decoder maps that to a clean
|
||||||
|
// io.EOF at frame start. That makes a truncated stream look like a
|
||||||
|
// valid empty one. Distinguish the two: on EOF, read once more from
|
||||||
|
// the age reader. A genuine end leaves it at (0, io.EOF); a truncated
|
||||||
|
// stream leaves its stored io.ErrUnexpectedEOF, which we surface.
|
||||||
|
if errors.Is(err, io.EOF) {
|
||||||
|
var probe [1]byte
|
||||||
|
|
||||||
|
m, ageErr := r.decryptor.Read(probe[:])
|
||||||
|
if m != 0 || !errors.Is(ageErr, io.EOF) {
|
||||||
|
return n, io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return n, err
|
return n, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -64,7 +83,9 @@ func (r *Reader) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Sum256 returns the SHA256 hash of all data read
|
// Sum256 returns the single SHA-256 of the plaintext read so far. This is the
|
||||||
|
// 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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,54 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,132 @@
|
|||||||
|
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)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,70 @@
|
|||||||
|
package blobgen_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"io"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"filippo.io/age"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestReaderRejectsHeaderNonceTruncation guards against a stream cut right
|
||||||
|
// after the age header plus its 16-byte nonce. age.Decrypt still succeeds on
|
||||||
|
// such an object, and the zstd decoder maps the age reader's
|
||||||
|
// io.ErrUnexpectedEOF to a clean io.EOF at frame start, so without the extra
|
||||||
|
// check the truncated stream would read as a valid empty one. Reading it must
|
||||||
|
// now fail.
|
||||||
|
func TestReaderRejectsHeaderNonceTruncation(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
identity, err := age.GenerateX25519Identity()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Encrypting empty plaintext yields header + nonce(16) + a single
|
||||||
|
// 16-byte final chunk tag. Dropping the trailing tag leaves exactly the
|
||||||
|
// age header plus its nonce — the truncation point that triggers the bug.
|
||||||
|
var full bytes.Buffer
|
||||||
|
|
||||||
|
w, err := age.Encrypt(&full, identity.Recipient())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, w.Close())
|
||||||
|
|
||||||
|
truncated := full.Bytes()[:full.Len()-16]
|
||||||
|
|
||||||
|
reader, err := blobgen.NewReader(bytes.NewReader(truncated), identity)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
defer func() { _ = reader.Close() }()
|
||||||
|
|
||||||
|
_, err = io.ReadAll(reader)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.ErrorIs(t, err, io.ErrUnexpectedEOF)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestReaderReadsGenuinelyEmptyBlob confirms the truncation check does not
|
||||||
|
// reject a legitimately empty payload: a blob written with no data must round
|
||||||
|
// trip back to zero bytes with no error.
|
||||||
|
func TestReaderReadsGenuinelyEmptyBlob(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
identity, err := age.GenerateX25519Identity()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var encrypted bytes.Buffer
|
||||||
|
|
||||||
|
writer, err := blobgen.NewWriter(
|
||||||
|
&encrypted, 3, []string{identity.Recipient().String()})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, writer.Close())
|
||||||
|
|
||||||
|
reader, err := blobgen.NewReader(bytes.NewReader(encrypted.Bytes()), identity)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
defer func() { _ = reader.Close() }()
|
||||||
|
|
||||||
|
data, err := io.ReadAll(reader)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Empty(t, data)
|
||||||
|
}
|
||||||
+30
-13
@@ -1,3 +1,6 @@
|
|||||||
|
// 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 (
|
||||||
@@ -12,6 +15,18 @@ 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
|
||||||
@@ -27,6 +42,11 @@ 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.
|
||||||
@@ -57,10 +77,12 @@ func NewWriter(
|
|||||||
// Parse recipients
|
// Parse recipients
|
||||||
var ageRecipients []age.Recipient
|
var ageRecipients []age.Recipient
|
||||||
|
|
||||||
for _, recipient := range recipients {
|
for i, 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("parsing recipient %s: %w", recipient, err)
|
return nil, fmt.Errorf("%w: recipient %d", errInvalidRecipient, i)
|
||||||
}
|
}
|
||||||
|
|
||||||
ageRecipients = append(ageRecipients, r)
|
ageRecipients = append(ageRecipients, r)
|
||||||
@@ -123,17 +145,12 @@ func (w *Writer) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Sum256 returns the double SHA256 hash of the uncompressed input data.
|
// ContentID returns the double SHA-256 of the uncompressed input data: the
|
||||||
// Double hashing (SHA256(SHA256(data))) prevents information leakage about
|
// name under which this content is stored. It is the second hash of the
|
||||||
// the plaintext - an attacker cannot confirm existence of known content
|
// running SHA-256, via DoubleSHA256; see that function for why content is
|
||||||
// by computing its hash and checking for a matching blob filename.
|
// named this way rather than by its plain SHA-256.
|
||||||
func (w *Writer) Sum256() []byte {
|
func (w *Writer) ContentID() []byte {
|
||||||
// First hash: SHA256(plaintext)
|
return DoubleSHA256(w.hasher.Sum(nil))
|
||||||
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,9 +12,10 @@ import (
|
|||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestWriterHashIsDoubleHash verifies that Writer.Sum256() returns
|
// TestWriterHashIsDoubleHash verifies that Writer.ContentID() returns
|
||||||
// the double hash SHA256(SHA256(plaintext)) for security.
|
// SHA256(SHA256(plaintext)). Stored objects are named by this second hash so a
|
||||||
// Double hashing prevents attackers from confirming existence of known content.
|
// name is not the plaintext's own SHA-256; this does not stop someone who
|
||||||
|
// already holds the plaintext from confirming it.
|
||||||
func TestWriterHashIsDoubleHash(t *testing.T) {
|
func TestWriterHashIsDoubleHash(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -43,7 +44,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.Sum256())
|
writerHash := hex.EncodeToString(writer.ContentID())
|
||||||
|
|
||||||
// Calculate the expected double hash: SHA256(SHA256(plaintext))
|
// Calculate the expected double hash: SHA256(SHA256(plaintext))
|
||||||
firstHash := sha256.Sum256(testData)
|
firstHash := sha256.Sum256(testData)
|
||||||
@@ -60,11 +61,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.Sum256() should return SHA256(SHA256(plaintext)) for security")
|
"Writer.ContentID() must be SHA256(SHA256(plaintext))")
|
||||||
|
|
||||||
// Verify it's NOT the single hash (would leak information)
|
// It must be the second hash, not the plaintext's own SHA-256.
|
||||||
assert.NotEqual(t, singleHashStr, writerHash,
|
assert.NotEqual(t, singleHashStr, writerHash,
|
||||||
"Writer hash should not be single hash (would allow content confirmation attacks)")
|
"Writer hash must be the double hash, not the single SHA-256")
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestWriterDeterministicHash verifies that the same input always produces
|
// TestWriterDeterministicHash verifies that the same input always produces
|
||||||
@@ -93,8 +94,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.Sum256())
|
hash1 := hex.EncodeToString(writer1.ContentID())
|
||||||
hash2 := hex.EncodeToString(writer2.Sum256())
|
hash2 := hex.EncodeToString(writer2.ContentID())
|
||||||
|
|
||||||
// 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")
|
||||||
@@ -108,3 +109,20 @@ 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,9 +33,10 @@ 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.
|
||||||
const chunkSizeSpread = 4
|
// The largest chunk the chunker can emit is therefore avg*ChunkSizeSpread.
|
||||||
|
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
|
||||||
@@ -45,8 +46,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),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+202
-114
@@ -7,11 +7,9 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
|
||||||
"os/signal"
|
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"syscall"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/adrg/xdg"
|
"github.com/adrg/xdg"
|
||||||
@@ -32,14 +30,33 @@ import (
|
|||||||
// may take before we give up.
|
// may take before we give up.
|
||||||
const shutdownTimeout = 30 * time.Second
|
const shutdownTimeout = 30 * time.Second
|
||||||
|
|
||||||
// AppOptions contains common options for creating the fx application.
|
// lockMode says whether a command mutates persistent state — the local
|
||||||
// It includes the configuration file path, logging options, and additional
|
// index database or the remote store — and so must hold the process-wide
|
||||||
// fx modules and invocations that should be included in the application.
|
// PID lock, or only reads that state and may run alongside a mutator.
|
||||||
|
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
|
||||||
@@ -48,6 +65,11 @@ 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,
|
||||||
) {
|
) {
|
||||||
@@ -55,7 +77,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 {
|
if opts.Cron || opts.Quiet || opts.JSON {
|
||||||
v.UI.SetQuiet(true)
|
v.UI.SetQuiet(true)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -136,75 +158,148 @@ func cleanStartupError(err error) error {
|
|||||||
return &startupError{msg: msg}
|
return &startupError{msg: msg}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RunApp starts and stops the fx application within the given context.
|
// RunApp starts the fx application, blocks until it is asked to stop, and
|
||||||
// It handles graceful shutdown on interrupt signals (SIGINT, SIGTERM) and
|
// then stops it. The app is asked to stop either by an OS interrupt
|
||||||
// ensures the application stops cleanly. The function blocks until the
|
// (SIGINT/SIGTERM — fx installs its own handler when app.Wait is called) or,
|
||||||
// application completes or is interrupted. Returns an error if startup fails.
|
// on normal completion, by the finished operation calling
|
||||||
|
// 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)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handle shutdown
|
// Block until an interrupt or the finished operation's
|
||||||
shutdownComplete := make(chan struct{})
|
// Shutdowner.Shutdown() arrives, then stop the app in this goroutine so we
|
||||||
|
// 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()
|
||||||
|
|
||||||
go func() {
|
shutdownCtx, cancel := context.WithTimeout(
|
||||||
defer close(shutdownComplete)
|
context.WithoutCancel(ctx), shutdownTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
<-sigChan
|
err = app.Stop(shutdownCtx)
|
||||||
log.Notice("Received interrupt signal, shutting down gracefully...")
|
if err != nil {
|
||||||
|
log.Error("Error during shutdown", "error", err)
|
||||||
// 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)
|
|
||||||
defer shutdownCancel()
|
|
||||||
|
|
||||||
err := app.Stop(shutdownCtx)
|
|
||||||
if err != nil {
|
|
||||||
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
|
|
||||||
case <-ctx.Done():
|
|
||||||
// Context cancelled (shouldn't happen in normal operation)
|
|
||||||
err := app.Stop(context.WithoutCancel(ctx))
|
|
||||||
if err != nil {
|
|
||||||
log.Error("Error stopping app", "error", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return ctx.Err()
|
|
||||||
case <-app.Done():
|
|
||||||
// App finished running (e.g., backup completed)
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// errReported marks a failure the operation has already shown the user
|
||||||
|
// (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 {
|
||||||
|
return 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 nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// runVaultikApp runs the standard single-operation command lifecycle
|
// runVaultikApp runs the standard single-operation command lifecycle
|
||||||
// shared by the list/purge/verify/remove/remote-info subcommands:
|
// shared by the snapshot list/purge/remove and remote nuke subcommands:
|
||||||
// resolve the config, start the fx app, run op against the Vaultik
|
// resolve the config, then run op against the Vaultik instance through
|
||||||
// instance in a goroutine, report a failure prefixed with failMsg
|
// RunOperation, reporting a failure prefixed with failMsg (suppressed
|
||||||
// (suppressed while suppressErrors is true, e.g. under --json), then
|
// while suppressErrors is true, e.g. under --json). mode says whether the
|
||||||
// trigger shutdown. The operation is cancelled when the app stops.
|
// command takes the PID lock. jsonOutput marks a command whose stdout is a
|
||||||
// extraQuiet is OR-ed into LogOptions.Quiet (e.g. --json output modes).
|
// JSON document: it quiets the UI but, unlike Quiet, leaves the stderr log
|
||||||
|
// level alone.
|
||||||
func runVaultikApp(
|
func runVaultikApp(
|
||||||
cmd *cobra.Command, extraQuiet, suppressErrors bool,
|
cmd *cobra.Command, mode lockMode, jsonOutput, 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()
|
||||||
@@ -214,75 +309,68 @@ func runVaultikApp(
|
|||||||
|
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
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,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet || extraQuiet,
|
Quiet: rootFlags.Quiet,
|
||||||
|
JSON: jsonOutput,
|
||||||
},
|
},
|
||||||
Modules: []fx.Option{},
|
Mode: mode,
|
||||||
Invokes: []fx.Option{
|
}, op, func(err error) {
|
||||||
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
if suppressErrors {
|
||||||
lc.Append(fx.Hook{
|
return
|
||||||
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)
|
|
||||||
ReportErrorf("%s: %v", failMsg, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
os.Exit(1)
|
log.Error(failMsg, "error", err)
|
||||||
}
|
ReportErrorf("%s: %v", failMsg, err)
|
||||||
}
|
|
||||||
|
|
||||||
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.
|
||||||
// It acquires a PID lock before starting to prevent concurrent instances.
|
// A mutating command takes the process-wide PID lock before starting so that
|
||||||
|
// 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 {
|
||||||
// Acquire PID lock to prevent concurrent instances
|
release, err := acquireLockIfMutating(opts.Mode,
|
||||||
lockDir := filepath.Join(xdg.DataHome, "vaultik")
|
filepath.Join(xdg.DataHome, "vaultik"))
|
||||||
|
|
||||||
lock, err := pidlock.Acquire(lockDir)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, pidlock.ErrAlreadyRunning) {
|
return err
|
||||||
return fmt.Errorf("cannot start: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return fmt.Errorf("failed to acquire lock: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
defer func() {
|
defer release()
|
||||||
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,7 +2,10 @@ 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) {
|
||||||
@@ -53,3 +56,42 @@ 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()
|
||||||
|
}
|
||||||
|
|||||||
+47
-31
@@ -4,6 +4,7 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -192,8 +193,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 ────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
@@ -212,6 +213,8 @@ 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
|
||||||
@@ -377,40 +380,53 @@ Examples:
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
root, err := loadYAMLFile(path)
|
return writeConfigSet(os.Stdout, path, args[0], args[1])
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
err = yamlPathSet(root, strings.Split(args[0], "."), args[1])
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
out, err := marshalConfigYAML(root)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("marshaling config: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
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 {
|
|
||||||
return fmt.Errorf("writing config file: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
_, _ = fmt.Fprintf(os.Stdout, "%s = %s\n", args[0], args[1])
|
|
||||||
|
|
||||||
return nil
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
err = yamlPathSet(root, strings.Split(key, "."), value)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
out, err := marshalConfigYAML(root)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("marshaling config: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = os.WriteFile(path, out, configFileMode)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("writing config file: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// os.WriteFile does not change the mode of a file that already exists,
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
// marshalConfigYAML renders a config document tree with 2-space indentation,
|
// marshalConfigYAML renders a config document tree with 2-space indentation,
|
||||||
// matching defaultConfigTemplate. yaml.Marshal defaults to 4 spaces, which
|
// matching defaultConfigTemplate. yaml.Marshal defaults to 4 spaces, which
|
||||||
// would reindent the whole file on the first `config set` despite the promise
|
// would reindent the whole file on the first `config set` despite the promise
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
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"
|
||||||
|
|
||||||
@@ -229,6 +232,68 @@ 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, ".")
|
||||||
}
|
}
|
||||||
|
|||||||
+18
-3
@@ -1,6 +1,7 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -19,7 +20,11 @@ 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()
|
||||||
@@ -27,9 +32,19 @@ func Entry() {
|
|||||||
|
|
||||||
err := rootCmd.Execute()
|
err := rootCmd.Execute()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
ReportErrorf("%s", err.Error())
|
// An operation that ran inside the fx app has already reported
|
||||||
os.Exit(1)
|
// 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())
|
||||||
|
}
|
||||||
|
|
||||||
|
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, Entry)
|
stdout := captureProcessStdout(t, func() { _ = Entry() })
|
||||||
|
|
||||||
requireExactlyOneJSONDocument(t, stdout)
|
requireExactlyOneJSONDocument(t, stdout)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,140 @@
|
|||||||
|
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, Entry)
|
stdout := captureProcessStdout(t, func() { _ = Entry() })
|
||||||
|
|
||||||
requireExactlyOneJSONDocument(t, stdout)
|
requireExactlyOneJSONDocument(t, stdout)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,58 @@
|
|||||||
|
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)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+7
-37
@@ -1,12 +1,7 @@
|
|||||||
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"
|
||||||
)
|
)
|
||||||
@@ -33,44 +28,19 @@ func NewInfoCommand() *cobra.Command {
|
|||||||
// Use the app framework
|
// Use the app framework
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
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,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet,
|
||||||
},
|
},
|
||||||
Modules: []fx.Option{},
|
Mode: readOnly,
|
||||||
Invokes: []fx.Option{
|
}, func(v *vaultik.Vaultik) error {
|
||||||
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
return v.ShowInfo()
|
||||||
lc.Append(fx.Hook{
|
}, func(err error) {
|
||||||
OnStart: func(_ context.Context) error {
|
log.Error("Failed to show info", "error", err)
|
||||||
go func() {
|
ReportErrorf("Failed to show info: %v", err)
|
||||||
err := v.ShowInfo()
|
|
||||||
if err != nil {
|
|
||||||
if !errors.Is(err, context.Canceled) {
|
|
||||||
log.Error("Failed to show info", "error", 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
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}),
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
+12
-44
@@ -1,12 +1,7 @@
|
|||||||
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"
|
||||||
)
|
)
|
||||||
@@ -41,51 +36,24 @@ 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 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,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet || opts.JSON,
|
Quiet: rootFlags.Quiet,
|
||||||
|
JSON: opts.JSON,
|
||||||
},
|
},
|
||||||
Modules: []fx.Option{},
|
Mode: mutating,
|
||||||
Invokes: []fx.Option{
|
}, func(v *vaultik.Vaultik) error {
|
||||||
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
return v.Prune(opts)
|
||||||
lc.Append(fx.Hook{
|
}, func(err error) {
|
||||||
OnStart: func(_ context.Context) error {
|
if opts.JSON {
|
||||||
// Start the prune operation in a goroutine
|
return
|
||||||
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)
|
|
||||||
ReportErrorf("Prune failed: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
os.Exit(1)
|
log.Error("Prune operation failed", "error", err)
|
||||||
}
|
ReportErrorf("Prune failed: %v", err)
|
||||||
}
|
|
||||||
|
|
||||||
// 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
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}),
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
+13
-39
@@ -1,12 +1,9 @@
|
|||||||
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"
|
||||||
)
|
)
|
||||||
@@ -48,7 +45,7 @@ This is destructive and irreversible. Requires --force.`,
|
|||||||
return errNukeNeedsForce
|
return errNukeNeedsForce
|
||||||
}
|
}
|
||||||
|
|
||||||
return runVaultikApp(cmd, false, false, "Remote nuke failed",
|
return runVaultikApp(cmd, mutating, false, false, "Remote nuke failed",
|
||||||
func(v *vaultik.Vaultik) error {
|
func(v *vaultik.Vaultik) error {
|
||||||
return v.NukeRemote(true)
|
return v.NukeRemote(true)
|
||||||
})
|
})
|
||||||
@@ -83,47 +80,24 @@ func newRemoteInfoCommand() *cobra.Command {
|
|||||||
|
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
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,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet || jsonOutput,
|
Quiet: rootFlags.Quiet,
|
||||||
|
JSON: jsonOutput,
|
||||||
},
|
},
|
||||||
Modules: []fx.Option{},
|
Mode: readOnly,
|
||||||
Invokes: []fx.Option{
|
}, func(v *vaultik.Vaultik) error {
|
||||||
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
return v.RemoteInfo(jsonOutput)
|
||||||
lc.Append(fx.Hook{
|
}, func(err error) {
|
||||||
OnStart: func(_ context.Context) error {
|
if jsonOutput {
|
||||||
go func() {
|
return
|
||||||
err := v.RemoteInfo(jsonOutput)
|
}
|
||||||
if err != nil {
|
|
||||||
if !errors.Is(err, context.Canceled) {
|
|
||||||
if !jsonOutput {
|
|
||||||
log.Error("Failed to get remote info", "error", err)
|
|
||||||
ReportErrorf("Failed to get remote info: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
os.Exit(1)
|
log.Error("Failed to get remote info", "error", err)
|
||||||
}
|
ReportErrorf("Failed to get remote info: %v", err)
|
||||||
}
|
|
||||||
|
|
||||||
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,8 +57,9 @@ 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,
|
||||||
"Continue past per-file errors instead of aborting "+
|
"Skip files that cannot be read when creating a snapshot, or "+
|
||||||
"(applies to snapshot create and restore)")
|
"that cannot be restored when restoring, instead of aborting "+
|
||||||
|
"(packing and storage errors still abort)")
|
||||||
|
|
||||||
// Add subcommands
|
// Add subcommands
|
||||||
cmd.AddCommand(
|
cmd.AddCommand(
|
||||||
|
|||||||
@@ -0,0 +1,97 @@
|
|||||||
|
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")
|
||||||
|
}
|
||||||
+28
-80
@@ -1,13 +1,10 @@
|
|||||||
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"
|
||||||
)
|
)
|
||||||
@@ -86,7 +83,8 @@ 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()
|
||||||
|
|
||||||
return RunWithApp(cmd.Context(), AppOptions{
|
// --cron suppression is wired through v.UI by setupGlobals.
|
||||||
|
return RunOperation(cmd.Context(), AppOptions{
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogOptions: log.Options{
|
LogOptions: log.Options{
|
||||||
Verbose: rootFlags.Verbose,
|
Verbose: rootFlags.Verbose,
|
||||||
@@ -94,42 +92,12 @@ 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,
|
||||||
},
|
},
|
||||||
Modules: []fx.Option{},
|
Mode: mutating,
|
||||||
Invokes: []fx.Option{
|
}, func(v *vaultik.Vaultik) error {
|
||||||
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
return v.CreateSnapshot(opts)
|
||||||
lc.Append(fx.Hook{
|
}, func(err error) {
|
||||||
OnStart: func(_ context.Context) error {
|
log.Error("Snapshot creation failed", "error", err)
|
||||||
// Start the snapshot creation in a goroutine
|
ReportErrorf("Snapshot creation failed: %v", err)
|
||||||
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)
|
|
||||||
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
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}),
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -158,7 +126,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, false, false,
|
return runVaultikApp(cmd, readOnly, 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)
|
||||||
@@ -194,7 +162,7 @@ restrict the operation to specific snapshot names.`,
|
|||||||
return errPurgeCriteriaBoth
|
return errPurgeCriteriaBoth
|
||||||
}
|
}
|
||||||
|
|
||||||
return runVaultikApp(cmd, false, false,
|
return runVaultikApp(cmd, mutating, 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)
|
||||||
@@ -220,8 +188,11 @@ func newSnapshotVerifyCommand() *cobra.Command {
|
|||||||
|
|
||||||
cmd := &cobra.Command{
|
cmd := &cobra.Command{
|
||||||
Use: "verify <snapshot-id>",
|
Use: "verify <snapshot-id>",
|
||||||
Short: "Verify snapshot integrity",
|
Short: "Check a snapshot's blobs are present with the listed size",
|
||||||
Long: "Verifies that all blobs referenced in a snapshot exist.\n\n" +
|
Long: "Checks that every blob the snapshot's manifest lists is present\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).",
|
||||||
@@ -237,47 +208,24 @@ func newSnapshotVerifyCommand() *cobra.Command {
|
|||||||
|
|
||||||
rootFlags := GetRootFlags()
|
rootFlags := GetRootFlags()
|
||||||
|
|
||||||
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,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet || opts.JSON,
|
Quiet: rootFlags.Quiet,
|
||||||
|
JSON: opts.JSON,
|
||||||
},
|
},
|
||||||
Modules: []fx.Option{},
|
Mode: readOnly,
|
||||||
Invokes: []fx.Option{
|
}, func(v *vaultik.Vaultik) error {
|
||||||
fx.Invoke(func(v *vaultik.Vaultik, lc fx.Lifecycle) {
|
return v.VerifySnapshotWithOptions(snapshotID, opts)
|
||||||
lc.Append(fx.Hook{
|
}, func(err error) {
|
||||||
OnStart: func(_ context.Context) error {
|
if opts.JSON {
|
||||||
go func() {
|
return
|
||||||
err := v.VerifySnapshotWithOptions(snapshotID, opts)
|
}
|
||||||
if err != nil {
|
|
||||||
if !errors.Is(err, context.Canceled) {
|
|
||||||
if !opts.JSON {
|
|
||||||
log.Error("Verification failed", "error", err)
|
|
||||||
ReportErrorf("Verification failed: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
os.Exit(1)
|
log.Error("Verification failed", "error", err)
|
||||||
}
|
ReportErrorf("Verification failed: %v", err)
|
||||||
}
|
|
||||||
|
|
||||||
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
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}),
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -316,7 +264,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, opts.JSON, opts.JSON,
|
return runVaultikApp(cmd, mutating, 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,16 +1,8 @@
|
|||||||
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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -25,15 +17,6 @@ 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{}
|
||||||
@@ -52,8 +35,12 @@ 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 VAULTIK_AGE_SECRET_KEY environment variable to be set with
|
Requires the age private key in the VAULTIK_AGE_SECRET_KEY environment
|
||||||
the age private key.
|
variable. The variable may hold the whole age-keygen file (comments and
|
||||||
|
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
|
||||||
@@ -81,7 +68,8 @@ Examples:
|
|||||||
return cmd
|
return cmd
|
||||||
}
|
}
|
||||||
|
|
||||||
// runRestore parses arguments and runs the restore operation through the app framework
|
// runRestore parses arguments and runs the restore operation through the
|
||||||
|
// 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]
|
||||||
|
|
||||||
@@ -90,87 +78,31 @@ 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 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,
|
||||||
Debug: rootFlags.Debug,
|
Debug: rootFlags.Debug,
|
||||||
Quiet: rootFlags.Quiet,
|
Quiet: rootFlags.Quiet,
|
||||||
},
|
},
|
||||||
Modules: buildRestoreModules(),
|
Mode: readOnly,
|
||||||
Invokes: buildRestoreInvokes(snapshotID, opts),
|
}, func(v *vaultik.Vaultik) error {
|
||||||
|
return v.Restore(&vaultik.RestoreOptions{
|
||||||
|
SnapshotID: snapshotID,
|
||||||
|
TargetDir: opts.TargetDir,
|
||||||
|
Paths: opts.Paths,
|
||||||
|
Verify: opts.Verify,
|
||||||
|
SkipErrors: rootFlags.SkipErrors,
|
||||||
|
})
|
||||||
|
}, func(err error) {
|
||||||
|
log.Error("Restore operation failed", "error", err)
|
||||||
|
ReportErrorf("Restore failed: %v", err)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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,
|
|
||||||
TargetDir: opts.TargetDir,
|
|
||||||
Paths: opts.Paths,
|
|
||||||
Verify: opts.Verify,
|
|
||||||
SkipErrors: GetRootFlags().SkipErrors,
|
|
||||||
}
|
|
||||||
|
|
||||||
err := app.Vaultik.Restore(restoreOpts)
|
|
||||||
if err != nil {
|
|
||||||
if !errors.Is(err, context.Canceled) {
|
|
||||||
log.Error("Restore operation failed", "error", 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
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -0,0 +1,37 @@
|
|||||||
|
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 ...)")
|
||||||
|
}
|
||||||
|
}
|
||||||
+106
-35
@@ -16,11 +16,19 @@ 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
|
||||||
@@ -37,13 +45,19 @@ 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("blob_size_limit must be at least chunk_size")
|
errBlobSizeTooSmall = errors.New(
|
||||||
errBadCompression = errors.New("compression_level must be between 1 and 19")
|
"blob_size_limit must be at least the largest chunk the chunker can " +
|
||||||
errBadStorageScheme = errors.New(
|
"emit (chunk_size times the FastCDC size spread)")
|
||||||
|
errBadCompression = errors.New("compression_level must be between 1 and 19")
|
||||||
|
errBadStorageScheme = errors.New(
|
||||||
"storage_url must start with s3://, file://, or rclone://")
|
"storage_url must start with s3://, file://, or rclone://")
|
||||||
errStorageNotConfigured = errors.New(
|
errStorageNotConfigured = errors.New(
|
||||||
"storage not configured; set storage_url or provide s3.endpoint + " +
|
"storage not configured; set storage_url or provide s3.endpoint + " +
|
||||||
@@ -121,6 +135,26 @@ 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.
|
||||||
@@ -130,8 +164,13 @@ func (c *Config) SnapshotNames() []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"`
|
||||||
BlobSizeLimit Size `yaml:"blob_size_limit"`
|
// AgeSecretKeySource names where AgeSecretKey was configured
|
||||||
ChunkSize Size `yaml:"chunk_size"`
|
// ("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"`
|
||||||
|
ChunkSize Size `yaml:"chunk_size"`
|
||||||
// Exclude holds global excludes applied to all snapshots.
|
// Exclude holds global excludes applied to all snapshots.
|
||||||
Exclude []string `yaml:"exclude"`
|
Exclude []string `yaml:"exclude"`
|
||||||
Hostname string `yaml:"hostname"`
|
Hostname string `yaml:"hostname"`
|
||||||
@@ -162,8 +201,10 @@ 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 bool `yaml:"use_ssl"`
|
// UseSSL selects HTTPS for a scheme-less endpoint. Omitted (nil) means
|
||||||
PartSize Size `yaml:"part_size"`
|
// the default, TLS; set it to false only to force plain HTTP.
|
||||||
|
UseSSL *bool `yaml:"use_ssl"`
|
||||||
|
PartSize Size `yaml:"part_size"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Path wraps the config file path for fx dependency injection.
|
// Path wraps the config file path for fx dependency injection.
|
||||||
@@ -238,10 +279,7 @@ func Load(path string) (*Config, error) {
|
|||||||
cfg.IndexPath = expandTilde(envIndexPath)
|
cfg.IndexPath = expandTilde(envIndexPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check for environment variable override for AgeSecretKey
|
cfg.setAgeSecretKey()
|
||||||
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 == "" {
|
||||||
@@ -285,18 +323,30 @@ 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
|
// - At least one age recipient must be specified, and every recipient must
|
||||||
// - At least one snapshot must be configured with at least one path
|
// parse as an X25519 age1... public key (so a bad entry fails at load, not
|
||||||
// - Storage must be configured (either storage_url or s3.* fields)
|
// mid-backup); errors name the position, never the value
|
||||||
// - Chunk size must be at least 1MB
|
// - At least one snapshot must be configured with at least one path
|
||||||
// - Blob size limit must be at least the chunk size
|
// - Storage must be configured (either storage_url or s3.* fields)
|
||||||
// - Compression level must be between 1 and 19
|
// - Chunk size must be at least 1MB
|
||||||
|
// - Blob size limit must be at least the largest chunk the chunker can emit
|
||||||
|
// (chunk_size times chunker.ChunkSizeSpread), so a single-chunk blob never
|
||||||
|
// exceeds the configured limit
|
||||||
|
// - 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
|
||||||
}
|
}
|
||||||
@@ -317,8 +367,13 @@ func (c *Config) Validate() error {
|
|||||||
return errChunkSizeTooSmall
|
return errChunkSizeTooSmall
|
||||||
}
|
}
|
||||||
|
|
||||||
if c.BlobSizeLimit.Int64() < c.ChunkSize.Int64() {
|
// The chunker can emit chunks up to chunk_size * ChunkSizeSpread, and the
|
||||||
return errBlobSizeTooSmall
|
// packer places a single such chunk into an otherwise empty blob. A limit
|
||||||
|
// 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 ||
|
||||||
@@ -329,6 +384,38 @@ 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.
|
||||||
@@ -385,22 +472,6 @@ 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.
|
||||||
//
|
//
|
||||||
|
|||||||
+213
-39
@@ -1,9 +1,13 @@
|
|||||||
package config //nolint:testpackage // exercises unexported extractAgeSecretKey
|
package config //nolint:testpackage // exercises unexported source constants
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"sneak.berlin/go/vaultik/internal/chunker"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -83,6 +87,48 @@ 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()
|
||||||
@@ -101,53 +147,57 @@ func TestConfigFromEnv(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestExtractAgeSecretKey tests extraction of AGE-SECRET-KEY from various inputs
|
// TestValidateBlobSizeLimit checks the blob_size_limit boundary: it must be at
|
||||||
func TestExtractAgeSecretKey(t *testing.T) {
|
// least the largest chunk the chunker can emit (chunk_size times
|
||||||
|
// 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
|
||||||
input string
|
blobLimit Size
|
||||||
expected string
|
wantErr bool
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "plain key",
|
name: "at chunk_size but below largest chunk is rejected",
|
||||||
input: testIntegrationAgePrivateKey,
|
blobLimit: chunkSize,
|
||||||
expected: testIntegrationAgePrivateKey,
|
wantErr: true,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "key with trailing newline",
|
name: "between chunk_size and largest chunk is rejected",
|
||||||
input: testIntegrationAgePrivateKey + "\n",
|
blobLimit: Size(chunkSize.Int64() * 2),
|
||||||
expected: testIntegrationAgePrivateKey,
|
wantErr: true,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "full age-keygen output",
|
name: "one byte below largest chunk is rejected",
|
||||||
input: "# created: 2025-01-14T12:00:00Z\n" +
|
blobLimit: Size(largestChunk - 1),
|
||||||
"# public key: " + testIntegrationAgePublicKey + "\n" +
|
wantErr: true,
|
||||||
testIntegrationAgePrivateKey + "\n",
|
|
||||||
expected: testIntegrationAgePrivateKey,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "age-keygen output with extra blank lines",
|
name: "exactly at largest chunk is accepted",
|
||||||
input: "# created: 2025-01-14T12:00:00Z\n" +
|
blobLimit: Size(largestChunk),
|
||||||
"# public key: " + testIntegrationAgePublicKey + "\n\n" +
|
wantErr: false,
|
||||||
testIntegrationAgePrivateKey + "\n\n",
|
|
||||||
expected: testIntegrationAgePrivateKey,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "key with leading whitespace",
|
name: "above largest chunk is accepted",
|
||||||
input: " " + testIntegrationAgePrivateKey + " ",
|
blobLimit: Size(largestChunk * 100),
|
||||||
expected: testIntegrationAgePrivateKey,
|
wantErr: false,
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "empty input",
|
|
||||||
input: "",
|
|
||||||
expected: "",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "only comments",
|
|
||||||
input: "# this is a comment\n# another comment",
|
|
||||||
expected: "# this is a comment\n# another comment",
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -155,10 +205,134 @@ func TestExtractAgeSecretKey(t *testing.T) {
|
|||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
result := extractAgeSecretKey(tt.input)
|
err := newConfig(tt.blobLimit).Validate()
|
||||||
if result != tt.expected {
|
if tt.wantErr {
|
||||||
t.Errorf("extractAgeSecretKey(%q) = %q, want %q",
|
if !errors.Is(err, errBlobSizeTooSmall) {
|
||||||
tt.input, result, tt.expected)
|
t.Fatalf("Validate() error = %v, want errBlobSizeTooSmall", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
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)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,224 +0,0 @@
|
|||||||
// 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")
|
|
||||||
@@ -1,178 +0,0 @@
|
|||||||
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,6 +208,30 @@ 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,12 +7,32 @@ 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) {
|
||||||
query := `
|
return r.list(ctx, `
|
||||||
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,6 +17,7 @@ import (
|
|||||||
"embed"
|
"embed"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net/url"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"sort"
|
"sort"
|
||||||
@@ -219,6 +220,135 @@ 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,6 +2,7 @@ package database
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -15,6 +16,44 @@ 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
|
||||||
@@ -34,6 +73,11 @@ 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)
|
||||||
|
|||||||
@@ -0,0 +1,100 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,145 @@
|
|||||||
|
//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,6 +11,18 @@ 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 {
|
||||||
@@ -206,6 +218,48 @@ 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,
|
||||||
|
|||||||
+15
-1
@@ -14,8 +14,18 @@ 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(opts)
|
return Config{
|
||||||
|
Verbose: opts.Verbose,
|
||||||
|
Debug: opts.Debug,
|
||||||
|
Cron: opts.Cron,
|
||||||
|
Quiet: opts.Quiet,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Options are provided by the CLI.
|
// Options are provided by the CLI.
|
||||||
@@ -24,4 +34,8 @@ 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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,39 @@
|
|||||||
|
package log_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"log/slog"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"sneak.berlin/go/vaultik/internal/log"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestTTYHandlerEscapesControlCharacters logs a message and an attribute
|
||||||
|
// value that each carry an ESC and a newline — the shape a crafted path or
|
||||||
|
// storage error from the destination would take — and checks neither raw
|
||||||
|
// byte reaches the output. The handler's own colour codes (ESC ... m) are
|
||||||
|
// stripped first; any ESC left after that came from the untrusted value.
|
||||||
|
func TestTTYHandlerEscapesControlCharacters(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
|
||||||
|
logger := slog.New(log.NewTTYHandler(&buf, debugHandlerOptions()))
|
||||||
|
logger.Info("start\x1b[31mZAP\nend", "target", "a\x1b[31mZAP\nb")
|
||||||
|
|
||||||
|
out := buf.String()
|
||||||
|
|
||||||
|
// The only newline is the line terminator; the injected ones were escaped.
|
||||||
|
require.Equal(t, 1, strings.Count(out, "\n"),
|
||||||
|
"a newline in the message or a value must be escaped, not emitted raw")
|
||||||
|
|
||||||
|
// After the handler's own colour codes are removed, no ESC survives.
|
||||||
|
stripped := ansiEscape.ReplaceAllString(out, "")
|
||||||
|
require.NotContains(t, stripped, "\x1b",
|
||||||
|
"a raw ESC from the message or a value must not reach the terminal")
|
||||||
|
|
||||||
|
// The escaped form is what appears instead.
|
||||||
|
require.Contains(t, out, `\x1b`)
|
||||||
|
}
|
||||||
@@ -5,9 +5,11 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
"unicode"
|
||||||
)
|
)
|
||||||
|
|
||||||
// groupSeparator joins an open group path to an attribute key. This
|
// groupSeparator joins an open group path to an attribute key. This
|
||||||
@@ -116,11 +118,14 @@ func (h *TTYHandler) Handle(_ context.Context, r slog.Record) error {
|
|||||||
levelColor = colorReset
|
levelColor = colorReset
|
||||||
}
|
}
|
||||||
|
|
||||||
// Print main message
|
// Print main message. The message is escaped before the colour codes
|
||||||
|
// are written around it: it can carry text from an untrusted source
|
||||||
|
// (a storage error, for one), and a raw control character would
|
||||||
|
// otherwise reach the terminal.
|
||||||
_, _ = fmt.Fprintf(h.out, "%s%s%s %s%s%s %s%s%s",
|
_, _ = fmt.Fprintf(h.out, "%s%s%s %s%s%s %s%s%s",
|
||||||
colorGray, timestamp, colorReset,
|
colorGray, timestamp, colorReset,
|
||||||
levelColor, level, colorReset,
|
levelColor, level, colorReset,
|
||||||
colorBold, r.Message, colorReset)
|
colorBold, sanitize(r.Message), colorReset)
|
||||||
|
|
||||||
// Attributes carried by the handler come first, then the record's
|
// Attributes carried by the handler come first, then the record's
|
||||||
// own. Handler attributes were qualified when they were added; the
|
// own. Handler attributes were qualified when they were added; the
|
||||||
@@ -260,9 +265,29 @@ func (h *TTYHandler) writeAttr(a slog.Attr) {
|
|||||||
// Future kinds also use the plain string form.
|
// Future kinds also use the plain string form.
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Escape the key and value before the colour codes are written around
|
||||||
|
// them. Both can carry text from an untrusted source — a manifest
|
||||||
|
// timestamp, a storage error, a path or symlink target read back from
|
||||||
|
// the snapshot database — so a control character in one of them must
|
||||||
|
// be rendered as an escape sequence rather than reaching the terminal,
|
||||||
|
// where it could move the cursor or inject its own colours.
|
||||||
_, _ = fmt.Fprintf(h.out, " %s%s%s=%s%s%s",
|
_, _ = fmt.Fprintf(h.out, " %s%s%s=%s%s%s",
|
||||||
colorCyan, a.Key, colorReset,
|
colorCyan, sanitize(a.Key), colorReset,
|
||||||
colorBlue, value, colorReset)
|
colorBlue, sanitize(value), colorReset)
|
||||||
|
}
|
||||||
|
|
||||||
|
// sanitize returns s unchanged when every rune in it is printable, and a
|
||||||
|
// double-quoted, backslash-escaped form (\n, \x1b, …) otherwise. It is
|
||||||
|
// applied to untrusted text before any colour code is written, so a
|
||||||
|
// control character can never reach the terminal raw.
|
||||||
|
func sanitize(s string) string {
|
||||||
|
for _, r := range s {
|
||||||
|
if !unicode.IsPrint(r) {
|
||||||
|
return strconv.Quote(s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return s
|
||||||
}
|
}
|
||||||
|
|
||||||
// formatDuration formats a duration in a human-readable way
|
// formatDuration formats a duration in a human-readable way
|
||||||
|
|||||||
@@ -0,0 +1,49 @@
|
|||||||
|
//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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -7,6 +7,19 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
|
|
||||||
"github.com/klauspost/compress/zstd"
|
"github.com/klauspost/compress/zstd"
|
||||||
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Manifest size bounds. A manifest lists one small entry per blob, and
|
||||||
|
// blobs are large (the default target is 10 GB), so even a manifest for a
|
||||||
|
// petabyte-scale backup is a few megabytes. These caps are far above any
|
||||||
|
// manifest the writer can emit, yet stop a crafted, highly compressible
|
||||||
|
// manifest from expanding without limit when decoded: the manifest is
|
||||||
|
// fetched from the store, which is not trusted, and json.Decode buffers
|
||||||
|
// the whole value in memory.
|
||||||
|
const (
|
||||||
|
manifestMaxCompressed = 256 * 1024 * 1024 // 256 MiB
|
||||||
|
manifestMaxDecompressed = 1024 * 1024 * 1024 // 1 GiB
|
||||||
)
|
)
|
||||||
|
|
||||||
// Manifest represents the structure of a snapshot's blob manifest
|
// Manifest represents the structure of a snapshot's blob manifest
|
||||||
@@ -28,19 +41,31 @@ type BlobInfo struct {
|
|||||||
CompressedSize int64 `json:"compressed_size"`
|
CompressedSize int64 `json:"compressed_size"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// DecodeManifest decodes a manifest from a reader containing compressed JSON
|
// DecodeManifest decodes a manifest from a reader containing compressed
|
||||||
|
// JSON, reading through byte limits on both the compressed input and the
|
||||||
|
// decompressed output so an untrusted manifest cannot exhaust memory.
|
||||||
func DecodeManifest(r io.Reader) (*Manifest, error) {
|
func DecodeManifest(r io.Reader) (*Manifest, error) {
|
||||||
// Decompress using zstd
|
return decodeManifest(r, manifestMaxCompressed, manifestMaxDecompressed)
|
||||||
zr, err := zstd.NewReader(r)
|
}
|
||||||
|
|
||||||
|
// decodeManifest is DecodeManifest with explicit limits, so tests can drive
|
||||||
|
// the bounds with small inputs instead of gigabyte-scale ones.
|
||||||
|
func decodeManifest(
|
||||||
|
r io.Reader, maxCompressed, maxDecompressed int64,
|
||||||
|
) (*Manifest, error) {
|
||||||
|
// Decompress using zstd, bounding how many compressed bytes are read.
|
||||||
|
zr, err := zstd.NewReader(blobgen.LimitReader(r, maxCompressed))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("creating zstd reader: %w", err)
|
return nil, fmt.Errorf("creating zstd reader: %w", err)
|
||||||
}
|
}
|
||||||
defer zr.Close()
|
defer zr.Close()
|
||||||
|
|
||||||
// Decode JSON manifest
|
// Decode JSON manifest, bounding how far the compressed input may
|
||||||
|
// expand: json.Decode buffers the whole value, so without this a
|
||||||
|
// small, highly compressible manifest could expand to gigabytes.
|
||||||
var manifest Manifest
|
var manifest Manifest
|
||||||
|
|
||||||
err = json.NewDecoder(zr).Decode(&manifest)
|
err = json.NewDecoder(blobgen.LimitReader(zr, maxDecompressed)).Decode(&manifest)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("decoding manifest: %w", err)
|
return nil, fmt.Errorf("decoding manifest: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,79 @@
|
|||||||
|
//nolint:testpackage // exercises the unexported decodeManifest bounds
|
||||||
|
package snapshot
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testSnapshotID is a stand-in snapshot ID reused across the bound cases.
|
||||||
|
const testSnapshotID = "host_home_2026-01-01T00:00:00Z"
|
||||||
|
|
||||||
|
// TestDecodeManifestRoundTrip is the baseline: with generous bounds a
|
||||||
|
// manifest the writer produced decodes back unchanged.
|
||||||
|
func TestDecodeManifestRoundTrip(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
want := &Manifest{
|
||||||
|
SnapshotID: testSnapshotID,
|
||||||
|
Timestamp: "2026-01-01T00:00:00Z",
|
||||||
|
BlobCount: 2,
|
||||||
|
TotalCompressedSize: 42,
|
||||||
|
Blobs: []BlobInfo{
|
||||||
|
{Hash: "aa", CompressedSize: 21},
|
||||||
|
{Hash: "bb", CompressedSize: 21},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
compressed, err := EncodeManifest(want, 3)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
got, err := decodeManifest(
|
||||||
|
bytes.NewReader(compressed), manifestMaxCompressed, manifestMaxDecompressed)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, want, got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDecodeManifestBoundsDecompressedOutput feeds a valid but highly
|
||||||
|
// compressible manifest — one whose timestamp is a megabyte of the same
|
||||||
|
// character — through a small decompressed bound. The compressed form is
|
||||||
|
// tiny, so only the decompressed bound stops it; decoding must fail within
|
||||||
|
// that bound rather than expanding the value in memory.
|
||||||
|
func TestDecodeManifestBoundsDecompressedOutput(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
bomb := &Manifest{
|
||||||
|
SnapshotID: testSnapshotID,
|
||||||
|
Timestamp: strings.Repeat("a", 1<<20),
|
||||||
|
}
|
||||||
|
|
||||||
|
compressed, err := EncodeManifest(bomb, 3)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Less(t, len(compressed), 4096,
|
||||||
|
"the compressible manifest must be small compressed")
|
||||||
|
|
||||||
|
_, err = decodeManifest(bytes.NewReader(compressed), 1<<20, 4096)
|
||||||
|
require.ErrorIs(t, err, blobgen.ErrOutputTooLarge)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDecodeManifestBoundsCompressedInput checks the compressed-input
|
||||||
|
// bound fires independently: a valid manifest with a generous decompressed
|
||||||
|
// bound but a tiny compressed bound still fails.
|
||||||
|
func TestDecodeManifestBoundsCompressedInput(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
manifest := &Manifest{
|
||||||
|
SnapshotID: testSnapshotID,
|
||||||
|
Timestamp: strings.Repeat("a", 4096),
|
||||||
|
}
|
||||||
|
|
||||||
|
compressed, err := EncodeManifest(manifest, 3)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
_, err = decodeManifest(bytes.NewReader(compressed), 8, manifestMaxDecompressed)
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
@@ -0,0 +1,64 @@
|
|||||||
|
//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,7 +63,9 @@ 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 bool // Skip file read errors (log loudly but continue)
|
// skipErrors skips files that cannot be opened or read (logged loudly);
|
||||||
|
// 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
|
||||||
|
|
||||||
@@ -121,7 +123,9 @@ 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 bool // Skip file read errors (log loudly but continue)
|
// SkipErrors skips files that cannot be opened or read (log loudly but
|
||||||
|
// 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
|
||||||
@@ -220,7 +224,14 @@ func (s *Scanner) Scan(
|
|||||||
defer s.progress.Stop()
|
defer s.progress.Stop()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Phase 0: Load known files and chunks from database into memory for fast lookup
|
// Phase 0: Repair any state left by an interrupted previous run, then
|
||||||
|
// 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
|
||||||
@@ -317,6 +328,38 @@ 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(
|
||||||
@@ -392,11 +435,14 @@ func (s *Scanner) loadKnownFiles(
|
|||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// loadKnownChunks loads all known chunk hashes from the database into a
|
// loadKnownChunks loads the chunk hashes safe to deduplicate against into
|
||||||
// map for fast lookup. This avoids per-chunk database queries during file
|
// an in-memory map for fast lookup, avoiding per-chunk database queries
|
||||||
// processing.
|
// during file processing. Only chunks held by a blob whose upload
|
||||||
|
// 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.List(ctx)
|
chunks, err := s.repos.Chunks.ListInUploadedBlobs(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("listing chunks: %w", err)
|
return fmt.Errorf("listing chunks: %w", err)
|
||||||
}
|
}
|
||||||
@@ -1294,6 +1340,15 @@ 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",
|
||||||
@@ -1303,7 +1358,7 @@ func (s *Scanner) processFileWithErrorHandling(
|
|||||||
|
|
||||||
return true, nil
|
return true, nil
|
||||||
}
|
}
|
||||||
// Skip file read errors if --skip-errors is enabled
|
// Skip open/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)
|
||||||
@@ -1401,7 +1456,17 @@ 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))
|
||||||
})
|
})
|
||||||
@@ -1660,6 +1725,20 @@ 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,
|
||||||
@@ -1710,7 +1789,11 @@ func (s *Scanner) processFileStreaming(
|
|||||||
if !chunkExists {
|
if !chunkExists {
|
||||||
err := s.addChunkToPacker(ctx, chunk)
|
err := s.addChunkToPacker(ctx, chunk)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
// Mark as a packer error so --skip-errors cannot swallow it:
|
||||||
|
// 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}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,216 @@
|
|||||||
|
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))
|
||||||
|
}
|
||||||
|
}
|
||||||
+32
-113
@@ -44,6 +44,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -294,68 +295,6 @@ 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.
|
||||||
@@ -758,12 +697,18 @@ func (sm *SnapshotManager) compressFile(inputPath, outputPath string) error {
|
|||||||
|
|
||||||
writerClosed = true
|
writerClosed = true
|
||||||
|
|
||||||
log.Debug("Compression complete", "hash", hex.EncodeToString(writer.Sum256()))
|
log.Debug("Compression complete", "hash", hex.EncodeToString(writer.ContentID()))
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// copyFile copies a file from src to dst
|
// exportCopyPerm restricts the exported snapshot database copy to the owning
|
||||||
|
// 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)
|
||||||
|
|
||||||
@@ -783,7 +728,9 @@ 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.Create(dst)
|
destFile, err := sm.fs.OpenFile(
|
||||||
|
dst, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, exportCopyPerm,
|
||||||
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -809,6 +756,11 @@ 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,
|
||||||
@@ -839,20 +791,26 @@ 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 {
|
||||||
log.Warn("Failed to get blob details", "hash", hash, "error", err)
|
return nil, fmt.Errorf("getting blob details for %s: %w", hash, err)
|
||||||
|
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if blob != nil {
|
if blob == nil {
|
||||||
blobs = append(blobs, BlobInfo{
|
return nil, fmt.Errorf("%w: blob %s, snapshot %s",
|
||||||
Hash: hash,
|
errBlobMissingFromDatabase, hash, snapshotID)
|
||||||
CompressedSize: blob.CompressedSize,
|
|
||||||
})
|
|
||||||
totalCompressedSize += blob.CompressedSize
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
blobs = append(blobs, BlobInfo{
|
||||||
|
Hash: hash,
|
||||||
|
CompressedSize: blob.CompressedSize,
|
||||||
|
})
|
||||||
|
totalCompressedSize += blob.CompressedSize
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create manifest. SnapshotID in the unencrypted manifest is the
|
// Create manifest. SnapshotID in the unencrypted manifest is the
|
||||||
@@ -913,45 +871,6 @@ 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,
|
||||||
|
|||||||
@@ -0,0 +1,198 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -0,0 +1,209 @@
|
|||||||
|
// 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
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
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,10 +111,11 @@ 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
|
// Ensure protocol is present. Absent an explicit use_ssl, default to TLS;
|
||||||
|
// 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 {
|
if cfg.S3.UseSSL == nil || *cfg.S3.UseSSL {
|
||||||
endpoint = "https://" + endpoint
|
endpoint = "https://" + endpoint
|
||||||
} else {
|
} else {
|
||||||
endpoint = "http://" + endpoint
|
endpoint = "http://" + endpoint
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+36
-14
@@ -13,18 +13,23 @@ import (
|
|||||||
"sneak.berlin/go/vaultik/internal/storage"
|
"sneak.berlin/go/vaultik/internal/storage"
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestS3StorerMissingKeyMapsToErrNotFound verifies that the s3 backend reports
|
// s3TestBucket is the bucket created for each in-process S3 server.
|
||||||
// a missing object as storage.ErrNotFound, matching the file and rclone
|
const s3TestBucket = "test-bucket"
|
||||||
// backends and the Storer contract. Without the mapping, Get and Stat leak the
|
|
||||||
// raw SDK error and errors.Is(err, storage.ErrNotFound) is false.
|
// newS3Storer builds an s3:// backend backed by a fresh in-process
|
||||||
|
// 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:paralleltest // shares an in-process S3 server via t.Cleanup
|
//nolint:ireturn // conformance runs against the Storer interface by design
|
||||||
func TestS3StorerMissingKeyMapsToErrNotFound(t *testing.T) {
|
func newS3Storer(t *testing.T) storage.Storer {
|
||||||
const bucket = "test-bucket"
|
t.Helper()
|
||||||
|
|
||||||
backend := s3mem.New()
|
backend := s3mem.New()
|
||||||
|
|
||||||
err := backend.CreateBucket(bucket)
|
err := backend.CreateBucket(s3TestBucket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("create bucket: %v", err)
|
t.Fatalf("create bucket: %v", err)
|
||||||
}
|
}
|
||||||
@@ -32,11 +37,9 @@ func TestS3StorerMissingKeyMapsToErrNotFound(t *testing.T) {
|
|||||||
srv := httptest.NewServer(gofakes3.New(backend).Server())
|
srv := httptest.NewServer(gofakes3.New(backend).Server())
|
||||||
t.Cleanup(srv.Close)
|
t.Cleanup(srv.Close)
|
||||||
|
|
||||||
ctx := context.Background()
|
client, err := s3.NewClient(context.Background(), s3.Config{
|
||||||
|
|
||||||
client, err := s3.NewClient(ctx, s3.Config{
|
|
||||||
Endpoint: srv.URL,
|
Endpoint: srv.URL,
|
||||||
Bucket: bucket,
|
Bucket: s3TestBucket,
|
||||||
AccessKeyID: "test",
|
AccessKeyID: "test",
|
||||||
SecretAccessKey: "test",
|
SecretAccessKey: "test",
|
||||||
Region: "us-east-1",
|
Region: "us-east-1",
|
||||||
@@ -45,9 +48,28 @@ func TestS3StorerMissingKeyMapsToErrNotFound(t *testing.T) {
|
|||||||
t.Fatalf("new client: %v", err)
|
t.Fatalf("new client: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
storer := storage.NewS3Storer(client)
|
return storage.NewS3Storer(client)
|
||||||
|
}
|
||||||
|
|
||||||
_, err = storer.Get(ctx, "does-not-exist")
|
// TestS3Storer runs the shared Storer contract against the s3:// backend,
|
||||||
|
// 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)
|
||||||
}
|
}
|
||||||
|
|||||||
+101
-46
@@ -4,6 +4,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -23,6 +24,10 @@ 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.
|
||||||
@@ -59,61 +64,111 @@ func ParseStorageURL(rawURL string) (*URL, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handle s3:// URLs
|
|
||||||
if strings.HasPrefix(rawURL, "s3://") {
|
if strings.HasPrefix(rawURL, "s3://") {
|
||||||
u, err := url.Parse(rawURL)
|
return parseS3URL(rawURL)
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("invalid URL: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
bucket := u.Host
|
|
||||||
if bucket == "" {
|
|
||||||
return nil, ErrMissingBucket
|
|
||||||
}
|
|
||||||
|
|
||||||
prefix := strings.TrimPrefix(u.Path, "/")
|
|
||||||
|
|
||||||
query := u.Query()
|
|
||||||
|
|
||||||
useSSL := true
|
|
||||||
if query.Get("ssl") == "false" {
|
|
||||||
useSSL = false
|
|
||||||
}
|
|
||||||
|
|
||||||
return &URL{
|
|
||||||
Scheme: schemeS3,
|
|
||||||
Bucket: bucket,
|
|
||||||
Prefix: prefix,
|
|
||||||
Endpoint: query.Get("endpoint"),
|
|
||||||
Region: query.Get("region"),
|
|
||||||
UseSSL: useSSL,
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handle rclone:// URLs
|
|
||||||
if strings.HasPrefix(rawURL, "rclone://") {
|
if strings.HasPrefix(rawURL, "rclone://") {
|
||||||
u, err := url.Parse(rawURL)
|
return parseRcloneURL(rawURL)
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("invalid URL: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
remote := u.Host
|
|
||||||
if remote == "" {
|
|
||||||
return nil, ErrMissingRemote
|
|
||||||
}
|
|
||||||
|
|
||||||
path := strings.TrimPrefix(u.Path, "/")
|
|
||||||
|
|
||||||
return &URL{
|
|
||||||
Scheme: schemeRclone,
|
|
||||||
Prefix: path,
|
|
||||||
RcloneRemote: remote,
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, ErrUnsupportedScheme
|
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)
|
||||||
|
if err != nil {
|
||||||
|
return nil, wrapParseError(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if u.User != nil {
|
||||||
|
return nil, ErrURLCredentials
|
||||||
|
}
|
||||||
|
|
||||||
|
bucket := u.Host
|
||||||
|
if bucket == "" {
|
||||||
|
return nil, ErrMissingBucket
|
||||||
|
}
|
||||||
|
|
||||||
|
query := u.Query()
|
||||||
|
|
||||||
|
err = rejectUnknownParams(query, "endpoint", "region", "ssl")
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &URL{
|
||||||
|
Scheme: schemeS3,
|
||||||
|
Bucket: bucket,
|
||||||
|
Prefix: strings.TrimPrefix(u.Path, "/"),
|
||||||
|
Endpoint: query.Get("endpoint"),
|
||||||
|
Region: query.Get("region"),
|
||||||
|
UseSSL: query.Get("ssl") != "false",
|
||||||
|
}, 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 {
|
||||||
|
return nil, ErrURLCredentials
|
||||||
|
}
|
||||||
|
|
||||||
|
remote := u.Host
|
||||||
|
if remote == "" {
|
||||||
|
return nil, ErrMissingRemote
|
||||||
|
}
|
||||||
|
|
||||||
|
err = rejectUnknownParams(u.Query())
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &URL{
|
||||||
|
Scheme: schemeRclone,
|
||||||
|
Prefix: strings.TrimPrefix(u.Path, "/"),
|
||||||
|
RcloneRemote: remote,
|
||||||
|
}, 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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.
|
||||||
func (u *URL) String() string {
|
func (u *URL) String() string {
|
||||||
switch u.Scheme {
|
switch u.Scheme {
|
||||||
|
|||||||
@@ -0,0 +1,208 @@
|
|||||||
|
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())
|
||||||
|
}
|
||||||
|
}
|
||||||
+14
-59
@@ -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, paths, and
|
// vaultik codebase. Using distinct types for IDs, hashes, and paths prevents
|
||||||
// credentials prevents accidental mixing of semantically different values
|
// accidental mixing of semantically different values that happen to share the
|
||||||
// that happen to share the same underlying type.
|
// 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,34 +157,6 @@ 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
|
||||||
|
|
||||||
@@ -199,31 +171,14 @@ type GlobPattern string
|
|||||||
|
|
||||||
// String methods for Stringer interface
|
// String methods for Stringer interface
|
||||||
|
|
||||||
func (id FileID) String() string { return uuid.UUID(id).String() }
|
func (id FileID) String() string { return uuid.UUID(id).String() }
|
||||||
func (id BlobID) String() string { return uuid.UUID(id).String() }
|
func (id BlobID) String() string { return uuid.UUID(id).String() }
|
||||||
func (id SnapshotID) String() string { return string(id) }
|
func (id SnapshotID) String() string { return string(id) }
|
||||||
func (h ChunkHash) String() string { return string(h) }
|
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 (h Hostname) String() string { return string(h) }
|
||||||
func (e S3Endpoint) String() string { return string(e) }
|
func (v Version) String() string { return string(v) }
|
||||||
func (b BucketName) String() string { return string(b) }
|
func (r GitRevision) String() string { return string(r) }
|
||||||
func (p S3Prefix) String() string { return string(p) }
|
func (p GlobPattern) 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 (v Version) String() string { return string(v) }
|
|
||||||
func (r GitRevision) String() string { return string(r) }
|
|
||||||
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) }
|
|
||||||
|
|||||||
@@ -0,0 +1,136 @@
|
|||||||
|
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())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+22
-3
@@ -23,7 +23,9 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
"unicode"
|
||||||
|
|
||||||
"github.com/dustin/go-humanize"
|
"github.com/dustin/go-humanize"
|
||||||
"golang.org/x/term"
|
"golang.org/x/term"
|
||||||
@@ -225,17 +227,17 @@ func (w *Writer) Hex(s string) string {
|
|||||||
short = s[:hexAbbrevLen] + "..."
|
short = s[:hexAbbrevLen] + "..."
|
||||||
}
|
}
|
||||||
|
|
||||||
return w.paint(ansiCyan, short)
|
return w.paint(ansiCyan, sanitize(short))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Snapshot colorizes a snapshot ID (full, no abbreviation).
|
// Snapshot colorizes a snapshot ID (full, no abbreviation).
|
||||||
func (w *Writer) Snapshot(id string) string {
|
func (w *Writer) Snapshot(id string) string {
|
||||||
return w.paint(ansiCyan+ansiBold, id)
|
return w.paint(ansiCyan+ansiBold, sanitize(id))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Path colorizes a filesystem path.
|
// Path colorizes a filesystem path.
|
||||||
func (w *Writer) Path(p string) string {
|
func (w *Writer) Path(p string) string {
|
||||||
return w.paint(ansiBlue, p)
|
return w.paint(ansiBlue, sanitize(p))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Size colorizes a byte count using humanize.Bytes.
|
// Size colorizes a byte count using humanize.Bytes.
|
||||||
@@ -310,6 +312,23 @@ func (w *Writer) paint(color, s string) string {
|
|||||||
return color + s + ansiReset
|
return color + s + ansiReset
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// sanitize returns s unchanged when every rune in it is printable, and a
|
||||||
|
// double-quoted, backslash-escaped form (\n, \x1b, …) otherwise. The
|
||||||
|
// string value formatters escape their argument through this before
|
||||||
|
// painting: identifiers, paths and symlink targets they render come from
|
||||||
|
// the snapshot database, which is not trusted, and escaping must happen
|
||||||
|
// before colour is applied — the painted result already contains the
|
||||||
|
// escape codes the raw text would otherwise be indistinguishable from.
|
||||||
|
func sanitize(s string) string {
|
||||||
|
for _, r := range s {
|
||||||
|
if !unicode.IsPrint(r) {
|
||||||
|
return strconv.Quote(s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
// emit writes "<prefix> <body>\n" with the prefix painted in prefixColor
|
// emit writes "<prefix> <body>\n" with the prefix painted in prefixColor
|
||||||
// and the body optionally painted in bodyColor (empty = no body color).
|
// and the body optionally painted in bodyColor (empty = no body color).
|
||||||
func (w *Writer) emit(prefixColor, prefix, bodyColor, format string, args []any) {
|
func (w *Writer) emit(prefixColor, prefix, bodyColor, format string, args []any) {
|
||||||
|
|||||||
@@ -0,0 +1,32 @@
|
|||||||
|
package ui_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestValueFormattersEscapeControlCharacters checks that a path carrying an
|
||||||
|
// ESC and a newline — the shape a symlink target read back from the
|
||||||
|
// snapshot database could take — is escaped before it reaches the output.
|
||||||
|
// Colour is off here, so the only way a control byte could appear is from
|
||||||
|
// the value itself.
|
||||||
|
func TestValueFormattersEscapeControlCharacters(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
w, buf := newTestWriter(false)
|
||||||
|
w.Infof("restoring %s", w.Path("a\x1b[31mZAP\nb"))
|
||||||
|
|
||||||
|
out := buf.String()
|
||||||
|
|
||||||
|
if strings.ContainsRune(out, '\x1b') {
|
||||||
|
t.Fatalf("raw ESC from a value survived in output: %q", out)
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Count(out, "\n") != 1 {
|
||||||
|
t.Fatalf("a newline in a value must be escaped, not emitted raw: %q", out)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !strings.Contains(out, `\x1b`) {
|
||||||
|
t.Fatalf("expected the escaped form of ESC in output: %q", out)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -2,35 +2,40 @@ package vaultik
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/sha256"
|
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"time"
|
|
||||||
|
|
||||||
"filippo.io/age"
|
"filippo.io/age"
|
||||||
"sneak.berlin/go/vaultik/internal/blobgen"
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
"sneak.berlin/go/vaultik/internal/log"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// errBlobHashMismatch is returned when a fetched blob's content hash does
|
// errBlobHashMismatch is returned when a fetched blob's content hash does
|
||||||
// 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
|
||||||
// redundant SHA-256 computation.
|
// redundant SHA-256 computation.
|
||||||
type hashVerifyReader struct {
|
type hashVerifyReader struct {
|
||||||
reader *blobgen.Reader // underlying decrypted blob reader (has internal hasher)
|
reader *blobgen.Reader // underlying decrypted blob reader (has internal hasher)
|
||||||
|
limited io.Reader // reader bounded to the blob's recorded plaintext size
|
||||||
fetcher io.ReadCloser // raw fetched stream (closed on Close)
|
fetcher io.ReadCloser // raw fetched stream (closed on Close)
|
||||||
blobHash string // expected double-SHA-256 hex
|
blobHash string // expected double-SHA-256 hex
|
||||||
done bool // EOF reached
|
done bool // EOF reached
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *hashVerifyReader) Read(p []byte) (int, error) {
|
func (h *hashVerifyReader) Read(p []byte) (int, error) {
|
||||||
n, err := h.reader.Read(p)
|
n, err := h.limited.Read(p)
|
||||||
if errors.Is(err, io.EOF) {
|
if errors.Is(err, io.EOF) {
|
||||||
h.done = true
|
h.done = true
|
||||||
}
|
}
|
||||||
@@ -38,21 +43,22 @@ func (h *hashVerifyReader) Read(p []byte) (int, error) {
|
|||||||
return n, err
|
return n, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close verifies the hash (if the stream was fully read) and closes underlying readers.
|
// Close closes the underlying readers and verifies the blob hash. The
|
||||||
|
// 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 {
|
||||||
firstHash := h.reader.Sum256()
|
return errBlobNotFullyRead
|
||||||
secondHasher := sha256.New()
|
}
|
||||||
secondHasher.Write(firstHash)
|
|
||||||
|
|
||||||
actualHashHex := hex.EncodeToString(secondHasher.Sum(nil))
|
actualHashHex := hex.EncodeToString(blobgen.DoubleSHA256(h.reader.Sum256()))
|
||||||
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, shortHash(h.blobHash), shortHash(actualHashHex))
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if readerErr != nil {
|
if readerErr != nil {
|
||||||
@@ -66,15 +72,22 @@ func (h *hashVerifyReader) Close() error {
|
|||||||
// returns a streaming reader that computes the double-SHA-256 hash on the fly.
|
// returns a streaming reader that computes the double-SHA-256 hash on the fly.
|
||||||
// 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.
|
||||||
|
//
|
||||||
|
// maxPlaintextSize is the blob's uncompressed_size as recorded in the
|
||||||
|
// snapshot database. Decompression stops with blobgen.ErrOutputTooLarge
|
||||||
|
// once the plaintext exceeds it, so a tampered blob cannot expand without
|
||||||
|
// limit — using the recorded size, not the restoring host's
|
||||||
|
// blob_size_limit, since that config may differ from the backup host's.
|
||||||
func (v *Vaultik) FetchAndDecryptBlob(
|
func (v *Vaultik) FetchAndDecryptBlob(
|
||||||
ctx context.Context, blobHash string, expectedSize int64, identity age.Identity,
|
ctx context.Context, blobHash string, maxPlaintextSize int64,
|
||||||
|
identities ...age.Identity,
|
||||||
) (io.ReadCloser, error) {
|
) (io.ReadCloser, error) {
|
||||||
rc, _, err := v.FetchBlob(ctx, blobHash, expectedSize)
|
rc, err := v.FetchBlob(ctx, blobHash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
reader, err := blobgen.NewReader(rc, identity)
|
reader, err := blobgen.NewReader(rc, identities...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = rc.Close()
|
_ = rc.Close()
|
||||||
|
|
||||||
@@ -83,45 +96,29 @@ func (v *Vaultik) FetchAndDecryptBlob(
|
|||||||
|
|
||||||
return &hashVerifyReader{
|
return &hashVerifyReader{
|
||||||
reader: reader,
|
reader: reader,
|
||||||
|
limited: blobgen.LimitReader(reader, maxPlaintextSize),
|
||||||
fetcher: rc,
|
fetcher: rc,
|
||||||
blobHash: blobHash,
|
blobHash: blobHash,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// FetchBlob downloads a blob and returns a reader for the encrypted data.
|
// FetchBlob downloads a blob and returns a reader for the encrypted data.
|
||||||
// Times the Storage.Get and Storage.Stat round-trips separately at
|
|
||||||
// debug level so we can see whether the size-only Stat (which is an
|
|
||||||
// extra request on every fetch) is hurting throughput.
|
|
||||||
func (v *Vaultik) FetchBlob(
|
func (v *Vaultik) FetchBlob(
|
||||||
ctx context.Context, blobHash string, expectedSize int64,
|
ctx context.Context, blobHash string,
|
||||||
) (io.ReadCloser, int64, error) {
|
) (io.ReadCloser, error) {
|
||||||
|
// blobHash reaches here from the snapshot database, which is not
|
||||||
|
// trusted. Reject a malformed hash before it is spliced into a storage
|
||||||
|
// path (blobHash[:2]/blobHash[2:4]) or a fetch is attempted.
|
||||||
|
if !isBlobHash(blobHash) {
|
||||||
|
return nil, fmt.Errorf("%w: %s", errInvalidBlobHash, shortHash(blobHash))
|
||||||
|
}
|
||||||
|
|
||||||
blobPath := fmt.Sprintf("blobs/%s/%s/%s", blobHash[:2], blobHash[2:4], blobHash)
|
blobPath := fmt.Sprintf("blobs/%s/%s/%s", blobHash[:2], blobHash[2:4], blobHash)
|
||||||
|
|
||||||
t0 := time.Now()
|
|
||||||
rc, err := v.Storage.Get(ctx, blobPath)
|
rc, err := v.Storage.Get(ctx, blobPath)
|
||||||
getDur := time.Since(t0)
|
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, 0, fmt.Errorf("downloading blob %s: %w", blobHash[:16], err)
|
return nil, fmt.Errorf("downloading blob %s: %w", shortHash(blobHash), err)
|
||||||
}
|
}
|
||||||
|
|
||||||
t0 = time.Now()
|
return rc, nil
|
||||||
info, err := v.Storage.Stat(ctx, blobPath)
|
|
||||||
statDur := time.Since(t0)
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
_ = rc.Close()
|
|
||||||
|
|
||||||
return nil, 0, fmt.Errorf("stat blob %s: %w", blobHash[:16], err)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Debug("FetchBlob round-trips",
|
|
||||||
"hash", blobHash[:16],
|
|
||||||
"ms_storage_get", getDur.Milliseconds(),
|
|
||||||
"ms_storage_stat", statDur.Milliseconds(),
|
|
||||||
"expected_size", expectedSize,
|
|
||||||
"stat_size", info.Size,
|
|
||||||
)
|
|
||||||
|
|
||||||
return rc, info.Size, nil
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,50 @@
|
|||||||
|
package vaultik_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"filippo.io/age"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
|
"sneak.berlin/go/vaultik/internal/vaultik"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestFetchAndDecryptBlobBoundsPlaintext feeds a small, highly
|
||||||
|
// compressible blob (256 KiB of zeros) whose decompressed size far exceeds
|
||||||
|
// the plaintext bound passed to FetchAndDecryptBlob. Decompression must
|
||||||
|
// stop with blobgen.ErrOutputTooLarge within the bound rather than
|
||||||
|
// expanding the whole blob into the restore cache.
|
||||||
|
func TestFetchAndDecryptBlobBoundsPlaintext(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
identity, err := age.GenerateX25519Identity()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
plaintext := make([]byte, 256*1024)
|
||||||
|
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)
|
||||||
|
|
||||||
|
const maxPlaintext = 1024
|
||||||
|
|
||||||
|
rc, err := tv.FetchAndDecryptBlob(
|
||||||
|
context.Background(), correctHash, maxPlaintext, identity)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
n, copyErr := io.Copy(io.Discard, rc)
|
||||||
|
_ = rc.Close()
|
||||||
|
|
||||||
|
require.ErrorIs(t, copyErr, blobgen.ErrOutputTooLarge)
|
||||||
|
require.LessOrEqual(t, n, int64(maxPlaintext)+1,
|
||||||
|
"decompression must stop within the recorded plaintext bound")
|
||||||
|
}
|
||||||
@@ -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.Sum256).
|
// blobgen.Writer.ContentID).
|
||||||
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.Sum256())
|
writerHash := hex.EncodeToString(writer.ContentID())
|
||||||
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)
|
||||||
@@ -55,6 +55,35 @@ func buildHashTestBlob(
|
|||||||
return encBuf.Bytes(), correctHash
|
return encBuf.Bytes(), correctHash
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestFetchBlobRejectsMalformedHash verifies FetchBlob refuses a blob hash
|
||||||
|
// that is not 64 lowercase hex characters before it builds a storage path
|
||||||
|
// or issues any request. The hash reaches FetchBlob from the snapshot
|
||||||
|
// database, which is not trusted, so a value such as one containing "/.."
|
||||||
|
// must never reach the store.
|
||||||
|
func TestFetchBlobRejectsMalformedHash(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
mockStorage := NewMockStorer()
|
||||||
|
tv := vaultik.NewForTesting(mockStorage)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
for _, bad := range []string{
|
||||||
|
"aa/../../../home/u/.profile",
|
||||||
|
"abc",
|
||||||
|
strings.Repeat("A", 64), // uppercase hex is not accepted
|
||||||
|
strings.Repeat("g", 64), // not hex
|
||||||
|
} {
|
||||||
|
_, err := tv.FetchBlob(ctx, bad)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatalf("expected error for malformed hash %q, got nil", bad)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if calls := mockStorage.GetCalls(); len(calls) != 0 {
|
||||||
|
t.Fatalf("storage was accessed for a malformed hash: %v", calls)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestFetchAndDecryptBlobVerifiesHash verifies that FetchAndDecryptBlob checks
|
// TestFetchAndDecryptBlobVerifiesHash verifies that FetchAndDecryptBlob checks
|
||||||
// the double-SHA-256 hash of the decrypted plaintext against the expected blob hash.
|
// the double-SHA-256 hash of the decrypted plaintext against the expected blob hash.
|
||||||
func TestFetchAndDecryptBlobVerifiesHash(t *testing.T) {
|
func TestFetchAndDecryptBlobVerifiesHash(t *testing.T) {
|
||||||
@@ -84,7 +113,7 @@ func TestFetchAndDecryptBlobVerifiesHash(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
rc, err := tv.FetchAndDecryptBlob(
|
rc, err := tv.FetchAndDecryptBlob(
|
||||||
ctx, correctHash, int64(len(encryptedData)), identity)
|
ctx, correctHash, int64(len(plaintext)), identity)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("expected success, got error: %v", err)
|
t.Fatalf("expected success, got error: %v", err)
|
||||||
}
|
}
|
||||||
@@ -116,7 +145,7 @@ func TestFetchAndDecryptBlobVerifiesHash(t *testing.T) {
|
|||||||
mockStorage.mu.Unlock()
|
mockStorage.mu.Unlock()
|
||||||
|
|
||||||
rc, err := tv.FetchAndDecryptBlob(
|
rc, err := tv.FetchAndDecryptBlob(
|
||||||
ctx, fakeHash, int64(len(encryptedData)), identity)
|
ctx, fakeHash, int64(len(plaintext)), identity)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error opening stream: %v", err)
|
t.Fatalf("unexpected error opening stream: %v", err)
|
||||||
}
|
}
|
||||||
@@ -133,3 +162,51 @@ 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(plaintext)), 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -6,13 +6,17 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Sentinel errors for blob cache lookups.
|
// Sentinel errors for blob cache lookups.
|
||||||
var (
|
var (
|
||||||
errCacheKeyMissing = errors.New("key not in cache")
|
errCacheKeyMissing = errors.New("key not in cache")
|
||||||
errCacheReadBeyondBlob = errors.New("read beyond blob size")
|
errCacheReadBeyondBlob = errors.New("read beyond blob size")
|
||||||
|
errCacheKeyHasSeparator = errors.New(
|
||||||
|
"cache key contains a path separator")
|
||||||
|
errCacheNegativeRead = errors.New("negative offset or length")
|
||||||
)
|
)
|
||||||
|
|
||||||
// blobCacheFileMode is the permission mode for cached blob files.
|
// blobCacheFileMode is the permission mode for cached blob files.
|
||||||
@@ -74,6 +78,11 @@ func newBlobDiskCache(maxBytes int64) (*blobDiskCache, error) {
|
|||||||
// Put writes blob data to disk cache. Entries larger than maxBytes are
|
// Put writes blob data to disk cache. Entries larger than maxBytes are
|
||||||
// silently skipped.
|
// silently skipped.
|
||||||
func (c *blobDiskCache) Put(key string, data []byte) error {
|
func (c *blobDiskCache) Put(key string, data []byte) error {
|
||||||
|
p, err := c.path(key)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
entrySize := int64(len(data))
|
entrySize := int64(len(data))
|
||||||
|
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
@@ -87,11 +96,12 @@ func (c *blobDiskCache) Put(key string, data []byte) error {
|
|||||||
if e, ok := c.items[key]; ok {
|
if e, ok := c.items[key]; ok {
|
||||||
c.unlink(e)
|
c.unlink(e)
|
||||||
c.curBytes -= e.size
|
c.curBytes -= e.size
|
||||||
_ = os.Remove(c.path(key))
|
_ = os.Remove(p)
|
||||||
|
|
||||||
delete(c.items, key)
|
delete(c.items, key)
|
||||||
}
|
}
|
||||||
|
|
||||||
err := os.WriteFile(c.path(key), data, blobCacheFileMode)
|
err = os.WriteFile(p, data, blobCacheFileMode)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("writing blob to cache: %w", err)
|
return fmt.Errorf("writing blob to cache: %w", err)
|
||||||
}
|
}
|
||||||
@@ -119,19 +129,26 @@ func (c *blobDiskCache) Put(key string, data []byte) error {
|
|||||||
// disk without buffering its entire plaintext (which may be tens of GB)
|
// disk without buffering its entire plaintext (which may be tens of GB)
|
||||||
// in RAM.
|
// in RAM.
|
||||||
func (c *blobDiskCache) PutFromReader(key string, r io.Reader) (int64, error) {
|
func (c *blobDiskCache) PutFromReader(key string, r io.Reader) (int64, error) {
|
||||||
|
p, err := c.path(key)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
// Remove any prior entry first; we'll re-link after the file is
|
// Remove any prior entry first; we'll re-link after the file is
|
||||||
// written successfully.
|
// written successfully.
|
||||||
if e, ok := c.items[key]; ok {
|
if e, ok := c.items[key]; ok {
|
||||||
c.unlink(e)
|
c.unlink(e)
|
||||||
c.curBytes -= e.size
|
c.curBytes -= e.size
|
||||||
_ = os.Remove(c.path(key))
|
_ = os.Remove(p)
|
||||||
|
|
||||||
delete(c.items, key)
|
delete(c.items, key)
|
||||||
}
|
}
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
|
//nolint:gosec // G304: path() rejects keys with a separator
|
||||||
f, err := os.OpenFile(
|
f, err := os.OpenFile(
|
||||||
c.path(key), os.O_CREATE|os.O_TRUNC|os.O_WRONLY, blobCacheFileMode)
|
p, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, blobCacheFileMode)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("creating cache file: %w", err)
|
return 0, fmt.Errorf("creating cache file: %w", err)
|
||||||
}
|
}
|
||||||
@@ -140,13 +157,13 @@ func (c *blobDiskCache) PutFromReader(key string, r io.Reader) (int64, error) {
|
|||||||
closeErr := f.Close()
|
closeErr := f.Close()
|
||||||
|
|
||||||
if copyErr != nil {
|
if copyErr != nil {
|
||||||
_ = os.Remove(c.path(key))
|
_ = os.Remove(p)
|
||||||
|
|
||||||
return written, fmt.Errorf("streaming to cache file: %w", copyErr)
|
return written, fmt.Errorf("streaming to cache file: %w", copyErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
if closeErr != nil {
|
if closeErr != nil {
|
||||||
_ = os.Remove(c.path(key))
|
_ = os.Remove(p)
|
||||||
|
|
||||||
return written, fmt.Errorf("closing cache file: %w", closeErr)
|
return written, fmt.Errorf("closing cache file: %w", closeErr)
|
||||||
}
|
}
|
||||||
@@ -158,7 +175,7 @@ func (c *blobDiskCache) PutFromReader(key string, r io.Reader) (int64, error) {
|
|||||||
// floor — but the restore path passes math.MaxInt64 as maxBytes
|
// floor — but the restore path passes math.MaxInt64 as maxBytes
|
||||||
// so this branch is effectively unreachable there.
|
// so this branch is effectively unreachable there.
|
||||||
if written > c.maxBytes {
|
if written > c.maxBytes {
|
||||||
_ = os.Remove(c.path(key))
|
_ = os.Remove(p)
|
||||||
|
|
||||||
return written, nil
|
return written, nil
|
||||||
}
|
}
|
||||||
@@ -181,6 +198,11 @@ func (c *blobDiskCache) PutFromReader(key string, r io.Reader) (int64, error) {
|
|||||||
|
|
||||||
// Get reads a cached blob from disk. Returns data and true on hit.
|
// Get reads a cached blob from disk. Returns data and true on hit.
|
||||||
func (c *blobDiskCache) Get(key string) ([]byte, bool) {
|
func (c *blobDiskCache) Get(key string) ([]byte, bool) {
|
||||||
|
p, err := c.path(key)
|
||||||
|
if err != nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
c.getCalls++
|
c.getCalls++
|
||||||
|
|
||||||
@@ -195,7 +217,8 @@ func (c *blobDiskCache) Get(key string) ([]byte, bool) {
|
|||||||
c.pushFront(e)
|
c.pushFront(e)
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
data, err := os.ReadFile(c.path(key))
|
//nolint:gosec // G304: path() rejects keys with a separator
|
||||||
|
data, err := os.ReadFile(p)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
if e2, ok2 := c.items[key]; ok2 && e2 == e {
|
if e2, ok2 := c.items[key]; ok2 && e2 == e {
|
||||||
@@ -213,6 +236,20 @@ func (c *blobDiskCache) Get(key string) ([]byte, bool) {
|
|||||||
|
|
||||||
// ReadAt reads a slice of a cached blob without loading the entire blob into memory.
|
// ReadAt reads a slice of a cached blob without loading the entire blob into memory.
|
||||||
func (c *blobDiskCache) ReadAt(key string, offset, length int64) ([]byte, error) {
|
func (c *blobDiskCache) ReadAt(key string, offset, length int64) ([]byte, error) {
|
||||||
|
p, err := c.path(key)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// offset and length come from a blob_chunks row read back from the
|
||||||
|
// destination. A negative value must be rejected outright; the upper
|
||||||
|
// bound is checked as length > size-offset (a subtraction) so a huge
|
||||||
|
// offset+length cannot overflow int64 and slip past the check.
|
||||||
|
if offset < 0 || length < 0 {
|
||||||
|
return nil, fmt.Errorf("%w: offset=%d length=%d",
|
||||||
|
errCacheNegativeRead, offset, length)
|
||||||
|
}
|
||||||
|
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
c.readAtCalls++
|
c.readAtCalls++
|
||||||
|
|
||||||
@@ -223,7 +260,7 @@ func (c *blobDiskCache) ReadAt(key string, offset, length int64) ([]byte, error)
|
|||||||
return nil, fmt.Errorf("%w: %q", errCacheKeyMissing, key)
|
return nil, fmt.Errorf("%w: %q", errCacheKeyMissing, key)
|
||||||
}
|
}
|
||||||
|
|
||||||
if offset+length > e.size {
|
if length > e.size-offset {
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
return nil, fmt.Errorf("%w: offset=%d length=%d size=%d",
|
return nil, fmt.Errorf("%w: offset=%d length=%d size=%d",
|
||||||
@@ -234,7 +271,7 @@ func (c *blobDiskCache) ReadAt(key string, offset, length int64) ([]byte, error)
|
|||||||
c.pushFront(e)
|
c.pushFront(e)
|
||||||
c.mu.Unlock()
|
c.mu.Unlock()
|
||||||
|
|
||||||
f, err := os.Open(c.path(key))
|
f, err := os.Open(p) //nolint:gosec // G304: path() rejects keys with a separator
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -276,7 +313,13 @@ func (c *blobDiskCache) Delete(key string) {
|
|||||||
c.unlink(e)
|
c.unlink(e)
|
||||||
delete(c.items, key)
|
delete(c.items, key)
|
||||||
c.curBytes -= e.size
|
c.curBytes -= e.size
|
||||||
_ = os.Remove(c.path(key))
|
|
||||||
|
// The key is already in the map, so it passed path() when it was
|
||||||
|
// inserted; the error cannot occur here.
|
||||||
|
p, err := c.path(key)
|
||||||
|
if err == nil {
|
||||||
|
_ = os.Remove(p)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Keys returns a snapshot of all cached keys. Safe for iteration without
|
// Keys returns a snapshot of all cached keys. Safe for iteration without
|
||||||
@@ -347,8 +390,18 @@ func (c *blobDiskCache) Close() error {
|
|||||||
return os.RemoveAll(c.dir)
|
return os.RemoveAll(c.dir)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *blobDiskCache) path(key string) string {
|
// path returns the on-disk location of the cache file for key. The key is
|
||||||
return filepath.Join(c.dir, key)
|
// a blob hash read back from the destination and is not trusted: a value
|
||||||
|
// such as "aa/../../../home/u/.profile" would otherwise make filepath.Join
|
||||||
|
// escape the cache directory, so a key containing a path separator is
|
||||||
|
// refused rather than joined.
|
||||||
|
func (c *blobDiskCache) path(key string) (string, error) {
|
||||||
|
if strings.ContainsRune(key, '/') ||
|
||||||
|
strings.ContainsRune(key, filepath.Separator) {
|
||||||
|
return "", fmt.Errorf("%w: %q", errCacheKeyHasSeparator, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
return filepath.Join(c.dir, key), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *blobDiskCache) unlink(e *blobDiskCacheEntry) {
|
func (c *blobDiskCache) unlink(e *blobDiskCacheEntry) {
|
||||||
@@ -391,5 +444,10 @@ func (c *blobDiskCache) evictLRU() {
|
|||||||
c.unlink(victim)
|
c.unlink(victim)
|
||||||
delete(c.items, victim.key)
|
delete(c.items, victim.key)
|
||||||
c.curBytes -= victim.size
|
c.curBytes -= victim.size
|
||||||
_ = os.Remove(c.path(victim.key))
|
|
||||||
|
// victim.key was validated by path() on insertion, so this cannot err.
|
||||||
|
p, err := c.path(victim.key)
|
||||||
|
if err == nil {
|
||||||
|
_ = os.Remove(p)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
package vaultik
|
||||||
|
|
||||||
|
import "errors"
|
||||||
|
|
||||||
|
// blobHashHexLen is the length of a blob hash written as lowercase hex: a
|
||||||
|
// SHA-256 digest is 32 bytes, so 64 characters. Remote snapshot keys are
|
||||||
|
// SHA-256 hashes too and share this exact form.
|
||||||
|
const blobHashHexLen = 64
|
||||||
|
|
||||||
|
// shortHashLen is how many leading characters of a hash appear in log and
|
||||||
|
// error text.
|
||||||
|
const shortHashLen = 16
|
||||||
|
|
||||||
|
// errInvalidBlobHash reports a value used as a blob hash that is not
|
||||||
|
// exactly 64 lowercase hex characters. Restore, verify and prune read
|
||||||
|
// these values back from the destination, which is not trusted, so each
|
||||||
|
// one is checked before it is used to build a path or drive a read.
|
||||||
|
var errInvalidBlobHash = errors.New(
|
||||||
|
"blob hash is not 64 lowercase hex characters")
|
||||||
|
|
||||||
|
// isBlobHash reports whether s is exactly 64 lowercase hex characters.
|
||||||
|
// Every real blob hash and remote snapshot key has this form.
|
||||||
|
//
|
||||||
|
// The check is a plain function, not a method on types.BlobHash: the
|
||||||
|
// packer stores "temp-placeholder-{uuid}" as the hash of an unfinished
|
||||||
|
// blob in the local index, so the type itself must keep accepting values
|
||||||
|
// that are not hashes.
|
||||||
|
func isBlobHash(s string) bool {
|
||||||
|
if len(s) != blobHashHexLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, r := range s {
|
||||||
|
if (r < '0' || r > '9') && (r < 'a' || r > 'f') {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// shortHash returns the leading part of a hash for log and error text. It
|
||||||
|
// never panics: a string shorter than the prefix is returned whole. A hash
|
||||||
|
// read from the destination may be malformed, and formatting one for a
|
||||||
|
// message must not crash the command.
|
||||||
|
func shortHash(s string) string {
|
||||||
|
if len(s) <= shortHashLen {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
return s[:shortHashLen]
|
||||||
|
}
|
||||||
@@ -0,0 +1,262 @@
|
|||||||
|
package vaultik //nolint:testpackage // drives unexported input validation
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"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/storage"
|
||||||
|
"sneak.berlin/go/vaultik/internal/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// These tests treat every hash, offset and length read back from the
|
||||||
|
// destination as hostile. A blob hash comes from the downloaded snapshot
|
||||||
|
// database or the store listing, neither of which is authenticated (see
|
||||||
|
// https://git.eeqj.de/sneak/vaultik/issues/155), so each is validated
|
||||||
|
// before it is used to build a path or size an allocation.
|
||||||
|
|
||||||
|
// TestBlobCacheRejectsKeyWithSeparator proves the arbitrary-file-write
|
||||||
|
// hole is closed: a blob hash that climbs out of the cache directory is
|
||||||
|
// refused and nothing is written outside it. This is the exact write the
|
||||||
|
// restore path performs, keyed by the hash from the snapshot database.
|
||||||
|
func TestBlobCacheRejectsKeyWithSeparator(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cache, err := newBlobDiskCache(1 << 20)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
defer func() { _ = cache.Close() }()
|
||||||
|
|
||||||
|
target := filepath.Join(t.TempDir(), "pwned")
|
||||||
|
|
||||||
|
// A hash whose relative form escapes the cache directory to target.
|
||||||
|
key, err := filepath.Rel(cache.dir, target)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Contains(t, key, "..")
|
||||||
|
|
||||||
|
err = cache.Put(key, []byte("secret"))
|
||||||
|
require.ErrorIs(t, err, errCacheKeyHasSeparator)
|
||||||
|
|
||||||
|
_, err = cache.PutFromReader(key, strings.NewReader("secret"))
|
||||||
|
require.ErrorIs(t, err, errCacheKeyHasSeparator)
|
||||||
|
|
||||||
|
_, statErr := os.Stat(target)
|
||||||
|
require.Truef(t, os.IsNotExist(statErr),
|
||||||
|
"cache wrote outside its directory at %s", target)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestBuildBlobIndexesRejectsHostileHash proves restore refuses a snapshot
|
||||||
|
// database whose blob_hash escapes the cache directory. buildBlobIndexes is
|
||||||
|
// the first place restore reads these hashes back, and it fails there, before
|
||||||
|
// any blob is fetched or written, so a hash containing /../ cannot steer a
|
||||||
|
// later write outside the cache directory.
|
||||||
|
func TestBuildBlobIndexesRejectsHostileHash(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
db, err := database.New(ctx, filepath.Join(t.TempDir(), "index.sqlite"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
// A blob cache and a target file just outside it. The hostile hash is
|
||||||
|
// the relative path from the cache to that target, so an unguarded
|
||||||
|
// restore keyed by this hash would write there.
|
||||||
|
cache, err := newBlobDiskCache(1 << 20)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
defer func() { _ = cache.Close() }()
|
||||||
|
|
||||||
|
target := filepath.Join(t.TempDir(), "pwned")
|
||||||
|
hostile, err := filepath.Rel(cache.dir, target)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Contains(t, hostile, "..")
|
||||||
|
|
||||||
|
repos := database.NewRepositories(db)
|
||||||
|
require.NoError(t, repos.Blobs.Create(ctx, nil, &database.Blob{
|
||||||
|
ID: types.NewBlobID(),
|
||||||
|
Hash: types.BlobHash(hostile),
|
||||||
|
CreatedTS: time.Now().UTC(),
|
||||||
|
}))
|
||||||
|
|
||||||
|
v := NewForTesting(nil)
|
||||||
|
v.SetContext(ctx)
|
||||||
|
|
||||||
|
_, _, err = v.buildBlobIndexes(repos)
|
||||||
|
require.ErrorIs(t, err, errInvalidBlobHash)
|
||||||
|
|
||||||
|
_, statErr := os.Stat(target)
|
||||||
|
require.Truef(t, os.IsNotExist(statErr),
|
||||||
|
"restore wrote outside the cache directory at %s", target)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestBlobCacheReadAtRejectsBadBounds proves a blob_chunks row cannot
|
||||||
|
// drive an out-of-range or negative read. offset/length reach ReadAt
|
||||||
|
// straight from the database.
|
||||||
|
func TestBlobCacheReadAtRejectsBadBounds(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cache, err := newBlobDiskCache(1 << 20)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
defer func() { _ = cache.Close() }()
|
||||||
|
|
||||||
|
require.NoError(t, cache.Put("blob", make([]byte, 100)))
|
||||||
|
|
||||||
|
_, err = cache.ReadAt("blob", -1, 10)
|
||||||
|
require.ErrorIs(t, err, errCacheNegativeRead)
|
||||||
|
|
||||||
|
_, err = cache.ReadAt("blob", 0, -1)
|
||||||
|
require.ErrorIs(t, err, errCacheNegativeRead)
|
||||||
|
|
||||||
|
// A length past the end is rejected via the subtraction bound, so a
|
||||||
|
// huge offset+length cannot overflow past the check.
|
||||||
|
_, err = cache.ReadAt("blob", 50, 60)
|
||||||
|
require.ErrorIs(t, err, errCacheReadBeyondBlob)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestListAllRemoteBlobsSkipsNonConformingName proves a bogus object name
|
||||||
|
// under blobs/ (here a three-character name) is skipped rather than
|
||||||
|
// entering the blob map, so prune's later hash[:2]/hash[2:4] path build
|
||||||
|
// cannot panic on it.
|
||||||
|
func TestListAllRemoteBlobsSkipsNonConformingName(t *testing.T) {
|
||||||
|
// Initialize the global logger before t.Parallel() so the write lands
|
||||||
|
// in the serial phase and cannot race other parallel tests reading it.
|
||||||
|
log.Initialize(log.Config{})
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
good := strings.Repeat("a", blobHashHexLen)
|
||||||
|
store := &stubLister{objects: []storage.ObjectInfo{
|
||||||
|
{Key: "blobs/" + good[:2] + "/" + good[2:4] + "/" + good, Size: 10},
|
||||||
|
{Key: "blobs/a/b/c", Size: 3},
|
||||||
|
}}
|
||||||
|
|
||||||
|
v := &Vaultik{Storage: store}
|
||||||
|
v.SetContext(context.Background())
|
||||||
|
|
||||||
|
blobs, err := v.listAllRemoteBlobs()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Contains(t, blobs, good)
|
||||||
|
require.NotContains(t, blobs, "c")
|
||||||
|
require.Len(t, blobs, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestVerifyManifestBlobsRejectsShortHash proves a manifest (which is not
|
||||||
|
// authenticated) with a short blob hash fails cleanly instead of panicking
|
||||||
|
// on blob.Hash[:2].
|
||||||
|
func TestVerifyManifestBlobsRejectsShortHash(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
v := &Vaultik{Stdout: io.Discard}
|
||||||
|
manifest := &snapshot.Manifest{
|
||||||
|
Blobs: []snapshot.BlobInfo{{Hash: "abc", CompressedSize: 1}},
|
||||||
|
}
|
||||||
|
|
||||||
|
verified, missing, mismatched, missingSize, err :=
|
||||||
|
v.verifyManifestBlobs(manifest, &VerifyOptions{JSON: true})
|
||||||
|
require.ErrorIs(t, err, errInvalidBlobHash)
|
||||||
|
require.Zero(t, verified)
|
||||||
|
require.Zero(t, missing)
|
||||||
|
require.Zero(t, mismatched)
|
||||||
|
require.Zero(t, missingSize)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestVerifyBlobChunksRejectsNegativeLength proves a blob_chunks row with a
|
||||||
|
// negative length returns an error rather than reaching make([]byte,
|
||||||
|
// length) or streaming an untrusted size.
|
||||||
|
func TestVerifyBlobChunksRejectsNegativeLength(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
db, err := database.New(ctx, filepath.Join(t.TempDir(), "index.sqlite"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
repos := database.NewRepositories(db)
|
||||||
|
blobHash := strings.Repeat("b", blobHashHexLen)
|
||||||
|
blob := &database.Blob{
|
||||||
|
ID: types.NewBlobID(),
|
||||||
|
Hash: types.BlobHash(blobHash),
|
||||||
|
CreatedTS: time.Now().UTC(),
|
||||||
|
}
|
||||||
|
require.NoError(t, repos.Blobs.Create(ctx, nil, blob))
|
||||||
|
|
||||||
|
chunkHash := strings.Repeat("c", blobHashHexLen)
|
||||||
|
require.NoError(t, repos.Chunks.Create(ctx, nil,
|
||||||
|
&database.Chunk{ChunkHash: types.ChunkHash(chunkHash), Size: 1024}))
|
||||||
|
require.NoError(t, repos.BlobChunks.Create(ctx, nil, &database.BlobChunk{
|
||||||
|
BlobID: blob.ID,
|
||||||
|
ChunkHash: types.ChunkHash(chunkHash),
|
||||||
|
Offset: 0,
|
||||||
|
Length: -1,
|
||||||
|
}))
|
||||||
|
|
||||||
|
v := NewForTesting(nil)
|
||||||
|
|
||||||
|
_, err = v.verifyBlobChunks(db.Conn(), blobHash, strings.NewReader(""))
|
||||||
|
require.ErrorIs(t, err, errNegativeChunkLength)
|
||||||
|
}
|
||||||
|
|
||||||
|
// errStubUnused marks a stubLister method a test never exercises.
|
||||||
|
var errStubUnused = errors.New("stubLister method not used in test")
|
||||||
|
|
||||||
|
// stubLister is a storage.Storer whose ListStream yields a fixed set of
|
||||||
|
// objects; every other method is unused by the tests here.
|
||||||
|
type stubLister struct {
|
||||||
|
objects []storage.ObjectInfo
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubLister) ListStream(
|
||||||
|
_ context.Context, prefix string,
|
||||||
|
) <-chan storage.ObjectInfo {
|
||||||
|
ch := make(chan storage.ObjectInfo, len(s.objects))
|
||||||
|
for _, o := range s.objects {
|
||||||
|
if strings.HasPrefix(o.Key, prefix) {
|
||||||
|
ch <- o
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
close(ch)
|
||||||
|
|
||||||
|
return ch
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubLister) Put(_ context.Context, _ string, _ io.Reader) error {
|
||||||
|
return errStubUnused
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubLister) PutWithProgress(
|
||||||
|
_ context.Context, _ string, _ io.Reader, _ int64, _ storage.ProgressCallback,
|
||||||
|
) error {
|
||||||
|
return errStubUnused
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubLister) Get(_ context.Context, _ string) (io.ReadCloser, error) {
|
||||||
|
return nil, errStubUnused
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubLister) Stat(_ context.Context, _ string) (*storage.ObjectInfo, error) {
|
||||||
|
return nil, errStubUnused
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubLister) Delete(_ context.Context, _ string) error {
|
||||||
|
return errStubUnused
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubLister) List(_ context.Context, _ string) ([]string, error) {
|
||||||
|
return nil, errStubUnused
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubLister) Info() storage.Info {
|
||||||
|
return storage.Info{}
|
||||||
|
}
|
||||||
@@ -0,0 +1,647 @@
|
|||||||
|
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,6 +167,12 @@ 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
|
||||||
@@ -179,27 +185,22 @@ 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 {
|
||||||
log.Error("Failed to download manifest", "remote_key", remoteKey, "error", err)
|
return nil, fmt.Errorf("reading manifest %s: %w", remoteKey, 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", manifestCount, "unique_blobs_referenced", len(allBlobsReferenced))
|
"count", len(remoteKeys), "unique_blobs_referenced", len(allBlobsReferenced))
|
||||||
|
|
||||||
return allBlobsReferenced, nil
|
return allBlobsReferenced, nil
|
||||||
}
|
}
|
||||||
@@ -246,9 +247,22 @@ func (v *Vaultik) listAllRemoteBlobs() (map[string]int64, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
parts := strings.Split(object.Key, "/")
|
parts := strings.Split(object.Key, "/")
|
||||||
if len(parts) == blobKeyParts && parts[0] == "blobs" {
|
if len(parts) != blobKeyParts || parts[0] != "blobs" {
|
||||||
allBlobs[parts[3]] = object.Size
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// The object name is read from the destination store and is not
|
||||||
|
// trusted. A name that is not a blob hash (e.g. a short or
|
||||||
|
// non-hex string) would panic the later hash[:2]/hash[2:4] path
|
||||||
|
// build, so skip it with a warning rather than delete it.
|
||||||
|
if !isBlobHash(parts[3]) {
|
||||||
|
log.Warn("Skipping non-conforming object under blobs/",
|
||||||
|
"key", object.Key)
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
allBlobs[parts[3]] = object.Size
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Info("Found blobs in storage", "count", len(allBlobs))
|
log.Info("Found blobs in storage", "count", len(allBlobs))
|
||||||
|
|||||||
@@ -0,0 +1,47 @@
|
|||||||
|
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")
|
||||||
|
}
|
||||||
@@ -0,0 +1,146 @@
|
|||||||
|
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,6 +12,7 @@ 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"
|
||||||
)
|
)
|
||||||
@@ -60,8 +61,11 @@ 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 remote metadata stub so syncWithRemote keeps it
|
// Create the remote metadata stub under the production layout so
|
||||||
metadataKey := "metadata/" + id + "/manifest.json.zst"
|
// syncWithRemote keeps the local row. Production stores metadata
|
||||||
|
// 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)
|
||||||
}
|
}
|
||||||
|
|||||||
+361
-95
@@ -1,7 +1,6 @@
|
|||||||
package vaultik
|
package vaultik
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"context"
|
"context"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
@@ -11,6 +10,7 @@ import (
|
|||||||
"math"
|
"math"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"filippo.io/age"
|
"filippo.io/age"
|
||||||
@@ -18,6 +18,7 @@ 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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -28,19 +29,47 @@ 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:\n" +
|
"age private key file:\n" +
|
||||||
" export VAULTIK_AGE_SECRET_KEY='AGE-SECRET-KEY-...'")
|
" export VAULTIK_AGE_SECRET_KEY=\"$(cat vaultik_backup_private_key.txt)\"")
|
||||||
errBlobMissingFromIndex = errors.New("blob hash missing from blob index")
|
// errInvalidAgeSecretKey is returned when the configured key does not
|
||||||
errChunkNotInAnyBlob = errors.New("chunk not found in any blob")
|
// parse as any age identity. It names the source but never the value,
|
||||||
errBlobIDNotInHashIndex = errors.New("blob id missing from hash index")
|
// which is secret, so the message is safe to print and log.
|
||||||
errShortChunkRead = errors.New("short read")
|
errInvalidAgeSecretKey = errors.New(
|
||||||
|
"configured age secret key holds no usable age identity")
|
||||||
|
errBlobMissingFromIndex = errors.New("blob hash missing from blob index")
|
||||||
|
errChunkNotInAnyBlob = errors.New("chunk not found in any blob")
|
||||||
|
errBlobIDNotInHashIndex = errors.New("blob id missing from hash index")
|
||||||
|
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")
|
||||||
|
// errEmptySnapshotDB is returned when the decrypted metadata database has
|
||||||
|
// zero length, which happens when the object was truncated or replaced
|
||||||
|
// with an empty payload. Rejected before any schema is built on it.
|
||||||
|
errEmptySnapshotDB = errors.New("decrypted snapshot database is empty")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// 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
|
||||||
@@ -76,7 +105,7 @@ type RestoreResult struct {
|
|||||||
func (v *Vaultik) Restore(opts *RestoreOptions) error {
|
func (v *Vaultik) Restore(opts *RestoreOptions) error {
|
||||||
startTime := time.Now()
|
startTime := time.Now()
|
||||||
|
|
||||||
identity, err := v.prepareRestoreIdentity()
|
identities, err := v.restoreIdentities()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -90,7 +119,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, err := v.downloadSnapshotDB(opts.SnapshotID, identity)
|
tempDB, tempDir, err := v.downloadSnapshotDB(opts.SnapshotID, identities)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("downloading snapshot database: %w", err)
|
return fmt.Errorf("downloading snapshot database: %w", err)
|
||||||
}
|
}
|
||||||
@@ -100,10 +129,11 @@ 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)
|
||||||
}
|
}
|
||||||
// Clean up temp file
|
// Remove the whole private directory, so the decrypted database
|
||||||
err = v.Fs.Remove(tempDB.Path())
|
// and any SQLite side files it produced are gone on every path.
|
||||||
|
err = v.Fs.RemoveAll(tempDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Debug("Failed to remove temp database", "error", err)
|
log.Debug("Failed to remove temp database directory", "error", err)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
@@ -138,7 +168,7 @@ func (v *Vaultik) Restore(opts *RestoreOptions) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Step 5: Restore files
|
// Step 5: Restore files
|
||||||
result, err := v.restoreAllFiles(files, repos, opts, identity, chunkToBlobMap)
|
result, err := v.restoreAllFiles(files, repos, opts, identities, chunkToBlobMap)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -199,21 +229,28 @@ func (v *Vaultik) finishRestore(
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// prepareRestoreIdentity validates that an age secret key is configured
|
// restoreIdentities parses the configured age secret key once into every
|
||||||
// and parses it.
|
// identity it contains. The value may be a single key line or a whole
|
||||||
//
|
// age-keygen file with several identities; all of them are returned so
|
||||||
//nolint:ireturn // age.Identity is the decryption abstraction by design
|
// blobgen (via age.Decrypt) can read a blob encrypted to any of their
|
||||||
func (v *Vaultik) prepareRestoreIdentity() (age.Identity, error) {
|
// recipients. This is the first step of both restore and deep verify, so
|
||||||
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
identity, err := age.ParseX25519Identity(v.Config.AgeSecretKey)
|
// age.ParseIdentities skips comment and blank lines and rejects a
|
||||||
|
// 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("parsing age secret key: %w", err)
|
return nil, fmt.Errorf("%w (source: %s)",
|
||||||
|
errInvalidAgeSecretKey, v.Config.AgeSecretKeySourceName())
|
||||||
}
|
}
|
||||||
|
|
||||||
return identity, nil
|
return identities, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// restoreAllFiles processes files in blob-locality order: drain every
|
// restoreAllFiles processes files in blob-locality order: drain every
|
||||||
@@ -226,7 +263,7 @@ func (v *Vaultik) restoreAllFiles(
|
|||||||
files []*database.File,
|
files []*database.File,
|
||||||
repos *database.Repositories,
|
repos *database.Repositories,
|
||||||
opts *RestoreOptions,
|
opts *RestoreOptions,
|
||||||
identity age.Identity,
|
identities []age.Identity,
|
||||||
chunkToBlobMap map[string]*database.BlobChunk,
|
chunkToBlobMap map[string]*database.BlobChunk,
|
||||||
) (*RestoreResult, error) {
|
) (*RestoreResult, error) {
|
||||||
result := &RestoreResult{}
|
result := &RestoreResult{}
|
||||||
@@ -280,7 +317,7 @@ func (v *Vaultik) restoreAllFiles(
|
|||||||
ctx: v.ctx,
|
ctx: v.ctx,
|
||||||
repos: repos,
|
repos: repos,
|
||||||
opts: opts,
|
opts: opts,
|
||||||
identity: identity,
|
identities: identities,
|
||||||
chunkToBlobMap: chunkToBlobMap,
|
chunkToBlobMap: chunkToBlobMap,
|
||||||
blobByHash: blobByHash,
|
blobByHash: blobByHash,
|
||||||
blobIDToHash: blobIDToHash,
|
blobIDToHash: blobIDToHash,
|
||||||
@@ -356,6 +393,13 @@ 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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -370,20 +414,26 @@ 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 := plan.pickNextDownload()
|
next, ok := plan.pickNextDownload()
|
||||||
if next.IsZero() {
|
if !ok {
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, hash := range plan.blobsNeeded(next) {
|
for _, hash := range plan.blobsNeeded(next) {
|
||||||
blob, ok := s.blobByHash[hash]
|
// Stop between blobs on cancel so an interrupt ends the download
|
||||||
if !ok {
|
// phase promptly rather than fetching the rest of the set.
|
||||||
return false, fmt.Errorf("%w: %s", errBlobMissingFromIndex, hash[:16])
|
if s.ctx.Err() != nil {
|
||||||
|
return false, s.ctx.Err()
|
||||||
}
|
}
|
||||||
|
|
||||||
err := s.downloadBlobToCache(hash, blob.CompressedSize)
|
blob, ok := s.blobByHash[hash]
|
||||||
|
if !ok {
|
||||||
|
return false, fmt.Errorf("%w: %s", errBlobMissingFromIndex, shortHash(hash))
|
||||||
|
}
|
||||||
|
|
||||||
|
err := s.downloadBlobToCache(hash, blob.CompressedSize, blob.UncompressedSize)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, fmt.Errorf("downloading blob %s: %w", hash[:16], err)
|
return false, fmt.Errorf("downloading blob %s: %w", shortHash(hash), err)
|
||||||
}
|
}
|
||||||
|
|
||||||
s.result.BlobsDownloaded++
|
s.result.BlobsDownloaded++
|
||||||
@@ -427,6 +477,16 @@ func (v *Vaultik) buildBlobIndexes(
|
|||||||
blobByHash := make(map[string]*database.Blob, len(blobsByID))
|
blobByHash := make(map[string]*database.Blob, len(blobsByID))
|
||||||
for id, blob := range blobsByID {
|
for id, blob := range blobsByID {
|
||||||
hash := blob.Hash.String()
|
hash := blob.Hash.String()
|
||||||
|
|
||||||
|
// The snapshot database is untrusted. A hash that is not 64
|
||||||
|
// lowercase hex characters could steer a later fetch to a path
|
||||||
|
// outside the cache directory, so reject it here, before any
|
||||||
|
// blob is downloaded.
|
||||||
|
if !isBlobHash(hash) {
|
||||||
|
return nil, nil, fmt.Errorf(
|
||||||
|
"%w: %s", errInvalidBlobHash, shortHash(hash))
|
||||||
|
}
|
||||||
|
|
||||||
blobIDToHash[id] = hash
|
blobIDToHash[id] = hash
|
||||||
blobByHash[hash] = blob
|
blobByHash[hash] = blob
|
||||||
}
|
}
|
||||||
@@ -581,11 +641,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, identity age.Identity,
|
snapshotID string, identities []age.Identity,
|
||||||
) (*database.DB, error) {
|
) (*database.DB, string, 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
|
||||||
@@ -593,69 +653,128 @@ 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() }()
|
||||||
|
|
||||||
// Read all data
|
// Decrypt and decompress straight from the storage stream, then stream
|
||||||
encryptedData, err := io.ReadAll(reader)
|
// the plaintext to a temp file. Neither the encrypted bytes nor the
|
||||||
|
// decrypted database is ever held whole in memory; a snapshot database
|
||||||
|
// can be large.
|
||||||
|
blobReader, err := blobgen.NewReader(reader, identities...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("reading encrypted data: %w", err)
|
return nil, "", fmt.Errorf("creating decryption reader: %w", err)
|
||||||
}
|
|
||||||
|
|
||||||
log.Debug("Downloaded encrypted database",
|
|
||||||
"size", ubytes(int64(len(encryptedData))))
|
|
||||||
|
|
||||||
// Decrypt and decompress using blobgen.Reader
|
|
||||||
blobReader, err := blobgen.NewReader(bytes.NewReader(encryptedData), identity)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("creating decryption reader: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
defer func() { _ = blobReader.Close() }()
|
defer func() { _ = blobReader.Close() }()
|
||||||
|
|
||||||
// Read the binary SQLite database
|
db, tempDir, err := v.materializeSnapshotDB(blobReader)
|
||||||
dbData, err := io.ReadAll(blobReader)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("decrypting and decompressing: %w", err)
|
return nil, "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Debug("Decrypted database", "size", ubytes(int64(len(dbData))))
|
// Confirm the decrypted database really is the snapshot named by
|
||||||
|
// remoteKey before any files are read from it. On mismatch, close the
|
||||||
// Create a temporary database file and write the binary SQLite data directly
|
// database and remove its private directory so nothing is left behind.
|
||||||
tempFile, err := afero.TempFile(v.Fs, "", "vaultik-restore-*.db")
|
err = v.verifySnapshotDBIdentity(db, snapshotID, remoteKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("creating temp file: %w", err)
|
_ = db.Close()
|
||||||
|
_ = v.Fs.RemoveAll(tempDir)
|
||||||
|
|
||||||
|
return nil, "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
tempPath := tempFile.Name()
|
return db, tempDir, nil
|
||||||
|
}
|
||||||
|
|
||||||
// Write the binary SQLite database directly
|
// verifySnapshotDBIdentity confirms the decrypted metadata database really
|
||||||
_, err = tempFile.Write(dbData)
|
// 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 {
|
||||||
_ = tempFile.Close()
|
return fmt.Errorf("checking identity of database for %s: %w", requested, err)
|
||||||
_ = v.Fs.Remove(tempPath)
|
|
||||||
|
|
||||||
return nil, fmt.Errorf("writing database file: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err = tempFile.Close()
|
if snapshot.RemoteSnapshotKey(snap.ID.String()) != remoteKey {
|
||||||
if err != nil {
|
return fmt.Errorf("%w: requested %s but the database is snapshot %s",
|
||||||
_ = v.Fs.Remove(tempPath)
|
errSnapshotDBMismatch, requested, snap.ID)
|
||||||
|
|
||||||
return nil, fmt.Errorf("closing temp file: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Debug("Created restore database", "path", tempPath)
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Open the database
|
// materializeSnapshotDB streams the decrypted snapshot database into a
|
||||||
db, err := database.New(v.ctx, tempPath)
|
// fresh private (0700) temp directory and opens the file read-only. The
|
||||||
|
// database is copied through an io.Copy buffer rather than read whole into
|
||||||
|
// memory. On any failure it removes the directory before returning, so no
|
||||||
|
// decrypted metadata is left on disk when the copy is interrupted or the
|
||||||
|
// payload is damaged. On success the returned directory is the caller's to
|
||||||
|
// remove.
|
||||||
|
func (v *Vaultik) materializeSnapshotDB(
|
||||||
|
dbReader io.Reader,
|
||||||
|
) (*database.DB, string, error) {
|
||||||
|
tempDir, err := afero.TempDir(v.Fs, "", "vaultik-restore-")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("opening restore database: %w", err)
|
return nil, "", fmt.Errorf("creating temp directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return db, nil
|
success := false
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
if !success {
|
||||||
|
_ = v.Fs.RemoveAll(tempDir)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
dbPath := filepath.Join(tempDir, snapshotDBFilename)
|
||||||
|
|
||||||
|
dbFile, err := v.Fs.OpenFile(
|
||||||
|
dbPath, os.O_CREATE|os.O_EXCL|os.O_WRONLY, restoreFileMode)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", fmt.Errorf("creating database file: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
written, copyErr := io.Copy(dbFile, dbReader)
|
||||||
|
closeErr := dbFile.Close()
|
||||||
|
|
||||||
|
if copyErr != nil {
|
||||||
|
return nil, "", fmt.Errorf("writing database file: %w", copyErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
if closeErr != nil {
|
||||||
|
return nil, "", fmt.Errorf("closing database file: %w", closeErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debug("Created restore database", "path", dbPath, "size", ubytes(written))
|
||||||
|
|
||||||
|
// Reject an empty database before OpenReadOnly builds a schema on it.
|
||||||
|
if written == 0 {
|
||||||
|
return nil, "", errEmptySnapshotDB
|
||||||
|
}
|
||||||
|
|
||||||
|
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
|
||||||
@@ -744,7 +863,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
|
||||||
identity age.Identity
|
identities []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
|
||||||
@@ -760,13 +879,85 @@ 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 := filepath.Join(s.opts.TargetDir, file.Path.String())
|
targetPath, err := containedRestorePath(
|
||||||
|
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)
|
||||||
}
|
}
|
||||||
@@ -814,6 +1005,13 @@ 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++
|
||||||
@@ -821,25 +1019,22 @@ func (s *restoreSession) restoreDirectory(
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// applyFileMetadata applies stored permissions, ownership (when running
|
// applyFileMetadata applies ownership (when running as root on a real
|
||||||
// as root on a real filesystem), and mtime to a restored path. Failures
|
// filesystem) and mtime to a restored path. Permission mode is applied
|
||||||
// are logged at debug level and do not abort the restore.
|
// separately by each caller, with different failure handling, so it is
|
||||||
|
// 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)
|
||||||
}
|
}
|
||||||
@@ -873,17 +1068,30 @@ func (s *restoreSession) restoreRegularFile(
|
|||||||
|
|
||||||
t0 = time.Now()
|
t0 = time.Now()
|
||||||
|
|
||||||
outFile, err := s.v.Fs.Create(targetPath)
|
// Remove any existing entry, then create the file with a restrictive
|
||||||
|
// 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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -901,9 +1109,12 @@ 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++
|
||||||
@@ -914,6 +1125,31 @@ 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.
|
||||||
@@ -926,6 +1162,12 @@ 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]
|
||||||
@@ -978,12 +1220,12 @@ func (s *restoreSession) writeFileChunks(
|
|||||||
// size, which is what makes multi-GB blobs tractable on machines with
|
// size, which is what makes multi-GB blobs tractable on machines with
|
||||||
// less RAM than the blob.
|
// less RAM than the blob.
|
||||||
func (s *restoreSession) downloadBlobToCache(
|
func (s *restoreSession) downloadBlobToCache(
|
||||||
blobHash string, expectedSize int64,
|
blobHash string, compressedSize, uncompressedSize int64,
|
||||||
) error {
|
) error {
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
|
|
||||||
t0 := time.Now()
|
t0 := time.Now()
|
||||||
rc, err := s.v.FetchAndDecryptBlob(s.ctx, blobHash, expectedSize, s.identity)
|
rc, err := s.v.FetchAndDecryptBlob(s.ctx, blobHash, uncompressedSize, s.identities...)
|
||||||
fetchSetupDur := time.Since(t0)
|
fetchSetupDur := time.Since(t0)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -995,17 +1237,25 @@ 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
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Debug("Streamed blob into disk cache",
|
log.Debug("Streamed blob into disk cache",
|
||||||
"hash", blobHash[:16],
|
"hash", blobHash[:16],
|
||||||
"compressed_bytes", expectedSize,
|
"compressed_bytes", compressedSize,
|
||||||
"plaintext_bytes", written,
|
"plaintext_bytes", written,
|
||||||
"ms_total", time.Since(start).Milliseconds(),
|
"ms_total", time.Since(start).Milliseconds(),
|
||||||
"ms_fetch_setup", fetchSetupDur.Milliseconds(),
|
"ms_fetch_setup", fetchSetupDur.Milliseconds(),
|
||||||
@@ -1062,17 +1312,22 @@ func (v *Vaultik) verifyRestoredFiles(
|
|||||||
return ctx.Err()
|
return ctx.Err()
|
||||||
}
|
}
|
||||||
|
|
||||||
targetPath := filepath.Join(targetDir, file.Path.String())
|
targetPath, err := containedRestorePath(v.Fs, 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
|
||||||
@@ -1162,6 +1417,17 @@ 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))
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,175 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
@@ -0,0 +1,108 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
@@ -0,0 +1,159 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
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")
|
||||||
|
}
|
||||||
@@ -0,0 +1,306 @@
|
|||||||
|
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,10 +171,13 @@ 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 zero FileID return means nothing is pending.
|
// The second return value is false when no file needs a download, so a
|
||||||
func (p *restorePlan) pickNextDownload() types.FileID {
|
// genuine file carrying the nil UUID is picked rather than mistaken for
|
||||||
|
// "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
|
||||||
@@ -188,14 +191,15 @@ func (p *restorePlan) pickNextDownload() types.FileID {
|
|||||||
}
|
}
|
||||||
|
|
||||||
idStr := id.String()
|
idStr := id.String()
|
||||||
if n < bestCount || (n == bestCount && (best.IsZero() || idStr < bestID)) {
|
if !found || n < bestCount || (n == bestCount && idStr < bestID) {
|
||||||
best = id
|
best = id
|
||||||
|
found = true
|
||||||
bestCount = n
|
bestCount = n
|
||||||
bestID = idStr
|
bestID = idStr
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return best
|
return best, found
|
||||||
}
|
}
|
||||||
|
|
||||||
// blobsNeeded returns the uncached blob hashes for fileID in any order.
|
// blobsNeeded returns the uncached blob hashes for fileID in any order.
|
||||||
|
|||||||
@@ -0,0 +1,88 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
@@ -0,0 +1,251 @@
|
|||||||
|
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))
|
||||||
|
}
|
||||||
@@ -0,0 +1,108 @@
|
|||||||
|
package vaultik //nolint:testpackage // inspects unexported snapshot-db materialization
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"filippo.io/age"
|
||||||
|
"github.com/spf13/afero"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"sneak.berlin/go/vaultik/internal/blobgen"
|
||||||
|
"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(bytes.NewReader(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")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMaterializeSnapshotDBRejectsCompleteEmptyStream proves the written == 0
|
||||||
|
// guard rejects a genuinely empty but complete metadata object: a real age
|
||||||
|
// header, nonce, and final tag encrypting zero plaintext bytes. The truncation
|
||||||
|
// case is stopped earlier by the reader (io.ErrUnexpectedEOF) and never reaches
|
||||||
|
// this branch, so it needs its own input. This complete stream decrypts to zero
|
||||||
|
// bytes with a clean EOF, passes the reader, and must be refused as empty rather
|
||||||
|
// than accepted as a valid zero-table database. Reverting the guard lets the
|
||||||
|
// empty file open as a fresh schema and the test fails.
|
||||||
|
func TestMaterializeSnapshotDBRejectsCompleteEmptyStream(t *testing.T) {
|
||||||
|
identity, err := age.GenerateX25519Identity()
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var stream bytes.Buffer
|
||||||
|
|
||||||
|
w, err := age.Encrypt(&stream, identity.Recipient())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, w.Close())
|
||||||
|
|
||||||
|
blobReader, err := blobgen.NewReader(bytes.NewReader(stream.Bytes()), identity)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
t.Cleanup(func() { _ = blobReader.Close() })
|
||||||
|
|
||||||
|
t.Setenv("TMPDIR", t.TempDir())
|
||||||
|
|
||||||
|
v := &Vaultik{ctx: context.Background(), Fs: afero.NewOsFs()}
|
||||||
|
|
||||||
|
_, _, err = v.materializeSnapshotDB(blobReader)
|
||||||
|
require.ErrorIs(t, err, errEmptySnapshotDB)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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(
|
||||||
|
bytes.NewReader([]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")
|
||||||
|
}
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
package vaultik_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"filippo.io/age"
|
||||||
|
"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"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestRestoreRejectsTruncatedMetadataDB backs up a real tree, then replaces
|
||||||
|
// the snapshot's db.zst.age with a stream cut right after the age header and
|
||||||
|
// its 16-byte nonce. age.Decrypt still accepts such an object and the zstd
|
||||||
|
// decoder turns the truncated read into a clean EOF, so before the fix restore
|
||||||
|
// built a fresh empty schema and reported success. Restore must now fail with
|
||||||
|
// io.ErrUnexpectedEOF, the error the reader raises for a truncated object.
|
||||||
|
// Asserting that specific error pins the reader fix: without it the truncation
|
||||||
|
// yields an empty database, which the identity check rejects for an unrelated
|
||||||
|
// reason, and this test would pass anyway.
|
||||||
|
func TestRestoreRejectsTruncatedMetadataDB(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")
|
||||||
|
restoreDir := filepath.Join(tempDir, "restored")
|
||||||
|
dbPath := filepath.Join(tempDir, "index.sqlite")
|
||||||
|
|
||||||
|
chunkSize := int64(64 * 1024)
|
||||||
|
maxBlobSize := int64(512 * 1024)
|
||||||
|
|
||||||
|
setupE2ESourceTree(t, fs, dataDir, chunkSize)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
cfg, storer, snapshotID := runFileStorageBackup(
|
||||||
|
ctx, t, fs, dataDir, storeDir, dbPath, chunkSize, maxBlobSize)
|
||||||
|
|
||||||
|
// Encrypting empty plaintext to the snapshot recipient yields
|
||||||
|
// header + nonce(16) + a single 16-byte final chunk tag. Dropping the
|
||||||
|
// trailing tag leaves exactly the age header plus its nonce — the
|
||||||
|
// truncation an attacker can write over metadata without any key.
|
||||||
|
recipient, err := age.ParseX25519Recipient(testAgePublicKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var full bytes.Buffer
|
||||||
|
|
||||||
|
w, err := age.Encrypt(&full, recipient)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, w.Close())
|
||||||
|
|
||||||
|
truncated := full.Bytes()[:full.Len()-16]
|
||||||
|
|
||||||
|
dbKeyPath := filepath.Join(storeDir, "metadata",
|
||||||
|
snapshot.RemoteSnapshotKey(snapshotID), "db.zst.age")
|
||||||
|
require.NoError(t, afero.WriteFile(fs, dbKeyPath, truncated, 0o644))
|
||||||
|
|
||||||
|
restoreVaultik := &vaultik.Vaultik{
|
||||||
|
Config: cfg,
|
||||||
|
Storage: storer,
|
||||||
|
Fs: fs,
|
||||||
|
Stdout: io.Discard,
|
||||||
|
Stderr: io.Discard,
|
||||||
|
UI: ui.NewWithColor(io.Discard, false),
|
||||||
|
}
|
||||||
|
restoreVaultik.SetContext(ctx)
|
||||||
|
|
||||||
|
err = restoreVaultik.Restore(&vaultik.RestoreOptions{
|
||||||
|
SnapshotID: snapshotID,
|
||||||
|
TargetDir: restoreDir,
|
||||||
|
Verify: true,
|
||||||
|
})
|
||||||
|
require.ErrorIs(t, err, io.ErrUnexpectedEOF)
|
||||||
|
}
|
||||||
@@ -0,0 +1,191 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
+162
-75
@@ -20,7 +20,6 @@ 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")
|
||||||
@@ -56,8 +55,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 - see CleanupIncompleteSnapshots for details
|
// This is critical for data safety; PruneDatabase below does it.
|
||||||
hostname := v.Config.Hostname
|
hostname := v.Config.Hostname
|
||||||
if hostname == "" {
|
if hostname == "" {
|
||||||
hostname, _ = os.Hostname()
|
hostname, _ = os.Hostname()
|
||||||
@@ -670,11 +669,24 @@ func (v *Vaultik) VerifySnapshotWithOptions(
|
|||||||
|
|
||||||
v.printVerifyHeader(snapshotID, opts)
|
v.printVerifyHeader(snapshotID, opts)
|
||||||
|
|
||||||
// Resolve the identifier to the snapshot's remote key and download the
|
// Resolve the identifier to the snapshot's remote key. A human ID is
|
||||||
// manifest. A human ID is hashed; a remote key (or its abbreviation,
|
// hashed; a remote key (or its abbreviation, as printed for a
|
||||||
// as printed for a remote-only snapshot) is used as-is, so a host with
|
// remote-only snapshot) is used as-is, so a host with no local index
|
||||||
// no local index can verify a snapshot it can only see on the store.
|
// can verify a snapshot it can only see on the store. The key is kept
|
||||||
manifest, err := v.resolveAndDownloadManifest(snapshotID)
|
// so we can also check for the snapshot's encrypted database below.
|
||||||
|
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
|
||||||
@@ -704,14 +716,34 @@ func (v *Vaultik) VerifySnapshotWithOptions(
|
|||||||
|
|
||||||
v.printlnStdout()
|
v.printlnStdout()
|
||||||
|
|
||||||
// Check each blob exists
|
// Check each blob is present with the size the manifest records.
|
||||||
v.stdoutf("Checking blob existence...\n")
|
v.stdoutf("Checking blob presence and sizes...\n")
|
||||||
}
|
}
|
||||||
|
|
||||||
result.Verified, result.Missing, result.MissingSize =
|
// A snapshot is only restorable if its encrypted database is present
|
||||||
v.verifyManifestBlobsExist(manifest, opts)
|
// alongside the blobs. Shallow verify checks that the object exists; it
|
||||||
|
// does not decrypt it (that is deep verify's job).
|
||||||
|
dbPath := fmt.Sprintf("metadata/%s/db.zst.age", remoteKey)
|
||||||
|
|
||||||
return v.formatVerifyResult(result, manifest, opts)
|
_, dbErr := v.Storage.Stat(v.ctx, dbPath)
|
||||||
|
if dbErr != nil {
|
||||||
|
result.DatabaseMissing = true
|
||||||
|
}
|
||||||
|
|
||||||
|
result.Verified, result.Missing, result.Mismatched, result.MissingSize, err =
|
||||||
|
v.verifyManifestBlobs(manifest, opts)
|
||||||
|
if err != nil {
|
||||||
|
if opts.JSON {
|
||||||
|
result.Status = verifyStatusFailed
|
||||||
|
result.ErrorMessage = fmt.Sprintf("verifying manifest blobs: %v", err)
|
||||||
|
|
||||||
|
return v.outputVerifyJSON(result)
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Errorf("verifying manifest blobs: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return v.formatVerifyResult(result, opts)
|
||||||
}
|
}
|
||||||
|
|
||||||
// printVerifyHeader prints the snapshot ID and parsed timestamp for
|
// printVerifyHeader prints the snapshot ID and parsed timestamp for
|
||||||
@@ -736,25 +768,34 @@ func (v *Vaultik) printVerifyHeader(snapshotID string, opts *VerifyOptions) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// verifyManifestBlobsExist checks that each blob in the manifest exists
|
// verifyManifestBlobs checks that each blob in the manifest is present in
|
||||||
// in storage, returning the verified count, missing count, and total
|
// storage with the size the manifest records, returning the counts of
|
||||||
// missing bytes.
|
// blobs that were present with the right size, absent, and present but the
|
||||||
func (v *Vaultik) verifyManifestBlobsExist(
|
// wrong size, plus the total bytes of the absent blobs. It does not read
|
||||||
|
// 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, int64) {
|
) (int, int, int, int64, error) {
|
||||||
var (
|
var (
|
||||||
verified, missing int
|
verified, missing, mismatched int
|
||||||
missingSize int64
|
missingSize int64
|
||||||
)
|
)
|
||||||
|
|
||||||
for _, blob := range manifest.Blobs {
|
for _, blob := range manifest.Blobs {
|
||||||
|
// The manifest is unauthenticated, so its blob hashes are checked
|
||||||
|
// before being spliced into a storage path.
|
||||||
|
if !isBlobHash(blob.Hash) {
|
||||||
|
return 0, 0, 0, 0, fmt.Errorf("%w: %s",
|
||||||
|
errInvalidBlobHash, shortHash(blob.Hash))
|
||||||
|
}
|
||||||
|
|
||||||
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)
|
||||||
|
|
||||||
// Shallow: check existence only (deep verification is handled
|
stat, err := v.Storage.Stat(v.ctx, blobPath)
|
||||||
// by RunDeepVerify).
|
switch {
|
||||||
_, err := v.Storage.Stat(v.ctx, blobPath)
|
case err != nil:
|
||||||
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))
|
||||||
@@ -762,23 +803,32 @@ func (v *Vaultik) verifyManifestBlobsExist(
|
|||||||
|
|
||||||
missing++
|
missing++
|
||||||
missingSize += blob.CompressedSize
|
missingSize += blob.CompressedSize
|
||||||
} else {
|
case stat.Size != blob.CompressedSize:
|
||||||
|
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, missingSize
|
return verified, missing, mismatched, missingSize, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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, manifest *snapshot.Manifest, opts *VerifyOptions,
|
result *VerifyResult, opts *VerifyOptions,
|
||||||
) error {
|
) error {
|
||||||
|
failure := shallowVerifyFailure(result)
|
||||||
|
|
||||||
if opts.JSON {
|
if opts.JSON {
|
||||||
if result.Missing > 0 {
|
if failure != "" {
|
||||||
result.Status = verifyStatusFailed
|
result.Status = verifyStatusFailed
|
||||||
result.ErrorMessage = fmt.Sprintf("%d blobs are missing", result.Missing)
|
result.ErrorMessage = failure
|
||||||
} else {
|
} else {
|
||||||
result.Status = "ok"
|
result.Status = "ok"
|
||||||
}
|
}
|
||||||
@@ -787,29 +837,57 @@ func (v *Vaultik) formatVerifyResult(
|
|||||||
}
|
}
|
||||||
|
|
||||||
v.stdoutf("\nVerification complete:\n")
|
v.stdoutf("\nVerification complete:\n")
|
||||||
v.stdoutf(" Verified: %d blobs (%s)\n", result.Verified,
|
v.stdoutf(" Present with listed size: %d blobs\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 result.Missing > 0 {
|
if failure != "" {
|
||||||
v.stdoutf("FAILED - %d blobs are missing\n", result.Missing)
|
v.stdoutf("FAILED - %s\n", failure)
|
||||||
|
|
||||||
return fmt.Errorf("%d %w", result.Missing, errBlobsMissing)
|
return fmt.Errorf("%w: %s", errSnapshotVerifyFailed, failure)
|
||||||
}
|
}
|
||||||
|
|
||||||
v.stdoutf("OK - All blobs verified\n")
|
// Report only what was actually checked: presence and size, not contents.
|
||||||
|
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)
|
||||||
@@ -935,29 +1013,23 @@ 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")
|
||||||
|
|
||||||
// Get all remote snapshot IDs
|
// Remote metadata lives under metadata/<remote-key>/, where the
|
||||||
remoteSnapshots := make(map[string]bool)
|
// directory name is snapshot.RemoteSnapshotKey(id), not the human
|
||||||
objectCh := v.Storage.ListStream(v.ctx, "metadata/")
|
// snapshot ID. Compare each local row's hashed key against that set
|
||||||
|
// so a row still backed by remote metadata is kept. Comparing human
|
||||||
for object := range objectCh {
|
// IDs against the hashed directory names matches nothing and deletes
|
||||||
if object.Err != nil {
|
// every local snapshot record (issue #160).
|
||||||
return fmt.Errorf("listing remote snapshots: %w", object.Err)
|
remoteKeys, err := v.listAllRemoteSnapshotKeys()
|
||||||
}
|
if err != nil {
|
||||||
|
return fmt.Errorf("listing remote snapshots: %w", err)
|
||||||
// Extract snapshot ID from paths like metadata/hostname-20240115-143052Z/
|
|
||||||
parts := strings.Split(object.Key, "/")
|
|
||||||
if len(parts) >= minSnapshotIDParts &&
|
|
||||||
parts[0] == metadataDirName && parts[1] != "" {
|
|
||||||
// Skip macOS resource fork files (._*) and other hidden files
|
|
||||||
if strings.HasPrefix(parts[1], ".") {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
remoteSnapshots[parts[1]] = true
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Debug("Found remote snapshots", "count", len(remoteSnapshots))
|
remoteKeySet := make(map[string]bool, len(remoteKeys))
|
||||||
|
for _, k := range remoteKeys {
|
||||||
|
remoteKeySet[k] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Debug("Found remote snapshots", "count", len(remoteKeySet))
|
||||||
|
|
||||||
// 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)
|
||||||
@@ -965,12 +1037,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 that don't exist remotely
|
// Remove local snapshots whose metadata is absent from the remote.
|
||||||
removedCount := 0
|
removedCount := 0
|
||||||
|
|
||||||
for _, snap := range localSnapshots {
|
for _, snap := range localSnapshots {
|
||||||
snapshotIDStr := snap.ID.String()
|
snapshotIDStr := snap.ID.String()
|
||||||
if !remoteSnapshots[snapshotIDStr] {
|
if !remoteKeySet[snapshot.RemoteSnapshotKey(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)
|
||||||
|
|
||||||
@@ -1258,21 +1330,36 @@ func (v *Vaultik) listAllRemoteSnapshotKeys() ([]string, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
parts := strings.Split(object.Key, "/")
|
parts := strings.Split(object.Key, "/")
|
||||||
if len(parts) >= minSnapshotIDParts &&
|
if len(parts) < minSnapshotIDParts ||
|
||||||
parts[0] == metadataDirName && parts[1] != "" {
|
parts[0] != metadataDirName || parts[1] == "" {
|
||||||
// Skip macOS resource fork files (._*) and other hidden files
|
continue
|
||||||
if strings.HasPrefix(parts[1], ".") {
|
}
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if strings.HasSuffix(object.Key, "/") ||
|
// Skip macOS resource fork files (._*) and other hidden files
|
||||||
strings.Contains(object.Key, "/manifest.json.zst") {
|
if strings.HasPrefix(parts[1], ".") {
|
||||||
key := parts[1]
|
continue
|
||||||
if !seen[key] {
|
}
|
||||||
seen[key] = true
|
|
||||||
keys = append(keys, key)
|
if !strings.HasSuffix(object.Key, "/") &&
|
||||||
}
|
!strings.Contains(object.Key, "/manifest.json.zst") {
|
||||||
}
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
key := parts[1]
|
||||||
|
|
||||||
|
// A remote snapshot key is a SHA-256 hash: 64 lowercase hex
|
||||||
|
// characters. The listing comes from the untrusted destination,
|
||||||
|
// so accept a key only in that form.
|
||||||
|
if !isBlobHash(key) {
|
||||||
|
log.Warn("Skipping non-conforming key under metadata/",
|
||||||
|
"key", object.Key)
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !seen[key] {
|
||||||
|
seen[key] = true
|
||||||
|
keys = append(keys, key)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user