Lock the state directory and write vault files atomically (closes #34)
check / check (push) Successful in 46s
check / check (push) Successful in 46s
Each command that changes the state directory holds one lock: flock(2) on `lock` in the state directory, dropped by the kernel if the process dies, or a process-wide mutex on the in-memory test filesystem. It covers the state directory, not each vault, because `currentvault`, `vault create` and cross-vault moves span vaults, and a lock file in a vault would be deleted by `vault remove` under a waiting command. Files go through `secret.WriteFileAtomic`; versions, new secrets and cross-vault copies are built in a temporary directory and renamed into place; removals rename out of the way first. Left for later: replacing an unlocker (#71) and deleting what an interrupted command leaves under a `.tmp-` name (#75). Model: opus-5-5
This commit is contained in:
@@ -62,4 +62,10 @@ var (
|
||||
// ErrUnlockerNotFound indicates no unlocker with the given ID exists.
|
||||
// Composed as "unlocker with ID <id> not found".
|
||||
ErrUnlockerNotFound = errors.New("not found")
|
||||
|
||||
// ErrNoLockForFilesystem indicates LockStateDir was given a filesystem
|
||||
// it cannot lock. Composed as "cannot lock the state directory on
|
||||
// filesystem <type>".
|
||||
ErrNoLockForFilesystem = errors.New(
|
||||
"cannot lock the state directory on filesystem")
|
||||
)
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
package vault
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// lockFileName is the file in the state directory that LockStateDir locks.
|
||||
const lockFileName = "lock"
|
||||
|
||||
// memFsLock stands in for the lock file on the in-memory filesystem, which
|
||||
// has no file locks. Every in-memory filesystem in the process shares it.
|
||||
//
|
||||
//nolint:gochecknoglobals // must outlive the call that takes it
|
||||
var memFsLock sync.Mutex
|
||||
|
||||
// LockStateDir takes the lock that a command changing anything under
|
||||
// stateDir holds until it returns, and returns the function that releases
|
||||
// it. While one command holds it, the next one waits here. Reads take no
|
||||
// lock: each file or directory a command changes is replaced in a single
|
||||
// rename, so a reader finds it as it was before or after, never half-made.
|
||||
//
|
||||
// On the real filesystem the lock is flock(2) on the file "lock" in
|
||||
// stateDir, which the kernel releases when the process dies, so a killed
|
||||
// command never leaves the tool locked. The in-memory filesystem the tests
|
||||
// use has no file locks, so a process-wide mutex stands in for flock there.
|
||||
// Any other filesystem is refused rather than left unlocked.
|
||||
func LockStateDir(fs afero.Fs, stateDir string) (func(), error) {
|
||||
switch fs.(type) {
|
||||
case *afero.OsFs:
|
||||
return flockStateDir(stateDir)
|
||||
case *afero.MemMapFs:
|
||||
memFsLock.Lock()
|
||||
|
||||
return memFsLock.Unlock, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("%w %T", ErrNoLockForFilesystem, fs)
|
||||
}
|
||||
}
|
||||
|
||||
// flockStateDir takes flock(2) on the lock file in stateDir, creating the
|
||||
// directory and the file if needed. Go opens files close-on-exec, so
|
||||
// programs the command runs, such as gpg, do not inherit the lock.
|
||||
func flockStateDir(stateDir string) (func(), error) {
|
||||
err := os.MkdirAll(stateDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create state directory: %w", err)
|
||||
}
|
||||
|
||||
lockPath := filepath.Join(stateDir, lockFileName)
|
||||
|
||||
//nolint:gosec // G304: the path is the lock file in the state directory
|
||||
file, err := os.OpenFile(lockPath, os.O_RDWR|os.O_CREATE, secret.FilePerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to open lock file: %w", err)
|
||||
}
|
||||
|
||||
err = syscall.Flock(int(file.Fd()), syscall.LOCK_EX)
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
|
||||
return nil, fmt.Errorf("failed to lock %s: %w", lockPath, err)
|
||||
}
|
||||
|
||||
// Closing the file releases the lock.
|
||||
return func() { _ = file.Close() }, nil
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"github.com/spf13/afero"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const (
|
||||
// lockWait is how long a test waits for the lock before deciding it
|
||||
// will never come free.
|
||||
lockWait = 10 * time.Second
|
||||
|
||||
// heldWait is how long a test watches a second holder fail to take a
|
||||
// lock that is held. Broken exclusion lets it in at once.
|
||||
heldWait = 100 * time.Millisecond
|
||||
)
|
||||
|
||||
// lockFilesystem is a filesystem LockStateDir can lock, with a state
|
||||
// directory on it.
|
||||
type lockFilesystem struct {
|
||||
name string
|
||||
fs afero.Fs
|
||||
stateDir string
|
||||
}
|
||||
|
||||
// lockFilesystems returns the real filesystem, locked with flock, and the
|
||||
// in-memory one, locked with a mutex.
|
||||
func lockFilesystems(t *testing.T) []lockFilesystem {
|
||||
t.Helper()
|
||||
|
||||
return []lockFilesystem{
|
||||
{"memory", afero.NewMemMapFs(), testStateDir},
|
||||
{"real", afero.NewOsFs(), t.TempDir()},
|
||||
}
|
||||
}
|
||||
|
||||
// lockInBackground starts taking the lock and returns a channel that
|
||||
// delivers the function releasing it once it has been taken.
|
||||
func lockInBackground(
|
||||
t *testing.T, fs afero.Fs, stateDir string,
|
||||
) <-chan func() {
|
||||
t.Helper()
|
||||
|
||||
taken := make(chan func(), 1)
|
||||
|
||||
go func() {
|
||||
release, err := vault.LockStateDir(fs, stateDir)
|
||||
if assert.NoError(t, err) {
|
||||
taken <- release
|
||||
}
|
||||
}()
|
||||
|
||||
return taken
|
||||
}
|
||||
|
||||
// TestLockStateDirExcludes checks that while the lock is held a second
|
||||
// holder, with its own open lock file on the real filesystem, waits, and
|
||||
// that it gets the lock once the first releases it.
|
||||
func TestLockStateDirExcludes(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, lfs := range lockFilesystems(t) {
|
||||
t.Run(lfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
release, err := vault.LockStateDir(lfs.fs, lfs.stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
taken := lockInBackground(t, lfs.fs, lfs.stateDir)
|
||||
|
||||
select {
|
||||
case second := <-taken:
|
||||
second()
|
||||
release()
|
||||
t.Fatal("a second holder took the lock while it was held")
|
||||
case <-time.After(heldWait):
|
||||
}
|
||||
|
||||
release()
|
||||
|
||||
select {
|
||||
case second := <-taken:
|
||||
second()
|
||||
case <-time.After(lockWait):
|
||||
t.Fatal("the second holder never got the lock")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLockStateDirFreeAfterPanic checks that a holder that panics, and
|
||||
// releases the lock with defer as every command does, leaves it free.
|
||||
func TestLockStateDirFreeAfterPanic(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, lfs := range lockFilesystems(t) {
|
||||
t.Run(lfs.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
assert.Panics(t, func() {
|
||||
release, err := vault.LockStateDir(lfs.fs, lfs.stateDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer release()
|
||||
|
||||
panic("the command failed")
|
||||
})
|
||||
|
||||
select {
|
||||
case release := <-lockInBackground(t, lfs.fs, lfs.stateDir):
|
||||
release()
|
||||
case <-time.After(lockWait):
|
||||
t.Fatal("the lock was still held after its holder panicked")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLockStateDirRefusesOtherFilesystems checks that a filesystem with no
|
||||
// lock implementation is refused instead of being used unlocked.
|
||||
func TestLockStateDirRefusesOtherFilesystems(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fs := afero.NewReadOnlyFs(afero.NewMemMapFs())
|
||||
|
||||
release, err := vault.LockStateDir(fs, testStateDir)
|
||||
require.ErrorIs(t, err, vault.ErrNoLockForFilesystem)
|
||||
assert.Nil(t, release)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
+105
-72
@@ -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
|
||||
}
|
||||
|
||||
+35
-48
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user