diff --git a/TODO.md b/TODO.md index 281cbac..fd4910f 100644 --- a/TODO.md +++ b/TODO.md @@ -25,6 +25,21 @@ Bring the repo into policy compliance in one commit: # Completed Steps +- 2026-10-03: Commands that change the state directory hold one lock + (`flock` on `lock` in the state directory; a mutex on the in-memory + test filesystem), so concurrent commands no longer lose versions or + race on the current pointers. Every file is written through + `secret.WriteFileAtomic` (temporary file, sync, rename); new + versions, new secrets and cross-vault copies are built in a + temporary directory and renamed into place, and removals rename out + of the way first, so an interrupted command leaves nothing + half-written, with one exception: an unlocker added under the + directory name of an existing one is rewritten in place, file by + file, and a crash part-way leaves it unable to open the vault. That + happens to a passphrase unlocker added to a vault that has one, and + to a PGP, keychain or Secure Enclave unlocker added on the same host + and day as another of its type + (https://git.eeqj.de/sneak/secret/issues/71). - 2026-10-02: A plain `docker build .` builds again: the size tests skip a case that needs more locked memory than the process can lock, and run every case under `script/cibuild`. The image stamps the @@ -94,8 +109,6 @@ Bring the repo into policy compliance in one commit: pgpunlocker.go:256, version.go:155); age secret key held in a plain string in cli/crypto.go:86,91,113; private keys exposed via buffer.Bytes() to GPGEncryptFunc and EncryptWithPassphrase. - - Race conditions: no file locking in vault/secrets.go:142-176; - non-atomic writes can leave the vault inconsistent. - Input validation: dots in secret names risk path traversal (vault/secrets.go:75-99); no maximum secret size (DoS). - Timing attacks: bytes.Equal passphrase compare (cli/init.go: diff --git a/internal/cli/crypto.go b/internal/cli/crypto.go index 723b472..6a1cbcc 100644 --- a/internal/cli/crypto.go +++ b/internal/cli/crypto.go @@ -72,10 +72,18 @@ func newDecryptCmd() *cobra.Command { // resolveEncryptionKey returns a secure buffer holding the age secret key // for the named secret, generating and storing a new key if the secret -// does not exist. The caller must destroy the returned buffer. +// does not exist. The caller must destroy the returned buffer. It holds the +// state directory lock itself, so that Encrypt streams its input and output +// unlocked and cannot block a secret command at the other end of a pipe. func (cli *Instance) resolveEncryptionKey( vlt *vault.Vault, secretName string, ) (*memguard.LockedBuffer, error) { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return nil, err + } + defer release() + // Check if secret exists secretObj := secret.NewSecret(vlt, secretName) diff --git a/internal/cli/generate.go b/internal/cli/generate.go index 6b15640..4234b42 100644 --- a/internal/cli/generate.go +++ b/internal/cli/generate.go @@ -155,6 +155,12 @@ func (cli *Instance) GenerateSecret( return fmt.Errorf("failed to generate random secret: %w", err) } + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Store the secret in the vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { diff --git a/internal/cli/init.go b/internal/cli/init.go index 6162c7e..b943ef2 100644 --- a/internal/cli/init.go +++ b/internal/cli/init.go @@ -103,8 +103,21 @@ func (cli *Instance) setupDefaultVault( return vlt, ltIdentity, nil } -// Init initializes the secret manager +// Init initializes the secret manager, holding the state directory lock +// while initialize runs func (cli *Instance) Init(cmd *cobra.Command) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + + return cli.initialize(cmd) +} + +// initialize creates the state directory, the default vault and its first +// unlocker +func (cli *Instance) initialize(cmd *cobra.Command) error { secret.Debug("Starting secret manager initialization") // Create state directory diff --git a/internal/cli/lock_test.go b/internal/cli/lock_test.go new file mode 100644 index 0000000..17d3875 --- /dev/null +++ b/internal/cli/lock_test.go @@ -0,0 +1,224 @@ +//nolint:testpackage // sets the unexported fields of Instance +package cli + +import ( + "io" + "path/filepath" + "strconv" + "strings" + "sync" + "testing" + "time" + + "git.eeqj.de/sneak/secret/internal/secret" + "git.eeqj.de/sneak/secret/internal/vault" + "github.com/spf13/afero" + "github.com/spf13/cobra" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// addAtOnce runs one add of the secret name per value, all at once, and +// returns their errors. +func addAtOnce( + fs afero.Fs, stateDir, name string, force bool, values []string, +) []error { + errs := make(chan error, len(values)) + + for _, value := range values { + go func() { + cli := NewCLIInstanceWithStateDir(fs, stateDir) + cli.cmd = &cobra.Command{} + cli.cmd.SetIn(strings.NewReader(value)) + + errs <- cli.AddSecret(name, force) + }() + } + + results := make([]error, 0, len(values)) + for range values { + results = append(results, <-errs) + } + + return results +} + +// numbered returns count distinct values starting with prefix. +func numbered(prefix string, count int) []string { + values := make([]string, 0, count) + for i := range count { + values = append(values, prefix+"-"+strconv.Itoa(i)) + } + + return values +} + +// TestConcurrentAddsKeepEveryVersion runs adds of one secret at once, on +// the in-memory and on the real filesystem. Without the state directory +// lock, adds of a new secret all find it absent and replace each other, and +// forced adds read the same highest version number and overwrite each +// other's version. With it they behave as if run one after another. +// +//nolint:paralleltest // t.Setenv forbids parallel subtests +func TestConcurrentAddsKeepEveryVersion(t *testing.T) { + t.Setenv(secret.EnvMnemonic, testMnemonic) + + const adds = 8 + + for _, tc := range []struct { + name string + fs afero.Fs + stateDir string + }{ + {"memory", afero.NewMemMapFs(), testStateDir}, + {"real", afero.NewOsFs(), t.TempDir()}, + } { + t.Run(tc.name, func(t *testing.T) { + _, err := vault.CreateVault(tc.fs, tc.stateDir, "default") + require.NoError(t, err) + + // One add creates the secret; the others find that it exists + created := 0 + + for _, err := range addAtOnce(tc.fs, tc.stateDir, "shared", false, + numbered("create", adds)) { + if err == nil { + created++ + } else { + require.ErrorIs(t, err, vault.ErrSecretExists) + } + } + + require.Equal(t, 1, created, "exactly one add creates the secret") + + // Every forced add stores a version of its own + for _, err := range addAtOnce(tc.fs, tc.stateDir, "shared", true, + numbered("force", adds)) { + require.NoError(t, err) + } + + vlt, err := vault.GetCurrentVault(tc.fs, tc.stateDir) + require.NoError(t, err) + + vaultDir, err := vlt.GetDirectory() + require.NoError(t, err) + + versions, err := secret.ListVersions(tc.fs, + filepath.Join(vaultDir, "secrets.d", "shared")) + require.NoError(t, err) + require.Len(t, versions, adds+1, "one version per successful add") + + values := make(map[string]bool, len(versions)) + + for _, version := range versions { + value, err := vlt.GetSecretVersion("shared", version) + require.NoError(t, err) + + values[string(value)] = true + } + + assert.Len(t, values, adds+1, "every add stored its own value") + }) + } +} + +// readNotifier passes reads through to Reader and closes reading at the +// first one. +type readNotifier struct { + io.Reader + + reading chan struct{} + once sync.Once +} + +func (r *readNotifier) Read(p []byte) (int, error) { + r.once.Do(func() { close(r.reading) }) + + return r.Reader.Read(p) +} + +// TestEncryptPipedIntoAdd runs `secret encrypt key | secret add name` in +// one process, starting encrypt once add is reading its input. Had add +// taken the state directory lock before reading, it would hold the lock +// while waiting for encrypt's output, and encrypt would wait for the lock +// to store its key: neither would finish. +func TestEncryptPipedIntoAdd(t *testing.T) { + t.Setenv(secret.EnvMnemonic, testMnemonic) + + fs := afero.NewMemMapFs() + _, err := vault.CreateVault(fs, testStateDir, "default") + require.NoError(t, err) + require.NoError(t, afero.WriteFile(fs, "/plaintext", []byte("piped"), 0o600)) + + pipeReader, pipeWriter := io.Pipe() + // If the test gives up, this makes add's read fail, so that both + // commands return and release the lock the other tests use + t.Cleanup(func() { _ = pipeReader.Close() }) + + const commands = 2 + + input := &readNotifier{Reader: pipeReader, reading: make(chan struct{})} + results := make(chan error, commands) + + go func() { + add := NewCLIInstanceWithStateDir(fs, testStateDir) + add.cmd = &cobra.Command{} + add.cmd.SetIn(input) + + results <- add.AddSecret("encrypted", false) + }() + + go func() { + <-input.reading + + encrypt := NewCLIInstanceWithStateDir(fs, testStateDir) + encrypt.cmd = &cobra.Command{} + encrypt.cmd.SetOut(pipeWriter) + + err := encrypt.Encrypt("key", "/plaintext", "") + // Ends add's input, as the end of the pipe does + _ = pipeWriter.CloseWithError(err) + + results <- err + }() + + timeout := time.After(10 * time.Second) + + for range commands { + select { + case err := <-results: + require.NoError(t, err) + case <-timeout: + t.Fatal("secret encrypt piped into secret add never finished") + } + } +} + +// TestFailedCommandReleasesLock checks that a command failing after it +// took the state directory lock leaves the lock free for the next command. +func TestFailedCommandReleasesLock(t *testing.T) { + t.Parallel() + + fs := afero.NewMemMapFs() + cli := NewCLIInstanceWithStateDir(fs, testStateDir) + + // Fails once it holds the lock: there is no current vault + err := cli.RemoveSecret(&cobra.Command{}, "missing", false) + require.Error(t, err) + + taken := make(chan func(), 1) + + go func() { + release, err := vault.LockStateDir(fs, testStateDir) + if assert.NoError(t, err) { + taken <- release + } + }() + + select { + case release := <-taken: + release() + case <-time.After(10 * time.Second): + t.Fatal("the failed command left the state directory locked") + } +} diff --git a/internal/cli/secrets.go b/internal/cli/secrets.go index f62ecf0..5574f77 100644 --- a/internal/cli/secrets.go +++ b/internal/cli/secrets.go @@ -377,6 +377,15 @@ func (cli *Instance) AddSecret(secretName string, force bool) error { valueBuffer := combineBuffers(buffers, totalSize) defer valueBuffer.Destroy() + // Locked only now that stdin has been read: in `secret encrypt key | + // secret add name`, holding the lock while reading would leave each + // command waiting for the other. + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Add the secret to the vault secret.Debug("Calling vault.AddSecret", "secret_name", secretName, "value_length", valueBuffer.Size(), "force", force) @@ -635,6 +644,14 @@ func (cli *Instance) ImportSecret( valueBuffer := combineBuffers(buffers, totalSize) defer valueBuffer.Destroy() + // Locked only now that the file has been read, as in AddSecret: the + // file may be a pipe written by another secret command. + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Store the secret in the vault err = vlt.AddSecret(secretName, valueBuffer, force) if err != nil { @@ -649,6 +666,12 @@ func (cli *Instance) ImportSecret( // RemoveSecret removes a secret from the vault func (cli *Instance) RemoveSecret(cmd *cobra.Command, secretName string, _ bool) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Get current vault currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -683,7 +706,7 @@ func (cli *Instance) RemoveSecret(cmd *cobra.Command, secretName string, _ bool) } // Remove the secret directory - err = cli.fs.RemoveAll(secretDir) + err = secret.RemoveDirAtomic(cli.fs, secretDir) if err != nil { return fmt.Errorf("failed to remove secret: %w", err) } @@ -698,6 +721,12 @@ func (cli *Instance) RemoveSecret(cmd *cobra.Command, secretName string, _ bool) func (cli *Instance) MoveSecret( cmd *cobra.Command, source, dest string, force bool, ) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Parse source and destination srcVaultName, srcSecretName, srcQualified := ParseVaultSecretRef(source) destVaultName, destSecretName, destQualified := ParseVaultSecretRef(dest) @@ -792,7 +821,7 @@ func (cli *Instance) moveSecretWithinVault( return fmt.Errorf("secret '%s' %w", dest, errSecretExistsNoForce) } - err = cli.fs.RemoveAll(destDir) + err = secret.RemoveDirAtomic(cli.fs, destDir) if err != nil { return fmt.Errorf("failed to remove existing destination: %w", err) } @@ -872,7 +901,7 @@ func (cli *Instance) moveSecretCrossVault( } // Delete source secret - err = cli.fs.RemoveAll(srcSecretDir) + err = secret.RemoveDirAtomic(cli.fs, srcSecretDir) if err != nil { // Copy succeeded but delete failed - warn but don't fail cmd.Printf("Warning: copied secret but failed to remove source: %v\n", err) diff --git a/internal/cli/unlockers.go b/internal/cli/unlockers.go index 593f462..ba5139e 100644 --- a/internal/cli/unlockers.go +++ b/internal/cli/unlockers.go @@ -534,6 +534,12 @@ func (cli *Instance) printUnlockersTable(unlockers []UnlockerInfo) error { // UnlockersAdd adds a new unlocker func (cli *Instance) UnlockersAdd(unlockerType string, cmd *cobra.Command) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + switch unlockerType { case unlockerTypePassphrase: return cli.addPassphraseUnlocker(cmd) @@ -714,6 +720,12 @@ func (cli *Instance) addPGPUnlocker(cmd *cobra.Command) error { func (cli *Instance) UnlockersRemove( unlockerID string, force bool, cmd *cobra.Command, ) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -763,6 +775,12 @@ func (cli *Instance) UnlockersRemove( // UnlockerSelect selects an unlocker as current func (cli *Instance) UnlockerSelect(unlockerID string) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { diff --git a/internal/cli/vault.go b/internal/cli/vault.go index 63781d4..32ce02a 100644 --- a/internal/cli/vault.go +++ b/internal/cli/vault.go @@ -267,6 +267,12 @@ func resolvePassphrase() (*memguard.LockedBuffer, error) { func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error { secret.Debug("Creating new vault", "name", name, "state_dir", cli.stateDir) + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Get or prompt for mnemonic var mnemonicStr string @@ -354,7 +360,13 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error { // SelectVault selects a vault as the current one func (cli *Instance) SelectVault(cmd *cobra.Command, name string) error { - err := vault.SelectVault(cli.fs, cli.stateDir, name) + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + + err = vault.SelectVault(cli.fs, cli.stateDir, name) if err != nil { return err } @@ -442,8 +454,21 @@ func updateVaultImportMetadata( return nil } -// VaultImport imports a mnemonic into a specific vault +// VaultImport imports a mnemonic into a specific vault, holding the state +// directory lock while importMnemonic runs func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + + return cli.importMnemonic(cmd, vaultName) +} + +// importMnemonic gives the vault a long-term key derived from the mnemonic +// and a passphrase unlocker +func (cli *Instance) importMnemonic(cmd *cobra.Command, vaultName string) error { secret.Debug("Importing mnemonic into vault", "vault_name", vaultName, "state_dir", cli.stateDir) @@ -478,7 +503,7 @@ func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error { secret.Debug("Storing long-term public key", "pubkey", ltPublicKey, "vault_dir", vaultDir) - err = afero.WriteFile(cli.fs, pubKeyPath, []byte(ltPublicKey), secret.FilePerms) + err = secret.WriteFileAtomic(cli.fs, pubKeyPath, []byte(ltPublicKey)) if err != nil { return fmt.Errorf("failed to store long-term public key: %w", err) } @@ -577,6 +602,12 @@ func (cli *Instance) switchAwayFromVault( // RemoveVault removes a vault with safety checks func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Get list of all vaults vaults, err := vault.ListVaults(cli.fs, cli.stateDir) if err != nil { @@ -626,7 +657,7 @@ func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) er } // Remove the vault directory - err = cli.fs.RemoveAll(vaultDir) + err = secret.RemoveDirAtomic(cli.fs, vaultDir) if err != nil { return fmt.Errorf("failed to remove vault directory: %w", err) } diff --git a/internal/cli/version.go b/internal/cli/version.go index 17f113f..35db7cf 100644 --- a/internal/cli/version.go +++ b/internal/cli/version.go @@ -239,6 +239,12 @@ func formatVersionTime(t *time.Time) string { func (cli *Instance) PromoteVersion( cmd *cobra.Command, secretName string, version string, ) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -282,6 +288,12 @@ func (cli *Instance) PromoteVersion( func (cli *Instance) RemoveVersion( cmd *cobra.Command, secretName string, version string, ) error { + release, err := vault.LockStateDir(cli.fs, cli.stateDir) + if err != nil { + return err + } + defer release() + // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -333,7 +345,7 @@ func (cli *Instance) RemoveVersion( } // Remove the version directory - err = cli.fs.RemoveAll(versionDir) + err = secret.RemoveDirAtomic(cli.fs, versionDir) if err != nil { return fmt.Errorf("failed to remove version: %w", err) } diff --git a/internal/secret/atomic.go b/internal/secret/atomic.go new file mode 100644 index 0000000..f9c85fb --- /dev/null +++ b/internal/secret/atomic.go @@ -0,0 +1,86 @@ +package secret + +import ( + "fmt" + "path/filepath" + + "github.com/spf13/afero" +) + +// WriteFileAtomic replaces the file at path with data so that a reader, or +// a crash at any moment, finds either the old content or the new, never a +// partial file. The data goes into a temporary file that afero.TempFile +// creates with mode 0600 in the same directory (a rename is only atomic +// within one filesystem), is synced to disk, and is renamed over path. The +// temporary file is removed if any step fails. +func WriteFileAtomic(fs afero.Fs, path string, data []byte) error { + tmp, err := afero.TempFile(fs, filepath.Dir(path), + "."+filepath.Base(path)+".tmp-*") + if err != nil { + return fmt.Errorf("failed to create temporary file for %s: %w", path, err) + } + + _, err = tmp.Write(data) + if err == nil { + err = tmp.Sync() + } + + closeErr := tmp.Close() + if err == nil { + err = closeErr + } + + if err == nil { + err = fs.Rename(tmp.Name(), path) + } + + if err != nil { + _ = fs.Remove(tmp.Name()) + + return fmt.Errorf("failed to write %s: %w", path, err) + } + + return nil +} + +// TempDirFor creates an empty temporary directory in which to build the +// directory target before renaming it into place, or into which to move +// target before deleting it. It is made in target's grandparent: on the +// same filesystem, so the rename is atomic, and outside target's parent, +// the directory that is listed to find vaults, secrets, versions and +// unlockers, so one left behind by a crash is never taken for one of them. +// Its name leaves out target's, which may already be as long as a file name +// can be. +func TempDirFor(fs afero.Fs, target string) (string, error) { + dir, err := afero.TempDir(fs, filepath.Dir(filepath.Dir(target)), ".tmp-") + if err != nil { + return "", fmt.Errorf( + "failed to create temporary directory for %s: %w", target, err) + } + + return dir, nil +} + +// RemoveDirAtomic deletes the directory dir so that it disappears in one +// rename: dir is moved into a new directory from TempDirFor, which is then +// deleted. A crash part-way leaves only that temporary directory behind. +func RemoveDirAtomic(fs afero.Fs, dir string) error { + tmp, err := TempDirFor(fs, dir) + if err != nil { + return err + } + + err = fs.Rename(dir, filepath.Join(tmp, filepath.Base(dir))) + if err != nil { + _ = fs.Remove(tmp) + + return fmt.Errorf("failed to remove %s: %w", dir, err) + } + + err = fs.RemoveAll(tmp) + if err != nil { + return fmt.Errorf("failed to remove %s: %w", dir, err) + } + + return nil +} diff --git a/internal/secret/atomic_test.go b/internal/secret/atomic_test.go new file mode 100644 index 0000000..1c8b786 --- /dev/null +++ b/internal/secret/atomic_test.go @@ -0,0 +1,506 @@ +package secret_test + +import ( + "errors" + "os" + "path/filepath" + "strings" + "testing" + + "filippo.io/age" + "git.eeqj.de/sneak/secret/internal/secret" + "git.eeqj.de/sneak/secret/internal/vault" + "github.com/awnumar/memguard" + "github.com/spf13/afero" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +var errInjected = errors.New("injected failure") + +// The kinds of change hookFs passes to before. +const ( + opCreate = "create" + opOpen = "open" + opMkdir = "mkdir" + opRemove = "remove" + opRename = "rename" +) + +// currentFile is the file in a secret's directory that names its current +// version. +const currentFile = "current" + +// hookFs passes every call through to Fs, but first calls before for each +// call that changes the filesystem, with the path it changes (the new path, +// for Rename). A test uses before to inspect the tree at every point where +// a crash could stop the code under test, or returns an error from it to +// make that call fail. +type hookFs struct { + afero.Fs + + before func(op, path string) error +} + +//nolint:ireturn // implements afero.Fs +func (h hookFs) Create(name string) (afero.File, error) { + err := h.before(opCreate, name) + if err != nil { + return nil, err + } + + return h.Fs.Create(name) +} + +//nolint:ireturn // implements afero.Fs +func (h hookFs) OpenFile( + name string, flag int, perm os.FileMode, +) (afero.File, error) { + err := h.before(opOpen, name) + if err != nil { + return nil, err + } + + return h.Fs.OpenFile(name, flag, perm) +} + +func (h hookFs) Mkdir(name string, perm os.FileMode) error { + err := h.before(opMkdir, name) + if err != nil { + return err + } + + return h.Fs.Mkdir(name, perm) +} + +func (h hookFs) MkdirAll(path string, perm os.FileMode) error { + err := h.before(opMkdir, path) + if err != nil { + return err + } + + return h.Fs.MkdirAll(path, perm) +} + +func (h hookFs) Remove(name string) error { + err := h.before(opRemove, name) + if err != nil { + return err + } + + return h.Fs.Remove(name) +} + +func (h hookFs) RemoveAll(path string) error { + err := h.before(opRemove, path) + if err != nil { + return err + } + + return h.Fs.RemoveAll(path) +} + +func (h hookFs) Rename(oldname, newname string) error { + err := h.before(opRename, newname) + if err != nil { + return err + } + + return h.Fs.Rename(oldname, newname) +} + +// testFilesystem is a filesystem to run a test on, with a directory in it +// to work in. +type testFilesystem struct { + name string + open func(t *testing.T) (afero.Fs, string) +} + +// testFilesystems are the in-memory filesystem that most tests use and the +// real one: every rename-based guarantee is checked on both. +// +//nolint:gochecknoglobals // read-only table shared by the tests below +var testFilesystems = []testFilesystem{ + {"memory", func(*testing.T) (afero.Fs, string) { + return afero.NewMemMapFs(), "/test" + }}, + {"real", func(t *testing.T) (afero.Fs, string) { + t.Helper() + + return afero.NewOsFs(), t.TempDir() + }}, +} + +// dirNames lists the names in dir. +func dirNames(t *testing.T, fs afero.Fs, dir string) []string { + t.Helper() + + entries, err := afero.ReadDir(fs, dir) + require.NoError(t, err) + + names := make([]string, 0, len(entries)) + for _, entry := range entries { + names = append(names, entry.Name()) + } + + return names +} + +// writeLongTermKey gives the test vault under stateDir a new long-term key +// and returns it. +func writeLongTermKey( + t *testing.T, fs afero.Fs, stateDir string, +) *age.X25519Identity { + t.Helper() + + vault := &MockVersionVault{Name: testVaultName, fs: fs, stateDir: stateDir} + + vaultDir, err := vault.GetDirectory() + require.NoError(t, err) + require.NoError(t, fs.MkdirAll(vaultDir, 0o700)) + + ltIdentity, err := age.GenerateX25519Identity() + require.NoError(t, err) + require.NoError(t, afero.WriteFile(fs, filepath.Join(vaultDir, "pub.age"), + []byte(ltIdentity.Recipient().String()), 0o600)) + + return ltIdentity +} + +// newVaultWithSecret creates the vault name under stateDir from the test +// mnemonic, with a secret "shared" in it that holds value. +func newVaultWithSecret( + t *testing.T, fs afero.Fs, stateDir, name, value string, +) *vault.Vault { + t.Helper() + + vlt, err := vault.CreateVault(fs, stateDir, name) + require.NoError(t, err) + + buffer := memguard.NewBufferFromBytes([]byte(value)) + defer buffer.Destroy() + + require.NoError(t, vlt.AddSecret("shared", buffer, false)) + + return vlt +} + +func TestWriteFileAtomicReplacesFile(t *testing.T) { + t.Parallel() + + for _, tfs := range testFilesystems { + t.Run(tfs.name, func(t *testing.T) { + t.Parallel() + + fs, dir := tfs.open(t) + path := filepath.Join(dir, currentFile) + + require.NoError(t, secret.WriteFileAtomic(fs, path, []byte("old"))) + require.NoError(t, secret.WriteFileAtomic(fs, path, []byte("new"))) + + data, err := afero.ReadFile(fs, path) + require.NoError(t, err) + assert.Equal(t, "new", string(data)) + + info, err := fs.Stat(path) + require.NoError(t, err) + assert.Equal(t, secret.FilePerms, info.Mode().Perm()) + + // No temporary file is left next to it + assert.Equal(t, []string{currentFile}, dirNames(t, fs, dir)) + }) + } +} + +func TestWriteFileAtomicFailureKeepsOldFile(t *testing.T) { + t.Parallel() + + for _, tfs := range testFilesystems { + t.Run(tfs.name, func(t *testing.T) { + t.Parallel() + + base, dir := tfs.open(t) + path := filepath.Join(dir, currentFile) + require.NoError(t, secret.WriteFileAtomic(base, path, []byte("old"))) + + fs := hookFs{Fs: base, before: func(op, _ string) error { + if op == opRename { + return errInjected + } + + return nil + }} + + err := secret.WriteFileAtomic(fs, path, []byte("new")) + require.ErrorIs(t, err, errInjected) + + data, err := afero.ReadFile(base, path) + require.NoError(t, err) + assert.Equal(t, "old", string(data)) + + // The temporary file is removed again + assert.Equal(t, []string{currentFile}, dirNames(t, base, dir)) + }) + } +} + +// TestRemoveDirAtomic checks that RemoveDirAtomic deletes nothing where the +// directory stands, which a crash could stop half-way, and that it leaves +// nothing behind. +func TestRemoveDirAtomic(t *testing.T) { + t.Parallel() + + for _, tfs := range testFilesystems { + t.Run(tfs.name, func(t *testing.T) { + t.Parallel() + + base, dir := tfs.open(t) + listed := filepath.Join(dir, "secrets.d") + target := filepath.Join(listed, "doomed") + + require.NoError(t, base.MkdirAll(filepath.Join(target, "versions"), 0o700)) + require.NoError(t, secret.WriteFileAtomic(base, + filepath.Join(target, currentFile), []byte("20231216.001"))) + + fs := hookFs{Fs: base, before: func(op, path string) error { + if op == opRemove && strings.HasPrefix(path, target) { + t.Errorf("deleted %s where it stands", path) + } + + return nil + }} + + require.NoError(t, secret.RemoveDirAtomic(fs, target)) + + // Gone, and no temporary directory is left in the directory + // that is listed or in the one above it + assert.Empty(t, dirNames(t, base, listed)) + assert.Equal(t, []string{"secrets.d"}, dirNames(t, base, dir)) + }) + } +} + +// TestLongestNames adds a secret to a vault and removes the vault, both +// named with 255 bytes, the most a file name may have, on the real +// filesystem: the temporary directories they use must fit that limit too. +func TestLongestNames(t *testing.T) { + t.Setenv(secret.EnvMnemonic, testMnemonic) + + const longestName = 255 + + fs := afero.NewOsFs() + name := strings.Repeat("a", longestName) + + vlt, err := vault.CreateVault(fs, t.TempDir(), name) + require.NoError(t, err) + + value := memguard.NewBufferFromBytes([]byte("long")) + defer value.Destroy() + + require.NoError(t, vlt.AddSecret(name, value, false)) + + got, err := vlt.GetSecret(name) + require.NoError(t, err) + assert.Equal(t, "long", string(got)) + + vaultDir, err := vlt.GetDirectory() + require.NoError(t, err) + require.NoError(t, secret.RemoveDirAtomic(fs, vaultDir)) + assert.NoDirExists(t, vaultDir) +} + +// TestForcedCopyKeepsDestinationUntilReplaced copies a secret over one in +// another vault, as a forced move between vaults does, and makes the last +// step that completes the copy fail. The secret it was to replace must +// still be there unchanged: it may go only once its replacement is whole. +// +//nolint:paralleltest // t.Setenv forbids t.Parallel +func TestForcedCopyKeepsDestinationUntilReplaced(t *testing.T) { + t.Setenv(secret.EnvMnemonic, testMnemonic) + + for _, tfs := range testFilesystems { + t.Run(tfs.name, func(t *testing.T) { + base, stateDir := tfs.open(t) + src := newVaultWithSecret(t, base, stateDir, "source", "new") + dest := newVaultWithSecret(t, base, stateDir, "dest", "old") + + // The copy is complete once its current file is written + fs := hookFs{Fs: base, before: func(op, path string) error { + if op == opRename && filepath.Base(path) == currentFile { + return errInjected + } + + return nil + }} + + err := vault.NewVault(fs, stateDir, "dest"). + CopySecretAllVersions(src, "shared", "shared", true) + require.ErrorIs(t, err, errInjected) + + value, err := dest.GetSecret("shared") + require.NoError(t, err) + assert.Equal(t, "old", string(value)) + }) + } +} + +// TestTempDirsStayOutOfListings adds a version, adds a secret, copies a +// secret over another and removes one, and checks that none of them makes a +// directory directly in secrets.d or in a versions directory. Those are +// listed to find secrets and versions, so a temporary directory made there +// would be listed while half-built, and one left by a crash would stay. +// +//nolint:paralleltest // t.Setenv forbids t.Parallel +func TestTempDirsStayOutOfListings(t *testing.T) { + t.Setenv(secret.EnvMnemonic, testMnemonic) + + for _, tfs := range testFilesystems { + t.Run(tfs.name, func(t *testing.T) { + base, stateDir := tfs.open(t) + newVaultWithSecret(t, base, stateDir, "default", "first") + + fs := hookFs{Fs: base, before: func(op, path string) error { + parent := filepath.Base(filepath.Dir(path)) + if op == opMkdir && (parent == "secrets.d" || parent == "versions") { + t.Errorf("made %s where it is listed", path) + } + + return nil + }} + vlt := vault.NewVault(fs, stateDir, "default") + + value := memguard.NewBufferFromBytes([]byte("second")) + defer value.Destroy() + + require.NoError(t, vlt.AddSecret("shared", value, true)) + require.NoError(t, vlt.AddSecret("other", value, false)) + require.NoError(t, vlt.CopySecretAllVersions(vlt, "shared", "other", true)) + + vaultDir, err := vlt.GetDirectory() + require.NoError(t, err) + require.NoError(t, secret.RemoveDirAtomic(fs, + filepath.Join(vaultDir, "secrets.d", "shared"))) + }) + } +} + +// TestVersionSaveIsWholeOrAbsent checks, before every change Save makes and +// once after it returns, that the version directory either does not exist +// or holds all of its files: a crash at any point leaves no version that +// cannot be decrypted. +func TestVersionSaveIsWholeOrAbsent(t *testing.T) { + t.Parallel() + + for _, tfs := range testFilesystems { + t.Run(tfs.name, func(t *testing.T) { + t.Parallel() + + base, stateDir := tfs.open(t) + ltIdentity := writeLongTermKey(t, base, stateDir) + + var versionDir string + + checkVersionDir := func(string, string) error { + exists, err := afero.DirExists(base, versionDir) + require.NoError(t, err) + + if exists { + assert.ElementsMatch(t, + []string{"pub.age", "value.age", "priv.age", "metadata.age"}, + dirNames(t, base, versionDir), + "version directory visible before it was complete") + } + + return nil + } + + fs := hookFs{Fs: base, before: checkVersionDir} + vault := &MockVersionVault{Name: testVaultName, fs: fs, stateDir: stateDir} + sv := secret.NewVersion(vault, "test/secret", "20231215.001") + versionDir = sv.Directory + + value := memguard.NewBufferFromBytes([]byte("whole or nothing")) + defer value.Destroy() + + require.NoError(t, sv.Save(value)) + require.NoError(t, checkVersionDir("", "")) + + got, err := sv.GetValue(ltIdentity) + require.NoError(t, err) + + defer got.Destroy() + + assert.Equal(t, "whole or nothing", got.String()) + }) + } +} + +// TestVersionSaveFailureLeavesNothing makes the write of the encrypted +// private key fail, after the value has been written, and checks that +// neither the version nor its temporary directory is left behind. +func TestVersionSaveFailureLeavesNothing(t *testing.T) { + t.Parallel() + + for _, tfs := range testFilesystems { + t.Run(tfs.name, func(t *testing.T) { + t.Parallel() + + base, stateDir := tfs.open(t) + writeLongTermKey(t, base, stateDir) + + fs := hookFs{Fs: base, before: func(op, path string) error { + if op == opRename && filepath.Base(path) == "priv.age" { + return errInjected + } + + return nil + }} + vault := &MockVersionVault{Name: testVaultName, fs: fs, stateDir: stateDir} + sv := secret.NewVersion(vault, "test/secret", "20231215.001") + + value := memguard.NewBufferFromBytes([]byte("never stored")) + defer value.Destroy() + + require.ErrorIs(t, sv.Save(value), errInjected) + + // The secret directory holds only the empty versions directory + versionsDir := filepath.Dir(sv.Directory) + assert.Equal(t, []string{"versions"}, + dirNames(t, base, filepath.Dir(versionsDir))) + assert.Empty(t, dirNames(t, base, versionsDir)) + }) + } +} + +// TestSetCurrentVersionNeverMissing checks, before every change +// SetCurrentVersion makes, that the current file exists: a reader or a crash +// never finds the secret without a current version. +func TestSetCurrentVersionNeverMissing(t *testing.T) { + t.Parallel() + + for _, tfs := range testFilesystems { + t.Run(tfs.name, func(t *testing.T) { + t.Parallel() + + base, dir := tfs.open(t) + secretDir := filepath.Join(dir, "secret") + require.NoError(t, base.MkdirAll(secretDir, 0o700)) + require.NoError(t, secret.SetCurrentVersion(base, secretDir, "20231216.001")) + + currentPath := filepath.Join(secretDir, currentFile) + fs := hookFs{Fs: base, before: func(string, string) error { + exists, err := afero.Exists(base, currentPath) + require.NoError(t, err) + assert.True(t, exists, "current is missing") + + return nil + }} + + require.NoError(t, secret.SetCurrentVersion(fs, secretDir, "20231216.002")) + + version, err := secret.GetCurrentVersion(base, secretDir) + require.NoError(t, err) + assert.Equal(t, "20231216.002", version) + }) + } +} diff --git a/internal/secret/keychainunlocker.go b/internal/secret/keychainunlocker.go index c544214..d0f3367 100644 --- a/internal/secret/keychainunlocker.go +++ b/internal/secret/keychainunlocker.go @@ -195,7 +195,7 @@ func (k *KeychainUnlocker) Remove() error { // Step 3: Remove directory Debug("Removing keychain unlocker directory", "directory", k.Directory) - if err := k.fs.RemoveAll(k.Directory); err != nil { + if err := RemoveDirAtomic(k.fs, k.Directory); err != nil { Debug("Failed to remove keychain unlocker directory", "error", err, "directory", k.Directory) return fmt.Errorf("failed to remove keychain unlocker directory: %w", err) @@ -373,7 +373,7 @@ func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, er // Step 3: Store age recipient as plaintext ageRecipient := ageIdentity.Recipient().String() recipientPath := filepath.Join(unlockerDir, "pub.txt") - if err := afero.WriteFile(fs, recipientPath, []byte(ageRecipient), FilePerms); err != nil { + if err := WriteFileAtomic(fs, recipientPath, []byte(ageRecipient)); err != nil { return nil, fmt.Errorf("failed to write age recipient: %w", err) } @@ -392,7 +392,7 @@ func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, er } agePrivKeyPath := filepath.Join(unlockerDir, "priv.age") - if err := afero.WriteFile(fs, agePrivKeyPath, encryptedAgePrivKey, FilePerms); err != nil { + if err := WriteFileAtomic(fs, agePrivKeyPath, encryptedAgePrivKey); err != nil { return nil, fmt.Errorf("failed to write encrypted age private key: %w", err) } @@ -411,7 +411,7 @@ func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, er // Write encrypted long-term private key ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age") - if err := afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKeyToAge, FilePerms); err != nil { + if err := WriteFileAtomic(fs, ltPrivKeyPath, encryptedLtPrivKeyToAge); err != nil { return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err) } @@ -451,9 +451,9 @@ func CreateKeychainUnlocker(fs afero.Fs, stateDir string) (*KeychainUnlocker, er return nil, fmt.Errorf("failed to marshal unlocker metadata: %w", err) } - if err := afero.WriteFile(fs, + if err := WriteFileAtomic(fs, filepath.Join(unlockerDir, "unlocker-metadata.json"), - metadataBytes, FilePerms); err != nil { + metadataBytes); err != nil { return nil, fmt.Errorf("failed to write unlocker metadata: %w", err) } diff --git a/internal/secret/passphraseunlocker.go b/internal/secret/passphraseunlocker.go index 9711d9a..19ebc43 100644 --- a/internal/secret/passphraseunlocker.go +++ b/internal/secret/passphraseunlocker.go @@ -127,7 +127,7 @@ func (p *PassphraseUnlocker) Remove() error { // For passphrase unlockers, we just need to remove the directory // No external resources (like keychain items) to clean up - err := p.fs.RemoveAll(p.Directory) + err := RemoveDirAtomic(p.fs, p.Directory) if err != nil { return fmt.Errorf("failed to remove passphrase unlocker directory: %w", err) } diff --git a/internal/secret/pgpunlocker.go b/internal/secret/pgpunlocker.go index c6fcbc3..2849049 100644 --- a/internal/secret/pgpunlocker.go +++ b/internal/secret/pgpunlocker.go @@ -172,7 +172,7 @@ func (p *PGPUnlocker) GetID() string { func (p *PGPUnlocker) Remove() error { // For PGP unlockers, we just need to remove the directory // No external resources (like keychain items) to clean up - err := p.fs.RemoveAll(p.Directory) + err := RemoveDirAtomic(p.fs, p.Directory) if err != nil { return fmt.Errorf("failed to remove PGP unlocker directory: %w", err) } @@ -275,7 +275,7 @@ func CreatePGPUnlocker( ageRecipient := ageIdentity.Recipient().String() recipientPath := filepath.Join(unlockerDir, "pub.txt") - err = afero.WriteFile(fs, recipientPath, []byte(ageRecipient), FilePerms) + err = WriteFileAtomic(fs, recipientPath, []byte(ageRecipient)) if err != nil { return nil, fmt.Errorf("failed to write age recipient: %w", err) } @@ -298,7 +298,7 @@ func CreatePGPUnlocker( // Write encrypted long-term private key ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age") - err = afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKeyToAge, FilePerms) + err = WriteFileAtomic(fs, ltPrivKeyPath, encryptedLtPrivKeyToAge) if err != nil { return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err) } @@ -315,7 +315,7 @@ func CreatePGPUnlocker( agePrivKeyPath := filepath.Join(unlockerDir, "priv.age.gpg") - err = afero.WriteFile(fs, agePrivKeyPath, encryptedAgePrivKey, FilePerms) + err = WriteFileAtomic(fs, agePrivKeyPath, encryptedAgePrivKey) if err != nil { return nil, fmt.Errorf("failed to write encrypted age private key: %w", err) } @@ -357,9 +357,8 @@ func writePGPUnlockerMetadata( return nil, fmt.Errorf("failed to marshal unlocker metadata: %w", err) } - err = afero.WriteFile(fs, - filepath.Join(unlockerDir, "unlocker-metadata.json"), - metadataBytes, FilePerms) + err = WriteFileAtomic(fs, + filepath.Join(unlockerDir, "unlocker-metadata.json"), metadataBytes) if err != nil { return nil, fmt.Errorf("failed to write unlocker metadata: %w", err) } diff --git a/internal/secret/seunlocker_darwin.go b/internal/secret/seunlocker_darwin.go index 9d92717..69e8ea8 100644 --- a/internal/secret/seunlocker_darwin.go +++ b/internal/secret/seunlocker_darwin.go @@ -148,7 +148,7 @@ func (s *SecureEnclaveUnlocker) Remove() error { } Debug("Removing SE unlocker directory", "directory", s.Directory) - if err := s.fs.RemoveAll(s.Directory); err != nil { + if err := RemoveDirAtomic(s.fs, s.Directory); err != nil { return fmt.Errorf("failed to remove SE unlocker directory: %w", err) } @@ -271,7 +271,7 @@ func CreateSecureEnclaveUnlocker( // Write SE-encrypted long-term key ltKeyPath := filepath.Join(unlockerDir, seLongtermFilename) - if err := afero.WriteFile(fs, ltKeyPath, encryptedLtKey, FilePerms); err != nil { + if err := WriteFileAtomic(fs, ltKeyPath, encryptedLtKey); err != nil { return nil, fmt.Errorf( "failed to write SE-encrypted long-term key: %w", err, @@ -295,7 +295,7 @@ func CreateSecureEnclaveUnlocker( } metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") - if err := afero.WriteFile(fs, metadataPath, metadataBytes, FilePerms); err != nil { + if err := WriteFileAtomic(fs, metadataPath, metadataBytes); err != nil { return nil, fmt.Errorf("failed to write metadata: %w", err) } diff --git a/internal/secret/version.go b/internal/secret/version.go index 39efed2..0d0d369 100644 --- a/internal/secret/version.go +++ b/internal/secret/version.go @@ -131,7 +131,10 @@ func GenerateVersionName(fs afero.Fs, secretDir string) (string, error) { return fmt.Sprintf("%s.%03d", today, newSerial), nil } -// Save saves the version metadata and value +// Save saves the version metadata and value. The files are written into a +// temporary directory that is renamed to sv.Directory once all of them are +// complete, so the version directory is either whole or absent, even if the +// process dies part-way. func (sv *Version) Save(value *memguard.LockedBuffer) error { if value == nil { return errNilValueBuffer @@ -145,14 +148,22 @@ func (sv *Version) Save(value *memguard.LockedBuffer) error { fs := sv.vault.GetFilesystem() - // Create version directory - err := fs.MkdirAll(sv.Directory, DirPerms) + // Create the versions directory the finished version is renamed into + err := fs.MkdirAll(filepath.Dir(sv.Directory), DirPerms) if err != nil { - Debug("Failed to create version directory", "error", err, "dir", sv.Directory) + Debug("Failed to create versions directory", "error", err, "dir", sv.Directory) - return fmt.Errorf("failed to create version directory: %w", err) + return fmt.Errorf("failed to create versions directory: %w", err) } + tmpDir, err := TempDirFor(fs, sv.Directory) + if err != nil { + return err + } + + // Once the rename below has moved it into place, this finds nothing. + defer func() { _ = fs.RemoveAll(tmpDir) }() + // Generate a new keypair for this version Debug("Generating version-specific keypair", "version", sv.Version) @@ -173,21 +184,28 @@ func (sv *Version) Save(value *memguard.LockedBuffer) error { slog.String("public_key", versionIdentity.Recipient().String()), ) - err = sv.writePublicKeyAndValue(fs, versionIdentity, value) + err = sv.writePublicKeyAndValue(fs, tmpDir, versionIdentity, value) if err != nil { return err } - err = sv.writeEncryptedPrivateKey(fs, versionPrivateKeyBuffer) + err = sv.writeEncryptedPrivateKey(fs, tmpDir, versionPrivateKeyBuffer) if err != nil { return err } - err = sv.writeEncryptedMetadata(fs, versionIdentity) + err = sv.writeEncryptedMetadata(fs, tmpDir, versionIdentity) if err != nil { return err } + err = fs.Rename(tmpDir, sv.Directory) + if err != nil { + Debug("Failed to move version into place", "error", err, "dir", sv.Directory) + + return fmt.Errorf("failed to move version into place: %w", err) + } + Debug("Successfully saved secret version", "version", sv.Version, "secret_name", sv.SecretName) @@ -358,17 +376,18 @@ func (sv *Version) GetValue( } // writePublicKeyAndValue stores the version's public key and the value -// encrypted to it. +// encrypted to it in dir. func (sv *Version) writePublicKeyAndValue( fs afero.Fs, + dir string, versionIdentity *age.X25519Identity, value *memguard.LockedBuffer, ) error { versionPublicKey := versionIdentity.Recipient().String() - pubKeyPath := filepath.Join(sv.Directory, "pub.age") + pubKeyPath := filepath.Join(dir, "pub.age") Debug("Writing version public key", "path", pubKeyPath) - err := afero.WriteFile(fs, pubKeyPath, []byte(versionPublicKey), FilePerms) + err := WriteFileAtomic(fs, pubKeyPath, []byte(versionPublicKey)) if err != nil { Debug("Failed to write version public key", "error", err, "path", pubKeyPath) @@ -385,10 +404,10 @@ func (sv *Version) writePublicKeyAndValue( return fmt.Errorf("failed to encrypt version value: %w", err) } - valuePath := filepath.Join(sv.Directory, "value.age") + valuePath := filepath.Join(dir, "value.age") Debug("Writing encrypted version value", "path", valuePath) - err = afero.WriteFile(fs, valuePath, encryptedValue, FilePerms) + err = WriteFileAtomic(fs, valuePath, encryptedValue) if err != nil { Debug("Failed to write encrypted version value", "error", err, "path", valuePath) @@ -399,9 +418,10 @@ func (sv *Version) writePublicKeyAndValue( } // writeEncryptedPrivateKey encrypts the version's private key to the -// vault's long-term public key and stores it. +// vault's long-term public key and stores it in dir. func (sv *Version) writeEncryptedPrivateKey( fs afero.Fs, + dir string, versionPrivateKeyBuffer *memguard.LockedBuffer, ) error { vaultDir, _ := sv.vault.GetDirectory() @@ -435,10 +455,10 @@ func (sv *Version) writeEncryptedPrivateKey( return fmt.Errorf("failed to encrypt version private key: %w", err) } - privKeyPath := filepath.Join(sv.Directory, "priv.age") + privKeyPath := filepath.Join(dir, "priv.age") Debug("Writing encrypted version private key", "path", privKeyPath) - err = afero.WriteFile(fs, privKeyPath, encryptedPrivKey, FilePerms) + err = WriteFileAtomic(fs, privKeyPath, encryptedPrivKey) if err != nil { Debug("Failed to write encrypted version private key", "error", err, "path", privKeyPath) @@ -450,9 +470,10 @@ func (sv *Version) writeEncryptedPrivateKey( } // writeEncryptedMetadata encrypts the version metadata to the version's -// public key and stores it. +// public key and stores it in dir. func (sv *Version) writeEncryptedMetadata( fs afero.Fs, + dir string, versionIdentity *age.X25519Identity, ) error { Debug("Encrypting version metadata", "version", sv.Version) @@ -476,10 +497,10 @@ func (sv *Version) writeEncryptedMetadata( return fmt.Errorf("failed to encrypt version metadata: %w", err) } - metadataPath := filepath.Join(sv.Directory, "metadata.age") + metadataPath := filepath.Join(dir, "metadata.age") Debug("Writing encrypted version metadata", "path", metadataPath) - err = afero.WriteFile(fs, metadataPath, encryptedMetadata, FilePerms) + err = WriteFileAtomic(fs, metadataPath, encryptedMetadata) if err != nil { Debug("Failed to write encrypted version metadata", "error", err, "path", metadataPath) @@ -540,15 +561,12 @@ func GetCurrentVersion(fs afero.Fs, secretDir string) (string, error) { } // SetCurrentVersion updates the "current" file to point to a specific version -// The file contains just the version name (e.g., "20231215.001") +// The file contains just the version name (e.g., "20231215.001"). It is +// replaced in one rename, so once written it always exists. func SetCurrentVersion(fs afero.Fs, secretDir string, version string) error { currentPath := filepath.Join(secretDir, "current") - // Remove existing file if it exists - _ = fs.Remove(currentPath) - - // Write just the version name to the file - err := afero.WriteFile(fs, currentPath, []byte(version), FilePerms) + err := WriteFileAtomic(fs, currentPath, []byte(version)) if err != nil { return fmt.Errorf("failed to create current version file: %w", err) } diff --git a/internal/vault/errors.go b/internal/vault/errors.go index 7d51e2e..cbae67e 100644 --- a/internal/vault/errors.go +++ b/internal/vault/errors.go @@ -62,4 +62,10 @@ var ( // ErrUnlockerNotFound indicates no unlocker with the given ID exists. // Composed as "unlocker with ID not found". ErrUnlockerNotFound = errors.New("not found") + + // ErrNoLockForFilesystem indicates LockStateDir was given a filesystem + // it cannot lock. Composed as "cannot lock the state directory on + // filesystem ". + ErrNoLockForFilesystem = errors.New( + "cannot lock the state directory on filesystem") ) diff --git a/internal/vault/lock.go b/internal/vault/lock.go new file mode 100644 index 0000000..3639c9c --- /dev/null +++ b/internal/vault/lock.go @@ -0,0 +1,73 @@ +package vault + +import ( + "fmt" + "os" + "path/filepath" + "sync" + "syscall" + + "git.eeqj.de/sneak/secret/internal/secret" + "github.com/spf13/afero" +) + +// lockFileName is the file in the state directory that LockStateDir locks. +const lockFileName = "lock" + +// memFsLock stands in for the lock file on the in-memory filesystem, which +// has no file locks. Every in-memory filesystem in the process shares it. +// +//nolint:gochecknoglobals // must outlive the call that takes it +var memFsLock sync.Mutex + +// LockStateDir takes the lock that a command changing anything under +// stateDir holds until it returns, and returns the function that releases +// it. While one command holds it, the next one waits here. Reads take no +// lock: each file or directory a command changes is replaced in a single +// rename, so a reader finds it as it was before or after, never half-made. +// +// On the real filesystem the lock is flock(2) on the file "lock" in +// stateDir, which the kernel releases when the process dies, so a killed +// command never leaves the tool locked. The in-memory filesystem the tests +// use has no file locks, so a process-wide mutex stands in for flock there. +// Any other filesystem is refused rather than left unlocked. +func LockStateDir(fs afero.Fs, stateDir string) (func(), error) { + switch fs.(type) { + case *afero.OsFs: + return flockStateDir(stateDir) + case *afero.MemMapFs: + memFsLock.Lock() + + return memFsLock.Unlock, nil + default: + return nil, fmt.Errorf("%w %T", ErrNoLockForFilesystem, fs) + } +} + +// flockStateDir takes flock(2) on the lock file in stateDir, creating the +// directory and the file if needed. Go opens files close-on-exec, so +// programs the command runs, such as gpg, do not inherit the lock. +func flockStateDir(stateDir string) (func(), error) { + err := os.MkdirAll(stateDir, secret.DirPerms) + if err != nil { + return nil, fmt.Errorf("failed to create state directory: %w", err) + } + + lockPath := filepath.Join(stateDir, lockFileName) + + //nolint:gosec // G304: the path is the lock file in the state directory + file, err := os.OpenFile(lockPath, os.O_RDWR|os.O_CREATE, secret.FilePerms) + if err != nil { + return nil, fmt.Errorf("failed to open lock file: %w", err) + } + + err = syscall.Flock(int(file.Fd()), syscall.LOCK_EX) + if err != nil { + _ = file.Close() + + return nil, fmt.Errorf("failed to lock %s: %w", lockPath, err) + } + + // Closing the file releases the lock. + return func() { _ = file.Close() }, nil +} diff --git a/internal/vault/lock_test.go b/internal/vault/lock_test.go new file mode 100644 index 0000000..d6882fc --- /dev/null +++ b/internal/vault/lock_test.go @@ -0,0 +1,134 @@ +package vault_test + +import ( + "testing" + "time" + + "git.eeqj.de/sneak/secret/internal/vault" + "github.com/spf13/afero" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const ( + // lockWait is how long a test waits for the lock before deciding it + // will never come free. + lockWait = 10 * time.Second + + // heldWait is how long a test watches a second holder fail to take a + // lock that is held. Broken exclusion lets it in at once. + heldWait = 100 * time.Millisecond +) + +// lockFilesystem is a filesystem LockStateDir can lock, with a state +// directory on it. +type lockFilesystem struct { + name string + fs afero.Fs + stateDir string +} + +// lockFilesystems returns the real filesystem, locked with flock, and the +// in-memory one, locked with a mutex. +func lockFilesystems(t *testing.T) []lockFilesystem { + t.Helper() + + return []lockFilesystem{ + {"memory", afero.NewMemMapFs(), testStateDir}, + {"real", afero.NewOsFs(), t.TempDir()}, + } +} + +// lockInBackground starts taking the lock and returns a channel that +// delivers the function releasing it once it has been taken. +func lockInBackground( + t *testing.T, fs afero.Fs, stateDir string, +) <-chan func() { + t.Helper() + + taken := make(chan func(), 1) + + go func() { + release, err := vault.LockStateDir(fs, stateDir) + if assert.NoError(t, err) { + taken <- release + } + }() + + return taken +} + +// TestLockStateDirExcludes checks that while the lock is held a second +// holder, with its own open lock file on the real filesystem, waits, and +// that it gets the lock once the first releases it. +func TestLockStateDirExcludes(t *testing.T) { + t.Parallel() + + for _, lfs := range lockFilesystems(t) { + t.Run(lfs.name, func(t *testing.T) { + t.Parallel() + + release, err := vault.LockStateDir(lfs.fs, lfs.stateDir) + require.NoError(t, err) + + taken := lockInBackground(t, lfs.fs, lfs.stateDir) + + select { + case second := <-taken: + second() + release() + t.Fatal("a second holder took the lock while it was held") + case <-time.After(heldWait): + } + + release() + + select { + case second := <-taken: + second() + case <-time.After(lockWait): + t.Fatal("the second holder never got the lock") + } + }) + } +} + +// TestLockStateDirFreeAfterPanic checks that a holder that panics, and +// releases the lock with defer as every command does, leaves it free. +func TestLockStateDirFreeAfterPanic(t *testing.T) { + t.Parallel() + + for _, lfs := range lockFilesystems(t) { + t.Run(lfs.name, func(t *testing.T) { + t.Parallel() + + assert.Panics(t, func() { + release, err := vault.LockStateDir(lfs.fs, lfs.stateDir) + require.NoError(t, err) + + defer release() + + panic("the command failed") + }) + + select { + case release := <-lockInBackground(t, lfs.fs, lfs.stateDir): + release() + case <-time.After(lockWait): + t.Fatal("the lock was still held after its holder panicked") + } + }) + } +} + +// TestLockStateDirRefusesOtherFilesystems checks that a filesystem with no +// lock implementation is refused instead of being used unlocked. +func TestLockStateDirRefusesOtherFilesystems(t *testing.T) { + t.Parallel() + + fs := afero.NewReadOnlyFs(afero.NewMemMapFs()) + + release, err := vault.LockStateDir(fs, testStateDir) + require.ErrorIs(t, err, vault.ErrNoLockForFilesystem) + assert.Nil(t, release) +} diff --git a/internal/vault/management.go b/internal/vault/management.go index f112da8..9cfa4d9 100644 --- a/internal/vault/management.go +++ b/internal/vault/management.go @@ -169,7 +169,7 @@ func processMnemonicForVault( ltPubKeyPath := filepath.Join(vaultDir, "pub.age") - err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltPubKey), secret.FilePerms) + err = secret.WriteFileAtomic(fs, ltPubKeyPath, []byte(ltPubKey)) if err != nil { return 0, "", "", fmt.Errorf("failed to write long-term public key: %w", err) } @@ -295,21 +295,13 @@ func SelectVault(fs afero.Fs, stateDir string, name string) error { return fmt.Errorf("vault %s %w", name, ErrVaultNotFound) } - // Create or update the currentvault file with just the vault name + // Create or replace the currentvault file with just the vault name. It + // is replaced in one rename, so it never goes missing. currentVaultPath := filepath.Join(stateDir, "currentvault") - // Remove existing file if it exists - _, err = fs.Stat(currentVaultPath) - if err == nil { - secret.Debug("Removing existing currentvault file", "path", currentVaultPath) - - _ = fs.Remove(currentVaultPath) - } - - // Write just the vault name to the file secret.Debug("Writing currentvault file", "vault_name", name) - err = afero.WriteFile(fs, currentVaultPath, []byte(name), secret.FilePerms) + err = secret.WriteFileAtomic(fs, currentVaultPath, []byte(name)) if err != nil { return fmt.Errorf("failed to select vault: %w", err) } diff --git a/internal/vault/metadata.go b/internal/vault/metadata.go index 0ac3a72..56daae3 100644 --- a/internal/vault/metadata.go +++ b/internal/vault/metadata.go @@ -113,7 +113,7 @@ func SaveVaultMetadata(fs afero.Fs, vaultDir string, metadata *Metadata) error { return fmt.Errorf("failed to marshal vault metadata: %w", err) } - err = afero.WriteFile(fs, metadataPath, metadataBytes, secret.FilePerms) + err = secret.WriteFileAtomic(fs, metadataPath, metadataBytes) if err != nil { return fmt.Errorf("failed to write vault metadata: %w", err) } diff --git a/internal/vault/secrets.go b/internal/vault/secrets.go index b8ab559..aff861c 100644 --- a/internal/vault/secrets.go +++ b/internal/vault/secrets.go @@ -156,17 +156,59 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool) slog.String("secret_dir", secretDir), ) - // Check for an existing secret and prepare its directory - exists, previousVersion, err := v.prepareSecretDir(name, secretDir, force) + // Check for an existing secret and the version the new one supersedes + exists, previousVersion, err := v.checkExistingSecret(name, secretDir, force) if err != nil { return err } + if exists { + return v.addVersion(name, secretDir, value, previousVersion) + } + + return v.addNewSecret(name, secretDir, value) +} + +// addNewSecret creates a secret by assembling its first version and current +// pointer in a temporary directory, then renaming that directory to +// secretDir, so an interrupted add leaves no half-made secret behind. +func (v *Vault) addNewSecret( + name, secretDir string, value *memguard.LockedBuffer, +) error { + buildDir, err := secret.TempDirFor(v.fs, secretDir) + if err != nil { + return err + } + + // Once the rename below has moved it into place, this finds nothing. + defer func() { _ = v.fs.RemoveAll(buildDir) }() + + err = v.addVersion(name, buildDir, value, nil) + if err != nil { + return err + } + + err = v.fs.Rename(buildDir, secretDir) + if err != nil { + return fmt.Errorf("failed to move new secret into place: %w", err) + } + + return nil +} + +// addVersion saves value as a new version under secretDir, sets the +// notAfter timestamp of the version it supersedes, if any, and then points +// current at the new version. Until that last step, current still names the +// previous version, which stays readable. +func (v *Vault) addVersion( + name, secretDir string, value *memguard.LockedBuffer, + previousVersion *secret.Version, +) error { now := time.Now() // Create the new version and save the encrypted value versionName, err := v.createAndSaveVersion( - name, secretDir, value, previousVersion, &now, exists) + name, secretDir, value, previousVersion, &now) if err != nil { return err } @@ -236,7 +278,7 @@ func updateVersionMetadata( // Write encrypted metadata metadataPath := filepath.Join(version.Directory, "metadata.age") - err = afero.WriteFile(fs, metadataPath, encryptedMetadata, secret.FilePerms) + err = secret.WriteFileAtomic(fs, metadataPath, encryptedMetadata) if err != nil { return fmt.Errorf("failed to write encrypted version metadata: %w", err) } @@ -394,12 +436,14 @@ func (v *Vault) GetSecretObject(name string) (*secret.Secret, error) { return secretObj, nil } -// CopySecretVersion copies a single version from source to this vault -// It decrypts the value using srcIdentity and re-encrypts for this vault +// CopySecretVersion copies a single version from source into destSecretDir +// in this vault. It decrypts the value using srcIdentity and re-encrypts +// for this vault. func (v *Vault) CopySecretVersion( srcVersion *secret.Version, srcIdentity *age.X25519Identity, destSecretName string, + destSecretDir string, destVersionName string, ) error { secret.DebugWith("Copying secret version to vault", @@ -425,6 +469,7 @@ func (v *Vault) CopySecretVersion( // Create destination version with same name destVersion := secret.NewVersion(v, destSecretName, destVersionName) + destVersion.Directory = filepath.Join(destSecretDir, "versions", destVersionName) // Copy metadata (preserve original timestamps) destVersion.Metadata = srcVersion.Metadata @@ -465,11 +510,11 @@ func (v *Vault) CopySecretAllVersions( return fmt.Errorf("failed to get destination vault directory: %w", err) } - // Check if destination secret already exists and clear it if forced + // Refuse to replace an existing destination secret unless forced destStorageName := strings.ReplaceAll(destSecretName, "/", "%") destSecretDir := filepath.Join(destVaultDir, "secrets.d", destStorageName) - err = v.prepareCopyDestination(destSecretDir, destSecretName, force) + err = v.checkCopyDestination(destSecretDir, destSecretName, force) if err != nil { return err } @@ -505,14 +550,8 @@ func (v *Vault) CopySecretAllVersions( return fmt.Errorf("failed to get current version: %w", err) } - // Create destination secret directory - err = v.fs.MkdirAll(destSecretDir, secret.DirPerms) - if err != nil { - return fmt.Errorf("failed to create destination secret directory: %w", err) - } - - // Copy each version and set the current pointer, rolling back on error - err = v.copyVersionsWithRollback(srcVault, srcIdentity, + // Copy each version and the current pointer, then move the copy into place + err = v.copyVersions(srcVault, srcIdentity, srcSecretName, destSecretName, destSecretDir, versions, currentVersion) if err != nil { return err @@ -527,10 +566,10 @@ func (v *Vault) CopySecretAllVersions( return nil } -// prepareSecretDir checks for an existing secret directory and prepares it -// for a new version. It returns whether the secret already existed and the -// current version to be superseded, if any. -func (v *Vault) prepareSecretDir( +// checkExistingSecret reports whether the secret already exists, refuses to +// overwrite it unless force is set, and returns its current version, which +// the new version supersedes, if any. +func (v *Vault) checkExistingSecret( name, secretDir string, force bool, ) (bool, *secret.Version, error) { // Check if secret already exists @@ -547,19 +586,6 @@ func (v *Vault) prepareSecretDir( secret.Debug("Secret existence check complete", "exists", exists) if !exists { - // Create secret directory for new secret - secret.Debug("Creating secret directory", "secret_dir", secretDir) - - err = v.fs.MkdirAll(secretDir, secret.DirPerms) - if err != nil { - secret.Debug("Failed to create secret directory", - "error", err, "secret_dir", secretDir) - - return false, nil, fmt.Errorf("failed to create secret directory: %w", err) - } - - secret.Debug("Created secret directory successfully") - return false, nil, nil } @@ -701,11 +727,11 @@ func (v *Vault) resolveSecretVersion(name, version string) (string, error) { } // createAndSaveVersion generates a new version name, sets the version -// timestamps, and saves the encrypted value. When saving fails for a newly -// created secret, the secret directory is removed again. +// timestamps, and saves the encrypted value under secretDir, which is a +// temporary directory while a new secret is being assembled. func (v *Vault) createAndSaveVersion( name, secretDir string, value *memguard.LockedBuffer, - previousVersion *secret.Version, now *time.Time, exists bool, + previousVersion *secret.Version, now *time.Time, ) (string, error) { // Generate new version name versionName, err := secret.GenerateVersionName(v.fs, secretDir) @@ -719,6 +745,7 @@ func (v *Vault) createAndSaveVersion( // Create new version newVersion := secret.NewVersion(v, name, versionName) + newVersion.Directory = filepath.Join(secretDir, "versions", versionName) // Set version timestamps if previousVersion == nil { @@ -738,57 +765,73 @@ func (v *Vault) createAndSaveVersion( if err != nil { secret.Debug("Failed to save new version", "error", err, "version", versionName) - // Clean up the secret directory if this was a new secret - if !exists { - secret.Debug("Cleaning up secret directory due to save failure", - "secret_dir", secretDir) - - _ = v.fs.RemoveAll(secretDir) - } - return "", fmt.Errorf("failed to save version: %w", err) } return versionName, nil } -// copyVersionsWithRollback copies each version of the source secret into the -// destination directory and sets the current version pointer, removing the -// partial copy when any step fails. -func (v *Vault) copyVersionsWithRollback( +// copyVersions copies each version of the source secret and its current +// pointer into a temporary directory, then moves that directory to +// destSecretDir, replacing a secret already there. Nothing in this vault +// changes until the copy is complete, so an interrupted copy leaves only a +// temporary directory behind. +func (v *Vault) copyVersions( srcVault *Vault, srcIdentity *age.X25519Identity, srcSecretName, destSecretName, destSecretDir string, versions []string, currentVersion string, ) error { - // Copy each version + buildDir, err := secret.TempDirFor(v.fs, destSecretDir) + if err != nil { + return err + } + + // Once the rename below has moved it into place, this finds nothing. + defer func() { _ = v.fs.RemoveAll(buildDir) }() + for _, versionName := range versions { srcVersion := secret.NewVersion(srcVault, srcSecretName, versionName) - err := v.CopySecretVersion(srcVersion, srcIdentity, destSecretName, versionName) + err = v.CopySecretVersion( + srcVersion, srcIdentity, destSecretName, buildDir, versionName) if err != nil { - // Rollback: remove partial copy - secret.Debug("Rolling back partial copy due to error", "error", err) - - _ = v.fs.RemoveAll(destSecretDir) - return fmt.Errorf("failed to copy version %s: %w", versionName, err) } } - // Set current version - err := secret.SetCurrentVersion(v.fs, destSecretDir, currentVersion) + err = secret.SetCurrentVersion(v.fs, buildDir, currentVersion) if err != nil { - _ = v.fs.RemoveAll(destSecretDir) - return fmt.Errorf("failed to set current version: %w", err) } + // With --force, the secret being replaced goes only now that its + // replacement is complete + exists, err := afero.DirExists(v.fs, destSecretDir) + if err != nil { + return fmt.Errorf("failed to check destination: %w", err) + } + + if exists { + secret.Debug("Removing existing destination secret", "path", destSecretDir) + + err = secret.RemoveDirAtomic(v.fs, destSecretDir) + if err != nil { + return fmt.Errorf("failed to remove existing destination secret: %w", err) + } + } + + err = v.fs.Rename(buildDir, destSecretDir) + if err != nil { + return fmt.Errorf("failed to move copied secret into place: %w", err) + } + return nil } -// prepareCopyDestination ensures the destination secret directory can be -// created, removing an existing secret when force is set. -func (v *Vault) prepareCopyDestination( +// checkCopyDestination refuses to copy over an existing secret unless force +// is set. A secret being replaced is removed by copyVersions, once its +// replacement is complete. +func (v *Vault) checkCopyDestination( destSecretDir, destSecretName string, force bool, ) error { exists, err := afero.DirExists(v.fs, destSecretDir) @@ -803,15 +846,5 @@ func (v *Vault) prepareCopyDestination( ) } - if exists && force { - // Remove existing secret - secret.Debug("Removing existing destination secret", "path", destSecretDir) - - err = v.fs.RemoveAll(destSecretDir) - if err != nil { - return fmt.Errorf("failed to remove existing destination secret: %w", err) - } - } - return nil } diff --git a/internal/vault/unlockers.go b/internal/vault/unlockers.go index e2c562b..bddb25d 100644 --- a/internal/vault/unlockers.go +++ b/internal/vault/unlockers.go @@ -310,30 +310,16 @@ func (v *Vault) SelectUnlocker(unlockerID string) error { return fmt.Errorf("unlocker with ID %s %w", unlockerID, ErrUnlockerNotFound) } - // Create/update current-unlocker file with just the unlocker name + // Create or replace the current-unlocker file with just the unlocker + // name. It is replaced in one rename, so it never goes missing. currentUnlockerPath := filepath.Join(vaultDir, "current-unlocker") - // Remove existing file if it exists - exists, err := afero.Exists(v.fs, currentUnlockerPath) - if err != nil { - return fmt.Errorf("failed to check if current-unlocker file exists: %w", err) - } - - if exists { - err = v.fs.Remove(currentUnlockerPath) - if err != nil { - return fmt.Errorf("failed to remove existing current-unlocker file: %w", err) - } - } - // Get just the unlocker name (basename of the directory) unlockerName := filepath.Base(targetUnlockerDir) - // Write just the unlocker name to the file secret.Debug("Writing current-unlocker file", "unlocker_name", unlockerName) - err = afero.WriteFile(v.fs, currentUnlockerPath, []byte(unlockerName), - secret.FilePerms) + err = secret.WriteFileAtomic(v.fs, currentUnlockerPath, []byte(unlockerName)) if err != nil { return fmt.Errorf("failed to create current-unlocker file: %w", err) } @@ -351,6 +337,14 @@ func (v *Vault) CreatePassphraseUnlocker( return nil, fmt.Errorf("failed to get vault directory: %w", err) } + // We need to get the long-term key (either from memory if unlocked, or + // derive it). Getting it before anything is written means failing to + // get it changes nothing, even when replacing the current unlocker. + ltIdentity, err := v.GetOrDeriveLongTermKey() + if err != nil { + return nil, fmt.Errorf("failed to get long-term key: %w", err) + } + // Create unlocker directory unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerTypePassphrase) @@ -371,33 +365,7 @@ func (v *Vault) CreatePassphraseUnlocker( return nil, err } - // Create metadata - metadata := UnlockerMetadata{ - Type: unlockerTypePassphrase, - CreatedAt: time.Now(), - Flags: []string{}, - } - - // Write metadata - metadataBytes, err := json.MarshalIndent(metadata, "", " ") - if err != nil { - return nil, fmt.Errorf("failed to marshal metadata: %w", err) - } - - metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") - - err = afero.WriteFile(v.fs, metadataPath, metadataBytes, secret.FilePerms) - if err != nil { - return nil, fmt.Errorf("failed to write unlocker metadata: %w", err) - } - // Encrypt long-term private key to this unlocker - // We need to get the long-term key (either from memory if unlocked, or derive it) - ltIdentity, err := v.GetOrDeriveLongTermKey() - if err != nil { - return nil, fmt.Errorf("failed to get long-term key: %w", err) - } - ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String())) defer ltPrivKeyBuffer.Destroy() @@ -409,11 +377,31 @@ func (v *Vault) CreatePassphraseUnlocker( ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age") - err = afero.WriteFile(v.fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms) + err = secret.WriteFileAtomic(v.fs, ltPrivKeyPath, encryptedLtPrivKey) if err != nil { return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err) } + // Write the metadata last: readers skip an unlocker directory without + // it, so an unlocker interrupted before this point is never used. + metadata := UnlockerMetadata{ + Type: unlockerTypePassphrase, + CreatedAt: time.Now(), + Flags: []string{}, + } + + metadataBytes, err := json.MarshalIndent(metadata, "", " ") + if err != nil { + return nil, fmt.Errorf("failed to marshal metadata: %w", err) + } + + metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") + + err = secret.WriteFileAtomic(v.fs, metadataPath, metadataBytes) + if err != nil { + return nil, fmt.Errorf("failed to write unlocker metadata: %w", err) + } + // Create the unlocker instance unlocker := secret.NewPassphraseUnlocker(v.fs, unlockerDir, metadata) @@ -467,9 +455,8 @@ func (v *Vault) writeUnlockerKeypair( // Write public key pubKeyPath := filepath.Join(unlockerDir, "pub.age") - err := afero.WriteFile(v.fs, pubKeyPath, - []byte(unlockerIdentity.Recipient().String()), - secret.FilePerms) + err := secret.WriteFileAtomic(v.fs, pubKeyPath, + []byte(unlockerIdentity.Recipient().String())) if err != nil { return fmt.Errorf("failed to write unlocker public key: %w", err) } @@ -488,7 +475,7 @@ func (v *Vault) writeUnlockerKeypair( // Write encrypted private key privKeyPath := filepath.Join(unlockerDir, "priv.age") - err = afero.WriteFile(v.fs, privKeyPath, encryptedPrivKey, secret.FilePerms) + err = secret.WriteFileAtomic(v.fs, privKeyPath, encryptedPrivKey) if err != nil { return fmt.Errorf("failed to write encrypted unlocker private key: %w", err) } diff --git a/internal/vault/vault_error_test.go b/internal/vault/vault_error_test.go index ad61e82..bec8833 100644 --- a/internal/vault/vault_error_test.go +++ b/internal/vault/vault_error_test.go @@ -90,4 +90,10 @@ func TestAddSecretCleansUpOnFailure(t *testing.T) { secretDir := filepath.Join(vaultDir, "secrets.d", testSecretName) exists, _ := afero.DirExists(fs, secretDir) assert.False(t, exists, "Secret directory should not exist after failed AddSecret") + + // Nor is the temporary directory the secret was assembled in left behind + entries, err := afero.ReadDir(fs, vaultDir) + require.NoError(t, err) + require.Len(t, entries, 1) + assert.Equal(t, "pub.age", entries[0].Name()) }