Read secret environment variables once per command, then unset them (closes #60) #94

Merged
clawbot merged 1 commits from issue-60-secrets-out-of-env into next 2026-10-04 16:07:58 +02:00
42 changed files with 727 additions and 428 deletions
+12
View File
@@ -318,6 +318,18 @@ Each vault maintains its own set of unlockers and one long-term key. The long-te
- `SB_UNLOCK_PASSPHRASE`: Pre-set unlock passphrase (avoids interactive prompt) - `SB_UNLOCK_PASSPHRASE`: Pre-set unlock passphrase (avoids interactive prompt)
- `SB_GPG_KEY_ID`: GPG key ID for PGP unlockers - `SB_GPG_KEY_ID`: GPG key ID for PGP unlockers
**Warning:** `SB_SECRET_MNEMONIC` and `SB_UNLOCK_PASSPHRASE` expose the secret
they hold. Other processes running as the same user can read a process's
environment (on Linux, from `/proc/<pid>/environ`). Every child process of the
shell or script that sets them inherits them, `gpg` included. Set on a command
line or in a CI job, they end up in shell history and CI logs. `secret` unsets
each one as soon as it has read it, so that the programs it runs itself, such
as `gpg`, do not inherit it, but that erases nothing: the environment the
process started with, and its memory, still hold the value. The interactive
prompt, which every command except `secret vault import` offers when the
variable is not set, is the safer default; `secret vault import` has no prompt
and needs both variables.
## Security Features ## Security Features
### Encryption ### Encryption
+14 -2
View File
@@ -25,6 +25,20 @@ Bring the repo into policy compliance in one commit:
# Completed Steps # Completed Steps
- 2026-10-04: `SB_SECRET_MNEMONIC` and `SB_UNLOCK_PASSPHRASE` are read once
per command, in its `RunE`, into locked buffers on the CLI `Instance`, and
unset at once, so that no program the command runs, `gpg` included,
inherits them (https://git.eeqj.de/sneak/secret/issues/60). Nothing below
the command reads the environment; the buffers are passed down:
`vault.CreateVault` takes the mnemonic (nil for none), a `Vault` derives its
long-term key from its `Mnemonic` and gives its `UnlockPassphrase` to a
passphrase unlocker, and the PGP, keychain and Secure Enclave unlocker
constructors take both. `CreatePGPUnlocker` sets both on the vault it
loads, through `SetMnemonic` and `SetUnlockPassphrase`, now part of
`VaultInterface`, before calling its `GetOrDeriveLongTermKey`. `init` and
`vault create` no longer put the mnemonic into the environment. Unsetting
erases nothing: the starting environment (`/proc/<pid>/environ`) and
memory still hold the value. The README warns against both variables.
- 2026-10-04: `.golangci.yml` is again the canonical file from - 2026-10-04: `.golangci.yml` is again the canonical file from
`sneak/prompts`, byte for byte `sneak/prompts`, byte for byte
(https://git.eeqj.de/sneak/secret/issues/66). It runs `gomodguard_v2` (https://git.eeqj.de/sneak/secret/issues/66). It runs `gomodguard_v2`
@@ -293,8 +307,6 @@ Bring the repo into policy compliance in one commit:
suggestions. suggestions.
- Validate GPG key existence before creating PGP unlock keys. - Validate GPG key existence before creating PGP unlock keys.
- Split oversized CLI functions. - Split oversized CLI functions.
- Document env var security (SB_UNLOCK_PASSPHRASE,
SB_SECRET_MNEMONIC); clear after use.
- mlock/munlock for sensitive allocations. - mlock/munlock for sensitive allocations.
- Cleanups: read statedir from environment or default instead of - Cleanups: read statedir from environment or default instead of
passing it around. passing it around.
+47
View File
@@ -3,8 +3,10 @@ package cli
import ( import (
"fmt" "fmt"
"os"
"git.eeqj.de/sneak/secret/internal/secret" "git.eeqj.de/sneak/secret/internal/secret"
"github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
"github.com/spf13/cobra" "github.com/spf13/cobra"
) )
@@ -14,6 +16,11 @@ type Instance struct {
fs afero.Fs fs afero.Fs
stateDir string stateDir string
cmd *cobra.Command cmd *cobra.Command
// Mnemonic and UnlockPassphrase hold the values of SB_SECRET_MNEMONIC
// and SB_UNLOCK_PASSPHRASE that readSecretEnv read, or nil when it found
// none.
Mnemonic *memguard.LockedBuffer
UnlockPassphrase *memguard.LockedBuffer
} }
// NewCLIInstance creates a new CLI instance with the real filesystem // NewCLIInstance creates a new CLI instance with the real filesystem
@@ -68,3 +75,43 @@ func (cli *Instance) SetStateDir(stateDir string) {
func (cli *Instance) GetStateDir() string { func (cli *Instance) GetStateDir() string {
return cli.stateDir return cli.stateDir
} }
// readSecretEnv reads SB_SECRET_MNEMONIC into cli.Mnemonic and
// SB_UNLOCK_PASSPHRASE into cli.UnlockPassphrase. A command that may need
// either calls it once, before anything else, and passes the buffers on
// from there: each variable is unset as soon as it is read, so that the
// processes this one starts, gpg among them, do not inherit it, and a
// second read would find nothing. The returned function destroys both
// buffers.
func (cli *Instance) readSecretEnv() func() {
cli.Mnemonic = readAndUnsetEnv(secret.EnvMnemonic)
cli.UnlockPassphrase = readAndUnsetEnv(secret.EnvUnlockPassphrase)
mnemonic, passphrase := cli.Mnemonic, cli.UnlockPassphrase
return func() {
if mnemonic != nil {
mnemonic.Destroy()
}
if passphrase != nil {
passphrase.Destroy()
}
}
}
// readAndUnsetEnv returns the value of the environment variable name in a
// locked buffer, or nil when it is unset or empty, and unsets the variable.
// Unsetting does not erase the value: it stays in this process's memory,
// and in /proc/<pid>/environ, which shows the environment the process
// started with. The caller must destroy the returned buffer.
func readAndUnsetEnv(name string) *memguard.LockedBuffer {
value := os.Getenv(name)
_ = os.Unsetenv(name)
if value == "" {
return nil
}
return memguard.NewBufferFromBytes([]byte(value))
}
+58 -16
View File
@@ -2,6 +2,7 @@ package cli_test
import ( import (
"bytes" "bytes"
"os"
"testing" "testing"
"git.eeqj.de/sneak/secret/internal/cli" "git.eeqj.de/sneak/secret/internal/cli"
@@ -20,16 +21,27 @@ import (
// decrypted any more. Each must refuse, change nothing, and leave every // decrypted any more. Each must refuse, change nothing, and leave every
// vault's secret readable through its passphrase unlocker. // vault's secret readable through its passphrase unlocker.
// //
//nolint:paralleltest // t.Setenv forbids parallel subtests //nolint:paralleltest // the cases share cmd
func TestCreateExistingVaultChangesNothing(t *testing.T) { func TestCreateExistingVaultChangesNothing(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) mnemonic := testMnemonicBuffer(t)
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase) passphrase := memguard.NewBufferFromBytes([]byte(testPassphrase))
t.Cleanup(passphrase.Destroy)
// newCLI returns an instance on fs given the mnemonic and the unlock
// passphrase, as from the environment
newCLI := func(fs afero.Fs) *cli.Instance {
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
c.Mnemonic = mnemonic
c.UnlockPassphrase = passphrase
return c
}
// `secret init`, `secret vault create work`, `secret vault select // `secret init`, `secret vault create work`, `secret vault select
// default`, and the secret "x" in each vault. "work" is then not the // default`, and the secret "x" in each vault. "work" is then not the
// current vault, which creating it again must not change. // current vault, which creating it again must not change.
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir) c := newCLI(fs)
cmd := &cobra.Command{} cmd := &cobra.Command{}
require.NoError(t, c.Init(cmd)) require.NoError(t, c.Init(cmd))
@@ -74,7 +86,7 @@ func TestCreateExistingVaultChangesNothing(t *testing.T) {
t.Run(tt.command, func(t *testing.T) { t.Run(tt.command, func(t *testing.T) {
fs := newFsFromSnapshot(t, before) fs := newFsFromSnapshot(t, before)
err := tt.run(cli.NewCLIInstanceWithStateDir(fs, testStateDir)) err := tt.run(newCLI(fs))
require.EqualError(t, err, tt.want) require.EqualError(t, err, tt.want)
require.Equal(t, before, snapshotStateDir(t, fs)) require.Equal(t, before, snapshotStateDir(t, fs))
@@ -85,10 +97,11 @@ func TestCreateExistingVaultChangesNothing(t *testing.T) {
// reading each vault's secret once from it shows that it still decrypts // reading each vault's secret once from it shows that it still decrypts
// after each case. Without the mnemonic, reading a secret goes through // after each case. Without the mnemonic, reading a secret goes through
// the vault's passphrase unlocker, which is slow. // the vault's passphrase unlocker, which is slow.
t.Setenv(secret.EnvMnemonic, "")
for _, name := range vaults { for _, name := range vaults {
value, err := vault.NewVault(fs, testStateDir, name).GetSecret("x") vlt := vault.NewVault(fs, testStateDir, name)
vlt.UnlockPassphrase = passphrase
value, err := vlt.GetSecret("x")
require.NoError(t, err) require.NoError(t, err)
unchanged := bytes.Equal([]byte("value"), value.Bytes()) unchanged := bytes.Equal([]byte("value"), value.Bytes())
@@ -98,19 +111,43 @@ func TestCreateExistingVaultChangesNothing(t *testing.T) {
} }
} }
// TestVaultCreationLeavesNoSecretInEnvironment is a regression test for
// https://git.eeqj.de/sneak/secret/issues/60, where `secret init` and
// `secret vault create` put the mnemonic into the process environment,
// which every program they ran inherited, and SB_SECRET_MNEMONIC and
// SB_UNLOCK_PASSPHRASE were never unset. Each command, given both, must
// leave neither in the environment.
func TestVaultCreationLeavesNoSecretInEnvironment(t *testing.T) {
t.Setenv(secret.EnvStateDir, t.TempDir())
run := func(args ...string) {
t.Setenv(secret.EnvMnemonic, testMnemonic)
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
// With no terminal to prompt on, this succeeds only if the command
// read both variables
_, err := cli.ExecuteCommandInProcess(args, "", nil)
require.NoError(t, err)
for _, name := range []string{secret.EnvMnemonic, secret.EnvUnlockPassphrase} {
_, set := os.LookupEnv(name)
require.False(t, set, "%s is set after %v", name, args)
}
}
run("init")
run("vault", "create", "work")
}
// TestStopAtPassphrasePromptLeavesNothing is a regression test for the // TestStopAtPassphrasePromptLeavesNothing is a regression test for the
// review of https://git.eeqj.de/sneak/secret/pulls/82: `secret init` or // review of https://git.eeqj.de/sneak/secret/pulls/82: `secret init` or
// `secret vault create` stopped at the passphrase prompt left a vault with // `secret vault create` stopped at the passphrase prompt left a vault with
// no unlocker, which neither command would then create again. Each must ask // no unlocker, which neither command would then create again. Each must ask
// for the passphrase before writing anything. // for the passphrase before writing anything.
// //
//nolint:paralleltest // t.Setenv forbids parallel subtests //nolint:paralleltest // the cases share cmd
func TestStopAtPassphrasePromptLeavesNothing(t *testing.T) { func TestStopAtPassphrasePromptLeavesNothing(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) mnemonic := testMnemonicBuffer(t)
// Without the passphrase in the environment, both commands prompt for
// it, which fails because the tests do not run in a terminal.
t.Setenv(secret.EnvUnlockPassphrase, "")
// An empty state directory for `secret init`, and one holding the vault // An empty state directory for `secret init`, and one holding the vault
// "default" for `secret vault create work`. // "default" for `secret vault create work`.
@@ -118,7 +155,7 @@ func TestStopAtPassphrasePromptLeavesNothing(t *testing.T) {
require.NoError(t, empty.MkdirAll(testStateDir, secret.DirPerms)) require.NoError(t, empty.MkdirAll(testStateDir, secret.DirPerms))
withDefault := afero.NewMemMapFs() withDefault := afero.NewMemMapFs()
_, err := vault.CreateVault(withDefault, testStateDir, "default") _, err := vault.CreateVault(withDefault, testStateDir, "default", mnemonic)
require.NoError(t, err) require.NoError(t, err)
cmd := &cobra.Command{} cmd := &cobra.Command{}
@@ -144,7 +181,12 @@ func TestStopAtPassphrasePromptLeavesNothing(t *testing.T) {
t.Run(tt.command, func(t *testing.T) { t.Run(tt.command, func(t *testing.T) {
before := snapshotStateDir(t, tt.fs) before := snapshotStateDir(t, tt.fs)
err := tt.run(cli.NewCLIInstanceWithStateDir(tt.fs, testStateDir)) // Given no unlock passphrase, both commands prompt for it, which
// fails because the tests do not run in a terminal.
c := cli.NewCLIInstanceWithStateDir(tt.fs, testStateDir)
c.Mnemonic = mnemonic
err := tt.run(c)
require.ErrorContains(t, err, "failed to read passphrase") require.ErrorContains(t, err, "failed to read passphrase")
require.Equal(t, before, snapshotStateDir(t, tt.fs)) require.Equal(t, before, snapshotStateDir(t, tt.fs))
+12 -5
View File
@@ -41,6 +41,9 @@ func newCryptoCmd(
cli.cmd = cmd cli.cmd = cmd
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
return run(cli, args[0], inputFile, outputFile) return run(cli, args[0], inputFile, outputFile)
}, },
} }
@@ -156,6 +159,8 @@ func (cli *Instance) Encrypt(secretName, inputFile, outputFile string) error {
return err return err
} }
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
// Get or create the age secret key for this secret // Get or create the age secret key for this secret
keyBuffer, err := cli.resolveEncryptionKey(vlt, secretName) keyBuffer, err := cli.resolveEncryptionKey(vlt, secretName)
if err != nil { if err != nil {
@@ -230,6 +235,8 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error {
return err return err
} }
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
// Check if secret exists // Check if secret exists
secretObj := secret.NewSecret(vlt, secretName) secretObj := secret.NewSecret(vlt, secretName)
@@ -308,13 +315,13 @@ func isValidAgeSecretKey(key string) bool {
return err == nil return err == nil
} }
// getSecretValue retrieves the value of a secret using the appropriate // getSecretValue retrieves the value of a secret with the vault's mnemonic
// unlocker // when it has one, else with the current unlocker
func (cli *Instance) getSecretValue( func (cli *Instance) getSecretValue(
vlt *vault.Vault, secretObj *secret.Secret, vlt *vault.Vault, secretObj *secret.Secret,
) (*memguard.LockedBuffer, error) { ) (*memguard.LockedBuffer, error) {
if os.Getenv(secret.EnvMnemonic) != "" { if vlt.Mnemonic != nil {
return secretObj.GetValue(nil) return secretObj.GetValue(nil, vlt.Mnemonic)
} }
unlocker, err := vlt.GetCurrentUnlocker() unlocker, err := vlt.GetCurrentUnlocker()
@@ -322,5 +329,5 @@ func (cli *Instance) getSecretValue(
return nil, fmt.Errorf("failed to get current unlocker: %w", err) return nil, fmt.Errorf("failed to get current unlocker: %w", err)
} }
return secretObj.GetValue(unlocker) return secretObj.GetValue(unlocker, nil)
} }
+5
View File
@@ -76,6 +76,9 @@ func newGenerateSecretCmd() *cobra.Command {
return fmt.Errorf("failed to initialize CLI: %w", err) return fmt.Errorf("failed to initialize CLI: %w", err)
} }
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
return cli.GenerateSecret(cmd, args[0], length, secretType, force) return cli.GenerateSecret(cmd, args[0], length, secretType, force)
}, },
} }
@@ -167,6 +170,8 @@ func (cli *Instance) GenerateSecret(
return err return err
} }
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
// Protect the generated secret immediately // Protect the generated secret immediately
secretBuffer := memguard.NewBufferFromBytes([]byte(secretValue)) secretBuffer := memguard.NewBufferFromBytes([]byte(secretValue))
defer secretBuffer.Destroy() defer secretBuffer.Destroy()
+19 -18
View File
@@ -39,16 +39,20 @@ func RunInit(cmd *cobra.Command, _ []string) error {
log.Fatalf("failed to initialize CLI: %v", err) log.Fatalf("failed to initialize CLI: %v", err)
} }
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
return cli.Init(cmd) return cli.Init(cmd)
} }
// promptMnemonic reads the mnemonic from the environment or interactively. // promptMnemonic returns the mnemonic from the environment, cli.Mnemonic,
// The returned cleanup function must be deferred by the caller. // or reads it interactively. The returned cleanup function must be deferred
func promptMnemonic() (string, func(), error) { // by the caller.
if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" { func (cli *Instance) promptMnemonic() (*memguard.LockedBuffer, func(), error) {
if cli.Mnemonic != nil {
secret.Debug("Using mnemonic from environment variable") secret.Debug("Using mnemonic from environment variable")
return envMnemonic, func() {}, nil return cli.Mnemonic, func() {}, nil
} }
secret.Debug("Prompting user for mnemonic phrase") secret.Debug("Prompting user for mnemonic phrase")
@@ -58,23 +62,23 @@ func promptMnemonic() (string, func(), error) {
if err != nil { if err != nil {
secret.Debug("Failed to read mnemonic from stdin", "error", err) secret.Debug("Failed to read mnemonic from stdin", "error", err)
return "", nil, fmt.Errorf("failed to read mnemonic: %w", err) return nil, nil, fmt.Errorf("failed to read mnemonic: %w", err)
} }
fmt.Fprintln(os.Stderr) // Add newline after hidden input fmt.Fprintln(os.Stderr) // Add newline after hidden input
return mnemonicBuffer.String(), mnemonicBuffer.Destroy, nil return mnemonicBuffer, mnemonicBuffer.Destroy, nil
} }
// setupDefaultVault creates the default vault and derives its long-term // setupDefaultVault creates the default vault and derives its long-term
// identity from the mnemonic // identity from the mnemonic
func (cli *Instance) setupDefaultVault( func (cli *Instance) setupDefaultVault(
stateDir, mnemonicStr string, stateDir string, mnemonic *memguard.LockedBuffer,
) (*vault.Vault, *age.X25519Identity, error) { ) (*vault.Vault, *age.X25519Identity, error) {
// Create the default vault - it will handle key derivation internally // Create the default vault - it will handle key derivation internally
secret.Debug("Creating default vault") secret.Debug("Creating default vault")
vlt, err := vault.CreateVault(cli.fs, cli.stateDir, "default") vlt, err := vault.CreateVault(cli.fs, cli.stateDir, "default", mnemonic)
if err != nil { if err != nil {
secret.Debug("Failed to create default vault", "error", err) secret.Debug("Failed to create default vault", "error", err)
@@ -92,7 +96,7 @@ func (cli *Instance) setupDefaultVault(
} }
// Derive the long-term key using the same index that CreateVault used // Derive the long-term key using the same index that CreateVault used
ltIdentity, err := agehd.DeriveIdentity(mnemonicStr, metadata.DerivationIndex) ltIdentity, err := agehd.DeriveIdentity(mnemonic.String(), metadata.DerivationIndex)
if err != nil { if err != nil {
secret.Debug("Failed to derive long-term key", "error", err) secret.Debug("Failed to derive long-term key", "error", err)
@@ -136,12 +140,13 @@ func (cli *Instance) initialize(cmd *cobra.Command) error {
} }
// Prompt for mnemonic // Prompt for mnemonic
mnemonicStr, cleanupMnemonic, err := promptMnemonic() mnemonic, cleanupMnemonic, err := cli.promptMnemonic()
if err != nil { if err != nil {
return err return err
} }
defer cleanupMnemonic() defer cleanupMnemonic()
mnemonicStr := mnemonic.String()
if mnemonicStr == "" { if mnemonicStr == "" {
secret.Debug("Empty mnemonic provided") secret.Debug("Empty mnemonic provided")
@@ -162,18 +167,14 @@ func (cli *Instance) initialize(cmd *cobra.Command) error {
// Ask for the unlocker passphrase before creating the vault, so that // Ask for the unlocker passphrase before creating the vault, so that
// stopping at the prompt leaves no vault without an unlocker behind // stopping at the prompt leaves no vault without an unlocker behind
passphraseBuffer, err := resolvePassphrase() passphraseBuffer, cleanupPassphrase, err := cli.resolvePassphrase()
if err != nil { if err != nil {
return err return err
} }
defer passphraseBuffer.Destroy() defer cleanupPassphrase()
// Set mnemonic in environment for CreateVault to use
restoreMnemonicEnv := setMnemonicEnv(mnemonicStr)
defer restoreMnemonicEnv()
// Create the default vault and derive its long-term key // Create the default vault and derive its long-term key
vlt, ltIdentity, err := cli.setupDefaultVault(stateDir, mnemonicStr) vlt, ltIdentity, err := cli.setupDefaultVault(stateDir, mnemonic)
if err != nil { if err != nil {
return err return err
} }
+14 -7
View File
@@ -286,7 +286,7 @@ func TestSecretManagerIntegration(t *testing.T) {
// Test 25: Concurrent operations // Test 25: Concurrent operations
// Purpose: Test multiple simultaneous operations // Purpose: Test multiple simultaneous operations
// Expected: Proper locking/synchronization, no corruption // Expected: Proper locking/synchronization, no corruption
test25ConcurrentOperations(t, testMnemonic, runSecret, runSecretWithEnv) test25ConcurrentOperations(t, tempDir, secretPath, testMnemonic, runSecret)
// Test 26: Large secret values // Test 26: Large secret values
// Purpose: Test with large secret values (e.g., certificates) // Purpose: Test with large secret values (e.g., certificates)
@@ -2009,28 +2009,35 @@ func test24EnvironmentVariables(t *testing.T, tempDir, secretPath, testMnemonic,
assert.Equal(t, "env-test-value", strings.TrimSpace(string(cmdOutput2))) assert.Equal(t, "env-test-value", strings.TrimSpace(string(cmdOutput2)))
} }
func test25ConcurrentOperations(t *testing.T, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { func test25ConcurrentOperations(t *testing.T, tempDir, secretPath, testMnemonic string, runSecret func(...string) (string, error)) {
t.Helper() t.Helper()
// Make sure we're in default vault // Make sure we're in default vault
_, err := runSecret("vault", "select", "default") _, err := runSecret("vault", "select", "default")
require.NoError(t, err, "vault select should succeed") require.NoError(t, err, "vault select should succeed")
// Run multiple concurrent reads // Run multiple concurrent reads, as separate processes: within one
// process the first command to read the mnemonic would unset it for
// the others
const numReaders = 5 const numReaders = 5
errCh := make(chan error, numReaders) errCh := make(chan error, numReaders)
for i := range numReaders { for i := range numReaders {
go func(id int) { go func(id int) {
output, err := runSecretWithEnv(map[string]string{ cmd := exec.CommandContext(t.Context(), secretPath, "get", "database/password")
secret.EnvMnemonic: testMnemonic, cmd.Env = []string{
}, "get", "database/password") secret.EnvStateDir + "=" + tempDir,
secret.EnvMnemonic + "=" + testMnemonic,
"PATH=" + os.Getenv("PATH"),
"HOME=" + os.Getenv("HOME"),
}
output, err := cmd.Output()
switch { switch {
case err != nil: case err != nil:
errCh <- fmt.Errorf("reader %d failed: %w", id, err) errCh <- fmt.Errorf("reader %d failed: %w", id, err)
case strings.TrimSpace(output) == "": case strings.TrimSpace(string(output)) == "":
errCh <- fmt.Errorf("%w: reader %d", errEmptyValue, id) errCh <- fmt.Errorf("%w: reader %d", errEmptyValue, id)
default: default:
errCh <- nil errCh <- nil
+32 -20
View File
@@ -52,15 +52,18 @@ func lockInBackground(t *testing.T, fs afero.Fs) <-chan func() {
} }
// addAtOnce runs one add of the secret name per value, all at once, and // addAtOnce runs one add of the secret name per value, all at once, and
// returns their errors. // returns their errors. Each add is given mnemonic, which a forced add
// needs.
func addAtOnce( func addAtOnce(
fs afero.Fs, stateDir, name string, force bool, values []string, fs afero.Fs, stateDir, name string, force bool, values []string,
mnemonic *memguard.LockedBuffer,
) []error { ) []error {
errs := make(chan error, len(values)) errs := make(chan error, len(values))
for _, value := range values { for _, value := range values {
go func() { go func() {
cli := NewCLIInstanceWithStateDir(fs, stateDir) cli := NewCLIInstanceWithStateDir(fs, stateDir)
cli.Mnemonic = mnemonic
cli.cmd = &cobra.Command{} cli.cmd = &cobra.Command{}
cli.cmd.SetIn(strings.NewReader(value)) cli.cmd.SetIn(strings.NewReader(value))
@@ -92,9 +95,9 @@ func numbered(prefix string, count int) []string {
// forced adds read the same highest version number and overwrite each // forced adds read the same highest version number and overwrite each
// other's version. With it they behave as if run one after another. // other's version. With it they behave as if run one after another.
// //
//nolint:paralleltest // t.Setenv forbids parallel subtests //nolint:paralleltest // times commands against the in-memory lock all tests share
func TestConcurrentAddsKeepEveryVersion(t *testing.T) { func TestConcurrentAddsKeepEveryVersion(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) mnemonic := testMnemonicBuffer(t)
const adds = 8 const adds = 8
@@ -107,14 +110,14 @@ func TestConcurrentAddsKeepEveryVersion(t *testing.T) {
{"real", afero.NewOsFs(), t.TempDir()}, {"real", afero.NewOsFs(), t.TempDir()},
} { } {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
_, err := vault.CreateVault(tc.fs, tc.stateDir, "default") _, err := vault.CreateVault(tc.fs, tc.stateDir, "default", mnemonic)
require.NoError(t, err) require.NoError(t, err)
// One add creates the secret; the others find that it exists // One add creates the secret; the others find that it exists
created := 0 created := 0
for _, err := range addAtOnce(tc.fs, tc.stateDir, "shared", false, for _, err := range addAtOnce(tc.fs, tc.stateDir, "shared", false,
numbered("create", adds)) { numbered("create", adds), mnemonic) {
if err == nil { if err == nil {
created++ created++
} else { } else {
@@ -126,13 +129,15 @@ func TestConcurrentAddsKeepEveryVersion(t *testing.T) {
// Every forced add stores a version of its own // Every forced add stores a version of its own
for _, err := range addAtOnce(tc.fs, tc.stateDir, "shared", true, for _, err := range addAtOnce(tc.fs, tc.stateDir, "shared", true,
numbered("force", adds)) { numbered("force", adds), mnemonic) {
require.NoError(t, err) require.NoError(t, err)
} }
vlt, err := vault.GetCurrentVault(tc.fs, tc.stateDir) vlt, err := vault.GetCurrentVault(tc.fs, tc.stateDir)
require.NoError(t, err) require.NoError(t, err)
vlt.Mnemonic = mnemonic
vaultDir, err := vlt.GetDirectory() vaultDir, err := vlt.GetDirectory()
require.NoError(t, err) require.NoError(t, err)
@@ -176,11 +181,11 @@ func (r *readNotifier) Read(p []byte) (int, error) {
// taken the state directory lock before reading, it would hold the lock // taken the state directory lock before reading, it would hold the lock
// while waiting for encrypt's output, and encrypt would wait for the lock // while waiting for encrypt's output, and encrypt would wait for the lock
// to store its key: neither would finish. // to store its key: neither would finish.
//
//nolint:paralleltest // times commands against the in-memory lock all tests share
func TestEncryptPipedIntoAdd(t *testing.T) { func TestEncryptPipedIntoAdd(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic)
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
_, err := vault.CreateVault(fs, testStateDir, "default") _, err := vault.CreateVault(fs, testStateDir, "default", testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, afero.WriteFile(fs, testInput, []byte("piped"), 0o600)) require.NoError(t, afero.WriteFile(fs, testInput, []byte("piped"), 0o600))
@@ -283,14 +288,16 @@ func setupEveryCommand(
) (string, string) { ) (string, string) {
t.Helper() t.Helper()
other, err := vault.CreateVault(fs, testStateDir, "other") mnemonic := testMnemonicBuffer(t)
other, err := vault.CreateVault(fs, testStateDir, "other", mnemonic)
require.NoError(t, err) require.NoError(t, err)
otherDir, err := other.GetDirectory() otherDir, err := other.GetDirectory()
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, fs.Remove(filepath.Join(otherDir, "pub.age"))) require.NoError(t, fs.Remove(filepath.Join(otherDir, "pub.age")))
vlt, err := vault.CreateVault(fs, testStateDir, "work") vlt, err := vault.CreateVault(fs, testStateDir, "work", mnemonic)
require.NoError(t, err) require.NoError(t, err)
addTestSecret(t, vlt, []byte("older"), false) addTestSecret(t, vlt, []byte("older"), false)
@@ -363,7 +370,12 @@ func requireWaitsForLock(
release = sync.OnceFunc(release) release = sync.OnceFunc(release)
defer release() defer release()
unlockPassphrase := memguard.NewBufferFromBytes([]byte(testPassphrase))
defer unlockPassphrase.Destroy()
cli := NewCLIInstanceWithStateDir(fs, testStateDir) cli := NewCLIInstanceWithStateDir(fs, testStateDir)
cli.Mnemonic = testMnemonicBuffer(t)
cli.UnlockPassphrase = unlockPassphrase
cli.cmd = &cobra.Command{} cli.cmd = &cobra.Command{}
cli.cmd.SetIn(strings.NewReader("value")) cli.cmd.SetIn(strings.NewReader("value"))
cli.cmd.SetOut(io.Discard) cli.cmd.SetOut(io.Discard)
@@ -400,11 +412,8 @@ func requireWaitsForLock(
// TestChangingCommandsWaitForLock checks that each command that changes the // TestChangingCommandsWaitForLock checks that each command that changes the
// state directory waits for its lock. // state directory waits for its lock.
// //
//nolint:paralleltest // t.Setenv forbids parallel subtests //nolint:paralleltest // waitingForLock sees any test's command waiting for the lock
func TestChangingCommandsWaitForLock(t *testing.T) { func TestChangingCommandsWaitForLock(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic)
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
for _, tc := range []struct { for _, tc := range []struct {
name string name string
withUnlocker bool withUnlocker bool
@@ -468,15 +477,18 @@ func TestChangingCommandsWaitForLock(t *testing.T) {
// TestEncryptWithExistingKeyTakesNoLock checks that secret encrypt with a // TestEncryptWithExistingKeyTakesNoLock checks that secret encrypt with a
// key that already exists, which only reads the state directory, finishes // key that already exists, which only reads the state directory, finishes
// while another command holds the state directory lock. // while another command holds the state directory lock.
//
//nolint:paralleltest // times commands against the in-memory lock all tests share
func TestEncryptWithExistingKeyTakesNoLock(t *testing.T) { func TestEncryptWithExistingKeyTakesNoLock(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) mnemonic := testMnemonicBuffer(t)
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
_, err := vault.CreateVault(fs, testStateDir, "default") _, err := vault.CreateVault(fs, testStateDir, "default", mnemonic)
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, afero.WriteFile(fs, testInput, []byte("input"), 0o600)) require.NoError(t, afero.WriteFile(fs, testInput, []byte("input"), 0o600))
encrypt := NewCLIInstanceWithStateDir(fs, testStateDir) encrypt := NewCLIInstanceWithStateDir(fs, testStateDir)
encrypt.Mnemonic = mnemonic
encrypt.cmd = &cobra.Command{} encrypt.cmd = &cobra.Command{}
encrypt.cmd.SetOut(io.Discard) encrypt.cmd.SetOut(io.Discard)
@@ -505,11 +517,11 @@ func TestEncryptWithExistingKeyTakesNoLock(t *testing.T) {
// state directory lock by the time it writes its output. Holding it while // state directory lock by the time it writes its output. Holding it while
// streaming would stall every other changing command for as long as the // streaming would stall every other changing command for as long as the
// stream lasts, and forever when the other end of the pipe is one of them. // stream lasts, and forever when the other end of the pipe is one of them.
//
//nolint:paralleltest // times commands against the in-memory lock all tests share
func TestEncryptStreamsUnlocked(t *testing.T) { func TestEncryptStreamsUnlocked(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic)
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
_, err := vault.CreateVault(fs, testStateDir, "default") _, err := vault.CreateVault(fs, testStateDir, "default", testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, afero.WriteFile(fs, testInput, []byte("streamed"), 0o600)) require.NoError(t, afero.WriteFile(fs, testInput, []byte("streamed"), 0o600))
+13 -12
View File
@@ -6,7 +6,6 @@ import (
"testing" "testing"
"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/vault" "git.eeqj.de/sneak/secret/internal/vault"
"github.com/awnumar/memguard" "github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
@@ -20,9 +19,9 @@ import (
// move within "work" left "work" the current vault. "default" is the current // move within "work" left "work" the current vault. "default" is the current
// vault in every case, and each case runs on its own copy of the state // vault in every case, and each case runs on its own copy of the state
// directory. // directory.
//
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
func TestRejectedMoveWithinVaultLeavesStateUnchanged(t *testing.T) { func TestRejectedMoveWithinVaultLeavesStateUnchanged(t *testing.T) {
t.Parallel()
before := snapshotStateDir(t, newTwoVaultFs(t)) before := snapshotStateDir(t, newTwoVaultFs(t))
require.Equal(t, "default", before[testStateDir+"/currentvault"]) require.Equal(t, "default", before[testStateDir+"/currentvault"])
@@ -72,6 +71,8 @@ func TestRejectedMoveWithinVaultLeavesStateUnchanged(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.command, func(t *testing.T) { t.Run(tt.command, func(t *testing.T) {
t.Parallel()
fs := newFsFromSnapshot(t, before) fs := newFsFromSnapshot(t, before)
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir) c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
@@ -86,9 +87,9 @@ func TestRejectedMoveWithinVaultLeavesStateUnchanged(t *testing.T) {
// TestMoveWithinOtherVaultKeepsCurrentVault checks that `secret mv work:x // TestMoveWithinOtherVaultKeepsCurrentVault checks that `secret mv work:x
// work:y`, with "default" the current vault, renames "x" to "y" in "work" and // work:y`, with "default" the current vault, renames "x" to "y" in "work" and
// leaves "default" the current vault. // leaves "default" the current vault.
//
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
func TestMoveWithinOtherVaultKeepsCurrentVault(t *testing.T) { func TestMoveWithinOtherVaultKeepsCurrentVault(t *testing.T) {
t.Parallel()
fs := newTwoVaultFs(t) fs := newTwoVaultFs(t)
c := cli.NewCLIInstanceWithStateDir(fs, testStateDir) c := cli.NewCLIInstanceWithStateDir(fs, testStateDir)
@@ -111,10 +112,8 @@ func TestMoveWithinOtherVaultKeepsCurrentVault(t *testing.T) {
// the secret "x", and the secrets.d of "other" is a link to that of // the secret "x", and the secrets.d of "other" is a link to that of
// "default", so other:x is default:x. Each move must be rejected and leave // "default", so other:x is default:x. Each move must be rejected and leave
// the secret and the links as they were. // the secret and the links as they were.
//
//nolint:paralleltest // t.Setenv
func TestMoveOntoSameSecretUnderAnotherNameIsRejected(t *testing.T) { func TestMoveOntoSameSecretUnderAnotherNameIsRejected(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
const isSame = "is the same secret on this filesystem" const isSame = "is the same secret on this filesystem"
@@ -150,15 +149,17 @@ func TestMoveOntoSameSecretUnderAnotherNameIsRejected(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.command, func(t *testing.T) { t.Run(tt.command, func(t *testing.T) {
t.Parallel()
fs := afero.NewOsFs() fs := afero.NewOsFs()
stateDir := t.TempDir() stateDir := t.TempDir()
vaultsDir := filepath.Join(stateDir, "vaults.d") vaultsDir := filepath.Join(stateDir, "vaults.d")
// "default" is created last, so it is the current vault. // "default" is created last, so it is the current vault.
_, err := vault.CreateVault(fs, stateDir, "other") _, err := vault.CreateVault(fs, stateDir, "other", testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
vlt, err := vault.CreateVault(fs, stateDir, "default") vlt, err := vault.CreateVault(fs, stateDir, "default", testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("value")), false) err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("value")), false)
@@ -199,12 +200,12 @@ func TestMoveOntoSameSecretUnderAnotherNameIsRejected(t *testing.T) {
// and "foo" are two secrets, `secret mv --force Foo foo` still replaces "foo" // and "foo" are two secrets, `secret mv --force Foo foo` still replaces "foo"
// with "Foo". // with "Foo".
func TestForcedCaseOnlyMoveOnCaseSensitiveFilesystem(t *testing.T) { func TestForcedCaseOnlyMoveOnCaseSensitiveFilesystem(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
fs := afero.NewOsFs() fs := afero.NewOsFs()
stateDir := t.TempDir() stateDir := t.TempDir()
vlt, err := vault.CreateVault(fs, stateDir, "default") vlt, err := vault.CreateVault(fs, stateDir, "default", testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
err = vlt.AddSecret("Foo", memguard.NewBufferFromBytes([]byte("upper")), false) err = vlt.AddSecret("Foo", memguard.NewBufferFromBytes([]byte("upper")), false)
+34 -15
View File
@@ -33,6 +33,17 @@ const (
missingFile = "/no/such/file" missingFile = "/no/such/file"
) )
// testMnemonicBuffer returns testMnemonic in a locked buffer that is
// destroyed when the test ends.
func testMnemonicBuffer(t *testing.T) *memguard.LockedBuffer {
t.Helper()
mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonic))
t.Cleanup(mnemonic.Destroy)
return mnemonic
}
// The state directory newTwoVaultFs copies, recorded by snapshotStateDir. // The state directory newTwoVaultFs copies, recorded by snapshotStateDir.
// Creating a passphrase unlocker is slow by design, so the vaults are made // Creating a passphrase unlocker is slow by design, so the vaults are made
// once, by the first test that needs them. // once, by the first test that needs them.
@@ -52,13 +63,12 @@ var (
func newTwoVaultFs(t *testing.T) afero.Fs { func newTwoVaultFs(t *testing.T) afero.Fs {
t.Helper() t.Helper()
t.Setenv(secret.EnvMnemonic, testMnemonic)
twoVaultsOnce.Do(func() { twoVaultsOnce.Do(func() {
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
mnemonic := testMnemonicBuffer(t)
for _, name := range []string{"work", "default"} { for _, name := range []string{"work", "default"} {
vlt, err := vault.CreateVault(fs, testStateDir, name) vlt, err := vault.CreateVault(fs, testStateDir, name, mnemonic)
require.NoError(t, err) require.NoError(t, err)
err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("value")), false) err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("value")), false)
@@ -163,7 +173,7 @@ func requireRejectedAndUnchanged(
// Moves and imports use --force, so that only the name check stands in // Moves and imports use --force, so that only the name check stands in
// the way. // the way.
// //
//nolint:paralleltest // newTwoVaultFs uses t.Setenv //nolint:paralleltest // the cases share cmd
func TestInvalidSecretNameLeavesVaultsUnchanged(t *testing.T) { func TestInvalidSecretNameLeavesVaultsUnchanged(t *testing.T) {
// Creating a passphrase unlocker is slow by design, so the vaults are // Creating a passphrase unlocker is slow by design, so the vaults are
// created once and each case runs on its own copy of them. // created once and each case runs on its own copy of them.
@@ -259,7 +269,7 @@ func TestInvalidSecretNameLeavesVaultsUnchanged(t *testing.T) {
// `secret version rm x ""` every version of x. A version argument is // `secret version rm x ""` every version of x. A version argument is
// accepted only if it is one of the versions `secret version list` lists. // accepted only if it is one of the versions `secret version list` lists.
// //
//nolint:paralleltest // newTwoVaultFs uses t.Setenv //nolint:paralleltest // the cases share cmd
func TestInvalidVersionLeavesVaultsUnchanged(t *testing.T) { func TestInvalidVersionLeavesVaultsUnchanged(t *testing.T) {
before := snapshotStateDir(t, newTwoVaultFs(t)) before := snapshotStateDir(t, newTwoVaultFs(t))
@@ -297,15 +307,17 @@ func TestInvalidVersionLeavesVaultsUnchanged(t *testing.T) {
// `secret vault import ..` wrote a long-term key and an unlocker into the // `secret vault import ..` wrote a long-term key and an unlocker into the
// state directory itself, and `secret vault select ..` made it the current // state directory itself, and `secret vault select ..` made it the current
// vault. Each command that takes a vault name must reject an invalid one // vault. Each command that takes a vault name must reject an invalid one
// before building a path from it. The mnemonic and the passphrase are set, // before building a path from it. The instance is given the mnemonic and
// and moves and removals use --force, so that only the name check stands // the passphrase, and moves and removals use --force, so that only the name
// in the way. // check stands in the way.
// //
//nolint:paralleltest // newTwoVaultFs uses t.Setenv //nolint:paralleltest // the cases share cmd
func TestInvalidVaultNameLeavesStateUnchanged(t *testing.T) { func TestInvalidVaultNameLeavesStateUnchanged(t *testing.T) {
before := snapshotStateDir(t, newTwoVaultFs(t)) before := snapshotStateDir(t, newTwoVaultFs(t))
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase) mnemonic := testMnemonicBuffer(t)
passphrase := memguard.NewBufferFromBytes([]byte(testPassphrase))
t.Cleanup(passphrase.Destroy)
cmd := &cobra.Command{} cmd := &cobra.Command{}
@@ -338,7 +350,12 @@ func TestInvalidVaultNameLeavesStateUnchanged(t *testing.T) {
for _, name := range []string{"", ".", "..", "a/b"} { for _, name := range []string{"", ".", "..", "a/b"} {
t.Run(fmt.Sprintf(tt.command, name), func(t *testing.T) { t.Run(fmt.Sprintf(tt.command, name), func(t *testing.T) {
requireRejectedAndUnchanged(t, before, vault.ValidateVaultName(name), requireRejectedAndUnchanged(t, before, vault.ValidateVaultName(name),
func(c *cli.Instance) error { return tt.run(c, name) }) func(c *cli.Instance) error {
c.Mnemonic = mnemonic
c.UnlockPassphrase = passphrase
return tt.run(c, name)
})
}) })
} }
} }
@@ -347,14 +364,16 @@ func TestInvalidVaultNameLeavesStateUnchanged(t *testing.T) {
// TestRemoveVersionRemovesOnlyThatVersion checks that `secret version rm` // TestRemoveVersionRemovesOnlyThatVersion checks that `secret version rm`
// with a version that is not the current one removes that version and // with a version that is not the current one removes that version and
// changes nothing else. // changes nothing else.
//
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
func TestRemoveVersionRemovesOnlyThatVersion(t *testing.T) { func TestRemoveVersionRemovesOnlyThatVersion(t *testing.T) {
t.Parallel()
fs := newTwoVaultFs(t) fs := newTwoVaultFs(t)
vlt, err := vault.GetCurrentVault(fs, testStateDir) vlt, err := vault.GetCurrentVault(fs, testStateDir)
require.NoError(t, err) require.NoError(t, err)
vlt.Mnemonic = testMnemonicBuffer(t)
// A second version of "x" becomes the current one. // A second version of "x" becomes the current one.
err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("new")), true) err = vlt.AddSecret("x", memguard.NewBufferFromBytes([]byte("new")), true)
require.NoError(t, err) require.NoError(t, err)
@@ -388,9 +407,9 @@ func TestRemoveVersionRemovesOnlyThatVersion(t *testing.T) {
// TestMoveToVaultNameRenamesInCurrentVault checks that `secret mv x work`, // TestMoveToVaultNameRenamesInCurrentVault checks that `secret mv x work`,
// where "work" is also the name of a vault, renames the secret "x" to "work" // where "work" is also the name of a vault, renames the secret "x" to "work"
// in the current vault and changes nothing else. // in the current vault and changes nothing else.
//
//nolint:paralleltest // newTwoVaultFs uses t.Setenv
func TestMoveToVaultNameRenamesInCurrentVault(t *testing.T) { func TestMoveToVaultNameRenamesInCurrentVault(t *testing.T) {
t.Parallel()
before := snapshotStateDir(t, newTwoVaultFs(t)) before := snapshotStateDir(t, newTwoVaultFs(t))
fs := newFsFromSnapshot(t, before) fs := newFsFromSnapshot(t, before)
+24
View File
@@ -81,6 +81,9 @@ func newAddCmd() *cobra.Command {
cli.cmd = cmd // Set the command for stdin access cli.cmd = cmd // Set the command for stdin access
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
secret.Debug("Created CLI instance, calling AddSecret") secret.Debug("Created CLI instance, calling AddSecret")
return cli.AddSecret(args[0], force) return cli.AddSecret(args[0], force)
@@ -111,6 +114,9 @@ func newGetCmd() *cobra.Command {
return fmt.Errorf("failed to initialize CLI: %w", err) return fmt.Errorf("failed to initialize CLI: %w", err)
} }
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
// Without --version, get the current version. A given // Without --version, get the current version. A given
// --version is checked as typed, so an empty one is rejected. // --version is checked as typed, so an empty one is rejected.
if !cmd.Flags().Changed("version") { if !cmd.Flags().Changed("version") {
@@ -174,6 +180,9 @@ func newImportCmd() *cobra.Command {
return fmt.Errorf("failed to initialize CLI: %w", err) return fmt.Errorf("failed to initialize CLI: %w", err)
} }
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
return cli.ImportSecret(cmd, args[0], sourceFile, force) return cli.ImportSecret(cmd, args[0], sourceFile, force)
}, },
} }
@@ -248,6 +257,9 @@ The source secret is deleted after successful copy.`,
return fmt.Errorf("failed to initialize CLI: %w", err) return fmt.Errorf("failed to initialize CLI: %w", err)
} }
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
return cli.MoveSecret(cmd, args[0], args[1], force) return cli.MoveSecret(cmd, args[0], args[1], force)
}, },
} }
@@ -354,6 +366,8 @@ func (cli *Instance) AddSecret(secretName string, force bool) error {
return err return err
} }
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
secret.Debug("Got current vault", "vault_name", vlt.GetName()) secret.Debug("Got current vault", "vault_name", vlt.GetName())
// Read secret value directly into protected buffers // Read secret value directly into protected buffers
@@ -420,6 +434,8 @@ func (cli *Instance) GetSecret(cmd *cobra.Command, secretName string) error {
return err return err
} }
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
value, err := vlt.GetSecret(secretName) value, err := vlt.GetSecret(secretName)
if err != nil { if err != nil {
return err return err
@@ -448,6 +464,8 @@ func (cli *Instance) GetSecretWithVersion(
return err return err
} }
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
// Get the secret value // Get the secret value
value, err := vlt.GetSecretVersion(secretName, version) value, err := vlt.GetSecretVersion(secretName, version)
if err != nil { if err != nil {
@@ -633,6 +651,8 @@ func (cli *Instance) ImportSecret(
return err return err
} }
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
// Read secret value from the source file into protected buffers // Read secret value from the source file into protected buffers
file, err := cli.fs.Open(sourceFile) file, err := cli.fs.Open(sourceFile)
if err != nil { if err != nil {
@@ -993,6 +1013,10 @@ func (cli *Instance) moveSecretCrossVault(
destVault.Name, destSecretName) destVault.Name, destSecretName)
} }
// Copying needs the long-term keys of both vaults
srcVault.Mnemonic, srcVault.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
destVault.Mnemonic, destVault.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
// Unlock destination vault (will fail if neither mnemonic nor unlocker available) // Unlock destination vault (will fail if neither mnemonic nor unlocker available)
_, err = destVault.GetOrDeriveLongTermKey() _, err = destVault.GetOrDeriveLongTermKey()
if err != nil { if err != nil {
+6 -10
View File
@@ -10,7 +10,6 @@ import (
"strings" "strings"
"testing" "testing"
"git.eeqj.de/sneak/secret/internal/secret"
"git.eeqj.de/sneak/secret/internal/vault" "git.eeqj.de/sneak/secret/internal/vault"
"git.eeqj.de/sneak/secret/pkg/agehd" "git.eeqj.de/sneak/secret/pkg/agehd"
"github.com/spf13/afero" "github.com/spf13/afero"
@@ -71,11 +70,8 @@ func newSizeTestVault(t *testing.T) (afero.Fs, *vault.Vault) {
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Set test mnemonic
t.Setenv(secret.EnvMnemonic, testMnemonic)
// Create vault // Create vault
_, err := vault.CreateVault(fs, testStateDir, testVaultName) _, err := vault.CreateVault(fs, testStateDir, testVaultName, testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
// Set current vault // Set current vault
@@ -205,7 +201,7 @@ func runImportSecretSizeCase(t *testing.T, size int, wantErr bool, errMsg string
// TestAddSecretVariousSizes tests adding secrets of various sizes through stdin // TestAddSecretVariousSizes tests adding secrets of various sizes through stdin
// //
//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault //nolint:paralleltest // together the subtests lock more than the memlock limit
func TestAddSecretVariousSizes(t *testing.T) { func TestAddSecretVariousSizes(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
@@ -265,7 +261,7 @@ func TestAddSecretVariousSizes(t *testing.T) {
// TestImportSecretVariousSizes tests importing secrets of various sizes from files // TestImportSecretVariousSizes tests importing secrets of various sizes from files
// //
//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault //nolint:paralleltest // together the subtests lock more than the memlock limit
func TestImportSecretVariousSizes(t *testing.T) { func TestImportSecretVariousSizes(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
@@ -325,7 +321,7 @@ func TestImportSecretVariousSizes(t *testing.T) {
// TestAddSecretBufferGrowth tests that our buffer growth strategy works correctly // TestAddSecretBufferGrowth tests that our buffer growth strategy works correctly
// //
//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault //nolint:paralleltest // together the subtests lock more than the memlock limit
func TestAddSecretBufferGrowth(t *testing.T) { func TestAddSecretBufferGrowth(t *testing.T) {
// Test various sizes that should trigger buffer growth // Test various sizes that should trigger buffer growth
sizes := []int{ sizes := []int{
@@ -392,9 +388,9 @@ func TestAddSecretBufferGrowth(t *testing.T) {
} }
// TestAddSecretStreamingBehavior tests that we handle streaming input correctly // TestAddSecretStreamingBehavior tests that we handle streaming input correctly
//
//nolint:paralleltest // uses t.Setenv via newSizeTestVault
func TestAddSecretStreamingBehavior(t *testing.T) { func TestAddSecretStreamingBehavior(t *testing.T) {
t.Parallel()
fs, vlt := newSizeTestVault(t) fs, vlt := newSizeTestVault(t)
// Create a custom reader that simulates slow streaming input // Create a custom reader that simulates slow streaming input
+15 -11
View File
@@ -16,7 +16,6 @@ import (
"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/internal/vault"
"github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
"github.com/spf13/cobra" "github.com/spf13/cobra"
) )
@@ -230,6 +229,9 @@ func newUnlockerAddCmd() *cobra.Command {
return fmt.Errorf("failed to initialize CLI: %w", err) return fmt.Errorf("failed to initialize CLI: %w", err)
} }
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
unlockerType := args[0] unlockerType := args[0]
// Validate unlocker type // Validate unlocker type
@@ -580,19 +582,19 @@ func (cli *Instance) addPassphraseUnlocker(cmd *cobra.Command) error {
// For passphrase unlockers, we don't need the vault to be unlocked // For passphrase unlockers, we don't need the vault to be unlocked
// The CreatePassphraseUnlocker method will handle getting the // The CreatePassphraseUnlocker method will handle getting the
// long-term key // long-term key
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
// Check if passphrase is set in environment variable // The new unlocker gets the passphrase from the environment, which also
var passphraseBuffer *memguard.LockedBuffer // unlocks the current passphrase unlocker, else the one entered here
if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" { passphraseBuffer := cli.UnlockPassphrase
passphraseBuffer = memguard.NewBufferFromBytes([]byte(envPassphrase)) if passphraseBuffer == nil {
} else {
// Use secure passphrase input with confirmation // Use secure passphrase input with confirmation
passphraseBuffer, err = readSecurePassphrase("Enter passphrase for unlocker: ") passphraseBuffer, err = readSecurePassphrase("Enter passphrase for unlocker: ")
if err != nil { if err != nil {
return fmt.Errorf("failed to read passphrase: %w", err) return fmt.Errorf("failed to read passphrase: %w", err)
} }
defer passphraseBuffer.Destroy()
} }
defer passphraseBuffer.Destroy()
passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer) passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
if err != nil { if err != nil {
@@ -613,7 +615,8 @@ func (cli *Instance) addKeychainUnlocker(cmd *cobra.Command) error {
return errKeychainMacOSOnly return errKeychainMacOSOnly
} }
keychainUnlocker, err := secret.CreateKeychainUnlocker(cli.fs, cli.stateDir) keychainUnlocker, err := secret.CreateKeychainUnlocker(
cli.fs, cli.stateDir, cli.Mnemonic, cli.UnlockPassphrase)
if err != nil { if err != nil {
return fmt.Errorf("failed to create macOS Keychain unlocker: %w", err) return fmt.Errorf("failed to create macOS Keychain unlocker: %w", err)
} }
@@ -643,7 +646,8 @@ func (cli *Instance) addSecureEnclaveUnlocker(cmd *cobra.Command) error {
return errSecureEnclaveMacOSOnly return errSecureEnclaveMacOSOnly
} }
seUnlocker, err := secret.CreateSecureEnclaveUnlocker(cli.fs, cli.stateDir) seUnlocker, err := secret.CreateSecureEnclaveUnlocker(
cli.fs, cli.stateDir, cli.Mnemonic, cli.UnlockPassphrase)
if err != nil { if err != nil {
return fmt.Errorf("failed to create Secure Enclave unlocker: %w", err) return fmt.Errorf("failed to create Secure Enclave unlocker: %w", err)
} }
@@ -707,8 +711,8 @@ func (cli *Instance) addPGPUnlocker(cmd *cobra.Command) error {
return fmt.Errorf("GPG key %s %w", gpgKeyID, errGPGKeyAlreadyUnlocker) return fmt.Errorf("GPG key %s %w", gpgKeyID, errGPGKeyAlreadyUnlocker)
} }
pgpUnlocker, err := secret.CreatePGPUnlocker( pgpUnlocker, err := secret.CreatePGPUnlocker(cli.fs, cli.stateDir,
cli.fs, cli.stateDir, gpgKeyID, fingerprint) gpgKeyID, fingerprint, cli.Mnemonic, cli.UnlockPassphrase)
if err != nil { if err != nil {
return err return err
} }
+16 -16
View File
@@ -5,7 +5,6 @@ import (
"path/filepath" "path/filepath"
"testing" "testing"
"git.eeqj.de/sneak/secret/internal/secret"
"git.eeqj.de/sneak/secret/internal/vault" "git.eeqj.de/sneak/secret/internal/vault"
"github.com/awnumar/memguard" "github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
@@ -24,28 +23,31 @@ const (
// TestAddPGPUnlocker adds a PGP unlocker for a throwaway GPG key to a vault // TestAddPGPUnlocker adds a PGP unlocker for a throwaway GPG key to a vault
// with a passphrase unlocker, getting the vault's long-term key from the // with a passphrase unlocker, getting the vault's long-term key from the
// mnemonic or, with the mnemonic unset, from the passphrase unlocker. It // mnemonic or, with no mnemonic given, from the passphrase unlocker. It
// then reads a secret with neither the mnemonic nor the passphrase set, so // then reads a secret with neither the mnemonic nor the passphrase given, so
// through the new unlocker, which the add selects. // through the new unlocker, which the add selects.
//
//nolint:paralleltest // t.Setenv (GNUPGHOME) forbids parallel tests
func TestAddPGPUnlocker(t *testing.T) { func TestAddPGPUnlocker(t *testing.T) {
newTestGPGKey(t) newTestGPGKey(t)
passphrase := memguard.NewBufferFromBytes([]byte(testPassphrase))
t.Cleanup(passphrase.Destroy)
tests := []struct { tests := []struct {
name string name string
// mnemonic is the mnemonic set while the unlocker is added. // mnemonic is the mnemonic given while the unlocker is added, or nil.
mnemonic string mnemonic *memguard.LockedBuffer
}{ }{
{"long-term key from the mnemonic", testMnemonic}, {"long-term key from the mnemonic", testMnemonicBuffer(t)},
{"long-term key from the current unlocker", ""}, {"long-term key from the current unlocker", nil},
} }
for _, test := range tests { for _, test := range tests {
t.Run(test.name, func(t *testing.T) { t.Run(test.name, func(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic)
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
vlt, err := vault.CreateVault(fs, listTestStateDir, listTestVaultName) vlt, err := vault.CreateVault(fs, listTestStateDir, listTestVaultName,
testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
err = vlt.AddSecret(addTestSecretName, err = vlt.AddSecret(addTestSecretName,
@@ -56,15 +58,13 @@ func TestAddPGPUnlocker(t *testing.T) {
memguard.NewBufferFromBytes([]byte(testPassphrase))) memguard.NewBufferFromBytes([]byte(testPassphrase)))
require.NoError(t, err) require.NoError(t, err)
t.Setenv(secret.EnvMnemonic, test.mnemonic)
instance, cmd := newTestInstance(fs) instance, cmd := newTestInstance(fs)
instance.Mnemonic = test.mnemonic
instance.UnlockPassphrase = passphrase
cmd.Flags().String("keyid", unreadableTestGPGUserID, "") cmd.Flags().String("keyid", unreadableTestGPGUserID, "")
require.NoError(t, instance.UnlockersAdd(unlockerTypePGP, cmd)) require.NoError(t, instance.UnlockersAdd(unlockerTypePGP, cmd))
t.Setenv(secret.EnvMnemonic, "")
t.Setenv(secret.EnvUnlockPassphrase, "")
reopened := vault.NewVault(fs, listTestStateDir, listTestVaultName) reopened := vault.NewVault(fs, listTestStateDir, listTestVaultName)
current, err := reopened.GetCurrentUnlocker() current, err := reopened.GetCurrentUnlocker()
+27 -60
View File
@@ -5,7 +5,6 @@ import (
"errors" "errors"
"fmt" "fmt"
"log" "log"
"os"
"path/filepath" "path/filepath"
"slices" "slices"
"strings" "strings"
@@ -85,6 +84,9 @@ func newVaultCreateCmd() *cobra.Command {
return fmt.Errorf("failed to initialize CLI: %w", err) return fmt.Errorf("failed to initialize CLI: %w", err)
} }
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
return cli.CreateVault(cmd, args[0]) return cli.CreateVault(cmd, args[0])
}, },
} }
@@ -136,6 +138,9 @@ func newVaultImportCmd() *cobra.Command {
return fmt.Errorf("failed to initialize CLI: %w", err) return fmt.Errorf("failed to initialize CLI: %w", err)
} }
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
return cli.VaultImport(cmd, vaultName) return cli.VaultImport(cmd, vaultName)
}, },
} }
@@ -228,28 +233,14 @@ func (cli *Instance) ListVaults(cmd *cobra.Command, jsonOutput bool) error {
return nil return nil
} }
// setMnemonicEnv sets the mnemonic environment variable and returns a // resolvePassphrase returns the unlock passphrase from the environment,
// function that restores the previous value // cli.UnlockPassphrase, or prompts the user for it with confirmation. The
func setMnemonicEnv(mnemonicStr string) func() { // returned cleanup function must be deferred by the caller.
originalMnemonic := os.Getenv(secret.EnvMnemonic) func (cli *Instance) resolvePassphrase() (*memguard.LockedBuffer, func(), error) {
_ = os.Setenv(secret.EnvMnemonic, mnemonicStr) if cli.UnlockPassphrase != nil {
return func() {
if originalMnemonic != "" {
_ = os.Setenv(secret.EnvMnemonic, originalMnemonic)
} else {
_ = os.Unsetenv(secret.EnvMnemonic)
}
}
}
// resolvePassphrase returns the unlock passphrase from the environment or
// prompts the user for it with confirmation
func resolvePassphrase() (*memguard.LockedBuffer, error) {
if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" {
secret.Debug("Using unlock passphrase from environment variable") secret.Debug("Using unlock passphrase from environment variable")
return memguard.NewBufferFromBytes([]byte(envPassphrase)), nil return cli.UnlockPassphrase, func() {}, nil
} }
secret.Debug("Prompting user for unlock passphrase") secret.Debug("Prompting user for unlock passphrase")
@@ -257,10 +248,10 @@ func resolvePassphrase() (*memguard.LockedBuffer, error) {
// Use secure passphrase input with confirmation // Use secure passphrase input with confirmation
passphraseBuffer, err := readSecurePassphrase("Enter passphrase for unlocker: ") passphraseBuffer, err := readSecurePassphrase("Enter passphrase for unlocker: ")
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to read passphrase: %w", err) return nil, nil, fmt.Errorf("failed to read passphrase: %w", err)
} }
return passphraseBuffer, nil return passphraseBuffer, passphraseBuffer.Destroy, nil
} }
// CreateVault creates a new vault // CreateVault creates a new vault
@@ -273,30 +264,13 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error {
} }
defer release() defer release()
// Get or prompt for mnemonic mnemonic, cleanupMnemonic, err := cli.promptMnemonic()
var mnemonicStr string if err != nil {
return err
if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" {
secret.Debug("Using mnemonic from environment variable")
mnemonicStr = envMnemonic
} else {
secret.Debug("Prompting user for mnemonic phrase")
// Read mnemonic securely without echo
mnemonicBuffer, err := secret.ReadPassphrase(
"Enter your BIP39 mnemonic phrase: ")
if err != nil {
secret.Debug("Failed to read mnemonic from stdin", "error", err)
return fmt.Errorf("failed to read mnemonic: %w", err)
}
defer mnemonicBuffer.Destroy()
mnemonicStr = mnemonicBuffer.String()
fmt.Fprintln(os.Stderr) // Add newline after hidden input
} }
defer cleanupMnemonic()
mnemonicStr := mnemonic.String()
if mnemonicStr == "" { if mnemonicStr == "" {
return errMnemonicEmpty return errMnemonicEmpty
} }
@@ -311,18 +285,14 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error {
// Ask for the unlocker passphrase before creating the vault, so that // Ask for the unlocker passphrase before creating the vault, so that
// stopping at the prompt leaves no vault without an unlocker behind // stopping at the prompt leaves no vault without an unlocker behind
passphraseBuffer, err := resolvePassphrase() passphraseBuffer, cleanupPassphrase, err := cli.resolvePassphrase()
if err != nil { if err != nil {
return err return err
} }
defer passphraseBuffer.Destroy() defer cleanupPassphrase()
// Set mnemonic in environment for CreateVault to use
restoreMnemonicEnv := setMnemonicEnv(mnemonicStr)
defer restoreMnemonicEnv()
// Create the vault - it will handle key derivation internally // Create the vault - it will handle key derivation internally
vlt, err := vault.CreateVault(cli.fs, cli.stateDir, name) vlt, err := vault.CreateVault(cli.fs, cli.stateDir, name, mnemonic)
if err != nil { if err != nil {
return err return err
} }
@@ -412,11 +382,12 @@ func (cli *Instance) vaultImportPreflight(
} }
// Get mnemonic from environment // Get mnemonic from environment
mnemonic := os.Getenv(secret.EnvMnemonic) if cli.Mnemonic == nil {
if mnemonic == "" {
return "", "", "", errMnemonicEnvNotSet return "", "", "", errMnemonicEnvNotSet
} }
mnemonic := cli.Mnemonic.String()
// Validate the mnemonic // Validate the mnemonic
mnemonicWords := strings.Fields(mnemonic) mnemonicWords := strings.Fields(mnemonic)
secret.Debug("Validating BIP39 mnemonic", "word_count", len(mnemonicWords)) secret.Debug("Validating BIP39 mnemonic", "word_count", len(mnemonicWords))
@@ -539,17 +510,13 @@ func (cli *Instance) importMnemonic(cmd *cobra.Command, vaultName string) error
} }
// Get passphrase from environment variable // Get passphrase from environment variable
passphraseStr := os.Getenv(secret.EnvUnlockPassphrase) passphraseBuffer := cli.UnlockPassphrase
if passphraseStr == "" { if passphraseBuffer == nil {
return errPassphraseEnvNotSet return errPassphraseEnvNotSet
} }
secret.Debug("Using unlock passphrase from environment variable") secret.Debug("Using unlock passphrase from environment variable")
// Create secure buffer for passphrase
passphraseBuffer := memguard.NewBufferFromBytes([]byte(passphraseStr))
defer passphraseBuffer.Destroy()
// Unlock the vault with the derived long-term key // Unlock the vault with the derived long-term key
vlt.Unlock(ltIdentity) vlt.Unlock(ltIdentity)
+5
View File
@@ -54,6 +54,9 @@ func VersionCommands(cli *Instance) *cobra.Command {
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 {
destroySecrets := cli.readSecretEnv()
defer destroySecrets()
return cli.ListVersions(cmd, args[0]) return cli.ListVersions(cmd, args[0])
}, },
} }
@@ -172,6 +175,8 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error {
currentVersion = "" currentVersion = ""
} }
vlt.Mnemonic, vlt.UnlockPassphrase = cli.Mnemonic, cli.UnlockPassphrase
// Get long-term key for decrypting metadata // Get long-term key for decrypting metadata
ltIdentity, err := vlt.GetOrDeriveLongTermKey() ltIdentity, err := vlt.GetOrDeriveLongTermKey()
if err != nil { if err != nil {
+35 -11
View File
@@ -45,6 +45,17 @@ const (
testStateDir = "/test/state" testStateDir = "/test/state"
) )
// testMnemonicBuffer returns testMnemonic in a locked buffer that is
// destroyed when the test ends.
func testMnemonicBuffer(t *testing.T) *memguard.LockedBuffer {
t.Helper()
mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonic))
t.Cleanup(mnemonic.Destroy)
return mnemonic
}
// Helper function to add a version of the "test/secret" secret to the // Helper function to add a version of the "test/secret" secret to the
// vault with proper buffer protection // vault with proper buffer protection
func addTestSecret(t *testing.T, vlt *vault.Vault, value []byte, force bool) { func addTestSecret(t *testing.T, vlt *vault.Vault, value []byte, force bool) {
@@ -61,11 +72,8 @@ func addTestSecret(t *testing.T, vlt *vault.Vault, value []byte, force bool) {
func setupTestVault(t *testing.T, fs afero.Fs) { func setupTestVault(t *testing.T, fs afero.Fs) {
t.Helper() t.Helper()
// Set mnemonic for testing
t.Setenv(secret.EnvMnemonic, testMnemonic)
// Create vault // Create vault
vlt, err := vault.CreateVault(fs, testStateDir, "default") vlt, err := vault.CreateVault(fs, testStateDir, "default", testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
// Derive and store long-term key from mnemonic // Derive and store long-term key from mnemonic
@@ -83,11 +91,13 @@ func setupTestVault(t *testing.T, fs afero.Fs) {
require.NoError(t, err) require.NoError(t, err)
} }
//nolint:paralleltest // uses t.Setenv via setupTestVault
func TestListVersionsCommand(t *testing.T) { func TestListVersionsCommand(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
stateDir := testStateDir stateDir := testStateDir
cli := NewCLIInstanceWithStateDir(fs, stateDir) cli := NewCLIInstanceWithStateDir(fs, stateDir)
cli.Mnemonic = testMnemonicBuffer(t)
// Set up vault with long-term key // Set up vault with long-term key
setupTestVault(t, fs) setupTestVault(t, fs)
@@ -96,6 +106,8 @@ func TestListVersionsCommand(t *testing.T) {
vlt, err := vault.GetCurrentVault(fs, stateDir) vlt, err := vault.GetCurrentVault(fs, stateDir)
require.NoError(t, err) require.NoError(t, err)
vlt.Mnemonic = cli.Mnemonic
addTestSecret(t, vlt, []byte("version-1"), false) addTestSecret(t, vlt, []byte("version-1"), false)
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
@@ -139,8 +151,9 @@ func TestListVersionsCommand(t *testing.T) {
assert.Equal(t, 2, versionLines) assert.Equal(t, 2, versionLines)
} }
//nolint:paralleltest // uses t.Setenv via setupTestVault
func TestListVersionsNonExistentSecret(t *testing.T) { func TestListVersionsNonExistentSecret(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
stateDir := testStateDir stateDir := testStateDir
cli := NewCLIInstanceWithStateDir(fs, stateDir) cli := NewCLIInstanceWithStateDir(fs, stateDir)
@@ -161,8 +174,9 @@ func TestListVersionsNonExistentSecret(t *testing.T) {
assert.Contains(t, err.Error(), "not found") assert.Contains(t, err.Error(), "not found")
} }
//nolint:paralleltest // uses t.Setenv via setupTestVault
func TestPromoteVersionCommand(t *testing.T) { func TestPromoteVersionCommand(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
stateDir := testStateDir stateDir := testStateDir
cli := NewCLIInstanceWithStateDir(fs, stateDir) cli := NewCLIInstanceWithStateDir(fs, stateDir)
@@ -174,6 +188,8 @@ func TestPromoteVersionCommand(t *testing.T) {
vlt, err := vault.GetCurrentVault(fs, stateDir) vlt, err := vault.GetCurrentVault(fs, stateDir)
require.NoError(t, err) require.NoError(t, err)
vlt.Mnemonic = testMnemonicBuffer(t)
addTestSecret(t, vlt, []byte("version-1"), false) addTestSecret(t, vlt, []byte("version-1"), false)
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
@@ -224,8 +240,9 @@ func TestPromoteVersionCommand(t *testing.T) {
assert.Equal(t, []byte("version-1"), promoted.Bytes()) assert.Equal(t, []byte("version-1"), promoted.Bytes())
} }
//nolint:paralleltest // uses t.Setenv via setupTestVault
func TestPromoteNonExistentVersion(t *testing.T) { func TestPromoteNonExistentVersion(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
stateDir := testStateDir stateDir := testStateDir
cli := NewCLIInstanceWithStateDir(fs, stateDir) cli := NewCLIInstanceWithStateDir(fs, stateDir)
@@ -252,11 +269,13 @@ func TestPromoteNonExistentVersion(t *testing.T) {
assert.Contains(t, err.Error(), "not found") assert.Contains(t, err.Error(), "not found")
} }
//nolint:paralleltest // uses t.Setenv via setupTestVault
func TestGetSecretWithVersion(t *testing.T) { func TestGetSecretWithVersion(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
stateDir := testStateDir stateDir := testStateDir
cli := NewCLIInstanceWithStateDir(fs, stateDir) cli := NewCLIInstanceWithStateDir(fs, stateDir)
cli.Mnemonic = testMnemonicBuffer(t)
// Set up vault with long-term key // Set up vault with long-term key
setupTestVault(t, fs) setupTestVault(t, fs)
@@ -265,6 +284,8 @@ func TestGetSecretWithVersion(t *testing.T) {
vlt, err := vault.GetCurrentVault(fs, stateDir) vlt, err := vault.GetCurrentVault(fs, stateDir)
require.NoError(t, err) require.NoError(t, err)
vlt.Mnemonic = cli.Mnemonic
addTestSecret(t, vlt, []byte("version-1"), false) addTestSecret(t, vlt, []byte("version-1"), false)
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
@@ -298,10 +319,12 @@ func TestGetSecretWithVersion(t *testing.T) {
assert.Equal(t, "version-1", buf.String()) assert.Equal(t, "version-1", buf.String())
} }
//nolint:paralleltest // uses t.Setenv via setupTestVault
func TestGetSecretWritesBinaryValue(t *testing.T) { func TestGetSecretWritesBinaryValue(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
cli := NewCLIInstanceWithStateDir(fs, testStateDir) cli := NewCLIInstanceWithStateDir(fs, testStateDir)
cli.Mnemonic = testMnemonicBuffer(t)
setupTestVault(t, fs) setupTestVault(t, fs)
@@ -361,8 +384,9 @@ func TestVersionCommandStructure(t *testing.T) {
assert.Equal(t, "Promote a specific version to current", promoteCmd.Short) assert.Equal(t, "Promote a specific version to current", promoteCmd.Short)
} }
//nolint:paralleltest // uses t.Setenv via setupTestVault
func TestListVersionsEmptyOutput(t *testing.T) { func TestListVersionsEmptyOutput(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
stateDir := testStateDir stateDir := testStateDir
cli := NewCLIInstanceWithStateDir(fs, stateDir) cli := NewCLIInstanceWithStateDir(fs, stateDir)
+25 -21
View File
@@ -219,7 +219,7 @@ func newVaultWithSecret(
) *vault.Vault { ) *vault.Vault {
t.Helper() t.Helper()
vlt, err := vault.CreateVault(fs, stateDir, name) vlt, err := vault.CreateVault(fs, stateDir, name, testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
buffer := memguard.NewBufferFromBytes([]byte(value)) buffer := memguard.NewBufferFromBytes([]byte(value))
@@ -329,14 +329,14 @@ func TestRemoveDirAtomic(t *testing.T) {
// named with 255 bytes, the most a file name may have, on the real // named with 255 bytes, the most a file name may have, on the real
// filesystem: the temporary directories they use must fit that limit too. // filesystem: the temporary directories they use must fit that limit too.
func TestLongestNames(t *testing.T) { func TestLongestNames(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
const longestName = 255 const longestName = 255
fs := afero.NewOsFs() fs := afero.NewOsFs()
name := strings.Repeat("a", longestName) name := strings.Repeat("a", longestName)
vlt, err := vault.CreateVault(fs, t.TempDir(), name) vlt, err := vault.CreateVault(fs, t.TempDir(), name, testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
value := memguard.NewBufferFromBytes([]byte("long")) value := memguard.NewBufferFromBytes([]byte("long"))
@@ -361,13 +361,13 @@ func TestLongestNames(t *testing.T) {
// another vault, as a forced move between vaults does, and makes the last // another vault, as a forced move between vaults does, and makes the last
// step that completes the copy fail. The secret it was to replace must // step that completes the copy fail. The secret it was to replace must
// still be there unchanged: it may go only once its replacement is whole. // still be there unchanged: it may go only once its replacement is whole.
//
//nolint:paralleltest // t.Setenv forbids t.Parallel
func TestForcedCopyKeepsDestinationUntilReplaced(t *testing.T) { func TestForcedCopyKeepsDestinationUntilReplaced(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
for _, tfs := range testFilesystems { for _, tfs := range testFilesystems {
t.Run(tfs.name, func(t *testing.T) { t.Run(tfs.name, func(t *testing.T) {
t.Parallel()
base, stateDir := tfs.open(t) base, stateDir := tfs.open(t)
src := newVaultWithSecret(t, base, stateDir, "source", "new") src := newVaultWithSecret(t, base, stateDir, "source", "new")
dest := newVaultWithSecret(t, base, stateDir, "dest", "old") dest := newVaultWithSecret(t, base, stateDir, "dest", "old")
@@ -400,13 +400,13 @@ func TestForcedCopyKeepsDestinationUntilReplaced(t *testing.T) {
// directory directly in secrets.d or in a versions directory. Those are // directory directly in secrets.d or in a versions directory. Those are
// listed to find secrets and versions, so a temporary directory made there // listed to find secrets and versions, so a temporary directory made there
// would be listed while half-built, and one left by a crash would stay. // would be listed while half-built, and one left by a crash would stay.
//
//nolint:paralleltest // t.Setenv forbids t.Parallel
func TestTempDirsStayOutOfListings(t *testing.T) { func TestTempDirsStayOutOfListings(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
for _, tfs := range testFilesystems { for _, tfs := range testFilesystems {
t.Run(tfs.name, func(t *testing.T) { t.Run(tfs.name, func(t *testing.T) {
t.Parallel()
base, stateDir := tfs.open(t) base, stateDir := tfs.open(t)
newVaultWithSecret(t, base, stateDir, "default", "first") newVaultWithSecret(t, base, stateDir, "default", "first")
@@ -419,6 +419,7 @@ func TestTempDirsStayOutOfListings(t *testing.T) {
return nil return nil
}} }}
vlt := vault.NewVault(fs, stateDir, "default") vlt := vault.NewVault(fs, stateDir, "default")
vlt.Mnemonic = testMnemonicBuffer(t)
value := memguard.NewBufferFromBytes([]byte("second")) value := memguard.NewBufferFromBytes([]byte("second"))
defer value.Destroy() defer value.Destroy()
@@ -527,13 +528,13 @@ func TestVersionSaveFailureLeavesNothing(t *testing.T) {
// unlocker again and checks, before each change this makes, that the file // unlocker again and checks, before each change this makes, that the file
// naming the current one exists: a reader or a crash never finds it // naming the current one exists: a reader or a crash never finds it
// missing. // missing.
//
//nolint:paralleltest // t.Setenv forbids t.Parallel
func TestCurrentFilesNeverMissing(t *testing.T) { func TestCurrentFilesNeverMissing(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
for _, tfs := range testFilesystems { for _, tfs := range testFilesystems {
t.Run(tfs.name, func(t *testing.T) { t.Run(tfs.name, func(t *testing.T) {
t.Parallel()
base, stateDir := tfs.open(t) base, stateDir := tfs.open(t)
vlt := newVaultWithSecret(t, base, stateDir, testVaultName, "value") vlt := newVaultWithSecret(t, base, stateDir, testVaultName, "value")
@@ -625,11 +626,11 @@ func TestWriteFileAtomicTempFile(t *testing.T) {
// anything, so that it never leaves a partial unlocker, nor breaks the one // anything, so that it never leaves a partial unlocker, nor breaks the one
// it would replace. // it would replace.
func TestPassphraseUnlockerGetsKeyFirst(t *testing.T) { func TestPassphraseUnlockerGetsKeyFirst(t *testing.T) {
// No mnemonic, and no current unlocker to get the key from t.Parallel()
t.Setenv(secret.EnvMnemonic, "")
// No mnemonic, and no current unlocker to get the key from
base := afero.NewMemMapFs() base := afero.NewMemMapFs()
_, err := vault.CreateVault(base, testVaultStateDir, testVaultName) _, err := vault.CreateVault(base, testVaultStateDir, testVaultName, nil)
require.NoError(t, err) require.NoError(t, err)
fs := hookFs{Fs: base, before: func(_, path string) error { fs := hookFs{Fs: base, before: func(_, path string) error {
@@ -650,17 +651,18 @@ func TestPassphraseUnlockerGetsKeyFirst(t *testing.T) {
// creating a passphrase unlocker makes, that the unlocker's directory either // creating a passphrase unlocker makes, that the unlocker's directory either
// does not exist or holds all of its files: a crash or a failure at any point // does not exist or holds all of its files: a crash or a failure at any point
// leaves no partial unlocker. // leaves no partial unlocker.
//
//nolint:paralleltest // t.Setenv forbids t.Parallel
func TestPassphraseUnlockerIsWholeOrAbsent(t *testing.T) { func TestPassphraseUnlockerIsWholeOrAbsent(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
files := []string{"pub.age", privKeyFile, "longterm.age", unlockerMetadataFile} files := []string{"pub.age", privKeyFile, "longterm.age", unlockerMetadataFile}
for _, tfs := range testFilesystems { for _, tfs := range testFilesystems {
t.Run(tfs.name, func(t *testing.T) { t.Run(tfs.name, func(t *testing.T) {
t.Parallel()
base, stateDir := tfs.open(t) base, stateDir := tfs.open(t)
vlt, err := vault.CreateVault(base, stateDir, testVaultName) vlt, err := vault.CreateVault(base, stateDir, testVaultName,
testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
vaultDir, err := vlt.GetDirectory() vaultDir, err := vlt.GetDirectory()
@@ -683,8 +685,10 @@ func TestPassphraseUnlockerIsWholeOrAbsent(t *testing.T) {
passphrase := memguard.NewBufferFromBytes([]byte(unlockerPassphrase)) passphrase := memguard.NewBufferFromBytes([]byte(unlockerPassphrase))
defer passphrase.Destroy() defer passphrase.Destroy()
_, err = vault.NewVault(fs, stateDir, testVaultName). hooked := vault.NewVault(fs, stateDir, testVaultName)
CreatePassphraseUnlocker(passphrase) hooked.Mnemonic = vlt.Mnemonic
_, err = hooked.CreatePassphraseUnlocker(passphrase)
require.NoError(t, err) require.NoError(t, err)
assert.ElementsMatch(t, files, dirNames(t, base, unlockerDir)) assert.ElementsMatch(t, files, dirNames(t, base, unlockerDir))
}) })
+7 -2
View File
@@ -34,6 +34,8 @@ func (v *realVault) GetFilesystem() afero.Fs { return v.fs }
func (v *realVault) AddSecret(string, *memguard.LockedBuffer, bool) error { panic("not used") } func (v *realVault) AddSecret(string, *memguard.LockedBuffer, bool) error { panic("not used") }
func (v *realVault) GetCurrentUnlocker() (Unlocker, error) { panic("not used") } func (v *realVault) GetCurrentUnlocker() (Unlocker, error) { panic("not used") }
func (v *realVault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) { panic("not used") } func (v *realVault) 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) { func (v *realVault) CreatePassphraseUnlocker(*memguard.LockedBuffer) (*PassphraseUnlocker, error) {
panic("not used") panic("not used")
} }
@@ -59,6 +61,8 @@ func createRealVault(t *testing.T, fs afero.Fs, stateDir, name string, derivatio
} }
func TestGetLongTermPrivateKeyUsesVaultDerivationIndex(t *testing.T) { func TestGetLongTermPrivateKeyUsesVaultDerivationIndex(t *testing.T) {
t.Parallel()
const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about" 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.
@@ -73,9 +77,10 @@ func TestGetLongTermPrivateKeyUsesVaultDerivationIndex(t *testing.T) {
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
vault := createRealVault(t, fs, "/state", "test-vault", 5) vault := createRealVault(t, fs, "/state", "test-vault", 5)
t.Setenv(EnvMnemonic, testMnemonic) mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonic))
defer mnemonic.Destroy()
result, err := getLongTermPrivateKey(fs, vault) result, err := getLongTermPrivateKey(fs, vault, mnemonic, nil)
require.NoError(t, err) require.NoError(t, err)
defer result.Destroy() defer result.Destroy()
+19 -9
View File
@@ -239,12 +239,14 @@ func generateKeychainUnlockerName(vaultName string) (string, error) {
return fmt.Sprintf("secret-%s-%s-%s", vaultName, hostname, enrollmentDate), nil return fmt.Sprintf("secret-%s-%s-%s", vaultName, hostname, enrollmentDate), nil
} }
// getLongTermPrivateKey retrieves the long-term private key either from environment or current unlocker // getLongTermPrivateKey derives the long-term private key from mnemonic when
// it is not nil, else gets it through the current unlocker, which is given
// passphrase when it is a passphrase unlocker.
// Returns a LockedBuffer to ensure the private key is protected in memory // Returns a LockedBuffer to ensure the private key is protected in memory
func getLongTermPrivateKey(fs afero.Fs, vault VaultInterface) (*memguard.LockedBuffer, error) { func getLongTermPrivateKey(
// Check if mnemonic is available in environment variable fs afero.Fs, vault VaultInterface, mnemonic, passphrase *memguard.LockedBuffer,
envMnemonic := os.Getenv(EnvMnemonic) ) (*memguard.LockedBuffer, error) {
if envMnemonic != "" { if mnemonic != nil {
// Read vault metadata to get the correct derivation index // Read vault metadata to get the correct derivation index
vaultDir, err := vault.GetDirectory() vaultDir, err := vault.GetDirectory()
if err != nil { if err != nil {
@@ -263,7 +265,7 @@ func getLongTermPrivateKey(fs afero.Fs, vault VaultInterface) (*memguard.LockedB
} }
// Use mnemonic with the vault's actual derivation index // Use mnemonic with the vault's actual derivation index
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex) ltIdentity, err := agehd.DeriveIdentity(mnemonic.String(), metadata.DerivationIndex)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err) return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
} }
@@ -278,6 +280,10 @@ func getLongTermPrivateKey(fs afero.Fs, vault VaultInterface) (*memguard.LockedB
return nil, fmt.Errorf("failed to get current unlocker: %w", err) return nil, fmt.Errorf("failed to get current unlocker: %w", err)
} }
if passphraseUnlocker, ok := currentUnlocker.(*PassphraseUnlocker); ok {
passphraseUnlocker.Passphrase = passphrase
}
// Get the current unlocker identity // Get the current unlocker identity
currentUnlockerIdentity, err := currentUnlocker.GetIdentity() currentUnlockerIdentity, err := currentUnlocker.GetIdentity()
if err != nil { if err != nil {
@@ -322,8 +328,12 @@ func getLongTermPrivateKey(fs afero.Fs, vault VaultInterface) (*memguard.LockedB
return ltPrivKeyBuffer, nil return ltPrivKeyBuffer, nil
} }
// CreateKeychainUnlocker creates a new keychain unlocker and stores it in the vault // CreateKeychainUnlocker creates a new keychain unlocker and stores it in the
func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, error) { // vault. The long-term key comes from mnemonic when it is not nil, else from
// the current unlocker, as getLongTermPrivateKey describes.
func CreateKeychainUnlocker(
fs afero.Fs, stateDir string, mnemonic, passphrase *memguard.LockedBuffer,
) (*KeychainUnlocker, error) {
// Check if we're on macOS // Check if we're on macOS
if err := checkMacOSAvailable(); err != nil { if err := checkMacOSAvailable(); err != nil {
return nil, err return nil, err
@@ -376,7 +386,7 @@ func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, er
} }
// Step 4: Get or derive the long-term private key // Step 4: Get or derive the long-term private key
ltPrivKeyData, err := getLongTermPrivateKey(fs, vault) ltPrivKeyData, err := getLongTermPrivateKey(fs, vault, mnemonic, passphrase)
if err != nil { if err != nil {
return nil, err return nil, err
} }
+4 -1
View File
@@ -6,6 +6,7 @@ import (
"errors" "errors"
"filippo.io/age" "filippo.io/age"
"github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
) )
@@ -75,6 +76,8 @@ func (k *KeychainUnlocker) Remove() error {
} }
// CreateKeychainUnlocker returns an error on non-Darwin platforms // CreateKeychainUnlocker returns an error on non-Darwin platforms
func CreateKeychainUnlocker(_ afero.Fs, _ string) (*KeychainUnlocker, error) { func CreateKeychainUnlocker(
_ afero.Fs, _ string, _, _ *memguard.LockedBuffer,
) (*KeychainUnlocker, error) {
return nil, errKeychainNotSupported return nil, errKeychainNotSupported
} }
+34 -19
View File
@@ -19,6 +19,17 @@ import (
const testMnemonic = "abandon abandon abandon abandon abandon abandon " + const testMnemonic = "abandon abandon abandon abandon abandon abandon " +
"abandon abandon abandon abandon abandon about" "abandon abandon abandon abandon abandon about"
// testMnemonicBuffer returns testMnemonic in a locked buffer that is
// destroyed when the test ends.
func testMnemonicBuffer(t *testing.T) *memguard.LockedBuffer {
t.Helper()
mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonic))
t.Cleanup(mnemonic.Destroy)
return mnemonic
}
// writeTestPublicKey writes the unlocker public key and verifies it exists. // writeTestPublicKey writes the unlocker public key and verifies it exists.
func writeTestPublicKey( func writeTestPublicKey(
t *testing.T, fs afero.Fs, unlockerDir string, agePublicKey string, t *testing.T, fs afero.Fs, unlockerDir string, agePublicKey string,
@@ -163,7 +174,7 @@ func newTestPassphraseUnlocker(
return unlocker, ageIdentity, unlockerDir return unlocker, ageIdentity, unlockerDir
} }
//nolint:paralleltest // subtests share real-FS state and t.Setenv, order matters //nolint:paralleltest // subtests share real-FS state, order matters
func TestPassphraseUnlockerWithRealFS(t *testing.T) { func TestPassphraseUnlockerWithRealFS(t *testing.T) {
// This test uses real filesystem // This test uses real filesystem
if os.Getenv("CI") == "true" { if os.Getenv("CI") == "true" {
@@ -195,38 +206,42 @@ func TestPassphraseUnlockerWithRealFS(t *testing.T) {
writeTestLongTermKey(t, fs, unlockerDir, agePublicKey) writeTestLongTermKey(t, fs, unlockerDir, agePublicKey)
}) })
// Set test environment variable (cleaned up automatically) passphrase := memguard.NewBufferFromBytes([]byte(testPassphrase))
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase) defer passphrase.Destroy()
// Test getting identity from environment variable unlocker.Passphrase = passphrase
t.Run("GetIdentityFromEnv", func(t *testing.T) {
identity, err := unlocker.GetIdentity()
if err != nil {
t.Fatalf("Failed to get identity from env: %v", err)
}
// Verify the identity matches what we expect // Test getting identity with the passphrase the unlocker was given,
expectedPubKey := ageIdentity.Recipient().String() // twice: using it must leave it intact for the next use
t.Run("GetIdentityWithPassphrase", func(t *testing.T) {
for range 2 {
identity, err := unlocker.GetIdentity()
if err != nil {
t.Fatalf("Failed to get identity with passphrase: %v", err)
}
actualPubKey := identity.Recipient().String() // Verify the identity matches what we expect
if actualPubKey != expectedPubKey { expectedPubKey := ageIdentity.Recipient().String()
t.Errorf("Public key mismatch. Expected %s, got %s",
expectedPubKey, actualPubKey) actualPubKey := identity.Recipient().String()
if actualPubKey != expectedPubKey {
t.Errorf("Public key mismatch. Expected %s, got %s",
expectedPubKey, actualPubKey)
}
} }
}) })
// Unset the environment variable to test interactive prompt unlocker.Passphrase = nil
_ = os.Unsetenv(secret.EnvUnlockPassphrase)
// Test getting identity from prompt (this would require mocking the // Test getting identity from prompt (this would require mocking the
// prompt). For real integration tests, we'd need a way to mock the // prompt). For real integration tests, we'd need a way to mock the
// passphrase input. Here we just verify the error is what we expect // passphrase input. Here we just verify the error is what we expect
// when no passphrase is available. // when no passphrase is available.
t.Run("GetIdentityWithoutEnv", func(t *testing.T) { t.Run("GetIdentityWithoutPassphrase", func(t *testing.T) {
// This should fail since we're not in an interactive terminal // This should fail since we're not in an interactive terminal
_, err := unlocker.GetIdentity() _, err := unlocker.GetIdentity()
if err == nil { if err == nil {
t.Errorf("Should have failed to get identity without passphrase env var") t.Errorf("Should have failed to get identity without a passphrase")
} }
}) })
+8 -18
View File
@@ -3,7 +3,6 @@ package secret
import ( import (
"fmt" "fmt"
"log/slog" "log/slog"
"os"
"path/filepath" "path/filepath"
"filippo.io/age" "filippo.io/age"
@@ -135,28 +134,19 @@ func (p *PassphraseUnlocker) Remove() error {
return nil return nil
} }
// getPassphrase retrieves the passphrase from memory, environment, or // getPassphrase returns a copy of p.Passphrase, or else asks the user for
// user input. Returns a LockedBuffer for secure memory handling // the passphrase. The caller must destroy the returned buffer.
func (p *PassphraseUnlocker) getPassphrase() (*memguard.LockedBuffer, error) { func (p *PassphraseUnlocker) getPassphrase() (*memguard.LockedBuffer, error) {
// First check if we already have the passphrase
if p.Passphrase != nil && p.Passphrase.IsAlive() { if p.Passphrase != nil && p.Passphrase.IsAlive() {
Debug("Using in-memory passphrase", "unlocker_id", p.GetID()) Debug("Using in-memory passphrase", "unlocker_id", p.GetID())
// Return a copy of the passphrase buffer // Not NewBufferFromBytes, which would wipe p.Passphrase
return memguard.NewBufferFromBytes(p.Passphrase.Bytes()), nil passphrase := memguard.NewBuffer(p.Passphrase.Size())
passphrase.Copy(p.Passphrase.Bytes())
return passphrase, nil
} }
Debug("No passphrase in memory, checking environment") Debug("No passphrase in memory, prompting user")
// Check environment variable for passphrase
passphraseStr := os.Getenv(EnvUnlockPassphrase)
if passphraseStr != "" {
Debug("Using passphrase from environment", "unlocker_id", p.GetID())
// Convert to secure buffer
secureBuffer := memguard.NewBufferFromBytes([]byte(passphraseStr))
return secureBuffer, nil
}
Debug("No passphrase in environment, prompting user")
// Prompt for passphrase // Prompt for passphrase
secureBuffer, err := ReadPassphrase("Enter unlock passphrase: ") secureBuffer, err := ReadPassphrase("Enter unlock passphrase: ")
if err != nil { if err != nil {
+5 -3
View File
@@ -227,8 +227,10 @@ Passphrase: ` + testPassphrase + `
// Test data // Test data
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about" 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.EnvMnemonic, testMnemonic)
t.Setenv(secret.EnvGPGKeyID, keyID) t.Setenv(secret.EnvGPGKeyID, keyID)
// Set up vault structure for testing // Set up vault structure for testing
@@ -244,7 +246,7 @@ Passphrase: ` + testPassphrase + `
defer timer.Stop() defer timer.Stop()
// Create a test vault directory structure // Create a test vault directory structure
vlt, err := vault.CreateVault(fs, stateDir, vaultName) vlt, err := vault.CreateVault(fs, stateDir, vaultName, mnemonic)
if err != nil { if err != nil {
t.Fatalf("Failed to create vault: %v", err) t.Fatalf("Failed to create vault: %v", err)
} }
@@ -290,7 +292,7 @@ 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) 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)
} }
+8 -1
View File
@@ -254,9 +254,12 @@ func pgpUnlockerDir(
// fingerprint as ResolveGPGKeyFingerprint returns it, in the metadata. // fingerprint as ResolveGPGKeyFingerprint returns it, in the metadata.
// Everything that can fail short of writing a file is done before anything // Everything that can fail short of writing a file is done before anything
// is written, and the files are written through WriteDir, so a failure // is written, and the files are written through WriteDir, so a failure
// leaves no partial unlocker. // leaves no partial unlocker. The long-term key comes from mnemonic when it
// is not nil, else from the current unlocker, which is given passphrase when
// it is a passphrase unlocker.
func CreatePGPUnlocker( func CreatePGPUnlocker(
fs afero.Fs, stateDir, gpgKeyID, fingerprint string, fs afero.Fs, stateDir, gpgKeyID, fingerprint string,
mnemonic, passphrase *memguard.LockedBuffer,
) (*PGPUnlocker, error) { ) (*PGPUnlocker, error) {
err := checkGPGAvailable() err := checkGPGAvailable()
if err != nil { if err != nil {
@@ -268,6 +271,10 @@ func CreatePGPUnlocker(
return nil, err return nil, err
} }
// The vault's GetOrDeriveLongTermKey, in step 2, uses both
vault.SetMnemonic(mnemonic)
vault.SetUnlockPassphrase(passphrase)
// Step 1: Generate a new age keypair for the PGP unlocker // Step 1: Generate a new age keypair for the PGP unlocker
ageIdentity, err := age.GenerateX25519Identity() ageIdentity, err := age.GenerateX25519Identity()
if err != nil { if err != nil {
+4 -3
View File
@@ -41,12 +41,13 @@ func installFakeGPG(t *testing.T) {
// getting the vault's long-term key, which used to come after part of the // getting the vault's long-term key, which used to come after part of the
// unlocker was written, and asserts that nothing is written. Getting the key // unlocker was written, and asserts that nothing is written. Getting the key
// fails because there is no mnemonic and no current unlocker. // fails because there is no mnemonic and no current unlocker.
//
//nolint:paralleltest // installFakeGPG uses t.Setenv
func TestCreatePGPUnlockerFailureWritesNothing(t *testing.T) { func TestCreatePGPUnlockerFailureWritesNothing(t *testing.T) {
installFakeGPG(t) installFakeGPG(t)
t.Setenv(secret.EnvMnemonic, "")
base := afero.NewMemMapFs() base := afero.NewMemMapFs()
vlt, err := vault.CreateVault(base, testVaultStateDir, testVaultName) vlt, err := vault.CreateVault(base, testVaultStateDir, testVaultName, nil)
require.NoError(t, err) require.NoError(t, err)
fs := hookFs{Fs: base, before: func(_, path string) error { fs := hookFs{Fs: base, before: func(_, path string) error {
@@ -56,7 +57,7 @@ func TestCreatePGPUnlockerFailureWritesNothing(t *testing.T) {
}} }}
_, err = secret.CreatePGPUnlocker( _, err = secret.CreatePGPUnlocker(
fs, testVaultStateDir, testGPGKeyID, testGPGFingerprint) fs, testVaultStateDir, testGPGKeyID, testGPGFingerprint, nil, nil)
require.Error(t, err) require.Error(t, err)
vaultDir, err := vlt.GetDirectory() vaultDir, err := vlt.GetDirectory()
+17 -11
View File
@@ -5,7 +5,6 @@ import (
"errors" "errors"
"fmt" "fmt"
"log/slog" "log/slog"
"os"
"path/filepath" "path/filepath"
"strings" "strings"
"time" "time"
@@ -36,6 +35,11 @@ type VaultInterface interface {
GetFilesystem() afero.Fs GetFilesystem() afero.Fs
GetCurrentUnlocker() (Unlocker, error) GetCurrentUnlocker() (Unlocker, error)
GetOrDeriveLongTermKey() (*age.X25519Identity, error) GetOrDeriveLongTermKey() (*age.X25519Identity, error)
// SetMnemonic and SetUnlockPassphrase give GetOrDeriveLongTermKey the
// mnemonic to derive the long-term key from, and the passphrase for a
// current passphrase unlocker; nil for none.
SetMnemonic(mnemonic *memguard.LockedBuffer)
SetUnlockPassphrase(passphrase *memguard.LockedBuffer)
CreatePassphraseUnlocker( CreatePassphraseUnlocker(
passphrase *memguard.LockedBuffer) (*PassphraseUnlocker, error) passphrase *memguard.LockedBuffer) (*PassphraseUnlocker, error)
} }
@@ -77,9 +81,12 @@ func NewSecret(vault VaultInterface, name string) *Secret {
} }
} }
// GetValue retrieves and decrypts the current version's value using the // GetValue retrieves and decrypts the current version's value, with the
// provided unlocker // vault's long-term key derived from mnemonic when it is not nil, else
func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) { // obtained through unlocker
func (s *Secret) GetValue(
unlocker Unlocker, mnemonic *memguard.LockedBuffer,
) (*memguard.LockedBuffer, error) {
DebugWith("Getting secret value", DebugWith("Getting secret value",
slog.String("secret_name", s.Name), slog.String("secret_name", s.Name),
slog.String("vault_name", s.vault.GetName()), slog.String("vault_name", s.vault.GetName()),
@@ -114,9 +121,8 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) {
// Create version object // Create version object
version := NewVersion(s.vault, s.Name, currentVersion) version := NewVersion(s.vault, s.Name, currentVersion)
// Check for SB_SECRET_MNEMONIC environment variable for direct decryption if mnemonic != nil {
if envMnemonic := os.Getenv(EnvMnemonic); envMnemonic != "" { return s.getValueViaMnemonic(version, mnemonic.String())
return s.getValueViaMnemonic(version, envMnemonic)
} }
Debug("Using unlocker for vault access", "secret_name", s.Name) Debug("Using unlocker for vault access", "secret_name", s.Name)
@@ -210,11 +216,11 @@ func (s *Secret) Exists() (bool, error) {
} }
// getValueViaMnemonic derives the vault's long-term key from the // getValueViaMnemonic derives the vault's long-term key from the
// mnemonic in the environment and decrypts the version value with it. // mnemonic and decrypts the version value with it.
func (s *Secret) getValueViaMnemonic( func (s *Secret) getValueViaMnemonic(
version *Version, envMnemonic string, version *Version, mnemonic string,
) (*memguard.LockedBuffer, error) { ) (*memguard.LockedBuffer, error) {
Debug("Using mnemonic from environment for direct long-term key derivation", Debug("Using mnemonic for direct long-term key derivation",
"secret_name", s.Name) "secret_name", s.Name)
// Get vault directory to read metadata // Get vault directory to read metadata
@@ -251,7 +257,7 @@ func (s *Secret) getValueViaMnemonic(
) )
// Use mnemonic with the vault's derivation index from metadata // Use mnemonic with the vault's derivation index from metadata
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex) ltIdentity, err := agehd.DeriveIdentity(mnemonic, metadata.DerivationIndex)
if err != nil { if err != nil {
Debug("Failed to derive long-term key from mnemonic for secret", Debug("Failed to derive long-term key from mnemonic for secret",
"error", err, "secret_name", s.Name) "error", err, "secret_name", s.Name)
+50 -23
View File
@@ -2,6 +2,7 @@
package secret package secret
import ( import (
"encoding/json"
"errors" "errors"
"os" "os"
"path/filepath" "path/filepath"
@@ -22,7 +23,7 @@ const testMnemonicValue = "abandon abandon abandon abandon abandon abandon " +
"abandon abandon abandon abandon abandon about" "abandon abandon abandon abandon abandon about"
var ( var (
errMnemonicNotSet = errors.New("SB_SECRET_MNEMONIC not set") errMnemonicNotSet = errors.New("mock vault has no mnemonic")
errNotImplementedInMock = errors.New("not implemented in mock") errNotImplementedInMock = errors.New("not implemented in mock")
) )
@@ -32,6 +33,7 @@ type MockVault struct {
fs afero.Fs fs afero.Fs
directory string directory string
derivationIndex uint32 derivationIndex uint32
mnemonic *memguard.LockedBuffer
} }
func (m *MockVault) GetDirectory() (string, error) { func (m *MockVault) GetDirectory() (string, error) {
@@ -61,12 +63,11 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool)
ltPubKeyPath := filepath.Join(m.directory, "pub.age") ltPubKeyPath := filepath.Join(m.directory, "pub.age")
// Derive long-term key using the vault's derivation index // Derive long-term key using the vault's derivation index
mnemonic := os.Getenv(EnvMnemonic) if m.mnemonic == nil {
if mnemonic == "" {
return errMnemonicNotSet return errMnemonicNotSet
} }
ltIdentity, err := agehd.DeriveIdentity(mnemonic, m.derivationIndex) ltIdentity, err := agehd.DeriveIdentity(m.mnemonic.String(), m.derivationIndex)
if err != nil { if err != nil {
return err return err
} }
@@ -111,6 +112,12 @@ func (m *MockVault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
return nil, errNotImplementedInMock return nil, errNotImplementedInMock
} }
func (m *MockVault) SetMnemonic(mnemonic *memguard.LockedBuffer) {
m.mnemonic = mnemonic
}
func (m *MockVault) SetUnlockPassphrase(_ *memguard.LockedBuffer) {}
func (m *MockVault) CreatePassphraseUnlocker( func (m *MockVault) CreatePassphraseUnlocker(
_ *memguard.LockedBuffer, _ *memguard.LockedBuffer,
) (*PassphraseUnlocker, error) { ) (*PassphraseUnlocker, error) {
@@ -238,13 +245,13 @@ func verifySecretFiles(t *testing.T, fs afero.Fs, vaultDir, secretName string) {
} }
} }
//nolint:paralleltest // uses t.Setenv (process-global environment) //nolint:paralleltest // subtests share one vault, order matters
func TestPerSecretKeyFunctionality(t *testing.T) { func TestPerSecretKeyFunctionality(t *testing.T) {
// Create an in-memory filesystem for testing // Create an in-memory filesystem for testing
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Set test mnemonic for direct encryption/decryption mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonicValue))
t.Setenv(EnvMnemonic, testMnemonicValue) defer mnemonic.Destroy()
// Set up a test vault structure // Set up a test vault structure
baseDir := "/test-config/berlin.sneak.pkg.secret" baseDir := "/test-config/berlin.sneak.pkg.secret"
@@ -258,6 +265,7 @@ func TestPerSecretKeyFunctionality(t *testing.T) {
fs: fs, fs: fs,
directory: vaultDir, directory: vaultDir,
derivationIndex: 0, derivationIndex: 0,
mnemonic: mnemonic,
} }
// Test data // Test data
@@ -314,26 +322,45 @@ func TestPerSecretKeyFunctionality(t *testing.T) {
}) })
} }
func TestSecretGetValueWithEnvMnemonicUsesVaultDerivationIndex(t *testing.T) { // TestSecretGetValueWithMnemonicUsesVaultDerivationIndex checks that
// This test demonstrates the bug where GetValue uses hardcoded index 0 // GetValue, given the mnemonic, derives the long-term key at the derivation
// instead of the vault's actual derivation index when using environment mnemonic // index in the vault's metadata. At index 0 it could not decrypt the secret,
// which was encrypted to the key at index 1.
func TestSecretGetValueWithMnemonicUsesVaultDerivationIndex(t *testing.T) {
t.Parallel()
// Set up test mnemonic fs := afero.NewMemMapFs()
t.Setenv(EnvMnemonic, testMnemonicValue) vaultDir := "/test-config/vaults.d/test-vault"
// Create temporary directory for vaults mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonicValue))
fs := afero.NewOsFs() defer mnemonic.Destroy()
tempDir, err := afero.TempDir(fs, "", "secret-test-")
vlt := &MockVault{
name: "test-vault",
fs: fs,
directory: vaultDir,
derivationIndex: 1,
mnemonic: mnemonic,
}
metadata, err := json.Marshal(VaultMetadata{DerivationIndex: vlt.derivationIndex})
require.NoError(t, err)
require.NoError(t, fs.MkdirAll(vaultDir, DirPerms))
err = afero.WriteFile(
fs, filepath.Join(vaultDir, "vault-metadata.json"), metadata, FilePerms)
require.NoError(t, err) require.NoError(t, err)
defer func() { secretName, secretValue := "x", "value"
_ = fs.RemoveAll(tempDir)
}()
stateDir := filepath.Join(tempDir, ".secret") err = vlt.AddSecret(secretName,
require.NoError(t, fs.MkdirAll(stateDir, 0o700)) memguard.NewBufferFromBytes([]byte(secretValue)), false)
require.NoError(t, err)
// This test is now in the integration test file where it can use real vaults value, err := NewSecret(vlt, secretName).GetValue(nil, mnemonic)
// The bug is demonstrated there - see test31EnvMnemonicUsesVaultDerivationIndex require.NoError(t, err)
t.Log("This test demonstrates the bug in the integration test file")
defer value.Destroy()
require.Equal(t, secretValue, value.String())
} }
+14 -6
View File
@@ -207,9 +207,12 @@ func generateSEKeyLabel(vaultName string) (string, error) {
// CreateSecureEnclaveUnlocker creates a new SE unlocker. // CreateSecureEnclaveUnlocker creates a new SE unlocker.
// The vault's long-term private key is encrypted directly by the Secure Enclave // The vault's long-term private key is encrypted directly by the Secure Enclave
// using ECIES. No intermediate age keypair is used. // using ECIES. No intermediate age keypair is used.
// The long-term key comes from mnemonic when it is not nil, else from the
// current unlocker, as getLongTermKeyForSE describes.
func CreateSecureEnclaveUnlocker( func CreateSecureEnclaveUnlocker(
fs afero.Fs, fs afero.Fs,
stateDir string, stateDir string,
mnemonic, passphrase *memguard.LockedBuffer,
) (*SecureEnclaveUnlocker, error) { ) (*SecureEnclaveUnlocker, error) {
if err := checkMacOSAvailable(); err != nil { if err := checkMacOSAvailable(); err != nil {
return nil, err return nil, err
@@ -236,7 +239,7 @@ func CreateSecureEnclaveUnlocker(
Debug("Created SE key", "label", seKeyLabel, "hash", seKeyHash) Debug("Created SE key", "label", seKeyLabel, "hash", seKeyHash)
// Step 2: Get the vault's long-term private key // Step 2: Get the vault's long-term private key
ltPrivKeyData, err := getLongTermKeyForSE(fs, vault) ltPrivKeyData, err := getLongTermKeyForSE(fs, vault, mnemonic, passphrase)
if err != nil { if err != nil {
return nil, fmt.Errorf( return nil, fmt.Errorf(
"failed to get long-term private key: %w", "failed to get long-term private key: %w",
@@ -306,14 +309,15 @@ func CreateSecureEnclaveUnlocker(
}, nil }, nil
} }
// getLongTermKeyForSE retrieves the vault's long-term private key // getLongTermKeyForSE retrieves the vault's long-term private key, derived
// either from the mnemonic env var or by unlocking via the current unlocker. // from mnemonic when it is not nil, else through the current unlocker, which
// is given passphrase when it is a passphrase unlocker.
func getLongTermKeyForSE( func getLongTermKeyForSE(
fs afero.Fs, fs afero.Fs,
vault VaultInterface, vault VaultInterface,
mnemonic, passphrase *memguard.LockedBuffer,
) (*memguard.LockedBuffer, error) { ) (*memguard.LockedBuffer, error) {
envMnemonic := os.Getenv(EnvMnemonic) if mnemonic != nil {
if envMnemonic != "" {
// Read vault metadata to get the correct derivation index // Read vault metadata to get the correct derivation index
vaultDir, err := vault.GetDirectory() vaultDir, err := vault.GetDirectory()
if err != nil { if err != nil {
@@ -333,7 +337,7 @@ func getLongTermKeyForSE(
// Use mnemonic with the vault's actual derivation index // Use mnemonic with the vault's actual derivation index
ltIdentity, err := agehd.DeriveIdentity( ltIdentity, err := agehd.DeriveIdentity(
envMnemonic, mnemonic.String(),
metadata.DerivationIndex, metadata.DerivationIndex,
) )
@@ -352,6 +356,10 @@ func getLongTermKeyForSE(
return nil, fmt.Errorf("failed to get current unlocker: %w", err) return nil, fmt.Errorf("failed to get current unlocker: %w", err)
} }
if passphraseUnlocker, ok := currentUnlocker.(*PassphraseUnlocker); ok {
passphraseUnlocker.Passphrase = passphrase
}
currentIdentity, err := currentUnlocker.GetIdentity() currentIdentity, err := currentUnlocker.GetIdentity()
if err != nil { if err != nil {
return nil, fmt.Errorf( return nil, fmt.Errorf(
+2
View File
@@ -6,6 +6,7 @@ import (
"errors" "errors"
"filippo.io/age" "filippo.io/age"
"github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
) )
@@ -80,6 +81,7 @@ func (s *SecureEnclaveUnlocker) Remove() error {
func CreateSecureEnclaveUnlocker( func CreateSecureEnclaveUnlocker(
_ afero.Fs, _ afero.Fs,
_ string, _ string,
_, _ *memguard.LockedBuffer,
) (*SecureEnclaveUnlocker, error) { ) (*SecureEnclaveUnlocker, error) {
return nil, errSENotSupported return nil, errSENotSupported
} }
+1 -1
View File
@@ -78,7 +78,7 @@ func TestCreateSecureEnclaveUnlockerReturnsError(t *testing.T) {
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
unlocker, err := CreateSecureEnclaveUnlocker(fs, "/tmp/test") unlocker, err := CreateSecureEnclaveUnlocker(fs, "/tmp/test", nil, nil)
assert.Nil(t, unlocker) assert.Nil(t, unlocker)
require.Error(t, err) require.Error(t, err)
require.ErrorIs(t, err, errSENotSupported) require.ErrorIs(t, err, errSENotSupported)
+4
View File
@@ -91,6 +91,10 @@ func (m *MockVersionVault) GetOrDeriveLongTermKey() (*age.X25519Identity, error)
return nil, errNotImplementedInMock return nil, errNotImplementedInMock
} }
func (m *MockVersionVault) SetMnemonic(_ *memguard.LockedBuffer) {}
func (m *MockVersionVault) SetUnlockPassphrase(_ *memguard.LockedBuffer) {}
func (m *MockVersionVault) CreatePassphraseUnlocker( func (m *MockVersionVault) CreatePassphraseUnlocker(
_ *memguard.LockedBuffer, _ *memguard.LockedBuffer,
) (*secret.PassphraseUnlocker, error) { ) (*secret.PassphraseUnlocker, error) {
+25 -20
View File
@@ -8,7 +8,6 @@ import (
"testing" "testing"
"filippo.io/age" "filippo.io/age"
"git.eeqj.de/sneak/secret/internal/secret"
"git.eeqj.de/sneak/secret/internal/vault" "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/awnumar/memguard"
@@ -41,46 +40,49 @@ func deriveVaultIdentity(
return ltIdentity return ltIdentity
} }
//nolint:paralleltest // t.Setenv forbids parallel subtests
func TestVaultWithRealFilesystem(t *testing.T) { func TestVaultWithRealFilesystem(t *testing.T) {
t.Parallel()
// Create a temporary directory for our tests // Create a temporary directory for our tests
tempDir := t.TempDir() tempDir := t.TempDir()
// Use the real filesystem // Use the real filesystem
fs := afero.NewOsFs() fs := afero.NewOsFs()
// Set test environment variables
t.Setenv(secret.EnvMnemonic, testMnemonic)
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
// Test currentvault file handling (plain file with relative path) // Test currentvault file handling (plain file with relative path)
t.Run("CurrentVaultFileHandling", func(t *testing.T) { t.Run("CurrentVaultFileHandling", func(t *testing.T) {
t.Parallel()
testCurrentVaultFileHandling(t, fs, tempDir) testCurrentVaultFileHandling(t, fs, tempDir)
}) })
// Test secret operations with deeply nested paths // Test secret operations with deeply nested paths
t.Run("DeepPathSecrets", func(t *testing.T) { t.Run("DeepPathSecrets", func(t *testing.T) {
t.Parallel()
testDeepPathSecrets(t, fs, tempDir) testDeepPathSecrets(t, fs, tempDir)
}) })
// Test key caching in GetOrDeriveLongTermKey // Test key caching in GetOrDeriveLongTermKey
t.Run("KeyCaching", func(t *testing.T) { t.Run("KeyCaching", func(t *testing.T) {
t.Parallel()
testKeyCaching(t, fs, tempDir) testKeyCaching(t, fs, tempDir)
}) })
// Test vault name validation // Test vault name validation
t.Run("VaultNameValidation", func(t *testing.T) { t.Run("VaultNameValidation", func(t *testing.T) {
t.Parallel()
testVaultNameValidation(t, fs, tempDir) testVaultNameValidation(t, fs, tempDir)
}) })
// Test multiple vaults and switching between them // Test multiple vaults and switching between them
t.Run("MultipleVaults", func(t *testing.T) { t.Run("MultipleVaults", func(t *testing.T) {
t.Parallel()
testMultipleVaults(t, fs, tempDir) testMultipleVaults(t, fs, tempDir)
}) })
// Test adding a secret in one vault and verifying it's not visible in // Test adding a secret in one vault and verifying it's not visible in
// another // another
t.Run("VaultIsolation", func(t *testing.T) { t.Run("VaultIsolation", func(t *testing.T) {
t.Parallel()
testVaultIsolation(t, fs, tempDir) testVaultIsolation(t, fs, tempDir)
}) })
} }
@@ -96,7 +98,8 @@ func testCurrentVaultFileHandling(t *testing.T, fs afero.Fs, tempDir string) {
} }
// Create a test vault // Create a test vault
vlt, err := vault.CreateVault(fs, stateDir, testVaultName) vlt, err := vault.CreateVault(fs, stateDir, testVaultName,
testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Fatalf("Failed to create vault: %v", err) t.Fatalf("Failed to create vault: %v", err)
} }
@@ -141,9 +144,10 @@ func testDeepPathSecrets(t *testing.T, fs afero.Fs, tempDir string) {
t.Fatalf("Failed to create state dir: %v", err) t.Fatalf("Failed to create state dir: %v", err)
} }
// Create a test vault - CreateVault now handles public key when // Create a test vault - CreateVault writes the public key derived from
// mnemonic is in env // the mnemonic
vlt, err := vault.CreateVault(fs, stateDir, testVaultName) vlt, err := vault.CreateVault(fs, stateDir, testVaultName,
testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Fatalf("Failed to create vault: %v", err) t.Fatalf("Failed to create vault: %v", err)
} }
@@ -216,9 +220,10 @@ func testKeyCaching(t *testing.T, fs afero.Fs, tempDir string) {
t.Fatalf("Failed to create state dir: %v", err) t.Fatalf("Failed to create state dir: %v", err)
} }
// Create a test vault - CreateVault now handles public key when // Create a test vault - CreateVault writes the public key derived from
// mnemonic is in env // the mnemonic
vlt, err := vault.CreateVault(fs, stateDir, testVaultName) vlt, err := vault.CreateVault(fs, stateDir, testVaultName,
testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Fatalf("Failed to create vault: %v", err) t.Fatalf("Failed to create vault: %v", err)
} }
@@ -319,7 +324,7 @@ func testVaultNameValidation(t *testing.T, fs afero.Fs, tempDir string) {
} }
for _, name := range validNames { for _, name := range validNames {
_, err := vault.CreateVault(fs, stateDir, name) _, err := vault.CreateVault(fs, stateDir, name, testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Errorf("Failed to create vault with valid name %q: %v", name, err) t.Errorf("Failed to create vault with valid name %q: %v", name, err)
} }
@@ -335,7 +340,7 @@ func testVaultNameValidation(t *testing.T, fs afero.Fs, tempDir string) {
} }
for _, name := range invalidNames { for _, name := range invalidNames {
_, err := vault.CreateVault(fs, stateDir, name) _, err := vault.CreateVault(fs, stateDir, name, testMnemonicBuffer(t))
if err == nil { if err == nil {
t.Errorf("Expected error creating vault with invalid name %q, "+ t.Errorf("Expected error creating vault with invalid name %q, "+
"but got none", name) "but got none", name)
@@ -356,7 +361,7 @@ func testMultipleVaults(t *testing.T, fs afero.Fs, tempDir string) {
// Create three vaults // Create three vaults
vaultNames := []string{"vault1", "vault2", "vault3"} vaultNames := []string{"vault1", "vault2", "vault3"}
for _, name := range vaultNames { for _, name := range vaultNames {
_, err := vault.CreateVault(fs, stateDir, name) _, err := vault.CreateVault(fs, stateDir, name, testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Fatalf("Failed to create vault %s: %v", name, err) t.Fatalf("Failed to create vault %s: %v", name, err)
} }
@@ -404,14 +409,14 @@ func testVaultIsolation(t *testing.T, fs afero.Fs, tempDir string) {
t.Fatalf("Failed to create state dir: %v", err) t.Fatalf("Failed to create state dir: %v", err)
} }
// Create two vaults - CreateVault now handles public key when mnemonic // Create two vaults - CreateVault writes the public key derived from
// is in env // the mnemonic
vault1, err := vault.CreateVault(fs, stateDir, "vault1") vault1, err := vault.CreateVault(fs, stateDir, "vault1", testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Fatalf("Failed to create vault1: %v", err) t.Fatalf("Failed to create vault1: %v", err)
} }
vault2, err := vault.CreateVault(fs, stateDir, "vault2") vault2, err := vault.CreateVault(fs, stateDir, "vault2", testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Fatalf("Failed to create vault2: %v", err) t.Fatalf("Failed to create vault2: %v", err)
} }
+9 -10
View File
@@ -44,15 +44,12 @@ var errUnexpectedValue = errors.New("unexpected value")
// TestVersionIntegrationWorkflow tests the complete version workflow // TestVersionIntegrationWorkflow tests the complete version workflow
// //
//nolint:paralleltest // t.Setenv forbids parallel subtests //nolint:paralleltest // the subtests are steps that build on each other
func TestVersionIntegrationWorkflow(t *testing.T) { func TestVersionIntegrationWorkflow(t *testing.T) {
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Set mnemonic for testing // Create vault without a long-term key, which is set up below
t.Setenv(secret.EnvMnemonic, testMnemonic) vault, err := CreateVault(fs, testStateDir, "test", nil)
// Create vault
vault, err := CreateVault(fs, testStateDir, "test")
require.NoError(t, err) require.NoError(t, err)
// Derive and store long-term key from mnemonic // Derive and store long-term key from mnemonic
@@ -351,9 +348,9 @@ func testVersionErrorCases(t *testing.T, vault *Vault, secretName string) {
} }
// TestVersionConcurrency tests concurrent version operations // TestVersionConcurrency tests concurrent version operations
//
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
func TestVersionConcurrency(t *testing.T) { func TestVersionConcurrency(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Set up vault // Set up vault
@@ -366,6 +363,8 @@ func TestVersionConcurrency(t *testing.T) {
// Test concurrent reads // Test concurrent reads
t.Run("concurrent_reads", func(t *testing.T) { t.Run("concurrent_reads", func(t *testing.T) {
t.Parallel()
done := make(chan bool, 10) done := make(chan bool, 10)
errCh := make(chan error, 10) errCh := make(chan error, 10)
@@ -403,9 +402,9 @@ func TestVersionConcurrency(t *testing.T) {
} }
// TestVersionCompatibility tests that old secrets without versions still work // TestVersionCompatibility tests that old secrets without versions still work
//
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
func TestVersionCompatibility(t *testing.T) { func TestVersionCompatibility(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Set up vault // Set up vault
+23 -16
View File
@@ -3,7 +3,6 @@ package vault
import ( import (
"fmt" "fmt"
"os"
"path/filepath" "path/filepath"
"regexp" "regexp"
"strings" "strings"
@@ -11,6 +10,7 @@ import (
"git.eeqj.de/sneak/secret/internal/secret" "git.eeqj.de/sneak/secret/internal/secret"
"git.eeqj.de/sneak/secret/pkg/agehd" "git.eeqj.de/sneak/secret/pkg/agehd"
"github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
) )
@@ -155,19 +155,18 @@ func ListVaults(fs afero.Fs, stateDir string) ([]string, error) {
// It returns the derivation index, public key hash, and family hash. // It returns the derivation index, public key hash, and family hash.
func processMnemonicForVault( func processMnemonicForVault(
fs afero.Fs, stateDir, vaultDir, vaultName string, fs afero.Fs, stateDir, vaultDir, vaultName string,
mnemonicBuffer *memguard.LockedBuffer,
) (uint32, string, string, error) { ) (uint32, string, string, error) {
// Check if mnemonic is available in environment if mnemonicBuffer == nil {
mnemonic := os.Getenv(secret.EnvMnemonic) secret.Debug("No mnemonic given, vault created without long-term key",
if mnemonic == "" {
secret.Debug("No mnemonic in environment, vault created without long-term key",
"vault", vaultName) "vault", vaultName)
// Use 0 for derivation index when no mnemonic is provided // Use 0 for derivation index when no mnemonic is provided
return 0, "", "", nil return 0, "", "", nil
} }
secret.Debug("Mnemonic found in environment, deriving long-term key", mnemonic := mnemonicBuffer.String()
"vault", vaultName)
secret.Debug("Mnemonic given, deriving long-term key", "vault", vaultName)
// Get the next available derivation index for this mnemonic // Get the next available derivation index for this mnemonic
derivationIndex, err := GetNextDerivationIndex(fs, stateDir, mnemonic) derivationIndex, err := GetNextDerivationIndex(fs, stateDir, mnemonic)
@@ -208,12 +207,17 @@ func processMnemonicForVault(
return derivationIndex, publicKeyHash, familyHash, nil return derivationIndex, publicKeyHash, familyHash, nil
} }
// CreateVault creates a new vault and selects it as the current vault. It // CreateVault creates a new vault and selects it as the current vault. When
// refuses a vault that already exists before writing anything: creating it // mnemonic is not nil, the vault's long-term key is derived from it, and the
// again would replace its keys, and its secrets could no longer be // returned vault has it as its Mnemonic; when it is nil, the vault has no
// decrypted. The commands that call it hold the state directory lock, so no // long-term key until one is imported. It refuses a vault that already
// other command can create the vault between the check and the writes. // exists before writing anything: creating it again would replace its keys,
func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) { // and its secrets could no longer be decrypted. The commands that call it
// hold the state directory lock, so no other command can create the vault
// between the check and the writes.
func CreateVault(
fs afero.Fs, stateDir string, name string, mnemonic *memguard.LockedBuffer,
) (*Vault, error) {
secret.Debug("Creating new vault", "name", name, "state_dir", stateDir) secret.Debug("Creating new vault", "name", name, "state_dir", stateDir)
err := ValidateVaultName(name) err := ValidateVaultName(name)
@@ -263,7 +267,7 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) {
// Process mnemonic if available // Process mnemonic if available
derivationIndex, publicKeyHash, familyHash, err := processMnemonicForVault( derivationIndex, publicKeyHash, familyHash, err := processMnemonicForVault(
fs, stateDir, vaultDir, name) fs, stateDir, vaultDir, name, mnemonic)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -292,7 +296,10 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) {
// Create and return the vault // Create and return the vault
secret.Debug("Successfully created vault", "name", name) secret.Debug("Successfully created vault", "name", name)
return NewVault(fs, stateDir, name), nil vlt := NewVault(fs, stateDir, name)
vlt.Mnemonic = mnemonic
return vlt, nil
} }
// SelectVault selects the given vault as the current vault // SelectVault selects the given vault as the current vault
+6 -10
View File
@@ -297,14 +297,14 @@ func TestSampleHashCalculation(t *testing.T) {
} }
func TestWorkflowMismatch(t *testing.T) { func TestWorkflowMismatch(t *testing.T) {
t.Parallel()
// Create a temporary directory for testing // Create a temporary directory for testing
tempDir := t.TempDir() tempDir := t.TempDir()
fs := afero.NewOsFs() fs := afero.NewOsFs()
// Test Case 1: Create vault WITH mnemonic (like init command) // Test Case 1: Create vault WITH mnemonic (like init command)
t.Setenv("SB_SECRET_MNEMONIC", testMnemonic) _, err := vault.CreateVault(fs, tempDir, "default", testMnemonicBuffer(t))
_, err := vault.CreateVault(fs, tempDir, "default")
if err != nil { if err != nil {
t.Fatalf("Failed to create vault with mnemonic: %v", err) t.Fatalf("Failed to create vault with mnemonic: %v", err)
} }
@@ -321,19 +321,15 @@ func TestWorkflowMismatch(t *testing.T) {
metadata1.DerivationIndex, metadata1.PublicKeyHash) metadata1.DerivationIndex, metadata1.PublicKeyHash)
// Test Case 2: Create vault WITHOUT mnemonic, then import (work vault) // Test Case 2: Create vault WITHOUT mnemonic, then import (work vault)
t.Setenv("SB_SECRET_MNEMONIC", "") _, err = vault.CreateVault(fs, tempDir, "work", nil)
_, err = vault.CreateVault(fs, tempDir, "work")
if err != nil { if err != nil {
t.Fatalf("Failed to create vault without mnemonic: %v", err) t.Fatalf("Failed to create vault without mnemonic: %v", err)
} }
vault2Dir := filepath.Join(tempDir, "vaults.d", "work") vault2Dir := filepath.Join(tempDir, "vaults.d", "work")
// Simulate the vault import process // Simulate the vault import process: get the next available derivation
t.Setenv("SB_SECRET_MNEMONIC", testMnemonic) // index for this mnemonic
// Get the next available derivation index for this mnemonic
derivationIndex, err := vault.GetNextDerivationIndex(fs, tempDir, testMnemonic) derivationIndex, err := vault.GetNextDerivationIndex(fs, tempDir, testMnemonic)
if err != nil { if err != nil {
t.Fatalf("Failed to get next derivation index: %v", err) t.Fatalf("Failed to get next derivation index: %v", err)
+13 -14
View File
@@ -3,7 +3,6 @@ package vault_test
import ( import (
"testing" "testing"
"git.eeqj.de/sneak/secret/internal/secret"
"git.eeqj.de/sneak/secret/internal/vault" "git.eeqj.de/sneak/secret/internal/vault"
"github.com/awnumar/memguard" "github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
@@ -13,15 +12,13 @@ import (
// TestGetSecretVersionRejectsPathTraversal verifies that GetSecretVersion // TestGetSecretVersionRejectsPathTraversal verifies that GetSecretVersion
// validates the secret name and rejects path traversal attempts. // validates the secret name and rejects path traversal attempts.
// This is a regression test for https://git.eeqj.de/sneak/secret/issues/13 // This is a regression test for https://git.eeqj.de/sneak/secret/issues/13
//
//nolint:paralleltest // t.Setenv in parent forbids parallel subtests
func TestGetSecretVersionRejectsPathTraversal(t *testing.T) { func TestGetSecretVersionRejectsPathTraversal(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName) vlt, err := vault.CreateVault(fs, testStateDir, testVaultName,
testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
// Add a legitimate secret so the vault is set up // Add a legitimate secret so the vault is set up
@@ -41,6 +38,8 @@ func TestGetSecretVersionRejectsPathTraversal(t *testing.T) {
for _, name := range maliciousNames { for _, name := range maliciousNames {
t.Run(name, func(t *testing.T) { t.Run(name, func(t *testing.T) {
t.Parallel()
_, err := vlt.GetSecretVersion(name, "") _, err := vlt.GetSecretVersion(name, "")
require.Error(t, err, require.Error(t, err,
"GetSecretVersion should reject malicious name: %s", name) "GetSecretVersion should reject malicious name: %s", name)
@@ -53,12 +52,12 @@ func TestGetSecretVersionRejectsPathTraversal(t *testing.T) {
// TestGetSecretRejectsPathTraversal verifies GetSecret (which calls // TestGetSecretRejectsPathTraversal verifies GetSecret (which calls
// GetSecretVersion) also rejects path traversal names. // GetSecretVersion) also rejects path traversal names.
func TestGetSecretRejectsPathTraversal(t *testing.T) { func TestGetSecretRejectsPathTraversal(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName) vlt, err := vault.CreateVault(fs, testStateDir, testVaultName,
testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
_, err = vlt.GetSecret("../../../etc/passwd") _, err = vlt.GetSecret("../../../etc/passwd")
@@ -68,15 +67,13 @@ func TestGetSecretRejectsPathTraversal(t *testing.T) {
// TestGetSecretObjectRejectsPathTraversal verifies GetSecretObject // TestGetSecretObjectRejectsPathTraversal verifies GetSecretObject
// also validates names and rejects path traversal attempts. // also validates names and rejects path traversal attempts.
//
//nolint:paralleltest // t.Setenv in parent forbids parallel subtests
func TestGetSecretObjectRejectsPathTraversal(t *testing.T) { func TestGetSecretObjectRejectsPathTraversal(t *testing.T) {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Parallel()
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName) vlt, err := vault.CreateVault(fs, testStateDir, testVaultName,
testMnemonicBuffer(t))
require.NoError(t, err) require.NoError(t, err)
maliciousNames := []string{ maliciousNames := []string{
@@ -87,6 +84,8 @@ func TestGetSecretObjectRejectsPathTraversal(t *testing.T) {
for _, name := range maliciousNames { for _, name := range maliciousNames {
t.Run(name, func(t *testing.T) { t.Run(name, func(t *testing.T) {
t.Parallel()
_, err := vlt.GetSecretObject(name) _, err := vlt.GetSecretObject(name)
require.Error(t, err, "GetSecretObject should reject: %s", name) require.Error(t, err, "GetSecretObject should reject: %s", name)
require.Contains(t, err.Error(), "invalid secret name") require.Contains(t, err.Error(), "invalid secret name")
+14 -19
View File
@@ -41,14 +41,6 @@ import (
const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon " + const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon " +
"abandon abandon abandon abandon about" "abandon abandon abandon abandon about"
// envTestMnemonic is the (deliberately different) mnemonic placed in the
// environment; the vault is unlocked manually with the derived key in
// createTestVaultWithKey.
//
//nolint:dupword // BIP39-style test mnemonic intentionally repeats a word
const envTestMnemonic = "abandon abandon abandon abandon abandon abandon " +
"abandon abandon abandon about"
// Shared fixtures for white-box tests in this package. // Shared fixtures for white-box tests in this package.
const ( const (
testStateDir = "/test/state" testStateDir = "/test/state"
@@ -73,11 +65,8 @@ func addTestSecretToVault(
func createTestVaultWithKey(t *testing.T, fs afero.Fs) *Vault { func createTestVaultWithKey(t *testing.T, fs afero.Fs) *Vault {
t.Helper() t.Helper()
// Set mnemonic for testing // Create vault without a long-term key, which is set up below
t.Setenv(secret.EnvMnemonic, envTestMnemonic) vault, err := CreateVault(fs, testStateDir, "test", nil)
// Create vault
vault, err := CreateVault(fs, testStateDir, "test")
require.NoError(t, err) require.NoError(t, err)
// Derive and store long-term key from mnemonic // Derive and store long-term key from mnemonic
@@ -98,8 +87,9 @@ func createTestVaultWithKey(t *testing.T, fs afero.Fs) *Vault {
return vault return vault
} }
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
func TestVaultAddSecretCreatesVersion(t *testing.T) { func TestVaultAddSecretCreatesVersion(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Create vault with long-term key // Create vault with long-term key
@@ -137,8 +127,9 @@ func TestVaultAddSecretCreatesVersion(t *testing.T) {
assert.Equal(t, expectedValue, retrievedValue.Bytes()) assert.Equal(t, expectedValue, retrievedValue.Bytes())
} }
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
func TestVaultAddSecretMultipleVersions(t *testing.T) { func TestVaultAddSecretMultipleVersions(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Create vault with long-term key // Create vault with long-term key
@@ -174,8 +165,9 @@ func TestVaultAddSecretMultipleVersions(t *testing.T) {
assert.Equal(t, []byte("version-2"), value.Bytes()) assert.Equal(t, []byte("version-2"), value.Bytes())
} }
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
func TestVaultGetSecretVersion(t *testing.T) { func TestVaultGetSecretVersion(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Create vault with long-term key // Create vault with long-term key
@@ -220,8 +212,9 @@ func TestVaultGetSecretVersion(t *testing.T) {
require.ErrorIs(t, err, ErrVersionNotFound) require.ErrorIs(t, err, ErrVersionNotFound)
} }
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
func TestVaultVersionTimestamps(t *testing.T) { func TestVaultVersionTimestamps(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Create vault with long-term key // Create vault with long-term key
@@ -303,8 +296,9 @@ func TestVaultVersionTimestamps(t *testing.T) {
assert.Nil(t, secondVersion.Metadata.NotAfter) // Current version assert.Nil(t, secondVersion.Metadata.NotAfter) // Current version
} }
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
func TestVaultGetNonExistentVersion(t *testing.T) { func TestVaultGetNonExistentVersion(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Create vault with long-term key // Create vault with long-term key
@@ -319,8 +313,9 @@ func TestVaultGetNonExistentVersion(t *testing.T) {
assert.Contains(t, err.Error(), "not found") assert.Contains(t, err.Error(), "not found")
} }
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
func TestUpdateVersionMetadata(t *testing.T) { func TestUpdateVersionMetadata(t *testing.T) {
t.Parallel()
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Create vault with long-term key // Create vault with long-term key
+3 -1
View File
@@ -70,7 +70,9 @@ func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
secret.Debug("Creating passphrase unlocker instance", secret.Debug("Creating passphrase unlocker instance",
"unlocker_type", metadata.Type) "unlocker_type", metadata.Type)
unlocker = secret.NewPassphraseUnlocker(v.fs, unlockerDir, metadata) passphraseUnlocker := secret.NewPassphraseUnlocker(v.fs, unlockerDir, metadata)
passphraseUnlocker.Passphrase = v.UnlockPassphrase
unlocker = passphraseUnlocker
case "pgp": case "pgp":
secret.Debug("Creating PGP unlocker instance", "unlocker_type", metadata.Type) secret.Debug("Creating PGP unlocker instance", "unlocker_type", metadata.Type)
+25 -7
View File
@@ -3,12 +3,12 @@ package vault
import ( import (
"fmt" "fmt"
"log/slog" "log/slog"
"os"
"path/filepath" "path/filepath"
"filippo.io/age" "filippo.io/age"
"git.eeqj.de/sneak/secret/internal/secret" "git.eeqj.de/sneak/secret/internal/secret"
"git.eeqj.de/sneak/secret/pkg/agehd" "git.eeqj.de/sneak/secret/pkg/agehd"
"github.com/awnumar/memguard"
"github.com/spf13/afero" "github.com/spf13/afero"
) )
@@ -18,6 +18,13 @@ type Vault struct {
fs afero.Fs fs afero.Fs
stateDir string stateDir string
longTermKey *age.X25519Identity // In-memory long-term key when unlocked longTermKey *age.X25519Identity // In-memory long-term key when unlocked
// Mnemonic, when not nil, is what the long-term key is derived from
// instead of the current unlocker. The caller destroys it.
Mnemonic *memguard.LockedBuffer
// UnlockPassphrase, when not nil, is given to the current unlocker
// when that is a passphrase unlocker, which otherwise prompts for it.
// The caller destroys it.
UnlockPassphrase *memguard.LockedBuffer
} }
// NewVault creates a new Vault instance // NewVault creates a new Vault instance
@@ -56,6 +63,18 @@ func (v *Vault) ClearLongTermKey() {
v.longTermKey = nil v.longTermKey = nil
} }
// SetMnemonic sets v.Mnemonic, for code that has v only as a
// secret.VaultInterface.
func (v *Vault) SetMnemonic(mnemonic *memguard.LockedBuffer) {
v.Mnemonic = mnemonic
}
// SetUnlockPassphrase sets v.UnlockPassphrase, for code that has v only as
// a secret.VaultInterface.
func (v *Vault) SetUnlockPassphrase(passphrase *memguard.LockedBuffer) {
v.UnlockPassphrase = passphrase
}
// GetOrDeriveLongTermKey gets the long-term key from memory or derives it // GetOrDeriveLongTermKey gets the long-term key from memory or derives it
// from available sources // from available sources
func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) { func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
@@ -66,9 +85,8 @@ func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
secret.Debug("Vault is locked, attempting to unlock", "vault_name", v.Name) secret.Debug("Vault is locked, attempting to unlock", "vault_name", v.Name)
// Try to derive from environment mnemonic first if v.Mnemonic != nil {
if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" { return v.deriveLongTermKeyFromMnemonic(v.Mnemonic.String())
return v.deriveLongTermKeyFromMnemonic(envMnemonic)
} }
// No mnemonic available, try to use current unlocker // No mnemonic available, try to use current unlocker
@@ -181,9 +199,9 @@ func (v *Vault) NumSecrets() (int, error) {
// deriveLongTermKeyFromMnemonic derives the long-term key from the given // deriveLongTermKeyFromMnemonic derives the long-term key from the given
// mnemonic, verifies it against the vault metadata, and caches it in memory. // mnemonic, verifies it against the vault metadata, and caches it in memory.
func (v *Vault) deriveLongTermKeyFromMnemonic( func (v *Vault) deriveLongTermKeyFromMnemonic(
envMnemonic string, mnemonic string,
) (*age.X25519Identity, error) { ) (*age.X25519Identity, error) {
secret.Debug("Using mnemonic from environment for long-term key derivation", secret.Debug("Using mnemonic for long-term key derivation",
"vault_name", v.Name) "vault_name", v.Name)
// Load vault metadata to get the derivation index // Load vault metadata to get the derivation index
@@ -199,7 +217,7 @@ func (v *Vault) deriveLongTermKeyFromMnemonic(
return nil, fmt.Errorf("failed to load vault metadata: %w", err) return nil, fmt.Errorf("failed to load vault metadata: %w", err)
} }
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex) ltIdentity, err := agehd.DeriveIdentity(mnemonic, metadata.DerivationIndex)
if err != nil { if err != nil {
secret.Debug("Failed to derive long-term key from mnemonic", secret.Debug("Failed to derive long-term key from mnemonic",
"error", err, "vault_name", v.Name) "error", err, "vault_name", v.Name)
+19 -10
View File
@@ -27,12 +27,19 @@ const (
testPassphrase = "test-passphrase" testPassphrase = "test-passphrase"
) )
//nolint:paralleltest // t.Setenv and order-dependent subtests forbid parallel // testMnemonicBuffer returns testMnemonic in a locked buffer that is
func TestVaultOperations(t *testing.T) { // destroyed when the test ends.
// Test environment will be cleaned up automatically by t.Setenv func testMnemonicBuffer(t *testing.T) *memguard.LockedBuffer {
t.Setenv(secret.EnvMnemonic, testMnemonic) t.Helper()
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
mnemonic := memguard.NewBufferFromBytes([]byte(testMnemonic))
t.Cleanup(mnemonic.Destroy)
return mnemonic
}
//nolint:paralleltest // order-dependent subtests forbid parallel
func TestVaultOperations(t *testing.T) {
// Use in-memory filesystem // Use in-memory filesystem
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
@@ -64,7 +71,8 @@ func TestVaultOperations(t *testing.T) {
func testCreateVault(t *testing.T, fs afero.Fs) { func testCreateVault(t *testing.T, fs afero.Fs) {
t.Helper() t.Helper()
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName) vlt, err := vault.CreateVault(fs, testStateDir, testVaultName,
testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Fatalf("Failed to create vault: %v", err) t.Fatalf("Failed to create vault: %v", err)
} }
@@ -221,6 +229,8 @@ func testUnlockerOperations(t *testing.T, fs afero.Fs) {
} }
// Test vault unlocking (should happen automatically via mnemonic) // Test vault unlocking (should happen automatically via mnemonic)
vlt.Mnemonic = testMnemonicBuffer(t)
if vlt.Locked() { if vlt.Locked() {
_, err := vlt.UnlockVault() _, err := vlt.UnlockVault()
if err != nil { if err != nil {
@@ -281,15 +291,14 @@ func testUnlockerOperations(t *testing.T, fs afero.Fs) {
} }
func TestListUnlockers_SkipsMissingMetadata(t *testing.T) { func TestListUnlockers_SkipsMissingMetadata(t *testing.T) {
// Set test environment variables t.Parallel()
t.Setenv(secret.EnvMnemonic, testMnemonic)
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
// Use in-memory filesystem // Use in-memory filesystem
fs := afero.NewMemMapFs() fs := afero.NewMemMapFs()
// Create vault // Create vault
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName) vlt, err := vault.CreateVault(fs, testStateDir, testVaultName,
testMnemonicBuffer(t))
if err != nil { if err != nil {
t.Fatalf("Failed to create vault: %v", err) t.Fatalf("Failed to create vault: %v", err)
} }