Compare commits

...
2 Commits
Author SHA1 Message Date
sneak 4a32252504 Type-check and lint the macOS build from Linux (closes #50)
check / check (push) Failing after 2s
script/lint-darwin (make lint-darwin; run by script/check, and its
commands by the Dockerfile lint stage) runs go vet and golangci-lint
with GOOS=darwin and cgo off. Compiling cgo for macOS needs Apple's SDK,
so the three functions that call go-keychain, which is cgo there, move
to keychainunlocker_cgo.go; a macOS build without cgo gets
keychainunlocker_nocgo.go and the macse stub, whose errors name the
missing macOS build with cgo. The rest of the keychain unlocker and its
plain-Go tests are now checked; their findings are fixed without
changing behaviour, and lines over 88 columns in the unchecked files
are wrapped.

Model: opus-5-5
2026-10-04 15:59:24 +00:00
clawbot 1cc8653981 Ask before removing a secret, version, vault or unlocker (closes #39)
check / check (push) Failing after 2s
secret rm, secret version rm, secret vault remove and secret unlocker
remove ask [y/N] on a terminal, naming what they remove, and go ahead
only on y or yes. Without --force, a command whose stdin is not a
terminal fails at once. --force, now also on rm and version rm,
removes without asking; it replaces the old refusals to remove a vault
with secrets or the last unlocker without --force. The checks run and
the question is asked before the state directory lock is taken; under
the lock the checks run again, and nothing is removed if they would
ask a different question.

Model: opus-5-5
2026-10-04 17:41:48 +02:00
34 changed files with 2202 additions and 1025 deletions
+4 -1
View File
@@ -14,8 +14,11 @@ ARG CHECK_EPOCH
COPY . . COPY . .
RUN make fmt-check RUN make fmt-check
# Not make lint: script/lint is a docker build, which cannot run in here. # Not make lint or make lint-darwin: script/lint and script/lint-darwin are
# docker builds, which cannot run in here. These are their commands.
RUN golangci-lint run --config .golangci.yml ./... RUN golangci-lint run --config .golangci.yml ./...
RUN GOOS=darwin CGO_ENABLED=0 go vet ./...
RUN GOOS=darwin CGO_ENABLED=0 golangci-lint run --config .golangci.yml ./...
# Build stage — tests and compilation # Build stage — tests and compilation
# golang 1.24.13-alpine (2026-03-10) # golang 1.24.13-alpine (2026-03-10)
+13 -3
View File
@@ -1,6 +1,6 @@
# Lint image, built by script/lint: golangci-lint runs as a build step, so a # Lint image, built by script/lint and script/lint-darwin: golangci-lint runs
# successful build is a clean lint. Works where the docker daemon is remote # as a build step, so a successful build is a clean lint. Works where the
# and bind mounts are impossible. # docker daemon is remote and bind mounts are impossible.
# golangci/golangci-lint:v2.12.2 (Debian-based), 2026-08-07 # golangci/golangci-lint:v2.12.2 (Debian-based), 2026-08-07
FROM golangci/golangci-lint:v2.12.2@sha256:5cceeef04e53efe1470638d4b4b4f5ceefd574955ab3941b2d9a68a8c9ad5240 AS deps FROM golangci/golangci-lint:v2.12.2@sha256:5cceeef04e53efe1470638d4b4b4f5ceefd574955ab3941b2d9a68a8c9ad5240 AS deps
@@ -17,3 +17,13 @@ FROM deps AS lint
COPY . . COPY . .
RUN golangci-lint run --config .golangci.yml ./... RUN golangci-lint run --config .golangci.yml ./...
# script/lint-darwin rebuilds this stage on every run, by this name. It
# checks the code as a macOS build compiles it, but with cgo off, which
# leaves out the files that need cgo on macOS (see script/lint-darwin).
FROM deps AS lint-darwin
COPY . .
RUN GOOS=darwin CGO_ENABLED=0 go vet ./...
RUN GOOS=darwin CGO_ENABLED=0 golangci-lint run --config .golangci.yml ./...
+6 -2
View File
@@ -1,7 +1,7 @@
export CGO_ENABLED=1 export CGO_ENABLED=1
.PHONY: default bootstrap setup build test lint fmt fmt-check check docker \ .PHONY: default bootstrap setup build test lint lint-darwin fmt fmt-check \
docker-run clean install hooks check docker docker-run clean install hooks
default: check default: check
@@ -24,6 +24,10 @@ fmt:
lint: lint:
@script/lint @script/lint
# Type-check and lint the macOS build from Linux (see script/lint-darwin)
lint-darwin:
@script/lint-darwin
check: check:
@script/check @script/check
+59 -23
View File
@@ -70,6 +70,24 @@ make build
## Commands Reference ## Commands Reference
### Confirmation Before Removal
`secret rm`, `secret version rm`, `secret vault remove` and
`secret unlocker remove` destroy data that exists nowhere else. On a terminal
each one first asks `[y/N]`, naming exactly what it is about to remove, and
goes ahead only on `y` or `yes`; any other answer, a bare Enter included,
cancels and removes nothing. The question is asked only after the command's
checks have passed, and before it changes anything.
Whether to ask is decided by stdin, where the answer is read from, so
`secret rm foo | tee log` still asks. When stdin is not a terminal, as in a
script or a CI job, nobody is there to answer: the command fails at once,
removes nothing, and says to pass `--force`.
`--force` (`-f`) removes without asking, whatever the command removes: a vault
that holds secrets and the last unlocker of a vault included. Scripts that
remove things pass `--force`.
### Initialization ### Initialization
#### `secret init` #### `secret init`
@@ -100,13 +118,13 @@ Switches to the specified vault for subsequent operations.
#### `secret vault remove <name> [--force]` / `secret vault rm` ⚠️ 🛑 #### `secret vault remove <name> [--force]` / `secret vault rm` ⚠️ 🛑
**DANGER**: Permanently removes a vault and all its secrets. Like Unix `rm`, **DANGER**: Permanently removes a vault and all its secrets. It first asks
this command does not ask for confirmation. for confirmation, naming the vault and how many secrets it holds (see
[Confirmation Before Removal](#confirmation-before-removal)). The last vault
cannot be removed. Removing the current vault makes another vault the current
one.
Requires --force if the vault contains secrets. With --force, will - `--force, -f`: Remove without asking, also a vault that contains secrets
automatically switch to another vault if removing the current one.
- `--force, -f`: Force removal even if vault contains secrets
- **NO RECOVERY**: All secrets in the vault will be permanently deleted - **NO RECOVERY**: All secrets in the vault will be permanently deleted
### Secret Management ### Secret Management
@@ -132,9 +150,12 @@ Retrieves and outputs a secret value to stdout.
Lists all secrets in the current vault. Optional filter for substring Lists all secrets in the current vault. Optional filter for substring
matching. matching.
#### `secret remove <secret-name>` / `secret rm` ⚠️ 🛑 #### `secret remove <secret-name> [--force]` / `secret rm` ⚠️ 🛑
**DANGER**: Permanently removes a secret and ALL its versions. Like Unix `rm`, this command does not ask for confirmation. **DANGER**: Permanently removes a secret and ALL its versions. It first asks
for confirmation, naming the secret, its vault and how many versions it has
(see [Confirmation Before Removal](#confirmation-before-removal)).
- `--force, -f`: Remove without asking
- **NO RECOVERY**: Once removed, the secret cannot be recovered - **NO RECOVERY**: Once removed, the secret cannot be recovered
- **ALL VERSIONS DELETED**: Every version of the secret will be permanently deleted - **ALL VERSIONS DELETED**: Every version of the secret will be permanently deleted
@@ -158,10 +179,12 @@ Lists all versions of a secret showing creation time, status, and validity perio
Promotes a specific version to current by updating the symlink. Does not Promotes a specific version to current by updating the symlink. Does not
modify any timestamps, allowing for rollback scenarios. modify any timestamps, allowing for rollback scenarios.
#### `secret version remove <secret-name> <version>` / `secret version rm` ⚠️ 🛑 #### `secret version remove <secret-name> <version> [--force]` / `secret version rm` ⚠️ 🛑
**DANGER**: Permanently removes a specific version of a secret. Like Unix **DANGER**: Permanently removes a specific version of a secret. It first asks
`rm`, this command does not ask for confirmation. for confirmation, naming the version, the secret and its vault (see
[Confirmation Before Removal](#confirmation-before-removal)).
- `--force, -f`: Remove without asking
- **NO RECOVERY**: Once removed, this version cannot be recovered - **NO RECOVERY**: Once removed, this version cannot be recovered
- Cannot remove the current version (must promote another version first) - Cannot remove the current version (must promote another version first)
@@ -202,12 +225,15 @@ has, which is removed only once the new one is the current unlocker.
#### `secret unlocker remove <unlocker-id> [--force]` / `secret unlocker rm` ⚠️ 🛑 #### `secret unlocker remove <unlocker-id> [--force]` / `secret unlocker rm` ⚠️ 🛑
**DANGER**: Permanently removes an unlocker. Like Unix `rm`, this command **DANGER**: Permanently removes an unlocker. It first asks for confirmation,
does not ask for confirmation. Cannot remove the last unlocker if the vault naming the unlocker and its vault and saying whether it is the vault's last
has secrets unless --force is used. An unlocker directory that unlocker; for the last one it says how many secrets the vault holds and warns
`secret unlocker list` skips with a warning, because its metadata cannot be that the vault then opens only with its mnemonic (see
read or parsed, is removed by the directory name the warning gives. [Confirmation Before Removal](#confirmation-before-removal)). An unlocker
- `--force, -f`: Force removal of last unlocker even if vault has secrets 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`: Remove without asking, even the last unlocker
- **CRITICAL WARNING**: Without unlockers and without your mnemonic phrase, - **CRITICAL WARNING**: Without unlockers and without your mnemonic phrase,
vault data will be PERMANENTLY INACCESSIBLE vault data will be PERMANENTLY INACCESSIBLE
- **NO RECOVERY**: Removing all unlockers without having your mnemonic means - **NO RECOVERY**: Removing all unlockers without having your mnemonic means
@@ -377,7 +403,7 @@ secret list
secret get database/prod/password secret get database/prod/password
secret get services/api/key secret get services/api/key
# Remove a secret ⚠️ 🛑 (NO CONFIRMATION - PERMANENT!) # Remove a secret ⚠️ 🛑 (asks first - PERMANENT!)
secret remove ssh/servers/web01 secret remove ssh/servers/web01
``` ```
@@ -400,7 +426,7 @@ echo "personal-email-pass" | secret add email/password
# List all vaults # List all vaults
secret vault list secret vault list
# Remove a vault ⚠️ 🛑 (NO CONFIRMATION - PERMANENT!) # Remove a vault ⚠️ 🛑 (--force: NO CONFIRMATION - PERMANENT!)
secret vault remove personal --force secret vault remove personal --force
``` ```
@@ -418,7 +444,7 @@ secret unlocker list
# Select a specific unlocker # Select a specific unlocker
secret unlocker select <unlocker-id> secret unlocker select <unlocker-id>
# Remove an unlocker ⚠️ 🛑 (NO CONFIRMATION!) # Remove an unlocker ⚠️ 🛑 (asks first!)
secret unlocker remove <unlocker-id> secret unlocker remove <unlocker-id>
``` ```
@@ -431,7 +457,7 @@ secret version list database/prod/password
# Promote an older version to current # Promote an older version to current
secret version promote database/prod/password 20231215.001 secret version promote database/prod/password 20231215.001
# Remove an old version ⚠️ 🛑 (NO CONFIRMATION - PERMANENT!) # Remove an old version ⚠️ 🛑 (asks first - PERMANENT!)
secret version remove database/prod/password 20231214.001 secret version remove database/prod/password 20231214.001
``` ```
@@ -472,6 +498,10 @@ secret decrypt encryption/mykey --input document.txt.age --output document.txt
- **macOS**: Full support including Keychain and Secure Enclave integration - **macOS**: Full support including Keychain and Secure Enclave integration
- **Linux**: Full support (excluding macOS-specific features) - **Linux**: Full support (excluding macOS-specific features)
The keychain and Secure Enclave unlockers need a macOS build with cgo. A macOS
build without cgo, such as one cross-compiled from Linux, offers them but fails
to add or use them.
## Security Considerations ## Security Considerations
### Threat Model ### Threat Model
@@ -534,10 +564,16 @@ them. We provide:
- `script/lint` — run `golangci-lint` in docker only: builds - `script/lint` — run `golangci-lint` in docker only: builds
`Dockerfile.lint`, where the linter is a build step that runs on every `Dockerfile.lint`, where the linter is a build step that runs on every
call, also on an unchanged tree call, also on an unchanged tree
- `script/lint-darwin` — run `go vet` and `golangci-lint` in docker on
the code as a macOS build compiles it (`GOOS=darwin`), which a Linux
build never compiles; cgo is off, so the keychain unlocker's calls into
the keychain (`internal/secret/keychainunlocker_cgo.go`, and
`keychainunlocker_test.go`) and the Secure Enclave bindings
(`internal/macse`) are not checked
- `script/fmt` — format all Go code (writes) - `script/fmt` — format all Go code (writes)
- `script/fmt-check` — check formatting without writing - `script/fmt-check` — check formatting without writing
- `script/check` — run `script/test`, `script/lint`, and - `script/check` — run `script/test`, `script/lint`,
`script/fmt-check` `script/lint-darwin`, and `script/fmt-check`
- `script/docker` — build the Docker image tagged with the project name - `script/docker` — build the Docker image tagged with the project name
- `script/cibuild` — CI entrypoint: `docker build --ulimit - `script/cibuild` — CI entrypoint: `docker build --ulimit
memlock=-1:-1 .` (memguard needs mlock; the Dockerfile runs the memlock=-1:-1 .` (memguard needs mlock; the Dockerfile runs the
+43 -7
View File
@@ -25,6 +25,41 @@ Bring the repo into policy compliance in one commit:
# Completed Steps # Completed Steps
- 2026-10-04: `script/lint-darwin` (`make lint-darwin`) runs `go vet` and
`golangci-lint` in docker on the code as a macOS build compiles it
(`GOOS=darwin`), with cgo off
(https://git.eeqj.de/sneak/secret/issues/50). `script/check` runs it, and
the `Dockerfile` lint stage runs its commands, so `script/cibuild` does too.
Before, CI on Linux never compiled the files built only for macOS. Compiling
cgo code for macOS needs Apple's SDK headers, and both `internal/macse` and
`github.com/keybase/go-keychain` are cgo on macOS. So the three functions
that call `go-keychain` moved from `keychainunlocker.go` to
`keychainunlocker_cgo.go`, built only with cgo on macOS like
`macse_darwin.go`. A macOS build without cgo, which before did not compile,
gets `keychainunlocker_nocgo.go` and the `macse` stub instead, whose errors
say the keychain or Secure Enclave needs a macOS build with cgo. The check
covers the rest of the keychain unlocker, the Secure Enclave unlocker and
the macOS-only tests other than `keychainunlocker_test.go`, whose lint
findings are fixed. For the length and complexity limits, parts of
`GetIdentity`, `getLongTermPrivateKey` and `CreateKeychainUnlocker` moved
into functions of their own, and the Secure Enclave unlocker derives the
long-term key from the mnemonic through the same function as the keychain
unlocker instead of a copy of it. Lines over 88 columns in the files the
check cannot see are wrapped.
- 2026-10-04: `secret rm`, `secret version rm`, `secret vault remove` and
`secret unlocker remove` ask `[y/N]` before removing anything
(https://git.eeqj.de/sneak/secret/issues/39), naming what they remove: the
secret, its vault and its version count; the version, secret and vault; the
vault and its secret count; the unlocker, its vault and whether it is the
last, and for the last the vault's secret count and that the vault then
opens only with its mnemonic. Only `y` or `yes` goes ahead. Without
`--force`, a command whose stdin is not a terminal fails at once. `--force`
(now also on `rm` and `version rm`) removes without asking; it replaces the
old refusals to remove a vault with secrets or the last unlocker of one
without `--force`, which the question now covers. The checks run, and the
question is asked, before the state directory lock is taken; under the
lock the checks run again, and if they would ask a different question,
nothing is removed. `secret rm` fails when it cannot count the versions.
- 2026-10-04: A crash while an unlocker is being replaced no longer leaves a - 2026-10-04: A crash while an unlocker is being replaced no longer leaves a
current unlocker that cannot open the vault current unlocker that cannot open the vault
(https://git.eeqj.de/sneak/secret/issues/71). Every new unlocker gets a (https://git.eeqj.de/sneak/secret/issues/71). Every new unlocker gets a
@@ -283,11 +318,14 @@ Bring the repo into policy compliance in one commit:
- Cover mnemonic-vs-xprv identity consistency in - Cover mnemonic-vs-xprv identity consistency in
`pkg/agehd/agehd_test.go` `TestMnemonicVsXPRVConsistency` (was an `pkg/agehd/agehd_test.go` `TestMnemonicVsXPRVConsistency` (was an
in-code FIXME removed for godox). in-code FIXME removed for godox).
- Darwin-gated files (`internal/secret/keychainunlocker.go`, - CI does not compile, lint or test the files built only with cgo on
`seunlocker_darwin.go`, `internal/macse/macse_darwin.go`, related macOS, since compiling them needs Apple's SDK:
tests) are not linted on the Linux CI runner and still contain lines `internal/secret/keychainunlocker_cgo.go` (the three functions that call
over the new 88-column limit; they will surface if lint ever runs on `go-keychain`) with `keychainunlocker_test.go`, and `internal/macse`
macOS. (`macse_darwin.go`, `macse_test.go`, the Objective-C sources). Lint has
never run on them, so it would likely find more there than the line
lengths. No macOS test runs in CI. A macOS runner would cover all of it
(asked on https://git.eeqj.de/sneak/secret/issues/50).
- Merge secure-enclave-unlocker to main once review is done. - Merge secure-enclave-unlocker to main once review is done.
- 1.0 critical security blockers (from repo TODO.md): - 1.0 critical security blockers (from repo TODO.md):
- Command injection: GPG key IDs passed unescaped to exec.Command - Command injection: GPG key IDs passed unescaped to exec.Command
@@ -304,8 +342,6 @@ Bring the repo into policy compliance in one commit:
- High priority: - High priority:
- Secure temporary file handling and cleanup. - Secure temporary file handling and cleanup.
- Initialize a default unlock key at vault creation. - Initialize a default unlock key at vault creation.
- Confirmation prompts for destructive operations (keys rm, vault
deletion).
- Add secret rm and vault deletion commands. - Add secret rm and vault deletion commands.
- Medium priority: - Medium priority:
- Standardize error messages; stop leaking internals. - Standardize error messages; stop leaking internals.
+1
View File
@@ -9,6 +9,7 @@ require (
github.com/btcsuite/btcd/btcec/v2 v2.1.3 github.com/btcsuite/btcd/btcec/v2 v2.1.3
github.com/btcsuite/btcd/btcutil v1.1.6 github.com/btcsuite/btcd/btcutil v1.1.6
github.com/btcsuite/btcutil v0.0.0-20190425235716-9e5f4b9a998d github.com/btcsuite/btcutil v0.0.0-20190425235716-9e5f4b9a998d
github.com/creack/pty v1.1.24
github.com/keybase/go-keychain v0.0.0-20230307172405-3e4884637dd1 github.com/keybase/go-keychain v0.0.0-20230307172405-3e4884637dd1
github.com/oklog/ulid/v2 v2.1.1 github.com/oklog/ulid/v2 v2.1.1
github.com/spf13/afero v1.14.0 github.com/spf13/afero v1.14.0
+2
View File
@@ -35,6 +35,8 @@ github.com/btcsuite/snappy-go v1.0.0/go.mod h1:8woku9dyThutzjeg+3xrA5iCpBRH8XEEg
github.com/btcsuite/websocket v0.0.0-20150119174127-31079b680792/go.mod h1:ghJtEyQwv5/p4Mg4C0fgbePVuGr935/5ddU9Z3TmDRY= github.com/btcsuite/websocket v0.0.0-20150119174127-31079b680792/go.mod h1:ghJtEyQwv5/p4Mg4C0fgbePVuGr935/5ddU9Z3TmDRY=
github.com/btcsuite/winsvc v1.0.0/go.mod h1:jsenWakMcC0zFBFurPLEAyrnc/teJEM1O46fmI40EZs= github.com/btcsuite/winsvc v1.0.0/go.mod h1:jsenWakMcC0zFBFurPLEAyrnc/teJEM1O46fmI40EZs=
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
github.com/creack/pty v1.1.24 h1:bJrF4RRfyJnbTJqzRLHzcGaZK1NeM5kTC9jGgovnR1s=
github.com/creack/pty v1.1.24/go.mod h1:08sCNb52WyoAwi2QDyzUCTgcvVFhUzewun7wtTfvcwE=
github.com/davecgh/go-spew v0.0.0-20171005155431-ecdeabc65495/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v0.0.0-20171005155431-ecdeabc65495/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
+5
View File
@@ -3,6 +3,7 @@ package cli
import ( import (
"fmt" "fmt"
"io"
"os" "os"
"git.eeqj.de/sneak/secret/internal/secret" "git.eeqj.de/sneak/secret/internal/secret"
@@ -21,6 +22,10 @@ type Instance struct {
// none. // none.
Mnemonic *memguard.LockedBuffer Mnemonic *memguard.LockedBuffer
UnlockPassphrase *memguard.LockedBuffer UnlockPassphrase *memguard.LockedBuffer
// terminal, when set, stands in for the terminal that confirm reads
// the user's answer from; only tests set it. When it is nil, confirm
// reads stdin, and only when stdin is a terminal.
terminal io.Reader
} }
// NewCLIInstance creates a new CLI instance with the real filesystem // NewCLIInstance creates a new CLI instance with the real filesystem
+108
View File
@@ -0,0 +1,108 @@
package cli
import (
"bufio"
"errors"
"fmt"
"io"
"os"
"strings"
"git.eeqj.de/sneak/secret/internal/vault"
"github.com/spf13/cobra"
"golang.org/x/term"
)
// Sentinel errors for asking the user to confirm a removal
var (
errNoTerminal = errors.New("stdin is not a terminal, so there is " +
"nobody to ask for confirmation; pass --force to remove without asking")
errNotConfirmed = errors.New("cancelled; nothing was removed")
errChangedWhileAsking = errors.New("what was to be removed changed " +
"while waiting for the answer; nothing was removed")
)
// askThenLock asks the user to confirm a removal, unless force is set, and
// then takes the state directory lock and returns the function that
// releases it. find makes the command's checks, keeps what it found for
// the caller to remove, and returns the question that names it. find runs
// before the question, which is asked without the lock so that no other
// command waits while the user answers, and runs again once the lock is
// taken. That run is the last, so the caller removes what find found under
// the lock. If its question then differs from the one the user answered,
// something changed in between, and askThenLock fails.
func (cli *Instance) askThenLock(
cmd *cobra.Command, force bool, find func() (string, error),
) (func(), error) {
asked := ""
if !force {
question, err := find()
if err != nil {
return nil, err
}
err = cli.confirm(cmd, question)
if err != nil {
return nil, err
}
asked = question
}
release, err := vault.LockStateDir(cli.fs, cli.stateDir)
if err != nil {
return nil, err
}
question, err := find()
if err == nil && !force && question != asked {
err = errChangedWhileAsking
}
if err != nil {
release()
return nil, err
}
return release, nil
}
// confirm asks question and returns nil only when the user answers y or
// yes; any other answer, a bare Enter included, cancels. When stdin is not
// a terminal it asks nothing and fails at once: nobody is there to answer,
// and waiting for an answer would hang a script. Stdin decides, not
// stdout, because the answer is read from stdin: `secret rm foo | tee log`
// still asks. The question goes to stderr.
func (cli *Instance) confirm(cmd *cobra.Command, question string) error {
answers := cli.terminal
if answers == nil {
answers = cmd.InOrStdin()
if !isTerminal(answers) {
return errNoTerminal
}
}
_, _ = fmt.Fprintf(cmd.ErrOrStderr(), "%s [y/N] ", question)
answer, err := bufio.NewReader(answers).ReadString('\n')
if err != nil && !errors.Is(err, io.EOF) {
return fmt.Errorf("failed to read the answer: %w", err)
}
switch strings.ToLower(strings.TrimSpace(answer)) {
case "y", "yes":
return nil
default:
return errNotConfirmed
}
}
// isTerminal reports whether r is a terminal.
func isTerminal(r io.Reader) bool {
file, ok := r.(*os.File)
return ok && term.IsTerminal(int(file.Fd()))
}
+410
View File
@@ -0,0 +1,410 @@
// Confirmation Tests
//
// `secret rm`, `secret version rm`, `secret vault remove` and
// `secret unlocker remove` ask the user to confirm on a terminal, naming
// what they are about to remove, and remove it only on y or yes. --force
// skips the question. Without --force, a command whose stdin is not a
// terminal fails at once, since nobody is there to answer.
//
// The tests answer through Instance.terminal, which stands in for a
// terminal. Without it, whether stdin is a terminal decides; the tests in
// integration_test.go that run `secret rm` on a pseudo-terminal cover that.
//nolint:testpackage // sets the unexported terminal field of Instance
package cli
import (
"bufio"
"bytes"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"testing"
"time"
"git.eeqj.de/sneak/secret/internal/secret"
"git.eeqj.de/sneak/secret/internal/vault"
"github.com/spf13/afero"
"github.com/spf13/cobra"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
const (
// confirmTestSecret is the secret the tests remove, or remove a
// version of, in the vault "work".
confirmTestSecret = "test/secret"
// lastUnlockerRemoval names the case that removes the only unlocker.
lastUnlockerRemoval = "unlocker rm, the last one"
)
// removal is one removal command, set up on its own state directory.
type removal struct {
fs afero.Fs
run func(cli *Instance, cmd *cobra.Command, force bool) error
// removed is the directory the command removes.
removed string
// question is the question the command asks.
question string
}
// newConfirmTestVaults returns an in-memory state directory with the
// vaults "other" and "work", the current one. "work" holds two versions of
// confirmTestSecret and the given number of PGP unlockers. It returns the
// directory of "work" and the older version.
func newConfirmTestVaults(
t *testing.T, unlockers int,
) (*afero.MemMapFs, string, string) {
t.Helper()
fs := &afero.MemMapFs{}
mnemonic := testMnemonicBuffer(t)
_, err := vault.CreateVault(fs, testStateDir, "other", mnemonic)
require.NoError(t, err)
vlt, err := vault.CreateVault(fs, testStateDir, "work", mnemonic)
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)
for i := range unlockers {
writePGPUnlocker(t, fs, filepath.Join(vaultDir, "unlockers.d"),
fmt.Sprintf("pgp-%d", i),
time.Date(2026, time.October, 4, 12, i, 0, 0, time.UTC),
listTestGPGKeyID+string(rune('A'+i)))
}
// ListVersions lists the newest version first.
return fs, vaultDir, versions[1]
}
// newRemoval sets up the removal the command names.
func newRemoval(t *testing.T, command string) removal {
t.Helper()
unlockers := 2
if command == lastUnlockerRemoval {
unlockers = 1
}
fs, workDir, older := newConfirmTestVaults(t, unlockers)
unlockerID := "pgp-" + listTestGPGKeyID + "A"
removeFirstUnlocker := func(cli *Instance, cmd *cobra.Command, force bool) error {
return cli.UnlockersRemove(unlockerID, force, cmd)
}
switch command {
case "rm":
return removal{
fs: fs,
run: func(cli *Instance, cmd *cobra.Command, force bool) error {
return cli.RemoveSecret(cmd, confirmTestSecret, force)
},
removed: filepath.Join(workDir, "secrets.d", "test%secret"),
question: "Permanently remove secret 'test/secret' and its 2 " +
"version(s) from vault 'work'?",
}
case "version rm":
return removal{
fs: fs,
run: func(cli *Instance, cmd *cobra.Command, force bool) error {
return cli.RemoveVersion(cmd, confirmTestSecret, older, force)
},
removed: filepath.Join(
workDir, "secrets.d", "test%secret", "versions", older),
question: "Permanently remove version " + older +
" of secret 'test/secret' from vault 'work'?",
}
case "vault rm":
return removal{
fs: fs,
run: func(cli *Instance, cmd *cobra.Command, force bool) error {
return cli.RemoveVault(cmd, "work", force)
},
removed: workDir,
question: "Permanently remove vault 'work' and its 1 secret(s)?",
}
case "unlocker rm":
return removal{
fs: fs,
run: removeFirstUnlocker,
removed: filepath.Join(workDir, "unlockers.d", "pgp-0"),
question: "Permanently remove unlocker '" + unlockerID +
"' from vault 'work'? It is not the vault's last unlocker.",
}
case lastUnlockerRemoval:
return removal{
fs: fs,
run: removeFirstUnlocker,
removed: filepath.Join(workDir, "unlockers.d", "pgp-0"),
question: "Permanently remove unlocker '" + unlockerID +
"', the last unlocker of vault 'work', which holds 1 " +
"secret(s)? Without an unlocker the vault opens only " +
"with its mnemonic.",
}
}
t.Fatalf("no removal %q", command)
return removal{}
}
// removalCommands lists the commands newRemoval sets up.
func removalCommands() []string {
return []string{
"rm", "version rm", "vault rm", "unlocker rm", lastUnlockerRemoval,
}
}
// newConfirmTestCommand returns a command whose output is discarded and
// whose stderr, where the question goes, is the returned buffer.
func newConfirmTestCommand() (*cobra.Command, *bytes.Buffer) {
var stderr bytes.Buffer
cmd := &cobra.Command{}
cmd.SetOut(io.Discard)
cmd.SetErr(&stderr)
return cmd, &stderr
}
// requireExists asserts whether the directory dir exists.
func requireExists(t *testing.T, fs afero.Fs, dir string, want bool) {
t.Helper()
exists, err := afero.DirExists(fs, dir)
require.NoError(t, err)
require.Equal(t, want, exists, dir)
}
// TestConfirmAnswers checks which answers confirm accepts: y or yes, in
// any case, around which spaces do not matter.
func TestConfirmAnswers(t *testing.T) {
t.Parallel()
for answer, want := range map[string]error{
"y\n": nil,
"Y\n": nil,
"yes\n": nil,
" YES \n": nil,
"y": nil,
"\n": errNotConfirmed,
"": errNotConfirmed,
"n\n": errNotConfirmed,
"yy\n": errNotConfirmed,
"no\ny\n": errNotConfirmed,
} {
t.Run(fmt.Sprintf("%q", answer), func(t *testing.T) {
t.Parallel()
cli := &Instance{terminal: strings.NewReader(answer)}
cmd, stderr := newConfirmTestCommand()
err := cli.confirm(cmd, "Remove it?")
require.ErrorIs(t, err, want)
assert.Equal(t, "Remove it? [y/N] ", stderr.String())
})
}
}
// TestRemovalAnsweredYesRemoves checks that each removal asks its question
// and removes what it names when the user answers y.
func TestRemovalAnsweredYesRemoves(t *testing.T) {
t.Parallel()
for _, command := range removalCommands() {
t.Run(command, func(t *testing.T) {
t.Parallel()
r := newRemoval(t, command)
requireExists(t, r.fs, r.removed, true)
cli := NewCLIInstanceWithStateDir(r.fs, testStateDir)
cli.terminal = strings.NewReader("y\n")
cmd, stderr := newConfirmTestCommand()
require.NoError(t, r.run(cli, cmd, false))
assert.Equal(t, r.question+" [y/N] ", stderr.String())
requireExists(t, r.fs, r.removed, false)
})
}
}
// TestRemovalDeclinedLeavesEverything checks that each removal changes
// nothing when the user answers anything but y or yes, a bare Enter
// included.
func TestRemovalDeclinedLeavesEverything(t *testing.T) {
t.Parallel()
for _, command := range removalCommands() {
for _, answer := range []string{"\n", "n\n", ""} {
t.Run(fmt.Sprintf("%s %q", command, answer), func(t *testing.T) {
t.Parallel()
r := newRemoval(t, command)
before := stateDirModTimes(t, r.fs)
cli := NewCLIInstanceWithStateDir(r.fs, testStateDir)
cli.terminal = strings.NewReader(answer)
cmd, stderr := newConfirmTestCommand()
err := r.run(cli, cmd, false)
require.ErrorIs(t, err, errNotConfirmed)
assert.Equal(t, r.question+" [y/N] ", stderr.String())
assert.Equal(t, before, stateDirModTimes(t, r.fs))
})
}
}
}
// TestRemovalForcedAsksNothing checks that each removal with --force
// removes what it would have named without asking, and without reading
// its input, which is not a terminal.
func TestRemovalForcedAsksNothing(t *testing.T) {
t.Parallel()
for _, command := range removalCommands() {
t.Run(command, func(t *testing.T) {
t.Parallel()
r := newRemoval(t, command)
input := strings.NewReader("n\n")
cli := NewCLIInstanceWithStateDir(r.fs, testStateDir)
cmd, stderr := newConfirmTestCommand()
cmd.SetIn(input)
require.NoError(t, r.run(cli, cmd, true))
assert.Empty(t, stderr.String(), "asked with --force")
assert.Equal(t, 2, input.Len(), "read its input with --force")
requireExists(t, r.fs, r.removed, false)
})
}
}
// TestRemovalWithoutTerminalFailsAtOnce checks that each removal without
// --force, whose input is not a terminal, fails at once telling the user
// to pass --force, and changes nothing. The input is a pipe that nobody
// writes to or closes, so reading it would block for good.
func TestRemovalWithoutTerminalFailsAtOnce(t *testing.T) {
t.Parallel()
for _, command := range removalCommands() {
t.Run(command, func(t *testing.T) {
t.Parallel()
r := newRemoval(t, command)
before := stateDirModTimes(t, r.fs)
input, inputWriter, err := os.Pipe()
require.NoError(t, err)
t.Cleanup(func() {
_ = inputWriter.Close()
_ = input.Close()
})
cli := NewCLIInstanceWithStateDir(r.fs, testStateDir)
cmd, stderr := newConfirmTestCommand()
cmd.SetIn(input)
done := make(chan error, 1)
go func() { done <- r.run(cli, cmd, false) }()
select {
case err := <-done:
require.ErrorIs(t, err, errNoTerminal)
assert.Contains(t, err.Error(), "pass --force")
case <-time.After(lockWait):
// Closing the pipe ends the read, and frees the lock if
// the command holds it.
_ = inputWriter.Close()
t.Fatal("waited for an answer on input that is not a terminal")
}
assert.Empty(t, stderr.String(), "asked without a terminal")
assert.Equal(t, before, stateDirModTimes(t, r.fs))
})
}
}
// TestRemovalAsksWithoutHoldingLock checks that while `secret rm` waits
// for its answer, another command can take the state directory lock and
// change the secret, and that the removal then removes nothing, since the
// secret is no longer what the question named.
func TestRemovalAsksWithoutHoldingLock(t *testing.T) {
t.Parallel()
r := newRemoval(t, "rm")
answers, answerWriter := io.Pipe()
questions, questionWriter := io.Pipe()
// Closing the answers ends the read if the test fails while the
// command waits for one.
t.Cleanup(func() { _ = answerWriter.Close() })
rm := NewCLIInstanceWithStateDir(r.fs, testStateDir)
rm.terminal = answers
cmd := &cobra.Command{}
cmd.SetOut(io.Discard)
cmd.SetErr(questionWriter)
done := make(chan error, 1)
go func() { done <- r.run(rm, cmd, false) }()
question, err := bufio.NewReader(questions).ReadString(']')
require.NoError(t, err)
require.Equal(t, r.question+" [y/N]", question)
// Adds a third version while rm waits for its answer.
add := NewCLIInstanceWithStateDir(r.fs, testStateDir)
add.Mnemonic = testMnemonicBuffer(t)
add.cmd = &cobra.Command{}
add.cmd.SetIn(strings.NewReader("newest"))
add.cmd.SetOut(io.Discard)
added := make(chan error, 1)
go func() { added <- add.AddSecret(confirmTestSecret, true) }()
select {
case err := <-added:
require.NoError(t, err)
case <-time.After(lockWait):
t.Fatal("secret add waited for the lock while secret rm asked")
}
_, err = answerWriter.Write([]byte("y\n"))
require.NoError(t, err)
select {
case err := <-done:
require.ErrorIs(t, err, errChangedWhileAsking)
case <-time.After(lockWait):
t.Fatal("secret rm did not finish once answered")
}
requireExists(t, r.fs, r.removed, true)
}
+149
View File
@@ -2,10 +2,13 @@
package cli_test package cli_test
import ( import (
"bufio"
"bytes"
"context" "context"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"io"
"os" "os"
"os/exec" "os/exec"
"path/filepath" "path/filepath"
@@ -16,7 +19,11 @@ import (
"git.eeqj.de/sneak/secret/internal/cli" "git.eeqj.de/sneak/secret/internal/cli"
"git.eeqj.de/sneak/secret/internal/secret" "git.eeqj.de/sneak/secret/internal/secret"
"git.eeqj.de/sneak/secret/internal/vault"
"git.eeqj.de/sneak/secret/pkg/agehd" "git.eeqj.de/sneak/secret/pkg/agehd"
"github.com/awnumar/memguard"
"github.com/creack/pty"
"github.com/spf13/afero"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
@@ -2543,3 +2550,145 @@ func copyFile(src, dst string) error {
return nil return nil
} }
// secretRmCommand makes a state directory whose vault "default" holds the
// secret "x", and returns `secret rm x` on the built binary against it, and
// the directory of "x". The vault has no unlocker, so making it derives no
// key from a passphrase.
func secretRmCommand(ctx context.Context, t *testing.T) (*exec.Cmd, string) {
t.Helper()
stateDir := t.TempDir()
mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonic))
defer mnemonic.Destroy()
vlt, err := vault.CreateVault(afero.NewOsFs(), stateDir, "default", mnemonic)
require.NoError(t, err)
value := memguard.NewBufferFromBytes([]byte("value"))
defer value.Destroy()
require.NoError(t, vlt.AddSecret("x", value, false))
//nolint:gosec // G204: test executes the freshly built secret binary
cmd := exec.CommandContext(ctx, secretBinaryPath(t), "rm", "x")
cmd.Env = []string{
secret.EnvStateDir + "=" + stateDir,
"PATH=" + os.Getenv("PATH"),
"HOME=" + os.Getenv("HOME"),
}
return cmd, filepath.Join(stateDir, "vaults.d", "default", "secrets.d", "x")
}
// TestRemoveWithoutTerminalFailsAtOnce runs `secret rm` without --force,
// with a stdin that is not a terminal and never delivers anything, as in a
// script or a CI job. It must fail at once, telling the user to pass
// --force, instead of waiting for an answer, and remove nothing.
func TestRemoveWithoutTerminalFailsAtOnce(t *testing.T) {
t.Parallel()
// Nobody writes to or closes the pipe, so reading it would block for good.
stdin, stdinWriter, err := os.Pipe()
require.NoError(t, err)
defer func() {
_ = stdinWriter.Close()
_ = stdin.Close()
}()
ctx, cancel := context.WithTimeout(t.Context(), time.Minute)
defer cancel()
cmd, secretDir := secretRmCommand(ctx, t)
cmd.Stdin = stdin
output, err := cmd.CombinedOutput()
require.NoError(t, ctx.Err(), "secret rm waited for an answer")
require.Error(t, err)
assert.Contains(t, string(output), "pass --force")
assert.DirExists(t, secretDir)
}
// The next two tests run `secret rm` with a terminal on stdin or on stdout
// and stderr, not both: whether it asks must depend on stdin alone, where
// the answer is read from. pty.Open returns the two ends of a new terminal:
// tty is the end a program uses as its terminal, and ptmx the end the test
// reads what the terminal shows from and types into.
// TestRemoveIgnoresTerminalOnStdout runs `echo y | secret rm x` at a
// terminal. stdin is a pipe, so nobody can answer there, and the command
// must fail as in a script, removing nothing.
func TestRemoveIgnoresTerminalOnStdout(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithTimeout(t.Context(), time.Minute)
defer cancel()
cmd, secretDir := secretRmCommand(ctx, t)
ptmx, tty, err := pty.Open()
require.NoError(t, err)
defer func() { _ = ptmx.Close() }()
cmd.Stdin = strings.NewReader("y\n")
cmd.Stdout = tty
cmd.Stderr = tty
require.NoError(t, cmd.Start())
_ = tty.Close()
// The read ends once secret rm has exited and so closed the terminal.
shown, _ := io.ReadAll(ptmx)
require.Error(t, cmd.Wait())
assert.Contains(t, string(shown), "pass --force")
assert.DirExists(t, secretDir)
}
// TestRemoveAsksAtTerminalOnStdin runs `secret rm x | cat` at a terminal.
// It must ask on the terminal, and remove the secret when y is typed there.
func TestRemoveAsksAtTerminalOnStdin(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithTimeout(t.Context(), time.Minute)
defer cancel()
cmd, secretDir := secretRmCommand(ctx, t)
ptmx, tty, err := pty.Open()
require.NoError(t, err)
defer func() { _ = ptmx.Close() }()
cmd.Stdin = tty
// Not a file, so exec.Cmd connects stdout through a pipe.
cmd.Stdout = io.Discard
cmd.Stderr = tty
require.NoError(t, cmd.Start())
_ = tty.Close()
var (
shown []byte
char byte
)
terminal := bufio.NewReader(ptmx)
for !bytes.HasSuffix(shown, []byte("[y/N] ")) {
char, err = terminal.ReadByte()
require.NoError(t, err, "secret rm ended without asking: %s", shown)
shown = append(shown, char)
}
_, err = ptmx.WriteString("y\n")
require.NoError(t, err)
require.NoError(t, cmd.Wait())
assert.NoDirExists(t, secretDir)
}
+11 -9
View File
@@ -241,8 +241,10 @@ func TestFailedCommandReleasesLock(t *testing.T) {
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
cli := NewCLIInstanceWithStateDir(fs, testStateDir) cli := NewCLIInstanceWithStateDir(fs, testStateDir)
// Fails once it holds the lock: there is no current vault // Fails once it holds the lock: there is no current vault. Without
err := cli.RemoveSecret(&cobra.Command{}, "missing", false) // --force it would fail before taking the lock, on the check it makes
// before asking.
err := cli.RemoveSecret(&cobra.Command{}, "missing", true)
require.Error(t, err) require.Error(t, err)
select { select {
@@ -431,8 +433,8 @@ func TestChangingCommandsWaitForLock(t *testing.T) {
{"encrypt", false, func(cli *Instance, _, _ string) error { {"encrypt", false, func(cli *Instance, _, _ string) error {
return cli.Encrypt("key", testInput, "") return cli.Encrypt("key", testInput, "")
}}, }},
{"rm", false, func(cli *Instance, _, _ string) error { {"rm --force", false, func(cli *Instance, _, _ string) error {
return cli.RemoveSecret(cli.cmd, "test/secret", false) return cli.RemoveSecret(cli.cmd, "test/secret", true)
}}, }},
{"move", false, func(cli *Instance, _, _ string) error { {"move", false, func(cli *Instance, _, _ string) error {
return cli.MoveSecret(cli.cmd, "test/secret", "moved", false) return cli.MoveSecret(cli.cmd, "test/secret", "moved", false)
@@ -440,8 +442,8 @@ func TestChangingCommandsWaitForLock(t *testing.T) {
{"version promote", false, func(cli *Instance, olderVersion, _ string) error { {"version promote", false, func(cli *Instance, olderVersion, _ string) error {
return cli.PromoteVersion(cli.cmd, "test/secret", olderVersion) return cli.PromoteVersion(cli.cmd, "test/secret", olderVersion)
}}, }},
{"version rm", false, func(cli *Instance, olderVersion, _ string) error { {"version rm --force", false, func(cli *Instance, olderVersion, _ string) error {
return cli.RemoveVersion(cli.cmd, "test/secret", olderVersion) return cli.RemoveVersion(cli.cmd, "test/secret", olderVersion, true)
}}, }},
{"vault create", false, func(cli *Instance, _, _ string) error { {"vault create", false, func(cli *Instance, _, _ string) error {
return cli.CreateVault(cli.cmd, "created") return cli.CreateVault(cli.cmd, "created")
@@ -452,13 +454,13 @@ func TestChangingCommandsWaitForLock(t *testing.T) {
{"vault import", false, func(cli *Instance, _, _ string) error { {"vault import", false, func(cli *Instance, _, _ string) error {
return cli.VaultImport(cli.cmd, "other") return cli.VaultImport(cli.cmd, "other")
}}, }},
{"vault rm", false, func(cli *Instance, _, _ string) error { {"vault rm --force", false, func(cli *Instance, _, _ string) error {
return cli.RemoveVault(cli.cmd, "other", false) return cli.RemoveVault(cli.cmd, "other", true)
}}, }},
{"unlocker add", false, func(cli *Instance, _, _ string) error { {"unlocker add", false, func(cli *Instance, _, _ string) error {
return cli.UnlockersAdd("passphrase", cli.cmd) return cli.UnlockersAdd("passphrase", cli.cmd)
}}, }},
{"unlocker rm", true, func(cli *Instance, _, unlockerID string) error { {"unlocker rm --force", true, func(cli *Instance, _, unlockerID string) error {
return cli.UnlockersRemove(unlockerID, true, cli.cmd) return cli.UnlockersRemove(unlockerID, true, cli.cmd)
}}, }},
{"unlocker select", true, func(cli *Instance, _, unlockerID string) error { {"unlocker select", true, func(cli *Instance, _, unlockerID string) error {
+18 -18
View File
@@ -170,8 +170,8 @@ func requireRejectedAndUnchanged(
// TestInvalidSecretNameLeavesVaultsUnchanged is a regression test for // TestInvalidSecretNameLeavesVaultsUnchanged is a regression test for
// https://git.eeqj.de/sneak/secret/issues/33, where `secret rm ..` deleted // 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. // 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 // Removals, moves and imports use --force, so that only the name check
// the way. // stands in the way.
// //
//nolint:paralleltest // the cases share cmd //nolint:paralleltest // the cases share cmd
func TestInvalidSecretNameLeavesVaultsUnchanged(t *testing.T) { func TestInvalidSecretNameLeavesVaultsUnchanged(t *testing.T) {
@@ -191,17 +191,17 @@ func TestInvalidSecretNameLeavesVaultsUnchanged(t *testing.T) {
rejected string // the secret name the command must reject rejected string // the secret name the command must reject
run func(c *cli.Instance) error run func(c *cli.Instance) error
}{ }{
{"rm ..", "..", func(c *cli.Instance) error { {"rm --force ..", "..", func(c *cli.Instance) error {
return c.RemoveSecret(cmd, "..", false) return c.RemoveSecret(cmd, "..", true)
}}, }},
{"rm .", ".", func(c *cli.Instance) error { {"rm --force .", ".", func(c *cli.Instance) error {
return c.RemoveSecret(cmd, ".", false) return c.RemoveSecret(cmd, ".", true)
}}, }},
{`rm ""`, "", func(c *cli.Instance) error { {`rm --force ""`, "", func(c *cli.Instance) error {
return c.RemoveSecret(cmd, "", false) return c.RemoveSecret(cmd, "", true)
}}, }},
{"rm ../../etc", "../../etc", func(c *cli.Instance) error { {"rm --force ../../etc", "../../etc", func(c *cli.Instance) error {
return c.RemoveSecret(cmd, "../../etc", false) return c.RemoveSecret(cmd, "../../etc", true)
}}, }},
{"mv --force .. x", "..", func(c *cli.Instance) error { {"mv --force .. x", "..", func(c *cli.Instance) error {
return c.MoveSecret(cmd, "..", "x", true) return c.MoveSecret(cmd, "..", "x", true)
@@ -244,8 +244,8 @@ func TestInvalidSecretNameLeavesVaultsUnchanged(t *testing.T) {
{"version promote ..", "..", func(c *cli.Instance) error { {"version promote ..", "..", func(c *cli.Instance) error {
return c.PromoteVersion(cmd, "..", testVersion) return c.PromoteVersion(cmd, "..", testVersion)
}}, }},
{"version rm ..", "..", func(c *cli.Instance) error { {"version rm --force ..", "..", func(c *cli.Instance) error {
return c.RemoveVersion(cmd, "..", testVersion) return c.RemoveVersion(cmd, "..", testVersion, true)
}}, }},
{"encrypt ..", "..", func(c *cli.Instance) error { {"encrypt ..", "..", func(c *cli.Instance) error {
return c.Encrypt("..", "", "") return c.Encrypt("..", "", "")
@@ -279,8 +279,8 @@ func TestInvalidVersionLeavesVaultsUnchanged(t *testing.T) {
command string command string
run func(c *cli.Instance, version string) error run func(c *cli.Instance, version string) error
}{ }{
{"version rm x", func(c *cli.Instance, version string) error { {"version rm --force x", func(c *cli.Instance, version string) error {
return c.RemoveVersion(cmd, "x", version) return c.RemoveVersion(cmd, "x", version, true)
}}, }},
{"version promote x", func(c *cli.Instance, version string) error { {"version promote x", func(c *cli.Instance, version string) error {
return c.PromoteVersion(cmd, "x", version) return c.PromoteVersion(cmd, "x", version)
@@ -361,9 +361,9 @@ func TestInvalidVaultNameLeavesStateUnchanged(t *testing.T) {
} }
} }
// TestRemoveVersionRemovesOnlyThatVersion checks that `secret version rm` // TestRemoveVersionRemovesOnlyThatVersion checks that
// with a version that is not the current one removes that version and // `secret version rm --force` with a version that is not the current one
// changes nothing else. // removes that version and changes nothing else.
func TestRemoveVersionRemovesOnlyThatVersion(t *testing.T) { func TestRemoveVersionRemovesOnlyThatVersion(t *testing.T) {
t.Parallel() t.Parallel()
@@ -389,7 +389,7 @@ func TestRemoveVersionRemovesOnlyThatVersion(t *testing.T) {
require.Contains(t, before, oldDir) require.Contains(t, before, oldDir)
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir) c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
err = c.RemoveVersion(&cobra.Command{}, "x", versions[1]) err = c.RemoveVersion(&cobra.Command{}, "x", versions[1], true)
require.NoError(t, err) require.NoError(t, err)
// Expected: the state as before without everything under oldDir. // Expected: the state as before without everything under oldDir.
+68 -29
View File
@@ -205,19 +205,25 @@ func newRemoveCmd() *cobra.Command {
Aliases: []string{"rm"}, Aliases: []string{"rm"},
Short: "Remove a secret from the vault", Short: "Remove a secret from the vault",
Long: `Remove a secret and all its versions from the current ` + Long: `Remove a secret and all its versions from the current ` +
`vault. This action is permanent and cannot be undone.`, `vault. This action is permanent and cannot be undone. ` +
`Asks for confirmation first; when stdin is not a terminal, ` +
`fails unless --force is given.`,
Args: cobra.ExactArgs(1), Args: cobra.ExactArgs(1),
ValidArgsFunction: getSecretNamesCompletionFunc(cli.fs, cli.stateDir), ValidArgsFunction: getSecretNamesCompletionFunc(cli.fs, cli.stateDir),
RunE: func(cmd *cobra.Command, args []string) error { RunE: func(cmd *cobra.Command, args []string) error {
force, _ := cmd.Flags().GetBool("force")
cli, err := NewCLIInstance() cli, err := NewCLIInstance()
if err != nil { if err != nil {
return fmt.Errorf("failed to initialize CLI: %w", err) return fmt.Errorf("failed to initialize CLI: %w", err)
} }
return cli.RemoveSecret(cmd, args[0], false) return cli.RemoveSecret(cmd, args[0], force)
}, },
} }
cmd.Flags().BoolP("force", "f", false, "Remove without asking for confirmation")
return cmd return cmd
} }
@@ -699,29 +705,64 @@ func (cli *Instance) ImportSecret(
return nil return nil
} }
// RemoveSecret removes a secret from the vault // RemoveSecret removes a secret and all its versions from the current
func (cli *Instance) RemoveSecret(cmd *cobra.Command, secretName string, _ bool) error { // vault, after asking the user to confirm unless force is set.
func (cli *Instance) RemoveSecret(
cmd *cobra.Command, secretName string, force bool,
) error {
err := vault.ValidateSecretName(secretName) err := vault.ValidateSecretName(secretName)
if err != nil { if err != nil {
return err return err
} }
release, err := vault.LockStateDir(cli.fs, cli.stateDir) var found secretToRemove
release, err := cli.askThenLock(cmd, force, func() (string, error) {
var err error
found, err = cli.findSecretToRemove(secretName)
return found.question, err
})
if err != nil { if err != nil {
return err return err
} }
defer release() defer release()
// Get current vault err = secret.RemoveDirAtomic(cli.fs, found.dir)
currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
if err != nil { if err != nil {
return err return fmt.Errorf("failed to remove secret: %w", err)
}
cmd.Printf("Removed secret '%s' (%d version(s) deleted)\n",
secretName, found.versions)
return nil
}
// secretToRemove is what removing a secret removes, as findSecretToRemove
// found it.
type secretToRemove struct {
// dir is the secret's directory, which holds all its versions.
dir string
versions int
// question names what is removed, for the user to confirm.
question string
}
// findSecretToRemove checks that the secret exists in the current vault
// and counts its versions.
func (cli *Instance) findSecretToRemove(
secretName string,
) (secretToRemove, error) {
currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
if err != nil {
return secretToRemove{}, err
} }
// Check if secret exists
vaultDir, err := currentVlt.GetDirectory() vaultDir, err := currentVlt.GetDirectory()
if err != nil { if err != nil {
return err return secretToRemove{}, err
} }
encodedName := strings.ReplaceAll(secretName, "/", "%") encodedName := strings.ReplaceAll(secretName, "/", "%")
@@ -729,32 +770,30 @@ func (cli *Instance) RemoveSecret(cmd *cobra.Command, secretName string, _ bool)
exists, err := afero.DirExists(cli.fs, secretDir) exists, err := afero.DirExists(cli.fs, secretDir)
if err != nil { if err != nil {
return fmt.Errorf("failed to check if secret exists: %w", err) return secretToRemove{},
fmt.Errorf("failed to check if secret exists: %w", err)
} }
if !exists { if !exists {
return fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound) return secretToRemove{},
fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound)
} }
// Count versions for information // A secret without a versions directory has no versions, and can
versionsDir := filepath.Join(secretDir, "versions") // still be removed.
versionCount := 0 versions, err := afero.ReadDir(cli.fs, filepath.Join(secretDir, "versions"))
if err != nil && !errors.Is(err, os.ErrNotExist) {
entries, err := afero.ReadDir(cli.fs, versionsDir) return secretToRemove{}, fmt.Errorf(
if err == nil { "failed to count the versions of secret '%s': %w", secretName, err)
versionCount = len(entries)
} }
// Remove the secret directory return secretToRemove{
err = secret.RemoveDirAtomic(cli.fs, secretDir) dir: secretDir,
if err != nil { versions: len(versions),
return fmt.Errorf("failed to remove secret: %w", err) question: fmt.Sprintf("Permanently remove secret '%s' and its %d "+
} "version(s) from vault '%s'?",
secretName, len(versions), currentVlt.GetName()),
cmd.Printf("Removed secret '%s' (%d version(s) deleted)\n", }, nil
secretName, versionCount)
return nil
} }
// MoveSecret moves or renames a secret (within or across vaults), holding // MoveSecret moves or renames a secret (within or across vaults), holding
+77 -39
View File
@@ -47,7 +47,6 @@ var (
errGPGKeyAlreadyUnlocker = errors.New( errGPGKeyAlreadyUnlocker = errors.New(
"is already added as an unlocker") "is already added as an unlocker")
errUnsupportedUnlockerType = errors.New("unsupported unlocker type") errUnsupportedUnlockerType = errors.New("unsupported unlocker type")
errLastUnlocker = errors.New("refusing to remove last unlocker")
) )
// UnlockerInfo represents unlocker information for display // UnlockerInfo represents unlocker information for display
@@ -267,10 +266,11 @@ func newUnlockerRemoveCmd() *cobra.Command {
Use: "remove <unlocker-id>", Use: "remove <unlocker-id>",
Aliases: []string{"rm"}, Aliases: []string{"rm"},
Short: "Remove an unlocker", Short: "Remove an unlocker",
Long: `Remove an unlocker from the current vault. Cannot remove ` + Long: `Remove an unlocker from the current vault. Asks for ` +
`the last unlocker if the vault has secrets unless --force is ` + `confirmation first, saying whether it is the vault's last ` +
`used. Warning: Without unlockers and without your mnemonic, ` + `unlocker; when stdin is not a terminal, fails unless --force ` +
`vault data will be permanently inaccessible.`, `is given. Warning: Without unlockers and without your ` +
`mnemonic, vault data will be permanently inaccessible.`,
Args: cobra.ExactArgs(1), Args: cobra.ExactArgs(1),
ValidArgsFunction: getUnlockerIDsCompletionFunc(cli.fs, cli.stateDir), ValidArgsFunction: getUnlockerIDsCompletionFunc(cli.fs, cli.stateDir),
RunE: func(cmd *cobra.Command, args []string) error { RunE: func(cmd *cobra.Command, args []string) error {
@@ -286,7 +286,7 @@ func newUnlockerRemoveCmd() *cobra.Command {
} }
cmd.Flags().BoolP("force", "f", false, cmd.Flags().BoolP("force", "f", false,
"Force removal of last unlocker even if vault has secrets") "Remove without asking for confirmation, even the last unlocker")
return cmd return cmd
} }
@@ -726,55 +726,91 @@ func (cli *Instance) addPGPUnlocker(cmd *cobra.Command) error {
return nil return nil
} }
// UnlockersRemove removes an unlocker, holding the state directory lock // UnlockersRemove removes an unlocker from the current vault, after asking
// while removeUnlocker runs // the user to confirm unless force is set.
func (cli *Instance) UnlockersRemove( func (cli *Instance) UnlockersRemove(
unlockerID string, force bool, cmd *cobra.Command, unlockerID string, force bool, cmd *cobra.Command,
) error { ) error {
release, err := vault.LockStateDir(cli.fs, cli.stateDir) var found unlockerToRemove
release, err := cli.askThenLock(cmd, force, func() (string, error) {
var err error
found, err = cli.findUnlockerToRemove(unlockerID)
return found.question, err
})
if err != nil { if err != nil {
return err return err
} }
defer release() defer release()
return cli.removeUnlocker(unlockerID, force, cmd) return cli.removeUnlocker(unlockerID, found, cmd)
} }
// removeUnlocker removes an unlocker with safety checks // unlockerToRemove is what removing an unlocker removes, as
func (cli *Instance) removeUnlocker( // findUnlockerToRemove found it.
unlockerID string, force bool, cmd *cobra.Command, type unlockerToRemove struct {
) error { vlt *vault.Vault
// Get current vault // last is set when the unlocker counts as the vault's last one, and
// secrets is then the number of secrets in the vault.
last bool
secrets int
// question names what is removed, for the user to confirm.
question string
}
// findUnlockerToRemove checks that the current vault has the unlocker and
// finds whether it is the vault's last one.
func (cli *Instance) findUnlockerToRemove(
unlockerID string,
) (unlockerToRemove, error) {
vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
if err != nil { if err != nil {
return err return unlockerToRemove{}, err
}
exists, err := vlt.HasUnlocker(unlockerID)
if err != nil {
return unlockerToRemove{}, err
}
if !exists {
return unlockerToRemove{}, fmt.Errorf("unlocker with ID %s %w",
unlockerID, vault.ErrUnlockerNotFound)
} }
// Get list of unlockers. It leaves out a directory whose metadata file // Get list of unlockers. It leaves out a directory whose metadata file
// is missing or cannot be checked for, read or parsed. // is missing or cannot be checked for, read or parsed.
unlockers, err := vlt.ListUnlockers() unlockers, err := vlt.ListUnlockers()
if err != nil { if err != nil {
return fmt.Errorf("failed to list unlockers: %w", err) return unlockerToRemove{},
fmt.Errorf("failed to list unlockers: %w", err)
} }
vaultDir, err := vlt.GetDirectory() vaultDir, err := vlt.GetDirectory()
if err != nil { if err != nil {
return fmt.Errorf("failed to get vault directory: %w", err) return unlockerToRemove{},
fmt.Errorf("failed to get vault directory: %w", err)
} }
unlockersDir := filepath.Join(vaultDir, "unlockers.d") unlockersDir := filepath.Join(vaultDir, "unlockers.d")
// Check if we're removing the last unlocker found := unlockerToRemove{
removingLast := false vlt: vlt,
question: fmt.Sprintf("Permanently remove unlocker '%s' from vault "+
"'%s'? It is not the vault's last unlocker.",
unlockerID, vlt.GetName()),
}
if len(unlockers) == 1 { if len(unlockers) == 1 {
lastID, err := findUnlockerIDByMetadata( lastID, err := findUnlockerIDByMetadata(
cli.fs, unlockersDir, unlockers[0], true) cli.fs, unlockersDir, unlockers[0], true)
if err != nil { if err != nil {
return err return unlockerToRemove{}, err
} }
removingLast = lastID == unlockerID found.last = lastID == unlockerID
} }
// unlockerID may instead name a directory left out of the list. If its // unlockerID may instead name a directory left out of the list. If its
@@ -783,34 +819,36 @@ func (cli *Instance) removeUnlocker(
// for or read, the unlocker may be the only working one, so removing it // for or read, the unlocker may be the only working one, so removing it
// counts as removing the last unlocker. // counts as removing the last unlocker.
if metadataUnreadable(cli.fs, filepath.Join(unlockersDir, unlockerID)) { if metadataUnreadable(cli.fs, filepath.Join(unlockersDir, unlockerID)) {
removingLast = true found.last = true
} }
if removingLast { if found.last {
// Check if vault has secrets found.secrets, err = vlt.NumSecrets()
numSecrets, err := vlt.NumSecrets()
if err != nil { if err != nil {
return fmt.Errorf("failed to count secrets: %w", err) return unlockerToRemove{},
fmt.Errorf("failed to count secrets: %w", err)
} }
if numSecrets > 0 && !force { found.question = fmt.Sprintf("Permanently remove unlocker '%s', "+
cmd.Println("ERROR: Cannot remove the last unlocker when the " + "the last unlocker of vault '%s', which holds %d secret(s)? "+
"vault contains secrets.") "Without an unlocker the vault opens only with its mnemonic.",
cmd.Println("WARNING: Without unlockers, you MUST have your " + unlockerID, vlt.GetName(), found.secrets)
"mnemonic phrase to decrypt the vault.")
cmd.Println("If you want to proceed anyway, use --force")
return errLastUnlocker
} }
if numSecrets > 0 && force { return found, nil
}
// removeUnlocker removes the unlocker that findUnlockerToRemove found. The
// caller holds the state directory lock.
func (cli *Instance) removeUnlocker(
unlockerID string, found unlockerToRemove, cmd *cobra.Command,
) error {
if found.last && found.secrets > 0 {
cmd.Println("WARNING: Removing the last unlocker. You MUST " + cmd.Println("WARNING: Removing the last unlocker. You MUST " +
"have your mnemonic phrase to access this vault again!") "have your mnemonic phrase to access this vault again!")
} }
}
// Remove the unlocker err := found.vlt.RemoveUnlocker(unlockerID)
err = vlt.RemoveUnlocker(unlockerID)
if err != nil { if err != nil {
return err return err
} }
+27 -32
View File
@@ -5,14 +5,15 @@
// one the commands act on, metadata that is not JSON, and check that 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 // commands step past it, and that it can itself be removed by its
// directory name, which `secret unlocker list` names in its warning. A // directory name, which `secret unlocker list` names in its warning. A
// last test checks that an unlocker whose metadata file cannot be read is // last test checks that an unlocker whose metadata file cannot be read
// removed by its directory name only as the last unlocker is. // counts as the last unlocker when it is removed by its directory name.
//nolint:testpackage // white-box test of unexported internals //nolint:testpackage // white-box test of unexported internals
package cli package cli
import ( import (
"path/filepath" "path/filepath"
"strings"
"testing" "testing"
"git.eeqj.de/sneak/secret/internal/vault" "git.eeqj.de/sneak/secret/internal/vault"
@@ -56,37 +57,27 @@ func TestUnlockerSelectSkipsCorruptUnlocker(t *testing.T) {
} }
// TestUnlockerRemoveWithCorruptUnlocker asserts that the second unlocker // TestUnlockerRemoveWithCorruptUnlocker asserts that the second unlocker
// can be removed, unless the vault holds secrets: the corrupt unlocker // counts as the vault's last one, since the corrupt unlocker cannot unlock
// cannot unlock the vault, so the second is its last. The corrupt one can // the vault, and that the corrupt one, removed by its directory name, does
// be removed by its directory name without --force even then. // not. Either is removed once the user confirms.
func TestUnlockerRemoveWithCorruptUnlocker(t *testing.T) { func TestUnlockerRemoveWithCorruptUnlocker(t *testing.T) {
t.Parallel() t.Parallel()
tests := []struct { tests := []struct {
name string name string
unlockerID string unlockerID string
withSecret bool wantLast bool
wantErr error
wantEntries []string wantEntries []string
}{ }{
{ {
name: "the other unlocker", name: "the other unlocker",
unlockerID: "pgp-" + listTestGPGKeyID + "B", unlockerID: "pgp-" + listTestGPGKeyID + "B",
wantLast: true,
wantEntries: []string{listTestUnlockerDirOne}, 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", name: "the corrupt unlocker by its directory name",
unlockerID: listTestUnlockerDirOne, unlockerID: listTestUnlockerDirOne,
withSecret: true,
wantEntries: []string{listTestUnlockerDirTwo}, wantEntries: []string{listTestUnlockerDirTwo},
}, },
} }
@@ -96,14 +87,16 @@ func TestUnlockerRemoveWithCorruptUnlocker(t *testing.T) {
t.Parallel() t.Parallel()
fs := newCorruptUnlockerVault(t) fs := newCorruptUnlockerVault(t)
if tt.withSecret {
writeTestSecret(t, fs, testVaultDir(listTestVaultName)) writeTestSecret(t, fs, testVaultDir(listTestVaultName))
}
instance, cmd := newTestInstance(fs) instance, cmd := newTestInstance(fs)
err := instance.UnlockersRemove(tt.unlockerID, false, cmd) found, err := instance.findUnlockerToRemove(tt.unlockerID)
require.ErrorIs(t, err, tt.wantErr) require.NoError(t, err)
assert.Equal(t, tt.wantLast, found.last)
instance.terminal = strings.NewReader("y\n")
require.NoError(t, instance.UnlockersRemove(tt.unlockerID, false, cmd))
assertDirEntries(t, fs, assertDirEntries(t, fs,
filepath.Join(testVaultDir(listTestVaultName), filepath.Join(testVaultDir(listTestVaultName),
@@ -113,13 +106,14 @@ func TestUnlockerRemoveWithCorruptUnlocker(t *testing.T) {
} }
} }
// TestUnlockerRemoveWithUnreadableMetadata asserts that removing the only // TestUnlockerRemoveWithUnreadableMetadata asserts that the only unlocker
// unlocker of a vault with secrets by its directory name, when its // of a vault with secrets, removed by its directory name when its metadata
// metadata file cannot be checked for or read, is refused without --force: // file cannot be checked for or read, counts as the vault's last unlocker,
// listing leaves it out, but it may still be the vault's only working // so the question warns that it is: listing leaves it out, but it may
// unlocker. With --force it is removed. The state directory lock refuses // still be the vault's only working unlocker. It is then removed. The
// the failing filesystem, so the test calls removeUnlocker, which // state directory lock refuses the failing filesystem, so the test calls
// UnlockersRemove runs once it holds the lock. // findUnlockerToRemove and removeUnlocker, which UnlockersRemove runs to
// make its checks and, once it holds the lock, to remove the unlocker.
func TestUnlockerRemoveWithUnreadableMetadata(t *testing.T) { func TestUnlockerRemoveWithUnreadableMetadata(t *testing.T) {
t.Parallel() t.Parallel()
@@ -155,12 +149,13 @@ func TestUnlockerRemoveWithUnreadableMetadata(t *testing.T) {
instance, cmd := newTestInstance(tt.wrap(base)) instance, cmd := newTestInstance(tt.wrap(base))
err := instance.removeUnlocker(listTestUnlockerDirOne, false, cmd) found, err := instance.findUnlockerToRemove(listTestUnlockerDirOne)
require.ErrorIs(t, err, errLastUnlocker) require.NoError(t, err)
assertDirEntries(t, base, unlockersDir, listTestUnlockerDirOne) assert.True(t, found.last)
assert.Contains(t, found.question, "the last unlocker")
require.NoError(t, require.NoError(t,
instance.removeUnlocker(listTestUnlockerDirOne, true, cmd)) instance.removeUnlocker(listTestUnlockerDirOne, found, cmd))
assertDirEntries(t, base, unlockersDir) assertDirEntries(t, base, unlockersDir)
}) })
} }
+36 -9
View File
@@ -2,15 +2,18 @@
// //
// The checks that guard adding a PGP unlocker (is this key already an // The checks that guard adding a PGP unlocker (is this key already an
// unlocker?), removing the last unlocker and removing a vault (does the // unlocker?), removing the last unlocker and removing a vault (does the
// vault hold secrets?), and importing a mnemonic (does the vault already // vault hold secrets?), removing a secret (how many versions does it
// have a long-term key?) each look at the vault on disk before acting. // have?), 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 // When that look fails they must refuse to act, not read the failure as
// "nothing there" and go ahead. // "nothing there" and go ahead.
// //
// The tests make the look fail with a wrapper around the in-memory // The tests make the look fail with a wrapper around the in-memory
// filesystem, which the state directory lock refuses. So they call the // filesystem, which the state directory lock refuses. So they call the
// function each command runs once it holds the lock, such as removeVault // function each command runs once it holds the lock, such as addPGPUnlocker
// for RemoveVault. // for UnlockersAdd, or, for a removal, the function that makes its checks,
// such as findVaultToRemove for RemoveVault, which runs again under the
// lock before anything is removed, with --force or without.
//nolint:testpackage // white-box test of unexported internals //nolint:testpackage // white-box test of unexported internals
package cli package cli
@@ -285,10 +288,9 @@ func TestRemoveLastUnlockerAbortsWhenSecretsUnreadable(t *testing.T) {
base := newListTestVault(t, 1) base := newListTestVault(t, 1)
writeTestSecret(t, base, vaultDir) writeTestSecret(t, base, vaultDir)
instance, cmd := newTestInstance(&statFailFs{Fs: base, path: path}) instance, _ := newTestInstance(&statFailFs{Fs: base, path: path})
err := instance.removeUnlocker( _, err := instance.findUnlockerToRemove("pgp-" + listTestGPGKeyID + "A")
"pgp-"+listTestGPGKeyID+"A", false, cmd)
require.ErrorIs(t, err, errStatFailed) require.ErrorIs(t, err, errStatFailed)
assertDirEntries(t, base, unlockersDir, listTestUnlockerDirOne) assertDirEntries(t, base, unlockersDir, listTestUnlockerDirOne)
@@ -332,9 +334,9 @@ func TestRemoveVaultAbortsWhenSecretsDirUnreadable(t *testing.T) {
base := newListTestVault(t, 1) base := newListTestVault(t, 1)
writeTestSecret(t, base, vaultDir) writeTestSecret(t, base, vaultDir)
instance, cmd := newTestInstance(tt.failFs(base)) instance, _ := newTestInstance(tt.failFs(base))
err := instance.removeVault(cmd, unreadableTestOtherVault, false) _, err := instance.findVaultToRemove(unreadableTestOtherVault)
require.ErrorIs(t, err, tt.wantErr) require.ErrorIs(t, err, tt.wantErr)
@@ -345,6 +347,31 @@ func TestRemoveVaultAbortsWhenSecretsDirUnreadable(t *testing.T) {
} }
} }
// TestRemoveSecretAbortsWhenVersionsUnreadable asserts that a secret is
// kept when its versions directory exists but cannot be listed, so that
// the question cannot say how many versions would be removed.
func TestRemoveSecretAbortsWhenVersionsUnreadable(t *testing.T) {
t.Parallel()
secretDir := filepath.Join(testVaultDir(listTestVaultName),
unreadableTestSecretsDirName, unreadableTestSecretName)
versionsDir := filepath.Join(secretDir, "versions")
base := newListTestVault(t, 1)
writeTestSecret(t, base, testVaultDir(listTestVaultName))
require.NoError(t, base.MkdirAll(versionsDir, listTestDirPerm))
instance, _ := newTestInstance(&openFailFs{Fs: base, path: versionsDir})
_, err := instance.findSecretToRemove(unreadableTestSecretName)
require.ErrorIs(t, err, errOpenFailed)
exists, err := afero.DirExists(base, secretDir)
require.NoError(t, err)
assert.True(t, exists, "the secret must not be removed")
}
// TestVaultImportAbortsWhenPubKeyUnreadable asserts that a mnemonic import // TestVaultImportAbortsWhenPubKeyUnreadable asserts that a mnemonic import
// stops when whether the vault already has a long-term key cannot be // stops when whether the vault already has a long-term key cannot be
// determined. // determined.
+88 -67
View File
@@ -31,8 +31,6 @@ var (
errPassphraseEnvNotSet = errors.New( errPassphraseEnvNotSet = errors.New(
"SB_UNLOCK_PASSPHRASE environment variable not set") "SB_UNLOCK_PASSPHRASE environment variable not set")
errCannotRemoveLastVault = errors.New("cannot remove the last vault") errCannotRemoveLastVault = errors.New("cannot remove the last vault")
errVaultContainsSecrets = errors.New(
"contains secrets; use --force to remove")
) )
func newVaultCmd() *cobra.Command { func newVaultCmd() *cobra.Command {
@@ -156,9 +154,12 @@ func newVaultRemoveCmd() *cobra.Command {
Use: "remove <name>", Use: "remove <name>",
Aliases: []string{"rm"}, Aliases: []string{"rm"},
Short: "Remove a vault", Short: "Remove a vault",
Long: `Remove a vault. Requires --force if the vault contains ` + Long: `Remove a vault and all its secrets. Asks for ` +
`secrets. Will automatically switch to another vault if ` + `confirmation first, naming how many secrets the vault ` +
`removing the currently selected one.`, `holds; when stdin is not a terminal, fails unless --force ` +
`is given. Will automatically switch to another vault if ` +
`removing the currently selected one. The last vault ` +
`cannot be removed.`,
Args: cobra.ExactArgs(1), Args: cobra.ExactArgs(1),
ValidArgsFunction: getVaultNamesCompletionFunc(cli.fs, cli.stateDir), ValidArgsFunction: getVaultNamesCompletionFunc(cli.fs, cli.stateDir),
RunE: func(cmd *cobra.Command, args []string) error { RunE: func(cmd *cobra.Command, args []string) error {
@@ -173,7 +174,8 @@ func newVaultRemoveCmd() *cobra.Command {
}, },
} }
cmd.Flags().BoolP("force", "f", false, "Force removal even if vault contains secrets") cmd.Flags().BoolP("force", "f", false,
"Remove without asking for confirmation, even a vault that contains secrets")
return cmd return cmd
} }
@@ -537,27 +539,27 @@ func (cli *Instance) importMnemonic(cmd *cobra.Command, vaultName string) error
return nil return nil
} }
// vaultHasSecrets reports whether the vault directory contains any secrets // countVaultSecrets returns the number of secrets in the vault directory
func (cli *Instance) vaultHasSecrets(vaultDir string) (bool, error) { func (cli *Instance) countVaultSecrets(vaultDir string) (int, error) {
secretsDir := filepath.Join(vaultDir, "secrets.d") secretsDir := filepath.Join(vaultDir, "secrets.d")
exists, err := afero.DirExists(cli.fs, secretsDir) exists, err := afero.DirExists(cli.fs, secretsDir)
if err != nil { if err != nil {
return false, fmt.Errorf("failed to check secrets directory %s: %w", return 0, fmt.Errorf("failed to check secrets directory %s: %w",
secretsDir, err) secretsDir, err)
} }
if !exists { if !exists {
return false, nil return 0, nil
} }
entries, err := afero.ReadDir(cli.fs, secretsDir) entries, err := afero.ReadDir(cli.fs, secretsDir)
if err != nil { if err != nil {
return false, fmt.Errorf("failed to read secrets directory %s: %w", return 0, fmt.Errorf("failed to read secrets directory %s: %w",
secretsDir, err) secretsDir, err)
} }
return len(entries) > 0, nil return len(entries), nil
} }
// switchAwayFromVault selects another vault as current before removal // switchAwayFromVault selects another vault as current before removal
@@ -586,88 +588,107 @@ func (cli *Instance) switchAwayFromVault(
return nil return nil
} }
// RemoveVault removes a vault, holding the state directory lock while // RemoveVault removes a vault and all its secrets, after asking the user
// removeVault runs // to confirm unless force is set.
func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) error { func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) error {
err := vault.ValidateVaultName(name) err := vault.ValidateVaultName(name)
if err != nil { if err != nil {
return err return err
} }
release, err := vault.LockStateDir(cli.fs, cli.stateDir) var found vaultToRemove
release, err := cli.askThenLock(cmd, force, func() (string, error) {
var err error
found, err = cli.findVaultToRemove(name)
return found.question, err
})
if err != nil { if err != nil {
return err return err
} }
defer release() 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 {
return fmt.Errorf("failed to list vaults: %w", err)
}
// Check if vault exists
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 errCannotRemoveLastVault
}
// Check if this is the current vault
currentVault, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
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
hasSecrets, err := cli.vaultHasSecrets(vaultDir)
if err != nil {
return err
}
// Require --force if vault has secrets
if hasSecrets && !force {
return fmt.Errorf("vault '%s' %w", name, errVaultContainsSecrets)
}
// If removing current vault, switch to another vault first // If removing current vault, switch to another vault first
if isCurrentVault { if found.isCurrent {
err = cli.switchAwayFromVault(cmd, vaults, name) err = cli.switchAwayFromVault(cmd, found.vaults, name)
if err != nil { if err != nil {
return err return err
} }
} }
// Remove the vault directory // Remove the vault directory
err = secret.RemoveDirAtomic(cli.fs, vaultDir) err = secret.RemoveDirAtomic(cli.fs, found.dir)
if err != nil { if err != nil {
return fmt.Errorf("failed to remove vault directory: %w", err) return fmt.Errorf("failed to remove vault directory: %w", err)
} }
cmd.Printf("Removed vault '%s'\n", name) cmd.Printf("Removed vault '%s'\n", name)
if hasSecrets { if found.secrets > 0 {
cmd.Printf("Warning: Vault contained secrets that have been " + cmd.Printf("Warning: Vault contained secrets that have been " +
"permanently deleted\n") "permanently deleted\n")
} }
return nil return nil
} }
// vaultToRemove is what removing a vault removes, as findVaultToRemove
// found it.
type vaultToRemove struct {
// dir is the vault's directory, which holds all its secrets.
dir string
secrets int
// vaults lists every vault, this one included, and isCurrent is set
// when this one is the current vault.
vaults []string
isCurrent bool
// question names what is removed, for the user to confirm.
question string
}
// findVaultToRemove checks that the vault exists and is not the last one,
// and counts its secrets.
func (cli *Instance) findVaultToRemove(name string) (vaultToRemove, error) {
vaults, err := vault.ListVaults(cli.fs, cli.stateDir)
if err != nil {
return vaultToRemove{}, fmt.Errorf("failed to list vaults: %w", err)
}
if !slices.Contains(vaults, name) {
return vaultToRemove{},
fmt.Errorf("vault '%s' %w", name, errVaultDoesNotExist)
}
if len(vaults) == 1 {
return vaultToRemove{}, errCannotRemoveLastVault
}
currentVault, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
if err != nil {
return vaultToRemove{},
fmt.Errorf("failed to get current vault: %w", err)
}
vaultDir, err := vault.NewVault(cli.fs, cli.stateDir, name).GetDirectory()
if err != nil {
return vaultToRemove{},
fmt.Errorf("failed to get vault directory: %w", err)
}
secrets, err := cli.countVaultSecrets(vaultDir)
if err != nil {
return vaultToRemove{}, err
}
return vaultToRemove{
dir: vaultDir,
secrets: secrets,
vaults: vaults,
isCurrent: currentVault.GetName() == name,
question: fmt.Sprintf(
"Permanently remove vault '%s' and its %d secret(s)?",
name, secrets),
}, nil
}
+62 -25
View File
@@ -89,7 +89,8 @@ func VersionCommands(cli *Instance) *cobra.Command {
Aliases: []string{"rm"}, Aliases: []string{"rm"},
Short: "Remove a specific version of a secret", Short: "Remove a specific version of a secret",
Long: "Remove a specific version of a secret. Cannot remove the " + Long: "Remove a specific version of a secret. Cannot remove the " +
"current version.", "current version. Asks for confirmation first; when stdin " +
"is not a terminal, fails unless --force is given.",
Args: cobra.ExactArgs(2), //nolint:mnd // secret-name and version args Args: cobra.ExactArgs(2), //nolint:mnd // secret-name and version args
ValidArgsFunction: func( ValidArgsFunction: func(
cmd *cobra.Command, args []string, toComplete string, cmd *cobra.Command, args []string, toComplete string,
@@ -102,10 +103,15 @@ func VersionCommands(cli *Instance) *cobra.Command {
return nil, cobra.ShellCompDirectiveNoFileComp return nil, cobra.ShellCompDirectiveNoFileComp
}, },
RunE: func(cmd *cobra.Command, args []string) error { RunE: func(cmd *cobra.Command, args []string) error {
return cli.RemoveVersion(cmd, args[0], args[1]) force, _ := cmd.Flags().GetBool("force")
return cli.RemoveVersion(cmd, args[0], args[1], force)
}, },
} }
removeCmd.Flags().BoolP("force", "f", false,
"Remove without asking for confirmation")
versionCmd.AddCommand(listCmd, promoteCmd, removeCmd) versionCmd.AddCommand(listCmd, promoteCmd, removeCmd)
return versionCmd return versionCmd
@@ -297,30 +303,62 @@ func (cli *Instance) PromoteVersion(
return nil return nil
} }
// RemoveVersion removes a specific version of a secret // RemoveVersion removes a specific version of a secret, after asking the
// user to confirm unless force is set.
func (cli *Instance) RemoveVersion( func (cli *Instance) RemoveVersion(
cmd *cobra.Command, secretName string, version string, cmd *cobra.Command, secretName string, version string, force bool,
) error { ) error {
err := vault.ValidateSecretName(secretName) err := vault.ValidateSecretName(secretName)
if err != nil { if err != nil {
return err return err
} }
release, err := vault.LockStateDir(cli.fs, cli.stateDir) var found versionToRemove
release, err := cli.askThenLock(cmd, force, func() (string, error) {
var err error
found, err = cli.findVersionToRemove(secretName, version)
return found.question, err
})
if err != nil { if err != nil {
return err return err
} }
defer release() defer release()
// Get current vault err = secret.RemoveDirAtomic(cli.fs, found.dir)
if err != nil {
return fmt.Errorf("failed to remove version: %w", err)
}
cmd.Printf("Removed version %s of secret '%s'\n", version, secretName)
return nil
}
// versionToRemove is what removing a version removes, as
// findVersionToRemove found it.
type versionToRemove struct {
// dir is the version's directory.
dir string
// question names what is removed, for the user to confirm.
question string
}
// findVersionToRemove checks that the version exists in the secret in the
// current vault and is not its current version.
func (cli *Instance) findVersionToRemove(
secretName, version string,
) (versionToRemove, error) {
vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir)
if err != nil { if err != nil {
return err return versionToRemove{}, err
} }
vaultDir, err := vlt.GetDirectory() vaultDir, err := vlt.GetDirectory()
if err != nil { if err != nil {
return err return versionToRemove{}, err
} }
// Get the encoded secret name // Get the encoded secret name
@@ -330,45 +368,44 @@ func (cli *Instance) RemoveVersion(
// Check if secret exists // Check if secret exists
exists, err := afero.DirExists(cli.fs, secretDir) exists, err := afero.DirExists(cli.fs, secretDir)
if err != nil { if err != nil {
return fmt.Errorf("failed to check if secret exists: %w", err) return versionToRemove{},
fmt.Errorf("failed to check if secret exists: %w", err)
} }
if !exists { if !exists {
return fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound) return versionToRemove{},
fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound)
} }
// Check if version exists // Check if version exists
exists, err = secret.VersionExists(cli.fs, secretDir, version) exists, err = secret.VersionExists(cli.fs, secretDir, version)
if err != nil { if err != nil {
return fmt.Errorf("failed to check if version exists: %w", err) return versionToRemove{},
fmt.Errorf("failed to check if version exists: %w", err)
} }
if !exists { if !exists {
return fmt.Errorf("version '%s' %w '%s'", return versionToRemove{}, fmt.Errorf("version '%s' %w '%s'",
version, errVersionNotFound, secretName) version, errVersionNotFound, secretName)
} }
// Get current version // Get current version
currentVersion, err := secret.GetCurrentVersion(cli.fs, secretDir) currentVersion, err := secret.GetCurrentVersion(cli.fs, secretDir)
if err != nil { if err != nil {
return fmt.Errorf("failed to get current version: %w", err) return versionToRemove{},
fmt.Errorf("failed to get current version: %w", err)
} }
// Don't allow removing the current version // Don't allow removing the current version
if version == currentVersion { if version == currentVersion {
return fmt.Errorf("cannot remove the current version '%s'; %w", return versionToRemove{}, fmt.Errorf(
"cannot remove the current version '%s'; %w",
version, errCannotRemoveCurrentVersion) version, errCannotRemoveCurrentVersion)
} }
// Remove the version directory return versionToRemove{
versionDir := filepath.Join(secretDir, "versions", version) dir: filepath.Join(secretDir, "versions", version),
question: fmt.Sprintf("Permanently remove version %s of secret "+
err = secret.RemoveDirAtomic(cli.fs, versionDir) "'%s' from vault '%s'?", version, secretName, vlt.GetName()),
if err != nil { }, nil
return fmt.Errorf("failed to remove version: %w", err)
}
cmd.Printf("Removed version %s of secret '%s'\n", version, secretName)
return nil
} }
+8 -4
View File
@@ -38,7 +38,8 @@ const (
) )
// CreateKey creates a new P-256 non-exportable key in the Secure Enclave via sc_auth. // 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). // Returns the uncompressed public key bytes (65 bytes) and the identity hash
// (for deletion).
func CreateKey(label string) (publicKey []byte, hash string, err error) { func CreateKey(label string) (publicKey []byte, hash string, err error) {
pubKeyBuf := make([]C.uint8_t, p256UncompressedKeySize) pubKeyBuf := make([]C.uint8_t, p256UncompressedKeySize)
pubKeyLen := C.int(p256UncompressedKeySize) pubKeyLen := C.int(p256UncompressedKeySize)
@@ -57,7 +58,8 @@ func CreateKey(label string) (publicKey []byte, hash string, err error) {
return nil, "", fmt.Errorf("secure enclave: %s", C.GoString(&errBuf[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 //nolint:nlreturn // CGo result extraction
pk := C.GoBytes(unsafe.Pointer(&pubKeyBuf[0]), pubKeyLen)
h := C.GoString(&hashBuf[0]) h := C.GoString(&hashBuf[0])
return pk, h, nil return pk, h, nil
@@ -83,7 +85,8 @@ func Encrypt(label string, plaintext []byte) ([]byte, error) {
return nil, fmt.Errorf("secure enclave: %s", C.GoString(&errBuf[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 //nolint:nlreturn // CGo result extraction
out := C.GoBytes(unsafe.Pointer(&ciphertextBuf[0]), ciphertextLen)
return out, nil return out, nil
} }
@@ -107,7 +110,8 @@ func Decrypt(label string, ciphertext []byte) ([]byte, error) {
return nil, fmt.Errorf("secure enclave: %s", C.GoString(&errBuf[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 //nolint:nlreturn // CGo result extraction
out := C.GoBytes(unsafe.Pointer(&plaintextBuf[0]), plaintextLen)
return out, nil return out, nil
} }
+6 -6
View File
@@ -1,28 +1,28 @@
//go:build !darwin //go:build !darwin || !cgo
// Package macse provides Go bindings for macOS Secure Enclave operations. // Package macse provides Go bindings for macOS Secure Enclave operations.
package macse package macse
import "errors" import "errors"
var errNotSupported = errors.New("secure enclave is only supported on macOS") var errNotSupported = errors.New("secure enclave needs a macOS build with cgo")
// CreateKey is not supported on non-darwin platforms. // CreateKey fails: the Secure Enclave needs a macOS build with cgo.
func CreateKey(_ string) ([]byte, string, error) { func CreateKey(_ string) ([]byte, string, error) {
return nil, "", errNotSupported return nil, "", errNotSupported
} }
// Encrypt is not supported on non-darwin platforms. // Encrypt fails: the Secure Enclave needs a macOS build with cgo.
func Encrypt(_ string, _ []byte) ([]byte, error) { func Encrypt(_ string, _ []byte) ([]byte, error) {
return nil, errNotSupported return nil, errNotSupported
} }
// Decrypt is not supported on non-darwin platforms. // Decrypt fails: the Secure Enclave needs a macOS build with cgo.
func Decrypt(_ string, _ []byte) ([]byte, error) { func Decrypt(_ string, _ []byte) ([]byte, error) {
return nil, errNotSupported return nil, errNotSupported
} }
// DeleteKey is not supported on non-darwin platforms. // DeleteKey fails: the Secure Enclave needs a macOS build with cgo.
func DeleteKey(_ string) error { func DeleteKey(_ string) error {
return errNotSupported return errNotSupported
} }
+5 -4
View File
@@ -1,5 +1,4 @@
//go:build darwin //go:build darwin && cgo
// +build darwin
package macse package macse
@@ -45,7 +44,8 @@ func TestCreateAndDeleteKey(t *testing.T) {
// Verify valid uncompressed P-256 public key // Verify valid uncompressed P-256 public key
if len(pubKey) != p256UncompressedKeySize { if len(pubKey) != p256UncompressedKeySize {
t.Fatalf("expected public key length %d, got %d", p256UncompressedKeySize, len(pubKey)) t.Fatalf("expected public key length %d, got %d",
p256UncompressedKeySize, len(pubKey))
} }
if pubKey[0] != 0x04 { if pubKey[0] != 0x04 {
@@ -83,7 +83,8 @@ func TestEncryptDecryptRoundTrip(t *testing.T) {
}() }()
// Test data simulating an age private key // Test data simulating an age private key
plaintext := []byte("AGE-SECRET-KEY-1QQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQ") plaintext := []byte("AGE-SECRET-KEY-1" +
"QQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQQ")
// Encrypt // Encrypt
ciphertext, err := Encrypt(testKeyLabel, plaintext) ciphertext, err := Encrypt(testKeyLabel, plaintext)
+36 -9
View File
@@ -1,5 +1,6 @@
//go:build darwin //go:build darwin
//nolint:testpackage // white-box test of unexported getLongTermPrivateKey
package secret package secret
import ( import (
@@ -28,21 +29,43 @@ func (v *realVault) GetDirectory() (string, error) {
return filepath.Join(v.stateDir, "vaults.d", v.name), nil return filepath.Join(v.stateDir, "vaults.d", v.name), nil
} }
func (v *realVault) GetName() string { return v.name } func (v *realVault) GetName() string { return v.name }
//nolint:ireturn // implements VaultInterface
func (v *realVault) GetFilesystem() afero.Fs { return v.fs } func (v *realVault) GetFilesystem() afero.Fs { return v.fs }
// Unused by getLongTermPrivateKey — these satisfy VaultInterface. // Unused by getLongTermPrivateKey — these satisfy VaultInterface.
func (v *realVault) AddSecret(string, *memguard.LockedBuffer, bool) error { panic("not used") } func (v *realVault) AddSecret(string, *memguard.LockedBuffer, bool) error {
func (v *realVault) GetCurrentUnlocker() (Unlocker, error) { panic("not used") } panic("not used")
func (v *realVault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) { panic("not used") } }
func (v *realVault) SetMnemonic(*memguard.LockedBuffer) { panic("not used") }
func (v *realVault) SetUnlockPassphrase(*memguard.LockedBuffer) { panic("not used") } //nolint:ireturn // implements VaultInterface
func (v *realVault) CreatePassphraseUnlocker(*memguard.LockedBuffer) (*PassphraseUnlocker, error) { func (v *realVault) GetCurrentUnlocker() (Unlocker, error) {
panic("not used")
}
func (v *realVault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
panic("not used")
}
func (v *realVault) SetMnemonic(*memguard.LockedBuffer) {
panic("not used")
}
func (v *realVault) SetUnlockPassphrase(*memguard.LockedBuffer) {
panic("not used")
}
func (v *realVault) CreatePassphraseUnlocker(
*memguard.LockedBuffer,
) (*PassphraseUnlocker, error) {
panic("not used") panic("not used")
} }
// createRealVault sets up a complete vault directory structure on an in-memory // createRealVault sets up a complete vault directory structure on an in-memory
// filesystem, identical to what vault.CreateVault produces. // filesystem, identical to what vault.CreateVault produces.
func createRealVault(t *testing.T, fs afero.Fs, stateDir, name string, derivationIndex uint32) *realVault { func createRealVault(
t *testing.T, fs afero.Fs, stateDir, name string, derivationIndex uint32,
) *realVault {
t.Helper() t.Helper()
vaultDir := filepath.Join(stateDir, "vaults.d", name) vaultDir := filepath.Join(stateDir, "vaults.d", name)
@@ -55,7 +78,8 @@ func createRealVault(t *testing.T, fs afero.Fs, stateDir, name string, derivatio
} }
metaBytes, err := json.Marshal(metadata) metaBytes, err := json.Marshal(metadata)
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, afero.WriteFile(fs, filepath.Join(vaultDir, "vault-metadata.json"), metaBytes, FilePerms)) require.NoError(t, afero.WriteFile(fs,
filepath.Join(vaultDir, "vault-metadata.json"), metaBytes, FilePerms))
return &realVault{name: name, stateDir: stateDir, fs: fs} return &realVault{name: name, stateDir: stateDir, fs: fs}
} }
@@ -63,7 +87,9 @@ func createRealVault(t *testing.T, fs afero.Fs, stateDir, name string, derivatio
func TestGetLongTermPrivateKeyUsesVaultDerivationIndex(t *testing.T) { func TestGetLongTermPrivateKeyUsesVaultDerivationIndex(t *testing.T) {
t.Parallel() t.Parallel()
const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about" //nolint:dupword // BIP39 test mnemonic repeats words by design
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. // Derive expected keys at two different indices to prove they differ.
key0, err := agehd.DeriveIdentity(testMnemonic, 0) key0, err := agehd.DeriveIdentity(testMnemonic, 0)
@@ -82,6 +108,7 @@ func TestGetLongTermPrivateKeyUsesVaultDerivationIndex(t *testing.T) {
result, err := getLongTermPrivateKey(fs, vault, mnemonic, nil) result, err := getLongTermPrivateKey(fs, vault, mnemonic, nil)
require.NoError(t, err) require.NoError(t, err)
defer result.Destroy() defer result.Destroy()
assert.Equal(t, key5.String(), string(result.Bytes()), assert.Equal(t, key5.String(), string(result.Bytes()),
+203 -203
View File
@@ -1,11 +1,11 @@
//go:build darwin //go:build darwin
// +build darwin
package secret package secret
import ( import (
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"log/slog" "log/slog"
"os" "os"
@@ -17,23 +17,40 @@ import (
"filippo.io/age" "filippo.io/age"
"git.eeqj.de/sneak/secret/pkg/agehd" "git.eeqj.de/sneak/secret/pkg/agehd"
"github.com/awnumar/memguard" "github.com/awnumar/memguard"
keychain "github.com/keybase/go-keychain"
"github.com/spf13/afero" "github.com/spf13/afero"
) )
const ( const (
agePrivKeyPassphraseLength = 64 agePrivKeyPassphraseLength = 64
// KEYCHAIN_APP_IDENTIFIER is the service name used for keychain items // KEYCHAIN_APP_IDENTIFIER is the service name used for keychain items
KEYCHAIN_APP_IDENTIFIER = "berlin.sneak.app.secret" //nolint:revive // ALL_CAPS is intentional for this constant //
//nolint:revive // ALL_CAPS is intentional for this constant
KEYCHAIN_APP_IDENTIFIER = "berlin.sneak.app.secret"
// keychainUnlockerType is the metadata type string for keychain unlockers.
keychainUnlockerType = "keychain"
// macOSFlag is the unlocker metadata flag of the macOS-only unlockers.
macOSFlag = "macos"
) )
// keychainItemNameRegex validates keychain item names // keychainItemNameRegex validates keychain item names
// Allows alphanumeric characters, dots, hyphens, and underscores only // Allows alphanumeric characters, dots, hyphens, and underscores only
var keychainItemNameRegex = regexp.MustCompile(`^[A-Za-z0-9._-]+$`) var keychainItemNameRegex = regexp.MustCompile(`^[A-Za-z0-9._-]+$`)
var (
errNotMacOS = errors.New(
"keychain unlockers are only supported on macOS")
errKeychainItemNameEmpty = errors.New("keychain item name cannot be empty")
errInvalidKeychainItemName = errors.New("invalid keychain item name format")
errUnsupportedCurrentUnlocker = errors.New(
"unsupported current unlocker type for keychain unlocker creation")
)
// KeychainUnlockerMetadata extends UnlockerMetadata with keychain-specific data // KeychainUnlockerMetadata extends UnlockerMetadata with keychain-specific data
type KeychainUnlockerMetadata struct { type KeychainUnlockerMetadata struct {
UnlockerMetadata UnlockerMetadata
// Keychain item name // Keychain item name
KeychainItemName string `json:"keychainItemName"` KeychainItemName string `json:"keychainItemName"`
} }
@@ -45,6 +62,17 @@ type KeychainUnlocker struct {
fs afero.Fs fs afero.Fs
} }
// NewKeychainUnlocker creates a new KeychainUnlocker instance
func NewKeychainUnlocker(
fs afero.Fs, directory string, metadata UnlockerMetadata,
) *KeychainUnlocker {
return &KeychainUnlocker{
Directory: directory,
Metadata: metadata,
fs: fs,
}
}
// GetIdentity implements Unlocker interface for Keychain-based unlockers // GetIdentity implements Unlocker interface for Keychain-based unlockers
func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) { func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) {
DebugWith("Getting keychain unlocker identity", DebugWith("Getting keychain unlocker identity",
@@ -52,50 +80,20 @@ func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) {
slog.String("unlocker_type", k.GetType()), slog.String("unlocker_type", k.GetType()),
) )
// Step 1: Get keychain item name keychainData, err := k.readKeychainData()
keychainItemName, err := k.GetKeychainItemName()
if err != nil { if err != nil {
Debug("Failed to get keychain item name", "error", err, "unlocker_id", k.GetID()) return nil, err
return nil, fmt.Errorf("failed to get keychain item name: %w", err)
}
// Step 2: Retrieve data from keychain
Debug("Retrieving data from macOS keychain", "keychain_item", keychainItemName)
keychainDataBytes, err := retrieveFromKeychain(keychainItemName)
if err != nil {
Debug("Failed to retrieve data from keychain", "error", err, "keychain_item", keychainItemName)
return nil, fmt.Errorf("failed to retrieve data from keychain: %w", err)
}
DebugWith("Retrieved data from keychain",
slog.String("unlocker_id", k.GetID()),
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
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() defer keychainData.AgePrivKeyPassphrase.Destroy()
Debug("Parsed keychain data successfully", "unlocker_id", k.GetID())
// Step 4: Read the encrypted age private key from filesystem // Step 4: Read the encrypted age private key from filesystem
agePrivKeyPath := filepath.Join(k.Directory, "priv.age") agePrivKeyPath := filepath.Join(k.Directory, "priv.age")
Debug("Reading encrypted age private key", "path", agePrivKeyPath) Debug("Reading encrypted age private key", "path", agePrivKeyPath)
encryptedAgePrivKeyData, err := afero.ReadFile(k.fs, agePrivKeyPath) encryptedAgePrivKeyData, err := afero.ReadFile(k.fs, agePrivKeyPath)
if err != nil { if err != nil {
Debug("Failed to read encrypted age private key", "error", err, "path", agePrivKeyPath) Debug("Failed to read encrypted age private key",
"error", err, "path", agePrivKeyPath)
return nil, fmt.Errorf("failed to read encrypted age private key: %w", err) return nil, fmt.Errorf("failed to read encrypted age private key: %w", err)
} }
@@ -106,12 +104,17 @@ func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) {
) )
// Step 5: Decrypt the age private key using the passphrase from keychain // Step 5: Decrypt the age private key using the passphrase from keychain
Debug("Decrypting age private key with keychain passphrase", "unlocker_id", k.GetID()) Debug("Decrypting age private key with keychain passphrase",
agePrivKeyBuffer, err := DecryptWithPassphrase(encryptedAgePrivKeyData, keychainData.AgePrivKeyPassphrase) "unlocker_id", k.GetID())
if err != nil {
Debug("Failed to decrypt age private key with keychain passphrase", "error", err, "unlocker_id", k.GetID())
return nil, fmt.Errorf("failed to decrypt age private key with keychain passphrase: %w", err) 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())
return nil, fmt.Errorf(
"failed to decrypt age private key with keychain passphrase: %w", err)
} }
defer agePrivKeyBuffer.Destroy() defer agePrivKeyBuffer.Destroy()
@@ -140,7 +143,7 @@ func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) {
// GetType implements Unlocker interface // GetType implements Unlocker interface
func (k *KeychainUnlocker) GetType() string { func (k *KeychainUnlocker) GetType() string {
return "keychain" return keychainUnlockerType
} }
// GetMetadata implements Unlocker interface // GetMetadata implements Unlocker interface
@@ -174,58 +177,105 @@ func (k *KeychainUnlocker) Remove() error {
// Step 1: Get keychain item name // Step 1: Get keychain item name
keychainItemName, err := k.GetKeychainItemName() keychainItemName, err := k.GetKeychainItemName()
if err != nil { if err != nil {
Debug("Failed to get keychain item name during removal", "error", err, "unlocker_id", k.GetID()) Debug("Failed to get keychain item name during removal",
"error", err, "unlocker_id", k.GetID())
return fmt.Errorf("failed to get keychain item name: %w", err) return fmt.Errorf("failed to get keychain item name: %w", err)
} }
// Step 2: Remove from keychain // Step 2: Remove from keychain
Debug("Removing keychain item", "keychain_item", keychainItemName) Debug("Removing keychain item", "keychain_item", keychainItemName)
if err := deleteFromKeychain(keychainItemName); err != nil {
Debug("Failed to remove keychain item", "error", err, "keychain_item", keychainItemName) err = deleteFromKeychain(keychainItemName)
if err != nil {
Debug("Failed to remove keychain item",
"error", err, "keychain_item", keychainItemName)
return fmt.Errorf("failed to remove keychain item: %w", err) return fmt.Errorf("failed to remove keychain item: %w", err)
} }
// Step 3: Remove directory // Step 3: Remove directory
Debug("Removing keychain unlocker directory", "directory", k.Directory) Debug("Removing keychain unlocker directory", "directory", k.Directory)
if err := RemoveDirAtomic(k.fs, k.Directory); err != nil {
Debug("Failed to remove keychain unlocker directory", "error", err, "directory", k.Directory) err = RemoveDirAtomic(k.fs, k.Directory)
if 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) return fmt.Errorf("failed to remove keychain unlocker directory: %w", err)
} }
Debug("Successfully removed keychain unlocker", "unlocker_id", k.GetID(), "keychain_item", keychainItemName) Debug("Successfully removed keychain unlocker",
"unlocker_id", k.GetID(), "keychain_item", keychainItemName)
return nil return nil
} }
// NewKeychainUnlocker creates a new KeychainUnlocker instance
func NewKeychainUnlocker(fs afero.Fs, directory string, metadata UnlockerMetadata) *KeychainUnlocker {
return &KeychainUnlocker{
Directory: directory,
Metadata: metadata,
fs: fs,
}
}
// GetKeychainItemName returns the keychain item name from metadata // GetKeychainItemName returns the keychain item name from metadata
func (k *KeychainUnlocker) GetKeychainItemName() (string, error) { func (k *KeychainUnlocker) GetKeychainItemName() (string, error) {
// Load the metadata // Load the metadata
metadataPath := filepath.Join(k.Directory, "unlocker-metadata.json") metadataPath := filepath.Join(k.Directory, "unlocker-metadata.json")
metadataData, err := afero.ReadFile(k.fs, metadataPath) metadataData, err := afero.ReadFile(k.fs, metadataPath)
if err != nil { if err != nil {
return "", fmt.Errorf("failed to read keychain metadata: %w", err) return "", fmt.Errorf("failed to read keychain metadata: %w", err)
} }
var keychainMetadata KeychainUnlockerMetadata var keychainMetadata KeychainUnlockerMetadata
if err := json.Unmarshal(metadataData, &keychainMetadata); err != nil {
err = json.Unmarshal(metadataData, &keychainMetadata)
if err != nil {
return "", fmt.Errorf("failed to parse keychain metadata: %w", err) return "", fmt.Errorf("failed to parse keychain metadata: %w", err)
} }
return keychainMetadata.KeychainItemName, nil return keychainMetadata.KeychainItemName, nil
} }
// readKeychainData reads and parses the data this unlocker keeps in the
// keychain (steps 1 to 3 of GetIdentity). The caller must destroy the
// returned AgePrivKeyPassphrase.
func (k *KeychainUnlocker) readKeychainData() (*KeychainData, error) {
// Step 1: Get keychain item name
keychainItemName, err := k.GetKeychainItemName()
if err != nil {
Debug("Failed to get keychain item name", "error", err, "unlocker_id", k.GetID())
return nil, fmt.Errorf("failed to get keychain item name: %w", err)
}
// Step 2: Retrieve data from keychain
Debug("Retrieving data from macOS keychain", "keychain_item", keychainItemName)
keychainDataBytes, err := retrieveFromKeychain(keychainItemName)
if err != nil {
Debug("Failed to retrieve data from keychain",
"error", err, "keychain_item", keychainItemName)
return nil, fmt.Errorf("failed to retrieve data from keychain: %w", err)
}
DebugWith("Retrieved data from keychain",
slog.String("unlocker_id", k.GetID()),
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
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)
}
Debug("Parsed keychain data successfully", "unlocker_id", k.GetID())
return keychainData, nil
}
// generateKeychainUnlockerName generates a unique name for the keychain unlocker // generateKeychainUnlockerName generates a unique name for the keychain unlocker
func generateKeychainUnlockerName(vaultName string) (string, error) { func generateKeychainUnlockerName(vaultName string) (string, error) {
hostname, err := os.Hostname() hostname, err := os.Hostname()
@@ -247,31 +297,7 @@ func getLongTermPrivateKey(
fs afero.Fs, vault VaultInterface, mnemonic, passphrase *memguard.LockedBuffer, fs afero.Fs, vault VaultInterface, mnemonic, passphrase *memguard.LockedBuffer,
) (*memguard.LockedBuffer, error) { ) (*memguard.LockedBuffer, error) {
if mnemonic != nil { if mnemonic != nil {
// Read vault metadata to get the correct derivation index return deriveLongTermPrivateKey(fs, vault, mnemonic)
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(mnemonic.String(), metadata.DerivationIndex)
if err != nil {
return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
}
// Return the private key in a secure buffer
return memguard.NewBufferFromBytes([]byte(ltIdentity.String())), nil
} }
// Get the vault to access current unlocker // Get the vault to access current unlocker
@@ -292,34 +318,43 @@ func getLongTermPrivateKey(
// Get encrypted long-term key from current unlocker, handling different types // Get encrypted long-term key from current unlocker, handling different types
var encryptedLtPrivKey []byte var encryptedLtPrivKey []byte
switch currentUnlocker := currentUnlocker.(type) { switch currentUnlocker := currentUnlocker.(type) {
case *PassphraseUnlocker: case *PassphraseUnlocker:
// Read the encrypted long-term private key from passphrase unlocker // Read the encrypted long-term private key from passphrase unlocker
encryptedLtPrivKey, err = afero.ReadFile(fs, filepath.Join(currentUnlocker.GetDirectory(), "longterm.age")) encryptedLtPrivKey, err = afero.ReadFile(fs,
filepath.Join(currentUnlocker.GetDirectory(), "longterm.age"))
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to read encrypted long-term key from current passphrase unlocker: %w", err) return nil, fmt.Errorf("failed to read encrypted long-term key "+
"from current passphrase unlocker: %w", err)
} }
case *PGPUnlocker: case *PGPUnlocker:
// Read the encrypted long-term private key from PGP unlocker // Read the encrypted long-term private key from PGP unlocker
encryptedLtPrivKey, err = afero.ReadFile(fs, filepath.Join(currentUnlocker.GetDirectory(), "longterm.age")) encryptedLtPrivKey, err = afero.ReadFile(fs,
filepath.Join(currentUnlocker.GetDirectory(), "longterm.age"))
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to read encrypted long-term key from current PGP unlocker: %w", err) return nil, fmt.Errorf("failed to read encrypted long-term key "+
"from current PGP unlocker: %w", err)
} }
case *KeychainUnlocker: case *KeychainUnlocker:
// Read the encrypted long-term private key from another keychain unlocker // Read the encrypted long-term private key from another keychain
encryptedLtPrivKey, err = afero.ReadFile(fs, filepath.Join(currentUnlocker.GetDirectory(), "longterm.age")) // unlocker
encryptedLtPrivKey, err = afero.ReadFile(fs,
filepath.Join(currentUnlocker.GetDirectory(), "longterm.age"))
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to read encrypted long-term key from current keychain unlocker: %w", err) return nil, fmt.Errorf("failed to read encrypted long-term key "+
"from current keychain unlocker: %w", err)
} }
default: default:
return nil, fmt.Errorf("unsupported current unlocker type for keychain unlocker creation") return nil, errUnsupportedCurrentUnlocker
} }
// Decrypt long-term private key using current unlocker // Decrypt long-term private key using current unlocker
ltPrivKeyBuffer, err := DecryptWithIdentity(encryptedLtPrivKey, currentUnlockerIdentity) ltPrivKeyBuffer, err := DecryptWithIdentity(
encryptedLtPrivKey, currentUnlockerIdentity)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err) return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err)
} }
@@ -328,6 +363,43 @@ func getLongTermPrivateKey(
return ltPrivKeyBuffer, nil return ltPrivKeyBuffer, nil
} }
// deriveLongTermPrivateKey derives the long-term private key from mnemonic at
// the vault's derivation index, for getLongTermPrivateKey and
// getLongTermKeyForSE.
func deriveLongTermPrivateKey(
fs afero.Fs, vault VaultInterface, mnemonic *memguard.LockedBuffer,
) (*memguard.LockedBuffer, error) {
// 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
err = json.Unmarshal(metadataBytes, &metadata)
if 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(mnemonic.String(), metadata.DerivationIndex)
if err != nil {
return nil, fmt.Errorf(
"failed to derive long-term key from mnemonic: %w", err)
}
// Return the private key in a secure buffer
return memguard.NewBufferFromBytes([]byte(ltIdentity.String())), nil
}
// CreateKeychainUnlocker creates a new keychain unlocker and stores it in the // CreateKeychainUnlocker creates a new keychain unlocker and stores it in the
// vault. The long-term key comes from mnemonic when it is not nil, else from // vault. The long-term key comes from mnemonic when it is not nil, else from
// the current unlocker, as getLongTermPrivateKey describes. // the current unlocker, as getLongTermPrivateKey describes.
@@ -335,7 +407,8 @@ func CreateKeychainUnlocker(
fs afero.Fs, stateDir string, mnemonic, passphrase *memguard.LockedBuffer, fs afero.Fs, stateDir string, mnemonic, passphrase *memguard.LockedBuffer,
) (*KeychainUnlocker, error) { ) (*KeychainUnlocker, error) {
// Check if we're on macOS // Check if we're on macOS
if err := checkMacOSAvailable(); err != nil { err := checkMacOSAvailable()
if err != nil {
return nil, err return nil, err
} }
@@ -377,10 +450,12 @@ func CreateKeychainUnlocker(
// Step 3: Encrypt age private key with the generated passphrase // Step 3: Encrypt age private key with the generated passphrase
// Create a secure buffer for the private key // Create a secure buffer for the private key
agePrivKeyStr := ageIdentity.String() agePrivKeyStr := ageIdentity.String()
agePrivKeyBuffer := memguard.NewBufferFromBytes([]byte(agePrivKeyStr)) agePrivKeyBuffer := memguard.NewBufferFromBytes([]byte(agePrivKeyStr))
defer agePrivKeyBuffer.Destroy() defer agePrivKeyBuffer.Destroy()
encryptedAgePrivKey, err := EncryptWithPassphrase(agePrivKeyBuffer, agePrivKeyPassphrase) encryptedAgePrivKey, err := EncryptWithPassphrase(
agePrivKeyBuffer, agePrivKeyPassphrase)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to encrypt age private key with passphrase: %w", err) return nil, fmt.Errorf("failed to encrypt age private key with passphrase: %w", err)
} }
@@ -393,9 +468,11 @@ func CreateKeychainUnlocker(
defer ltPrivKeyData.Destroy() defer ltPrivKeyData.Destroy()
// Step 5: 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()) encryptedLtPrivKeyToAge, err := EncryptToRecipient(
ltPrivKeyData, ageIdentity.Recipient())
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to encrypt long-term private key to age unlocker: %w", err) return nil, fmt.Errorf(
"failed to encrypt long-term private key to age unlocker: %w", err)
} }
// Step 6: Prepare keychain data // Step 6: Prepare keychain data
@@ -411,12 +488,23 @@ func CreateKeychainUnlocker(
} }
defer keychainDataBuffer.Destroy() defer keychainDataBuffer.Destroy()
return writeKeychainUnlocker(fs, unlockerDir, keychainItemName, ageRecipient,
encryptedAgePrivKey, encryptedLtPrivKeyToAge, keychainDataBuffer)
}
// writeKeychainUnlocker writes a new keychain unlocker into unlockerDir and
// stores its data in the keychain (steps 7 and 8 of CreateKeychainUnlocker).
func writeKeychainUnlocker(
fs afero.Fs, unlockerDir, keychainItemName, ageRecipient string,
encryptedAgePrivKey, encryptedLtPrivKey []byte,
keychainDataBuffer *memguard.LockedBuffer,
) (*KeychainUnlocker, error) {
// Step 7: Prepare enhanced metadata // Step 7: Prepare enhanced metadata
keychainMetadata := KeychainUnlockerMetadata{ keychainMetadata := KeychainUnlockerMetadata{
UnlockerMetadata: UnlockerMetadata{ UnlockerMetadata: UnlockerMetadata{
Type: "keychain", Type: keychainUnlockerType,
CreatedAt: time.Now(), CreatedAt: time.Now(),
Flags: []string{"keychain", "macos"}, Flags: []string{keychainUnlockerType, macOSFlag},
}, },
KeychainItemName: keychainItemName, KeychainItemName: keychainItemName,
} }
@@ -429,27 +517,29 @@ func CreateKeychainUnlocker(
// Step 8: Write the unlocker's files and store the data in the keychain, // Step 8: Write the unlocker's files and store the data in the keychain,
// the metadata last // the metadata last
err = WriteDir(fs, unlockerDir, func(dir string) error { err = WriteDir(fs, unlockerDir, func(dir string) error {
pubPath := filepath.Join(dir, "pub.txt") err := WriteFileAtomic(fs, filepath.Join(dir, "pub.txt"), []byte(ageRecipient))
if err := WriteFileAtomic(fs, pubPath, []byte(ageRecipient)); err != nil { if err != nil {
return fmt.Errorf("failed to write age recipient: %w", err) return fmt.Errorf("failed to write age recipient: %w", err)
} }
privPath := filepath.Join(dir, "priv.age") err = WriteFileAtomic(fs, filepath.Join(dir, "priv.age"), encryptedAgePrivKey)
if err := WriteFileAtomic(fs, privPath, encryptedAgePrivKey); err != nil { if err != nil {
return fmt.Errorf("failed to write encrypted age private key: %w", err) return fmt.Errorf("failed to write encrypted age private key: %w", err)
} }
ltKeyPath := filepath.Join(dir, "longterm.age") err = WriteFileAtomic(fs, filepath.Join(dir, "longterm.age"), encryptedLtPrivKey)
if err := WriteFileAtomic(fs, ltKeyPath, encryptedLtPrivKeyToAge); err != nil { if err != nil {
return fmt.Errorf("failed to write encrypted long-term private key: %w", err) return fmt.Errorf("failed to write encrypted long-term private key: %w", err)
} }
if err := storeInKeychain(keychainItemName, keychainDataBuffer); err != nil { err = storeInKeychain(keychainItemName, keychainDataBuffer)
if err != nil {
return fmt.Errorf("failed to store data in keychain: %w", err) return fmt.Errorf("failed to store data in keychain: %w", err)
} }
metadataPath := filepath.Join(dir, "unlocker-metadata.json") err = WriteFileAtomic(fs, filepath.Join(dir, "unlocker-metadata.json"),
if err := WriteFileAtomic(fs, metadataPath, metadataBytes); err != nil { metadataBytes)
if err != nil {
return fmt.Errorf("failed to write unlocker metadata: %w", err) return fmt.Errorf("failed to write unlocker metadata: %w", err)
} }
@@ -469,111 +559,21 @@ func CreateKeychainUnlocker(
// checkMacOSAvailable verifies that we're running on macOS // checkMacOSAvailable verifies that we're running on macOS
func checkMacOSAvailable() error { func checkMacOSAvailable() error {
if runtime.GOOS != "darwin" { if runtime.GOOS != "darwin" {
return fmt.Errorf("keychain unlockers are only supported on macOS, current OS: %s", runtime.GOOS) return fmt.Errorf("%w, current OS: %s", errNotMacOS, runtime.GOOS)
} }
return nil return nil
} }
// validateKeychainItemName validates that a keychain item name is safe for command execution // validateKeychainItemName validates that a keychain item name is safe for
// command execution
func validateKeychainItemName(itemName string) error { func validateKeychainItemName(itemName string) error {
if itemName == "" { if itemName == "" {
return fmt.Errorf("keychain item name cannot be empty") return errKeychainItemNameEmpty
} }
if !keychainItemNameRegex.MatchString(itemName) { if !keychainItemNameRegex.MatchString(itemName) {
return fmt.Errorf("invalid keychain item name format: %s", itemName) return fmt.Errorf("%w: %s", errInvalidKeychainItemName, itemName)
}
return nil
}
// storeInKeychain stores data in the macOS keychain using keybase/go-keychain
func storeInKeychain(itemName string, data *memguard.LockedBuffer) error {
if data == nil {
return fmt.Errorf("data buffer is nil")
}
if err := validateKeychainItemName(itemName); err != nil {
return fmt.Errorf("invalid keychain item name: %w", err)
}
item := keychain.NewItem()
item.SetSecClass(keychain.SecClassGenericPassword)
item.SetService(KEYCHAIN_APP_IDENTIFIER)
item.SetAccount(itemName)
item.SetLabel(fmt.Sprintf("%s - %s", KEYCHAIN_APP_IDENTIFIER, itemName))
item.SetDescription("Secret vault keychain data")
item.SetData(data.Bytes())
item.SetSynchronizable(keychain.SynchronizableNo)
// Use AccessibleWhenUnlockedThisDeviceOnly for better security and to trigger auth
item.SetAccessible(keychain.AccessibleWhenUnlockedThisDeviceOnly)
// First try to delete any existing item
deleteItem := keychain.NewItem()
deleteItem.SetSecClass(keychain.SecClassGenericPassword)
deleteItem.SetService(KEYCHAIN_APP_IDENTIFIER)
deleteItem.SetAccount(itemName)
_ = keychain.DeleteItem(deleteItem) // Ignore error as item might not exist
// Add the new item
if err := keychain.AddItem(item); err != nil {
return fmt.Errorf("failed to store item in keychain: %w", err)
}
return nil
}
// retrieveFromKeychain retrieves data from the macOS keychain using keybase/go-keychain
func retrieveFromKeychain(itemName string) ([]byte, error) {
if err := validateKeychainItemName(itemName); err != nil {
return nil, fmt.Errorf("invalid keychain item name: %w", err)
}
query := keychain.NewItem()
query.SetSecClass(keychain.SecClassGenericPassword)
query.SetService(KEYCHAIN_APP_IDENTIFIER)
query.SetAccount(itemName)
query.SetMatchLimit(keychain.MatchLimitOne)
query.SetReturnData(true)
results, err := keychain.QueryItem(query)
if err != nil {
return nil, fmt.Errorf("failed to retrieve item from keychain: %w", err)
}
if len(results) == 0 {
return nil, fmt.Errorf("keychain item not found: %s", itemName)
}
return results[0].Data, nil
}
// deleteFromKeychain removes an item from the macOS keychain using keybase/go-keychain
// If the item doesn't exist, this function returns nil (not an error) since the goal
// is to ensure the item is gone, and it already being gone satisfies that goal.
func deleteFromKeychain(itemName string) error {
if err := validateKeychainItemName(itemName); err != nil {
return fmt.Errorf("invalid keychain item name: %w", err)
}
item := keychain.NewItem()
item.SetSecClass(keychain.SecClassGenericPassword)
item.SetService(KEYCHAIN_APP_IDENTIFIER)
item.SetAccount(itemName)
if err := keychain.DeleteItem(item); err != nil {
// If the item doesn't exist, that's not an error - the goal is to ensure
// the item is gone, and it already being gone satisfies that goal.
// This is important for cleaning up unlocker directories when the keychain
// item has already been removed (e.g., manually by user, or synced vault
// from a different machine).
if err == keychain.ErrorItemNotFound {
Debug("Keychain item not found during deletion, ignoring", "item_name", itemName)
return nil
}
return fmt.Errorf("failed to delete item from keychain: %w", err)
} }
return nil return nil
+104
View File
@@ -0,0 +1,104 @@
//go:build darwin && cgo
package secret
import (
"fmt"
"github.com/awnumar/memguard"
keychain "github.com/keybase/go-keychain"
)
// The keychain unlocker's only calls into go-keychain, which is cgo on macOS.
// A macOS build without cgo gets keychainunlocker_nocgo.go instead.
// storeInKeychain stores data in the macOS keychain using keybase/go-keychain
func storeInKeychain(itemName string, data *memguard.LockedBuffer) error {
if data == nil {
return fmt.Errorf("data buffer is nil")
}
if err := validateKeychainItemName(itemName); err != nil {
return fmt.Errorf("invalid keychain item name: %w", err)
}
item := keychain.NewItem()
item.SetSecClass(keychain.SecClassGenericPassword)
item.SetService(KEYCHAIN_APP_IDENTIFIER)
item.SetAccount(itemName)
item.SetLabel(fmt.Sprintf("%s - %s", KEYCHAIN_APP_IDENTIFIER, itemName))
item.SetDescription("Secret vault keychain data")
item.SetData(data.Bytes())
item.SetSynchronizable(keychain.SynchronizableNo)
// Use AccessibleWhenUnlockedThisDeviceOnly for better security and to trigger auth
item.SetAccessible(keychain.AccessibleWhenUnlockedThisDeviceOnly)
// First try to delete any existing item
deleteItem := keychain.NewItem()
deleteItem.SetSecClass(keychain.SecClassGenericPassword)
deleteItem.SetService(KEYCHAIN_APP_IDENTIFIER)
deleteItem.SetAccount(itemName)
_ = keychain.DeleteItem(deleteItem) // Ignore error as item might not exist
// Add the new item
if err := keychain.AddItem(item); err != nil {
return fmt.Errorf("failed to store item in keychain: %w", err)
}
return nil
}
// retrieveFromKeychain retrieves data from the macOS keychain using keybase/go-keychain
func retrieveFromKeychain(itemName string) ([]byte, error) {
if err := validateKeychainItemName(itemName); err != nil {
return nil, fmt.Errorf("invalid keychain item name: %w", err)
}
query := keychain.NewItem()
query.SetSecClass(keychain.SecClassGenericPassword)
query.SetService(KEYCHAIN_APP_IDENTIFIER)
query.SetAccount(itemName)
query.SetMatchLimit(keychain.MatchLimitOne)
query.SetReturnData(true)
results, err := keychain.QueryItem(query)
if err != nil {
return nil, fmt.Errorf("failed to retrieve item from keychain: %w", err)
}
if len(results) == 0 {
return nil, fmt.Errorf("keychain item not found: %s", itemName)
}
return results[0].Data, nil
}
// deleteFromKeychain removes an item from the macOS keychain using keybase/go-keychain
// If the item doesn't exist, this function returns nil (not an error) since the goal
// is to ensure the item is gone, and it already being gone satisfies that goal.
func deleteFromKeychain(itemName string) error {
if err := validateKeychainItemName(itemName); err != nil {
return fmt.Errorf("invalid keychain item name: %w", err)
}
item := keychain.NewItem()
item.SetSecClass(keychain.SecClassGenericPassword)
item.SetService(KEYCHAIN_APP_IDENTIFIER)
item.SetAccount(itemName)
if err := keychain.DeleteItem(item); err != nil {
// If the item doesn't exist, that's not an error - the goal is to ensure
// the item is gone, and it already being gone satisfies that goal.
// This is important for cleaning up unlocker directories when the keychain
// item has already been removed (e.g., manually by user, or synced vault
// from a different machine).
if err == keychain.ErrorItemNotFound {
Debug("Keychain item not found during deletion, ignoring", "item_name", itemName)
return nil
}
return fmt.Errorf("failed to delete item from keychain: %w", err)
}
return nil
}
+30
View File
@@ -0,0 +1,30 @@
//go:build darwin && !cgo
package secret
import (
"errors"
"github.com/awnumar/memguard"
)
// In a macOS build without cgo, these take the place of the functions in
// keychainunlocker_cgo.go: go-keychain is cgo on macOS, so they can only fail.
var errKeychainNotSupported = errors.New(
"keychain unlockers need a macOS build with cgo")
// storeInKeychain fails: the keychain needs a macOS build with cgo.
func storeInKeychain(_ string, _ *memguard.LockedBuffer) error {
return errKeychainNotSupported
}
// retrieveFromKeychain fails: the keychain needs a macOS build with cgo.
func retrieveFromKeychain(_ string) ([]byte, error) {
return nil, errKeychainNotSupported
}
// deleteFromKeychain fails: the keychain needs a macOS build with cgo.
func deleteFromKeychain(_ string) error {
return errKeychainNotSupported
}
+9 -6
View File
@@ -1,5 +1,4 @@
//go:build darwin //go:build darwin && cgo
// +build darwin
package secret package secret
@@ -35,7 +34,8 @@ func TestKeychainStoreRetrieveDelete(t *testing.T) {
// Test 2: Retrieve data from keychain // Test 2: Retrieve data from keychain
retrievedData, err := retrieveFromKeychain(testItemName) retrievedData, err := retrieveFromKeychain(testItemName)
require.NoError(t, err, "Failed to retrieve data from keychain") require.NoError(t, err, "Failed to retrieve data from keychain")
assert.Equal(t, testData, string(retrievedData), "Retrieved data doesn't match stored data") assert.Equal(t, testData, string(retrievedData),
"Retrieved data doesn't match stored data")
// Test 3: Update existing item (store again with different data) // Test 3: Update existing item (store again with different data)
newTestData := "updated-test-data-67890" newTestData := "updated-test-data-67890"
@@ -48,7 +48,8 @@ func TestKeychainStoreRetrieveDelete(t *testing.T) {
// Verify updated data // Verify updated data
retrievedData, err = retrieveFromKeychain(testItemName) retrievedData, err = retrieveFromKeychain(testItemName)
require.NoError(t, err, "Failed to retrieve updated data from keychain") require.NoError(t, err, "Failed to retrieve updated data from keychain")
assert.Equal(t, newTestData, string(retrievedData), "Retrieved data doesn't match updated data") assert.Equal(t, newTestData, string(retrievedData),
"Retrieved data doesn't match updated data")
// Test 4: Delete from keychain // Test 4: Delete from keychain
err = deleteFromKeychain(testItemName) err = deleteFromKeychain(testItemName)
@@ -93,7 +94,8 @@ func TestKeychainInvalidItemName(t *testing.T) {
for _, name := range invalidNames { for _, name := range invalidNames {
err := storeInKeychain(name, testData) err := storeInKeychain(name, testData)
assert.Error(t, err, "Expected error for invalid name: %s", name) assert.Error(t, err, "Expected error for invalid name: %s", name)
assert.Contains(t, err.Error(), "invalid keychain item name", "Error should mention invalid name for: %s", name) assert.Contains(t, err.Error(), "invalid keychain item name",
"Error should mention invalid name for: %s", name)
} }
// Test valid names (should not error on validation) // Test valid names (should not error on validation)
@@ -180,5 +182,6 @@ func TestDeleteNonExistentKeychainItem(t *testing.T) {
// This is important for cleaning up unlocker directories when the keychain item // This is important for cleaning up unlocker directories when the keychain item
// has already been removed (e.g., manually by user, or on a different machine) // has already been removed (e.g., manually by user, or on a different machine)
err := deleteFromKeychain(testItemName) err := deleteFromKeychain(testItemName)
assert.NoError(t, err, "Deleting non-existent keychain item should not return an error") assert.NoError(t, err,
"Deleting non-existent keychain item should not return an error")
} }
+252 -122
View File
@@ -4,7 +4,9 @@ package secret_test
import ( import (
"bytes" "bytes"
"context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"os" "os"
@@ -22,23 +24,24 @@ import (
"github.com/spf13/afero" "github.com/spf13/afero"
) )
// Register vault with secret package for testing // pgpUnlockerType is the type of a PGP unlocker.
func init() { const pgpUnlockerType = "pgp"
// Register the vault.GetCurrentVault function with the secret package
secret.RegisterGetCurrentVaultFunc(func(fs afero.Fs, stateDir string) (secret.VaultInterface, error) { var errNilDataBuffer = errors.New("data buffer is nil")
return vault.GetCurrentVault(fs, stateDir)
})
}
// setupNonInteractiveGPG creates a custom GPG environment for testing // setupNonInteractiveGPG creates a custom GPG environment for testing
func setupNonInteractiveGPG(t *testing.T, _, passphrase, gnupgHomeDir string) { func setupNonInteractiveGPG(t *testing.T, _, passphrase, gnupgHomeDir string) {
t.Helper()
// Create GPG config file for non-interactive operation // Create GPG config file for non-interactive operation
gpgConfPath := filepath.Join(gnupgHomeDir, "gpg.conf") gpgConfPath := filepath.Join(gnupgHomeDir, "gpg.conf")
gpgConfContent := `batch gpgConfContent := `batch
no-tty no-tty
pinentry-mode loopback pinentry-mode loopback
` `
if err := os.WriteFile(gpgConfPath, []byte(gpgConfContent), 0o600); err != nil {
err := os.WriteFile(gpgConfPath, []byte(gpgConfContent), 0o600)
if err != nil {
t.Fatalf("Failed to write GPG config file: %v", err) t.Fatalf("Failed to write GPG config file: %v", err)
} }
@@ -47,11 +50,15 @@ pinentry-mode loopback
origDecryptFunc := secret.GPGDecryptFunc origDecryptFunc := secret.GPGDecryptFunc
// Set custom GPG functions for this test // Set custom GPG functions for this test
secret.GPGEncryptFunc = func(data *memguard.LockedBuffer, keyID string) ([]byte, error) { secret.GPGEncryptFunc = func(
data *memguard.LockedBuffer, keyID string,
) ([]byte, error) {
if data == nil { if data == nil {
return nil, fmt.Errorf("data buffer is nil") return nil, errNilDataBuffer
} }
cmd := exec.Command("gpg",
//nolint:gosec // G204: test runs gpg with test-controlled arguments
cmd := exec.CommandContext(t.Context(), "gpg",
"--homedir", gnupgHomeDir, "--homedir", gnupgHomeDir,
"--batch", "--batch",
"--yes", "--yes",
@@ -63,11 +70,13 @@ pinentry-mode loopback
"-r", keyID) "-r", keyID)
var stdout, stderr bytes.Buffer var stdout, stderr bytes.Buffer
cmd.Stdout = &stdout cmd.Stdout = &stdout
cmd.Stderr = &stderr cmd.Stderr = &stderr
cmd.Stdin = bytes.NewReader(data.Bytes()) cmd.Stdin = bytes.NewReader(data.Bytes())
if err := cmd.Run(); err != nil { err := cmd.Run()
if err != nil {
return nil, fmt.Errorf("GPG encryption failed: %w\nStderr: %s", err, stderr.String()) return nil, fmt.Errorf("GPG encryption failed: %w\nStderr: %s", err, stderr.String())
} }
@@ -75,7 +84,8 @@ pinentry-mode loopback
} }
secret.GPGDecryptFunc = func(encryptedData []byte) (*memguard.LockedBuffer, error) { secret.GPGDecryptFunc = func(encryptedData []byte) (*memguard.LockedBuffer, error) {
cmd := exec.Command("gpg", //nolint:gosec // G204: test runs gpg with test-controlled arguments
cmd := exec.CommandContext(t.Context(), "gpg",
"--homedir", gnupgHomeDir, "--homedir", gnupgHomeDir,
"--batch", "--batch",
"--yes", "--yes",
@@ -85,11 +95,13 @@ pinentry-mode loopback
"--decrypt") "--decrypt")
var stdout, stderr bytes.Buffer var stdout, stderr bytes.Buffer
cmd.Stdout = &stdout cmd.Stdout = &stdout
cmd.Stderr = &stderr cmd.Stderr = &stderr
cmd.Stdin = bytes.NewReader(encryptedData) cmd.Stdin = bytes.NewReader(encryptedData)
if err := cmd.Run(); err != nil { err := cmd.Run()
if err != nil {
return nil, fmt.Errorf("GPG decryption failed: %w\nStderr: %s", err, stderr.String()) return nil, fmt.Errorf("GPG decryption failed: %w\nStderr: %s", err, stderr.String())
} }
@@ -105,20 +117,24 @@ pinentry-mode loopback
} }
// runGPGWithPassphrase executes a GPG command with the specified passphrase // runGPGWithPassphrase executes a GPG command with the specified passphrase
func runGPGWithPassphrase(gnupgHome, passphrase string, args []string, input io.Reader) ([]byte, error) { func runGPGWithPassphrase(
cmdArgs := []string{ ctx context.Context,
gnupgHome, passphrase string, args []string, input io.Reader,
) ([]byte, error) {
cmdArgs := append([]string{
"--homedir=" + gnupgHome, "--homedir=" + gnupgHome,
"--batch", "--batch",
"--yes", "--yes",
"--pinentry-mode", "loopback", "--pinentry-mode", "loopback",
"--passphrase", passphrase, "--passphrase", passphrase,
} }, args...)
cmdArgs = append(cmdArgs, args...)
cmd := exec.Command("gpg", cmdArgs...) //nolint:gosec // G204: test runs gpg with test-controlled arguments
cmd := exec.CommandContext(ctx, "gpg", cmdArgs...)
cmd.Stdin = input cmd.Stdin = input
var stdout, stderr bytes.Buffer var stdout, stderr bytes.Buffer
cmd.Stdout = &stdout cmd.Stdout = &stdout
cmd.Stderr = &stderr cmd.Stderr = &stderr
@@ -130,14 +146,96 @@ func runGPGWithPassphrase(gnupgHome, passphrase string, args []string, input io.
return stdout.Bytes(), nil return stdout.Bytes(), nil
} }
// generateTestGPGKey generates a GPG key protected by passphrase in
// gnupgHomeDir and returns its key ID and fingerprint.
func generateTestGPGKey(
t *testing.T, tempDir, gnupgHomeDir, passphrase string,
) (string, string) {
t.Helper()
// Create GPG batch file for key generation
batchFile := filepath.Join(tempDir, "gen-key-batch")
batchContent := `%echo Generating a test key
Key-Type: RSA
Key-Length: 2048
Name-Real: Test User
Name-Email: test@example.com
Expire-Date: 0
Passphrase: ` + passphrase + `
%commit
%echo Key generation completed
`
err := os.WriteFile(batchFile, []byte(batchContent), 0o600)
if err != nil {
t.Fatalf("Failed to write batch file: %v", err)
}
// Generate GPG key with batch mode
t.Log("Generating GPG key...")
_, err = runGPGWithPassphrase(t.Context(), gnupgHomeDir, passphrase,
[]string{"--gen-key", batchFile}, nil)
if err != nil {
t.Fatalf("Failed to generate GPG key: %v", err)
}
t.Log("GPG key generated successfully")
// Get the key ID and fingerprint
output, err := runGPGWithPassphrase(t.Context(), gnupgHomeDir, passphrase,
[]string{"--list-secret-keys", "--with-colons", "--fingerprint"}, nil)
if err != nil {
t.Fatalf("Failed to list GPG keys: %v", err)
}
// Parse output to get key ID and fingerprint
var keyID, fingerprint string
for line := range strings.SplitSeq(string(output), "\n") {
if strings.HasPrefix(line, "sec:") {
fields := strings.Split(line, ":")
if len(fields) >= 5 {
keyID = fields[4]
}
} else if strings.HasPrefix(line, "fpr:") {
fields := strings.Split(line, ":")
if len(fields) >= 10 && fields[9] != "" {
fingerprint = fields[9]
break
}
}
}
if keyID == "" {
t.Fatalf("Failed to find GPG key ID in output: %s", output)
}
if fingerprint == "" {
t.Fatalf("Failed to find GPG fingerprint in output: %s", output)
}
t.Logf("Generated GPG key ID: %s", keyID)
t.Logf("Generated GPG fingerprint: %s", fingerprint)
return keyID, fingerprint
}
//nolint:paralleltest // t.Setenv forbids parallel subtests
func TestPGPUnlockerWithRealFS(t *testing.T) { func TestPGPUnlockerWithRealFS(t *testing.T) {
// Check if gpg is available // Check if gpg is available
if _, err := exec.LookPath("gpg"); err != nil { _, err := exec.LookPath("gpg")
if err != nil {
t.Log("GPG not available, PGP unlock key tests may not fully function") t.Log("GPG not available, PGP unlock key tests may not fully function")
// Continue anyway to test what we can // Continue anyway to test what we can
} }
// Create a temporary directory for our tests // Create a temporary directory for our tests. Not t.TempDir: its longer
// path would put gpg-agent's socket in GNUPGHOME past the 104-byte limit
// macOS sets on socket paths.
//
//nolint:usetesting // see the comment above
tempDir, err := os.MkdirTemp("", "secret-pgp-test-") tempDir, err := os.MkdirTemp("", "secret-pgp-test-")
if err != nil { if err != nil {
t.Fatalf("Failed to create temp dir: %v", err) t.Fatalf("Failed to create temp dir: %v", err)
@@ -146,7 +244,9 @@ func TestPGPUnlockerWithRealFS(t *testing.T) {
// Create a temporary GNUPGHOME // Create a temporary GNUPGHOME
gnupgHomeDir := filepath.Join(tempDir, "gnupg") gnupgHomeDir := filepath.Join(tempDir, "gnupg")
if err := os.MkdirAll(gnupgHomeDir, 0o700); err != nil {
err = os.MkdirAll(gnupgHomeDir, 0o700)
if err != nil {
t.Fatalf("Failed to create GNUPGHOME: %v", err) t.Fatalf("Failed to create GNUPGHOME: %v", err)
} }
@@ -159,64 +259,7 @@ func TestPGPUnlockerWithRealFS(t *testing.T) {
// Setup non-interactive GPG with custom functions // Setup non-interactive GPG with custom functions
setupNonInteractiveGPG(t, tempDir, testPassphrase, gnupgHomeDir) setupNonInteractiveGPG(t, tempDir, testPassphrase, gnupgHomeDir)
// Create GPG batch file for key generation keyID, fingerprint := generateTestGPGKey(t, tempDir, gnupgHomeDir, testPassphrase)
batchFile := filepath.Join(tempDir, "gen-key-batch")
batchContent := `%echo Generating a test key
Key-Type: RSA
Key-Length: 2048
Name-Real: Test User
Name-Email: test@example.com
Expire-Date: 0
Passphrase: ` + testPassphrase + `
%commit
%echo Key generation completed
`
if err := os.WriteFile(batchFile, []byte(batchContent), 0o600); err != nil {
t.Fatalf("Failed to write batch file: %v", err)
}
// Generate GPG key with batch mode
t.Log("Generating GPG key...")
_, err = runGPGWithPassphrase(gnupgHomeDir, testPassphrase,
[]string{"--gen-key", batchFile}, nil)
if err != nil {
t.Fatalf("Failed to generate GPG key: %v", err)
}
t.Log("GPG key generated successfully")
// Get the key ID and fingerprint
output, err := runGPGWithPassphrase(gnupgHomeDir, testPassphrase,
[]string{"--list-secret-keys", "--with-colons", "--fingerprint"}, nil)
if err != nil {
t.Fatalf("Failed to list GPG keys: %v", err)
}
// Parse output to get key ID and fingerprint
var keyID, fingerprint string
lines := strings.Split(string(output), "\n")
for _, line := range lines {
if strings.HasPrefix(line, "sec:") {
fields := strings.Split(line, ":")
if len(fields) >= 5 {
keyID = fields[4]
}
} else if strings.HasPrefix(line, "fpr:") {
fields := strings.Split(line, ":")
if len(fields) >= 10 && fields[9] != "" {
fingerprint = fields[9]
break
}
}
}
if keyID == "" {
t.Fatalf("Failed to find GPG key ID in output: %s", output)
}
if fingerprint == "" {
t.Fatalf("Failed to find GPG fingerprint in output: %s", output)
}
t.Logf("Generated GPG key ID: %s", keyID)
t.Logf("Generated GPG fingerprint: %s", fingerprint)
// Set the GPG_AGENT_INFO to empty to ensure gpg-agent doesn't interfere // Set the GPG_AGENT_INFO to empty to ensure gpg-agent doesn't interfere
t.Setenv("GPG_AGENT_INFO", "") t.Setenv("GPG_AGENT_INFO", "")
@@ -224,12 +267,6 @@ Passphrase: ` + testPassphrase + `
// Use the real filesystem // Use the real filesystem
fs := afero.NewOsFs() fs := afero.NewOsFs()
// Test data
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonic))
defer mnemonic.Destroy()
// Set test environment variables // Set test environment variables
t.Setenv(secret.EnvGPGKeyID, keyID) t.Setenv(secret.EnvGPGKeyID, keyID)
@@ -239,12 +276,59 @@ Passphrase: ` + testPassphrase + `
// Test creation of a PGP unlock key through a vault // Test creation of a PGP unlock key through a vault
t.Run("CreatePGPUnlocker", func(t *testing.T) { t.Run("CreatePGPUnlocker", func(t *testing.T) {
testCreatePGPUnlocker(t, fs, stateDir, vaultName, keyID, fingerprint)
})
// Set up key directory for individual tests
unlockerDir := filepath.Join(tempDir, "unlocker")
err = os.MkdirAll(unlockerDir, secret.DirPerms)
if err != nil {
t.Fatalf("Failed to create unlocker directory: %v", err)
}
// Set up test metadata
metadata := secret.UnlockerMetadata{
Type: pgpUnlockerType,
CreatedAt: time.Now(),
Flags: []string{"gpg", "encrypted"},
}
// Create a PGP unlocker for the remaining tests
unlocker := secret.NewPGPUnlocker(fs, unlockerDir, metadata)
// Test getting GPG key ID
t.Run("GetGPGKeyID", func(t *testing.T) {
testGetGPGKeyID(t, fs, unlocker, unlockerDir, metadata, fingerprint)
})
// Test getting identity from PGP unlocker
t.Run("GetIdentity", func(t *testing.T) {
testPGPUnlockerGetIdentity(t, fs, unlocker, unlockerDir, keyID)
})
// Test removing the unlocker
t.Run("RemoveUnlocker", func(t *testing.T) {
testRemovePGPUnlocker(t, fs, unlocker, unlockerDir)
})
}
// testCreatePGPUnlocker creates a vault with a passphrase unlocker, then a
// PGP unlocker for the GPG key keyID, and checks the PGP unlocker's files
// and metadata.
func testCreatePGPUnlocker(
t *testing.T, fs afero.Fs, stateDir, vaultName, keyID, fingerprint string,
) {
t.Helper()
// Set a limited test timeout to avoid hanging // Set a limited test timeout to avoid hanging
timer := time.AfterFunc(30*time.Second, func() { timer := time.AfterFunc(30*time.Second, func() {
t.Fatalf("Test timed out after 30 seconds") t.Fatalf("Test timed out after 30 seconds")
}) })
defer timer.Stop() defer timer.Stop()
mnemonic := testMnemonicBuffer(t)
// Create a test vault directory structure // Create a test vault directory structure
vlt, err := vault.CreateVault(fs, stateDir, vaultName, mnemonic) vlt, err := vault.CreateVault(fs, stateDir, vaultName, mnemonic)
if err != nil { if err != nil {
@@ -271,7 +355,10 @@ Passphrase: ` + testPassphrase + `
// Write long-term public key // Write long-term public key
ltPubKeyPath := filepath.Join(vaultDir, "pub.age") ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
if err := afero.WriteFile(fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), secret.FilePerms); err != nil {
err = afero.WriteFile(fs, ltPubKeyPath,
[]byte(ltIdentity.Recipient().String()), secret.FilePerms)
if err != nil {
t.Fatalf("Failed to write long-term public key: %v", err) t.Fatalf("Failed to write long-term public key: %v", err)
} }
@@ -281,6 +368,7 @@ Passphrase: ` + testPassphrase + `
// Create a passphrase unlocker first (to have current unlocker) // Create a passphrase unlocker first (to have current unlocker)
passphraseBuffer := memguard.NewBufferFromBytes([]byte("test-passphrase")) passphraseBuffer := memguard.NewBufferFromBytes([]byte("test-passphrase"))
defer passphraseBuffer.Destroy() defer passphraseBuffer.Destroy()
passUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer) passUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
if err != nil { if err != nil {
t.Fatalf("Failed to create passphrase unlocker: %v", err) t.Fatalf("Failed to create passphrase unlocker: %v", err)
@@ -292,7 +380,8 @@ Passphrase: ` + testPassphrase + `
} }
// Now create a PGP unlock key (this will use our custom GPGEncryptFunc) // Now create a PGP unlock key (this will use our custom GPGEncryptFunc)
pgpUnlocker, err := secret.CreatePGPUnlocker(fs, stateDir, keyID, fingerprint, mnemonic, nil) pgpUnlocker, err := secret.CreatePGPUnlocker(
fs, stateDir, keyID, fingerprint, mnemonic, nil)
if err != nil { if err != nil {
t.Fatalf("Failed to create PGP unlock key: %v", err) t.Fatalf("Failed to create PGP unlock key: %v", err)
} }
@@ -303,63 +392,91 @@ Passphrase: ` + testPassphrase + `
} }
// Check if the key has the correct type // Check if the key has the correct type
if pgpUnlocker.GetType() != "pgp" { if pgpUnlocker.GetType() != pgpUnlockerType {
t.Errorf("Expected PGP unlock key type 'pgp', got '%s'", pgpUnlocker.GetType()) t.Errorf("Expected PGP unlock key type 'pgp', got '%s'", pgpUnlocker.GetType())
} }
// Check if the key ID includes the GPG fingerprint // Check if the key ID includes the GPG fingerprint
if !strings.Contains(pgpUnlocker.GetID(), fingerprint) { if !strings.Contains(pgpUnlocker.GetID(), fingerprint) {
t.Errorf("PGP unlock key ID '%s' does not contain GPG fingerprint '%s'", pgpUnlocker.GetID(), fingerprint) t.Errorf("PGP unlock key ID '%s' does not contain GPG fingerprint '%s'",
pgpUnlocker.GetID(), fingerprint)
} }
checkPGPUnlockerFiles(t, fs, pgpUnlocker.GetDirectory())
checkPGPUnlockerMetadata(t, fs, pgpUnlocker.GetDirectory(), fingerprint)
}
// checkPGPUnlockerFiles checks that the PGP unlocker in unlockerDir has all
// its files.
func checkPGPUnlockerFiles(t *testing.T, fs afero.Fs, unlockerDir string) {
t.Helper()
// Check if the key directory exists // Check if the key directory exists
unlockerDir := pgpUnlocker.GetDirectory()
keyExists, err := afero.DirExists(fs, unlockerDir) keyExists, err := afero.DirExists(fs, unlockerDir)
if err != nil { if err != nil {
t.Fatalf("Failed to check if PGP key directory exists: %v", err) t.Fatalf("Failed to check if PGP key directory exists: %v", err)
} }
if !keyExists { if !keyExists {
t.Errorf("PGP unlock key directory does not exist: %s", unlockerDir) t.Errorf("PGP unlock key directory does not exist: %s", unlockerDir)
} }
// Check if required files exist // Check if required files exist
recipientPath := filepath.Join(unlockerDir, "pub.txt") recipientPath := filepath.Join(unlockerDir, "pub.txt")
recipientExists, err := afero.Exists(fs, recipientPath) recipientExists, err := afero.Exists(fs, recipientPath)
if err != nil { if err != nil {
t.Fatalf("Failed to check if recipient file exists: %v", err) t.Fatalf("Failed to check if recipient file exists: %v", err)
} }
if !recipientExists { if !recipientExists {
t.Errorf("PGP unlock key recipient file does not exist: %s", recipientPath) t.Errorf("PGP unlock key recipient file does not exist: %s", recipientPath)
} }
privKeyPath := filepath.Join(unlockerDir, "priv.age.gpg") privKeyPath := filepath.Join(unlockerDir, "priv.age.gpg")
privKeyExists, err := afero.Exists(fs, privKeyPath) privKeyExists, err := afero.Exists(fs, privKeyPath)
if err != nil { if err != nil {
t.Fatalf("Failed to check if private key file exists: %v", err) t.Fatalf("Failed to check if private key file exists: %v", err)
} }
if !privKeyExists { if !privKeyExists {
t.Errorf("PGP unlock key private key file does not exist: %s", privKeyPath) t.Errorf("PGP unlock key private key file does not exist: %s", privKeyPath)
} }
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") metadataPath := filepath.Join(unlockerDir, unlockerMetadataFile)
metadataExists, err := afero.Exists(fs, metadataPath) metadataExists, err := afero.Exists(fs, metadataPath)
if err != nil { if err != nil {
t.Fatalf("Failed to check if metadata file exists: %v", err) t.Fatalf("Failed to check if metadata file exists: %v", err)
} }
if !metadataExists { if !metadataExists {
t.Errorf("PGP unlock key metadata file does not exist: %s", metadataPath) t.Errorf("PGP unlock key metadata file does not exist: %s", metadataPath)
} }
longtermPath := filepath.Join(unlockerDir, "longterm.age") longtermPath := filepath.Join(unlockerDir, "longterm.age")
longtermExists, err := afero.Exists(fs, longtermPath) longtermExists, err := afero.Exists(fs, longtermPath)
if err != nil { if err != nil {
t.Fatalf("Failed to check if longterm key file exists: %v", err) t.Fatalf("Failed to check if longterm key file exists: %v", err)
} }
if !longtermExists { if !longtermExists {
t.Errorf("PGP unlock key longterm key file does not exist: %s", longtermPath) t.Errorf("PGP unlock key longterm key file does not exist: %s", longtermPath)
} }
}
// checkPGPUnlockerMetadata checks that the metadata of the PGP unlocker in
// unlockerDir names its type and the GPG key by fingerprint.
func checkPGPUnlockerMetadata(
t *testing.T, fs afero.Fs, unlockerDir, fingerprint string,
) {
t.Helper()
// Read and verify metadata // Read and verify metadata
metadataPath := filepath.Join(unlockerDir, unlockerMetadataFile)
metadataBytes, err := afero.ReadFile(fs, metadataPath) metadataBytes, err := afero.ReadFile(fs, metadataPath)
if err != nil { if err != nil {
t.Fatalf("Failed to read metadata: %v", err) t.Fatalf("Failed to read metadata: %v", err)
@@ -373,40 +490,32 @@ Passphrase: ` + testPassphrase + `
GPGKeyID string `json:"gpgKeyId"` GPGKeyID string `json:"gpgKeyId"`
} }
if err := json.Unmarshal(metadataBytes, &metadata); err != nil { err = json.Unmarshal(metadataBytes, &metadata)
if err != nil {
t.Fatalf("Failed to parse metadata: %v", err) t.Fatalf("Failed to parse metadata: %v", err)
} }
if metadata.Type != "pgp" { if metadata.Type != pgpUnlockerType {
t.Errorf("Expected metadata type 'pgp', got '%s'", metadata.Type) t.Errorf("Expected metadata type 'pgp', got '%s'", metadata.Type)
} }
if metadata.GPGKeyID != fingerprint { if metadata.GPGKeyID != fingerprint {
t.Errorf("Expected GPG fingerprint '%s', got '%s'", fingerprint, metadata.GPGKeyID) t.Errorf("Expected GPG fingerprint '%s', got '%s'", fingerprint, metadata.GPGKeyID)
} }
})
// Set up key directory for individual tests
unlockerDir := filepath.Join(tempDir, "unlocker")
if err := os.MkdirAll(unlockerDir, secret.DirPerms); err != nil {
t.Fatalf("Failed to create unlocker directory: %v", err)
} }
// Set up test metadata // testGetGPGKeyID writes PGP unlocker metadata holding the GPG fingerprint
metadata := secret.UnlockerMetadata{ // into unlockerDir and checks that unlocker reads it back.
Type: "pgp", func testGetGPGKeyID(
CreatedAt: time.Now(), t *testing.T, fs afero.Fs, unlocker *secret.PGPUnlocker,
Flags: []string{"gpg", "encrypted"}, unlockerDir string, metadata secret.UnlockerMetadata, fingerprint string,
} ) {
t.Helper()
// Create a PGP unlocker for the remaining tests
unlocker := secret.NewPGPUnlocker(fs, unlockerDir, metadata)
// Test getting GPG key ID
t.Run("GetGPGKeyID", func(t *testing.T) {
// Create PGP metadata with GPG key ID // Create PGP metadata with GPG key ID
type PGPUnlockerMetadata struct { type PGPUnlockerMetadata struct {
secret.UnlockerMetadata secret.UnlockerMetadata
GPGKeyID string `json:"gpgKeyId"` GPGKeyID string `json:"gpgKeyId"`
} }
@@ -416,12 +525,15 @@ Passphrase: ` + testPassphrase + `
} }
// Write metadata file // Write metadata file
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") metadataPath := filepath.Join(unlockerDir, unlockerMetadataFile)
metadataBytes, err := json.MarshalIndent(pgpMetadata, "", " ") metadataBytes, err := json.MarshalIndent(pgpMetadata, "", " ")
if err != nil { if err != nil {
t.Fatalf("Failed to marshal metadata: %v", err) t.Fatalf("Failed to marshal metadata: %v", err)
} }
if err := afero.WriteFile(fs, metadataPath, metadataBytes, secret.FilePerms); err != nil {
err = afero.WriteFile(fs, metadataPath, metadataBytes, secret.FilePerms)
if err != nil {
t.Fatalf("Failed to write metadata: %v", err) t.Fatalf("Failed to write metadata: %v", err)
} }
@@ -435,10 +547,16 @@ Passphrase: ` + testPassphrase + `
if retrievedKeyID != fingerprint { if retrievedKeyID != fingerprint {
t.Errorf("Expected GPG fingerprint '%s', got '%s'", fingerprint, retrievedKeyID) t.Errorf("Expected GPG fingerprint '%s', got '%s'", fingerprint, retrievedKeyID)
} }
}) }
// testPGPUnlockerGetIdentity writes an age identity encrypted to the GPG key
// keyID into unlockerDir and checks that unlocker decrypts it.
func testPGPUnlockerGetIdentity(
t *testing.T, fs afero.Fs, unlocker *secret.PGPUnlocker,
unlockerDir, keyID string,
) {
t.Helper()
// Test getting identity from PGP unlocker
t.Run("GetIdentity", func(t *testing.T) {
// Generate an age identity for testing // Generate an age identity for testing
ageIdentity, err := age.GenerateX25519Identity() ageIdentity, err := age.GenerateX25519Identity()
if err != nil { if err != nil {
@@ -447,13 +565,17 @@ Passphrase: ` + testPassphrase + `
// Write the recipient // Write the recipient
recipientPath := filepath.Join(unlockerDir, "pub.txt") recipientPath := filepath.Join(unlockerDir, "pub.txt")
if err := afero.WriteFile(fs, recipientPath, []byte(ageIdentity.Recipient().String()), secret.FilePerms); err != nil {
err = afero.WriteFile(fs, recipientPath,
[]byte(ageIdentity.Recipient().String()), secret.FilePerms)
if err != nil {
t.Fatalf("Failed to write recipient: %v", err) t.Fatalf("Failed to write recipient: %v", err)
} }
// GPG encrypt the private key using our custom encrypt function // GPG encrypt the private key using our custom encrypt function
privKeyBuffer := memguard.NewBufferFromBytes([]byte(ageIdentity.String())) privKeyBuffer := memguard.NewBufferFromBytes([]byte(ageIdentity.String()))
defer privKeyBuffer.Destroy() defer privKeyBuffer.Destroy()
encryptedOutput, err := secret.GPGEncryptFunc(privKeyBuffer, keyID) encryptedOutput, err := secret.GPGEncryptFunc(privKeyBuffer, keyID)
if err != nil { if err != nil {
t.Fatalf("Failed to encrypt with GPG: %v", err) t.Fatalf("Failed to encrypt with GPG: %v", err)
@@ -461,7 +583,9 @@ Passphrase: ` + testPassphrase + `
// Write the encrypted data to a file // Write the encrypted data to a file
encryptedPath := filepath.Join(unlockerDir, "priv.age.gpg") encryptedPath := filepath.Join(unlockerDir, "priv.age.gpg")
if err := afero.WriteFile(fs, encryptedPath, encryptedOutput, secret.FilePerms); err != nil {
err = afero.WriteFile(fs, encryptedPath, encryptedOutput, secret.FilePerms)
if err != nil {
t.Fatalf("Failed to write encrypted private key: %v", err) t.Fatalf("Failed to write encrypted private key: %v", err)
} }
@@ -474,18 +598,24 @@ Passphrase: ` + testPassphrase + `
// Verify the identity matches // Verify the identity matches
expectedPubKey := ageIdentity.Recipient().String() expectedPubKey := ageIdentity.Recipient().String()
actualPubKey := identity.Recipient().String() actualPubKey := identity.Recipient().String()
if actualPubKey != expectedPubKey { if actualPubKey != expectedPubKey {
t.Errorf("Expected public key '%s', got '%s'", expectedPubKey, actualPubKey) t.Errorf("Expected public key '%s', got '%s'", expectedPubKey, actualPubKey)
} }
}) }
// testRemovePGPUnlocker removes unlocker and checks that unlockerDir is gone.
func testRemovePGPUnlocker(
t *testing.T, fs afero.Fs, unlocker *secret.PGPUnlocker, unlockerDir string,
) {
t.Helper()
// Test removing the unlocker
t.Run("RemoveUnlocker", func(t *testing.T) {
// Ensure unlocker directory exists before removal // Ensure unlocker directory exists before removal
keyExists, err := afero.DirExists(fs, unlockerDir) keyExists, err := afero.DirExists(fs, unlockerDir)
if err != nil { if err != nil {
t.Fatalf("Failed to check if unlocker directory exists: %v", err) t.Fatalf("Failed to check if unlocker directory exists: %v", err)
} }
if !keyExists { if !keyExists {
t.Fatalf("Unlocker directory does not exist: %s", unlockerDir) t.Fatalf("Unlocker directory does not exist: %s", unlockerDir)
} }
@@ -501,8 +631,8 @@ Passphrase: ` + testPassphrase + `
if err != nil { if err != nil {
t.Fatalf("Failed to check if unlocker directory exists: %v", err) t.Fatalf("Failed to check if unlocker directory exists: %v", err)
} }
if keyExists { if keyExists {
t.Errorf("Unlocker directory still exists after removal: %s", unlockerDir) t.Errorf("Unlocker directory still exists after removal: %s", unlockerDir)
} }
})
} }
+57 -67
View File
@@ -1,5 +1,4 @@
//go:build darwin //go:build darwin
// +build darwin
package secret package secret
@@ -13,7 +12,6 @@ import (
"filippo.io/age" "filippo.io/age"
"git.eeqj.de/sneak/secret/internal/macse" "git.eeqj.de/sneak/secret/internal/macse"
"git.eeqj.de/sneak/secret/pkg/agehd"
"github.com/awnumar/memguard" "github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
) )
@@ -32,6 +30,7 @@ const (
// SecureEnclaveUnlockerMetadata extends UnlockerMetadata with SE-specific data. // SecureEnclaveUnlockerMetadata extends UnlockerMetadata with SE-specific data.
type SecureEnclaveUnlockerMetadata struct { type SecureEnclaveUnlockerMetadata struct {
UnlockerMetadata UnlockerMetadata
SEKeyLabel string `json:"seKeyLabel"` SEKeyLabel string `json:"seKeyLabel"`
SEKeyHash string `json:"seKeyHash"` SEKeyHash string `json:"seKeyHash"`
} }
@@ -43,6 +42,19 @@ type SecureEnclaveUnlocker struct {
fs afero.Fs fs afero.Fs
} }
// NewSecureEnclaveUnlocker creates a new SecureEnclaveUnlocker instance.
func NewSecureEnclaveUnlocker(
fs afero.Fs,
directory string,
metadata UnlockerMetadata,
) *SecureEnclaveUnlocker {
return &SecureEnclaveUnlocker{
Directory: directory,
Metadata: metadata,
fs: fs,
}
}
// GetIdentity implements Unlocker interface for SE-based unlockers. // GetIdentity implements Unlocker interface for SE-based unlockers.
// Decrypts the vault's long-term private key directly using the Secure Enclave. // Decrypts the vault's long-term private key directly using the Secure Enclave.
func (s *SecureEnclaveUnlocker) GetIdentity() (*age.X25519Identity, error) { func (s *SecureEnclaveUnlocker) GetIdentity() (*age.X25519Identity, error) {
@@ -58,6 +70,7 @@ func (s *SecureEnclaveUnlocker) GetIdentity() (*age.X25519Identity, error) {
// Read ECIES-encrypted long-term private key from disk // Read ECIES-encrypted long-term private key from disk
encryptedPath := filepath.Join(s.Directory, seLongtermFilename) encryptedPath := filepath.Join(s.Directory, seLongtermFilename)
encryptedData, err := afero.ReadFile(s.fs, encryptedPath) encryptedData, err := afero.ReadFile(s.fs, encryptedPath)
if err != nil { if err != nil {
return nil, fmt.Errorf( return nil, fmt.Errorf(
@@ -140,7 +153,9 @@ func (s *SecureEnclaveUnlocker) Remove() error {
if seKeyHash != "" { if seKeyHash != "" {
Debug("Deleting SE key", "hash", seKeyHash) Debug("Deleting SE key", "hash", seKeyHash)
if err := macse.DeleteKey(seKeyHash); err != nil {
err = macse.DeleteKey(seKeyHash)
if err != nil {
Debug("Failed to delete SE key", "error", err, "hash", seKeyHash) Debug("Failed to delete SE key", "error", err, "hash", seKeyHash)
return fmt.Errorf("failed to delete SE key: %w", err) return fmt.Errorf("failed to delete SE key: %w", err)
@@ -148,7 +163,9 @@ func (s *SecureEnclaveUnlocker) Remove() error {
} }
Debug("Removing SE unlocker directory", "directory", s.Directory) Debug("Removing SE unlocker directory", "directory", s.Directory)
if err := RemoveDirAtomic(s.fs, s.Directory); err != nil {
err = RemoveDirAtomic(s.fs, s.Directory)
if err != nil {
return fmt.Errorf("failed to remove SE unlocker directory: %w", err) return fmt.Errorf("failed to remove SE unlocker directory: %w", err)
} }
@@ -158,34 +175,24 @@ func (s *SecureEnclaveUnlocker) Remove() error {
} }
// getSEKeyInfo reads the SE key label and hash from metadata. // getSEKeyInfo reads the SE key label and hash from metadata.
func (s *SecureEnclaveUnlocker) getSEKeyInfo() (label string, hash string, err error) { func (s *SecureEnclaveUnlocker) getSEKeyInfo() (string, string, error) {
metadataPath := filepath.Join(s.Directory, "unlocker-metadata.json") metadataPath := filepath.Join(s.Directory, "unlocker-metadata.json")
metadataData, err := afero.ReadFile(s.fs, metadataPath) metadataData, err := afero.ReadFile(s.fs, metadataPath)
if err != nil { if err != nil {
return "", "", fmt.Errorf("failed to read SE metadata: %w", err) return "", "", fmt.Errorf("failed to read SE metadata: %w", err)
} }
var seMetadata SecureEnclaveUnlockerMetadata var seMetadata SecureEnclaveUnlockerMetadata
if err := json.Unmarshal(metadataData, &seMetadata); err != nil {
err = json.Unmarshal(metadataData, &seMetadata)
if err != nil {
return "", "", fmt.Errorf("failed to parse SE metadata: %w", err) return "", "", fmt.Errorf("failed to parse SE metadata: %w", err)
} }
return seMetadata.SEKeyLabel, seMetadata.SEKeyHash, nil 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. // generateSEKeyLabel generates a unique label for the SE CTK identity.
func generateSEKeyLabel(vaultName string) (string, error) { func generateSEKeyLabel(vaultName string) (string, error) {
hostname, err := os.Hostname() hostname, err := os.Hostname()
@@ -214,7 +221,8 @@ func CreateSecureEnclaveUnlocker(
stateDir string, stateDir string,
mnemonic, passphrase *memguard.LockedBuffer, mnemonic, passphrase *memguard.LockedBuffer,
) (*SecureEnclaveUnlocker, error) { ) (*SecureEnclaveUnlocker, error) {
if err := checkMacOSAvailable(); err != nil { err := checkMacOSAvailable()
if err != nil {
return nil, err return nil, err
} }
@@ -231,6 +239,7 @@ func CreateSecureEnclaveUnlocker(
// Step 1: Create P-256 key in the Secure Enclave via sc_auth // Step 1: Create P-256 key in the Secure Enclave via sc_auth
Debug("Creating Secure Enclave key", "label", seKeyLabel) Debug("Creating Secure Enclave key", "label", seKeyLabel)
_, seKeyHash, err := macse.CreateKey(seKeyLabel) _, seKeyHash, err := macse.CreateKey(seKeyLabel)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to create SE key: %w", err) return nil, fmt.Errorf("failed to create SE key: %w", err)
@@ -263,14 +272,14 @@ func CreateSecureEnclaveUnlocker(
return nil, fmt.Errorf("failed to get vault directory: %w", err) return nil, fmt.Errorf("failed to get vault directory: %w", err)
} }
unlockerDirName := fmt.Sprintf("se-%s", filepath.Base(seKeyLabel)) unlockerDirName := "se-" + filepath.Base(seKeyLabel)
unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerDirName) unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerDirName)
seMetadata := SecureEnclaveUnlockerMetadata{ seMetadata := SecureEnclaveUnlockerMetadata{
UnlockerMetadata: UnlockerMetadata{ UnlockerMetadata: UnlockerMetadata{
Type: seUnlockerType, Type: seUnlockerType,
CreatedAt: time.Now().UTC(), CreatedAt: time.Now().UTC(),
Flags: []string{seUnlockerType, "macos"}, Flags: []string{seUnlockerType, macOSFlag},
}, },
SEKeyLabel: seKeyLabel, SEKeyLabel: seKeyLabel,
SEKeyHash: seKeyHash, SEKeyHash: seKeyHash,
@@ -283,20 +292,7 @@ func CreateSecureEnclaveUnlocker(
// Step 5: Write the SE-encrypted long-term key, then the metadata // Step 5: Write the SE-encrypted long-term key, then the metadata
err = WriteDir(fs, unlockerDir, func(dir string) error { err = WriteDir(fs, unlockerDir, func(dir string) error {
ltKeyPath := filepath.Join(dir, seLongtermFilename) return writeSEUnlockerFiles(fs, dir, encryptedLtKey, metadataBytes)
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 { if err != nil {
return nil, err return nil, err
@@ -309,6 +305,29 @@ func CreateSecureEnclaveUnlocker(
}, nil }, nil
} }
// writeSEUnlockerFiles writes the files of a new SE unlocker into dir: the
// SE-encrypted long-term key, then the metadata.
func writeSEUnlockerFiles(
fs afero.Fs, dir string, encryptedLtKey, metadataBytes []byte,
) error {
err := WriteFileAtomic(fs, filepath.Join(dir, seLongtermFilename),
encryptedLtKey)
if err != nil {
return fmt.Errorf(
"failed to write SE-encrypted long-term key: %w",
err,
)
}
err = WriteFileAtomic(fs,
filepath.Join(dir, "unlocker-metadata.json"), metadataBytes)
if err != nil {
return fmt.Errorf("failed to write metadata: %w", err)
}
return nil
}
// getLongTermKeyForSE retrieves the vault's long-term private key, derived // getLongTermKeyForSE retrieves the vault's long-term private key, derived
// from mnemonic when it is not nil, else through the current unlocker, which // from mnemonic when it is not nil, else through the current unlocker, which
// is given passphrase when it is a passphrase unlocker. // is given passphrase when it is a passphrase unlocker.
@@ -318,37 +337,7 @@ func getLongTermKeyForSE(
mnemonic, passphrase *memguard.LockedBuffer, mnemonic, passphrase *memguard.LockedBuffer,
) (*memguard.LockedBuffer, error) { ) (*memguard.LockedBuffer, error) {
if mnemonic != nil { if mnemonic != nil {
// Read vault metadata to get the correct derivation index return deriveLongTermPrivateKey(fs, vault, mnemonic)
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(
mnemonic.String(),
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() currentUnlocker, err := vault.GetCurrentUnlocker()
@@ -373,6 +362,7 @@ func getLongTermKeyForSE(
currentUnlocker.GetDirectory(), currentUnlocker.GetDirectory(),
"longterm.age", "longterm.age",
) )
encryptedLtKey, err := afero.ReadFile(fs, longtermPath) encryptedLtKey, err := afero.ReadFile(fs, longtermPath)
if err != nil { if err != nil {
return nil, fmt.Errorf( return nil, fmt.Errorf(
+20 -8
View File
@@ -1,6 +1,6 @@
//go:build darwin //go:build darwin
// +build darwin
//nolint:testpackage // white-box test of unexported Secure Enclave helpers
package secret package secret
import ( import (
@@ -13,12 +13,14 @@ import (
) )
func TestNewSecureEnclaveUnlocker(t *testing.T) { func TestNewSecureEnclaveUnlocker(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
dir := "/tmp/test-se-unlocker" dir := "/tmp/test-se-unlocker"
metadata := UnlockerMetadata{ metadata := UnlockerMetadata{
Type: "secure-enclave", Type: seUnlockerType,
CreatedAt: time.Date(2026, 1, 15, 10, 30, 0, 0, time.UTC), CreatedAt: time.Date(2026, 1, 15, 10, 30, 0, 0, time.UTC),
Flags: []string{"secure-enclave", "macos"}, Flags: []string{seUnlockerType, "macos"},
} }
unlocker := NewSecureEnclaveUnlocker(fs, dir, metadata) unlocker := NewSecureEnclaveUnlocker(fs, dir, metadata)
@@ -35,9 +37,11 @@ func TestNewSecureEnclaveUnlocker(t *testing.T) {
} }
func TestSecureEnclaveUnlockerImplementsInterface(t *testing.T) { func TestSecureEnclaveUnlockerImplementsInterface(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
metadata := UnlockerMetadata{ metadata := UnlockerMetadata{
Type: "secure-enclave", Type: seUnlockerType,
CreatedAt: time.Now().UTC(), CreatedAt: time.Now().UTC(),
} }
@@ -48,9 +52,11 @@ func TestSecureEnclaveUnlockerImplementsInterface(t *testing.T) {
} }
func TestSecureEnclaveUnlockerGetIDFormat(t *testing.T) { func TestSecureEnclaveUnlockerGetIDFormat(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
metadata := UnlockerMetadata{ metadata := UnlockerMetadata{
Type: "secure-enclave", Type: seUnlockerType,
CreatedAt: time.Date(2026, 3, 10, 14, 30, 0, 0, time.UTC), CreatedAt: time.Date(2026, 3, 10, 14, 30, 0, 0, time.UTC),
} }
@@ -63,6 +69,8 @@ func TestSecureEnclaveUnlockerGetIDFormat(t *testing.T) {
} }
func TestGenerateSEKeyLabel(t *testing.T) { func TestGenerateSEKeyLabel(t *testing.T) {
t.Parallel()
label, err := generateSEKeyLabel("test-vault") label, err := generateSEKeyLabel("test-vault")
require.NoError(t, err) require.NoError(t, err)
@@ -72,6 +80,8 @@ func TestGenerateSEKeyLabel(t *testing.T) {
} }
func TestSecureEnclaveUnlockerGetIdentityMissingFile(t *testing.T) { func TestSecureEnclaveUnlockerGetIdentityMissingFile(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
dir := "/tmp/test-se-unlocker-missing" dir := "/tmp/test-se-unlocker-missing"
@@ -84,10 +94,12 @@ func TestSecureEnclaveUnlockerGetIdentityMissingFile(t *testing.T) {
"seKeyLabel": "berlin.sneak.app.secret.se.test", "seKeyLabel": "berlin.sneak.app.secret.se.test",
"seKeyHash": "abc123" "seKeyHash": "abc123"
}` }`
require.NoError(t, afero.WriteFile(fs, dir+"/unlocker-metadata.json", []byte(metadataJSON), FilePerms)) require.NoError(t, afero.WriteFile(
fs, dir+"/unlocker-metadata.json", []byte(metadataJSON), FilePerms,
))
metadata := UnlockerMetadata{ metadata := UnlockerMetadata{
Type: "secure-enclave", Type: seUnlockerType,
CreatedAt: time.Date(2026, 1, 15, 10, 30, 0, 0, time.UTC), CreatedAt: time.Date(2026, 1, 15, 10, 30, 0, 0, time.UTC),
} }
@@ -96,6 +108,6 @@ func TestSecureEnclaveUnlockerGetIdentityMissingFile(t *testing.T) {
// GetIdentity should fail because the encrypted longterm key file is missing // GetIdentity should fail because the encrypted longterm key file is missing
identity, err := unlocker.GetIdentity() identity, err := unlocker.GetIdentity()
assert.Nil(t, identity) assert.Nil(t, identity)
assert.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "failed to read SE-encrypted long-term key") assert.Contains(t, err.Error(), "failed to read SE-encrypted long-term key")
} }
+29 -120
View File
@@ -1,5 +1,6 @@
//go:build darwin //go:build darwin
//nolint:testpackage // white-box test of unexported validateKeychainItemName
package secret package secret
import ( import (
@@ -7,138 +8,46 @@ import (
) )
func TestValidateKeychainItemName(t *testing.T) { func TestValidateKeychainItemName(t *testing.T) {
t.Parallel()
tests := []struct { tests := []struct {
name string name string
itemName string itemName string
wantErr bool wantErr bool
}{ }{
// Valid cases // Valid cases
{ {name: "valid simple name", itemName: "my-secret-key", wantErr: false},
name: "valid simple name", {name: "valid name with dots", itemName: "com.example.app.key", wantErr: false},
itemName: "my-secret-key", {name: "valid name with underscores", itemName: "my_secret_key_123", wantErr: false},
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},
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 // Invalid cases
{ {name: "empty item name", itemName: "", wantErr: true},
name: "empty item name", {name: "item name with spaces", itemName: "my secret key", wantErr: true},
itemName: "", {name: "item name with semicolon", itemName: "key;rm -rf /", wantErr: true},
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 spaces", {name: "item name with quotes", itemName: "key\"name", wantErr: true},
itemName: "my secret key", {name: "item name with single quotes", itemName: "key'name", wantErr: true},
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 semicolon", {name: "item name with ampersand", itemName: "key&echo test", wantErr: true},
itemName: "key;rm -rf /", {name: "item name with redirect", itemName: "key>/tmp/test", wantErr: true},
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 pipe", {name: "item name with asterisk", itemName: "key*", wantErr: true},
itemName: "key|cat /etc/passwd", {name: "item name with question mark", itemName: "key?", wantErr: true},
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 { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
t.Parallel()
err := validateKeychainItemName(tt.itemName) err := validateKeychainItemName(tt.itemName)
if (err != nil) != tt.wantErr { if (err != nil) != tt.wantErr {
t.Errorf("validateKeychainItemName() error = %v, wantErr %v", err, tt.wantErr) t.Errorf("validateKeychainItemName() error = %v, wantErr %v", err, tt.wantErr)
+18
View File
@@ -274,6 +274,24 @@ func (v *Vault) readUnlockerMetadataOrWarn(
return metadata, true return metadata, true
} }
// HasUnlocker reports whether RemoveUnlocker finds something to remove by
// the ID unlockerID: an unlocker with that ID, or an unlocker directory of
// that name that ListUnlockers skips.
func (v *Vault) HasUnlocker(unlockerID string) (bool, error) {
vaultDir, err := v.GetDirectory()
if err != nil {
return false, err
}
_, unlockerDir, err := v.findUnlockerByID(
filepath.Join(vaultDir, "unlockers.d"), unlockerID)
if err != nil {
return false, err
}
return unlockerDir != "", nil
}
// RemoveUnlocker removes an unlocker from this vault. An unlocker // RemoveUnlocker removes an unlocker from this vault. An unlocker
// directory that ListUnlockers skips is removed by its directory name; its // directory that ListUnlockers skips is removed by its directory name; its
// type is unknown, so only the directory is removed. // type is unknown, so only the directory is removed.
+3 -3
View File
@@ -1,7 +1,6 @@
#!/bin/sh #!/bin/sh
# script/check: run all checks (test, lint, fmt-check). Our own # script/check: run all checks (test, lint, lint-darwin, fmt-check). Our
# extension to scripts-to-rule-them-all. Must not modify any files. # own extension to scripts-to-rule-them-all. Must not modify any files.
# Generic: usually needs no adaptation.
set -eu set -eu
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -9,6 +8,7 @@ SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
main() { main() {
"$SCRIPT_DIR/test" "$SCRIPT_DIR/test"
"$SCRIPT_DIR/lint" "$SCRIPT_DIR/lint"
"$SCRIPT_DIR/lint-darwin"
"$SCRIPT_DIR/fmt-check" "$SCRIPT_DIR/fmt-check"
} }
+26
View File
@@ -0,0 +1,26 @@
#!/bin/sh
# script/lint-darwin: type-check (go vet) and lint the code as a macOS
# build compiles it, from Linux, in docker only. CI runs on Linux, which
# never compiles the files built only for macOS. Builds the lint-darwin
# stage of Dockerfile.lint, rebuilt on every run as script/lint does.
#
# Cgo is off: compiling cgo code for macOS needs Apple's SDK headers. That
# leaves out the files built only with cgo on macOS: the keychain unlocker's
# calls into the keychain (keychainunlocker_cgo.go, and
# keychainunlocker_test.go) and the Secure Enclave bindings (internal/macse).
# Nothing on Linux checks those.
set -eu
ROOT="$(cd "$(dirname "$0")/.." && pwd -P)"
main() {
cd "$ROOT"
docker build \
--progress=plain \
--target lint-darwin \
--no-cache-filter=lint-darwin \
--output=type=cacheonly \
-f Dockerfile.lint .
}
main "$@"