Files
secret/internal/vault/integration_test.go
T
sneak 10a709e897
check / check (push) Failing after 2s
Create a vault whole in a temporary directory, then select it (closes #105)
vault.CreateVault takes the unlocker passphrase and writes the vault
directory, its metadata, long-term public key and passphrase unlocker
into a temporary directory, renames that into vaults.d once complete,
and only then makes the vault current. secret init and secret vault
create call it once instead of adding the unlocker afterwards, so a
kill part-way leaves either no vault, whose temporary directory the
next command that takes the lock deletes, or a complete one. A test
records the state directory before every change the call makes and
checks each state, and the command run again from it.

Model: opus-5-5
2026-10-04 18:25:43 +00:00

465 lines
12 KiB
Go

package vault_test
import (
"bytes"
"os"
"path/filepath"
"slices"
"testing"
"filippo.io/age"
"git.eeqj.de/sneak/secret/internal/vault"
"git.eeqj.de/sneak/secret/pkg/agehd"
"github.com/awnumar/memguard"
"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
}
func TestVaultWithRealFilesystem(t *testing.T) {
t.Parallel()
// Create a temporary directory for our tests
tempDir := t.TempDir()
// Use the real filesystem
fs := afero.NewOsFs()
// Test currentvault file handling (plain file with relative path)
t.Run("CurrentVaultFileHandling", func(t *testing.T) {
t.Parallel()
testCurrentVaultFileHandling(t, fs, tempDir)
})
// Test secret operations with deeply nested paths
t.Run("DeepPathSecrets", func(t *testing.T) {
t.Parallel()
testDeepPathSecrets(t, fs, tempDir)
})
// Test key caching in GetOrDeriveLongTermKey
t.Run("KeyCaching", func(t *testing.T) {
t.Parallel()
testKeyCaching(t, fs, tempDir)
})
// Test vault name validation
t.Run("VaultNameValidation", func(t *testing.T) {
t.Parallel()
testVaultNameValidation(t, fs, tempDir)
})
// Test multiple vaults and switching between them
t.Run("MultipleVaults", func(t *testing.T) {
t.Parallel()
testMultipleVaults(t, fs, tempDir)
})
// Test adding a secret in one vault and verifying it's not visible in
// another
t.Run("VaultIsolation", func(t *testing.T) {
t.Parallel()
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,
testMnemonicBuffer(t), nil)
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 writes the public key derived from
// the mnemonic
vlt, err := vault.CreateVault(fs, stateDir, testVaultName,
testMnemonicBuffer(t), nil)
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)
}
defer retrievedValue.Destroy()
if !bytes.Equal(retrievedValue.Bytes(), expectedValue) {
t.Errorf("Retrieved value doesn't match. Expected %q, got %q",
expectedValue, retrievedValue.Bytes())
}
}
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 writes the public key derived from
// the mnemonic
vlt, err := vault.CreateVault(fs, stateDir, testVaultName,
testMnemonicBuffer(t), nil)
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, testMnemonicBuffer(t), nil)
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, testMnemonicBuffer(t), nil)
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, testMnemonicBuffer(t), nil)
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 writes the public key derived from
// the mnemonic
vault1, err := vault.CreateVault(fs, stateDir, "vault1", testMnemonicBuffer(t), nil)
if err != nil {
t.Fatalf("Failed to create vault1: %v", err)
}
vault2, err := vault.CreateVault(fs, stateDir, "vault2", testMnemonicBuffer(t), nil)
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")
}
}