Compare commits
69
Commits
e10b4cec82
...
next
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5ca615a7a6 | ||
|
|
62967f28d0 | ||
|
|
007254a1f0 | ||
|
|
eb596b8be6 | ||
|
|
596b978cb1 | ||
|
|
24d99819a3 | ||
|
|
1ec0423e6e | ||
|
|
4ed77902d1 | ||
|
|
00713b8677 | ||
|
|
71c386ecbf | ||
|
|
cba526d33f | ||
|
|
fb4481b4f7 | ||
|
|
5ec59862ff | ||
|
|
e640d10964 | ||
|
|
4e562f834f | ||
|
|
641d5659ec | ||
|
|
663986f551 | ||
|
|
32a61ff963 | ||
|
|
bdb1c7ec18 | ||
|
|
51e3731076 | ||
|
|
a5faec0466 | ||
|
|
7c6531eaf7 | ||
|
|
d52b4f1240 | ||
|
|
41cea400a7 | ||
|
|
6e5e0db999 | ||
|
|
e0e5ae68a4 | ||
|
|
1fc11529ed | ||
|
|
b090b3f86b | ||
|
|
a3d3fb3b69 | ||
|
|
4dc26c9394 | ||
|
|
7546cb094f | ||
|
|
797d2678c8 | ||
|
|
78015afb35 | ||
|
|
1c330c697f | ||
|
|
d18e286377 | ||
|
|
f49fde3a06 | ||
|
|
206651f89a | ||
|
|
c0f221b1ca | ||
|
|
09be20a044 | ||
|
|
2e1ba7d2e0 | ||
|
|
1a23016df1 | ||
|
|
ebe3c17618 | ||
|
|
1a96360f6a | ||
|
|
4f5d2126d6 | ||
|
|
6be4601763 | ||
|
|
36ece2fca7 | ||
|
|
dc225bd0b1 | ||
|
|
6acd57d0ec | ||
|
|
596027f210 | ||
|
|
0aa9a52497 | ||
|
|
09ec79c57e | ||
|
|
e8339f4d12 | ||
|
|
4f984cd9c6 | ||
|
|
d1caf0a208 | ||
|
|
8eb25b98fd | ||
|
|
6211b8e768 | ||
|
|
0307f23024 | ||
|
|
3fd30bb9e6 | ||
|
|
6ff00c696a | ||
|
|
c6551e4901 | ||
|
|
b06d7fa3f4 | ||
|
|
16d5b237d2 | ||
|
|
660de5716a | ||
|
|
51fb2805fd | ||
|
|
6ffb24b544 | ||
|
|
4419ef7730 | ||
|
|
991b1a5a0b | ||
|
|
fd77a047f9 | ||
|
|
341428d9ca |
@@ -1,3 +0,0 @@
|
||||
EXTREMELY IMPORTANT: Read and follow the policies, procedures, and
|
||||
instructions in the `AGENTS.md` file in the root of the repository. Make
|
||||
sure you follow *all* of the instructions meticulously.
|
||||
+10
-2
@@ -1,3 +1,9 @@
|
||||
# .git is sent without its config. Without a VERSION build argument the
|
||||
# stage that compiles runs `git describe --tags --always` on .git, which
|
||||
# does not need .git/config; that file can hold a credential, such as a
|
||||
# password in a remote URL or the token the CI checkout step stores there.
|
||||
.git/config
|
||||
|
||||
# Build artifacts
|
||||
secret
|
||||
coverage.out
|
||||
@@ -10,6 +16,9 @@ coverage.out
|
||||
*.swo
|
||||
*~
|
||||
|
||||
# Dependencies
|
||||
node_modules
|
||||
|
||||
# macOS
|
||||
.DS_Store
|
||||
|
||||
@@ -17,5 +26,4 @@ coverage.out
|
||||
.claude/
|
||||
|
||||
# Local settings
|
||||
.golangci.yml
|
||||
.claude/settings.local.json
|
||||
.claude/settings.local.json
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
root = true
|
||||
|
||||
[*]
|
||||
indent_style = space
|
||||
indent_size = 4
|
||||
end_of_line = lf
|
||||
charset = utf-8
|
||||
trim_trailing_whitespace = true
|
||||
insert_final_newline = true
|
||||
|
||||
[Makefile]
|
||||
indent_style = tab
|
||||
@@ -0,0 +1,9 @@
|
||||
name: check
|
||||
on: [push]
|
||||
jobs:
|
||||
check:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
# actions/checkout v4.2.2, 2026-02-28
|
||||
- uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683
|
||||
- run: script/cibuild
|
||||
+29
-3
@@ -1,8 +1,34 @@
|
||||
# OS
|
||||
.DS_Store
|
||||
**/.DS_Store
|
||||
Thumbs.db
|
||||
|
||||
# Editors
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
*.bak
|
||||
.idea/
|
||||
.vscode/
|
||||
*.sublime-*
|
||||
|
||||
# Agent scratch (worktrees of this repo, created and destroyed by
|
||||
# in-flight tooling). Unanchored: .gitignore patterns already match at
|
||||
# every depth, so no prefix is wanted here. This is not a .dockerignore
|
||||
# entry and must not be given a `**/` prefix on the way into one.
|
||||
.claude/
|
||||
|
||||
# Node
|
||||
node_modules/
|
||||
|
||||
# Environment / secrets
|
||||
.env
|
||||
.env.*
|
||||
*.pem
|
||||
*.key
|
||||
|
||||
# This repo. /secret is the built binary, anchored so that it does not
|
||||
# also match the internal/secret/ package directory.
|
||||
/secret
|
||||
*.log
|
||||
cli.test
|
||||
vault.test
|
||||
*.test
|
||||
settings.local.json
|
||||
|
||||
+87
-117
@@ -1,128 +1,98 @@
|
||||
version: "2"
|
||||
|
||||
# Config schema uses the golangci-lint v2 layout (settings live under
|
||||
# linters.settings, not top-level linters-settings) so that the
|
||||
# thresholds below are actually applied by golangci-lint >= v2.
|
||||
|
||||
run:
|
||||
go: "1.24"
|
||||
tests: false
|
||||
timeout: 5m
|
||||
modules-download-mode: readonly
|
||||
|
||||
linters:
|
||||
default: all
|
||||
enable:
|
||||
# Additional linters requested
|
||||
- testifylint # Checks usage of github.com/stretchr/testify
|
||||
- usetesting # usetesting is an analyzer that detects using os.Setenv instead of t.Setenv since Go 1.17
|
||||
- tagliatelle # Checks the struct tags
|
||||
- nlreturn # nlreturn checks for a new line before return and branch statements
|
||||
- nilnil # Checks that there is no simultaneous return of nil error and an invalid value
|
||||
- nestif # Reports deeply nested if statements
|
||||
- mnd # An analyzer to detect magic numbers
|
||||
- lll # Reports long lines
|
||||
- intrange # intrange is a linter to find places where for loops could make use of an integer range
|
||||
- gochecknoglobals # Check that no global variables exist
|
||||
|
||||
# Default/existing linters that are commonly useful
|
||||
- govet
|
||||
- errcheck
|
||||
- staticcheck
|
||||
- unused
|
||||
- ineffassign
|
||||
- misspell
|
||||
- revive
|
||||
- gosec
|
||||
- unconvert
|
||||
- unparam
|
||||
|
||||
linters-settings:
|
||||
lll:
|
||||
line-length: 120
|
||||
|
||||
mnd:
|
||||
# List of enabled checks, see https://github.com/tommy-muehle/go-mnd/#checks for description.
|
||||
checks:
|
||||
- argument
|
||||
- case
|
||||
- condition
|
||||
- operation
|
||||
- return
|
||||
- assign
|
||||
ignored-numbers:
|
||||
- '0'
|
||||
- '1'
|
||||
- '2'
|
||||
- '8'
|
||||
- '16'
|
||||
- '40' # GPG fingerprint length
|
||||
- '64'
|
||||
- '128'
|
||||
- '256'
|
||||
- '512'
|
||||
- '1024'
|
||||
- '2048'
|
||||
- '4096'
|
||||
|
||||
nestif:
|
||||
min-complexity: 4
|
||||
|
||||
nlreturn:
|
||||
block-size: 2
|
||||
|
||||
revive:
|
||||
rules:
|
||||
- name: var-naming
|
||||
arguments:
|
||||
- []
|
||||
- []
|
||||
- "upperCaseConst=true"
|
||||
|
||||
tagliatelle:
|
||||
case:
|
||||
# Successor to the deprecated gomodguard. Named explicitly, rather than
|
||||
# left to `default: all`, because it carries the module policy below.
|
||||
- gomodguard_v2
|
||||
disable:
|
||||
# Genuinely incompatible with project patterns
|
||||
- exhaustruct # Requires all struct fields
|
||||
- godot # Requires comments to end with periods
|
||||
- wrapcheck # Too verbose for internal packages
|
||||
- varnamelen # Short names like db, id are idiomatic Go
|
||||
# Deprecated: the warning is attached to the old name, so it is
|
||||
# silenced by disabling that name, not by enabling the successor.
|
||||
- wsl # Deprecated, replaced by wsl_v5
|
||||
- gomodguard # Deprecated, replaced by gomodguard_v2
|
||||
settings:
|
||||
lll:
|
||||
line-length: 88
|
||||
funlen:
|
||||
lines: 80
|
||||
statements: 50
|
||||
cyclop:
|
||||
max-complexity: 15
|
||||
dupl:
|
||||
threshold: 100
|
||||
depguard:
|
||||
# Test-support code must not be compiled into the shipped binary. A
|
||||
# test-support package exists to hand a test privileges the program
|
||||
# itself must never have, so a file that is not a test must not import
|
||||
# one. Test files, and the files inside a package whose directory name
|
||||
# ends in `test`, are where that code belongs, and are exempt.
|
||||
#
|
||||
# The deny list below is the one part of this file a repository is
|
||||
# expected to extend, and the only part it may. depguard matches an
|
||||
# import path against a list of prefixes, so it cannot be told "any path
|
||||
# whose last segment ends in test"; a repository's own test-support
|
||||
# packages have to be named here one at a time, by full import path,
|
||||
# under a module path that differs from repository to repository. Add
|
||||
# them; change nothing else.
|
||||
rules:
|
||||
json: snake
|
||||
yaml: snake
|
||||
xml: snake
|
||||
bson: snake
|
||||
|
||||
testifylint:
|
||||
enable-all: true
|
||||
|
||||
usetesting: {}
|
||||
test-support:
|
||||
list-mode: lax
|
||||
files:
|
||||
- "$all"
|
||||
- "!$test"
|
||||
- "!**/*test/**"
|
||||
deny:
|
||||
- pkg: net/http/httptest
|
||||
desc: >-
|
||||
Test-support code belongs in test files and in packages whose
|
||||
directory name ends in test, not in the shipped binary.
|
||||
# Only decisions already recorded in the Go package defaults are
|
||||
# listed here. Every entry matches the module path exactly.
|
||||
gomodguard_v2:
|
||||
blocked:
|
||||
- module: github.com/rs/zerolog
|
||||
recommendations:
|
||||
- log/slog
|
||||
reason: "Structured logging is stdlib log/slog."
|
||||
# One entry per pre-fork module path, because the later releases
|
||||
# are separate paths. A prefix match would be shorter but would
|
||||
# also reach github.com/go-redis/redismock, the test double for
|
||||
# the successor these entries recommend.
|
||||
- module: github.com/go-redis/redis
|
||||
recommendations:
|
||||
- github.com/redis/go-redis/v9
|
||||
reason: "Pre-fork module; use the maintained go-redis v9."
|
||||
- module: github.com/go-redis/redis/v7
|
||||
recommendations:
|
||||
- github.com/redis/go-redis/v9
|
||||
reason: "Pre-fork module; use the maintained go-redis v9."
|
||||
- module: github.com/go-redis/redis/v8
|
||||
recommendations:
|
||||
- github.com/redis/go-redis/v9
|
||||
reason: "Pre-fork module; use the maintained go-redis v9."
|
||||
- module: github.com/sergi/go-diff
|
||||
recommendations:
|
||||
- github.com/aymanbagabas/go-udiff
|
||||
reason: "No unified diff output; use go-udiff."
|
||||
- module: github.com/hexops/gotextdiff
|
||||
recommendations:
|
||||
- github.com/aymanbagabas/go-udiff
|
||||
reason: "Unmaintained fork; use go-udiff."
|
||||
|
||||
issues:
|
||||
max-issues-per-linter: 0
|
||||
max-same-issues: 0
|
||||
exclude-rules:
|
||||
- path: ".*_gen\\.go"
|
||||
linters:
|
||||
- lll
|
||||
|
||||
# Exclude unused parameter warnings for cobra command signatures
|
||||
- text: "parameter '(args|cmd)' seems to be unused"
|
||||
linters:
|
||||
- revive
|
||||
|
||||
# Allow ALL_CAPS constant names
|
||||
- text: "don't use ALL_CAPS in Go names"
|
||||
linters:
|
||||
- revive
|
||||
|
||||
# Exclude all linters for internal/macse directory
|
||||
- path: "internal/macse/.*"
|
||||
linters:
|
||||
- errcheck
|
||||
- lll
|
||||
- mnd
|
||||
- nestif
|
||||
- nlreturn
|
||||
- revive
|
||||
- unconvert
|
||||
- govet
|
||||
- staticcheck
|
||||
- unused
|
||||
- ineffassign
|
||||
- misspell
|
||||
- gosec
|
||||
- unparam
|
||||
- testifylint
|
||||
- usetesting
|
||||
- tagliatelle
|
||||
- nilnil
|
||||
- intrange
|
||||
- gochecknoglobals
|
||||
|
||||
@@ -141,3 +141,17 @@ Version: 2025-06-08
|
||||
- Local application imports
|
||||
|
||||
Each group should be separated by a blank line.
|
||||
|
||||
## Go-Specific Guidelines
|
||||
|
||||
1. **No `panic`, `log.Fatal`, or `os.Exit` in library code.** Always propagate errors via return values.
|
||||
|
||||
2. **Constructors return `(*T, error)`, not just `*T`.** Callers must handle errors, not crash.
|
||||
|
||||
3. **Wrap errors** with `fmt.Errorf("context: %w", err)` for debuggability.
|
||||
|
||||
4. **Never modify linter config** (`.golangci.yml`) to suppress findings. Fix the code.
|
||||
|
||||
5. **All PRs must pass `make check` with zero failures.** No exceptions, no "pre-existing issue" excuses.
|
||||
|
||||
6. **Pin external dependencies by commit hash**, not mutable tags.
|
||||
|
||||
+50
-32
@@ -1,50 +1,68 @@
|
||||
# Build stage
|
||||
FROM golang:1.24-alpine AS builder
|
||||
# Lint stage — fast feedback on formatting and lint issues
|
||||
# golangci/golangci-lint:v2.12.2 (Debian-based), 2026-08-07
|
||||
FROM golangci/golangci-lint:v2.12.2@sha256:5cceeef04e53efe1470638d4b4b4f5ceefd574955ab3941b2d9a68a8c9ad5240 AS lint
|
||||
|
||||
# Install build dependencies
|
||||
RUN apk add --no-cache \
|
||||
gcc \
|
||||
musl-dev \
|
||||
make \
|
||||
git
|
||||
|
||||
# Set working directory
|
||||
WORKDIR /build
|
||||
|
||||
# Copy go mod files
|
||||
WORKDIR /src
|
||||
COPY go.mod go.sum ./
|
||||
|
||||
# Download dependencies
|
||||
RUN go mod download
|
||||
|
||||
# Copy source code
|
||||
# script/cibuild sets CHECK_EPOCH to the current time, so the RUN steps
|
||||
# below run again on each build, an unchanged tree included, while the
|
||||
# steps above stay cached. ARG is per stage: the build stage declares it too.
|
||||
ARG CHECK_EPOCH
|
||||
|
||||
COPY . .
|
||||
|
||||
# Build the binary
|
||||
RUN CGO_ENABLED=1 go build -v -o secret cmd/secret/main.go
|
||||
RUN make fmt-check
|
||||
# Not make lint: script/lint is a docker build, which cannot run in here.
|
||||
RUN golangci-lint run --config .golangci.yml ./...
|
||||
|
||||
# Build stage — tests and compilation
|
||||
# golang 1.24.13-alpine (2026-03-10)
|
||||
FROM golang@sha256:8bee1901f1e530bfb4a7850aa7a479d17ae3a18beb6e09064ed54cfd245b7191 AS builder
|
||||
|
||||
# Force BuildKit to run the lint stage
|
||||
COPY --from=lint /src/go.sum /dev/null
|
||||
|
||||
RUN apk add --no-cache gcc musl-dev make git gnupg
|
||||
|
||||
WORKDIR /build
|
||||
COPY go.mod go.sum ./
|
||||
RUN go mod download
|
||||
|
||||
# As in the lint stage: the RUN steps below run again on each script/cibuild.
|
||||
ARG CHECK_EPOCH
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN make test
|
||||
|
||||
# The version stamped into the binary: the VERSION build argument when one
|
||||
# is given, otherwise `git describe --tags --always` of the .git the build
|
||||
# context carries: the tag on a tagged commit, tag-N-gHASH on a commit after
|
||||
# one, the short commit when no tag is reachable. A context that carries .git
|
||||
# and still yields no version fails the build.
|
||||
ARG VERSION
|
||||
RUN version="${VERSION:-$(git describe --tags --always)}"; \
|
||||
if [ -e .git ] && { [ -z "$version" ] || [ "$version" = dev ] || \
|
||||
[ "$version" = unknown ]; }; then \
|
||||
echo "no version could be derived although the build context carries .git" >&2; \
|
||||
exit 1; \
|
||||
fi; \
|
||||
make build VERSION="${version:-dev}"
|
||||
|
||||
# Runtime stage
|
||||
FROM alpine:latest
|
||||
# alpine 3.23 (2026-03-10)
|
||||
FROM alpine@sha256:25109184c71bdad752c8312a8623239686a9a2071e8825f20acb8f2198c3f659
|
||||
|
||||
# Install runtime dependencies
|
||||
RUN apk add --no-cache \
|
||||
ca-certificates \
|
||||
gnupg
|
||||
RUN apk add --no-cache ca-certificates gnupg
|
||||
|
||||
# Create non-root user
|
||||
RUN adduser -D -s /bin/sh secret
|
||||
|
||||
# Copy binary from builder
|
||||
COPY --from=builder /build/secret /usr/local/bin/secret
|
||||
|
||||
# Ensure binary is executable
|
||||
RUN chmod +x /usr/local/bin/secret
|
||||
|
||||
# Switch to non-root user
|
||||
USER secret
|
||||
|
||||
# Set working directory
|
||||
WORKDIR /home/secret
|
||||
|
||||
# Set entrypoint
|
||||
ENTRYPOINT ["secret"]
|
||||
ENTRYPOINT ["secret"]
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
# Lint image, built by script/lint: golangci-lint runs as a build step, so a
|
||||
# successful build is a clean lint. Works where the docker daemon is remote
|
||||
# and bind mounts are impossible.
|
||||
|
||||
# golangci/golangci-lint:v2.12.2 (Debian-based), 2026-08-07
|
||||
FROM golangci/golangci-lint:v2.12.2@sha256:5cceeef04e53efe1470638d4b4b4f5ceefd574955ab3941b2d9a68a8c9ad5240 AS deps
|
||||
|
||||
WORKDIR /src
|
||||
|
||||
COPY go.mod go.sum ./
|
||||
RUN go mod download
|
||||
|
||||
# script/lint rebuilds this stage on every run, by this name; the module
|
||||
# download above stays cached.
|
||||
FROM deps AS lint
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN golangci-lint run --config .golangci.yml ./...
|
||||
@@ -1,44 +1,49 @@
|
||||
export CGO_ENABLED=1
|
||||
export DOCKER_HOST := ssh://root@ber1app1.local
|
||||
|
||||
# Version information
|
||||
VERSION := 0.1.0
|
||||
GIT_COMMIT := $(shell git rev-parse HEAD 2>/dev/null || echo "unknown")
|
||||
LDFLAGS := -X 'git.eeqj.de/sneak/secret/internal/cli.Version=$(VERSION)' \
|
||||
-X 'git.eeqj.de/sneak/secret/internal/cli.GitCommit=$(GIT_COMMIT)'
|
||||
.PHONY: default bootstrap setup build test lint fmt fmt-check check docker \
|
||||
docker-run clean install hooks
|
||||
|
||||
default: check
|
||||
|
||||
build: ./secret
|
||||
bootstrap:
|
||||
@script/bootstrap
|
||||
|
||||
./secret: ./internal/*/*.go ./pkg/*/*.go ./cmd/*/*.go ./go.*
|
||||
go build -v -ldflags "$(LDFLAGS)" -o $@ cmd/secret/main.go
|
||||
setup:
|
||||
@script/setup
|
||||
|
||||
vet:
|
||||
go vet ./...
|
||||
# Build ./secret; `make build VERSION=x` stamps x instead of `git describe`
|
||||
build:
|
||||
@script/build
|
||||
|
||||
test: lint vet
|
||||
go test ./... || go test -v ./...
|
||||
test:
|
||||
@script/test
|
||||
|
||||
fmt:
|
||||
go fmt ./...
|
||||
@script/fmt
|
||||
|
||||
lint:
|
||||
golangci-lint run --timeout 5m
|
||||
@script/lint
|
||||
|
||||
check: build test
|
||||
check:
|
||||
@script/check
|
||||
|
||||
# Build Docker container
|
||||
docker:
|
||||
docker build -t sneak/secret .
|
||||
@script/docker
|
||||
|
||||
# Run Docker container interactively
|
||||
docker-run:
|
||||
docker run --rm -it sneak/secret
|
||||
docker run --rm -it "$$(./script/projectname)"
|
||||
|
||||
# Clean build artifacts
|
||||
clean:
|
||||
rm -f ./secret
|
||||
|
||||
install: ./secret
|
||||
install: build
|
||||
cp ./secret $(HOME)/bin/secret
|
||||
|
||||
fmt-check:
|
||||
@script/fmt-check
|
||||
|
||||
hooks:
|
||||
@script/install-precommit
|
||||
|
||||
@@ -91,6 +91,9 @@ Lists all available vaults. The current vault is marked.
|
||||
|
||||
Creates a new vault with the specified name.
|
||||
|
||||
**Vault Name Format:** only lowercase ASCII letters, digits, `.`, `-` and `_`
|
||||
are allowed, and a name must not be empty, `.` or `..`.
|
||||
|
||||
#### `secret vault select <name>`
|
||||
|
||||
Switches to the specified vault for subsequent operations.
|
||||
@@ -113,7 +116,9 @@ automatically switch to another vault if removing the current one.
|
||||
Adds a secret to the current vault. Reads the secret value from stdin.
|
||||
- `--force, -f`: Overwrite existing secret
|
||||
|
||||
**Secret Name Format:** `[a-z0-9\.\-\_\/]+`
|
||||
**Secret Name Format:** only ASCII letters, digits, `.`, `-`, `_` and `/`
|
||||
are allowed, and a name must not be empty, start with `.` or `/`, end with
|
||||
`/`, contain `//`, or have `..` as a path segment.
|
||||
- Forward slashes (`/`) are converted to percent signs (`%`) for storage
|
||||
- Examples: `database/password`, `api.key`, `ssh_private_key`
|
||||
|
||||
@@ -137,6 +142,9 @@ matching.
|
||||
|
||||
Moves or renames a secret within the current vault.
|
||||
- Fails if the destination already exists
|
||||
- Fails if the destination is the source under another name, such as `foo`
|
||||
for `Foo` on a case-insensitive filesystem (the macOS default); there, to
|
||||
change only the case of a name, move the secret to a third name first
|
||||
- Preserves all versions and metadata
|
||||
|
||||
### Version Management
|
||||
@@ -184,6 +192,7 @@ Creates a new unlocker of the specified type:
|
||||
- `passphrase`: Traditional passphrase-protected unlocker
|
||||
- `pgp`: Uses an existing GPG key for encryption/decryption
|
||||
- `keychain`: macOS Keychain integration (macOS only)
|
||||
- `secure-enclave`: Hardware-backed Secure Enclave protection (macOS only)
|
||||
|
||||
**Options:**
|
||||
- `--keyid <id>`: GPG key ID (optional for PGP type, uses default key if not specified)
|
||||
@@ -192,7 +201,9 @@ Creates a new unlocker of the specified type:
|
||||
|
||||
**DANGER**: Permanently removes an unlocker. Like Unix `rm`, this command
|
||||
does not ask for confirmation. Cannot remove the last unlocker if the vault
|
||||
has secrets unless --force is used.
|
||||
has secrets unless --force is used. An unlocker directory that
|
||||
`secret unlocker list` skips with a warning, because its metadata cannot be
|
||||
read or parsed, is removed by the directory name the warning gives.
|
||||
- `--force, -f`: Force removal of last unlocker even if vault has secrets
|
||||
- **CRITICAL WARNING**: Without unlockers and without your mnemonic phrase,
|
||||
vault data will be PERMANENTLY INACCESSIBLE
|
||||
@@ -286,11 +297,11 @@ Unlockers provide different authentication methods to access the long-term keys:
|
||||
- Automatic unlocking when Keychain is unlocked
|
||||
- Cross-application integration
|
||||
|
||||
4. **Secure Enclave Unlockers** (macOS - planned):
|
||||
4. **Secure Enclave Unlockers** (macOS):
|
||||
- Hardware-backed key storage using Apple Secure Enclave
|
||||
- Currently partially implemented but non-functional
|
||||
- Requires Apple Developer Program membership and code signing entitlements
|
||||
- Full implementation blocked by entitlement requirements
|
||||
- Uses `sc_auth` / CryptoTokenKit for SE key management (no Apple Developer Program required)
|
||||
- ECIES encryption: vault long-term key encrypted directly by SE hardware
|
||||
- Protected by biometric authentication (Touch ID) or system password
|
||||
|
||||
Each vault maintains its own set of unlockers and one long-term key. The long-term key is encrypted to each unlocker, allowing any authorized unlocker to access vault secrets.
|
||||
|
||||
@@ -330,8 +341,7 @@ Each vault maintains its own set of unlockers and one long-term key. The long-te
|
||||
|
||||
- Hardware token support via PGP/GPG integration
|
||||
- macOS Keychain integration for system-level security
|
||||
- Secure Enclave support planned (requires paid Apple Developer Program for
|
||||
signed entitlements to access the SEP and doxxing myself to Apple)
|
||||
- Secure Enclave integration for hardware-backed key protection (macOS, via `sc_auth` / CryptoTokenKit)
|
||||
|
||||
## Examples
|
||||
|
||||
@@ -385,6 +395,7 @@ secret vault remove personal --force
|
||||
secret unlocker add passphrase # Password-based
|
||||
secret unlocker add pgp --keyid ABCD1234 # GPG key
|
||||
secret unlocker add keychain # macOS Keychain (macOS only)
|
||||
secret unlocker add secure-enclave # macOS Secure Enclave (macOS only)
|
||||
|
||||
# List unlockers
|
||||
secret unlocker list
|
||||
@@ -443,7 +454,7 @@ secret decrypt encryption/mykey --input document.txt.age --output document.txt
|
||||
|
||||
### Cross-Platform Support
|
||||
|
||||
- **macOS**: Full support including Keychain and planned Secure Enclave integration
|
||||
- **macOS**: Full support including Keychain and Secure Enclave integration
|
||||
- **Linux**: Full support (excluding macOS-specific features)
|
||||
|
||||
## Security Considerations
|
||||
@@ -485,9 +496,46 @@ go test ./... # Unit tests
|
||||
go test -tags=integration -v ./internal/cli # Integration tests
|
||||
```
|
||||
|
||||
## Entrypoints
|
||||
|
||||
This repository adheres to the
|
||||
[Scripts to Rule Them All](https://github.com/github/scripts-to-rule-them-all)
|
||||
standard: normalized scripts in `script/` are the entrypoints for the
|
||||
development workflow, and the Makefile targets are thin shims that call
|
||||
them. We provide:
|
||||
|
||||
- `script/bootstrap` — install all dependencies (Go, Go module
|
||||
download), idempotently; golangci-lint is not installed, it runs in
|
||||
docker
|
||||
- `script/setup` — make a fresh clone ready for development: runs
|
||||
`script/bootstrap`, then `script/install-precommit`
|
||||
- `script/projectname` — output the project name (`secret`); used by
|
||||
other scripts such as `script/docker`
|
||||
- `script/build` — build the `secret` binary into the repo root, stamping
|
||||
the version (`VERSION` from the environment, else `git describe`) and
|
||||
the git commit
|
||||
- `script/test` — run `go vet` and the test suite (verbose rerun on
|
||||
failure)
|
||||
- `script/lint` — run `golangci-lint` in docker only: builds
|
||||
`Dockerfile.lint`, where the linter is a build step that runs on every
|
||||
call, also on an unchanged tree
|
||||
- `script/fmt` — format all Go code (writes)
|
||||
- `script/fmt-check` — check formatting without writing
|
||||
- `script/check` — run `script/test`, `script/lint`, and
|
||||
`script/fmt-check`
|
||||
- `script/docker` — build the Docker image tagged with the project name
|
||||
- `script/cibuild` — CI entrypoint: `docker build --ulimit
|
||||
memlock=-1:-1 .` (memguard needs mlock; the Dockerfile runs the
|
||||
checks), with a new `CHECK_EPOCH` build argument on every run so the
|
||||
checks run again on an unchanged tree
|
||||
- `script/precommit` — pre-commit checks: `go mod tidy` verification,
|
||||
then `script/check`
|
||||
- `script/install-precommit` — install the git pre-commit hook that
|
||||
runs `script/precommit`
|
||||
|
||||
## Features
|
||||
|
||||
- **Multiple Authentication Methods**: Supports passphrase, PGP, and macOS Keychain unlockers
|
||||
- **Multiple Authentication Methods**: Supports passphrase, PGP, macOS Keychain, and Secure Enclave unlockers
|
||||
- **Vault Isolation**: Complete separation between different vaults
|
||||
- **Per-Secret Encryption**: Each secret has its own encryption key
|
||||
- **BIP39 Mnemonic Support**: Keyless operation using mnemonic phrases
|
||||
|
||||
@@ -0,0 +1,408 @@
|
||||
---
|
||||
title: Repository Policies
|
||||
last_modified: 2026-07-06
|
||||
---
|
||||
|
||||
This document covers repository structure, tooling, and workflow standards. Code
|
||||
style conventions are in separate documents:
|
||||
|
||||
- [Code Styleguide](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/CODE_STYLEGUIDE.md)
|
||||
(general, bash, Docker)
|
||||
- [Go](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/CODE_STYLEGUIDE_GO.md)
|
||||
- [JavaScript](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/CODE_STYLEGUIDE_JS.md)
|
||||
- [Python](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/CODE_STYLEGUIDE_PYTHON.md)
|
||||
- [Go HTTP Server Conventions](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/GO_HTTP_SERVER_CONVENTIONS.md)
|
||||
|
||||
---
|
||||
|
||||
- Cross-project documentation (such as this file) must include
|
||||
`last_modified: YYYY-MM-DD` in the YAML front matter so it can be kept in sync
|
||||
with the authoritative source as policies evolve.
|
||||
|
||||
- **ALL external references must be pinned by cryptographic hash.** This
|
||||
includes Docker base images, Go modules, npm packages, GitHub Actions, and
|
||||
anything else fetched from a remote source. Version tags (`@v4`, `@latest`,
|
||||
`:3.21`, etc.) are server-mutable and therefore remote code execution
|
||||
vulnerabilities. The ONLY acceptable way to reference an external dependency
|
||||
is by its content hash (Docker `@sha256:...`, Go module hash in `go.sum`, npm
|
||||
integrity hash in lockfile, GitHub Actions `@<commit-sha>`). No exceptions.
|
||||
This also means never `curl | bash` to install tools like pyenv, nvm, rustup,
|
||||
etc. Instead, download a specific release archive from GitHub, verify its hash
|
||||
(hardcoded in the Dockerfile or script), and only then install. Unverified
|
||||
install scripts are arbitrary remote code execution. This is the single most
|
||||
important rule in this document. Double-check every external reference in
|
||||
every file before committing. There are zero exceptions to this rule.
|
||||
|
||||
- Every repo with software must have a root `Makefile` with these targets:
|
||||
`make bootstrap`, `make setup`, `make test`, `make lint`, `make fmt` (writes),
|
||||
`make fmt-check` (read-only), `make check` (runs `test`, `lint`, `fmt-check`),
|
||||
`make docker`, and `make hooks` (installs pre-commit hook). A model Makefile
|
||||
is at `https://git.eeqj.de/sneak/prompts/raw/branch/main/Makefile`.
|
||||
|
||||
- Repos follow the
|
||||
[Scripts to Rule Them All](https://github.com/github/scripts-to-rule-them-all)
|
||||
pattern: the implementation of each Makefile target lives in an executable
|
||||
script in `script/` (`script/bootstrap`, `script/setup`, `script/test`,
|
||||
`script/lint`, `script/fmt`, `script/fmt-check`, `script/check`,
|
||||
`script/docker`), and the Makefile targets are thin shims that call them. The
|
||||
scripts must be POSIX sh (`#!/bin/sh`, `set -eu`, no bashisms) so they run in
|
||||
minimal containers (e.g. alpine images have no bash); locate the repo root
|
||||
with `$(cd "$(dirname "$0")/.." && pwd -P)` and `cd` there before acting. From
|
||||
the standard's canonical set we use `bootstrap`, `setup` (make the repo ready
|
||||
for development after a fresh clone: runs `bootstrap`, then
|
||||
`install-precommit`, plus any repo-specific initialization), `test`, and
|
||||
`cibuild`. `script/bootstrap` installs all dependencies idempotently and
|
||||
assumes nothing is present: base tools come from nix, apt, brew, or apk
|
||||
(detected in that order; apt runs noninteractive). For node it uses the
|
||||
installed node if present; otherwise it installs a PINNED node version via
|
||||
nvm, first installing nvm itself if missing — from a hash-verified GitHub
|
||||
release archive (never `curl | sh`), with bash installed as an explicit
|
||||
prerequisite since nvm requires bash. yarn is then pinned via
|
||||
`corepack prepare yarn@<version> --activate`. Never install "latest" or "lts";
|
||||
always exact versions. `script/cibuild` runs the CI build: it changes to the
|
||||
repo root and runs `docker build .`; the Gitea workflow calls it. Four further
|
||||
scripts are our own extensions to the standard: `script/check` runs
|
||||
`script/test`, `script/lint`, and `script/fmt-check`; `script/precommit` is
|
||||
what the git pre-commit hook runs, and it calls `script/check`;
|
||||
`script/install-precommit` installs the git pre-commit hook (the `make hooks`
|
||||
target shims to it); and `script/projectname` (literally that filename) simply
|
||||
outputs the project's name. Scripts that need the name call
|
||||
`script/projectname` — e.g. `script/docker` assembles its image tag from it —
|
||||
so those scripts stay byte-identical across all repos. Repo-type-specific
|
||||
pre-commit extras (e.g. `go mod tidy` verification in Go repos) belong in
|
||||
`script/precommit`, not in the hook itself. Model scripts are at
|
||||
`https://git.eeqj.de/sneak/prompts/raw/branch/main/script/<name>`. The README
|
||||
must document the provided scripts in an **Entrypoints** section (see the
|
||||
README requirements below).
|
||||
|
||||
- Always use Makefile targets (`make fmt`, `make test`, `make lint`, etc.)
|
||||
instead of invoking the underlying tools directly. The Makefile is the single
|
||||
source of truth for how these operations are run.
|
||||
|
||||
- The Makefile is authoritative documentation for how the repo is used. Beyond
|
||||
the required targets above, it should have targets for every common operation:
|
||||
running a local development server (`make run`, `make dev`), re-initializing
|
||||
or migrating the database (`make db-reset`, `make migrate`), building
|
||||
artifacts (`make build`), generating code, seeding data, or anything else a
|
||||
developer would do regularly. If someone checks out the repo and types
|
||||
`make<tab>`, they should see every meaningful operation available. A new
|
||||
contributor should be able to understand the entire development workflow by
|
||||
reading the Makefile.
|
||||
|
||||
- Every repo should have a `Dockerfile`. All Dockerfiles must run `make check`
|
||||
as a build step so the build fails if the branch is not green. For non-server
|
||||
repos, the Dockerfile should bring up a development environment and run
|
||||
`make check`. For server repos, `make check` should run as an early build
|
||||
stage before the final image is assembled. Dockerfiles install development
|
||||
prerequisites by running `script/bootstrap` rather than duplicating installs
|
||||
inline; COPY `script/` and the dependency manifests (`package.json` +
|
||||
`yarn.lock`, `go.mod` + `go.sum`, etc.) before running it so the bootstrap
|
||||
layer stays cached until dependencies change.
|
||||
|
||||
- **Dockerfiles must use a separate lint stage for fail-fast feedback.** Go
|
||||
repos use a multistage build where linting runs in an independent stage based
|
||||
on the `golangci/golangci-lint` image (pinned by hash). This stage runs
|
||||
`make fmt-check` and `make lint` before the full build begins. The build stage
|
||||
then declares an explicit dependency on the lint stage via
|
||||
`COPY --from=lint /src/go.sum /dev/null`, which forces BuildKit to complete
|
||||
linting before proceeding to compilation and tests. This ensures lint failures
|
||||
surface in seconds rather than minutes, without blocking on dependency
|
||||
download or compilation in the build stage.
|
||||
|
||||
The standard pattern for a Go repo Dockerfile is:
|
||||
|
||||
```dockerfile
|
||||
# Lint stage — fast feedback on formatting and lint issues
|
||||
# golangci/golangci-lint:v2.x.x, YYYY-MM-DD
|
||||
FROM golangci/golangci-lint@sha256:... AS lint
|
||||
WORKDIR /src
|
||||
COPY go.mod go.sum ./
|
||||
RUN go mod download
|
||||
COPY . .
|
||||
RUN make fmt-check
|
||||
RUN make lint
|
||||
|
||||
# Build stage
|
||||
# golang:1.x-alpine, YYYY-MM-DD
|
||||
FROM golang@sha256:... AS builder
|
||||
WORKDIR /src
|
||||
|
||||
# Force BuildKit to run the lint stage before proceeding
|
||||
COPY --from=lint /src/go.sum /dev/null
|
||||
|
||||
COPY go.mod go.sum ./
|
||||
RUN go mod download
|
||||
COPY . .
|
||||
RUN make test
|
||||
|
||||
ARG VERSION=dev
|
||||
RUN CGO_ENABLED=0 go build -trimpath \
|
||||
-ldflags="-s -w -X main.Version=${VERSION}" \
|
||||
-o /app ./cmd/app/
|
||||
|
||||
# Runtime stage
|
||||
FROM alpine@sha256:...
|
||||
COPY --from=builder /app /usr/local/bin/app
|
||||
ENTRYPOINT ["app"]
|
||||
```
|
||||
|
||||
Key points:
|
||||
- The lint stage uses the `golangci/golangci-lint` image directly (it
|
||||
includes both Go and the linter), so there is no need to install the
|
||||
linter separately.
|
||||
- `COPY --from=lint /src/go.sum /dev/null` is a no-op file copy that creates
|
||||
a stage dependency. BuildKit runs stages in parallel by default; without
|
||||
this line, the build stage would not wait for lint to finish and a lint
|
||||
failure might not fail the overall build.
|
||||
- If the project uses `//go:embed` directives that reference build artifacts
|
||||
(e.g. a web frontend compiled in a separate stage), the lint stage must
|
||||
create placeholder files so the embed directives resolve. Example:
|
||||
`RUN mkdir -p web/dist && touch web/dist/index.html web/dist/style.css`.
|
||||
The lint stage should not depend on the actual build output — it exists to
|
||||
fail fast.
|
||||
- If the project requires CGO or system libraries for linting (e.g.
|
||||
`vips-dev`), install them in the lint stage with `apk add`.
|
||||
- The build stage runs `make test` after compilation setup. Tests run in the
|
||||
build stage, not the lint stage, because they may require compiled
|
||||
artifacts or heavier dependencies.
|
||||
|
||||
- Every repo should have a Gitea Actions workflow (`.gitea/workflows/`) that
|
||||
runs `script/cibuild` (which runs `docker build .`) on push. Since the
|
||||
Dockerfile already runs `make check`, a successful build implies all checks
|
||||
pass.
|
||||
|
||||
- Use platform-standard formatters: `black` for Python, `prettier` for
|
||||
JS/CSS/Markdown/HTML, `go fmt` for Go. Always use default configuration with
|
||||
two exceptions: four-space indents (except Go), and `proseWrap: always` for
|
||||
Markdown (hard-wrap at 80 columns). Documentation and writing repos (Markdown,
|
||||
HTML, CSS) should also have `.prettierrc` and `.prettierignore`.
|
||||
|
||||
- Pre-commit hook: runs `script/precommit`, which calls `script/check`. If local
|
||||
testing is not possible in the repo, `script/precommit` may skip `script/test`
|
||||
and run only `script/lint` and `script/fmt-check`. The hook is installed by
|
||||
`script/install-precommit`; the Makefile must provide a `make hooks` target
|
||||
that shims to it.
|
||||
|
||||
- All repos with software must have tests that run via the platform-standard
|
||||
test framework (`go test`, `pytest`, `jest`/`vitest`, etc.). If no meaningful
|
||||
tests exist yet, add the most minimal test possible — e.g. importing the
|
||||
module under test to verify it compiles/parses. There is no excuse for
|
||||
`make test` to be a no-op.
|
||||
|
||||
- `make test` must complete in under 20 seconds. Add a 30-second timeout in the
|
||||
Makefile.
|
||||
|
||||
- **`make test` should use the conditional verbose rerun pattern.** Run tests
|
||||
without `-v` (verbose) first. If tests fail, automatically rerun with `-v` to
|
||||
show full output. This keeps CI logs and `docker build` output clean on
|
||||
success (just package/suite summaries) while providing full diagnostic detail
|
||||
on failure (every test case, every assertion). The general shell pattern:
|
||||
|
||||
```makefile
|
||||
test:
|
||||
@<test-command> || \
|
||||
{ echo "--- Rerunning with -v for details ---"; \
|
||||
<test-command-with-v>; exit 1; }
|
||||
```
|
||||
|
||||
Go example:
|
||||
|
||||
```makefile
|
||||
test:
|
||||
@go test -timeout 30s -race -cover ./... || \
|
||||
{ echo "--- Rerunning with -v for details ---"; \
|
||||
go test -timeout 30s -race -v ./...; exit 1; }
|
||||
```
|
||||
|
||||
Python example:
|
||||
|
||||
```makefile
|
||||
test:
|
||||
@python -m pytest || \
|
||||
{ echo "--- Rerunning with -v for details ---"; \
|
||||
python -m pytest -v; exit 1; }
|
||||
```
|
||||
|
||||
The `exit 1` ensures the target always fails after a rerun — the first run
|
||||
already proved the tests are broken, so the build must not pass even if a
|
||||
flaky test happens to succeed on the second attempt. The rerun exists solely
|
||||
for diagnostic output.
|
||||
|
||||
- Docker builds must complete in under 5 minutes.
|
||||
|
||||
- `make check` must not modify any files in the repo. Tests may use temporary
|
||||
directories.
|
||||
|
||||
- `main` must always pass `make check`, no exceptions.
|
||||
|
||||
- Never commit secrets. `.env` files, credentials, API keys, and private keys
|
||||
must be in `.gitignore`. No exceptions.
|
||||
|
||||
- `.gitignore` should be comprehensive from the start: OS files (`.DS_Store`),
|
||||
editor files (`.swp`, `*~`), language build artifacts, and `node_modules/`.
|
||||
Fetch the standard `.gitignore` from
|
||||
`https://git.eeqj.de/sneak/prompts/raw/branch/main/.gitignore` when setting up
|
||||
a new repo.
|
||||
|
||||
- **No build artifacts in version control.** Code-derived data (compiled
|
||||
bundles, minified output, generated assets) must never be committed to the
|
||||
repository if it can be avoided. The build process (e.g. Dockerfile, Makefile)
|
||||
should generate these at build time. Notable exception: Go protobuf generated
|
||||
files (`.pb.go`) ARE committed because repos need to work with `go get`, which
|
||||
downloads code but does not execute code generation.
|
||||
|
||||
- Never use `git add -A` or `git add .`. Always stage files explicitly by name.
|
||||
|
||||
- Never force-push to `main`.
|
||||
|
||||
- Make all changes on a feature branch. You can do whatever you want on a
|
||||
feature branch.
|
||||
|
||||
- `.golangci.yml` is standardized and must _NEVER_ be modified by an agent, only
|
||||
manually by the user. Fetch from
|
||||
`https://git.eeqj.de/sneak/prompts/raw/branch/main/.golangci.yml`.
|
||||
|
||||
- When pinning images or packages by hash, add a comment above the reference
|
||||
with the version and date (YYYY-MM-DD).
|
||||
|
||||
- Use `yarn`, not `npm`.
|
||||
|
||||
- Write all dates as YYYY-MM-DD (ISO 8601).
|
||||
|
||||
- Simple projects should be configured with environment variables.
|
||||
|
||||
- Dockerized web services listen on port 8080 by default, overridable with
|
||||
`PORT`.
|
||||
|
||||
- **HTTP/web services must be hardened for production internet exposure before
|
||||
tagging 1.0.** This means full compliance with security best practices
|
||||
including, without limitation, all of the following:
|
||||
- **Security headers** on every response:
|
||||
- `Strict-Transport-Security` (HSTS) with `max-age` of at least one year
|
||||
and `includeSubDomains`.
|
||||
- `Content-Security-Policy` (CSP) with a restrictive default policy
|
||||
(`default-src 'self'` as a baseline, tightened per-resource as
|
||||
needed). Never use `unsafe-inline` or `unsafe-eval` unless
|
||||
unavoidable, and document the reason.
|
||||
- `X-Frame-Options: DENY` (or `SAMEORIGIN` if framing is required).
|
||||
Prefer the `frame-ancestors` CSP directive as the primary control.
|
||||
- `X-Content-Type-Options: nosniff`.
|
||||
- `Referrer-Policy: strict-origin-when-cross-origin` (or stricter).
|
||||
- `Permissions-Policy` restricting access to browser features the
|
||||
application does not use (camera, microphone, geolocation, etc.).
|
||||
- **Request and response limits:**
|
||||
- Maximum request body size enforced on all endpoints (e.g. Go
|
||||
`http.MaxBytesReader`). Choose a sane default per-route; never accept
|
||||
unbounded input.
|
||||
- Maximum response body size where applicable (e.g. paginated APIs).
|
||||
- `ReadTimeout` and `ReadHeaderTimeout` on the `http.Server` to defend
|
||||
against slowloris attacks.
|
||||
- `WriteTimeout` on the `http.Server`.
|
||||
- `IdleTimeout` on the `http.Server`.
|
||||
- Per-handler execution time limits via `context.WithTimeout` or
|
||||
chi/stdlib `middleware.Timeout`.
|
||||
- **Authentication and session security:**
|
||||
- Rate limiting on password-based authentication endpoints. API keys are
|
||||
high-entropy and not susceptible to brute force, so they are exempt.
|
||||
- CSRF tokens on all state-mutating HTML forms. API endpoints
|
||||
authenticated via `Authorization` header (Bearer token, API key) are
|
||||
exempt because the browser does not attach these automatically.
|
||||
- Passwords stored using bcrypt, scrypt, or argon2 — never plain-text,
|
||||
MD5, or SHA.
|
||||
- Session cookies set with `HttpOnly`, `Secure`, and `SameSite=Lax` (or
|
||||
`Strict`) attributes.
|
||||
- **Reverse proxy awareness:**
|
||||
- True client IP detection when behind a reverse proxy
|
||||
(`X-Forwarded-For`, `X-Real-IP`). The application must accept
|
||||
forwarded headers only from a configured set of trusted proxy
|
||||
addresses — never trust `X-Forwarded-For` unconditionally.
|
||||
- **CORS:**
|
||||
- Authenticated endpoints must restrict `Access-Control-Allow-Origin` to
|
||||
an explicit allowlist of known origins. Wildcard (`*`) is acceptable
|
||||
only for public, unauthenticated read-only APIs.
|
||||
- **Error handling:**
|
||||
- Internal errors must never leak stack traces, SQL queries, file paths,
|
||||
or other implementation details to the client. Return generic error
|
||||
messages in production; detailed errors only when `DEBUG` is enabled.
|
||||
- **TLS:**
|
||||
- Services never terminate TLS directly. They are always deployed behind
|
||||
a TLS-terminating reverse proxy. The service itself listens on plain
|
||||
HTTP. However, HSTS headers and `Secure` cookie flags must still be
|
||||
set by the application so that the browser enforces HTTPS end-to-end.
|
||||
|
||||
This list is non-exhaustive. Apply defense-in-depth: if a standard security
|
||||
hardening measure exists for HTTP services and is not listed here, it is
|
||||
still expected. When in doubt, harden.
|
||||
|
||||
- `README.md` is the primary documentation. Required sections:
|
||||
- **Description**: First line must include the project name, purpose,
|
||||
category (web server, SPA, CLI tool, etc.), license, and author. Example:
|
||||
"µPaaS is an MIT-licensed Go web application by @sneak that receives
|
||||
git-frontend webhooks and deploys applications via Docker in realtime."
|
||||
- **Getting Started**: Copy-pasteable install/usage code block.
|
||||
- **Entrypoints**: Opens by stating that the repo adheres to the
|
||||
[Scripts to Rule Them All](https://github.com/github/scripts-to-rule-them-all)
|
||||
standard (with that link), then documents each provided `script/`
|
||||
entrypoint and its purpose.
|
||||
- **Rationale**: Why does this exist?
|
||||
- **Design**: How is the program structured?
|
||||
- **TODO**: Update meticulously, even between commits. When planning, put
|
||||
the todo list in the README so a new agent can pick up where the last one
|
||||
left off.
|
||||
- **License**: MIT, GPL, or WTFPL. Ask the user for new projects. Include a
|
||||
`LICENSE` file in the repo root and a License section in the README.
|
||||
- **Author**: [@sneak](https://sneak.berlin).
|
||||
|
||||
- First commit of a new repo should contain only `README.md`.
|
||||
|
||||
- Go module root: `sneak.berlin/go/<name>`. Always run `go mod tidy` before
|
||||
committing.
|
||||
|
||||
- Use SemVer.
|
||||
|
||||
- Database migrations live in `internal/db/migrations/` and must be embedded in
|
||||
the binary.
|
||||
- `000_migration.sql` — contains ONLY the creation of the migrations
|
||||
tracking table itself. Nothing else.
|
||||
- `001_schema.sql` — the full application schema.
|
||||
- **Pre-1.0.0:** never add additional migration files (002, 003, etc.).
|
||||
There is no installed base to migrate. Edit `001_schema.sql` directly.
|
||||
- **Post-1.0.0:** add new numbered migration files for each schema change.
|
||||
Never edit existing migrations after release.
|
||||
|
||||
- All repos should have an `.editorconfig` enforcing the project's indentation
|
||||
settings.
|
||||
|
||||
- Avoid putting files in the repo root unless necessary. Root should contain
|
||||
only project-level config files (`README.md`, `Makefile`, `Dockerfile`,
|
||||
`LICENSE`, `.gitignore`, `.editorconfig`, `REPO_POLICIES.md`, and
|
||||
language-specific config). Everything else goes in a subdirectory. Canonical
|
||||
subdirectory names:
|
||||
- `bin/` — executable scripts and tools
|
||||
- `cmd/` — Go command entrypoints
|
||||
- `configs/` — configuration templates and examples
|
||||
- `deploy/` — deployment manifests (k8s, compose, terraform)
|
||||
- `docs/` — documentation and markdown (README.md stays in root)
|
||||
- `internal/` — Go internal packages
|
||||
- `internal/db/migrations/` — database migrations
|
||||
- `pkg/` — Go library packages
|
||||
- `share/` — systemd units, data files
|
||||
- `static/` — static assets (images, fonts, etc.)
|
||||
- `web/` — web frontend source
|
||||
|
||||
- When setting up a new repo, files from the `prompts` repo may be used as
|
||||
templates. Fetch them from
|
||||
`https://git.eeqj.de/sneak/prompts/raw/branch/main/<path>`.
|
||||
|
||||
- New repos must contain at minimum:
|
||||
- `README.md`, `.git`, `.gitignore`, `.editorconfig`
|
||||
- `LICENSE`, `REPO_POLICIES.md` (copy from the `prompts` repo)
|
||||
- `Makefile`
|
||||
- `script/` entrypoints (`bootstrap`, `setup`, `projectname`, `test`,
|
||||
`lint`, `fmt`, `fmt-check`, `check`, `docker`, `cibuild`, `precommit`,
|
||||
`install-precommit`)
|
||||
- `Dockerfile`, `.dockerignore`
|
||||
- `.gitea/workflows/check.yml`
|
||||
- Go: `go.mod`, `go.sum`, `.golangci.yml`
|
||||
- JS: `package.json`, `yarn.lock`, `.prettierrc`, `.prettierignore`
|
||||
- Python: `pyproject.toml`
|
||||
@@ -1,147 +1,305 @@
|
||||
# TODO for 1.0 Release
|
||||
# Workflow
|
||||
|
||||
This document outlines the bugs, issues, and improvements that need to be
|
||||
addressed before the 1.0 release of the secret manager. Items are
|
||||
prioritized from most critical (top) to least critical (bottom).
|
||||
* branch (from `main`)
|
||||
* do the work in Next Step
|
||||
* move Next Step to the top of Completed Steps
|
||||
* move the top item of Future Steps into Next Step
|
||||
* commit (`TODO.md` changes in the same commit as the work)
|
||||
* merge to `main` if the branch is not protected, otherwise open a PR
|
||||
* push
|
||||
|
||||
## CRITICAL BLOCKERS FOR 1.0 RELEASE
|
||||
# Status
|
||||
|
||||
### Command Injection Vulnerabilities
|
||||
- [ ] **1. PGP command injection risk**: `internal/secret/pgpunlocker.go:323-327` - GPG key IDs passed directly to exec.Command without proper escaping
|
||||
- [ ] **2. Keychain command injection risk**: `internal/secret/keychainunlocker.go:472-476` - data.String() passed to security command without escaping
|
||||
pre-1.0. No git tags. TODO.md carries open 1.0 security blockers. Work in
|
||||
flight on branch secure-enclave-unlocker (clean tree as of 2026-07-06).
|
||||
|
||||
### Memory Security Critical Issues
|
||||
- [ ] **3. Plain text passphrase in memory**: `internal/secret/keychainunlocker.go:342,393-396` - KeychainData struct stores AgePrivKeyPassphrase as unprotected string
|
||||
- [ ] **4. Sensitive string conversions**: `internal/secret/keychainunlocker.go:356`, `internal/secret/pgpunlocker.go:256`, `internal/secret/version.go:155` - Age identity .String() creates unprotected copies
|
||||
# Next Step
|
||||
|
||||
### Race Conditions (Data Corruption Risk)
|
||||
- [ ] **5. No file locking mechanism**: `internal/vault/secrets.go:142-176` - Multiple concurrent operations can corrupt vault state
|
||||
- [ ] **6. Non-atomic file operations**: Various locations - Interrupted writes leave vault inconsistent
|
||||
Bring the repo into policy compliance in one commit:
|
||||
|
||||
### Input Validation Vulnerabilities
|
||||
- [ ] **7. Path traversal risk**: `internal/vault/secrets.go:75-99` - Secret names allow dots which could enable traversal attacks with encoding
|
||||
- [ ] **8. Missing size limits**: `internal/vault/secrets.go:102` - No maximum secret size allows DoS via memory exhaustion
|
||||
- Add fmt-check and hooks targets to the Makefile (test/lint/fmt/check/
|
||||
docker already exist).
|
||||
- Add REPO_POLICIES.md and .editorconfig.
|
||||
- Add .gitea/workflows/check.yml running make check.
|
||||
- Verify Dockerfile base images are pinned by sha256.
|
||||
|
||||
### Timing Attack Vulnerabilities
|
||||
- [ ] **9. Non-constant-time passphrase comparison**: `internal/cli/init.go:209-216` - bytes.Equal() vulnerable to timing attacks
|
||||
- [ ] **10. Non-constant-time key validation**: `internal/vault/vault.go:95-100` - Public key comparison leaks timing information
|
||||
# Completed Steps
|
||||
|
||||
## CRITICAL MEMORY SECURITY ISSUES
|
||||
- 2026-10-04: `.golangci.yml` is again the canonical file from
|
||||
`sneak/prompts`, byte for byte
|
||||
(https://git.eeqj.de/sneak/secret/issues/66). It runs `gomodguard_v2`
|
||||
in place of the deprecated `gomodguard`, so the lint no longer warns,
|
||||
and enables `depguard` with a rule that keeps `net/http/httptest` out of
|
||||
non-test files. Neither raised a finding in this repo.
|
||||
- 2026-10-04: `secret unlocker add pgp` works on Linux
|
||||
(https://git.eeqj.de/sneak/secret/issues/88). `CreatePGPUnlocker` gets
|
||||
the vault's long-term key as adding a passphrase unlocker does, with the
|
||||
vault's `GetOrDeriveLongTermKey`, now part of `VaultInterface`: from the
|
||||
mnemonic, checked against the vault, or else from the current unlocker.
|
||||
Before, it used the keychain unlocker's helper, which on every platform
|
||||
but macOS always failed. A test adds a PGP unlocker for a throwaway GPG
|
||||
key, getting the long-term key once from the mnemonic and once from a
|
||||
passphrase unlocker, and reads a secret through the new unlocker.
|
||||
- 2026-10-04: A vault name may use only lowercase ASCII letters, digits,
|
||||
`.`, `-` and `_`, and must not be empty, `.` or `..`
|
||||
(https://git.eeqj.de/sneak/secret/issues/68); the error and `README.md`
|
||||
state the rule. `vault create`, `vault import`, `vault select`,
|
||||
`vault remove`, both vault names of `mv` and shell completion of a
|
||||
`vault:secret` argument check the name as typed with
|
||||
`vault.ValidateVaultName` before building any path from it. Before,
|
||||
`vault import ..` wrote a long-term key and an unlocker into the state
|
||||
directory itself, and `vault select ..` made that the current vault.
|
||||
- 2026-10-04: `script/cibuild` runs the checks again on an unchanged
|
||||
tree (https://git.eeqj.de/sneak/secret/issues/54). It passes the
|
||||
current time as the `CHECK_EPOCH` build argument, which both the lint
|
||||
and the build stage of the `Dockerfile` declare after their module
|
||||
download, so the `RUN` steps below the argument run again on each
|
||||
build while the base images and module downloads stay cached. Before,
|
||||
a second run on the same tree took every check from the build cache
|
||||
and reported success having run nothing.
|
||||
- 2026-10-04: A failed unlocker add no longer leaves a partial unlocker
|
||||
directory (https://git.eeqj.de/sneak/secret/issues/48).
|
||||
`secret unlocker add pgp` resolves the GPG key's fingerprint once, for
|
||||
its duplicate check, and passes it to `CreatePGPUnlocker` to record.
|
||||
`CreatePGPUnlocker` and `CreateKeychainUnlocker` get the long-term key
|
||||
and encrypt everything before writing anything. All four unlocker
|
||||
types write their files through `secret.WriteDir`: a new unlocker is
|
||||
built in a temporary directory, renamed into place when complete and
|
||||
removed on a failure. One added under the directory name of an
|
||||
existing unlocker is still written into that directory in place
|
||||
(https://git.eeqj.de/sneak/secret/issues/71).
|
||||
- 2026-10-04: `secret unlocker select` and `secret unlocker remove`
|
||||
skip, with the warning `unlocker list` gives, an unlocker directory
|
||||
whose metadata file cannot be checked for, read or parsed, instead of
|
||||
failing when it sorts before the unlocker asked for. Such a directory,
|
||||
or one without a metadata file, is removed by its directory name, the
|
||||
name the warning gives; only the directory is removed, since its type
|
||||
is unknown. Removing one whose metadata file is missing or corrupt
|
||||
never counts as removing the last unlocker. Removing one whose metadata
|
||||
file cannot be checked for or read always does, since it may be the
|
||||
only working unlocker, so in a vault with secrets it needs `--force`.
|
||||
- 2026-10-04: A failed command prints its error once, without the usage
|
||||
text after it (https://git.eeqj.de/sneak/secret/issues/41). Usage is
|
||||
still printed for a command called wrongly: wrong number of arguments,
|
||||
unknown flag, bad flag value, missing required flag, or flags that
|
||||
break a flag group (mutually exclusive, required together, one
|
||||
required). The root command's `PersistentPreRunE` turns usage off.
|
||||
Cobra checks arguments and flag values before that hook but required
|
||||
flags and flag groups only after it, so the hook checks those two
|
||||
first. Root `SilenceUsage` would have hidden usage for all of these.
|
||||
- 2026-10-04: `secret get` keeps the secret in locked memory until it
|
||||
writes it out (https://git.eeqj.de/sneak/secret/issues/37):
|
||||
`Vault.GetSecret` and `Vault.GetSecretVersion` return a
|
||||
`*memguard.LockedBuffer`, which every caller destroys, and `secret get`
|
||||
writes its bytes straight to stdout, still with no trailing newline.
|
||||
Before, the value was copied into ordinary memory that nothing wiped,
|
||||
and `get --version` also wrote it to the debug log.
|
||||
- 2026-10-04: The `Makefile` no longer sets `DOCKER_HOST`, so its docker
|
||||
targets use the local docker daemon, or whatever `DOCKER_HOST` the
|
||||
environment sets. `make build` calls the new `script/build`, which
|
||||
stamps the version (`VERSION` from the environment, else
|
||||
`git describe`) and the git commit as before. `build`, `clean`,
|
||||
`install` and `docker-run` are in `.PHONY`; `make install` depends on
|
||||
`build`. The `vet` target is gone: `script/test` runs `go vet` first.
|
||||
- 2026-10-04: `.gitignore` is the org's standard file, which ignores
|
||||
`.env`, `.env.*`, `*.pem` and `*.key` and editor and OS files, plus
|
||||
this repo's `/secret`, `*.log`, `*.test` and `settings.local.json`
|
||||
(https://git.eeqj.de/sneak/secret/issues/40). `.dockerignore` also
|
||||
leaves out `node_modules`; `.git` stays in the build context for the
|
||||
version stamp.
|
||||
- 2026-10-04: `secret init` refuses when the default vault exists, and
|
||||
`secret vault create NAME` when `NAME` does, with "vault NAME already
|
||||
exists", before writing anything. The check is in `vault.CreateVault`,
|
||||
which both commands call while holding the state directory lock, so two
|
||||
creates of one vault at once cannot both pass the check. Before, either
|
||||
command replaced the vault's metadata, passphrase unlocker and
|
||||
`longterm.age`, so none of its secrets could be decrypted any more. Both
|
||||
commands now ask for the unlocker passphrase before creating the vault,
|
||||
so one stopped at that prompt leaves no vault behind.
|
||||
- 2026-10-04: The `internal/cli` tests are back to about their time
|
||||
before the state directory lock
|
||||
(https://git.eeqj.de/sneak/secret/issues/80). The test that each
|
||||
changing command waits for the lock releases it as soon as it sees the
|
||||
command waiting there, instead of after a fixed 100 ms. The two vaults
|
||||
with passphrase unlockers that the path and move tests start from are
|
||||
made once and copied for each test.
|
||||
- 2026-10-04: `secret mv` rejects a move whose destination is the source
|
||||
under another name, such as `foo` for `Foo` on a case-insensitive
|
||||
filesystem (the macOS default) or a name reached through a symbolic
|
||||
link, before changing anything, with or without `--force`, within a
|
||||
vault and between vaults; before, `--force` removed the destination and
|
||||
so deleted the secret. A rename that changes only letter case works on a
|
||||
case-sensitive filesystem as before.
|
||||
- 2026-10-04: Lint runs only in docker: `script/lint` builds
|
||||
`Dockerfile.lint`, where golangci-lint is a build step rebuilt on
|
||||
every run (`--no-cache-filter`), so an unchanged tree is linted too;
|
||||
the module download stays cached. `script/bootstrap` no longer
|
||||
installs golangci-lint, and the `Dockerfile` lint stage calls it
|
||||
directly instead of `make lint`. `golangci-lint config verify` is not
|
||||
run: it fetches its schema live over unpinned HTTPS.
|
||||
- 2026-10-04: A PGP unlocker whose metadata has no usable GPG key ID
|
||||
no longer panics: `GetID()` warns with the unlocker's directory and
|
||||
returns `pgp-unknown`. `ListUnlockers` skips, with a warning, an
|
||||
unlocker whose metadata file cannot be checked for, read or parsed
|
||||
instead of failing, so `secret unlocker list` still lists the others;
|
||||
the listing's ID lookup no longer warns about that directory again.
|
||||
- 2026-10-03: `secret mv` rejects a move whose destination is the
|
||||
source (`mv --force x x`, `mv --force work:x work:`, or an empty
|
||||
destination, which defaults to the source name) before changing
|
||||
anything; before, `--force` removed the destination first and so
|
||||
deleted the secret. Every vault name given with `vault:` must be one
|
||||
of the existing vaults by exact name, so `work:x work/:x` is rejected
|
||||
instead of being taken for a move between two vaults. A move within a
|
||||
named vault no longer makes that vault the current one, whether it
|
||||
succeeds or fails.
|
||||
- 2026-10-03: Commands that change the state directory hold one lock
|
||||
(`flock` on `lock` in the state directory; a mutex on the in-memory
|
||||
test filesystem), so concurrent commands no longer lose versions or
|
||||
race on the current pointers. Every file is written through
|
||||
`secret.WriteFileAtomic` (temporary file, sync, rename), so no file
|
||||
is ever half-written and `current`, `currentvault` and
|
||||
`current-unlocker` never go missing. New versions, new secrets and
|
||||
cross-vault copies are built in a temporary directory and renamed
|
||||
into place, and removals rename out of the way first, so a version
|
||||
or secret is never half-added and never half-removed. An
|
||||
interrupted command can still leave:
|
||||
- a broken unlocker, when it was replacing one: an unlocker added
|
||||
under the directory name of an existing one is rewritten file by
|
||||
file. That happens to a passphrase unlocker added to a vault that
|
||||
has one, and to a PGP, keychain or Secure Enclave unlocker added
|
||||
on the same host and day as another of its type
|
||||
(https://git.eeqj.de/sneak/secret/issues/71);
|
||||
- from `init` or `vault create` killed after the passphrase prompt
|
||||
but before the unlocker is written, a vault with no unlocker,
|
||||
which `vault create` has already made the current vault;
|
||||
- data under a `.tmp-` name in the state directory: a secret,
|
||||
version or unlocker being added, or the secret, version, unlocker
|
||||
or vault being removed, encrypted keys included. Nothing deletes
|
||||
it; it must be deleted by hand
|
||||
(https://git.eeqj.de/sneak/secret/issues/75).
|
||||
- 2026-10-03: The checks run before changing a vault now stop with an
|
||||
error naming the path and cause when they cannot read what they
|
||||
inspect, instead of reading the failure as "nothing there": the
|
||||
duplicate check before `unlocker add pgp` (an unreadable
|
||||
`unlockers.d` or unlocker metadata file), the secret count that
|
||||
guards removing the last unlocker and removing a vault, and the
|
||||
existing long-term key check before `vault import`.
|
||||
- 2026-10-03: `version rm`, `version promote` and `get --version`
|
||||
accept a version only if it is one of the versions `version list`
|
||||
lists for that secret, compared as typed before any path is built
|
||||
(`secret.VersionExists`), and touch nothing otherwise. An empty
|
||||
`--version` is rejected instead of meaning the current version.
|
||||
Before, `secret version rm x ../../..` deleted the whole vault,
|
||||
`secret version rm x ..` the secret, and `.` or `""` every version.
|
||||
- 2026-10-03: Key material is wiped on every exit: `Entry()` returns
|
||||
the exit code after its deferred `memguard.Purge()` has run, and only
|
||||
`main` calls `os.Exit`. SIGINT and SIGTERM go through memguard's
|
||||
handler, which wipes every buffer before exiting; when the process is
|
||||
in the terminal's foreground process group it first restores the
|
||||
terminal settings from startup, so an interrupted passphrase prompt no
|
||||
longer leaves echo off.
|
||||
- 2026-10-03: Every command that builds a path from a secret name
|
||||
checks the name first with `vault.ValidateSecretName` and touches
|
||||
nothing when it is invalid: `rm`, `mv` (both names, within a vault
|
||||
and between vaults, before switching the current vault), `import`,
|
||||
`version list`/`promote`/`rm`, `encrypt` and `decrypt`. The error
|
||||
and `README.md` state the naming rule. Before, `secret rm ..`
|
||||
deleted the whole vault and `secret rm .` every secret in it.
|
||||
- 2026-10-03: The keychain unlocker's age key passphrase stays in
|
||||
locked memory: it is generated into a locked buffer, and the
|
||||
keychain JSON is written and read by `KeychainData` code in
|
||||
`internal/secret/keychaindata.go` (tested on Linux) without
|
||||
`encoding/json` holding it; the JSON field names are unchanged.
|
||||
- 2026-10-02: A plain `docker build .` builds again: the size tests
|
||||
skip a case that needs more locked memory than the process can
|
||||
lock, and run every case under `script/cibuild`. The image stamps the
|
||||
`VERSION` build argument, else `git describe --tags --always`, into
|
||||
`Version`, and fails if `.git` is present but yields no version;
|
||||
`make build` stamps `git describe` too, not a fixed `0.1.0`.
|
||||
`.dockerignore` keeps `.git/config` out; `script/docker` is the
|
||||
canonical copy.
|
||||
- 2026-08-07: Updated golangci-lint to v2.12.2 with the canonical
|
||||
`.golangci.yml` (all linters enabled minus the standard disable
|
||||
list, `lll` 88, tests linted); bumped the `Dockerfile` lint-stage
|
||||
image to the tagged v2.12.2 Debian digest; fixed all ~1550 new
|
||||
findings across `internal/` and `pkg/` (line wrapping, `wsl_v5`
|
||||
blank lines, sentinel errors for `err113`, `t.Parallel()` where
|
||||
safe, `_test` package conversions, complexity/`dupl` helper
|
||||
extraction) on branch `golangci-v2.12.2`. Reworked after review:
|
||||
the `err113` sentinels in `internal/vault`, `internal/secret`,
|
||||
`internal/cli` and `pkg/bip85` were reshaped so every composed
|
||||
error message is byte-identical to `main`, and
|
||||
`findUnlockerIDByMetadata` now returns an error so `unlocker list`
|
||||
skips an unreadable `unlockers.d` entry with a warning instead of
|
||||
emitting a fabricated fallback ID.
|
||||
- 2026-07-07 Adopted scripts-to-rule-them-all: `script/` entrypoints,
|
||||
Makefile shims, README Entrypoints section
|
||||
- 2026-03-11: Secure Enclave unlocker for hardware-backed secret
|
||||
protection, plus review fixes (stub panics, derivation index, tests,
|
||||
README) on branch secure-enclave-unlocker.
|
||||
- 2026-02-28: Repo cleanup, removed stale .cursorrules and coverage.out.
|
||||
- Audit fix wave (issues #1, #2, #3, #13, #14): skip unlockers with
|
||||
missing metadata, allow uppercase secret names, fix hardcoded
|
||||
derivation index, validate names in GetSecretVersion against path
|
||||
traversal, return errors instead of panicking, add Warn() on silent
|
||||
anomalies.
|
||||
- Memory security hardening: LockedBuffer used through encrypt/decrypt
|
||||
paths (Save/EncryptWithPassphrase/GetValue/gpg helpers), deprecated
|
||||
bare-[]byte APIs removed.
|
||||
- Per-secret keypair architecture, vault package refactor, versioning
|
||||
with --version, comprehensive test suite with in-memory filesystem.
|
||||
- Debug logging system (slog, GODEBUG flag, TTY-aware output).
|
||||
- Renamed SEP unlocker to Keychain, reorganized import commands.
|
||||
- 2025-05-28: Initial implementation (vault, age encryption, mnemonic,
|
||||
CLI).
|
||||
|
||||
### Functions accepting bare []byte for sensitive data
|
||||
- [x] **1. Secret.Save accepts unprotected data**: `internal/secret/secret.go:67` - `Save(value []byte, force bool)` - ✓ REMOVED - deprecated function deleted
|
||||
- [x] **2. EncryptWithPassphrase accepts unprotected data**: `internal/secret/crypto.go:73` - `EncryptWithPassphrase(data []byte, passphrase *memguard.LockedBuffer)` - ✓ FIXED - now accepts LockedBuffer for data
|
||||
- [x] **3. storeInKeychain accepts unprotected data**: `internal/secret/keychainunlocker.go:469` - `storeInKeychain(itemName string, data []byte)` - ✓ FIXED - now accepts LockedBuffer for data
|
||||
- [x] **4. gpgEncryptDefault accepts unprotected data**: `internal/secret/pgpunlocker.go:351` - `gpgEncryptDefault(data []byte, keyID string)` - ✓ FIXED - now accepts LockedBuffer for data
|
||||
# Future Steps
|
||||
|
||||
### Functions returning unprotected secrets
|
||||
- [x] **5. GetValue returns unprotected secret**: `internal/secret/secret.go:93` - `GetValue(unlocker Unlocker) ([]byte, error)` - ✓ FIXED - now returns LockedBuffer internally
|
||||
- [x] **6. DecryptWithIdentity returns unprotected data**: `internal/secret/crypto.go:57` - `DecryptWithIdentity(data []byte, identity age.Identity) ([]byte, error)` - ✓ FIXED - now returns LockedBuffer
|
||||
- [x] **7. DecryptWithPassphrase returns unprotected data**: `internal/secret/crypto.go:94` - `DecryptWithPassphrase(encryptedData []byte, passphrase *memguard.LockedBuffer) ([]byte, error)` - ✓ FIXED - now returns LockedBuffer
|
||||
- [x] **8. gpgDecryptDefault returns unprotected data**: `internal/secret/pgpunlocker.go:368` - `gpgDecryptDefault(encryptedData []byte) ([]byte, error)` - ✓ FIXED - now returns LockedBuffer
|
||||
- [x] **9. getSecretValue returns unprotected data**: `internal/cli/crypto.go:269` - `getSecretValue()` returns bare []byte - ✓ ALREADY FIXED - returns LockedBuffer
|
||||
|
||||
### Intermediate string variables for passphrases
|
||||
- [x] **10. Passphrase extracted to string**: `internal/secret/crypto.go:79,100` - `passphraseStr := passphrase.String()` - ✓ UNAVOIDABLE - age library requires string parameter
|
||||
- [ ] **11. Age secret key in plain string**: `internal/cli/crypto.go:86,91,113` - Age secret key stored in plain string variable before conversion back to secure buffer
|
||||
|
||||
### Unprotected buffer.Bytes() usage
|
||||
- [ ] **12. GPG encrypt exposes private key**: `internal/secret/pgpunlocker.go:256` - `GPGEncryptFunc(agePrivateKeyBuffer.Bytes(), gpgKeyID)` - private key exposed to external function
|
||||
- [ ] **13. Keychain encrypt exposes private key**: `internal/secret/keychainunlocker.go:371` - `EncryptWithPassphrase(agePrivKeyBuffer.Bytes(), passphraseBuffer)` - private key passed as bare bytes
|
||||
|
||||
## Code Cleanups
|
||||
|
||||
* we shouldn't be passing around a statedir, it should be read from the
|
||||
environment or default.
|
||||
|
||||
## HIGH PRIORITY SECURITY ISSUES
|
||||
|
||||
- [ ] **4. Application crashes on corrupted metadata**: Code panics instead
|
||||
of returning errors when metadata is corrupt, causing denial of service.
|
||||
Found in pgpunlocker.go:116 and keychainunlocker.go:141.
|
||||
|
||||
- [ ] **5. Insufficient input validation**: Secret names allow potentially
|
||||
dangerous patterns including dots that could enable path traversal attacks
|
||||
(vault/secrets.go:70-93).
|
||||
|
||||
- [ ] **6. Race conditions in file operations**: Multiple concurrent
|
||||
operations could corrupt the vault state due to lack of file locking
|
||||
mechanisms.
|
||||
|
||||
- [ ] **7. Insecure temporary file handling**: Temporary files containing
|
||||
sensitive data may not be properly cleaned up or secured.
|
||||
|
||||
## HIGH PRIORITY FUNCTIONALITY ISSUES
|
||||
|
||||
- [ ] **8. Inappropriate Cobra usage printing**: Commands currently print
|
||||
usage information for all errors, including internal program failures.
|
||||
Usage should only be printed when the user provides incorrect arguments or
|
||||
invalid commands.
|
||||
|
||||
- [ ] **9. Missing current unlock key initialization**: When creating
|
||||
vaults, no default unlock key is selected, which can cause operations to
|
||||
fail.
|
||||
|
||||
- [ ] **10. Add confirmation prompts for destructive operations**:
|
||||
Operations like `keys rm` and vault deletion should require confirmation.
|
||||
|
||||
- [ ] **11. No secret deletion command**: Missing `secret rm <secret-name>`
|
||||
functionality.
|
||||
|
||||
- [ ] **12. Missing vault deletion command**: No way to delete vaults that
|
||||
are no longer needed.
|
||||
|
||||
## MEDIUM PRIORITY ISSUES
|
||||
|
||||
- [ ] **13. Inconsistent error messages**: Error messages need
|
||||
standardization and should be user-friendly. Many errors currently expose
|
||||
internal implementation details.
|
||||
|
||||
- [ ] **14. No graceful handling of corrupted state**: If key files are
|
||||
corrupted or missing, the tool should provide clear error messages and
|
||||
recovery suggestions.
|
||||
|
||||
- [ ] **15. No validation of GPG key existence**: Should verify the
|
||||
specified GPG key exists before creating PGP unlock keys.
|
||||
|
||||
- [ ] **16. Better separation of concerns**: Some functions in CLI do too
|
||||
much and should be split.
|
||||
|
||||
- [ ] **17. Environment variable security**: Sensitive data read from
|
||||
environment variables (SB_UNLOCK_PASSPHRASE, SB_SECRET_MNEMONIC) without
|
||||
proper clearing. Document security implications.
|
||||
|
||||
- [ ] **18. No secure memory allocation**: No use of mlock/munlock to
|
||||
prevent sensitive data from being swapped to disk.
|
||||
|
||||
## LOWER PRIORITY ENHANCEMENTS
|
||||
|
||||
- [ ] **19. Add `--help` examples**: Command help should include practical examples for each operation.
|
||||
|
||||
- [ ] **20. Add shell completion**: Bash/Zsh completion for commands and secret names.
|
||||
|
||||
- [ ] **21. Colored output**: Use colors to improve readability of lists and error messages.
|
||||
|
||||
- [ ] **22. Add `--quiet` flag**: Option to suppress non-essential output.
|
||||
|
||||
- [ ] **23. Smart secret name suggestions**: When a secret name is not found, suggest similar names.
|
||||
|
||||
- [ ] **24. Audit logging**: Log all secret access and modifications for security auditing.
|
||||
|
||||
- [ ] **25. Integration tests for hardware features**: Automated testing of Keychain and GPG functionality.
|
||||
|
||||
- [ ] **26. Consistent naming conventions**: Some variables and functions use inconsistent naming patterns.
|
||||
|
||||
- [ ] **27. Export/import functionality**: Add ability to export/import entire vaults, not just individual secrets.
|
||||
|
||||
- [ ] **28. Batch operations**: Add commands to process multiple secrets at once.
|
||||
|
||||
- [ ] **29. Search functionality**: Add ability to search secret names and potentially contents.
|
||||
|
||||
- [ ] **30. Secret metadata**: Add support for descriptions, tags, or other metadata with secrets.
|
||||
|
||||
## COMPLETED ITEMS ✓
|
||||
|
||||
- [x] **Missing secret history/versioning**: ✓ Implemented - versioning system exists with --version flag support
|
||||
- [x] **XDG compliance on Linux**: ✓ Implemented - uses os.UserConfigDir() which respects XDG_CONFIG_HOME
|
||||
- [x] **Consistent interface implementation**: ✓ Implemented - Unlocker interface is well-defined and consistently implemented
|
||||
- Compliance (after Next Step lands): keep main green under the new
|
||||
.gitea workflow; run make check before every merge.
|
||||
- Implement version-number shell completion for the second arg of
|
||||
`secret version promote` and `secret version rm`
|
||||
(`internal/cli/version.go`; was an in-code TODO removed for godox).
|
||||
- Cover mnemonic-vs-xprv identity consistency in
|
||||
`pkg/agehd/agehd_test.go` `TestMnemonicVsXPRVConsistency` (was an
|
||||
in-code FIXME removed for godox).
|
||||
- Darwin-gated files (`internal/secret/keychainunlocker.go`,
|
||||
`seunlocker_darwin.go`, `internal/macse/macse_darwin.go`, related
|
||||
tests) are not linted on the Linux CI runner and still contain lines
|
||||
over the new 88-column limit; they will surface if lint ever runs on
|
||||
macOS.
|
||||
- Merge secure-enclave-unlocker to main once review is done.
|
||||
- 1.0 critical security blockers (from repo TODO.md):
|
||||
- Command injection: GPG key IDs passed unescaped to exec.Command
|
||||
(pgpunlocker.go:323-327); data.String() passed unescaped to the
|
||||
security command (keychainunlocker.go:472-476).
|
||||
- Memory security: age identity .String() creates unprotected
|
||||
copies (keychainunlocker.go:356, pgpunlocker.go:256,
|
||||
version.go:155); age secret key held in a plain string in
|
||||
cli/crypto.go:86,91,113; private keys exposed via buffer.Bytes()
|
||||
to GPGEncryptFunc and EncryptWithPassphrase.
|
||||
- Input validation: no maximum secret size (DoS).
|
||||
- Timing attacks: bytes.Equal passphrase compare (cli/init.go:
|
||||
209-216); non-constant-time public key compare (vault.go:95-100).
|
||||
- High priority:
|
||||
- Secure temporary file handling and cleanup.
|
||||
- Initialize a default unlock key at vault creation.
|
||||
- Confirmation prompts for destructive operations (keys rm, vault
|
||||
deletion).
|
||||
- Add secret rm and vault deletion commands.
|
||||
- Medium priority:
|
||||
- Standardize error messages; stop leaking internals.
|
||||
- Graceful handling of corrupted or missing key files with recovery
|
||||
suggestions.
|
||||
- Validate GPG key existence before creating PGP unlock keys.
|
||||
- Split oversized CLI functions.
|
||||
- Document env var security (SB_UNLOCK_PASSPHRASE,
|
||||
SB_SECRET_MNEMONIC); clear after use.
|
||||
- mlock/munlock for sensitive allocations.
|
||||
- Cleanups: read statedir from environment or default instead of
|
||||
passing it around.
|
||||
- Enhancements: help examples, shell completion, colored output,
|
||||
--quiet flag, name suggestions on miss, audit logging, hardware
|
||||
integration tests (Keychain, GPG), naming consistency, vault
|
||||
export/import, batch operations, search, secret metadata
|
||||
(descriptions, tags).
|
||||
|
||||
+6
-2
@@ -1,8 +1,12 @@
|
||||
// Package main is the entry point for the secret CLI application.
|
||||
package main
|
||||
|
||||
import "git.eeqj.de/sneak/secret/internal/cli"
|
||||
import (
|
||||
"os"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/cli"
|
||||
)
|
||||
|
||||
func main() {
|
||||
cli.Entry()
|
||||
os.Exit(cli.Entry())
|
||||
}
|
||||
|
||||
-102
@@ -1,102 +0,0 @@
|
||||
mode: set
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:57.41,60.38 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:60.38,61.41 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:65.2,70.3 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:74.50,76.2 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:79.85,81.28 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:81.28,83.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:86.2,87.16 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:87.16,89.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:92.2,93.16 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:93.16,95.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:98.2,98.35 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:102.89,105.16 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:105.16,107.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:110.2,114.21 4 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:118.99,119.46 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:119.46,121.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:124.2,134.39 5 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:134.39,137.15 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:137.15,140.4 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:143.3,145.17 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:145.17,147.4 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:150.3,150.15 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:150.15,152.4 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:155.3,156.17 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:156.17,158.4 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:160.3,160.14 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:163.2,163.17 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:167.107,171.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:171.16,173.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:177.2,186.15 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:187.15,188.13 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:189.15,190.13 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:191.15,192.13 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:193.15,194.13 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:195.15,196.13 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:197.10,198.64 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:202.2,204.21 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:208.84,212.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:212.16,214.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:217.2,222.16 4 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:222.16,224.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:226.2,226.26 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:230.99,234.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:234.16,236.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:239.2,251.45 6 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:251.45,253.3 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:256.2,275.45 12 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:279.39,284.2 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:287.91,288.36 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:288.36,290.3 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:292.2,295.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:295.16,297.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:300.2,302.41 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:306.100,307.32 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:307.32,309.3 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:311.2,314.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:314.16,316.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:319.2,325.35 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:325.35,327.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:329.2,329.33 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:333.100,334.32 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:334.32,336.3 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:338.2,341.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:341.16,343.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:346.2,349.32 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:349.32,351.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:353.2,353.30 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:357.57,375.52 7 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:375.52,381.46 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:381.46,385.4 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:387.3,387.20 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:390.2,390.21 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/bip85/bip85.go:394.67,396.2 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:32.22,36.2 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:40.67,41.31 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:41.31,43.3 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:46.2,55.16 6 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:55.16,57.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:58.2,59.16 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:59.16,61.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:63.2,63.52 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:68.63,74.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:74.16,76.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:79.2,83.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:83.16,85.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:88.2,91.16 4 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:91.16,93.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:95.2,95.17 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:100.67,103.16 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:103.16,105.3 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:108.2,112.16 3 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:112.16,114.3 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:117.2,120.16 4 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:120.16,122.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:124.2,124.17 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:129.77,131.16 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:131.16,133.3 1 0
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:135.2,135.33 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:140.81,142.16 2 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:142.16,144.3 1 1
|
||||
git.eeqj.de/sneak/secret/pkg/agehd/agehd.go:146.2,146.33 1 1
|
||||
@@ -16,6 +16,7 @@ require (
|
||||
github.com/stretchr/testify v1.8.4
|
||||
github.com/tyler-smith/go-bip39 v1.1.0
|
||||
golang.org/x/crypto v0.38.0
|
||||
golang.org/x/sys v0.33.0
|
||||
golang.org/x/term v0.32.0
|
||||
)
|
||||
|
||||
@@ -31,7 +32,6 @@ require (
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/spf13/pflag v1.0.6 // indirect
|
||||
golang.org/x/sys v0.33.0 // indirect
|
||||
golang.org/x/text v0.25.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
)
|
||||
|
||||
+17
-13
@@ -17,27 +17,36 @@ type Instance struct {
|
||||
}
|
||||
|
||||
// NewCLIInstance creates a new CLI instance with the real filesystem
|
||||
func NewCLIInstance() *Instance {
|
||||
func NewCLIInstance() (*Instance, error) {
|
||||
fs := afero.NewOsFs()
|
||||
stateDir := secret.DetermineStateDir("")
|
||||
|
||||
stateDir, err := secret.DetermineStateDir("")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cannot determine state directory: %w", err)
|
||||
}
|
||||
|
||||
return &Instance{
|
||||
fs: fs,
|
||||
stateDir: stateDir,
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
// NewCLIInstanceWithFs creates a new CLI instance with the given filesystem (for testing)
|
||||
func NewCLIInstanceWithFs(fs afero.Fs) *Instance {
|
||||
stateDir := secret.DetermineStateDir("")
|
||||
// NewCLIInstanceWithFs creates a new CLI instance with the given
|
||||
// filesystem (for testing)
|
||||
func NewCLIInstanceWithFs(fs afero.Fs) (*Instance, error) {
|
||||
stateDir, err := secret.DetermineStateDir("")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cannot determine state directory: %w", err)
|
||||
}
|
||||
|
||||
return &Instance{
|
||||
fs: fs,
|
||||
stateDir: stateDir,
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
// NewCLIInstanceWithStateDir creates a new CLI instance with custom state directory (for testing)
|
||||
// NewCLIInstanceWithStateDir creates a new CLI instance with custom state
|
||||
// directory (for testing)
|
||||
func NewCLIInstanceWithStateDir(fs afero.Fs, stateDir string) *Instance {
|
||||
return &Instance{
|
||||
fs: fs,
|
||||
@@ -59,8 +68,3 @@ func (cli *Instance) SetStateDir(stateDir string) {
|
||||
func (cli *Instance) GetStateDir() string {
|
||||
return cli.stateDir
|
||||
}
|
||||
|
||||
// Print outputs to the command's configured output writer
|
||||
func (cli *Instance) Print(a ...interface{}) (n int, err error) {
|
||||
return fmt.Fprint(cli.cmd.OutOrStdout(), a...)
|
||||
}
|
||||
|
||||
@@ -1,34 +1,43 @@
|
||||
package cli
|
||||
package cli_test
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/cli"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
func TestCLIInstanceStateDir(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Test the CLI instance state directory functionality
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
// Create a test state directory
|
||||
testStateDir := "/test-state-dir"
|
||||
cli := NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
instance := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
|
||||
if cli.GetStateDir() != testStateDir {
|
||||
t.Errorf("Expected state directory %q, got %q", testStateDir, cli.GetStateDir())
|
||||
got := instance.GetStateDir()
|
||||
if got != testStateDir {
|
||||
t.Errorf("Expected state directory %q, got %q", testStateDir, got)
|
||||
}
|
||||
}
|
||||
|
||||
//nolint:paralleltest // reads process environment to determine the state dir
|
||||
func TestCLIInstanceWithFs(t *testing.T) {
|
||||
// Test creating CLI instance with custom filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
cli := NewCLIInstanceWithFs(fs)
|
||||
|
||||
instance, err := cli.NewCLIInstanceWithFs(fs)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
// The state directory should be determined automatically
|
||||
stateDir := cli.GetStateDir()
|
||||
stateDir := instance.GetStateDir()
|
||||
if stateDir == "" {
|
||||
t.Error("Expected non-empty state directory")
|
||||
}
|
||||
@@ -41,7 +50,11 @@ func TestDetermineStateDir(t *testing.T) {
|
||||
testEnvDir := "/test-env-dir"
|
||||
t.Setenv(secret.EnvStateDir, testEnvDir)
|
||||
|
||||
stateDir := secret.DetermineStateDir("")
|
||||
stateDir, err := secret.DetermineStateDir("")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if stateDir != testEnvDir {
|
||||
t.Errorf("Expected state directory %q from environment, got %q", testEnvDir, stateDir)
|
||||
}
|
||||
@@ -49,9 +62,15 @@ func TestDetermineStateDir(t *testing.T) {
|
||||
// Test with custom config dir
|
||||
_ = os.Unsetenv(secret.EnvStateDir)
|
||||
customConfigDir := "/custom-config"
|
||||
stateDir = secret.DetermineStateDir(customConfigDir)
|
||||
|
||||
stateDir, err = secret.DetermineStateDir(customConfigDir)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
expectedDir := filepath.Join(customConfigDir, secret.AppID)
|
||||
if stateDir != expectedDir {
|
||||
t.Errorf("Expected state directory %q with custom config, got %q", expectedDir, stateDir)
|
||||
t.Errorf("Expected state directory %q with custom config, got %q",
|
||||
expectedDir, stateDir)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,12 +1,16 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
// errUnsupportedShell is returned for unknown shell completion targets
|
||||
var errUnsupportedShell = errors.New("unsupported shell type")
|
||||
|
||||
func newCompletionCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "completion [bash|zsh|fish|powershell]",
|
||||
@@ -55,7 +59,7 @@ PowerShell:
|
||||
case "powershell":
|
||||
return cmd.Root().GenPowerShellCompletionWithDesc(os.Stdout)
|
||||
default:
|
||||
return fmt.Errorf("unsupported shell type: %s", args[0])
|
||||
return fmt.Errorf("%w: %s", errUnsupportedShell, args[0])
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
+104
-95
@@ -1,7 +1,6 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
@@ -11,11 +10,14 @@ import (
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
// getSecretNamesCompletionFunc returns a completion function that provides secret names
|
||||
// getSecretNamesCompletionFunc returns a completion function that provides
|
||||
// secret names
|
||||
func getSecretNamesCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
cmd *cobra.Command, args []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
return func(_ *cobra.Command, _ []string, toComplete string) ([]string, cobra.ShellCompDirective) {
|
||||
return func(
|
||||
_ *cobra.Command, _ []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
// Get current vault
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
@@ -30,6 +32,7 @@ func getSecretNamesCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
|
||||
// Filter secrets based on what user has typed
|
||||
var completions []string
|
||||
|
||||
for _, secret := range secrets {
|
||||
if strings.HasPrefix(secret, toComplete) {
|
||||
completions = append(completions, secret)
|
||||
@@ -40,11 +43,14 @@ func getSecretNamesCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
}
|
||||
}
|
||||
|
||||
// getUnlockerIDsCompletionFunc returns a completion function that provides unlocker IDs
|
||||
// getUnlockerIDsCompletionFunc returns a completion function that provides
|
||||
// unlocker IDs
|
||||
func getUnlockerIDsCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
cmd *cobra.Command, args []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
return func(_ *cobra.Command, _ []string, toComplete string) ([]string, cobra.ShellCompDirective) {
|
||||
return func(
|
||||
_ *cobra.Command, _ []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
// Get current vault
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
@@ -66,55 +72,24 @@ func getUnlockerIDsCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
// Collect unlocker IDs
|
||||
var completions []string
|
||||
|
||||
unlockersDir := filepath.Join(vaultDir, "unlockers.d")
|
||||
|
||||
for _, metadata := range unlockerMetadataList {
|
||||
// Get the actual unlocker ID by creating the unlocker instance
|
||||
unlockersDir := filepath.Join(vaultDir, "unlockers.d")
|
||||
files, err := afero.ReadDir(fs, unlockersDir)
|
||||
id, err := findUnlockerIDByMetadata(
|
||||
fs, unlockersDir, metadata, false,
|
||||
)
|
||||
if err != nil {
|
||||
secret.Warn(
|
||||
"Could not read unlockers directory during completion, "+
|
||||
"skipping unlocker",
|
||||
"unlockers_dir", unlockersDir, "error", err)
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
for _, file := range files {
|
||||
if !file.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
unlockerDir := filepath.Join(unlockersDir, file.Name())
|
||||
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json")
|
||||
|
||||
// Check if this is the right unlocker by comparing metadata
|
||||
metadataBytes, err := afero.ReadFile(fs, metadataPath)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
var diskMetadata secret.UnlockerMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &diskMetadata); err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
// Match by type and creation time
|
||||
if diskMetadata.Type == metadata.Type && diskMetadata.CreatedAt.Equal(metadata.CreatedAt) {
|
||||
// Create the appropriate unlocker instance
|
||||
var unlocker secret.Unlocker
|
||||
switch metadata.Type {
|
||||
case "passphrase":
|
||||
unlocker = secret.NewPassphraseUnlocker(fs, unlockerDir, diskMetadata)
|
||||
case "keychain":
|
||||
unlocker = secret.NewKeychainUnlocker(fs, unlockerDir, diskMetadata)
|
||||
case "pgp":
|
||||
unlocker = secret.NewPGPUnlocker(fs, unlockerDir, diskMetadata)
|
||||
}
|
||||
|
||||
if unlocker != nil {
|
||||
id := unlocker.GetID()
|
||||
if strings.HasPrefix(id, toComplete) {
|
||||
completions = append(completions, id)
|
||||
}
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
if id != "" && strings.HasPrefix(id, toComplete) {
|
||||
completions = append(completions, id)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -122,17 +97,21 @@ func getUnlockerIDsCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
}
|
||||
}
|
||||
|
||||
// getVaultNamesCompletionFunc returns a completion function that provides vault names
|
||||
// getVaultNamesCompletionFunc returns a completion function that provides
|
||||
// vault names
|
||||
func getVaultNamesCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
cmd *cobra.Command, args []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
return func(_ *cobra.Command, _ []string, toComplete string) ([]string, cobra.ShellCompDirective) {
|
||||
return func(
|
||||
_ *cobra.Command, _ []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
vaults, err := vault.ListVaults(fs, stateDir)
|
||||
if err != nil {
|
||||
return nil, cobra.ShellCompDirectiveNoFileComp
|
||||
}
|
||||
|
||||
var completions []string
|
||||
|
||||
for _, v := range vaults {
|
||||
if strings.HasPrefix(v, toComplete) {
|
||||
completions = append(completions, v)
|
||||
@@ -143,57 +122,87 @@ func getVaultNamesCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
}
|
||||
}
|
||||
|
||||
// getVaultSecretCompletionFunc returns a completion function for vault:secret format
|
||||
// It completes vault names with ":" suffix, and after ":" it completes secrets from that vault
|
||||
// completeVaultQualifiedSecrets completes "vault:secret" references once a
|
||||
// colon is present in the input. It completes nothing when the vault part
|
||||
// is not a valid vault name, so that a name such as ".." cannot list a
|
||||
// directory outside vaults.d.
|
||||
func completeVaultQualifiedSecrets(
|
||||
fs afero.Fs, stateDir, toComplete string,
|
||||
) []string {
|
||||
var completions []string
|
||||
|
||||
// Complete secret names for the specified vault
|
||||
parts := strings.SplitN(toComplete, ":", vaultSecretParts)
|
||||
vaultName := parts[0]
|
||||
secretPrefix := parts[1]
|
||||
|
||||
if vault.ValidateVaultName(vaultName) != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
vlt := vault.NewVault(fs, stateDir, vaultName)
|
||||
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err == nil {
|
||||
for _, secretName := range secrets {
|
||||
if strings.HasPrefix(secretName, secretPrefix) {
|
||||
completions = append(completions, vaultName+":"+secretName)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return completions
|
||||
}
|
||||
|
||||
// completeUnqualifiedVaultSecrets completes vault names (with a ":"
|
||||
// suffix) and secrets from the current vault
|
||||
func completeUnqualifiedVaultSecrets(
|
||||
fs afero.Fs, stateDir, toComplete string,
|
||||
) []string {
|
||||
var completions []string
|
||||
|
||||
// Complete vault names with ":" suffix
|
||||
vaults, err := vault.ListVaults(fs, stateDir)
|
||||
if err == nil {
|
||||
for _, v := range vaults {
|
||||
if strings.HasPrefix(v, toComplete) {
|
||||
completions = append(completions, v+":")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Also complete secrets from current vault (for within-vault moves)
|
||||
currentVlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
if err == nil {
|
||||
secrets, err := currentVlt.ListSecrets()
|
||||
if err == nil {
|
||||
for _, secretName := range secrets {
|
||||
if strings.HasPrefix(secretName, toComplete) {
|
||||
completions = append(completions, secretName)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return completions
|
||||
}
|
||||
|
||||
// getVaultSecretCompletionFunc returns a completion function for the
|
||||
// vault:secret format. It completes vault names with ":" suffix, and
|
||||
// after ":" it completes secrets from that vault.
|
||||
func getVaultSecretCompletionFunc(fs afero.Fs, stateDir string) func(
|
||||
cmd *cobra.Command, args []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
return func(_ *cobra.Command, _ []string, toComplete string) ([]string, cobra.ShellCompDirective) {
|
||||
var completions []string
|
||||
|
||||
return func(
|
||||
_ *cobra.Command, _ []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
// Check if we're completing after a vault: prefix
|
||||
if strings.Contains(toComplete, ":") {
|
||||
// Complete secret names for the specified vault
|
||||
const vaultSecretParts = 2
|
||||
parts := strings.SplitN(toComplete, ":", vaultSecretParts)
|
||||
vaultName := parts[0]
|
||||
secretPrefix := parts[1]
|
||||
|
||||
vlt := vault.NewVault(fs, stateDir, vaultName)
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err == nil {
|
||||
for _, secretName := range secrets {
|
||||
if strings.HasPrefix(secretName, secretPrefix) {
|
||||
completions = append(completions, vaultName+":"+secretName)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return completions, cobra.ShellCompDirectiveNoFileComp
|
||||
return completeVaultQualifiedSecrets(fs, stateDir, toComplete),
|
||||
cobra.ShellCompDirectiveNoFileComp
|
||||
}
|
||||
|
||||
// Complete vault names with ":" suffix
|
||||
vaults, err := vault.ListVaults(fs, stateDir)
|
||||
if err == nil {
|
||||
for _, v := range vaults {
|
||||
if strings.HasPrefix(v, toComplete) {
|
||||
completions = append(completions, v+":")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Also complete secrets from current vault (for within-vault moves)
|
||||
if currentVlt, err := vault.GetCurrentVault(fs, stateDir); err == nil {
|
||||
secrets, err := currentVlt.ListSecrets()
|
||||
if err == nil {
|
||||
for _, secretName := range secrets {
|
||||
if strings.HasPrefix(secretName, toComplete) {
|
||||
completions = append(completions, secretName)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return completions, cobra.ShellCompDirectiveNoSpace
|
||||
return completeUnqualifiedVaultSecrets(fs, stateDir, toComplete),
|
||||
cobra.ShellCompDirectiveNoSpace
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package cli
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestVaultSecretCompletionRejectsInvalidVaultName is a regression test for
|
||||
// https://git.eeqj.de/sneak/secret/issues/68: completing a `vault:secret`
|
||||
// argument lists nothing when the vault part is not a valid vault name, even
|
||||
// where that name, joined onto vaults.d, leads to a secrets.d directory.
|
||||
func TestVaultSecretCompletionRejectsInvalidVaultName(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const (
|
||||
stateDir = "/state"
|
||||
dirPerm = 0o700
|
||||
)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
// The vault "work" holds the secret "x". So does every directory an
|
||||
// invalid name below would lead to from vaults.d.
|
||||
for _, vaultName := range []string{"work", ".", "..", "a/b"} {
|
||||
secretDir := filepath.Join(stateDir, "vaults.d", vaultName, "secrets.d", "x")
|
||||
require.NoError(t, fs.MkdirAll(secretDir, dirPerm))
|
||||
}
|
||||
|
||||
assert.Equal(t, []string{"work:x"},
|
||||
completeVaultQualifiedSecrets(fs, stateDir, "work:"))
|
||||
|
||||
for _, toComplete := range []string{".:", "..:", "a/b:"} {
|
||||
assert.Empty(t, completeVaultQualifiedSecrets(fs, stateDir, toComplete),
|
||||
"completing %q", toComplete)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,153 @@
|
||||
package cli_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/cli"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestCreateExistingVaultChangesNothing is a regression test for
|
||||
// https://git.eeqj.de/sneak/secret/issues/74, where running `secret init`
|
||||
// a second time, or `secret vault create` with the name of an existing
|
||||
// vault, replaced that vault's keys, so that none of its secrets could be
|
||||
// decrypted any more. Each must refuse, change nothing, and leave every
|
||||
// vault's secret readable through its passphrase unlocker.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids parallel subtests
|
||||
func TestCreateExistingVaultChangesNothing(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
// `secret init`, `secret vault create work`, `secret vault select
|
||||
// default`, and the secret "x" in each vault. "work" is then not the
|
||||
// current vault, which creating it again must not change.
|
||||
fs := afero.NewMemMapFs()
|
||||
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
cmd := &cobra.Command{}
|
||||
|
||||
require.NoError(t, c.Init(cmd))
|
||||
require.NoError(t, c.CreateVault(cmd, "work"))
|
||||
require.NoError(t, c.SelectVault(cmd, "default"))
|
||||
|
||||
vaults, err := vault.ListVaults(fs, testStateDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, vaults, 2)
|
||||
|
||||
for _, name := range vaults {
|
||||
value := memguard.NewBufferFromBytes([]byte("value"))
|
||||
err := vault.NewVault(fs, testStateDir, name).AddSecret("x", value, false)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
before := snapshotStateDir(t, fs)
|
||||
|
||||
tests := []struct {
|
||||
command string
|
||||
want string
|
||||
run func(c *cli.Instance) error
|
||||
}{
|
||||
{
|
||||
"init",
|
||||
"failed to create default vault: vault default already exists",
|
||||
func(c *cli.Instance) error { return c.Init(cmd) },
|
||||
},
|
||||
{
|
||||
"vault create default",
|
||||
"vault default already exists",
|
||||
func(c *cli.Instance) error { return c.CreateVault(cmd, "default") },
|
||||
},
|
||||
{
|
||||
"vault create work",
|
||||
"vault work already exists",
|
||||
func(c *cli.Instance) error { return c.CreateVault(cmd, "work") },
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.command, func(t *testing.T) {
|
||||
fs := newFsFromSnapshot(t, before)
|
||||
|
||||
err := tt.run(cli.NewCLIInstanceWithStateDir(fs, testStateDir))
|
||||
|
||||
require.EqualError(t, err, tt.want)
|
||||
require.Equal(t, before, snapshotStateDir(t, fs))
|
||||
})
|
||||
}
|
||||
|
||||
// Every case left the state directory exactly as recorded in before, so
|
||||
// reading each vault's secret once from it shows that it still decrypts
|
||||
// after each case. Without the mnemonic, reading a secret goes through
|
||||
// the vault's passphrase unlocker, which is slow.
|
||||
t.Setenv(secret.EnvMnemonic, "")
|
||||
|
||||
for _, name := range vaults {
|
||||
value, err := vault.NewVault(fs, testStateDir, name).GetSecret("x")
|
||||
require.NoError(t, err)
|
||||
|
||||
unchanged := bytes.Equal([]byte("value"), value.Bytes())
|
||||
value.Destroy()
|
||||
|
||||
require.True(t, unchanged, "vault %q kept its secret", name)
|
||||
}
|
||||
}
|
||||
|
||||
// TestStopAtPassphrasePromptLeavesNothing is a regression test for the
|
||||
// review of https://git.eeqj.de/sneak/secret/pulls/82: `secret init` or
|
||||
// `secret vault create` stopped at the passphrase prompt left a vault with
|
||||
// no unlocker, which neither command would then create again. Each must ask
|
||||
// for the passphrase before writing anything.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids parallel subtests
|
||||
func TestStopAtPassphrasePromptLeavesNothing(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
// Without the passphrase in the environment, both commands prompt for
|
||||
// it, which fails because the tests do not run in a terminal.
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "")
|
||||
|
||||
// An empty state directory for `secret init`, and one holding the vault
|
||||
// "default" for `secret vault create work`.
|
||||
empty := afero.NewMemMapFs()
|
||||
require.NoError(t, empty.MkdirAll(testStateDir, secret.DirPerms))
|
||||
|
||||
withDefault := afero.NewMemMapFs()
|
||||
_, err := vault.CreateVault(withDefault, testStateDir, "default")
|
||||
require.NoError(t, err)
|
||||
|
||||
cmd := &cobra.Command{}
|
||||
|
||||
tests := []struct {
|
||||
command string
|
||||
fs afero.Fs
|
||||
run func(c *cli.Instance) error
|
||||
}{
|
||||
{
|
||||
"init",
|
||||
empty,
|
||||
func(c *cli.Instance) error { return c.Init(cmd) },
|
||||
},
|
||||
{
|
||||
"vault create work",
|
||||
withDefault,
|
||||
func(c *cli.Instance) error { return c.CreateVault(cmd, "work") },
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.command, func(t *testing.T) {
|
||||
before := snapshotStateDir(t, tt.fs)
|
||||
|
||||
err := tt.run(cli.NewCLIInstanceWithStateDir(tt.fs, testStateDir))
|
||||
|
||||
require.ErrorContains(t, err, "failed to read passphrase")
|
||||
require.Equal(t, before, snapshotStateDir(t, tt.fs))
|
||||
})
|
||||
}
|
||||
}
|
||||
+147
-85
@@ -1,6 +1,7 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
@@ -12,20 +13,35 @@ import (
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newEncryptCmd() *cobra.Command {
|
||||
// Sentinel errors for encrypt/decrypt operations
|
||||
var (
|
||||
errNotAgeSecretKey = errors.New(
|
||||
"does not contain a valid age secret key")
|
||||
errSecretDoesNotExist = errors.New("does not exist")
|
||||
)
|
||||
|
||||
// newCryptoCmd builds an encrypt/decrypt command with input/output flags
|
||||
func newCryptoCmd(
|
||||
use, short, long string,
|
||||
run func(cli *Instance, secretName, inputFile, outputFile string) error,
|
||||
) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "encrypt <secret-name>",
|
||||
Short: "Encrypt data using an age secret key stored in a secret",
|
||||
Long: `Encrypt data using an age secret key. If the secret doesn't exist, a new age key is generated and stored.`,
|
||||
Use: use,
|
||||
Short: short,
|
||||
Long: long,
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
inputFile, _ := cmd.Flags().GetString("input")
|
||||
outputFile, _ := cmd.Flags().GetString("output")
|
||||
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize CLI: %w", err)
|
||||
}
|
||||
|
||||
cli.cmd = cmd
|
||||
|
||||
return cli.Encrypt(args[0], inputFile, outputFile)
|
||||
return run(cli, args[0], inputFile, outputFile)
|
||||
},
|
||||
}
|
||||
|
||||
@@ -35,86 +51,120 @@ func newEncryptCmd() *cobra.Command {
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newEncryptCmd() *cobra.Command {
|
||||
return newCryptoCmd(
|
||||
"encrypt <secret-name>",
|
||||
"Encrypt data using an age secret key stored in a secret",
|
||||
"Encrypt data using an age secret key. If the secret doesn't "+
|
||||
"exist, a new age key is generated and stored.",
|
||||
(*Instance).Encrypt,
|
||||
)
|
||||
}
|
||||
|
||||
func newDecryptCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "decrypt <secret-name>",
|
||||
Short: "Decrypt data using an age secret key stored in a secret",
|
||||
Long: `Decrypt data using an age secret key stored in the specified secret.`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
inputFile, _ := cmd.Flags().GetString("input")
|
||||
outputFile, _ := cmd.Flags().GetString("output")
|
||||
return newCryptoCmd(
|
||||
"decrypt <secret-name>",
|
||||
"Decrypt data using an age secret key stored in a secret",
|
||||
"Decrypt data using an age secret key stored in the specified secret.",
|
||||
(*Instance).Decrypt,
|
||||
)
|
||||
}
|
||||
|
||||
cli := NewCLIInstance()
|
||||
cli.cmd = cmd
|
||||
// storeNewEncryptionKey generates an age secret key and stores it as the
|
||||
// named secret, holding the state directory lock while it does. It fails
|
||||
// with vault.ErrSecretExists if another command stored the secret first.
|
||||
// The caller must destroy the returned buffer.
|
||||
func (cli *Instance) storeNewEncryptionKey(
|
||||
vlt *vault.Vault, secretName string,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
return cli.Decrypt(args[0], inputFile, outputFile)
|
||||
},
|
||||
identity, err := age.GenerateX25519Identity()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate age key: %w", err)
|
||||
}
|
||||
|
||||
cmd.Flags().StringP("input", "i", "", "Input file (default: stdin)")
|
||||
cmd.Flags().StringP("output", "o", "", "Output file (default: stdout)")
|
||||
// Store the generated key directly in a secure buffer
|
||||
secureBuffer := memguard.NewBufferFromBytes([]byte(identity.String()))
|
||||
|
||||
return cmd
|
||||
err = vlt.AddSecret(secretName, secureBuffer, false)
|
||||
if err != nil {
|
||||
secureBuffer.Destroy()
|
||||
|
||||
return nil, fmt.Errorf("failed to store age key: %w", err)
|
||||
}
|
||||
|
||||
return secureBuffer, nil
|
||||
}
|
||||
|
||||
// resolveEncryptionKey returns a secure buffer holding the age secret key
|
||||
// for the named secret, generating and storing a new key if the secret
|
||||
// does not exist. The caller must destroy the returned buffer. Only storing
|
||||
// a new key takes the state directory lock, so that reading an existing key
|
||||
// works on a read-only state directory and keeps no other command waiting
|
||||
// at the passphrase prompt, and Encrypt streams its input and output
|
||||
// unlocked.
|
||||
func (cli *Instance) resolveEncryptionKey(
|
||||
vlt *vault.Vault, secretName string,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
// Check if secret exists
|
||||
secretObj := secret.NewSecret(vlt, secretName)
|
||||
|
||||
exists, err := secretObj.Exists()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
key, err := cli.storeNewEncryptionKey(vlt, secretName)
|
||||
if !errors.Is(err, vault.ErrSecretExists) {
|
||||
return key, err
|
||||
}
|
||||
// Another command stored the key since the check above: read it
|
||||
}
|
||||
|
||||
// Secret exists, get the age secret key from it
|
||||
secretBuffer, err := cli.getSecretValue(vlt, secretObj)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get secret value: %w", err)
|
||||
}
|
||||
|
||||
// Validate that it's a valid age secret key
|
||||
if !isValidAgeSecretKey(secretBuffer.String()) {
|
||||
secretBuffer.Destroy()
|
||||
|
||||
return nil, fmt.Errorf("secret '%s' %w", secretName, errNotAgeSecretKey)
|
||||
}
|
||||
|
||||
return secretBuffer, nil
|
||||
}
|
||||
|
||||
// Encrypt encrypts data using an age secret key stored in a secret
|
||||
func (cli *Instance) Encrypt(secretName, inputFile, outputFile string) error {
|
||||
err := vault.ValidateSecretName(secretName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get current vault
|
||||
vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var ageSecretKey string
|
||||
|
||||
// Check if secret exists
|
||||
secretObj := secret.NewSecret(vlt, secretName)
|
||||
exists, err := secretObj.Exists()
|
||||
// Get or create the age secret key for this secret
|
||||
keyBuffer, err := cli.resolveEncryptionKey(vlt, secretName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
return err
|
||||
}
|
||||
defer keyBuffer.Destroy()
|
||||
|
||||
if !exists { //nolint:nestif // Clear conditional logic for secret generation vs retrieval
|
||||
// Secret doesn't exist, generate new age key and store it
|
||||
identity, err := age.GenerateX25519Identity()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate age key: %w", err)
|
||||
}
|
||||
|
||||
// Store the generated key directly in a secure buffer
|
||||
identityStr := identity.String()
|
||||
secureBuffer := memguard.NewBufferFromBytes([]byte(identityStr))
|
||||
defer secureBuffer.Destroy()
|
||||
|
||||
// Set ageSecretKey for later use (we need it for encryption)
|
||||
ageSecretKey = identityStr
|
||||
|
||||
err = vlt.AddSecret(secretName, secureBuffer, false)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to store age key: %w", err)
|
||||
}
|
||||
} else {
|
||||
// Secret exists, get the age secret key from it
|
||||
secretBuffer, err := cli.getSecretValue(vlt, secretObj)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get secret value: %w", err)
|
||||
}
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
ageSecretKey = secretBuffer.String()
|
||||
|
||||
// Validate that it's a valid age secret key
|
||||
if !isValidAgeSecretKey(ageSecretKey) {
|
||||
return fmt.Errorf("secret '%s' does not contain a valid age secret key", secretName)
|
||||
}
|
||||
}
|
||||
|
||||
// Parse the secret key using secure buffer
|
||||
finalSecureBuffer := memguard.NewBufferFromBytes([]byte(ageSecretKey))
|
||||
defer finalSecureBuffer.Destroy()
|
||||
|
||||
identity, err := age.ParseX25519Identity(finalSecureBuffer.String())
|
||||
// Parse the secret key
|
||||
identity, err := age.ParseX25519Identity(keyBuffer.String())
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to parse age secret key: %w", err)
|
||||
}
|
||||
@@ -124,23 +174,27 @@ func (cli *Instance) Encrypt(secretName, inputFile, outputFile string) error {
|
||||
|
||||
// Set up input reader
|
||||
var input io.Reader = os.Stdin
|
||||
|
||||
if inputFile != "" {
|
||||
file, err := cli.fs.Open(inputFile)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open input file: %w", err)
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
input = file
|
||||
}
|
||||
|
||||
// Set up output writer
|
||||
output := cli.cmd.OutOrStdout()
|
||||
|
||||
if outputFile != "" {
|
||||
file, err := cli.fs.Create(outputFile)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create output file: %w", err)
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
output = file
|
||||
}
|
||||
|
||||
@@ -150,11 +204,13 @@ func (cli *Instance) Encrypt(secretName, inputFile, outputFile string) error {
|
||||
return fmt.Errorf("failed to create age encryptor: %w", err)
|
||||
}
|
||||
|
||||
if _, err := io.Copy(encryptor, input); err != nil {
|
||||
_, err = io.Copy(encryptor, input)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to encrypt data: %w", err)
|
||||
}
|
||||
|
||||
if err := encryptor.Close(); err != nil {
|
||||
err = encryptor.Close()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to finalize encryption: %w", err)
|
||||
}
|
||||
|
||||
@@ -163,6 +219,11 @@ func (cli *Instance) Encrypt(secretName, inputFile, outputFile string) error {
|
||||
|
||||
// Decrypt decrypts data using an age secret key stored in a secret
|
||||
func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error {
|
||||
err := vault.ValidateSecretName(secretName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get current vault
|
||||
vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
@@ -171,26 +232,18 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error {
|
||||
|
||||
// Check if secret exists
|
||||
secretObj := secret.NewSecret(vlt, secretName)
|
||||
|
||||
exists, err := secretObj.Exists()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("secret '%s' does not exist", secretName)
|
||||
return fmt.Errorf("secret '%s' %w", secretName, errSecretDoesNotExist)
|
||||
}
|
||||
|
||||
// Get the age secret key from the secret
|
||||
var secretBuffer *memguard.LockedBuffer
|
||||
if os.Getenv(secret.EnvMnemonic) != "" {
|
||||
secretBuffer, err = secretObj.GetValue(nil)
|
||||
} else {
|
||||
unlocker, unlockErr := vlt.GetCurrentUnlocker()
|
||||
if unlockErr != nil {
|
||||
return fmt.Errorf("failed to get current unlocker: %w", unlockErr)
|
||||
}
|
||||
secretBuffer, err = secretObj.GetValue(unlocker)
|
||||
}
|
||||
secretBuffer, err := cli.getSecretValue(vlt, secretObj)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get secret value: %w", err)
|
||||
}
|
||||
@@ -198,7 +251,7 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error {
|
||||
|
||||
// Validate that it's a valid age secret key
|
||||
if !isValidAgeSecretKey(secretBuffer.String()) {
|
||||
return fmt.Errorf("secret '%s' does not contain a valid age secret key", secretName)
|
||||
return fmt.Errorf("secret '%s' %w", secretName, errNotAgeSecretKey)
|
||||
}
|
||||
|
||||
// Parse the age secret key to get the identity
|
||||
@@ -209,23 +262,27 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error {
|
||||
|
||||
// Set up input reader
|
||||
var input io.Reader = os.Stdin
|
||||
|
||||
if inputFile != "" {
|
||||
file, err := cli.fs.Open(inputFile)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open input file: %w", err)
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
input = file
|
||||
}
|
||||
|
||||
// Set up output writer
|
||||
output := cli.cmd.OutOrStdout()
|
||||
|
||||
if outputFile != "" {
|
||||
file, err := cli.fs.Create(outputFile)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create output file: %w", err)
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
output = file
|
||||
}
|
||||
|
||||
@@ -235,22 +292,27 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error {
|
||||
return fmt.Errorf("failed to create age decryptor: %w", err)
|
||||
}
|
||||
|
||||
if _, err := io.Copy(output, decryptor); err != nil {
|
||||
_, err = io.Copy(output, decryptor)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to decrypt data: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// isValidAgeSecretKey checks if a string is a valid age secret key by attempting to parse it
|
||||
// isValidAgeSecretKey checks if a string is a valid age secret key by
|
||||
// attempting to parse it
|
||||
func isValidAgeSecretKey(key string) bool {
|
||||
_, err := age.ParseX25519Identity(key)
|
||||
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// getSecretValue retrieves the value of a secret using the appropriate unlocker
|
||||
func (cli *Instance) getSecretValue(vlt *vault.Vault, secretObj *secret.Secret) (*memguard.LockedBuffer, error) {
|
||||
// getSecretValue retrieves the value of a secret using the appropriate
|
||||
// unlocker
|
||||
func (cli *Instance) getSecretValue(
|
||||
vlt *vault.Vault, secretObj *secret.Secret,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
if os.Getenv(secret.EnvMnemonic) != "" {
|
||||
return secretObj.GetValue(nil)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
package cli_test
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/cli"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// Entry must return its exit code rather than exit, so that its deferred
|
||||
// memguard purge runs on the success and the error path alike.
|
||||
//
|
||||
//nolint:paralleltest // sets os.Args, and Entry wipes every buffer in the process
|
||||
func TestEntryWipesBuffersAndReturnsExitCode(t *testing.T) {
|
||||
savedArgs := os.Args
|
||||
|
||||
t.Cleanup(func() { os.Args = savedArgs })
|
||||
|
||||
tests := []struct {
|
||||
args []string
|
||||
exitCode int
|
||||
}{
|
||||
{args: []string{"secret", "--help"}, exitCode: 0},
|
||||
{args: []string{"secret", "no-such-command"}, exitCode: 1},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
buf := memguard.NewBufferFromBytes([]byte("key material"))
|
||||
os.Args = tt.args
|
||||
|
||||
assert.Equal(t, tt.exitCode, cli.Entry(), "exit code for %v", tt.args)
|
||||
assert.False(t, buf.IsAlive(), "Entry left a buffer unwiped for %v", tt.args)
|
||||
}
|
||||
}
|
||||
|
||||
// Ctrl-C while `secret add` waits for the value on stdin must end the
|
||||
// process through memguard's signal handler, which wipes every buffer and
|
||||
// exits with status 1, not through Go's default handling, which kills the
|
||||
// process with the buffers intact.
|
||||
func TestInterruptExitsThroughMemguard(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const waitingForValue = "Reading secret value from stdin"
|
||||
|
||||
ctx, cancel := context.WithTimeout(t.Context(), time.Minute)
|
||||
defer cancel()
|
||||
|
||||
wd, err := filepath.Abs("../..")
|
||||
require.NoError(t, err)
|
||||
|
||||
secretPath := filepath.Join(wd, "secret")
|
||||
env := []string{
|
||||
secret.EnvStateDir + "=" + t.TempDir(),
|
||||
secret.EnvMnemonic + "=" + testMnemonic,
|
||||
secret.EnvUnlockPassphrase + "=test-passphrase",
|
||||
"PATH=/usr/bin:/bin",
|
||||
// The debug log on stderr shows when add starts waiting for the value.
|
||||
"GODEBUG=berlin.sneak.pkg.secret",
|
||||
}
|
||||
|
||||
//nolint:gosec // G204: test executes the freshly built secret binary
|
||||
initCmd := exec.CommandContext(ctx, secretPath, "init")
|
||||
initCmd.Env = env
|
||||
|
||||
output, err := initCmd.CombinedOutput()
|
||||
require.NoError(t, err, "init should succeed: %s", output)
|
||||
|
||||
//nolint:gosec // G204: test executes the freshly built secret binary
|
||||
addCmd := exec.CommandContext(ctx, secretPath, "add", "test/secret")
|
||||
addCmd.Env = env
|
||||
|
||||
// Held open and never written, so add keeps waiting for the value.
|
||||
stdin, err := addCmd.StdinPipe()
|
||||
require.NoError(t, err)
|
||||
|
||||
defer func() { _ = stdin.Close() }()
|
||||
|
||||
stderr, err := addCmd.StderrPipe()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, addCmd.Start())
|
||||
|
||||
waiting := false
|
||||
|
||||
scanner := bufio.NewScanner(stderr)
|
||||
for !waiting && scanner.Scan() {
|
||||
waiting = strings.Contains(scanner.Text(), waitingForValue)
|
||||
}
|
||||
|
||||
require.True(t, waiting, "add never logged %q", waitingForValue)
|
||||
require.NoError(t, addCmd.Process.Signal(os.Interrupt))
|
||||
|
||||
err = addCmd.Wait()
|
||||
|
||||
var exitErr *exec.ExitError
|
||||
|
||||
require.ErrorAs(t, err, &exitErr)
|
||||
assert.Equal(t, 1, exitErr.ExitCode(), "add ended with %v", err)
|
||||
}
|
||||
+50
-16
@@ -2,6 +2,7 @@ package cli
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"os"
|
||||
@@ -17,6 +18,16 @@ const (
|
||||
mnemonicEntropyBits = 128
|
||||
)
|
||||
|
||||
// Sentinel errors for secret generation
|
||||
var (
|
||||
errLengthTooSmall = errors.New("length must be at least 1")
|
||||
errLengthNotPositive = errors.New("length must be positive")
|
||||
errMnemonicTypeNotSupported = errors.New(
|
||||
"mnemonic type not supported for secret generation, " +
|
||||
"use 'secret generate mnemonic' instead")
|
||||
errUnsupportedSecretType = errors.New("unsupported type")
|
||||
)
|
||||
|
||||
func newGenerateCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "generate",
|
||||
@@ -38,7 +49,10 @@ func newGenerateMnemonicCmd() *cobra.Command {
|
||||
`mnemonic phrase that can be used with 'secret init' ` +
|
||||
`or 'secret import'.`,
|
||||
RunE: func(cmd *cobra.Command, _ []string) error {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize CLI: %w", err)
|
||||
}
|
||||
|
||||
return cli.GenerateMnemonic(cmd)
|
||||
},
|
||||
@@ -49,21 +63,27 @@ func newGenerateSecretCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "secret <name>",
|
||||
Short: "Generate a random secret and store it in the vault",
|
||||
Long: `Generate a cryptographically secure random secret and store it in the current vault under the given name.`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
Long: `Generate a cryptographically secure random secret and ` +
|
||||
`store it in the current vault under the given name.`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
length, _ := cmd.Flags().GetInt("length")
|
||||
secretType, _ := cmd.Flags().GetString("type")
|
||||
force, _ := cmd.Flags().GetBool("force")
|
||||
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize CLI: %w", err)
|
||||
}
|
||||
|
||||
return cli.GenerateSecret(cmd, args[0], length, secretType, force)
|
||||
},
|
||||
}
|
||||
|
||||
cmd.Flags().IntP("length", "l", defaultSecretLength, "Length of the generated secret (default 16)")
|
||||
cmd.Flags().StringP("type", "t", "base58", "Type of secret to generate (base58, alnum)")
|
||||
cmd.Flags().IntP("length", "l", defaultSecretLength,
|
||||
"Length of the generated secret (default 16)")
|
||||
cmd.Flags().StringP("type", "t", "base58",
|
||||
"Type of secret to generate (base58, alnum)")
|
||||
cmd.Flags().BoolP("force", "f", false, "Overwrite existing secret")
|
||||
|
||||
return cmd
|
||||
@@ -92,7 +112,8 @@ func (cli *Instance) GenerateMnemonic(cmd *cobra.Command) error {
|
||||
fmt.Fprintln(os.Stderr, " • Write it down on paper and store it safely")
|
||||
fmt.Fprintln(os.Stderr, " • Do not store it digitally or share it with anyone")
|
||||
fmt.Fprintln(os.Stderr, " • You will need this phrase to recover your secrets")
|
||||
fmt.Fprintln(os.Stderr, " • If you lose this phrase, your secrets cannot be recovered")
|
||||
fmt.Fprintln(os.Stderr,
|
||||
" • If you lose this phrase, your secrets cannot be recovered")
|
||||
fmt.Fprintln(os.Stderr, "")
|
||||
fmt.Fprintln(os.Stderr, "Use this mnemonic with:")
|
||||
fmt.Fprintln(os.Stderr, " secret init (to initialize a new secret manager)")
|
||||
@@ -110,11 +131,13 @@ func (cli *Instance) GenerateSecret(
|
||||
force bool,
|
||||
) error {
|
||||
if length < 1 {
|
||||
return fmt.Errorf("length must be at least 1")
|
||||
return errLengthTooSmall
|
||||
}
|
||||
|
||||
var secretValue string
|
||||
var err error
|
||||
var (
|
||||
secretValue string
|
||||
err error
|
||||
)
|
||||
|
||||
switch secretType {
|
||||
case "base58":
|
||||
@@ -122,15 +145,22 @@ func (cli *Instance) GenerateSecret(
|
||||
case "alnum":
|
||||
secretValue, err = generateRandomAlnum(length)
|
||||
case "mnemonic":
|
||||
return fmt.Errorf("mnemonic type not supported for secret generation, use 'secret generate mnemonic' instead")
|
||||
return errMnemonicTypeNotSupported
|
||||
default:
|
||||
return fmt.Errorf("unsupported type: %s (supported: base58, alnum)", secretType)
|
||||
return fmt.Errorf("%w: %s (supported: base58, alnum)",
|
||||
errUnsupportedSecretType, secretType)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to generate random secret: %w", err)
|
||||
}
|
||||
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
// Store the secret in the vault
|
||||
vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
@@ -141,11 +171,13 @@ func (cli *Instance) GenerateSecret(
|
||||
secretBuffer := memguard.NewBufferFromBytes([]byte(secretValue))
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
if err := vlt.AddSecret(secretName, secretBuffer, force); err != nil {
|
||||
err = vlt.AddSecret(secretName, secretBuffer, force)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cmd.Printf("Generated and stored %d-character %s secret: %s\n", length, secretType, secretName)
|
||||
cmd.Printf("Generated and stored %d-character %s secret: %s\n",
|
||||
length, secretType, secretName)
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -164,10 +196,11 @@ func generateRandomAlnum(length int) (string, error) {
|
||||
return generateRandomString(length, alnumChars)
|
||||
}
|
||||
|
||||
// generateRandomString generates a random string of the specified length using the given character set
|
||||
// generateRandomString generates a random string of the specified length
|
||||
// using the given character set
|
||||
func generateRandomString(length int, charset string) (string, error) {
|
||||
if length <= 0 {
|
||||
return "", fmt.Errorf("length must be positive")
|
||||
return "", errLengthNotPositive
|
||||
}
|
||||
|
||||
result := make([]byte, length)
|
||||
@@ -178,6 +211,7 @@ func generateRandomString(length int, charset string) (string, error) {
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to generate random number: %w", err)
|
||||
}
|
||||
|
||||
result[i] = charset[randomIndex.Int64()]
|
||||
}
|
||||
|
||||
|
||||
+24
-10
@@ -4,6 +4,7 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
@@ -17,7 +18,7 @@ import (
|
||||
)
|
||||
|
||||
// Version info - these are set at build time
|
||||
var ( //nolint:gochecknoglobals // Set at build time
|
||||
var (
|
||||
Version = "dev" //nolint:gochecknoglobals // Set at build time
|
||||
GitCommit = "unknown" //nolint:gochecknoglobals // Set at build time
|
||||
)
|
||||
@@ -34,20 +35,24 @@ type InfoOutput struct {
|
||||
NumVaults int `json:"numVaults"`
|
||||
NumSecrets int `json:"numSecrets"`
|
||||
TotalSize int64 `json:"totalSizeBytes"`
|
||||
OldestSecret time.Time `json:"oldestSecret,omitempty"`
|
||||
LatestSecret time.Time `json:"latestSecret,omitempty"`
|
||||
OldestSecret time.Time `json:"oldestSecret"`
|
||||
LatestSecret time.Time `json:"latestSecret"`
|
||||
}
|
||||
|
||||
// newInfoCmd returns the info command
|
||||
func newInfoCmd() *cobra.Command {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
log.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
var jsonOutput bool
|
||||
|
||||
cmd := &cobra.Command{
|
||||
Use: "info",
|
||||
Short: "Display system information",
|
||||
Long: "Display information about the secret system including version, vault statistics, and storage usage",
|
||||
Long: "Display information about the secret system including " +
|
||||
"version, vault statistics, and storage usage",
|
||||
RunE: func(cmd *cobra.Command, _ []string) error {
|
||||
return cli.Info(cmd, jsonOutput)
|
||||
},
|
||||
@@ -77,6 +82,7 @@ func (cli *Instance) Info(cmd *cobra.Command, jsonOutput bool) error {
|
||||
|
||||
// Count vaults
|
||||
vaultsDir := filepath.Join(cli.stateDir, "vaults.d")
|
||||
|
||||
vaultEntries, err := afero.ReadDir(cli.fs, vaultsDir)
|
||||
if err == nil {
|
||||
for _, entry := range vaultEntries {
|
||||
@@ -88,12 +94,15 @@ func (cli *Instance) Info(cmd *cobra.Command, jsonOutput bool) error {
|
||||
|
||||
// Gather statistics from all vaults
|
||||
if info.NumVaults > 0 {
|
||||
totalSecrets, totalSize, oldestTime, latestTime, _ := gatherVaultStats(cli.fs, vaultsDir)
|
||||
totalSecrets, totalSize, oldestTime, latestTime, _ := gatherVaultStats(
|
||||
cli.fs, vaultsDir)
|
||||
info.NumSecrets = totalSecrets
|
||||
info.TotalSize = totalSize
|
||||
|
||||
if !oldestTime.IsZero() {
|
||||
info.OldestSecret = oldestTime
|
||||
}
|
||||
|
||||
if !latestTime.IsZero() {
|
||||
info.LatestSecret = latestTime
|
||||
}
|
||||
@@ -140,19 +149,24 @@ func prettyPrintInfo(w io.Writer, info InfoOutput) error {
|
||||
_, _ = fmt.Fprintln(w, strings.Repeat("─", separatorLength))
|
||||
|
||||
_, _ = fmt.Fprintf(w, "🗂️ Vaults: %s\n", bold.Sprint(info.NumVaults))
|
||||
|
||||
_, _ = fmt.Fprintf(w, "🔑 Secrets: %s\n", bold.Sprint(info.NumSecrets))
|
||||
|
||||
if info.TotalSize >= 0 {
|
||||
//nolint:gosec // TotalSize is always >= 0
|
||||
_, _ = fmt.Fprintf(w, "💾 Total Size: %s\n", bold.Sprint(humanize.Bytes(uint64(info.TotalSize))))
|
||||
_, _ = fmt.Fprintf(w, "💾 Total Size: %s\n",
|
||||
bold.Sprint(humanize.Bytes(uint64(info.TotalSize))))
|
||||
} else {
|
||||
_, _ = fmt.Fprintf(w, "💾 Total Size: %s\n", bold.Sprint("0 B"))
|
||||
}
|
||||
|
||||
if !info.OldestSecret.IsZero() {
|
||||
_, _ = fmt.Fprintf(w, "🕰️ Oldest Secret: %s\n", info.OldestSecret.Format("2006-01-02 15:04:05"))
|
||||
_, _ = fmt.Fprintf(w, "🕰️ Oldest Secret: %s\n",
|
||||
info.OldestSecret.Format("2006-01-02 15:04:05"))
|
||||
}
|
||||
|
||||
if !info.LatestSecret.IsZero() {
|
||||
_, _ = fmt.Fprintf(w, "✨ Latest Secret: %s\n", info.LatestSecret.Format("2006-01-02 15:04:05"))
|
||||
_, _ = fmt.Fprintf(w, "✨ Latest Secret: %s\n",
|
||||
info.LatestSecret.Format("2006-01-02 15:04:05"))
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(w)
|
||||
|
||||
+97
-58
@@ -4,80 +4,119 @@ import (
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// gatherVaultStats collects statistics from all vaults
|
||||
// vaultStats accumulates statistics while walking vault directories
|
||||
type vaultStats struct {
|
||||
totalSecrets int
|
||||
totalSize int64
|
||||
oldestTime time.Time
|
||||
latestTime time.Time
|
||||
}
|
||||
|
||||
// addVersion accumulates size and timestamp info for one version directory
|
||||
func (s *vaultStats) addVersion(fs afero.Fs, versionPath string) {
|
||||
// Add size of encrypted data
|
||||
dataPath := filepath.Join(versionPath, "data.age")
|
||||
|
||||
stat, err := fs.Stat(dataPath)
|
||||
if err == nil {
|
||||
s.totalSize += stat.Size()
|
||||
}
|
||||
|
||||
// Add size of metadata
|
||||
metaPath := filepath.Join(versionPath, "metadata.age")
|
||||
|
||||
stat, err = fs.Stat(metaPath)
|
||||
if err == nil {
|
||||
s.totalSize += stat.Size()
|
||||
}
|
||||
|
||||
// Track timestamps
|
||||
stat, err = fs.Stat(versionPath)
|
||||
if err == nil {
|
||||
modTime := stat.ModTime()
|
||||
if s.oldestTime.IsZero() || modTime.Before(s.oldestTime) {
|
||||
s.oldestTime = modTime
|
||||
}
|
||||
|
||||
if s.latestTime.IsZero() || modTime.After(s.latestTime) {
|
||||
s.latestTime = modTime
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// addSecret accumulates stats for one secret directory
|
||||
func (s *vaultStats) addSecret(fs afero.Fs, secretsPath, secretName string) {
|
||||
s.totalSecrets++
|
||||
secretPath := filepath.Join(secretsPath, secretName)
|
||||
|
||||
// Get size and timestamps from all versions
|
||||
versionsPath := filepath.Join(secretPath, "versions")
|
||||
|
||||
versionEntries, err := afero.ReadDir(fs, versionsPath)
|
||||
if err != nil {
|
||||
secret.Warn("Could not read versions directory for secret",
|
||||
"secret", secretName, "error", err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
for _, versionEntry := range versionEntries {
|
||||
if !versionEntry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
s.addVersion(fs, filepath.Join(versionsPath, versionEntry.Name()))
|
||||
}
|
||||
}
|
||||
|
||||
// addVault accumulates stats for one vault directory
|
||||
func (s *vaultStats) addVault(fs afero.Fs, vaultsDir, vaultName string) {
|
||||
vaultPath := filepath.Join(vaultsDir, vaultName)
|
||||
secretsPath := filepath.Join(vaultPath, "secrets.d")
|
||||
|
||||
// Count secrets in this vault
|
||||
secretEntries, err := afero.ReadDir(fs, secretsPath)
|
||||
if err != nil {
|
||||
secret.Warn("Could not read secrets directory for vault",
|
||||
"vault", vaultName, "error", err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
for _, secretEntry := range secretEntries {
|
||||
if !secretEntry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
s.addSecret(fs, secretsPath, secretEntry.Name())
|
||||
}
|
||||
}
|
||||
|
||||
// gatherVaultStats collects statistics from all vaults, returning the
|
||||
// total secret count, total size, and oldest/latest secret timestamps
|
||||
func gatherVaultStats(
|
||||
fs afero.Fs,
|
||||
vaultsDir string,
|
||||
) (totalSecrets int, totalSize int64, oldestTime, latestTime time.Time, err error) {
|
||||
) (int, int64, time.Time, time.Time, error) {
|
||||
vaultEntries, err := afero.ReadDir(fs, vaultsDir)
|
||||
if err != nil {
|
||||
return 0, 0, time.Time{}, time.Time{}, err
|
||||
}
|
||||
|
||||
var stats vaultStats
|
||||
|
||||
for _, vaultEntry := range vaultEntries {
|
||||
if !vaultEntry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
vaultPath := filepath.Join(vaultsDir, vaultEntry.Name())
|
||||
secretsPath := filepath.Join(vaultPath, "secrets.d")
|
||||
|
||||
// Count secrets in this vault
|
||||
secretEntries, err := afero.ReadDir(fs, secretsPath)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
for _, secretEntry := range secretEntries {
|
||||
if !secretEntry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
totalSecrets++
|
||||
secretPath := filepath.Join(secretsPath, secretEntry.Name())
|
||||
|
||||
// Get size and timestamps from all versions
|
||||
versionsPath := filepath.Join(secretPath, "versions")
|
||||
versionEntries, err := afero.ReadDir(fs, versionsPath)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
for _, versionEntry := range versionEntries {
|
||||
if !versionEntry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
versionPath := filepath.Join(versionsPath, versionEntry.Name())
|
||||
|
||||
// Add size of encrypted data
|
||||
dataPath := filepath.Join(versionPath, "data.age")
|
||||
if stat, err := fs.Stat(dataPath); err == nil {
|
||||
totalSize += stat.Size()
|
||||
}
|
||||
|
||||
// Add size of metadata
|
||||
metaPath := filepath.Join(versionPath, "metadata.age")
|
||||
if stat, err := fs.Stat(metaPath); err == nil {
|
||||
totalSize += stat.Size()
|
||||
}
|
||||
|
||||
// Track timestamps
|
||||
if stat, err := fs.Stat(versionPath); err == nil {
|
||||
modTime := stat.ModTime()
|
||||
if oldestTime.IsZero() || modTime.Before(oldestTime) {
|
||||
oldestTime = modTime
|
||||
}
|
||||
if latestTime.IsZero() || modTime.After(latestTime) {
|
||||
latestTime = modTime
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
stats.addVault(fs, vaultsDir, vaultEntry.Name())
|
||||
}
|
||||
|
||||
return totalSecrets, totalSize, oldestTime, latestTime, nil
|
||||
return stats.totalSecrets, stats.totalSize,
|
||||
stats.oldestTime, stats.latestTime, nil
|
||||
}
|
||||
|
||||
+117
-105
@@ -1,7 +1,9 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -12,37 +14,118 @@ import (
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/tyler-smith/go-bip39"
|
||||
)
|
||||
|
||||
// errPassphraseMismatch is returned when passphrase confirmation fails
|
||||
var errPassphraseMismatch = errors.New("passphrases do not match")
|
||||
|
||||
// NewInitCmd creates the init command
|
||||
func NewInitCmd() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "init",
|
||||
Short: "Initialize the secrets manager",
|
||||
Long: `Create the necessary directory structure for storing secrets and generate encryption keys.`,
|
||||
RunE: RunInit,
|
||||
Long: `Create the necessary directory structure for storing ` +
|
||||
`secrets and generate encryption keys.`,
|
||||
RunE: RunInit,
|
||||
}
|
||||
}
|
||||
|
||||
// RunInit is the exported function that handles the init command
|
||||
func RunInit(cmd *cobra.Command, _ []string) error {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
log.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
return cli.Init(cmd)
|
||||
}
|
||||
|
||||
// Init initializes the secret manager
|
||||
// promptMnemonic reads the mnemonic from the environment or interactively.
|
||||
// The returned cleanup function must be deferred by the caller.
|
||||
func promptMnemonic() (string, func(), error) {
|
||||
if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" {
|
||||
secret.Debug("Using mnemonic from environment variable")
|
||||
|
||||
return envMnemonic, func() {}, nil
|
||||
}
|
||||
|
||||
secret.Debug("Prompting user for mnemonic phrase")
|
||||
|
||||
// Read mnemonic securely without echo
|
||||
mnemonicBuffer, err := secret.ReadPassphrase("Enter your BIP39 mnemonic phrase: ")
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read mnemonic from stdin", "error", err)
|
||||
|
||||
return "", nil, fmt.Errorf("failed to read mnemonic: %w", err)
|
||||
}
|
||||
|
||||
fmt.Fprintln(os.Stderr) // Add newline after hidden input
|
||||
|
||||
return mnemonicBuffer.String(), mnemonicBuffer.Destroy, nil
|
||||
}
|
||||
|
||||
// setupDefaultVault creates the default vault and derives its long-term
|
||||
// identity from the mnemonic
|
||||
func (cli *Instance) setupDefaultVault(
|
||||
stateDir, mnemonicStr string,
|
||||
) (*vault.Vault, *age.X25519Identity, error) {
|
||||
// Create the default vault - it will handle key derivation internally
|
||||
secret.Debug("Creating default vault")
|
||||
|
||||
vlt, err := vault.CreateVault(cli.fs, cli.stateDir, "default")
|
||||
if err != nil {
|
||||
secret.Debug("Failed to create default vault", "error", err)
|
||||
|
||||
return nil, nil, fmt.Errorf("failed to create default vault: %w", err)
|
||||
}
|
||||
|
||||
// Get the vault metadata to retrieve the derivation index
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "default")
|
||||
|
||||
metadata, err := vault.LoadVaultMetadata(cli.fs, vaultDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to load vault metadata", "error", err)
|
||||
|
||||
return nil, nil, fmt.Errorf("failed to load vault metadata: %w", err)
|
||||
}
|
||||
|
||||
// Derive the long-term key using the same index that CreateVault used
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonicStr, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to derive long-term key", "error", err)
|
||||
|
||||
return nil, nil, fmt.Errorf(
|
||||
"failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
|
||||
return vlt, ltIdentity, nil
|
||||
}
|
||||
|
||||
// Init initializes the secret manager, holding the state directory lock
|
||||
// while initialize runs
|
||||
func (cli *Instance) Init(cmd *cobra.Command) error {
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
return cli.initialize(cmd)
|
||||
}
|
||||
|
||||
// initialize creates the state directory, the default vault and its first
|
||||
// unlocker
|
||||
func (cli *Instance) initialize(cmd *cobra.Command) error {
|
||||
secret.Debug("Starting secret manager initialization")
|
||||
|
||||
// Create state directory
|
||||
stateDir := cli.GetStateDir()
|
||||
secret.DebugWith("Creating state directory", slog.String("path", stateDir))
|
||||
|
||||
if err := cli.fs.MkdirAll(stateDir, secret.DirPerms); err != nil {
|
||||
err := cli.fs.MkdirAll(stateDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to create state directory", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to create state directory: %w", err)
|
||||
@@ -53,100 +136,56 @@ func (cli *Instance) Init(cmd *cobra.Command) error {
|
||||
}
|
||||
|
||||
// Prompt for mnemonic
|
||||
var mnemonicStr string
|
||||
|
||||
if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" {
|
||||
secret.Debug("Using mnemonic from environment variable")
|
||||
mnemonicStr = envMnemonic
|
||||
} else {
|
||||
secret.Debug("Prompting user for mnemonic phrase")
|
||||
// Read mnemonic securely without echo
|
||||
mnemonicBuffer, err := secret.ReadPassphrase("Enter your BIP39 mnemonic phrase: ")
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read mnemonic from stdin", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to read mnemonic: %w", err)
|
||||
}
|
||||
defer mnemonicBuffer.Destroy()
|
||||
|
||||
mnemonicStr = mnemonicBuffer.String()
|
||||
fmt.Fprintln(os.Stderr) // Add newline after hidden input
|
||||
mnemonicStr, cleanupMnemonic, err := promptMnemonic()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer cleanupMnemonic()
|
||||
|
||||
if mnemonicStr == "" {
|
||||
secret.Debug("Empty mnemonic provided")
|
||||
|
||||
return fmt.Errorf("mnemonic cannot be empty")
|
||||
return errMnemonicEmpty
|
||||
}
|
||||
|
||||
// Validate the mnemonic using BIP39
|
||||
secret.DebugWith("Validating BIP39 mnemonic", slog.Int("word_count", len(strings.Fields(mnemonicStr))))
|
||||
secret.DebugWith("Validating BIP39 mnemonic",
|
||||
slog.Int("word_count", len(strings.Fields(mnemonicStr))))
|
||||
|
||||
if !bip39.IsMnemonicValid(mnemonicStr) {
|
||||
secret.Debug("Invalid BIP39 mnemonic provided")
|
||||
|
||||
return fmt.Errorf("invalid BIP39 mnemonic phrase\nRun 'secret generate mnemonic' to create a valid mnemonic")
|
||||
return fmt.Errorf(
|
||||
"%w\nRun 'secret generate mnemonic' to create a valid mnemonic",
|
||||
errInvalidMnemonicPhrase)
|
||||
}
|
||||
|
||||
// Ask for the unlocker passphrase before creating the vault, so that
|
||||
// stopping at the prompt leaves no vault without an unlocker behind
|
||||
passphraseBuffer, err := resolvePassphrase()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
// Set mnemonic in environment for CreateVault to use
|
||||
originalMnemonic := os.Getenv(secret.EnvMnemonic)
|
||||
_ = os.Setenv(secret.EnvMnemonic, mnemonicStr)
|
||||
defer func() {
|
||||
if originalMnemonic != "" {
|
||||
_ = os.Setenv(secret.EnvMnemonic, originalMnemonic)
|
||||
} else {
|
||||
_ = os.Unsetenv(secret.EnvMnemonic)
|
||||
}
|
||||
}()
|
||||
restoreMnemonicEnv := setMnemonicEnv(mnemonicStr)
|
||||
defer restoreMnemonicEnv()
|
||||
|
||||
// Create the default vault - it will handle key derivation internally
|
||||
secret.Debug("Creating default vault")
|
||||
vlt, err := vault.CreateVault(cli.fs, cli.stateDir, "default")
|
||||
// Create the default vault and derive its long-term key
|
||||
vlt, ltIdentity, err := cli.setupDefaultVault(stateDir, mnemonicStr)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to create default vault", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to create default vault: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
// Get the vault metadata to retrieve the derivation index
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "default")
|
||||
metadata, err := vault.LoadVaultMetadata(cli.fs, vaultDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to load vault metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to load vault metadata: %w", err)
|
||||
}
|
||||
|
||||
// Derive the long-term key using the same index that CreateVault used
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonicStr, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to derive long-term key", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
ltPubKey := ltIdentity.Recipient().String()
|
||||
|
||||
// Unlock the vault with the derived long-term key
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Prompt for passphrase for unlocker
|
||||
var passphraseBuffer *memguard.LockedBuffer
|
||||
if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" {
|
||||
secret.Debug("Using unlock passphrase from environment variable")
|
||||
passphraseBuffer = memguard.NewBufferFromBytes([]byte(envPassphrase))
|
||||
} else {
|
||||
secret.Debug("Prompting user for unlock passphrase")
|
||||
// Use secure passphrase input with confirmation
|
||||
passphraseBuffer, err = readSecurePassphrase("Enter passphrase for unlocker: ")
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read unlock passphrase", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to read passphrase: %w", err)
|
||||
}
|
||||
}
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
// Create passphrase-protected unlocker
|
||||
secret.Debug("Creating passphrase-protected unlocker")
|
||||
|
||||
passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to create unlocker", "error", err)
|
||||
@@ -154,35 +193,8 @@ func (cli *Instance) Init(cmd *cobra.Command) error {
|
||||
return fmt.Errorf("failed to create unlocker: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt long-term private key to the unlocker
|
||||
unlockerDir := passphraseUnlocker.GetDirectory()
|
||||
|
||||
// Read unlocker public key
|
||||
unlockerPubKeyData, err := afero.ReadFile(cli.fs, filepath.Join(unlockerDir, "pub.age"))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read unlocker public key: %w", err)
|
||||
}
|
||||
|
||||
unlockerRecipient, err := age.ParseX25519Recipient(string(unlockerPubKeyData))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to parse unlocker public key: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt long-term private key to unlocker
|
||||
// Use memguard to protect the private key in memory
|
||||
ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String()))
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, unlockerRecipient)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to encrypt long-term private key: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted long-term private key
|
||||
ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age")
|
||||
if err := afero.WriteFile(cli.fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms); err != nil {
|
||||
return fmt.Errorf("failed to write encrypted long-term private key: %w", err)
|
||||
}
|
||||
// Note: CreatePassphraseUnlocker already encrypts and writes the long-term
|
||||
// private key to longterm.age, so no need to do it again here.
|
||||
|
||||
if cmd != nil {
|
||||
cmd.Printf("\nDefault vault created and configured\n")
|
||||
@@ -219,7 +231,7 @@ func readSecurePassphrase(prompt string) (*memguard.LockedBuffer, error) {
|
||||
passphraseBuffer1.Destroy()
|
||||
passphraseBuffer2.Destroy()
|
||||
|
||||
return nil, fmt.Errorf("passphrases do not match")
|
||||
return nil, errPassphraseMismatch
|
||||
}
|
||||
|
||||
// Clean up the second buffer, we'll return the first
|
||||
|
||||
+488
-316
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,552 @@
|
||||
//nolint:testpackage // sets the unexported fields of Instance
|
||||
package cli
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const (
|
||||
// lockWait is how long a test waits for something that must happen
|
||||
// once the lock is free.
|
||||
lockWait = 10 * time.Second
|
||||
|
||||
// testPassphrase protects the passphrase unlockers the tests create.
|
||||
testPassphrase = "test-passphrase"
|
||||
|
||||
// testInput is a file outside the state directory that commands read.
|
||||
testInput = "/input"
|
||||
)
|
||||
|
||||
// lockInBackground starts taking the state directory lock and returns a
|
||||
// channel that delivers the function releasing it once it has been taken.
|
||||
func lockInBackground(t *testing.T, fs afero.Fs) <-chan func() {
|
||||
t.Helper()
|
||||
|
||||
taken := make(chan func(), 1)
|
||||
|
||||
go func() {
|
||||
release, err := vault.LockStateDir(fs, testStateDir)
|
||||
if assert.NoError(t, err) {
|
||||
taken <- release
|
||||
}
|
||||
}()
|
||||
|
||||
return taken
|
||||
}
|
||||
|
||||
// addAtOnce runs one add of the secret name per value, all at once, and
|
||||
// returns their errors.
|
||||
func addAtOnce(
|
||||
fs afero.Fs, stateDir, name string, force bool, values []string,
|
||||
) []error {
|
||||
errs := make(chan error, len(values))
|
||||
|
||||
for _, value := range values {
|
||||
go func() {
|
||||
cli := NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
cli.cmd = &cobra.Command{}
|
||||
cli.cmd.SetIn(strings.NewReader(value))
|
||||
|
||||
errs <- cli.AddSecret(name, force)
|
||||
}()
|
||||
}
|
||||
|
||||
results := make([]error, 0, len(values))
|
||||
for range values {
|
||||
results = append(results, <-errs)
|
||||
}
|
||||
|
||||
return results
|
||||
}
|
||||
|
||||
// numbered returns count distinct values starting with prefix.
|
||||
func numbered(prefix string, count int) []string {
|
||||
values := make([]string, 0, count)
|
||||
for i := range count {
|
||||
values = append(values, prefix+"-"+strconv.Itoa(i))
|
||||
}
|
||||
|
||||
return values
|
||||
}
|
||||
|
||||
// TestConcurrentAddsKeepEveryVersion runs adds of one secret at once, on
|
||||
// the in-memory and on the real filesystem. Without the state directory
|
||||
// lock, adds of a new secret all find it absent and replace each other, and
|
||||
// forced adds read the same highest version number and overwrite each
|
||||
// other's version. With it they behave as if run one after another.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids parallel subtests
|
||||
func TestConcurrentAddsKeepEveryVersion(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
const adds = 8
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
fs afero.Fs
|
||||
stateDir string
|
||||
}{
|
||||
{"memory", afero.NewMemMapFs(), testStateDir},
|
||||
{"real", afero.NewOsFs(), t.TempDir()},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := vault.CreateVault(tc.fs, tc.stateDir, "default")
|
||||
require.NoError(t, err)
|
||||
|
||||
// One add creates the secret; the others find that it exists
|
||||
created := 0
|
||||
|
||||
for _, err := range addAtOnce(tc.fs, tc.stateDir, "shared", false,
|
||||
numbered("create", adds)) {
|
||||
if err == nil {
|
||||
created++
|
||||
} else {
|
||||
require.ErrorIs(t, err, vault.ErrSecretExists)
|
||||
}
|
||||
}
|
||||
|
||||
require.Equal(t, 1, created, "exactly one add creates the secret")
|
||||
|
||||
// Every forced add stores a version of its own
|
||||
for _, err := range addAtOnce(tc.fs, tc.stateDir, "shared", true,
|
||||
numbered("force", adds)) {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
vlt, err := vault.GetCurrentVault(tc.fs, tc.stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
|
||||
versions, err := secret.ListVersions(tc.fs,
|
||||
filepath.Join(vaultDir, "secrets.d", "shared"))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, adds+1, "one version per successful add")
|
||||
|
||||
values := make(map[string]bool, len(versions))
|
||||
|
||||
for _, version := range versions {
|
||||
value, err := vlt.GetSecretVersion("shared", version)
|
||||
require.NoError(t, err)
|
||||
|
||||
values[string(value.Bytes())] = true
|
||||
value.Destroy()
|
||||
}
|
||||
|
||||
assert.Len(t, values, adds+1, "every add stored its own value")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// readNotifier passes reads through to Reader and closes reading at the
|
||||
// first one.
|
||||
type readNotifier struct {
|
||||
io.Reader
|
||||
|
||||
reading chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (r *readNotifier) Read(p []byte) (int, error) {
|
||||
r.once.Do(func() { close(r.reading) })
|
||||
|
||||
return r.Reader.Read(p)
|
||||
}
|
||||
|
||||
// TestEncryptPipedIntoAdd runs `secret encrypt key | secret add name` in
|
||||
// one process, starting encrypt once add is reading its input. Had add
|
||||
// taken the state directory lock before reading, it would hold the lock
|
||||
// while waiting for encrypt's output, and encrypt would wait for the lock
|
||||
// to store its key: neither would finish.
|
||||
func TestEncryptPipedIntoAdd(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
_, err := vault.CreateVault(fs, testStateDir, "default")
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, afero.WriteFile(fs, testInput, []byte("piped"), 0o600))
|
||||
|
||||
pipeReader, pipeWriter := io.Pipe()
|
||||
// If the test gives up, this makes add's read fail, so that both
|
||||
// commands return and release the lock the other tests use
|
||||
t.Cleanup(func() { _ = pipeReader.Close() })
|
||||
|
||||
const commands = 2
|
||||
|
||||
input := &readNotifier{Reader: pipeReader, reading: make(chan struct{})}
|
||||
results := make(chan error, commands)
|
||||
|
||||
go func() {
|
||||
add := NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
add.cmd = &cobra.Command{}
|
||||
add.cmd.SetIn(input)
|
||||
|
||||
results <- add.AddSecret("encrypted", false)
|
||||
}()
|
||||
|
||||
go func() {
|
||||
<-input.reading
|
||||
|
||||
encrypt := NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
encrypt.cmd = &cobra.Command{}
|
||||
encrypt.cmd.SetOut(pipeWriter)
|
||||
|
||||
err := encrypt.Encrypt("key", testInput, "")
|
||||
// Ends add's input, as the end of the pipe does
|
||||
_ = pipeWriter.CloseWithError(err)
|
||||
|
||||
results <- err
|
||||
}()
|
||||
|
||||
timeout := time.After(lockWait)
|
||||
|
||||
for range commands {
|
||||
select {
|
||||
case err := <-results:
|
||||
require.NoError(t, err)
|
||||
case <-timeout:
|
||||
t.Fatal("secret encrypt piped into secret add never finished")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestFailedCommandReleasesLock checks that a command failing after it
|
||||
// took the state directory lock leaves the lock free for the next command.
|
||||
func TestFailedCommandReleasesLock(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
cli := NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
|
||||
// Fails once it holds the lock: there is no current vault
|
||||
err := cli.RemoveSecret(&cobra.Command{}, "missing", false)
|
||||
require.Error(t, err)
|
||||
|
||||
select {
|
||||
case release := <-lockInBackground(t, fs):
|
||||
release()
|
||||
case <-time.After(lockWait):
|
||||
t.Fatal("the failed command left the state directory locked")
|
||||
}
|
||||
}
|
||||
|
||||
// stateDirModTimes returns the modification time of every file and
|
||||
// directory under the test state directory. Any change a command makes, even
|
||||
// rewriting a file with the same content, changes it.
|
||||
func stateDirModTimes(t *testing.T, fs afero.Fs) map[string]int64 {
|
||||
t.Helper()
|
||||
|
||||
modTimes := make(map[string]int64)
|
||||
|
||||
err := afero.Walk(fs, testStateDir,
|
||||
func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
modTimes[path] = info.ModTime().UnixNano()
|
||||
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
return modTimes
|
||||
}
|
||||
|
||||
// setupEveryCommand makes what each command in
|
||||
// TestChangingCommandsWaitForLock needs: the current vault "work" with two
|
||||
// versions of "test/secret", the vault "other" without a long-term key, for
|
||||
// vault import, and the file testInput. There is no vault "default", which
|
||||
// init creates. If withUnlocker is set, it also gives "work" a passphrase
|
||||
// unlocker, which is slow. It returns the older version and the unlocker's
|
||||
// ID.
|
||||
func setupEveryCommand(
|
||||
t *testing.T, fs afero.Fs, withUnlocker bool,
|
||||
) (string, string) {
|
||||
t.Helper()
|
||||
|
||||
other, err := vault.CreateVault(fs, testStateDir, "other")
|
||||
require.NoError(t, err)
|
||||
|
||||
otherDir, err := other.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, fs.Remove(filepath.Join(otherDir, "pub.age")))
|
||||
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, "work")
|
||||
require.NoError(t, err)
|
||||
|
||||
addTestSecret(t, vlt, []byte("older"), false)
|
||||
addTestSecret(t, vlt, []byte("newer"), true)
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
|
||||
versions, err := secret.ListVersions(fs,
|
||||
filepath.Join(vaultDir, "secrets.d", "test%secret"))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 2)
|
||||
|
||||
unlockerID := ""
|
||||
|
||||
if withUnlocker {
|
||||
passphrase := memguard.NewBufferFromBytes([]byte(testPassphrase))
|
||||
defer passphrase.Destroy()
|
||||
|
||||
unlocker, err := vlt.CreatePassphraseUnlocker(passphrase)
|
||||
require.NoError(t, err)
|
||||
|
||||
unlockerID = unlocker.GetID()
|
||||
}
|
||||
|
||||
require.NoError(t, afero.WriteFile(fs, testInput, []byte("input"), 0o600))
|
||||
|
||||
// Newest first
|
||||
return versions[1], unlockerID
|
||||
}
|
||||
|
||||
// waitingForLock reports whether a goroutine is stopped in
|
||||
// vault.LockStateDir, waiting for the in-memory filesystem's lock. The
|
||||
// stack trace of such a goroutine starts with the reason it waits,
|
||||
// "[sync.Mutex.Lock]", and names LockStateDir.
|
||||
func waitingForLock() bool {
|
||||
stacks := make([]byte, 1<<20)
|
||||
stacks = stacks[:runtime.Stack(stacks, true)]
|
||||
|
||||
for goroutine := range bytes.SplitSeq(stacks, []byte("\n\n")) {
|
||||
if bytes.Contains(goroutine, []byte("[sync.Mutex.Lock")) &&
|
||||
bytes.Contains(goroutine, []byte("vault.LockStateDir(")) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// requireWaitsForLock runs a command, given what setupEveryCommand made,
|
||||
// while holding the state directory lock. The command must neither finish
|
||||
// nor change anything before it waits for the lock, and must succeed once
|
||||
// the lock is released.
|
||||
func requireWaitsForLock(
|
||||
t *testing.T,
|
||||
withUnlocker bool,
|
||||
run func(cli *Instance, olderVersion, unlockerID string) error,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
olderVersion, unlockerID := setupEveryCommand(t, fs, withUnlocker)
|
||||
before := stateDirModTimes(t, fs)
|
||||
|
||||
release, err := vault.LockStateDir(fs, testStateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Released at most once, and also if the test fails while holding it,
|
||||
// so that later tests can take it
|
||||
release = sync.OnceFunc(release)
|
||||
defer release()
|
||||
|
||||
cli := NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
cli.cmd = &cobra.Command{}
|
||||
cli.cmd.SetIn(strings.NewReader("value"))
|
||||
cli.cmd.SetOut(io.Discard)
|
||||
|
||||
done := make(chan error, 1)
|
||||
|
||||
go func() { done <- run(cli, olderVersion, unlockerID) }()
|
||||
|
||||
timeout := time.After(lockWait)
|
||||
|
||||
for !waitingForLock() {
|
||||
select {
|
||||
case err := <-done:
|
||||
t.Fatalf("finished while the lock was held, with error %v", err)
|
||||
case <-timeout:
|
||||
t.Fatal("never waited for the lock")
|
||||
case <-time.After(time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
assert.Equal(t, before, stateDirModTimes(t, fs),
|
||||
"changed the state directory before waiting for the lock")
|
||||
|
||||
release()
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
require.NoError(t, err)
|
||||
case <-time.After(lockWait):
|
||||
t.Fatal("did not finish once the lock was released")
|
||||
}
|
||||
}
|
||||
|
||||
// TestChangingCommandsWaitForLock checks that each command that changes the
|
||||
// state directory waits for its lock.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids parallel subtests
|
||||
func TestChangingCommandsWaitForLock(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
withUnlocker bool
|
||||
run func(cli *Instance, olderVersion, unlockerID string) error
|
||||
}{
|
||||
{"add", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.AddSecret("added", false)
|
||||
}},
|
||||
{"import", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.ImportSecret(cli.cmd, "imported", testInput, false)
|
||||
}},
|
||||
{"generate secret", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.GenerateSecret(cli.cmd, "generated", 16, "base58", false)
|
||||
}},
|
||||
{"encrypt", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.Encrypt("key", testInput, "")
|
||||
}},
|
||||
{"rm", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.RemoveSecret(cli.cmd, "test/secret", false)
|
||||
}},
|
||||
{"move", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.MoveSecret(cli.cmd, "test/secret", "moved", false)
|
||||
}},
|
||||
{"version promote", false, func(cli *Instance, olderVersion, _ string) error {
|
||||
return cli.PromoteVersion(cli.cmd, "test/secret", olderVersion)
|
||||
}},
|
||||
{"version rm", false, func(cli *Instance, olderVersion, _ string) error {
|
||||
return cli.RemoveVersion(cli.cmd, "test/secret", olderVersion)
|
||||
}},
|
||||
{"vault create", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.CreateVault(cli.cmd, "created")
|
||||
}},
|
||||
{"vault select", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.SelectVault(cli.cmd, "other")
|
||||
}},
|
||||
{"vault import", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.VaultImport(cli.cmd, "other")
|
||||
}},
|
||||
{"vault rm", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.RemoveVault(cli.cmd, "other", false)
|
||||
}},
|
||||
{"unlocker add", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.UnlockersAdd("passphrase", cli.cmd)
|
||||
}},
|
||||
{"unlocker rm", true, func(cli *Instance, _, unlockerID string) error {
|
||||
return cli.UnlockersRemove(unlockerID, true, cli.cmd)
|
||||
}},
|
||||
{"unlocker select", true, func(cli *Instance, _, unlockerID string) error {
|
||||
return cli.UnlockerSelect(unlockerID)
|
||||
}},
|
||||
{"init", false, func(cli *Instance, _, _ string) error {
|
||||
return cli.Init(cli.cmd)
|
||||
}},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
requireWaitsForLock(t, tc.withUnlocker, tc.run)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestEncryptWithExistingKeyTakesNoLock checks that secret encrypt with a
|
||||
// key that already exists, which only reads the state directory, finishes
|
||||
// while another command holds the state directory lock.
|
||||
func TestEncryptWithExistingKeyTakesNoLock(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
_, err := vault.CreateVault(fs, testStateDir, "default")
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, afero.WriteFile(fs, testInput, []byte("input"), 0o600))
|
||||
|
||||
encrypt := NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
encrypt.cmd = &cobra.Command{}
|
||||
encrypt.cmd.SetOut(io.Discard)
|
||||
|
||||
// Stores the key
|
||||
require.NoError(t, encrypt.Encrypt("key", testInput, ""))
|
||||
|
||||
release, err := vault.LockStateDir(fs, testStateDir)
|
||||
require.NoError(t, err)
|
||||
// Also frees a waiting encrypt if the test fails, so that it releases
|
||||
// the lock the other tests use
|
||||
defer release()
|
||||
|
||||
done := make(chan error, 1)
|
||||
|
||||
go func() { done <- encrypt.Encrypt("key", testInput, "") }()
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
require.NoError(t, err)
|
||||
case <-time.After(lockWait):
|
||||
t.Fatal("secret encrypt with an existing key waited for the lock")
|
||||
}
|
||||
}
|
||||
|
||||
// TestEncryptStreamsUnlocked checks that secret encrypt has released the
|
||||
// state directory lock by the time it writes its output. Holding it while
|
||||
// streaming would stall every other changing command for as long as the
|
||||
// stream lasts, and forever when the other end of the pipe is one of them.
|
||||
func TestEncryptStreamsUnlocked(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
_, err := vault.CreateVault(fs, testStateDir, "default")
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, afero.WriteFile(fs, testInput, []byte("streamed"), 0o600))
|
||||
|
||||
outputReader, outputWriter := io.Pipe()
|
||||
done := make(chan error, 1)
|
||||
|
||||
go func() {
|
||||
encrypt := NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
encrypt.cmd = &cobra.Command{}
|
||||
encrypt.cmd.SetOut(outputWriter)
|
||||
|
||||
err := encrypt.Encrypt("key", testInput, "")
|
||||
_ = outputWriter.CloseWithError(err)
|
||||
|
||||
done <- err
|
||||
}()
|
||||
|
||||
// The first byte of output: encrypt is streaming now, and blocked
|
||||
// writing until it is read
|
||||
_, err = io.ReadFull(outputReader, make([]byte, 1))
|
||||
require.NoError(t, err)
|
||||
|
||||
taken := lockInBackground(t, fs)
|
||||
|
||||
select {
|
||||
case release := <-taken:
|
||||
release()
|
||||
case <-time.After(lockWait):
|
||||
// Let encrypt finish, so that it releases the lock, then free it
|
||||
// again for the tests that follow
|
||||
_, _ = io.Copy(io.Discard, outputReader)
|
||||
|
||||
(<-taken)()
|
||||
t.Fatal("secret encrypt held the lock while streaming")
|
||||
}
|
||||
|
||||
_, err = io.Copy(io.Discard, outputReader)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, <-done)
|
||||
}
|
||||
@@ -0,0 +1,234 @@
|
||||
package cli_test
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/cli"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestRejectedMoveWithinVaultLeavesStateUnchanged is a regression test for
|
||||
// https://git.eeqj.de/sneak/secret/issues/73, where a forced move of a secret
|
||||
// onto itself deleted it, also when "work" was spelled two ways, and a failed
|
||||
// move within "work" left "work" the current vault. "default" is the current
|
||||
// vault in every case, and each case runs on its own copy of the state
|
||||
// directory.
|
||||
//
|
||||
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
|
||||
func TestRejectedMoveWithinVaultLeavesStateUnchanged(t *testing.T) {
|
||||
before := snapshotStateDir(t, newTwoVaultFs(t))
|
||||
require.Equal(t, "default", before[testStateDir+"/currentvault"])
|
||||
|
||||
const (
|
||||
ontoItself = "secret 'x' cannot be moved onto itself"
|
||||
workX = "work:x"
|
||||
)
|
||||
|
||||
tests := []struct {
|
||||
command string
|
||||
source, dest string
|
||||
force bool
|
||||
wantErr string
|
||||
}{
|
||||
{"mv x x", "x", "x", false, ontoItself},
|
||||
{"mv --force x x", "x", "x", true, ontoItself},
|
||||
{"mv --force work:x work:", workX, "work:", true, ontoItself},
|
||||
// An empty destination name defaults to the source name.
|
||||
{`mv --force work:x ""`, workX, "", true, ontoItself},
|
||||
// "work" is a vault name, so the destination is work:x.
|
||||
{"mv --force work:x work", workX, "work", true, ontoItself},
|
||||
{
|
||||
"mv work:nosuch work:y", "work:nosuch", "work:y", false,
|
||||
"secret 'nosuch' not found",
|
||||
},
|
||||
// Only an existing vault is used.
|
||||
{
|
||||
"mv --force nosuch:x nosuch:y", "nosuch:x", "nosuch:y", true,
|
||||
"vault 'nosuch' does not exist",
|
||||
},
|
||||
// Each of these spells "work" a second way. The spelling is not a
|
||||
// valid vault name, so the move is not taken for a move between two
|
||||
// vaults, which would delete the destination, here the source.
|
||||
{
|
||||
"mv --force work:x work/:x", workX, "work/:x", true,
|
||||
vault.ValidateVaultName("work/").Error(),
|
||||
},
|
||||
{
|
||||
"mv --force work/:x work:", "work/:x", "work:", true,
|
||||
vault.ValidateVaultName("work/").Error(),
|
||||
},
|
||||
{
|
||||
"mv --force work:x ./work:x", workX, "./work:x", true,
|
||||
vault.ValidateVaultName("./work").Error(),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.command, func(t *testing.T) {
|
||||
fs := newFsFromSnapshot(t, before)
|
||||
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
|
||||
err := c.MoveSecret(&cobra.Command{}, tt.source, tt.dest, tt.force)
|
||||
|
||||
require.Equal(t, before, snapshotStateDir(t, fs))
|
||||
require.EqualError(t, err, tt.wantErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestMoveWithinOtherVaultKeepsCurrentVault checks that `secret mv work:x
|
||||
// work:y`, with "default" the current vault, renames "x" to "y" in "work" and
|
||||
// leaves "default" the current vault.
|
||||
//
|
||||
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
|
||||
func TestMoveWithinOtherVaultKeepsCurrentVault(t *testing.T) {
|
||||
fs := newTwoVaultFs(t)
|
||||
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
|
||||
err := c.MoveSecret(&cobra.Command{}, "work:x", "work:y", false)
|
||||
require.NoError(t, err)
|
||||
|
||||
after := snapshotStateDir(t, fs)
|
||||
workSecrets := testStateDir + "/vaults.d/work/secrets.d/"
|
||||
|
||||
require.Equal(t, "default", after[testStateDir+"/currentvault"])
|
||||
require.Contains(t, after, workSecrets+"y/")
|
||||
require.NotContains(t, after, workSecrets+"x/")
|
||||
}
|
||||
|
||||
// TestMoveOntoSameSecretUnderAnotherNameIsRejected is a regression test for
|
||||
// https://git.eeqj.de/sneak/secret/issues/78: on a case-insensitive
|
||||
// filesystem "Foo" and "foo" are one secret, and `secret mv --force Foo foo`
|
||||
// removed the destination, which was the source. Symbolic links on the real
|
||||
// filesystem give one secret two names here: in "default", "y" is a link to
|
||||
// the secret "x", and the secrets.d of "other" is a link to that of
|
||||
// "default", so other:x is default:x. Each move must be rejected and leave
|
||||
// the secret and the links as they were.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv
|
||||
func TestMoveOntoSameSecretUnderAnotherNameIsRejected(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
const isSame = "is the same secret on this filesystem"
|
||||
|
||||
tests := []struct {
|
||||
command string
|
||||
source, dest string
|
||||
force bool
|
||||
wantErr string
|
||||
}{
|
||||
{
|
||||
"mv --force y x", "y", "x", true,
|
||||
"secret 'y' cannot be moved onto itself: 'x' " + isSame,
|
||||
},
|
||||
{
|
||||
"mv --force x y", "x", "y", true,
|
||||
"secret 'x' cannot be moved onto itself: 'y' " + isSame,
|
||||
},
|
||||
{
|
||||
"mv x y", "x", "y", false,
|
||||
"secret 'x' cannot be moved onto itself: 'y' " + isSame,
|
||||
},
|
||||
{
|
||||
"mv --force default:x other:x", "default:x", "other:x", true,
|
||||
"secret 'default:x' cannot be moved onto itself: 'other:x' " +
|
||||
isSame,
|
||||
},
|
||||
{
|
||||
"mv default:x other", "default:x", "other", false,
|
||||
"secret 'default:x' cannot be moved onto itself: 'other:x' " +
|
||||
isSame,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.command, func(t *testing.T) {
|
||||
fs := afero.NewOsFs()
|
||||
stateDir := t.TempDir()
|
||||
vaultsDir := filepath.Join(stateDir, "vaults.d")
|
||||
|
||||
// "default" is created last, so it is the current vault.
|
||||
_, err := vault.CreateVault(fs, stateDir, "other")
|
||||
require.NoError(t, err)
|
||||
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "default")
|
||||
require.NoError(t, err)
|
||||
|
||||
err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("value")), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
defaultSecrets := filepath.Join(vaultsDir, "default", "secrets.d")
|
||||
otherSecrets := filepath.Join(vaultsDir, "other", "secrets.d")
|
||||
link := filepath.Join(defaultSecrets, "y")
|
||||
|
||||
require.NoError(t, os.Symlink("x", link))
|
||||
require.NoError(t, os.Remove(otherSecrets))
|
||||
require.NoError(t, os.Symlink(defaultSecrets, otherSecrets))
|
||||
|
||||
c := cli.NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
moveErr := c.MoveSecret(&cobra.Command{}, tt.source, tt.dest, tt.force)
|
||||
|
||||
value, err := vlt.GetSecret("x")
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
require.Equal(t, []byte("value"), value.Bytes())
|
||||
|
||||
target, err := os.Readlink(link)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "x", target)
|
||||
|
||||
target, err = os.Readlink(otherSecrets)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, defaultSecrets, target)
|
||||
|
||||
require.EqualError(t, moveErr, tt.wantErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestForcedCaseOnlyMoveOnCaseSensitiveFilesystem checks that where "Foo"
|
||||
// and "foo" are two secrets, `secret mv --force Foo foo` still replaces "foo"
|
||||
// with "Foo".
|
||||
func TestForcedCaseOnlyMoveOnCaseSensitiveFilesystem(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
fs := afero.NewOsFs()
|
||||
stateDir := t.TempDir()
|
||||
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "default")
|
||||
require.NoError(t, err)
|
||||
|
||||
err = vlt.AddSecret("Foo", memguard.NewBufferFromBytes([]byte("upper")), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = os.Stat(filepath.Join(stateDir, "vaults.d", "default", "secrets.d", "foo"))
|
||||
if err == nil {
|
||||
t.Skip("the temporary directory is on a case-insensitive filesystem")
|
||||
}
|
||||
|
||||
err = vlt.AddSecret("foo", memguard.NewBufferFromBytes([]byte("lower")), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
c := cli.NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
err = c.MoveSecret(&cobra.Command{}, "Foo", "foo", true)
|
||||
require.NoError(t, err)
|
||||
|
||||
value, err := vlt.GetSecret("foo")
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
require.Equal(t, []byte("upper"), value.Bytes())
|
||||
|
||||
_, err = vlt.GetSecret("Foo")
|
||||
require.ErrorIs(t, err, vault.ErrSecretNotFound)
|
||||
}
|
||||
@@ -0,0 +1,418 @@
|
||||
package cli_test
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"maps"
|
||||
"os"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/cli"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const (
|
||||
// testStateDir is the in-memory state directory of the test vaults.
|
||||
testStateDir = "/test/state"
|
||||
|
||||
// testPassphrase protects the passphrase unlocker of each test vault.
|
||||
testPassphrase = "test-passphrase"
|
||||
|
||||
// testVersion is a version name in the format the vault uses.
|
||||
testVersion = "20260101.001"
|
||||
|
||||
// missingFile is an import source that does not exist, so an import
|
||||
// that opened it before checking the name would fail with another error.
|
||||
missingFile = "/no/such/file"
|
||||
)
|
||||
|
||||
// The state directory newTwoVaultFs copies, recorded by snapshotStateDir.
|
||||
// Creating a passphrase unlocker is slow by design, so the vaults are made
|
||||
// once, by the first test that needs them.
|
||||
//
|
||||
//nolint:gochecknoglobals // shared by the tests that use newTwoVaultFs
|
||||
var (
|
||||
twoVaultsOnce sync.Once
|
||||
twoVaults map[string]string
|
||||
)
|
||||
|
||||
// newTwoVaultFs returns an in-memory filesystem holding the vaults "work"
|
||||
// and "default", the current one. Each holds the secret "x" and a
|
||||
// passphrase unlocker, so both secrets.d and unlockers.d have contents.
|
||||
// Every call returns a new copy of the same vaults.
|
||||
//
|
||||
//nolint:ireturn // afero.Fs is the filesystem abstraction used throughout
|
||||
func newTwoVaultFs(t *testing.T) afero.Fs {
|
||||
t.Helper()
|
||||
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
twoVaultsOnce.Do(func() {
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
for _, name := range []string{"work", "default"} {
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, name)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("value")), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = vlt.CreatePassphraseUnlocker(
|
||||
memguard.NewBufferFromBytes([]byte(testPassphrase)))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
twoVaults = snapshotStateDir(t, fs)
|
||||
})
|
||||
|
||||
require.NotNil(t, twoVaults, "making the vaults failed in an earlier test")
|
||||
|
||||
return newFsFromSnapshot(t, twoVaults)
|
||||
}
|
||||
|
||||
// snapshotStateDir maps every file under the state directory to its
|
||||
// contents, and every directory, written with a trailing "/", to "". Two
|
||||
// snapshots are equal only if nothing in it was added, removed or changed.
|
||||
func snapshotStateDir(t *testing.T, fs afero.Fs) map[string]string {
|
||||
t.Helper()
|
||||
|
||||
tree := map[string]string{}
|
||||
|
||||
err := afero.Walk(fs, testStateDir, func(
|
||||
path string, info os.FileInfo, err error,
|
||||
) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if info.IsDir() {
|
||||
tree[path+"/"] = ""
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
content, err := afero.ReadFile(fs, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
tree[path] = string(content)
|
||||
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
return tree
|
||||
}
|
||||
|
||||
// newFsFromSnapshot returns a new in-memory filesystem holding exactly the
|
||||
// directories and files recorded by snapshotStateDir.
|
||||
//
|
||||
//nolint:ireturn // afero.Fs is the filesystem abstraction used throughout
|
||||
func newFsFromSnapshot(t *testing.T, tree map[string]string) afero.Fs {
|
||||
t.Helper()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
// In sorted order every directory comes before its contents.
|
||||
for _, path := range slices.Sorted(maps.Keys(tree)) {
|
||||
dir, isDir := strings.CutSuffix(path, "/")
|
||||
if isDir {
|
||||
require.NoError(t, fs.MkdirAll(dir, secret.DirPerms))
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
err := afero.WriteFile(fs, path, []byte(tree[path]), secret.FilePerms)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
return fs
|
||||
}
|
||||
|
||||
// requireRejectedAndUnchanged runs a command on a copy of the state
|
||||
// directory recorded in before. It requires an error with exactly the
|
||||
// message of want, so that a later check rejecting the argument does not
|
||||
// count, and everything under the state directory as it was: the error
|
||||
// alone proves nothing, since it could come after the vault had already
|
||||
// been deleted.
|
||||
func requireRejectedAndUnchanged(
|
||||
t *testing.T, before map[string]string, want error,
|
||||
run func(c *cli.Instance) error,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
fs := newFsFromSnapshot(t, before)
|
||||
|
||||
err := run(cli.NewCLIInstanceWithStateDir(fs, testStateDir))
|
||||
|
||||
require.Equal(t, before, snapshotStateDir(t, fs))
|
||||
require.EqualError(t, err, want.Error())
|
||||
}
|
||||
|
||||
// TestInvalidSecretNameLeavesVaultsUnchanged is a regression test for
|
||||
// https://git.eeqj.de/sneak/secret/issues/33, where `secret rm ..` deleted
|
||||
// the whole vault, and `secret rm .` or `secret rm ""` every secret in it.
|
||||
// Moves and imports use --force, so that only the name check stands in
|
||||
// the way.
|
||||
//
|
||||
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
|
||||
func TestInvalidSecretNameLeavesVaultsUnchanged(t *testing.T) {
|
||||
// Creating a passphrase unlocker is slow by design, so the vaults are
|
||||
// created once and each case runs on its own copy of them.
|
||||
before := snapshotStateDir(t, newTwoVaultFs(t))
|
||||
|
||||
vaultDir := testStateDir + "/vaults.d/default"
|
||||
require.Contains(t, before, vaultDir+"/secrets.d/x/")
|
||||
require.Contains(t, before, vaultDir+"/unlockers.d/passphrase/")
|
||||
require.Equal(t, "default", before[testStateDir+"/currentvault"])
|
||||
|
||||
cmd := &cobra.Command{}
|
||||
|
||||
tests := []struct {
|
||||
command string
|
||||
rejected string // the secret name the command must reject
|
||||
run func(c *cli.Instance) error
|
||||
}{
|
||||
{"rm ..", "..", func(c *cli.Instance) error {
|
||||
return c.RemoveSecret(cmd, "..", false)
|
||||
}},
|
||||
{"rm .", ".", func(c *cli.Instance) error {
|
||||
return c.RemoveSecret(cmd, ".", false)
|
||||
}},
|
||||
{`rm ""`, "", func(c *cli.Instance) error {
|
||||
return c.RemoveSecret(cmd, "", false)
|
||||
}},
|
||||
{"rm ../../etc", "../../etc", func(c *cli.Instance) error {
|
||||
return c.RemoveSecret(cmd, "../../etc", false)
|
||||
}},
|
||||
{"mv --force .. x", "..", func(c *cli.Instance) error {
|
||||
return c.MoveSecret(cmd, "..", "x", true)
|
||||
}},
|
||||
{"mv --force x ..", "..", func(c *cli.Instance) error {
|
||||
return c.MoveSecret(cmd, "x", "..", true)
|
||||
}},
|
||||
{`mv --force x ""`, "", func(c *cli.Instance) error {
|
||||
return c.MoveSecret(cmd, "x", "", true)
|
||||
}},
|
||||
// "work" is not the current vault: a move within it must not
|
||||
// select it when a name is rejected.
|
||||
{"mv --force work:.. work:x", "..", func(c *cli.Instance) error {
|
||||
return c.MoveSecret(cmd, "work:..", "work:x", true)
|
||||
}},
|
||||
{"mv --force work:x work:..", "..", func(c *cli.Instance) error {
|
||||
return c.MoveSecret(cmd, "work:x", "work:..", true)
|
||||
}},
|
||||
{"mv --force default:.. work", "..", func(c *cli.Instance) error {
|
||||
return c.MoveSecret(cmd, "default:..", "work", true)
|
||||
}},
|
||||
{"mv --force default:.. work:y", "..", func(c *cli.Instance) error {
|
||||
return c.MoveSecret(cmd, "default:..", "work:y", true)
|
||||
}},
|
||||
{"mv --force default:x work:..", "..", func(c *cli.Instance) error {
|
||||
return c.MoveSecret(cmd, "default:x", "work:..", true)
|
||||
}},
|
||||
{"import --force ..", "..", func(c *cli.Instance) error {
|
||||
return c.ImportSecret(cmd, "..", missingFile, true)
|
||||
}},
|
||||
{"import --force .", ".", func(c *cli.Instance) error {
|
||||
return c.ImportSecret(cmd, ".", missingFile, true)
|
||||
}},
|
||||
{"import --force ../../etc", "../../etc", func(c *cli.Instance) error {
|
||||
return c.ImportSecret(cmd, "../../etc", missingFile, true)
|
||||
}},
|
||||
{"version list ..", "..", func(c *cli.Instance) error {
|
||||
return c.ListVersions(cmd, "..")
|
||||
}},
|
||||
{"version promote ..", "..", func(c *cli.Instance) error {
|
||||
return c.PromoteVersion(cmd, "..", testVersion)
|
||||
}},
|
||||
{"version rm ..", "..", func(c *cli.Instance) error {
|
||||
return c.RemoveVersion(cmd, "..", testVersion)
|
||||
}},
|
||||
{"encrypt ..", "..", func(c *cli.Instance) error {
|
||||
return c.Encrypt("..", "", "")
|
||||
}},
|
||||
{"decrypt ..", "..", func(c *cli.Instance) error {
|
||||
return c.Decrypt("..", "", "")
|
||||
}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.command, func(t *testing.T) {
|
||||
requireRejectedAndUnchanged(t, before, vault.ValidateSecretName(tt.rejected), tt.run)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestInvalidVersionLeavesVaultsUnchanged is a regression test for
|
||||
// https://git.eeqj.de/sneak/secret/issues/67, where
|
||||
// `secret version rm x ../../..` deleted the whole vault,
|
||||
// `secret version rm x ..` the secret x, and `secret version rm x .` or
|
||||
// `secret version rm x ""` every version of x. A version argument is
|
||||
// accepted only if it is one of the versions `secret version list` lists.
|
||||
//
|
||||
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
|
||||
func TestInvalidVersionLeavesVaultsUnchanged(t *testing.T) {
|
||||
before := snapshotStateDir(t, newTwoVaultFs(t))
|
||||
|
||||
cmd := &cobra.Command{}
|
||||
|
||||
commands := []struct {
|
||||
command string
|
||||
run func(c *cli.Instance, version string) error
|
||||
}{
|
||||
{"version rm x", func(c *cli.Instance, version string) error {
|
||||
return c.RemoveVersion(cmd, "x", version)
|
||||
}},
|
||||
{"version promote x", func(c *cli.Instance, version string) error {
|
||||
return c.PromoteVersion(cmd, "x", version)
|
||||
}},
|
||||
{"get x --version", func(c *cli.Instance, version string) error {
|
||||
return c.GetSecretWithVersion(cmd, "x", version)
|
||||
}},
|
||||
}
|
||||
|
||||
for _, tt := range commands {
|
||||
for _, version := range []string{"", ".", "..", "../../..", "a/b"} {
|
||||
t.Run(fmt.Sprintf("%s %q", tt.command, version), func(t *testing.T) {
|
||||
want := fmt.Errorf("version '%s' %w '%s'",
|
||||
version, vault.ErrVersionNotFound, "x")
|
||||
requireRejectedAndUnchanged(t, before, want,
|
||||
func(c *cli.Instance) error { return tt.run(c, version) })
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestInvalidVaultNameLeavesStateUnchanged is a regression test for
|
||||
// https://git.eeqj.de/sneak/secret/issues/68, where
|
||||
// `secret vault import ..` wrote a long-term key and an unlocker into the
|
||||
// state directory itself, and `secret vault select ..` made it the current
|
||||
// vault. Each command that takes a vault name must reject an invalid one
|
||||
// before building a path from it. The mnemonic and the passphrase are set,
|
||||
// and moves and removals use --force, so that only the name check stands
|
||||
// in the way.
|
||||
//
|
||||
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
|
||||
func TestInvalidVaultNameLeavesStateUnchanged(t *testing.T) {
|
||||
before := snapshotStateDir(t, newTwoVaultFs(t))
|
||||
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
cmd := &cobra.Command{}
|
||||
|
||||
// Each command is a format with %q where the vault name goes.
|
||||
commands := []struct {
|
||||
command string
|
||||
run func(c *cli.Instance, name string) error
|
||||
}{
|
||||
{"vault create %q", func(c *cli.Instance, name string) error {
|
||||
return c.CreateVault(cmd, name)
|
||||
}},
|
||||
{"vault import %q", func(c *cli.Instance, name string) error {
|
||||
return c.VaultImport(cmd, name)
|
||||
}},
|
||||
{"vault select %q", func(c *cli.Instance, name string) error {
|
||||
return c.SelectVault(cmd, name)
|
||||
}},
|
||||
{"vault remove --force %q", func(c *cli.Instance, name string) error {
|
||||
return c.RemoveVault(cmd, name, true)
|
||||
}},
|
||||
{"mv --force %q:x work:x", func(c *cli.Instance, name string) error {
|
||||
return c.MoveSecret(cmd, name+":x", "work:x", true)
|
||||
}},
|
||||
{"mv --force default:x %q:x", func(c *cli.Instance, name string) error {
|
||||
return c.MoveSecret(cmd, "default:x", name+":x", true)
|
||||
}},
|
||||
}
|
||||
|
||||
for _, tt := range commands {
|
||||
for _, name := range []string{"", ".", "..", "a/b"} {
|
||||
t.Run(fmt.Sprintf(tt.command, name), func(t *testing.T) {
|
||||
requireRejectedAndUnchanged(t, before, vault.ValidateVaultName(name),
|
||||
func(c *cli.Instance) error { return tt.run(c, name) })
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestRemoveVersionRemovesOnlyThatVersion checks that `secret version rm`
|
||||
// with a version that is not the current one removes that version and
|
||||
// changes nothing else.
|
||||
//
|
||||
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
|
||||
func TestRemoveVersionRemovesOnlyThatVersion(t *testing.T) {
|
||||
fs := newTwoVaultFs(t)
|
||||
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
// A second version of "x" becomes the current one.
|
||||
err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("new")), true)
|
||||
require.NoError(t, err)
|
||||
|
||||
secretDir := testStateDir + "/vaults.d/default/secrets.d/x"
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 2)
|
||||
|
||||
// ListVersions lists the newest version first.
|
||||
oldDir := secretDir + "/versions/" + versions[1] + "/"
|
||||
before := snapshotStateDir(t, fs)
|
||||
require.Contains(t, before, oldDir)
|
||||
|
||||
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
err = c.RemoveVersion(&cobra.Command{}, "x", versions[1])
|
||||
require.NoError(t, err)
|
||||
|
||||
// Expected: the state as before without everything under oldDir.
|
||||
want := map[string]string{}
|
||||
|
||||
for path, content := range before {
|
||||
if !strings.HasPrefix(path, oldDir) {
|
||||
want[path] = content
|
||||
}
|
||||
}
|
||||
|
||||
require.Equal(t, want, snapshotStateDir(t, fs))
|
||||
}
|
||||
|
||||
// TestMoveToVaultNameRenamesInCurrentVault checks that `secret mv x work`,
|
||||
// where "work" is also the name of a vault, renames the secret "x" to "work"
|
||||
// in the current vault and changes nothing else.
|
||||
//
|
||||
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
|
||||
func TestMoveToVaultNameRenamesInCurrentVault(t *testing.T) {
|
||||
before := snapshotStateDir(t, newTwoVaultFs(t))
|
||||
fs := newFsFromSnapshot(t, before)
|
||||
|
||||
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
err := c.MoveSecret(&cobra.Command{}, "x", "work", false)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Expected: the state as before, with everything under the current
|
||||
// vault's secrets.d/x/ now under secrets.d/work/.
|
||||
oldDir := testStateDir + "/vaults.d/default/secrets.d/x/"
|
||||
newDir := testStateDir + "/vaults.d/default/secrets.d/work/"
|
||||
want := map[string]string{}
|
||||
|
||||
for path, content := range before {
|
||||
rest, found := strings.CutPrefix(path, oldDir)
|
||||
if found {
|
||||
path = newDir + rest
|
||||
}
|
||||
|
||||
want[path] = content
|
||||
}
|
||||
|
||||
require.Contains(t, want, newDir)
|
||||
require.Equal(t, want, snapshotStateDir(t, fs))
|
||||
}
|
||||
+54
-8
@@ -4,26 +4,72 @@ import (
|
||||
"os"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/cobra"
|
||||
"golang.org/x/sys/unix"
|
||||
"golang.org/x/term"
|
||||
)
|
||||
|
||||
// Entry is the entry point for the secret CLI application
|
||||
func Entry() {
|
||||
cmd := newRootCmd()
|
||||
if err := cmd.Execute(); err != nil {
|
||||
os.Exit(1)
|
||||
// Entry runs the secret CLI and returns the process exit code. It wipes
|
||||
// every memguard buffer before it returns, so the caller must do nothing
|
||||
// but exit with the code.
|
||||
func Entry() int {
|
||||
// On SIGINT or SIGTERM memguard runs this function, wipes every buffer
|
||||
// and exits with status 1. The passphrase prompt turns terminal echo
|
||||
// off until the read finishes, so a signal there would leave echo off.
|
||||
// Only a process in the terminal's foreground process group may reset
|
||||
// it: one in the background that tries is stopped instead of exiting.
|
||||
terminalState, terminalErr := term.GetState(unix.Stdin)
|
||||
|
||||
memguard.CatchSignal(func(os.Signal) {
|
||||
foreground, err := unix.IoctlGetInt(unix.Stdin, unix.TIOCGPGRP)
|
||||
if terminalErr == nil && err == nil && foreground == unix.Getpgrp() {
|
||||
_ = term.Restore(unix.Stdin, terminalState)
|
||||
}
|
||||
}, os.Interrupt, unix.SIGTERM)
|
||||
|
||||
defer memguard.Purge()
|
||||
|
||||
err := newRootCmd().Execute()
|
||||
if err != nil {
|
||||
return 1
|
||||
}
|
||||
|
||||
return 0
|
||||
}
|
||||
|
||||
func newRootCmd() *cobra.Command {
|
||||
secret.Debug("newRootCmd starting")
|
||||
|
||||
cmd := &cobra.Command{
|
||||
Use: "secret",
|
||||
Short: "A simple secrets manager",
|
||||
Long: `A simple secrets manager to store and retrieve sensitive information securely.`,
|
||||
// Ensure usage is shown after errors
|
||||
SilenceUsage: false,
|
||||
Long: `A simple secrets manager to store and retrieve sensitive ` +
|
||||
`information securely.`,
|
||||
// Cobra prints the error a command returns; Entry does not.
|
||||
SilenceErrors: false,
|
||||
// Usage belongs only to a command called wrongly. Cobra has
|
||||
// checked its arguments and flag values before this runs, but
|
||||
// checks required flags (ValidateRequiredFlags) and flag groups
|
||||
// (ValidateFlagGroups) only after it, so both are checked here
|
||||
// to keep usage for them. An error after that comes from running
|
||||
// the command, and usage would only bury it. A subcommand that
|
||||
// sets its own PersistentPreRun replaces this one.
|
||||
PersistentPreRunE: func(cmd *cobra.Command, _ []string) error {
|
||||
err := cmd.ValidateRequiredFlags()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = cmd.ValidateFlagGroups()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cmd.SilenceUsage = true
|
||||
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
secret.Debug("Adding subcommands to root command")
|
||||
|
||||
+557
-309
File diff suppressed because it is too large
Load Diff
+231
-195
@@ -1,3 +1,4 @@
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package cli
|
||||
|
||||
import (
|
||||
@@ -16,9 +17,195 @@ import (
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// testVaultName is the vault name used by the size tests.
|
||||
const testVaultName = "test-vault"
|
||||
|
||||
// lockedBytesPerSecretByte bounds the locked memory that storing a secret
|
||||
// holds at once: the buffers it is read into reach up to 1.5 times its
|
||||
// size, and they are then copied into one more buffer of its size.
|
||||
const lockedBytesPerSecretByte = 3
|
||||
|
||||
// skipIfLockedMemoryTooLow skips the test when this process cannot lock
|
||||
// the memory a secret of size bytes needs, found by locking a buffer of
|
||||
// that size and releasing it. memguard panics, ending the whole test run,
|
||||
// when it cannot lock a buffer, and a plain `docker build .` runs the
|
||||
// tests under an 8 MiB locked-memory limit (RLIMIT_MEMLOCK). A process
|
||||
// allowed to lock past that limit runs every case.
|
||||
func skipIfLockedMemoryTooLow(t *testing.T, size int) {
|
||||
t.Helper()
|
||||
|
||||
need := lockedBytesPerSecretByte * size
|
||||
|
||||
buf, err := unix.Mmap(-1, 0, need,
|
||||
unix.PROT_READ|unix.PROT_WRITE, unix.MAP_PRIVATE|unix.MAP_ANON)
|
||||
require.NoError(t, err)
|
||||
|
||||
lockErr := unix.Mlock(buf)
|
||||
|
||||
// Unmapping the buffer also unlocks it.
|
||||
err = unix.Munmap(buf)
|
||||
require.NoError(t, err)
|
||||
|
||||
if lockErr != nil {
|
||||
var limit unix.Rlimit
|
||||
|
||||
err = unix.Getrlimit(unix.RLIMIT_MEMLOCK, &limit)
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Skipf("a %d-byte secret needs up to %d bytes of locked memory, "+
|
||||
"which could not be locked under the locked-memory limit "+
|
||||
"(RLIMIT_MEMLOCK) of %d bytes: %v",
|
||||
size, need, limit.Cur, lockErr)
|
||||
}
|
||||
}
|
||||
|
||||
// newSizeTestVault creates an in-memory vault unlocked with the test
|
||||
// mnemonic and returns the filesystem and vault.
|
||||
//
|
||||
//nolint:ireturn // afero.Fs is the filesystem abstraction used throughout
|
||||
func newSizeTestVault(t *testing.T) (afero.Fs, *vault.Vault) {
|
||||
t.Helper()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
// Set test mnemonic
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
// Create vault
|
||||
_, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set current vault
|
||||
currentVaultPath := filepath.Join(testStateDir, "currentvault")
|
||||
vaultPath := filepath.Join(testStateDir, "vaults.d", testVaultName)
|
||||
err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Get vault and set up long-term key
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
require.NoError(t, err)
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
return fs, vlt
|
||||
}
|
||||
|
||||
// runAddSecretSizeCase adds a secret of the given size through stdin and
|
||||
// verifies the outcome.
|
||||
func runAddSecretSizeCase(t *testing.T, size int, wantErr bool, errMsg string) {
|
||||
t.Helper()
|
||||
skipIfLockedMemoryTooLow(t, size)
|
||||
|
||||
fs, vlt := newSizeTestVault(t)
|
||||
|
||||
// Generate test data of specified size
|
||||
testData := make([]byte, size)
|
||||
_, err := rand.Read(testData)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add newline that will be stripped
|
||||
testDataWithNewline := make([]byte, 0, len(testData)+1)
|
||||
testDataWithNewline = append(testDataWithNewline, testData...)
|
||||
testDataWithNewline = append(testDataWithNewline, '\n')
|
||||
|
||||
// Create command with fake stdin
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetIn(bytes.NewReader(testDataWithNewline))
|
||||
|
||||
// Create CLI instance
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
cli.fs = fs
|
||||
cli.stateDir = testStateDir
|
||||
cli.cmd = cmd
|
||||
|
||||
// Test adding the secret
|
||||
secretName := fmt.Sprintf("test-secret-%d", size)
|
||||
err = cli.AddSecret(secretName, false)
|
||||
|
||||
if wantErr {
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), errMsg)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify the secret was stored correctly
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer retrievedValue.Destroy()
|
||||
|
||||
assert.Equal(t, testData, retrievedValue.Bytes(),
|
||||
"Retrieved secret should match original (without newline)")
|
||||
}
|
||||
|
||||
// runImportSecretSizeCase imports a secret file of the given size and
|
||||
// verifies the outcome.
|
||||
func runImportSecretSizeCase(t *testing.T, size int, wantErr bool, errMsg string) {
|
||||
t.Helper()
|
||||
skipIfLockedMemoryTooLow(t, size)
|
||||
|
||||
fs, vlt := newSizeTestVault(t)
|
||||
|
||||
// Generate test data of specified size
|
||||
testData := make([]byte, size)
|
||||
_, err := rand.Read(testData)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Write test data to file
|
||||
testFile := fmt.Sprintf("/test/secret-%d.bin", size)
|
||||
err = afero.WriteFile(fs, testFile, testData, 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create command
|
||||
cmd := &cobra.Command{}
|
||||
|
||||
// Create CLI instance
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
cli.fs = fs
|
||||
cli.stateDir = testStateDir
|
||||
|
||||
// Test importing the secret
|
||||
secretName := fmt.Sprintf("imported-secret-%d", size)
|
||||
err = cli.ImportSecret(cmd, secretName, testFile, false)
|
||||
|
||||
if wantErr {
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), errMsg)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify the secret was stored correctly
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer retrievedValue.Destroy()
|
||||
|
||||
assert.Equal(t, testData, retrievedValue.Bytes(),
|
||||
"Retrieved secret should match original")
|
||||
}
|
||||
|
||||
// TestAddSecretVariousSizes tests adding secrets of various sizes through stdin
|
||||
//
|
||||
//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault
|
||||
func TestAddSecretVariousSizes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -71,73 +258,14 @@ func TestAddSecretVariousSizes(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// Set up test environment
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set test mnemonic
|
||||
t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about")
|
||||
|
||||
// Create vault
|
||||
vaultName := "test-vault"
|
||||
_, err := vault.CreateVault(fs, stateDir, vaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set current vault
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
vaultPath := filepath.Join(stateDir, "vaults.d", vaultName)
|
||||
err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Get vault and set up long-term key
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about", 0)
|
||||
require.NoError(t, err)
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Generate test data of specified size
|
||||
testData := make([]byte, tt.size)
|
||||
_, err = rand.Read(testData)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add newline that will be stripped
|
||||
testDataWithNewline := append(testData, '\n')
|
||||
|
||||
// Create fake stdin
|
||||
stdin := bytes.NewReader(testDataWithNewline)
|
||||
|
||||
// Create command with fake stdin
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetIn(stdin)
|
||||
|
||||
// Create CLI instance
|
||||
cli := NewCLIInstance()
|
||||
cli.fs = fs
|
||||
cli.stateDir = stateDir
|
||||
cli.cmd = cmd
|
||||
|
||||
// Test adding the secret
|
||||
secretName := fmt.Sprintf("test-secret-%d", tt.size)
|
||||
err = cli.AddSecret(secretName, false)
|
||||
|
||||
if tt.shouldError {
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tt.errorMsg)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify the secret was stored correctly
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testData, retrievedValue, "Retrieved secret should match original (without newline)")
|
||||
}
|
||||
runAddSecretSizeCase(t, tt.size, tt.shouldError, tt.errorMsg)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestImportSecretVariousSizes tests importing secrets of various sizes from files
|
||||
//
|
||||
//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault
|
||||
func TestImportSecretVariousSizes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -190,70 +318,14 @@ func TestImportSecretVariousSizes(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// Set up test environment
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set test mnemonic
|
||||
t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about")
|
||||
|
||||
// Create vault
|
||||
vaultName := "test-vault"
|
||||
_, err := vault.CreateVault(fs, stateDir, vaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set current vault
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
vaultPath := filepath.Join(stateDir, "vaults.d", vaultName)
|
||||
err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Get vault and set up long-term key
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about", 0)
|
||||
require.NoError(t, err)
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Generate test data of specified size
|
||||
testData := make([]byte, tt.size)
|
||||
_, err = rand.Read(testData)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Write test data to file
|
||||
testFile := fmt.Sprintf("/test/secret-%d.bin", tt.size)
|
||||
err = afero.WriteFile(fs, testFile, testData, 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create command
|
||||
cmd := &cobra.Command{}
|
||||
|
||||
// Create CLI instance
|
||||
cli := NewCLIInstance()
|
||||
cli.fs = fs
|
||||
cli.stateDir = stateDir
|
||||
|
||||
// Test importing the secret
|
||||
secretName := fmt.Sprintf("imported-secret-%d", tt.size)
|
||||
err = cli.ImportSecret(cmd, secretName, testFile, false)
|
||||
|
||||
if tt.shouldError {
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tt.errorMsg)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify the secret was stored correctly
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testData, retrievedValue, "Retrieved secret should match original")
|
||||
}
|
||||
runImportSecretSizeCase(t, tt.size, tt.shouldError, tt.errorMsg)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAddSecretBufferGrowth tests that our buffer growth strategy works correctly
|
||||
//
|
||||
//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault
|
||||
func TestAddSecretBufferGrowth(t *testing.T) {
|
||||
// Test various sizes that should trigger buffer growth
|
||||
sizes := []int{
|
||||
@@ -277,31 +349,9 @@ func TestAddSecretBufferGrowth(t *testing.T) {
|
||||
|
||||
for _, size := range sizes {
|
||||
t.Run(fmt.Sprintf("size_%d", size), func(t *testing.T) {
|
||||
// Set up test environment
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
skipIfLockedMemoryTooLow(t, size)
|
||||
|
||||
// Set test mnemonic
|
||||
t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about")
|
||||
|
||||
// Create vault
|
||||
vaultName := "test-vault"
|
||||
_, err := vault.CreateVault(fs, stateDir, vaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set current vault
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
vaultPath := filepath.Join(stateDir, "vaults.d", vaultName)
|
||||
err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Get vault and set up long-term key
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about", 0)
|
||||
require.NoError(t, err)
|
||||
vlt.Unlock(ltIdentity)
|
||||
fs, vlt := newSizeTestVault(t)
|
||||
|
||||
// Create test data of exactly the specified size
|
||||
// Use a pattern that's easy to verify
|
||||
@@ -310,17 +360,18 @@ func TestAddSecretBufferGrowth(t *testing.T) {
|
||||
testData[i] = byte(i % 256)
|
||||
}
|
||||
|
||||
// Create fake stdin without newline
|
||||
stdin := bytes.NewReader(testData)
|
||||
|
||||
// Create command with fake stdin
|
||||
// Create command with fake stdin (no newline)
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetIn(stdin)
|
||||
cmd.SetIn(bytes.NewReader(testData))
|
||||
|
||||
// Create CLI instance
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
cli.fs = fs
|
||||
cli.stateDir = stateDir
|
||||
cli.stateDir = testStateDir
|
||||
cli.cmd = cmd
|
||||
|
||||
// Test adding the secret
|
||||
@@ -331,55 +382,41 @@ func TestAddSecretBufferGrowth(t *testing.T) {
|
||||
// Verify the secret was stored correctly
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testData, retrievedValue, "Retrieved secret should match original exactly")
|
||||
|
||||
defer retrievedValue.Destroy()
|
||||
|
||||
assert.Equal(t, testData, retrievedValue.Bytes(),
|
||||
"Retrieved secret should match original exactly")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAddSecretStreamingBehavior tests that we handle streaming input correctly
|
||||
//
|
||||
//nolint:paralleltest // uses t.Setenv via newSizeTestVault
|
||||
func TestAddSecretStreamingBehavior(t *testing.T) {
|
||||
// Set up test environment
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set test mnemonic
|
||||
t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about")
|
||||
|
||||
// Create vault
|
||||
vaultName := "test-vault"
|
||||
_, err := vault.CreateVault(fs, stateDir, vaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set current vault
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
vaultPath := filepath.Join(stateDir, "vaults.d", vaultName)
|
||||
err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Get vault and set up long-term key
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about", 0)
|
||||
require.NoError(t, err)
|
||||
vlt.Unlock(ltIdentity)
|
||||
fs, vlt := newSizeTestVault(t)
|
||||
|
||||
// Create a custom reader that simulates slow streaming input
|
||||
// This will help verify our buffer handling works correctly with partial reads
|
||||
testData := []byte(strings.Repeat("Hello, World! ", 1000)) // ~14KB
|
||||
slowReader := &slowReader{
|
||||
streamingStdin := &slowReader{
|
||||
data: testData,
|
||||
chunkSize: 1000, // Read 1KB at a time
|
||||
}
|
||||
|
||||
// Create command with slow reader as stdin
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetIn(slowReader)
|
||||
cmd.SetIn(streamingStdin)
|
||||
|
||||
// Create CLI instance
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
cli.fs = fs
|
||||
cli.stateDir = stateDir
|
||||
cli.stateDir = testStateDir
|
||||
cli.cmd = cmd
|
||||
|
||||
// Test adding the secret
|
||||
@@ -389,7 +426,11 @@ func TestAddSecretStreamingBehavior(t *testing.T) {
|
||||
// Verify the secret was stored correctly
|
||||
retrievedValue, err := vlt.GetSecret("streaming-test")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testData, retrievedValue, "Retrieved secret should match original")
|
||||
|
||||
defer retrievedValue.Destroy()
|
||||
|
||||
assert.Equal(t, testData, retrievedValue.Bytes(),
|
||||
"Retrieved secret should match original")
|
||||
}
|
||||
|
||||
// slowReader simulates a reader that returns data in small chunks
|
||||
@@ -399,27 +440,22 @@ type slowReader struct {
|
||||
chunkSize int
|
||||
}
|
||||
|
||||
func (r *slowReader) Read(p []byte) (n int, err error) {
|
||||
func (r *slowReader) Read(p []byte) (int, error) {
|
||||
if r.offset >= len(r.data) {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
// Read at most chunkSize bytes
|
||||
// Read at most chunkSize bytes, bounded by the remaining data and
|
||||
// the destination buffer
|
||||
remaining := len(r.data) - r.offset
|
||||
toRead := r.chunkSize
|
||||
if toRead > remaining {
|
||||
toRead = remaining
|
||||
}
|
||||
if toRead > len(p) {
|
||||
toRead = len(p)
|
||||
}
|
||||
toRead := min(r.chunkSize, remaining, len(p))
|
||||
|
||||
n = copy(p, r.data[r.offset:r.offset+toRead])
|
||||
n := copy(p, r.data[r.offset:r.offset+toRead])
|
||||
r.offset += n
|
||||
|
||||
if r.offset >= len(r.data) {
|
||||
err = io.EOF
|
||||
return n, io.EOF
|
||||
}
|
||||
|
||||
return n, err
|
||||
return n, nil
|
||||
}
|
||||
|
||||
@@ -7,57 +7,64 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestGetCommandOutputsToStdout tests that 'secret get' outputs the secret value to stdout, not stderr
|
||||
// TestGetCommandOutputsToStdout tests that 'secret get' outputs the secret
|
||||
// value to stdout, not stderr
|
||||
func TestGetCommandOutputsToStdout(t *testing.T) {
|
||||
// Create a temporary directory for our vault
|
||||
tempDir := t.TempDir()
|
||||
|
||||
// Set environment variables for the test
|
||||
t.Setenv("SB_SECRET_STATE_DIR", tempDir)
|
||||
t.Setenv(secret.EnvStateDir, tempDir)
|
||||
|
||||
// Find the secret binary path
|
||||
wd, err := filepath.Abs("../..")
|
||||
require.NoError(t, err, "should get working directory")
|
||||
secretPath := filepath.Join(wd, "secret")
|
||||
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
secretPath := filepath.Join(wd, "secret")
|
||||
testPassphrase := "test-passphrase"
|
||||
|
||||
// Initialize vault
|
||||
cmd := exec.Command(secretPath, "init")
|
||||
//nolint:gosec // G204: test executes the freshly built secret binary
|
||||
cmd := exec.CommandContext(t.Context(), secretPath, "init")
|
||||
cmd.Env = []string{
|
||||
"SB_SECRET_STATE_DIR=" + tempDir,
|
||||
"SB_SECRET_MNEMONIC=" + testMnemonic,
|
||||
"SB_UNLOCK_PASSPHRASE=" + testPassphrase,
|
||||
secret.EnvStateDir + "=" + tempDir,
|
||||
secret.EnvMnemonic + "=" + testMnemonic,
|
||||
secret.EnvUnlockPassphrase + "=" + testPassphrase,
|
||||
"PATH=" + "/usr/bin:/bin",
|
||||
}
|
||||
|
||||
output, err := cmd.CombinedOutput()
|
||||
require.NoError(t, err, "init should succeed: %s", string(output))
|
||||
|
||||
// Add a secret
|
||||
cmd = exec.Command(secretPath, "add", "test/secret")
|
||||
//nolint:gosec // G204: test executes the freshly built secret binary
|
||||
cmd = exec.CommandContext(t.Context(), secretPath, "add", "test/secret")
|
||||
cmd.Env = []string{
|
||||
"SB_SECRET_STATE_DIR=" + tempDir,
|
||||
"SB_SECRET_MNEMONIC=" + testMnemonic,
|
||||
secret.EnvStateDir + "=" + tempDir,
|
||||
secret.EnvMnemonic + "=" + testMnemonic,
|
||||
"PATH=" + "/usr/bin:/bin",
|
||||
}
|
||||
cmd.Stdin = strings.NewReader("test-secret-value")
|
||||
|
||||
output, err = cmd.CombinedOutput()
|
||||
require.NoError(t, err, "add should succeed: %s", string(output))
|
||||
|
||||
// Test that 'secret get' outputs to stdout, not stderr
|
||||
cmd = exec.Command(secretPath, "get", "test/secret")
|
||||
//nolint:gosec // G204: test executes the freshly built secret binary
|
||||
cmd = exec.CommandContext(t.Context(), secretPath, "get", "test/secret")
|
||||
cmd.Env = []string{
|
||||
"SB_SECRET_STATE_DIR=" + tempDir,
|
||||
"SB_SECRET_MNEMONIC=" + testMnemonic,
|
||||
secret.EnvStateDir + "=" + tempDir,
|
||||
secret.EnvMnemonic + "=" + testMnemonic,
|
||||
"PATH=" + "/usr/bin:/bin",
|
||||
}
|
||||
|
||||
var stdout, stderr bytes.Buffer
|
||||
|
||||
cmd.Stdout = &stdout
|
||||
cmd.Stderr = &stderr
|
||||
|
||||
@@ -65,7 +72,8 @@ func TestGetCommandOutputsToStdout(t *testing.T) {
|
||||
require.NoError(t, err, "get should succeed")
|
||||
|
||||
// The secret value should be in stdout
|
||||
assert.Equal(t, "test-secret-value", strings.TrimSpace(stdout.String()), "secret value should be in stdout")
|
||||
assert.Equal(t, "test-secret-value", strings.TrimSpace(stdout.String()),
|
||||
"secret value should be in stdout")
|
||||
|
||||
// Nothing should be in stderr
|
||||
assert.Empty(t, stderr.String(), "stderr should be empty")
|
||||
|
||||
@@ -9,7 +9,9 @@ import (
|
||||
)
|
||||
|
||||
// ExecuteCommandInProcess executes a CLI command in-process for testing
|
||||
func ExecuteCommandInProcess(args []string, stdin string, env map[string]string) (string, error) {
|
||||
func ExecuteCommandInProcess(
|
||||
args []string, stdin string, env map[string]string,
|
||||
) (string, error) {
|
||||
secret.Debug("ExecuteCommandInProcess called", "args", args)
|
||||
|
||||
// Save current environment
|
||||
@@ -43,11 +45,13 @@ func ExecuteCommandInProcess(args []string, stdin string, env map[string]string)
|
||||
err := rootCmd.Execute()
|
||||
|
||||
output := buf.String()
|
||||
secret.Debug("Command execution completed", "error", err, "outputLength", len(output), "output", output)
|
||||
secret.Debug("Command execution completed",
|
||||
"error", err, "outputLength", len(output), "output", output)
|
||||
|
||||
// Add debug info for troubleshooting
|
||||
if len(output) == 0 && err == nil {
|
||||
secret.Debug("Warning: Command executed successfully but produced no output", "args", args)
|
||||
secret.Debug("Warning: Command executed successfully but produced no output",
|
||||
"args", args)
|
||||
}
|
||||
|
||||
// Restore environment
|
||||
|
||||
@@ -1,21 +1,23 @@
|
||||
package cli
|
||||
package cli_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/cli"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
//nolint:paralleltest // executes the CLI in-process against shared state
|
||||
func TestOutputCapture(t *testing.T) {
|
||||
// Test vault list command which we fixed
|
||||
output, err := ExecuteCommandInProcess([]string{"vault", "list"}, "", nil)
|
||||
output, err := cli.ExecuteCommandInProcess([]string{"vault", "list"}, "", nil)
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, output, "Available vaults", "should capture vault list output")
|
||||
t.Logf("vault list output: %q", output)
|
||||
|
||||
// Test help command
|
||||
output, err = ExecuteCommandInProcess([]string{"--help"}, "", nil)
|
||||
output, err = cli.ExecuteCommandInProcess([]string{"--help"}, "", nil)
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, output, "help output should not be empty")
|
||||
t.Logf("help output length: %d", len(output))
|
||||
|
||||
+571
-293
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,105 @@
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package cli
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// unknownTestGPGUserID is a GPG user ID that no key in the test keyring has.
|
||||
const unknownTestGPGUserID = "not-in-keyring@example.com"
|
||||
|
||||
// The secret TestAddPGPUnlocker stores, then reads through the new unlocker.
|
||||
const (
|
||||
addTestSecretName = "api-key"
|
||||
addTestSecretValue = "value"
|
||||
)
|
||||
|
||||
// TestAddPGPUnlocker adds a PGP unlocker for a throwaway GPG key to a vault
|
||||
// with a passphrase unlocker, getting the vault's long-term key from the
|
||||
// mnemonic or, with the mnemonic unset, from the passphrase unlocker. It
|
||||
// then reads a secret with neither the mnemonic nor the passphrase set, so
|
||||
// through the new unlocker, which the add selects.
|
||||
func TestAddPGPUnlocker(t *testing.T) {
|
||||
newTestGPGKey(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
// mnemonic is the mnemonic set while the unlocker is added.
|
||||
mnemonic string
|
||||
}{
|
||||
{"long-term key from the mnemonic", testMnemonic},
|
||||
{"long-term key from the current unlocker", ""},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
vlt, err := vault.CreateVault(fs, listTestStateDir, listTestVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = vlt.AddSecret(addTestSecretName,
|
||||
memguard.NewBufferFromBytes([]byte(addTestSecretValue)), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = vlt.CreatePassphraseUnlocker(
|
||||
memguard.NewBufferFromBytes([]byte(testPassphrase)))
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Setenv(secret.EnvMnemonic, test.mnemonic)
|
||||
|
||||
instance, cmd := newTestInstance(fs)
|
||||
cmd.Flags().String("keyid", unreadableTestGPGUserID, "")
|
||||
require.NoError(t, instance.UnlockersAdd(unlockerTypePGP, cmd))
|
||||
|
||||
t.Setenv(secret.EnvMnemonic, "")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "")
|
||||
|
||||
reopened := vault.NewVault(fs, listTestStateDir, listTestVaultName)
|
||||
|
||||
current, err := reopened.GetCurrentUnlocker()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, unlockerTypePGP, current.GetType())
|
||||
|
||||
value, err := reopened.GetSecret(addTestSecretName)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
assert.Equal(t, addTestSecretValue, value.String())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAddPGPUnlockerUnknownKey asserts that adding a PGP unlocker for a key
|
||||
// the keyring does not hold fails at looking up the key's fingerprint and
|
||||
// leaves no new unlocker directory. The error must come from the lookup: a
|
||||
// lookup moved after anything is written would also come after getting the
|
||||
// vault's long-term key, which fails first here: this vault's unlockers hold
|
||||
// no keys.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv (GNUPGHOME) forbids parallel tests
|
||||
func TestAddPGPUnlockerUnknownKey(t *testing.T) {
|
||||
newTestGPGKey(t)
|
||||
|
||||
base := newListTestVault(t, 1)
|
||||
instance, cmd := newTestInstance(base)
|
||||
cmd.Flags().String("keyid", unknownTestGPGUserID, "")
|
||||
|
||||
err := instance.addPGPUnlocker(cmd)
|
||||
|
||||
require.ErrorContains(t, err, "failed to resolve GPG key fingerprint")
|
||||
assertDirEntries(t, base,
|
||||
filepath.Join(testVaultDir(listTestVaultName), listTestUnlockersDirName),
|
||||
listTestUnlockerDirOne)
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
// Corrupt Unlocker Tests
|
||||
//
|
||||
// `secret unlocker select` and `secret unlocker remove` find an unlocker
|
||||
// by its ID. These tests give the first unlocker, which sorts before the
|
||||
// one the commands act on, metadata that is not JSON, and check that the
|
||||
// commands step past it, and that it can itself be removed by its
|
||||
// directory name, which `secret unlocker list` names in its warning. A
|
||||
// last test checks that an unlocker whose metadata file cannot be read is
|
||||
// removed by its directory name only as the last unlocker is.
|
||||
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package cli
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// newCorruptUnlockerVault returns the two-unlocker test vault with the
|
||||
// metadata of the first unlocker replaced by text that is not JSON.
|
||||
func newCorruptUnlockerVault(t *testing.T) *afero.MemMapFs {
|
||||
t.Helper()
|
||||
|
||||
fs := newListTestVault(t, 2)
|
||||
require.NoError(t, afero.WriteFile(fs,
|
||||
filepath.Join(testVaultDir(listTestVaultName), listTestUnlockersDirName,
|
||||
listTestUnlockerDirOne, listTestMetadataFileName),
|
||||
[]byte("not json"), listTestFilePerm))
|
||||
|
||||
return fs
|
||||
}
|
||||
|
||||
// TestUnlockerSelectSkipsCorruptUnlocker asserts that the second unlocker
|
||||
// can be selected, and that the corrupt one, having no type to be used as,
|
||||
// cannot be selected by its directory name.
|
||||
func TestUnlockerSelectSkipsCorruptUnlocker(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := newCorruptUnlockerVault(t)
|
||||
instance, _ := newTestInstance(fs)
|
||||
|
||||
require.NoError(t, instance.UnlockerSelect("pgp-"+listTestGPGKeyID+"B"))
|
||||
|
||||
current, err := afero.ReadFile(fs,
|
||||
filepath.Join(testVaultDir(listTestVaultName), "current-unlocker"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, listTestUnlockerDirTwo, string(current))
|
||||
|
||||
err = instance.UnlockerSelect(listTestUnlockerDirOne)
|
||||
require.ErrorIs(t, err, vault.ErrUnlockerNotFound)
|
||||
}
|
||||
|
||||
// TestUnlockerRemoveWithCorruptUnlocker asserts that the second unlocker
|
||||
// can be removed, unless the vault holds secrets: the corrupt unlocker
|
||||
// cannot unlock the vault, so the second is its last. The corrupt one can
|
||||
// be removed by its directory name without --force even then.
|
||||
func TestUnlockerRemoveWithCorruptUnlocker(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
unlockerID string
|
||||
withSecret bool
|
||||
wantErr error
|
||||
wantEntries []string
|
||||
}{
|
||||
{
|
||||
name: "the other unlocker",
|
||||
unlockerID: "pgp-" + listTestGPGKeyID + "B",
|
||||
wantEntries: []string{listTestUnlockerDirOne},
|
||||
},
|
||||
{
|
||||
name: "the other unlocker, the last one, with secrets",
|
||||
unlockerID: "pgp-" + listTestGPGKeyID + "B",
|
||||
withSecret: true,
|
||||
wantErr: errLastUnlocker,
|
||||
wantEntries: []string{
|
||||
listTestUnlockerDirOne, listTestUnlockerDirTwo,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "the corrupt unlocker by its directory name",
|
||||
unlockerID: listTestUnlockerDirOne,
|
||||
withSecret: true,
|
||||
wantEntries: []string{listTestUnlockerDirTwo},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := newCorruptUnlockerVault(t)
|
||||
if tt.withSecret {
|
||||
writeTestSecret(t, fs, testVaultDir(listTestVaultName))
|
||||
}
|
||||
|
||||
instance, cmd := newTestInstance(fs)
|
||||
|
||||
err := instance.UnlockersRemove(tt.unlockerID, false, cmd)
|
||||
require.ErrorIs(t, err, tt.wantErr)
|
||||
|
||||
assertDirEntries(t, fs,
|
||||
filepath.Join(testVaultDir(listTestVaultName),
|
||||
listTestUnlockersDirName),
|
||||
tt.wantEntries...)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestUnlockerRemoveWithUnreadableMetadata asserts that removing the only
|
||||
// unlocker of a vault with secrets by its directory name, when its
|
||||
// metadata file cannot be checked for or read, is refused without --force:
|
||||
// listing leaves it out, but it may still be the vault's only working
|
||||
// unlocker. With --force it is removed. The state directory lock refuses
|
||||
// the failing filesystem, so the test calls removeUnlocker, which
|
||||
// UnlockersRemove runs once it holds the lock.
|
||||
func TestUnlockerRemoveWithUnreadableMetadata(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
vaultDir := testVaultDir(listTestVaultName)
|
||||
unlockersDir := filepath.Join(vaultDir, listTestUnlockersDirName)
|
||||
failingPath := filepath.Join(unlockersDir, listTestUnlockerDirOne,
|
||||
listTestMetadataFileName)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
wrap func(base afero.Fs) afero.Fs
|
||||
}{
|
||||
{
|
||||
name: "checking for the file fails",
|
||||
wrap: func(base afero.Fs) afero.Fs {
|
||||
return &metadataStatFailFs{Fs: base, uncheckablePath: failingPath}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "reading the file fails",
|
||||
wrap: func(base afero.Fs) afero.Fs {
|
||||
return &metadataReadFailFs{Fs: base, unreadablePath: failingPath}
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base := newListTestVault(t, 1)
|
||||
writeTestSecret(t, base, vaultDir)
|
||||
|
||||
instance, cmd := newTestInstance(tt.wrap(base))
|
||||
|
||||
err := instance.removeUnlocker(listTestUnlockerDirOne, false, cmd)
|
||||
require.ErrorIs(t, err, errLastUnlocker)
|
||||
assertDirEntries(t, base, unlockersDir, listTestUnlockerDirOne)
|
||||
|
||||
require.NoError(t,
|
||||
instance.removeUnlocker(listTestUnlockerDirOne, true, cmd))
|
||||
assertDirEntries(t, base, unlockersDir)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,378 @@
|
||||
// Unlocker List Tests
|
||||
//
|
||||
// Tests for `secret unlocker list` behavior when the unlockers.d directory,
|
||||
// or an unlocker's metadata in it, cannot be read while the listing is
|
||||
// being rendered:
|
||||
//
|
||||
// - TestUnlockersListSkipsUnreadableUnlockersDir: an unreadable
|
||||
// unlockers.d yields no rows rather than rows bearing synthesized IDs.
|
||||
// - TestUnlockersListSkipsOnlyUnreadableEntries: a readable entry is
|
||||
// still listed, with its real ID and its current-unlocker marker,
|
||||
// when a later entry's scan fails.
|
||||
// - TestUnlockersListToleratesCorruptMetadata: one unlocker's corrupt
|
||||
// metadata does not stop the others from being listed.
|
||||
// - TestUnlockersListSkipsUnreadableMetadata: an unlocker whose metadata
|
||||
// file cannot be checked for or read is left out, and the other is
|
||||
// still listed.
|
||||
//
|
||||
// The listing resolves each unlocker's real ID by rescanning unlockers.d
|
||||
// after the vault has already enumerated it. If that rescan fails the ID
|
||||
// is unknowable, so the entry must be skipped: a synthesized ID matches
|
||||
// no `unlocker remove` or `unlocker select` argument and would also
|
||||
// suppress the current-unlocker marker.
|
||||
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package cli
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const (
|
||||
// listTestStateDir is the state directory of the synthetic vault used
|
||||
// by the unlocker listing tests.
|
||||
listTestStateDir = "/state"
|
||||
|
||||
// listTestVaultName is the name of that synthetic vault.
|
||||
listTestVaultName = "default"
|
||||
|
||||
// listTestGPGKeyID is the GPG key ID recorded in the readable PGP
|
||||
// unlocker's metadata. The unlocker's real ID is derived from it, and
|
||||
// differs from the timestamp-derived fallback ID.
|
||||
listTestGPGKeyID = "DEADBEEFDEADBEEF"
|
||||
|
||||
// listTestUnlockerDirOne and listTestUnlockerDirTwo are the unlocker
|
||||
// directory names under unlockers.d.
|
||||
listTestUnlockerDirOne = "host-pgp-2026-08-09"
|
||||
listTestUnlockerDirTwo = "host-pgp-2026-08-10"
|
||||
|
||||
// listTestUnlockersDirName is the directory the listing rescans to
|
||||
// resolve unlocker IDs.
|
||||
listTestUnlockersDirName = "unlockers.d"
|
||||
|
||||
// listTestMetadataFileName is the per-unlocker metadata file name.
|
||||
listTestMetadataFileName = "unlocker-metadata.json"
|
||||
|
||||
// listTestDirPerm and listTestFilePerm are the fixture permissions.
|
||||
listTestDirPerm = 0o700
|
||||
listTestFilePerm = 0o600
|
||||
)
|
||||
|
||||
// errUnlockersDirUnreadable is returned by the test filesystem in place of
|
||||
// a successful open of unlockers.d.
|
||||
var errUnlockersDirUnreadable = errors.New("permission denied")
|
||||
|
||||
// unlockersDirFailFs makes unlockers.d unreadable once it has been opened
|
||||
// successfully openBudget times. This reproduces the directory becoming
|
||||
// unreadable (permission change, partially restored backup, EIO) between
|
||||
// the vault's own enumeration and the per-entry rescan that resolves
|
||||
// unlocker IDs.
|
||||
type unlockersDirFailFs struct {
|
||||
afero.Fs
|
||||
|
||||
openBudget int
|
||||
opens int
|
||||
}
|
||||
|
||||
//nolint:ireturn // afero.File is the interface required by afero.Fs
|
||||
func (f *unlockersDirFailFs) Open(name string) (afero.File, error) {
|
||||
if filepath.Base(name) == listTestUnlockersDirName {
|
||||
f.opens++
|
||||
if f.opens > f.openBudget {
|
||||
return nil, errUnlockersDirUnreadable
|
||||
}
|
||||
}
|
||||
|
||||
//nolint:wrapcheck // test double must return the wrapped Fs error as-is
|
||||
return f.Fs.Open(name)
|
||||
}
|
||||
|
||||
// errMetadataUnreadable is returned by the test filesystem in place of a
|
||||
// successful open of one unlocker's metadata file.
|
||||
var errMetadataUnreadable = errors.New("input/output error")
|
||||
|
||||
// metadataReadFailFs fails every open of the file at unreadablePath. The
|
||||
// file still exists, so checking for it succeeds and only reading it fails.
|
||||
type metadataReadFailFs struct {
|
||||
afero.Fs
|
||||
|
||||
unreadablePath string
|
||||
}
|
||||
|
||||
//nolint:ireturn // afero.File is the interface required by afero.Fs
|
||||
func (f *metadataReadFailFs) Open(name string) (afero.File, error) {
|
||||
if name == f.unreadablePath {
|
||||
return nil, errMetadataUnreadable
|
||||
}
|
||||
|
||||
//nolint:wrapcheck // test double must return the wrapped Fs error as-is
|
||||
return f.Fs.Open(name)
|
||||
}
|
||||
|
||||
// errMetadataUncheckable is returned by the test filesystem in place of a
|
||||
// successful check for one unlocker's metadata file.
|
||||
var errMetadataUncheckable = errors.New("permission denied")
|
||||
|
||||
// metadataStatFailFs fails every check for whether the file at
|
||||
// uncheckablePath exists, as when its unlocker directory cannot be entered.
|
||||
type metadataStatFailFs struct {
|
||||
afero.Fs
|
||||
|
||||
uncheckablePath string
|
||||
}
|
||||
|
||||
func (f *metadataStatFailFs) Stat(name string) (os.FileInfo, error) {
|
||||
if name == f.uncheckablePath {
|
||||
return nil, errMetadataUncheckable
|
||||
}
|
||||
|
||||
//nolint:wrapcheck // test double must return the wrapped Fs error as-is
|
||||
return f.Fs.Stat(name)
|
||||
}
|
||||
|
||||
// writePGPUnlocker writes a PGP unlocker directory with metadata that
|
||||
// yields the real ID "pgp-<keyID>".
|
||||
func writePGPUnlocker(
|
||||
t *testing.T, fs afero.Fs, unlockersDir, dirName string,
|
||||
createdAt time.Time, keyID string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
metadata := secret.PGPUnlockerMetadata{
|
||||
UnlockerMetadata: secret.UnlockerMetadata{
|
||||
Type: unlockerTypePGP,
|
||||
CreatedAt: createdAt,
|
||||
},
|
||||
GPGKeyID: keyID,
|
||||
}
|
||||
|
||||
encoded, err := json.Marshal(metadata)
|
||||
require.NoError(t, err)
|
||||
|
||||
dir := filepath.Join(unlockersDir, dirName)
|
||||
require.NoError(t, fs.MkdirAll(dir, listTestDirPerm))
|
||||
require.NoError(t, afero.WriteFile(
|
||||
fs, filepath.Join(dir, listTestMetadataFileName), encoded,
|
||||
listTestFilePerm,
|
||||
))
|
||||
}
|
||||
|
||||
// newListTestVault builds a synthetic vault on a MemMapFs containing the
|
||||
// given number of PGP unlockers, with the first one selected as current.
|
||||
func newListTestVault(t *testing.T, unlockerCount int) *afero.MemMapFs {
|
||||
t.Helper()
|
||||
|
||||
base := &afero.MemMapFs{}
|
||||
vaultDir := filepath.Join(listTestStateDir, "vaults.d", listTestVaultName)
|
||||
unlockersDir := filepath.Join(vaultDir, listTestUnlockersDirName)
|
||||
|
||||
require.NoError(t, afero.WriteFile(
|
||||
base, filepath.Join(listTestStateDir, "currentvault"),
|
||||
[]byte(listTestVaultName), listTestFilePerm,
|
||||
))
|
||||
|
||||
names := []string{listTestUnlockerDirOne, listTestUnlockerDirTwo}
|
||||
names = names[:unlockerCount]
|
||||
|
||||
for i, name := range names {
|
||||
writePGPUnlocker(t, base, unlockersDir, name,
|
||||
time.Date(2026, time.August, 9+i, 12, 30, 0, 0, time.UTC),
|
||||
listTestGPGKeyID+string(rune('A'+i)),
|
||||
)
|
||||
}
|
||||
|
||||
require.NoError(t, afero.WriteFile(
|
||||
base, filepath.Join(vaultDir, "current-unlocker"),
|
||||
[]byte(names[0]), listTestFilePerm,
|
||||
))
|
||||
|
||||
return base
|
||||
}
|
||||
|
||||
// listUnlockersJSON runs UnlockersList in JSON mode against the given
|
||||
// filesystem and decodes the emitted unlocker rows.
|
||||
func listUnlockersJSON(t *testing.T, fs afero.Fs) []UnlockerInfo {
|
||||
t.Helper()
|
||||
|
||||
var buf bytes.Buffer
|
||||
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetOut(&buf)
|
||||
cmd.SetErr(&buf)
|
||||
|
||||
instance := &Instance{fs: fs, stateDir: listTestStateDir, cmd: cmd}
|
||||
require.NoError(t, instance.UnlockersList(true))
|
||||
|
||||
var decoded struct {
|
||||
Unlockers []UnlockerInfo `json:"unlockers"`
|
||||
}
|
||||
|
||||
require.NoError(t, json.Unmarshal(buf.Bytes(), &decoded))
|
||||
|
||||
return decoded.Unlockers
|
||||
}
|
||||
|
||||
// TestUnlockersListSkipsUnreadableUnlockersDir asserts that an unlockers.d
|
||||
// which becomes unreadable after the vault enumerated it produces no rows,
|
||||
// rather than rows carrying fabricated fallback IDs.
|
||||
func TestUnlockersListSkipsUnreadableUnlockersDir(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base := newListTestVault(t, 1)
|
||||
// Budget of one: the vault's own ListUnlockers scan succeeds, the
|
||||
// per-entry rescan that resolves the ID fails.
|
||||
fs := &unlockersDirFailFs{Fs: base, openBudget: 1}
|
||||
|
||||
unlockers := listUnlockersJSON(t, fs)
|
||||
|
||||
assert.Empty(t, unlockers,
|
||||
"an unreadable unlockers.d must yield no rows, not fabricated IDs")
|
||||
}
|
||||
|
||||
// TestUnlockersListSkipsOnlyUnreadableEntries asserts that a readable
|
||||
// entry survives with its real ID and current-unlocker marker when a later
|
||||
// entry's rescan fails.
|
||||
func TestUnlockersListSkipsOnlyUnreadableEntries(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base := newListTestVault(t, 2)
|
||||
// Budget of two: ListUnlockers plus the first entry's rescan succeed,
|
||||
// the second entry's rescan fails.
|
||||
fs := &unlockersDirFailFs{Fs: base, openBudget: 2}
|
||||
|
||||
unlockers := listUnlockersJSON(t, fs)
|
||||
|
||||
require.Len(t, unlockers, 1,
|
||||
"only the entry whose directory was readable may be listed")
|
||||
assert.Equal(t, "pgp-"+listTestGPGKeyID+"A", unlockers[0].ID,
|
||||
"the surviving row must carry the real unlocker ID")
|
||||
assert.True(t, unlockers[0].IsCurrent,
|
||||
"the current-unlocker marker must survive the skip")
|
||||
}
|
||||
|
||||
// TestUnlockersListReadableEntriesAreListed is the control case: with a
|
||||
// fully readable unlockers.d every entry is listed with its real ID.
|
||||
func TestUnlockersListReadableEntriesAreListed(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base := newListTestVault(t, 2)
|
||||
|
||||
unlockers := listUnlockersJSON(t, base)
|
||||
|
||||
require.Len(t, unlockers, 2)
|
||||
assert.Equal(t, "pgp-"+listTestGPGKeyID+"A", unlockers[0].ID)
|
||||
assert.Equal(t, "pgp-"+listTestGPGKeyID+"B", unlockers[1].ID)
|
||||
assert.True(t, unlockers[0].IsCurrent)
|
||||
assert.False(t, unlockers[1].IsCurrent)
|
||||
}
|
||||
|
||||
// TestUnlockersListToleratesCorruptMetadata asserts that one unlocker with
|
||||
// corrupt metadata does not stop the listing. Metadata that is not JSON
|
||||
// leaves that unlocker out; PGP metadata without a usable GPG key ID lists
|
||||
// it as "pgp-unknown". The healthy unlocker is listed with its real ID.
|
||||
func TestUnlockersListToleratesCorruptMetadata(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
healthyID := "pgp-" + listTestGPGKeyID + "A"
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
metadata string
|
||||
wantIDs []string
|
||||
}{
|
||||
{
|
||||
name: "not JSON",
|
||||
metadata: "not json",
|
||||
wantIDs: []string{healthyID},
|
||||
},
|
||||
{
|
||||
name: "GPG key ID of the wrong type",
|
||||
metadata: `{"type": "pgp", "gpgKeyId": 42}`,
|
||||
wantIDs: []string{healthyID, "pgp-unknown"},
|
||||
},
|
||||
{
|
||||
name: "GPG key ID missing",
|
||||
metadata: `{"type": "pgp"}`,
|
||||
wantIDs: []string{healthyID, "pgp-unknown"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := newListTestVault(t, 2)
|
||||
metadataPath := filepath.Join(listTestStateDir, "vaults.d",
|
||||
listTestVaultName, listTestUnlockersDirName,
|
||||
listTestUnlockerDirTwo, listTestMetadataFileName)
|
||||
require.NoError(t, afero.WriteFile(
|
||||
fs, metadataPath, []byte(tt.metadata), listTestFilePerm,
|
||||
))
|
||||
|
||||
unlockers := listUnlockersJSON(t, fs)
|
||||
require.Len(t, unlockers, len(tt.wantIDs))
|
||||
|
||||
for i, wantID := range tt.wantIDs {
|
||||
assert.Equal(t, wantID, unlockers[i].ID)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestUnlockersListSkipsUnreadableMetadata asserts that an unlocker whose
|
||||
// metadata file cannot be checked for or cannot be read is left out of the
|
||||
// listing, and the other unlocker is still listed with its real ID. The
|
||||
// failing one sorts first, so finding the other's ID has to step past it
|
||||
// as well.
|
||||
func TestUnlockersListSkipsUnreadableMetadata(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
failingPath := filepath.Join(listTestStateDir, "vaults.d",
|
||||
listTestVaultName, listTestUnlockersDirName,
|
||||
listTestUnlockerDirOne, listTestMetadataFileName)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
wrap func(base afero.Fs) afero.Fs
|
||||
}{
|
||||
{
|
||||
name: "checking for the file fails",
|
||||
wrap: func(base afero.Fs) afero.Fs {
|
||||
return &metadataStatFailFs{Fs: base, uncheckablePath: failingPath}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "reading the file fails",
|
||||
wrap: func(base afero.Fs) afero.Fs {
|
||||
return &metadataReadFailFs{Fs: base, unreadablePath: failingPath}
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := tt.wrap(newListTestVault(t, 2))
|
||||
|
||||
unlockers := listUnlockersJSON(t, fs)
|
||||
|
||||
require.Len(t, unlockers, 1,
|
||||
"only the unlocker with usable metadata may be listed")
|
||||
assert.Equal(t, "pgp-"+listTestGPGKeyID+"B", unlockers[0].ID,
|
||||
"the listed row must carry the real unlocker ID")
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,362 @@
|
||||
// Unreadable Directory Tests
|
||||
//
|
||||
// The checks that guard adding a PGP unlocker (is this key already an
|
||||
// unlocker?), removing the last unlocker and removing a vault (does the
|
||||
// vault hold secrets?), and importing a mnemonic (does the vault already
|
||||
// have a long-term key?) each look at the vault on disk before acting.
|
||||
// When that look fails they must refuse to act, not read the failure as
|
||||
// "nothing there" and go ahead.
|
||||
//
|
||||
// The tests make the look fail with a wrapper around the in-memory
|
||||
// filesystem, which the state directory lock refuses. So they call the
|
||||
// function each command runs once it holds the lock, such as removeVault
|
||||
// for RemoveVault.
|
||||
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const (
|
||||
// unreadableTestGPGUserID is the user ID of the throwaway GPG key the
|
||||
// PGP unlocker tests generate, and the --keyid they pass.
|
||||
unreadableTestGPGUserID = "unlocker-test@example.com"
|
||||
|
||||
// unreadableTestSecretName is the secret stored in the vaults the
|
||||
// removal tests remove from.
|
||||
unreadableTestSecretName = "api-key"
|
||||
|
||||
// unreadableTestOtherVault is a second vault for the vault removal
|
||||
// test, since the last vault can never be removed.
|
||||
unreadableTestOtherVault = "work"
|
||||
|
||||
// unreadableTestSecretsDirName is the directory holding a vault's
|
||||
// secrets, and unreadableTestCurrentFileName the per-secret file
|
||||
// naming its current version.
|
||||
unreadableTestSecretsDirName = "secrets.d"
|
||||
unreadableTestCurrentFileName = "current"
|
||||
)
|
||||
|
||||
// errStatFailed is returned by statFailFs in place of a successful stat.
|
||||
var errStatFailed = errors.New("input/output error")
|
||||
|
||||
// statFailFs fails every Stat of one path, as an I/O or permission error
|
||||
// on that path would.
|
||||
type statFailFs struct {
|
||||
afero.Fs
|
||||
|
||||
path string
|
||||
}
|
||||
|
||||
func (f *statFailFs) Stat(name string) (os.FileInfo, error) {
|
||||
if name == f.path {
|
||||
return nil, errStatFailed
|
||||
}
|
||||
|
||||
return f.Fs.Stat(name)
|
||||
}
|
||||
|
||||
// errOpenFailed is returned by openFailFs in place of a successful open.
|
||||
var errOpenFailed = errors.New("permission denied")
|
||||
|
||||
// openFailFs fails every Open of one path, as a directory without read
|
||||
// permission does: checking that it exists succeeds, listing it fails.
|
||||
type openFailFs struct {
|
||||
afero.Fs
|
||||
|
||||
path string
|
||||
}
|
||||
|
||||
//nolint:ireturn // afero.File is the interface required by afero.Fs
|
||||
func (f *openFailFs) Open(name string) (afero.File, error) {
|
||||
if name == f.path {
|
||||
return nil, errOpenFailed
|
||||
}
|
||||
|
||||
return f.Fs.Open(name)
|
||||
}
|
||||
|
||||
// testVaultDir returns the directory of the named vault in the synthetic
|
||||
// state directory built by newListTestVault.
|
||||
func testVaultDir(vaultName string) string {
|
||||
return filepath.Join(listTestStateDir, "vaults.d", vaultName)
|
||||
}
|
||||
|
||||
// newTestInstance returns a CLI instance on fs whose output is discarded.
|
||||
func newTestInstance(fs afero.Fs) (*Instance, *cobra.Command) {
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetOut(io.Discard)
|
||||
cmd.SetErr(io.Discard)
|
||||
|
||||
return &Instance{fs: fs, stateDir: listTestStateDir, cmd: cmd}, cmd
|
||||
}
|
||||
|
||||
// assertDirEntries asserts that dir holds exactly the named entries.
|
||||
func assertDirEntries(t *testing.T, fs afero.Fs, dir string, want ...string) {
|
||||
t.Helper()
|
||||
|
||||
entries, err := afero.ReadDir(fs, dir)
|
||||
require.NoError(t, err)
|
||||
|
||||
names := make([]string, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
names = append(names, entry.Name())
|
||||
}
|
||||
|
||||
assert.ElementsMatch(t, want, names)
|
||||
}
|
||||
|
||||
// newTestGPGKey points GNUPGHOME at a fresh directory, generates a GPG key
|
||||
// without a passphrase there, with a subkey for encryption, and returns the
|
||||
// key's fingerprint.
|
||||
func newTestGPGKey(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
// Not t.TempDir(): on macOS its path is too long for the gpg-agent
|
||||
// socket, which is created inside GNUPGHOME there.
|
||||
gnupgHome, err := os.MkdirTemp("", "gpg") //nolint:usetesting // short path
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Cleanup(func() { _ = os.RemoveAll(gnupgHome) })
|
||||
t.Setenv("GNUPGHOME", gnupgHome)
|
||||
|
||||
t.Cleanup(func() {
|
||||
// Stop the gpg-agent that key generation starts; cleanups run in
|
||||
// reverse order, so this happens before its directory is removed.
|
||||
// t.Context is already canceled when cleanup runs.
|
||||
ctx := context.WithoutCancel(t.Context())
|
||||
_ = exec.CommandContext(ctx, "gpgconf", "--kill", "gpg-agent").Run()
|
||||
})
|
||||
|
||||
output, err := exec.CommandContext(t.Context(), "gpg", "--batch",
|
||||
"--pinentry-mode", "loopback", "--passphrase", "",
|
||||
"--quick-gen-key", unreadableTestGPGUserID, "ed25519", "sign", "never",
|
||||
).CombinedOutput()
|
||||
require.NoError(t, err, "generating the test GPG key: %s", output)
|
||||
|
||||
fingerprint, err := secret.ResolveGPGKeyFingerprint(unreadableTestGPGUserID)
|
||||
require.NoError(t, err)
|
||||
|
||||
//nolint:gosec // G204: fingerprint is the test key's, as gpg printed it
|
||||
output, err = exec.CommandContext(t.Context(), "gpg", "--batch",
|
||||
"--pinentry-mode", "loopback", "--passphrase", "",
|
||||
"--quick-add-key", fingerprint, "cv25519", "encr", "never",
|
||||
).CombinedOutput()
|
||||
require.NoError(t, err, "adding the test GPG key's encryption subkey: %s",
|
||||
output)
|
||||
|
||||
return fingerprint
|
||||
}
|
||||
|
||||
// addTestPGPUnlocker runs `secret unlocker add pgp` for the test key
|
||||
// against fs.
|
||||
func addTestPGPUnlocker(fs afero.Fs) error {
|
||||
instance, cmd := newTestInstance(fs)
|
||||
cmd.Flags().String("keyid", unreadableTestGPGUserID, "")
|
||||
|
||||
return instance.addPGPUnlocker(cmd)
|
||||
}
|
||||
|
||||
// TestAddPGPUnlockerDuplicateCheck asserts that adding a PGP unlocker for
|
||||
// a key that already has one fails, and creates no unlocker directory,
|
||||
// when unlockers.d or the existing unlocker's metadata file cannot be
|
||||
// read; and, as the control case, that the existing unlocker is refused
|
||||
// as a duplicate when everything can be read.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv (GNUPGHOME) forbids parallel tests
|
||||
func TestAddPGPUnlockerDuplicateCheck(t *testing.T) {
|
||||
fingerprint := newTestGPGKey(t)
|
||||
unlockersDir := filepath.Join(
|
||||
testVaultDir(listTestVaultName), listTestUnlockersDirName)
|
||||
duplicateDir := filepath.Join(unlockersDir, listTestUnlockerDirTwo)
|
||||
|
||||
// newVaultWithDuplicate returns a vault holding an unlocker for the
|
||||
// test key, beside the one newListTestVault writes.
|
||||
newVaultWithDuplicate := func(t *testing.T) afero.Fs {
|
||||
t.Helper()
|
||||
|
||||
base := newListTestVault(t, 1)
|
||||
writePGPUnlocker(t, base, unlockersDir, listTestUnlockerDirTwo,
|
||||
time.Date(2026, time.August, 10, 12, 30, 0, 0, time.UTC),
|
||||
fingerprint)
|
||||
|
||||
return base
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
failFs func(base afero.Fs) afero.Fs
|
||||
wantErr error
|
||||
// wantPath is the path the error must name.
|
||||
wantPath string
|
||||
}{
|
||||
{
|
||||
name: "unlockers.d unreadable",
|
||||
failFs: func(base afero.Fs) afero.Fs {
|
||||
return &unlockersDirFailFs{Fs: base}
|
||||
},
|
||||
wantErr: errUnlockersDirUnreadable,
|
||||
wantPath: unlockersDir,
|
||||
},
|
||||
{
|
||||
name: "existing unlocker's metadata unreadable",
|
||||
failFs: func(base afero.Fs) afero.Fs {
|
||||
return &metadataReadFailFs{
|
||||
Fs: base,
|
||||
unreadablePath: filepath.Join(
|
||||
duplicateDir, listTestMetadataFileName),
|
||||
}
|
||||
},
|
||||
wantErr: errMetadataUnreadable,
|
||||
wantPath: duplicateDir,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
base := newVaultWithDuplicate(t)
|
||||
|
||||
err := addTestPGPUnlocker(tt.failFs(base))
|
||||
|
||||
require.ErrorIs(t, err, tt.wantErr)
|
||||
require.NotErrorIs(t, err, errGPGKeyAlreadyUnlocker)
|
||||
assert.Contains(t, err.Error(), tt.wantPath,
|
||||
"the error must name what it could not read")
|
||||
assertDirEntries(t, base, unlockersDir,
|
||||
listTestUnlockerDirOne, listTestUnlockerDirTwo)
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("duplicate refused", func(t *testing.T) {
|
||||
base := newVaultWithDuplicate(t)
|
||||
|
||||
err := addTestPGPUnlocker(base)
|
||||
|
||||
require.ErrorIs(t, err, errGPGKeyAlreadyUnlocker)
|
||||
assertDirEntries(t, base, unlockersDir,
|
||||
listTestUnlockerDirOne, listTestUnlockerDirTwo)
|
||||
})
|
||||
}
|
||||
|
||||
// writeTestSecret stores a secret with a current-version pointer, which is
|
||||
// what makes it count as a secret, in the given vault directory.
|
||||
func writeTestSecret(t *testing.T, fs afero.Fs, vaultDir string) {
|
||||
t.Helper()
|
||||
|
||||
secretDir := filepath.Join(
|
||||
vaultDir, unreadableTestSecretsDirName, unreadableTestSecretName)
|
||||
require.NoError(t, fs.MkdirAll(secretDir, listTestDirPerm))
|
||||
require.NoError(t, afero.WriteFile(fs,
|
||||
filepath.Join(secretDir, unreadableTestCurrentFileName),
|
||||
[]byte("20260809.001"), listTestFilePerm))
|
||||
}
|
||||
|
||||
// TestRemoveLastUnlockerAbortsWhenSecretsUnreadable asserts that the last
|
||||
// unlocker is kept when the secrets it protects cannot be counted.
|
||||
func TestRemoveLastUnlockerAbortsWhenSecretsUnreadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
vaultDir := testVaultDir(listTestVaultName)
|
||||
unlockersDir := filepath.Join(vaultDir, listTestUnlockersDirName)
|
||||
secretsDir := filepath.Join(vaultDir, unreadableTestSecretsDirName)
|
||||
|
||||
for _, path := range []string{
|
||||
secretsDir,
|
||||
filepath.Join(secretsDir, unreadableTestSecretName,
|
||||
unreadableTestCurrentFileName),
|
||||
} {
|
||||
t.Run(filepath.Base(path), func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base := newListTestVault(t, 1)
|
||||
writeTestSecret(t, base, vaultDir)
|
||||
instance, cmd := newTestInstance(&statFailFs{Fs: base, path: path})
|
||||
|
||||
err := instance.removeUnlocker(
|
||||
"pgp-"+listTestGPGKeyID+"A", false, cmd)
|
||||
|
||||
require.ErrorIs(t, err, errStatFailed)
|
||||
assertDirEntries(t, base, unlockersDir, listTestUnlockerDirOne)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestRemoveVaultAbortsWhenSecretsDirUnreadable asserts that a vault is
|
||||
// kept when whether it holds secrets cannot be determined: when checking
|
||||
// that secrets.d exists fails, and when it exists but cannot be listed.
|
||||
func TestRemoveVaultAbortsWhenSecretsDirUnreadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
vaultDir := testVaultDir(unreadableTestOtherVault)
|
||||
secretsDir := filepath.Join(vaultDir, unreadableTestSecretsDirName)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
failFs func(base afero.Fs) afero.Fs
|
||||
wantErr error
|
||||
}{
|
||||
{
|
||||
name: "check fails",
|
||||
failFs: func(base afero.Fs) afero.Fs {
|
||||
return &statFailFs{Fs: base, path: secretsDir}
|
||||
},
|
||||
wantErr: errStatFailed,
|
||||
},
|
||||
{
|
||||
name: "listing fails",
|
||||
failFs: func(base afero.Fs) afero.Fs {
|
||||
return &openFailFs{Fs: base, path: secretsDir}
|
||||
},
|
||||
wantErr: errOpenFailed,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base := newListTestVault(t, 1)
|
||||
writeTestSecret(t, base, vaultDir)
|
||||
instance, cmd := newTestInstance(tt.failFs(base))
|
||||
|
||||
err := instance.removeVault(cmd, unreadableTestOtherVault, false)
|
||||
|
||||
require.ErrorIs(t, err, tt.wantErr)
|
||||
|
||||
exists, err := afero.DirExists(base, vaultDir)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, exists, "the vault must not be removed")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestVaultImportAbortsWhenPubKeyUnreadable asserts that a mnemonic import
|
||||
// stops when whether the vault already has a long-term key cannot be
|
||||
// determined.
|
||||
func TestVaultImportAbortsWhenPubKeyUnreadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base := newListTestVault(t, 1)
|
||||
instance, cmd := newTestInstance(&statFailFs{
|
||||
Fs: base, path: filepath.Join(testVaultDir(listTestVaultName), "pub.age"),
|
||||
})
|
||||
|
||||
err := instance.importMnemonic(cmd, listTestVaultName)
|
||||
|
||||
require.ErrorIs(t, err, errStatFailed)
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package cli_test
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/cli"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// usageHeading starts the usage text cobra prints after an error.
|
||||
const usageHeading = "Usage:"
|
||||
|
||||
// A command called wrongly gets usage after its error; a command that
|
||||
// fails while running gets its error alone. Either way the command fails
|
||||
// and its error is shown exactly once.
|
||||
//
|
||||
//nolint:paralleltest // executes the CLI in-process and sets the environment
|
||||
func TestUsageOnlyForCallErrors(t *testing.T) {
|
||||
// No vault in the state directory, so `get x` fails while running.
|
||||
env := map[string]string{secret.EnvStateDir: t.TempDir()}
|
||||
|
||||
tests := []struct {
|
||||
call string
|
||||
wantUsage bool
|
||||
}{
|
||||
{call: "get", wantUsage: true},
|
||||
{call: "get x y", wantUsage: true},
|
||||
{call: "get --no-such-flag x", wantUsage: true},
|
||||
{call: "generate secret x --length abc", wantUsage: true},
|
||||
{call: "import x", wantUsage: true},
|
||||
{call: "get x", wantUsage: false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
output, err := cli.ExecuteCommandInProcess(strings.Fields(tt.call), "", env)
|
||||
require.Error(t, err, "%q should fail", tt.call)
|
||||
|
||||
assert.Equal(t, 1, strings.Count(output, err.Error()),
|
||||
"%q should show its error once:\n%s", tt.call, output)
|
||||
assert.Equal(t, tt.wantUsage, strings.Contains(output, usageHeading),
|
||||
"usage shown for %q:\n%s", tt.call, output)
|
||||
}
|
||||
}
|
||||
+341
-134
@@ -2,9 +2,12 @@ package cli
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -17,6 +20,22 @@ import (
|
||||
"github.com/tyler-smith/go-bip39"
|
||||
)
|
||||
|
||||
// Sentinel errors for vault operations
|
||||
var (
|
||||
errMnemonicEmpty = errors.New("mnemonic cannot be empty")
|
||||
errInvalidMnemonicPhrase = errors.New("invalid BIP39 mnemonic phrase")
|
||||
errInvalidMnemonic = errors.New("invalid BIP39 mnemonic")
|
||||
errVaultHasLongTermKey = errors.New(
|
||||
"already has a long-term key configured")
|
||||
errMnemonicEnvNotSet = errors.New(
|
||||
"SB_SECRET_MNEMONIC environment variable not set")
|
||||
errPassphraseEnvNotSet = errors.New(
|
||||
"SB_UNLOCK_PASSPHRASE environment variable not set")
|
||||
errCannotRemoveLastVault = errors.New("cannot remove the last vault")
|
||||
errVaultContainsSecrets = errors.New(
|
||||
"contains secrets; use --force to remove")
|
||||
)
|
||||
|
||||
func newVaultCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "vault",
|
||||
@@ -35,13 +54,16 @@ func newVaultCmd() *cobra.Command {
|
||||
|
||||
func newVaultListCmd() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Use: cmdUseList,
|
||||
Aliases: []string{"ls"},
|
||||
Short: "List available vaults",
|
||||
RunE: func(cmd *cobra.Command, _ []string) error {
|
||||
jsonOutput, _ := cmd.Flags().GetBool("json")
|
||||
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize CLI: %w", err)
|
||||
}
|
||||
|
||||
return cli.ListVaults(cmd, jsonOutput)
|
||||
},
|
||||
@@ -58,7 +80,10 @@ func newVaultCreateCmd() *cobra.Command {
|
||||
Short: "Create a new vault",
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize CLI: %w", err)
|
||||
}
|
||||
|
||||
return cli.CreateVault(cmd, args[0])
|
||||
},
|
||||
@@ -66,7 +91,10 @@ func newVaultCreateCmd() *cobra.Command {
|
||||
}
|
||||
|
||||
func newVaultSelectCmd() *cobra.Command {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
log.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
return &cobra.Command{
|
||||
Use: "select <name>",
|
||||
@@ -74,7 +102,10 @@ func newVaultSelectCmd() *cobra.Command {
|
||||
Args: cobra.ExactArgs(1),
|
||||
ValidArgsFunction: getVaultNamesCompletionFunc(cli.fs, cli.stateDir),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize CLI: %w", err)
|
||||
}
|
||||
|
||||
return cli.SelectVault(cmd, args[0])
|
||||
},
|
||||
@@ -82,12 +113,16 @@ func newVaultSelectCmd() *cobra.Command {
|
||||
}
|
||||
|
||||
func newVaultImportCmd() *cobra.Command {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
log.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
return &cobra.Command{
|
||||
Use: "import <vault-name>",
|
||||
Short: "Import a mnemonic into a vault",
|
||||
Long: `Import a BIP39 mnemonic phrase into the specified vault (default if not specified).`,
|
||||
Use: "import <vault-name>",
|
||||
Short: "Import a mnemonic into a vault",
|
||||
Long: `Import a BIP39 mnemonic phrase into the specified vault ` +
|
||||
`(default if not specified).`,
|
||||
Args: cobra.MaximumNArgs(1),
|
||||
ValidArgsFunction: getVaultNamesCompletionFunc(cli.fs, cli.stateDir),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
@@ -96,7 +131,10 @@ func newVaultImportCmd() *cobra.Command {
|
||||
vaultName = args[0]
|
||||
}
|
||||
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize CLI: %w", err)
|
||||
}
|
||||
|
||||
return cli.VaultImport(cmd, vaultName)
|
||||
},
|
||||
@@ -104,18 +142,27 @@ func newVaultImportCmd() *cobra.Command {
|
||||
}
|
||||
|
||||
func newVaultRemoveCmd() *cobra.Command {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
log.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
cmd := &cobra.Command{
|
||||
Use: "remove <name>",
|
||||
Aliases: []string{"rm"},
|
||||
Short: "Remove a vault",
|
||||
Long: `Remove a vault. Requires --force if the vault contains secrets. Will automatically ` +
|
||||
`switch to another vault if removing the currently selected one.`,
|
||||
Long: `Remove a vault. Requires --force if the vault contains ` +
|
||||
`secrets. Will automatically switch to another vault if ` +
|
||||
`removing the currently selected one.`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
ValidArgsFunction: getVaultNamesCompletionFunc(cli.fs, cli.stateDir),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
force, _ := cmd.Flags().GetBool("force")
|
||||
cli := NewCLIInstance()
|
||||
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initialize CLI: %w", err)
|
||||
}
|
||||
|
||||
return cli.RemoveVault(cmd, args[0], force)
|
||||
},
|
||||
@@ -136,11 +183,13 @@ func (cli *Instance) ListVaults(cmd *cobra.Command, jsonOutput bool) error {
|
||||
if jsonOutput { //nolint:nestif // Separate JSON and text output formatting logic
|
||||
// Get current vault name for context
|
||||
currentVault := ""
|
||||
if currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir); err == nil {
|
||||
|
||||
currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
|
||||
if err == nil {
|
||||
currentVault = currentVlt.GetName()
|
||||
}
|
||||
|
||||
result := map[string]interface{}{
|
||||
result := map[string]any{
|
||||
"vaults": vaults,
|
||||
"currentVault": currentVault,
|
||||
}
|
||||
@@ -149,16 +198,20 @@ func (cli *Instance) ListVaults(cmd *cobra.Command, jsonOutput bool) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
cmd.Println(string(jsonBytes))
|
||||
} else {
|
||||
// Text output
|
||||
cmd.Println("Available vaults:")
|
||||
|
||||
if len(vaults) == 0 {
|
||||
cmd.Println(" (none)")
|
||||
} else {
|
||||
// Try to get current vault for marking
|
||||
currentVault := ""
|
||||
if currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir); err == nil {
|
||||
|
||||
currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
|
||||
if err == nil {
|
||||
currentVault = currentVlt.GetName()
|
||||
}
|
||||
|
||||
@@ -175,19 +228,63 @@ func (cli *Instance) ListVaults(cmd *cobra.Command, jsonOutput bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// setMnemonicEnv sets the mnemonic environment variable and returns a
|
||||
// function that restores the previous value
|
||||
func setMnemonicEnv(mnemonicStr string) func() {
|
||||
originalMnemonic := os.Getenv(secret.EnvMnemonic)
|
||||
_ = os.Setenv(secret.EnvMnemonic, mnemonicStr)
|
||||
|
||||
return func() {
|
||||
if originalMnemonic != "" {
|
||||
_ = os.Setenv(secret.EnvMnemonic, originalMnemonic)
|
||||
} else {
|
||||
_ = os.Unsetenv(secret.EnvMnemonic)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// resolvePassphrase returns the unlock passphrase from the environment or
|
||||
// prompts the user for it with confirmation
|
||||
func resolvePassphrase() (*memguard.LockedBuffer, error) {
|
||||
if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" {
|
||||
secret.Debug("Using unlock passphrase from environment variable")
|
||||
|
||||
return memguard.NewBufferFromBytes([]byte(envPassphrase)), nil
|
||||
}
|
||||
|
||||
secret.Debug("Prompting user for unlock passphrase")
|
||||
|
||||
// Use secure passphrase input with confirmation
|
||||
passphraseBuffer, err := readSecurePassphrase("Enter passphrase for unlocker: ")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read passphrase: %w", err)
|
||||
}
|
||||
|
||||
return passphraseBuffer, nil
|
||||
}
|
||||
|
||||
// CreateVault creates a new vault
|
||||
func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error {
|
||||
secret.Debug("Creating new vault", "name", name, "state_dir", cli.stateDir)
|
||||
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
// Get or prompt for mnemonic
|
||||
var mnemonicStr string
|
||||
|
||||
if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" {
|
||||
secret.Debug("Using mnemonic from environment variable")
|
||||
|
||||
mnemonicStr = envMnemonic
|
||||
} else {
|
||||
secret.Debug("Prompting user for mnemonic phrase")
|
||||
// Read mnemonic securely without echo
|
||||
mnemonicBuffer, err := secret.ReadPassphrase("Enter your BIP39 mnemonic phrase: ")
|
||||
mnemonicBuffer, err := secret.ReadPassphrase(
|
||||
"Enter your BIP39 mnemonic phrase: ")
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read mnemonic from stdin", "error", err)
|
||||
|
||||
@@ -196,30 +293,33 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error {
|
||||
defer mnemonicBuffer.Destroy()
|
||||
|
||||
mnemonicStr = mnemonicBuffer.String()
|
||||
|
||||
fmt.Fprintln(os.Stderr) // Add newline after hidden input
|
||||
}
|
||||
|
||||
if mnemonicStr == "" {
|
||||
return fmt.Errorf("mnemonic cannot be empty")
|
||||
return errMnemonicEmpty
|
||||
}
|
||||
|
||||
// Validate the mnemonic
|
||||
mnemonicWords := strings.Fields(mnemonicStr)
|
||||
secret.Debug("Validating BIP39 mnemonic", "word_count", len(mnemonicWords))
|
||||
|
||||
if !bip39.IsMnemonicValid(mnemonicStr) {
|
||||
return fmt.Errorf("invalid BIP39 mnemonic phrase")
|
||||
return errInvalidMnemonicPhrase
|
||||
}
|
||||
|
||||
// Ask for the unlocker passphrase before creating the vault, so that
|
||||
// stopping at the prompt leaves no vault without an unlocker behind
|
||||
passphraseBuffer, err := resolvePassphrase()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
// Set mnemonic in environment for CreateVault to use
|
||||
originalMnemonic := os.Getenv(secret.EnvMnemonic)
|
||||
_ = os.Setenv(secret.EnvMnemonic, mnemonicStr)
|
||||
defer func() {
|
||||
if originalMnemonic != "" {
|
||||
_ = os.Setenv(secret.EnvMnemonic, originalMnemonic)
|
||||
} else {
|
||||
_ = os.Unsetenv(secret.EnvMnemonic)
|
||||
}
|
||||
}()
|
||||
restoreMnemonicEnv := setMnemonicEnv(mnemonicStr)
|
||||
defer restoreMnemonicEnv()
|
||||
|
||||
// Create the vault - it will handle key derivation internally
|
||||
vlt, err := vault.CreateVault(cli.fs, cli.stateDir, name)
|
||||
@@ -229,6 +329,7 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error {
|
||||
|
||||
// Get the vault metadata to retrieve the derivation index
|
||||
vaultDir := filepath.Join(cli.stateDir, "vaults.d", name)
|
||||
|
||||
metadata, err := vault.LoadVaultMetadata(cli.fs, vaultDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to load vault metadata: %w", err)
|
||||
@@ -243,23 +344,9 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error {
|
||||
// Unlock the vault with the derived long-term key
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Get or prompt for passphrase
|
||||
var passphraseBuffer *memguard.LockedBuffer
|
||||
if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" {
|
||||
secret.Debug("Using unlock passphrase from environment variable")
|
||||
passphraseBuffer = memguard.NewBufferFromBytes([]byte(envPassphrase))
|
||||
} else {
|
||||
secret.Debug("Prompting user for unlock passphrase")
|
||||
// Use secure passphrase input with confirmation
|
||||
passphraseBuffer, err = readSecurePassphrase("Enter passphrase for unlocker: ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read passphrase: %w", err)
|
||||
}
|
||||
}
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
// Create passphrase-protected unlocker
|
||||
secret.Debug("Creating passphrase-protected unlocker")
|
||||
|
||||
passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create unlocker: %w", err)
|
||||
@@ -274,7 +361,14 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error {
|
||||
|
||||
// SelectVault selects a vault as the current one
|
||||
func (cli *Instance) SelectVault(cmd *cobra.Command, name string) error {
|
||||
if err := vault.SelectVault(cli.fs, cli.stateDir, name); err != nil {
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
err = vault.SelectVault(cli.fs, cli.stateDir, name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -283,84 +377,64 @@ func (cli *Instance) SelectVault(cmd *cobra.Command, name string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// VaultImport imports a mnemonic into a specific vault
|
||||
func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error {
|
||||
secret.Debug("Importing mnemonic into vault", "vault_name", vaultName, "state_dir", cli.stateDir)
|
||||
|
||||
// Get the specific vault by name
|
||||
vlt := vault.NewVault(cli.fs, cli.stateDir, vaultName)
|
||||
|
||||
// vaultImportPreflight verifies the vault exists without a long-term key
|
||||
// and returns the vault directory, public key path, and validated mnemonic
|
||||
func (cli *Instance) vaultImportPreflight(
|
||||
vlt *vault.Vault, vaultName string,
|
||||
) (string, string, string, error) {
|
||||
// Check if vault exists
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
return err
|
||||
return "", "", "", err
|
||||
}
|
||||
|
||||
exists, err := afero.DirExists(cli.fs, vaultDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if vault exists: %w", err)
|
||||
return "", "", "", fmt.Errorf("failed to check if vault exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("vault '%s' does not exist", vaultName)
|
||||
return "", "", "", fmt.Errorf("vault '%s' %w",
|
||||
vaultName, errVaultDoesNotExist)
|
||||
}
|
||||
|
||||
// Check if vault already has a public key
|
||||
pubKeyPath := fmt.Sprintf("%s/pub.age", vaultDir)
|
||||
if _, err := cli.fs.Stat(pubKeyPath); err == nil {
|
||||
return fmt.Errorf("vault '%s' already has a long-term key configured", vaultName)
|
||||
pubKeyPath := vaultDir + "/pub.age"
|
||||
|
||||
exists, err = afero.Exists(cli.fs, pubKeyPath)
|
||||
if err != nil {
|
||||
return "", "", "", fmt.Errorf("failed to check %s: %w", pubKeyPath, err)
|
||||
}
|
||||
|
||||
if exists {
|
||||
return "", "", "", fmt.Errorf("vault '%s' %w",
|
||||
vaultName, errVaultHasLongTermKey)
|
||||
}
|
||||
|
||||
// Get mnemonic from environment
|
||||
mnemonic := os.Getenv(secret.EnvMnemonic)
|
||||
if mnemonic == "" {
|
||||
return fmt.Errorf("SB_SECRET_MNEMONIC environment variable not set")
|
||||
return "", "", "", errMnemonicEnvNotSet
|
||||
}
|
||||
|
||||
// Validate the mnemonic
|
||||
mnemonicWords := strings.Fields(mnemonic)
|
||||
secret.Debug("Validating BIP39 mnemonic", "word_count", len(mnemonicWords))
|
||||
|
||||
if !bip39.IsMnemonicValid(mnemonic) {
|
||||
return fmt.Errorf("invalid BIP39 mnemonic")
|
||||
return "", "", "", errInvalidMnemonic
|
||||
}
|
||||
|
||||
// Get the next available derivation index for this mnemonic
|
||||
derivationIndex, err := vault.GetNextDerivationIndex(cli.fs, cli.stateDir, mnemonic)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get next derivation index", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to get next derivation index: %w", err)
|
||||
}
|
||||
secret.Debug("Using derivation index", "index", derivationIndex)
|
||||
|
||||
// Derive long-term key from mnemonic with the appropriate index
|
||||
secret.Debug("Deriving long-term key from mnemonic", "index", derivationIndex)
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonic, derivationIndex)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to derive long-term key: %w", err)
|
||||
}
|
||||
|
||||
// Store long-term public key in vault
|
||||
ltPublicKey := ltIdentity.Recipient().String()
|
||||
secret.Debug("Storing long-term public key", "pubkey", ltPublicKey, "vault_dir", vaultDir)
|
||||
|
||||
if err := afero.WriteFile(cli.fs, pubKeyPath, []byte(ltPublicKey), secret.FilePerms); err != nil {
|
||||
return fmt.Errorf("failed to store long-term public key: %w", err)
|
||||
}
|
||||
|
||||
// Calculate public key hash from the actual derivation index being used
|
||||
// This is used to verify that the derived key matches what was stored
|
||||
publicKeyHash := vault.ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
|
||||
// Calculate family hash from index 0 (same for all vaults with this mnemonic)
|
||||
// This is used to identify which vaults belong to the same mnemonic family
|
||||
identity0, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to derive identity for index 0: %w", err)
|
||||
}
|
||||
familyHash := vault.ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
return vaultDir, pubKeyPath, mnemonic, nil
|
||||
}
|
||||
|
||||
// updateVaultImportMetadata stores the derivation info in vault metadata
|
||||
func updateVaultImportMetadata(
|
||||
fs afero.Fs, vaultDir string, derivationIndex uint32,
|
||||
publicKeyHash, familyHash string,
|
||||
) error {
|
||||
// Load existing metadata
|
||||
existingMetadata, err := vault.LoadVaultMetadata(cli.fs, vaultDir)
|
||||
existingMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
// If metadata doesn't exist, create new
|
||||
existingMetadata = &vault.Metadata{
|
||||
@@ -373,17 +447,101 @@ func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error {
|
||||
existingMetadata.PublicKeyHash = publicKeyHash
|
||||
existingMetadata.MnemonicFamilyHash = familyHash
|
||||
|
||||
if err := vault.SaveVaultMetadata(cli.fs, vaultDir, existingMetadata); err != nil {
|
||||
err = vault.SaveVaultMetadata(fs, vaultDir, existingMetadata)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to save vault metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to save vault metadata: %w", err)
|
||||
}
|
||||
|
||||
secret.Debug("Saved vault metadata with derivation index and public key hash")
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// VaultImport imports a mnemonic into a specific vault, holding the state
|
||||
// directory lock while importMnemonic runs
|
||||
func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error {
|
||||
err := vault.ValidateVaultName(vaultName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
return cli.importMnemonic(cmd, vaultName)
|
||||
}
|
||||
|
||||
// importMnemonic gives the vault a long-term key derived from the mnemonic
|
||||
// and a passphrase unlocker
|
||||
func (cli *Instance) importMnemonic(cmd *cobra.Command, vaultName string) error {
|
||||
secret.Debug("Importing mnemonic into vault",
|
||||
"vault_name", vaultName, "state_dir", cli.stateDir)
|
||||
|
||||
// Get the specific vault by name
|
||||
vlt := vault.NewVault(cli.fs, cli.stateDir, vaultName)
|
||||
|
||||
vaultDir, pubKeyPath, mnemonic, err := cli.vaultImportPreflight(vlt, vaultName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get the next available derivation index for this mnemonic
|
||||
derivationIndex, err := vault.GetNextDerivationIndex(cli.fs, cli.stateDir, mnemonic)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get next derivation index", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to get next derivation index: %w", err)
|
||||
}
|
||||
|
||||
secret.Debug("Using derivation index", "index", derivationIndex)
|
||||
|
||||
// Derive long-term key from mnemonic with the appropriate index
|
||||
secret.Debug("Deriving long-term key from mnemonic", "index", derivationIndex)
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonic, derivationIndex)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to derive long-term key: %w", err)
|
||||
}
|
||||
|
||||
// Store long-term public key in vault
|
||||
ltPublicKey := ltIdentity.Recipient().String()
|
||||
secret.Debug("Storing long-term public key",
|
||||
"pubkey", ltPublicKey, "vault_dir", vaultDir)
|
||||
|
||||
err = secret.WriteFileAtomic(cli.fs, pubKeyPath, []byte(ltPublicKey))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to store long-term public key: %w", err)
|
||||
}
|
||||
|
||||
// Calculate public key hash from the actual derivation index being used
|
||||
// This is used to verify that the derived key matches what was stored
|
||||
publicKeyHash := vault.ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
|
||||
// Calculate family hash from index 0 (same for all vaults with this
|
||||
// mnemonic). This is used to identify which vaults belong to the same
|
||||
// mnemonic family.
|
||||
identity0, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to derive identity for index 0: %w", err)
|
||||
}
|
||||
|
||||
familyHash := vault.ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
err = updateVaultImportMetadata(
|
||||
cli.fs, vaultDir, derivationIndex, publicKeyHash, familyHash)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get passphrase from environment variable
|
||||
passphraseStr := os.Getenv(secret.EnvUnlockPassphrase)
|
||||
if passphraseStr == "" {
|
||||
return fmt.Errorf("SB_UNLOCK_PASSPHRASE environment variable not set")
|
||||
return errPassphraseEnvNotSet
|
||||
}
|
||||
|
||||
secret.Debug("Using unlock passphrase from environment variable")
|
||||
@@ -397,6 +555,7 @@ func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error {
|
||||
|
||||
// Create passphrase-protected unlocker
|
||||
secret.Debug("Creating passphrase-protected unlocker")
|
||||
|
||||
passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to create unlocker", "error", err)
|
||||
@@ -411,8 +570,74 @@ func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveVault removes a vault with safety checks
|
||||
// vaultHasSecrets reports whether the vault directory contains any secrets
|
||||
func (cli *Instance) vaultHasSecrets(vaultDir string) (bool, error) {
|
||||
secretsDir := filepath.Join(vaultDir, "secrets.d")
|
||||
|
||||
exists, err := afero.DirExists(cli.fs, secretsDir)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed to check secrets directory %s: %w",
|
||||
secretsDir, err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
entries, err := afero.ReadDir(cli.fs, secretsDir)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed to read secrets directory %s: %w",
|
||||
secretsDir, err)
|
||||
}
|
||||
|
||||
return len(entries) > 0, nil
|
||||
}
|
||||
|
||||
// switchAwayFromVault selects another vault as current before removal
|
||||
func (cli *Instance) switchAwayFromVault(
|
||||
cmd *cobra.Command, vaults []string, name string,
|
||||
) error {
|
||||
// Find another vault to switch to
|
||||
var newVault string
|
||||
|
||||
for _, v := range vaults {
|
||||
if v != name {
|
||||
newVault = v
|
||||
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Switch to the new vault
|
||||
err := vault.SelectVault(cli.fs, cli.stateDir, newVault)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to switch to vault '%s': %w", newVault, err)
|
||||
}
|
||||
|
||||
cmd.Printf("Switched current vault to '%s'\n", newVault)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveVault removes a vault, holding the state directory lock while
|
||||
// removeVault runs
|
||||
func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) error {
|
||||
err := vault.ValidateVaultName(name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
return cli.removeVault(cmd, name, force)
|
||||
}
|
||||
|
||||
// removeVault removes a vault with safety checks
|
||||
func (cli *Instance) removeVault(cmd *cobra.Command, name string, force bool) error {
|
||||
// Get list of all vaults
|
||||
vaults, err := vault.ListVaults(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
@@ -420,21 +645,13 @@ func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) er
|
||||
}
|
||||
|
||||
// Check if vault exists
|
||||
vaultExists := false
|
||||
for _, v := range vaults {
|
||||
if v == name {
|
||||
vaultExists = true
|
||||
|
||||
break
|
||||
}
|
||||
}
|
||||
if !vaultExists {
|
||||
return fmt.Errorf("vault '%s' does not exist", name)
|
||||
if !slices.Contains(vaults, name) {
|
||||
return fmt.Errorf("vault '%s' %w", name, errVaultDoesNotExist)
|
||||
}
|
||||
|
||||
// Don't allow removing the last vault
|
||||
if len(vaults) == 1 {
|
||||
return fmt.Errorf("cannot remove the last vault")
|
||||
return errCannotRemoveLastVault
|
||||
}
|
||||
|
||||
// Check if this is the current vault
|
||||
@@ -442,57 +659,47 @@ func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) er
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get current vault: %w", err)
|
||||
}
|
||||
|
||||
isCurrentVault := currentVault.GetName() == name
|
||||
|
||||
// Load the vault to check for secrets
|
||||
vlt := vault.NewVault(cli.fs, cli.stateDir, name)
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Check if vault has secrets
|
||||
secretsDir := filepath.Join(vaultDir, "secrets.d")
|
||||
hasSecrets := false
|
||||
if exists, _ := afero.DirExists(cli.fs, secretsDir); exists {
|
||||
entries, err := afero.ReadDir(cli.fs, secretsDir)
|
||||
if err == nil && len(entries) > 0 {
|
||||
hasSecrets = true
|
||||
}
|
||||
hasSecrets, err := cli.vaultHasSecrets(vaultDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Require --force if vault has secrets
|
||||
if hasSecrets && !force {
|
||||
return fmt.Errorf("vault '%s' contains secrets; use --force to remove", name)
|
||||
return fmt.Errorf("vault '%s' %w", name, errVaultContainsSecrets)
|
||||
}
|
||||
|
||||
// If removing current vault, switch to another vault first
|
||||
if isCurrentVault {
|
||||
// Find another vault to switch to
|
||||
var newVault string
|
||||
for _, v := range vaults {
|
||||
if v != name {
|
||||
newVault = v
|
||||
|
||||
break
|
||||
}
|
||||
err = cli.switchAwayFromVault(cmd, vaults, name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Switch to the new vault
|
||||
if err := vault.SelectVault(cli.fs, cli.stateDir, newVault); err != nil {
|
||||
return fmt.Errorf("failed to switch to vault '%s': %w", newVault, err)
|
||||
}
|
||||
cmd.Printf("Switched current vault to '%s'\n", newVault)
|
||||
}
|
||||
|
||||
// Remove the vault directory
|
||||
if err := cli.fs.RemoveAll(vaultDir); err != nil {
|
||||
err = secret.RemoveDirAtomic(cli.fs, vaultDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to remove vault directory: %w", err)
|
||||
}
|
||||
|
||||
cmd.Printf("Removed vault '%s'\n", name)
|
||||
|
||||
if hasSecrets {
|
||||
cmd.Printf("Warning: Vault contained secrets that have been permanently deleted\n")
|
||||
cmd.Printf("Warning: Vault contained secrets that have been " +
|
||||
"permanently deleted\n")
|
||||
}
|
||||
|
||||
return nil
|
||||
|
||||
+134
-61
@@ -1,11 +1,16 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"text/tabwriter"
|
||||
"time"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/spf13/afero"
|
||||
@@ -16,9 +21,18 @@ const (
|
||||
tabWriterPadding = 2
|
||||
)
|
||||
|
||||
// Sentinel errors for version operations
|
||||
var (
|
||||
errVersionNotFound = errors.New("not found for secret")
|
||||
errCannotRemoveCurrentVersion = errors.New("promote another version first")
|
||||
)
|
||||
|
||||
// newVersionCmd returns the version management command
|
||||
func newVersionCmd() *cobra.Command {
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
log.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
return VersionCommands(cli)
|
||||
}
|
||||
@@ -28,7 +42,8 @@ func VersionCommands(cli *Instance) *cobra.Command {
|
||||
versionCmd := &cobra.Command{
|
||||
Use: "version",
|
||||
Short: "Manage secret versions",
|
||||
Long: "Commands for managing secret versions including listing, promoting, and retrieving specific versions",
|
||||
Long: "Commands for managing secret versions including listing, " +
|
||||
"promoting, and retrieving specific versions",
|
||||
}
|
||||
|
||||
// List versions command
|
||||
@@ -47,14 +62,17 @@ func VersionCommands(cli *Instance) *cobra.Command {
|
||||
promoteCmd := &cobra.Command{
|
||||
Use: "promote <secret-name> <version>",
|
||||
Short: "Promote a specific version to current",
|
||||
Long: "Updates the current symlink to point to the specified version without modifying timestamps",
|
||||
Args: cobra.ExactArgs(2), //nolint:mnd // Command requires exactly 2 arguments: secret-name and version
|
||||
ValidArgsFunction: func(cmd *cobra.Command, args []string, toComplete string) ([]string, cobra.ShellCompDirective) {
|
||||
Long: "Updates the current symlink to point to the specified " +
|
||||
"version without modifying timestamps",
|
||||
Args: cobra.ExactArgs(2), //nolint:mnd // secret-name and version args
|
||||
ValidArgsFunction: func(
|
||||
cmd *cobra.Command, args []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
// Complete secret name for first arg
|
||||
if len(args) == 0 {
|
||||
return getSecretNamesCompletionFunc(cli.fs, cli.stateDir)(cmd, args, toComplete)
|
||||
}
|
||||
// TODO: Complete version numbers for second arg
|
||||
// Version number completion for the second arg is not implemented
|
||||
return nil, cobra.ShellCompDirectiveNoFileComp
|
||||
},
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
@@ -67,14 +85,17 @@ func VersionCommands(cli *Instance) *cobra.Command {
|
||||
Use: "remove <secret-name> <version>",
|
||||
Aliases: []string{"rm"},
|
||||
Short: "Remove a specific version of a secret",
|
||||
Long: "Remove a specific version of a secret. Cannot remove the current version.",
|
||||
Args: cobra.ExactArgs(2), //nolint:mnd // Command requires exactly 2 arguments: secret-name and version
|
||||
ValidArgsFunction: func(cmd *cobra.Command, args []string, toComplete string) ([]string, cobra.ShellCompDirective) {
|
||||
Long: "Remove a specific version of a secret. Cannot remove the " +
|
||||
"current version.",
|
||||
Args: cobra.ExactArgs(2), //nolint:mnd // secret-name and version args
|
||||
ValidArgsFunction: func(
|
||||
cmd *cobra.Command, args []string, toComplete string,
|
||||
) ([]string, cobra.ShellCompDirective) {
|
||||
// Complete secret name for first arg
|
||||
if len(args) == 0 {
|
||||
return getSecretNamesCompletionFunc(cli.fs, cli.stateDir)(cmd, args, toComplete)
|
||||
}
|
||||
// TODO: Complete version numbers for second arg
|
||||
// Version number completion for the second arg is not implemented
|
||||
return nil, cobra.ShellCompDirectiveNoFileComp
|
||||
},
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
@@ -91,6 +112,11 @@ func VersionCommands(cli *Instance) *cobra.Command {
|
||||
func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error {
|
||||
secret.Debug("ListVersions called", "secret_name", secretName)
|
||||
|
||||
err := vault.ValidateSecretName(secretName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get current vault
|
||||
vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
@@ -117,10 +143,11 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error {
|
||||
|
||||
return fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Debug("Secret not found", "secret_name", secretName)
|
||||
|
||||
return fmt.Errorf("secret '%s' not found", secretName)
|
||||
return fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound)
|
||||
}
|
||||
|
||||
// List all versions
|
||||
@@ -141,6 +168,7 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error {
|
||||
currentVersion, err := secret.GetCurrentVersion(cli.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get current version", "error", err)
|
||||
|
||||
currentVersion = ""
|
||||
}
|
||||
|
||||
@@ -156,44 +184,7 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error {
|
||||
|
||||
// Load and display each version's metadata
|
||||
for _, version := range versions {
|
||||
sv := secret.NewVersion(vlt, secretName, version)
|
||||
|
||||
// Load metadata
|
||||
if err := sv.LoadMetadata(ltIdentity); err != nil {
|
||||
secret.Debug("Failed to load version metadata", "version", version, "error", err)
|
||||
// Display version with error
|
||||
status := "error"
|
||||
if version == currentVersion {
|
||||
status = "current (error)"
|
||||
}
|
||||
_, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n", version, "-", status, "-", "-")
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
// Determine status
|
||||
status := "expired"
|
||||
if version == currentVersion {
|
||||
status = "current"
|
||||
}
|
||||
|
||||
// Format timestamps
|
||||
createdAt := "-"
|
||||
if sv.Metadata.CreatedAt != nil {
|
||||
createdAt = sv.Metadata.CreatedAt.Format("2006-01-02 15:04:05")
|
||||
}
|
||||
|
||||
notBefore := "-"
|
||||
if sv.Metadata.NotBefore != nil {
|
||||
notBefore = sv.Metadata.NotBefore.Format("2006-01-02 15:04:05")
|
||||
}
|
||||
|
||||
notAfter := "-"
|
||||
if sv.Metadata.NotAfter != nil {
|
||||
notAfter = sv.Metadata.NotAfter.Format("2006-01-02 15:04:05")
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n", version, createdAt, status, notBefore, notAfter)
|
||||
printVersionRow(w, vlt, secretName, version, currentVersion, ltIdentity)
|
||||
}
|
||||
|
||||
_ = w.Flush()
|
||||
@@ -201,8 +192,69 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// printVersionRow loads one version's metadata and writes its table row
|
||||
func printVersionRow(
|
||||
w io.Writer, vlt *vault.Vault,
|
||||
secretName, version, currentVersion string,
|
||||
ltIdentity *age.X25519Identity,
|
||||
) {
|
||||
sv := secret.NewVersion(vlt, secretName, version)
|
||||
|
||||
// Load metadata
|
||||
err := sv.LoadMetadata(ltIdentity)
|
||||
if err != nil {
|
||||
secret.Warn("Failed to load version metadata",
|
||||
"version", version, "error", err)
|
||||
// Display version with error
|
||||
status := "error"
|
||||
if version == currentVersion {
|
||||
status = "current (error)"
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n", version, "-", status, "-", "-")
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
// Determine status
|
||||
status := "expired"
|
||||
if version == currentVersion {
|
||||
status = "current"
|
||||
}
|
||||
|
||||
// Format timestamps
|
||||
createdAt := formatVersionTime(sv.Metadata.CreatedAt)
|
||||
notBefore := formatVersionTime(sv.Metadata.NotBefore)
|
||||
notAfter := formatVersionTime(sv.Metadata.NotAfter)
|
||||
|
||||
_, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n",
|
||||
version, createdAt, status, notBefore, notAfter)
|
||||
}
|
||||
|
||||
// formatVersionTime formats an optional version timestamp, "-" when unset
|
||||
func formatVersionTime(t *time.Time) string {
|
||||
if t == nil {
|
||||
return "-"
|
||||
}
|
||||
|
||||
return t.Format("2006-01-02 15:04:05")
|
||||
}
|
||||
|
||||
// PromoteVersion promotes a specific version to current
|
||||
func (cli *Instance) PromoteVersion(cmd *cobra.Command, secretName string, version string) error {
|
||||
func (cli *Instance) PromoteVersion(
|
||||
cmd *cobra.Command, secretName string, version string,
|
||||
) error {
|
||||
err := vault.ValidateSecretName(secretName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
// Get current vault
|
||||
vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
@@ -219,17 +271,19 @@ func (cli *Instance) PromoteVersion(cmd *cobra.Command, secretName string, versi
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", encodedName)
|
||||
|
||||
// Check if version exists
|
||||
versionDir := filepath.Join(secretDir, "versions", version)
|
||||
exists, err := afero.DirExists(cli.fs, versionDir)
|
||||
exists, err := secret.VersionExists(cli.fs, secretDir, version)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if version exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("version '%s' not found for secret '%s'", version, secretName)
|
||||
return fmt.Errorf("version '%s' %w '%s'",
|
||||
version, errVersionNotFound, secretName)
|
||||
}
|
||||
|
||||
// Update the current symlink using the proper function
|
||||
if err := secret.SetCurrentVersion(cli.fs, secretDir, version); err != nil {
|
||||
err = secret.SetCurrentVersion(cli.fs, secretDir, version)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update current version: %w", err)
|
||||
}
|
||||
|
||||
@@ -239,7 +293,20 @@ func (cli *Instance) PromoteVersion(cmd *cobra.Command, secretName string, versi
|
||||
}
|
||||
|
||||
// RemoveVersion removes a specific version of a secret
|
||||
func (cli *Instance) RemoveVersion(cmd *cobra.Command, secretName string, version string) error {
|
||||
func (cli *Instance) RemoveVersion(
|
||||
cmd *cobra.Command, secretName string, version string,
|
||||
) error {
|
||||
err := vault.ValidateSecretName(secretName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
|
||||
// Get current vault
|
||||
vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
|
||||
if err != nil {
|
||||
@@ -260,18 +327,20 @@ func (cli *Instance) RemoveVersion(cmd *cobra.Command, secretName string, versio
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("secret '%s' not found", secretName)
|
||||
return fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound)
|
||||
}
|
||||
|
||||
// Check if version exists
|
||||
versionDir := filepath.Join(secretDir, "versions", version)
|
||||
exists, err = afero.DirExists(cli.fs, versionDir)
|
||||
exists, err = secret.VersionExists(cli.fs, secretDir, version)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if version exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("version '%s' not found for secret '%s'", version, secretName)
|
||||
return fmt.Errorf("version '%s' %w '%s'",
|
||||
version, errVersionNotFound, secretName)
|
||||
}
|
||||
|
||||
// Get current version
|
||||
@@ -282,11 +351,15 @@ func (cli *Instance) RemoveVersion(cmd *cobra.Command, secretName string, versio
|
||||
|
||||
// Don't allow removing the current version
|
||||
if version == currentVersion {
|
||||
return fmt.Errorf("cannot remove the current version '%s'; promote another version first", version)
|
||||
return fmt.Errorf("cannot remove the current version '%s'; %w",
|
||||
version, errCannotRemoveCurrentVersion)
|
||||
}
|
||||
|
||||
// Remove the version directory
|
||||
if err := cli.fs.RemoveAll(versionDir); err != nil {
|
||||
versionDir := filepath.Join(secretDir, "versions", version)
|
||||
|
||||
err = secret.RemoveDirAtomic(cli.fs, versionDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to remove version: %w", err)
|
||||
}
|
||||
|
||||
|
||||
+117
-36
@@ -7,6 +7,7 @@
|
||||
// - TestPromoteVersionCommand: Tests `secret version promote` command
|
||||
// - TestPromoteNonExistentVersion: Tests error handling for invalid promotion
|
||||
// - TestGetSecretWithVersion: Tests `secret get --version` flag functionality
|
||||
// - TestGetSecretWritesBinaryValue: Tests `secret get` output of binary values
|
||||
// - TestVersionCommandStructure: Tests command structure and help text
|
||||
// - TestListVersionsEmptyOutput: Tests edge case with no versions
|
||||
//
|
||||
@@ -14,6 +15,7 @@
|
||||
// - setupTestVault(): CLI test helper for vault initialization
|
||||
// - Uses consistent test mnemonic for reproducible testing
|
||||
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package cli
|
||||
|
||||
import (
|
||||
@@ -22,6 +24,7 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
@@ -32,29 +35,41 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// Helper function to add a secret to vault with proper buffer protection
|
||||
func addTestSecret(t *testing.T, vlt *vault.Vault, name string, value []byte, force bool) {
|
||||
const (
|
||||
// testMnemonic is the standard BIP39 mnemonic used for CLI tests.
|
||||
//nolint:dupword // BIP39 test mnemonic intentionally repeats a word
|
||||
testMnemonic = "abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon abandon abandon about"
|
||||
|
||||
// testStateDir is the in-memory state directory used by CLI tests.
|
||||
testStateDir = "/test/state"
|
||||
)
|
||||
|
||||
// Helper function to add a version of the "test/secret" secret to the
|
||||
// vault with proper buffer protection
|
||||
func addTestSecret(t *testing.T, vlt *vault.Vault, value []byte, force bool) {
|
||||
t.Helper()
|
||||
|
||||
buffer := memguard.NewBufferFromBytes(value)
|
||||
defer buffer.Destroy()
|
||||
err := vlt.AddSecret(name, buffer, force)
|
||||
|
||||
err := vlt.AddSecret("test/secret", buffer, force)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// Helper function to set up a vault with long-term key
|
||||
func setupTestVault(t *testing.T, fs afero.Fs, stateDir string) {
|
||||
// Helper function to set up a vault with long-term key in testStateDir
|
||||
func setupTestVault(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
// Set mnemonic for testing
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
// Create vault
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "default")
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, "default")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Derive and store long-term key from mnemonic
|
||||
mnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Store long-term public key in vault
|
||||
@@ -64,30 +79,32 @@ func setupTestVault(t *testing.T, fs afero.Fs, stateDir string) {
|
||||
require.NoError(t, err)
|
||||
|
||||
// Select vault
|
||||
err = vault.SelectVault(fs, stateDir, "default")
|
||||
err = vault.SelectVault(fs, testStateDir, "default")
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
//nolint:paralleltest // uses t.Setenv via setupTestVault
|
||||
func TestListVersionsCommand(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
stateDir := testStateDir
|
||||
cli := NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
|
||||
// Set up vault with long-term key
|
||||
setupTestVault(t, fs, stateDir)
|
||||
setupTestVault(t, fs)
|
||||
|
||||
// Add a secret with multiple versions
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
addTestSecret(t, vlt, "test/secret", []byte("version-1"), false)
|
||||
addTestSecret(t, vlt, []byte("version-1"), false)
|
||||
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecret(t, vlt, "test/secret", []byte("version-2"), true)
|
||||
addTestSecret(t, vlt, []byte("version-2"), true)
|
||||
|
||||
// Create a command for output capture
|
||||
cmd := newRootCmd()
|
||||
|
||||
var buf bytes.Buffer
|
||||
cmd.SetOut(&buf)
|
||||
cmd.SetErr(&buf)
|
||||
@@ -112,24 +129,28 @@ func TestListVersionsCommand(t *testing.T) {
|
||||
// Should have two version entries
|
||||
lines := strings.Split(outputStr, "\n")
|
||||
versionLines := 0
|
||||
|
||||
for _, line := range lines {
|
||||
if strings.Contains(line, ".001") || strings.Contains(line, ".002") {
|
||||
versionLines++
|
||||
}
|
||||
}
|
||||
|
||||
assert.Equal(t, 2, versionLines)
|
||||
}
|
||||
|
||||
//nolint:paralleltest // uses t.Setenv via setupTestVault
|
||||
func TestListVersionsNonExistentSecret(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
stateDir := testStateDir
|
||||
cli := NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
|
||||
// Set up vault with long-term key
|
||||
setupTestVault(t, fs, stateDir)
|
||||
setupTestVault(t, fs)
|
||||
|
||||
// Create a command for output capture
|
||||
cmd := newRootCmd()
|
||||
|
||||
var buf bytes.Buffer
|
||||
cmd.SetOut(&buf)
|
||||
cmd.SetErr(&buf)
|
||||
@@ -140,23 +161,24 @@ func TestListVersionsNonExistentSecret(t *testing.T) {
|
||||
assert.Contains(t, err.Error(), "not found")
|
||||
}
|
||||
|
||||
//nolint:paralleltest // uses t.Setenv via setupTestVault
|
||||
func TestPromoteVersionCommand(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
stateDir := testStateDir
|
||||
cli := NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
|
||||
// Set up vault with long-term key
|
||||
setupTestVault(t, fs, stateDir)
|
||||
setupTestVault(t, fs)
|
||||
|
||||
// Add a secret with multiple versions
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
addTestSecret(t, vlt, "test/secret", []byte("version-1"), false)
|
||||
addTestSecret(t, vlt, []byte("version-1"), false)
|
||||
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecret(t, vlt, "test/secret", []byte("version-2"), true)
|
||||
addTestSecret(t, vlt, []byte("version-2"), true)
|
||||
|
||||
// Get versions
|
||||
vaultDir, _ := vlt.GetDirectory()
|
||||
@@ -168,13 +190,17 @@ func TestPromoteVersionCommand(t *testing.T) {
|
||||
// Current should be version-2
|
||||
value, err := vlt.GetSecret("test/secret")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2"), value)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-2"), value.Bytes())
|
||||
|
||||
// Promote first version
|
||||
firstVersion := versions[1] // Older version
|
||||
|
||||
// Create a command for output capture
|
||||
cmd := newRootCmd()
|
||||
|
||||
var buf bytes.Buffer
|
||||
cmd.SetOut(&buf)
|
||||
cmd.SetErr(&buf)
|
||||
@@ -190,27 +216,32 @@ func TestPromoteVersionCommand(t *testing.T) {
|
||||
assert.Contains(t, outputStr, firstVersion)
|
||||
|
||||
// Verify current is now version-1
|
||||
value, err = vlt.GetSecret("test/secret")
|
||||
promoted, err := vlt.GetSecret("test/secret")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1"), value)
|
||||
|
||||
defer promoted.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-1"), promoted.Bytes())
|
||||
}
|
||||
|
||||
//nolint:paralleltest // uses t.Setenv via setupTestVault
|
||||
func TestPromoteNonExistentVersion(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
stateDir := testStateDir
|
||||
cli := NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
|
||||
// Set up vault with long-term key
|
||||
setupTestVault(t, fs, stateDir)
|
||||
setupTestVault(t, fs)
|
||||
|
||||
// Add a secret
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
addTestSecret(t, vlt, "test/secret", []byte("value"), false)
|
||||
addTestSecret(t, vlt, []byte("value"), false)
|
||||
|
||||
// Create a command for output capture
|
||||
cmd := newRootCmd()
|
||||
|
||||
var buf bytes.Buffer
|
||||
cmd.SetOut(&buf)
|
||||
cmd.SetErr(&buf)
|
||||
@@ -221,23 +252,24 @@ func TestPromoteNonExistentVersion(t *testing.T) {
|
||||
assert.Contains(t, err.Error(), "not found")
|
||||
}
|
||||
|
||||
//nolint:paralleltest // uses t.Setenv via setupTestVault
|
||||
func TestGetSecretWithVersion(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
stateDir := testStateDir
|
||||
cli := NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
|
||||
// Set up vault with long-term key
|
||||
setupTestVault(t, fs, stateDir)
|
||||
setupTestVault(t, fs)
|
||||
|
||||
// Add a secret with multiple versions
|
||||
vlt, err := vault.GetCurrentVault(fs, stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
addTestSecret(t, vlt, "test/secret", []byte("version-1"), false)
|
||||
addTestSecret(t, vlt, []byte("version-1"), false)
|
||||
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecret(t, vlt, "test/secret", []byte("version-2"), true)
|
||||
addTestSecret(t, vlt, []byte("version-2"), true)
|
||||
|
||||
// Get versions
|
||||
vaultDir, _ := vlt.GetDirectory()
|
||||
@@ -248,25 +280,72 @@ func TestGetSecretWithVersion(t *testing.T) {
|
||||
|
||||
// Create a command for output capture
|
||||
cmd := newRootCmd()
|
||||
|
||||
var buf bytes.Buffer
|
||||
cmd.SetOut(&buf)
|
||||
|
||||
// Test getting current version (empty version string)
|
||||
err = cli.GetSecretWithVersion(cmd, "test/secret", "")
|
||||
// Test getting the current version
|
||||
err = cli.GetSecret(cmd, "test/secret")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "version-2", buf.String())
|
||||
|
||||
// Test getting specific version
|
||||
buf.Reset()
|
||||
|
||||
firstVersion := versions[1] // Older version
|
||||
err = cli.GetSecretWithVersion(cmd, "test/secret", firstVersion)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "version-1", buf.String())
|
||||
}
|
||||
|
||||
//nolint:paralleltest // uses t.Setenv via setupTestVault
|
||||
func TestGetSecretWritesBinaryValue(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
cli := NewCLIInstanceWithStateDir(fs, testStateDir)
|
||||
|
||||
setupTestVault(t, fs)
|
||||
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
value := []byte{0x00, 'a', 0xff, 0xfe, 0x00, 0xc3, 0x28, 'z', 0x00}
|
||||
require.False(t, utf8.Valid(value))
|
||||
// A copy, since storing a value wipes the slice it came from
|
||||
addTestSecret(t, vlt, bytes.Clone(value), false)
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
|
||||
versions, err := secret.ListVersions(fs,
|
||||
filepath.Join(vaultDir, "secrets.d", "test%secret"))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 1)
|
||||
|
||||
cmd := newRootCmd()
|
||||
|
||||
var buf bytes.Buffer
|
||||
cmd.SetOut(&buf)
|
||||
|
||||
// Each writes exactly the stored bytes, with no trailing newline
|
||||
err = cli.GetSecret(cmd, "test/secret")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, value, buf.Bytes())
|
||||
|
||||
buf.Reset()
|
||||
|
||||
err = cli.GetSecretWithVersion(cmd, "test/secret", versions[0])
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, value, buf.Bytes())
|
||||
}
|
||||
|
||||
//nolint:paralleltest // reads process environment to determine the state dir
|
||||
func TestVersionCommandStructure(t *testing.T) {
|
||||
// Test that version commands are properly structured
|
||||
cli := NewCLIInstance()
|
||||
cli, err := NewCLIInstance()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to initialize CLI: %v", err)
|
||||
}
|
||||
|
||||
cmd := VersionCommands(cli)
|
||||
|
||||
assert.Equal(t, "version", cmd.Use)
|
||||
@@ -282,13 +361,14 @@ func TestVersionCommandStructure(t *testing.T) {
|
||||
assert.Equal(t, "Promote a specific version to current", promoteCmd.Short)
|
||||
}
|
||||
|
||||
//nolint:paralleltest // uses t.Setenv via setupTestVault
|
||||
func TestListVersionsEmptyOutput(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
stateDir := testStateDir
|
||||
cli := NewCLIInstanceWithStateDir(fs, stateDir)
|
||||
|
||||
// Set up vault with long-term key
|
||||
setupTestVault(t, fs, stateDir)
|
||||
setupTestVault(t, fs)
|
||||
|
||||
// Create a secret directory without versions (edge case)
|
||||
vaultDir := stateDir + "/vaults.d/default"
|
||||
@@ -298,6 +378,7 @@ func TestListVersionsEmptyOutput(t *testing.T) {
|
||||
|
||||
// Create a command for output capture
|
||||
cmd := newRootCmd()
|
||||
|
||||
var buf bytes.Buffer
|
||||
cmd.SetOut(&buf)
|
||||
cmd.SetErr(&buf)
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
//go:build darwin
|
||||
|
||||
// Package macse provides Go bindings for macOS Secure Enclave operations
|
||||
// using CryptoTokenKit identities created via sc_auth.
|
||||
// Key creation and deletion shell out to sc_auth (which has SE entitlements).
|
||||
// Encrypt/decrypt use Security.framework ECIES directly (works unsigned).
|
||||
package macse
|
||||
|
||||
/*
|
||||
#cgo CFLAGS: -x objective-c -fobjc-arc
|
||||
#cgo LDFLAGS: -framework Security -framework Foundation -framework CoreFoundation
|
||||
#include <stdlib.h>
|
||||
#include "secure_enclave.h"
|
||||
*/
|
||||
import "C"
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
const (
|
||||
// p256UncompressedKeySize is the size of an uncompressed P-256 public key.
|
||||
p256UncompressedKeySize = 65
|
||||
|
||||
// errorBufferSize is the size of the C error message buffer.
|
||||
errorBufferSize = 512
|
||||
|
||||
// hashBufferSize is the size of the hash output buffer.
|
||||
hashBufferSize = 128
|
||||
|
||||
// maxCiphertextSize is the max buffer for ECIES ciphertext.
|
||||
// ECIES overhead for P-256: 65 (ephemeral pub) + 16 (GCM tag) + 16 (IV) + plaintext.
|
||||
maxCiphertextSize = 8192
|
||||
|
||||
// maxPlaintextSize is the max buffer for decrypted plaintext.
|
||||
maxPlaintextSize = 8192
|
||||
)
|
||||
|
||||
// CreateKey creates a new P-256 non-exportable key in the Secure Enclave via sc_auth.
|
||||
// Returns the uncompressed public key bytes (65 bytes) and the identity hash (for deletion).
|
||||
func CreateKey(label string) (publicKey []byte, hash string, err error) {
|
||||
pubKeyBuf := make([]C.uint8_t, p256UncompressedKeySize)
|
||||
pubKeyLen := C.int(p256UncompressedKeySize)
|
||||
var hashBuf [hashBufferSize]C.char
|
||||
var errBuf [errorBufferSize]C.char
|
||||
|
||||
cLabel := C.CString(label)
|
||||
defer C.free(unsafe.Pointer(cLabel)) //nolint:nlreturn // CGo free pattern
|
||||
|
||||
result := C.se_create_key(cLabel,
|
||||
&pubKeyBuf[0], &pubKeyLen,
|
||||
&hashBuf[0], C.int(hashBufferSize),
|
||||
&errBuf[0], C.int(errorBufferSize))
|
||||
|
||||
if result != 0 {
|
||||
return nil, "", fmt.Errorf("secure enclave: %s", C.GoString(&errBuf[0]))
|
||||
}
|
||||
|
||||
pk := C.GoBytes(unsafe.Pointer(&pubKeyBuf[0]), pubKeyLen) //nolint:nlreturn // CGo result extraction
|
||||
h := C.GoString(&hashBuf[0])
|
||||
|
||||
return pk, h, nil
|
||||
}
|
||||
|
||||
// Encrypt encrypts plaintext using the SE-backed public key via ECIES
|
||||
// (eciesEncryptionStandardVariableIVX963SHA256AESGCM).
|
||||
// Encryption uses only the public key; no SE interaction required.
|
||||
func Encrypt(label string, plaintext []byte) ([]byte, error) {
|
||||
ciphertextBuf := make([]C.uint8_t, maxCiphertextSize)
|
||||
ciphertextLen := C.int(maxCiphertextSize)
|
||||
var errBuf [errorBufferSize]C.char
|
||||
|
||||
cLabel := C.CString(label)
|
||||
defer C.free(unsafe.Pointer(cLabel)) //nolint:nlreturn // CGo free pattern
|
||||
|
||||
result := C.se_encrypt(cLabel,
|
||||
(*C.uint8_t)(unsafe.Pointer(&plaintext[0])), C.int(len(plaintext)),
|
||||
&ciphertextBuf[0], &ciphertextLen,
|
||||
&errBuf[0], C.int(errorBufferSize))
|
||||
|
||||
if result != 0 {
|
||||
return nil, fmt.Errorf("secure enclave: %s", C.GoString(&errBuf[0]))
|
||||
}
|
||||
|
||||
out := C.GoBytes(unsafe.Pointer(&ciphertextBuf[0]), ciphertextLen) //nolint:nlreturn // CGo result extraction
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Decrypt decrypts ECIES ciphertext using the SE-backed private key.
|
||||
// The ECDH portion of decryption is performed inside the Secure Enclave.
|
||||
func Decrypt(label string, ciphertext []byte) ([]byte, error) {
|
||||
plaintextBuf := make([]C.uint8_t, maxPlaintextSize)
|
||||
plaintextLen := C.int(maxPlaintextSize)
|
||||
var errBuf [errorBufferSize]C.char
|
||||
|
||||
cLabel := C.CString(label)
|
||||
defer C.free(unsafe.Pointer(cLabel)) //nolint:nlreturn // CGo free pattern
|
||||
|
||||
result := C.se_decrypt(cLabel,
|
||||
(*C.uint8_t)(unsafe.Pointer(&ciphertext[0])), C.int(len(ciphertext)),
|
||||
&plaintextBuf[0], &plaintextLen,
|
||||
&errBuf[0], C.int(errorBufferSize))
|
||||
|
||||
if result != 0 {
|
||||
return nil, fmt.Errorf("secure enclave: %s", C.GoString(&errBuf[0]))
|
||||
}
|
||||
|
||||
out := C.GoBytes(unsafe.Pointer(&plaintextBuf[0]), plaintextLen) //nolint:nlreturn // CGo result extraction
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// DeleteKey removes a CTK identity from the Secure Enclave via sc_auth.
|
||||
func DeleteKey(hash string) error {
|
||||
var errBuf [errorBufferSize]C.char
|
||||
|
||||
cHash := C.CString(hash)
|
||||
defer C.free(unsafe.Pointer(cHash)) //nolint:nlreturn // CGo free pattern
|
||||
|
||||
result := C.se_delete_key(cHash, &errBuf[0], C.int(errorBufferSize))
|
||||
|
||||
if result != 0 {
|
||||
return fmt.Errorf("secure enclave: %s", C.GoString(&errBuf[0]))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
//go:build !darwin
|
||||
|
||||
// Package macse provides Go bindings for macOS Secure Enclave operations.
|
||||
package macse
|
||||
|
||||
import "errors"
|
||||
|
||||
var errNotSupported = errors.New("secure enclave is only supported on macOS")
|
||||
|
||||
// CreateKey is not supported on non-darwin platforms.
|
||||
func CreateKey(_ string) ([]byte, string, error) {
|
||||
return nil, "", errNotSupported
|
||||
}
|
||||
|
||||
// Encrypt is not supported on non-darwin platforms.
|
||||
func Encrypt(_ string, _ []byte) ([]byte, error) {
|
||||
return nil, errNotSupported
|
||||
}
|
||||
|
||||
// Decrypt is not supported on non-darwin platforms.
|
||||
func Decrypt(_ string, _ []byte) ([]byte, error) {
|
||||
return nil, errNotSupported
|
||||
}
|
||||
|
||||
// DeleteKey is not supported on non-darwin platforms.
|
||||
func DeleteKey(_ string) error {
|
||||
return errNotSupported
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
//go:build darwin
|
||||
// +build darwin
|
||||
|
||||
package macse
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
)
|
||||
|
||||
const testKeyLabel = "berlin.sneak.app.secret.test.se-key"
|
||||
|
||||
// testKeyHash stores the hash of the created test key for cleanup.
|
||||
var testKeyHash string //nolint:gochecknoglobals
|
||||
|
||||
// skipIfNoSecureEnclave skips the test if SE access is unavailable.
|
||||
func skipIfNoSecureEnclave(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
probeLabel := "berlin.sneak.app.secret.test.se-probe"
|
||||
_, hash, err := CreateKey(probeLabel)
|
||||
if err != nil {
|
||||
t.Skipf("Secure Enclave unavailable (skipping): %v", err)
|
||||
}
|
||||
|
||||
if hash != "" {
|
||||
_ = DeleteKey(hash)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateAndDeleteKey(t *testing.T) {
|
||||
skipIfNoSecureEnclave(t)
|
||||
|
||||
if testKeyHash != "" {
|
||||
_ = DeleteKey(testKeyHash)
|
||||
}
|
||||
|
||||
pubKey, hash, err := CreateKey(testKeyLabel)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateKey failed: %v", err)
|
||||
}
|
||||
|
||||
testKeyHash = hash
|
||||
t.Logf("Created key with hash: %s", hash)
|
||||
|
||||
// Verify valid uncompressed P-256 public key
|
||||
if len(pubKey) != p256UncompressedKeySize {
|
||||
t.Fatalf("expected public key length %d, got %d", p256UncompressedKeySize, len(pubKey))
|
||||
}
|
||||
|
||||
if pubKey[0] != 0x04 {
|
||||
t.Fatalf("expected uncompressed point prefix 0x04, got 0x%02x", pubKey[0])
|
||||
}
|
||||
|
||||
if hash == "" {
|
||||
t.Fatal("expected non-empty hash")
|
||||
}
|
||||
|
||||
// Delete the key
|
||||
if err := DeleteKey(hash); err != nil {
|
||||
t.Fatalf("DeleteKey failed: %v", err)
|
||||
}
|
||||
|
||||
testKeyHash = ""
|
||||
t.Log("Key created, verified, and deleted successfully")
|
||||
}
|
||||
|
||||
func TestEncryptDecryptRoundTrip(t *testing.T) {
|
||||
skipIfNoSecureEnclave(t)
|
||||
|
||||
_, hash, err := CreateKey(testKeyLabel)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateKey failed: %v", err)
|
||||
}
|
||||
|
||||
testKeyHash = hash
|
||||
|
||||
defer func() {
|
||||
if testKeyHash != "" {
|
||||
_ = DeleteKey(testKeyHash)
|
||||
testKeyHash = ""
|
||||
}
|
||||
}()
|
||||
|
||||
// Test data simulating an age private key
|
||||
plaintext := []byte("AGE-SECRET-KEY-1QQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQ")
|
||||
|
||||
// Encrypt
|
||||
ciphertext, err := Encrypt(testKeyLabel, plaintext)
|
||||
if err != nil {
|
||||
t.Fatalf("Encrypt failed: %v", err)
|
||||
}
|
||||
|
||||
t.Logf("Plaintext: %d bytes, Ciphertext: %d bytes", len(plaintext), len(ciphertext))
|
||||
|
||||
if bytes.Equal(ciphertext, plaintext) {
|
||||
t.Fatal("ciphertext should differ from plaintext")
|
||||
}
|
||||
|
||||
// Decrypt
|
||||
decrypted, err := Decrypt(testKeyLabel, ciphertext)
|
||||
if err != nil {
|
||||
t.Fatalf("Decrypt failed: %v", err)
|
||||
}
|
||||
|
||||
if !bytes.Equal(decrypted, plaintext) {
|
||||
t.Fatalf("decrypted data does not match original plaintext")
|
||||
}
|
||||
|
||||
t.Log("ECIES encrypt/decrypt round-trip successful")
|
||||
}
|
||||
|
||||
func TestEncryptProducesDifferentCiphertexts(t *testing.T) {
|
||||
skipIfNoSecureEnclave(t)
|
||||
|
||||
_, hash, err := CreateKey(testKeyLabel)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateKey failed: %v", err)
|
||||
}
|
||||
|
||||
testKeyHash = hash
|
||||
|
||||
defer func() {
|
||||
if testKeyHash != "" {
|
||||
_ = DeleteKey(testKeyHash)
|
||||
testKeyHash = ""
|
||||
}
|
||||
}()
|
||||
|
||||
plaintext := []byte("test-secret-data")
|
||||
|
||||
ct1, err := Encrypt(testKeyLabel, plaintext)
|
||||
if err != nil {
|
||||
t.Fatalf("first Encrypt failed: %v", err)
|
||||
}
|
||||
|
||||
ct2, err := Encrypt(testKeyLabel, plaintext)
|
||||
if err != nil {
|
||||
t.Fatalf("second Encrypt failed: %v", err)
|
||||
}
|
||||
|
||||
// ECIES uses a random ephemeral key each time, so ciphertexts should differ
|
||||
if bytes.Equal(ct1, ct2) {
|
||||
t.Fatal("two encryptions of same plaintext should produce different ciphertexts")
|
||||
}
|
||||
|
||||
// Both should decrypt to the same plaintext
|
||||
dec1, err := Decrypt(testKeyLabel, ct1)
|
||||
if err != nil {
|
||||
t.Fatalf("first Decrypt failed: %v", err)
|
||||
}
|
||||
|
||||
dec2, err := Decrypt(testKeyLabel, ct2)
|
||||
if err != nil {
|
||||
t.Fatalf("second Decrypt failed: %v", err)
|
||||
}
|
||||
|
||||
if !bytes.Equal(dec1, plaintext) || !bytes.Equal(dec2, plaintext) {
|
||||
t.Fatal("both ciphertexts should decrypt to original plaintext")
|
||||
}
|
||||
|
||||
t.Log("ECIES correctly produces different ciphertexts that decrypt to same plaintext")
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
//go:build darwin
|
||||
|
||||
#ifndef SECURE_ENCLAVE_H
|
||||
#define SECURE_ENCLAVE_H
|
||||
|
||||
#include <stdint.h>
|
||||
|
||||
// se_create_key creates a new P-256 key in the Secure Enclave via sc_auth.
|
||||
// label: unique identifier for the CTK identity (UTF-8 C string)
|
||||
// pub_key_out: output buffer for the uncompressed public key (65 bytes for P-256)
|
||||
// pub_key_len: on input, size of pub_key_out; on output, actual size written
|
||||
// hash_out: output buffer for the identity hash (for deletion)
|
||||
// hash_out_len: size of hash_out buffer
|
||||
// error_out: output buffer for error message
|
||||
// error_out_len: size of error_out buffer
|
||||
// Returns 0 on success, -1 on failure.
|
||||
int se_create_key(const char *label,
|
||||
uint8_t *pub_key_out, int *pub_key_len,
|
||||
char *hash_out, int hash_out_len,
|
||||
char *error_out, int error_out_len);
|
||||
|
||||
// se_encrypt encrypts data using the SE-backed public key (ECIES).
|
||||
// label: label of the CTK identity whose public key to use
|
||||
// plaintext: data to encrypt
|
||||
// plaintext_len: length of plaintext
|
||||
// ciphertext_out: output buffer for the ECIES ciphertext
|
||||
// ciphertext_len: on input, size of buffer; on output, actual size written
|
||||
// error_out: output buffer for error message
|
||||
// error_out_len: size of error_out buffer
|
||||
// Returns 0 on success, -1 on failure.
|
||||
int se_encrypt(const char *label,
|
||||
const uint8_t *plaintext, int plaintext_len,
|
||||
uint8_t *ciphertext_out, int *ciphertext_len,
|
||||
char *error_out, int error_out_len);
|
||||
|
||||
// se_decrypt decrypts ECIES ciphertext using the SE-backed private key.
|
||||
// The ECDH portion of decryption is performed inside the Secure Enclave.
|
||||
// label: label of the CTK identity whose private key to use
|
||||
// ciphertext: ECIES ciphertext produced by se_encrypt
|
||||
// ciphertext_len: length of ciphertext
|
||||
// plaintext_out: output buffer for decrypted data
|
||||
// plaintext_len: on input, size of buffer; on output, actual size written
|
||||
// error_out: output buffer for error message
|
||||
// error_out_len: size of error_out buffer
|
||||
// Returns 0 on success, -1 on failure.
|
||||
int se_decrypt(const char *label,
|
||||
const uint8_t *ciphertext, int ciphertext_len,
|
||||
uint8_t *plaintext_out, int *plaintext_len,
|
||||
char *error_out, int error_out_len);
|
||||
|
||||
// se_delete_key removes a CTK identity from the Secure Enclave via sc_auth.
|
||||
// hash: the identity hash returned by se_create_key
|
||||
// error_out: output buffer for error message
|
||||
// error_out_len: size of error_out buffer
|
||||
// Returns 0 on success, -1 on failure.
|
||||
int se_delete_key(const char *hash,
|
||||
char *error_out, int error_out_len);
|
||||
|
||||
#endif // SECURE_ENCLAVE_H
|
||||
@@ -0,0 +1,302 @@
|
||||
//go:build darwin
|
||||
|
||||
#import <Foundation/Foundation.h>
|
||||
#import <Security/Security.h>
|
||||
#include "secure_enclave.h"
|
||||
#include <string.h>
|
||||
|
||||
// snprintf_error writes an error message string to the output buffer.
|
||||
static void snprintf_error(char *error_out, int error_out_len, NSString *msg) {
|
||||
if (error_out && error_out_len > 0) {
|
||||
snprintf(error_out, error_out_len, "%s", msg.UTF8String);
|
||||
}
|
||||
}
|
||||
|
||||
// lookup_ctk_identity finds a CTK identity by label and returns the private key.
|
||||
static SecKeyRef lookup_ctk_private_key(const char *label, char *error_out, int error_out_len) {
|
||||
NSDictionary *query = @{
|
||||
(id)kSecClass: (id)kSecClassIdentity,
|
||||
(id)kSecAttrLabel: [NSString stringWithUTF8String:label],
|
||||
(id)kSecMatchLimit: (id)kSecMatchLimitOne,
|
||||
(id)kSecReturnRef: @YES,
|
||||
};
|
||||
|
||||
SecIdentityRef identity = NULL;
|
||||
OSStatus status = SecItemCopyMatching((__bridge CFDictionaryRef)query, (CFTypeRef *)&identity);
|
||||
|
||||
if (status != errSecSuccess || !identity) {
|
||||
NSString *msg = [NSString stringWithFormat:@"CTK identity '%s' not found: OSStatus %d",
|
||||
label, (int)status];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
SecKeyRef privateKey = NULL;
|
||||
status = SecIdentityCopyPrivateKey(identity, &privateKey);
|
||||
CFRelease(identity);
|
||||
|
||||
if (status != errSecSuccess || !privateKey) {
|
||||
NSString *msg = [NSString stringWithFormat:
|
||||
@"failed to get private key from CTK identity '%s': OSStatus %d",
|
||||
label, (int)status];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
return privateKey;
|
||||
}
|
||||
|
||||
int se_create_key(const char *label,
|
||||
uint8_t *pub_key_out, int *pub_key_len,
|
||||
char *hash_out, int hash_out_len,
|
||||
char *error_out, int error_out_len) {
|
||||
@autoreleasepool {
|
||||
NSString *labelStr = [NSString stringWithUTF8String:label];
|
||||
|
||||
// Shell out to sc_auth (which has SE entitlements) to create the key
|
||||
NSTask *task = [[NSTask alloc] init];
|
||||
task.executableURL = [NSURL fileURLWithPath:@"/usr/sbin/sc_auth"];
|
||||
task.arguments = @[
|
||||
@"create-ctk-identity",
|
||||
@"-k", @"p-256-ne",
|
||||
@"-t", @"none",
|
||||
@"-l", labelStr,
|
||||
];
|
||||
|
||||
NSPipe *stderrPipe = [NSPipe pipe];
|
||||
task.standardOutput = [NSPipe pipe];
|
||||
task.standardError = stderrPipe;
|
||||
|
||||
NSError *nsError = nil;
|
||||
if (![task launchAndReturnError:&nsError]) {
|
||||
NSString *msg = [NSString stringWithFormat:@"failed to launch sc_auth: %@",
|
||||
nsError.localizedDescription];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return -1;
|
||||
}
|
||||
|
||||
[task waitUntilExit];
|
||||
|
||||
if (task.terminationStatus != 0) {
|
||||
NSData *stderrData = [stderrPipe.fileHandleForReading readDataToEndOfFile];
|
||||
NSString *stderrStr = [[NSString alloc] initWithData:stderrData
|
||||
encoding:NSUTF8StringEncoding];
|
||||
NSString *msg = [NSString stringWithFormat:@"sc_auth failed: %@",
|
||||
stderrStr ?: @"unknown error"];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return -1;
|
||||
}
|
||||
|
||||
// Retrieve the public key from the created identity
|
||||
SecKeyRef privateKey = lookup_ctk_private_key(label, error_out, error_out_len);
|
||||
if (!privateKey) {
|
||||
return -1;
|
||||
}
|
||||
|
||||
SecKeyRef publicKey = SecKeyCopyPublicKey(privateKey);
|
||||
CFRelease(privateKey);
|
||||
|
||||
if (!publicKey) {
|
||||
snprintf_error(error_out, error_out_len, @"failed to get public key");
|
||||
return -1;
|
||||
}
|
||||
|
||||
CFErrorRef cfError = NULL;
|
||||
CFDataRef pubKeyData = SecKeyCopyExternalRepresentation(publicKey, &cfError);
|
||||
CFRelease(publicKey);
|
||||
|
||||
if (!pubKeyData) {
|
||||
NSError *err = (__bridge_transfer NSError *)cfError;
|
||||
NSString *msg = [NSString stringWithFormat:@"failed to export public key: %@",
|
||||
err.localizedDescription];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return -1;
|
||||
}
|
||||
|
||||
const UInt8 *bytes = CFDataGetBytePtr(pubKeyData);
|
||||
CFIndex length = CFDataGetLength(pubKeyData);
|
||||
|
||||
if (length > *pub_key_len) {
|
||||
CFRelease(pubKeyData);
|
||||
snprintf_error(error_out, error_out_len, @"public key buffer too small");
|
||||
return -1;
|
||||
}
|
||||
|
||||
memcpy(pub_key_out, bytes, length);
|
||||
*pub_key_len = (int)length;
|
||||
CFRelease(pubKeyData);
|
||||
|
||||
// Get the identity hash by parsing sc_auth list output
|
||||
hash_out[0] = '\0';
|
||||
NSTask *listTask = [[NSTask alloc] init];
|
||||
listTask.executableURL = [NSURL fileURLWithPath:@"/usr/sbin/sc_auth"];
|
||||
listTask.arguments = @[@"list-ctk-identities"];
|
||||
|
||||
NSPipe *listPipe = [NSPipe pipe];
|
||||
listTask.standardOutput = listPipe;
|
||||
listTask.standardError = [NSPipe pipe];
|
||||
|
||||
if ([listTask launchAndReturnError:&nsError]) {
|
||||
[listTask waitUntilExit];
|
||||
NSData *listData = [listPipe.fileHandleForReading readDataToEndOfFile];
|
||||
NSString *listStr = [[NSString alloc] initWithData:listData
|
||||
encoding:NSUTF8StringEncoding];
|
||||
|
||||
for (NSString *line in [listStr componentsSeparatedByString:@"\n"]) {
|
||||
if ([line containsString:labelStr]) {
|
||||
NSMutableArray *tokens = [NSMutableArray array];
|
||||
for (NSString *part in [line componentsSeparatedByCharactersInSet:
|
||||
[NSCharacterSet whitespaceCharacterSet]]) {
|
||||
if (part.length > 0) {
|
||||
[tokens addObject:part];
|
||||
}
|
||||
}
|
||||
if (tokens.count > 1) {
|
||||
snprintf(hash_out, hash_out_len, "%s", [tokens[1] UTF8String]);
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
int se_encrypt(const char *label,
|
||||
const uint8_t *plaintext, int plaintext_len,
|
||||
uint8_t *ciphertext_out, int *ciphertext_len,
|
||||
char *error_out, int error_out_len) {
|
||||
@autoreleasepool {
|
||||
SecKeyRef privateKey = lookup_ctk_private_key(label, error_out, error_out_len);
|
||||
if (!privateKey) {
|
||||
return -1;
|
||||
}
|
||||
|
||||
SecKeyRef publicKey = SecKeyCopyPublicKey(privateKey);
|
||||
CFRelease(privateKey);
|
||||
|
||||
if (!publicKey) {
|
||||
snprintf_error(error_out, error_out_len, @"failed to get public key for encryption");
|
||||
return -1;
|
||||
}
|
||||
|
||||
NSData *plaintextData = [NSData dataWithBytes:plaintext length:plaintext_len];
|
||||
|
||||
CFErrorRef cfError = NULL;
|
||||
CFDataRef encrypted = SecKeyCreateEncryptedData(
|
||||
publicKey,
|
||||
kSecKeyAlgorithmECIESEncryptionStandardVariableIVX963SHA256AESGCM,
|
||||
(__bridge CFDataRef)plaintextData,
|
||||
&cfError
|
||||
);
|
||||
CFRelease(publicKey);
|
||||
|
||||
if (!encrypted) {
|
||||
NSError *nsError = (__bridge_transfer NSError *)cfError;
|
||||
NSString *msg = [NSString stringWithFormat:@"ECIES encryption failed: %@",
|
||||
nsError.localizedDescription];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return -1;
|
||||
}
|
||||
|
||||
const UInt8 *encBytes = CFDataGetBytePtr(encrypted);
|
||||
CFIndex encLength = CFDataGetLength(encrypted);
|
||||
|
||||
if (encLength > *ciphertext_len) {
|
||||
CFRelease(encrypted);
|
||||
snprintf_error(error_out, error_out_len, @"ciphertext buffer too small");
|
||||
return -1;
|
||||
}
|
||||
|
||||
memcpy(ciphertext_out, encBytes, encLength);
|
||||
*ciphertext_len = (int)encLength;
|
||||
CFRelease(encrypted);
|
||||
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
int se_decrypt(const char *label,
|
||||
const uint8_t *ciphertext, int ciphertext_len,
|
||||
uint8_t *plaintext_out, int *plaintext_len,
|
||||
char *error_out, int error_out_len) {
|
||||
@autoreleasepool {
|
||||
SecKeyRef privateKey = lookup_ctk_private_key(label, error_out, error_out_len);
|
||||
if (!privateKey) {
|
||||
return -1;
|
||||
}
|
||||
|
||||
NSData *ciphertextData = [NSData dataWithBytes:ciphertext length:ciphertext_len];
|
||||
|
||||
CFErrorRef cfError = NULL;
|
||||
CFDataRef decrypted = SecKeyCreateDecryptedData(
|
||||
privateKey,
|
||||
kSecKeyAlgorithmECIESEncryptionStandardVariableIVX963SHA256AESGCM,
|
||||
(__bridge CFDataRef)ciphertextData,
|
||||
&cfError
|
||||
);
|
||||
CFRelease(privateKey);
|
||||
|
||||
if (!decrypted) {
|
||||
NSError *nsError = (__bridge_transfer NSError *)cfError;
|
||||
NSString *msg = [NSString stringWithFormat:@"ECIES decryption failed: %@",
|
||||
nsError.localizedDescription];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return -1;
|
||||
}
|
||||
|
||||
const UInt8 *decBytes = CFDataGetBytePtr(decrypted);
|
||||
CFIndex decLength = CFDataGetLength(decrypted);
|
||||
|
||||
if (decLength > *plaintext_len) {
|
||||
CFRelease(decrypted);
|
||||
snprintf_error(error_out, error_out_len, @"plaintext buffer too small");
|
||||
return -1;
|
||||
}
|
||||
|
||||
memcpy(plaintext_out, decBytes, decLength);
|
||||
*plaintext_len = (int)decLength;
|
||||
CFRelease(decrypted);
|
||||
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
int se_delete_key(const char *hash,
|
||||
char *error_out, int error_out_len) {
|
||||
@autoreleasepool {
|
||||
NSTask *task = [[NSTask alloc] init];
|
||||
task.executableURL = [NSURL fileURLWithPath:@"/usr/sbin/sc_auth"];
|
||||
task.arguments = @[
|
||||
@"delete-ctk-identity",
|
||||
@"-h", [NSString stringWithUTF8String:hash],
|
||||
];
|
||||
|
||||
NSPipe *stderrPipe = [NSPipe pipe];
|
||||
task.standardOutput = [NSPipe pipe];
|
||||
task.standardError = stderrPipe;
|
||||
|
||||
NSError *nsError = nil;
|
||||
if (![task launchAndReturnError:&nsError]) {
|
||||
NSString *msg = [NSString stringWithFormat:@"failed to launch sc_auth: %@",
|
||||
nsError.localizedDescription];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return -1;
|
||||
}
|
||||
|
||||
[task waitUntilExit];
|
||||
|
||||
if (task.terminationStatus != 0) {
|
||||
NSData *stderrData = [stderrPipe.fileHandleForReading readDataToEndOfFile];
|
||||
NSString *stderrStr = [[NSString alloc] initWithData:stderrData
|
||||
encoding:NSUTF8StringEncoding];
|
||||
NSString *msg = [NSString stringWithFormat:@"sc_auth delete failed: %@",
|
||||
stderrStr ?: @"unknown error"];
|
||||
snprintf_error(error_out, error_out_len, msg);
|
||||
return -1;
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package secret
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// WriteFileAtomic replaces the file at path with data so that a reader, or
|
||||
// a crash at any moment, finds either the old content or the new, never a
|
||||
// partial file. The data goes into a temporary file that afero.TempFile
|
||||
// creates with mode 0600 in the same directory (a rename is only atomic
|
||||
// within one filesystem), is synced to disk, and is renamed over path. The
|
||||
// temporary file is removed if any step fails.
|
||||
func WriteFileAtomic(fs afero.Fs, path string, data []byte) error {
|
||||
tmp, err := afero.TempFile(fs, filepath.Dir(path),
|
||||
"."+filepath.Base(path)+".tmp-*")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create temporary file for %s: %w", path, err)
|
||||
}
|
||||
|
||||
_, err = tmp.Write(data)
|
||||
if err == nil {
|
||||
err = tmp.Sync()
|
||||
}
|
||||
|
||||
closeErr := tmp.Close()
|
||||
if err == nil {
|
||||
err = closeErr
|
||||
}
|
||||
|
||||
if err == nil {
|
||||
err = fs.Rename(tmp.Name(), path)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
_ = fs.Remove(tmp.Name())
|
||||
|
||||
return fmt.Errorf("failed to write %s: %w", path, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// TempDirFor creates an empty temporary directory in which to build the
|
||||
// directory target before renaming it into place, or into which to move
|
||||
// target before deleting it. It is made in target's grandparent: on the
|
||||
// same filesystem, so the rename is atomic, and outside target's parent,
|
||||
// the directory that is listed to find vaults, secrets, versions and
|
||||
// unlockers, so one left behind by a crash is never taken for one of them.
|
||||
// Its name leaves out target's, which may already be as long as a file name
|
||||
// can be.
|
||||
func TempDirFor(fs afero.Fs, target string) (string, error) {
|
||||
dir, err := afero.TempDir(fs, filepath.Dir(filepath.Dir(target)), ".tmp-")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(
|
||||
"failed to create temporary directory for %s: %w", target, err)
|
||||
}
|
||||
|
||||
return dir, nil
|
||||
}
|
||||
|
||||
// WriteDir calls write to write the files of the directory dir. When dir does
|
||||
// not exist yet, write writes them into a temporary directory from TempDirFor,
|
||||
// which is then renamed to dir, so that neither a failure nor a crash leaves
|
||||
// dir half-written; on a failure the temporary directory is removed, and a
|
||||
// failure to remove it is returned along with the first. A directory cannot be
|
||||
// renamed over one that has files in it, so when dir already exists, write
|
||||
// writes into it in place; dir is then never removed.
|
||||
func WriteDir(fs afero.Fs, dir string, write func(dir string) error) error {
|
||||
exists, err := afero.Exists(fs, dir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check for %s: %w", dir, err)
|
||||
}
|
||||
|
||||
if exists {
|
||||
return write(dir)
|
||||
}
|
||||
|
||||
// Create the directory the finished one is renamed into
|
||||
err = fs.MkdirAll(filepath.Dir(dir), DirPerms)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create %s: %w", filepath.Dir(dir), err)
|
||||
}
|
||||
|
||||
tmp, err := TempDirFor(fs, dir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = write(tmp)
|
||||
if err == nil {
|
||||
err = fs.Rename(tmp, dir)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
removeErr := fs.RemoveAll(tmp)
|
||||
if removeErr != nil {
|
||||
err = errors.Join(err,
|
||||
fmt.Errorf("failed to remove %s: %w", tmp, removeErr))
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveDirAtomic deletes the directory dir so that it disappears in one
|
||||
// rename: dir is moved into a new directory from TempDirFor, which is then
|
||||
// deleted. A crash part-way leaves only that temporary directory behind.
|
||||
func RemoveDirAtomic(fs afero.Fs, dir string) error {
|
||||
tmp, err := TempDirFor(fs, dir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = fs.Rename(dir, filepath.Join(tmp, filepath.Base(dir)))
|
||||
if err != nil {
|
||||
_ = fs.Remove(tmp)
|
||||
|
||||
return fmt.Errorf("failed to remove %s: %w", dir, err)
|
||||
}
|
||||
|
||||
err = fs.RemoveAll(tmp)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to remove %s: %w", dir, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,764 @@
|
||||
package secret_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
var errInjected = errors.New("injected failure")
|
||||
|
||||
// The kinds of change hookFs passes to before.
|
||||
const (
|
||||
opCreate = "create"
|
||||
opOpen = "open"
|
||||
opSync = "sync"
|
||||
opMkdir = "mkdir"
|
||||
opRemove = "remove"
|
||||
opRename = "rename"
|
||||
)
|
||||
|
||||
// currentFile is the file in a secret's directory that names its current
|
||||
// version.
|
||||
const currentFile = "current"
|
||||
|
||||
// unlockerMetadataFile is the file a new unlocker writes last.
|
||||
const unlockerMetadataFile = "unlocker-metadata.json"
|
||||
|
||||
// privKeyFile is the file that holds the encrypted private key of a version
|
||||
// or of a passphrase unlocker.
|
||||
const privKeyFile = "priv.age"
|
||||
|
||||
// unlockerPassphrase protects the passphrase unlockers the tests create.
|
||||
//
|
||||
//nolint:gosec // G101: test data, not a real credential
|
||||
const unlockerPassphrase = "unlocker passphrase"
|
||||
|
||||
// hookFs passes every call through to Fs, but first calls before for each
|
||||
// call that changes the filesystem, and for each Sync of a file opened
|
||||
// through it, with the path it changes (the new path, for Rename). A test
|
||||
// uses before to inspect the tree at every point where a crash could stop
|
||||
// the code under test, or returns an error from it to make that call fail.
|
||||
// If opened is set, OpenFile also tells it the mode it opens each file with.
|
||||
type hookFs struct {
|
||||
afero.Fs
|
||||
|
||||
before func(op, path string) error
|
||||
opened func(path string, perm os.FileMode)
|
||||
}
|
||||
|
||||
// hookFile is a file opened through hookFs.
|
||||
type hookFile struct {
|
||||
afero.File
|
||||
|
||||
before func(op, path string) error
|
||||
}
|
||||
|
||||
func (f hookFile) Sync() error {
|
||||
err := f.before(opSync, f.Name())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return f.File.Sync()
|
||||
}
|
||||
|
||||
//nolint:ireturn // implements afero.Fs
|
||||
func (h hookFs) Create(name string) (afero.File, error) {
|
||||
err := h.before(opCreate, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
file, err := h.Fs.Create(name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return hookFile{File: file, before: h.before}, nil
|
||||
}
|
||||
|
||||
//nolint:ireturn // implements afero.Fs
|
||||
func (h hookFs) OpenFile(
|
||||
name string, flag int, perm os.FileMode,
|
||||
) (afero.File, error) {
|
||||
err := h.before(opOpen, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if h.opened != nil {
|
||||
h.opened(name, perm)
|
||||
}
|
||||
|
||||
file, err := h.Fs.OpenFile(name, flag, perm)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return hookFile{File: file, before: h.before}, nil
|
||||
}
|
||||
|
||||
func (h hookFs) Mkdir(name string, perm os.FileMode) error {
|
||||
err := h.before(opMkdir, name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return h.Fs.Mkdir(name, perm)
|
||||
}
|
||||
|
||||
func (h hookFs) MkdirAll(path string, perm os.FileMode) error {
|
||||
err := h.before(opMkdir, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return h.Fs.MkdirAll(path, perm)
|
||||
}
|
||||
|
||||
func (h hookFs) Remove(name string) error {
|
||||
err := h.before(opRemove, name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return h.Fs.Remove(name)
|
||||
}
|
||||
|
||||
func (h hookFs) RemoveAll(path string) error {
|
||||
err := h.before(opRemove, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return h.Fs.RemoveAll(path)
|
||||
}
|
||||
|
||||
func (h hookFs) Rename(oldname, newname string) error {
|
||||
err := h.before(opRename, newname)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return h.Fs.Rename(oldname, newname)
|
||||
}
|
||||
|
||||
// testFilesystem is a filesystem to run a test on, with a directory in it
|
||||
// to work in.
|
||||
type testFilesystem struct {
|
||||
name string
|
||||
open func(t *testing.T) (afero.Fs, string)
|
||||
}
|
||||
|
||||
// testFilesystems are the in-memory filesystem that most tests use and the
|
||||
// real one: every rename-based guarantee is checked on both.
|
||||
//
|
||||
//nolint:gochecknoglobals // read-only table shared by the tests below
|
||||
var testFilesystems = []testFilesystem{
|
||||
{"memory", func(*testing.T) (afero.Fs, string) {
|
||||
return afero.NewMemMapFs(), "/test"
|
||||
}},
|
||||
{"real", func(t *testing.T) (afero.Fs, string) {
|
||||
t.Helper()
|
||||
|
||||
return afero.NewOsFs(), t.TempDir()
|
||||
}},
|
||||
}
|
||||
|
||||
// dirNames lists the names in dir.
|
||||
func dirNames(t *testing.T, fs afero.Fs, dir string) []string {
|
||||
t.Helper()
|
||||
|
||||
entries, err := afero.ReadDir(fs, dir)
|
||||
require.NoError(t, err)
|
||||
|
||||
names := make([]string, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
names = append(names, entry.Name())
|
||||
}
|
||||
|
||||
return names
|
||||
}
|
||||
|
||||
// writeLongTermKey gives the test vault under stateDir a new long-term key
|
||||
// and returns it.
|
||||
func writeLongTermKey(
|
||||
t *testing.T, fs afero.Fs, stateDir string,
|
||||
) *age.X25519Identity {
|
||||
t.Helper()
|
||||
|
||||
vault := &MockVersionVault{Name: testVaultName, fs: fs, stateDir: stateDir}
|
||||
|
||||
vaultDir, err := vault.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, fs.MkdirAll(vaultDir, 0o700))
|
||||
|
||||
ltIdentity, err := age.GenerateX25519Identity()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, afero.WriteFile(fs, filepath.Join(vaultDir, "pub.age"),
|
||||
[]byte(ltIdentity.Recipient().String()), 0o600))
|
||||
|
||||
return ltIdentity
|
||||
}
|
||||
|
||||
// newVaultWithSecret creates the vault name under stateDir from the test
|
||||
// mnemonic, with a secret "shared" in it that holds value.
|
||||
func newVaultWithSecret(
|
||||
t *testing.T, fs afero.Fs, stateDir, name, value string,
|
||||
) *vault.Vault {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.CreateVault(fs, stateDir, name)
|
||||
require.NoError(t, err)
|
||||
|
||||
buffer := memguard.NewBufferFromBytes([]byte(value))
|
||||
defer buffer.Destroy()
|
||||
|
||||
require.NoError(t, vlt.AddSecret("shared", buffer, false))
|
||||
|
||||
return vlt
|
||||
}
|
||||
|
||||
func TestWriteFileAtomicReplacesFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs, dir := tfs.open(t)
|
||||
path := filepath.Join(dir, currentFile)
|
||||
|
||||
require.NoError(t, secret.WriteFileAtomic(fs, path, []byte("old")))
|
||||
require.NoError(t, secret.WriteFileAtomic(fs, path, []byte("new")))
|
||||
|
||||
data, err := afero.ReadFile(fs, path)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "new", string(data))
|
||||
|
||||
info, err := fs.Stat(path)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, secret.FilePerms, info.Mode().Perm())
|
||||
|
||||
// No temporary file is left next to it
|
||||
assert.Equal(t, []string{currentFile}, dirNames(t, fs, dir))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteFileAtomicFailureKeepsOldFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base, dir := tfs.open(t)
|
||||
path := filepath.Join(dir, currentFile)
|
||||
require.NoError(t, secret.WriteFileAtomic(base, path, []byte("old")))
|
||||
|
||||
fs := hookFs{Fs: base, before: func(op, _ string) error {
|
||||
if op == opRename {
|
||||
return errInjected
|
||||
}
|
||||
|
||||
return nil
|
||||
}}
|
||||
|
||||
err := secret.WriteFileAtomic(fs, path, []byte("new"))
|
||||
require.ErrorIs(t, err, errInjected)
|
||||
|
||||
data, err := afero.ReadFile(base, path)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "old", string(data))
|
||||
|
||||
// The temporary file is removed again
|
||||
assert.Equal(t, []string{currentFile}, dirNames(t, base, dir))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestRemoveDirAtomic checks that RemoveDirAtomic deletes nothing where the
|
||||
// directory stands, which a crash could stop half-way, and that it leaves
|
||||
// nothing behind.
|
||||
func TestRemoveDirAtomic(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base, dir := tfs.open(t)
|
||||
listed := filepath.Join(dir, "secrets.d")
|
||||
target := filepath.Join(listed, "doomed")
|
||||
|
||||
require.NoError(t, base.MkdirAll(filepath.Join(target, "versions"), 0o700))
|
||||
require.NoError(t, secret.WriteFileAtomic(base,
|
||||
filepath.Join(target, currentFile), []byte("20231216.001")))
|
||||
|
||||
fs := hookFs{Fs: base, before: func(op, path string) error {
|
||||
if op == opRemove && strings.HasPrefix(path, target) {
|
||||
t.Errorf("deleted %s where it stands", path)
|
||||
}
|
||||
|
||||
return nil
|
||||
}}
|
||||
|
||||
require.NoError(t, secret.RemoveDirAtomic(fs, target))
|
||||
|
||||
// Gone, and no temporary directory is left in the directory
|
||||
// that is listed or in the one above it
|
||||
assert.Empty(t, dirNames(t, base, listed))
|
||||
assert.Equal(t, []string{"secrets.d"}, dirNames(t, base, dir))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLongestNames adds a secret to a vault and removes the vault, both
|
||||
// named with 255 bytes, the most a file name may have, on the real
|
||||
// filesystem: the temporary directories they use must fit that limit too.
|
||||
func TestLongestNames(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
const longestName = 255
|
||||
|
||||
fs := afero.NewOsFs()
|
||||
name := strings.Repeat("a", longestName)
|
||||
|
||||
vlt, err := vault.CreateVault(fs, t.TempDir(), name)
|
||||
require.NoError(t, err)
|
||||
|
||||
value := memguard.NewBufferFromBytes([]byte("long"))
|
||||
defer value.Destroy()
|
||||
|
||||
require.NoError(t, vlt.AddSecret(name, value, false))
|
||||
|
||||
got, err := vlt.GetSecret(name)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer got.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("long"), got.Bytes())
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, secret.RemoveDirAtomic(fs, vaultDir))
|
||||
assert.NoDirExists(t, vaultDir)
|
||||
}
|
||||
|
||||
// TestForcedCopyKeepsDestinationUntilReplaced copies a secret over one in
|
||||
// another vault, as a forced move between vaults does, and makes the last
|
||||
// step that completes the copy fail. The secret it was to replace must
|
||||
// still be there unchanged: it may go only once its replacement is whole.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids t.Parallel
|
||||
func TestForcedCopyKeepsDestinationUntilReplaced(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
base, stateDir := tfs.open(t)
|
||||
src := newVaultWithSecret(t, base, stateDir, "source", "new")
|
||||
dest := newVaultWithSecret(t, base, stateDir, "dest", "old")
|
||||
|
||||
// The copy is complete once its current file is written
|
||||
fs := hookFs{Fs: base, before: func(op, path string) error {
|
||||
if op == opRename && filepath.Base(path) == currentFile {
|
||||
return errInjected
|
||||
}
|
||||
|
||||
return nil
|
||||
}}
|
||||
|
||||
err := vault.NewVault(fs, stateDir, "dest").
|
||||
CopySecretAllVersions(src, "shared", "shared", true)
|
||||
require.ErrorIs(t, err, errInjected)
|
||||
|
||||
value, err := dest.GetSecret("shared")
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("old"), value.Bytes())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestTempDirsStayOutOfListings adds a version, adds a secret, copies a
|
||||
// secret over another and removes one, and checks that none of them makes a
|
||||
// directory directly in secrets.d or in a versions directory. Those are
|
||||
// listed to find secrets and versions, so a temporary directory made there
|
||||
// would be listed while half-built, and one left by a crash would stay.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids t.Parallel
|
||||
func TestTempDirsStayOutOfListings(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
base, stateDir := tfs.open(t)
|
||||
newVaultWithSecret(t, base, stateDir, "default", "first")
|
||||
|
||||
fs := hookFs{Fs: base, before: func(op, path string) error {
|
||||
parent := filepath.Base(filepath.Dir(path))
|
||||
if op == opMkdir && (parent == "secrets.d" || parent == "versions") {
|
||||
t.Errorf("made %s where it is listed", path)
|
||||
}
|
||||
|
||||
return nil
|
||||
}}
|
||||
vlt := vault.NewVault(fs, stateDir, "default")
|
||||
|
||||
value := memguard.NewBufferFromBytes([]byte("second"))
|
||||
defer value.Destroy()
|
||||
|
||||
require.NoError(t, vlt.AddSecret("shared", value, true))
|
||||
require.NoError(t, vlt.AddSecret("other", value, false))
|
||||
require.NoError(t, vlt.CopySecretAllVersions(vlt, "shared", "other", true))
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, secret.RemoveDirAtomic(fs,
|
||||
filepath.Join(vaultDir, "secrets.d", "shared")))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestVersionSaveIsWholeOrAbsent checks, before every change Save makes and
|
||||
// once after it returns, that the version directory either does not exist
|
||||
// or holds all of its files: a crash at any point leaves no version that
|
||||
// cannot be decrypted.
|
||||
func TestVersionSaveIsWholeOrAbsent(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base, stateDir := tfs.open(t)
|
||||
ltIdentity := writeLongTermKey(t, base, stateDir)
|
||||
|
||||
var versionDir string
|
||||
|
||||
checkVersionDir := func(string, string) error {
|
||||
exists, err := afero.DirExists(base, versionDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
if exists {
|
||||
assert.ElementsMatch(t,
|
||||
[]string{"pub.age", "value.age", privKeyFile, "metadata.age"},
|
||||
dirNames(t, base, versionDir),
|
||||
"version directory visible before it was complete")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
fs := hookFs{Fs: base, before: checkVersionDir}
|
||||
vault := &MockVersionVault{Name: testVaultName, fs: fs, stateDir: stateDir}
|
||||
sv := secret.NewVersion(vault, "test/secret", "20231215.001")
|
||||
versionDir = sv.Directory
|
||||
|
||||
value := memguard.NewBufferFromBytes([]byte("whole or nothing"))
|
||||
defer value.Destroy()
|
||||
|
||||
require.NoError(t, sv.Save(value))
|
||||
require.NoError(t, checkVersionDir("", ""))
|
||||
|
||||
got, err := sv.GetValue(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer got.Destroy()
|
||||
|
||||
assert.Equal(t, "whole or nothing", got.String())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestVersionSaveFailureLeavesNothing makes the write of the encrypted
|
||||
// private key fail, after the value has been written, and checks that
|
||||
// neither the version nor its temporary directory is left behind.
|
||||
func TestVersionSaveFailureLeavesNothing(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base, stateDir := tfs.open(t)
|
||||
writeLongTermKey(t, base, stateDir)
|
||||
|
||||
fs := hookFs{Fs: base, before: func(op, path string) error {
|
||||
if op == opRename && filepath.Base(path) == privKeyFile {
|
||||
return errInjected
|
||||
}
|
||||
|
||||
return nil
|
||||
}}
|
||||
vault := &MockVersionVault{Name: testVaultName, fs: fs, stateDir: stateDir}
|
||||
sv := secret.NewVersion(vault, "test/secret", "20231215.001")
|
||||
|
||||
value := memguard.NewBufferFromBytes([]byte("never stored"))
|
||||
defer value.Destroy()
|
||||
|
||||
require.ErrorIs(t, sv.Save(value), errInjected)
|
||||
|
||||
// The secret directory holds only the empty versions directory
|
||||
versionsDir := filepath.Dir(sv.Directory)
|
||||
assert.Equal(t, []string{"versions"},
|
||||
dirNames(t, base, filepath.Dir(versionsDir)))
|
||||
assert.Empty(t, dirNames(t, base, versionsDir))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCurrentFilesNeverMissing selects the current version, vault and
|
||||
// unlocker again and checks, before each change this makes, that the file
|
||||
// naming the current one exists: a reader or a crash never finds it
|
||||
// missing.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids t.Parallel
|
||||
func TestCurrentFilesNeverMissing(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
base, stateDir := tfs.open(t)
|
||||
vlt := newVaultWithSecret(t, base, stateDir, testVaultName, "value")
|
||||
|
||||
passphrase := memguard.NewBufferFromBytes([]byte(unlockerPassphrase))
|
||||
defer passphrase.Destroy()
|
||||
|
||||
// Created as the current unlocker
|
||||
unlocker, err := vlt.CreatePassphraseUnlocker(passphrase)
|
||||
require.NoError(t, err)
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "shared")
|
||||
version, err := secret.GetCurrentVersion(base, secretDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, tc := range []struct {
|
||||
path string
|
||||
reselect func(fs afero.Fs) error
|
||||
}{
|
||||
{filepath.Join(secretDir, currentFile), func(fs afero.Fs) error {
|
||||
return secret.SetCurrentVersion(fs, secretDir, version)
|
||||
}},
|
||||
{filepath.Join(stateDir, "currentvault"), func(fs afero.Fs) error {
|
||||
return vault.SelectVault(fs, stateDir, testVaultName)
|
||||
}},
|
||||
{filepath.Join(vaultDir, "current-unlocker"), func(fs afero.Fs) error {
|
||||
return vault.NewVault(fs, stateDir, testVaultName).
|
||||
SelectUnlocker(unlocker.GetID())
|
||||
}},
|
||||
} {
|
||||
fs := hookFs{Fs: base, before: func(string, string) error {
|
||||
exists, err := afero.Exists(base, tc.path)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, exists, "%s is missing", filepath.Base(tc.path))
|
||||
|
||||
return nil
|
||||
}}
|
||||
|
||||
require.NoError(t, tc.reselect(fs))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteFileAtomicTempFile checks that WriteFileAtomic creates its
|
||||
// temporary file with mode 0600, rather than wider and narrowed later, so
|
||||
// that no other user can ever read it, and syncs it before renaming it into
|
||||
// place, so that a crash cannot leave the file named but its data lost.
|
||||
func TestWriteFileAtomicTempFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base, dir := tfs.open(t)
|
||||
|
||||
var modes []os.FileMode
|
||||
|
||||
synced := false
|
||||
fs := hookFs{
|
||||
Fs: base,
|
||||
before: func(op, _ string) error {
|
||||
switch op {
|
||||
case opSync:
|
||||
synced = true
|
||||
case opRename:
|
||||
assert.True(t, synced, "renamed before syncing")
|
||||
}
|
||||
|
||||
return nil
|
||||
},
|
||||
opened: func(_ string, perm os.FileMode) {
|
||||
modes = append(modes, perm)
|
||||
},
|
||||
}
|
||||
|
||||
require.NoError(t, secret.WriteFileAtomic(fs,
|
||||
filepath.Join(dir, currentFile), []byte("new")))
|
||||
assert.Equal(t, []os.FileMode{secret.FilePerms}, modes)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestPassphraseUnlockerGetsKeyFirst creates a passphrase unlocker in a
|
||||
// vault whose long-term key cannot be had: it must fail without writing
|
||||
// anything, so that it never leaves a partial unlocker, nor breaks the one
|
||||
// it would replace.
|
||||
func TestPassphraseUnlockerGetsKeyFirst(t *testing.T) {
|
||||
// No mnemonic, and no current unlocker to get the key from
|
||||
t.Setenv(secret.EnvMnemonic, "")
|
||||
|
||||
base := afero.NewMemMapFs()
|
||||
_, err := vault.CreateVault(base, testVaultStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
fs := hookFs{Fs: base, before: func(_, path string) error {
|
||||
t.Errorf("changed %s before getting the long-term key", path)
|
||||
|
||||
return nil
|
||||
}}
|
||||
|
||||
passphrase := memguard.NewBufferFromBytes([]byte(unlockerPassphrase))
|
||||
defer passphrase.Destroy()
|
||||
|
||||
_, err = vault.NewVault(fs, testVaultStateDir, testVaultName).
|
||||
CreatePassphraseUnlocker(passphrase)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
// TestPassphraseUnlockerIsWholeOrAbsent checks, before every change that
|
||||
// creating a passphrase unlocker makes, that the unlocker's directory either
|
||||
// does not exist or holds all of its files: a crash or a failure at any point
|
||||
// leaves no partial unlocker.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids t.Parallel
|
||||
func TestPassphraseUnlockerIsWholeOrAbsent(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
files := []string{"pub.age", privKeyFile, "longterm.age", unlockerMetadataFile}
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
base, stateDir := tfs.open(t)
|
||||
vlt, err := vault.CreateVault(base, stateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
|
||||
unlockerDir := filepath.Join(vaultDir, "unlockers.d", "passphrase")
|
||||
|
||||
fs := hookFs{Fs: base, before: func(string, string) error {
|
||||
exists, err := afero.DirExists(base, unlockerDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
if exists {
|
||||
assert.ElementsMatch(t, files, dirNames(t, base, unlockerDir),
|
||||
"unlocker directory visible before it was complete")
|
||||
}
|
||||
|
||||
return nil
|
||||
}}
|
||||
|
||||
passphrase := memguard.NewBufferFromBytes([]byte(unlockerPassphrase))
|
||||
defer passphrase.Destroy()
|
||||
|
||||
_, err = vault.NewVault(fs, stateDir, testVaultName).
|
||||
CreatePassphraseUnlocker(passphrase)
|
||||
require.NoError(t, err)
|
||||
assert.ElementsMatch(t, files, dirNames(t, base, unlockerDir))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteDirFailureLeavesNothing makes writing a new directory fail after
|
||||
// a file has been written in it, and checks that neither the directory nor
|
||||
// its temporary directory is left behind; and, when the temporary directory
|
||||
// cannot be removed either, that both failures are reported.
|
||||
func TestWriteDirFailureLeavesNothing(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base, dir := tfs.open(t)
|
||||
listed := filepath.Join(dir, "unlockers.d")
|
||||
target := filepath.Join(listed, "new")
|
||||
|
||||
writeThenFail := func(tmp string) error {
|
||||
require.NoError(t, secret.WriteFileAtomic(base,
|
||||
filepath.Join(tmp, unlockerMetadataFile), []byte("{}")))
|
||||
|
||||
return errInjected
|
||||
}
|
||||
|
||||
err := secret.WriteDir(base, target, writeThenFail)
|
||||
require.ErrorIs(t, err, errInjected)
|
||||
|
||||
// Nothing in the directory that is listed, nor beside it
|
||||
assert.Empty(t, dirNames(t, base, listed))
|
||||
assert.Equal(t, []string{"unlockers.d"}, dirNames(t, base, dir))
|
||||
|
||||
fs := hookFs{Fs: base, before: func(op, _ string) error {
|
||||
if op == opRemove {
|
||||
return os.ErrPermission
|
||||
}
|
||||
|
||||
return nil
|
||||
}}
|
||||
|
||||
err = secret.WriteDir(fs, target, writeThenFail)
|
||||
require.ErrorIs(t, err, errInjected)
|
||||
require.ErrorIs(t, err, os.ErrPermission)
|
||||
assert.Empty(t, dirNames(t, base, listed))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWriteDirKeepsExistingDir makes writing into a directory that already
|
||||
// exists fail, and checks that the directory, with what was in it, is still
|
||||
// there: WriteDir writes into it in place and never removes it.
|
||||
func TestWriteDirKeepsExistingDir(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tfs := range testFilesystems {
|
||||
t.Run(tfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs, dir := tfs.open(t)
|
||||
target := filepath.Join(dir, "unlockers.d", "passphrase")
|
||||
require.NoError(t, fs.MkdirAll(target, secret.DirPerms))
|
||||
require.NoError(t, secret.WriteFileAtomic(fs,
|
||||
filepath.Join(target, unlockerMetadataFile), []byte("{}")))
|
||||
|
||||
err := secret.WriteDir(fs, target, func(got string) error {
|
||||
assert.Equal(t, target, got)
|
||||
|
||||
return errInjected
|
||||
})
|
||||
require.ErrorIs(t, err, errInjected)
|
||||
assert.Equal(t, []string{unlockerMetadataFile}, dirNames(t, fs, target))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -12,7 +12,8 @@ const (
|
||||
// EnvMnemonic is the environment variable for providing the mnemonic phrase
|
||||
EnvMnemonic = "SB_SECRET_MNEMONIC"
|
||||
// EnvUnlockPassphrase is the environment variable for providing the unlock passphrase
|
||||
EnvUnlockPassphrase = "SB_UNLOCK_PASSPHRASE" //nolint:gosec // G101: This is an env var name, not a credential
|
||||
//nolint:gosec // G101: env var name, not a credential
|
||||
EnvUnlockPassphrase = "SB_UNLOCK_PASSPHRASE"
|
||||
// EnvGPGKeyID is the environment variable for providing the GPG key ID
|
||||
EnvGPGKeyID = "SB_GPG_KEY_ID"
|
||||
)
|
||||
|
||||
+71
-30
@@ -2,6 +2,7 @@ package secret
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
@@ -12,39 +13,61 @@ import (
|
||||
"golang.org/x/term"
|
||||
)
|
||||
|
||||
var (
|
||||
errNilPassphraseBuffer = errors.New("passphrase buffer is nil")
|
||||
errStdinNotTerminal = errors.New(
|
||||
"cannot read passphrase from non-terminal stdin " +
|
||||
"(piped input or script). Please set the SB_UNLOCK_PASSPHRASE " +
|
||||
"environment variable or run interactively")
|
||||
errStderrNotTerminal = errors.New(
|
||||
"cannot prompt for passphrase: stderr is not a terminal " +
|
||||
"(running in non-interactive mode). Please set the " +
|
||||
"SB_UNLOCK_PASSPHRASE environment variable")
|
||||
errEmptyPassphrase = errors.New("passphrase cannot be empty")
|
||||
)
|
||||
|
||||
// EncryptToRecipient encrypts data to a recipient using age
|
||||
// The data parameter should be a LockedBuffer for secure memory handling
|
||||
func EncryptToRecipient(data *memguard.LockedBuffer, recipient age.Recipient) ([]byte, error) {
|
||||
func EncryptToRecipient(
|
||||
data *memguard.LockedBuffer, recipient age.Recipient,
|
||||
) ([]byte, error) {
|
||||
if data == nil {
|
||||
return nil, fmt.Errorf("data buffer is nil")
|
||||
return nil, errNilDataBuffer
|
||||
}
|
||||
|
||||
Debug("EncryptToRecipient starting", "data_length", data.Size())
|
||||
|
||||
var buf bytes.Buffer
|
||||
|
||||
Debug("Creating age encryptor")
|
||||
|
||||
w, err := age.Encrypt(&buf, recipient)
|
||||
if err != nil {
|
||||
Debug("Failed to create encryptor", "error", err)
|
||||
|
||||
return nil, fmt.Errorf("failed to create encryptor: %w", err)
|
||||
}
|
||||
Debug("Created age encryptor successfully")
|
||||
|
||||
Debug("Created age encryptor successfully")
|
||||
Debug("Writing data to encryptor")
|
||||
if _, err := w.Write(data.Bytes()); err != nil {
|
||||
|
||||
_, err = w.Write(data.Bytes())
|
||||
if err != nil {
|
||||
Debug("Failed to write data to encryptor", "error", err)
|
||||
|
||||
return nil, fmt.Errorf("failed to write data: %w", err)
|
||||
}
|
||||
Debug("Wrote data to encryptor successfully")
|
||||
|
||||
Debug("Wrote data to encryptor successfully")
|
||||
Debug("Closing encryptor")
|
||||
if err := w.Close(); err != nil {
|
||||
|
||||
err = w.Close()
|
||||
if err != nil {
|
||||
Debug("Failed to close encryptor", "error", err)
|
||||
|
||||
return nil, fmt.Errorf("failed to close encryptor: %w", err)
|
||||
}
|
||||
|
||||
Debug("Closed encryptor successfully")
|
||||
|
||||
result := buf.Bytes()
|
||||
@@ -54,7 +77,9 @@ func EncryptToRecipient(data *memguard.LockedBuffer, recipient age.Recipient) ([
|
||||
}
|
||||
|
||||
// DecryptWithIdentity decrypts data with an identity using age
|
||||
func DecryptWithIdentity(data []byte, identity age.Identity) (*memguard.LockedBuffer, error) {
|
||||
func DecryptWithIdentity(
|
||||
data []byte, identity age.Identity,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
r, err := age.Decrypt(bytes.NewReader(data), identity)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create decryptor: %w", err)
|
||||
@@ -68,20 +93,31 @@ func DecryptWithIdentity(data []byte, identity age.Identity) (*memguard.LockedBu
|
||||
// Create a secure buffer for the decrypted data
|
||||
resultBuffer := memguard.NewBufferFromBytes(result)
|
||||
|
||||
// Zero out the original slice to prevent plaintext from lingering
|
||||
// in unprotected memory
|
||||
for i := range result {
|
||||
result[i] = 0
|
||||
}
|
||||
|
||||
return resultBuffer, nil
|
||||
}
|
||||
|
||||
// EncryptWithPassphrase encrypts data using a passphrase with age's scrypt-based encryption
|
||||
// Both data and passphrase parameters should be LockedBuffers for secure memory handling
|
||||
func EncryptWithPassphrase(data *memguard.LockedBuffer, passphrase *memguard.LockedBuffer) ([]byte, error) {
|
||||
// EncryptWithPassphrase encrypts data using a passphrase with age's
|
||||
// scrypt-based encryption. Both data and passphrase parameters should
|
||||
// be LockedBuffers for secure memory handling
|
||||
func EncryptWithPassphrase(
|
||||
data *memguard.LockedBuffer, passphrase *memguard.LockedBuffer,
|
||||
) ([]byte, error) {
|
||||
if data == nil {
|
||||
return nil, fmt.Errorf("data buffer is nil")
|
||||
}
|
||||
if passphrase == nil {
|
||||
return nil, fmt.Errorf("passphrase buffer is nil")
|
||||
return nil, errNilDataBuffer
|
||||
}
|
||||
|
||||
// Create recipient directly from passphrase - unavoidable string conversion due to age API
|
||||
if passphrase == nil {
|
||||
return nil, errNilPassphraseBuffer
|
||||
}
|
||||
|
||||
// Create recipient directly from passphrase - unavoidable string
|
||||
// conversion due to age API
|
||||
recipient, err := age.NewScryptRecipient(passphrase.String())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create scrypt recipient: %w", err)
|
||||
@@ -90,14 +126,18 @@ func EncryptWithPassphrase(data *memguard.LockedBuffer, passphrase *memguard.Loc
|
||||
return EncryptToRecipient(data, recipient)
|
||||
}
|
||||
|
||||
// DecryptWithPassphrase decrypts data using a passphrase with age's scrypt-based decryption
|
||||
// The passphrase parameter should be a LockedBuffer for secure memory handling
|
||||
func DecryptWithPassphrase(encryptedData []byte, passphrase *memguard.LockedBuffer) (*memguard.LockedBuffer, error) {
|
||||
// DecryptWithPassphrase decrypts data using a passphrase with age's
|
||||
// scrypt-based decryption. The passphrase parameter should be a
|
||||
// LockedBuffer for secure memory handling
|
||||
func DecryptWithPassphrase(
|
||||
encryptedData []byte, passphrase *memguard.LockedBuffer,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
if passphrase == nil {
|
||||
return nil, fmt.Errorf("passphrase buffer is nil")
|
||||
return nil, errNilPassphraseBuffer
|
||||
}
|
||||
|
||||
// Create identity directly from passphrase - unavoidable string conversion due to age API
|
||||
// Create identity directly from passphrase - unavoidable string
|
||||
// conversion due to age API
|
||||
identity, err := age.NewScryptIdentity(passphrase.String())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create scrypt identity: %w", err)
|
||||
@@ -112,29 +152,30 @@ func DecryptWithPassphrase(encryptedData []byte, passphrase *memguard.LockedBuff
|
||||
func ReadPassphrase(prompt string) (*memguard.LockedBuffer, error) {
|
||||
// Check if stdin is a terminal
|
||||
if !term.IsTerminal(syscall.Stdin) {
|
||||
// Not a terminal - never read passphrases from piped input for security reasons
|
||||
return nil, fmt.Errorf("cannot read passphrase from non-terminal stdin " +
|
||||
"(piped input or script). Please set the SB_UNLOCK_PASSPHRASE " +
|
||||
"environment variable or run interactively")
|
||||
// Not a terminal - never read passphrases from piped input
|
||||
// for security reasons
|
||||
return nil, errStdinNotTerminal
|
||||
}
|
||||
|
||||
// stdin is a terminal, check if stderr is also a terminal for interactive prompting
|
||||
// stdin is a terminal, check if stderr is also a terminal for
|
||||
// interactive prompting
|
||||
if !term.IsTerminal(syscall.Stderr) {
|
||||
return nil, fmt.Errorf("cannot prompt for passphrase: stderr is not a terminal " +
|
||||
"(running in non-interactive mode). Please set the SB_UNLOCK_PASSPHRASE " +
|
||||
"environment variable")
|
||||
return nil, errStderrNotTerminal
|
||||
}
|
||||
|
||||
// Both stdin and stderr are terminals - use secure password reading
|
||||
fmt.Fprint(os.Stderr, prompt) // Write prompt to stderr, not stdout
|
||||
|
||||
passphrase, err := term.ReadPassword(syscall.Stdin)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read passphrase: %w", err)
|
||||
}
|
||||
fmt.Fprintln(os.Stderr) // Print newline to stderr since ReadPassword doesn't echo
|
||||
|
||||
// Print newline to stderr since ReadPassword doesn't echo
|
||||
fmt.Fprintln(os.Stderr)
|
||||
|
||||
if len(passphrase) == 0 {
|
||||
return nil, fmt.Errorf("passphrase cannot be empty")
|
||||
return nil, errEmptyPassphrase
|
||||
}
|
||||
|
||||
// Create a secure buffer and copy the passphrase
|
||||
|
||||
@@ -13,28 +13,33 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
debugEnabled bool //nolint:gochecknoglobals // Package-wide debug state is necessary
|
||||
debugLogger *slog.Logger //nolint:gochecknoglobals // Package-wide logger instance is necessary
|
||||
debugEnabled bool //nolint:gochecknoglobals // package debug state
|
||||
debugLogger *slog.Logger //nolint:gochecknoglobals // package debug logger
|
||||
)
|
||||
|
||||
//nolint:gochecknoinits // debug logging must be ready before any package use
|
||||
func init() {
|
||||
InitDebugLogging()
|
||||
}
|
||||
|
||||
// InitDebugLogging initializes the debug logging system based on current GODEBUG environment variable
|
||||
// InitDebugLogging initializes the debug logging system based on the
|
||||
// current GODEBUG environment variable
|
||||
func InitDebugLogging() {
|
||||
godebug := os.Getenv("GODEBUG")
|
||||
debugEnabled = strings.Contains(godebug, "berlin.sneak.pkg.secret")
|
||||
|
||||
if !debugEnabled {
|
||||
// Create a no-op logger that discards all output
|
||||
debugLogger = slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
debugLogger = slog.New(slog.DiscardHandler)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
// Disable stderr buffering for immediate debug output when debugging is enabled
|
||||
_, _, _ = syscall.Syscall(syscall.SYS_FCNTL, os.Stderr.Fd(), syscall.F_SETFL, syscall.O_SYNC)
|
||||
// Disable stderr buffering for immediate debug output when
|
||||
// debugging is enabled
|
||||
//nolint:dogsled // syscall.Syscall returns three values, none needed
|
||||
_, _, _ = syscall.Syscall(
|
||||
syscall.SYS_FCNTL, os.Stderr.Fd(), syscall.F_SETFL, syscall.O_SYNC)
|
||||
|
||||
// Check if STDERR is a TTY
|
||||
isTTY := term.IsTerminal(syscall.Stderr)
|
||||
@@ -58,19 +63,36 @@ func IsDebugEnabled() bool {
|
||||
return debugEnabled
|
||||
}
|
||||
|
||||
// Warn logs a warning message to stderr unconditionally (visible
|
||||
// without --verbose or debug flags)
|
||||
func Warn(msg string, args ...any) {
|
||||
var output strings.Builder
|
||||
|
||||
output.WriteString("WARNING: " + msg)
|
||||
|
||||
for i := 0; i+1 < len(args); i += 2 {
|
||||
fmt.Fprintf(&output, " %s=%v", args[i], args[i+1])
|
||||
}
|
||||
|
||||
output.WriteString("\n")
|
||||
fmt.Fprint(os.Stderr, output.String())
|
||||
}
|
||||
|
||||
// Debug logs a debug message with optional attributes
|
||||
func Debug(msg string, args ...any) {
|
||||
if !debugEnabled {
|
||||
return
|
||||
}
|
||||
|
||||
debugLogger.Debug(msg, args...)
|
||||
}
|
||||
|
||||
// DebugF logs a formatted debug message with optional attributes
|
||||
func DebugF(format string, args ...any) {
|
||||
// Debugf logs a formatted debug message with optional attributes
|
||||
func Debugf(format string, args ...any) {
|
||||
if !debugEnabled {
|
||||
return
|
||||
}
|
||||
|
||||
debugLogger.Debug(fmt.Sprintf(format, args...))
|
||||
}
|
||||
|
||||
@@ -79,6 +101,7 @@ func DebugWith(msg string, attrs ...slog.Attr) {
|
||||
if !debugEnabled {
|
||||
return
|
||||
}
|
||||
|
||||
debugLogger.LogAttrs(context.Background(), slog.LevelDebug, msg, attrs...)
|
||||
}
|
||||
|
||||
@@ -108,15 +131,18 @@ func (h *colorizedHandler) Handle(_ context.Context, record slog.Record) error {
|
||||
if record.NumAttrs() > 0 {
|
||||
output += " \033[33m{"
|
||||
first := true
|
||||
|
||||
record.Attrs(func(attr slog.Attr) bool {
|
||||
if !first {
|
||||
output += ", "
|
||||
}
|
||||
|
||||
first = false
|
||||
output += fmt.Sprintf("%s=%#v", attr.Key, attr.Value.Any())
|
||||
|
||||
return true
|
||||
})
|
||||
|
||||
output += "}\033[0m"
|
||||
}
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
//nolint:testpackage // white-box test of unexported debug internals
|
||||
package secret
|
||||
|
||||
import (
|
||||
@@ -90,9 +91,11 @@ func TestDebugLogging(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
//nolint:paralleltest // exercises process-global debug logger state
|
||||
func TestDebugFunctions(t *testing.T) {
|
||||
// Enable debug for testing
|
||||
t.Setenv("GODEBUG", "berlin.sneak.pkg.secret")
|
||||
|
||||
defer InitDebugLogging() // Re-initialize after test
|
||||
|
||||
InitDebugLogging()
|
||||
@@ -107,8 +110,8 @@ func TestDebugFunctions(t *testing.T) {
|
||||
Debug("test with args", "key", "value", "number", 42)
|
||||
})
|
||||
|
||||
t.Run("DebugF", func(_ *testing.T) {
|
||||
DebugF("formatted message: %s %d", "test", 123)
|
||||
t.Run("Debugf", func(_ *testing.T) {
|
||||
Debugf("formatted message: %s %d", "test", 123)
|
||||
})
|
||||
|
||||
t.Run("DebugWith", func(_ *testing.T) {
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
//go:build darwin
|
||||
|
||||
package secret
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// realVault is a minimal VaultInterface backed by a real afero filesystem,
|
||||
// using the same directory layout as vault.Vault.
|
||||
type realVault struct {
|
||||
name string
|
||||
stateDir string
|
||||
fs afero.Fs
|
||||
}
|
||||
|
||||
func (v *realVault) GetDirectory() (string, error) {
|
||||
return filepath.Join(v.stateDir, "vaults.d", v.name), nil
|
||||
}
|
||||
func (v *realVault) GetName() string { return v.name }
|
||||
func (v *realVault) GetFilesystem() afero.Fs { return v.fs }
|
||||
|
||||
// Unused by getLongTermPrivateKey — these satisfy VaultInterface.
|
||||
func (v *realVault) AddSecret(string, *memguard.LockedBuffer, bool) error { panic("not used") }
|
||||
func (v *realVault) GetCurrentUnlocker() (Unlocker, error) { panic("not used") }
|
||||
func (v *realVault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) { panic("not used") }
|
||||
func (v *realVault) CreatePassphraseUnlocker(*memguard.LockedBuffer) (*PassphraseUnlocker, error) {
|
||||
panic("not used")
|
||||
}
|
||||
|
||||
// createRealVault sets up a complete vault directory structure on an in-memory
|
||||
// filesystem, identical to what vault.CreateVault produces.
|
||||
func createRealVault(t *testing.T, fs afero.Fs, stateDir, name string, derivationIndex uint32) *realVault {
|
||||
t.Helper()
|
||||
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", name)
|
||||
require.NoError(t, fs.MkdirAll(filepath.Join(vaultDir, "secrets.d"), DirPerms))
|
||||
require.NoError(t, fs.MkdirAll(filepath.Join(vaultDir, "unlockers.d"), DirPerms))
|
||||
|
||||
metadata := VaultMetadata{
|
||||
CreatedAt: time.Now(),
|
||||
DerivationIndex: derivationIndex,
|
||||
}
|
||||
metaBytes, err := json.Marshal(metadata)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, afero.WriteFile(fs, filepath.Join(vaultDir, "vault-metadata.json"), metaBytes, FilePerms))
|
||||
|
||||
return &realVault{name: name, stateDir: stateDir, fs: fs}
|
||||
}
|
||||
|
||||
func TestGetLongTermPrivateKeyUsesVaultDerivationIndex(t *testing.T) {
|
||||
const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
// Derive expected keys at two different indices to prove they differ.
|
||||
key0, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
require.NoError(t, err)
|
||||
key5, err := agehd.DeriveIdentity(testMnemonic, 5)
|
||||
require.NoError(t, err)
|
||||
require.NotEqual(t, key0.String(), key5.String(),
|
||||
"sanity check: different derivation indices must produce different keys")
|
||||
|
||||
// Build a real vault with DerivationIndex=5 on an in-memory filesystem.
|
||||
fs := afero.NewMemMapFs()
|
||||
vault := createRealVault(t, fs, "/state", "test-vault", 5)
|
||||
|
||||
t.Setenv(EnvMnemonic, testMnemonic)
|
||||
|
||||
result, err := getLongTermPrivateKey(fs, vault)
|
||||
require.NoError(t, err)
|
||||
defer result.Destroy()
|
||||
|
||||
assert.Equal(t, key5.String(), string(result.Bytes()),
|
||||
"getLongTermPrivateKey should derive at vault's DerivationIndex (5)")
|
||||
assert.NotEqual(t, key0.String(), string(result.Bytes()),
|
||||
"getLongTermPrivateKey must not use hardcoded index 0")
|
||||
}
|
||||
+18
-29
@@ -1,43 +1,23 @@
|
||||
package secret
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"os"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
// generateRandomString generates a random string of the specified length using the given character set
|
||||
func generateRandomString(length int, charset string) (string, error) {
|
||||
if length <= 0 {
|
||||
return "", fmt.Errorf("length must be positive")
|
||||
}
|
||||
|
||||
result := make([]byte, length)
|
||||
charsetLen := big.NewInt(int64(len(charset)))
|
||||
|
||||
for i := range length {
|
||||
randomIndex, err := rand.Int(rand.Reader, charsetLen)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to generate random number: %w", err)
|
||||
}
|
||||
result[i] = charset[randomIndex.Int64()]
|
||||
}
|
||||
|
||||
return string(result), nil
|
||||
}
|
||||
|
||||
// DetermineStateDir determines the state directory based on environment variables and OS
|
||||
func DetermineStateDir(customConfigDir string) string {
|
||||
// DetermineStateDir determines the state directory based on environment
|
||||
// variables and OS.
|
||||
// It returns an error if no usable directory can be determined.
|
||||
func DetermineStateDir(customConfigDir string) (string, error) {
|
||||
// Check for environment variable first
|
||||
if envStateDir := os.Getenv(EnvStateDir); envStateDir != "" {
|
||||
return envStateDir
|
||||
return envStateDir, nil
|
||||
}
|
||||
|
||||
// Use custom config dir if provided
|
||||
if customConfigDir != "" {
|
||||
return filepath.Join(customConfigDir, AppID)
|
||||
return filepath.Join(customConfigDir, AppID), nil
|
||||
}
|
||||
|
||||
// Use os.UserConfigDir() which handles platform-specific directories:
|
||||
@@ -47,10 +27,19 @@ func DetermineStateDir(customConfigDir string) string {
|
||||
configDir, err := os.UserConfigDir()
|
||||
if err != nil {
|
||||
// Fallback to a reasonable default if we can't determine user config dir
|
||||
homeDir, _ := os.UserHomeDir()
|
||||
homeDir, homeErr := os.UserHomeDir()
|
||||
if homeErr != nil {
|
||||
return "", fmt.Errorf(
|
||||
"unable to determine state directory: config dir: %w, home dir: %w",
|
||||
err, homeErr)
|
||||
}
|
||||
|
||||
return filepath.Join(homeDir, ".config", AppID)
|
||||
fallbackDir := filepath.Join(homeDir, ".config", AppID)
|
||||
Warn("Could not determine user config directory, falling back to default",
|
||||
"fallback", fallbackDir, "error", err)
|
||||
|
||||
return fallbackDir, nil
|
||||
}
|
||||
|
||||
return filepath.Join(configDir, AppID)
|
||||
return filepath.Join(configDir, AppID), nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package secret_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
)
|
||||
|
||||
func TestDetermineStateDir_ErrorsWhenHomeDirUnavailable(t *testing.T) {
|
||||
// Clear all env vars that could provide a home/config directory.
|
||||
// On Darwin, os.UserHomeDir may still succeed via the password
|
||||
// database, so we also test via an explicit empty-customConfigDir
|
||||
// path to exercise the fallback branch.
|
||||
t.Setenv(secret.EnvStateDir, "")
|
||||
t.Setenv("HOME", "")
|
||||
t.Setenv("XDG_CONFIG_HOME", "")
|
||||
|
||||
result, err := secret.DetermineStateDir("")
|
||||
// On systems where both lookups fail, we must get an error.
|
||||
// On systems where the OS provides a fallback (e.g. macOS pw db),
|
||||
// result should still be valid (non-empty, not root-relative).
|
||||
if err != nil {
|
||||
// Good — the error case is handled.
|
||||
return
|
||||
}
|
||||
|
||||
if result == "/.config/"+secret.AppID || result == "" {
|
||||
t.Errorf(
|
||||
"DetermineStateDir returned dangerous/empty path %q without error",
|
||||
result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetermineStateDir_UsesEnvVar(t *testing.T) {
|
||||
t.Setenv(secret.EnvStateDir, "/custom/state")
|
||||
|
||||
result, err := secret.DetermineStateDir("")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if result != "/custom/state" {
|
||||
t.Errorf("expected /custom/state, got %q", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetermineStateDir_UsesCustomConfigDir(t *testing.T) {
|
||||
t.Setenv(secret.EnvStateDir, "")
|
||||
|
||||
result, err := secret.DetermineStateDir("/my/config")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
expected := "/my/config/" + secret.AppID
|
||||
if result != expected {
|
||||
t.Errorf("expected %q, got %q", expected, result)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,142 @@
|
||||
package secret
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/awnumar/memguard"
|
||||
)
|
||||
|
||||
var (
|
||||
errPassphraseLength = errors.New(
|
||||
"passphrase length must be a positive even number")
|
||||
errPassphraseNotHex = errors.New(
|
||||
"keychain passphrase must be lowercase hex")
|
||||
errNoKeychainPassphrase = errors.New(
|
||||
"keychain data has no agePrivKeyPassphrase string")
|
||||
)
|
||||
|
||||
// KeychainData is what a keychain unlocker stores in the macOS keychain.
|
||||
// It is stored as JSON, but encode and decodeKeychainData keep the
|
||||
// passphrase out of encoding/json, which would leave copies of it in
|
||||
// ordinary memory.
|
||||
type KeychainData struct {
|
||||
AgePublicKey string
|
||||
AgePrivKeyPassphrase *memguard.LockedBuffer
|
||||
EncryptedLongtermKey string
|
||||
}
|
||||
|
||||
// generateRandomPassphrase returns length random lowercase hex characters
|
||||
// in a locked buffer. The caller must destroy it.
|
||||
func generateRandomPassphrase(length int) (*memguard.LockedBuffer, error) {
|
||||
// Each random byte becomes two hex characters.
|
||||
randomBytes := hex.DecodedLen(length)
|
||||
if length <= 0 || hex.EncodedLen(randomBytes) != length {
|
||||
return nil, errPassphraseLength
|
||||
}
|
||||
|
||||
random := memguard.NewBufferRandom(randomBytes)
|
||||
defer random.Destroy()
|
||||
|
||||
passphrase := memguard.NewBuffer(length)
|
||||
hex.Encode(passphrase.Bytes(), random.Bytes())
|
||||
passphrase.Freeze()
|
||||
|
||||
return passphrase, nil
|
||||
}
|
||||
|
||||
// encode returns d as JSON in a locked buffer:
|
||||
// {"agePublicKey":"...","agePrivKeyPassphrase":"...","encryptedLongtermKey":"..."}.
|
||||
// The passphrase is copied straight into the buffer, so it must be hex,
|
||||
// which JSON does not escape. The caller must destroy the returned buffer.
|
||||
func (d *KeychainData) encode() (*memguard.LockedBuffer, error) {
|
||||
if d.AgePrivKeyPassphrase == nil {
|
||||
return nil, errNilPassphraseBuffer
|
||||
}
|
||||
|
||||
if d.AgePrivKeyPassphrase.Size() == 0 {
|
||||
return nil, errEmptyPassphrase
|
||||
}
|
||||
|
||||
for _, c := range d.AgePrivKeyPassphrase.Bytes() {
|
||||
if strings.IndexByte("0123456789abcdef", c) < 0 {
|
||||
return nil, errPassphraseNotHex
|
||||
}
|
||||
}
|
||||
|
||||
publicKey, err := json.Marshal(d.AgePublicKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encode age public key: %w", err)
|
||||
}
|
||||
|
||||
longtermKey, err := json.Marshal(d.EncryptedLongtermKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encode long-term key: %w", err)
|
||||
}
|
||||
|
||||
parts := [][]byte{
|
||||
[]byte(`{"agePublicKey":`), publicKey,
|
||||
[]byte(`,"agePrivKeyPassphrase":"`), d.AgePrivKeyPassphrase.Bytes(),
|
||||
[]byte(`","encryptedLongtermKey":`), longtermKey,
|
||||
[]byte(`}`),
|
||||
}
|
||||
|
||||
size := 0
|
||||
for _, part := range parts {
|
||||
size += len(part)
|
||||
}
|
||||
|
||||
encoded := memguard.NewBuffer(size)
|
||||
|
||||
written := 0
|
||||
for _, part := range parts {
|
||||
written += copy(encoded.Bytes()[written:], part)
|
||||
}
|
||||
|
||||
encoded.Freeze()
|
||||
|
||||
return encoded, nil
|
||||
}
|
||||
|
||||
// decodeKeychainData parses keychain data written by encode. The caller
|
||||
// must destroy the returned AgePrivKeyPassphrase.
|
||||
func decodeKeychainData(data *memguard.LockedBuffer) (*KeychainData, error) {
|
||||
if data == nil {
|
||||
return nil, errNilDataBuffer
|
||||
}
|
||||
|
||||
// json.Unmarshal gives a json.RawMessage field the field's JSON text
|
||||
// unchanged, in the one copy RawMessage makes; it is wiped on return.
|
||||
var fields struct {
|
||||
AgePublicKey string `json:"agePublicKey"`
|
||||
AgePrivKeyPassphrase json.RawMessage `json:"agePrivKeyPassphrase"`
|
||||
EncryptedLongtermKey string `json:"encryptedLongtermKey"`
|
||||
}
|
||||
|
||||
defer func() { memguard.WipeBytes(fields.AgePrivKeyPassphrase) }()
|
||||
|
||||
err := json.Unmarshal(data.Bytes(), &fields)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse keychain data: %w", err)
|
||||
}
|
||||
|
||||
// json.Unmarshal accepted the JSON, so text that starts with a quote is
|
||||
// a whole string. The passphrase is hex, so it is the text between the
|
||||
// quotes.
|
||||
quoted := fields.AgePrivKeyPassphrase
|
||||
if !bytes.HasPrefix(quoted, []byte(`"`)) {
|
||||
return nil, errNoKeychainPassphrase
|
||||
}
|
||||
|
||||
return &KeychainData{
|
||||
AgePublicKey: fields.AgePublicKey,
|
||||
// NewBufferFromBytes wipes the bytes it copies.
|
||||
AgePrivKeyPassphrase: memguard.NewBufferFromBytes(
|
||||
quoted[1 : len(quoted)-1]),
|
||||
EncryptedLongtermKey: fields.EncryptedLongtermKey,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package secret
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestGenerateRandomPassphrase(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
first, err := generateRandomPassphrase(64)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer first.Destroy()
|
||||
|
||||
second, err := generateRandomPassphrase(64)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer second.Destroy()
|
||||
|
||||
assert.Regexp(t, `^[0-9a-f]{64}$`, first.String())
|
||||
assert.NotEqual(t, first.String(), second.String())
|
||||
assert.False(t, first.IsMutable())
|
||||
|
||||
for _, length := range []int{0, -2, 63} {
|
||||
_, err := generateRandomPassphrase(length)
|
||||
require.ErrorIs(t, err, errPassphraseLength, "length %d", length)
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeychainDataEncodeDecode(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
passphrase := memguard.NewBufferFromBytes([]byte("0a1b2c3d"))
|
||||
defer passphrase.Destroy()
|
||||
|
||||
data := KeychainData{
|
||||
AgePublicKey: "age1example",
|
||||
AgePrivKeyPassphrase: passphrase,
|
||||
EncryptedLongtermKey: "beef",
|
||||
}
|
||||
|
||||
encoded, err := data.encode()
|
||||
require.NoError(t, err)
|
||||
|
||||
defer encoded.Destroy()
|
||||
|
||||
assert.JSONEq(t,
|
||||
`{"agePublicKey":"age1example",`+
|
||||
`"agePrivKeyPassphrase":"0a1b2c3d",`+
|
||||
`"encryptedLongtermKey":"beef"}`,
|
||||
encoded.String())
|
||||
assert.False(t, encoded.IsMutable())
|
||||
|
||||
decoded, err := decodeKeychainData(encoded)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer decoded.AgePrivKeyPassphrase.Destroy()
|
||||
|
||||
assert.Equal(t, "age1example", decoded.AgePublicKey)
|
||||
assert.Equal(t, "0a1b2c3d", decoded.AgePrivKeyPassphrase.String())
|
||||
assert.Equal(t, "beef", decoded.EncryptedLongtermKey)
|
||||
}
|
||||
|
||||
func TestKeychainDataEncodeRejectsBadPassphrase(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
passphrase *memguard.LockedBuffer
|
||||
wantErr error
|
||||
}{
|
||||
{"nil", nil, errNilPassphraseBuffer},
|
||||
{"empty", memguard.NewBuffer(0), errEmptyPassphrase},
|
||||
{
|
||||
"not hex",
|
||||
memguard.NewBufferFromBytes([]byte(`abc"def`)),
|
||||
errPassphraseNotHex,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
data := KeychainData{AgePrivKeyPassphrase: tt.passphrase}
|
||||
_, err := data.encode()
|
||||
require.ErrorIs(t, err, tt.wantErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeKeychainDataRejectsBadData(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, text := range []string{
|
||||
`{"agePublicKey":"age1example"}`,
|
||||
`{"agePrivKeyPassphrase":42}`,
|
||||
} {
|
||||
data := memguard.NewBufferFromBytes([]byte(text))
|
||||
_, err := decodeKeychainData(data)
|
||||
data.Destroy()
|
||||
require.ErrorIs(t, err, errNoKeychainPassphrase, text)
|
||||
}
|
||||
|
||||
notJSON := memguard.NewBufferFromBytes([]byte(`{"agePrivKeyPassphrase":`))
|
||||
defer notJSON.Destroy()
|
||||
|
||||
_, err := decodeKeychainData(notJSON)
|
||||
|
||||
var syntaxError *json.SyntaxError
|
||||
require.ErrorAs(t, err, &syntaxError)
|
||||
}
|
||||
@@ -45,13 +45,6 @@ type KeychainUnlocker struct {
|
||||
fs afero.Fs
|
||||
}
|
||||
|
||||
// KeychainData represents the data stored in the macOS keychain
|
||||
type KeychainData struct {
|
||||
AgePublicKey string `json:"agePublicKey"`
|
||||
AgePrivKeyPassphrase string `json:"agePrivKeyPassphrase"`
|
||||
EncryptedLongtermKey string `json:"encryptedLongtermKey"`
|
||||
}
|
||||
|
||||
// GetIdentity implements Unlocker interface for Keychain-based unlockers
|
||||
func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
DebugWith("Getting keychain unlocker identity",
|
||||
@@ -81,13 +74,18 @@ func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
slog.Int("data_length", len(keychainDataBytes)),
|
||||
)
|
||||
|
||||
// Move the keychain data into locked memory; this wipes keychainDataBytes
|
||||
keychainDataBuffer := memguard.NewBufferFromBytes(keychainDataBytes)
|
||||
defer keychainDataBuffer.Destroy()
|
||||
|
||||
// Step 3: Parse keychain data
|
||||
var keychainData KeychainData
|
||||
if err := json.Unmarshal(keychainDataBytes, &keychainData); err != nil {
|
||||
keychainData, err := decodeKeychainData(keychainDataBuffer)
|
||||
if err != nil {
|
||||
Debug("Failed to parse keychain data", "error", err, "unlocker_id", k.GetID())
|
||||
|
||||
return nil, fmt.Errorf("failed to parse keychain data: %w", err)
|
||||
}
|
||||
defer keychainData.AgePrivKeyPassphrase.Destroy()
|
||||
|
||||
Debug("Parsed keychain data successfully", "unlocker_id", k.GetID())
|
||||
|
||||
@@ -109,11 +107,7 @@ func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
|
||||
// Step 5: Decrypt the age private key using the passphrase from keychain
|
||||
Debug("Decrypting age private key with keychain passphrase", "unlocker_id", k.GetID())
|
||||
// Create secure buffer for the keychain passphrase
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte(keychainData.AgePrivKeyPassphrase))
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
agePrivKeyBuffer, err := DecryptWithPassphrase(encryptedAgePrivKeyData, passphraseBuffer)
|
||||
agePrivKeyBuffer, err := DecryptWithPassphrase(encryptedAgePrivKeyData, keychainData.AgePrivKeyPassphrase)
|
||||
if err != nil {
|
||||
Debug("Failed to decrypt age private key with keychain passphrase", "error", err, "unlocker_id", k.GetID())
|
||||
|
||||
@@ -195,7 +189,7 @@ func (k *KeychainUnlocker) Remove() error {
|
||||
|
||||
// Step 3: Remove directory
|
||||
Debug("Removing keychain unlocker directory", "directory", k.Directory)
|
||||
if err := k.fs.RemoveAll(k.Directory); err != nil {
|
||||
if err := RemoveDirAtomic(k.fs, k.Directory); err != nil {
|
||||
Debug("Failed to remove keychain unlocker directory", "error", err, "directory", k.Directory)
|
||||
|
||||
return fmt.Errorf("failed to remove keychain unlocker directory: %w", err)
|
||||
@@ -251,8 +245,25 @@ func getLongTermPrivateKey(fs afero.Fs, vault VaultInterface) (*memguard.LockedB
|
||||
// Check if mnemonic is available in environment variable
|
||||
envMnemonic := os.Getenv(EnvMnemonic)
|
||||
if envMnemonic != "" {
|
||||
// Use mnemonic directly to derive long-term key
|
||||
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, 0)
|
||||
// Read vault metadata to get the correct derivation index
|
||||
vaultDir, err := vault.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
metadataPath := filepath.Join(vaultDir, "vault-metadata.json")
|
||||
metadataBytes, err := afero.ReadFile(fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read vault metadata: %w", err)
|
||||
}
|
||||
|
||||
var metadata VaultMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse vault metadata: %w", err)
|
||||
}
|
||||
|
||||
// Use mnemonic with the vault's actual derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
@@ -330,16 +341,13 @@ func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, er
|
||||
return nil, fmt.Errorf("failed to generate keychain item name: %w", err)
|
||||
}
|
||||
|
||||
// Create unlocker directory using the keychain item name as the directory name
|
||||
// The unlocker directory is named after the keychain item
|
||||
vaultDir, err := vault.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
unlockerDir := filepath.Join(vaultDir, "unlockers.d", keychainItemName)
|
||||
if err := fs.MkdirAll(unlockerDir, DirPerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to create unlocker directory: %w", err)
|
||||
}
|
||||
|
||||
// Step 1: Generate a new age keypair for the keychain unlocker
|
||||
ageIdentity, err := age.GenerateX25519Identity()
|
||||
@@ -347,79 +355,53 @@ func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, er
|
||||
return nil, fmt.Errorf("failed to generate age keypair: %w", err)
|
||||
}
|
||||
|
||||
ageRecipient := ageIdentity.Recipient().String()
|
||||
|
||||
// Step 2: Generate a random passphrase for encrypting the age private key
|
||||
agePrivKeyPassphrase, err := generateRandomPassphrase(agePrivKeyPassphraseLength)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate age private key passphrase: %w", err)
|
||||
}
|
||||
defer agePrivKeyPassphrase.Destroy()
|
||||
|
||||
// Step 3: Store age recipient as plaintext
|
||||
ageRecipient := ageIdentity.Recipient().String()
|
||||
recipientPath := filepath.Join(unlockerDir, "pub.txt")
|
||||
if err := afero.WriteFile(fs, recipientPath, []byte(ageRecipient), FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write age recipient: %w", err)
|
||||
}
|
||||
|
||||
// Step 4: Encrypt age private key with the generated passphrase and store on disk
|
||||
// Create secure buffers for both the private key and passphrase
|
||||
// Step 3: Encrypt age private key with the generated passphrase
|
||||
// Create a secure buffer for the private key
|
||||
agePrivKeyStr := ageIdentity.String()
|
||||
agePrivKeyBuffer := memguard.NewBufferFromBytes([]byte(agePrivKeyStr))
|
||||
defer agePrivKeyBuffer.Destroy()
|
||||
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte(agePrivKeyPassphrase))
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
encryptedAgePrivKey, err := EncryptWithPassphrase(agePrivKeyBuffer, passphraseBuffer)
|
||||
encryptedAgePrivKey, err := EncryptWithPassphrase(agePrivKeyBuffer, agePrivKeyPassphrase)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt age private key with passphrase: %w", err)
|
||||
}
|
||||
|
||||
agePrivKeyPath := filepath.Join(unlockerDir, "priv.age")
|
||||
if err := afero.WriteFile(fs, agePrivKeyPath, encryptedAgePrivKey, FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write encrypted age private key: %w", err)
|
||||
}
|
||||
|
||||
// Step 5: Get or derive the long-term private key
|
||||
// Step 4: Get or derive the long-term private key
|
||||
ltPrivKeyData, err := getLongTermPrivateKey(fs, vault)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer ltPrivKeyData.Destroy()
|
||||
|
||||
// Step 6: Encrypt long-term private key to the new age unlocker
|
||||
// Step 5: Encrypt long-term private key to the new age unlocker
|
||||
encryptedLtPrivKeyToAge, err := EncryptToRecipient(ltPrivKeyData, ageIdentity.Recipient())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt long-term private key to age unlocker: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted long-term private key
|
||||
ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age")
|
||||
if err := afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKeyToAge, FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
// Step 7: Prepare keychain data
|
||||
// Step 6: Prepare keychain data
|
||||
keychainData := KeychainData{
|
||||
AgePublicKey: ageRecipient,
|
||||
AgePrivKeyPassphrase: agePrivKeyPassphrase,
|
||||
EncryptedLongtermKey: hex.EncodeToString(encryptedLtPrivKeyToAge),
|
||||
}
|
||||
|
||||
keychainDataBytes, err := json.Marshal(keychainData)
|
||||
keychainDataBuffer, err := keychainData.encode()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal keychain data: %w", err)
|
||||
return nil, fmt.Errorf("failed to encode keychain data: %w", err)
|
||||
}
|
||||
|
||||
// Create a secure buffer for keychain data
|
||||
keychainDataBuffer := memguard.NewBufferFromBytes(keychainDataBytes)
|
||||
defer keychainDataBuffer.Destroy()
|
||||
|
||||
// Step 8: Store data in keychain
|
||||
if err := storeInKeychain(keychainItemName, keychainDataBuffer); err != nil {
|
||||
return nil, fmt.Errorf("failed to store data in keychain: %w", err)
|
||||
}
|
||||
|
||||
// Step 9: Create and write enhanced metadata
|
||||
// Step 7: Prepare enhanced metadata
|
||||
keychainMetadata := KeychainUnlockerMetadata{
|
||||
UnlockerMetadata: UnlockerMetadata{
|
||||
Type: "keychain",
|
||||
@@ -434,10 +416,37 @@ func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, er
|
||||
return nil, fmt.Errorf("failed to marshal unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
if err := afero.WriteFile(fs,
|
||||
filepath.Join(unlockerDir, "unlocker-metadata.json"),
|
||||
metadataBytes, FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write unlocker metadata: %w", err)
|
||||
// Step 8: Write the unlocker's files and store the data in the keychain,
|
||||
// the metadata last
|
||||
err = WriteDir(fs, unlockerDir, func(dir string) error {
|
||||
pubPath := filepath.Join(dir, "pub.txt")
|
||||
if err := WriteFileAtomic(fs, pubPath, []byte(ageRecipient)); err != nil {
|
||||
return fmt.Errorf("failed to write age recipient: %w", err)
|
||||
}
|
||||
|
||||
privPath := filepath.Join(dir, "priv.age")
|
||||
if err := WriteFileAtomic(fs, privPath, encryptedAgePrivKey); err != nil {
|
||||
return fmt.Errorf("failed to write encrypted age private key: %w", err)
|
||||
}
|
||||
|
||||
ltKeyPath := filepath.Join(dir, "longterm.age")
|
||||
if err := WriteFileAtomic(fs, ltKeyPath, encryptedLtPrivKeyToAge); err != nil {
|
||||
return fmt.Errorf("failed to write encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
if err := storeInKeychain(keychainItemName, keychainDataBuffer); err != nil {
|
||||
return fmt.Errorf("failed to store data in keychain: %w", err)
|
||||
}
|
||||
|
||||
metadataPath := filepath.Join(dir, "unlocker-metadata.json")
|
||||
if err := WriteFileAtomic(fs, metadataPath, metadataBytes); err != nil {
|
||||
return fmt.Errorf("failed to write unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &KeychainUnlocker{
|
||||
@@ -484,7 +493,7 @@ func storeInKeychain(itemName string, data *memguard.LockedBuffer) error {
|
||||
item.SetAccount(itemName)
|
||||
item.SetLabel(fmt.Sprintf("%s - %s", KEYCHAIN_APP_IDENTIFIER, itemName))
|
||||
item.SetDescription("Secret vault keychain data")
|
||||
item.SetData([]byte(data.String()))
|
||||
item.SetData(data.Bytes())
|
||||
item.SetSynchronizable(keychain.SynchronizableNo)
|
||||
// Use AccessibleWhenUnlockedThisDeviceOnly for better security and to trigger auth
|
||||
item.SetAccessible(keychain.AccessibleWhenUnlockedThisDeviceOnly)
|
||||
@@ -559,8 +568,3 @@ func deleteFromKeychain(itemName string) error {
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// generateRandomPassphrase generates a random passphrase for encrypting the age private key
|
||||
func generateRandomPassphrase(length int) (string, error) {
|
||||
return generateRandomString(length, "0123456789abcdef")
|
||||
}
|
||||
|
||||
@@ -1,17 +1,18 @@
|
||||
//go:build !darwin
|
||||
// +build !darwin
|
||||
|
||||
package secret
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"filippo.io/age"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// KeychainUnlockerMetadata is a stub for non-Darwin platforms
|
||||
type KeychainUnlockerMetadata struct {
|
||||
UnlockerMetadata
|
||||
|
||||
KeychainItemName string `json:"keychainItemName"`
|
||||
}
|
||||
|
||||
@@ -22,52 +23,58 @@ type KeychainUnlocker struct {
|
||||
fs afero.Fs
|
||||
}
|
||||
|
||||
// GetIdentity panics on non-Darwin platforms
|
||||
var errKeychainNotSupported = errors.New(
|
||||
"keychain unlockers are only supported on macOS")
|
||||
|
||||
// NewKeychainUnlocker creates a stub KeychainUnlocker on non-Darwin
|
||||
// platforms. The returned instance's methods that require macOS
|
||||
// functionality will return errors.
|
||||
func NewKeychainUnlocker(
|
||||
fs afero.Fs, directory string, metadata UnlockerMetadata,
|
||||
) *KeychainUnlocker {
|
||||
return &KeychainUnlocker{
|
||||
Directory: directory,
|
||||
Metadata: metadata,
|
||||
fs: fs,
|
||||
}
|
||||
}
|
||||
|
||||
// GetIdentity returns an error on non-Darwin platforms
|
||||
func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
return nil, errKeychainNotSupported
|
||||
}
|
||||
|
||||
// GetType panics on non-Darwin platforms
|
||||
// GetType returns the unlocker type
|
||||
func (k *KeychainUnlocker) GetType() string {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
return "keychain"
|
||||
}
|
||||
|
||||
// GetMetadata panics on non-Darwin platforms
|
||||
// GetMetadata returns the unlocker metadata
|
||||
func (k *KeychainUnlocker) GetMetadata() UnlockerMetadata {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
return k.Metadata
|
||||
}
|
||||
|
||||
// GetDirectory panics on non-Darwin platforms
|
||||
// GetDirectory returns the unlocker directory
|
||||
func (k *KeychainUnlocker) GetDirectory() string {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
return k.Directory
|
||||
}
|
||||
|
||||
// GetID returns the unlocker ID
|
||||
func (k *KeychainUnlocker) GetID() string {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
return k.Metadata.CreatedAt.Format("2006-01-02.15.04") + "-keychain"
|
||||
}
|
||||
|
||||
// GetKeychainItemName panics on non-Darwin platforms
|
||||
// GetKeychainItemName returns an error on non-Darwin platforms
|
||||
func (k *KeychainUnlocker) GetKeychainItemName() (string, error) {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
return "", errKeychainNotSupported
|
||||
}
|
||||
|
||||
// Remove panics on non-Darwin platforms
|
||||
// Remove returns an error on non-Darwin platforms
|
||||
func (k *KeychainUnlocker) Remove() error {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
return errKeychainNotSupported
|
||||
}
|
||||
|
||||
// NewKeychainUnlocker panics on non-Darwin platforms
|
||||
func NewKeychainUnlocker(fs afero.Fs, directory string, metadata UnlockerMetadata) *KeychainUnlocker {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
}
|
||||
|
||||
// CreateKeychainUnlocker panics on non-Darwin platforms
|
||||
func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, error) {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
}
|
||||
|
||||
// getLongTermPrivateKey panics on non-Darwin platforms
|
||||
func getLongTermPrivateKey(fs afero.Fs, vault VaultInterface) (*memguard.LockedBuffer, error) {
|
||||
panic("keychain unlockers are only supported on macOS")
|
||||
// CreateKeychainUnlocker returns an error on non-Darwin platforms
|
||||
func CreateKeychainUnlocker(_ afero.Fs, _ string) (*KeychainUnlocker, error) {
|
||||
return nil, errKeychainNotSupported
|
||||
}
|
||||
|
||||
@@ -13,29 +13,134 @@ import (
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
func TestPassphraseUnlockerWithRealFS(t *testing.T) {
|
||||
// This test uses real filesystem
|
||||
if os.Getenv("CI") == "true" {
|
||||
t.Log("Running in CI environment with real filesystem")
|
||||
}
|
||||
// testMnemonic is the standard BIP39 test vector mnemonic.
|
||||
//
|
||||
//nolint:dupword // BIP39 test mnemonic repeats words by design
|
||||
const testMnemonic = "abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon abandon abandon about"
|
||||
|
||||
// Create a temporary directory for our tests
|
||||
tempDir, err := os.MkdirTemp("", "secret-passphrase-test-")
|
||||
// writeTestPublicKey writes the unlocker public key and verifies it exists.
|
||||
func writeTestPublicKey(
|
||||
t *testing.T, fs afero.Fs, unlockerDir string, agePublicKey string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
pubKeyPath := filepath.Join(unlockerDir, "pub.age")
|
||||
|
||||
err := afero.WriteFile(fs, pubKeyPath, []byte(agePublicKey), secret.FilePerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create temp dir: %v", err)
|
||||
t.Fatalf("Failed to write public key: %v", err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir) // Clean up after test
|
||||
|
||||
// Use the real filesystem
|
||||
fs := afero.NewOsFs()
|
||||
// Verify the file exists
|
||||
exists, err := afero.Exists(fs, pubKeyPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check if public key exists: %v", err)
|
||||
}
|
||||
|
||||
// Test data
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
testPassphrase := "test-passphrase-123"
|
||||
if !exists {
|
||||
t.Errorf("Public key file should exist at %s", pubKeyPath)
|
||||
}
|
||||
}
|
||||
|
||||
// Create the directory structure
|
||||
unlockerDir := filepath.Join(tempDir, "unlocker")
|
||||
if err := os.MkdirAll(unlockerDir, secret.DirPerms); err != nil {
|
||||
// writeTestPrivateKey encrypts the private key with the passphrase,
|
||||
// writes it, and verifies it exists.
|
||||
func writeTestPrivateKey(
|
||||
t *testing.T,
|
||||
fs afero.Fs,
|
||||
unlockerDir string,
|
||||
agePrivateKey string,
|
||||
testPassphrase string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
privKeyBuffer := memguard.NewBufferFromBytes([]byte(agePrivateKey))
|
||||
defer privKeyBuffer.Destroy()
|
||||
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte(testPassphrase))
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
encryptedPrivKey, err := secret.EncryptWithPassphrase(
|
||||
privKeyBuffer, passphraseBuffer)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to encrypt private key: %v", err)
|
||||
}
|
||||
|
||||
privKeyPath := filepath.Join(unlockerDir, "priv.age")
|
||||
|
||||
err = afero.WriteFile(fs, privKeyPath, encryptedPrivKey, secret.FilePerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to write encrypted private key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the file exists
|
||||
exists, err := afero.Exists(fs, privKeyPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check if private key exists: %v", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
t.Errorf("Encrypted private key file should exist at %s", privKeyPath)
|
||||
}
|
||||
}
|
||||
|
||||
// writeTestLongTermKey encrypts the derived long-term key to the
|
||||
// unlocker's recipient, writes it, and verifies it exists.
|
||||
func writeTestLongTermKey(
|
||||
t *testing.T, fs afero.Fs, unlockerDir string, agePublicKey string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
// Derive a long-term identity from the test mnemonic
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term identity: %v", err)
|
||||
}
|
||||
|
||||
// Encrypt long-term private key to the unlocker's recipient
|
||||
recipient, err := age.ParseX25519Recipient(agePublicKey)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to parse recipient: %v", err)
|
||||
}
|
||||
|
||||
ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String()))
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, recipient)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to encrypt long-term private key: %v", err)
|
||||
}
|
||||
|
||||
ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age")
|
||||
|
||||
err = afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to write encrypted long-term private key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the file exists
|
||||
exists, err := afero.Exists(fs, ltPrivKeyPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check if long-term key exists: %v", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
t.Errorf("Encrypted long-term key file should exist at %s", ltPrivKeyPath)
|
||||
}
|
||||
}
|
||||
|
||||
// newTestPassphraseUnlocker creates a temp unlocker directory and a
|
||||
// passphrase unlocker with a fresh age identity for testing.
|
||||
func newTestPassphraseUnlocker(
|
||||
t *testing.T, fs afero.Fs,
|
||||
) (*secret.PassphraseUnlocker, *age.X25519Identity, string) {
|
||||
t.Helper()
|
||||
|
||||
// Create the directory structure in a temp dir
|
||||
unlockerDir := filepath.Join(t.TempDir(), "unlocker")
|
||||
|
||||
err := os.MkdirAll(unlockerDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create unlocker directory: %v", err)
|
||||
}
|
||||
|
||||
@@ -54,86 +159,40 @@ func TestPassphraseUnlockerWithRealFS(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to generate age identity: %v", err)
|
||||
}
|
||||
|
||||
return unlocker, ageIdentity, unlockerDir
|
||||
}
|
||||
|
||||
//nolint:paralleltest // subtests share real-FS state and t.Setenv, order matters
|
||||
func TestPassphraseUnlockerWithRealFS(t *testing.T) {
|
||||
// This test uses real filesystem
|
||||
if os.Getenv("CI") == "true" {
|
||||
t.Log("Running in CI environment with real filesystem")
|
||||
}
|
||||
|
||||
// Use the real filesystem
|
||||
fs := afero.NewOsFs()
|
||||
|
||||
// Test data
|
||||
testPassphrase := "test-passphrase-123"
|
||||
|
||||
unlocker, ageIdentity, unlockerDir := newTestPassphraseUnlocker(t, fs)
|
||||
agePrivateKey := ageIdentity.String()
|
||||
agePublicKey := ageIdentity.Recipient().String()
|
||||
|
||||
// Test writing public key
|
||||
t.Run("WritePublicKey", func(t *testing.T) {
|
||||
pubKeyPath := filepath.Join(unlockerDir, "pub.age")
|
||||
if err := afero.WriteFile(fs, pubKeyPath, []byte(agePublicKey), secret.FilePerms); err != nil {
|
||||
t.Fatalf("Failed to write public key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the file exists
|
||||
exists, err := afero.Exists(fs, pubKeyPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check if public key exists: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Errorf("Public key file should exist at %s", pubKeyPath)
|
||||
}
|
||||
writeTestPublicKey(t, fs, unlockerDir, agePublicKey)
|
||||
})
|
||||
|
||||
// Test encrypting private key with passphrase
|
||||
t.Run("EncryptPrivateKey", func(t *testing.T) {
|
||||
privKeyBuffer := memguard.NewBufferFromBytes([]byte(agePrivateKey))
|
||||
defer privKeyBuffer.Destroy()
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte(testPassphrase))
|
||||
defer passphraseBuffer.Destroy()
|
||||
encryptedPrivKey, err := secret.EncryptWithPassphrase(privKeyBuffer, passphraseBuffer)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to encrypt private key: %v", err)
|
||||
}
|
||||
|
||||
privKeyPath := filepath.Join(unlockerDir, "priv.age")
|
||||
if err := afero.WriteFile(fs, privKeyPath, encryptedPrivKey, secret.FilePerms); err != nil {
|
||||
t.Fatalf("Failed to write encrypted private key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the file exists
|
||||
exists, err := afero.Exists(fs, privKeyPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check if private key exists: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Errorf("Encrypted private key file should exist at %s", privKeyPath)
|
||||
}
|
||||
writeTestPrivateKey(t, fs, unlockerDir, agePrivateKey, testPassphrase)
|
||||
})
|
||||
|
||||
// Test writing long-term key
|
||||
t.Run("WriteLongTermKey", func(t *testing.T) {
|
||||
// Derive a long-term identity from the test mnemonic
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term identity: %v", err)
|
||||
}
|
||||
|
||||
// Encrypt long-term private key to the unlocker's recipient
|
||||
recipient, err := age.ParseX25519Recipient(agePublicKey)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to parse recipient: %v", err)
|
||||
}
|
||||
|
||||
ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String()))
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, recipient)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to encrypt long-term private key: %v", err)
|
||||
}
|
||||
|
||||
ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age")
|
||||
if err := afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms); err != nil {
|
||||
t.Fatalf("Failed to write encrypted long-term private key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the file exists
|
||||
exists, err := afero.Exists(fs, ltPrivKeyPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check if long-term key exists: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
t.Errorf("Encrypted long-term key file should exist at %s", ltPrivKeyPath)
|
||||
}
|
||||
writeTestLongTermKey(t, fs, unlockerDir, agePublicKey)
|
||||
})
|
||||
|
||||
// Set test environment variable (cleaned up automatically)
|
||||
@@ -148,18 +207,21 @@ func TestPassphraseUnlockerWithRealFS(t *testing.T) {
|
||||
|
||||
// Verify the identity matches what we expect
|
||||
expectedPubKey := ageIdentity.Recipient().String()
|
||||
|
||||
actualPubKey := identity.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Public key mismatch. Expected %s, got %s", expectedPubKey, actualPubKey)
|
||||
t.Errorf("Public key mismatch. Expected %s, got %s",
|
||||
expectedPubKey, actualPubKey)
|
||||
}
|
||||
})
|
||||
|
||||
// Unset the environment variable to test interactive prompt
|
||||
os.Unsetenv(secret.EnvUnlockPassphrase)
|
||||
_ = os.Unsetenv(secret.EnvUnlockPassphrase)
|
||||
|
||||
// Test getting identity from prompt (this would require mocking the prompt)
|
||||
// For real integration tests, we'd need to provide a way to mock the passphrase input
|
||||
// Here we'll just verify the error is what we expect when no passphrase is available
|
||||
// Test getting identity from prompt (this would require mocking the
|
||||
// prompt). For real integration tests, we'd need a way to mock the
|
||||
// passphrase input. Here we just verify the error is what we expect
|
||||
// when no passphrase is available.
|
||||
t.Run("GetIdentityWithoutEnv", func(t *testing.T) {
|
||||
// This should fail since we're not in an interactive terminal
|
||||
_, err := unlocker.GetIdentity()
|
||||
@@ -180,6 +242,7 @@ func TestPassphraseUnlockerWithRealFS(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check if unlocker directory exists: %v", err)
|
||||
}
|
||||
|
||||
if exists {
|
||||
t.Errorf("Unlocker directory should not exist after removal")
|
||||
}
|
||||
|
||||
@@ -19,37 +19,15 @@ type PassphraseUnlocker struct {
|
||||
Passphrase *memguard.LockedBuffer // Secure buffer for passphrase
|
||||
}
|
||||
|
||||
// getPassphrase retrieves the passphrase from memory, environment, or user input
|
||||
// Returns a LockedBuffer for secure memory handling
|
||||
func (p *PassphraseUnlocker) getPassphrase() (*memguard.LockedBuffer, error) {
|
||||
// First check if we already have the passphrase
|
||||
if p.Passphrase != nil && p.Passphrase.IsAlive() {
|
||||
Debug("Using in-memory passphrase", "unlocker_id", p.GetID())
|
||||
// Return a copy of the passphrase buffer
|
||||
return memguard.NewBufferFromBytes(p.Passphrase.Bytes()), nil
|
||||
// NewPassphraseUnlocker creates a new PassphraseUnlocker instance
|
||||
func NewPassphraseUnlocker(
|
||||
fs afero.Fs, directory string, metadata UnlockerMetadata,
|
||||
) *PassphraseUnlocker {
|
||||
return &PassphraseUnlocker{
|
||||
Directory: directory,
|
||||
Metadata: metadata,
|
||||
fs: fs,
|
||||
}
|
||||
|
||||
Debug("No passphrase in memory, checking environment")
|
||||
// Check environment variable for passphrase
|
||||
passphraseStr := os.Getenv(EnvUnlockPassphrase)
|
||||
if passphraseStr != "" {
|
||||
Debug("Using passphrase from environment", "unlocker_id", p.GetID())
|
||||
// Convert to secure buffer
|
||||
secureBuffer := memguard.NewBufferFromBytes([]byte(passphraseStr))
|
||||
|
||||
return secureBuffer, nil
|
||||
}
|
||||
|
||||
Debug("No passphrase in environment, prompting user")
|
||||
// Prompt for passphrase
|
||||
secureBuffer, err := ReadPassphrase("Enter unlock passphrase: ")
|
||||
if err != nil {
|
||||
Debug("Failed to read passphrase", "error", err, "unlocker_id", p.GetID())
|
||||
|
||||
return nil, fmt.Errorf("failed to read passphrase: %w", err)
|
||||
}
|
||||
|
||||
return secureBuffer, nil
|
||||
}
|
||||
|
||||
// GetIdentity implements Unlocker interface for passphrase-based unlockers
|
||||
@@ -71,7 +49,8 @@ func (p *PassphraseUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
|
||||
encryptedPrivKeyData, err := afero.ReadFile(p.fs, unlockerPrivPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read passphrase unlocker private key", "error", err, "path", unlockerPrivPath)
|
||||
Debug("Failed to read passphrase unlocker private key",
|
||||
"error", err, "path", unlockerPrivPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read unlocker private key: %w", err)
|
||||
}
|
||||
@@ -86,7 +65,8 @@ func (p *PassphraseUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
// Decrypt the unlocker private key with passphrase
|
||||
privKeyBuffer, err := DecryptWithPassphrase(encryptedPrivKeyData, passphraseBuffer)
|
||||
if err != nil {
|
||||
Debug("Failed to decrypt unlocker private key", "error", err, "unlocker_id", p.GetID())
|
||||
Debug("Failed to decrypt unlocker private key",
|
||||
"error", err, "unlocker_id", p.GetID())
|
||||
|
||||
return nil, fmt.Errorf("failed to decrypt unlocker private key: %w", err)
|
||||
}
|
||||
@@ -135,7 +115,7 @@ func (p *PassphraseUnlocker) GetID() string {
|
||||
// Generate ID using creation timestamp: YYYY-MM-DD.HH.mm-passphrase
|
||||
createdAt := p.Metadata.CreatedAt
|
||||
|
||||
return fmt.Sprintf("%s-passphrase", createdAt.Format("2006-01-02.15.04"))
|
||||
return createdAt.Format("2006-01-02.15.04") + "-passphrase"
|
||||
}
|
||||
|
||||
// Remove implements Unlocker interface - removes the passphrase unlocker
|
||||
@@ -147,20 +127,45 @@ func (p *PassphraseUnlocker) Remove() error {
|
||||
|
||||
// For passphrase unlockers, we just need to remove the directory
|
||||
// No external resources (like keychain items) to clean up
|
||||
if err := p.fs.RemoveAll(p.Directory); err != nil {
|
||||
err := RemoveDirAtomic(p.fs, p.Directory)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to remove passphrase unlocker directory: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// NewPassphraseUnlocker creates a new PassphraseUnlocker instance
|
||||
func NewPassphraseUnlocker(fs afero.Fs, directory string, metadata UnlockerMetadata) *PassphraseUnlocker {
|
||||
return &PassphraseUnlocker{
|
||||
Directory: directory,
|
||||
Metadata: metadata,
|
||||
fs: fs,
|
||||
// getPassphrase retrieves the passphrase from memory, environment, or
|
||||
// user input. Returns a LockedBuffer for secure memory handling
|
||||
func (p *PassphraseUnlocker) getPassphrase() (*memguard.LockedBuffer, error) {
|
||||
// First check if we already have the passphrase
|
||||
if p.Passphrase != nil && p.Passphrase.IsAlive() {
|
||||
Debug("Using in-memory passphrase", "unlocker_id", p.GetID())
|
||||
// Return a copy of the passphrase buffer
|
||||
return memguard.NewBufferFromBytes(p.Passphrase.Bytes()), nil
|
||||
}
|
||||
|
||||
Debug("No passphrase in memory, checking environment")
|
||||
// Check environment variable for passphrase
|
||||
passphraseStr := os.Getenv(EnvUnlockPassphrase)
|
||||
if passphraseStr != "" {
|
||||
Debug("Using passphrase from environment", "unlocker_id", p.GetID())
|
||||
// Convert to secure buffer
|
||||
secureBuffer := memguard.NewBufferFromBytes([]byte(passphraseStr))
|
||||
|
||||
return secureBuffer, nil
|
||||
}
|
||||
|
||||
Debug("No passphrase in environment, prompting user")
|
||||
// Prompt for passphrase
|
||||
secureBuffer, err := ReadPassphrase("Enter unlock passphrase: ")
|
||||
if err != nil {
|
||||
Debug("Failed to read passphrase", "error", err, "unlocker_id", p.GetID())
|
||||
|
||||
return nil, fmt.Errorf("failed to read passphrase: %w", err)
|
||||
}
|
||||
|
||||
return secureBuffer, nil
|
||||
}
|
||||
|
||||
// CreatePassphraseUnlocker creates a new passphrase-protected unlocker
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
//go:build darwin
|
||||
|
||||
package secret_test
|
||||
|
||||
import (
|
||||
@@ -140,7 +142,7 @@ func TestPGPUnlockerWithRealFS(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create temp dir: %v", err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir) // Clean up after test
|
||||
defer func() { _ = os.RemoveAll(tempDir) }() // Clean up after test
|
||||
|
||||
// Create a temporary GNUPGHOME
|
||||
gnupgHomeDir := filepath.Join(tempDir, "gnupg")
|
||||
@@ -288,7 +290,7 @@ Passphrase: ` + testPassphrase + `
|
||||
}
|
||||
|
||||
// Now create a PGP unlock key (this will use our custom GPGEncryptFunc)
|
||||
pgpUnlocker, err := secret.CreatePGPUnlocker(fs, stateDir, keyID)
|
||||
pgpUnlocker, err := secret.CreatePGPUnlocker(fs, stateDir, keyID, fingerprint)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create PGP unlock key: %v", err)
|
||||
}
|
||||
|
||||
+188
-97
@@ -1,7 +1,9 @@
|
||||
package secret
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
@@ -16,17 +18,28 @@ import (
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
var (
|
||||
errGPGKeyIDEmpty = errors.New("GPG key ID cannot be empty")
|
||||
errInvalidGPGKeyID = errors.New("invalid GPG key ID format")
|
||||
errNoGPGFingerprint = errors.New("could not find fingerprint for GPG key")
|
||||
errNilDataBuffer = errors.New("data buffer is nil")
|
||||
)
|
||||
|
||||
// Variables to allow overriding in tests
|
||||
var (
|
||||
// GPGEncryptFunc is the function used for GPG encryption
|
||||
// Can be overridden in tests to provide a non-interactive implementation
|
||||
//nolint:gochecknoglobals // Required for test mocking
|
||||
GPGEncryptFunc func(data *memguard.LockedBuffer, keyID string) ([]byte, error) = gpgEncryptDefault
|
||||
GPGEncryptFunc func(
|
||||
data *memguard.LockedBuffer, keyID string,
|
||||
) ([]byte, error) = gpgEncryptDefault
|
||||
|
||||
// GPGDecryptFunc is the function used for GPG decryption
|
||||
// Can be overridden in tests to provide a non-interactive implementation
|
||||
//nolint:gochecknoglobals // Required for test mocking
|
||||
GPGDecryptFunc func(encryptedData []byte) (*memguard.LockedBuffer, error) = gpgDecryptDefault
|
||||
GPGDecryptFunc func(
|
||||
encryptedData []byte,
|
||||
) (*memguard.LockedBuffer, error) = gpgDecryptDefault
|
||||
|
||||
// gpgKeyIDRegex validates GPG key IDs
|
||||
// Allows either:
|
||||
@@ -45,6 +58,7 @@ var (
|
||||
// PGPUnlockerMetadata extends UnlockerMetadata with PGP-specific data
|
||||
type PGPUnlockerMetadata struct {
|
||||
UnlockerMetadata
|
||||
|
||||
// GPG key ID used for encryption
|
||||
GPGKeyID string `json:"gpgKeyId"`
|
||||
}
|
||||
@@ -56,6 +70,17 @@ type PGPUnlocker struct {
|
||||
fs afero.Fs
|
||||
}
|
||||
|
||||
// NewPGPUnlocker creates a new PGPUnlocker instance
|
||||
func NewPGPUnlocker(
|
||||
fs afero.Fs, directory string, metadata UnlockerMetadata,
|
||||
) *PGPUnlocker {
|
||||
return &PGPUnlocker{
|
||||
Directory: directory,
|
||||
Metadata: metadata,
|
||||
fs: fs,
|
||||
}
|
||||
}
|
||||
|
||||
// GetIdentity implements Unlocker interface for PGP-based unlockers
|
||||
func (p *PGPUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
DebugWith("Getting PGP unlocker identity",
|
||||
@@ -69,7 +94,8 @@ func (p *PGPUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
|
||||
encryptedAgePrivKeyData, err := afero.ReadFile(p.fs, agePrivKeyPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read PGP-encrypted age private key", "error", err, "path", agePrivKeyPath)
|
||||
Debug("Failed to read PGP-encrypted age private key",
|
||||
"error", err, "path", agePrivKeyPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read encrypted age private key: %w", err)
|
||||
}
|
||||
@@ -81,9 +107,11 @@ func (p *PGPUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
|
||||
// Step 2: Decrypt the age private key using GPG
|
||||
Debug("Decrypting age private key with GPG", "unlocker_id", p.GetID())
|
||||
|
||||
agePrivKeyBuffer, err := GPGDecryptFunc(encryptedAgePrivKeyData)
|
||||
if err != nil {
|
||||
Debug("Failed to decrypt age private key with GPG", "error", err, "unlocker_id", p.GetID())
|
||||
Debug("Failed to decrypt age private key with GPG",
|
||||
"error", err, "unlocker_id", p.GetID())
|
||||
|
||||
return nil, fmt.Errorf("failed to decrypt age private key with GPG: %w", err)
|
||||
}
|
||||
@@ -96,6 +124,7 @@ func (p *PGPUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
|
||||
// Step 3: Parse the decrypted age private key
|
||||
Debug("Parsing decrypted age private key", "unlocker_id", p.GetID())
|
||||
|
||||
ageIdentity, err := age.ParseX25519Identity(agePrivKeyBuffer.String())
|
||||
if err != nil {
|
||||
Debug("Failed to parse age private key", "error", err, "unlocker_id", p.GetID())
|
||||
@@ -126,57 +155,61 @@ func (p *PGPUnlocker) GetDirectory() string {
|
||||
return p.Directory
|
||||
}
|
||||
|
||||
// GetID implements Unlocker interface - generates ID from GPG key ID
|
||||
// GetID implements Unlocker interface - generates ID from GPG key ID.
|
||||
// If the metadata has no usable GPG key ID, it warns with the unlocker's
|
||||
// directory and returns "pgp-unknown", so listing the other unlockers
|
||||
// still works.
|
||||
func (p *PGPUnlocker) GetID() string {
|
||||
// Generate ID using GPG key ID: pgp-<keyid>
|
||||
gpgKeyID, err := p.GetGPGKeyID()
|
||||
if err != nil {
|
||||
// The vault metadata is corrupt - this is a fatal error
|
||||
// We cannot continue with a fallback ID as that would mask data corruption
|
||||
panic(fmt.Sprintf("PGP unlocker metadata is corrupt or missing GPG key ID: %v", err))
|
||||
Warn("PGP unlocker metadata is corrupt or missing its GPG key ID",
|
||||
"directory", p.Directory, "error", err)
|
||||
|
||||
return "pgp-unknown"
|
||||
}
|
||||
|
||||
return fmt.Sprintf("pgp-%s", gpgKeyID)
|
||||
return "pgp-" + gpgKeyID
|
||||
}
|
||||
|
||||
// Remove implements Unlocker interface - removes the PGP unlocker
|
||||
func (p *PGPUnlocker) Remove() error {
|
||||
// For PGP unlockers, we just need to remove the directory
|
||||
// No external resources (like keychain items) to clean up
|
||||
if err := p.fs.RemoveAll(p.Directory); err != nil {
|
||||
err := RemoveDirAtomic(p.fs, p.Directory)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to remove PGP unlocker directory: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// NewPGPUnlocker creates a new PGPUnlocker instance
|
||||
func NewPGPUnlocker(fs afero.Fs, directory string, metadata UnlockerMetadata) *PGPUnlocker {
|
||||
return &PGPUnlocker{
|
||||
Directory: directory,
|
||||
Metadata: metadata,
|
||||
fs: fs,
|
||||
}
|
||||
}
|
||||
|
||||
// GetGPGKeyID returns the GPG key ID from metadata
|
||||
func (p *PGPUnlocker) GetGPGKeyID() (string, error) {
|
||||
// Load the metadata
|
||||
metadataPath := filepath.Join(p.Directory, "unlocker-metadata.json")
|
||||
|
||||
metadataData, err := afero.ReadFile(p.fs, metadataPath)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to read PGP metadata: %w", err)
|
||||
}
|
||||
|
||||
var pgpMetadata PGPUnlockerMetadata
|
||||
if err := json.Unmarshal(metadataData, &pgpMetadata); err != nil {
|
||||
|
||||
err = json.Unmarshal(metadataData, &pgpMetadata)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to parse PGP metadata: %w", err)
|
||||
}
|
||||
|
||||
if pgpMetadata.GPGKeyID == "" {
|
||||
return "", fmt.Errorf("PGP metadata: %w", errGPGKeyIDEmpty)
|
||||
}
|
||||
|
||||
return pgpMetadata.GPGKeyID, nil
|
||||
}
|
||||
|
||||
// generatePGPUnlockerName generates a unique name for the PGP unlocker based on hostname and date
|
||||
// generatePGPUnlockerName generates a unique name for the PGP unlocker
|
||||
// based on hostname and date
|
||||
func generatePGPUnlockerName() (string, error) {
|
||||
hostname, err := os.Hostname()
|
||||
if err != nil {
|
||||
@@ -189,34 +222,50 @@ func generatePGPUnlockerName() (string, error) {
|
||||
return fmt.Sprintf("%s-pgp-%s", hostname, enrollmentDate), nil
|
||||
}
|
||||
|
||||
// CreatePGPUnlocker creates a new PGP unlocker and stores it in the vault
|
||||
func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnlocker, error) {
|
||||
// Check if GPG is available
|
||||
if err := checkGPGAvailable(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// pgpUnlockerDir returns the current vault and the directory in it for a
|
||||
// new PGP unlocker, named after the host and the day.
|
||||
//
|
||||
//nolint:ireturn // the vault is only available behind VaultInterface
|
||||
func pgpUnlockerDir(
|
||||
fs afero.Fs, stateDir string,
|
||||
) (VaultInterface, string, error) {
|
||||
// Get current vault
|
||||
vault, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get current vault: %w", err)
|
||||
return nil, "", fmt.Errorf("failed to get current vault: %w", err)
|
||||
}
|
||||
|
||||
// Generate the unlocker name based on hostname and date
|
||||
unlockerName, err := generatePGPUnlockerName()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate unlocker name: %w", err)
|
||||
return nil, "", fmt.Errorf("failed to generate unlocker name: %w", err)
|
||||
}
|
||||
|
||||
// Create unlocker directory using the generated name
|
||||
vaultDir, err := vault.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
return nil, "", fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerName)
|
||||
if err := fs.MkdirAll(unlockerDir, DirPerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to create unlocker directory: %w", err)
|
||||
return vault, filepath.Join(vaultDir, "unlockers.d", unlockerName), nil
|
||||
}
|
||||
|
||||
// CreatePGPUnlocker creates a new PGP unlocker and stores it in the vault.
|
||||
// It encrypts to the GPG key gpgKeyID and records fingerprint, that key's
|
||||
// fingerprint as ResolveGPGKeyFingerprint returns it, in the metadata.
|
||||
// Everything that can fail short of writing a file is done before anything
|
||||
// is written, and the files are written through WriteDir, so a failure
|
||||
// leaves no partial unlocker.
|
||||
func CreatePGPUnlocker(
|
||||
fs afero.Fs, stateDir, gpgKeyID, fingerprint string,
|
||||
) (*PGPUnlocker, error) {
|
||||
err := checkGPGAvailable()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
vault, unlockerDir, err := pgpUnlockerDir(fs, stateDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Step 1: Generate a new age keypair for the PGP unlocker
|
||||
@@ -225,54 +274,14 @@ func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnloc
|
||||
return nil, fmt.Errorf("failed to generate age keypair: %w", err)
|
||||
}
|
||||
|
||||
// Step 2: Store age recipient as plaintext
|
||||
ageRecipient := ageIdentity.Recipient().String()
|
||||
recipientPath := filepath.Join(unlockerDir, "pub.txt")
|
||||
if err := afero.WriteFile(fs, recipientPath, []byte(ageRecipient), FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write age recipient: %w", err)
|
||||
}
|
||||
|
||||
// Step 3: Get or derive the long-term private key
|
||||
ltPrivKeyData, err := getLongTermPrivateKey(fs, vault)
|
||||
// Step 2: Encrypt the long-term private key to the new keypair, and the
|
||||
// keypair's private key to the GPG key
|
||||
encryptedLtPrivKey, encryptedAgePrivKey, err := encryptPGPUnlockerKeys(
|
||||
vault, ageIdentity, gpgKeyID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer ltPrivKeyData.Destroy()
|
||||
|
||||
// Step 7: Encrypt long-term private key to the new age unlocker
|
||||
encryptedLtPrivKeyToAge, err := EncryptToRecipient(ltPrivKeyData, ageIdentity.Recipient())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt long-term private key to age unlocker: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted long-term private key
|
||||
ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age")
|
||||
if err := afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKeyToAge, FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
// Step 8: Encrypt age private key to the GPG key ID
|
||||
// Use memguard to protect the private key in memory
|
||||
agePrivateKeyBuffer := memguard.NewBufferFromBytes([]byte(ageIdentity.String()))
|
||||
defer agePrivateKeyBuffer.Destroy()
|
||||
|
||||
encryptedAgePrivKey, err := GPGEncryptFunc(agePrivateKeyBuffer, gpgKeyID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt age private key with GPG: %w", err)
|
||||
}
|
||||
|
||||
agePrivKeyPath := filepath.Join(unlockerDir, "priv.age.gpg")
|
||||
if err := afero.WriteFile(fs, agePrivKeyPath, encryptedAgePrivKey, FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write encrypted age private key: %w", err)
|
||||
}
|
||||
|
||||
// Step 9: Resolve the GPG key ID to its full fingerprint
|
||||
fingerprint, err := ResolveGPGKeyFingerprint(gpgKeyID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to resolve GPG key fingerprint: %w", err)
|
||||
}
|
||||
|
||||
// Step 10: Create and write enhanced metadata with full fingerprint
|
||||
pgpMetadata := PGPUnlockerMetadata{
|
||||
UnlockerMetadata: UnlockerMetadata{
|
||||
Type: "pgp",
|
||||
@@ -287,10 +296,13 @@ func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnloc
|
||||
return nil, fmt.Errorf("failed to marshal unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
if err := afero.WriteFile(fs,
|
||||
filepath.Join(unlockerDir, "unlocker-metadata.json"),
|
||||
metadataBytes, FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write unlocker metadata: %w", err)
|
||||
// Step 3: Write the unlocker's files, the metadata last
|
||||
err = WriteDir(fs, unlockerDir, func(dir string) error {
|
||||
return writePGPUnlockerFiles(fs, dir, ageIdentity.Recipient(),
|
||||
encryptedLtPrivKey, encryptedAgePrivKey, metadataBytes)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &PGPUnlocker{
|
||||
@@ -300,14 +312,80 @@ func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnloc
|
||||
}, nil
|
||||
}
|
||||
|
||||
// encryptPGPUnlockerKeys returns the vault's long-term private key encrypted
|
||||
// to the new PGP unlocker's age keypair, and that keypair's private key
|
||||
// encrypted to the GPG key gpgKeyID.
|
||||
func encryptPGPUnlockerKeys(
|
||||
vault VaultInterface, ageIdentity *age.X25519Identity, gpgKeyID string,
|
||||
) ([]byte, []byte, error) {
|
||||
// From the mnemonic or the current unlocker, as for a passphrase unlocker
|
||||
ltIdentity, err := vault.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
|
||||
ltPrivKeyData := memguard.NewBufferFromBytes([]byte(ltIdentity.String()))
|
||||
defer ltPrivKeyData.Destroy()
|
||||
|
||||
encryptedLtPrivKey, err := EncryptToRecipient(
|
||||
ltPrivKeyData, ageIdentity.Recipient())
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf(
|
||||
"failed to encrypt long-term private key to age unlocker: %w", err)
|
||||
}
|
||||
|
||||
// Use memguard to protect the private key in memory
|
||||
agePrivateKeyBuffer := memguard.NewBufferFromBytes([]byte(ageIdentity.String()))
|
||||
defer agePrivateKeyBuffer.Destroy()
|
||||
|
||||
encryptedAgePrivKey, err := GPGEncryptFunc(agePrivateKeyBuffer, gpgKeyID)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf(
|
||||
"failed to encrypt age private key with GPG: %w", err)
|
||||
}
|
||||
|
||||
return encryptedLtPrivKey, encryptedAgePrivKey, nil
|
||||
}
|
||||
|
||||
// writePGPUnlockerFiles writes the files of a PGP unlocker into dir, the
|
||||
// metadata last.
|
||||
func writePGPUnlockerFiles(
|
||||
fs afero.Fs, dir string, ageRecipient *age.X25519Recipient,
|
||||
encryptedLtPrivKey, encryptedAgePrivKey, metadataBytes []byte,
|
||||
) error {
|
||||
err := WriteFileAtomic(fs, filepath.Join(dir, "pub.txt"),
|
||||
[]byte(ageRecipient.String()))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write age recipient: %w", err)
|
||||
}
|
||||
|
||||
err = WriteFileAtomic(fs, filepath.Join(dir, "longterm.age"), encryptedLtPrivKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
err = WriteFileAtomic(fs, filepath.Join(dir, "priv.age.gpg"), encryptedAgePrivKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write encrypted age private key: %w", err)
|
||||
}
|
||||
|
||||
err = WriteFileAtomic(fs,
|
||||
filepath.Join(dir, "unlocker-metadata.json"), metadataBytes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateGPGKeyID validates that a GPG key ID is safe for command execution
|
||||
func validateGPGKeyID(keyID string) error {
|
||||
if keyID == "" {
|
||||
return fmt.Errorf("GPG key ID cannot be empty")
|
||||
return errGPGKeyIDEmpty
|
||||
}
|
||||
|
||||
if !gpgKeyIDRegex.MatchString(keyID) {
|
||||
return fmt.Errorf("invalid GPG key ID format: %s", keyID)
|
||||
return fmt.Errorf("%w: %s", errInvalidGPGKeyID, keyID)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -315,20 +393,24 @@ func validateGPGKeyID(keyID string) error {
|
||||
|
||||
// ResolveGPGKeyFingerprint resolves any GPG key identifier to its full fingerprint
|
||||
func ResolveGPGKeyFingerprint(keyID string) (string, error) {
|
||||
if err := validateGPGKeyID(keyID); err != nil {
|
||||
err := validateGPGKeyID(keyID)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("invalid GPG key ID: %w", err)
|
||||
}
|
||||
|
||||
// Use GPG to get the full fingerprint for the key
|
||||
cmd := exec.Command("gpg", "--list-keys", "--with-colons", "--fingerprint", keyID)
|
||||
cmd := exec.CommandContext( //nolint:gosec // G204: keyID validated above
|
||||
context.Background(),
|
||||
"gpg", "--list-keys", "--with-colons", "--fingerprint", keyID,
|
||||
)
|
||||
|
||||
output, err := cmd.Output()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to resolve GPG key fingerprint: %w", err)
|
||||
}
|
||||
|
||||
// Parse the output to extract the fingerprint
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
for line := range strings.SplitSeq(string(output), "\n") {
|
||||
if strings.HasPrefix(line, "fpr:") {
|
||||
fields := strings.Split(line, ":")
|
||||
if len(fields) >= 10 && fields[9] != "" {
|
||||
@@ -337,14 +419,18 @@ func ResolveGPGKeyFingerprint(keyID string) (string, error) {
|
||||
}
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("could not find fingerprint for GPG key: %s", keyID)
|
||||
return "", fmt.Errorf("%w: %s", errNoGPGFingerprint, keyID)
|
||||
}
|
||||
|
||||
// checkGPGAvailable verifies that GPG is available
|
||||
func checkGPGAvailable() error {
|
||||
cmd := exec.Command("gpg", "--version")
|
||||
if err := cmd.Run(); err != nil {
|
||||
return fmt.Errorf("GPG not available: %w (make sure 'gpg' command is installed and in PATH)", err)
|
||||
cmd := exec.CommandContext(context.Background(), "gpg", "--version")
|
||||
|
||||
err := cmd.Run()
|
||||
if err != nil {
|
||||
return fmt.Errorf(
|
||||
"GPG not available: %w (make sure 'gpg' command is installed and in PATH)",
|
||||
err)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -353,13 +439,18 @@ func checkGPGAvailable() error {
|
||||
// gpgEncryptDefault is the default implementation of GPG encryption
|
||||
func gpgEncryptDefault(data *memguard.LockedBuffer, keyID string) ([]byte, error) {
|
||||
if data == nil {
|
||||
return nil, fmt.Errorf("data buffer is nil")
|
||||
return nil, errNilDataBuffer
|
||||
}
|
||||
if err := validateGPGKeyID(keyID); err != nil {
|
||||
|
||||
err := validateGPGKeyID(keyID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid GPG key ID: %w", err)
|
||||
}
|
||||
|
||||
cmd := exec.Command("gpg", "--trust-model", "always", "--armor", "--encrypt", "-r", keyID)
|
||||
cmd := exec.CommandContext( //nolint:gosec // G204: keyID validated above
|
||||
context.Background(),
|
||||
"gpg", "--trust-model", "always", "--armor", "--encrypt", "-r", keyID,
|
||||
)
|
||||
cmd.Stdin = strings.NewReader(data.String())
|
||||
|
||||
output, err := cmd.Output()
|
||||
@@ -372,7 +463,7 @@ func gpgEncryptDefault(data *memguard.LockedBuffer, keyID string) ([]byte, error
|
||||
|
||||
// gpgDecryptDefault is the default implementation of GPG decryption
|
||||
func gpgDecryptDefault(encryptedData []byte) (*memguard.LockedBuffer, error) {
|
||||
cmd := exec.Command("gpg", "--quiet", "--decrypt")
|
||||
cmd := exec.CommandContext(context.Background(), "gpg", "--quiet", "--decrypt")
|
||||
cmd.Stdin = strings.NewReader(string(encryptedData))
|
||||
|
||||
output, err := cmd.Output()
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
package secret_test
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// The GPG key ID and fingerprint passed to CreatePGPUnlocker.
|
||||
const (
|
||||
testGPGKeyID = "0123456789ABCDEF"
|
||||
testGPGFingerprint = "0123456789ABCDEF0123456789ABCDEF01234567"
|
||||
)
|
||||
|
||||
// fakeGPGScript is a gpg for which `gpg --version` succeeds and anything
|
||||
// else fails.
|
||||
const fakeGPGScript = `#!/bin/sh
|
||||
[ "$*" = --version ]
|
||||
`
|
||||
|
||||
// installFakeGPG makes fakeGPGScript the only gpg on PATH for the test.
|
||||
func installFakeGPG(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
dir := t.TempDir()
|
||||
|
||||
//nolint:gosec // G306: the script must be executable
|
||||
err := os.WriteFile(filepath.Join(dir, "gpg"), []byte(fakeGPGScript), 0o700)
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Setenv("PATH", dir)
|
||||
}
|
||||
|
||||
// TestCreatePGPUnlockerFailureWritesNothing makes CreatePGPUnlocker fail at
|
||||
// getting the vault's long-term key, which used to come after part of the
|
||||
// unlocker was written, and asserts that nothing is written. Getting the key
|
||||
// fails because there is no mnemonic and no current unlocker.
|
||||
func TestCreatePGPUnlockerFailureWritesNothing(t *testing.T) {
|
||||
installFakeGPG(t)
|
||||
t.Setenv(secret.EnvMnemonic, "")
|
||||
|
||||
base := afero.NewMemMapFs()
|
||||
vlt, err := vault.CreateVault(base, testVaultStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
fs := hookFs{Fs: base, before: func(_, path string) error {
|
||||
t.Errorf("changed %s", path)
|
||||
|
||||
return nil
|
||||
}}
|
||||
|
||||
_, err = secret.CreatePGPUnlocker(
|
||||
fs, testVaultStateDir, testGPGKeyID, testGPGFingerprint)
|
||||
require.Error(t, err)
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, dirNames(t, base, filepath.Join(vaultDir, "unlockers.d")))
|
||||
}
|
||||
+163
-101
@@ -2,6 +2,7 @@ package secret
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
@@ -15,6 +16,18 @@ import (
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
var (
|
||||
// errSecretNotFound carries only the message tail; callers compose
|
||||
// "secret <name> not found" around it so the emitted text is
|
||||
// unchanged.
|
||||
errSecretNotFound = errors.New("not found")
|
||||
errUnlockerRequired = errors.New("unlocker required to decrypt secret")
|
||||
errGetEncryptedDataDeprecated = errors.New(
|
||||
"GetEncryptedData is deprecated - use version-specific methods")
|
||||
errGetCurrentVaultNotRegistered = errors.New(
|
||||
"GetCurrentVault function not registered")
|
||||
)
|
||||
|
||||
// VaultInterface defines the interface that vault implementations must satisfy
|
||||
type VaultInterface interface {
|
||||
GetDirectory() (string, error)
|
||||
@@ -22,7 +35,9 @@ type VaultInterface interface {
|
||||
GetName() string
|
||||
GetFilesystem() afero.Fs
|
||||
GetCurrentUnlocker() (Unlocker, error)
|
||||
CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*PassphraseUnlocker, error)
|
||||
GetOrDeriveLongTermKey() (*age.X25519Identity, error)
|
||||
CreatePassphraseUnlocker(
|
||||
passphrase *memguard.LockedBuffer) (*PassphraseUnlocker, error)
|
||||
}
|
||||
|
||||
// Secret represents a secret in a vault
|
||||
@@ -62,7 +77,8 @@ func NewSecret(vault VaultInterface, name string) *Secret {
|
||||
}
|
||||
}
|
||||
|
||||
// GetValue retrieves and decrypts the current version's value using the provided unlocker
|
||||
// GetValue retrieves and decrypts the current version's value using the
|
||||
// provided unlocker
|
||||
func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) {
|
||||
DebugWith("Getting secret value",
|
||||
slog.String("secret_name", s.Name),
|
||||
@@ -72,14 +88,17 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) {
|
||||
// Check if secret exists
|
||||
exists, err := s.Exists()
|
||||
if err != nil {
|
||||
Debug("Failed to check if secret exists during GetValue", "error", err, "secret_name", s.Name)
|
||||
Debug("Failed to check if secret exists during GetValue",
|
||||
"error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
if !exists {
|
||||
Debug("Secret not found during GetValue", "secret_name", s.Name, "vault_name", s.vault.GetName())
|
||||
|
||||
return nil, fmt.Errorf("secret %s not found", s.Name)
|
||||
if !exists {
|
||||
Debug("Secret not found during GetValue",
|
||||
"secret_name", s.Name, "vault_name", s.vault.GetName())
|
||||
|
||||
return nil, fmt.Errorf("secret %s %w", s.Name, errSecretNotFound)
|
||||
}
|
||||
|
||||
Debug("Secret exists, getting current version", "secret_name", s.Name)
|
||||
@@ -95,52 +114,9 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) {
|
||||
// Create version object
|
||||
version := NewVersion(s.vault, s.Name, currentVersion)
|
||||
|
||||
// Check if we have SB_SECRET_MNEMONIC environment variable for direct decryption
|
||||
// Check for SB_SECRET_MNEMONIC environment variable for direct decryption
|
||||
if envMnemonic := os.Getenv(EnvMnemonic); envMnemonic != "" {
|
||||
Debug("Using mnemonic from environment for direct long-term key derivation", "secret_name", s.Name)
|
||||
|
||||
// Get vault directory to read metadata
|
||||
vaultDir, err := s.vault.GetDirectory()
|
||||
if err != nil {
|
||||
Debug("Failed to get vault directory", "error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get the correct derivation index
|
||||
metadataPath := filepath.Join(vaultDir, "vault-metadata.json")
|
||||
metadataBytes, err := afero.ReadFile(s.vault.GetFilesystem(), metadataPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read vault metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read vault metadata: %w", err)
|
||||
}
|
||||
|
||||
var metadata VaultMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
Debug("Failed to parse vault metadata", "error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to parse vault metadata: %w", err)
|
||||
}
|
||||
|
||||
DebugWith("Using vault derivation index from metadata",
|
||||
slog.String("secret_name", s.Name),
|
||||
slog.String("vault_name", s.vault.GetName()),
|
||||
slog.Uint64("derivation_index", uint64(metadata.DerivationIndex)),
|
||||
)
|
||||
|
||||
// Use mnemonic with the vault's derivation index from metadata
|
||||
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
Debug("Failed to derive long-term key from mnemonic for secret", "error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
|
||||
Debug("Successfully derived long-term key from mnemonic", "secret_name", s.Name)
|
||||
|
||||
// Use the long-term key to decrypt the version
|
||||
return version.GetValue(ltIdentity)
|
||||
return s.getValueViaMnemonic(version, envMnemonic)
|
||||
}
|
||||
|
||||
Debug("Using unlocker for vault access", "secret_name", s.Name)
|
||||
@@ -149,51 +125,12 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) {
|
||||
if unlocker == nil {
|
||||
Debug("No unlocker provided for secret decryption", "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("unlocker required to decrypt secret")
|
||||
return nil, errUnlockerRequired
|
||||
}
|
||||
|
||||
DebugWith("Getting vault's long-term key using unlocker",
|
||||
slog.String("secret_name", s.Name),
|
||||
slog.String("unlocker_type", unlocker.GetType()),
|
||||
slog.String("unlocker_id", unlocker.GetID()),
|
||||
)
|
||||
|
||||
// Step 1: Use the unlocker to get the vault's long-term private key
|
||||
unlockIdentity, err := unlocker.GetIdentity()
|
||||
ltIdentity, err := s.getLongTermIdentityFromUnlocker(unlocker)
|
||||
if err != nil {
|
||||
Debug("Failed to get unlocker identity", "error", err, "secret_name", s.Name, "unlocker_type", unlocker.GetType())
|
||||
|
||||
return nil, fmt.Errorf("failed to get unlocker identity: %w", err)
|
||||
}
|
||||
|
||||
// Read the encrypted long-term private key from the unlocker directory
|
||||
encryptedLtPrivKeyPath := filepath.Join(unlocker.GetDirectory(), "longterm.age")
|
||||
Debug("Reading encrypted long-term private key", "path", encryptedLtPrivKeyPath)
|
||||
|
||||
encryptedLtPrivKey, err := afero.ReadFile(s.vault.GetFilesystem(), encryptedLtPrivKeyPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read encrypted long-term private key", "error", err, "path", encryptedLtPrivKeyPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
// Decrypt the encrypted long-term private key using the unlocker
|
||||
Debug("Decrypting long-term private key using unlocker", "secret_name", s.Name)
|
||||
ltPrivKeyBuffer, err := DecryptWithIdentity(encryptedLtPrivKey, unlockIdentity)
|
||||
if err != nil {
|
||||
Debug("Failed to decrypt long-term private key", "error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err)
|
||||
}
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
// Parse the long-term private key
|
||||
Debug("Parsing long-term private key", "secret_name", s.Name)
|
||||
ltIdentity, err := age.ParseX25519Identity(ltPrivKeyBuffer.String())
|
||||
if err != nil {
|
||||
Debug("Failed to parse long-term private key", "error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to parse long-term private key: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
DebugWith("Successfully obtained vault's long-term key",
|
||||
@@ -207,7 +144,8 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) {
|
||||
|
||||
// LoadMetadata is deprecated - metadata is now per-version and encrypted
|
||||
func (s *Secret) LoadMetadata() error {
|
||||
Debug("LoadMetadata called but is deprecated in versioned model", "secret_name", s.Name)
|
||||
Debug("LoadMetadata called but is deprecated in versioned model",
|
||||
"secret_name", s.Name)
|
||||
// For backward compatibility, we'll populate with basic info
|
||||
now := time.Now()
|
||||
s.Metadata = Metadata{
|
||||
@@ -227,9 +165,10 @@ func (s *Secret) GetMetadata() Metadata {
|
||||
|
||||
// GetEncryptedData is deprecated - data is now stored in versions
|
||||
func (s *Secret) GetEncryptedData() ([]byte, error) {
|
||||
Debug("GetEncryptedData called but is deprecated in versioned model", "secret_name", s.Name)
|
||||
Debug("GetEncryptedData called but is deprecated in versioned model",
|
||||
"secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("GetEncryptedData is deprecated - use version-specific methods")
|
||||
return nil, errGetEncryptedDataDeprecated
|
||||
}
|
||||
|
||||
// Exists checks if the secret exists on disk
|
||||
@@ -242,7 +181,8 @@ func (s *Secret) Exists() (bool, error) {
|
||||
// Check if the secret directory exists and has a current symlink
|
||||
exists, err := afero.DirExists(s.vault.GetFilesystem(), s.Directory)
|
||||
if err != nil {
|
||||
Debug("Failed to check secret directory existence", "error", err, "secret_dir", s.Directory)
|
||||
Debug("Failed to check secret directory existence",
|
||||
"error", err, "secret_dir", s.Directory)
|
||||
|
||||
return false, err
|
||||
}
|
||||
@@ -269,14 +209,134 @@ func (s *Secret) Exists() (bool, error) {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// getValueViaMnemonic derives the vault's long-term key from the
|
||||
// mnemonic in the environment and decrypts the version value with it.
|
||||
func (s *Secret) getValueViaMnemonic(
|
||||
version *Version, envMnemonic string,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
Debug("Using mnemonic from environment for direct long-term key derivation",
|
||||
"secret_name", s.Name)
|
||||
|
||||
// Get vault directory to read metadata
|
||||
vaultDir, err := s.vault.GetDirectory()
|
||||
if err != nil {
|
||||
Debug("Failed to get vault directory", "error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get the correct derivation index
|
||||
metadataPath := filepath.Join(vaultDir, "vault-metadata.json")
|
||||
|
||||
metadataBytes, err := afero.ReadFile(s.vault.GetFilesystem(), metadataPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read vault metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read vault metadata: %w", err)
|
||||
}
|
||||
|
||||
var metadata VaultMetadata
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
Debug("Failed to parse vault metadata", "error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to parse vault metadata: %w", err)
|
||||
}
|
||||
|
||||
DebugWith("Using vault derivation index from metadata",
|
||||
slog.String("secret_name", s.Name),
|
||||
slog.String("vault_name", s.vault.GetName()),
|
||||
slog.Uint64("derivation_index", uint64(metadata.DerivationIndex)),
|
||||
)
|
||||
|
||||
// Use mnemonic with the vault's derivation index from metadata
|
||||
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
Debug("Failed to derive long-term key from mnemonic for secret",
|
||||
"error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf(
|
||||
"failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
|
||||
Debug("Successfully derived long-term key from mnemonic", "secret_name", s.Name)
|
||||
|
||||
// Use the long-term key to decrypt the version
|
||||
return version.GetValue(ltIdentity)
|
||||
}
|
||||
|
||||
// getLongTermIdentityFromUnlocker uses the unlocker to obtain and parse
|
||||
// the vault's long-term private key.
|
||||
func (s *Secret) getLongTermIdentityFromUnlocker(
|
||||
unlocker Unlocker,
|
||||
) (*age.X25519Identity, error) {
|
||||
DebugWith("Getting vault's long-term key using unlocker",
|
||||
slog.String("secret_name", s.Name),
|
||||
slog.String("unlocker_type", unlocker.GetType()),
|
||||
slog.String("unlocker_id", unlocker.GetID()),
|
||||
)
|
||||
|
||||
// Step 1: Use the unlocker to get the vault's long-term private key
|
||||
unlockIdentity, err := unlocker.GetIdentity()
|
||||
if err != nil {
|
||||
Debug("Failed to get unlocker identity",
|
||||
"error", err, "secret_name", s.Name,
|
||||
"unlocker_type", unlocker.GetType())
|
||||
|
||||
return nil, fmt.Errorf("failed to get unlocker identity: %w", err)
|
||||
}
|
||||
|
||||
// Read the encrypted long-term private key from the unlocker directory
|
||||
encryptedLtPrivKeyPath := filepath.Join(unlocker.GetDirectory(), "longterm.age")
|
||||
Debug("Reading encrypted long-term private key", "path", encryptedLtPrivKeyPath)
|
||||
|
||||
encryptedLtPrivKey, err := afero.ReadFile(
|
||||
s.vault.GetFilesystem(), encryptedLtPrivKeyPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read encrypted long-term private key",
|
||||
"error", err, "path", encryptedLtPrivKeyPath)
|
||||
|
||||
return nil, fmt.Errorf(
|
||||
"failed to read encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
// Decrypt the encrypted long-term private key using the unlocker
|
||||
Debug("Decrypting long-term private key using unlocker", "secret_name", s.Name)
|
||||
|
||||
ltPrivKeyBuffer, err := DecryptWithIdentity(encryptedLtPrivKey, unlockIdentity)
|
||||
if err != nil {
|
||||
Debug("Failed to decrypt long-term private key",
|
||||
"error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err)
|
||||
}
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
// Parse the long-term private key
|
||||
Debug("Parsing long-term private key", "secret_name", s.Name)
|
||||
|
||||
ltIdentity, err := age.ParseX25519Identity(ltPrivKeyBuffer.String())
|
||||
if err != nil {
|
||||
Debug("Failed to parse long-term private key",
|
||||
"error", err, "secret_name", s.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to parse long-term private key: %w", err)
|
||||
}
|
||||
|
||||
return ltIdentity, nil
|
||||
}
|
||||
|
||||
// GetCurrentVault gets the current vault from the file system
|
||||
// This function is a wrapper around the actual implementation in the vault package
|
||||
// and exists to break the import cycle.
|
||||
//
|
||||
//nolint:ireturn // must return the interface to break the import cycle
|
||||
func GetCurrentVault(fs afero.Fs, stateDir string) (VaultInterface, error) {
|
||||
// This is a forward declaration. The actual implementation is provided
|
||||
// by the vault package when it calls RegisterGetCurrentVaultFunc.
|
||||
if getCurrentVaultFunc == nil {
|
||||
return nil, fmt.Errorf("GetCurrentVault function not registered")
|
||||
return nil, errGetCurrentVaultNotRegistered
|
||||
}
|
||||
|
||||
return getCurrentVaultFunc(fs, stateDir)
|
||||
@@ -288,8 +348,10 @@ func GetCurrentVault(fs afero.Fs, stateDir string) (VaultInterface, error) {
|
||||
//nolint:gochecknoglobals // Required to break import cycle
|
||||
var getCurrentVaultFunc func(fs afero.Fs, stateDir string) (VaultInterface, error)
|
||||
|
||||
// RegisterGetCurrentVaultFunc allows the vault package to register its implementation
|
||||
// of GetCurrentVault to break the import cycle
|
||||
func RegisterGetCurrentVaultFunc(fn func(fs afero.Fs, stateDir string) (VaultInterface, error)) {
|
||||
// RegisterGetCurrentVaultFunc allows the vault package to register its
|
||||
// implementation of GetCurrentVault to break the import cycle
|
||||
func RegisterGetCurrentVaultFunc(
|
||||
fn func(fs afero.Fs, stateDir string) (VaultInterface, error),
|
||||
) {
|
||||
getCurrentVaultFunc = fn
|
||||
}
|
||||
|
||||
+135
-124
@@ -1,7 +1,8 @@
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package secret
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
@@ -14,6 +15,17 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// testMnemonicValue is the standard BIP39 test vector mnemonic.
|
||||
//
|
||||
//nolint:dupword // BIP39 test mnemonic repeats words by design
|
||||
const testMnemonicValue = "abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon abandon abandon about"
|
||||
|
||||
var (
|
||||
errMnemonicNotSet = errors.New("SB_SECRET_MNEMONIC not set")
|
||||
errNotImplementedInMock = errors.New("not implemented in mock")
|
||||
)
|
||||
|
||||
// MockVault is a test implementation of the VaultInterface
|
||||
type MockVault struct {
|
||||
name string
|
||||
@@ -30,14 +42,18 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool)
|
||||
// Create secret directory with proper storage name conversion
|
||||
storageName := strings.ReplaceAll(name, "/", "%")
|
||||
secretDir := filepath.Join(m.directory, "secrets.d", storageName)
|
||||
if err := m.fs.MkdirAll(secretDir, 0o700); err != nil {
|
||||
|
||||
err := m.fs.MkdirAll(secretDir, 0o700)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Create version directory with proper path
|
||||
versionName := "20240101.001" // Use a fixed version name for testing
|
||||
versionDir := filepath.Join(secretDir, "versions", versionName)
|
||||
if err := m.fs.MkdirAll(versionDir, 0o700); err != nil {
|
||||
|
||||
err = m.fs.MkdirAll(versionDir, 0o700)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -47,7 +63,7 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool)
|
||||
// Derive long-term key using the vault's derivation index
|
||||
mnemonic := os.Getenv(EnvMnemonic)
|
||||
if mnemonic == "" {
|
||||
return fmt.Errorf("SB_SECRET_MNEMONIC not set")
|
||||
return errMnemonicNotSet
|
||||
}
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonic, m.derivationIndex)
|
||||
@@ -56,13 +72,58 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool)
|
||||
}
|
||||
|
||||
// Write long-term public key if it doesn't exist
|
||||
if _, err := m.fs.Stat(ltPubKeyPath); os.IsNotExist(err) {
|
||||
_, err = m.fs.Stat(ltPubKeyPath)
|
||||
if os.IsNotExist(err) {
|
||||
pubKey := ltIdentity.Recipient().String()
|
||||
if err := afero.WriteFile(m.fs, ltPubKeyPath, []byte(pubKey), 0o600); err != nil {
|
||||
|
||||
err = afero.WriteFile(m.fs, ltPubKeyPath, []byte(pubKey), 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
err = m.writeVersionFiles(versionDir, value, ltIdentity)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Create current file pointing to the version (just the version name)
|
||||
currentLink := filepath.Join(secretDir, "current")
|
||||
|
||||
return afero.WriteFile(m.fs, currentLink, []byte(versionName), 0o600)
|
||||
}
|
||||
|
||||
func (m *MockVault) GetName() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
//nolint:ireturn // implements VaultInterface
|
||||
func (m *MockVault) GetFilesystem() afero.Fs {
|
||||
return m.fs
|
||||
}
|
||||
|
||||
//nolint:ireturn // implements VaultInterface
|
||||
func (m *MockVault) GetCurrentUnlocker() (Unlocker, error) {
|
||||
return nil, errNotImplementedInMock
|
||||
}
|
||||
|
||||
func (m *MockVault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
|
||||
return nil, errNotImplementedInMock
|
||||
}
|
||||
|
||||
func (m *MockVault) CreatePassphraseUnlocker(
|
||||
_ *memguard.LockedBuffer,
|
||||
) (*PassphraseUnlocker, error) {
|
||||
return nil, errNotImplementedInMock
|
||||
}
|
||||
|
||||
// writeVersionFiles generates a version keypair and writes the version
|
||||
// key and value files for the mock vault.
|
||||
func (m *MockVault) writeVersionFiles(
|
||||
versionDir string,
|
||||
value *memguard.LockedBuffer,
|
||||
ltIdentity *age.X25519Identity,
|
||||
) error {
|
||||
// Generate version-specific keypair
|
||||
versionIdentity, err := age.GenerateX25519Identity()
|
||||
if err != nil {
|
||||
@@ -71,7 +132,10 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool)
|
||||
|
||||
// Write version public key
|
||||
pubKeyPath := filepath.Join(versionDir, "pub.age")
|
||||
if err := afero.WriteFile(m.fs, pubKeyPath, []byte(versionIdentity.Recipient().String()), 0o600); err != nil {
|
||||
|
||||
err = afero.WriteFile(
|
||||
m.fs, pubKeyPath, []byte(versionIdentity.Recipient().String()), 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -83,60 +147,32 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool)
|
||||
|
||||
// Write encrypted value
|
||||
valuePath := filepath.Join(versionDir, "value.age")
|
||||
if err := afero.WriteFile(m.fs, valuePath, encryptedValue, 0o600); err != nil {
|
||||
|
||||
err = afero.WriteFile(m.fs, valuePath, encryptedValue, 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Encrypt version private key to long-term public key
|
||||
versionPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(versionIdentity.String()))
|
||||
defer versionPrivKeyBuffer.Destroy()
|
||||
encryptedPrivKey, err := EncryptToRecipient(versionPrivKeyBuffer, ltIdentity.Recipient())
|
||||
|
||||
encryptedPrivKey, err := EncryptToRecipient(
|
||||
versionPrivKeyBuffer, ltIdentity.Recipient())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Write encrypted version private key
|
||||
privKeyPath := filepath.Join(versionDir, "priv.age")
|
||||
if err := afero.WriteFile(m.fs, privKeyPath, encryptedPrivKey, 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Create current file pointing to the version (just the version name)
|
||||
currentLink := filepath.Join(secretDir, "current")
|
||||
if err := afero.WriteFile(m.fs, currentLink, []byte(versionName), 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
return afero.WriteFile(m.fs, privKeyPath, encryptedPrivKey, 0o600)
|
||||
}
|
||||
|
||||
func (m *MockVault) GetName() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
func (m *MockVault) GetFilesystem() afero.Fs {
|
||||
return m.fs
|
||||
}
|
||||
|
||||
func (m *MockVault) GetCurrentUnlocker() (Unlocker, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *MockVault) CreatePassphraseUnlocker(_ *memguard.LockedBuffer) (*PassphraseUnlocker, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func TestPerSecretKeyFunctionality(t *testing.T) {
|
||||
// Create an in-memory filesystem for testing
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
// Set test mnemonic for direct encryption/decryption
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(EnvMnemonic, testMnemonic)
|
||||
|
||||
// Set up a test vault structure
|
||||
baseDir := "/test-config/berlin.sneak.pkg.secret"
|
||||
vaultDir := filepath.Join(baseDir, "vaults.d", "test-vault")
|
||||
// setupMockVaultDirs creates the vault directory structure, long-term
|
||||
// public key, and current vault pointer for tests.
|
||||
func setupMockVaultDirs(t *testing.T, fs afero.Fs, baseDir, vaultDir string) {
|
||||
t.Helper()
|
||||
|
||||
// Create vault directory structure
|
||||
err := fs.MkdirAll(filepath.Join(vaultDir, "secrets.d"), DirPerms)
|
||||
@@ -145,13 +181,14 @@ func TestPerSecretKeyFunctionality(t *testing.T) {
|
||||
}
|
||||
|
||||
// Generate a long-term keypair for the vault using the test mnemonic
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonicValue, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to generate long-term identity: %v", err)
|
||||
}
|
||||
|
||||
// Write long-term public key
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
|
||||
err = afero.WriteFile(
|
||||
fs,
|
||||
ltPubKeyPath,
|
||||
@@ -164,10 +201,56 @@ func TestPerSecretKeyFunctionality(t *testing.T) {
|
||||
|
||||
// Set current vault
|
||||
currentVaultPath := filepath.Join(baseDir, "currentvault")
|
||||
|
||||
err = afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), FilePerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to set current vault: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// verifySecretFiles checks that AddSecret created the expected version
|
||||
// files for the secret.
|
||||
func verifySecretFiles(t *testing.T, fs afero.Fs, vaultDir, secretName string) {
|
||||
t.Helper()
|
||||
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", secretName)
|
||||
|
||||
// Check versions directory exists
|
||||
versionsDir := filepath.Join(secretDir, "versions")
|
||||
|
||||
versionsDirExists, err := afero.DirExists(fs, versionsDir)
|
||||
if err != nil || !versionsDirExists {
|
||||
t.Fatalf("versions directory was not created")
|
||||
}
|
||||
|
||||
// Check current file exists and points at a version
|
||||
currentVersion, err := GetCurrentVersion(fs, secretDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current version: %v", err)
|
||||
}
|
||||
|
||||
// Check value.age exists in the version directory
|
||||
versionDir := filepath.Join(versionsDir, currentVersion)
|
||||
|
||||
valueExists, err := afero.Exists(fs, filepath.Join(versionDir, "value.age"))
|
||||
if err != nil || !valueExists {
|
||||
t.Fatalf("value.age file was not created in version directory")
|
||||
}
|
||||
}
|
||||
|
||||
//nolint:paralleltest // uses t.Setenv (process-global environment)
|
||||
func TestPerSecretKeyFunctionality(t *testing.T) {
|
||||
// Create an in-memory filesystem for testing
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
// Set test mnemonic for direct encryption/decryption
|
||||
t.Setenv(EnvMnemonic, testMnemonicValue)
|
||||
|
||||
// Set up a test vault structure
|
||||
baseDir := "/test-config/berlin.sneak.pkg.secret"
|
||||
vaultDir := filepath.Join(baseDir, "vaults.d", "test-vault")
|
||||
|
||||
setupMockVaultDirs(t, fs, baseDir, vaultDir)
|
||||
|
||||
// Create vault instance using the mock vault
|
||||
vault := &MockVault{
|
||||
@@ -193,30 +276,7 @@ func TestPerSecretKeyFunctionality(t *testing.T) {
|
||||
}
|
||||
|
||||
// Verify that all expected files were created
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", secretName)
|
||||
|
||||
// Check versions directory exists
|
||||
versionsDir := filepath.Join(secretDir, "versions")
|
||||
versionsDirExists, err := afero.DirExists(fs, versionsDir)
|
||||
if err != nil || !versionsDirExists {
|
||||
t.Fatalf("versions directory was not created")
|
||||
}
|
||||
|
||||
// Check current symlink exists
|
||||
currentVersion, err := GetCurrentVersion(fs, secretDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current version: %v", err)
|
||||
}
|
||||
|
||||
// Check value.age exists in the version directory
|
||||
versionDir := filepath.Join(versionsDir, currentVersion)
|
||||
valueExists, err := afero.Exists(
|
||||
fs,
|
||||
filepath.Join(versionDir, "value.age"),
|
||||
)
|
||||
if err != nil || !valueExists {
|
||||
t.Fatalf("value.age file was not created in version directory")
|
||||
}
|
||||
verifySecretFiles(t, fs, vaultDir, secretName)
|
||||
|
||||
t.Logf("All expected files created successfully with versioning")
|
||||
})
|
||||
@@ -245,76 +305,27 @@ func TestPerSecretKeyFunctionality(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Error checking if secret exists: %v", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
t.Fatalf("Secret should exist but Exists() returned false")
|
||||
}
|
||||
|
||||
t.Logf("Secret.Exists() works correctly")
|
||||
})
|
||||
}
|
||||
|
||||
// For testing purposes only
|
||||
func isValidSecretName(name string) bool {
|
||||
if name == "" {
|
||||
return false
|
||||
}
|
||||
// Valid characters for secret names: lowercase letters, numbers, dash, dot, underscore, slash
|
||||
for _, char := range name {
|
||||
if (char < 'a' || char > 'z') && // lowercase letters
|
||||
(char < '0' || char > '9') && // numbers
|
||||
char != '-' && // dash
|
||||
char != '.' && // dot
|
||||
char != '_' && // underscore
|
||||
char != '/' { // slash
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func TestSecretNameValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
valid bool
|
||||
}{
|
||||
{"valid-name", true},
|
||||
{"valid.name", true},
|
||||
{"valid_name", true},
|
||||
{"valid/path/name", true},
|
||||
{"123valid", true},
|
||||
{"", false},
|
||||
{"Invalid-Name", false}, // uppercase not allowed
|
||||
{"invalid name", false}, // space not allowed
|
||||
{"invalid@name", false}, // @ not allowed
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
result := isValidSecretName(test.name)
|
||||
if result != test.valid {
|
||||
t.Errorf(
|
||||
"isValidSecretName(%q) = %v, want %v",
|
||||
test.name,
|
||||
result,
|
||||
test.valid,
|
||||
)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecretGetValueWithEnvMnemonicUsesVaultDerivationIndex(t *testing.T) {
|
||||
// This test demonstrates the bug where GetValue uses hardcoded index 0
|
||||
// instead of the vault's actual derivation index when using environment mnemonic
|
||||
|
||||
// Set up test mnemonic
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(EnvMnemonic, testMnemonic)
|
||||
t.Setenv(EnvMnemonic, testMnemonicValue)
|
||||
|
||||
// Create temporary directory for vaults
|
||||
fs := afero.NewOsFs()
|
||||
tempDir, err := afero.TempDir(fs, "", "secret-test-")
|
||||
require.NoError(t, err)
|
||||
|
||||
defer func() {
|
||||
_ = fs.RemoveAll(tempDir)
|
||||
}()
|
||||
|
||||
@@ -0,0 +1,385 @@
|
||||
//go:build darwin
|
||||
// +build darwin
|
||||
|
||||
package secret
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/internal/macse"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
const (
|
||||
// seKeyLabelPrefix is the prefix for Secure Enclave CTK identity labels.
|
||||
seKeyLabelPrefix = "berlin.sneak.app.secret.se"
|
||||
|
||||
// seUnlockerType is the metadata type string for Secure Enclave unlockers.
|
||||
seUnlockerType = "secure-enclave"
|
||||
|
||||
// seLongtermFilename is the filename for the SE-encrypted vault long-term private key.
|
||||
seLongtermFilename = "longterm.age.se"
|
||||
)
|
||||
|
||||
// SecureEnclaveUnlockerMetadata extends UnlockerMetadata with SE-specific data.
|
||||
type SecureEnclaveUnlockerMetadata struct {
|
||||
UnlockerMetadata
|
||||
SEKeyLabel string `json:"seKeyLabel"`
|
||||
SEKeyHash string `json:"seKeyHash"`
|
||||
}
|
||||
|
||||
// SecureEnclaveUnlocker represents a Secure Enclave-protected unlocker.
|
||||
type SecureEnclaveUnlocker struct {
|
||||
Directory string
|
||||
Metadata UnlockerMetadata
|
||||
fs afero.Fs
|
||||
}
|
||||
|
||||
// GetIdentity implements Unlocker interface for SE-based unlockers.
|
||||
// Decrypts the vault's long-term private key directly using the Secure Enclave.
|
||||
func (s *SecureEnclaveUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
DebugWith("Getting SE unlocker identity",
|
||||
slog.String("unlocker_id", s.GetID()),
|
||||
)
|
||||
|
||||
// Get SE key label from metadata
|
||||
seKeyLabel, _, err := s.getSEKeyInfo()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get SE key info: %w", err)
|
||||
}
|
||||
|
||||
// Read ECIES-encrypted long-term private key from disk
|
||||
encryptedPath := filepath.Join(s.Directory, seLongtermFilename)
|
||||
encryptedData, err := afero.ReadFile(s.fs, encryptedPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to read SE-encrypted long-term key: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
||||
DebugWith("Read SE-encrypted long-term key",
|
||||
slog.Int("encrypted_length", len(encryptedData)),
|
||||
)
|
||||
|
||||
// Decrypt using the Secure Enclave (ECDH happens inside SE hardware)
|
||||
decryptedData, err := macse.Decrypt(seKeyLabel, encryptedData)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to decrypt long-term key with SE: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
||||
// Parse the decrypted long-term private key
|
||||
ltIdentity, err := age.ParseX25519Identity(string(decryptedData))
|
||||
|
||||
// Clear sensitive data immediately
|
||||
for i := range decryptedData {
|
||||
decryptedData[i] = 0
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to parse long-term private key: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
||||
DebugWith("Successfully decrypted long-term key via SE",
|
||||
slog.String("unlocker_id", s.GetID()),
|
||||
)
|
||||
|
||||
return ltIdentity, nil
|
||||
}
|
||||
|
||||
// GetType implements Unlocker interface.
|
||||
func (s *SecureEnclaveUnlocker) GetType() string {
|
||||
return seUnlockerType
|
||||
}
|
||||
|
||||
// GetMetadata implements Unlocker interface.
|
||||
func (s *SecureEnclaveUnlocker) GetMetadata() UnlockerMetadata {
|
||||
return s.Metadata
|
||||
}
|
||||
|
||||
// GetDirectory implements Unlocker interface.
|
||||
func (s *SecureEnclaveUnlocker) GetDirectory() string {
|
||||
return s.Directory
|
||||
}
|
||||
|
||||
// GetID implements Unlocker interface.
|
||||
func (s *SecureEnclaveUnlocker) GetID() string {
|
||||
hostname, err := os.Hostname()
|
||||
if err != nil {
|
||||
hostname = "unknown"
|
||||
}
|
||||
|
||||
createdAt := s.Metadata.CreatedAt
|
||||
timestamp := createdAt.Format("2006-01-02.15.04")
|
||||
|
||||
return fmt.Sprintf("%s-%s-%s", timestamp, hostname, seUnlockerType)
|
||||
}
|
||||
|
||||
// Remove implements Unlocker interface.
|
||||
func (s *SecureEnclaveUnlocker) Remove() error {
|
||||
_, seKeyHash, err := s.getSEKeyInfo()
|
||||
if err != nil {
|
||||
Debug("Failed to get SE key info during removal", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to get SE key info: %w", err)
|
||||
}
|
||||
|
||||
if seKeyHash != "" {
|
||||
Debug("Deleting SE key", "hash", seKeyHash)
|
||||
if err := macse.DeleteKey(seKeyHash); err != nil {
|
||||
Debug("Failed to delete SE key", "error", err, "hash", seKeyHash)
|
||||
|
||||
return fmt.Errorf("failed to delete SE key: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
Debug("Removing SE unlocker directory", "directory", s.Directory)
|
||||
if err := RemoveDirAtomic(s.fs, s.Directory); err != nil {
|
||||
return fmt.Errorf("failed to remove SE unlocker directory: %w", err)
|
||||
}
|
||||
|
||||
Debug("Successfully removed SE unlocker", "unlocker_id", s.GetID())
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// getSEKeyInfo reads the SE key label and hash from metadata.
|
||||
func (s *SecureEnclaveUnlocker) getSEKeyInfo() (label string, hash string, err error) {
|
||||
metadataPath := filepath.Join(s.Directory, "unlocker-metadata.json")
|
||||
metadataData, err := afero.ReadFile(s.fs, metadataPath)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("failed to read SE metadata: %w", err)
|
||||
}
|
||||
|
||||
var seMetadata SecureEnclaveUnlockerMetadata
|
||||
if err := json.Unmarshal(metadataData, &seMetadata); err != nil {
|
||||
return "", "", fmt.Errorf("failed to parse SE metadata: %w", err)
|
||||
}
|
||||
|
||||
return seMetadata.SEKeyLabel, seMetadata.SEKeyHash, nil
|
||||
}
|
||||
|
||||
// NewSecureEnclaveUnlocker creates a new SecureEnclaveUnlocker instance.
|
||||
func NewSecureEnclaveUnlocker(
|
||||
fs afero.Fs,
|
||||
directory string,
|
||||
metadata UnlockerMetadata,
|
||||
) *SecureEnclaveUnlocker {
|
||||
return &SecureEnclaveUnlocker{
|
||||
Directory: directory,
|
||||
Metadata: metadata,
|
||||
fs: fs,
|
||||
}
|
||||
}
|
||||
|
||||
// generateSEKeyLabel generates a unique label for the SE CTK identity.
|
||||
func generateSEKeyLabel(vaultName string) (string, error) {
|
||||
hostname, err := os.Hostname()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to get hostname: %w", err)
|
||||
}
|
||||
|
||||
enrollmentDate := time.Now().UTC().Format("2006-01-02")
|
||||
|
||||
return fmt.Sprintf(
|
||||
"%s.%s-%s-%s",
|
||||
seKeyLabelPrefix,
|
||||
vaultName,
|
||||
hostname,
|
||||
enrollmentDate,
|
||||
), nil
|
||||
}
|
||||
|
||||
// CreateSecureEnclaveUnlocker creates a new SE unlocker.
|
||||
// The vault's long-term private key is encrypted directly by the Secure Enclave
|
||||
// using ECIES. No intermediate age keypair is used.
|
||||
func CreateSecureEnclaveUnlocker(
|
||||
fs afero.Fs,
|
||||
stateDir string,
|
||||
) (*SecureEnclaveUnlocker, error) {
|
||||
if err := checkMacOSAvailable(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
vault, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get current vault: %w", err)
|
||||
}
|
||||
|
||||
// Generate SE key label
|
||||
seKeyLabel, err := generateSEKeyLabel(vault.GetName())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate SE key label: %w", err)
|
||||
}
|
||||
|
||||
// Step 1: Create P-256 key in the Secure Enclave via sc_auth
|
||||
Debug("Creating Secure Enclave key", "label", seKeyLabel)
|
||||
_, seKeyHash, err := macse.CreateKey(seKeyLabel)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create SE key: %w", err)
|
||||
}
|
||||
|
||||
Debug("Created SE key", "label", seKeyLabel, "hash", seKeyHash)
|
||||
|
||||
// Step 2: Get the vault's long-term private key
|
||||
ltPrivKeyData, err := getLongTermKeyForSE(fs, vault)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to get long-term private key: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
defer ltPrivKeyData.Destroy()
|
||||
|
||||
// Step 3: Encrypt the long-term key directly with the SE (ECIES)
|
||||
encryptedLtKey, err := macse.Encrypt(seKeyLabel, ltPrivKeyData.Bytes())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to encrypt long-term key with SE: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
||||
// Step 4: Prepare the unlocker directory's path and metadata
|
||||
vaultDir, err := vault.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
unlockerDirName := fmt.Sprintf("se-%s", filepath.Base(seKeyLabel))
|
||||
unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerDirName)
|
||||
|
||||
seMetadata := SecureEnclaveUnlockerMetadata{
|
||||
UnlockerMetadata: UnlockerMetadata{
|
||||
Type: seUnlockerType,
|
||||
CreatedAt: time.Now().UTC(),
|
||||
Flags: []string{seUnlockerType, "macos"},
|
||||
},
|
||||
SEKeyLabel: seKeyLabel,
|
||||
SEKeyHash: seKeyHash,
|
||||
}
|
||||
|
||||
metadataBytes, err := json.MarshalIndent(seMetadata, "", " ")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal metadata: %w", err)
|
||||
}
|
||||
|
||||
// Step 5: Write the SE-encrypted long-term key, then the metadata
|
||||
err = WriteDir(fs, unlockerDir, func(dir string) error {
|
||||
ltKeyPath := filepath.Join(dir, seLongtermFilename)
|
||||
if err := WriteFileAtomic(fs, ltKeyPath, encryptedLtKey); err != nil {
|
||||
return fmt.Errorf(
|
||||
"failed to write SE-encrypted long-term key: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
||||
metadataPath := filepath.Join(dir, "unlocker-metadata.json")
|
||||
if err := WriteFileAtomic(fs, metadataPath, metadataBytes); err != nil {
|
||||
return fmt.Errorf("failed to write metadata: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &SecureEnclaveUnlocker{
|
||||
Directory: unlockerDir,
|
||||
Metadata: seMetadata.UnlockerMetadata,
|
||||
fs: fs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// getLongTermKeyForSE retrieves the vault's long-term private key
|
||||
// either from the mnemonic env var or by unlocking via the current unlocker.
|
||||
func getLongTermKeyForSE(
|
||||
fs afero.Fs,
|
||||
vault VaultInterface,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
envMnemonic := os.Getenv(EnvMnemonic)
|
||||
if envMnemonic != "" {
|
||||
// Read vault metadata to get the correct derivation index
|
||||
vaultDir, err := vault.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
metadataPath := filepath.Join(vaultDir, "vault-metadata.json")
|
||||
metadataBytes, err := afero.ReadFile(fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read vault metadata: %w", err)
|
||||
}
|
||||
|
||||
var metadata VaultMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse vault metadata: %w", err)
|
||||
}
|
||||
|
||||
// Use mnemonic with the vault's actual derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(
|
||||
envMnemonic,
|
||||
metadata.DerivationIndex,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to derive long-term key from mnemonic: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
||||
return memguard.NewBufferFromBytes([]byte(ltIdentity.String())), nil
|
||||
}
|
||||
|
||||
currentUnlocker, err := vault.GetCurrentUnlocker()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get current unlocker: %w", err)
|
||||
}
|
||||
|
||||
currentIdentity, err := currentUnlocker.GetIdentity()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to get current unlocker identity: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
||||
// All unlocker types store longterm.age in their directory
|
||||
longtermPath := filepath.Join(
|
||||
currentUnlocker.GetDirectory(),
|
||||
"longterm.age",
|
||||
)
|
||||
encryptedLtKey, err := afero.ReadFile(fs, longtermPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to read encrypted long-term key: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
|
||||
ltPrivKeyBuffer, err := DecryptWithIdentity(
|
||||
encryptedLtKey,
|
||||
currentIdentity,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decrypt long-term key: %w", err)
|
||||
}
|
||||
|
||||
return ltPrivKeyBuffer, nil
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
//go:build !darwin
|
||||
|
||||
package secret
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"filippo.io/age"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// seUnlockerType is the type string for Secure Enclave unlockers.
|
||||
const seUnlockerType = "secure-enclave"
|
||||
|
||||
var errSENotSupported = errors.New(
|
||||
"secure enclave unlockers are only supported on macOS",
|
||||
)
|
||||
|
||||
// SecureEnclaveUnlockerMetadata is a stub for non-Darwin platforms.
|
||||
type SecureEnclaveUnlockerMetadata struct {
|
||||
UnlockerMetadata
|
||||
|
||||
SEKeyLabel string `json:"seKeyLabel"`
|
||||
SEKeyHash string `json:"seKeyHash"`
|
||||
}
|
||||
|
||||
// SecureEnclaveUnlocker is a stub for non-Darwin platforms.
|
||||
type SecureEnclaveUnlocker struct {
|
||||
Directory string
|
||||
Metadata UnlockerMetadata
|
||||
fs afero.Fs
|
||||
}
|
||||
|
||||
// NewSecureEnclaveUnlocker creates a stub SecureEnclaveUnlocker on
|
||||
// non-Darwin platforms. The returned instance's methods that require
|
||||
// macOS functionality will return errors.
|
||||
func NewSecureEnclaveUnlocker(
|
||||
fs afero.Fs,
|
||||
directory string,
|
||||
metadata UnlockerMetadata,
|
||||
) *SecureEnclaveUnlocker {
|
||||
return &SecureEnclaveUnlocker{
|
||||
Directory: directory,
|
||||
Metadata: metadata,
|
||||
fs: fs,
|
||||
}
|
||||
}
|
||||
|
||||
// GetIdentity returns an error on non-Darwin platforms.
|
||||
func (s *SecureEnclaveUnlocker) GetIdentity() (*age.X25519Identity, error) {
|
||||
return nil, errSENotSupported
|
||||
}
|
||||
|
||||
// GetType returns the unlocker type.
|
||||
func (s *SecureEnclaveUnlocker) GetType() string {
|
||||
return seUnlockerType
|
||||
}
|
||||
|
||||
// GetMetadata returns the unlocker metadata.
|
||||
func (s *SecureEnclaveUnlocker) GetMetadata() UnlockerMetadata {
|
||||
return s.Metadata
|
||||
}
|
||||
|
||||
// GetDirectory returns the unlocker directory.
|
||||
func (s *SecureEnclaveUnlocker) GetDirectory() string {
|
||||
return s.Directory
|
||||
}
|
||||
|
||||
// GetID returns the unlocker ID.
|
||||
func (s *SecureEnclaveUnlocker) GetID() string {
|
||||
return s.Metadata.CreatedAt.Format("2006-01-02.15.04") + "-" + seUnlockerType
|
||||
}
|
||||
|
||||
// Remove returns an error on non-Darwin platforms.
|
||||
func (s *SecureEnclaveUnlocker) Remove() error {
|
||||
return errSENotSupported
|
||||
}
|
||||
|
||||
// CreateSecureEnclaveUnlocker returns an error on non-Darwin platforms.
|
||||
func CreateSecureEnclaveUnlocker(
|
||||
_ afero.Fs,
|
||||
_ string,
|
||||
) (*SecureEnclaveUnlocker, error) {
|
||||
return nil, errSENotSupported
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
//go:build !darwin
|
||||
|
||||
//nolint:testpackage // white-box test asserting unexported sentinel errors
|
||||
package secret
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNewSecureEnclaveUnlocker(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
dir := "/tmp/test-se-unlocker"
|
||||
metadata := UnlockerMetadata{
|
||||
Type: seUnlockerType,
|
||||
CreatedAt: time.Date(2026, 1, 15, 10, 30, 0, 0, time.UTC),
|
||||
Flags: []string{seUnlockerType, "macos"},
|
||||
}
|
||||
|
||||
unlocker := NewSecureEnclaveUnlocker(fs, dir, metadata)
|
||||
require.NotNil(t, unlocker, "NewSecureEnclaveUnlocker should return a valid instance")
|
||||
|
||||
// Test GetType returns correct type
|
||||
assert.Equal(t, seUnlockerType, unlocker.GetType())
|
||||
|
||||
// Test GetMetadata returns the metadata we passed in
|
||||
assert.Equal(t, metadata, unlocker.GetMetadata())
|
||||
|
||||
// Test GetDirectory returns the directory we passed in
|
||||
assert.Equal(t, dir, unlocker.GetDirectory())
|
||||
|
||||
// Test GetID returns a formatted string with the creation timestamp
|
||||
expectedID := "2026-01-15.10.30-secure-enclave"
|
||||
assert.Equal(t, expectedID, unlocker.GetID())
|
||||
}
|
||||
|
||||
func TestSecureEnclaveUnlockerGetIdentityReturnsError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
metadata := UnlockerMetadata{
|
||||
Type: seUnlockerType,
|
||||
CreatedAt: time.Now().UTC(),
|
||||
}
|
||||
|
||||
unlocker := NewSecureEnclaveUnlocker(fs, "/tmp/test", metadata)
|
||||
|
||||
identity, err := unlocker.GetIdentity()
|
||||
assert.Nil(t, identity)
|
||||
require.Error(t, err)
|
||||
require.ErrorIs(t, err, errSENotSupported)
|
||||
}
|
||||
|
||||
func TestSecureEnclaveUnlockerRemoveReturnsError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
metadata := UnlockerMetadata{
|
||||
Type: seUnlockerType,
|
||||
CreatedAt: time.Now().UTC(),
|
||||
}
|
||||
|
||||
unlocker := NewSecureEnclaveUnlocker(fs, "/tmp/test", metadata)
|
||||
|
||||
err := unlocker.Remove()
|
||||
require.Error(t, err)
|
||||
require.ErrorIs(t, err, errSENotSupported)
|
||||
}
|
||||
|
||||
func TestCreateSecureEnclaveUnlockerReturnsError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
unlocker, err := CreateSecureEnclaveUnlocker(fs, "/tmp/test")
|
||||
assert.Nil(t, unlocker)
|
||||
require.Error(t, err)
|
||||
require.ErrorIs(t, err, errSENotSupported)
|
||||
}
|
||||
|
||||
func TestSecureEnclaveUnlockerImplementsInterface(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
metadata := UnlockerMetadata{
|
||||
Type: seUnlockerType,
|
||||
CreatedAt: time.Now().UTC(),
|
||||
}
|
||||
|
||||
unlocker := NewSecureEnclaveUnlocker(fs, "/tmp/test", metadata)
|
||||
|
||||
// Verify the stub implements the Unlocker interface
|
||||
var _ Unlocker = unlocker
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
//go:build darwin
|
||||
// +build darwin
|
||||
|
||||
package secret
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNewSecureEnclaveUnlocker(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
dir := "/tmp/test-se-unlocker"
|
||||
metadata := UnlockerMetadata{
|
||||
Type: "secure-enclave",
|
||||
CreatedAt: time.Date(2026, 1, 15, 10, 30, 0, 0, time.UTC),
|
||||
Flags: []string{"secure-enclave", "macos"},
|
||||
}
|
||||
|
||||
unlocker := NewSecureEnclaveUnlocker(fs, dir, metadata)
|
||||
require.NotNil(t, unlocker, "NewSecureEnclaveUnlocker should return a valid instance")
|
||||
|
||||
// Test GetType returns correct type
|
||||
assert.Equal(t, seUnlockerType, unlocker.GetType())
|
||||
|
||||
// Test GetMetadata returns the metadata we passed in
|
||||
assert.Equal(t, metadata, unlocker.GetMetadata())
|
||||
|
||||
// Test GetDirectory returns the directory we passed in
|
||||
assert.Equal(t, dir, unlocker.GetDirectory())
|
||||
}
|
||||
|
||||
func TestSecureEnclaveUnlockerImplementsInterface(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
metadata := UnlockerMetadata{
|
||||
Type: "secure-enclave",
|
||||
CreatedAt: time.Now().UTC(),
|
||||
}
|
||||
|
||||
unlocker := NewSecureEnclaveUnlocker(fs, "/tmp/test", metadata)
|
||||
|
||||
// Verify the darwin implementation implements the Unlocker interface
|
||||
var _ Unlocker = unlocker
|
||||
}
|
||||
|
||||
func TestSecureEnclaveUnlockerGetIDFormat(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
metadata := UnlockerMetadata{
|
||||
Type: "secure-enclave",
|
||||
CreatedAt: time.Date(2026, 3, 10, 14, 30, 0, 0, time.UTC),
|
||||
}
|
||||
|
||||
unlocker := NewSecureEnclaveUnlocker(fs, "/tmp/test", metadata)
|
||||
id := unlocker.GetID()
|
||||
|
||||
// ID should contain the timestamp and "secure-enclave" type
|
||||
assert.Contains(t, id, "2026-03-10.14.30")
|
||||
assert.Contains(t, id, seUnlockerType)
|
||||
}
|
||||
|
||||
func TestGenerateSEKeyLabel(t *testing.T) {
|
||||
label, err := generateSEKeyLabel("test-vault")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Label should contain the prefix and vault name
|
||||
assert.Contains(t, label, seKeyLabelPrefix)
|
||||
assert.Contains(t, label, "test-vault")
|
||||
}
|
||||
|
||||
func TestSecureEnclaveUnlockerGetIdentityMissingFile(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
dir := "/tmp/test-se-unlocker-missing"
|
||||
|
||||
// Create unlocker directory with metadata but no encrypted key file
|
||||
require.NoError(t, fs.MkdirAll(dir, DirPerms))
|
||||
|
||||
metadataJSON := `{
|
||||
"type": "secure-enclave",
|
||||
"createdAt": "2026-01-15T10:30:00Z",
|
||||
"seKeyLabel": "berlin.sneak.app.secret.se.test",
|
||||
"seKeyHash": "abc123"
|
||||
}`
|
||||
require.NoError(t, afero.WriteFile(fs, dir+"/unlocker-metadata.json", []byte(metadataJSON), FilePerms))
|
||||
|
||||
metadata := UnlockerMetadata{
|
||||
Type: "secure-enclave",
|
||||
CreatedAt: time.Date(2026, 1, 15, 10, 30, 0, 0, time.UTC),
|
||||
}
|
||||
|
||||
unlocker := NewSecureEnclaveUnlocker(fs, dir, metadata)
|
||||
|
||||
// GetIdentity should fail because the encrypted longterm key file is missing
|
||||
identity, err := unlocker.GetIdentity()
|
||||
assert.Nil(t, identity)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "failed to read SE-encrypted long-term key")
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
//go:build darwin
|
||||
|
||||
package secret
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateKeychainItemName(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
itemName string
|
||||
wantErr bool
|
||||
}{
|
||||
// Valid cases
|
||||
{
|
||||
name: "valid simple name",
|
||||
itemName: "my-secret-key",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid name with dots",
|
||||
itemName: "com.example.app.key",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid name with underscores",
|
||||
itemName: "my_secret_key_123",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid alphanumeric",
|
||||
itemName: "Secret123Key",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid with hyphen at start",
|
||||
itemName: "-my-key",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid with dot at start",
|
||||
itemName: ".hidden-key",
|
||||
wantErr: false,
|
||||
},
|
||||
|
||||
// Invalid cases
|
||||
{
|
||||
name: "empty item name",
|
||||
itemName: "",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with spaces",
|
||||
itemName: "my secret key",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with semicolon",
|
||||
itemName: "key;rm -rf /",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with pipe",
|
||||
itemName: "key|cat /etc/passwd",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with backticks",
|
||||
itemName: "key`whoami`",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with dollar sign",
|
||||
itemName: "key$(whoami)",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with quotes",
|
||||
itemName: "key\"name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with single quotes",
|
||||
itemName: "key'name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with backslash",
|
||||
itemName: "key\\name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with newline",
|
||||
itemName: "key\nname",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with carriage return",
|
||||
itemName: "key\rname",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with ampersand",
|
||||
itemName: "key&echo test",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with redirect",
|
||||
itemName: "key>/tmp/test",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with null byte",
|
||||
itemName: "key\x00name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with parentheses",
|
||||
itemName: "key(test)",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with brackets",
|
||||
itemName: "key[test]",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with asterisk",
|
||||
itemName: "key*",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with question mark",
|
||||
itemName: "key?",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := validateKeychainItemName(tt.itemName)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("validateKeychainItemName() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
//nolint:testpackage // white-box test of unexported internals
|
||||
package secret
|
||||
|
||||
import (
|
||||
@@ -5,148 +6,60 @@ import (
|
||||
)
|
||||
|
||||
func TestValidateGPGKeyID(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
keyID string
|
||||
wantErr bool
|
||||
}{
|
||||
// Valid cases
|
||||
{"valid email address", "test@example.com", false},
|
||||
{"valid email with dots and hyphens", "test.user-name@example-domain.co.uk", false},
|
||||
{"valid email with plus", "test+tag@example.com", false},
|
||||
{"valid short key ID (8 hex chars)", "ABCDEF12", false},
|
||||
{"valid long key ID (16 hex chars)", "ABCDEF1234567890", false},
|
||||
{
|
||||
name: "valid email address",
|
||||
keyID: "test@example.com",
|
||||
wantErr: false,
|
||||
"valid fingerprint (40 hex chars)",
|
||||
"ABCDEF1234567890ABCDEF1234567890ABCDEF12", false,
|
||||
},
|
||||
{
|
||||
name: "valid email with dots and hyphens",
|
||||
keyID: "test.user-name@example-domain.co.uk",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid email with plus",
|
||||
keyID: "test+tag@example.com",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid short key ID (8 hex chars)",
|
||||
keyID: "ABCDEF12",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid long key ID (16 hex chars)",
|
||||
keyID: "ABCDEF1234567890",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid fingerprint (40 hex chars)",
|
||||
keyID: "ABCDEF1234567890ABCDEF1234567890ABCDEF12",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid lowercase hex fingerprint",
|
||||
keyID: "abcdef1234567890abcdef1234567890abcdef12",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid mixed case hex",
|
||||
keyID: "AbCdEf1234567890",
|
||||
wantErr: false,
|
||||
"valid lowercase hex fingerprint",
|
||||
"abcdef1234567890abcdef1234567890abcdef12", false,
|
||||
},
|
||||
{"valid mixed case hex", "AbCdEf1234567890", false},
|
||||
|
||||
// Invalid cases
|
||||
{"empty key ID", "", true},
|
||||
{"key ID with spaces", "test user@example.com", true},
|
||||
{"key ID with semicolon (command injection)", "test@example.com; rm -rf /", true},
|
||||
{
|
||||
name: "empty key ID",
|
||||
keyID: "",
|
||||
wantErr: true,
|
||||
"key ID with pipe (command injection)",
|
||||
"test@example.com | cat /etc/passwd", true,
|
||||
},
|
||||
{"key ID with backticks (command injection)", "test@example.com`whoami`", true},
|
||||
{
|
||||
name: "key ID with spaces",
|
||||
keyID: "test user@example.com",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with semicolon (command injection)",
|
||||
keyID: "test@example.com; rm -rf /",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with pipe (command injection)",
|
||||
keyID: "test@example.com | cat /etc/passwd",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with backticks (command injection)",
|
||||
keyID: "test@example.com`whoami`",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with dollar sign (command injection)",
|
||||
keyID: "test@example.com$(whoami)",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with quotes",
|
||||
keyID: "test\"@example.com",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with single quotes",
|
||||
keyID: "test'@example.com",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with backslash",
|
||||
keyID: "test\\@example.com",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with newline",
|
||||
keyID: "test@example.com\nrm -rf /",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with carriage return",
|
||||
keyID: "test@example.com\rrm -rf /",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "hex with invalid length (7 chars)",
|
||||
keyID: "ABCDEF1",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "hex with invalid length (9 chars)",
|
||||
keyID: "ABCDEF123",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "hex with non-hex characters",
|
||||
keyID: "ABCDEFGH",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "mixed format (email with hex)",
|
||||
keyID: "test@ABCDEF12",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with ampersand",
|
||||
keyID: "test@example.com & echo test",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with redirect",
|
||||
keyID: "test@example.com > /tmp/test",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key ID with null byte",
|
||||
keyID: "test@example.com\x00",
|
||||
wantErr: true,
|
||||
"key ID with dollar sign (command injection)",
|
||||
"test@example.com$(whoami)", true,
|
||||
},
|
||||
{"key ID with quotes", "test\"@example.com", true},
|
||||
{"key ID with single quotes", "test'@example.com", true},
|
||||
{"key ID with backslash", "test\\@example.com", true},
|
||||
{"key ID with newline", "test@example.com\nrm -rf /", true},
|
||||
{"key ID with carriage return", "test@example.com\rrm -rf /", true},
|
||||
{"hex with invalid length (7 chars)", "ABCDEF1", true},
|
||||
{"hex with invalid length (9 chars)", "ABCDEF123", true},
|
||||
{"hex with non-hex characters", "ABCDEFGH", true},
|
||||
{"mixed format (email with hex)", "test@ABCDEF12", true},
|
||||
{"key ID with ampersand", "test@example.com & echo test", true},
|
||||
{"key ID with redirect", "test@example.com > /tmp/test", true},
|
||||
{"key ID with null byte", "test@example.com\x00", true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err := validateGPGKeyID(tt.keyID)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("validateGPGKeyID() error = %v, wantErr %v", err, tt.wantErr)
|
||||
@@ -154,144 +67,3 @@ func TestValidateGPGKeyID(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateKeychainItemName(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
itemName string
|
||||
wantErr bool
|
||||
}{
|
||||
// Valid cases
|
||||
{
|
||||
name: "valid simple name",
|
||||
itemName: "my-secret-key",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid name with dots",
|
||||
itemName: "com.example.app.key",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid name with underscores",
|
||||
itemName: "my_secret_key_123",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid alphanumeric",
|
||||
itemName: "Secret123Key",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid with hyphen at start",
|
||||
itemName: "-my-key",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "valid with dot at start",
|
||||
itemName: ".hidden-key",
|
||||
wantErr: false,
|
||||
},
|
||||
|
||||
// Invalid cases
|
||||
{
|
||||
name: "empty item name",
|
||||
itemName: "",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with spaces",
|
||||
itemName: "my secret key",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with semicolon",
|
||||
itemName: "key;rm -rf /",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with pipe",
|
||||
itemName: "key|cat /etc/passwd",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with backticks",
|
||||
itemName: "key`whoami`",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with dollar sign",
|
||||
itemName: "key$(whoami)",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with quotes",
|
||||
itemName: "key\"name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with single quotes",
|
||||
itemName: "key'name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with backslash",
|
||||
itemName: "key\\name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with newline",
|
||||
itemName: "key\nname",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with carriage return",
|
||||
itemName: "key\rname",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with ampersand",
|
||||
itemName: "key&echo test",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with redirect",
|
||||
itemName: "key>/tmp/test",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with null byte",
|
||||
itemName: "key\x00name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with parentheses",
|
||||
itemName: "key(test)",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with brackets",
|
||||
itemName: "key[test]",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with asterisk",
|
||||
itemName: "key*",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "item name with question mark",
|
||||
itemName: "key?",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := validateKeychainItemName(tt.itemName)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("validateKeychainItemName() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+237
-111
@@ -2,9 +2,11 @@ package secret
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -20,12 +22,17 @@ const (
|
||||
maxVersionsPerDay = 999
|
||||
)
|
||||
|
||||
var (
|
||||
errMaxVersionsPerDay = errors.New("exceeded maximum versions per day (999)")
|
||||
errNilValueBuffer = errors.New("value buffer is nil")
|
||||
)
|
||||
|
||||
// VersionMetadata contains information about a secret version
|
||||
type VersionMetadata struct {
|
||||
ID string `json:"id"` // ULID
|
||||
CreatedAt *time.Time `json:"createdAt,omitempty"` // When version was created
|
||||
NotBefore *time.Time `json:"notBefore,omitempty"` // When this version becomes active
|
||||
NotAfter *time.Time `json:"notAfter,omitempty"` // When this version expires (nil = current)
|
||||
NotAfter *time.Time `json:"notAfter,omitempty"` // Expiry (nil = current)
|
||||
}
|
||||
|
||||
// Version represents a version of a secret
|
||||
@@ -75,7 +82,8 @@ func GenerateVersionName(fs afero.Fs, secretDir string) (string, error) {
|
||||
versionsDir := filepath.Join(secretDir, "versions")
|
||||
|
||||
// Ensure versions directory exists
|
||||
if err := fs.MkdirAll(versionsDir, DirPerms); err != nil {
|
||||
err := fs.MkdirAll(versionsDir, DirPerms)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to create versions directory: %w", err)
|
||||
}
|
||||
|
||||
@@ -101,7 +109,12 @@ func GenerateVersionName(fs afero.Fs, secretDir string) (string, error) {
|
||||
}
|
||||
|
||||
var serial int
|
||||
if _, err := fmt.Sscanf(parts[1], "%03d", &serial); err != nil {
|
||||
|
||||
_, err := fmt.Sscanf(parts[1], "%03d", &serial)
|
||||
if err != nil {
|
||||
Warn("Skipping malformed version directory name",
|
||||
"name", entry.Name(), "error", err)
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -113,16 +126,19 @@ func GenerateVersionName(fs afero.Fs, secretDir string) (string, error) {
|
||||
// Generate new version name
|
||||
newSerial := maxSerial + 1
|
||||
if newSerial > maxVersionsPerDay {
|
||||
return "", fmt.Errorf("exceeded maximum versions per day (999)")
|
||||
return "", errMaxVersionsPerDay
|
||||
}
|
||||
|
||||
return fmt.Sprintf("%s.%03d", today, newSerial), nil
|
||||
}
|
||||
|
||||
// Save saves the version metadata and value
|
||||
// Save saves the version metadata and value. The files are written into a
|
||||
// temporary directory that is renamed to sv.Directory once all of them are
|
||||
// complete, so the version directory is either whole or absent, even if the
|
||||
// process dies part-way.
|
||||
func (sv *Version) Save(value *memguard.LockedBuffer) error {
|
||||
if value == nil {
|
||||
return fmt.Errorf("value buffer is nil")
|
||||
return errNilValueBuffer
|
||||
}
|
||||
|
||||
DebugWith("Saving secret version",
|
||||
@@ -133,15 +149,25 @@ func (sv *Version) Save(value *memguard.LockedBuffer) error {
|
||||
|
||||
fs := sv.vault.GetFilesystem()
|
||||
|
||||
// Create version directory
|
||||
if err := fs.MkdirAll(sv.Directory, DirPerms); err != nil {
|
||||
Debug("Failed to create version directory", "error", err, "dir", sv.Directory)
|
||||
// Create the versions directory the finished version is renamed into
|
||||
err := fs.MkdirAll(filepath.Dir(sv.Directory), DirPerms)
|
||||
if err != nil {
|
||||
Debug("Failed to create versions directory", "error", err, "dir", sv.Directory)
|
||||
|
||||
return fmt.Errorf("failed to create version directory: %w", err)
|
||||
return fmt.Errorf("failed to create versions directory: %w", err)
|
||||
}
|
||||
|
||||
// Step 1: Generate a new keypair for this version
|
||||
tmpDir, err := TempDirFor(fs, sv.Directory)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Once the rename below has moved it into place, this finds nothing.
|
||||
defer func() { _ = fs.RemoveAll(tmpDir) }()
|
||||
|
||||
// Generate a new keypair for this version
|
||||
Debug("Generating version-specific keypair", "version", sv.Version)
|
||||
|
||||
versionIdentity, err := age.GenerateX25519Identity()
|
||||
if err != nil {
|
||||
Debug("Failed to generate version keypair", "error", err, "version", sv.Version)
|
||||
@@ -149,110 +175,40 @@ func (sv *Version) Save(value *memguard.LockedBuffer) error {
|
||||
return fmt.Errorf("failed to generate version keypair: %w", err)
|
||||
}
|
||||
|
||||
versionPublicKey := versionIdentity.Recipient().String()
|
||||
// Store private key in memguard buffer immediately
|
||||
versionPrivateKeyBuffer := memguard.NewBufferFromBytes([]byte(versionIdentity.String()))
|
||||
versionPrivateKeyBuffer := memguard.NewBufferFromBytes(
|
||||
[]byte(versionIdentity.String()))
|
||||
defer versionPrivateKeyBuffer.Destroy()
|
||||
|
||||
DebugWith("Generated version keypair",
|
||||
slog.String("version", sv.Version),
|
||||
slog.String("public_key", versionPublicKey),
|
||||
slog.String("public_key", versionIdentity.Recipient().String()),
|
||||
)
|
||||
|
||||
// Step 2: Store the version's public key
|
||||
pubKeyPath := filepath.Join(sv.Directory, "pub.age")
|
||||
Debug("Writing version public key", "path", pubKeyPath)
|
||||
if err := afero.WriteFile(fs, pubKeyPath, []byte(versionPublicKey), FilePerms); err != nil {
|
||||
Debug("Failed to write version public key", "error", err, "path", pubKeyPath)
|
||||
|
||||
return fmt.Errorf("failed to write version public key: %w", err)
|
||||
}
|
||||
|
||||
// Step 3: Encrypt the value to the version's public key
|
||||
Debug("Encrypting value to version's public key", "version", sv.Version)
|
||||
encryptedValue, err := EncryptToRecipient(value, versionIdentity.Recipient())
|
||||
err = sv.writePublicKeyAndValue(fs, tmpDir, versionIdentity, value)
|
||||
if err != nil {
|
||||
Debug("Failed to encrypt version value", "error", err, "version", sv.Version)
|
||||
|
||||
return fmt.Errorf("failed to encrypt version value: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
// Step 4: Store the encrypted value
|
||||
valuePath := filepath.Join(sv.Directory, "value.age")
|
||||
Debug("Writing encrypted version value", "path", valuePath)
|
||||
if err := afero.WriteFile(fs, valuePath, encryptedValue, FilePerms); err != nil {
|
||||
Debug("Failed to write encrypted version value", "error", err, "path", valuePath)
|
||||
|
||||
return fmt.Errorf("failed to write encrypted version value: %w", err)
|
||||
}
|
||||
|
||||
// Step 5: Get vault's long-term public key for encrypting the version's private key
|
||||
vaultDir, _ := sv.vault.GetDirectory()
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
Debug("Reading long-term public key", "path", ltPubKeyPath)
|
||||
|
||||
ltPubKeyData, err := afero.ReadFile(fs, ltPubKeyPath)
|
||||
err = sv.writeEncryptedPrivateKey(fs, tmpDir, versionPrivateKeyBuffer)
|
||||
if err != nil {
|
||||
Debug("Failed to read long-term public key", "error", err, "path", ltPubKeyPath)
|
||||
|
||||
return fmt.Errorf("failed to read long-term public key: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
Debug("Parsing long-term public key")
|
||||
ltRecipient, err := age.ParseX25519Recipient(string(ltPubKeyData))
|
||||
err = sv.writeEncryptedMetadata(fs, tmpDir, versionIdentity)
|
||||
if err != nil {
|
||||
Debug("Failed to parse long-term public key", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to parse long-term public key: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
// Step 6: Encrypt the version's private key to the long-term public key
|
||||
Debug("Encrypting version private key to long-term public key", "version", sv.Version)
|
||||
encryptedPrivKey, err := EncryptToRecipient(versionPrivateKeyBuffer, ltRecipient)
|
||||
err = fs.Rename(tmpDir, sv.Directory)
|
||||
if err != nil {
|
||||
Debug("Failed to encrypt version private key", "error", err, "version", sv.Version)
|
||||
Debug("Failed to move version into place", "error", err, "dir", sv.Directory)
|
||||
|
||||
return fmt.Errorf("failed to encrypt version private key: %w", err)
|
||||
return fmt.Errorf("failed to move version into place: %w", err)
|
||||
}
|
||||
|
||||
// Step 7: Store the encrypted private key
|
||||
privKeyPath := filepath.Join(sv.Directory, "priv.age")
|
||||
Debug("Writing encrypted version private key", "path", privKeyPath)
|
||||
if err := afero.WriteFile(fs, privKeyPath, encryptedPrivKey, FilePerms); err != nil {
|
||||
Debug("Failed to write encrypted version private key", "error", err, "path", privKeyPath)
|
||||
|
||||
return fmt.Errorf("failed to write encrypted version private key: %w", err)
|
||||
}
|
||||
|
||||
// Step 8: Encrypt and store metadata
|
||||
Debug("Encrypting version metadata", "version", sv.Version)
|
||||
metadataBytes, err := json.MarshalIndent(sv.Metadata, "", " ")
|
||||
if err != nil {
|
||||
Debug("Failed to marshal version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to marshal version metadata: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt metadata to the version's public key
|
||||
metadataBuffer := memguard.NewBufferFromBytes(metadataBytes)
|
||||
defer metadataBuffer.Destroy()
|
||||
|
||||
encryptedMetadata, err := EncryptToRecipient(metadataBuffer, versionIdentity.Recipient())
|
||||
if err != nil {
|
||||
Debug("Failed to encrypt version metadata", "error", err, "version", sv.Version)
|
||||
|
||||
return fmt.Errorf("failed to encrypt version metadata: %w", err)
|
||||
}
|
||||
|
||||
metadataPath := filepath.Join(sv.Directory, "metadata.age")
|
||||
Debug("Writing encrypted version metadata", "path", metadataPath)
|
||||
if err := afero.WriteFile(fs, metadataPath, encryptedMetadata, FilePerms); err != nil {
|
||||
Debug("Failed to write encrypted version metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return fmt.Errorf("failed to write encrypted version metadata: %w", err)
|
||||
}
|
||||
|
||||
Debug("Successfully saved secret version", "version", sv.Version, "secret_name", sv.SecretName)
|
||||
Debug("Successfully saved secret version",
|
||||
"version", sv.Version, "secret_name", sv.SecretName)
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -268,9 +224,11 @@ func (sv *Version) LoadMetadata(ltIdentity *age.X25519Identity) error {
|
||||
|
||||
// Step 1: Read encrypted version private key
|
||||
encryptedPrivKeyPath := filepath.Join(sv.Directory, "priv.age")
|
||||
|
||||
encryptedPrivKey, err := afero.ReadFile(fs, encryptedPrivKeyPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read encrypted version private key", "error", err, "path", encryptedPrivKeyPath)
|
||||
Debug("Failed to read encrypted version private key",
|
||||
"error", err, "path", encryptedPrivKeyPath)
|
||||
|
||||
return fmt.Errorf("failed to read encrypted version private key: %w", err)
|
||||
}
|
||||
@@ -294,9 +252,11 @@ func (sv *Version) LoadMetadata(ltIdentity *age.X25519Identity) error {
|
||||
|
||||
// Step 4: Read encrypted metadata
|
||||
encryptedMetadataPath := filepath.Join(sv.Directory, "metadata.age")
|
||||
|
||||
encryptedMetadata, err := afero.ReadFile(fs, encryptedMetadataPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read encrypted version metadata", "error", err, "path", encryptedMetadataPath)
|
||||
Debug("Failed to read encrypted version metadata",
|
||||
"error", err, "path", encryptedMetadataPath)
|
||||
|
||||
return fmt.Errorf("failed to read encrypted version metadata: %w", err)
|
||||
}
|
||||
@@ -312,20 +272,25 @@ func (sv *Version) LoadMetadata(ltIdentity *age.X25519Identity) error {
|
||||
|
||||
// Step 6: Unmarshal metadata
|
||||
var metadata VersionMetadata
|
||||
if err := json.Unmarshal(metadataBuffer.Bytes(), &metadata); err != nil {
|
||||
|
||||
err = json.Unmarshal(metadataBuffer.Bytes(), &metadata)
|
||||
if err != nil {
|
||||
Debug("Failed to unmarshal version metadata", "error", err, "version", sv.Version)
|
||||
|
||||
return fmt.Errorf("failed to unmarshal version metadata: %w", err)
|
||||
}
|
||||
|
||||
sv.Metadata = metadata
|
||||
|
||||
Debug("Successfully loaded version metadata", "version", sv.Version)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetValue retrieves and decrypts the version value
|
||||
func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuffer, error) {
|
||||
func (sv *Version) GetValue(
|
||||
ltIdentity *age.X25519Identity,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
DebugWith("Getting version value",
|
||||
slog.String("secret_name", sv.SecretName),
|
||||
slog.String("version", sv.Version),
|
||||
@@ -343,16 +308,22 @@ func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuf
|
||||
// Step 1: Read encrypted version private key
|
||||
encryptedPrivKeyPath := filepath.Join(sv.Directory, "priv.age")
|
||||
Debug("Reading encrypted version private key", "path", encryptedPrivKeyPath)
|
||||
|
||||
encryptedPrivKey, err := afero.ReadFile(fs, encryptedPrivKeyPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read encrypted version private key", "error", err, "path", encryptedPrivKeyPath)
|
||||
Debug("Failed to read encrypted version private key",
|
||||
"error", err, "path", encryptedPrivKeyPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read encrypted version private key: %w", err)
|
||||
return nil, fmt.Errorf(
|
||||
"failed to read encrypted version private key: %w", err)
|
||||
}
|
||||
Debug("Successfully read encrypted version private key", "path", encryptedPrivKeyPath, "size", len(encryptedPrivKey))
|
||||
|
||||
Debug("Successfully read encrypted version private key",
|
||||
"path", encryptedPrivKeyPath, "size", len(encryptedPrivKey))
|
||||
|
||||
// Step 2: Decrypt version private key using long-term key
|
||||
Debug("Decrypting version private key with long-term identity", "version", sv.Version)
|
||||
|
||||
versionPrivKeyBuffer, err := DecryptWithIdentity(encryptedPrivKey, ltIdentity)
|
||||
if err != nil {
|
||||
Debug("Failed to decrypt version private key", "error", err, "version", sv.Version)
|
||||
@@ -360,7 +331,9 @@ func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuf
|
||||
return nil, fmt.Errorf("failed to decrypt version private key: %w", err)
|
||||
}
|
||||
defer versionPrivKeyBuffer.Destroy()
|
||||
Debug("Successfully decrypted version private key", "version", sv.Version, "size", versionPrivKeyBuffer.Size())
|
||||
|
||||
Debug("Successfully decrypted version private key",
|
||||
"version", sv.Version, "size", versionPrivKeyBuffer.Size())
|
||||
|
||||
// Step 3: Parse version private key
|
||||
versionIdentity, err := age.ParseX25519Identity(versionPrivKeyBuffer.String())
|
||||
@@ -373,16 +346,21 @@ func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuf
|
||||
// Step 4: Read encrypted value
|
||||
encryptedValuePath := filepath.Join(sv.Directory, "value.age")
|
||||
Debug("Reading encrypted value", "path", encryptedValuePath)
|
||||
|
||||
encryptedValue, err := afero.ReadFile(fs, encryptedValuePath)
|
||||
if err != nil {
|
||||
Debug("Failed to read encrypted version value", "error", err, "path", encryptedValuePath)
|
||||
Debug("Failed to read encrypted version value",
|
||||
"error", err, "path", encryptedValuePath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read encrypted version value: %w", err)
|
||||
}
|
||||
Debug("Successfully read encrypted value", "path", encryptedValuePath, "size", len(encryptedValue))
|
||||
|
||||
Debug("Successfully read encrypted value",
|
||||
"path", encryptedValuePath, "size", len(encryptedValue))
|
||||
|
||||
// Step 5: Decrypt value using version key
|
||||
Debug("Decrypting value with version identity", "version", sv.Version)
|
||||
|
||||
valueBuffer, err := DecryptWithIdentity(encryptedValue, versionIdentity)
|
||||
if err != nil {
|
||||
Debug("Failed to decrypt version value", "error", err, "version", sv.Version)
|
||||
@@ -398,6 +376,142 @@ func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuf
|
||||
return valueBuffer, nil
|
||||
}
|
||||
|
||||
// writePublicKeyAndValue stores the version's public key and the value
|
||||
// encrypted to it in dir.
|
||||
func (sv *Version) writePublicKeyAndValue(
|
||||
fs afero.Fs,
|
||||
dir string,
|
||||
versionIdentity *age.X25519Identity,
|
||||
value *memguard.LockedBuffer,
|
||||
) error {
|
||||
versionPublicKey := versionIdentity.Recipient().String()
|
||||
pubKeyPath := filepath.Join(dir, "pub.age")
|
||||
Debug("Writing version public key", "path", pubKeyPath)
|
||||
|
||||
err := WriteFileAtomic(fs, pubKeyPath, []byte(versionPublicKey))
|
||||
if err != nil {
|
||||
Debug("Failed to write version public key", "error", err, "path", pubKeyPath)
|
||||
|
||||
return fmt.Errorf("failed to write version public key: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt the value to the version's public key
|
||||
Debug("Encrypting value to version's public key", "version", sv.Version)
|
||||
|
||||
encryptedValue, err := EncryptToRecipient(value, versionIdentity.Recipient())
|
||||
if err != nil {
|
||||
Debug("Failed to encrypt version value", "error", err, "version", sv.Version)
|
||||
|
||||
return fmt.Errorf("failed to encrypt version value: %w", err)
|
||||
}
|
||||
|
||||
valuePath := filepath.Join(dir, "value.age")
|
||||
Debug("Writing encrypted version value", "path", valuePath)
|
||||
|
||||
err = WriteFileAtomic(fs, valuePath, encryptedValue)
|
||||
if err != nil {
|
||||
Debug("Failed to write encrypted version value", "error", err, "path", valuePath)
|
||||
|
||||
return fmt.Errorf("failed to write encrypted version value: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeEncryptedPrivateKey encrypts the version's private key to the
|
||||
// vault's long-term public key and stores it in dir.
|
||||
func (sv *Version) writeEncryptedPrivateKey(
|
||||
fs afero.Fs,
|
||||
dir string,
|
||||
versionPrivateKeyBuffer *memguard.LockedBuffer,
|
||||
) error {
|
||||
vaultDir, _ := sv.vault.GetDirectory()
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
Debug("Reading long-term public key", "path", ltPubKeyPath)
|
||||
|
||||
ltPubKeyData, err := afero.ReadFile(fs, ltPubKeyPath)
|
||||
if err != nil {
|
||||
Debug("Failed to read long-term public key", "error", err, "path", ltPubKeyPath)
|
||||
|
||||
return fmt.Errorf("failed to read long-term public key: %w", err)
|
||||
}
|
||||
|
||||
Debug("Parsing long-term public key")
|
||||
|
||||
ltRecipient, err := age.ParseX25519Recipient(string(ltPubKeyData))
|
||||
if err != nil {
|
||||
Debug("Failed to parse long-term public key", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to parse long-term public key: %w", err)
|
||||
}
|
||||
|
||||
Debug("Encrypting version private key to long-term public key",
|
||||
"version", sv.Version)
|
||||
|
||||
encryptedPrivKey, err := EncryptToRecipient(versionPrivateKeyBuffer, ltRecipient)
|
||||
if err != nil {
|
||||
Debug("Failed to encrypt version private key",
|
||||
"error", err, "version", sv.Version)
|
||||
|
||||
return fmt.Errorf("failed to encrypt version private key: %w", err)
|
||||
}
|
||||
|
||||
privKeyPath := filepath.Join(dir, "priv.age")
|
||||
Debug("Writing encrypted version private key", "path", privKeyPath)
|
||||
|
||||
err = WriteFileAtomic(fs, privKeyPath, encryptedPrivKey)
|
||||
if err != nil {
|
||||
Debug("Failed to write encrypted version private key",
|
||||
"error", err, "path", privKeyPath)
|
||||
|
||||
return fmt.Errorf("failed to write encrypted version private key: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeEncryptedMetadata encrypts the version metadata to the version's
|
||||
// public key and stores it in dir.
|
||||
func (sv *Version) writeEncryptedMetadata(
|
||||
fs afero.Fs,
|
||||
dir string,
|
||||
versionIdentity *age.X25519Identity,
|
||||
) error {
|
||||
Debug("Encrypting version metadata", "version", sv.Version)
|
||||
|
||||
metadataBytes, err := json.MarshalIndent(sv.Metadata, "", " ")
|
||||
if err != nil {
|
||||
Debug("Failed to marshal version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to marshal version metadata: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt metadata to the version's public key
|
||||
metadataBuffer := memguard.NewBufferFromBytes(metadataBytes)
|
||||
defer metadataBuffer.Destroy()
|
||||
|
||||
encryptedMetadata, err := EncryptToRecipient(
|
||||
metadataBuffer, versionIdentity.Recipient())
|
||||
if err != nil {
|
||||
Debug("Failed to encrypt version metadata", "error", err, "version", sv.Version)
|
||||
|
||||
return fmt.Errorf("failed to encrypt version metadata: %w", err)
|
||||
}
|
||||
|
||||
metadataPath := filepath.Join(dir, "metadata.age")
|
||||
Debug("Writing encrypted version metadata", "path", metadataPath)
|
||||
|
||||
err = WriteFileAtomic(fs, metadataPath, encryptedMetadata)
|
||||
if err != nil {
|
||||
Debug("Failed to write encrypted version metadata",
|
||||
"error", err, "path", metadataPath)
|
||||
|
||||
return fmt.Errorf("failed to write encrypted version metadata: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListVersions lists all versions of a secret
|
||||
func ListVersions(fs afero.Fs, secretDir string) ([]string, error) {
|
||||
versionsDir := filepath.Join(secretDir, "versions")
|
||||
@@ -407,6 +521,7 @@ func ListVersions(fs afero.Fs, secretDir string) ([]string, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check versions directory: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return []string{}, nil
|
||||
}
|
||||
@@ -418,6 +533,7 @@ func ListVersions(fs afero.Fs, secretDir string) ([]string, error) {
|
||||
}
|
||||
|
||||
var versions []string
|
||||
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
versions = append(versions, entry.Name())
|
||||
@@ -430,6 +546,18 @@ func ListVersions(fs afero.Fs, secretDir string) ([]string, error) {
|
||||
return versions, nil
|
||||
}
|
||||
|
||||
// VersionExists reports whether version is one of the versions ListVersions
|
||||
// lists for the secret in secretDir. It only compares names, so a version
|
||||
// the user typed can be checked with it before any path is built from it.
|
||||
func VersionExists(fs afero.Fs, secretDir string, version string) (bool, error) {
|
||||
versions, err := ListVersions(fs, secretDir)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
return slices.Contains(versions, version), nil
|
||||
}
|
||||
|
||||
// GetCurrentVersion returns the version that the "current" file points to
|
||||
// The file contains just the version name (e.g., "20231215.001")
|
||||
func GetCurrentVersion(fs afero.Fs, secretDir string) (string, error) {
|
||||
@@ -446,15 +574,13 @@ func GetCurrentVersion(fs afero.Fs, secretDir string) (string, error) {
|
||||
}
|
||||
|
||||
// SetCurrentVersion updates the "current" file to point to a specific version
|
||||
// The file contains just the version name (e.g., "20231215.001")
|
||||
// The file contains just the version name (e.g., "20231215.001"). It is
|
||||
// replaced in one rename, so once written it always exists.
|
||||
func SetCurrentVersion(fs afero.Fs, secretDir string, version string) error {
|
||||
currentPath := filepath.Join(secretDir, "current")
|
||||
|
||||
// Remove existing file if it exists
|
||||
_ = fs.Remove(currentPath)
|
||||
|
||||
// Write just the version name to the file
|
||||
if err := afero.WriteFile(fs, currentPath, []byte(version), FilePerms); err != nil {
|
||||
err := WriteFileAtomic(fs, currentPath, []byte(version))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create current version file: %w", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -32,22 +32,32 @@
|
||||
// - Long-term key required for all operations
|
||||
// - Concurrent reads handled safely
|
||||
|
||||
package secret
|
||||
package secret_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// MockVault implements VaultInterface for testing
|
||||
const (
|
||||
testSecretDir = "/test/secret"
|
||||
testVaultName = "test"
|
||||
testVaultStateDir = "/test"
|
||||
)
|
||||
|
||||
var errNotImplementedInMock = errors.New("not implemented in mock")
|
||||
|
||||
// MockVersionVault implements VaultInterface for testing
|
||||
type MockVersionVault struct {
|
||||
Name string
|
||||
fs afero.Fs
|
||||
@@ -60,31 +70,41 @@ func (m *MockVersionVault) GetDirectory() (string, error) {
|
||||
}
|
||||
|
||||
func (m *MockVersionVault) AddSecret(_ string, _ *memguard.LockedBuffer, _ bool) error {
|
||||
return fmt.Errorf("not implemented in mock")
|
||||
return errNotImplementedInMock
|
||||
}
|
||||
|
||||
func (m *MockVersionVault) GetName() string {
|
||||
return m.Name
|
||||
}
|
||||
|
||||
//nolint:ireturn // implements VaultInterface
|
||||
func (m *MockVersionVault) GetFilesystem() afero.Fs {
|
||||
return m.fs
|
||||
}
|
||||
|
||||
func (m *MockVersionVault) GetCurrentUnlocker() (Unlocker, error) {
|
||||
return nil, fmt.Errorf("not implemented in mock")
|
||||
//nolint:ireturn // implements VaultInterface
|
||||
func (m *MockVersionVault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
return nil, errNotImplementedInMock
|
||||
}
|
||||
|
||||
func (m *MockVersionVault) CreatePassphraseUnlocker(_ *memguard.LockedBuffer) (*PassphraseUnlocker, error) {
|
||||
return nil, fmt.Errorf("not implemented in mock")
|
||||
func (m *MockVersionVault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
|
||||
return nil, errNotImplementedInMock
|
||||
}
|
||||
|
||||
func (m *MockVersionVault) CreatePassphraseUnlocker(
|
||||
_ *memguard.LockedBuffer,
|
||||
) (*secret.PassphraseUnlocker, error) {
|
||||
return nil, errNotImplementedInMock
|
||||
}
|
||||
|
||||
func TestGenerateVersionName(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
secretDir := "/test/secret"
|
||||
secretDir := testSecretDir
|
||||
|
||||
// Test first version generation
|
||||
version1, err := GenerateVersionName(fs, secretDir)
|
||||
version1, err := secret.GenerateVersionName(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Regexp(t, `^\d{8}\.001$`, version1)
|
||||
|
||||
@@ -94,7 +114,7 @@ func TestGenerateVersionName(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
// Test second version generation on same day
|
||||
version2, err := GenerateVersionName(fs, secretDir)
|
||||
version2, err := secret.GenerateVersionName(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Regexp(t, `^\d{8}\.002$`, version2)
|
||||
|
||||
@@ -104,8 +124,10 @@ func TestGenerateVersionName(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGenerateVersionNameMaxSerial(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
secretDir := "/test/secret"
|
||||
secretDir := testSecretDir
|
||||
versionsDir := filepath.Join(secretDir, "versions")
|
||||
|
||||
// Create 999 versions
|
||||
@@ -117,20 +139,22 @@ func TestGenerateVersionNameMaxSerial(t *testing.T) {
|
||||
}
|
||||
|
||||
// Try to create one more - should fail
|
||||
_, err := GenerateVersionName(fs, secretDir)
|
||||
assert.Error(t, err)
|
||||
_, err := secret.GenerateVersionName(fs, secretDir)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "exceeded maximum versions per day")
|
||||
}
|
||||
|
||||
func TestNewVersion(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
vault := &MockVersionVault{
|
||||
Name: "test",
|
||||
Name: testVaultName,
|
||||
fs: fs,
|
||||
stateDir: "/test",
|
||||
stateDir: testVaultStateDir,
|
||||
}
|
||||
|
||||
sv := NewVersion(vault, "test/secret", "20231215.001")
|
||||
sv := secret.NewVersion(vault, "test/secret", "20231215.001")
|
||||
|
||||
assert.Equal(t, "test/secret", sv.SecretName)
|
||||
assert.Equal(t, "20231215.001", sv.Version)
|
||||
@@ -140,11 +164,13 @@ func TestNewVersion(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestSecretVersionSave(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
vault := &MockVersionVault{
|
||||
Name: "test",
|
||||
Name: testVaultName,
|
||||
fs: fs,
|
||||
stateDir: "/test",
|
||||
stateDir: testVaultStateDir,
|
||||
}
|
||||
|
||||
// Create vault directory structure and long-term key
|
||||
@@ -155,18 +181,21 @@ func TestSecretVersionSave(t *testing.T) {
|
||||
// Generate and store long-term public key
|
||||
ltIdentity, err := age.GenerateX25519Identity()
|
||||
require.NoError(t, err)
|
||||
|
||||
vault.longTermKey = ltIdentity
|
||||
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
err = afero.WriteFile(
|
||||
fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create and save a version
|
||||
sv := NewVersion(vault, "test/secret", "20231215.001")
|
||||
sv := secret.NewVersion(vault, "test/secret", "20231215.001")
|
||||
testValue := []byte("test-secret-value")
|
||||
|
||||
testBuffer := memguard.NewBufferFromBytes(testValue)
|
||||
defer testBuffer.Destroy()
|
||||
|
||||
err = sv.Save(testBuffer)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -178,11 +207,13 @@ func TestSecretVersionSave(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestSecretVersionLoadMetadata(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
vault := &MockVersionVault{
|
||||
Name: "test",
|
||||
Name: testVaultName,
|
||||
fs: fs,
|
||||
stateDir: "/test",
|
||||
stateDir: testVaultStateDir,
|
||||
}
|
||||
|
||||
// Setup vault with long-term key
|
||||
@@ -192,14 +223,16 @@ func TestSecretVersionLoadMetadata(t *testing.T) {
|
||||
|
||||
ltIdentity, err := age.GenerateX25519Identity()
|
||||
require.NoError(t, err)
|
||||
|
||||
vault.longTermKey = ltIdentity
|
||||
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
err = afero.WriteFile(
|
||||
fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create and save a version with custom metadata
|
||||
sv := NewVersion(vault, "test/secret", "20231215.001")
|
||||
sv := secret.NewVersion(vault, "test/secret", "20231215.001")
|
||||
now := time.Now()
|
||||
epochPlusOne := time.Unix(1, 0)
|
||||
sv.Metadata.NotBefore = &epochPlusOne
|
||||
@@ -207,11 +240,12 @@ func TestSecretVersionLoadMetadata(t *testing.T) {
|
||||
|
||||
testBuffer := memguard.NewBufferFromBytes([]byte("test-value"))
|
||||
defer testBuffer.Destroy()
|
||||
|
||||
err = sv.Save(testBuffer)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create new version object and load metadata
|
||||
sv2 := NewVersion(vault, "test/secret", "20231215.001")
|
||||
sv2 := secret.NewVersion(vault, "test/secret", "20231215.001")
|
||||
err = sv2.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -223,11 +257,13 @@ func TestSecretVersionLoadMetadata(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestSecretVersionGetValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
vault := &MockVersionVault{
|
||||
Name: "test",
|
||||
Name: testVaultName,
|
||||
fs: fs,
|
||||
stateDir: "/test",
|
||||
stateDir: testVaultStateDir,
|
||||
}
|
||||
|
||||
// Setup vault with long-term key
|
||||
@@ -237,64 +273,77 @@ func TestSecretVersionGetValue(t *testing.T) {
|
||||
|
||||
ltIdentity, err := age.GenerateX25519Identity()
|
||||
require.NoError(t, err)
|
||||
|
||||
vault.longTermKey = ltIdentity
|
||||
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
err = afero.WriteFile(
|
||||
fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create and save a version
|
||||
sv := NewVersion(vault, "test/secret", "20231215.001")
|
||||
sv := secret.NewVersion(vault, "test/secret", "20231215.001")
|
||||
originalValue := []byte("test-secret-value-12345")
|
||||
expectedValue := make([]byte, len(originalValue))
|
||||
copy(expectedValue, originalValue)
|
||||
|
||||
originalBuffer := memguard.NewBufferFromBytes(originalValue)
|
||||
defer originalBuffer.Destroy()
|
||||
|
||||
err = sv.Save(originalBuffer)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Retrieve the value
|
||||
retrievedBuffer, err := sv.GetValue(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer retrievedBuffer.Destroy()
|
||||
|
||||
assert.Equal(t, expectedValue, retrievedBuffer.Bytes())
|
||||
}
|
||||
|
||||
func TestListVersions(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
secretDir := "/test/secret"
|
||||
secretDir := testSecretDir
|
||||
versionsDir := filepath.Join(secretDir, "versions")
|
||||
|
||||
// No versions directory
|
||||
versions, err := ListVersions(fs, secretDir)
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, versions)
|
||||
|
||||
// Create some versions
|
||||
testVersions := []string{"20231215.001", "20231215.002", "20231216.001", "20231214.001"}
|
||||
testVersions := []string{
|
||||
"20231215.001", "20231215.002", "20231216.001", "20231214.001",
|
||||
}
|
||||
for _, v := range testVersions {
|
||||
err := fs.MkdirAll(filepath.Join(versionsDir, v), 0o755)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// Create a file (not directory) that should be ignored
|
||||
err = afero.WriteFile(fs, filepath.Join(versionsDir, "ignore.txt"), []byte("test"), 0o600)
|
||||
err = afero.WriteFile(
|
||||
fs, filepath.Join(versionsDir, "ignore.txt"), []byte("test"), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// List versions
|
||||
versions, err = ListVersions(fs, secretDir)
|
||||
versions, err = secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Should be sorted in reverse chronological order
|
||||
expected := []string{"20231216.001", "20231215.002", "20231215.001", "20231214.001"}
|
||||
expected := []string{
|
||||
"20231216.001", "20231215.002", "20231215.001", "20231214.001",
|
||||
}
|
||||
assert.Equal(t, expected, versions)
|
||||
}
|
||||
|
||||
func TestGetCurrentVersion(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
secretDir := "/test/secret"
|
||||
secretDir := testSecretDir
|
||||
|
||||
// The current file contains just the version name
|
||||
currentPath := filepath.Join(secretDir, "current")
|
||||
@@ -304,39 +353,43 @@ func TestGetCurrentVersion(t *testing.T) {
|
||||
err = afero.WriteFile(fs, currentPath, []byte("20231216.001"), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
version, err := GetCurrentVersion(fs, secretDir)
|
||||
version, err := secret.GetCurrentVersion(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "20231216.001", version)
|
||||
}
|
||||
|
||||
func TestSetCurrentVersion(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
secretDir := "/test/secret"
|
||||
secretDir := testSecretDir
|
||||
|
||||
err := fs.MkdirAll(secretDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Set current version
|
||||
err = SetCurrentVersion(fs, secretDir, "20231216.002")
|
||||
err = secret.SetCurrentVersion(fs, secretDir, "20231216.002")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify it was set
|
||||
version, err := GetCurrentVersion(fs, secretDir)
|
||||
version, err := secret.GetCurrentVersion(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "20231216.002", version)
|
||||
|
||||
// Update to different version
|
||||
err = SetCurrentVersion(fs, secretDir, "20231217.001")
|
||||
err = secret.SetCurrentVersion(fs, secretDir, "20231217.001")
|
||||
require.NoError(t, err)
|
||||
|
||||
version, err = GetCurrentVersion(fs, secretDir)
|
||||
version, err = secret.GetCurrentVersion(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "20231217.001", version)
|
||||
}
|
||||
|
||||
func TestVersionMetadataTimestamps(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Test that all timestamp fields behave consistently as pointers
|
||||
vm := VersionMetadata{
|
||||
vm := secret.VersionMetadata{
|
||||
ID: "test-id",
|
||||
}
|
||||
|
||||
@@ -368,5 +421,6 @@ func TestVersionMetadataTimestamps(t *testing.T) {
|
||||
// Helper function
|
||||
func fileExists(fs afero.Fs, path string) bool {
|
||||
exists, _ := afero.Exists(fs, path)
|
||||
|
||||
return exists
|
||||
}
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
package vault
|
||||
|
||||
import "errors"
|
||||
|
||||
// Sentinel errors returned by vault operations.
|
||||
//
|
||||
// Several of these carry deliberately partial text: the message a caller
|
||||
// composes with fmt.Errorf places the interpolated value where it has
|
||||
// always appeared, and the sentinel supplies only the surrounding fixed
|
||||
// words. This keeps every composed message byte-identical to the dynamic
|
||||
// errors these sentinels replaced. Each such sentinel notes the message it
|
||||
// participates in.
|
||||
var (
|
||||
// ErrMnemonicMismatch indicates the mnemonic-derived public key does
|
||||
// not match the vault's stored public key hash.
|
||||
ErrMnemonicMismatch = errors.New(
|
||||
"derived public key does not match vault: mnemonic may be incorrect",
|
||||
)
|
||||
|
||||
// ErrInvalidVaultName indicates a vault name that breaks the naming
|
||||
// rule: only lowercase ASCII letters, digits, '.', '-' and '_'; not
|
||||
// empty, "." or "..". Composed by ValidateVaultName as
|
||||
// "invalid vault name '<name>': <the rule>".
|
||||
ErrInvalidVaultName = errors.New("invalid vault name")
|
||||
|
||||
// ErrVaultNotFound indicates the named vault does not exist. Composed
|
||||
// as "vault <name> does not exist".
|
||||
ErrVaultNotFound = errors.New("does not exist")
|
||||
|
||||
// ErrVaultExists indicates that a vault to be created already exists.
|
||||
// Composed as "vault <name> already exists".
|
||||
ErrVaultExists = errors.New("already exists")
|
||||
|
||||
// ErrNilValueBuffer indicates a nil value buffer was supplied.
|
||||
ErrNilValueBuffer = errors.New("value buffer is nil")
|
||||
|
||||
// ErrInvalidSecretName indicates a secret name that breaks the naming
|
||||
// rule: only ASCII letters, digits, '.', '-', '_' and '/'; not empty;
|
||||
// no leading '.' or '/', no trailing '/', no '//', no '..' path segment.
|
||||
// Composed by ValidateSecretName as
|
||||
// "invalid secret name '<name>': <the rule>".
|
||||
ErrInvalidSecretName = errors.New("invalid secret name")
|
||||
|
||||
// ErrSecretExists indicates the secret already exists and --force
|
||||
// was not supplied. Composed as
|
||||
// "secret <name> already exists (use --force to overwrite)", or as
|
||||
// "secret '<name>' already exists in vault '<vault>' (use --force to
|
||||
// overwrite)" when copying between vaults.
|
||||
ErrSecretExists = errors.New("already exists")
|
||||
|
||||
// ErrSecretNotFound indicates the named secret does not exist.
|
||||
// Composed as "secret <name> not found".
|
||||
ErrSecretNotFound = errors.New("not found")
|
||||
|
||||
// ErrVersionNotFound indicates the requested secret version does not
|
||||
// exist. Composed as
|
||||
// "version '<version>' not found for secret '<name>'".
|
||||
ErrVersionNotFound = errors.New("not found for secret")
|
||||
|
||||
// ErrNoVersions indicates the source secret has no versions. Composed
|
||||
// as "source secret '<name>' has no versions".
|
||||
ErrNoVersions = errors.New("has no versions")
|
||||
|
||||
// ErrUnsupportedUnlockerType indicates an unlocker metadata type
|
||||
// that this build does not support.
|
||||
ErrUnsupportedUnlockerType = errors.New("unsupported unlocker type")
|
||||
|
||||
// ErrUnlockerNotFound indicates no unlocker with the given ID exists.
|
||||
// Composed as "unlocker with ID <id> not found".
|
||||
ErrUnlockerNotFound = errors.New("not found")
|
||||
|
||||
// ErrNoLockForFilesystem indicates LockStateDir was given a filesystem
|
||||
// it cannot lock. Composed as "cannot lock the state directory on
|
||||
// filesystem <type>".
|
||||
ErrNoLockForFilesystem = errors.New(
|
||||
"cannot lock the state directory on filesystem")
|
||||
)
|
||||
+412
-369
@@ -1,10 +1,13 @@
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
@@ -12,6 +15,33 @@ import (
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// deriveVaultIdentity derives the long-term identity for the given vault
|
||||
// from testMnemonic using the derivation index stored in its metadata.
|
||||
func deriveVaultIdentity(
|
||||
t *testing.T, fs afero.Fs, vlt *vault.Vault,
|
||||
) *age.X25519Identity {
|
||||
t.Helper()
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic,
|
||||
vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
return ltIdentity
|
||||
}
|
||||
|
||||
//nolint:paralleltest // t.Setenv forbids parallel subtests
|
||||
func TestVaultWithRealFilesystem(t *testing.T) {
|
||||
// Create a temporary directory for our tests
|
||||
tempDir := t.TempDir()
|
||||
@@ -19,398 +49,411 @@ func TestVaultWithRealFilesystem(t *testing.T) {
|
||||
// Use the real filesystem
|
||||
fs := afero.NewOsFs()
|
||||
|
||||
// Test mnemonic
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
// Set test environment variables
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
// Test currentvault file handling (plain file with relative path)
|
||||
t.Run("CurrentVaultFileHandling", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "currentvault-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Get the vault directory
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Verify the currentvault file exists and contains just the vault name
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
currentVaultContents, err := os.ReadFile(currentVaultPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to read currentvault file: %v", err)
|
||||
}
|
||||
|
||||
expectedVaultName := "test-vault"
|
||||
if string(currentVaultContents) != expectedVaultName {
|
||||
t.Errorf("Expected currentvault to contain %q, got %q", expectedVaultName, string(currentVaultContents))
|
||||
}
|
||||
|
||||
// Test that ResolveVaultSymlink correctly resolves the path
|
||||
resolvedPath, err := vault.ResolveVaultSymlink(fs, currentVaultPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve currentvault path: %v", err)
|
||||
}
|
||||
|
||||
if resolvedPath != vaultDir {
|
||||
t.Errorf("Expected resolved path to be %s, got %s", vaultDir, resolvedPath)
|
||||
}
|
||||
testCurrentVaultFileHandling(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test secret operations with deeply nested paths
|
||||
t.Run("DeepPathSecrets", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "deep-path-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault - CreateVault now handles public key when mnemonic is in env
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get its derivation index
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic using the vault's derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vault
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Create a secret with a deeply nested path
|
||||
deepPath := "api/credentials/production/database/primary"
|
||||
secretValue := []byte("supersecretdbpassword")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vlt.AddSecret(deepPath, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret with deep path: %v", err)
|
||||
}
|
||||
|
||||
// List secrets and verify our deep path secret is there
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets: %v", err)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, s := range secrets {
|
||||
if s == deepPath {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Errorf("Deep path secret not found in listed secrets")
|
||||
}
|
||||
|
||||
// Retrieve the secret and verify its value
|
||||
retrievedValue, err := vlt.GetSecret(deepPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to retrieve deep path secret: %v", err)
|
||||
}
|
||||
|
||||
if string(retrievedValue) != string(expectedValue) {
|
||||
t.Errorf("Retrieved value doesn't match. Expected %q, got %q",
|
||||
string(expectedValue), string(retrievedValue))
|
||||
}
|
||||
testDeepPathSecrets(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test key caching in GetOrDeriveLongTermKey
|
||||
t.Run("KeyCaching", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "key-cache-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault - CreateVault now handles public key when mnemonic is in env
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get its derivation index
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic for verification using the vault's derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the vault is locked initially
|
||||
if !vlt.Locked() {
|
||||
t.Errorf("Vault should be locked initially")
|
||||
}
|
||||
|
||||
// First call to GetOrDeriveLongTermKey should derive and cache the key
|
||||
firstKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the vault is now unlocked
|
||||
if vlt.Locked() {
|
||||
t.Errorf("Vault should be unlocked after GetOrDeriveLongTermKey")
|
||||
}
|
||||
|
||||
// Second call should return the cached key without re-deriving
|
||||
secondKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get cached long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify both keys are the same instance
|
||||
if firstKey != secondKey {
|
||||
t.Errorf("Second key call should return same instance as first call")
|
||||
}
|
||||
|
||||
// Verify the public key matches what we expect
|
||||
expectedPubKey := ltIdentity.Recipient().String()
|
||||
actualPubKey := firstKey.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Public key mismatch. Expected %s, got %s", expectedPubKey, actualPubKey)
|
||||
}
|
||||
|
||||
// Now clear the key and verify it's locked again
|
||||
vlt.ClearLongTermKey()
|
||||
if !vlt.Locked() {
|
||||
t.Errorf("Vault should be locked after clearing key")
|
||||
}
|
||||
|
||||
// Get the key again and verify it works
|
||||
thirdKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to re-derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the public key still matches
|
||||
actualPubKey = thirdKey.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Re-derived public key mismatch. Expected %s, got %s", expectedPubKey, actualPubKey)
|
||||
}
|
||||
testKeyCaching(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test vault name validation
|
||||
t.Run("VaultNameValidation", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "name-validation-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Test valid vault names
|
||||
validNames := []string{
|
||||
"default",
|
||||
"test-vault",
|
||||
"production.vault",
|
||||
"vault_123",
|
||||
"a-very-long-vault-name-with-dashes",
|
||||
}
|
||||
|
||||
for _, name := range validNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Errorf("Failed to create vault with valid name %q: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Test invalid vault names
|
||||
invalidNames := []string{
|
||||
"", // Empty
|
||||
"UPPERCASE", // Uppercase not allowed
|
||||
"invalid/name", // Slashes not allowed in vault names
|
||||
"invalid name", // Spaces not allowed
|
||||
"invalid@name", // Special chars not allowed
|
||||
}
|
||||
|
||||
for _, name := range invalidNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err == nil {
|
||||
t.Errorf("Expected error creating vault with invalid name %q, but got none", name)
|
||||
}
|
||||
}
|
||||
testVaultNameValidation(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test multiple vaults and switching between them
|
||||
t.Run("MultipleVaults", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "multi-vault-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create three vaults
|
||||
vaultNames := []string{"vault1", "vault2", "vault3"}
|
||||
for _, name := range vaultNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// List vaults and verify all three are there
|
||||
vaults, err := vault.ListVaults(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list vaults: %v", err)
|
||||
}
|
||||
|
||||
if len(vaults) != 3 {
|
||||
t.Errorf("Expected 3 vaults, got %d", len(vaults))
|
||||
}
|
||||
|
||||
// Test switching between vaults
|
||||
for _, name := range vaultNames {
|
||||
// Select the vault
|
||||
if err := vault.SelectVault(fs, stateDir, name); err != nil {
|
||||
t.Fatalf("Failed to select vault %s: %v", name, err)
|
||||
}
|
||||
|
||||
// Get current vault and verify it's the one we selected
|
||||
currentVault, err := vault.GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault after selecting %s: %v", name, err)
|
||||
}
|
||||
|
||||
if currentVault.GetName() != name {
|
||||
t.Errorf("Expected current vault to be %s, got %s", name, currentVault.GetName())
|
||||
}
|
||||
}
|
||||
testMultipleVaults(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test adding a secret in one vault and verifying it's not visible in another
|
||||
// Test adding a secret in one vault and verifying it's not visible in
|
||||
// another
|
||||
t.Run("VaultIsolation", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "isolation-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create two vaults - CreateVault now handles public key when mnemonic is in env
|
||||
vault1, err := vault.CreateVault(fs, stateDir, "vault1")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault1: %v", err)
|
||||
}
|
||||
|
||||
vault2, err := vault.CreateVault(fs, stateDir, "vault2")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault2: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic
|
||||
// Note: Both vaults will have different derivation indexes due to GetNextDerivationIndex
|
||||
|
||||
// Load vault1 metadata to get its derivation index
|
||||
vault1Dir, err := vault1.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault1 directory: %v", err)
|
||||
}
|
||||
vault1Metadata, err := vault.LoadVaultMetadata(fs, vault1Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault1 metadata: %v", err)
|
||||
}
|
||||
|
||||
ltIdentity1, err := agehd.DeriveIdentity(testMnemonic, vault1Metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key for vault1: %v", err)
|
||||
}
|
||||
|
||||
// Load vault2 metadata to get its derivation index
|
||||
vault2Dir, err := vault2.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault2 directory: %v", err)
|
||||
}
|
||||
vault2Metadata, err := vault.LoadVaultMetadata(fs, vault2Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault2 metadata: %v", err)
|
||||
}
|
||||
|
||||
ltIdentity2, err := agehd.DeriveIdentity(testMnemonic, vault2Metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key for vault2: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vaults with their respective keys
|
||||
vault1.Unlock(ltIdentity1)
|
||||
vault2.Unlock(ltIdentity2)
|
||||
|
||||
// Add a secret to vault1
|
||||
secretName := "test-secret"
|
||||
secretValue := []byte("secret in vault1")
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
if err := vault1.AddSecret(secretName, secretBuffer, false); err != nil {
|
||||
t.Fatalf("Failed to add secret to vault1: %v", err)
|
||||
}
|
||||
|
||||
// Verify the secret exists in vault1
|
||||
vault1Secrets, err := vault1.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets in vault1: %v", err)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, s := range vault1Secrets {
|
||||
if s == secretName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Errorf("Secret not found in vault1")
|
||||
}
|
||||
|
||||
// Verify the secret does NOT exist in vault2
|
||||
vault2Secrets, err := vault2.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets in vault2: %v", err)
|
||||
}
|
||||
|
||||
found = false
|
||||
for _, s := range vault2Secrets {
|
||||
if s == secretName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if found {
|
||||
t.Errorf("Secret from vault1 should not be visible in vault2")
|
||||
}
|
||||
testVaultIsolation(t, fs, tempDir)
|
||||
})
|
||||
}
|
||||
|
||||
func testCurrentVaultFileHandling(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "currentvault-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault
|
||||
vlt, err := vault.CreateVault(fs, stateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Get the vault directory
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Verify the currentvault file exists and contains just the vault name
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
|
||||
currentVaultContents, err := os.ReadFile(filepath.Clean(currentVaultPath))
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to read currentvault file: %v", err)
|
||||
}
|
||||
|
||||
if string(currentVaultContents) != testVaultName {
|
||||
t.Errorf("Expected currentvault to contain %q, got %q",
|
||||
testVaultName, string(currentVaultContents))
|
||||
}
|
||||
|
||||
// Test that ResolveVaultSymlink correctly resolves the path
|
||||
resolvedPath, err := vault.ResolveVaultSymlink(fs, currentVaultPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve currentvault path: %v", err)
|
||||
}
|
||||
|
||||
if resolvedPath != vaultDir {
|
||||
t.Errorf("Expected resolved path to be %s, got %s", vaultDir, resolvedPath)
|
||||
}
|
||||
}
|
||||
|
||||
func testDeepPathSecrets(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "deep-path-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault - CreateVault now handles public key when
|
||||
// mnemonic is in env
|
||||
vlt, err := vault.CreateVault(fs, stateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get its derivation index
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic using the vault's derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic,
|
||||
vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vault
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Create a secret with a deeply nested path
|
||||
deepPath := "api/credentials/production/database/primary"
|
||||
secretValue := []byte("supersecretdbpassword")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vlt.AddSecret(deepPath, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret with deep path: %v", err)
|
||||
}
|
||||
|
||||
// List secrets and verify our deep path secret is there
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(secrets, deepPath) {
|
||||
t.Errorf("Deep path secret not found in listed secrets")
|
||||
}
|
||||
|
||||
// Retrieve the secret and verify its value
|
||||
retrievedValue, err := vlt.GetSecret(deepPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to retrieve deep path secret: %v", err)
|
||||
}
|
||||
defer retrievedValue.Destroy()
|
||||
|
||||
if !bytes.Equal(retrievedValue.Bytes(), expectedValue) {
|
||||
t.Errorf("Retrieved value doesn't match. Expected %q, got %q",
|
||||
expectedValue, retrievedValue.Bytes())
|
||||
}
|
||||
}
|
||||
|
||||
func testKeyCaching(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "key-cache-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault - CreateVault now handles public key when
|
||||
// mnemonic is in env
|
||||
vlt, err := vault.CreateVault(fs, stateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get its derivation index
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic for verification using the
|
||||
// vault's derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic,
|
||||
vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the vault is locked initially
|
||||
if !vlt.Locked() {
|
||||
t.Errorf("Vault should be locked initially")
|
||||
}
|
||||
|
||||
// First call to GetOrDeriveLongTermKey should derive and cache the key
|
||||
firstKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the vault is now unlocked
|
||||
if vlt.Locked() {
|
||||
t.Errorf("Vault should be unlocked after GetOrDeriveLongTermKey")
|
||||
}
|
||||
|
||||
// Second call should return the cached key without re-deriving
|
||||
secondKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get cached long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify both keys are the same instance
|
||||
if firstKey != secondKey {
|
||||
t.Errorf("Second key call should return same instance as first call")
|
||||
}
|
||||
|
||||
// Verify the public key matches what we expect
|
||||
expectedPubKey := ltIdentity.Recipient().String()
|
||||
|
||||
actualPubKey := firstKey.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Public key mismatch. Expected %s, got %s",
|
||||
expectedPubKey, actualPubKey)
|
||||
}
|
||||
|
||||
// Now clear the key and verify it's locked again
|
||||
vlt.ClearLongTermKey()
|
||||
|
||||
if !vlt.Locked() {
|
||||
t.Errorf("Vault should be locked after clearing key")
|
||||
}
|
||||
|
||||
// Get the key again and verify it works
|
||||
thirdKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to re-derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the public key still matches
|
||||
actualPubKey = thirdKey.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Re-derived public key mismatch. Expected %s, got %s",
|
||||
expectedPubKey, actualPubKey)
|
||||
}
|
||||
}
|
||||
|
||||
func testVaultNameValidation(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "name-validation-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Test valid vault names
|
||||
validNames := []string{
|
||||
"default",
|
||||
"test-vault",
|
||||
"production.vault",
|
||||
"vault_123",
|
||||
"a-very-long-vault-name-with-dashes",
|
||||
}
|
||||
|
||||
for _, name := range validNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Errorf("Failed to create vault with valid name %q: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Test invalid vault names
|
||||
invalidNames := []string{
|
||||
"", // Empty
|
||||
"UPPERCASE", // Uppercase not allowed
|
||||
"invalid/name", // Slashes not allowed in vault names
|
||||
"invalid name", // Spaces not allowed
|
||||
"invalid@name", // Special chars not allowed
|
||||
}
|
||||
|
||||
for _, name := range invalidNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err == nil {
|
||||
t.Errorf("Expected error creating vault with invalid name %q, "+
|
||||
"but got none", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func testMultipleVaults(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "multi-vault-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create three vaults
|
||||
vaultNames := []string{"vault1", "vault2", "vault3"}
|
||||
for _, name := range vaultNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// List vaults and verify all three are there
|
||||
vaults, err := vault.ListVaults(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list vaults: %v", err)
|
||||
}
|
||||
|
||||
if len(vaults) != 3 {
|
||||
t.Errorf("Expected 3 vaults, got %d", len(vaults))
|
||||
}
|
||||
|
||||
// Test switching between vaults
|
||||
for _, name := range vaultNames {
|
||||
// Select the vault
|
||||
err := vault.SelectVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select vault %s: %v", name, err)
|
||||
}
|
||||
|
||||
// Get current vault and verify it's the one we selected
|
||||
currentVault, err := vault.GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault after selecting %s: %v",
|
||||
name, err)
|
||||
}
|
||||
|
||||
if currentVault.GetName() != name {
|
||||
t.Errorf("Expected current vault to be %s, got %s",
|
||||
name, currentVault.GetName())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func testVaultIsolation(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "isolation-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create two vaults - CreateVault now handles public key when mnemonic
|
||||
// is in env
|
||||
vault1, err := vault.CreateVault(fs, stateDir, "vault1")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault1: %v", err)
|
||||
}
|
||||
|
||||
vault2, err := vault.CreateVault(fs, stateDir, "vault2")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault2: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term keys from mnemonic
|
||||
// Note: Both vaults will have different derivation indexes due to
|
||||
// GetNextDerivationIndex
|
||||
ltIdentity1 := deriveVaultIdentity(t, fs, vault1)
|
||||
ltIdentity2 := deriveVaultIdentity(t, fs, vault2)
|
||||
|
||||
// Unlock the vaults with their respective keys
|
||||
vault1.Unlock(ltIdentity1)
|
||||
vault2.Unlock(ltIdentity2)
|
||||
|
||||
// Add a secret to vault1
|
||||
secretValue := []byte("secret in vault1")
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vault1.AddSecret(testSecretName, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret to vault1: %v", err)
|
||||
}
|
||||
|
||||
// Verify the secret exists in vault1
|
||||
vault1Secrets, err := vault1.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets in vault1: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(vault1Secrets, testSecretName) {
|
||||
t.Errorf("Secret not found in vault1")
|
||||
}
|
||||
|
||||
// Verify the secret does NOT exist in vault2
|
||||
vault2Secrets, err := vault2.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets in vault2: %v", err)
|
||||
}
|
||||
|
||||
if slices.Contains(vault2Secrets, testSecretName) {
|
||||
t.Errorf("Secret from vault1 should not be visible in vault2")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -19,14 +19,17 @@
|
||||
// - Consistent test mnemonic for reproducible keys
|
||||
// - Proper cleanup and isolation between tests
|
||||
|
||||
//nolint:testpackage // uses white-box test helpers shared with this package
|
||||
package vault
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/awnumar/memguard"
|
||||
@@ -35,38 +38,33 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// Helper function to add a secret to vault with proper buffer protection
|
||||
func addTestSecret(t *testing.T, vault *Vault, name string, value []byte, force bool) {
|
||||
t.Helper()
|
||||
buffer := memguard.NewBufferFromBytes(value)
|
||||
defer buffer.Destroy()
|
||||
err := vault.AddSecret(name, buffer, force)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
// errUnexpectedValue is returned by concurrent readers when a secret value
|
||||
// does not match the expected contents.
|
||||
var errUnexpectedValue = errors.New("unexpected value")
|
||||
|
||||
// TestVersionIntegrationWorkflow tests the complete version workflow
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids parallel subtests
|
||||
func TestVersionIntegrationWorkflow(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set mnemonic for testing
|
||||
t.Setenv(secret.EnvMnemonic,
|
||||
"abandon abandon abandon abandon abandon abandon "+
|
||||
"abandon abandon abandon abandon abandon about")
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
// Create vault
|
||||
vault, err := CreateVault(fs, stateDir, "test")
|
||||
vault, err := CreateVault(fs, testStateDir, "test")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Derive and store long-term key from mnemonic
|
||||
mnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Store long-term public key in vault
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
|
||||
err = afero.WriteFile(fs, ltPubKeyPath,
|
||||
[]byte(ltIdentity.Recipient().String()), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Unlock the vault
|
||||
@@ -76,225 +74,315 @@ func TestVersionIntegrationWorkflow(t *testing.T) {
|
||||
|
||||
// Step 1: Create initial version
|
||||
t.Run("create_initial_version", func(t *testing.T) {
|
||||
addTestSecret(t, vault, secretName, []byte("version-1-data"), false)
|
||||
|
||||
// Verify secret can be retrieved
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value)
|
||||
|
||||
// Verify version directory structure
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 1)
|
||||
|
||||
// Verify current symlink exists
|
||||
currentVersion, err := secret.GetCurrentVersion(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, versions[0], currentVersion)
|
||||
|
||||
// Verify metadata
|
||||
version := secret.NewVersion(vault, secretName, versions[0])
|
||||
err = version.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, version.Metadata.CreatedAt)
|
||||
assert.NotNil(t, version.Metadata.NotBefore)
|
||||
assert.Equal(t, int64(1), version.Metadata.NotBefore.Unix()) // epoch + 1
|
||||
assert.Nil(t, version.Metadata.NotAfter) // should be nil for current version
|
||||
testCreateInitialVersion(t, fs, vault, ltIdentity, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 2: Create second version
|
||||
var firstVersionName string
|
||||
t.Run("create_second_version", func(t *testing.T) {
|
||||
// Small delay to ensure different timestamps
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
// Get first version name before creating second
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
firstVersionName = versions[0]
|
||||
|
||||
// Create second version
|
||||
addTestSecret(t, vault, secretName, []byte("version-2-data"), true)
|
||||
|
||||
// Verify new value is current
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2-data"), value)
|
||||
|
||||
// Verify we now have two versions
|
||||
versions, err = secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 2)
|
||||
|
||||
// Verify first version metadata was updated with notAfter
|
||||
firstVersion := secret.NewVersion(vault, secretName, firstVersionName)
|
||||
err = firstVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, firstVersion.Metadata.NotAfter)
|
||||
|
||||
// Verify second version metadata
|
||||
secondVersion := secret.NewVersion(vault, secretName, versions[0])
|
||||
err = secondVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, secondVersion.Metadata.NotBefore)
|
||||
assert.Nil(t, secondVersion.Metadata.NotAfter)
|
||||
|
||||
// NotBefore of second should equal NotAfter of first
|
||||
assert.Equal(t, firstVersion.Metadata.NotAfter.Unix(), secondVersion.Metadata.NotBefore.Unix())
|
||||
testCreateSecondVersion(t, fs, vault, ltIdentity, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 3: Create third version
|
||||
t.Run("create_third_version", func(t *testing.T) {
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecret(t, vault, secretName, []byte("version-3-data"), true)
|
||||
|
||||
// Verify we now have three versions
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 3)
|
||||
|
||||
// Current should be version-3
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), value)
|
||||
testCreateThirdVersion(t, fs, vault, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 4: Retrieve specific versions
|
||||
t.Run("retrieve_specific_versions", func(t *testing.T) {
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 3)
|
||||
|
||||
// Get each version by its name
|
||||
value1, err := vault.GetSecretVersion(secretName, versions[2]) // oldest
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value1)
|
||||
|
||||
value2, err := vault.GetSecretVersion(secretName, versions[1]) // middle
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2-data"), value2)
|
||||
|
||||
value3, err := vault.GetSecretVersion(secretName, versions[0]) // newest
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), value3)
|
||||
|
||||
// Empty version should return current
|
||||
valueCurrent, err := vault.GetSecretVersion(secretName, "")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), valueCurrent)
|
||||
testRetrieveSpecificVersions(t, fs, vault, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 5: Promote old version to current
|
||||
t.Run("promote_old_version", func(t *testing.T) {
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Promote the first version (oldest) to current
|
||||
oldestVersion := versions[2]
|
||||
err = secret.SetCurrentVersion(fs, secretDir, oldestVersion)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify current now returns the old version's value
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value)
|
||||
|
||||
// Verify the version metadata hasn't changed
|
||||
// (promoting shouldn't modify timestamps)
|
||||
version := secret.NewVersion(vault, secretName, oldestVersion)
|
||||
err = version.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, version.Metadata.NotAfter) // should still have its old notAfter
|
||||
testPromoteOldVersion(t, fs, vault, ltIdentity, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 6: Test version limits
|
||||
t.Run("version_serial_limits", func(t *testing.T) {
|
||||
// Create a new secret for this test
|
||||
limitSecretName := "limit/test"
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "limit%test", "versions")
|
||||
|
||||
// Create 998 versions (we already have one from the first AddSecret)
|
||||
addTestSecret(t, vault, limitSecretName, []byte("initial"), false)
|
||||
|
||||
// Get today's date for consistent version names
|
||||
today := time.Now().Format("20060102")
|
||||
|
||||
// Manually create many versions with same date
|
||||
for i := 2; i <= 998; i++ {
|
||||
versionName := fmt.Sprintf("%s.%03d", today, i)
|
||||
versionDir := filepath.Join(secretDir, versionName)
|
||||
err := fs.MkdirAll(versionDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// Should be able to create one more (999)
|
||||
versionName, err := secret.GenerateVersionName(fs, filepath.Dir(secretDir))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, fmt.Sprintf("%s.999", today), versionName)
|
||||
|
||||
// Create the 999th version directory
|
||||
err = fs.MkdirAll(filepath.Join(secretDir, versionName), 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Should fail to create 1000th version
|
||||
_, err = secret.GenerateVersionName(fs, filepath.Dir(secretDir))
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "exceeded maximum versions per day")
|
||||
testVersionSerialLimits(t, fs, vault, vaultDir)
|
||||
})
|
||||
|
||||
// Step 7: Test error cases
|
||||
t.Run("error_cases", func(t *testing.T) {
|
||||
// Try to get non-existent version
|
||||
_, err := vault.GetSecretVersion(secretName, "99991231.999")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "not found")
|
||||
|
||||
// Try to get version of non-existent secret
|
||||
_, err = vault.GetSecretVersion("nonexistent/secret", "")
|
||||
assert.Error(t, err)
|
||||
|
||||
// Try to add secret without force when it exists
|
||||
failBuffer := memguard.NewBufferFromBytes([]byte("should-fail"))
|
||||
defer failBuffer.Destroy()
|
||||
err = vault.AddSecret(secretName, failBuffer, false)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already exists")
|
||||
testVersionErrorCases(t, vault, secretName)
|
||||
})
|
||||
}
|
||||
|
||||
func testCreateInitialVersion(
|
||||
t *testing.T, fs afero.Fs, vault *Vault,
|
||||
ltIdentity *age.X25519Identity, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-1-data"), false)
|
||||
|
||||
// Verify secret can be retrieved
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-1-data"), value.Bytes())
|
||||
|
||||
// Verify version directory structure
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 1)
|
||||
|
||||
// Verify current symlink exists
|
||||
currentVersion, err := secret.GetCurrentVersion(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, versions[0], currentVersion)
|
||||
|
||||
// Verify metadata
|
||||
version := secret.NewVersion(vault, secretName, versions[0])
|
||||
err = version.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, version.Metadata.CreatedAt)
|
||||
assert.NotNil(t, version.Metadata.NotBefore)
|
||||
assert.Equal(t, int64(1), version.Metadata.NotBefore.Unix()) // epoch + 1
|
||||
// NotAfter should be nil for current version
|
||||
assert.Nil(t, version.Metadata.NotAfter)
|
||||
}
|
||||
|
||||
func testCreateSecondVersion(
|
||||
t *testing.T, fs afero.Fs, vault *Vault,
|
||||
ltIdentity *age.X25519Identity, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
// Small delay to ensure different timestamps
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
// Get first version name before creating second
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
firstVersionName := versions[0]
|
||||
|
||||
// Create second version
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-2-data"), true)
|
||||
|
||||
// Verify new value is current
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-2-data"), value.Bytes())
|
||||
|
||||
// Verify we now have two versions
|
||||
versions, err = secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 2)
|
||||
|
||||
// Verify first version metadata was updated with notAfter
|
||||
firstVersion := secret.NewVersion(vault, secretName, firstVersionName)
|
||||
err = firstVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, firstVersion.Metadata.NotAfter)
|
||||
|
||||
// Verify second version metadata
|
||||
secondVersion := secret.NewVersion(vault, secretName, versions[0])
|
||||
err = secondVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, secondVersion.Metadata.NotBefore)
|
||||
assert.Nil(t, secondVersion.Metadata.NotAfter)
|
||||
|
||||
// NotBefore of second should equal NotAfter of first
|
||||
assert.Equal(t, firstVersion.Metadata.NotAfter.Unix(),
|
||||
secondVersion.Metadata.NotBefore.Unix())
|
||||
}
|
||||
|
||||
func testCreateThirdVersion(
|
||||
t *testing.T, fs afero.Fs, vault *Vault, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-3-data"), true)
|
||||
|
||||
// Verify we now have three versions
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 3)
|
||||
|
||||
// Current should be version-3
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-3-data"), value.Bytes())
|
||||
}
|
||||
|
||||
func testRetrieveSpecificVersions(
|
||||
t *testing.T, fs afero.Fs, vault *Vault, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 3)
|
||||
|
||||
// Get each version by its name
|
||||
value1, err := vault.GetSecretVersion(secretName, versions[2]) // oldest
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value1.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-1-data"), value1.Bytes())
|
||||
|
||||
value2, err := vault.GetSecretVersion(secretName, versions[1]) // middle
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value2.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-2-data"), value2.Bytes())
|
||||
|
||||
value3, err := vault.GetSecretVersion(secretName, versions[0]) // newest
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value3.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-3-data"), value3.Bytes())
|
||||
|
||||
// An empty version is not one of the versions; GetSecret gets the
|
||||
// current one
|
||||
_, err = vault.GetSecretVersion(secretName, "")
|
||||
require.ErrorIs(t, err, ErrVersionNotFound)
|
||||
}
|
||||
|
||||
func testPromoteOldVersion(
|
||||
t *testing.T, fs afero.Fs, vault *Vault,
|
||||
ltIdentity *age.X25519Identity, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Promote the first version (oldest) to current
|
||||
oldestVersion := versions[2]
|
||||
err = secret.SetCurrentVersion(fs, secretDir, oldestVersion)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify current now returns the old version's value
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-1-data"), value.Bytes())
|
||||
|
||||
// Verify the version metadata hasn't changed
|
||||
// (promoting shouldn't modify timestamps)
|
||||
version := secret.NewVersion(vault, secretName, oldestVersion)
|
||||
err = version.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
// should still have its old notAfter
|
||||
assert.NotNil(t, version.Metadata.NotAfter)
|
||||
}
|
||||
|
||||
func testVersionSerialLimits(
|
||||
t *testing.T, fs afero.Fs, vault *Vault, vaultDir string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
// Create a new secret for this test
|
||||
limitSecretName := "limit/test"
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "limit%test", "versions")
|
||||
|
||||
// Create 998 versions (we already have one from the first AddSecret)
|
||||
addTestSecretToVault(t, vault, limitSecretName, []byte("initial"), false)
|
||||
|
||||
// Get today's date for consistent version names
|
||||
today := time.Now().Format("20060102")
|
||||
|
||||
// Manually create many versions with same date
|
||||
for i := 2; i <= 998; i++ {
|
||||
versionName := fmt.Sprintf("%s.%03d", today, i)
|
||||
versionDir := filepath.Join(secretDir, versionName)
|
||||
err := fs.MkdirAll(versionDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// Should be able to create one more (999)
|
||||
versionName, err := secret.GenerateVersionName(fs, filepath.Dir(secretDir))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, today+".999", versionName)
|
||||
|
||||
// Create the 999th version directory
|
||||
err = fs.MkdirAll(filepath.Join(secretDir, versionName), 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Should fail to create 1000th version
|
||||
_, err = secret.GenerateVersionName(fs, filepath.Dir(secretDir))
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "exceeded maximum versions per day")
|
||||
}
|
||||
|
||||
func testVersionErrorCases(t *testing.T, vault *Vault, secretName string) {
|
||||
t.Helper()
|
||||
|
||||
// Try to get non-existent version
|
||||
_, err := vault.GetSecretVersion(secretName, "99991231.999")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "not found")
|
||||
|
||||
// Try to get version of non-existent secret
|
||||
_, err = vault.GetSecretVersion("nonexistent/secret", "")
|
||||
require.Error(t, err)
|
||||
|
||||
// Try to add secret without force when it exists
|
||||
failBuffer := memguard.NewBufferFromBytes([]byte("should-fail"))
|
||||
defer failBuffer.Destroy()
|
||||
|
||||
err = vault.AddSecret(secretName, failBuffer, false)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already exists")
|
||||
}
|
||||
|
||||
// TestVersionConcurrency tests concurrent version operations
|
||||
//
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVersionConcurrency(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set up vault
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
secretName := "concurrent/test"
|
||||
|
||||
// Create initial version
|
||||
addTestSecret(t, vault, secretName, []byte("initial"), false)
|
||||
addTestSecretToVault(t, vault, secretName, []byte("initial"), false)
|
||||
|
||||
// Test concurrent reads
|
||||
t.Run("concurrent_reads", func(t *testing.T) {
|
||||
done := make(chan bool, 10)
|
||||
errors := make(chan error, 10)
|
||||
errCh := make(chan error, 10)
|
||||
|
||||
for range 10 {
|
||||
go func() {
|
||||
value, err := vault.GetSecret(secretName)
|
||||
if err != nil {
|
||||
errors <- err
|
||||
} else if string(value) != "initial" {
|
||||
errors <- fmt.Errorf("unexpected value: %s", value)
|
||||
errCh <- err
|
||||
} else {
|
||||
if value.String() != "initial" {
|
||||
errCh <- fmt.Errorf("%w: %s",
|
||||
errUnexpectedValue, value.Bytes())
|
||||
}
|
||||
|
||||
value.Destroy()
|
||||
}
|
||||
|
||||
done <- true
|
||||
}()
|
||||
}
|
||||
@@ -306,7 +394,7 @@ func TestVersionConcurrency(t *testing.T) {
|
||||
|
||||
// Check for errors
|
||||
select {
|
||||
case err := <-errors:
|
||||
case err := <-errCh:
|
||||
t.Fatalf("concurrent read failed: %v", err)
|
||||
default:
|
||||
// No errors
|
||||
@@ -315,12 +403,14 @@ func TestVersionConcurrency(t *testing.T) {
|
||||
}
|
||||
|
||||
// TestVersionCompatibility tests that old secrets without versions still work
|
||||
//
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVersionCompatibility(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set up vault
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
ltIdentity, err := vault.GetOrDeriveLongTermKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -333,9 +423,12 @@ func TestVersionCompatibility(t *testing.T) {
|
||||
|
||||
// Create old-style encrypted value directly in secret directory
|
||||
testValue := []byte("legacy-value")
|
||||
|
||||
testValueBuffer := memguard.NewBufferFromBytes(testValue)
|
||||
defer testValueBuffer.Destroy()
|
||||
|
||||
ltRecipient := ltIdentity.Recipient()
|
||||
|
||||
encrypted, err := secret.EncryptToRecipient(testValueBuffer, ltRecipient)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -345,7 +438,7 @@ func TestVersionCompatibility(t *testing.T) {
|
||||
|
||||
// Should fail to get with version-aware methods
|
||||
_, err = vault.GetSecret(secretName)
|
||||
assert.Error(t, err)
|
||||
require.Error(t, err)
|
||||
|
||||
// List versions should return empty
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
package vault
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// lockFileName is the file in the state directory that LockStateDir locks.
|
||||
const lockFileName = "lock"
|
||||
|
||||
// memFsLock stands in for the lock file on the in-memory filesystem, which
|
||||
// has no file locks. Every in-memory filesystem in the process shares it.
|
||||
//
|
||||
//nolint:gochecknoglobals // must outlive the call that takes it
|
||||
var memFsLock sync.Mutex
|
||||
|
||||
// LockStateDir takes the lock that a command changing anything under
|
||||
// stateDir holds until it returns, and returns the function that releases
|
||||
// it. While one command holds it, the next one waits here. Reads take no
|
||||
// lock: each file or directory a command changes is replaced in a single
|
||||
// rename, so a reader finds it as it was before or after, never half-made.
|
||||
//
|
||||
// On the real filesystem the lock is flock(2) on the file "lock" in
|
||||
// stateDir, which the kernel releases when the process dies, so a killed
|
||||
// command never leaves the tool locked. The in-memory filesystem the tests
|
||||
// use has no file locks, so a process-wide mutex stands in for flock there.
|
||||
// Any other filesystem is refused rather than left unlocked.
|
||||
func LockStateDir(fs afero.Fs, stateDir string) (func(), error) {
|
||||
switch fs.(type) {
|
||||
case *afero.OsFs:
|
||||
return flockStateDir(stateDir)
|
||||
case *afero.MemMapFs:
|
||||
memFsLock.Lock()
|
||||
|
||||
return memFsLock.Unlock, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("%w %T", ErrNoLockForFilesystem, fs)
|
||||
}
|
||||
}
|
||||
|
||||
// flockStateDir takes flock(2) on the lock file in stateDir, creating the
|
||||
// directory and the file if needed. Go opens files close-on-exec, so
|
||||
// programs the command runs, such as gpg, do not inherit the lock.
|
||||
func flockStateDir(stateDir string) (func(), error) {
|
||||
err := os.MkdirAll(stateDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create state directory: %w", err)
|
||||
}
|
||||
|
||||
lockPath := filepath.Join(stateDir, lockFileName)
|
||||
|
||||
//nolint:gosec // G304: the path is the lock file in the state directory
|
||||
file, err := os.OpenFile(lockPath, os.O_RDWR|os.O_CREATE, secret.FilePerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to open lock file: %w", err)
|
||||
}
|
||||
|
||||
err = syscall.Flock(int(file.Fd()), syscall.LOCK_EX)
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
|
||||
return nil, fmt.Errorf("failed to lock %s: %w", lockPath, err)
|
||||
}
|
||||
|
||||
// Closing the file releases the lock.
|
||||
return func() { _ = file.Close() }, nil
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const (
|
||||
// lockWait is how long a test waits for the lock before deciding it
|
||||
// will never come free.
|
||||
lockWait = 10 * time.Second
|
||||
|
||||
// heldWait is how long a test watches a second holder fail to take a
|
||||
// lock that is held. Broken exclusion lets it in at once.
|
||||
heldWait = 100 * time.Millisecond
|
||||
)
|
||||
|
||||
// lockFilesystem is a filesystem LockStateDir can lock, with a state
|
||||
// directory on it.
|
||||
type lockFilesystem struct {
|
||||
name string
|
||||
fs afero.Fs
|
||||
stateDir string
|
||||
}
|
||||
|
||||
// lockFilesystems returns the real filesystem, locked with flock, and the
|
||||
// in-memory one, locked with a mutex.
|
||||
func lockFilesystems(t *testing.T) []lockFilesystem {
|
||||
t.Helper()
|
||||
|
||||
return []lockFilesystem{
|
||||
{"memory", afero.NewMemMapFs(), testStateDir},
|
||||
{"real", afero.NewOsFs(), t.TempDir()},
|
||||
}
|
||||
}
|
||||
|
||||
// lockInBackground starts taking the lock and returns a channel that
|
||||
// delivers the function releasing it once it has been taken.
|
||||
func lockInBackground(
|
||||
t *testing.T, fs afero.Fs, stateDir string,
|
||||
) <-chan func() {
|
||||
t.Helper()
|
||||
|
||||
taken := make(chan func(), 1)
|
||||
|
||||
go func() {
|
||||
release, err := vault.LockStateDir(fs, stateDir)
|
||||
if assert.NoError(t, err) {
|
||||
taken <- release
|
||||
}
|
||||
}()
|
||||
|
||||
return taken
|
||||
}
|
||||
|
||||
// TestLockStateDirExcludes checks that while the lock is held a second
|
||||
// holder, with its own open lock file on the real filesystem, waits, and
|
||||
// that it gets the lock once the first releases it.
|
||||
func TestLockStateDirExcludes(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, lfs := range lockFilesystems(t) {
|
||||
t.Run(lfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
release, err := vault.LockStateDir(lfs.fs, lfs.stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
taken := lockInBackground(t, lfs.fs, lfs.stateDir)
|
||||
|
||||
select {
|
||||
case second := <-taken:
|
||||
second()
|
||||
release()
|
||||
t.Fatal("a second holder took the lock while it was held")
|
||||
case <-time.After(heldWait):
|
||||
}
|
||||
|
||||
release()
|
||||
|
||||
select {
|
||||
case second := <-taken:
|
||||
second()
|
||||
case <-time.After(lockWait):
|
||||
t.Fatal("the second holder never got the lock")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLockStateDirFreeAfterPanic checks that a holder that panics, and
|
||||
// releases the lock with defer as every command does, leaves it free.
|
||||
func TestLockStateDirFreeAfterPanic(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, lfs := range lockFilesystems(t) {
|
||||
t.Run(lfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
assert.Panics(t, func() {
|
||||
release, err := vault.LockStateDir(lfs.fs, lfs.stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer release()
|
||||
|
||||
panic("the command failed")
|
||||
})
|
||||
|
||||
select {
|
||||
case release := <-lockInBackground(t, lfs.fs, lfs.stateDir):
|
||||
release()
|
||||
case <-time.After(lockWait):
|
||||
t.Fatal("the lock was still held after its holder panicked")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLockStateDirRefusesOtherFilesystems checks that a filesystem with no
|
||||
// lock implementation is refused instead of being used unlocked.
|
||||
func TestLockStateDirRefusesOtherFilesystems(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewReadOnlyFs(afero.NewMemMapFs())
|
||||
|
||||
release, err := vault.LockStateDir(fs, testStateDir)
|
||||
require.ErrorIs(t, err, vault.ErrNoLockForFilesystem)
|
||||
assert.Nil(t, release)
|
||||
}
|
||||
+101
-41
@@ -15,23 +15,44 @@ import (
|
||||
)
|
||||
|
||||
// Register the GetCurrentVault function with the secret package
|
||||
//
|
||||
//nolint:gochecknoinits // registers the vault accessor with the secret package
|
||||
func init() {
|
||||
secret.RegisterGetCurrentVaultFunc(func(fs afero.Fs, stateDir string) (secret.VaultInterface, error) {
|
||||
return GetCurrentVault(fs, stateDir)
|
||||
})
|
||||
secret.RegisterGetCurrentVaultFunc(
|
||||
func(fs afero.Fs, stateDir string) (secret.VaultInterface, error) {
|
||||
return GetCurrentVault(fs, stateDir)
|
||||
})
|
||||
}
|
||||
|
||||
// isValidVaultName validates vault names according to the format [a-z0-9\.\-\_]+
|
||||
// Note: We don't allow slashes in vault names unlike secret names
|
||||
// isValidVaultName reports whether name is a valid vault name: only
|
||||
// lowercase ASCII letters, digits, '.', '-' and '_', and not empty, "." or
|
||||
// "..". With no path separator allowed, a vault is always one directory
|
||||
// directly under vaults.d.
|
||||
func isValidVaultName(name string) bool {
|
||||
if name == "" {
|
||||
if name == "" || name == "." || name == ".." {
|
||||
return false
|
||||
}
|
||||
|
||||
matched, _ := regexp.MatchString(`^[a-z0-9\.\-\_]+$`, name)
|
||||
|
||||
return matched
|
||||
}
|
||||
|
||||
// ValidateVaultName returns an error wrapping ErrInvalidVaultName when name
|
||||
// is not a valid vault name. Call it on the name exactly as the user gave it,
|
||||
// before building any path from it.
|
||||
func ValidateVaultName(name string) error {
|
||||
if !isValidVaultName(name) {
|
||||
return fmt.Errorf(
|
||||
"%w '%s': only lowercase ASCII letters, digits, '.', '-' and '_' "+
|
||||
"are allowed, and a name must not be empty, '.' or '..'",
|
||||
ErrInvalidVaultName, name,
|
||||
)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ResolveVaultSymlink reads the currentvault file to get the path to the current vault
|
||||
// The file contains just the vault name (e.g., "default")
|
||||
func ResolveVaultSymlink(fs afero.Fs, currentVaultPath string) (string, error) {
|
||||
@@ -65,9 +86,11 @@ func GetCurrentVault(fs afero.Fs, stateDir string) (*Vault, error) {
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
|
||||
secret.Debug("Checking current vault symlink", "path", currentVaultPath)
|
||||
|
||||
_, err := fs.Stat(currentVaultPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to stat current vault symlink", "error", err, "path", currentVaultPath)
|
||||
secret.Debug("Failed to stat current vault symlink",
|
||||
"error", err, "path", currentVaultPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read current vault symlink: %w", err)
|
||||
}
|
||||
@@ -76,6 +99,7 @@ func GetCurrentVault(fs afero.Fs, stateDir string) (*Vault, error) {
|
||||
|
||||
// Resolve the symlink to get the actual vault directory
|
||||
secret.Debug("Resolving vault symlink")
|
||||
|
||||
targetPath, err := ResolveVaultSymlink(fs, currentVaultPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -88,7 +112,8 @@ func GetCurrentVault(fs afero.Fs, stateDir string) (*Vault, error) {
|
||||
vaultName := filepath.Base(targetPath)
|
||||
secret.Debug("Extracted vault name", "vault_name", vaultName)
|
||||
|
||||
secret.Debug("Current vault resolved", "vault_name", vaultName, "target_path", targetPath)
|
||||
secret.Debug("Current vault resolved",
|
||||
"vault_name", vaultName, "target_path", targetPath)
|
||||
|
||||
// Create and return the vault
|
||||
return NewVault(fs, stateDir, vaultName), nil
|
||||
@@ -103,6 +128,7 @@ func ListVaults(fs afero.Fs, stateDir string) ([]string, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if vaults directory exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return []string{}, nil
|
||||
}
|
||||
@@ -115,6 +141,7 @@ func ListVaults(fs afero.Fs, stateDir string) ([]string, error) {
|
||||
|
||||
// Extract vault names
|
||||
var vaults []string
|
||||
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
vaults = append(vaults, entry.Name())
|
||||
@@ -124,22 +151,26 @@ func ListVaults(fs afero.Fs, stateDir string) ([]string, error) {
|
||||
return vaults, nil
|
||||
}
|
||||
|
||||
// processMnemonicForVault handles mnemonic processing for vault creation
|
||||
func processMnemonicForVault(fs afero.Fs, stateDir, vaultDir, vaultName string) (
|
||||
derivationIndex uint32, publicKeyHash string, familyHash string, err error) {
|
||||
// processMnemonicForVault handles mnemonic processing for vault creation.
|
||||
// It returns the derivation index, public key hash, and family hash.
|
||||
func processMnemonicForVault(
|
||||
fs afero.Fs, stateDir, vaultDir, vaultName string,
|
||||
) (uint32, string, string, error) {
|
||||
// Check if mnemonic is available in environment
|
||||
mnemonic := os.Getenv(secret.EnvMnemonic)
|
||||
|
||||
if mnemonic == "" {
|
||||
secret.Debug("No mnemonic in environment, vault created without long-term key", "vault", vaultName)
|
||||
secret.Debug("No mnemonic in environment, vault created without long-term key",
|
||||
"vault", vaultName)
|
||||
// Use 0 for derivation index when no mnemonic is provided
|
||||
return 0, "", "", nil
|
||||
}
|
||||
|
||||
secret.Debug("Mnemonic found in environment, deriving long-term key", "vault", vaultName)
|
||||
secret.Debug("Mnemonic found in environment, deriving long-term key",
|
||||
"vault", vaultName)
|
||||
|
||||
// Get the next available derivation index for this mnemonic
|
||||
derivationIndex, err = GetNextDerivationIndex(fs, stateDir, mnemonic)
|
||||
derivationIndex, err := GetNextDerivationIndex(fs, stateDir, mnemonic)
|
||||
if err != nil {
|
||||
return 0, "", "", fmt.Errorf("failed to get next derivation index: %w", err)
|
||||
}
|
||||
@@ -152,14 +183,18 @@ func processMnemonicForVault(fs afero.Fs, stateDir, vaultDir, vaultName string)
|
||||
|
||||
// Write the public key
|
||||
ltPubKey := ltIdentity.Recipient().String()
|
||||
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
if err := afero.WriteFile(fs, ltPubKeyPath, []byte(ltPubKey), secret.FilePerms); err != nil {
|
||||
|
||||
err = secret.WriteFileAtomic(fs, ltPubKeyPath, []byte(ltPubKey))
|
||||
if err != nil {
|
||||
return 0, "", "", fmt.Errorf("failed to write long-term public key: %w", err)
|
||||
}
|
||||
|
||||
secret.Debug("Wrote long-term public key", "path", ltPubKeyPath)
|
||||
|
||||
// Compute verification hash from actual derivation index
|
||||
publicKeyHash = ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
publicKeyHash := ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
|
||||
// Compute family hash from index 0 (same for all vaults with this mnemonic)
|
||||
// This is used to identify which vaults belong to the same mnemonic family
|
||||
@@ -167,46 +202,68 @@ func processMnemonicForVault(fs afero.Fs, stateDir, vaultDir, vaultName string)
|
||||
if err != nil {
|
||||
return 0, "", "", fmt.Errorf("failed to derive identity for index 0: %w", err)
|
||||
}
|
||||
familyHash = ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
familyHash := ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
return derivationIndex, publicKeyHash, familyHash, nil
|
||||
}
|
||||
|
||||
// CreateVault creates a new vault
|
||||
// CreateVault creates a new vault and selects it as the current vault. It
|
||||
// refuses a vault that already exists before writing anything: creating it
|
||||
// again would replace its keys, and its secrets could no longer be
|
||||
// decrypted. The commands that call it hold the state directory lock, so no
|
||||
// other command can create the vault between the check and the writes.
|
||||
func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) {
|
||||
secret.Debug("Creating new vault", "name", name, "state_dir", stateDir)
|
||||
|
||||
// Validate vault name
|
||||
if !isValidVaultName(name) {
|
||||
err := ValidateVaultName(name)
|
||||
if err != nil {
|
||||
secret.Debug("Invalid vault name provided", "vault_name", name)
|
||||
|
||||
return nil, fmt.Errorf("invalid vault name '%s': must match pattern [a-z0-9.\\-_]+", name)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
secret.Debug("Vault name validation passed", "vault_name", name)
|
||||
|
||||
// Create vault directory structure
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", name)
|
||||
|
||||
exists, err := afero.DirExists(fs, vaultDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if vault exists: %w", err)
|
||||
}
|
||||
|
||||
if exists {
|
||||
return nil, fmt.Errorf("vault %s %w", name, ErrVaultExists)
|
||||
}
|
||||
|
||||
// Create vault directory structure
|
||||
secret.Debug("Creating vault directory structure", "vault_dir", vaultDir)
|
||||
|
||||
// Create main vault directory
|
||||
if err := fs.MkdirAll(vaultDir, secret.DirPerms); err != nil {
|
||||
err = fs.MkdirAll(vaultDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Create secrets directory
|
||||
secretsDir := filepath.Join(vaultDir, "secrets.d")
|
||||
if err := fs.MkdirAll(secretsDir, secret.DirPerms); err != nil {
|
||||
|
||||
err = fs.MkdirAll(secretsDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create secrets directory: %w", err)
|
||||
}
|
||||
|
||||
// Create unlockers directory
|
||||
unlockersDir := filepath.Join(vaultDir, "unlockers.d")
|
||||
if err := fs.MkdirAll(unlockersDir, secret.DirPerms); err != nil {
|
||||
|
||||
err = fs.MkdirAll(unlockersDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create unlockers directory: %w", err)
|
||||
}
|
||||
|
||||
// Process mnemonic if available
|
||||
derivationIndex, publicKeyHash, familyHash, err := processMnemonicForVault(fs, stateDir, vaultDir, name)
|
||||
derivationIndex, publicKeyHash, familyHash, err := processMnemonicForVault(
|
||||
fs, stateDir, vaultDir, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -218,13 +275,17 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) {
|
||||
PublicKeyHash: publicKeyHash,
|
||||
MnemonicFamilyHash: familyHash,
|
||||
}
|
||||
if err := SaveVaultMetadata(fs, vaultDir, metadata); err != nil {
|
||||
|
||||
err = SaveVaultMetadata(fs, vaultDir, metadata)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to save vault metadata: %w", err)
|
||||
}
|
||||
|
||||
// Select the newly created vault as current
|
||||
secret.Debug("Selecting newly created vault as current", "name", name)
|
||||
if err := SelectVault(fs, stateDir, name); err != nil {
|
||||
|
||||
err = SelectVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to select vault: %w", err)
|
||||
}
|
||||
|
||||
@@ -238,36 +299,35 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) {
|
||||
func SelectVault(fs afero.Fs, stateDir string, name string) error {
|
||||
secret.Debug("Selecting vault", "vault_name", name, "state_dir", stateDir)
|
||||
|
||||
// Validate vault name
|
||||
if !isValidVaultName(name) {
|
||||
err := ValidateVaultName(name)
|
||||
if err != nil {
|
||||
secret.Debug("Invalid vault name provided", "vault_name", name)
|
||||
|
||||
return fmt.Errorf("invalid vault name '%s': must match pattern [a-z0-9.\\-_]+", name)
|
||||
return err
|
||||
}
|
||||
|
||||
secret.Debug("Vault name validation passed", "vault_name", name)
|
||||
|
||||
// Check if vault exists
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", name)
|
||||
|
||||
exists, err := afero.DirExists(fs, vaultDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if vault exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("vault %s does not exist", name)
|
||||
return fmt.Errorf("vault %s %w", name, ErrVaultNotFound)
|
||||
}
|
||||
|
||||
// Create or update the currentvault file with just the vault name
|
||||
// Create or replace the currentvault file with just the vault name. It
|
||||
// is replaced in one rename, so it never goes missing.
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
|
||||
// Remove existing file if it exists
|
||||
if _, err := fs.Stat(currentVaultPath); err == nil {
|
||||
secret.Debug("Removing existing currentvault file", "path", currentVaultPath)
|
||||
_ = fs.Remove(currentVaultPath)
|
||||
}
|
||||
|
||||
// Write just the vault name to the file
|
||||
secret.Debug("Writing currentvault file", "vault_name", name)
|
||||
if err := afero.WriteFile(fs, currentVaultPath, []byte(name), secret.FilePerms); err != nil {
|
||||
|
||||
err = secret.WriteFileAtomic(fs, currentVaultPath, []byte(name))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to select vault: %w", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -34,12 +34,15 @@ func ComputeDoubleSHA256(data []byte) string {
|
||||
|
||||
// GetNextDerivationIndex finds the next available derivation index for a given mnemonic
|
||||
// by deriving the public key for index 0 and using its hash to identify related vaults
|
||||
func GetNextDerivationIndex(fs afero.Fs, stateDir string, mnemonic string) (uint32, error) {
|
||||
func GetNextDerivationIndex(
|
||||
fs afero.Fs, stateDir string, mnemonic string,
|
||||
) (uint32, error) {
|
||||
// First, derive the public key for index 0 to get our identifier
|
||||
identity0, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to derive identity for index 0: %w", err)
|
||||
}
|
||||
|
||||
pubKeyHash := ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
vaultsDir := filepath.Join(stateDir, "vaults.d")
|
||||
@@ -49,6 +52,7 @@ func GetNextDerivationIndex(fs afero.Fs, stateDir string, mnemonic string) (uint
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to check if vaults directory exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
// No vaults yet, start with index 0
|
||||
return 0, nil
|
||||
@@ -70,6 +74,7 @@ func GetNextDerivationIndex(fs afero.Fs, stateDir string, mnemonic string) (uint
|
||||
|
||||
// Try to read vault metadata
|
||||
metadataPath := filepath.Join(vaultsDir, entry.Name(), "vault-metadata.json")
|
||||
|
||||
metadataBytes, err := afero.ReadFile(fs, metadataPath)
|
||||
if err != nil {
|
||||
// Skip vaults without metadata
|
||||
@@ -77,7 +82,9 @@ func GetNextDerivationIndex(fs afero.Fs, stateDir string, mnemonic string) (uint
|
||||
}
|
||||
|
||||
var metadata Metadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
// Skip vaults with invalid metadata
|
||||
continue
|
||||
}
|
||||
@@ -106,7 +113,8 @@ func SaveVaultMetadata(fs afero.Fs, vaultDir string, metadata *Metadata) error {
|
||||
return fmt.Errorf("failed to marshal vault metadata: %w", err)
|
||||
}
|
||||
|
||||
if err := afero.WriteFile(fs, metadataPath, metadataBytes, secret.FilePerms); err != nil {
|
||||
err = secret.WriteFileAtomic(fs, metadataPath, metadataBytes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write vault metadata: %w", err)
|
||||
}
|
||||
|
||||
@@ -123,7 +131,9 @@ func LoadVaultMetadata(fs afero.Fs, vaultDir string) (*Metadata, error) {
|
||||
}
|
||||
|
||||
var metadata Metadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal vault metadata: %w", err)
|
||||
}
|
||||
|
||||
|
||||
+258
-212
@@ -1,208 +1,243 @@
|
||||
package vault
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
//nolint:paralleltest // subtests share an in-memory filesystem sequentially
|
||||
func TestVaultMetadata(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Test mnemonic for consistent testing
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
t.Run("ComputeDoubleSHA256", func(t *testing.T) {
|
||||
// Test data
|
||||
data := []byte("test data")
|
||||
hash := ComputeDoubleSHA256(data)
|
||||
|
||||
// Verify it's a valid hex string of 64 characters (32 bytes * 2)
|
||||
if len(hash) != 64 {
|
||||
t.Errorf("Expected hash length of 64, got %d", len(hash))
|
||||
}
|
||||
|
||||
// Verify consistency
|
||||
hash2 := ComputeDoubleSHA256(data)
|
||||
if hash != hash2 {
|
||||
t.Errorf("Hash should be consistent for same input")
|
||||
}
|
||||
|
||||
// Verify different input produces different hash
|
||||
hash3 := ComputeDoubleSHA256([]byte("different data"))
|
||||
if hash == hash3 {
|
||||
t.Errorf("Different input should produce different hash")
|
||||
}
|
||||
testComputeDoubleSHA256(t)
|
||||
})
|
||||
|
||||
t.Run("GetNextDerivationIndex", func(t *testing.T) {
|
||||
// Test with no existing vaults
|
||||
index, err := GetNextDerivationIndex(fs, stateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
if index != 0 {
|
||||
t.Errorf("Expected index 0 for first vault, got %d", index)
|
||||
}
|
||||
|
||||
// Create a vault with metadata and matching public key
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "vault1")
|
||||
if err := fs.MkdirAll(vaultDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Derive identity for index 0
|
||||
identity0, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity: %v", err)
|
||||
}
|
||||
pubKey0 := identity0.Recipient().String()
|
||||
pubKeyHash0 := ComputeDoubleSHA256([]byte(pubKey0))
|
||||
|
||||
// Write public key
|
||||
if err := afero.WriteFile(fs, filepath.Join(vaultDir, "pub.age"), []byte(pubKey0), 0o600); err != nil {
|
||||
t.Fatalf("Failed to write public key: %v", err)
|
||||
}
|
||||
|
||||
metadata1 := &Metadata{
|
||||
DerivationIndex: 0,
|
||||
PublicKeyHash: pubKeyHash0, // Hash of the actual key (index 0)
|
||||
MnemonicFamilyHash: pubKeyHash0, // Hash of index 0 key (for family identification)
|
||||
}
|
||||
if err := SaveVaultMetadata(fs, vaultDir, metadata1); err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
// Next index for same mnemonic should be 1
|
||||
index, err = GetNextDerivationIndex(fs, stateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
if index != 1 {
|
||||
t.Errorf("Expected index 1 for second vault with same mnemonic, got %d", index)
|
||||
}
|
||||
|
||||
// Different mnemonic should start at 0
|
||||
differentMnemonic := "zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo wrong"
|
||||
index, err = GetNextDerivationIndex(fs, stateDir, differentMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
if index != 0 {
|
||||
t.Errorf("Expected index 0 for first vault with different mnemonic, got %d", index)
|
||||
}
|
||||
|
||||
// Add another vault with same mnemonic but higher index
|
||||
vaultDir2 := filepath.Join(stateDir, "vaults.d", "vault2")
|
||||
if err := fs.MkdirAll(vaultDir2, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Derive identity for index 5
|
||||
identity5, err := agehd.DeriveIdentity(testMnemonic, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity: %v", err)
|
||||
}
|
||||
pubKey5 := identity5.Recipient().String()
|
||||
|
||||
// Write public key
|
||||
if err := afero.WriteFile(fs, filepath.Join(vaultDir2, "pub.age"), []byte(pubKey5), 0o600); err != nil {
|
||||
t.Fatalf("Failed to write public key: %v", err)
|
||||
}
|
||||
|
||||
// Compute the hash for index 5 key
|
||||
pubKeyHash5 := ComputeDoubleSHA256([]byte(pubKey5))
|
||||
|
||||
metadata2 := &Metadata{
|
||||
DerivationIndex: 5,
|
||||
PublicKeyHash: pubKeyHash5, // Hash of the actual key (index 5)
|
||||
MnemonicFamilyHash: pubKeyHash0, // Same family hash since it's from the same mnemonic
|
||||
}
|
||||
if err := SaveVaultMetadata(fs, vaultDir2, metadata2); err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
// Next index should be 1 (not 6) because we look for the first available slot
|
||||
index, err = GetNextDerivationIndex(fs, stateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
if index != 1 {
|
||||
t.Errorf("Expected index 1 (first available), got %d", index)
|
||||
}
|
||||
testGetNextDerivationIndex(t, fs)
|
||||
})
|
||||
|
||||
t.Run("MetadataPersistence", func(t *testing.T) {
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "test-vault")
|
||||
if err := fs.MkdirAll(vaultDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Create and save metadata
|
||||
metadata := &Metadata{
|
||||
DerivationIndex: 3,
|
||||
PublicKeyHash: "test-public-key-hash",
|
||||
}
|
||||
|
||||
if err := SaveVaultMetadata(fs, vaultDir, metadata); err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
// Load and verify
|
||||
loaded, err := LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load metadata: %v", err)
|
||||
}
|
||||
|
||||
if loaded.DerivationIndex != metadata.DerivationIndex {
|
||||
t.Errorf("DerivationIndex mismatch: expected %d, got %d", metadata.DerivationIndex, loaded.DerivationIndex)
|
||||
}
|
||||
if loaded.PublicKeyHash != metadata.PublicKeyHash {
|
||||
t.Errorf("PublicKeyHash mismatch: expected %s, got %s", metadata.PublicKeyHash, loaded.PublicKeyHash)
|
||||
}
|
||||
testMetadataPersistence(t, fs)
|
||||
})
|
||||
|
||||
t.Run("DifferentKeysForDifferentIndices", func(t *testing.T) {
|
||||
// Derive keys with different indices
|
||||
identity0, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity with index 0: %v", err)
|
||||
}
|
||||
|
||||
identity1, err := agehd.DeriveIdentity(testMnemonic, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity with index 1: %v", err)
|
||||
}
|
||||
|
||||
// Compute public key hashes
|
||||
pubKey0 := identity0.Recipient().String()
|
||||
pubKey1 := identity1.Recipient().String()
|
||||
hash0 := ComputeDoubleSHA256([]byte(pubKey0))
|
||||
|
||||
// Verify different indices produce different public keys
|
||||
if pubKey0 == pubKey1 {
|
||||
t.Errorf("Different derivation indices should produce different public keys")
|
||||
}
|
||||
|
||||
// But the hash of index 0's public key should be the same for the same mnemonic
|
||||
// This is what we use as the identifier
|
||||
identity0Again, _ := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
pubKey0Again := identity0Again.Recipient().String()
|
||||
hash0Again := ComputeDoubleSHA256([]byte(pubKey0Again))
|
||||
|
||||
if hash0 != hash0Again {
|
||||
t.Errorf("Same mnemonic should produce same public key hash for index 0")
|
||||
}
|
||||
testDifferentKeysForDifferentIndices(t)
|
||||
})
|
||||
}
|
||||
|
||||
func testComputeDoubleSHA256(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
// Test data
|
||||
data := []byte("test data")
|
||||
hash := vault.ComputeDoubleSHA256(data)
|
||||
|
||||
// Verify it's a valid hex string of 64 characters (32 bytes * 2)
|
||||
if len(hash) != 64 {
|
||||
t.Errorf("Expected hash length of 64, got %d", len(hash))
|
||||
}
|
||||
|
||||
// Verify consistency
|
||||
hash2 := vault.ComputeDoubleSHA256(data)
|
||||
if hash != hash2 {
|
||||
t.Errorf("Hash should be consistent for same input")
|
||||
}
|
||||
|
||||
// Verify different input produces different hash
|
||||
hash3 := vault.ComputeDoubleSHA256([]byte("different data"))
|
||||
if hash == hash3 {
|
||||
t.Errorf("Different input should produce different hash")
|
||||
}
|
||||
}
|
||||
|
||||
// createVaultDirWithMetadata creates a vault directory containing a public
|
||||
// key derived from testMnemonic at the given index plus saved metadata, and
|
||||
// returns the derived public key hash. An empty familyHash defaults to the
|
||||
// derived key's own hash.
|
||||
func createVaultDirWithMetadata(
|
||||
t *testing.T, fs afero.Fs, vaultName string,
|
||||
derivationIndex uint32, familyHash string,
|
||||
) string {
|
||||
t.Helper()
|
||||
|
||||
vaultDir := filepath.Join(testStateDir, "vaults.d", vaultName)
|
||||
|
||||
err := fs.MkdirAll(vaultDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Derive identity for the requested index
|
||||
identity, err := agehd.DeriveIdentity(testMnemonic, derivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity: %v", err)
|
||||
}
|
||||
|
||||
pubKey := identity.Recipient().String()
|
||||
pubKeyHash := vault.ComputeDoubleSHA256([]byte(pubKey))
|
||||
|
||||
// Write public key
|
||||
err = afero.WriteFile(fs, filepath.Join(vaultDir, "pub.age"),
|
||||
[]byte(pubKey), 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to write public key: %v", err)
|
||||
}
|
||||
|
||||
if familyHash == "" {
|
||||
familyHash = pubKeyHash
|
||||
}
|
||||
|
||||
metadata := &vault.Metadata{
|
||||
DerivationIndex: derivationIndex,
|
||||
PublicKeyHash: pubKeyHash,
|
||||
MnemonicFamilyHash: familyHash,
|
||||
}
|
||||
|
||||
err = vault.SaveVaultMetadata(fs, vaultDir, metadata)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
return pubKeyHash
|
||||
}
|
||||
|
||||
func testGetNextDerivationIndex(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
// Test with no existing vaults
|
||||
index, err := vault.GetNextDerivationIndex(fs, testStateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
|
||||
if index != 0 {
|
||||
t.Errorf("Expected index 0 for first vault, got %d", index)
|
||||
}
|
||||
|
||||
// Create a vault with metadata and matching public key (index 0; the
|
||||
// family hash is the index 0 key hash)
|
||||
pubKeyHash0 := createVaultDirWithMetadata(t, fs, "vault1", 0, "")
|
||||
|
||||
// Next index for same mnemonic should be 1
|
||||
index, err = vault.GetNextDerivationIndex(fs, testStateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
|
||||
if index != 1 {
|
||||
t.Errorf("Expected index 1 for second vault with same mnemonic, got %d", index)
|
||||
}
|
||||
|
||||
// Different mnemonic should start at 0
|
||||
//nolint:dupword // BIP39-style test mnemonic
|
||||
differentMnemonic := "zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo wrong"
|
||||
|
||||
index, err = vault.GetNextDerivationIndex(fs, testStateDir, differentMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
|
||||
if index != 0 {
|
||||
t.Errorf("Expected index 0 for first vault with different mnemonic, got %d",
|
||||
index)
|
||||
}
|
||||
|
||||
// Add another vault with same mnemonic but higher index (5), sharing
|
||||
// the same family hash since it's from the same mnemonic
|
||||
createVaultDirWithMetadata(t, fs, "vault2", 5, pubKeyHash0)
|
||||
|
||||
// Next index should be 1 (not 6): we look for the first available slot
|
||||
index, err = vault.GetNextDerivationIndex(fs, testStateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
|
||||
if index != 1 {
|
||||
t.Errorf("Expected index 1 (first available), got %d", index)
|
||||
}
|
||||
}
|
||||
|
||||
func testMetadataPersistence(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vaultDir := filepath.Join(testStateDir, "vaults.d", testVaultName)
|
||||
|
||||
err := fs.MkdirAll(vaultDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Create and save metadata
|
||||
metadata := &vault.Metadata{
|
||||
DerivationIndex: 3,
|
||||
PublicKeyHash: "test-public-key-hash",
|
||||
}
|
||||
|
||||
err = vault.SaveVaultMetadata(fs, vaultDir, metadata)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
// Load and verify
|
||||
loaded, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load metadata: %v", err)
|
||||
}
|
||||
|
||||
if loaded.DerivationIndex != metadata.DerivationIndex {
|
||||
t.Errorf("DerivationIndex mismatch: expected %d, got %d",
|
||||
metadata.DerivationIndex, loaded.DerivationIndex)
|
||||
}
|
||||
|
||||
if loaded.PublicKeyHash != metadata.PublicKeyHash {
|
||||
t.Errorf("PublicKeyHash mismatch: expected %s, got %s",
|
||||
metadata.PublicKeyHash, loaded.PublicKeyHash)
|
||||
}
|
||||
}
|
||||
|
||||
func testDifferentKeysForDifferentIndices(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
// Derive keys with different indices
|
||||
identity0, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity with index 0: %v", err)
|
||||
}
|
||||
|
||||
identity1, err := agehd.DeriveIdentity(testMnemonic, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity with index 1: %v", err)
|
||||
}
|
||||
|
||||
// Compute public key hashes
|
||||
pubKey0 := identity0.Recipient().String()
|
||||
pubKey1 := identity1.Recipient().String()
|
||||
hash0 := vault.ComputeDoubleSHA256([]byte(pubKey0))
|
||||
|
||||
// Verify different indices produce different public keys
|
||||
if pubKey0 == pubKey1 {
|
||||
t.Errorf("Different derivation indices should produce different public keys")
|
||||
}
|
||||
|
||||
// But the hash of index 0's public key should be the same for the same
|
||||
// mnemonic. This is what we use as the identifier
|
||||
identity0Again, _ := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
pubKey0Again := identity0Again.Recipient().String()
|
||||
hash0Again := vault.ComputeDoubleSHA256([]byte(pubKey0Again))
|
||||
|
||||
if hash0 != hash0Again {
|
||||
t.Errorf("Same mnemonic should produce same public key hash for index 0")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicKeyHashConsistency(t *testing.T) {
|
||||
// Use the same test mnemonic that the integration test uses
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Parallel()
|
||||
|
||||
// Derive identity from index 0 multiple times
|
||||
identity1, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
@@ -223,8 +258,8 @@ func TestPublicKeyHashConsistency(t *testing.T) {
|
||||
}
|
||||
|
||||
// Compute public key hashes
|
||||
hash1 := ComputeDoubleSHA256([]byte(identity1.Recipient().String()))
|
||||
hash2 := ComputeDoubleSHA256([]byte(identity2.Recipient().String()))
|
||||
hash1 := vault.ComputeDoubleSHA256([]byte(identity1.Recipient().String()))
|
||||
hash2 := vault.ComputeDoubleSHA256([]byte(identity2.Recipient().String()))
|
||||
|
||||
// Verify hashes are the same
|
||||
if hash1 != hash2 {
|
||||
@@ -237,11 +272,15 @@ func TestPublicKeyHashConsistency(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestSampleHashCalculation(t *testing.T) {
|
||||
// Test with the exact mnemonic from integration test if available
|
||||
// We'll also test with a few different mnemonics to make sure they produce different hashes
|
||||
t.Parallel()
|
||||
|
||||
// Test with the exact mnemonic from integration test if available. We
|
||||
// also test with a few different mnemonics to make sure they produce
|
||||
// different hashes
|
||||
mnemonics := []string{
|
||||
"abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about",
|
||||
testMnemonic,
|
||||
"legal winner thank year wave sausage worth useful legal winner thank yellow",
|
||||
//nolint:dupword // BIP39-style test mnemonic
|
||||
"zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo wrong",
|
||||
}
|
||||
|
||||
@@ -251,29 +290,29 @@ func TestSampleHashCalculation(t *testing.T) {
|
||||
t.Fatalf("Failed to derive identity for mnemonic %d: %v", i, err)
|
||||
}
|
||||
|
||||
hash := ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
hash := vault.ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
t.Logf("Mnemonic %d hash (index 0): %s", i, hash)
|
||||
t.Logf(" Recipient: %s", identity.Recipient().String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestWorkflowMismatch(t *testing.T) {
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
// Create a temporary directory for testing
|
||||
tempDir := t.TempDir()
|
||||
fs := afero.NewOsFs()
|
||||
|
||||
// Test Case 1: Create vault WITH mnemonic (like init command)
|
||||
t.Setenv("SB_SECRET_MNEMONIC", testMnemonic)
|
||||
_, err := CreateVault(fs, tempDir, "default")
|
||||
|
||||
_, err := vault.CreateVault(fs, tempDir, "default")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault with mnemonic: %v", err)
|
||||
}
|
||||
|
||||
// Load metadata for vault1
|
||||
vault1Dir := filepath.Join(tempDir, "vaults.d", "default")
|
||||
metadata1, err := LoadVaultMetadata(fs, vault1Dir)
|
||||
|
||||
metadata1, err := vault.LoadVaultMetadata(fs, vault1Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault1 metadata: %v", err)
|
||||
}
|
||||
@@ -281,9 +320,10 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
t.Logf("Vault1 (with mnemonic) - DerivationIndex: %d, PublicKeyHash: %s",
|
||||
metadata1.DerivationIndex, metadata1.PublicKeyHash)
|
||||
|
||||
// Test Case 2: Create vault WITHOUT mnemonic, then import (like work vault)
|
||||
// Test Case 2: Create vault WITHOUT mnemonic, then import (work vault)
|
||||
t.Setenv("SB_SECRET_MNEMONIC", "")
|
||||
_, err = CreateVault(fs, tempDir, "work")
|
||||
|
||||
_, err = vault.CreateVault(fs, tempDir, "work")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault without mnemonic: %v", err)
|
||||
}
|
||||
@@ -294,7 +334,7 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
t.Setenv("SB_SECRET_MNEMONIC", testMnemonic)
|
||||
|
||||
// Get the next available derivation index for this mnemonic
|
||||
derivationIndex, err := GetNextDerivationIndex(fs, tempDir, testMnemonic)
|
||||
derivationIndex, err := vault.GetNextDerivationIndex(fs, tempDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get next derivation index: %v", err)
|
||||
}
|
||||
@@ -306,10 +346,12 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity for index 0: %v", err)
|
||||
}
|
||||
publicKeyHash := ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
publicKeyHash := vault.ComputeDoubleSHA256(
|
||||
[]byte(identity0.Recipient().String()))
|
||||
|
||||
// Load existing metadata and update it (same as in VaultImport)
|
||||
existingMetadata, err := LoadVaultMetadata(fs, vault2Dir)
|
||||
existingMetadata, err := vault.LoadVaultMetadata(fs, vault2Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load existing metadata: %v", err)
|
||||
}
|
||||
@@ -318,12 +360,13 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
existingMetadata.DerivationIndex = derivationIndex
|
||||
existingMetadata.PublicKeyHash = publicKeyHash
|
||||
|
||||
if err := SaveVaultMetadata(fs, vault2Dir, existingMetadata); err != nil {
|
||||
err = vault.SaveVaultMetadata(fs, vault2Dir, existingMetadata)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to save vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Load updated metadata for vault2
|
||||
metadata2, err := LoadVaultMetadata(fs, vault2Dir)
|
||||
metadata2, err := vault.LoadVaultMetadata(fs, vault2Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault2 metadata: %v", err)
|
||||
}
|
||||
@@ -337,57 +380,59 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
t.Logf("Vault1 hash: %s", metadata1.PublicKeyHash)
|
||||
t.Logf("Vault2 hash: %s", metadata2.PublicKeyHash)
|
||||
} else {
|
||||
t.Logf("SUCCESS: Both vaults have the same public key hash: %s", metadata1.PublicKeyHash)
|
||||
t.Logf("SUCCESS: Both vaults have the same public key hash: %s",
|
||||
metadata1.PublicKeyHash)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReverseEngineerHash(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// This is the hash that the work vault is getting in the failing test
|
||||
wrongHash := "e34a2f500e395d8934a90a99ee9311edcfffd68cb701079575e50cbac7bb9417"
|
||||
correctHash := "992552b00b3879dfae461fab9a084b47784a032771c7a9accaebdde05ec7a7d1"
|
||||
|
||||
// Test mnemonic from integration test
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
// Calculate hash for test mnemonic
|
||||
identity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity: %v", err)
|
||||
}
|
||||
|
||||
calculatedHash := ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
calculatedHash := vault.ComputeDoubleSHA256(
|
||||
[]byte(identity.Recipient().String()))
|
||||
t.Logf("Test mnemonic hash: %s", calculatedHash)
|
||||
|
||||
if calculatedHash == correctHash {
|
||||
t.Logf("✓ Test mnemonic produces the correct hash")
|
||||
t.Logf("Test mnemonic produces the correct hash")
|
||||
} else {
|
||||
t.Errorf("✗ Test mnemonic does not produce the correct hash")
|
||||
t.Errorf("Test mnemonic does not produce the correct hash")
|
||||
}
|
||||
|
||||
if calculatedHash == wrongHash {
|
||||
t.Logf("✗ Test mnemonic unexpectedly produces the wrong hash")
|
||||
t.Logf("Test mnemonic unexpectedly produces the wrong hash")
|
||||
}
|
||||
|
||||
// Let's try some other possibilities - maybe there's a string normalization issue?
|
||||
// Try some other possibilities: maybe a string normalization issue?
|
||||
variations := []string{
|
||||
"abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about",
|
||||
" abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about ",
|
||||
"abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about\n",
|
||||
strings.TrimSpace("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"),
|
||||
testMnemonic,
|
||||
" " + testMnemonic + " ",
|
||||
testMnemonic + "\n",
|
||||
strings.TrimSpace(testMnemonic),
|
||||
}
|
||||
|
||||
for i, variation := range variations {
|
||||
identity, err := agehd.DeriveIdentity(variation, 0)
|
||||
if err != nil {
|
||||
t.Logf("Variation %d failed: %v", i, err)
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
hash := ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
hash := vault.ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
t.Logf("Variation %d hash: %s", i, hash)
|
||||
|
||||
if hash == wrongHash {
|
||||
t.Logf("✗ Found variation that produces wrong hash: '%s'", variation)
|
||||
t.Logf("Found variation that produces wrong hash: '%s'", variation)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -401,14 +446,15 @@ func TestReverseEngineerHash(t *testing.T) {
|
||||
identity, err := agehd.DeriveIdentity(emptyMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Logf("Empty mnemonic %d failed (expected): %v", i, err)
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
hash := ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
hash := vault.ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
t.Logf("Empty mnemonic %d hash: %s", i, hash)
|
||||
|
||||
if hash == wrongHash {
|
||||
t.Logf("✗ Empty mnemonic produces wrong hash!")
|
||||
t.Logf("Empty mnemonic produces wrong hash!")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestGetSecretVersionRejectsPathTraversal verifies that GetSecretVersion
|
||||
// validates the secret name and rejects path traversal attempts.
|
||||
// This is a regression test for https://git.eeqj.de/sneak/secret/issues/13
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv in parent forbids parallel subtests
|
||||
func TestGetSecretVersionRejectsPathTraversal(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add a legitimate secret so the vault is set up
|
||||
value := memguard.NewBufferFromBytes([]byte("legitimate-secret"))
|
||||
err = vlt.AddSecret("legit", value, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
// These names contain path traversal and should be rejected
|
||||
maliciousNames := []string{
|
||||
"../../../etc/passwd",
|
||||
"..%2f..%2fetc/passwd",
|
||||
".secret",
|
||||
"../sibling-vault/secrets.d/target",
|
||||
"foo/../bar",
|
||||
"a/../../etc/passwd",
|
||||
}
|
||||
|
||||
for _, name := range maliciousNames {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
_, err := vlt.GetSecretVersion(name, "")
|
||||
require.Error(t, err,
|
||||
"GetSecretVersion should reject malicious name: %s", name)
|
||||
require.Contains(t, err.Error(), "invalid secret name",
|
||||
"error should indicate invalid name for: %s", name)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGetSecretRejectsPathTraversal verifies GetSecret (which calls
|
||||
// GetSecretVersion) also rejects path traversal names.
|
||||
func TestGetSecretRejectsPathTraversal(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = vlt.GetSecret("../../../etc/passwd")
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid secret name")
|
||||
}
|
||||
|
||||
// TestGetSecretObjectRejectsPathTraversal verifies GetSecretObject
|
||||
// also validates names and rejects path traversal attempts.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv in parent forbids parallel subtests
|
||||
func TestGetSecretObjectRejectsPathTraversal(t *testing.T) {
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
maliciousNames := []string{
|
||||
"../../../etc/passwd",
|
||||
"foo/../bar",
|
||||
"a/../../etc/passwd",
|
||||
}
|
||||
|
||||
for _, name := range maliciousNames {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
_, err := vlt.GetSecretObject(name)
|
||||
require.Error(t, err, "GetSecretObject should reject: %s", name)
|
||||
require.Contains(t, err.Error(), "invalid secret name")
|
||||
})
|
||||
}
|
||||
}
|
||||
+432
-223
@@ -6,6 +6,7 @@ import (
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -21,7 +22,8 @@ func (v *Vault) ListSecrets() ([]string, error) {
|
||||
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory for secret listing", "error", err, "vault_name", v.Name)
|
||||
secret.Debug("Failed to get vault directory for secret listing",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, err
|
||||
}
|
||||
@@ -31,12 +33,15 @@ func (v *Vault) ListSecrets() ([]string, error) {
|
||||
// Check if secrets directory exists
|
||||
exists, err := afero.DirExists(v.fs, secretsDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check secrets directory", "error", err, "secrets_dir", secretsDir)
|
||||
secret.Debug("Failed to check secrets directory",
|
||||
"error", err, "secrets_dir", secretsDir)
|
||||
|
||||
return nil, fmt.Errorf("failed to check if secrets directory exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Debug("Secrets directory does not exist", "secrets_dir", secretsDir, "vault_name", v.Name)
|
||||
secret.Debug("Secrets directory does not exist",
|
||||
"secrets_dir", secretsDir, "vault_name", v.Name)
|
||||
|
||||
return []string{}, nil
|
||||
}
|
||||
@@ -44,12 +49,14 @@ func (v *Vault) ListSecrets() ([]string, error) {
|
||||
// List directories in secrets.d
|
||||
files, err := afero.ReadDir(v.fs, secretsDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read secrets directory", "error", err, "secrets_dir", secretsDir)
|
||||
secret.Debug("Failed to read secrets directory",
|
||||
"error", err, "secrets_dir", secretsDir)
|
||||
|
||||
return nil, fmt.Errorf("failed to read secrets directory: %w", err)
|
||||
}
|
||||
|
||||
var secrets []string
|
||||
|
||||
for _, file := range files {
|
||||
if file.IsDir() {
|
||||
// Convert storage name back to secret name
|
||||
@@ -67,11 +74,12 @@ func (v *Vault) ListSecrets() ([]string, error) {
|
||||
return secrets, nil
|
||||
}
|
||||
|
||||
// isValidSecretName validates secret names according to the format [a-z0-9\.\-\_\/]+
|
||||
// isValidSecretName validates secret names according to the format [a-zA-Z0-9\.\-\_\/]+
|
||||
// but with additional restrictions:
|
||||
// - No leading or trailing slashes
|
||||
// - No double slashes
|
||||
// - No names starting with dots
|
||||
// - No ".." path segments
|
||||
func isValidSecretName(name string) bool {
|
||||
if name == "" {
|
||||
return false
|
||||
@@ -92,16 +100,37 @@ func isValidSecretName(name string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check for path traversal via ".." components
|
||||
if slices.Contains(strings.Split(name, "/"), "..") {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check the basic pattern
|
||||
matched, _ := regexp.MatchString(`^[a-z0-9\.\-\_\/]+$`, name)
|
||||
matched, _ := regexp.MatchString(`^[a-zA-Z0-9\.\-\_\/]+$`, name)
|
||||
|
||||
return matched
|
||||
}
|
||||
|
||||
// ValidateSecretName returns an error wrapping ErrInvalidSecretName when
|
||||
// name is not a valid secret name. Call it on the name exactly as the user
|
||||
// gave it, before building any path from it.
|
||||
func ValidateSecretName(name string) error {
|
||||
if !isValidSecretName(name) {
|
||||
return fmt.Errorf(
|
||||
"%w '%s': only ASCII letters, digits, '.', '-', '_' and '/' are allowed, "+
|
||||
"and a name must not be empty, start with '.' or '/', end with '/', "+
|
||||
"contain '//', or have '..' as a path segment",
|
||||
ErrInvalidSecretName, name,
|
||||
)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddSecret adds a secret to this vault
|
||||
func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool) error {
|
||||
if value == nil {
|
||||
return fmt.Errorf("value buffer is nil")
|
||||
return ErrNilValueBuffer
|
||||
}
|
||||
|
||||
secret.DebugWith("Adding secret to vault",
|
||||
@@ -112,20 +141,25 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool)
|
||||
)
|
||||
|
||||
// Validate secret name
|
||||
if !isValidSecretName(name) {
|
||||
secret.Debug("Invalid secret name provided", "secret_name", name)
|
||||
|
||||
return fmt.Errorf("invalid secret name '%s': must match pattern [a-z0-9.\\-_/]+", name)
|
||||
}
|
||||
secret.Debug("Secret name validation passed", "secret_name", name)
|
||||
|
||||
secret.Debug("Getting vault directory")
|
||||
vaultDir, err := v.GetDirectory()
|
||||
err := ValidateSecretName(name)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory for secret addition", "error", err, "vault_name", v.Name)
|
||||
secret.Debug("Invalid secret name provided", "secret_name", name)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
secret.Debug("Secret name validation passed", "secret_name", name)
|
||||
|
||||
secret.Debug("Getting vault directory")
|
||||
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory for secret addition",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
secret.Debug("Got vault directory", "vault_dir", vaultDir)
|
||||
|
||||
// Convert slashes to percent signs for storage
|
||||
@@ -137,112 +171,72 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool)
|
||||
slog.String("secret_dir", secretDir),
|
||||
)
|
||||
|
||||
// Check if secret already exists
|
||||
secret.Debug("Checking if secret already exists", "secret_dir", secretDir)
|
||||
exists, err := afero.DirExists(v.fs, secretDir)
|
||||
// Check for an existing secret and the version the new one supersedes
|
||||
exists, previousVersion, err := v.checkExistingSecret(name, secretDir, force)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if secret exists", "error", err, "secret_dir", secretDir)
|
||||
|
||||
return fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
return err
|
||||
}
|
||||
secret.Debug("Secret existence check complete", "exists", exists)
|
||||
|
||||
// Handle existing secret case
|
||||
now := time.Now()
|
||||
var previousVersion *secret.Version
|
||||
|
||||
if exists {
|
||||
if !force {
|
||||
secret.Debug("Secret already exists and force not specified", "secret_name", name, "secret_dir", secretDir)
|
||||
|
||||
return fmt.Errorf("secret %s already exists (use --force to overwrite)", name)
|
||||
}
|
||||
|
||||
// Get the current version to update its notAfter timestamp
|
||||
currentVersionName, err := secret.GetCurrentVersion(v.fs, secretDir)
|
||||
if err == nil && currentVersionName != "" {
|
||||
previousVersion = secret.NewVersion(v, name, currentVersionName)
|
||||
// We'll need to load and update its metadata after we unlock the vault
|
||||
}
|
||||
} else {
|
||||
// Create secret directory for new secret
|
||||
secret.Debug("Creating secret directory", "secret_dir", secretDir)
|
||||
if err := v.fs.MkdirAll(secretDir, secret.DirPerms); err != nil {
|
||||
secret.Debug("Failed to create secret directory", "error", err, "secret_dir", secretDir)
|
||||
|
||||
return fmt.Errorf("failed to create secret directory: %w", err)
|
||||
}
|
||||
secret.Debug("Created secret directory successfully")
|
||||
return v.addVersion(name, secretDir, value, previousVersion)
|
||||
}
|
||||
|
||||
// Generate new version name
|
||||
versionName, err := secret.GenerateVersionName(v.fs, secretDir)
|
||||
return v.addNewSecret(name, secretDir, value)
|
||||
}
|
||||
|
||||
// addNewSecret creates a secret by assembling its first version and current
|
||||
// pointer in a temporary directory, then renaming that directory to
|
||||
// secretDir, so an interrupted add leaves no half-made secret behind.
|
||||
func (v *Vault) addNewSecret(
|
||||
name, secretDir string, value *memguard.LockedBuffer,
|
||||
) error {
|
||||
buildDir, err := secret.TempDirFor(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to generate version name", "error", err, "secret_name", name)
|
||||
|
||||
return fmt.Errorf("failed to generate version name: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
secret.Debug("Generated new version name", "version", versionName, "secret_name", name)
|
||||
// Once the rename below has moved it into place, this finds nothing.
|
||||
defer func() { _ = v.fs.RemoveAll(buildDir) }()
|
||||
|
||||
// Create new version
|
||||
newVersion := secret.NewVersion(v, name, versionName)
|
||||
|
||||
// Set version timestamps
|
||||
if previousVersion == nil {
|
||||
// First version: notBefore = epoch + 1 second
|
||||
epochPlusOne := time.Unix(1, 0)
|
||||
newVersion.Metadata.NotBefore = &epochPlusOne
|
||||
} else {
|
||||
// New version: notBefore = now
|
||||
newVersion.Metadata.NotBefore = &now
|
||||
|
||||
// We'll update the previous version's notAfter after we save the new version
|
||||
err = v.addVersion(name, buildDir, value, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Save the new version - pass the LockedBuffer directly
|
||||
if err := newVersion.Save(value); err != nil {
|
||||
secret.Debug("Failed to save new version", "error", err, "version", versionName)
|
||||
|
||||
// Clean up the secret directory if this was a new secret
|
||||
if !exists {
|
||||
secret.Debug("Cleaning up secret directory due to save failure", "secret_dir", secretDir)
|
||||
_ = v.fs.RemoveAll(secretDir)
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to save version: %w", err)
|
||||
err = v.fs.Rename(buildDir, secretDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to move new secret into place: %w", err)
|
||||
}
|
||||
|
||||
// Update previous version if it exists
|
||||
if previousVersion != nil {
|
||||
// Get long-term key to decrypt/encrypt metadata
|
||||
ltIdentity, err := v.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get long-term key for metadata update", "error", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
// addVersion saves value as a new version under secretDir, sets the
|
||||
// notAfter timestamp of the version it supersedes, if any, and then points
|
||||
// current at the new version. Until that last step, current still names the
|
||||
// previous version, which stays readable.
|
||||
func (v *Vault) addVersion(
|
||||
name, secretDir string, value *memguard.LockedBuffer,
|
||||
previousVersion *secret.Version,
|
||||
) error {
|
||||
now := time.Now()
|
||||
|
||||
// Load previous version metadata
|
||||
if err := previousVersion.LoadMetadata(ltIdentity); err != nil {
|
||||
secret.Debug("Failed to load previous version metadata", "error", err)
|
||||
// Create the new version and save the encrypted value
|
||||
versionName, err := v.createAndSaveVersion(
|
||||
name, secretDir, value, previousVersion, &now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to load previous version metadata: %w", err)
|
||||
}
|
||||
|
||||
// Update notAfter timestamp
|
||||
previousVersion.Metadata.NotAfter = &now
|
||||
|
||||
// Re-save the metadata (we need to implement an update method)
|
||||
if err := updateVersionMetadata(v.fs, previousVersion, ltIdentity); err != nil {
|
||||
secret.Debug("Failed to update previous version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to update previous version metadata: %w", err)
|
||||
}
|
||||
// Update previous version's notAfter timestamp if it exists
|
||||
err = v.updatePreviousVersion(previousVersion, &now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Set current symlink to new version
|
||||
if err := secret.SetCurrentVersion(v.fs, secretDir, versionName); err != nil {
|
||||
err = secret.SetCurrentVersion(v.fs, secretDir, versionName)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to set current version", "error", err, "version", versionName)
|
||||
|
||||
return fmt.Errorf("failed to set current version: %w", err)
|
||||
@@ -256,9 +250,12 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool)
|
||||
}
|
||||
|
||||
// updateVersionMetadata updates the metadata of an existing version
|
||||
func updateVersionMetadata(fs afero.Fs, version *secret.Version, ltIdentity *age.X25519Identity) error {
|
||||
func updateVersionMetadata(
|
||||
fs afero.Fs, version *secret.Version, ltIdentity *age.X25519Identity,
|
||||
) error {
|
||||
// Read the version's encrypted private key
|
||||
encryptedPrivKeyPath := filepath.Join(version.Directory, "priv.age")
|
||||
|
||||
encryptedPrivKey, err := afero.ReadFile(fs, encryptedPrivKeyPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read encrypted version private key: %w", err)
|
||||
@@ -287,94 +284,70 @@ func updateVersionMetadata(fs afero.Fs, version *secret.Version, ltIdentity *age
|
||||
metadataBuffer := memguard.NewBufferFromBytes(metadataBytes)
|
||||
defer metadataBuffer.Destroy()
|
||||
|
||||
encryptedMetadata, err := secret.EncryptToRecipient(metadataBuffer, versionIdentity.Recipient())
|
||||
encryptedMetadata, err := secret.EncryptToRecipient(metadataBuffer,
|
||||
versionIdentity.Recipient())
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to encrypt version metadata: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted metadata
|
||||
metadataPath := filepath.Join(version.Directory, "metadata.age")
|
||||
if err := afero.WriteFile(fs, metadataPath, encryptedMetadata, secret.FilePerms); err != nil {
|
||||
|
||||
err = secret.WriteFileAtomic(fs, metadataPath, encryptedMetadata)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write encrypted version metadata: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSecret retrieves a secret from this vault
|
||||
func (v *Vault) GetSecret(name string) ([]byte, error) {
|
||||
// GetSecret retrieves the current version of a secret from this vault.
|
||||
// The caller must destroy the returned buffer.
|
||||
func (v *Vault) GetSecret(name string) (*memguard.LockedBuffer, error) {
|
||||
secret.DebugWith("Getting secret from vault",
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.String("secret_name", name),
|
||||
)
|
||||
|
||||
return v.GetSecretVersion(name, "")
|
||||
// GetSecretObject validates the name and checks that the secret exists
|
||||
secretObj, err := v.GetSecretObject(name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
currentVersion, err := secret.GetCurrentVersion(v.fs, secretObj.Directory)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get current version", "error", err, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("failed to get current version: %w", err)
|
||||
}
|
||||
|
||||
return v.GetSecretVersion(name, currentVersion)
|
||||
}
|
||||
|
||||
// GetSecretVersion retrieves a specific version of a secret (empty version means current)
|
||||
func (v *Vault) GetSecretVersion(name string, version string) ([]byte, error) {
|
||||
// GetSecretVersion retrieves a specific version of a secret. The version
|
||||
// must be one of the secret's versions; GetSecret gets the current one.
|
||||
// The caller must destroy the returned buffer.
|
||||
func (v *Vault) GetSecretVersion(
|
||||
name string, version string,
|
||||
) (*memguard.LockedBuffer, error) {
|
||||
secret.DebugWith("Getting secret version from vault",
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.String("secret_name", name),
|
||||
slog.String("version", version),
|
||||
)
|
||||
|
||||
// Get vault directory
|
||||
vaultDir, err := v.GetDirectory()
|
||||
// Validate the name and check that the version exists
|
||||
err := v.checkSecretVersion(name, version)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Convert slashes to percent signs for storage
|
||||
storageName := strings.ReplaceAll(name, "/", "%")
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", storageName)
|
||||
|
||||
// Check if secret exists
|
||||
exists, err := afero.DirExists(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if secret exists", "error", err, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
if !exists {
|
||||
secret.Debug("Secret not found in vault", "secret_name", name, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("secret %s not found", name)
|
||||
}
|
||||
|
||||
// Determine which version to get
|
||||
if version == "" {
|
||||
// Get current version
|
||||
currentVersion, err := secret.GetCurrentVersion(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get current version", "error", err, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("failed to get current version: %w", err)
|
||||
}
|
||||
version = currentVersion
|
||||
secret.Debug("Using current version", "version", version, "secret_name", name)
|
||||
}
|
||||
|
||||
// Create version object
|
||||
secretVersion := secret.NewVersion(v, name, version)
|
||||
|
||||
// Check if version exists
|
||||
versionPath := filepath.Join(secretDir, "versions", version)
|
||||
exists, err = afero.DirExists(v.fs, versionPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if version exists", "error", err, "version", version)
|
||||
|
||||
return nil, fmt.Errorf("failed to check if version exists: %w", err)
|
||||
}
|
||||
if !exists {
|
||||
secret.Debug("Version not found", "version", version, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("version %s not found for secret %s", version, name)
|
||||
}
|
||||
|
||||
secret.Debug("Version exists, proceeding with vault unlock and decryption", "version", version, "secret_name", name)
|
||||
secret.Debug("Version exists, proceeding with vault unlock and decryption",
|
||||
"version", version, "secret_name", name)
|
||||
|
||||
// Unlock the vault (get long-term key in memory)
|
||||
longTermIdentity, err := v.UnlockVault()
|
||||
@@ -392,34 +365,25 @@ func (v *Vault) GetSecretVersion(name string, version string) ([]byte, error) {
|
||||
)
|
||||
|
||||
// Get the version's value
|
||||
secret.Debug("About to call secretVersion.GetValue", "version", version, "secret_name", name)
|
||||
secret.Debug("About to call secretVersion.GetValue",
|
||||
"version", version, "secret_name", name)
|
||||
|
||||
decryptedValue, err := secretVersion.GetValue(longTermIdentity)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to decrypt version value", "error", err, "version", version, "secret_name", name)
|
||||
secret.Debug("Failed to decrypt version value",
|
||||
"error", err, "version", version, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("failed to decrypt version: %w", err)
|
||||
}
|
||||
|
||||
// Create a copy to return since the buffer will be destroyed
|
||||
result := make([]byte, decryptedValue.Size())
|
||||
copy(result, decryptedValue.Bytes())
|
||||
decryptedValue.Destroy()
|
||||
|
||||
secret.DebugWith("Successfully decrypted secret version",
|
||||
slog.String("secret_name", name),
|
||||
slog.String("version", version),
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.Int("decrypted_length", len(result)),
|
||||
slog.Int("decrypted_length", decryptedValue.Size()),
|
||||
)
|
||||
|
||||
// Debug: Log metadata about the decrypted value without exposing the actual secret
|
||||
secret.Debug("Vault secret decryption debug info",
|
||||
"secret_name", name,
|
||||
"version", version,
|
||||
"decrypted_value_length", len(result),
|
||||
"is_empty", len(result) == 0)
|
||||
|
||||
return result, nil
|
||||
return decryptedValue, nil
|
||||
}
|
||||
|
||||
// UnlockVault unlocks the vault and returns the long-term private key
|
||||
@@ -428,7 +392,8 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) {
|
||||
|
||||
// If vault is already unlocked, return the cached key
|
||||
if !v.Locked() {
|
||||
secret.Debug("Vault already unlocked, returning cached long-term key", "vault_name", v.Name)
|
||||
secret.Debug("Vault already unlocked, returning cached long-term key",
|
||||
"vault_name", v.Name)
|
||||
|
||||
return v.longTermKey, nil
|
||||
}
|
||||
@@ -436,7 +401,8 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) {
|
||||
// Get or derive the long-term key (but don't store it yet)
|
||||
longTermIdentity, err := v.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get or derive long-term key", "error", err, "vault_name", v.Name)
|
||||
secret.Debug("Failed to get or derive long-term key",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
@@ -454,6 +420,11 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) {
|
||||
|
||||
// GetSecretObject retrieves a Secret object with metadata loaded from this vault
|
||||
func (v *Vault) GetSecretObject(name string) (*secret.Secret, error) {
|
||||
err := ValidateSecretName(name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// First check if the secret exists by checking for the metadata file
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
@@ -469,27 +440,31 @@ func (v *Vault) GetSecretObject(name string) (*secret.Secret, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("secret %s not found", name)
|
||||
return nil, fmt.Errorf("secret %s %w", name, ErrSecretNotFound)
|
||||
}
|
||||
|
||||
// Create a Secret object
|
||||
secretObj := secret.NewSecret(v, name)
|
||||
|
||||
// Load the metadata from disk
|
||||
if err := secretObj.LoadMetadata(); err != nil {
|
||||
err = secretObj.LoadMetadata()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return secretObj, nil
|
||||
}
|
||||
|
||||
// CopySecretVersion copies a single version from source to this vault
|
||||
// It decrypts the value using srcIdentity and re-encrypts for this vault
|
||||
// CopySecretVersion copies a single version from source into destSecretDir
|
||||
// in this vault. It decrypts the value using srcIdentity and re-encrypts
|
||||
// for this vault.
|
||||
func (v *Vault) CopySecretVersion(
|
||||
srcVersion *secret.Version,
|
||||
srcIdentity *age.X25519Identity,
|
||||
destSecretName string,
|
||||
destSecretDir string,
|
||||
destVersionName string,
|
||||
) error {
|
||||
secret.DebugWith("Copying secret version to vault",
|
||||
@@ -508,18 +483,21 @@ func (v *Vault) CopySecretVersion(
|
||||
defer valueBuffer.Destroy()
|
||||
|
||||
// Load source metadata
|
||||
if err := srcVersion.LoadMetadata(srcIdentity); err != nil {
|
||||
err = srcVersion.LoadMetadata(srcIdentity)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to load source metadata: %w", err)
|
||||
}
|
||||
|
||||
// Create destination version with same name
|
||||
destVersion := secret.NewVersion(v, destSecretName, destVersionName)
|
||||
destVersion.Directory = filepath.Join(destSecretDir, "versions", destVersionName)
|
||||
|
||||
// Copy metadata (preserve original timestamps)
|
||||
destVersion.Metadata = srcVersion.Metadata
|
||||
|
||||
// Save the version (encrypts to this vault's LT key)
|
||||
if err := destVersion.Save(valueBuffer); err != nil {
|
||||
err = destVersion.Save(valueBuffer)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to save destination version: %w", err)
|
||||
}
|
||||
|
||||
@@ -553,26 +531,13 @@ func (v *Vault) CopySecretAllVersions(
|
||||
return fmt.Errorf("failed to get destination vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Check if destination secret already exists
|
||||
// Refuse to replace an existing destination secret unless forced
|
||||
destStorageName := strings.ReplaceAll(destSecretName, "/", "%")
|
||||
destSecretDir := filepath.Join(destVaultDir, "secrets.d", destStorageName)
|
||||
|
||||
exists, err := afero.DirExists(v.fs, destSecretDir)
|
||||
err = v.checkCopyDestination(destSecretDir, destSecretName, force)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check destination: %w", err)
|
||||
}
|
||||
|
||||
if exists && !force {
|
||||
return fmt.Errorf("secret '%s' already exists in vault '%s' (use --force to overwrite)",
|
||||
destSecretName, v.Name)
|
||||
}
|
||||
|
||||
if exists && force {
|
||||
// Remove existing secret
|
||||
secret.Debug("Removing existing destination secret", "path", destSecretDir)
|
||||
if err := v.fs.RemoveAll(destSecretDir); err != nil {
|
||||
return fmt.Errorf("failed to remove existing destination secret: %w", err)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Get source vault's long-term key
|
||||
@@ -597,7 +562,7 @@ func (v *Vault) CopySecretAllVersions(
|
||||
}
|
||||
|
||||
if len(versions) == 0 {
|
||||
return fmt.Errorf("source secret '%s' has no versions", srcSecretName)
|
||||
return fmt.Errorf("source secret '%s' %w", srcSecretName, ErrNoVersions)
|
||||
}
|
||||
|
||||
// Get current version name
|
||||
@@ -606,28 +571,11 @@ func (v *Vault) CopySecretAllVersions(
|
||||
return fmt.Errorf("failed to get current version: %w", err)
|
||||
}
|
||||
|
||||
// Create destination secret directory
|
||||
if err := v.fs.MkdirAll(destSecretDir, secret.DirPerms); err != nil {
|
||||
return fmt.Errorf("failed to create destination secret directory: %w", err)
|
||||
}
|
||||
|
||||
// Copy each version
|
||||
for _, versionName := range versions {
|
||||
srcVersion := secret.NewVersion(srcVault, srcSecretName, versionName)
|
||||
if err := v.CopySecretVersion(srcVersion, srcIdentity, destSecretName, versionName); err != nil {
|
||||
// Rollback: remove partial copy
|
||||
secret.Debug("Rolling back partial copy due to error", "error", err)
|
||||
_ = v.fs.RemoveAll(destSecretDir)
|
||||
|
||||
return fmt.Errorf("failed to copy version %s: %w", versionName, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Set current version
|
||||
if err := secret.SetCurrentVersion(v.fs, destSecretDir, currentVersion); err != nil {
|
||||
_ = v.fs.RemoveAll(destSecretDir)
|
||||
|
||||
return fmt.Errorf("failed to set current version: %w", err)
|
||||
// Copy each version and the current pointer, then move the copy into place
|
||||
err = v.copyVersions(srcVault, srcIdentity,
|
||||
srcSecretName, destSecretName, destSecretDir, versions, currentVersion)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully copied all secret versions",
|
||||
@@ -638,3 +586,264 @@ func (v *Vault) CopySecretAllVersions(
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkExistingSecret reports whether the secret already exists, refuses to
|
||||
// overwrite it unless force is set, and returns its current version, which
|
||||
// the new version supersedes, if any.
|
||||
func (v *Vault) checkExistingSecret(
|
||||
name, secretDir string, force bool,
|
||||
) (bool, *secret.Version, error) {
|
||||
// Check if secret already exists
|
||||
secret.Debug("Checking if secret already exists", "secret_dir", secretDir)
|
||||
|
||||
exists, err := afero.DirExists(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if secret exists",
|
||||
"error", err, "secret_dir", secretDir)
|
||||
|
||||
return false, nil, fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
secret.Debug("Secret existence check complete", "exists", exists)
|
||||
|
||||
if !exists {
|
||||
return false, nil, nil
|
||||
}
|
||||
|
||||
if !force {
|
||||
secret.Debug("Secret already exists and force not specified",
|
||||
"secret_name", name, "secret_dir", secretDir)
|
||||
|
||||
return true, nil, fmt.Errorf(
|
||||
"secret %s %w (use --force to overwrite)",
|
||||
name, ErrSecretExists,
|
||||
)
|
||||
}
|
||||
|
||||
// Get the current version to update its notAfter timestamp
|
||||
var previousVersion *secret.Version
|
||||
|
||||
currentVersionName, err := secret.GetCurrentVersion(v.fs, secretDir)
|
||||
if err == nil && currentVersionName != "" {
|
||||
previousVersion = secret.NewVersion(v, name, currentVersionName)
|
||||
// We'll need to load and update its metadata after we unlock the vault
|
||||
}
|
||||
|
||||
return true, previousVersion, nil
|
||||
}
|
||||
|
||||
// updatePreviousVersion sets the notAfter timestamp on the version being
|
||||
// superseded. It is a no-op when previousVersion is nil.
|
||||
func (v *Vault) updatePreviousVersion(
|
||||
previousVersion *secret.Version, now *time.Time,
|
||||
) error {
|
||||
if previousVersion == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get long-term key to decrypt/encrypt metadata
|
||||
ltIdentity, err := v.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get long-term key for metadata update", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
|
||||
// Load previous version metadata
|
||||
err = previousVersion.LoadMetadata(ltIdentity)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to load previous version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to load previous version metadata: %w", err)
|
||||
}
|
||||
|
||||
// Update notAfter timestamp
|
||||
previousVersion.Metadata.NotAfter = now
|
||||
|
||||
// Re-save the metadata (we need to implement an update method)
|
||||
err = updateVersionMetadata(v.fs, previousVersion, ltIdentity)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to update previous version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to update previous version metadata: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkSecretVersion validates the secret name and verifies that the secret
|
||||
// exists and that version is one of its versions.
|
||||
func (v *Vault) checkSecretVersion(name, version string) error {
|
||||
// Validate secret name to prevent path traversal
|
||||
err := ValidateSecretName(name)
|
||||
if err != nil {
|
||||
secret.Debug("Invalid secret name provided", "secret_name", name)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// Get vault directory
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory", "error", err, "vault_name", v.Name)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// Convert slashes to percent signs for storage
|
||||
storageName := strings.ReplaceAll(name, "/", "%")
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", storageName)
|
||||
|
||||
// Check if secret exists
|
||||
exists, err := afero.DirExists(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if secret exists", "error", err, "secret_name", name)
|
||||
|
||||
return fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Debug("Secret not found in vault", "secret_name", name, "vault_name", v.Name)
|
||||
|
||||
return fmt.Errorf("secret %s %w", name, ErrSecretNotFound)
|
||||
}
|
||||
|
||||
// Check if version exists
|
||||
exists, err = secret.VersionExists(v.fs, secretDir, version)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if version exists", "error", err, "version", version)
|
||||
|
||||
return fmt.Errorf("failed to check if version exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Debug("Version not found", "version", version, "secret_name", name)
|
||||
|
||||
return fmt.Errorf("version '%s' %w '%s'", version, ErrVersionNotFound, name)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// createAndSaveVersion generates a new version name, sets the version
|
||||
// timestamps, and saves the encrypted value under secretDir, which is a
|
||||
// temporary directory while a new secret is being assembled.
|
||||
func (v *Vault) createAndSaveVersion(
|
||||
name, secretDir string, value *memguard.LockedBuffer,
|
||||
previousVersion *secret.Version, now *time.Time,
|
||||
) (string, error) {
|
||||
// Generate new version name
|
||||
versionName, err := secret.GenerateVersionName(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to generate version name", "error", err, "secret_name", name)
|
||||
|
||||
return "", fmt.Errorf("failed to generate version name: %w", err)
|
||||
}
|
||||
|
||||
secret.Debug("Generated new version name", "version", versionName, "secret_name", name)
|
||||
|
||||
// Create new version
|
||||
newVersion := secret.NewVersion(v, name, versionName)
|
||||
newVersion.Directory = filepath.Join(secretDir, "versions", versionName)
|
||||
|
||||
// Set version timestamps
|
||||
if previousVersion == nil {
|
||||
// First version: notBefore = epoch + 1 second
|
||||
epochPlusOne := time.Unix(1, 0)
|
||||
newVersion.Metadata.NotBefore = &epochPlusOne
|
||||
} else {
|
||||
// New version: notBefore = now
|
||||
newVersion.Metadata.NotBefore = now
|
||||
|
||||
// We'll update the previous version's notAfter after we save the
|
||||
// new version
|
||||
}
|
||||
|
||||
// Save the new version - pass the LockedBuffer directly
|
||||
err = newVersion.Save(value)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to save new version", "error", err, "version", versionName)
|
||||
|
||||
return "", fmt.Errorf("failed to save version: %w", err)
|
||||
}
|
||||
|
||||
return versionName, nil
|
||||
}
|
||||
|
||||
// copyVersions copies each version of the source secret and its current
|
||||
// pointer into a temporary directory, then moves that directory to
|
||||
// destSecretDir, replacing a secret already there. Nothing in this vault
|
||||
// changes until the copy is complete, so an interrupted copy leaves only a
|
||||
// temporary directory behind.
|
||||
func (v *Vault) copyVersions(
|
||||
srcVault *Vault, srcIdentity *age.X25519Identity,
|
||||
srcSecretName, destSecretName, destSecretDir string,
|
||||
versions []string, currentVersion string,
|
||||
) error {
|
||||
buildDir, err := secret.TempDirFor(v.fs, destSecretDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Once the rename below has moved it into place, this finds nothing.
|
||||
defer func() { _ = v.fs.RemoveAll(buildDir) }()
|
||||
|
||||
for _, versionName := range versions {
|
||||
srcVersion := secret.NewVersion(srcVault, srcSecretName, versionName)
|
||||
|
||||
err = v.CopySecretVersion(
|
||||
srcVersion, srcIdentity, destSecretName, buildDir, versionName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to copy version %s: %w", versionName, err)
|
||||
}
|
||||
}
|
||||
|
||||
err = secret.SetCurrentVersion(v.fs, buildDir, currentVersion)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to set current version: %w", err)
|
||||
}
|
||||
|
||||
// With --force, the secret being replaced goes only now that its
|
||||
// replacement is complete
|
||||
exists, err := afero.DirExists(v.fs, destSecretDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check destination: %w", err)
|
||||
}
|
||||
|
||||
if exists {
|
||||
secret.Debug("Removing existing destination secret", "path", destSecretDir)
|
||||
|
||||
err = secret.RemoveDirAtomic(v.fs, destSecretDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to remove existing destination secret: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
err = v.fs.Rename(buildDir, destSecretDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to move copied secret into place: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkCopyDestination refuses to copy over an existing secret unless force
|
||||
// is set. A secret being replaced is removed by copyVersions, once its
|
||||
// replacement is complete.
|
||||
func (v *Vault) checkCopyDestination(
|
||||
destSecretDir, destSecretName string, force bool,
|
||||
) error {
|
||||
exists, err := afero.DirExists(v.fs, destSecretDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check destination: %w", err)
|
||||
}
|
||||
|
||||
if exists && !force {
|
||||
return fmt.Errorf(
|
||||
"secret '%s' %w in vault '%s' (use --force to overwrite)",
|
||||
destSecretName, ErrSecretExists, v.Name,
|
||||
)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
//nolint:testpackage // white-box test of unexported isValidSecretName
|
||||
package vault
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestIsValidSecretNameUppercase(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
valid bool
|
||||
}{
|
||||
// Lowercase (existing behavior)
|
||||
{"valid-name", true},
|
||||
{"valid.name", true},
|
||||
{"valid_name", true},
|
||||
{"valid/path/name", true},
|
||||
{"123valid", true},
|
||||
|
||||
// Uppercase (new behavior - issue #2)
|
||||
{"Valid-Upper-Name", true},
|
||||
{"2025-11-21-ber1app1-vaultik-test-bucket-AKI", true},
|
||||
{"MixedCase/Path/Name", true},
|
||||
{"ALLUPPERCASE", true},
|
||||
{"ABC123", true},
|
||||
|
||||
// Still invalid
|
||||
{"", false},
|
||||
{"invalid name", false},
|
||||
{"invalid@name", false},
|
||||
{".dotstart", false},
|
||||
{"/leading-slash", false},
|
||||
{"trailing-slash/", false},
|
||||
{"double//slash", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
result := isValidSecretName(tt.name)
|
||||
if result != tt.valid {
|
||||
t.Errorf("isValidSecretName(%q) = %v, want %v", tt.name, result, tt.valid)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -2,10 +2,14 @@
|
||||
//
|
||||
// Integration tests for vault-level version operations:
|
||||
//
|
||||
// - TestVaultAddSecretCreatesVersion: Tests that AddSecret creates proper version structure
|
||||
// - TestVaultAddSecretMultipleVersions: Tests creating multiple versions with force flag
|
||||
// - TestVaultGetSecretVersion: Tests retrieving specific versions and current version
|
||||
// - TestVaultVersionTimestamps: Tests timestamp logic (notBefore/notAfter) across versions
|
||||
// - TestVaultAddSecretCreatesVersion: Tests that AddSecret creates proper
|
||||
// version structure
|
||||
// - TestVaultAddSecretMultipleVersions: Tests creating multiple versions with
|
||||
// force flag
|
||||
// - TestVaultGetSecretVersion: Tests retrieving specific versions and current
|
||||
// version
|
||||
// - TestVaultVersionTimestamps: Tests timestamp logic (notBefore/notAfter)
|
||||
// across versions
|
||||
// - TestVaultGetNonExistentVersion: Tests error handling for invalid versions
|
||||
// - TestUpdateVersionMetadata: Tests metadata update functionality
|
||||
//
|
||||
@@ -15,6 +19,7 @@
|
||||
// - Promotion doesn't modify timestamps
|
||||
// - Metadata remains encrypted and intact
|
||||
|
||||
//nolint:testpackage // white-box test of unexported updateVersionMetadata
|
||||
package vault
|
||||
|
||||
import (
|
||||
@@ -30,33 +35,61 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// testMnemonic is the mnemonic used to derive the vault long-term key.
|
||||
//
|
||||
//nolint:dupword // BIP39 test mnemonic intentionally repeats a word
|
||||
const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon abandon about"
|
||||
|
||||
// envTestMnemonic is the (deliberately different) mnemonic placed in the
|
||||
// environment; the vault is unlocked manually with the derived key in
|
||||
// createTestVaultWithKey.
|
||||
//
|
||||
//nolint:dupword // BIP39-style test mnemonic intentionally repeats a word
|
||||
const envTestMnemonic = "abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon about"
|
||||
|
||||
// Shared fixtures for white-box tests in this package.
|
||||
const (
|
||||
testStateDir = "/test/state"
|
||||
testSecretPath = "test/secret"
|
||||
)
|
||||
|
||||
// Helper function to add a secret to vault with proper buffer protection
|
||||
func addTestSecretToVault(t *testing.T, vault *Vault, name string, value []byte, force bool) {
|
||||
func addTestSecretToVault(
|
||||
t *testing.T, vault *Vault, name string, value []byte, force bool,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
buffer := memguard.NewBufferFromBytes(value)
|
||||
defer buffer.Destroy()
|
||||
|
||||
err := vault.AddSecret(name, buffer, force)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// Helper function to create a vault with long-term key set up
|
||||
func createTestVaultWithKey(t *testing.T, fs afero.Fs, stateDir, vaultName string) *Vault {
|
||||
// Helper function to create a vault named "test" with its long-term key set
|
||||
// up and unlocked
|
||||
func createTestVaultWithKey(t *testing.T, fs afero.Fs) *Vault {
|
||||
t.Helper()
|
||||
|
||||
// Set mnemonic for testing
|
||||
t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon about")
|
||||
t.Setenv(secret.EnvMnemonic, envTestMnemonic)
|
||||
|
||||
// Create vault
|
||||
vault, err := CreateVault(fs, stateDir, vaultName)
|
||||
vault, err := CreateVault(fs, testStateDir, "test")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Derive and store long-term key from mnemonic
|
||||
mnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Store long-term public key in vault
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
|
||||
err = afero.WriteFile(fs, ltPubKeyPath,
|
||||
[]byte(ltIdentity.Recipient().String()), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Unlock the vault with the derived key
|
||||
@@ -65,20 +98,19 @@ func createTestVaultWithKey(t *testing.T, fs afero.Fs, stateDir, vaultName strin
|
||||
return vault
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultAddSecretCreatesVersion(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Add a secret
|
||||
secretName := "test/secret"
|
||||
secretValue := []byte("initial-value")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
addTestSecretToVault(t, vault, secretName, secretValue, false)
|
||||
addTestSecretToVault(t, vault, testSecretPath, secretValue, false)
|
||||
|
||||
// Check that version directory was created
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
@@ -97,32 +129,34 @@ func TestVaultAddSecretCreatesVersion(t *testing.T) {
|
||||
assert.True(t, exists)
|
||||
|
||||
// Get the secret value
|
||||
retrievedValue, err := vault.GetSecret(secretName)
|
||||
retrievedValue, err := vault.GetSecret(testSecretPath)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, expectedValue, retrievedValue)
|
||||
|
||||
defer retrievedValue.Destroy()
|
||||
|
||||
assert.Equal(t, expectedValue, retrievedValue.Bytes())
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultAddSecretMultipleVersions(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
|
||||
secretName := "test/secret"
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Add first version
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-1"), false)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-1"), false)
|
||||
|
||||
// Try to add again without force - should fail
|
||||
failBuffer := memguard.NewBufferFromBytes([]byte("version-2"))
|
||||
defer failBuffer.Destroy()
|
||||
err := vault.AddSecret(secretName, failBuffer, false)
|
||||
assert.Error(t, err)
|
||||
|
||||
err := vault.AddSecret(testSecretPath, failBuffer, false)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already exists")
|
||||
|
||||
// Add with force - should create new version
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-2"), true)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-2"), true)
|
||||
|
||||
// Check that we have two versions
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
@@ -132,27 +166,28 @@ func TestVaultAddSecretMultipleVersions(t *testing.T) {
|
||||
assert.Len(t, entries, 2)
|
||||
|
||||
// Current value should be version-2
|
||||
value, err := vault.GetSecret(secretName)
|
||||
value, err := vault.GetSecret(testSecretPath)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2"), value)
|
||||
|
||||
defer value.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-2"), value.Bytes())
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultGetSecretVersion(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
|
||||
secretName := "test/secret"
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Add multiple versions
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-1"), false)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-1"), false)
|
||||
|
||||
// Small delay to ensure different version names
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-2"), true)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-2"), true)
|
||||
|
||||
// Get versions list
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
@@ -163,58 +198,68 @@ func TestVaultGetSecretVersion(t *testing.T) {
|
||||
|
||||
// Get specific version (first one)
|
||||
firstVersion := versions[1] // Last in list is first created
|
||||
value, err := vault.GetSecretVersion(secretName, firstVersion)
|
||||
first, err := vault.GetSecretVersion(testSecretPath, firstVersion)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1"), value)
|
||||
|
||||
defer first.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-1"), first.Bytes())
|
||||
|
||||
// Get specific version (second one)
|
||||
secondVersion := versions[0] // First in list is most recent
|
||||
value, err = vault.GetSecretVersion(secretName, secondVersion)
|
||||
second, err := vault.GetSecretVersion(testSecretPath, secondVersion)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2"), value)
|
||||
|
||||
// Get current (empty version)
|
||||
value, err = vault.GetSecretVersion(secretName, "")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2"), value)
|
||||
defer second.Destroy()
|
||||
|
||||
assert.Equal(t, []byte("version-2"), second.Bytes())
|
||||
|
||||
// An empty version is not one of the versions; GetSecret gets the
|
||||
// current one
|
||||
_, err = vault.GetSecretVersion(testSecretPath, "")
|
||||
require.ErrorIs(t, err, ErrVersionNotFound)
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultVersionTimestamps(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Get long-term key
|
||||
ltIdentity, err := vault.GetOrDeriveLongTermKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
secretName := "test/secret"
|
||||
|
||||
// Add first version
|
||||
beforeFirst := time.Now()
|
||||
|
||||
v1Buffer := memguard.NewBufferFromBytes([]byte("version-1"))
|
||||
defer v1Buffer.Destroy()
|
||||
err = vault.AddSecret(secretName, v1Buffer, false)
|
||||
|
||||
err = vault.AddSecret(testSecretPath, v1Buffer, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
afterFirst := time.Now()
|
||||
|
||||
// Get first version metadata
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
secretDir := vaultDir + "/secrets.d/test%secret"
|
||||
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 1)
|
||||
|
||||
firstVersion := secret.NewVersion(vault, secretName, versions[0])
|
||||
firstVersion := secret.NewVersion(vault, testSecretPath, versions[0])
|
||||
err = firstVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Check first version timestamps
|
||||
assert.NotNil(t, firstVersion.Metadata.CreatedAt)
|
||||
assert.True(t, firstVersion.Metadata.CreatedAt.After(beforeFirst.Add(-time.Second)))
|
||||
assert.True(t, firstVersion.Metadata.CreatedAt.Before(afterFirst.Add(time.Second)))
|
||||
assert.True(t,
|
||||
firstVersion.Metadata.CreatedAt.After(beforeFirst.Add(-time.Second)))
|
||||
assert.True(t,
|
||||
firstVersion.Metadata.CreatedAt.Before(afterFirst.Add(time.Second)))
|
||||
|
||||
assert.NotNil(t, firstVersion.Metadata.NotBefore)
|
||||
assert.Equal(t, int64(1), firstVersion.Metadata.NotBefore.Unix()) // Epoch + 1
|
||||
@@ -222,8 +267,11 @@ func TestVaultVersionTimestamps(t *testing.T) {
|
||||
|
||||
// Add second version
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
beforeSecond := time.Now()
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-2"), true)
|
||||
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-2"), true)
|
||||
|
||||
afterSecond := time.Now()
|
||||
|
||||
// Get updated versions
|
||||
@@ -232,56 +280,59 @@ func TestVaultVersionTimestamps(t *testing.T) {
|
||||
require.Len(t, versions, 2)
|
||||
|
||||
// Reload first version metadata (should have notAfter now)
|
||||
firstVersion = secret.NewVersion(vault, secretName, versions[1])
|
||||
firstVersion = secret.NewVersion(vault, testSecretPath, versions[1])
|
||||
err = firstVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.NotNil(t, firstVersion.Metadata.NotAfter)
|
||||
assert.True(t, firstVersion.Metadata.NotAfter.After(beforeSecond.Add(-time.Second)))
|
||||
assert.True(t, firstVersion.Metadata.NotAfter.Before(afterSecond.Add(time.Second)))
|
||||
assert.True(t,
|
||||
firstVersion.Metadata.NotAfter.After(beforeSecond.Add(-time.Second)))
|
||||
assert.True(t,
|
||||
firstVersion.Metadata.NotAfter.Before(afterSecond.Add(time.Second)))
|
||||
|
||||
// Check second version timestamps
|
||||
secondVersion := secret.NewVersion(vault, secretName, versions[0])
|
||||
secondVersion := secret.NewVersion(vault, testSecretPath, versions[0])
|
||||
err = secondVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.NotNil(t, secondVersion.Metadata.NotBefore)
|
||||
assert.True(t, secondVersion.Metadata.NotBefore.After(beforeSecond.Add(-time.Second)))
|
||||
assert.True(t, secondVersion.Metadata.NotBefore.Before(afterSecond.Add(time.Second)))
|
||||
assert.True(t,
|
||||
secondVersion.Metadata.NotBefore.After(beforeSecond.Add(-time.Second)))
|
||||
assert.True(t,
|
||||
secondVersion.Metadata.NotBefore.Before(afterSecond.Add(time.Second)))
|
||||
assert.Nil(t, secondVersion.Metadata.NotAfter) // Current version
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultGetNonExistentVersion(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Add a secret
|
||||
addTestSecretToVault(t, vault, "test/secret", []byte("value"), false)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("value"), false)
|
||||
|
||||
// Try to get non-existent version
|
||||
_, err := vault.GetSecretVersion("test/secret", "20991231.999")
|
||||
assert.Error(t, err)
|
||||
_, err := vault.GetSecretVersion(testSecretPath, "20991231.999")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "not found")
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestUpdateVersionMetadata(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Get long-term key
|
||||
ltIdentity, err := vault.GetOrDeriveLongTermKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create a version manually to test updateVersionMetadata
|
||||
secretName := "test/secret"
|
||||
versionName := "20231215.001"
|
||||
version := secret.NewVersion(vault, secretName, versionName)
|
||||
version := secret.NewVersion(vault, testSecretPath, versionName)
|
||||
|
||||
// Set initial metadata
|
||||
now := time.Now()
|
||||
@@ -292,6 +343,7 @@ func TestUpdateVersionMetadata(t *testing.T) {
|
||||
// Save version
|
||||
testBuffer := memguard.NewBufferFromBytes([]byte("test-value"))
|
||||
defer testBuffer.Destroy()
|
||||
|
||||
err = version.Save(testBuffer)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -301,7 +353,7 @@ func TestUpdateVersionMetadata(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
// Load and verify
|
||||
version2 := secret.NewVersion(vault, secretName, versionName)
|
||||
version2 := secret.NewVersion(vault, testSecretPath, versionName)
|
||||
err = version2.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
+236
-139
@@ -14,13 +14,22 @@ import (
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// Unlocker metadata type strings.
|
||||
const (
|
||||
unlockerTypePassphrase = "passphrase"
|
||||
unlockerTypeSecureEnclave = "secure-enclave"
|
||||
)
|
||||
|
||||
// GetCurrentUnlocker returns the current unlocker for this vault
|
||||
//
|
||||
//nolint:ireturn // returns one of several concrete unlocker implementations
|
||||
func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
secret.DebugWith("Getting current unlocker", slog.String("vault_name", v.Name))
|
||||
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory for unlocker", "error", err, "vault_name", v.Name)
|
||||
secret.Debug("Failed to get vault directory for unlocker",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, err
|
||||
}
|
||||
@@ -30,7 +39,8 @@ func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
// Check if the symlink exists
|
||||
_, err = v.fs.Stat(currentUnlockerPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to stat current unlocker symlink", "error", err, "path", currentUnlockerPath)
|
||||
secret.Debug("Failed to stat current unlocker symlink",
|
||||
"error", err, "path", currentUnlockerPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read current unlocker: %w", err)
|
||||
}
|
||||
@@ -47,46 +57,37 @@ func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
)
|
||||
|
||||
// Read unlocker metadata
|
||||
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json")
|
||||
secret.Debug("Reading unlocker metadata", "path", metadataPath)
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
metadata, err := v.readUnlockerMetadata(unlockerDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read unlocker metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read unlocker metadata: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
secret.Debug("Failed to parse unlocker metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to parse unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
secret.DebugWith("Parsed unlocker metadata",
|
||||
slog.String("unlocker_type", metadata.Type),
|
||||
slog.Time("created_at", metadata.CreatedAt),
|
||||
slog.Any("flags", metadata.Flags),
|
||||
)
|
||||
|
||||
// Create unlocker instance using direct constructors with filesystem
|
||||
var unlocker secret.Unlocker
|
||||
// Use metadata directly as it's already the correct type
|
||||
switch metadata.Type {
|
||||
case "passphrase":
|
||||
secret.Debug("Creating passphrase unlocker instance", "unlocker_type", metadata.Type)
|
||||
case unlockerTypePassphrase:
|
||||
secret.Debug("Creating passphrase unlocker instance",
|
||||
"unlocker_type", metadata.Type)
|
||||
|
||||
unlocker = secret.NewPassphraseUnlocker(v.fs, unlockerDir, metadata)
|
||||
case "pgp":
|
||||
secret.Debug("Creating PGP unlocker instance", "unlocker_type", metadata.Type)
|
||||
|
||||
unlocker = secret.NewPGPUnlocker(v.fs, unlockerDir, metadata)
|
||||
case "keychain":
|
||||
secret.Debug("Creating keychain unlocker instance", "unlocker_type", metadata.Type)
|
||||
|
||||
unlocker = secret.NewKeychainUnlocker(v.fs, unlockerDir, metadata)
|
||||
case unlockerTypeSecureEnclave:
|
||||
secret.Debug("Creating secure enclave unlocker instance",
|
||||
"unlocker_type", metadata.Type)
|
||||
|
||||
unlocker = secret.NewSecureEnclaveUnlocker(v.fs, unlockerDir, metadata)
|
||||
default:
|
||||
secret.Debug("Unsupported unlocker type", "type", metadata.Type)
|
||||
|
||||
return nil, fmt.Errorf("unsupported unlocker type: %s", metadata.Type)
|
||||
return nil, fmt.Errorf("%w: %s", ErrUnsupportedUnlockerType, metadata.Type)
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully created unlocker instance",
|
||||
@@ -98,14 +99,16 @@ func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
return unlocker, nil
|
||||
}
|
||||
|
||||
// resolveUnlockerDirectory reads the current-unlocker file to get the unlocker directory path
|
||||
// resolveUnlockerDirectory reads the current-unlocker file to get the
|
||||
// unlocker directory path
|
||||
// The file contains just the unlocker name (e.g., "passphrase")
|
||||
func (v *Vault) resolveUnlockerDirectory(currentUnlockerPath string) (string, error) {
|
||||
secret.Debug("Reading current-unlocker file", "path", currentUnlockerPath)
|
||||
|
||||
unlockerNameBytes, err := afero.ReadFile(v.fs, currentUnlockerPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read current-unlocker file", "error", err, "path", currentUnlockerPath)
|
||||
secret.Debug("Failed to read current-unlocker file",
|
||||
"error", err, "path", currentUnlockerPath)
|
||||
|
||||
return "", fmt.Errorf("failed to read current unlocker: %w", err)
|
||||
}
|
||||
@@ -122,50 +125,52 @@ func (v *Vault) resolveUnlockerDirectory(currentUnlockerPath string) (string, er
|
||||
return absolutePath, nil
|
||||
}
|
||||
|
||||
// findUnlockerByID finds an unlocker by its ID and returns the unlocker instance and its directory path
|
||||
func (v *Vault) findUnlockerByID(unlockersDir, unlockerID string) (secret.Unlocker, string, error) {
|
||||
// findUnlockerByID finds an unlocker by its ID and returns the unlocker
|
||||
// instance and its directory path. A directory that ListUnlockers skips is
|
||||
// skipped here too, with the same warning. Such a directory has no ID: if
|
||||
// no unlocker has the ID unlockerID but such a directory is named
|
||||
// unlockerID, that directory is returned with a nil unlocker, so that
|
||||
// RemoveUnlocker can remove it.
|
||||
//
|
||||
//nolint:ireturn // returns one of several concrete unlocker implementations
|
||||
func (v *Vault) findUnlockerByID(
|
||||
unlockersDir, unlockerID string,
|
||||
) (secret.Unlocker, string, error) {
|
||||
files, err := afero.ReadDir(v.fs, unlockersDir)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("failed to read unlockers directory: %w", err)
|
||||
}
|
||||
|
||||
skippedDirPath := ""
|
||||
|
||||
for _, file := range files {
|
||||
if !file.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
// Read metadata file
|
||||
metadataPath := filepath.Join(unlockersDir, file.Name(), "unlocker-metadata.json")
|
||||
exists, err := afero.Exists(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("failed to check if metadata exists for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
if !exists {
|
||||
// Skip directories without metadata - they might not be unlockers
|
||||
unlockerDirPath := filepath.Join(unlockersDir, file.Name())
|
||||
|
||||
metadata, ok := v.readUnlockerMetadataOrWarn(unlockersDir, file.Name())
|
||||
if !ok {
|
||||
if file.Name() == unlockerID {
|
||||
skippedDirPath = unlockerDirPath
|
||||
}
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("failed to read metadata for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
return nil, "", fmt.Errorf("failed to parse metadata for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
|
||||
unlockerDirPath := filepath.Join(unlockersDir, file.Name())
|
||||
|
||||
// Create the appropriate unlocker instance
|
||||
var tempUnlocker secret.Unlocker
|
||||
|
||||
switch metadata.Type {
|
||||
case "passphrase":
|
||||
case unlockerTypePassphrase:
|
||||
tempUnlocker = secret.NewPassphraseUnlocker(v.fs, unlockerDirPath, metadata)
|
||||
case "pgp":
|
||||
tempUnlocker = secret.NewPGPUnlocker(v.fs, unlockerDirPath, metadata)
|
||||
case "keychain":
|
||||
tempUnlocker = secret.NewKeychainUnlocker(v.fs, unlockerDirPath, metadata)
|
||||
case unlockerTypeSecureEnclave:
|
||||
tempUnlocker = secret.NewSecureEnclaveUnlocker(v.fs, unlockerDirPath, metadata)
|
||||
default:
|
||||
continue
|
||||
}
|
||||
@@ -176,7 +181,7 @@ func (v *Vault) findUnlockerByID(unlockersDir, unlockerID string) (secret.Unlock
|
||||
}
|
||||
}
|
||||
|
||||
return nil, "", nil
|
||||
return nil, skippedDirPath, nil
|
||||
}
|
||||
|
||||
// ListUnlockers returns a list of available unlockers for this vault
|
||||
@@ -193,6 +198,7 @@ func (v *Vault) ListUnlockers() ([]UnlockerMetadata, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if unlockers directory exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return []UnlockerMetadata{}, nil
|
||||
}
|
||||
@@ -204,28 +210,14 @@ func (v *Vault) ListUnlockers() ([]UnlockerMetadata, error) {
|
||||
}
|
||||
|
||||
var unlockers []UnlockerMetadata
|
||||
|
||||
for _, file := range files {
|
||||
if file.IsDir() {
|
||||
// Read metadata file
|
||||
metadataPath := filepath.Join(unlockersDir, file.Name(), "unlocker-metadata.json")
|
||||
exists, err := afero.Exists(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if metadata exists for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("unlocker directory %s is missing metadata file", file.Name())
|
||||
}
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read metadata for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse metadata for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
if !file.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
metadata, ok := v.readUnlockerMetadataOrWarn(unlockersDir, file.Name())
|
||||
if ok {
|
||||
unlockers = append(unlockers, metadata)
|
||||
}
|
||||
}
|
||||
@@ -233,7 +225,54 @@ func (v *Vault) ListUnlockers() ([]UnlockerMetadata, error) {
|
||||
return unlockers, nil
|
||||
}
|
||||
|
||||
// RemoveUnlocker removes an unlocker from this vault
|
||||
// readUnlockerMetadataOrWarn reads the metadata of the unlocker directory
|
||||
// name in unlockersDir. If the metadata file cannot be checked for, is
|
||||
// missing, or cannot be read or parsed, it warns, naming the directory,
|
||||
// and returns false: the caller skips that directory.
|
||||
func (v *Vault) readUnlockerMetadataOrWarn(
|
||||
unlockersDir, name string,
|
||||
) (UnlockerMetadata, bool) {
|
||||
metadataPath := filepath.Join(unlockersDir, name, "unlocker-metadata.json")
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
|
||||
exists, err := afero.Exists(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
secret.Warn("Skipping unlocker directory whose metadata file cannot be checked",
|
||||
"directory", name, "error", err)
|
||||
|
||||
return metadata, false
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Warn("Skipping unlocker directory with missing metadata file",
|
||||
"directory", name)
|
||||
|
||||
return metadata, false
|
||||
}
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
secret.Warn("Skipping unlocker directory with unreadable metadata file",
|
||||
"directory", name, "error", err)
|
||||
|
||||
return metadata, false
|
||||
}
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
secret.Warn("Skipping unlocker directory with corrupt metadata file",
|
||||
"directory", name, "error", err)
|
||||
|
||||
return metadata, false
|
||||
}
|
||||
|
||||
return metadata, true
|
||||
}
|
||||
|
||||
// RemoveUnlocker removes an unlocker from this vault. An unlocker
|
||||
// directory that ListUnlockers skips is removed by its directory name; its
|
||||
// type is unknown, so only the directory is removed.
|
||||
func (v *Vault) RemoveUnlocker(unlockerID string) error {
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
@@ -244,13 +283,17 @@ func (v *Vault) RemoveUnlocker(unlockerID string) error {
|
||||
unlockersDir := filepath.Join(vaultDir, "unlockers.d")
|
||||
|
||||
// Find the unlocker by ID
|
||||
unlocker, _, err := v.findUnlockerByID(unlockersDir, unlockerID)
|
||||
unlocker, unlockerDir, err := v.findUnlockerByID(unlockersDir, unlockerID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if unlockerDir == "" {
|
||||
return fmt.Errorf("unlocker with ID %s %w", unlockerID, ErrUnlockerNotFound)
|
||||
}
|
||||
|
||||
if unlocker == nil {
|
||||
return fmt.Errorf("unlocker with ID %s not found", unlockerID)
|
||||
return secret.RemoveDirAtomic(v.fs, unlockerDir)
|
||||
}
|
||||
|
||||
// Use the unlocker's Remove method
|
||||
@@ -268,33 +311,28 @@ func (v *Vault) SelectUnlocker(unlockerID string) error {
|
||||
unlockersDir := filepath.Join(vaultDir, "unlockers.d")
|
||||
|
||||
// Find the unlocker by ID
|
||||
_, targetUnlockerDir, err := v.findUnlockerByID(unlockersDir, unlockerID)
|
||||
unlocker, targetUnlockerDir, err := v.findUnlockerByID(unlockersDir, unlockerID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if targetUnlockerDir == "" {
|
||||
return fmt.Errorf("unlocker with ID %s not found", unlockerID)
|
||||
// A directory found without an unlocker is one ListUnlockers skips; it
|
||||
// cannot be selected.
|
||||
if unlocker == nil {
|
||||
return fmt.Errorf("unlocker with ID %s %w", unlockerID, ErrUnlockerNotFound)
|
||||
}
|
||||
|
||||
// Create/update current-unlocker file with just the unlocker name
|
||||
// Create or replace the current-unlocker file with just the unlocker
|
||||
// name. It is replaced in one rename, so it never goes missing.
|
||||
currentUnlockerPath := filepath.Join(vaultDir, "current-unlocker")
|
||||
|
||||
// Remove existing file if it exists
|
||||
if exists, err := afero.Exists(v.fs, currentUnlockerPath); err != nil {
|
||||
return fmt.Errorf("failed to check if current-unlocker file exists: %w", err)
|
||||
} else if exists {
|
||||
if err := v.fs.Remove(currentUnlockerPath); err != nil {
|
||||
return fmt.Errorf("failed to remove existing current-unlocker file: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Get just the unlocker name (basename of the directory)
|
||||
unlockerName := filepath.Base(targetUnlockerDir)
|
||||
|
||||
// Write just the unlocker name to the file
|
||||
secret.Debug("Writing current-unlocker file", "unlocker_name", unlockerName)
|
||||
if err := afero.WriteFile(v.fs, currentUnlockerPath, []byte(unlockerName), secret.FilePerms); err != nil {
|
||||
|
||||
err = secret.WriteFileAtomic(v.fs, currentUnlockerPath, []byte(unlockerName))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create current-unlocker file: %w", err)
|
||||
}
|
||||
|
||||
@@ -303,92 +341,151 @@ func (v *Vault) SelectUnlocker(unlockerID string) error {
|
||||
|
||||
// CreatePassphraseUnlocker creates a new passphrase-protected unlocker
|
||||
// The passphrase must be provided as a LockedBuffer for security
|
||||
func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*secret.PassphraseUnlocker, error) {
|
||||
func (v *Vault) CreatePassphraseUnlocker(
|
||||
passphrase *memguard.LockedBuffer,
|
||||
) (*secret.PassphraseUnlocker, error) {
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Create unlocker directory
|
||||
unlockerDir := filepath.Join(vaultDir, "unlockers.d", "passphrase")
|
||||
if err := v.fs.MkdirAll(unlockerDir, secret.DirPerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to create unlocker directory: %w", err)
|
||||
// We need to get the long-term key (either from memory if unlocked, or
|
||||
// derive it). Getting it before anything is written means failing to
|
||||
// get it changes nothing, even when replacing the current unlocker.
|
||||
ltIdentity, err := v.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
|
||||
unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerTypePassphrase)
|
||||
|
||||
// Generate new age keypair for unlocker
|
||||
unlockerIdentity, err := age.GenerateX25519Identity()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate unlocker: %w", err)
|
||||
}
|
||||
|
||||
// Write public key
|
||||
pubKeyPath := filepath.Join(unlockerDir, "pub.age")
|
||||
if err := afero.WriteFile(v.fs, pubKeyPath,
|
||||
[]byte(unlockerIdentity.Recipient().String()),
|
||||
secret.FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write unlocker public key: %w", err)
|
||||
}
|
||||
// Encrypt long-term private key to this unlocker
|
||||
ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String()))
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
// Encrypt private key with passphrase
|
||||
privKeyStr := unlockerIdentity.String()
|
||||
privKeyBuffer := memguard.NewBufferFromBytes([]byte(privKeyStr))
|
||||
defer privKeyBuffer.Destroy()
|
||||
encryptedPrivKey, err := secret.EncryptWithPassphrase(privKeyBuffer, passphrase)
|
||||
encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer,
|
||||
unlockerIdentity.Recipient())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt unlocker private key: %w", err)
|
||||
return nil, fmt.Errorf("failed to encrypt long-term private key: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted private key
|
||||
privKeyPath := filepath.Join(unlockerDir, "priv.age")
|
||||
if err := afero.WriteFile(v.fs, privKeyPath, encryptedPrivKey, secret.FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write encrypted unlocker private key: %w", err)
|
||||
}
|
||||
|
||||
// Create metadata
|
||||
metadata := UnlockerMetadata{
|
||||
Type: "passphrase",
|
||||
Type: unlockerTypePassphrase,
|
||||
CreatedAt: time.Now(),
|
||||
Flags: []string{},
|
||||
}
|
||||
|
||||
// Write metadata
|
||||
metadataBytes, err := json.MarshalIndent(metadata, "", " ")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal metadata: %w", err)
|
||||
}
|
||||
|
||||
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json")
|
||||
if err := afero.WriteFile(v.fs, metadataPath, metadataBytes, secret.FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt long-term private key to this unlocker
|
||||
// We need to get the long-term key (either from memory if unlocked, or derive it)
|
||||
ltIdentity, err := v.GetOrDeriveLongTermKey()
|
||||
// Write the unlocker's files, the metadata last
|
||||
err = secret.WriteDir(v.fs, unlockerDir, func(dir string) error {
|
||||
return v.writeUnlockerFiles(dir, unlockerIdentity, passphrase,
|
||||
encryptedLtPrivKey, metadataBytes)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
|
||||
ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String()))
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, unlockerIdentity.Recipient())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt long-term private key: %w", err)
|
||||
}
|
||||
|
||||
ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age")
|
||||
if err := afero.WriteFile(v.fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Create the unlocker instance
|
||||
unlocker := secret.NewPassphraseUnlocker(v.fs, unlockerDir, metadata)
|
||||
|
||||
// Select this unlocker as current
|
||||
if err := v.SelectUnlocker(unlocker.GetID()); err != nil {
|
||||
err = v.SelectUnlocker(unlocker.GetID())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to select new unlocker: %w", err)
|
||||
}
|
||||
|
||||
return unlocker, nil
|
||||
}
|
||||
|
||||
// readUnlockerMetadata reads and parses the unlocker-metadata.json file in
|
||||
// the given unlocker directory.
|
||||
func (v *Vault) readUnlockerMetadata(unlockerDir string) (UnlockerMetadata, error) {
|
||||
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json")
|
||||
secret.Debug("Reading unlocker metadata", "path", metadataPath)
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read unlocker metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return metadata, fmt.Errorf("failed to read unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to parse unlocker metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return metadata, fmt.Errorf("failed to parse unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
secret.DebugWith("Parsed unlocker metadata",
|
||||
slog.String("unlocker_type", metadata.Type),
|
||||
slog.Time("created_at", metadata.CreatedAt),
|
||||
slog.Any("flags", metadata.Flags),
|
||||
)
|
||||
|
||||
return metadata, nil
|
||||
}
|
||||
|
||||
// writeUnlockerFiles writes the files of a passphrase unlocker into
|
||||
// unlockerDir: its public key, its passphrase-encrypted private key, the
|
||||
// long-term private key encrypted to it, and its metadata, last.
|
||||
func (v *Vault) writeUnlockerFiles(
|
||||
unlockerDir string,
|
||||
unlockerIdentity *age.X25519Identity,
|
||||
passphrase *memguard.LockedBuffer,
|
||||
encryptedLtPrivKey, metadataBytes []byte,
|
||||
) error {
|
||||
// Write public key
|
||||
pubKeyPath := filepath.Join(unlockerDir, "pub.age")
|
||||
|
||||
err := secret.WriteFileAtomic(v.fs, pubKeyPath,
|
||||
[]byte(unlockerIdentity.Recipient().String()))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write unlocker public key: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt private key with passphrase
|
||||
privKeyStr := unlockerIdentity.String()
|
||||
|
||||
privKeyBuffer := memguard.NewBufferFromBytes([]byte(privKeyStr))
|
||||
defer privKeyBuffer.Destroy()
|
||||
|
||||
encryptedPrivKey, err := secret.EncryptWithPassphrase(privKeyBuffer, passphrase)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to encrypt unlocker private key: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted private key
|
||||
privKeyPath := filepath.Join(unlockerDir, "priv.age")
|
||||
|
||||
err = secret.WriteFileAtomic(v.fs, privKeyPath, encryptedPrivKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write encrypted unlocker private key: %w", err)
|
||||
}
|
||||
|
||||
err = secret.WriteFileAtomic(v.fs,
|
||||
filepath.Join(unlockerDir, "longterm.age"), encryptedLtPrivKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
err = secret.WriteFileAtomic(v.fs,
|
||||
filepath.Join(unlockerDir, "unlocker-metadata.json"), metadataBytes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
+132
-108
@@ -23,12 +23,14 @@ type Vault struct {
|
||||
// NewVault creates a new Vault instance
|
||||
func NewVault(fs afero.Fs, stateDir string, name string) *Vault {
|
||||
secret.Debug("Creating NewVault instance")
|
||||
|
||||
v := &Vault{
|
||||
Name: name,
|
||||
fs: fs,
|
||||
stateDir: stateDir,
|
||||
longTermKey: nil,
|
||||
}
|
||||
|
||||
secret.Debug("Created NewVault instance successfully")
|
||||
|
||||
return v
|
||||
@@ -54,7 +56,8 @@ func (v *Vault) ClearLongTermKey() {
|
||||
v.longTermKey = nil
|
||||
}
|
||||
|
||||
// GetOrDeriveLongTermKey gets the long-term key from memory or derives it from available sources
|
||||
// GetOrDeriveLongTermKey gets the long-term key from memory or derives it
|
||||
// from available sources
|
||||
func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
|
||||
// If we have it in memory, return it
|
||||
if !v.Locked() {
|
||||
@@ -65,55 +68,12 @@ func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
|
||||
|
||||
// Try to derive from environment mnemonic first
|
||||
if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" {
|
||||
secret.Debug("Using mnemonic from environment for long-term key derivation", "vault_name", v.Name)
|
||||
|
||||
// Load vault metadata to get the derivation index
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
metadata, err := LoadVaultMetadata(v.fs, vaultDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to load vault metadata", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to load vault metadata: %w", err)
|
||||
}
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to derive long-term key from mnemonic", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
|
||||
// Verify that the derived key matches the stored public key hash
|
||||
derivedPubKeyHash := ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
if derivedPubKeyHash != metadata.PublicKeyHash {
|
||||
secret.Debug("Derived public key hash does not match stored hash",
|
||||
"vault_name", v.Name,
|
||||
"derived_hash", derivedPubKeyHash,
|
||||
"stored_hash", metadata.PublicKeyHash,
|
||||
"derivation_index", metadata.DerivationIndex)
|
||||
|
||||
return nil, fmt.Errorf("derived public key does not match vault: mnemonic may be incorrect")
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully derived long-term key from mnemonic",
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.String("public_key", ltIdentity.Recipient().String()),
|
||||
slog.Uint64("derivation_index", uint64(metadata.DerivationIndex)),
|
||||
)
|
||||
|
||||
// Cache the derived key by unlocking the vault
|
||||
v.Unlock(ltIdentity)
|
||||
secret.Debug("Vault is unlocked (lt key in memory) via mnemonic", "vault_name", v.Name)
|
||||
|
||||
return ltIdentity, nil
|
||||
return v.deriveLongTermKeyFromMnemonic(envMnemonic)
|
||||
}
|
||||
|
||||
// No mnemonic available, try to use current unlocker
|
||||
secret.Debug("No mnemonic available, using current unlocker to unlock vault", "vault_name", v.Name)
|
||||
secret.Debug("No mnemonic available, using current unlocker to unlock vault",
|
||||
"vault_name", v.Name)
|
||||
|
||||
// Get current unlocker
|
||||
unlocker, err := v.GetCurrentUnlocker()
|
||||
@@ -129,55 +89,12 @@ func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
|
||||
slog.String("unlocker_id", unlocker.GetID()),
|
||||
)
|
||||
|
||||
// Get unlocker identity
|
||||
unlockerIdentity, err := unlocker.GetIdentity()
|
||||
// Get the long-term key via the unlocker.
|
||||
// SE unlockers return the long-term key directly from GetIdentity().
|
||||
// Other unlockers return their own identity, used to decrypt longterm.age.
|
||||
ltIdentity, err := v.unlockLongTermKey(unlocker)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get unlocker identity", "error", err, "unlocker_type", unlocker.GetType())
|
||||
|
||||
return nil, fmt.Errorf("failed to get unlocker identity: %w", err)
|
||||
}
|
||||
|
||||
// Read encrypted long-term private key from unlocker directory
|
||||
unlockerDir := unlocker.GetDirectory()
|
||||
encryptedLtPrivKeyPath := filepath.Join(unlockerDir, "longterm.age")
|
||||
secret.Debug("Reading encrypted long-term private key", "path", encryptedLtPrivKeyPath)
|
||||
|
||||
encryptedLtPrivKey, err := afero.ReadFile(v.fs, encryptedLtPrivKeyPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read encrypted long-term private key", "error", err, "path", encryptedLtPrivKeyPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
secret.DebugWith("Read encrypted long-term private key",
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.String("unlocker_type", unlocker.GetType()),
|
||||
slog.Int("encrypted_length", len(encryptedLtPrivKey)),
|
||||
)
|
||||
|
||||
// Decrypt long-term private key using unlocker
|
||||
secret.Debug("Decrypting long-term private key with unlocker", "unlocker_type", unlocker.GetType())
|
||||
ltPrivKeyBuffer, err := secret.DecryptWithIdentity(encryptedLtPrivKey, unlockerIdentity)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to decrypt long-term private key", "error", err, "unlocker_type", unlocker.GetType())
|
||||
|
||||
return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err)
|
||||
}
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
secret.DebugWith("Successfully decrypted long-term private key",
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.String("unlocker_type", unlocker.GetType()),
|
||||
slog.Int("decrypted_length", ltPrivKeyBuffer.Size()),
|
||||
)
|
||||
|
||||
// Parse long-term private key
|
||||
secret.Debug("Parsing long-term private key", "vault_name", v.Name)
|
||||
ltIdentity, err := age.ParseX25519Identity(ltPrivKeyBuffer.String())
|
||||
if err != nil {
|
||||
secret.Debug("Failed to parse long-term private key", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to parse long-term private key: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully obtained long-term identity via unlocker",
|
||||
@@ -204,7 +121,10 @@ func (v *Vault) GetName() string {
|
||||
return v.Name
|
||||
}
|
||||
|
||||
// GetFilesystem returns the vault's filesystem (for VaultInterface compatibility)
|
||||
// GetFilesystem returns the vault's filesystem (for VaultInterface
|
||||
// compatibility)
|
||||
//
|
||||
//nolint:ireturn // afero.Fs is the interface required by VaultInterface
|
||||
func (v *Vault) GetFilesystem() afero.Fs {
|
||||
return v.fs
|
||||
}
|
||||
@@ -217,7 +137,13 @@ func (v *Vault) NumSecrets() (int, error) {
|
||||
}
|
||||
|
||||
secretsDir := filepath.Join(vaultDir, "secrets.d")
|
||||
exists, _ := afero.DirExists(v.fs, secretsDir)
|
||||
|
||||
exists, err := afero.DirExists(v.fs, secretsDir)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to check secrets directory %s: %w",
|
||||
secretsDir, err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return 0, nil
|
||||
}
|
||||
@@ -227,29 +153,127 @@ func (v *Vault) NumSecrets() (int, error) {
|
||||
return 0, fmt.Errorf("failed to read secrets directory: %w", err)
|
||||
}
|
||||
|
||||
// Count only directories that contain at least one version file
|
||||
// Count only directories that have a "current" version pointer file
|
||||
count := 0
|
||||
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
// Check if this secret directory contains any version files
|
||||
// A valid secret has a "current" file pointing to the active version
|
||||
secretDir := filepath.Join(secretsDir, entry.Name())
|
||||
versionFiles, err := afero.ReadDir(v.fs, secretDir)
|
||||
currentFile := filepath.Join(secretDir, "current")
|
||||
|
||||
exists, err := afero.Exists(v.fs, currentFile)
|
||||
if err != nil {
|
||||
continue // Skip directories we can't read
|
||||
return 0, fmt.Errorf("failed to check %s: %w", currentFile, err)
|
||||
}
|
||||
|
||||
// Look for at least one version file (excluding "current" symlink)
|
||||
for _, vFile := range versionFiles {
|
||||
if !vFile.IsDir() && vFile.Name() != "current" {
|
||||
count++
|
||||
|
||||
break // Found at least one version, count this secret
|
||||
}
|
||||
if exists {
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// deriveLongTermKeyFromMnemonic derives the long-term key from the given
|
||||
// mnemonic, verifies it against the vault metadata, and caches it in memory.
|
||||
func (v *Vault) deriveLongTermKeyFromMnemonic(
|
||||
envMnemonic string,
|
||||
) (*age.X25519Identity, error) {
|
||||
secret.Debug("Using mnemonic from environment for long-term key derivation",
|
||||
"vault_name", v.Name)
|
||||
|
||||
// Load vault metadata to get the derivation index
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
metadata, err := LoadVaultMetadata(v.fs, vaultDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to load vault metadata", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to load vault metadata: %w", err)
|
||||
}
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to derive long-term key from mnemonic",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
|
||||
// Verify that the derived key matches the stored public key hash
|
||||
derivedPubKeyHash := ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
if derivedPubKeyHash != metadata.PublicKeyHash {
|
||||
secret.Debug("Derived public key hash does not match stored hash",
|
||||
"vault_name", v.Name,
|
||||
"derived_hash", derivedPubKeyHash,
|
||||
"stored_hash", metadata.PublicKeyHash,
|
||||
"derivation_index", metadata.DerivationIndex)
|
||||
|
||||
return nil, ErrMnemonicMismatch
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully derived long-term key from mnemonic",
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.String("public_key", ltIdentity.Recipient().String()),
|
||||
slog.Uint64("derivation_index", uint64(metadata.DerivationIndex)),
|
||||
)
|
||||
|
||||
// Cache the derived key by unlocking the vault
|
||||
v.Unlock(ltIdentity)
|
||||
secret.Debug("Vault is unlocked (lt key in memory) via mnemonic",
|
||||
"vault_name", v.Name)
|
||||
|
||||
return ltIdentity, nil
|
||||
}
|
||||
|
||||
// unlockLongTermKey extracts the vault's long-term key using the given
|
||||
// unlocker. SE unlockers decrypt the long-term key directly; other unlockers
|
||||
// use an intermediate identity.
|
||||
func (v *Vault) unlockLongTermKey(
|
||||
unlocker secret.Unlocker,
|
||||
) (*age.X25519Identity, error) {
|
||||
if unlocker.GetType() == unlockerTypeSecureEnclave {
|
||||
secret.Debug("SE unlocker: decrypting long-term key directly via Secure Enclave")
|
||||
|
||||
ltIdentity, err := unlocker.GetIdentity()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decrypt long-term key via SE: %w", err)
|
||||
}
|
||||
|
||||
return ltIdentity, nil
|
||||
}
|
||||
|
||||
// Standard unlockers: get unlocker identity, then decrypt longterm.age
|
||||
unlockerIdentity, err := unlocker.GetIdentity()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get unlocker identity: %w", err)
|
||||
}
|
||||
|
||||
encryptedLtPrivKeyPath := filepath.Join(unlocker.GetDirectory(), "longterm.age")
|
||||
|
||||
encryptedLtPrivKey, err := afero.ReadFile(v.fs, encryptedLtPrivKeyPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
ltPrivKeyBuffer, err := secret.DecryptWithIdentity(
|
||||
encryptedLtPrivKey, unlockerIdentity)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err)
|
||||
}
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
ltIdentity, err := age.ParseX25519Identity(ltPrivKeyBuffer.String())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse long-term private key: %w", err)
|
||||
}
|
||||
|
||||
return ltIdentity, nil
|
||||
}
|
||||
|
||||
@@ -13,32 +13,34 @@ import (
|
||||
)
|
||||
|
||||
func TestAddSecretFailsWithMissingPublicKey(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Create in-memory filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create a vault directory without a public key (simulating the error condition)
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "broken")
|
||||
// Create a vault directory without a public key (simulating the error
|
||||
// condition)
|
||||
vaultDir := filepath.Join(testStateDir, "vaults.d", "broken")
|
||||
require.NoError(t, fs.MkdirAll(vaultDir, secret.DirPerms))
|
||||
|
||||
// Create currentvault symlink
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
require.NoError(t, afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), secret.FilePerms))
|
||||
currentVaultPath := filepath.Join(testStateDir, "currentvault")
|
||||
require.NoError(t,
|
||||
afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), secret.FilePerms))
|
||||
|
||||
// Create vault instance
|
||||
vlt := vault.NewVault(fs, stateDir, "broken")
|
||||
vlt := vault.NewVault(fs, testStateDir, "broken")
|
||||
|
||||
// Try to add a secret - this should fail
|
||||
secretName := "test-secret"
|
||||
value := memguard.NewBufferFromBytes([]byte("test-value"))
|
||||
defer value.Destroy()
|
||||
|
||||
err := vlt.AddSecret(secretName, value, false)
|
||||
err := vlt.AddSecret(testSecretName, value, false)
|
||||
require.Error(t, err, "AddSecret should fail when public key is missing")
|
||||
assert.Contains(t, err.Error(), "failed to read long-term public key")
|
||||
|
||||
// Verify that the secret directory was NOT created
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", secretName)
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", testSecretName)
|
||||
exists, _ := afero.DirExists(fs, secretDir)
|
||||
assert.False(t, exists, "Secret directory should not exist after failed AddSecret")
|
||||
|
||||
@@ -47,41 +49,51 @@ func TestAddSecretFailsWithMissingPublicKey(t *testing.T) {
|
||||
if exists, _ := afero.DirExists(fs, secretsDir); exists {
|
||||
entries, err := afero.ReadDir(fs, secretsDir)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, entries, "secrets.d directory should be empty after failed AddSecret")
|
||||
assert.Empty(t, entries,
|
||||
"secrets.d directory should be empty after failed AddSecret")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddSecretCleansUpOnFailure(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Create in-memory filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create a vault directory with public key
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "test")
|
||||
vaultDir := filepath.Join(testStateDir, "vaults.d", "test")
|
||||
require.NoError(t, fs.MkdirAll(vaultDir, secret.DirPerms))
|
||||
|
||||
// Create a mock public key that will cause encryption to fail
|
||||
// by using an invalid age public key format
|
||||
pubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
require.NoError(t, afero.WriteFile(fs, pubKeyPath, []byte("invalid-public-key"), secret.FilePerms))
|
||||
require.NoError(t,
|
||||
afero.WriteFile(fs, pubKeyPath, []byte("invalid-public-key"),
|
||||
secret.FilePerms))
|
||||
|
||||
// Create currentvault symlink
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
require.NoError(t, afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), secret.FilePerms))
|
||||
currentVaultPath := filepath.Join(testStateDir, "currentvault")
|
||||
require.NoError(t,
|
||||
afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), secret.FilePerms))
|
||||
|
||||
// Create vault instance
|
||||
vlt := vault.NewVault(fs, stateDir, "test")
|
||||
vlt := vault.NewVault(fs, testStateDir, "test")
|
||||
|
||||
// Try to add a secret - this should fail during encryption
|
||||
secretName := "test-secret"
|
||||
value := memguard.NewBufferFromBytes([]byte("test-value"))
|
||||
defer value.Destroy()
|
||||
|
||||
err := vlt.AddSecret(secretName, value, false)
|
||||
err := vlt.AddSecret(testSecretName, value, false)
|
||||
require.Error(t, err, "AddSecret should fail with invalid public key")
|
||||
|
||||
// Verify that the secret directory was NOT created
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", secretName)
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", testSecretName)
|
||||
exists, _ := afero.DirExists(fs, secretDir)
|
||||
assert.False(t, exists, "Secret directory should not exist after failed AddSecret")
|
||||
|
||||
// Nor is the temporary directory the secret was assembled in left behind
|
||||
entries, err := afero.ReadDir(fs, vaultDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, entries, 1)
|
||||
assert.Equal(t, "pub.age", entries[0].Name())
|
||||
}
|
||||
|
||||
+303
-193
@@ -1,227 +1,337 @@
|
||||
package vault
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// testMnemonic is the shared BIP39 test mnemonic for tests in this package.
|
||||
//
|
||||
//nolint:dupword // BIP39 test mnemonic intentionally repeats a word
|
||||
const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon abandon about"
|
||||
|
||||
// Shared fixtures for tests in this package.
|
||||
const (
|
||||
testStateDir = "/test/state"
|
||||
testVaultName = "test-vault"
|
||||
testSecretName = "test-secret"
|
||||
testPassphrase = "test-passphrase"
|
||||
)
|
||||
|
||||
//nolint:paralleltest // t.Setenv and order-dependent subtests forbid parallel
|
||||
func TestVaultOperations(t *testing.T) {
|
||||
// Test environment will be cleaned up automatically by t.Setenv
|
||||
|
||||
// Set test environment variables
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
// Use in-memory filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Test vault creation
|
||||
t.Run("CreateVault", func(t *testing.T) {
|
||||
vlt, err := CreateVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
if vlt.GetName() != "test-vault" {
|
||||
t.Errorf("Expected vault name 'test-vault', got '%s'", vlt.GetName())
|
||||
}
|
||||
|
||||
// Check vault directory exists
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
exists, err := afero.DirExists(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check vault directory: %v", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
t.Errorf("Vault directory should exist")
|
||||
}
|
||||
testCreateVault(t, fs)
|
||||
})
|
||||
|
||||
// Test vault listing
|
||||
t.Run("ListVaults", func(t *testing.T) {
|
||||
vaults, err := ListVaults(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list vaults: %v", err)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, vault := range vaults {
|
||||
if vault == "test-vault" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Errorf("Expected to find 'test-vault' in vault list")
|
||||
}
|
||||
testListVaults(t, fs)
|
||||
})
|
||||
|
||||
// Test vault selection
|
||||
t.Run("SelectVault", func(t *testing.T) {
|
||||
err := SelectVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select vault: %v", err)
|
||||
}
|
||||
|
||||
// Test getting current vault
|
||||
currentVault, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
if currentVault.GetName() != "test-vault" {
|
||||
t.Errorf("Expected current vault 'test-vault', got '%s'", currentVault.GetName())
|
||||
}
|
||||
testSelectVault(t, fs)
|
||||
})
|
||||
|
||||
// Test secret operations
|
||||
t.Run("SecretOperations", func(t *testing.T) {
|
||||
vlt, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
testSecretOperations(t, fs)
|
||||
})
|
||||
|
||||
// First, derive the long-term key from the test mnemonic
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Get the public key from the derived identity
|
||||
ltPublicKey := ltIdentity.Recipient().String()
|
||||
|
||||
// Get the vault directory
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Write the correct public key to the pub.age file
|
||||
pubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, pubKeyPath, []byte(ltPublicKey), secret.FilePerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to write long-term public key: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vault with the derived identity
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Now add a secret
|
||||
secretName := "test/secret"
|
||||
secretValue := []byte("test-secret-value")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vlt.AddSecret(secretName, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret: %v", err)
|
||||
}
|
||||
|
||||
// List secrets
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets: %v", err)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, secret := range secrets {
|
||||
if secret == secretName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Errorf("Expected to find secret '%s' in list", secretName)
|
||||
}
|
||||
|
||||
// Get secret value
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get secret: %v", err)
|
||||
}
|
||||
|
||||
if string(retrievedValue) != string(expectedValue) {
|
||||
t.Errorf("Expected secret value '%s', got '%s'", string(expectedValue), string(retrievedValue))
|
||||
}
|
||||
t.Run("NumSecrets", func(t *testing.T) {
|
||||
testNumSecrets(t, fs)
|
||||
})
|
||||
|
||||
// Test unlocker operations
|
||||
t.Run("UnlockerOperations", func(t *testing.T) {
|
||||
vlt, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
// Test vault unlocking (should happen automatically via mnemonic)
|
||||
if vlt.Locked() {
|
||||
_, err := vlt.UnlockVault()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unlock vault: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Create a passphrase unlocker
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte("test-passphrase"))
|
||||
defer passphraseBuffer.Destroy()
|
||||
passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create passphrase unlocker: %v", err)
|
||||
}
|
||||
|
||||
// List unlockers
|
||||
unlockers, err := vlt.ListUnlockers()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list unlockers: %v", err)
|
||||
}
|
||||
|
||||
if len(unlockers) == 0 {
|
||||
t.Errorf("Expected at least one unlocker")
|
||||
}
|
||||
|
||||
// Check key type
|
||||
keyFound := false
|
||||
for _, key := range unlockers {
|
||||
if key.Type == "passphrase" {
|
||||
keyFound = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !keyFound {
|
||||
t.Errorf("Expected to find passphrase unlocker")
|
||||
}
|
||||
|
||||
// Test selecting unlocker
|
||||
err = vlt.SelectUnlocker(passphraseUnlocker.GetID())
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select unlocker: %v", err)
|
||||
}
|
||||
|
||||
// Test getting current unlocker
|
||||
currentUnlocker, err := vlt.GetCurrentUnlocker()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current unlocker: %v", err)
|
||||
}
|
||||
|
||||
if currentUnlocker.GetID() != passphraseUnlocker.GetID() {
|
||||
t.Errorf("Expected current unlocker ID '%s', got '%s'", passphraseUnlocker.GetID(), currentUnlocker.GetID())
|
||||
}
|
||||
testUnlockerOperations(t, fs)
|
||||
})
|
||||
}
|
||||
|
||||
func testCreateVault(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
if vlt.GetName() != testVaultName {
|
||||
t.Errorf("Expected vault name '%s', got '%s'", testVaultName, vlt.GetName())
|
||||
}
|
||||
|
||||
// Check vault directory exists
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
exists, err := afero.DirExists(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check vault directory: %v", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
t.Errorf("Vault directory should exist")
|
||||
}
|
||||
}
|
||||
|
||||
func testListVaults(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vaults, err := vault.ListVaults(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list vaults: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(vaults, testVaultName) {
|
||||
t.Errorf("Expected to find '%s' in vault list", testVaultName)
|
||||
}
|
||||
}
|
||||
|
||||
func testSelectVault(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
err := vault.SelectVault(fs, testStateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select vault: %v", err)
|
||||
}
|
||||
|
||||
// Test getting current vault
|
||||
currentVault, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
if currentVault.GetName() != testVaultName {
|
||||
t.Errorf("Expected current vault '%s', got '%s'",
|
||||
testVaultName, currentVault.GetName())
|
||||
}
|
||||
}
|
||||
|
||||
func testSecretOperations(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
// First, derive the long-term key from the test mnemonic
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Get the public key from the derived identity
|
||||
ltPublicKey := ltIdentity.Recipient().String()
|
||||
|
||||
// Get the vault directory
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Write the correct public key to the pub.age file
|
||||
pubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
|
||||
err = afero.WriteFile(fs, pubKeyPath, []byte(ltPublicKey), secret.FilePerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to write long-term public key: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vault with the derived identity
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Now add a secret
|
||||
secretName := "test/secret"
|
||||
secretValue := []byte("test-secret-value")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vlt.AddSecret(secretName, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret: %v", err)
|
||||
}
|
||||
|
||||
// List secrets
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(secrets, secretName) {
|
||||
t.Errorf("Expected to find secret '%s' in list", secretName)
|
||||
}
|
||||
|
||||
// Get secret value
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get secret: %v", err)
|
||||
}
|
||||
defer retrievedValue.Destroy()
|
||||
|
||||
if !bytes.Equal(retrievedValue.Bytes(), expectedValue) {
|
||||
t.Errorf("Expected secret value '%s', got '%s'",
|
||||
expectedValue, retrievedValue.Bytes())
|
||||
}
|
||||
}
|
||||
|
||||
func testNumSecrets(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
numSecrets, err := vlt.NumSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to count secrets: %v", err)
|
||||
}
|
||||
|
||||
// We added one secret in SecretOperations
|
||||
if numSecrets != 1 {
|
||||
t.Errorf("Expected 1 secret, got %d", numSecrets)
|
||||
}
|
||||
}
|
||||
|
||||
func testUnlockerOperations(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
// Test vault unlocking (should happen automatically via mnemonic)
|
||||
if vlt.Locked() {
|
||||
_, err := vlt.UnlockVault()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unlock vault: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Create a passphrase unlocker
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte(testPassphrase))
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create passphrase unlocker: %v", err)
|
||||
}
|
||||
|
||||
// List unlockers
|
||||
unlockers, err := vlt.ListUnlockers()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list unlockers: %v", err)
|
||||
}
|
||||
|
||||
if len(unlockers) == 0 {
|
||||
t.Errorf("Expected at least one unlocker")
|
||||
}
|
||||
|
||||
// Check key type
|
||||
keyFound := false
|
||||
|
||||
for _, key := range unlockers {
|
||||
if key.Type == "passphrase" {
|
||||
keyFound = true
|
||||
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !keyFound {
|
||||
t.Errorf("Expected to find passphrase unlocker")
|
||||
}
|
||||
|
||||
// Test selecting unlocker
|
||||
err = vlt.SelectUnlocker(passphraseUnlocker.GetID())
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select unlocker: %v", err)
|
||||
}
|
||||
|
||||
// Test getting current unlocker
|
||||
currentUnlocker, err := vlt.GetCurrentUnlocker()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current unlocker: %v", err)
|
||||
}
|
||||
|
||||
if currentUnlocker.GetID() != passphraseUnlocker.GetID() {
|
||||
t.Errorf("Expected current unlocker ID '%s', got '%s'",
|
||||
passphraseUnlocker.GetID(), currentUnlocker.GetID())
|
||||
}
|
||||
}
|
||||
|
||||
func TestListUnlockers_SkipsMissingMetadata(t *testing.T) {
|
||||
// Set test environment variables
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
// Use in-memory filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
|
||||
// Create vault
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Create a passphrase unlocker so we have at least one valid unlocker
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte(testPassphrase))
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
_, err = vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create passphrase unlocker: %v", err)
|
||||
}
|
||||
|
||||
// Create a bogus unlocker directory with no metadata file
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
bogusDir := filepath.Join(vaultDir, "unlockers.d", "bogus-no-metadata")
|
||||
|
||||
err = fs.MkdirAll(bogusDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create bogus directory: %v", err)
|
||||
}
|
||||
|
||||
// ListUnlockers should succeed, skipping the bogus directory
|
||||
unlockers, err := vlt.ListUnlockers()
|
||||
if err != nil {
|
||||
t.Fatalf("ListUnlockers returned error when it should have skipped "+
|
||||
"bad directory: %v", err)
|
||||
}
|
||||
|
||||
// Should still have the valid passphrase unlocker
|
||||
if len(unlockers) == 0 {
|
||||
t.Errorf("Expected at least one unlocker, got none")
|
||||
}
|
||||
|
||||
// Verify we only got the valid unlocker(s), not the bogus one
|
||||
for _, u := range unlockers {
|
||||
if u.Type == "" {
|
||||
t.Errorf("Got unlocker with empty type, likely from bogus directory")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+10
-1
@@ -9,6 +9,7 @@
|
||||
package agehd
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
@@ -28,6 +29,10 @@ const (
|
||||
x25519KeySize = 32 // 256-bit key size for X25519
|
||||
)
|
||||
|
||||
// errInvalidScalarSize is returned when the entropy is not exactly 32
|
||||
// bytes long.
|
||||
var errInvalidScalarSize = errors.New("need 32-byte scalar")
|
||||
|
||||
// clamp applies RFC-7748 clamping to a 32-byte scalar.
|
||||
func clamp(k []byte) {
|
||||
k[0] &= 248
|
||||
@@ -39,7 +44,7 @@ func clamp(k []byte) {
|
||||
// *age.X25519Identity by round-tripping through Bech32.
|
||||
func IdentityFromEntropy(ent []byte) (*age.X25519Identity, error) {
|
||||
if len(ent) != x25519KeySize {
|
||||
return nil, fmt.Errorf("need 32-byte scalar, got %d", len(ent))
|
||||
return nil, fmt.Errorf("%w, got %d", errInvalidScalarSize, len(ent))
|
||||
}
|
||||
|
||||
// Make a copy to avoid modifying the original
|
||||
@@ -51,10 +56,12 @@ func IdentityFromEntropy(ent []byte) (*age.X25519Identity, error) {
|
||||
bech32BitSize8 = 8 // Standard 8-bit encoding
|
||||
bech32BitSize5 = 5 // Bech32 5-bit encoding
|
||||
)
|
||||
|
||||
data, err := bech32.ConvertBits(key, bech32BitSize8, bech32BitSize5, true)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bech32 convert: %w", err)
|
||||
}
|
||||
|
||||
s, err := bech32.Encode(hrp, data)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bech32 encode: %w", err)
|
||||
@@ -87,6 +94,7 @@ func DeriveEntropy(mnemonic string, n uint32) ([]byte, error) {
|
||||
// Use BIP85 DRNG to generate deterministic 32 bytes for the age key
|
||||
drng := bip85.NewBIP85DRNG(entropy)
|
||||
key := make([]byte, x25519KeySize)
|
||||
|
||||
_, err = drng.Read(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read from DRNG: %w", err)
|
||||
@@ -116,6 +124,7 @@ func DeriveEntropyFromXPRV(xprv string, n uint32) ([]byte, error) {
|
||||
// Use BIP85 DRNG to generate deterministic 32 bytes for the age key
|
||||
drng := bip85.NewBIP85DRNG(entropy)
|
||||
key := make([]byte, x25519KeySize)
|
||||
|
||||
_, err = drng.Read(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read from DRNG: %w", err)
|
||||
|
||||
+310
-328
File diff suppressed because it is too large
Load Diff
+103
-30
@@ -9,6 +9,7 @@ import (
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
@@ -23,10 +24,10 @@ import (
|
||||
|
||||
const (
|
||||
// BIP85_MASTER_PATH is the derivation path prefix for all BIP85 applications
|
||||
BIP85_MASTER_PATH = "m/83696968'" //nolint:revive // ALL_CAPS used for BIP85 constants
|
||||
BIP85_MASTER_PATH = "m/83696968'" //nolint:revive // BIP85 spec naming
|
||||
|
||||
// BIP85_KEY_HMAC_KEY is the HMAC key used for deriving the entropy
|
||||
BIP85_KEY_HMAC_KEY = "bip-entropy-from-k" //nolint:revive // ALL_CAPS used for BIP85 constants
|
||||
BIP85_KEY_HMAC_KEY = "bip-entropy-from-k" //nolint:revive // BIP85 spec naming
|
||||
|
||||
// AppBIP39 is the application number for BIP39 mnemonics
|
||||
AppBIP39 = 39
|
||||
@@ -34,18 +35,50 @@ const (
|
||||
AppHDWIF = 2
|
||||
// AppXPRV is the application number for extended private key
|
||||
AppXPRV = 32
|
||||
APP_HEX = 128169 //nolint:revive // ALL_CAPS used for BIP85 constants
|
||||
APP_PWD64 = 707764 // Base64 passwords //nolint:revive // ALL_CAPS used for BIP85 constants
|
||||
APP_HEX = 128169 //nolint:revive // BIP85 spec naming
|
||||
APP_PWD64 = 707764 // Base64 passwords //nolint:revive // BIP85 spec naming
|
||||
AppPWD85 = 707785 // Base85 passwords
|
||||
APP_RSA = 828365 //nolint:revive // ALL_CAPS used for BIP85 constants
|
||||
APP_RSA = 828365 //nolint:revive // BIP85 spec naming
|
||||
)
|
||||
|
||||
// Sentinel errors for BIP85 derivation.
|
||||
var (
|
||||
// ErrNotPrivateKey is returned when the supplied master key is not a
|
||||
// private key.
|
||||
ErrNotPrivateKey = errors.New("master key must be a private key")
|
||||
// ErrInvalidPathComponent is returned when a derivation path component
|
||||
// cannot be parsed.
|
||||
ErrInvalidPathComponent = errors.New("invalid path component")
|
||||
// ErrInvalidWordCount is returned for unsupported BIP39 word counts.
|
||||
ErrInvalidWordCount = errors.New("invalid BIP39 word count")
|
||||
// ErrInvalidNumBytes is returned when numBytes is out of range.
|
||||
ErrInvalidNumBytes = errors.New("numBytes must be between 16 and 64")
|
||||
// ErrInvalidBase64PwdLen is returned when the Base64 password length
|
||||
// is out of range.
|
||||
ErrInvalidBase64PwdLen = errors.New("pwdLen must be between 20 and 86")
|
||||
// ErrInvalidBase85PwdLen is returned when the Base85 password length
|
||||
// is out of range.
|
||||
ErrInvalidBase85PwdLen = errors.New("pwdLen must be between 10 and 80")
|
||||
// ErrPasswordTooShort is returned when the derived material is
|
||||
// shorter than the requested password length. It carries only the
|
||||
// middle of the message, which the caller composes as
|
||||
// "derived password length <n> is shorter than requested length <m>",
|
||||
// so the emitted text is unchanged.
|
||||
ErrPasswordTooShort = errors.New("is shorter than requested length")
|
||||
// ErrEncodedTooShort is returned when the encoded material is shorter
|
||||
// than the requested password length. Composed as
|
||||
// "encoded length <n> is less than requested length <m>".
|
||||
ErrEncodedTooShort = errors.New("is less than requested length")
|
||||
)
|
||||
|
||||
// Version bytes for extended keys
|
||||
//
|
||||
//nolint:gochecknoglobals // standard BIP32 version constants
|
||||
var (
|
||||
// MainNetPrivateKey is the version for mainnet private keys
|
||||
MainNetPrivateKey = []byte{0x04, 0x88, 0xAD, 0xE4} //nolint:gochecknoglobals // Standard BIP32 constant
|
||||
MainNetPrivateKey = []byte{0x04, 0x88, 0xAD, 0xE4}
|
||||
// TestNetPrivateKey is the version for testnet private keys
|
||||
TestNetPrivateKey = []byte{0x04, 0x35, 0x83, 0x94} //nolint:gochecknoglobals // Standard BIP32 constant
|
||||
TestNetPrivateKey = []byte{0x04, 0x35, 0x83, 0x94}
|
||||
)
|
||||
|
||||
// DRNG is a deterministic random number generator seeded by BIP85 entropy
|
||||
@@ -71,7 +104,7 @@ func NewBIP85DRNG(entropy []byte) *DRNG {
|
||||
}
|
||||
|
||||
// Read implements the io.Reader interface
|
||||
func (d *DRNG) Read(p []byte) (n int, err error) {
|
||||
func (d *DRNG) Read(p []byte) (int, error) {
|
||||
return d.shake.Read(p)
|
||||
}
|
||||
|
||||
@@ -79,7 +112,7 @@ func (d *DRNG) Read(p []byte) (n int, err error) {
|
||||
func DeriveChildKey(masterKey *hdkeychain.ExtendedKey, path string) ([]byte, error) {
|
||||
// Validate the masterKey is a private key
|
||||
if !masterKey.IsPrivate() {
|
||||
return nil, fmt.Errorf("master key must be a private key")
|
||||
return nil, ErrNotPrivateKey
|
||||
}
|
||||
|
||||
// Derive the child key at the specified path
|
||||
@@ -98,8 +131,12 @@ func DeriveChildKey(masterKey *hdkeychain.ExtendedKey, path string) ([]byte, err
|
||||
return ecPrivKey.Serialize(), nil
|
||||
}
|
||||
|
||||
// DeriveBIP85Entropy derives entropy from a BIP32 master key using the BIP85 method
|
||||
func DeriveBIP85Entropy(masterKey *hdkeychain.ExtendedKey, path string) ([]byte, error) {
|
||||
// DeriveBIP85Entropy derives entropy from a BIP32 master key using the
|
||||
// BIP85 method
|
||||
func DeriveBIP85Entropy(
|
||||
masterKey *hdkeychain.ExtendedKey,
|
||||
path string,
|
||||
) ([]byte, error) {
|
||||
// Get the child key bytes
|
||||
privKeyBytes, err := DeriveChildKey(masterKey, path)
|
||||
if err != nil {
|
||||
@@ -115,7 +152,10 @@ func DeriveBIP85Entropy(masterKey *hdkeychain.ExtendedKey, path string) ([]byte,
|
||||
}
|
||||
|
||||
// deriveChildKey derives a child key from a parent key using the given path
|
||||
func deriveChildKey(parent *hdkeychain.ExtendedKey, path string) (*hdkeychain.ExtendedKey, error) {
|
||||
func deriveChildKey(
|
||||
parent *hdkeychain.ExtendedKey,
|
||||
path string,
|
||||
) (*hdkeychain.ExtendedKey, error) {
|
||||
if path == "" || path == "m" || path == "/" {
|
||||
return parent, nil
|
||||
}
|
||||
@@ -141,9 +181,12 @@ func deriveChildKey(parent *hdkeychain.ExtendedKey, path string) (*hdkeychain.Ex
|
||||
|
||||
// Parse the index
|
||||
var index uint32
|
||||
|
||||
_, err := fmt.Sscanf(component, "%d", &index)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid path component: %s", component)
|
||||
return nil, fmt.Errorf(
|
||||
"%w: %s", ErrInvalidPathComponent, component,
|
||||
)
|
||||
}
|
||||
|
||||
// Apply hardening if needed
|
||||
@@ -164,8 +207,14 @@ func deriveChildKey(parent *hdkeychain.ExtendedKey, path string) (*hdkeychain.Ex
|
||||
}
|
||||
|
||||
// DeriveBIP39Entropy derives entropy for a BIP39 mnemonic
|
||||
func DeriveBIP39Entropy(masterKey *hdkeychain.ExtendedKey, language, words, index uint32) ([]byte, error) {
|
||||
path := fmt.Sprintf("%s/%d'/%d'/%d'/%d'", BIP85_MASTER_PATH, AppBIP39, language, words, index)
|
||||
func DeriveBIP39Entropy(
|
||||
masterKey *hdkeychain.ExtendedKey,
|
||||
language, words, index uint32,
|
||||
) ([]byte, error) {
|
||||
path := fmt.Sprintf(
|
||||
"%s/%d'/%d'/%d'/%d'",
|
||||
BIP85_MASTER_PATH, AppBIP39, language, words, index,
|
||||
)
|
||||
|
||||
entropy, err := DeriveBIP85Entropy(masterKey, path)
|
||||
if err != nil {
|
||||
@@ -183,6 +232,7 @@ func DeriveBIP39Entropy(masterKey *hdkeychain.ExtendedKey, language, words, inde
|
||||
)
|
||||
|
||||
var bits int
|
||||
|
||||
switch words {
|
||||
case words12:
|
||||
bits = 128
|
||||
@@ -195,7 +245,7 @@ func DeriveBIP39Entropy(masterKey *hdkeychain.ExtendedKey, language, words, inde
|
||||
case words24:
|
||||
bits = 256
|
||||
default:
|
||||
return nil, fmt.Errorf("invalid BIP39 word count: %d", words)
|
||||
return nil, fmt.Errorf("%w: %d", ErrInvalidWordCount, words)
|
||||
}
|
||||
|
||||
// Truncate to the required number of bits (bytes = bits / 8)
|
||||
@@ -218,6 +268,7 @@ func DeriveWIFKey(masterKey *hdkeychain.ExtendedKey, index uint32) (string, erro
|
||||
|
||||
// Convert to WIF format
|
||||
privKey, _ := btcec.PrivKeyFromBytes(keyBytes)
|
||||
|
||||
wif, err := btcutil.NewWIF(privKey, &chaincfg.MainNetParams, true) // compressed=true
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to create WIF: %w", err)
|
||||
@@ -227,7 +278,10 @@ func DeriveWIFKey(masterKey *hdkeychain.ExtendedKey, index uint32) (string, erro
|
||||
}
|
||||
|
||||
// DeriveXPRV derives an extended private key (XPRV)
|
||||
func DeriveXPRV(masterKey *hdkeychain.ExtendedKey, index uint32) (*hdkeychain.ExtendedKey, error) {
|
||||
func DeriveXPRV(
|
||||
masterKey *hdkeychain.ExtendedKey,
|
||||
index uint32,
|
||||
) (*hdkeychain.ExtendedKey, error) {
|
||||
path := fmt.Sprintf("%s/%d'/%d'", BIP85_MASTER_PATH, AppXPRV, index)
|
||||
|
||||
entropy, err := DeriveBIP85Entropy(masterKey, path)
|
||||
@@ -266,10 +320,10 @@ func DeriveXPRV(masterKey *hdkeychain.ExtendedKey, index uint32) (*hdkeychain.Ex
|
||||
checksum := doubleSHA256(serializedBytes)[:4]
|
||||
|
||||
// Append checksum
|
||||
serializedWithChecksum := append(serializedBytes, checksum...)
|
||||
serializedBytes = append(serializedBytes, checksum...)
|
||||
|
||||
// Base58 encode
|
||||
xprvStr := base58.Encode(serializedWithChecksum)
|
||||
xprvStr := base58.Encode(serializedBytes)
|
||||
|
||||
// Parse the serialized xprv back to an ExtendedKey
|
||||
return hdkeychain.NewKeyFromString(xprvStr)
|
||||
@@ -284,9 +338,12 @@ func doubleSHA256(data []byte) []byte {
|
||||
}
|
||||
|
||||
// DeriveHex derives a raw hex string of specified length
|
||||
func DeriveHex(masterKey *hdkeychain.ExtendedKey, numBytes, index uint32) (string, error) {
|
||||
func DeriveHex(
|
||||
masterKey *hdkeychain.ExtendedKey,
|
||||
numBytes, index uint32,
|
||||
) (string, error) {
|
||||
if numBytes < 16 || numBytes > 64 {
|
||||
return "", fmt.Errorf("numBytes must be between 16 and 64")
|
||||
return "", ErrInvalidNumBytes
|
||||
}
|
||||
|
||||
path := fmt.Sprintf("%s/%d'/%d'/%d'", BIP85_MASTER_PATH, APP_HEX, numBytes, index)
|
||||
@@ -303,9 +360,12 @@ func DeriveHex(masterKey *hdkeychain.ExtendedKey, numBytes, index uint32) (strin
|
||||
}
|
||||
|
||||
// DeriveBase64Password derives a password encoded in Base64
|
||||
func DeriveBase64Password(masterKey *hdkeychain.ExtendedKey, pwdLen, index uint32) (string, error) {
|
||||
func DeriveBase64Password(
|
||||
masterKey *hdkeychain.ExtendedKey,
|
||||
pwdLen, index uint32,
|
||||
) (string, error) {
|
||||
if pwdLen < 20 || pwdLen > 86 {
|
||||
return "", fmt.Errorf("pwdLen must be between 20 and 86")
|
||||
return "", ErrInvalidBase64PwdLen
|
||||
}
|
||||
|
||||
path := fmt.Sprintf("%s/%d'/%d'/%d'", BIP85_MASTER_PATH, APP_PWD64, pwdLen, index)
|
||||
@@ -323,16 +383,22 @@ func DeriveBase64Password(masterKey *hdkeychain.ExtendedKey, pwdLen, index uint3
|
||||
|
||||
// Slice to the desired password length
|
||||
if len(encodedStr) < int(pwdLen) {
|
||||
return "", fmt.Errorf("derived password length %d is shorter than requested length %d", len(encodedStr), pwdLen)
|
||||
return "", fmt.Errorf(
|
||||
"derived password length %d %w %d",
|
||||
len(encodedStr), ErrPasswordTooShort, pwdLen,
|
||||
)
|
||||
}
|
||||
|
||||
return encodedStr[:pwdLen], nil
|
||||
}
|
||||
|
||||
// DeriveBase85Password derives a password encoded in Base85
|
||||
func DeriveBase85Password(masterKey *hdkeychain.ExtendedKey, pwdLen, index uint32) (string, error) {
|
||||
func DeriveBase85Password(
|
||||
masterKey *hdkeychain.ExtendedKey,
|
||||
pwdLen, index uint32,
|
||||
) (string, error) {
|
||||
if pwdLen < 10 || pwdLen > 80 {
|
||||
return "", fmt.Errorf("pwdLen must be between 10 and 80")
|
||||
return "", ErrInvalidBase85PwdLen
|
||||
}
|
||||
|
||||
path := fmt.Sprintf("%s/%d'/%d'/%d'", BIP85_MASTER_PATH, AppPWD85, pwdLen, index)
|
||||
@@ -347,16 +413,21 @@ func DeriveBase85Password(masterKey *hdkeychain.ExtendedKey, pwdLen, index uint3
|
||||
|
||||
// Slice to the desired password length
|
||||
if len(encoded) < int(pwdLen) {
|
||||
return "", fmt.Errorf("encoded length %d is less than requested length %d", len(encoded), pwdLen)
|
||||
return "", fmt.Errorf(
|
||||
"encoded length %d %w %d",
|
||||
len(encoded), ErrEncodedTooShort, pwdLen,
|
||||
)
|
||||
}
|
||||
|
||||
return encoded[:pwdLen], nil
|
||||
}
|
||||
|
||||
// encodeBase85WithRFC1924Charset encodes data using Base85 with the RFC1924 character set
|
||||
// encodeBase85WithRFC1924Charset encodes data using Base85 with the
|
||||
// RFC1924 character set
|
||||
func encodeBase85WithRFC1924Charset(data []byte) string {
|
||||
// RFC1924 character set
|
||||
charset := "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz!#$%&()*+-;<=>?@^_`{|}~"
|
||||
charset := "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ" +
|
||||
"abcdefghijklmnopqrstuvwxyz!#$%&()*+-;<=>?@^_`{|}~"
|
||||
|
||||
const (
|
||||
base85ChunkSize = 4 // Process 4 bytes at a time
|
||||
@@ -369,7 +440,9 @@ func encodeBase85WithRFC1924Charset(data []byte) string {
|
||||
copy(padded, data)
|
||||
|
||||
var buf strings.Builder
|
||||
buf.Grow(len(padded) * base85DigitCount / base85ChunkSize) // Each 4 bytes becomes 5 Base85 characters
|
||||
|
||||
// Each 4 bytes becomes 5 Base85 characters
|
||||
buf.Grow(len(padded) * base85DigitCount / base85ChunkSize)
|
||||
|
||||
// Process in 4-byte chunks
|
||||
for i := 0; i < len(padded); i += base85ChunkSize {
|
||||
|
||||
+566
-499
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user