Read secret environment variables once per command, then unset them (closes #60)
check / check (push) Failing after 2s

init and vault create put the mnemonic into the process environment for
vault.CreateVault to read back, so every program they ran, gpg included,
inherited it, and SB_SECRET_MNEMONIC and SB_UNLOCK_PASSPHRASE were read
at 13 places and never unset. Each command that may need them now reads
both once, in its RunE, into locked buffers on the CLI Instance, and
unsets them at once. The buffers are passed down: vault.CreateVault
takes the mnemonic, a Vault carries Mnemonic and UnlockPassphrase, and
the PGP, keychain and Secure Enclave unlocker constructors take both;
CreatePGPUnlocker sets them on the vault it loads through SetMnemonic
and SetUnlockPassphrase, new in VaultInterface. README warns against
both variables.

Model: opus-5-5
This commit is contained in:
2026-10-04 13:06:03 +00:00
parent 62967f28d0
commit 66a0714f92
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: `secret unlocker add pgp` works on Linux - 2026-10-04: `secret unlocker add pgp` works on Linux
(https://git.eeqj.de/sneak/secret/issues/88). `CreatePGPUnlocker` gets (https://git.eeqj.de/sneak/secret/issues/88). `CreatePGPUnlocker` gets
the vault's long-term key as adding a passphrase unlocker does, with the the vault's long-term key as adding a passphrase unlocker does, with the
@@ -287,8 +301,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)
} }