Update golangci-lint to v2.12.2 with canonical config (#29)
All checks were successful
check / check (push) Successful in 43s
All checks were successful
check / check (push) Successful in 43s
Bumps golangci-lint from v2.1.6 (digest-only pin in the `Dockerfile` lint stage) to v2.12.2, pinned by tag and digest (Debian-based image). Replaces `.golangci.yml` with the canonical strict config: all linters enabled except the standard disable list (`exhaustruct`, `depguard`, `godot`, `wsl`, `wrapcheck`, `varnamelen`), `lll` at 88, `funlen` 80/50, `cyclop` 15, `dupl` 100, and test files are now linted (the old config had `tests: false`, an enable-only list of ~20 linters, `lll` 120, and a blanket exclusion of `internal/macse`). The stricter config surfaced ~1550 findings, all fixed: - `wsl_v5` (439) / `nlreturn` (24): blank-line insertions - `lll` (309): line wrapping at 88 columns; long literals split with `+` concatenation, values unchanged - `noinlineerr` (130): `if err := ...` split into assignment plus check - `paralleltest` (116): `t.Parallel()` added to tests without shared state; reasoned `//nolint` where `t.Setenv` or shared fixtures forbid it - `err113` (97): package-level sentinel errors (new `internal/vault/errors.go`), `%w` wrapping, `errors.Is` - `perfsprint` (74) / `modernize` (39) / `intrange`: `strconv`, `errors.New`, `slices.Contains`, `any`, `SplitSeq` - `goconst` (40) / `dupword` (41) / `testifylint` (42) / `thelper` (33): constants, assertion fixes, `t.Helper()` - `noctx` (22): `exec.CommandContext` for gpg/CLI invocations - `testpackage` (18): black-box tests moved to `_test` packages where they use only exported identifiers; white-box files carry a reasoned `//nolint` - `funlen`/`cyclop`/`gocognit`/`nestif`/`dupl`: behavior-preserving helper extraction - assorted singletons: `gosec`, `gosmopolitan`, `funcorder`, `nonamedreturns`, `makezero`, `prealloc`, `godox`, `nolintlint`, `ireturn`, `nilnil`, `gochecknoinits` ## User-visible strings **None changed.** Every error message this branch composes is byte-identical to the one `main` composes. The `err113` sentinels are shaped so `fmt.Errorf` reassembles the original text around them: the sentinel carries the fixed words and the caller supplies the interpolated value in the position it has always occupied. Where the value sits mid-sentence the sentinel holds only a fragment (e.g. `vault.ErrVaultNotFound` is `"does not exist"`, composed by its caller as `vault <name> does not exist`); each such sentinel documents the message it participates in. Verified mechanically, not by inspection: every `fmt.Errorf` and `errors.New` call site in both trees is parsed, the `Error()` text of any sentinel passed to `%w` is substituted in, and the resulting sets of composed message templates are compared. All 350 templates `main` produces are still produced, character for character. The set of lost or altered messages is empty. ## `unlocker list` `findUnlockerIDByMetadata` returns `(string, error)` rather than signalling failure with an empty ID, so an unreadable `unlockers.d` is no longer indistinguishable from "no matching entry". `UnlockersList` skips such an entry with a warning naming the directory — its behavior before the scan was extracted into a helper — instead of emitting a row under a synthesized fallback ID that no `unlocker remove` or `unlocker select` can match and that suppresses the current-unlocker marker. The duplicate-check and shell-completion callers skip on the same condition, matching their pre-extraction behavior. Covered by `internal/cli/unlockers_list_test.go`. `TODO.md` records the change plus follow-ups (version-completion TODOs formerly in code comments, darwin-gated files exceeding 88 columns that Linux CI does not lint). `make check` is green and the pinned v2.12.2 image reports `0 issues.` Note the test suite needs the memlock ulimit from `script/cibuild` for the 10MB memguard test; that requirement is pre-existing. Not changed: `script/bootstrap` installs golangci-lint via the system package manager (no version pin to bump), and `script/lint` invokes whatever `golangci-lint` is on PATH. golangci-lint v2.12 deprecates `gomodguard` in favor of `gomodguard_v2` (warning only); the canonical config owns that decision. Co-authored-by: sneak <sneak@sneak.berlin> Reviewed-on: #29 Co-authored-by: clawbot <clawbot@noreply.example.org> Co-committed-by: clawbot <clawbot@noreply.example.org>
This commit was merged in pull request #29.
This commit is contained in:
65
internal/vault/errors.go
Normal file
65
internal/vault/errors.go
Normal file
@@ -0,0 +1,65 @@
|
||||
package vault
|
||||
|
||||
import "errors"
|
||||
|
||||
// Sentinel errors returned by vault operations.
|
||||
//
|
||||
// Several of these carry deliberately partial text: the message a caller
|
||||
// composes with fmt.Errorf places the interpolated value where it has
|
||||
// always appeared, and the sentinel supplies only the surrounding fixed
|
||||
// words. This keeps every composed message byte-identical to the dynamic
|
||||
// errors these sentinels replaced. Each such sentinel notes the message it
|
||||
// participates in.
|
||||
var (
|
||||
// ErrMnemonicMismatch indicates the mnemonic-derived public key does
|
||||
// not match the vault's stored public key hash.
|
||||
ErrMnemonicMismatch = errors.New(
|
||||
"derived public key does not match vault: mnemonic may be incorrect",
|
||||
)
|
||||
|
||||
// ErrInvalidVaultName indicates a vault name that does not match the
|
||||
// allowed pattern [a-z0-9.\-_]+. Composed as
|
||||
// "invalid vault name '<name>': must match pattern [a-z0-9.\-_]+".
|
||||
ErrInvalidVaultName = errors.New("invalid vault name")
|
||||
|
||||
// ErrVaultNotFound indicates the named vault does not exist. Composed
|
||||
// as "vault <name> does not exist".
|
||||
ErrVaultNotFound = errors.New("does not exist")
|
||||
|
||||
// ErrNilValueBuffer indicates a nil value buffer was supplied.
|
||||
ErrNilValueBuffer = errors.New("value buffer is nil")
|
||||
|
||||
// ErrInvalidSecretName indicates a secret name that does not match
|
||||
// the allowed pattern [a-z0-9.\-_/]+. Composed as
|
||||
// "invalid secret name '<name>': must match pattern [a-z0-9.\-_/]+",
|
||||
// or as "invalid secret name: <name>" by GetSecretObject.
|
||||
ErrInvalidSecretName = errors.New("invalid secret name")
|
||||
|
||||
// ErrSecretExists indicates the secret already exists and --force
|
||||
// was not supplied. Composed as
|
||||
// "secret <name> already exists (use --force to overwrite)", or as
|
||||
// "secret '<name>' already exists in vault '<vault>' (use --force to
|
||||
// overwrite)" when copying between vaults.
|
||||
ErrSecretExists = errors.New("already exists")
|
||||
|
||||
// ErrSecretNotFound indicates the named secret does not exist.
|
||||
// Composed as "secret <name> not found".
|
||||
ErrSecretNotFound = errors.New("not found")
|
||||
|
||||
// ErrVersionNotFound indicates the requested secret version does not
|
||||
// exist. Composed as
|
||||
// "version <version> not found for secret <name>".
|
||||
ErrVersionNotFound = errors.New("not found for secret")
|
||||
|
||||
// ErrNoVersions indicates the source secret has no versions. Composed
|
||||
// as "source secret '<name>' has no versions".
|
||||
ErrNoVersions = errors.New("has no versions")
|
||||
|
||||
// ErrUnsupportedUnlockerType indicates an unlocker metadata type
|
||||
// that this build does not support.
|
||||
ErrUnsupportedUnlockerType = errors.New("unsupported unlocker type")
|
||||
|
||||
// ErrUnlockerNotFound indicates no unlocker with the given ID exists.
|
||||
// Composed as "unlocker with ID <id> not found".
|
||||
ErrUnlockerNotFound = errors.New("not found")
|
||||
)
|
||||
@@ -3,8 +3,10 @@ package vault_test
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
@@ -12,6 +14,33 @@ import (
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// deriveVaultIdentity derives the long-term identity for the given vault
|
||||
// from testMnemonic using the derivation index stored in its metadata.
|
||||
func deriveVaultIdentity(
|
||||
t *testing.T, fs afero.Fs, vlt *vault.Vault,
|
||||
) *age.X25519Identity {
|
||||
t.Helper()
|
||||
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic,
|
||||
vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
return ltIdentity
|
||||
}
|
||||
|
||||
//nolint:paralleltest // t.Setenv forbids parallel subtests
|
||||
func TestVaultWithRealFilesystem(t *testing.T) {
|
||||
// Create a temporary directory for our tests
|
||||
tempDir := t.TempDir()
|
||||
@@ -19,398 +48,410 @@ func TestVaultWithRealFilesystem(t *testing.T) {
|
||||
// Use the real filesystem
|
||||
fs := afero.NewOsFs()
|
||||
|
||||
// Test mnemonic
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
// Set test environment variables
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
// Test currentvault file handling (plain file with relative path)
|
||||
t.Run("CurrentVaultFileHandling", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "currentvault-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Get the vault directory
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Verify the currentvault file exists and contains just the vault name
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
currentVaultContents, err := os.ReadFile(currentVaultPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to read currentvault file: %v", err)
|
||||
}
|
||||
|
||||
expectedVaultName := "test-vault"
|
||||
if string(currentVaultContents) != expectedVaultName {
|
||||
t.Errorf("Expected currentvault to contain %q, got %q", expectedVaultName, string(currentVaultContents))
|
||||
}
|
||||
|
||||
// Test that ResolveVaultSymlink correctly resolves the path
|
||||
resolvedPath, err := vault.ResolveVaultSymlink(fs, currentVaultPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve currentvault path: %v", err)
|
||||
}
|
||||
|
||||
if resolvedPath != vaultDir {
|
||||
t.Errorf("Expected resolved path to be %s, got %s", vaultDir, resolvedPath)
|
||||
}
|
||||
testCurrentVaultFileHandling(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test secret operations with deeply nested paths
|
||||
t.Run("DeepPathSecrets", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "deep-path-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault - CreateVault now handles public key when mnemonic is in env
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get its derivation index
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic using the vault's derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vault
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Create a secret with a deeply nested path
|
||||
deepPath := "api/credentials/production/database/primary"
|
||||
secretValue := []byte("supersecretdbpassword")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vlt.AddSecret(deepPath, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret with deep path: %v", err)
|
||||
}
|
||||
|
||||
// List secrets and verify our deep path secret is there
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets: %v", err)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, s := range secrets {
|
||||
if s == deepPath {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Errorf("Deep path secret not found in listed secrets")
|
||||
}
|
||||
|
||||
// Retrieve the secret and verify its value
|
||||
retrievedValue, err := vlt.GetSecret(deepPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to retrieve deep path secret: %v", err)
|
||||
}
|
||||
|
||||
if string(retrievedValue) != string(expectedValue) {
|
||||
t.Errorf("Retrieved value doesn't match. Expected %q, got %q",
|
||||
string(expectedValue), string(retrievedValue))
|
||||
}
|
||||
testDeepPathSecrets(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test key caching in GetOrDeriveLongTermKey
|
||||
t.Run("KeyCaching", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "key-cache-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault - CreateVault now handles public key when mnemonic is in env
|
||||
vlt, err := vault.CreateVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get its derivation index
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic for verification using the vault's derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the vault is locked initially
|
||||
if !vlt.Locked() {
|
||||
t.Errorf("Vault should be locked initially")
|
||||
}
|
||||
|
||||
// First call to GetOrDeriveLongTermKey should derive and cache the key
|
||||
firstKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the vault is now unlocked
|
||||
if vlt.Locked() {
|
||||
t.Errorf("Vault should be unlocked after GetOrDeriveLongTermKey")
|
||||
}
|
||||
|
||||
// Second call should return the cached key without re-deriving
|
||||
secondKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get cached long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify both keys are the same instance
|
||||
if firstKey != secondKey {
|
||||
t.Errorf("Second key call should return same instance as first call")
|
||||
}
|
||||
|
||||
// Verify the public key matches what we expect
|
||||
expectedPubKey := ltIdentity.Recipient().String()
|
||||
actualPubKey := firstKey.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Public key mismatch. Expected %s, got %s", expectedPubKey, actualPubKey)
|
||||
}
|
||||
|
||||
// Now clear the key and verify it's locked again
|
||||
vlt.ClearLongTermKey()
|
||||
if !vlt.Locked() {
|
||||
t.Errorf("Vault should be locked after clearing key")
|
||||
}
|
||||
|
||||
// Get the key again and verify it works
|
||||
thirdKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to re-derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the public key still matches
|
||||
actualPubKey = thirdKey.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Re-derived public key mismatch. Expected %s, got %s", expectedPubKey, actualPubKey)
|
||||
}
|
||||
testKeyCaching(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test vault name validation
|
||||
t.Run("VaultNameValidation", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "name-validation-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Test valid vault names
|
||||
validNames := []string{
|
||||
"default",
|
||||
"test-vault",
|
||||
"production.vault",
|
||||
"vault_123",
|
||||
"a-very-long-vault-name-with-dashes",
|
||||
}
|
||||
|
||||
for _, name := range validNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Errorf("Failed to create vault with valid name %q: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Test invalid vault names
|
||||
invalidNames := []string{
|
||||
"", // Empty
|
||||
"UPPERCASE", // Uppercase not allowed
|
||||
"invalid/name", // Slashes not allowed in vault names
|
||||
"invalid name", // Spaces not allowed
|
||||
"invalid@name", // Special chars not allowed
|
||||
}
|
||||
|
||||
for _, name := range invalidNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err == nil {
|
||||
t.Errorf("Expected error creating vault with invalid name %q, but got none", name)
|
||||
}
|
||||
}
|
||||
testVaultNameValidation(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test multiple vaults and switching between them
|
||||
t.Run("MultipleVaults", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "multi-vault-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create three vaults
|
||||
vaultNames := []string{"vault1", "vault2", "vault3"}
|
||||
for _, name := range vaultNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// List vaults and verify all three are there
|
||||
vaults, err := vault.ListVaults(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list vaults: %v", err)
|
||||
}
|
||||
|
||||
if len(vaults) != 3 {
|
||||
t.Errorf("Expected 3 vaults, got %d", len(vaults))
|
||||
}
|
||||
|
||||
// Test switching between vaults
|
||||
for _, name := range vaultNames {
|
||||
// Select the vault
|
||||
if err := vault.SelectVault(fs, stateDir, name); err != nil {
|
||||
t.Fatalf("Failed to select vault %s: %v", name, err)
|
||||
}
|
||||
|
||||
// Get current vault and verify it's the one we selected
|
||||
currentVault, err := vault.GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault after selecting %s: %v", name, err)
|
||||
}
|
||||
|
||||
if currentVault.GetName() != name {
|
||||
t.Errorf("Expected current vault to be %s, got %s", name, currentVault.GetName())
|
||||
}
|
||||
}
|
||||
testMultipleVaults(t, fs, tempDir)
|
||||
})
|
||||
|
||||
// Test adding a secret in one vault and verifying it's not visible in another
|
||||
// Test adding a secret in one vault and verifying it's not visible in
|
||||
// another
|
||||
t.Run("VaultIsolation", func(t *testing.T) {
|
||||
stateDir := filepath.Join(tempDir, "isolation-test")
|
||||
if err := os.MkdirAll(stateDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create two vaults - CreateVault now handles public key when mnemonic is in env
|
||||
vault1, err := vault.CreateVault(fs, stateDir, "vault1")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault1: %v", err)
|
||||
}
|
||||
|
||||
vault2, err := vault.CreateVault(fs, stateDir, "vault2")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault2: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic
|
||||
// Note: Both vaults will have different derivation indexes due to GetNextDerivationIndex
|
||||
|
||||
// Load vault1 metadata to get its derivation index
|
||||
vault1Dir, err := vault1.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault1 directory: %v", err)
|
||||
}
|
||||
vault1Metadata, err := vault.LoadVaultMetadata(fs, vault1Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault1 metadata: %v", err)
|
||||
}
|
||||
|
||||
ltIdentity1, err := agehd.DeriveIdentity(testMnemonic, vault1Metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key for vault1: %v", err)
|
||||
}
|
||||
|
||||
// Load vault2 metadata to get its derivation index
|
||||
vault2Dir, err := vault2.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault2 directory: %v", err)
|
||||
}
|
||||
vault2Metadata, err := vault.LoadVaultMetadata(fs, vault2Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault2 metadata: %v", err)
|
||||
}
|
||||
|
||||
ltIdentity2, err := agehd.DeriveIdentity(testMnemonic, vault2Metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key for vault2: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vaults with their respective keys
|
||||
vault1.Unlock(ltIdentity1)
|
||||
vault2.Unlock(ltIdentity2)
|
||||
|
||||
// Add a secret to vault1
|
||||
secretName := "test-secret"
|
||||
secretValue := []byte("secret in vault1")
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
if err := vault1.AddSecret(secretName, secretBuffer, false); err != nil {
|
||||
t.Fatalf("Failed to add secret to vault1: %v", err)
|
||||
}
|
||||
|
||||
// Verify the secret exists in vault1
|
||||
vault1Secrets, err := vault1.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets in vault1: %v", err)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, s := range vault1Secrets {
|
||||
if s == secretName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Errorf("Secret not found in vault1")
|
||||
}
|
||||
|
||||
// Verify the secret does NOT exist in vault2
|
||||
vault2Secrets, err := vault2.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets in vault2: %v", err)
|
||||
}
|
||||
|
||||
found = false
|
||||
for _, s := range vault2Secrets {
|
||||
if s == secretName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if found {
|
||||
t.Errorf("Secret from vault1 should not be visible in vault2")
|
||||
}
|
||||
testVaultIsolation(t, fs, tempDir)
|
||||
})
|
||||
}
|
||||
|
||||
func testCurrentVaultFileHandling(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "currentvault-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault
|
||||
vlt, err := vault.CreateVault(fs, stateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Get the vault directory
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Verify the currentvault file exists and contains just the vault name
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
|
||||
currentVaultContents, err := os.ReadFile(filepath.Clean(currentVaultPath))
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to read currentvault file: %v", err)
|
||||
}
|
||||
|
||||
if string(currentVaultContents) != testVaultName {
|
||||
t.Errorf("Expected currentvault to contain %q, got %q",
|
||||
testVaultName, string(currentVaultContents))
|
||||
}
|
||||
|
||||
// Test that ResolveVaultSymlink correctly resolves the path
|
||||
resolvedPath, err := vault.ResolveVaultSymlink(fs, currentVaultPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to resolve currentvault path: %v", err)
|
||||
}
|
||||
|
||||
if resolvedPath != vaultDir {
|
||||
t.Errorf("Expected resolved path to be %s, got %s", vaultDir, resolvedPath)
|
||||
}
|
||||
}
|
||||
|
||||
func testDeepPathSecrets(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "deep-path-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault - CreateVault now handles public key when
|
||||
// mnemonic is in env
|
||||
vlt, err := vault.CreateVault(fs, stateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get its derivation index
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic using the vault's derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic,
|
||||
vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vault
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Create a secret with a deeply nested path
|
||||
deepPath := "api/credentials/production/database/primary"
|
||||
secretValue := []byte("supersecretdbpassword")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vlt.AddSecret(deepPath, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret with deep path: %v", err)
|
||||
}
|
||||
|
||||
// List secrets and verify our deep path secret is there
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(secrets, deepPath) {
|
||||
t.Errorf("Deep path secret not found in listed secrets")
|
||||
}
|
||||
|
||||
// Retrieve the secret and verify its value
|
||||
retrievedValue, err := vlt.GetSecret(deepPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to retrieve deep path secret: %v", err)
|
||||
}
|
||||
|
||||
if string(retrievedValue) != string(expectedValue) {
|
||||
t.Errorf("Retrieved value doesn't match. Expected %q, got %q",
|
||||
string(expectedValue), string(retrievedValue))
|
||||
}
|
||||
}
|
||||
|
||||
func testKeyCaching(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "key-cache-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create a test vault - CreateVault now handles public key when
|
||||
// mnemonic is in env
|
||||
vlt, err := vault.CreateVault(fs, stateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Load vault metadata to get its derivation index
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
vaultMetadata, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term key from mnemonic for verification using the
|
||||
// vault's derivation index
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic,
|
||||
vaultMetadata.DerivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the vault is locked initially
|
||||
if !vlt.Locked() {
|
||||
t.Errorf("Vault should be locked initially")
|
||||
}
|
||||
|
||||
// First call to GetOrDeriveLongTermKey should derive and cache the key
|
||||
firstKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the vault is now unlocked
|
||||
if vlt.Locked() {
|
||||
t.Errorf("Vault should be unlocked after GetOrDeriveLongTermKey")
|
||||
}
|
||||
|
||||
// Second call should return the cached key without re-deriving
|
||||
secondKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get cached long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify both keys are the same instance
|
||||
if firstKey != secondKey {
|
||||
t.Errorf("Second key call should return same instance as first call")
|
||||
}
|
||||
|
||||
// Verify the public key matches what we expect
|
||||
expectedPubKey := ltIdentity.Recipient().String()
|
||||
|
||||
actualPubKey := firstKey.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Public key mismatch. Expected %s, got %s",
|
||||
expectedPubKey, actualPubKey)
|
||||
}
|
||||
|
||||
// Now clear the key and verify it's locked again
|
||||
vlt.ClearLongTermKey()
|
||||
|
||||
if !vlt.Locked() {
|
||||
t.Errorf("Vault should be locked after clearing key")
|
||||
}
|
||||
|
||||
// Get the key again and verify it works
|
||||
thirdKey, err := vlt.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to re-derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Verify the public key still matches
|
||||
actualPubKey = thirdKey.Recipient().String()
|
||||
if actualPubKey != expectedPubKey {
|
||||
t.Errorf("Re-derived public key mismatch. Expected %s, got %s",
|
||||
expectedPubKey, actualPubKey)
|
||||
}
|
||||
}
|
||||
|
||||
func testVaultNameValidation(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "name-validation-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Test valid vault names
|
||||
validNames := []string{
|
||||
"default",
|
||||
"test-vault",
|
||||
"production.vault",
|
||||
"vault_123",
|
||||
"a-very-long-vault-name-with-dashes",
|
||||
}
|
||||
|
||||
for _, name := range validNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Errorf("Failed to create vault with valid name %q: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Test invalid vault names
|
||||
invalidNames := []string{
|
||||
"", // Empty
|
||||
"UPPERCASE", // Uppercase not allowed
|
||||
"invalid/name", // Slashes not allowed in vault names
|
||||
"invalid name", // Spaces not allowed
|
||||
"invalid@name", // Special chars not allowed
|
||||
}
|
||||
|
||||
for _, name := range invalidNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err == nil {
|
||||
t.Errorf("Expected error creating vault with invalid name %q, "+
|
||||
"but got none", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func testMultipleVaults(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "multi-vault-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create three vaults
|
||||
vaultNames := []string{"vault1", "vault2", "vault3"}
|
||||
for _, name := range vaultNames {
|
||||
_, err := vault.CreateVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// List vaults and verify all three are there
|
||||
vaults, err := vault.ListVaults(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list vaults: %v", err)
|
||||
}
|
||||
|
||||
if len(vaults) != 3 {
|
||||
t.Errorf("Expected 3 vaults, got %d", len(vaults))
|
||||
}
|
||||
|
||||
// Test switching between vaults
|
||||
for _, name := range vaultNames {
|
||||
// Select the vault
|
||||
err := vault.SelectVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select vault %s: %v", name, err)
|
||||
}
|
||||
|
||||
// Get current vault and verify it's the one we selected
|
||||
currentVault, err := vault.GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault after selecting %s: %v",
|
||||
name, err)
|
||||
}
|
||||
|
||||
if currentVault.GetName() != name {
|
||||
t.Errorf("Expected current vault to be %s, got %s",
|
||||
name, currentVault.GetName())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func testVaultIsolation(t *testing.T, fs afero.Fs, tempDir string) {
|
||||
t.Helper()
|
||||
|
||||
stateDir := filepath.Join(tempDir, "isolation-test")
|
||||
|
||||
err := os.MkdirAll(stateDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create state dir: %v", err)
|
||||
}
|
||||
|
||||
// Create two vaults - CreateVault now handles public key when mnemonic
|
||||
// is in env
|
||||
vault1, err := vault.CreateVault(fs, stateDir, "vault1")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault1: %v", err)
|
||||
}
|
||||
|
||||
vault2, err := vault.CreateVault(fs, stateDir, "vault2")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault2: %v", err)
|
||||
}
|
||||
|
||||
// Derive long-term keys from mnemonic
|
||||
// Note: Both vaults will have different derivation indexes due to
|
||||
// GetNextDerivationIndex
|
||||
ltIdentity1 := deriveVaultIdentity(t, fs, vault1)
|
||||
ltIdentity2 := deriveVaultIdentity(t, fs, vault2)
|
||||
|
||||
// Unlock the vaults with their respective keys
|
||||
vault1.Unlock(ltIdentity1)
|
||||
vault2.Unlock(ltIdentity2)
|
||||
|
||||
// Add a secret to vault1
|
||||
secretValue := []byte("secret in vault1")
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vault1.AddSecret(testSecretName, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret to vault1: %v", err)
|
||||
}
|
||||
|
||||
// Verify the secret exists in vault1
|
||||
vault1Secrets, err := vault1.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets in vault1: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(vault1Secrets, testSecretName) {
|
||||
t.Errorf("Secret not found in vault1")
|
||||
}
|
||||
|
||||
// Verify the secret does NOT exist in vault2
|
||||
vault2Secrets, err := vault2.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets in vault2: %v", err)
|
||||
}
|
||||
|
||||
if slices.Contains(vault2Secrets, testSecretName) {
|
||||
t.Errorf("Secret from vault1 should not be visible in vault2")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -19,14 +19,17 @@
|
||||
// - Consistent test mnemonic for reproducible keys
|
||||
// - Proper cleanup and isolation between tests
|
||||
|
||||
//nolint:testpackage // uses white-box test helpers shared with this package
|
||||
package vault
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"filippo.io/age"
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/awnumar/memguard"
|
||||
@@ -35,38 +38,33 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// Helper function to add a secret to vault with proper buffer protection
|
||||
func addTestSecret(t *testing.T, vault *Vault, name string, value []byte, force bool) {
|
||||
t.Helper()
|
||||
buffer := memguard.NewBufferFromBytes(value)
|
||||
defer buffer.Destroy()
|
||||
err := vault.AddSecret(name, buffer, force)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
// errUnexpectedValue is returned by concurrent readers when a secret value
|
||||
// does not match the expected contents.
|
||||
var errUnexpectedValue = errors.New("unexpected value")
|
||||
|
||||
// TestVersionIntegrationWorkflow tests the complete version workflow
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv forbids parallel subtests
|
||||
func TestVersionIntegrationWorkflow(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set mnemonic for testing
|
||||
t.Setenv(secret.EnvMnemonic,
|
||||
"abandon abandon abandon abandon abandon abandon "+
|
||||
"abandon abandon abandon abandon abandon about")
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
|
||||
// Create vault
|
||||
vault, err := CreateVault(fs, stateDir, "test")
|
||||
vault, err := CreateVault(fs, testStateDir, "test")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Derive and store long-term key from mnemonic
|
||||
mnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Store long-term public key in vault
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
|
||||
err = afero.WriteFile(fs, ltPubKeyPath,
|
||||
[]byte(ltIdentity.Recipient().String()), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Unlock the vault
|
||||
@@ -76,225 +74,289 @@ func TestVersionIntegrationWorkflow(t *testing.T) {
|
||||
|
||||
// Step 1: Create initial version
|
||||
t.Run("create_initial_version", func(t *testing.T) {
|
||||
addTestSecret(t, vault, secretName, []byte("version-1-data"), false)
|
||||
|
||||
// Verify secret can be retrieved
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value)
|
||||
|
||||
// Verify version directory structure
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 1)
|
||||
|
||||
// Verify current symlink exists
|
||||
currentVersion, err := secret.GetCurrentVersion(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, versions[0], currentVersion)
|
||||
|
||||
// Verify metadata
|
||||
version := secret.NewVersion(vault, secretName, versions[0])
|
||||
err = version.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, version.Metadata.CreatedAt)
|
||||
assert.NotNil(t, version.Metadata.NotBefore)
|
||||
assert.Equal(t, int64(1), version.Metadata.NotBefore.Unix()) // epoch + 1
|
||||
assert.Nil(t, version.Metadata.NotAfter) // should be nil for current version
|
||||
testCreateInitialVersion(t, fs, vault, ltIdentity, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 2: Create second version
|
||||
var firstVersionName string
|
||||
t.Run("create_second_version", func(t *testing.T) {
|
||||
// Small delay to ensure different timestamps
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
// Get first version name before creating second
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
firstVersionName = versions[0]
|
||||
|
||||
// Create second version
|
||||
addTestSecret(t, vault, secretName, []byte("version-2-data"), true)
|
||||
|
||||
// Verify new value is current
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2-data"), value)
|
||||
|
||||
// Verify we now have two versions
|
||||
versions, err = secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 2)
|
||||
|
||||
// Verify first version metadata was updated with notAfter
|
||||
firstVersion := secret.NewVersion(vault, secretName, firstVersionName)
|
||||
err = firstVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, firstVersion.Metadata.NotAfter)
|
||||
|
||||
// Verify second version metadata
|
||||
secondVersion := secret.NewVersion(vault, secretName, versions[0])
|
||||
err = secondVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, secondVersion.Metadata.NotBefore)
|
||||
assert.Nil(t, secondVersion.Metadata.NotAfter)
|
||||
|
||||
// NotBefore of second should equal NotAfter of first
|
||||
assert.Equal(t, firstVersion.Metadata.NotAfter.Unix(), secondVersion.Metadata.NotBefore.Unix())
|
||||
testCreateSecondVersion(t, fs, vault, ltIdentity, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 3: Create third version
|
||||
t.Run("create_third_version", func(t *testing.T) {
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecret(t, vault, secretName, []byte("version-3-data"), true)
|
||||
|
||||
// Verify we now have three versions
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 3)
|
||||
|
||||
// Current should be version-3
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), value)
|
||||
testCreateThirdVersion(t, fs, vault, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 4: Retrieve specific versions
|
||||
t.Run("retrieve_specific_versions", func(t *testing.T) {
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 3)
|
||||
|
||||
// Get each version by its name
|
||||
value1, err := vault.GetSecretVersion(secretName, versions[2]) // oldest
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value1)
|
||||
|
||||
value2, err := vault.GetSecretVersion(secretName, versions[1]) // middle
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2-data"), value2)
|
||||
|
||||
value3, err := vault.GetSecretVersion(secretName, versions[0]) // newest
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), value3)
|
||||
|
||||
// Empty version should return current
|
||||
valueCurrent, err := vault.GetSecretVersion(secretName, "")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), valueCurrent)
|
||||
testRetrieveSpecificVersions(t, fs, vault, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 5: Promote old version to current
|
||||
t.Run("promote_old_version", func(t *testing.T) {
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Promote the first version (oldest) to current
|
||||
oldestVersion := versions[2]
|
||||
err = secret.SetCurrentVersion(fs, secretDir, oldestVersion)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify current now returns the old version's value
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value)
|
||||
|
||||
// Verify the version metadata hasn't changed
|
||||
// (promoting shouldn't modify timestamps)
|
||||
version := secret.NewVersion(vault, secretName, oldestVersion)
|
||||
err = version.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, version.Metadata.NotAfter) // should still have its old notAfter
|
||||
testPromoteOldVersion(t, fs, vault, ltIdentity, vaultDir, secretName)
|
||||
})
|
||||
|
||||
// Step 6: Test version limits
|
||||
t.Run("version_serial_limits", func(t *testing.T) {
|
||||
// Create a new secret for this test
|
||||
limitSecretName := "limit/test"
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "limit%test", "versions")
|
||||
|
||||
// Create 998 versions (we already have one from the first AddSecret)
|
||||
addTestSecret(t, vault, limitSecretName, []byte("initial"), false)
|
||||
|
||||
// Get today's date for consistent version names
|
||||
today := time.Now().Format("20060102")
|
||||
|
||||
// Manually create many versions with same date
|
||||
for i := 2; i <= 998; i++ {
|
||||
versionName := fmt.Sprintf("%s.%03d", today, i)
|
||||
versionDir := filepath.Join(secretDir, versionName)
|
||||
err := fs.MkdirAll(versionDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// Should be able to create one more (999)
|
||||
versionName, err := secret.GenerateVersionName(fs, filepath.Dir(secretDir))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, fmt.Sprintf("%s.999", today), versionName)
|
||||
|
||||
// Create the 999th version directory
|
||||
err = fs.MkdirAll(filepath.Join(secretDir, versionName), 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Should fail to create 1000th version
|
||||
_, err = secret.GenerateVersionName(fs, filepath.Dir(secretDir))
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "exceeded maximum versions per day")
|
||||
testVersionSerialLimits(t, fs, vault, vaultDir)
|
||||
})
|
||||
|
||||
// Step 7: Test error cases
|
||||
t.Run("error_cases", func(t *testing.T) {
|
||||
// Try to get non-existent version
|
||||
_, err := vault.GetSecretVersion(secretName, "99991231.999")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "not found")
|
||||
|
||||
// Try to get version of non-existent secret
|
||||
_, err = vault.GetSecretVersion("nonexistent/secret", "")
|
||||
assert.Error(t, err)
|
||||
|
||||
// Try to add secret without force when it exists
|
||||
failBuffer := memguard.NewBufferFromBytes([]byte("should-fail"))
|
||||
defer failBuffer.Destroy()
|
||||
err = vault.AddSecret(secretName, failBuffer, false)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already exists")
|
||||
testVersionErrorCases(t, vault, secretName)
|
||||
})
|
||||
}
|
||||
|
||||
func testCreateInitialVersion(
|
||||
t *testing.T, fs afero.Fs, vault *Vault,
|
||||
ltIdentity *age.X25519Identity, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-1-data"), false)
|
||||
|
||||
// Verify secret can be retrieved
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value)
|
||||
|
||||
// Verify version directory structure
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 1)
|
||||
|
||||
// Verify current symlink exists
|
||||
currentVersion, err := secret.GetCurrentVersion(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, versions[0], currentVersion)
|
||||
|
||||
// Verify metadata
|
||||
version := secret.NewVersion(vault, secretName, versions[0])
|
||||
err = version.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, version.Metadata.CreatedAt)
|
||||
assert.NotNil(t, version.Metadata.NotBefore)
|
||||
assert.Equal(t, int64(1), version.Metadata.NotBefore.Unix()) // epoch + 1
|
||||
// NotAfter should be nil for current version
|
||||
assert.Nil(t, version.Metadata.NotAfter)
|
||||
}
|
||||
|
||||
func testCreateSecondVersion(
|
||||
t *testing.T, fs afero.Fs, vault *Vault,
|
||||
ltIdentity *age.X25519Identity, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
// Small delay to ensure different timestamps
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
// Get first version name before creating second
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
firstVersionName := versions[0]
|
||||
|
||||
// Create second version
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-2-data"), true)
|
||||
|
||||
// Verify new value is current
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2-data"), value)
|
||||
|
||||
// Verify we now have two versions
|
||||
versions, err = secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 2)
|
||||
|
||||
// Verify first version metadata was updated with notAfter
|
||||
firstVersion := secret.NewVersion(vault, secretName, firstVersionName)
|
||||
err = firstVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, firstVersion.Metadata.NotAfter)
|
||||
|
||||
// Verify second version metadata
|
||||
secondVersion := secret.NewVersion(vault, secretName, versions[0])
|
||||
err = secondVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, secondVersion.Metadata.NotBefore)
|
||||
assert.Nil(t, secondVersion.Metadata.NotAfter)
|
||||
|
||||
// NotBefore of second should equal NotAfter of first
|
||||
assert.Equal(t, firstVersion.Metadata.NotAfter.Unix(),
|
||||
secondVersion.Metadata.NotBefore.Unix())
|
||||
}
|
||||
|
||||
func testCreateThirdVersion(
|
||||
t *testing.T, fs afero.Fs, vault *Vault, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-3-data"), true)
|
||||
|
||||
// Verify we now have three versions
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, versions, 3)
|
||||
|
||||
// Current should be version-3
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), value)
|
||||
}
|
||||
|
||||
func testRetrieveSpecificVersions(
|
||||
t *testing.T, fs afero.Fs, vault *Vault, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 3)
|
||||
|
||||
// Get each version by its name
|
||||
value1, err := vault.GetSecretVersion(secretName, versions[2]) // oldest
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value1)
|
||||
|
||||
value2, err := vault.GetSecretVersion(secretName, versions[1]) // middle
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2-data"), value2)
|
||||
|
||||
value3, err := vault.GetSecretVersion(secretName, versions[0]) // newest
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), value3)
|
||||
|
||||
// Empty version should return current
|
||||
valueCurrent, err := vault.GetSecretVersion(secretName, "")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-3-data"), valueCurrent)
|
||||
}
|
||||
|
||||
func testPromoteOldVersion(
|
||||
t *testing.T, fs afero.Fs, vault *Vault,
|
||||
ltIdentity *age.X25519Identity, vaultDir, secretName string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "integration%test")
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Promote the first version (oldest) to current
|
||||
oldestVersion := versions[2]
|
||||
err = secret.SetCurrentVersion(fs, secretDir, oldestVersion)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Verify current now returns the old version's value
|
||||
value, err := vault.GetSecret(secretName)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1-data"), value)
|
||||
|
||||
// Verify the version metadata hasn't changed
|
||||
// (promoting shouldn't modify timestamps)
|
||||
version := secret.NewVersion(vault, secretName, oldestVersion)
|
||||
err = version.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
// should still have its old notAfter
|
||||
assert.NotNil(t, version.Metadata.NotAfter)
|
||||
}
|
||||
|
||||
func testVersionSerialLimits(
|
||||
t *testing.T, fs afero.Fs, vault *Vault, vaultDir string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
// Create a new secret for this test
|
||||
limitSecretName := "limit/test"
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", "limit%test", "versions")
|
||||
|
||||
// Create 998 versions (we already have one from the first AddSecret)
|
||||
addTestSecretToVault(t, vault, limitSecretName, []byte("initial"), false)
|
||||
|
||||
// Get today's date for consistent version names
|
||||
today := time.Now().Format("20060102")
|
||||
|
||||
// Manually create many versions with same date
|
||||
for i := 2; i <= 998; i++ {
|
||||
versionName := fmt.Sprintf("%s.%03d", today, i)
|
||||
versionDir := filepath.Join(secretDir, versionName)
|
||||
err := fs.MkdirAll(versionDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// Should be able to create one more (999)
|
||||
versionName, err := secret.GenerateVersionName(fs, filepath.Dir(secretDir))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, today+".999", versionName)
|
||||
|
||||
// Create the 999th version directory
|
||||
err = fs.MkdirAll(filepath.Join(secretDir, versionName), 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Should fail to create 1000th version
|
||||
_, err = secret.GenerateVersionName(fs, filepath.Dir(secretDir))
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "exceeded maximum versions per day")
|
||||
}
|
||||
|
||||
func testVersionErrorCases(t *testing.T, vault *Vault, secretName string) {
|
||||
t.Helper()
|
||||
|
||||
// Try to get non-existent version
|
||||
_, err := vault.GetSecretVersion(secretName, "99991231.999")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "not found")
|
||||
|
||||
// Try to get version of non-existent secret
|
||||
_, err = vault.GetSecretVersion("nonexistent/secret", "")
|
||||
require.Error(t, err)
|
||||
|
||||
// Try to add secret without force when it exists
|
||||
failBuffer := memguard.NewBufferFromBytes([]byte("should-fail"))
|
||||
defer failBuffer.Destroy()
|
||||
|
||||
err = vault.AddSecret(secretName, failBuffer, false)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already exists")
|
||||
}
|
||||
|
||||
// TestVersionConcurrency tests concurrent version operations
|
||||
//
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVersionConcurrency(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set up vault
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
secretName := "concurrent/test"
|
||||
|
||||
// Create initial version
|
||||
addTestSecret(t, vault, secretName, []byte("initial"), false)
|
||||
addTestSecretToVault(t, vault, secretName, []byte("initial"), false)
|
||||
|
||||
// Test concurrent reads
|
||||
t.Run("concurrent_reads", func(t *testing.T) {
|
||||
done := make(chan bool, 10)
|
||||
errors := make(chan error, 10)
|
||||
errCh := make(chan error, 10)
|
||||
|
||||
for range 10 {
|
||||
go func() {
|
||||
value, err := vault.GetSecret(secretName)
|
||||
if err != nil {
|
||||
errors <- err
|
||||
errCh <- err
|
||||
} else if string(value) != "initial" {
|
||||
errors <- fmt.Errorf("unexpected value: %s", value)
|
||||
errCh <- fmt.Errorf("%w: %s", errUnexpectedValue, value)
|
||||
}
|
||||
|
||||
done <- true
|
||||
}()
|
||||
}
|
||||
@@ -306,7 +368,7 @@ func TestVersionConcurrency(t *testing.T) {
|
||||
|
||||
// Check for errors
|
||||
select {
|
||||
case err := <-errors:
|
||||
case err := <-errCh:
|
||||
t.Fatalf("concurrent read failed: %v", err)
|
||||
default:
|
||||
// No errors
|
||||
@@ -315,12 +377,14 @@ func TestVersionConcurrency(t *testing.T) {
|
||||
}
|
||||
|
||||
// TestVersionCompatibility tests that old secrets without versions still work
|
||||
//
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVersionCompatibility(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Set up vault
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
ltIdentity, err := vault.GetOrDeriveLongTermKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -333,9 +397,12 @@ func TestVersionCompatibility(t *testing.T) {
|
||||
|
||||
// Create old-style encrypted value directly in secret directory
|
||||
testValue := []byte("legacy-value")
|
||||
|
||||
testValueBuffer := memguard.NewBufferFromBytes(testValue)
|
||||
defer testValueBuffer.Destroy()
|
||||
|
||||
ltRecipient := ltIdentity.Recipient()
|
||||
|
||||
encrypted, err := secret.EncryptToRecipient(testValueBuffer, ltRecipient)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -345,7 +412,7 @@ func TestVersionCompatibility(t *testing.T) {
|
||||
|
||||
// Should fail to get with version-aware methods
|
||||
_, err = vault.GetSecret(secretName)
|
||||
assert.Error(t, err)
|
||||
require.Error(t, err)
|
||||
|
||||
// List versions should return empty
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
|
||||
@@ -15,10 +15,13 @@ import (
|
||||
)
|
||||
|
||||
// Register the GetCurrentVault function with the secret package
|
||||
//
|
||||
//nolint:gochecknoinits // registers the vault accessor with the secret package
|
||||
func init() {
|
||||
secret.RegisterGetCurrentVaultFunc(func(fs afero.Fs, stateDir string) (secret.VaultInterface, error) {
|
||||
return GetCurrentVault(fs, stateDir)
|
||||
})
|
||||
secret.RegisterGetCurrentVaultFunc(
|
||||
func(fs afero.Fs, stateDir string) (secret.VaultInterface, error) {
|
||||
return GetCurrentVault(fs, stateDir)
|
||||
})
|
||||
}
|
||||
|
||||
// isValidVaultName validates vault names according to the format [a-z0-9\.\-\_]+
|
||||
@@ -27,6 +30,7 @@ func isValidVaultName(name string) bool {
|
||||
if name == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
matched, _ := regexp.MatchString(`^[a-z0-9\.\-\_]+$`, name)
|
||||
|
||||
return matched
|
||||
@@ -65,9 +69,11 @@ func GetCurrentVault(fs afero.Fs, stateDir string) (*Vault, error) {
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
|
||||
secret.Debug("Checking current vault symlink", "path", currentVaultPath)
|
||||
|
||||
_, err := fs.Stat(currentVaultPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to stat current vault symlink", "error", err, "path", currentVaultPath)
|
||||
secret.Debug("Failed to stat current vault symlink",
|
||||
"error", err, "path", currentVaultPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read current vault symlink: %w", err)
|
||||
}
|
||||
@@ -76,6 +82,7 @@ func GetCurrentVault(fs afero.Fs, stateDir string) (*Vault, error) {
|
||||
|
||||
// Resolve the symlink to get the actual vault directory
|
||||
secret.Debug("Resolving vault symlink")
|
||||
|
||||
targetPath, err := ResolveVaultSymlink(fs, currentVaultPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -88,7 +95,8 @@ func GetCurrentVault(fs afero.Fs, stateDir string) (*Vault, error) {
|
||||
vaultName := filepath.Base(targetPath)
|
||||
secret.Debug("Extracted vault name", "vault_name", vaultName)
|
||||
|
||||
secret.Debug("Current vault resolved", "vault_name", vaultName, "target_path", targetPath)
|
||||
secret.Debug("Current vault resolved",
|
||||
"vault_name", vaultName, "target_path", targetPath)
|
||||
|
||||
// Create and return the vault
|
||||
return NewVault(fs, stateDir, vaultName), nil
|
||||
@@ -103,6 +111,7 @@ func ListVaults(fs afero.Fs, stateDir string) ([]string, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if vaults directory exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return []string{}, nil
|
||||
}
|
||||
@@ -115,6 +124,7 @@ func ListVaults(fs afero.Fs, stateDir string) ([]string, error) {
|
||||
|
||||
// Extract vault names
|
||||
var vaults []string
|
||||
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
vaults = append(vaults, entry.Name())
|
||||
@@ -124,22 +134,26 @@ func ListVaults(fs afero.Fs, stateDir string) ([]string, error) {
|
||||
return vaults, nil
|
||||
}
|
||||
|
||||
// processMnemonicForVault handles mnemonic processing for vault creation
|
||||
func processMnemonicForVault(fs afero.Fs, stateDir, vaultDir, vaultName string) (
|
||||
derivationIndex uint32, publicKeyHash string, familyHash string, err error) {
|
||||
// processMnemonicForVault handles mnemonic processing for vault creation.
|
||||
// It returns the derivation index, public key hash, and family hash.
|
||||
func processMnemonicForVault(
|
||||
fs afero.Fs, stateDir, vaultDir, vaultName string,
|
||||
) (uint32, string, string, error) {
|
||||
// Check if mnemonic is available in environment
|
||||
mnemonic := os.Getenv(secret.EnvMnemonic)
|
||||
|
||||
if mnemonic == "" {
|
||||
secret.Debug("No mnemonic in environment, vault created without long-term key", "vault", vaultName)
|
||||
secret.Debug("No mnemonic in environment, vault created without long-term key",
|
||||
"vault", vaultName)
|
||||
// Use 0 for derivation index when no mnemonic is provided
|
||||
return 0, "", "", nil
|
||||
}
|
||||
|
||||
secret.Debug("Mnemonic found in environment, deriving long-term key", "vault", vaultName)
|
||||
secret.Debug("Mnemonic found in environment, deriving long-term key",
|
||||
"vault", vaultName)
|
||||
|
||||
// Get the next available derivation index for this mnemonic
|
||||
derivationIndex, err = GetNextDerivationIndex(fs, stateDir, mnemonic)
|
||||
derivationIndex, err := GetNextDerivationIndex(fs, stateDir, mnemonic)
|
||||
if err != nil {
|
||||
return 0, "", "", fmt.Errorf("failed to get next derivation index: %w", err)
|
||||
}
|
||||
@@ -152,14 +166,18 @@ func processMnemonicForVault(fs afero.Fs, stateDir, vaultDir, vaultName string)
|
||||
|
||||
// Write the public key
|
||||
ltPubKey := ltIdentity.Recipient().String()
|
||||
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
if err := afero.WriteFile(fs, ltPubKeyPath, []byte(ltPubKey), secret.FilePerms); err != nil {
|
||||
|
||||
err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltPubKey), secret.FilePerms)
|
||||
if err != nil {
|
||||
return 0, "", "", fmt.Errorf("failed to write long-term public key: %w", err)
|
||||
}
|
||||
|
||||
secret.Debug("Wrote long-term public key", "path", ltPubKeyPath)
|
||||
|
||||
// Compute verification hash from actual derivation index
|
||||
publicKeyHash = ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
publicKeyHash := ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
|
||||
// Compute family hash from index 0 (same for all vaults with this mnemonic)
|
||||
// This is used to identify which vaults belong to the same mnemonic family
|
||||
@@ -167,7 +185,8 @@ func processMnemonicForVault(fs afero.Fs, stateDir, vaultDir, vaultName string)
|
||||
if err != nil {
|
||||
return 0, "", "", fmt.Errorf("failed to derive identity for index 0: %w", err)
|
||||
}
|
||||
familyHash = ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
familyHash := ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
return derivationIndex, publicKeyHash, familyHash, nil
|
||||
}
|
||||
@@ -180,8 +199,12 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) {
|
||||
if !isValidVaultName(name) {
|
||||
secret.Debug("Invalid vault name provided", "vault_name", name)
|
||||
|
||||
return nil, fmt.Errorf("invalid vault name '%s': must match pattern [a-z0-9.\\-_]+", name)
|
||||
return nil, fmt.Errorf(
|
||||
"%w '%s': must match pattern [a-z0-9.\\-_]+",
|
||||
ErrInvalidVaultName, name,
|
||||
)
|
||||
}
|
||||
|
||||
secret.Debug("Vault name validation passed", "vault_name", name)
|
||||
|
||||
// Create vault directory structure
|
||||
@@ -189,24 +212,30 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) {
|
||||
secret.Debug("Creating vault directory structure", "vault_dir", vaultDir)
|
||||
|
||||
// Create main vault directory
|
||||
if err := fs.MkdirAll(vaultDir, secret.DirPerms); err != nil {
|
||||
err := fs.MkdirAll(vaultDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Create secrets directory
|
||||
secretsDir := filepath.Join(vaultDir, "secrets.d")
|
||||
if err := fs.MkdirAll(secretsDir, secret.DirPerms); err != nil {
|
||||
|
||||
err = fs.MkdirAll(secretsDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create secrets directory: %w", err)
|
||||
}
|
||||
|
||||
// Create unlockers directory
|
||||
unlockersDir := filepath.Join(vaultDir, "unlockers.d")
|
||||
if err := fs.MkdirAll(unlockersDir, secret.DirPerms); err != nil {
|
||||
|
||||
err = fs.MkdirAll(unlockersDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create unlockers directory: %w", err)
|
||||
}
|
||||
|
||||
// Process mnemonic if available
|
||||
derivationIndex, publicKeyHash, familyHash, err := processMnemonicForVault(fs, stateDir, vaultDir, name)
|
||||
derivationIndex, publicKeyHash, familyHash, err := processMnemonicForVault(
|
||||
fs, stateDir, vaultDir, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -218,13 +247,17 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) {
|
||||
PublicKeyHash: publicKeyHash,
|
||||
MnemonicFamilyHash: familyHash,
|
||||
}
|
||||
if err := SaveVaultMetadata(fs, vaultDir, metadata); err != nil {
|
||||
|
||||
err = SaveVaultMetadata(fs, vaultDir, metadata)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to save vault metadata: %w", err)
|
||||
}
|
||||
|
||||
// Select the newly created vault as current
|
||||
secret.Debug("Selecting newly created vault as current", "name", name)
|
||||
if err := SelectVault(fs, stateDir, name); err != nil {
|
||||
|
||||
err = SelectVault(fs, stateDir, name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to select vault: %w", err)
|
||||
}
|
||||
|
||||
@@ -242,32 +275,42 @@ func SelectVault(fs afero.Fs, stateDir string, name string) error {
|
||||
if !isValidVaultName(name) {
|
||||
secret.Debug("Invalid vault name provided", "vault_name", name)
|
||||
|
||||
return fmt.Errorf("invalid vault name '%s': must match pattern [a-z0-9.\\-_]+", name)
|
||||
return fmt.Errorf(
|
||||
"%w '%s': must match pattern [a-z0-9.\\-_]+",
|
||||
ErrInvalidVaultName, name,
|
||||
)
|
||||
}
|
||||
|
||||
secret.Debug("Vault name validation passed", "vault_name", name)
|
||||
|
||||
// Check if vault exists
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", name)
|
||||
|
||||
exists, err := afero.DirExists(fs, vaultDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if vault exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return fmt.Errorf("vault %s does not exist", name)
|
||||
return fmt.Errorf("vault %s %w", name, ErrVaultNotFound)
|
||||
}
|
||||
|
||||
// Create or update the currentvault file with just the vault name
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
|
||||
// Remove existing file if it exists
|
||||
if _, err := fs.Stat(currentVaultPath); err == nil {
|
||||
_, 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)
|
||||
if err := afero.WriteFile(fs, currentVaultPath, []byte(name), secret.FilePerms); err != nil {
|
||||
|
||||
err = afero.WriteFile(fs, currentVaultPath, []byte(name), secret.FilePerms)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to select vault: %w", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -34,12 +34,15 @@ func ComputeDoubleSHA256(data []byte) string {
|
||||
|
||||
// GetNextDerivationIndex finds the next available derivation index for a given mnemonic
|
||||
// by deriving the public key for index 0 and using its hash to identify related vaults
|
||||
func GetNextDerivationIndex(fs afero.Fs, stateDir string, mnemonic string) (uint32, error) {
|
||||
func GetNextDerivationIndex(
|
||||
fs afero.Fs, stateDir string, mnemonic string,
|
||||
) (uint32, error) {
|
||||
// First, derive the public key for index 0 to get our identifier
|
||||
identity0, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to derive identity for index 0: %w", err)
|
||||
}
|
||||
|
||||
pubKeyHash := ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
vaultsDir := filepath.Join(stateDir, "vaults.d")
|
||||
@@ -49,6 +52,7 @@ func GetNextDerivationIndex(fs afero.Fs, stateDir string, mnemonic string) (uint
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to check if vaults directory exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
// No vaults yet, start with index 0
|
||||
return 0, nil
|
||||
@@ -70,6 +74,7 @@ func GetNextDerivationIndex(fs afero.Fs, stateDir string, mnemonic string) (uint
|
||||
|
||||
// Try to read vault metadata
|
||||
metadataPath := filepath.Join(vaultsDir, entry.Name(), "vault-metadata.json")
|
||||
|
||||
metadataBytes, err := afero.ReadFile(fs, metadataPath)
|
||||
if err != nil {
|
||||
// Skip vaults without metadata
|
||||
@@ -77,7 +82,9 @@ func GetNextDerivationIndex(fs afero.Fs, stateDir string, mnemonic string) (uint
|
||||
}
|
||||
|
||||
var metadata Metadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
// Skip vaults with invalid metadata
|
||||
continue
|
||||
}
|
||||
@@ -106,7 +113,8 @@ func SaveVaultMetadata(fs afero.Fs, vaultDir string, metadata *Metadata) error {
|
||||
return fmt.Errorf("failed to marshal vault metadata: %w", err)
|
||||
}
|
||||
|
||||
if err := afero.WriteFile(fs, metadataPath, metadataBytes, secret.FilePerms); err != nil {
|
||||
err = afero.WriteFile(fs, metadataPath, metadataBytes, secret.FilePerms)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write vault metadata: %w", err)
|
||||
}
|
||||
|
||||
@@ -123,7 +131,9 @@ func LoadVaultMetadata(fs afero.Fs, vaultDir string) (*Metadata, error) {
|
||||
}
|
||||
|
||||
var metadata Metadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal vault metadata: %w", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,208 +1,243 @@
|
||||
package vault
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
//nolint:paralleltest // subtests share an in-memory filesystem sequentially
|
||||
func TestVaultMetadata(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Test mnemonic for consistent testing
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
t.Run("ComputeDoubleSHA256", func(t *testing.T) {
|
||||
// Test data
|
||||
data := []byte("test data")
|
||||
hash := ComputeDoubleSHA256(data)
|
||||
|
||||
// Verify it's a valid hex string of 64 characters (32 bytes * 2)
|
||||
if len(hash) != 64 {
|
||||
t.Errorf("Expected hash length of 64, got %d", len(hash))
|
||||
}
|
||||
|
||||
// Verify consistency
|
||||
hash2 := ComputeDoubleSHA256(data)
|
||||
if hash != hash2 {
|
||||
t.Errorf("Hash should be consistent for same input")
|
||||
}
|
||||
|
||||
// Verify different input produces different hash
|
||||
hash3 := ComputeDoubleSHA256([]byte("different data"))
|
||||
if hash == hash3 {
|
||||
t.Errorf("Different input should produce different hash")
|
||||
}
|
||||
testComputeDoubleSHA256(t)
|
||||
})
|
||||
|
||||
t.Run("GetNextDerivationIndex", func(t *testing.T) {
|
||||
// Test with no existing vaults
|
||||
index, err := GetNextDerivationIndex(fs, stateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
if index != 0 {
|
||||
t.Errorf("Expected index 0 for first vault, got %d", index)
|
||||
}
|
||||
|
||||
// Create a vault with metadata and matching public key
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "vault1")
|
||||
if err := fs.MkdirAll(vaultDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Derive identity for index 0
|
||||
identity0, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity: %v", err)
|
||||
}
|
||||
pubKey0 := identity0.Recipient().String()
|
||||
pubKeyHash0 := ComputeDoubleSHA256([]byte(pubKey0))
|
||||
|
||||
// Write public key
|
||||
if err := afero.WriteFile(fs, filepath.Join(vaultDir, "pub.age"), []byte(pubKey0), 0o600); err != nil {
|
||||
t.Fatalf("Failed to write public key: %v", err)
|
||||
}
|
||||
|
||||
metadata1 := &Metadata{
|
||||
DerivationIndex: 0,
|
||||
PublicKeyHash: pubKeyHash0, // Hash of the actual key (index 0)
|
||||
MnemonicFamilyHash: pubKeyHash0, // Hash of index 0 key (for family identification)
|
||||
}
|
||||
if err := SaveVaultMetadata(fs, vaultDir, metadata1); err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
// Next index for same mnemonic should be 1
|
||||
index, err = GetNextDerivationIndex(fs, stateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
if index != 1 {
|
||||
t.Errorf("Expected index 1 for second vault with same mnemonic, got %d", index)
|
||||
}
|
||||
|
||||
// Different mnemonic should start at 0
|
||||
differentMnemonic := "zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo wrong"
|
||||
index, err = GetNextDerivationIndex(fs, stateDir, differentMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
if index != 0 {
|
||||
t.Errorf("Expected index 0 for first vault with different mnemonic, got %d", index)
|
||||
}
|
||||
|
||||
// Add another vault with same mnemonic but higher index
|
||||
vaultDir2 := filepath.Join(stateDir, "vaults.d", "vault2")
|
||||
if err := fs.MkdirAll(vaultDir2, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Derive identity for index 5
|
||||
identity5, err := agehd.DeriveIdentity(testMnemonic, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity: %v", err)
|
||||
}
|
||||
pubKey5 := identity5.Recipient().String()
|
||||
|
||||
// Write public key
|
||||
if err := afero.WriteFile(fs, filepath.Join(vaultDir2, "pub.age"), []byte(pubKey5), 0o600); err != nil {
|
||||
t.Fatalf("Failed to write public key: %v", err)
|
||||
}
|
||||
|
||||
// Compute the hash for index 5 key
|
||||
pubKeyHash5 := ComputeDoubleSHA256([]byte(pubKey5))
|
||||
|
||||
metadata2 := &Metadata{
|
||||
DerivationIndex: 5,
|
||||
PublicKeyHash: pubKeyHash5, // Hash of the actual key (index 5)
|
||||
MnemonicFamilyHash: pubKeyHash0, // Same family hash since it's from the same mnemonic
|
||||
}
|
||||
if err := SaveVaultMetadata(fs, vaultDir2, metadata2); err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
// Next index should be 1 (not 6) because we look for the first available slot
|
||||
index, err = GetNextDerivationIndex(fs, stateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
if index != 1 {
|
||||
t.Errorf("Expected index 1 (first available), got %d", index)
|
||||
}
|
||||
testGetNextDerivationIndex(t, fs)
|
||||
})
|
||||
|
||||
t.Run("MetadataPersistence", func(t *testing.T) {
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "test-vault")
|
||||
if err := fs.MkdirAll(vaultDir, 0o700); err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Create and save metadata
|
||||
metadata := &Metadata{
|
||||
DerivationIndex: 3,
|
||||
PublicKeyHash: "test-public-key-hash",
|
||||
}
|
||||
|
||||
if err := SaveVaultMetadata(fs, vaultDir, metadata); err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
// Load and verify
|
||||
loaded, err := LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load metadata: %v", err)
|
||||
}
|
||||
|
||||
if loaded.DerivationIndex != metadata.DerivationIndex {
|
||||
t.Errorf("DerivationIndex mismatch: expected %d, got %d", metadata.DerivationIndex, loaded.DerivationIndex)
|
||||
}
|
||||
if loaded.PublicKeyHash != metadata.PublicKeyHash {
|
||||
t.Errorf("PublicKeyHash mismatch: expected %s, got %s", metadata.PublicKeyHash, loaded.PublicKeyHash)
|
||||
}
|
||||
testMetadataPersistence(t, fs)
|
||||
})
|
||||
|
||||
t.Run("DifferentKeysForDifferentIndices", func(t *testing.T) {
|
||||
// Derive keys with different indices
|
||||
identity0, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity with index 0: %v", err)
|
||||
}
|
||||
|
||||
identity1, err := agehd.DeriveIdentity(testMnemonic, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity with index 1: %v", err)
|
||||
}
|
||||
|
||||
// Compute public key hashes
|
||||
pubKey0 := identity0.Recipient().String()
|
||||
pubKey1 := identity1.Recipient().String()
|
||||
hash0 := ComputeDoubleSHA256([]byte(pubKey0))
|
||||
|
||||
// Verify different indices produce different public keys
|
||||
if pubKey0 == pubKey1 {
|
||||
t.Errorf("Different derivation indices should produce different public keys")
|
||||
}
|
||||
|
||||
// But the hash of index 0's public key should be the same for the same mnemonic
|
||||
// This is what we use as the identifier
|
||||
identity0Again, _ := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
pubKey0Again := identity0Again.Recipient().String()
|
||||
hash0Again := ComputeDoubleSHA256([]byte(pubKey0Again))
|
||||
|
||||
if hash0 != hash0Again {
|
||||
t.Errorf("Same mnemonic should produce same public key hash for index 0")
|
||||
}
|
||||
testDifferentKeysForDifferentIndices(t)
|
||||
})
|
||||
}
|
||||
|
||||
func testComputeDoubleSHA256(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
// Test data
|
||||
data := []byte("test data")
|
||||
hash := vault.ComputeDoubleSHA256(data)
|
||||
|
||||
// Verify it's a valid hex string of 64 characters (32 bytes * 2)
|
||||
if len(hash) != 64 {
|
||||
t.Errorf("Expected hash length of 64, got %d", len(hash))
|
||||
}
|
||||
|
||||
// Verify consistency
|
||||
hash2 := vault.ComputeDoubleSHA256(data)
|
||||
if hash != hash2 {
|
||||
t.Errorf("Hash should be consistent for same input")
|
||||
}
|
||||
|
||||
// Verify different input produces different hash
|
||||
hash3 := vault.ComputeDoubleSHA256([]byte("different data"))
|
||||
if hash == hash3 {
|
||||
t.Errorf("Different input should produce different hash")
|
||||
}
|
||||
}
|
||||
|
||||
// createVaultDirWithMetadata creates a vault directory containing a public
|
||||
// key derived from testMnemonic at the given index plus saved metadata, and
|
||||
// returns the derived public key hash. An empty familyHash defaults to the
|
||||
// derived key's own hash.
|
||||
func createVaultDirWithMetadata(
|
||||
t *testing.T, fs afero.Fs, vaultName string,
|
||||
derivationIndex uint32, familyHash string,
|
||||
) string {
|
||||
t.Helper()
|
||||
|
||||
vaultDir := filepath.Join(testStateDir, "vaults.d", vaultName)
|
||||
|
||||
err := fs.MkdirAll(vaultDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Derive identity for the requested index
|
||||
identity, err := agehd.DeriveIdentity(testMnemonic, derivationIndex)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity: %v", err)
|
||||
}
|
||||
|
||||
pubKey := identity.Recipient().String()
|
||||
pubKeyHash := vault.ComputeDoubleSHA256([]byte(pubKey))
|
||||
|
||||
// Write public key
|
||||
err = afero.WriteFile(fs, filepath.Join(vaultDir, "pub.age"),
|
||||
[]byte(pubKey), 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to write public key: %v", err)
|
||||
}
|
||||
|
||||
if familyHash == "" {
|
||||
familyHash = pubKeyHash
|
||||
}
|
||||
|
||||
metadata := &vault.Metadata{
|
||||
DerivationIndex: derivationIndex,
|
||||
PublicKeyHash: pubKeyHash,
|
||||
MnemonicFamilyHash: familyHash,
|
||||
}
|
||||
|
||||
err = vault.SaveVaultMetadata(fs, vaultDir, metadata)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
return pubKeyHash
|
||||
}
|
||||
|
||||
func testGetNextDerivationIndex(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
// Test with no existing vaults
|
||||
index, err := vault.GetNextDerivationIndex(fs, testStateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
|
||||
if index != 0 {
|
||||
t.Errorf("Expected index 0 for first vault, got %d", index)
|
||||
}
|
||||
|
||||
// Create a vault with metadata and matching public key (index 0; the
|
||||
// family hash is the index 0 key hash)
|
||||
pubKeyHash0 := createVaultDirWithMetadata(t, fs, "vault1", 0, "")
|
||||
|
||||
// Next index for same mnemonic should be 1
|
||||
index, err = vault.GetNextDerivationIndex(fs, testStateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
|
||||
if index != 1 {
|
||||
t.Errorf("Expected index 1 for second vault with same mnemonic, got %d", index)
|
||||
}
|
||||
|
||||
// Different mnemonic should start at 0
|
||||
//nolint:dupword // BIP39-style test mnemonic
|
||||
differentMnemonic := "zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo wrong"
|
||||
|
||||
index, err = vault.GetNextDerivationIndex(fs, testStateDir, differentMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
|
||||
if index != 0 {
|
||||
t.Errorf("Expected index 0 for first vault with different mnemonic, got %d",
|
||||
index)
|
||||
}
|
||||
|
||||
// Add another vault with same mnemonic but higher index (5), sharing
|
||||
// the same family hash since it's from the same mnemonic
|
||||
createVaultDirWithMetadata(t, fs, "vault2", 5, pubKeyHash0)
|
||||
|
||||
// Next index should be 1 (not 6): we look for the first available slot
|
||||
index, err = vault.GetNextDerivationIndex(fs, testStateDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get derivation index: %v", err)
|
||||
}
|
||||
|
||||
if index != 1 {
|
||||
t.Errorf("Expected index 1 (first available), got %d", index)
|
||||
}
|
||||
}
|
||||
|
||||
func testMetadataPersistence(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vaultDir := filepath.Join(testStateDir, "vaults.d", testVaultName)
|
||||
|
||||
err := fs.MkdirAll(vaultDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Create and save metadata
|
||||
metadata := &vault.Metadata{
|
||||
DerivationIndex: 3,
|
||||
PublicKeyHash: "test-public-key-hash",
|
||||
}
|
||||
|
||||
err = vault.SaveVaultMetadata(fs, vaultDir, metadata)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to save metadata: %v", err)
|
||||
}
|
||||
|
||||
// Load and verify
|
||||
loaded, err := vault.LoadVaultMetadata(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load metadata: %v", err)
|
||||
}
|
||||
|
||||
if loaded.DerivationIndex != metadata.DerivationIndex {
|
||||
t.Errorf("DerivationIndex mismatch: expected %d, got %d",
|
||||
metadata.DerivationIndex, loaded.DerivationIndex)
|
||||
}
|
||||
|
||||
if loaded.PublicKeyHash != metadata.PublicKeyHash {
|
||||
t.Errorf("PublicKeyHash mismatch: expected %s, got %s",
|
||||
metadata.PublicKeyHash, loaded.PublicKeyHash)
|
||||
}
|
||||
}
|
||||
|
||||
func testDifferentKeysForDifferentIndices(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
// Derive keys with different indices
|
||||
identity0, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity with index 0: %v", err)
|
||||
}
|
||||
|
||||
identity1, err := agehd.DeriveIdentity(testMnemonic, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity with index 1: %v", err)
|
||||
}
|
||||
|
||||
// Compute public key hashes
|
||||
pubKey0 := identity0.Recipient().String()
|
||||
pubKey1 := identity1.Recipient().String()
|
||||
hash0 := vault.ComputeDoubleSHA256([]byte(pubKey0))
|
||||
|
||||
// Verify different indices produce different public keys
|
||||
if pubKey0 == pubKey1 {
|
||||
t.Errorf("Different derivation indices should produce different public keys")
|
||||
}
|
||||
|
||||
// But the hash of index 0's public key should be the same for the same
|
||||
// mnemonic. This is what we use as the identifier
|
||||
identity0Again, _ := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
pubKey0Again := identity0Again.Recipient().String()
|
||||
hash0Again := vault.ComputeDoubleSHA256([]byte(pubKey0Again))
|
||||
|
||||
if hash0 != hash0Again {
|
||||
t.Errorf("Same mnemonic should produce same public key hash for index 0")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicKeyHashConsistency(t *testing.T) {
|
||||
// Use the same test mnemonic that the integration test uses
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Parallel()
|
||||
|
||||
// Derive identity from index 0 multiple times
|
||||
identity1, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
@@ -223,8 +258,8 @@ func TestPublicKeyHashConsistency(t *testing.T) {
|
||||
}
|
||||
|
||||
// Compute public key hashes
|
||||
hash1 := ComputeDoubleSHA256([]byte(identity1.Recipient().String()))
|
||||
hash2 := ComputeDoubleSHA256([]byte(identity2.Recipient().String()))
|
||||
hash1 := vault.ComputeDoubleSHA256([]byte(identity1.Recipient().String()))
|
||||
hash2 := vault.ComputeDoubleSHA256([]byte(identity2.Recipient().String()))
|
||||
|
||||
// Verify hashes are the same
|
||||
if hash1 != hash2 {
|
||||
@@ -237,11 +272,15 @@ func TestPublicKeyHashConsistency(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestSampleHashCalculation(t *testing.T) {
|
||||
// Test with the exact mnemonic from integration test if available
|
||||
// We'll also test with a few different mnemonics to make sure they produce different hashes
|
||||
t.Parallel()
|
||||
|
||||
// Test with the exact mnemonic from integration test if available. We
|
||||
// also test with a few different mnemonics to make sure they produce
|
||||
// different hashes
|
||||
mnemonics := []string{
|
||||
"abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about",
|
||||
testMnemonic,
|
||||
"legal winner thank year wave sausage worth useful legal winner thank yellow",
|
||||
//nolint:dupword // BIP39-style test mnemonic
|
||||
"zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo zoo wrong",
|
||||
}
|
||||
|
||||
@@ -251,29 +290,29 @@ func TestSampleHashCalculation(t *testing.T) {
|
||||
t.Fatalf("Failed to derive identity for mnemonic %d: %v", i, err)
|
||||
}
|
||||
|
||||
hash := ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
hash := vault.ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
t.Logf("Mnemonic %d hash (index 0): %s", i, hash)
|
||||
t.Logf(" Recipient: %s", identity.Recipient().String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestWorkflowMismatch(t *testing.T) {
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
// Create a temporary directory for testing
|
||||
tempDir := t.TempDir()
|
||||
fs := afero.NewOsFs()
|
||||
|
||||
// Test Case 1: Create vault WITH mnemonic (like init command)
|
||||
t.Setenv("SB_SECRET_MNEMONIC", testMnemonic)
|
||||
_, err := CreateVault(fs, tempDir, "default")
|
||||
|
||||
_, err := vault.CreateVault(fs, tempDir, "default")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault with mnemonic: %v", err)
|
||||
}
|
||||
|
||||
// Load metadata for vault1
|
||||
vault1Dir := filepath.Join(tempDir, "vaults.d", "default")
|
||||
metadata1, err := LoadVaultMetadata(fs, vault1Dir)
|
||||
|
||||
metadata1, err := vault.LoadVaultMetadata(fs, vault1Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault1 metadata: %v", err)
|
||||
}
|
||||
@@ -281,9 +320,10 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
t.Logf("Vault1 (with mnemonic) - DerivationIndex: %d, PublicKeyHash: %s",
|
||||
metadata1.DerivationIndex, metadata1.PublicKeyHash)
|
||||
|
||||
// Test Case 2: Create vault WITHOUT mnemonic, then import (like work vault)
|
||||
// Test Case 2: Create vault WITHOUT mnemonic, then import (work vault)
|
||||
t.Setenv("SB_SECRET_MNEMONIC", "")
|
||||
_, err = CreateVault(fs, tempDir, "work")
|
||||
|
||||
_, err = vault.CreateVault(fs, tempDir, "work")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault without mnemonic: %v", err)
|
||||
}
|
||||
@@ -294,7 +334,7 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
t.Setenv("SB_SECRET_MNEMONIC", testMnemonic)
|
||||
|
||||
// Get the next available derivation index for this mnemonic
|
||||
derivationIndex, err := GetNextDerivationIndex(fs, tempDir, testMnemonic)
|
||||
derivationIndex, err := vault.GetNextDerivationIndex(fs, tempDir, testMnemonic)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get next derivation index: %v", err)
|
||||
}
|
||||
@@ -306,10 +346,12 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity for index 0: %v", err)
|
||||
}
|
||||
publicKeyHash := ComputeDoubleSHA256([]byte(identity0.Recipient().String()))
|
||||
|
||||
publicKeyHash := vault.ComputeDoubleSHA256(
|
||||
[]byte(identity0.Recipient().String()))
|
||||
|
||||
// Load existing metadata and update it (same as in VaultImport)
|
||||
existingMetadata, err := LoadVaultMetadata(fs, vault2Dir)
|
||||
existingMetadata, err := vault.LoadVaultMetadata(fs, vault2Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load existing metadata: %v", err)
|
||||
}
|
||||
@@ -318,12 +360,13 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
existingMetadata.DerivationIndex = derivationIndex
|
||||
existingMetadata.PublicKeyHash = publicKeyHash
|
||||
|
||||
if err := SaveVaultMetadata(fs, vault2Dir, existingMetadata); err != nil {
|
||||
err = vault.SaveVaultMetadata(fs, vault2Dir, existingMetadata)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to save vault metadata: %v", err)
|
||||
}
|
||||
|
||||
// Load updated metadata for vault2
|
||||
metadata2, err := LoadVaultMetadata(fs, vault2Dir)
|
||||
metadata2, err := vault.LoadVaultMetadata(fs, vault2Dir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to load vault2 metadata: %v", err)
|
||||
}
|
||||
@@ -337,57 +380,59 @@ func TestWorkflowMismatch(t *testing.T) {
|
||||
t.Logf("Vault1 hash: %s", metadata1.PublicKeyHash)
|
||||
t.Logf("Vault2 hash: %s", metadata2.PublicKeyHash)
|
||||
} else {
|
||||
t.Logf("SUCCESS: Both vaults have the same public key hash: %s", metadata1.PublicKeyHash)
|
||||
t.Logf("SUCCESS: Both vaults have the same public key hash: %s",
|
||||
metadata1.PublicKeyHash)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReverseEngineerHash(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// This is the hash that the work vault is getting in the failing test
|
||||
wrongHash := "e34a2f500e395d8934a90a99ee9311edcfffd68cb701079575e50cbac7bb9417"
|
||||
correctHash := "992552b00b3879dfae461fab9a084b47784a032771c7a9accaebdde05ec7a7d1"
|
||||
|
||||
// Test mnemonic from integration test
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
|
||||
// Calculate hash for test mnemonic
|
||||
identity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive identity: %v", err)
|
||||
}
|
||||
|
||||
calculatedHash := ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
calculatedHash := vault.ComputeDoubleSHA256(
|
||||
[]byte(identity.Recipient().String()))
|
||||
t.Logf("Test mnemonic hash: %s", calculatedHash)
|
||||
|
||||
if calculatedHash == correctHash {
|
||||
t.Logf("✓ Test mnemonic produces the correct hash")
|
||||
t.Logf("Test mnemonic produces the correct hash")
|
||||
} else {
|
||||
t.Errorf("✗ Test mnemonic does not produce the correct hash")
|
||||
t.Errorf("Test mnemonic does not produce the correct hash")
|
||||
}
|
||||
|
||||
if calculatedHash == wrongHash {
|
||||
t.Logf("✗ Test mnemonic unexpectedly produces the wrong hash")
|
||||
t.Logf("Test mnemonic unexpectedly produces the wrong hash")
|
||||
}
|
||||
|
||||
// Let's try some other possibilities - maybe there's a string normalization issue?
|
||||
// Try some other possibilities: maybe a string normalization issue?
|
||||
variations := []string{
|
||||
"abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about",
|
||||
" abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about ",
|
||||
"abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about\n",
|
||||
strings.TrimSpace("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"),
|
||||
testMnemonic,
|
||||
" " + testMnemonic + " ",
|
||||
testMnemonic + "\n",
|
||||
strings.TrimSpace(testMnemonic),
|
||||
}
|
||||
|
||||
for i, variation := range variations {
|
||||
identity, err := agehd.DeriveIdentity(variation, 0)
|
||||
if err != nil {
|
||||
t.Logf("Variation %d failed: %v", i, err)
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
hash := ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
hash := vault.ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
t.Logf("Variation %d hash: %s", i, hash)
|
||||
|
||||
if hash == wrongHash {
|
||||
t.Logf("✗ Found variation that produces wrong hash: '%s'", variation)
|
||||
t.Logf("Found variation that produces wrong hash: '%s'", variation)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -401,14 +446,15 @@ func TestReverseEngineerHash(t *testing.T) {
|
||||
identity, err := agehd.DeriveIdentity(emptyMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Logf("Empty mnemonic %d failed (expected): %v", i, err)
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
hash := ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
hash := vault.ComputeDoubleSHA256([]byte(identity.Recipient().String()))
|
||||
t.Logf("Empty mnemonic %d hash: %s", i, hash)
|
||||
|
||||
if hash == wrongHash {
|
||||
t.Logf("✗ Empty mnemonic produces wrong hash!")
|
||||
t.Logf("Empty mnemonic produces wrong hash!")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,27 +1,27 @@
|
||||
package vault
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"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"
|
||||
)
|
||||
|
||||
// TestGetSecretVersionRejectsPathTraversal verifies that GetSecretVersion
|
||||
// validates the secret name and rejects path traversal attempts.
|
||||
// This is a regression test for https://git.eeqj.de/sneak/secret/issues/13
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv in parent forbids parallel subtests
|
||||
func TestGetSecretVersionRejectsPathTraversal(t *testing.T) {
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
vlt, err := CreateVault(fs, stateDir, "test-vault")
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add a legitimate secret so the vault is set up
|
||||
@@ -42,42 +42,41 @@ func TestGetSecretVersionRejectsPathTraversal(t *testing.T) {
|
||||
for _, name := range maliciousNames {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
_, err := vlt.GetSecretVersion(name, "")
|
||||
assert.Error(t, err, "GetSecretVersion should reject malicious name: %s", name)
|
||||
assert.Contains(t, err.Error(), "invalid secret name",
|
||||
require.Error(t, err,
|
||||
"GetSecretVersion should reject malicious name: %s", name)
|
||||
require.Contains(t, err.Error(), "invalid secret name",
|
||||
"error should indicate invalid name for: %s", name)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGetSecretRejectsPathTraversal verifies GetSecret (which calls GetSecretVersion)
|
||||
// also rejects path traversal names.
|
||||
// TestGetSecretRejectsPathTraversal verifies GetSecret (which calls
|
||||
// GetSecretVersion) also rejects path traversal names.
|
||||
func TestGetSecretRejectsPathTraversal(t *testing.T) {
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
vlt, err := CreateVault(fs, stateDir, "test-vault")
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = vlt.GetSecret("../../../etc/passwd")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid secret name")
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "invalid secret name")
|
||||
}
|
||||
|
||||
// TestGetSecretObjectRejectsPathTraversal verifies GetSecretObject
|
||||
// also validates names and rejects path traversal attempts.
|
||||
//
|
||||
//nolint:paralleltest // t.Setenv in parent forbids parallel subtests
|
||||
func TestGetSecretObjectRejectsPathTraversal(t *testing.T) {
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
vlt, err := CreateVault(fs, stateDir, "test-vault")
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
require.NoError(t, err)
|
||||
|
||||
maliciousNames := []string{
|
||||
@@ -89,8 +88,8 @@ func TestGetSecretObjectRejectsPathTraversal(t *testing.T) {
|
||||
for _, name := range maliciousNames {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
_, err := vlt.GetSecretObject(name)
|
||||
assert.Error(t, err, "GetSecretObject should reject: %s", name)
|
||||
assert.Contains(t, err.Error(), "invalid secret name")
|
||||
require.Error(t, err, "GetSecretObject should reject: %s", name)
|
||||
require.Contains(t, err.Error(), "invalid secret name")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -21,7 +22,8 @@ func (v *Vault) ListSecrets() ([]string, error) {
|
||||
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory for secret listing", "error", err, "vault_name", v.Name)
|
||||
secret.Debug("Failed to get vault directory for secret listing",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, err
|
||||
}
|
||||
@@ -31,12 +33,15 @@ func (v *Vault) ListSecrets() ([]string, error) {
|
||||
// Check if secrets directory exists
|
||||
exists, err := afero.DirExists(v.fs, secretsDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check secrets directory", "error", err, "secrets_dir", secretsDir)
|
||||
secret.Debug("Failed to check secrets directory",
|
||||
"error", err, "secrets_dir", secretsDir)
|
||||
|
||||
return nil, fmt.Errorf("failed to check if secrets directory exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Debug("Secrets directory does not exist", "secrets_dir", secretsDir, "vault_name", v.Name)
|
||||
secret.Debug("Secrets directory does not exist",
|
||||
"secrets_dir", secretsDir, "vault_name", v.Name)
|
||||
|
||||
return []string{}, nil
|
||||
}
|
||||
@@ -44,12 +49,14 @@ func (v *Vault) ListSecrets() ([]string, error) {
|
||||
// List directories in secrets.d
|
||||
files, err := afero.ReadDir(v.fs, secretsDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read secrets directory", "error", err, "secrets_dir", secretsDir)
|
||||
secret.Debug("Failed to read secrets directory",
|
||||
"error", err, "secrets_dir", secretsDir)
|
||||
|
||||
return nil, fmt.Errorf("failed to read secrets directory: %w", err)
|
||||
}
|
||||
|
||||
var secrets []string
|
||||
|
||||
for _, file := range files {
|
||||
if file.IsDir() {
|
||||
// Convert storage name back to secret name
|
||||
@@ -93,10 +100,8 @@ func isValidSecretName(name string) bool {
|
||||
}
|
||||
|
||||
// Check for path traversal via ".." components
|
||||
for _, part := range strings.Split(name, "/") {
|
||||
if part == ".." {
|
||||
return false
|
||||
}
|
||||
if slices.Contains(strings.Split(name, "/"), "..") {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check the basic pattern
|
||||
@@ -108,7 +113,7 @@ func isValidSecretName(name string) bool {
|
||||
// AddSecret adds a secret to this vault
|
||||
func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool) error {
|
||||
if value == nil {
|
||||
return fmt.Errorf("value buffer is nil")
|
||||
return ErrNilValueBuffer
|
||||
}
|
||||
|
||||
secret.DebugWith("Adding secret to vault",
|
||||
@@ -122,17 +127,24 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool)
|
||||
if !isValidSecretName(name) {
|
||||
secret.Debug("Invalid secret name provided", "secret_name", name)
|
||||
|
||||
return fmt.Errorf("invalid secret name '%s': must match pattern [a-z0-9.\\-_/]+", name)
|
||||
return fmt.Errorf(
|
||||
"%w '%s': must match pattern [a-z0-9.\\-_/]+",
|
||||
ErrInvalidSecretName, name,
|
||||
)
|
||||
}
|
||||
|
||||
secret.Debug("Secret name validation passed", "secret_name", name)
|
||||
|
||||
secret.Debug("Getting vault directory")
|
||||
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory for secret addition", "error", err, "vault_name", v.Name)
|
||||
secret.Debug("Failed to get vault directory for secret addition",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
secret.Debug("Got vault directory", "vault_dir", vaultDir)
|
||||
|
||||
// Convert slashes to percent signs for storage
|
||||
@@ -144,112 +156,30 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool)
|
||||
slog.String("secret_dir", secretDir),
|
||||
)
|
||||
|
||||
// Check if secret already exists
|
||||
secret.Debug("Checking if secret already exists", "secret_dir", secretDir)
|
||||
exists, err := afero.DirExists(v.fs, secretDir)
|
||||
// Check for an existing secret and prepare its directory
|
||||
exists, previousVersion, err := v.prepareSecretDir(name, secretDir, force)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if secret exists", "error", err, "secret_dir", secretDir)
|
||||
|
||||
return fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
return err
|
||||
}
|
||||
secret.Debug("Secret existence check complete", "exists", exists)
|
||||
|
||||
// Handle existing secret case
|
||||
now := time.Now()
|
||||
var previousVersion *secret.Version
|
||||
|
||||
if exists {
|
||||
if !force {
|
||||
secret.Debug("Secret already exists and force not specified", "secret_name", name, "secret_dir", secretDir)
|
||||
|
||||
return fmt.Errorf("secret %s already exists (use --force to overwrite)", name)
|
||||
}
|
||||
|
||||
// Get the current version to update its notAfter timestamp
|
||||
currentVersionName, err := secret.GetCurrentVersion(v.fs, secretDir)
|
||||
if err == nil && currentVersionName != "" {
|
||||
previousVersion = secret.NewVersion(v, name, currentVersionName)
|
||||
// We'll need to load and update its metadata after we unlock the vault
|
||||
}
|
||||
} else {
|
||||
// Create secret directory for new secret
|
||||
secret.Debug("Creating secret directory", "secret_dir", secretDir)
|
||||
if err := v.fs.MkdirAll(secretDir, secret.DirPerms); err != nil {
|
||||
secret.Debug("Failed to create secret directory", "error", err, "secret_dir", secretDir)
|
||||
|
||||
return fmt.Errorf("failed to create secret directory: %w", err)
|
||||
}
|
||||
secret.Debug("Created secret directory successfully")
|
||||
}
|
||||
|
||||
// Generate new version name
|
||||
versionName, err := secret.GenerateVersionName(v.fs, secretDir)
|
||||
// Create the new version and save the encrypted value
|
||||
versionName, err := v.createAndSaveVersion(
|
||||
name, secretDir, value, previousVersion, &now, exists)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to generate version name", "error", err, "secret_name", name)
|
||||
|
||||
return fmt.Errorf("failed to generate version name: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
secret.Debug("Generated new version name", "version", versionName, "secret_name", name)
|
||||
|
||||
// Create new version
|
||||
newVersion := secret.NewVersion(v, name, versionName)
|
||||
|
||||
// Set version timestamps
|
||||
if previousVersion == nil {
|
||||
// First version: notBefore = epoch + 1 second
|
||||
epochPlusOne := time.Unix(1, 0)
|
||||
newVersion.Metadata.NotBefore = &epochPlusOne
|
||||
} else {
|
||||
// New version: notBefore = now
|
||||
newVersion.Metadata.NotBefore = &now
|
||||
|
||||
// We'll update the previous version's notAfter after we save the new version
|
||||
}
|
||||
|
||||
// Save the new version - pass the LockedBuffer directly
|
||||
if err := newVersion.Save(value); 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)
|
||||
}
|
||||
|
||||
// Update previous version if it exists
|
||||
if previousVersion != nil {
|
||||
// Get long-term key to decrypt/encrypt metadata
|
||||
ltIdentity, err := v.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get long-term key for metadata update", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
|
||||
// Load previous version metadata
|
||||
if err := previousVersion.LoadMetadata(ltIdentity); err != nil {
|
||||
secret.Debug("Failed to load previous version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to load previous version metadata: %w", err)
|
||||
}
|
||||
|
||||
// Update notAfter timestamp
|
||||
previousVersion.Metadata.NotAfter = &now
|
||||
|
||||
// Re-save the metadata (we need to implement an update method)
|
||||
if err := updateVersionMetadata(v.fs, previousVersion, ltIdentity); err != nil {
|
||||
secret.Debug("Failed to update previous version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to update previous version metadata: %w", err)
|
||||
}
|
||||
// Update previous version's notAfter timestamp if it exists
|
||||
err = v.updatePreviousVersion(previousVersion, &now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Set current symlink to new version
|
||||
if err := secret.SetCurrentVersion(v.fs, secretDir, versionName); err != nil {
|
||||
err = secret.SetCurrentVersion(v.fs, secretDir, versionName)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to set current version", "error", err, "version", versionName)
|
||||
|
||||
return fmt.Errorf("failed to set current version: %w", err)
|
||||
@@ -263,9 +193,12 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool)
|
||||
}
|
||||
|
||||
// updateVersionMetadata updates the metadata of an existing version
|
||||
func updateVersionMetadata(fs afero.Fs, version *secret.Version, ltIdentity *age.X25519Identity) error {
|
||||
func updateVersionMetadata(
|
||||
fs afero.Fs, version *secret.Version, ltIdentity *age.X25519Identity,
|
||||
) error {
|
||||
// Read the version's encrypted private key
|
||||
encryptedPrivKeyPath := filepath.Join(version.Directory, "priv.age")
|
||||
|
||||
encryptedPrivKey, err := afero.ReadFile(fs, encryptedPrivKeyPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read encrypted version private key: %w", err)
|
||||
@@ -294,14 +227,17 @@ func updateVersionMetadata(fs afero.Fs, version *secret.Version, ltIdentity *age
|
||||
metadataBuffer := memguard.NewBufferFromBytes(metadataBytes)
|
||||
defer metadataBuffer.Destroy()
|
||||
|
||||
encryptedMetadata, err := secret.EncryptToRecipient(metadataBuffer, versionIdentity.Recipient())
|
||||
encryptedMetadata, err := secret.EncryptToRecipient(metadataBuffer,
|
||||
versionIdentity.Recipient())
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to encrypt version metadata: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted metadata
|
||||
metadataPath := filepath.Join(version.Directory, "metadata.age")
|
||||
if err := afero.WriteFile(fs, metadataPath, encryptedMetadata, secret.FilePerms); err != nil {
|
||||
|
||||
err = afero.WriteFile(fs, metadataPath, encryptedMetadata, secret.FilePerms)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write encrypted version metadata: %w", err)
|
||||
}
|
||||
|
||||
@@ -318,7 +254,8 @@ func (v *Vault) GetSecret(name string) ([]byte, error) {
|
||||
return v.GetSecretVersion(name, "")
|
||||
}
|
||||
|
||||
// GetSecretVersion retrieves a specific version of a secret (empty version means current)
|
||||
// GetSecretVersion retrieves a specific version of a secret (empty version
|
||||
// means current)
|
||||
func (v *Vault) GetSecretVersion(name string, version string) ([]byte, error) {
|
||||
secret.DebugWith("Getting secret version from vault",
|
||||
slog.String("vault_name", v.Name),
|
||||
@@ -326,69 +263,17 @@ func (v *Vault) GetSecretVersion(name string, version string) ([]byte, error) {
|
||||
slog.String("version", version),
|
||||
)
|
||||
|
||||
// Validate secret name to prevent path traversal
|
||||
if !isValidSecretName(name) {
|
||||
secret.Debug("Invalid secret name provided", "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("invalid secret name '%s': must match pattern [a-z0-9.\\-_/]+", name)
|
||||
}
|
||||
|
||||
// Get vault directory
|
||||
vaultDir, err := v.GetDirectory()
|
||||
// Validate the name and resolve the version to fetch
|
||||
version, err := v.resolveSecretVersion(name, version)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Convert slashes to percent signs for storage
|
||||
storageName := strings.ReplaceAll(name, "/", "%")
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", storageName)
|
||||
|
||||
// Check if secret exists
|
||||
exists, err := afero.DirExists(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if secret exists", "error", err, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
if !exists {
|
||||
secret.Debug("Secret not found in vault", "secret_name", name, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("secret %s not found", name)
|
||||
}
|
||||
|
||||
// Determine which version to get
|
||||
if version == "" {
|
||||
// Get current version
|
||||
currentVersion, err := secret.GetCurrentVersion(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get current version", "error", err, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("failed to get current version: %w", err)
|
||||
}
|
||||
version = currentVersion
|
||||
secret.Debug("Using current version", "version", version, "secret_name", name)
|
||||
}
|
||||
|
||||
// Create version object
|
||||
secretVersion := secret.NewVersion(v, name, version)
|
||||
|
||||
// Check if version exists
|
||||
versionPath := filepath.Join(secretDir, "versions", version)
|
||||
exists, err = afero.DirExists(v.fs, versionPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if version exists", "error", err, "version", version)
|
||||
|
||||
return nil, fmt.Errorf("failed to check if version exists: %w", err)
|
||||
}
|
||||
if !exists {
|
||||
secret.Debug("Version not found", "version", version, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("version %s not found for secret %s", version, name)
|
||||
}
|
||||
|
||||
secret.Debug("Version exists, proceeding with vault unlock and decryption", "version", version, "secret_name", name)
|
||||
secret.Debug("Version exists, proceeding with vault unlock and decryption",
|
||||
"version", version, "secret_name", name)
|
||||
|
||||
// Unlock the vault (get long-term key in memory)
|
||||
longTermIdentity, err := v.UnlockVault()
|
||||
@@ -406,10 +291,13 @@ func (v *Vault) GetSecretVersion(name string, version string) ([]byte, error) {
|
||||
)
|
||||
|
||||
// Get the version's value
|
||||
secret.Debug("About to call secretVersion.GetValue", "version", version, "secret_name", name)
|
||||
secret.Debug("About to call secretVersion.GetValue",
|
||||
"version", version, "secret_name", name)
|
||||
|
||||
decryptedValue, err := secretVersion.GetValue(longTermIdentity)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to decrypt version value", "error", err, "version", version, "secret_name", name)
|
||||
secret.Debug("Failed to decrypt version value",
|
||||
"error", err, "version", version, "secret_name", name)
|
||||
|
||||
return nil, fmt.Errorf("failed to decrypt version: %w", err)
|
||||
}
|
||||
@@ -442,7 +330,8 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) {
|
||||
|
||||
// If vault is already unlocked, return the cached key
|
||||
if !v.Locked() {
|
||||
secret.Debug("Vault already unlocked, returning cached long-term key", "vault_name", v.Name)
|
||||
secret.Debug("Vault already unlocked, returning cached long-term key",
|
||||
"vault_name", v.Name)
|
||||
|
||||
return v.longTermKey, nil
|
||||
}
|
||||
@@ -450,7 +339,8 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) {
|
||||
// Get or derive the long-term key (but don't store it yet)
|
||||
longTermIdentity, err := v.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get or derive long-term key", "error", err, "vault_name", v.Name)
|
||||
secret.Debug("Failed to get or derive long-term key",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
@@ -469,7 +359,7 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) {
|
||||
// GetSecretObject retrieves a Secret object with metadata loaded from this vault
|
||||
func (v *Vault) GetSecretObject(name string) (*secret.Secret, error) {
|
||||
if !isValidSecretName(name) {
|
||||
return nil, fmt.Errorf("invalid secret name: %s", name)
|
||||
return nil, fmt.Errorf("%w: %s", ErrInvalidSecretName, name)
|
||||
}
|
||||
|
||||
// First check if the secret exists by checking for the metadata file
|
||||
@@ -487,15 +377,17 @@ func (v *Vault) GetSecretObject(name string) (*secret.Secret, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("secret %s not found", name)
|
||||
return nil, fmt.Errorf("secret %s %w", name, ErrSecretNotFound)
|
||||
}
|
||||
|
||||
// Create a Secret object
|
||||
secretObj := secret.NewSecret(v, name)
|
||||
|
||||
// Load the metadata from disk
|
||||
if err := secretObj.LoadMetadata(); err != nil {
|
||||
err = secretObj.LoadMetadata()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -526,7 +418,8 @@ func (v *Vault) CopySecretVersion(
|
||||
defer valueBuffer.Destroy()
|
||||
|
||||
// Load source metadata
|
||||
if err := srcVersion.LoadMetadata(srcIdentity); err != nil {
|
||||
err = srcVersion.LoadMetadata(srcIdentity)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to load source metadata: %w", err)
|
||||
}
|
||||
|
||||
@@ -537,7 +430,8 @@ func (v *Vault) CopySecretVersion(
|
||||
destVersion.Metadata = srcVersion.Metadata
|
||||
|
||||
// Save the version (encrypts to this vault's LT key)
|
||||
if err := destVersion.Save(valueBuffer); err != nil {
|
||||
err = destVersion.Save(valueBuffer)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to save destination version: %w", err)
|
||||
}
|
||||
|
||||
@@ -571,26 +465,13 @@ func (v *Vault) CopySecretAllVersions(
|
||||
return fmt.Errorf("failed to get destination vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Check if destination secret already exists
|
||||
// Check if destination secret already exists and clear it if forced
|
||||
destStorageName := strings.ReplaceAll(destSecretName, "/", "%")
|
||||
destSecretDir := filepath.Join(destVaultDir, "secrets.d", destStorageName)
|
||||
|
||||
exists, err := afero.DirExists(v.fs, destSecretDir)
|
||||
err = v.prepareCopyDestination(destSecretDir, destSecretName, force)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check destination: %w", err)
|
||||
}
|
||||
|
||||
if exists && !force {
|
||||
return fmt.Errorf("secret '%s' already exists in vault '%s' (use --force to overwrite)",
|
||||
destSecretName, v.Name)
|
||||
}
|
||||
|
||||
if exists && force {
|
||||
// Remove existing secret
|
||||
secret.Debug("Removing existing destination secret", "path", destSecretDir)
|
||||
if err := v.fs.RemoveAll(destSecretDir); err != nil {
|
||||
return fmt.Errorf("failed to remove existing destination secret: %w", err)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Get source vault's long-term key
|
||||
@@ -615,7 +496,7 @@ func (v *Vault) CopySecretAllVersions(
|
||||
}
|
||||
|
||||
if len(versions) == 0 {
|
||||
return fmt.Errorf("source secret '%s' has no versions", srcSecretName)
|
||||
return fmt.Errorf("source secret '%s' %w", srcSecretName, ErrNoVersions)
|
||||
}
|
||||
|
||||
// Get current version name
|
||||
@@ -625,27 +506,16 @@ func (v *Vault) CopySecretAllVersions(
|
||||
}
|
||||
|
||||
// Create destination secret directory
|
||||
if err := v.fs.MkdirAll(destSecretDir, secret.DirPerms); err != nil {
|
||||
err = v.fs.MkdirAll(destSecretDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create destination secret directory: %w", err)
|
||||
}
|
||||
|
||||
// Copy each version
|
||||
for _, versionName := range versions {
|
||||
srcVersion := secret.NewVersion(srcVault, srcSecretName, versionName)
|
||||
if err := v.CopySecretVersion(srcVersion, srcIdentity, destSecretName, versionName); 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
|
||||
if err := secret.SetCurrentVersion(v.fs, destSecretDir, currentVersion); err != nil {
|
||||
_ = v.fs.RemoveAll(destSecretDir)
|
||||
|
||||
return fmt.Errorf("failed to set current version: %w", err)
|
||||
// Copy each version and set the current pointer, rolling back on error
|
||||
err = v.copyVersionsWithRollback(srcVault, srcIdentity,
|
||||
srcSecretName, destSecretName, destSecretDir, versions, currentVersion)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully copied all secret versions",
|
||||
@@ -656,3 +526,292 @@ 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(
|
||||
name, secretDir string, force bool,
|
||||
) (bool, *secret.Version, error) {
|
||||
// Check if secret already exists
|
||||
secret.Debug("Checking if secret already exists", "secret_dir", secretDir)
|
||||
|
||||
exists, err := afero.DirExists(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if secret exists",
|
||||
"error", err, "secret_dir", secretDir)
|
||||
|
||||
return false, nil, fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
if !force {
|
||||
secret.Debug("Secret already exists and force not specified",
|
||||
"secret_name", name, "secret_dir", secretDir)
|
||||
|
||||
return true, nil, fmt.Errorf(
|
||||
"secret %s %w (use --force to overwrite)",
|
||||
name, ErrSecretExists,
|
||||
)
|
||||
}
|
||||
|
||||
// Get the current version to update its notAfter timestamp
|
||||
var previousVersion *secret.Version
|
||||
|
||||
currentVersionName, err := secret.GetCurrentVersion(v.fs, secretDir)
|
||||
if err == nil && currentVersionName != "" {
|
||||
previousVersion = secret.NewVersion(v, name, currentVersionName)
|
||||
// We'll need to load and update its metadata after we unlock the vault
|
||||
}
|
||||
|
||||
return true, previousVersion, nil
|
||||
}
|
||||
|
||||
// updatePreviousVersion sets the notAfter timestamp on the version being
|
||||
// superseded. It is a no-op when previousVersion is nil.
|
||||
func (v *Vault) updatePreviousVersion(
|
||||
previousVersion *secret.Version, now *time.Time,
|
||||
) error {
|
||||
if previousVersion == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get long-term key to decrypt/encrypt metadata
|
||||
ltIdentity, err := v.GetOrDeriveLongTermKey()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get long-term key for metadata update", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to get long-term key: %w", err)
|
||||
}
|
||||
|
||||
// Load previous version metadata
|
||||
err = previousVersion.LoadMetadata(ltIdentity)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to load previous version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to load previous version metadata: %w", err)
|
||||
}
|
||||
|
||||
// Update notAfter timestamp
|
||||
previousVersion.Metadata.NotAfter = now
|
||||
|
||||
// Re-save the metadata (we need to implement an update method)
|
||||
err = updateVersionMetadata(v.fs, previousVersion, ltIdentity)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to update previous version metadata", "error", err)
|
||||
|
||||
return fmt.Errorf("failed to update previous version metadata: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolveSecretVersion validates the secret name, verifies the secret and
|
||||
// version exist, and resolves an empty version to the current one.
|
||||
func (v *Vault) resolveSecretVersion(name, version string) (string, error) {
|
||||
// Validate secret name to prevent path traversal
|
||||
if !isValidSecretName(name) {
|
||||
secret.Debug("Invalid secret name provided", "secret_name", name)
|
||||
|
||||
return "", fmt.Errorf(
|
||||
"%w '%s': must match pattern [a-z0-9.\\-_/]+",
|
||||
ErrInvalidSecretName, name,
|
||||
)
|
||||
}
|
||||
|
||||
// Get vault directory
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory", "error", err, "vault_name", v.Name)
|
||||
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Convert slashes to percent signs for storage
|
||||
storageName := strings.ReplaceAll(name, "/", "%")
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", storageName)
|
||||
|
||||
// Check if secret exists
|
||||
exists, err := afero.DirExists(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if secret exists", "error", err, "secret_name", name)
|
||||
|
||||
return "", fmt.Errorf("failed to check if secret exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Debug("Secret not found in vault", "secret_name", name, "vault_name", v.Name)
|
||||
|
||||
return "", fmt.Errorf("secret %s %w", name, ErrSecretNotFound)
|
||||
}
|
||||
|
||||
// Determine which version to get
|
||||
if version == "" {
|
||||
// Get current version
|
||||
currentVersion, err := secret.GetCurrentVersion(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get current version", "error", err, "secret_name", name)
|
||||
|
||||
return "", fmt.Errorf("failed to get current version: %w", err)
|
||||
}
|
||||
|
||||
version = currentVersion
|
||||
|
||||
secret.Debug("Using current version", "version", version, "secret_name", name)
|
||||
}
|
||||
|
||||
// Check if version exists
|
||||
versionPath := filepath.Join(secretDir, "versions", version)
|
||||
|
||||
exists, err = afero.DirExists(v.fs, versionPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to check if version exists", "error", err, "version", version)
|
||||
|
||||
return "", fmt.Errorf("failed to check if version exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Debug("Version not found", "version", version, "secret_name", name)
|
||||
|
||||
return "", fmt.Errorf(
|
||||
"version %s %w %s",
|
||||
version, ErrVersionNotFound, name,
|
||||
)
|
||||
}
|
||||
|
||||
return version, nil
|
||||
}
|
||||
|
||||
// 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.
|
||||
func (v *Vault) createAndSaveVersion(
|
||||
name, secretDir string, value *memguard.LockedBuffer,
|
||||
previousVersion *secret.Version, now *time.Time, exists bool,
|
||||
) (string, error) {
|
||||
// Generate new version name
|
||||
versionName, err := secret.GenerateVersionName(v.fs, secretDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to generate version name", "error", err, "secret_name", name)
|
||||
|
||||
return "", fmt.Errorf("failed to generate version name: %w", err)
|
||||
}
|
||||
|
||||
secret.Debug("Generated new version name", "version", versionName, "secret_name", name)
|
||||
|
||||
// Create new version
|
||||
newVersion := secret.NewVersion(v, name, versionName)
|
||||
|
||||
// Set version timestamps
|
||||
if previousVersion == nil {
|
||||
// First version: notBefore = epoch + 1 second
|
||||
epochPlusOne := time.Unix(1, 0)
|
||||
newVersion.Metadata.NotBefore = &epochPlusOne
|
||||
} else {
|
||||
// New version: notBefore = now
|
||||
newVersion.Metadata.NotBefore = now
|
||||
|
||||
// We'll update the previous version's notAfter after we save the
|
||||
// new version
|
||||
}
|
||||
|
||||
// Save the new version - pass the LockedBuffer directly
|
||||
err = newVersion.Save(value)
|
||||
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(
|
||||
srcVault *Vault, srcIdentity *age.X25519Identity,
|
||||
srcSecretName, destSecretName, destSecretDir string,
|
||||
versions []string, currentVersion string,
|
||||
) error {
|
||||
// Copy each version
|
||||
for _, versionName := range versions {
|
||||
srcVersion := secret.NewVersion(srcVault, srcSecretName, versionName)
|
||||
|
||||
err := v.CopySecretVersion(srcVersion, srcIdentity, destSecretName, 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)
|
||||
if err != nil {
|
||||
_ = v.fs.RemoveAll(destSecretDir)
|
||||
|
||||
return fmt.Errorf("failed to set current version: %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(
|
||||
destSecretDir, destSecretName string, force bool,
|
||||
) error {
|
||||
exists, err := afero.DirExists(v.fs, destSecretDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check destination: %w", err)
|
||||
}
|
||||
|
||||
if exists && !force {
|
||||
return fmt.Errorf(
|
||||
"secret '%s' %w in vault '%s' (use --force to overwrite)",
|
||||
destSecretName, ErrSecretExists, v.Name,
|
||||
)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
//nolint:testpackage // white-box test of unexported isValidSecretName
|
||||
package vault
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestIsValidSecretNameUppercase(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
valid bool
|
||||
@@ -33,6 +36,8 @@ func TestIsValidSecretNameUppercase(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
result := isValidSecretName(tt.name)
|
||||
if result != tt.valid {
|
||||
t.Errorf("isValidSecretName(%q) = %v, want %v", tt.name, result, tt.valid)
|
||||
|
||||
@@ -2,10 +2,14 @@
|
||||
//
|
||||
// Integration tests for vault-level version operations:
|
||||
//
|
||||
// - TestVaultAddSecretCreatesVersion: Tests that AddSecret creates proper version structure
|
||||
// - TestVaultAddSecretMultipleVersions: Tests creating multiple versions with force flag
|
||||
// - TestVaultGetSecretVersion: Tests retrieving specific versions and current version
|
||||
// - TestVaultVersionTimestamps: Tests timestamp logic (notBefore/notAfter) across versions
|
||||
// - TestVaultAddSecretCreatesVersion: Tests that AddSecret creates proper
|
||||
// version structure
|
||||
// - TestVaultAddSecretMultipleVersions: Tests creating multiple versions with
|
||||
// force flag
|
||||
// - TestVaultGetSecretVersion: Tests retrieving specific versions and current
|
||||
// version
|
||||
// - TestVaultVersionTimestamps: Tests timestamp logic (notBefore/notAfter)
|
||||
// across versions
|
||||
// - TestVaultGetNonExistentVersion: Tests error handling for invalid versions
|
||||
// - TestUpdateVersionMetadata: Tests metadata update functionality
|
||||
//
|
||||
@@ -15,6 +19,7 @@
|
||||
// - Promotion doesn't modify timestamps
|
||||
// - Metadata remains encrypted and intact
|
||||
|
||||
//nolint:testpackage // white-box test of unexported updateVersionMetadata
|
||||
package vault
|
||||
|
||||
import (
|
||||
@@ -30,33 +35,61 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// testMnemonic is the mnemonic used to derive the vault long-term key.
|
||||
//
|
||||
//nolint:dupword // BIP39 test mnemonic intentionally repeats a word
|
||||
const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon abandon about"
|
||||
|
||||
// envTestMnemonic is the (deliberately different) mnemonic placed in the
|
||||
// environment; the vault is unlocked manually with the derived key in
|
||||
// createTestVaultWithKey.
|
||||
//
|
||||
//nolint:dupword // BIP39-style test mnemonic intentionally repeats a word
|
||||
const envTestMnemonic = "abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon about"
|
||||
|
||||
// Shared fixtures for white-box tests in this package.
|
||||
const (
|
||||
testStateDir = "/test/state"
|
||||
testSecretPath = "test/secret"
|
||||
)
|
||||
|
||||
// Helper function to add a secret to vault with proper buffer protection
|
||||
func addTestSecretToVault(t *testing.T, vault *Vault, name string, value []byte, force bool) {
|
||||
func addTestSecretToVault(
|
||||
t *testing.T, vault *Vault, name string, value []byte, force bool,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
buffer := memguard.NewBufferFromBytes(value)
|
||||
defer buffer.Destroy()
|
||||
|
||||
err := vault.AddSecret(name, buffer, force)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
// Helper function to create a vault with long-term key set up
|
||||
func createTestVaultWithKey(t *testing.T, fs afero.Fs, stateDir, vaultName string) *Vault {
|
||||
// Helper function to create a vault named "test" with its long-term key set
|
||||
// up and unlocked
|
||||
func createTestVaultWithKey(t *testing.T, fs afero.Fs) *Vault {
|
||||
t.Helper()
|
||||
|
||||
// Set mnemonic for testing
|
||||
t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon about")
|
||||
t.Setenv(secret.EnvMnemonic, envTestMnemonic)
|
||||
|
||||
// Create vault
|
||||
vault, err := CreateVault(fs, stateDir, vaultName)
|
||||
vault, err := CreateVault(fs, testStateDir, "test")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Derive and store long-term key from mnemonic
|
||||
mnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
ltIdentity, err := agehd.DeriveIdentity(mnemonic, 0)
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Store long-term public key in vault
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
ltPubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, ltPubKeyPath, []byte(ltIdentity.Recipient().String()), 0o600)
|
||||
|
||||
err = afero.WriteFile(fs, ltPubKeyPath,
|
||||
[]byte(ltIdentity.Recipient().String()), 0o600)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Unlock the vault with the derived key
|
||||
@@ -65,20 +98,19 @@ func createTestVaultWithKey(t *testing.T, fs afero.Fs, stateDir, vaultName strin
|
||||
return vault
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultAddSecretCreatesVersion(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Add a secret
|
||||
secretName := "test/secret"
|
||||
secretValue := []byte("initial-value")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
addTestSecretToVault(t, vault, secretName, secretValue, false)
|
||||
addTestSecretToVault(t, vault, testSecretPath, secretValue, false)
|
||||
|
||||
// Check that version directory was created
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
@@ -97,32 +129,31 @@ func TestVaultAddSecretCreatesVersion(t *testing.T) {
|
||||
assert.True(t, exists)
|
||||
|
||||
// Get the secret value
|
||||
retrievedValue, err := vault.GetSecret(secretName)
|
||||
retrievedValue, err := vault.GetSecret(testSecretPath)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, expectedValue, retrievedValue)
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultAddSecretMultipleVersions(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
|
||||
secretName := "test/secret"
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Add first version
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-1"), false)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-1"), false)
|
||||
|
||||
// Try to add again without force - should fail
|
||||
failBuffer := memguard.NewBufferFromBytes([]byte("version-2"))
|
||||
defer failBuffer.Destroy()
|
||||
err := vault.AddSecret(secretName, failBuffer, false)
|
||||
assert.Error(t, err)
|
||||
|
||||
err := vault.AddSecret(testSecretPath, failBuffer, false)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already exists")
|
||||
|
||||
// Add with force - should create new version
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-2"), true)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-2"), true)
|
||||
|
||||
// Check that we have two versions
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
@@ -132,27 +163,25 @@ func TestVaultAddSecretMultipleVersions(t *testing.T) {
|
||||
assert.Len(t, entries, 2)
|
||||
|
||||
// Current value should be version-2
|
||||
value, err := vault.GetSecret(secretName)
|
||||
value, err := vault.GetSecret(testSecretPath)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2"), value)
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultGetSecretVersion(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
|
||||
secretName := "test/secret"
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Add multiple versions
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-1"), false)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-1"), false)
|
||||
|
||||
// Small delay to ensure different version names
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-2"), true)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-2"), true)
|
||||
|
||||
// Get versions list
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
@@ -163,58 +192,62 @@ func TestVaultGetSecretVersion(t *testing.T) {
|
||||
|
||||
// Get specific version (first one)
|
||||
firstVersion := versions[1] // Last in list is first created
|
||||
value, err := vault.GetSecretVersion(secretName, firstVersion)
|
||||
value, err := vault.GetSecretVersion(testSecretPath, firstVersion)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-1"), value)
|
||||
|
||||
// Get specific version (second one)
|
||||
secondVersion := versions[0] // First in list is most recent
|
||||
value, err = vault.GetSecretVersion(secretName, secondVersion)
|
||||
value, err = vault.GetSecretVersion(testSecretPath, secondVersion)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2"), value)
|
||||
|
||||
// Get current (empty version)
|
||||
value, err = vault.GetSecretVersion(secretName, "")
|
||||
value, err = vault.GetSecretVersion(testSecretPath, "")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("version-2"), value)
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultVersionTimestamps(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Get long-term key
|
||||
ltIdentity, err := vault.GetOrDeriveLongTermKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
secretName := "test/secret"
|
||||
|
||||
// Add first version
|
||||
beforeFirst := time.Now()
|
||||
|
||||
v1Buffer := memguard.NewBufferFromBytes([]byte("version-1"))
|
||||
defer v1Buffer.Destroy()
|
||||
err = vault.AddSecret(secretName, v1Buffer, false)
|
||||
|
||||
err = vault.AddSecret(testSecretPath, v1Buffer, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
afterFirst := time.Now()
|
||||
|
||||
// Get first version metadata
|
||||
vaultDir, _ := vault.GetDirectory()
|
||||
secretDir := vaultDir + "/secrets.d/test%secret"
|
||||
|
||||
versions, err := secret.ListVersions(fs, secretDir)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, versions, 1)
|
||||
|
||||
firstVersion := secret.NewVersion(vault, secretName, versions[0])
|
||||
firstVersion := secret.NewVersion(vault, testSecretPath, versions[0])
|
||||
err = firstVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Check first version timestamps
|
||||
assert.NotNil(t, firstVersion.Metadata.CreatedAt)
|
||||
assert.True(t, firstVersion.Metadata.CreatedAt.After(beforeFirst.Add(-time.Second)))
|
||||
assert.True(t, firstVersion.Metadata.CreatedAt.Before(afterFirst.Add(time.Second)))
|
||||
assert.True(t,
|
||||
firstVersion.Metadata.CreatedAt.After(beforeFirst.Add(-time.Second)))
|
||||
assert.True(t,
|
||||
firstVersion.Metadata.CreatedAt.Before(afterFirst.Add(time.Second)))
|
||||
|
||||
assert.NotNil(t, firstVersion.Metadata.NotBefore)
|
||||
assert.Equal(t, int64(1), firstVersion.Metadata.NotBefore.Unix()) // Epoch + 1
|
||||
@@ -222,8 +255,11 @@ func TestVaultVersionTimestamps(t *testing.T) {
|
||||
|
||||
// Add second version
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
beforeSecond := time.Now()
|
||||
addTestSecretToVault(t, vault, secretName, []byte("version-2"), true)
|
||||
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("version-2"), true)
|
||||
|
||||
afterSecond := time.Now()
|
||||
|
||||
// Get updated versions
|
||||
@@ -232,56 +268,59 @@ func TestVaultVersionTimestamps(t *testing.T) {
|
||||
require.Len(t, versions, 2)
|
||||
|
||||
// Reload first version metadata (should have notAfter now)
|
||||
firstVersion = secret.NewVersion(vault, secretName, versions[1])
|
||||
firstVersion = secret.NewVersion(vault, testSecretPath, versions[1])
|
||||
err = firstVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.NotNil(t, firstVersion.Metadata.NotAfter)
|
||||
assert.True(t, firstVersion.Metadata.NotAfter.After(beforeSecond.Add(-time.Second)))
|
||||
assert.True(t, firstVersion.Metadata.NotAfter.Before(afterSecond.Add(time.Second)))
|
||||
assert.True(t,
|
||||
firstVersion.Metadata.NotAfter.After(beforeSecond.Add(-time.Second)))
|
||||
assert.True(t,
|
||||
firstVersion.Metadata.NotAfter.Before(afterSecond.Add(time.Second)))
|
||||
|
||||
// Check second version timestamps
|
||||
secondVersion := secret.NewVersion(vault, secretName, versions[0])
|
||||
secondVersion := secret.NewVersion(vault, testSecretPath, versions[0])
|
||||
err = secondVersion.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.NotNil(t, secondVersion.Metadata.NotBefore)
|
||||
assert.True(t, secondVersion.Metadata.NotBefore.After(beforeSecond.Add(-time.Second)))
|
||||
assert.True(t, secondVersion.Metadata.NotBefore.Before(afterSecond.Add(time.Second)))
|
||||
assert.True(t,
|
||||
secondVersion.Metadata.NotBefore.After(beforeSecond.Add(-time.Second)))
|
||||
assert.True(t,
|
||||
secondVersion.Metadata.NotBefore.Before(afterSecond.Add(time.Second)))
|
||||
assert.Nil(t, secondVersion.Metadata.NotAfter) // Current version
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestVaultGetNonExistentVersion(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Add a secret
|
||||
addTestSecretToVault(t, vault, "test/secret", []byte("value"), false)
|
||||
addTestSecretToVault(t, vault, testSecretPath, []byte("value"), false)
|
||||
|
||||
// Try to get non-existent version
|
||||
_, err := vault.GetSecretVersion("test/secret", "20991231.999")
|
||||
assert.Error(t, err)
|
||||
_, err := vault.GetSecretVersion(testSecretPath, "20991231.999")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "not found")
|
||||
}
|
||||
|
||||
//nolint:paralleltest // createTestVaultWithKey uses t.Setenv
|
||||
func TestUpdateVersionMetadata(t *testing.T) {
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault with long-term key
|
||||
vault := createTestVaultWithKey(t, fs, stateDir, "test")
|
||||
vault := createTestVaultWithKey(t, fs)
|
||||
|
||||
// Get long-term key
|
||||
ltIdentity, err := vault.GetOrDeriveLongTermKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create a version manually to test updateVersionMetadata
|
||||
secretName := "test/secret"
|
||||
versionName := "20231215.001"
|
||||
version := secret.NewVersion(vault, secretName, versionName)
|
||||
version := secret.NewVersion(vault, testSecretPath, versionName)
|
||||
|
||||
// Set initial metadata
|
||||
now := time.Now()
|
||||
@@ -292,6 +331,7 @@ func TestUpdateVersionMetadata(t *testing.T) {
|
||||
// Save version
|
||||
testBuffer := memguard.NewBufferFromBytes([]byte("test-value"))
|
||||
defer testBuffer.Destroy()
|
||||
|
||||
err = version.Save(testBuffer)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -301,7 +341,7 @@ func TestUpdateVersionMetadata(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
// Load and verify
|
||||
version2 := secret.NewVersion(vault, secretName, versionName)
|
||||
version2 := secret.NewVersion(vault, testSecretPath, versionName)
|
||||
err = version2.LoadMetadata(ltIdentity)
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
@@ -14,13 +14,22 @@ import (
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// Unlocker metadata type strings.
|
||||
const (
|
||||
unlockerTypePassphrase = "passphrase"
|
||||
unlockerTypeSecureEnclave = "secure-enclave"
|
||||
)
|
||||
|
||||
// GetCurrentUnlocker returns the current unlocker for this vault
|
||||
//
|
||||
//nolint:ireturn // returns one of several concrete unlocker implementations
|
||||
func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
secret.DebugWith("Getting current unlocker", slog.String("vault_name", v.Name))
|
||||
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
secret.Debug("Failed to get vault directory for unlocker", "error", err, "vault_name", v.Name)
|
||||
secret.Debug("Failed to get vault directory for unlocker",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, err
|
||||
}
|
||||
@@ -30,7 +39,8 @@ func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
// Check if the symlink exists
|
||||
_, err = v.fs.Stat(currentUnlockerPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to stat current unlocker symlink", "error", err, "path", currentUnlockerPath)
|
||||
secret.Debug("Failed to stat current unlocker symlink",
|
||||
"error", err, "path", currentUnlockerPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read current unlocker: %w", err)
|
||||
}
|
||||
@@ -47,49 +57,37 @@ func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
)
|
||||
|
||||
// Read unlocker metadata
|
||||
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json")
|
||||
secret.Debug("Reading unlocker metadata", "path", metadataPath)
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
metadata, err := v.readUnlockerMetadata(unlockerDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read unlocker metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to read unlocker metadata: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
secret.Debug("Failed to parse unlocker metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return nil, fmt.Errorf("failed to parse unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
secret.DebugWith("Parsed unlocker metadata",
|
||||
slog.String("unlocker_type", metadata.Type),
|
||||
slog.Time("created_at", metadata.CreatedAt),
|
||||
slog.Any("flags", metadata.Flags),
|
||||
)
|
||||
|
||||
// Create unlocker instance using direct constructors with filesystem
|
||||
var unlocker secret.Unlocker
|
||||
// Use metadata directly as it's already the correct type
|
||||
switch metadata.Type {
|
||||
case "passphrase":
|
||||
secret.Debug("Creating passphrase unlocker instance", "unlocker_type", metadata.Type)
|
||||
case unlockerTypePassphrase:
|
||||
secret.Debug("Creating passphrase unlocker instance",
|
||||
"unlocker_type", metadata.Type)
|
||||
|
||||
unlocker = secret.NewPassphraseUnlocker(v.fs, unlockerDir, metadata)
|
||||
case "pgp":
|
||||
secret.Debug("Creating PGP unlocker instance", "unlocker_type", metadata.Type)
|
||||
|
||||
unlocker = secret.NewPGPUnlocker(v.fs, unlockerDir, metadata)
|
||||
case "keychain":
|
||||
secret.Debug("Creating keychain unlocker instance", "unlocker_type", metadata.Type)
|
||||
|
||||
unlocker = secret.NewKeychainUnlocker(v.fs, unlockerDir, metadata)
|
||||
case "secure-enclave":
|
||||
secret.Debug("Creating secure enclave unlocker instance", "unlocker_type", metadata.Type)
|
||||
case unlockerTypeSecureEnclave:
|
||||
secret.Debug("Creating secure enclave unlocker instance",
|
||||
"unlocker_type", metadata.Type)
|
||||
|
||||
unlocker = secret.NewSecureEnclaveUnlocker(v.fs, unlockerDir, metadata)
|
||||
default:
|
||||
secret.Debug("Unsupported unlocker type", "type", metadata.Type)
|
||||
|
||||
return nil, fmt.Errorf("unsupported unlocker type: %s", metadata.Type)
|
||||
return nil, fmt.Errorf("%w: %s", ErrUnsupportedUnlockerType, metadata.Type)
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully created unlocker instance",
|
||||
@@ -101,14 +99,16 @@ func (v *Vault) GetCurrentUnlocker() (secret.Unlocker, error) {
|
||||
return unlocker, nil
|
||||
}
|
||||
|
||||
// resolveUnlockerDirectory reads the current-unlocker file to get the unlocker directory path
|
||||
// resolveUnlockerDirectory reads the current-unlocker file to get the
|
||||
// unlocker directory path
|
||||
// The file contains just the unlocker name (e.g., "passphrase")
|
||||
func (v *Vault) resolveUnlockerDirectory(currentUnlockerPath string) (string, error) {
|
||||
secret.Debug("Reading current-unlocker file", "path", currentUnlockerPath)
|
||||
|
||||
unlockerNameBytes, err := afero.ReadFile(v.fs, currentUnlockerPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read current-unlocker file", "error", err, "path", currentUnlockerPath)
|
||||
secret.Debug("Failed to read current-unlocker file",
|
||||
"error", err, "path", currentUnlockerPath)
|
||||
|
||||
return "", fmt.Errorf("failed to read current unlocker: %w", err)
|
||||
}
|
||||
@@ -125,8 +125,13 @@ func (v *Vault) resolveUnlockerDirectory(currentUnlockerPath string) (string, er
|
||||
return absolutePath, nil
|
||||
}
|
||||
|
||||
// findUnlockerByID finds an unlocker by its ID and returns the unlocker instance and its directory path
|
||||
func (v *Vault) findUnlockerByID(unlockersDir, unlockerID string) (secret.Unlocker, string, error) {
|
||||
// findUnlockerByID finds an unlocker by its ID and returns the unlocker
|
||||
// instance and its directory path
|
||||
//
|
||||
//nolint:ireturn // returns one of several concrete unlocker implementations
|
||||
func (v *Vault) findUnlockerByID(
|
||||
unlockersDir, unlockerID string,
|
||||
) (secret.Unlocker, string, error) {
|
||||
files, err := afero.ReadDir(v.fs, unlockersDir)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("failed to read unlockers directory: %w", err)
|
||||
@@ -139,10 +144,14 @@ func (v *Vault) findUnlockerByID(unlockersDir, unlockerID string) (secret.Unlock
|
||||
|
||||
// Read metadata file
|
||||
metadataPath := filepath.Join(unlockersDir, file.Name(), "unlocker-metadata.json")
|
||||
|
||||
exists, err := afero.Exists(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("failed to check if metadata exists for unlocker %s: %w", file.Name(), err)
|
||||
return nil, "", fmt.Errorf(
|
||||
"failed to check if metadata exists for unlocker %s: %w",
|
||||
file.Name(), err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
// Skip directories without metadata - they might not be unlockers
|
||||
continue
|
||||
@@ -150,26 +159,31 @@ func (v *Vault) findUnlockerByID(unlockersDir, unlockerID string) (secret.Unlock
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("failed to read metadata for unlocker %s: %w", file.Name(), err)
|
||||
return nil, "", fmt.Errorf(
|
||||
"failed to read metadata for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
return nil, "", fmt.Errorf("failed to parse metadata for unlocker %s: %w", file.Name(), err)
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf(
|
||||
"failed to parse metadata for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
|
||||
unlockerDirPath := filepath.Join(unlockersDir, file.Name())
|
||||
|
||||
// Create the appropriate unlocker instance
|
||||
var tempUnlocker secret.Unlocker
|
||||
|
||||
switch metadata.Type {
|
||||
case "passphrase":
|
||||
case unlockerTypePassphrase:
|
||||
tempUnlocker = secret.NewPassphraseUnlocker(v.fs, unlockerDirPath, metadata)
|
||||
case "pgp":
|
||||
tempUnlocker = secret.NewPGPUnlocker(v.fs, unlockerDirPath, metadata)
|
||||
case "keychain":
|
||||
tempUnlocker = secret.NewKeychainUnlocker(v.fs, unlockerDirPath, metadata)
|
||||
case "secure-enclave":
|
||||
case unlockerTypeSecureEnclave:
|
||||
tempUnlocker = secret.NewSecureEnclaveUnlocker(v.fs, unlockerDirPath, metadata)
|
||||
default:
|
||||
continue
|
||||
@@ -198,6 +212,7 @@ func (v *Vault) ListUnlockers() ([]UnlockerMetadata, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if unlockers directory exists: %w", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
return []UnlockerMetadata{}, nil
|
||||
}
|
||||
@@ -209,28 +224,39 @@ func (v *Vault) ListUnlockers() ([]UnlockerMetadata, error) {
|
||||
}
|
||||
|
||||
var unlockers []UnlockerMetadata
|
||||
|
||||
for _, file := range files {
|
||||
if file.IsDir() {
|
||||
// Read metadata file
|
||||
metadataPath := filepath.Join(unlockersDir, file.Name(), "unlocker-metadata.json")
|
||||
metadataPath := filepath.Join(unlockersDir, file.Name(),
|
||||
"unlocker-metadata.json")
|
||||
|
||||
exists, err := afero.Exists(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to check if metadata exists for unlocker %s: %w", file.Name(), err)
|
||||
return nil, fmt.Errorf(
|
||||
"failed to check if metadata exists for unlocker %s: %w",
|
||||
file.Name(), err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
secret.Warn("Skipping unlocker directory with missing metadata file", "directory", file.Name())
|
||||
secret.Warn("Skipping unlocker directory with missing metadata file",
|
||||
"directory", file.Name())
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read metadata for unlocker %s: %w", file.Name(), err)
|
||||
return nil, fmt.Errorf(
|
||||
"failed to read metadata for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
if err := json.Unmarshal(metadataBytes, &metadata); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse metadata for unlocker %s: %w", file.Name(), err)
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(
|
||||
"failed to parse metadata for unlocker %s: %w", file.Name(), err)
|
||||
}
|
||||
|
||||
unlockers = append(unlockers, metadata)
|
||||
@@ -257,7 +283,7 @@ func (v *Vault) RemoveUnlocker(unlockerID string) error {
|
||||
}
|
||||
|
||||
if unlocker == nil {
|
||||
return fmt.Errorf("unlocker with ID %s not found", unlockerID)
|
||||
return fmt.Errorf("unlocker with ID %s %w", unlockerID, ErrUnlockerNotFound)
|
||||
}
|
||||
|
||||
// Use the unlocker's Remove method
|
||||
@@ -281,17 +307,21 @@ func (v *Vault) SelectUnlocker(unlockerID string) error {
|
||||
}
|
||||
|
||||
if targetUnlockerDir == "" {
|
||||
return fmt.Errorf("unlocker with ID %s not found", unlockerID)
|
||||
return fmt.Errorf("unlocker with ID %s %w", unlockerID, ErrUnlockerNotFound)
|
||||
}
|
||||
|
||||
// Create/update current-unlocker file with just the unlocker name
|
||||
currentUnlockerPath := filepath.Join(vaultDir, "current-unlocker")
|
||||
|
||||
// Remove existing file if it exists
|
||||
if exists, err := afero.Exists(v.fs, currentUnlockerPath); err != nil {
|
||||
exists, err := afero.Exists(v.fs, currentUnlockerPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check if current-unlocker file exists: %w", err)
|
||||
} else if exists {
|
||||
if err := v.fs.Remove(currentUnlockerPath); err != nil {
|
||||
}
|
||||
|
||||
if exists {
|
||||
err = v.fs.Remove(currentUnlockerPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to remove existing current-unlocker file: %w", err)
|
||||
}
|
||||
}
|
||||
@@ -301,7 +331,10 @@ func (v *Vault) SelectUnlocker(unlockerID string) error {
|
||||
|
||||
// Write just the unlocker name to the file
|
||||
secret.Debug("Writing current-unlocker file", "unlocker_name", unlockerName)
|
||||
if err := afero.WriteFile(v.fs, currentUnlockerPath, []byte(unlockerName), secret.FilePerms); err != nil {
|
||||
|
||||
err = afero.WriteFile(v.fs, currentUnlockerPath, []byte(unlockerName),
|
||||
secret.FilePerms)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create current-unlocker file: %w", err)
|
||||
}
|
||||
|
||||
@@ -310,15 +343,19 @@ func (v *Vault) SelectUnlocker(unlockerID string) error {
|
||||
|
||||
// CreatePassphraseUnlocker creates a new passphrase-protected unlocker
|
||||
// The passphrase must be provided as a LockedBuffer for security
|
||||
func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*secret.PassphraseUnlocker, error) {
|
||||
func (v *Vault) CreatePassphraseUnlocker(
|
||||
passphrase *memguard.LockedBuffer,
|
||||
) (*secret.PassphraseUnlocker, error) {
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
// Create unlocker directory
|
||||
unlockerDir := filepath.Join(vaultDir, "unlockers.d", "passphrase")
|
||||
if err := v.fs.MkdirAll(unlockerDir, secret.DirPerms); err != nil {
|
||||
unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerTypePassphrase)
|
||||
|
||||
err = v.fs.MkdirAll(unlockerDir, secret.DirPerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create unlocker directory: %w", err)
|
||||
}
|
||||
|
||||
@@ -328,32 +365,15 @@ func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*se
|
||||
return nil, fmt.Errorf("failed to generate unlocker: %w", err)
|
||||
}
|
||||
|
||||
// Write public key
|
||||
pubKeyPath := filepath.Join(unlockerDir, "pub.age")
|
||||
if err := afero.WriteFile(v.fs, pubKeyPath,
|
||||
[]byte(unlockerIdentity.Recipient().String()),
|
||||
secret.FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write unlocker public key: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt private key with passphrase
|
||||
privKeyStr := unlockerIdentity.String()
|
||||
privKeyBuffer := memguard.NewBufferFromBytes([]byte(privKeyStr))
|
||||
defer privKeyBuffer.Destroy()
|
||||
encryptedPrivKey, err := secret.EncryptWithPassphrase(privKeyBuffer, passphrase)
|
||||
// Write the unlocker keypair (public and passphrase-encrypted private)
|
||||
err = v.writeUnlockerKeypair(unlockerDir, unlockerIdentity, passphrase)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt unlocker private key: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted private key
|
||||
privKeyPath := filepath.Join(unlockerDir, "priv.age")
|
||||
if err := afero.WriteFile(v.fs, privKeyPath, encryptedPrivKey, secret.FilePerms); err != nil {
|
||||
return nil, fmt.Errorf("failed to write encrypted unlocker private key: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Create metadata
|
||||
metadata := UnlockerMetadata{
|
||||
Type: "passphrase",
|
||||
Type: unlockerTypePassphrase,
|
||||
CreatedAt: time.Now(),
|
||||
Flags: []string{},
|
||||
}
|
||||
@@ -365,7 +385,9 @@ func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*se
|
||||
}
|
||||
|
||||
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json")
|
||||
if err := afero.WriteFile(v.fs, metadataPath, metadataBytes, secret.FilePerms); err != nil {
|
||||
|
||||
err = afero.WriteFile(v.fs, metadataPath, metadataBytes, secret.FilePerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to write unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
@@ -379,13 +401,16 @@ func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*se
|
||||
ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String()))
|
||||
defer ltPrivKeyBuffer.Destroy()
|
||||
|
||||
encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, unlockerIdentity.Recipient())
|
||||
encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer,
|
||||
unlockerIdentity.Recipient())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt long-term private key: %w", err)
|
||||
}
|
||||
|
||||
ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age")
|
||||
if err := afero.WriteFile(v.fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms); err != nil {
|
||||
|
||||
err = afero.WriteFile(v.fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
@@ -393,9 +418,80 @@ func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*se
|
||||
unlocker := secret.NewPassphraseUnlocker(v.fs, unlockerDir, metadata)
|
||||
|
||||
// Select this unlocker as current
|
||||
if err := v.SelectUnlocker(unlocker.GetID()); err != nil {
|
||||
err = v.SelectUnlocker(unlocker.GetID())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to select new unlocker: %w", err)
|
||||
}
|
||||
|
||||
return unlocker, nil
|
||||
}
|
||||
|
||||
// readUnlockerMetadata reads and parses the unlocker-metadata.json file in
|
||||
// the given unlocker directory.
|
||||
func (v *Vault) readUnlockerMetadata(unlockerDir string) (UnlockerMetadata, error) {
|
||||
metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json")
|
||||
secret.Debug("Reading unlocker metadata", "path", metadataPath)
|
||||
|
||||
var metadata UnlockerMetadata
|
||||
|
||||
metadataBytes, err := afero.ReadFile(v.fs, metadataPath)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to read unlocker metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return metadata, fmt.Errorf("failed to read unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
err = json.Unmarshal(metadataBytes, &metadata)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to parse unlocker metadata", "error", err, "path", metadataPath)
|
||||
|
||||
return metadata, fmt.Errorf("failed to parse unlocker metadata: %w", err)
|
||||
}
|
||||
|
||||
secret.DebugWith("Parsed unlocker metadata",
|
||||
slog.String("unlocker_type", metadata.Type),
|
||||
slog.Time("created_at", metadata.CreatedAt),
|
||||
slog.Any("flags", metadata.Flags),
|
||||
)
|
||||
|
||||
return metadata, nil
|
||||
}
|
||||
|
||||
// writeUnlockerKeypair writes the unlocker's public key and its
|
||||
// passphrase-encrypted private key into the unlocker directory.
|
||||
func (v *Vault) writeUnlockerKeypair(
|
||||
unlockerDir string,
|
||||
unlockerIdentity *age.X25519Identity,
|
||||
passphrase *memguard.LockedBuffer,
|
||||
) error {
|
||||
// Write public key
|
||||
pubKeyPath := filepath.Join(unlockerDir, "pub.age")
|
||||
|
||||
err := afero.WriteFile(v.fs, pubKeyPath,
|
||||
[]byte(unlockerIdentity.Recipient().String()),
|
||||
secret.FilePerms)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write unlocker public key: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt private key with passphrase
|
||||
privKeyStr := unlockerIdentity.String()
|
||||
|
||||
privKeyBuffer := memguard.NewBufferFromBytes([]byte(privKeyStr))
|
||||
defer privKeyBuffer.Destroy()
|
||||
|
||||
encryptedPrivKey, err := secret.EncryptWithPassphrase(privKeyBuffer, passphrase)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to encrypt unlocker private key: %w", err)
|
||||
}
|
||||
|
||||
// Write encrypted private key
|
||||
privKeyPath := filepath.Join(unlockerDir, "priv.age")
|
||||
|
||||
err = afero.WriteFile(v.fs, privKeyPath, encryptedPrivKey, secret.FilePerms)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write encrypted unlocker private key: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -23,12 +23,14 @@ type Vault struct {
|
||||
// NewVault creates a new Vault instance
|
||||
func NewVault(fs afero.Fs, stateDir string, name string) *Vault {
|
||||
secret.Debug("Creating NewVault instance")
|
||||
|
||||
v := &Vault{
|
||||
Name: name,
|
||||
fs: fs,
|
||||
stateDir: stateDir,
|
||||
longTermKey: nil,
|
||||
}
|
||||
|
||||
secret.Debug("Created NewVault instance successfully")
|
||||
|
||||
return v
|
||||
@@ -54,7 +56,8 @@ func (v *Vault) ClearLongTermKey() {
|
||||
v.longTermKey = nil
|
||||
}
|
||||
|
||||
// GetOrDeriveLongTermKey gets the long-term key from memory or derives it from available sources
|
||||
// GetOrDeriveLongTermKey gets the long-term key from memory or derives it
|
||||
// from available sources
|
||||
func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
|
||||
// If we have it in memory, return it
|
||||
if !v.Locked() {
|
||||
@@ -65,55 +68,12 @@ func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
|
||||
|
||||
// Try to derive from environment mnemonic first
|
||||
if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" {
|
||||
secret.Debug("Using mnemonic from environment for long-term key derivation", "vault_name", v.Name)
|
||||
|
||||
// Load vault metadata to get the derivation index
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
metadata, err := LoadVaultMetadata(v.fs, vaultDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to load vault metadata", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to load vault metadata: %w", err)
|
||||
}
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to derive long-term key from mnemonic", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
|
||||
// Verify that the derived key matches the stored public key hash
|
||||
derivedPubKeyHash := ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
if derivedPubKeyHash != metadata.PublicKeyHash {
|
||||
secret.Debug("Derived public key hash does not match stored hash",
|
||||
"vault_name", v.Name,
|
||||
"derived_hash", derivedPubKeyHash,
|
||||
"stored_hash", metadata.PublicKeyHash,
|
||||
"derivation_index", metadata.DerivationIndex)
|
||||
|
||||
return nil, fmt.Errorf("derived public key does not match vault: mnemonic may be incorrect")
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully derived long-term key from mnemonic",
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.String("public_key", ltIdentity.Recipient().String()),
|
||||
slog.Uint64("derivation_index", uint64(metadata.DerivationIndex)),
|
||||
)
|
||||
|
||||
// Cache the derived key by unlocking the vault
|
||||
v.Unlock(ltIdentity)
|
||||
secret.Debug("Vault is unlocked (lt key in memory) via mnemonic", "vault_name", v.Name)
|
||||
|
||||
return ltIdentity, nil
|
||||
return v.deriveLongTermKeyFromMnemonic(envMnemonic)
|
||||
}
|
||||
|
||||
// No mnemonic available, try to use current unlocker
|
||||
secret.Debug("No mnemonic available, using current unlocker to unlock vault", "vault_name", v.Name)
|
||||
secret.Debug("No mnemonic available, using current unlocker to unlock vault",
|
||||
"vault_name", v.Name)
|
||||
|
||||
// Get current unlocker
|
||||
unlocker, err := v.GetCurrentUnlocker()
|
||||
@@ -151,10 +111,130 @@ func (v *Vault) GetOrDeriveLongTermKey() (*age.X25519Identity, error) {
|
||||
return ltIdentity, nil
|
||||
}
|
||||
|
||||
// unlockLongTermKey extracts the vault's long-term key using the given unlocker.
|
||||
// SE unlockers decrypt the long-term key directly; other unlockers use an intermediate identity.
|
||||
func (v *Vault) unlockLongTermKey(unlocker secret.Unlocker) (*age.X25519Identity, error) {
|
||||
if unlocker.GetType() == "secure-enclave" {
|
||||
// GetDirectory returns the vault's directory path
|
||||
func (v *Vault) GetDirectory() (string, error) {
|
||||
return filepath.Join(v.stateDir, "vaults.d", v.Name), nil
|
||||
}
|
||||
|
||||
// GetName returns the vault's name (for VaultInterface compatibility)
|
||||
func (v *Vault) GetName() string {
|
||||
return v.Name
|
||||
}
|
||||
|
||||
// GetFilesystem returns the vault's filesystem (for VaultInterface
|
||||
// compatibility)
|
||||
//
|
||||
//nolint:ireturn // afero.Fs is the interface required by VaultInterface
|
||||
func (v *Vault) GetFilesystem() afero.Fs {
|
||||
return v.fs
|
||||
}
|
||||
|
||||
// NumSecrets returns the number of secrets in the vault
|
||||
func (v *Vault) NumSecrets() (int, error) {
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
secretsDir := filepath.Join(vaultDir, "secrets.d")
|
||||
|
||||
exists, _ := afero.DirExists(v.fs, secretsDir)
|
||||
if !exists {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
entries, err := afero.ReadDir(v.fs, secretsDir)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to read secrets directory: %w", err)
|
||||
}
|
||||
|
||||
// Count only directories that have a "current" version pointer file
|
||||
count := 0
|
||||
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
// A valid secret has a "current" file pointing to the active version
|
||||
secretDir := filepath.Join(secretsDir, entry.Name())
|
||||
currentFile := filepath.Join(secretDir, "current")
|
||||
|
||||
exists, err := afero.Exists(v.fs, currentFile)
|
||||
if err != nil {
|
||||
continue // Skip directories we can't read
|
||||
}
|
||||
|
||||
if exists {
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// deriveLongTermKeyFromMnemonic derives the long-term key from the given
|
||||
// mnemonic, verifies it against the vault metadata, and caches it in memory.
|
||||
func (v *Vault) deriveLongTermKeyFromMnemonic(
|
||||
envMnemonic string,
|
||||
) (*age.X25519Identity, error) {
|
||||
secret.Debug("Using mnemonic from environment for long-term key derivation",
|
||||
"vault_name", v.Name)
|
||||
|
||||
// Load vault metadata to get the derivation index
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
metadata, err := LoadVaultMetadata(v.fs, vaultDir)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to load vault metadata", "error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to load vault metadata: %w", err)
|
||||
}
|
||||
|
||||
ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex)
|
||||
if err != nil {
|
||||
secret.Debug("Failed to derive long-term key from mnemonic",
|
||||
"error", err, "vault_name", v.Name)
|
||||
|
||||
return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err)
|
||||
}
|
||||
|
||||
// Verify that the derived key matches the stored public key hash
|
||||
derivedPubKeyHash := ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String()))
|
||||
if derivedPubKeyHash != metadata.PublicKeyHash {
|
||||
secret.Debug("Derived public key hash does not match stored hash",
|
||||
"vault_name", v.Name,
|
||||
"derived_hash", derivedPubKeyHash,
|
||||
"stored_hash", metadata.PublicKeyHash,
|
||||
"derivation_index", metadata.DerivationIndex)
|
||||
|
||||
return nil, ErrMnemonicMismatch
|
||||
}
|
||||
|
||||
secret.DebugWith("Successfully derived long-term key from mnemonic",
|
||||
slog.String("vault_name", v.Name),
|
||||
slog.String("public_key", ltIdentity.Recipient().String()),
|
||||
slog.Uint64("derivation_index", uint64(metadata.DerivationIndex)),
|
||||
)
|
||||
|
||||
// Cache the derived key by unlocking the vault
|
||||
v.Unlock(ltIdentity)
|
||||
secret.Debug("Vault is unlocked (lt key in memory) via mnemonic",
|
||||
"vault_name", v.Name)
|
||||
|
||||
return ltIdentity, nil
|
||||
}
|
||||
|
||||
// unlockLongTermKey extracts the vault's long-term key using the given
|
||||
// unlocker. SE unlockers decrypt the long-term key directly; other unlockers
|
||||
// use an intermediate identity.
|
||||
func (v *Vault) unlockLongTermKey(
|
||||
unlocker secret.Unlocker,
|
||||
) (*age.X25519Identity, error) {
|
||||
if unlocker.GetType() == unlockerTypeSecureEnclave {
|
||||
secret.Debug("SE unlocker: decrypting long-term key directly via Secure Enclave")
|
||||
|
||||
ltIdentity, err := unlocker.GetIdentity()
|
||||
@@ -178,7 +258,8 @@ func (v *Vault) unlockLongTermKey(unlocker secret.Unlocker) (*age.X25519Identity
|
||||
return nil, fmt.Errorf("failed to read encrypted long-term private key: %w", err)
|
||||
}
|
||||
|
||||
ltPrivKeyBuffer, err := secret.DecryptWithIdentity(encryptedLtPrivKey, unlockerIdentity)
|
||||
ltPrivKeyBuffer, err := secret.DecryptWithIdentity(
|
||||
encryptedLtPrivKey, unlockerIdentity)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err)
|
||||
}
|
||||
@@ -191,59 +272,3 @@ func (v *Vault) unlockLongTermKey(unlocker secret.Unlocker) (*age.X25519Identity
|
||||
|
||||
return ltIdentity, nil
|
||||
}
|
||||
|
||||
// GetDirectory returns the vault's directory path
|
||||
func (v *Vault) GetDirectory() (string, error) {
|
||||
return filepath.Join(v.stateDir, "vaults.d", v.Name), nil
|
||||
}
|
||||
|
||||
// GetName returns the vault's name (for VaultInterface compatibility)
|
||||
func (v *Vault) GetName() string {
|
||||
return v.Name
|
||||
}
|
||||
|
||||
// GetFilesystem returns the vault's filesystem (for VaultInterface compatibility)
|
||||
func (v *Vault) GetFilesystem() afero.Fs {
|
||||
return v.fs
|
||||
}
|
||||
|
||||
// NumSecrets returns the number of secrets in the vault
|
||||
func (v *Vault) NumSecrets() (int, error) {
|
||||
vaultDir, err := v.GetDirectory()
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to get vault directory: %w", err)
|
||||
}
|
||||
|
||||
secretsDir := filepath.Join(vaultDir, "secrets.d")
|
||||
exists, _ := afero.DirExists(v.fs, secretsDir)
|
||||
if !exists {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
entries, err := afero.ReadDir(v.fs, secretsDir)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to read secrets directory: %w", err)
|
||||
}
|
||||
|
||||
// Count only directories that have a "current" version pointer file
|
||||
count := 0
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
|
||||
// A valid secret has a "current" file pointing to the active version
|
||||
secretDir := filepath.Join(secretsDir, entry.Name())
|
||||
currentFile := filepath.Join(secretDir, "current")
|
||||
exists, err := afero.Exists(v.fs, currentFile)
|
||||
if err != nil {
|
||||
continue // Skip directories we can't read
|
||||
}
|
||||
|
||||
if exists {
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
@@ -13,32 +13,34 @@ import (
|
||||
)
|
||||
|
||||
func TestAddSecretFailsWithMissingPublicKey(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Create in-memory filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create a vault directory without a public key (simulating the error condition)
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "broken")
|
||||
// Create a vault directory without a public key (simulating the error
|
||||
// condition)
|
||||
vaultDir := filepath.Join(testStateDir, "vaults.d", "broken")
|
||||
require.NoError(t, fs.MkdirAll(vaultDir, secret.DirPerms))
|
||||
|
||||
// Create currentvault symlink
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
require.NoError(t, afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), secret.FilePerms))
|
||||
currentVaultPath := filepath.Join(testStateDir, "currentvault")
|
||||
require.NoError(t,
|
||||
afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), secret.FilePerms))
|
||||
|
||||
// Create vault instance
|
||||
vlt := vault.NewVault(fs, stateDir, "broken")
|
||||
vlt := vault.NewVault(fs, testStateDir, "broken")
|
||||
|
||||
// Try to add a secret - this should fail
|
||||
secretName := "test-secret"
|
||||
value := memguard.NewBufferFromBytes([]byte("test-value"))
|
||||
defer value.Destroy()
|
||||
|
||||
err := vlt.AddSecret(secretName, value, false)
|
||||
err := vlt.AddSecret(testSecretName, value, false)
|
||||
require.Error(t, err, "AddSecret should fail when public key is missing")
|
||||
assert.Contains(t, err.Error(), "failed to read long-term public key")
|
||||
|
||||
// Verify that the secret directory was NOT created
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", secretName)
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", testSecretName)
|
||||
exists, _ := afero.DirExists(fs, secretDir)
|
||||
assert.False(t, exists, "Secret directory should not exist after failed AddSecret")
|
||||
|
||||
@@ -47,41 +49,45 @@ func TestAddSecretFailsWithMissingPublicKey(t *testing.T) {
|
||||
if exists, _ := afero.DirExists(fs, secretsDir); exists {
|
||||
entries, err := afero.ReadDir(fs, secretsDir)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, entries, "secrets.d directory should be empty after failed AddSecret")
|
||||
assert.Empty(t, entries,
|
||||
"secrets.d directory should be empty after failed AddSecret")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddSecretCleansUpOnFailure(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Create in-memory filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create a vault directory with public key
|
||||
vaultDir := filepath.Join(stateDir, "vaults.d", "test")
|
||||
vaultDir := filepath.Join(testStateDir, "vaults.d", "test")
|
||||
require.NoError(t, fs.MkdirAll(vaultDir, secret.DirPerms))
|
||||
|
||||
// Create a mock public key that will cause encryption to fail
|
||||
// by using an invalid age public key format
|
||||
pubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
require.NoError(t, afero.WriteFile(fs, pubKeyPath, []byte("invalid-public-key"), secret.FilePerms))
|
||||
require.NoError(t,
|
||||
afero.WriteFile(fs, pubKeyPath, []byte("invalid-public-key"),
|
||||
secret.FilePerms))
|
||||
|
||||
// Create currentvault symlink
|
||||
currentVaultPath := filepath.Join(stateDir, "currentvault")
|
||||
require.NoError(t, afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), secret.FilePerms))
|
||||
currentVaultPath := filepath.Join(testStateDir, "currentvault")
|
||||
require.NoError(t,
|
||||
afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), secret.FilePerms))
|
||||
|
||||
// Create vault instance
|
||||
vlt := vault.NewVault(fs, stateDir, "test")
|
||||
vlt := vault.NewVault(fs, testStateDir, "test")
|
||||
|
||||
// Try to add a secret - this should fail during encryption
|
||||
secretName := "test-secret"
|
||||
value := memguard.NewBufferFromBytes([]byte("test-value"))
|
||||
defer value.Destroy()
|
||||
|
||||
err := vlt.AddSecret(secretName, value, false)
|
||||
err := vlt.AddSecret(testSecretName, value, false)
|
||||
require.Error(t, err, "AddSecret should fail with invalid public key")
|
||||
|
||||
// Verify that the secret directory was NOT created
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", secretName)
|
||||
secretDir := filepath.Join(vaultDir, "secrets.d", testSecretName)
|
||||
exists, _ := afero.DirExists(fs, secretDir)
|
||||
assert.False(t, exists, "Secret directory should not exist after failed AddSecret")
|
||||
}
|
||||
|
||||
@@ -1,268 +1,301 @@
|
||||
package vault
|
||||
package vault_test
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"git.eeqj.de/sneak/secret/internal/secret"
|
||||
"git.eeqj.de/sneak/secret/internal/vault"
|
||||
"git.eeqj.de/sneak/secret/pkg/agehd"
|
||||
"github.com/awnumar/memguard"
|
||||
"github.com/spf13/afero"
|
||||
)
|
||||
|
||||
// testMnemonic is the shared BIP39 test mnemonic for tests in this package.
|
||||
//
|
||||
//nolint:dupword // BIP39 test mnemonic intentionally repeats a word
|
||||
const testMnemonic = "abandon abandon abandon abandon abandon abandon abandon " +
|
||||
"abandon abandon abandon abandon about"
|
||||
|
||||
// Shared fixtures for tests in this package.
|
||||
const (
|
||||
testStateDir = "/test/state"
|
||||
testVaultName = "test-vault"
|
||||
testSecretName = "test-secret"
|
||||
testPassphrase = "test-passphrase"
|
||||
)
|
||||
|
||||
//nolint:paralleltest // t.Setenv and order-dependent subtests forbid parallel
|
||||
func TestVaultOperations(t *testing.T) {
|
||||
// Test environment will be cleaned up automatically by t.Setenv
|
||||
|
||||
// Set test environment variables
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
// Use in-memory filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Test vault creation
|
||||
t.Run("CreateVault", func(t *testing.T) {
|
||||
vlt, err := CreateVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
if vlt.GetName() != "test-vault" {
|
||||
t.Errorf("Expected vault name 'test-vault', got '%s'", vlt.GetName())
|
||||
}
|
||||
|
||||
// Check vault directory exists
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
exists, err := afero.DirExists(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check vault directory: %v", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
t.Errorf("Vault directory should exist")
|
||||
}
|
||||
testCreateVault(t, fs)
|
||||
})
|
||||
|
||||
// Test vault listing
|
||||
t.Run("ListVaults", func(t *testing.T) {
|
||||
vaults, err := ListVaults(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list vaults: %v", err)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, vault := range vaults {
|
||||
if vault == "test-vault" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Errorf("Expected to find 'test-vault' in vault list")
|
||||
}
|
||||
testListVaults(t, fs)
|
||||
})
|
||||
|
||||
// Test vault selection
|
||||
t.Run("SelectVault", func(t *testing.T) {
|
||||
err := SelectVault(fs, stateDir, "test-vault")
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select vault: %v", err)
|
||||
}
|
||||
|
||||
// Test getting current vault
|
||||
currentVault, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
if currentVault.GetName() != "test-vault" {
|
||||
t.Errorf("Expected current vault 'test-vault', got '%s'", currentVault.GetName())
|
||||
}
|
||||
testSelectVault(t, fs)
|
||||
})
|
||||
|
||||
// Test secret operations
|
||||
t.Run("SecretOperations", func(t *testing.T) {
|
||||
vlt, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
// First, derive the long-term key from the test mnemonic
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Get the public key from the derived identity
|
||||
ltPublicKey := ltIdentity.Recipient().String()
|
||||
|
||||
// Get the vault directory
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Write the correct public key to the pub.age file
|
||||
pubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
err = afero.WriteFile(fs, pubKeyPath, []byte(ltPublicKey), secret.FilePerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to write long-term public key: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vault with the derived identity
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Now add a secret
|
||||
secretName := "test/secret"
|
||||
secretValue := []byte("test-secret-value")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vlt.AddSecret(secretName, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret: %v", err)
|
||||
}
|
||||
|
||||
// List secrets
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets: %v", err)
|
||||
}
|
||||
|
||||
found := false
|
||||
for _, secret := range secrets {
|
||||
if secret == secretName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Errorf("Expected to find secret '%s' in list", secretName)
|
||||
}
|
||||
|
||||
// Get secret value
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get secret: %v", err)
|
||||
}
|
||||
|
||||
if string(retrievedValue) != string(expectedValue) {
|
||||
t.Errorf("Expected secret value '%s', got '%s'", string(expectedValue), string(retrievedValue))
|
||||
}
|
||||
testSecretOperations(t, fs)
|
||||
})
|
||||
|
||||
// Test NumSecrets
|
||||
t.Run("NumSecrets", func(t *testing.T) {
|
||||
vlt, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
numSecrets, err := vlt.NumSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to count secrets: %v", err)
|
||||
}
|
||||
|
||||
// We added one secret in SecretOperations
|
||||
if numSecrets != 1 {
|
||||
t.Errorf("Expected 1 secret, got %d", numSecrets)
|
||||
}
|
||||
testNumSecrets(t, fs)
|
||||
})
|
||||
|
||||
// Test unlocker operations
|
||||
t.Run("UnlockerOperations", func(t *testing.T) {
|
||||
vlt, err := GetCurrentVault(fs, stateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
// Test vault unlocking (should happen automatically via mnemonic)
|
||||
if vlt.Locked() {
|
||||
_, err := vlt.UnlockVault()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unlock vault: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Create a passphrase unlocker
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte("test-passphrase"))
|
||||
defer passphraseBuffer.Destroy()
|
||||
passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create passphrase unlocker: %v", err)
|
||||
}
|
||||
|
||||
// List unlockers
|
||||
unlockers, err := vlt.ListUnlockers()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list unlockers: %v", err)
|
||||
}
|
||||
|
||||
if len(unlockers) == 0 {
|
||||
t.Errorf("Expected at least one unlocker")
|
||||
}
|
||||
|
||||
// Check key type
|
||||
keyFound := false
|
||||
for _, key := range unlockers {
|
||||
if key.Type == "passphrase" {
|
||||
keyFound = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !keyFound {
|
||||
t.Errorf("Expected to find passphrase unlocker")
|
||||
}
|
||||
|
||||
// Test selecting unlocker
|
||||
err = vlt.SelectUnlocker(passphraseUnlocker.GetID())
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select unlocker: %v", err)
|
||||
}
|
||||
|
||||
// Test getting current unlocker
|
||||
currentUnlocker, err := vlt.GetCurrentUnlocker()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current unlocker: %v", err)
|
||||
}
|
||||
|
||||
if currentUnlocker.GetID() != passphraseUnlocker.GetID() {
|
||||
t.Errorf("Expected current unlocker ID '%s', got '%s'", passphraseUnlocker.GetID(), currentUnlocker.GetID())
|
||||
}
|
||||
testUnlockerOperations(t, fs)
|
||||
})
|
||||
}
|
||||
|
||||
func testCreateVault(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
if vlt.GetName() != testVaultName {
|
||||
t.Errorf("Expected vault name '%s', got '%s'", testVaultName, vlt.GetName())
|
||||
}
|
||||
|
||||
// Check vault directory exists
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
exists, err := afero.DirExists(fs, vaultDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to check vault directory: %v", err)
|
||||
}
|
||||
|
||||
if !exists {
|
||||
t.Errorf("Vault directory should exist")
|
||||
}
|
||||
}
|
||||
|
||||
func testListVaults(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vaults, err := vault.ListVaults(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list vaults: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(vaults, testVaultName) {
|
||||
t.Errorf("Expected to find '%s' in vault list", testVaultName)
|
||||
}
|
||||
}
|
||||
|
||||
func testSelectVault(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
err := vault.SelectVault(fs, testStateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select vault: %v", err)
|
||||
}
|
||||
|
||||
// Test getting current vault
|
||||
currentVault, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
if currentVault.GetName() != testVaultName {
|
||||
t.Errorf("Expected current vault '%s', got '%s'",
|
||||
testVaultName, currentVault.GetName())
|
||||
}
|
||||
}
|
||||
|
||||
func testSecretOperations(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
// First, derive the long-term key from the test mnemonic
|
||||
ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to derive long-term key: %v", err)
|
||||
}
|
||||
|
||||
// Get the public key from the derived identity
|
||||
ltPublicKey := ltIdentity.Recipient().String()
|
||||
|
||||
// Get the vault directory
|
||||
vaultDir, err := vlt.GetDirectory()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
// Write the correct public key to the pub.age file
|
||||
pubKeyPath := filepath.Join(vaultDir, "pub.age")
|
||||
|
||||
err = afero.WriteFile(fs, pubKeyPath, []byte(ltPublicKey), secret.FilePerms)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to write long-term public key: %v", err)
|
||||
}
|
||||
|
||||
// Unlock the vault with the derived identity
|
||||
vlt.Unlock(ltIdentity)
|
||||
|
||||
// Now add a secret
|
||||
secretName := "test/secret"
|
||||
secretValue := []byte("test-secret-value")
|
||||
expectedValue := make([]byte, len(secretValue))
|
||||
copy(expectedValue, secretValue)
|
||||
|
||||
secretBuffer := memguard.NewBufferFromBytes(secretValue)
|
||||
defer secretBuffer.Destroy()
|
||||
|
||||
err = vlt.AddSecret(secretName, secretBuffer, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to add secret: %v", err)
|
||||
}
|
||||
|
||||
// List secrets
|
||||
secrets, err := vlt.ListSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list secrets: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(secrets, secretName) {
|
||||
t.Errorf("Expected to find secret '%s' in list", secretName)
|
||||
}
|
||||
|
||||
// Get secret value
|
||||
retrievedValue, err := vlt.GetSecret(secretName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get secret: %v", err)
|
||||
}
|
||||
|
||||
if string(retrievedValue) != string(expectedValue) {
|
||||
t.Errorf("Expected secret value '%s', got '%s'",
|
||||
string(expectedValue), string(retrievedValue))
|
||||
}
|
||||
}
|
||||
|
||||
func testNumSecrets(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
numSecrets, err := vlt.NumSecrets()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to count secrets: %v", err)
|
||||
}
|
||||
|
||||
// We added one secret in SecretOperations
|
||||
if numSecrets != 1 {
|
||||
t.Errorf("Expected 1 secret, got %d", numSecrets)
|
||||
}
|
||||
}
|
||||
|
||||
func testUnlockerOperations(t *testing.T, fs afero.Fs) {
|
||||
t.Helper()
|
||||
|
||||
vlt, err := vault.GetCurrentVault(fs, testStateDir)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current vault: %v", err)
|
||||
}
|
||||
|
||||
// Test vault unlocking (should happen automatically via mnemonic)
|
||||
if vlt.Locked() {
|
||||
_, err := vlt.UnlockVault()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unlock vault: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Create a passphrase unlocker
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte(testPassphrase))
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create passphrase unlocker: %v", err)
|
||||
}
|
||||
|
||||
// List unlockers
|
||||
unlockers, err := vlt.ListUnlockers()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to list unlockers: %v", err)
|
||||
}
|
||||
|
||||
if len(unlockers) == 0 {
|
||||
t.Errorf("Expected at least one unlocker")
|
||||
}
|
||||
|
||||
// Check key type
|
||||
keyFound := false
|
||||
|
||||
for _, key := range unlockers {
|
||||
if key.Type == "passphrase" {
|
||||
keyFound = true
|
||||
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !keyFound {
|
||||
t.Errorf("Expected to find passphrase unlocker")
|
||||
}
|
||||
|
||||
// Test selecting unlocker
|
||||
err = vlt.SelectUnlocker(passphraseUnlocker.GetID())
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to select unlocker: %v", err)
|
||||
}
|
||||
|
||||
// Test getting current unlocker
|
||||
currentUnlocker, err := vlt.GetCurrentUnlocker()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get current unlocker: %v", err)
|
||||
}
|
||||
|
||||
if currentUnlocker.GetID() != passphraseUnlocker.GetID() {
|
||||
t.Errorf("Expected current unlocker ID '%s', got '%s'",
|
||||
passphraseUnlocker.GetID(), currentUnlocker.GetID())
|
||||
}
|
||||
}
|
||||
|
||||
func TestListUnlockers_SkipsMissingMetadata(t *testing.T) {
|
||||
// Set test environment variables
|
||||
testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about"
|
||||
t.Setenv(secret.EnvMnemonic, testMnemonic)
|
||||
t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase")
|
||||
t.Setenv(secret.EnvUnlockPassphrase, testPassphrase)
|
||||
|
||||
// Use in-memory filesystem
|
||||
fs := afero.NewMemMapFs()
|
||||
stateDir := "/test/state"
|
||||
|
||||
// Create vault
|
||||
vlt, err := CreateVault(fs, stateDir, "test-vault")
|
||||
vlt, err := vault.CreateVault(fs, testStateDir, testVaultName)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create vault: %v", err)
|
||||
}
|
||||
|
||||
// Create a passphrase unlocker so we have at least one valid unlocker
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte("test-passphrase"))
|
||||
passphraseBuffer := memguard.NewBufferFromBytes([]byte(testPassphrase))
|
||||
defer passphraseBuffer.Destroy()
|
||||
|
||||
_, err = vlt.CreatePassphraseUnlocker(passphraseBuffer)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create passphrase unlocker: %v", err)
|
||||
@@ -273,7 +306,9 @@ func TestListUnlockers_SkipsMissingMetadata(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get vault directory: %v", err)
|
||||
}
|
||||
|
||||
bogusDir := filepath.Join(vaultDir, "unlockers.d", "bogus-no-metadata")
|
||||
|
||||
err = fs.MkdirAll(bogusDir, 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create bogus directory: %v", err)
|
||||
@@ -282,7 +317,8 @@ func TestListUnlockers_SkipsMissingMetadata(t *testing.T) {
|
||||
// ListUnlockers should succeed, skipping the bogus directory
|
||||
unlockers, err := vlt.ListUnlockers()
|
||||
if err != nil {
|
||||
t.Fatalf("ListUnlockers returned error when it should have skipped bad directory: %v", err)
|
||||
t.Fatalf("ListUnlockers returned error when it should have skipped "+
|
||||
"bad directory: %v", err)
|
||||
}
|
||||
|
||||
// Should still have the valid passphrase unlocker
|
||||
|
||||
Reference in New Issue
Block a user