diff --git a/.golangci.yml b/.golangci.yml index 265013a..26b1610 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -1,128 +1,34 @@ version: "2" +# Config schema uses the golangci-lint v2 layout (settings live under +# linters.settings, not top-level linters-settings) so that the +# thresholds below are actually applied by golangci-lint >= v2. + run: - go: "1.24" - tests: false + timeout: 5m + modules-download-mode: readonly linters: - enable: - # Additional linters requested - - testifylint # Checks usage of github.com/stretchr/testify - - usetesting # usetesting is an analyzer that detects using os.Setenv instead of t.Setenv since Go 1.17 - - tagliatelle # Checks the struct tags - - nlreturn # nlreturn checks for a new line before return and branch statements - - nilnil # Checks that there is no simultaneous return of nil error and an invalid value - - nestif # Reports deeply nested if statements - - mnd # An analyzer to detect magic numbers - - lll # Reports long lines - - intrange # intrange is a linter to find places where for loops could make use of an integer range - - gochecknoglobals # Check that no global variables exist - - # Default/existing linters that are commonly useful - - govet - - errcheck - - staticcheck - - unused - - ineffassign - - misspell - - revive - - gosec - - unconvert - - unparam - -linters-settings: - lll: - line-length: 120 - - mnd: - # List of enabled checks, see https://github.com/tommy-muehle/go-mnd/#checks for description. - checks: - - argument - - case - - condition - - operation - - return - - assign - ignored-numbers: - - '0' - - '1' - - '2' - - '8' - - '16' - - '40' # GPG fingerprint length - - '64' - - '128' - - '256' - - '512' - - '1024' - - '2048' - - '4096' - - nestif: - min-complexity: 4 - - nlreturn: - block-size: 2 - - revive: - rules: - - name: var-naming - arguments: - - [] - - [] - - "upperCaseConst=true" - - tagliatelle: - case: - rules: - json: snake - yaml: snake - xml: snake - bson: snake - - testifylint: - enable-all: true - - usetesting: {} + default: all + disable: + # Genuinely incompatible with project patterns + - exhaustruct # Requires all struct fields + - depguard # Dependency allow/block lists + - godot # Requires comments to end with periods + - wsl # Deprecated, replaced by wsl_v5 + - wrapcheck # Too verbose for internal packages + - varnamelen # Short names like db, id are idiomatic Go + settings: + lll: + line-length: 88 + funlen: + lines: 80 + statements: 50 + cyclop: + max-complexity: 15 + dupl: + threshold: 100 issues: max-issues-per-linter: 0 max-same-issues: 0 - exclude-rules: - - path: ".*_gen\\.go" - linters: - - lll - - # Exclude unused parameter warnings for cobra command signatures - - text: "parameter '(args|cmd)' seems to be unused" - linters: - - revive - - # Allow ALL_CAPS constant names - - text: "don't use ALL_CAPS in Go names" - linters: - - revive - - # Exclude all linters for internal/macse directory - - path: "internal/macse/.*" - linters: - - errcheck - - lll - - mnd - - nestif - - nlreturn - - revive - - unconvert - - govet - - staticcheck - - unused - - ineffassign - - misspell - - gosec - - unparam - - testifylint - - usetesting - - tagliatelle - - nilnil - - intrange - - gochecknoglobals diff --git a/Dockerfile b/Dockerfile index 6e3c3a3..bfe4d5f 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,6 +1,6 @@ # Lint stage — fast feedback on formatting and lint issues -# golangci/golangci-lint v2.1.6 (2026-03-10) -FROM golangci/golangci-lint@sha256:568ee1c1c53493575fa9494e280e579ac9ca865787bafe4df3023ae59ecf299b AS lint +# golangci/golangci-lint:v2.12.2 (Debian-based), 2026-08-07 +FROM golangci/golangci-lint:v2.12.2@sha256:5cceeef04e53efe1470638d4b4b4f5ceefd574955ab3941b2d9a68a8c9ad5240 AS lint WORKDIR /src COPY go.mod go.sum ./ diff --git a/TODO.md b/TODO.md index 909a91e..9589ac3 100644 --- a/TODO.md +++ b/TODO.md @@ -25,6 +25,14 @@ Bring the repo into policy compliance in one commit: # Completed Steps +- 2026-08-07: Updated golangci-lint to v2.12.2 with the canonical + `.golangci.yml` (all linters enabled minus the standard disable + list, `lll` 88, tests linted); bumped the `Dockerfile` lint-stage + image to the tagged v2.12.2 Debian digest; fixed all ~1550 new + findings across `internal/` and `pkg/` (line wrapping, `wsl_v5` + blank lines, sentinel errors for `err113`, `t.Parallel()` where + safe, `_test` package conversions, complexity/`dupl` helper + extraction) on branch `golangci-v2.12.2`. - 2026-07-07 Adopted scripts-to-rule-them-all: `script/` entrypoints, Makefile shims, README Entrypoints section - 2026-03-11: Secure Enclave unlocker for hardware-backed secret @@ -50,6 +58,17 @@ Bring the repo into policy compliance in one commit: - Compliance (after Next Step lands): keep main green under the new .gitea workflow; run make check before every merge. +- Implement version-number shell completion for the second arg of + `secret version promote` and `secret version rm` + (`internal/cli/version.go`; was an in-code TODO removed for godox). +- Cover mnemonic-vs-xprv identity consistency in + `pkg/agehd/agehd_test.go` `TestMnemonicVsXPRVConsistency` (was an + in-code FIXME removed for godox). +- Darwin-gated files (`internal/secret/keychainunlocker.go`, + `seunlocker_darwin.go`, `internal/macse/macse_darwin.go`, related + tests) are not linted on the Linux CI runner and still contain lines + over the new 88-column limit; they will surface if lint ever runs on + macOS. - Merge secure-enclave-unlocker to main once review is done. - 1.0 critical security blockers (from repo TODO.md): - Command injection: GPG key IDs passed unescaped to exec.Command diff --git a/internal/cli/cli.go b/internal/cli/cli.go index 5141c83..071d287 100644 --- a/internal/cli/cli.go +++ b/internal/cli/cli.go @@ -19,6 +19,7 @@ type Instance struct { // NewCLIInstance creates a new CLI instance with the real filesystem func NewCLIInstance() (*Instance, error) { fs := afero.NewOsFs() + stateDir, err := secret.DetermineStateDir("") if err != nil { return nil, fmt.Errorf("cannot determine state directory: %w", err) @@ -30,7 +31,8 @@ func NewCLIInstance() (*Instance, error) { }, nil } -// NewCLIInstanceWithFs creates a new CLI instance with the given filesystem (for testing) +// NewCLIInstanceWithFs creates a new CLI instance with the given +// filesystem (for testing) func NewCLIInstanceWithFs(fs afero.Fs) (*Instance, error) { stateDir, err := secret.DetermineStateDir("") if err != nil { @@ -43,7 +45,8 @@ func NewCLIInstanceWithFs(fs afero.Fs) (*Instance, error) { }, nil } -// NewCLIInstanceWithStateDir creates a new CLI instance with custom state directory (for testing) +// NewCLIInstanceWithStateDir creates a new CLI instance with custom state +// directory (for testing) func NewCLIInstanceWithStateDir(fs afero.Fs, stateDir string) *Instance { return &Instance{ fs: fs, @@ -67,6 +70,6 @@ func (cli *Instance) GetStateDir() string { } // Print outputs to the command's configured output writer -func (cli *Instance) Print(a ...interface{}) (n int, err error) { +func (cli *Instance) Print(a ...any) (int, error) { return fmt.Fprint(cli.cmd.OutOrStdout(), a...) } diff --git a/internal/cli/cli_test.go b/internal/cli/cli_test.go index 32fd44b..edd3a04 100644 --- a/internal/cli/cli_test.go +++ b/internal/cli/cli_test.go @@ -1,37 +1,43 @@ -package cli +package cli_test import ( "os" "path/filepath" "testing" + "git.eeqj.de/sneak/secret/internal/cli" "git.eeqj.de/sneak/secret/internal/secret" "github.com/spf13/afero" ) func TestCLIInstanceStateDir(t *testing.T) { + t.Parallel() + // Test the CLI instance state directory functionality fs := afero.NewMemMapFs() // Create a test state directory testStateDir := "/test-state-dir" - cli := NewCLIInstanceWithStateDir(fs, testStateDir) + instance := cli.NewCLIInstanceWithStateDir(fs, testStateDir) - if cli.GetStateDir() != testStateDir { - t.Errorf("Expected state directory %q, got %q", testStateDir, cli.GetStateDir()) + got := instance.GetStateDir() + if got != testStateDir { + t.Errorf("Expected state directory %q, got %q", testStateDir, got) } } +//nolint:paralleltest // reads process environment to determine the state dir func TestCLIInstanceWithFs(t *testing.T) { // Test creating CLI instance with custom filesystem fs := afero.NewMemMapFs() - cli, err := NewCLIInstanceWithFs(fs) + + instance, err := cli.NewCLIInstanceWithFs(fs) if err != nil { t.Fatalf("failed to initialize CLI: %v", err) } // The state directory should be determined automatically - stateDir := cli.GetStateDir() + stateDir := instance.GetStateDir() if stateDir == "" { t.Error("Expected non-empty state directory") } @@ -48,6 +54,7 @@ func TestDetermineStateDir(t *testing.T) { if err != nil { t.Fatalf("unexpected error: %v", err) } + if stateDir != testEnvDir { t.Errorf("Expected state directory %q from environment, got %q", testEnvDir, stateDir) } @@ -55,12 +62,15 @@ func TestDetermineStateDir(t *testing.T) { // Test with custom config dir _ = os.Unsetenv(secret.EnvStateDir) customConfigDir := "/custom-config" + stateDir, err = secret.DetermineStateDir(customConfigDir) if err != nil { t.Fatalf("unexpected error: %v", err) } + expectedDir := filepath.Join(customConfigDir, secret.AppID) if stateDir != expectedDir { - t.Errorf("Expected state directory %q with custom config, got %q", expectedDir, stateDir) + t.Errorf("Expected state directory %q with custom config, got %q", + expectedDir, stateDir) } } diff --git a/internal/cli/completion.go b/internal/cli/completion.go index 5178631..5c9d5ad 100644 --- a/internal/cli/completion.go +++ b/internal/cli/completion.go @@ -1,12 +1,16 @@ package cli import ( + "errors" "fmt" "os" "github.com/spf13/cobra" ) +// errUnsupportedShell is returned for unknown shell completion targets +var errUnsupportedShell = errors.New("unsupported shell type") + func newCompletionCmd() *cobra.Command { cmd := &cobra.Command{ Use: "completion [bash|zsh|fish|powershell]", @@ -55,7 +59,7 @@ PowerShell: case "powershell": return cmd.Root().GenPowerShellCompletionWithDesc(os.Stdout) default: - return fmt.Errorf("unsupported shell type: %s", args[0]) + return fmt.Errorf("%w: %s", errUnsupportedShell, args[0]) } }, } diff --git a/internal/cli/completions.go b/internal/cli/completions.go index 371f632..67b2a70 100644 --- a/internal/cli/completions.go +++ b/internal/cli/completions.go @@ -1,21 +1,22 @@ package cli import ( - "encoding/json" "path/filepath" "strings" - "git.eeqj.de/sneak/secret/internal/secret" "git.eeqj.de/sneak/secret/internal/vault" "github.com/spf13/afero" "github.com/spf13/cobra" ) -// getSecretNamesCompletionFunc returns a completion function that provides secret names +// getSecretNamesCompletionFunc returns a completion function that provides +// secret names func getSecretNamesCompletionFunc(fs afero.Fs, stateDir string) func( cmd *cobra.Command, args []string, toComplete string, ) ([]string, cobra.ShellCompDirective) { - return func(_ *cobra.Command, _ []string, toComplete string) ([]string, cobra.ShellCompDirective) { + return func( + _ *cobra.Command, _ []string, toComplete string, + ) ([]string, cobra.ShellCompDirective) { // Get current vault vlt, err := vault.GetCurrentVault(fs, stateDir) if err != nil { @@ -30,6 +31,7 @@ func getSecretNamesCompletionFunc(fs afero.Fs, stateDir string) func( // Filter secrets based on what user has typed var completions []string + for _, secret := range secrets { if strings.HasPrefix(secret, toComplete) { completions = append(completions, secret) @@ -40,11 +42,14 @@ func getSecretNamesCompletionFunc(fs afero.Fs, stateDir string) func( } } -// getUnlockerIDsCompletionFunc returns a completion function that provides unlocker IDs +// getUnlockerIDsCompletionFunc returns a completion function that provides +// unlocker IDs func getUnlockerIDsCompletionFunc(fs afero.Fs, stateDir string) func( cmd *cobra.Command, args []string, toComplete string, ) ([]string, cobra.ShellCompDirective) { - return func(_ *cobra.Command, _ []string, toComplete string) ([]string, cobra.ShellCompDirective) { + return func( + _ *cobra.Command, _ []string, toComplete string, + ) ([]string, cobra.ShellCompDirective) { // Get current vault vlt, err := vault.GetCurrentVault(fs, stateDir) if err != nil { @@ -66,61 +71,13 @@ func getUnlockerIDsCompletionFunc(fs afero.Fs, stateDir string) func( // Collect unlocker IDs var completions []string + unlockersDir := filepath.Join(vaultDir, "unlockers.d") + for _, metadata := range unlockerMetadataList { // Get the actual unlocker ID by creating the unlocker instance - unlockersDir := filepath.Join(vaultDir, "unlockers.d") - files, err := afero.ReadDir(fs, unlockersDir) - if err != nil { - secret.Warn("Could not read unlockers directory during completion", "error", err) - - continue - } - - for _, file := range files { - if !file.IsDir() { - continue - } - - unlockerDir := filepath.Join(unlockersDir, file.Name()) - metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") - - // Check if this is the right unlocker by comparing metadata - metadataBytes, err := afero.ReadFile(fs, metadataPath) - if err != nil { - secret.Warn("Could not read unlocker metadata during completion", "path", metadataPath, "error", err) - - continue - } - - var diskMetadata secret.UnlockerMetadata - if err := json.Unmarshal(metadataBytes, &diskMetadata); err != nil { - secret.Warn("Could not parse unlocker metadata during completion", "path", metadataPath, "error", err) - - continue - } - - // Match by type and creation time - if diskMetadata.Type == metadata.Type && diskMetadata.CreatedAt.Equal(metadata.CreatedAt) { - // Create the appropriate unlocker instance - var unlocker secret.Unlocker - switch metadata.Type { - case "passphrase": - unlocker = secret.NewPassphraseUnlocker(fs, unlockerDir, diskMetadata) - case "keychain": - unlocker = secret.NewKeychainUnlocker(fs, unlockerDir, diskMetadata) - case "pgp": - unlocker = secret.NewPGPUnlocker(fs, unlockerDir, diskMetadata) - } - - if unlocker != nil { - id := unlocker.GetID() - if strings.HasPrefix(id, toComplete) { - completions = append(completions, id) - } - } - - break - } + id := findUnlockerIDByMetadata(fs, unlockersDir, metadata, false) + if id != "" && strings.HasPrefix(id, toComplete) { + completions = append(completions, id) } } @@ -128,17 +85,21 @@ func getUnlockerIDsCompletionFunc(fs afero.Fs, stateDir string) func( } } -// getVaultNamesCompletionFunc returns a completion function that provides vault names +// getVaultNamesCompletionFunc returns a completion function that provides +// vault names func getVaultNamesCompletionFunc(fs afero.Fs, stateDir string) func( cmd *cobra.Command, args []string, toComplete string, ) ([]string, cobra.ShellCompDirective) { - return func(_ *cobra.Command, _ []string, toComplete string) ([]string, cobra.ShellCompDirective) { + return func( + _ *cobra.Command, _ []string, toComplete string, + ) ([]string, cobra.ShellCompDirective) { vaults, err := vault.ListVaults(fs, stateDir) if err != nil { return nil, cobra.ShellCompDirectiveNoFileComp } var completions []string + for _, v := range vaults { if strings.HasPrefix(v, toComplete) { completions = append(completions, v) @@ -149,57 +110,81 @@ func getVaultNamesCompletionFunc(fs afero.Fs, stateDir string) func( } } -// getVaultSecretCompletionFunc returns a completion function for vault:secret format -// It completes vault names with ":" suffix, and after ":" it completes secrets from that vault +// completeVaultQualifiedSecrets completes "vault:secret" references once a +// colon is present in the input +func completeVaultQualifiedSecrets( + fs afero.Fs, stateDir, toComplete string, +) []string { + var completions []string + + // Complete secret names for the specified vault + parts := strings.SplitN(toComplete, ":", vaultSecretParts) + vaultName := parts[0] + secretPrefix := parts[1] + + vlt := vault.NewVault(fs, stateDir, vaultName) + + secrets, err := vlt.ListSecrets() + if err == nil { + for _, secretName := range secrets { + if strings.HasPrefix(secretName, secretPrefix) { + completions = append(completions, vaultName+":"+secretName) + } + } + } + + return completions +} + +// completeUnqualifiedVaultSecrets completes vault names (with a ":" +// suffix) and secrets from the current vault +func completeUnqualifiedVaultSecrets( + fs afero.Fs, stateDir, toComplete string, +) []string { + var completions []string + + // Complete vault names with ":" suffix + vaults, err := vault.ListVaults(fs, stateDir) + if err == nil { + for _, v := range vaults { + if strings.HasPrefix(v, toComplete) { + completions = append(completions, v+":") + } + } + } + + // Also complete secrets from current vault (for within-vault moves) + currentVlt, err := vault.GetCurrentVault(fs, stateDir) + if err == nil { + secrets, err := currentVlt.ListSecrets() + if err == nil { + for _, secretName := range secrets { + if strings.HasPrefix(secretName, toComplete) { + completions = append(completions, secretName) + } + } + } + } + + return completions +} + +// getVaultSecretCompletionFunc returns a completion function for the +// vault:secret format. It completes vault names with ":" suffix, and +// after ":" it completes secrets from that vault. func getVaultSecretCompletionFunc(fs afero.Fs, stateDir string) func( cmd *cobra.Command, args []string, toComplete string, ) ([]string, cobra.ShellCompDirective) { - return func(_ *cobra.Command, _ []string, toComplete string) ([]string, cobra.ShellCompDirective) { - var completions []string - + return func( + _ *cobra.Command, _ []string, toComplete string, + ) ([]string, cobra.ShellCompDirective) { // Check if we're completing after a vault: prefix if strings.Contains(toComplete, ":") { - // Complete secret names for the specified vault - const vaultSecretParts = 2 - parts := strings.SplitN(toComplete, ":", vaultSecretParts) - vaultName := parts[0] - secretPrefix := parts[1] - - vlt := vault.NewVault(fs, stateDir, vaultName) - secrets, err := vlt.ListSecrets() - if err == nil { - for _, secretName := range secrets { - if strings.HasPrefix(secretName, secretPrefix) { - completions = append(completions, vaultName+":"+secretName) - } - } - } - - return completions, cobra.ShellCompDirectiveNoFileComp + return completeVaultQualifiedSecrets(fs, stateDir, toComplete), + cobra.ShellCompDirectiveNoFileComp } - // Complete vault names with ":" suffix - vaults, err := vault.ListVaults(fs, stateDir) - if err == nil { - for _, v := range vaults { - if strings.HasPrefix(v, toComplete) { - completions = append(completions, v+":") - } - } - } - - // Also complete secrets from current vault (for within-vault moves) - if currentVlt, err := vault.GetCurrentVault(fs, stateDir); err == nil { - secrets, err := currentVlt.ListSecrets() - if err == nil { - for _, secretName := range secrets { - if strings.HasPrefix(secretName, toComplete) { - completions = append(completions, secretName) - } - } - } - } - - return completions, cobra.ShellCompDirectiveNoSpace + return completeUnqualifiedVaultSecrets(fs, stateDir, toComplete), + cobra.ShellCompDirectiveNoSpace } } diff --git a/internal/cli/crypto.go b/internal/cli/crypto.go index 263a9e0..723b472 100644 --- a/internal/cli/crypto.go +++ b/internal/cli/crypto.go @@ -1,6 +1,7 @@ package cli import ( + "errors" "fmt" "io" "os" @@ -12,11 +13,22 @@ import ( "github.com/spf13/cobra" ) -func newEncryptCmd() *cobra.Command { +// Sentinel errors for encrypt/decrypt operations +var ( + errNotAgeSecretKey = errors.New( + "does not contain a valid age secret key") + errSecretDoesNotExist = errors.New("does not exist") +) + +// newCryptoCmd builds an encrypt/decrypt command with input/output flags +func newCryptoCmd( + use, short, long string, + run func(cli *Instance, secretName, inputFile, outputFile string) error, +) *cobra.Command { cmd := &cobra.Command{ - Use: "encrypt ", - Short: "Encrypt data using an age secret key stored in a secret", - Long: `Encrypt data using an age secret key. If the secret doesn't exist, a new age key is generated and stored.`, + Use: use, + Short: short, + Long: long, Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { inputFile, _ := cmd.Flags().GetString("input") @@ -26,9 +38,10 @@ func newEncryptCmd() *cobra.Command { if err != nil { return fmt.Errorf("failed to initialize CLI: %w", err) } + cli.cmd = cmd - return cli.Encrypt(args[0], inputFile, outputFile) + return run(cli, args[0], inputFile, outputFile) }, } @@ -38,30 +51,73 @@ func newEncryptCmd() *cobra.Command { return cmd } +func newEncryptCmd() *cobra.Command { + return newCryptoCmd( + "encrypt ", + "Encrypt data using an age secret key stored in a secret", + "Encrypt data using an age secret key. If the secret doesn't "+ + "exist, a new age key is generated and stored.", + (*Instance).Encrypt, + ) +} + func newDecryptCmd() *cobra.Command { - cmd := &cobra.Command{ - Use: "decrypt ", - Short: "Decrypt data using an age secret key stored in a secret", - Long: `Decrypt data using an age secret key stored in the specified secret.`, - Args: cobra.ExactArgs(1), - RunE: func(cmd *cobra.Command, args []string) error { - inputFile, _ := cmd.Flags().GetString("input") - outputFile, _ := cmd.Flags().GetString("output") + return newCryptoCmd( + "decrypt ", + "Decrypt data using an age secret key stored in a secret", + "Decrypt data using an age secret key stored in the specified secret.", + (*Instance).Decrypt, + ) +} - cli, err := NewCLIInstance() - if err != nil { - return fmt.Errorf("failed to initialize CLI: %w", err) - } - cli.cmd = cmd +// resolveEncryptionKey returns a secure buffer holding the age secret key +// for the named secret, generating and storing a new key if the secret +// does not exist. The caller must destroy the returned buffer. +func (cli *Instance) resolveEncryptionKey( + vlt *vault.Vault, secretName string, +) (*memguard.LockedBuffer, error) { + // Check if secret exists + secretObj := secret.NewSecret(vlt, secretName) - return cli.Decrypt(args[0], inputFile, outputFile) - }, + exists, err := secretObj.Exists() + if err != nil { + return nil, fmt.Errorf("failed to check if secret exists: %w", err) } - cmd.Flags().StringP("input", "i", "", "Input file (default: stdin)") - cmd.Flags().StringP("output", "o", "", "Output file (default: stdout)") + if !exists { + // Secret doesn't exist, generate new age key and store it + identity, err := age.GenerateX25519Identity() + if err != nil { + return nil, fmt.Errorf("failed to generate age key: %w", err) + } - return cmd + // Store the generated key directly in a secure buffer + secureBuffer := memguard.NewBufferFromBytes([]byte(identity.String())) + + err = vlt.AddSecret(secretName, secureBuffer, false) + if err != nil { + secureBuffer.Destroy() + + return nil, fmt.Errorf("failed to store age key: %w", err) + } + + return secureBuffer, nil + } + + // Secret exists, get the age secret key from it + secretBuffer, err := cli.getSecretValue(vlt, secretObj) + if err != nil { + return nil, fmt.Errorf("failed to get secret value: %w", err) + } + + // Validate that it's a valid age secret key + if !isValidAgeSecretKey(secretBuffer.String()) { + secretBuffer.Destroy() + + return nil, fmt.Errorf("secret '%s' %w", secretName, errNotAgeSecretKey) + } + + return secretBuffer, nil } // Encrypt encrypts data using an age secret key stored in a secret @@ -72,55 +128,15 @@ func (cli *Instance) Encrypt(secretName, inputFile, outputFile string) error { return err } - var ageSecretKey string - - // Check if secret exists - secretObj := secret.NewSecret(vlt, secretName) - exists, err := secretObj.Exists() + // Get or create the age secret key for this secret + keyBuffer, err := cli.resolveEncryptionKey(vlt, secretName) if err != nil { - return fmt.Errorf("failed to check if secret exists: %w", err) + return err } + defer keyBuffer.Destroy() - if !exists { //nolint:nestif // Clear conditional logic for secret generation vs retrieval - // Secret doesn't exist, generate new age key and store it - identity, err := age.GenerateX25519Identity() - if err != nil { - return fmt.Errorf("failed to generate age key: %w", err) - } - - // Store the generated key directly in a secure buffer - identityStr := identity.String() - secureBuffer := memguard.NewBufferFromBytes([]byte(identityStr)) - defer secureBuffer.Destroy() - - // Set ageSecretKey for later use (we need it for encryption) - ageSecretKey = identityStr - - err = vlt.AddSecret(secretName, secureBuffer, false) - if err != nil { - return fmt.Errorf("failed to store age key: %w", err) - } - } else { - // Secret exists, get the age secret key from it - secretBuffer, err := cli.getSecretValue(vlt, secretObj) - if err != nil { - return fmt.Errorf("failed to get secret value: %w", err) - } - defer secretBuffer.Destroy() - - ageSecretKey = secretBuffer.String() - - // Validate that it's a valid age secret key - if !isValidAgeSecretKey(ageSecretKey) { - return fmt.Errorf("secret '%s' does not contain a valid age secret key", secretName) - } - } - - // Parse the secret key using secure buffer - finalSecureBuffer := memguard.NewBufferFromBytes([]byte(ageSecretKey)) - defer finalSecureBuffer.Destroy() - - identity, err := age.ParseX25519Identity(finalSecureBuffer.String()) + // Parse the secret key + identity, err := age.ParseX25519Identity(keyBuffer.String()) if err != nil { return fmt.Errorf("failed to parse age secret key: %w", err) } @@ -130,23 +146,27 @@ func (cli *Instance) Encrypt(secretName, inputFile, outputFile string) error { // Set up input reader var input io.Reader = os.Stdin + if inputFile != "" { file, err := cli.fs.Open(inputFile) if err != nil { return fmt.Errorf("failed to open input file: %w", err) } defer func() { _ = file.Close() }() + input = file } // Set up output writer output := cli.cmd.OutOrStdout() + if outputFile != "" { file, err := cli.fs.Create(outputFile) if err != nil { return fmt.Errorf("failed to create output file: %w", err) } defer func() { _ = file.Close() }() + output = file } @@ -156,11 +176,13 @@ func (cli *Instance) Encrypt(secretName, inputFile, outputFile string) error { return fmt.Errorf("failed to create age encryptor: %w", err) } - if _, err := io.Copy(encryptor, input); err != nil { + _, err = io.Copy(encryptor, input) + if err != nil { return fmt.Errorf("failed to encrypt data: %w", err) } - if err := encryptor.Close(); err != nil { + err = encryptor.Close() + if err != nil { return fmt.Errorf("failed to finalize encryption: %w", err) } @@ -177,26 +199,18 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error { // Check if secret exists secretObj := secret.NewSecret(vlt, secretName) + exists, err := secretObj.Exists() if err != nil { return fmt.Errorf("failed to check if secret exists: %w", err) } if !exists { - return fmt.Errorf("secret '%s' does not exist", secretName) + return fmt.Errorf("secret '%s' %w", secretName, errSecretDoesNotExist) } // Get the age secret key from the secret - var secretBuffer *memguard.LockedBuffer - if os.Getenv(secret.EnvMnemonic) != "" { - secretBuffer, err = secretObj.GetValue(nil) - } else { - unlocker, unlockErr := vlt.GetCurrentUnlocker() - if unlockErr != nil { - return fmt.Errorf("failed to get current unlocker: %w", unlockErr) - } - secretBuffer, err = secretObj.GetValue(unlocker) - } + secretBuffer, err := cli.getSecretValue(vlt, secretObj) if err != nil { return fmt.Errorf("failed to get secret value: %w", err) } @@ -204,7 +218,7 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error { // Validate that it's a valid age secret key if !isValidAgeSecretKey(secretBuffer.String()) { - return fmt.Errorf("secret '%s' does not contain a valid age secret key", secretName) + return fmt.Errorf("secret '%s' %w", secretName, errNotAgeSecretKey) } // Parse the age secret key to get the identity @@ -215,23 +229,27 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error { // Set up input reader var input io.Reader = os.Stdin + if inputFile != "" { file, err := cli.fs.Open(inputFile) if err != nil { return fmt.Errorf("failed to open input file: %w", err) } defer func() { _ = file.Close() }() + input = file } // Set up output writer output := cli.cmd.OutOrStdout() + if outputFile != "" { file, err := cli.fs.Create(outputFile) if err != nil { return fmt.Errorf("failed to create output file: %w", err) } defer func() { _ = file.Close() }() + output = file } @@ -241,22 +259,27 @@ func (cli *Instance) Decrypt(secretName, inputFile, outputFile string) error { return fmt.Errorf("failed to create age decryptor: %w", err) } - if _, err := io.Copy(output, decryptor); err != nil { + _, err = io.Copy(output, decryptor) + if err != nil { return fmt.Errorf("failed to decrypt data: %w", err) } return nil } -// isValidAgeSecretKey checks if a string is a valid age secret key by attempting to parse it +// isValidAgeSecretKey checks if a string is a valid age secret key by +// attempting to parse it func isValidAgeSecretKey(key string) bool { _, err := age.ParseX25519Identity(key) return err == nil } -// getSecretValue retrieves the value of a secret using the appropriate unlocker -func (cli *Instance) getSecretValue(vlt *vault.Vault, secretObj *secret.Secret) (*memguard.LockedBuffer, error) { +// getSecretValue retrieves the value of a secret using the appropriate +// unlocker +func (cli *Instance) getSecretValue( + vlt *vault.Vault, secretObj *secret.Secret, +) (*memguard.LockedBuffer, error) { if os.Getenv(secret.EnvMnemonic) != "" { return secretObj.GetValue(nil) } diff --git a/internal/cli/generate.go b/internal/cli/generate.go index 623ccbd..6b15640 100644 --- a/internal/cli/generate.go +++ b/internal/cli/generate.go @@ -2,6 +2,7 @@ package cli import ( "crypto/rand" + "errors" "fmt" "math/big" "os" @@ -17,6 +18,16 @@ const ( mnemonicEntropyBits = 128 ) +// Sentinel errors for secret generation +var ( + errLengthTooSmall = errors.New("length must be at least 1") + errLengthNotPositive = errors.New("length must be positive") + errMnemonicTypeNotSupported = errors.New( + "mnemonic type not supported for secret generation, " + + "use 'secret generate mnemonic' instead") + errUnsupportedSecretType = errors.New("unsupported type") +) + func newGenerateCmd() *cobra.Command { cmd := &cobra.Command{ Use: "generate", @@ -52,8 +63,9 @@ func newGenerateSecretCmd() *cobra.Command { cmd := &cobra.Command{ Use: "secret ", Short: "Generate a random secret and store it in the vault", - Long: `Generate a cryptographically secure random secret and store it in the current vault under the given name.`, - Args: cobra.ExactArgs(1), + Long: `Generate a cryptographically secure random secret and ` + + `store it in the current vault under the given name.`, + Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { length, _ := cmd.Flags().GetInt("length") secretType, _ := cmd.Flags().GetString("type") @@ -68,8 +80,10 @@ func newGenerateSecretCmd() *cobra.Command { }, } - cmd.Flags().IntP("length", "l", defaultSecretLength, "Length of the generated secret (default 16)") - cmd.Flags().StringP("type", "t", "base58", "Type of secret to generate (base58, alnum)") + cmd.Flags().IntP("length", "l", defaultSecretLength, + "Length of the generated secret (default 16)") + cmd.Flags().StringP("type", "t", "base58", + "Type of secret to generate (base58, alnum)") cmd.Flags().BoolP("force", "f", false, "Overwrite existing secret") return cmd @@ -98,7 +112,8 @@ func (cli *Instance) GenerateMnemonic(cmd *cobra.Command) error { fmt.Fprintln(os.Stderr, " • Write it down on paper and store it safely") fmt.Fprintln(os.Stderr, " • Do not store it digitally or share it with anyone") fmt.Fprintln(os.Stderr, " • You will need this phrase to recover your secrets") - fmt.Fprintln(os.Stderr, " • If you lose this phrase, your secrets cannot be recovered") + fmt.Fprintln(os.Stderr, + " • If you lose this phrase, your secrets cannot be recovered") fmt.Fprintln(os.Stderr, "") fmt.Fprintln(os.Stderr, "Use this mnemonic with:") fmt.Fprintln(os.Stderr, " secret init (to initialize a new secret manager)") @@ -116,11 +131,13 @@ func (cli *Instance) GenerateSecret( force bool, ) error { if length < 1 { - return fmt.Errorf("length must be at least 1") + return errLengthTooSmall } - var secretValue string - var err error + var ( + secretValue string + err error + ) switch secretType { case "base58": @@ -128,9 +145,10 @@ func (cli *Instance) GenerateSecret( case "alnum": secretValue, err = generateRandomAlnum(length) case "mnemonic": - return fmt.Errorf("mnemonic type not supported for secret generation, use 'secret generate mnemonic' instead") + return errMnemonicTypeNotSupported default: - return fmt.Errorf("unsupported type: %s (supported: base58, alnum)", secretType) + return fmt.Errorf("%w: %s (supported: base58, alnum)", + errUnsupportedSecretType, secretType) } if err != nil { @@ -147,11 +165,13 @@ func (cli *Instance) GenerateSecret( secretBuffer := memguard.NewBufferFromBytes([]byte(secretValue)) defer secretBuffer.Destroy() - if err := vlt.AddSecret(secretName, secretBuffer, force); err != nil { + err = vlt.AddSecret(secretName, secretBuffer, force) + if err != nil { return err } - cmd.Printf("Generated and stored %d-character %s secret: %s\n", length, secretType, secretName) + cmd.Printf("Generated and stored %d-character %s secret: %s\n", + length, secretType, secretName) return nil } @@ -170,10 +190,11 @@ func generateRandomAlnum(length int) (string, error) { return generateRandomString(length, alnumChars) } -// generateRandomString generates a random string of the specified length using the given character set +// generateRandomString generates a random string of the specified length +// using the given character set func generateRandomString(length int, charset string) (string, error) { if length <= 0 { - return "", fmt.Errorf("length must be positive") + return "", errLengthNotPositive } result := make([]byte, length) @@ -184,6 +205,7 @@ func generateRandomString(length int, charset string) (string, error) { if err != nil { return "", fmt.Errorf("failed to generate random number: %w", err) } + result[i] = charset[randomIndex.Int64()] } diff --git a/internal/cli/info.go b/internal/cli/info.go index e993d91..6298e27 100644 --- a/internal/cli/info.go +++ b/internal/cli/info.go @@ -18,7 +18,7 @@ import ( ) // Version info - these are set at build time -var ( //nolint:gochecknoglobals // Set at build time +var ( Version = "dev" //nolint:gochecknoglobals // Set at build time GitCommit = "unknown" //nolint:gochecknoglobals // Set at build time ) @@ -35,8 +35,8 @@ type InfoOutput struct { NumVaults int `json:"numVaults"` NumSecrets int `json:"numSecrets"` TotalSize int64 `json:"totalSizeBytes"` - OldestSecret time.Time `json:"oldestSecret,omitempty"` - LatestSecret time.Time `json:"latestSecret,omitempty"` + OldestSecret time.Time `json:"oldestSecret"` + LatestSecret time.Time `json:"latestSecret"` } // newInfoCmd returns the info command @@ -51,7 +51,8 @@ func newInfoCmd() *cobra.Command { cmd := &cobra.Command{ Use: "info", Short: "Display system information", - Long: "Display information about the secret system including version, vault statistics, and storage usage", + Long: "Display information about the secret system including " + + "version, vault statistics, and storage usage", RunE: func(cmd *cobra.Command, _ []string) error { return cli.Info(cmd, jsonOutput) }, @@ -81,6 +82,7 @@ func (cli *Instance) Info(cmd *cobra.Command, jsonOutput bool) error { // Count vaults vaultsDir := filepath.Join(cli.stateDir, "vaults.d") + vaultEntries, err := afero.ReadDir(cli.fs, vaultsDir) if err == nil { for _, entry := range vaultEntries { @@ -92,12 +94,15 @@ func (cli *Instance) Info(cmd *cobra.Command, jsonOutput bool) error { // Gather statistics from all vaults if info.NumVaults > 0 { - totalSecrets, totalSize, oldestTime, latestTime, _ := gatherVaultStats(cli.fs, vaultsDir) + totalSecrets, totalSize, oldestTime, latestTime, _ := gatherVaultStats( + cli.fs, vaultsDir) info.NumSecrets = totalSecrets info.TotalSize = totalSize + if !oldestTime.IsZero() { info.OldestSecret = oldestTime } + if !latestTime.IsZero() { info.LatestSecret = latestTime } @@ -144,19 +149,24 @@ func prettyPrintInfo(w io.Writer, info InfoOutput) error { _, _ = fmt.Fprintln(w, strings.Repeat("─", separatorLength)) _, _ = fmt.Fprintf(w, "🗂️ Vaults: %s\n", bold.Sprint(info.NumVaults)) + _, _ = fmt.Fprintf(w, "🔑 Secrets: %s\n", bold.Sprint(info.NumSecrets)) + if info.TotalSize >= 0 { - //nolint:gosec // TotalSize is always >= 0 - _, _ = fmt.Fprintf(w, "💾 Total Size: %s\n", bold.Sprint(humanize.Bytes(uint64(info.TotalSize)))) + _, _ = fmt.Fprintf(w, "💾 Total Size: %s\n", + bold.Sprint(humanize.Bytes(uint64(info.TotalSize)))) } else { _, _ = fmt.Fprintf(w, "💾 Total Size: %s\n", bold.Sprint("0 B")) } if !info.OldestSecret.IsZero() { - _, _ = fmt.Fprintf(w, "🕰️ Oldest Secret: %s\n", info.OldestSecret.Format("2006-01-02 15:04:05")) + _, _ = fmt.Fprintf(w, "🕰️ Oldest Secret: %s\n", + info.OldestSecret.Format("2006-01-02 15:04:05")) } + if !info.LatestSecret.IsZero() { - _, _ = fmt.Fprintf(w, "✨ Latest Secret: %s\n", info.LatestSecret.Format("2006-01-02 15:04:05")) + _, _ = fmt.Fprintf(w, "✨ Latest Secret: %s\n", + info.LatestSecret.Format("2006-01-02 15:04:05")) } _, _ = fmt.Fprintln(w) diff --git a/internal/cli/info_helper.go b/internal/cli/info_helper.go index 4ed174c..3cb65fa 100644 --- a/internal/cli/info_helper.go +++ b/internal/cli/info_helper.go @@ -8,81 +8,115 @@ import ( "github.com/spf13/afero" ) -// gatherVaultStats collects statistics from all vaults +// vaultStats accumulates statistics while walking vault directories +type vaultStats struct { + totalSecrets int + totalSize int64 + oldestTime time.Time + latestTime time.Time +} + +// addVersion accumulates size and timestamp info for one version directory +func (s *vaultStats) addVersion(fs afero.Fs, versionPath string) { + // Add size of encrypted data + dataPath := filepath.Join(versionPath, "data.age") + + stat, err := fs.Stat(dataPath) + if err == nil { + s.totalSize += stat.Size() + } + + // Add size of metadata + metaPath := filepath.Join(versionPath, "metadata.age") + + stat, err = fs.Stat(metaPath) + if err == nil { + s.totalSize += stat.Size() + } + + // Track timestamps + stat, err = fs.Stat(versionPath) + if err == nil { + modTime := stat.ModTime() + if s.oldestTime.IsZero() || modTime.Before(s.oldestTime) { + s.oldestTime = modTime + } + + if s.latestTime.IsZero() || modTime.After(s.latestTime) { + s.latestTime = modTime + } + } +} + +// addSecret accumulates stats for one secret directory +func (s *vaultStats) addSecret(fs afero.Fs, secretsPath, secretName string) { + s.totalSecrets++ + secretPath := filepath.Join(secretsPath, secretName) + + // Get size and timestamps from all versions + versionsPath := filepath.Join(secretPath, "versions") + + versionEntries, err := afero.ReadDir(fs, versionsPath) + if err != nil { + secret.Warn("Could not read versions directory for secret", + "secret", secretName, "error", err) + + return + } + + for _, versionEntry := range versionEntries { + if !versionEntry.IsDir() { + continue + } + + s.addVersion(fs, filepath.Join(versionsPath, versionEntry.Name())) + } +} + +// addVault accumulates stats for one vault directory +func (s *vaultStats) addVault(fs afero.Fs, vaultsDir, vaultName string) { + vaultPath := filepath.Join(vaultsDir, vaultName) + secretsPath := filepath.Join(vaultPath, "secrets.d") + + // Count secrets in this vault + secretEntries, err := afero.ReadDir(fs, secretsPath) + if err != nil { + secret.Warn("Could not read secrets directory for vault", + "vault", vaultName, "error", err) + + return + } + + for _, secretEntry := range secretEntries { + if !secretEntry.IsDir() { + continue + } + + s.addSecret(fs, secretsPath, secretEntry.Name()) + } +} + +// gatherVaultStats collects statistics from all vaults, returning the +// total secret count, total size, and oldest/latest secret timestamps func gatherVaultStats( fs afero.Fs, vaultsDir string, -) (totalSecrets int, totalSize int64, oldestTime, latestTime time.Time, err error) { +) (int, int64, time.Time, time.Time, error) { vaultEntries, err := afero.ReadDir(fs, vaultsDir) if err != nil { return 0, 0, time.Time{}, time.Time{}, err } + var stats vaultStats + for _, vaultEntry := range vaultEntries { if !vaultEntry.IsDir() { continue } - vaultPath := filepath.Join(vaultsDir, vaultEntry.Name()) - secretsPath := filepath.Join(vaultPath, "secrets.d") - - // Count secrets in this vault - secretEntries, err := afero.ReadDir(fs, secretsPath) - if err != nil { - secret.Warn("Could not read secrets directory for vault", "vault", vaultEntry.Name(), "error", err) - - continue - } - - for _, secretEntry := range secretEntries { - if !secretEntry.IsDir() { - continue - } - - totalSecrets++ - secretPath := filepath.Join(secretsPath, secretEntry.Name()) - - // Get size and timestamps from all versions - versionsPath := filepath.Join(secretPath, "versions") - versionEntries, err := afero.ReadDir(fs, versionsPath) - if err != nil { - secret.Warn("Could not read versions directory for secret", "secret", secretEntry.Name(), "error", err) - - continue - } - - for _, versionEntry := range versionEntries { - if !versionEntry.IsDir() { - continue - } - - versionPath := filepath.Join(versionsPath, versionEntry.Name()) - - // Add size of encrypted data - dataPath := filepath.Join(versionPath, "data.age") - if stat, err := fs.Stat(dataPath); err == nil { - totalSize += stat.Size() - } - - // Add size of metadata - metaPath := filepath.Join(versionPath, "metadata.age") - if stat, err := fs.Stat(metaPath); err == nil { - totalSize += stat.Size() - } - - // Track timestamps - if stat, err := fs.Stat(versionPath); err == nil { - modTime := stat.ModTime() - if oldestTime.IsZero() || modTime.Before(oldestTime) { - oldestTime = modTime - } - if latestTime.IsZero() || modTime.After(latestTime) { - latestTime = modTime - } - } - } - } + stats.addVault(fs, vaultsDir, vaultEntry.Name()) } - return totalSecrets, totalSize, oldestTime, latestTime, nil + return stats.totalSecrets, stats.totalSize, + stats.oldestTime, stats.latestTime, nil } diff --git a/internal/cli/init.go b/internal/cli/init.go index 14590bc..6162c7e 100644 --- a/internal/cli/init.go +++ b/internal/cli/init.go @@ -1,6 +1,7 @@ package cli import ( + "errors" "fmt" "log" "log/slog" @@ -8,6 +9,7 @@ import ( "path/filepath" "strings" + "filippo.io/age" "git.eeqj.de/sneak/secret/internal/secret" "git.eeqj.de/sneak/secret/internal/vault" "git.eeqj.de/sneak/secret/pkg/agehd" @@ -16,13 +18,17 @@ import ( "github.com/tyler-smith/go-bip39" ) +// errPassphraseMismatch is returned when passphrase confirmation fails +var errPassphraseMismatch = errors.New("passphrases do not match") + // NewInitCmd creates the init command func NewInitCmd() *cobra.Command { return &cobra.Command{ Use: "init", Short: "Initialize the secrets manager", - Long: `Create the necessary directory structure for storing secrets and generate encryption keys.`, - RunE: RunInit, + Long: `Create the necessary directory structure for storing ` + + `secrets and generate encryption keys.`, + RunE: RunInit, } } @@ -36,6 +42,67 @@ func RunInit(cmd *cobra.Command, _ []string) error { return cli.Init(cmd) } +// promptMnemonic reads the mnemonic from the environment or interactively. +// The returned cleanup function must be deferred by the caller. +func promptMnemonic() (string, func(), error) { + if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" { + secret.Debug("Using mnemonic from environment variable") + + return envMnemonic, func() {}, nil + } + + secret.Debug("Prompting user for mnemonic phrase") + + // Read mnemonic securely without echo + mnemonicBuffer, err := secret.ReadPassphrase("Enter your BIP39 mnemonic phrase: ") + if err != nil { + secret.Debug("Failed to read mnemonic from stdin", "error", err) + + return "", nil, fmt.Errorf("failed to read mnemonic: %w", err) + } + + fmt.Fprintln(os.Stderr) // Add newline after hidden input + + return mnemonicBuffer.String(), mnemonicBuffer.Destroy, nil +} + +// setupDefaultVault creates the default vault and derives its long-term +// identity from the mnemonic +func (cli *Instance) setupDefaultVault( + stateDir, mnemonicStr string, +) (*vault.Vault, *age.X25519Identity, error) { + // Create the default vault - it will handle key derivation internally + secret.Debug("Creating default vault") + + vlt, err := vault.CreateVault(cli.fs, cli.stateDir, "default") + if err != nil { + secret.Debug("Failed to create default vault", "error", err) + + return nil, nil, fmt.Errorf("failed to create default vault: %w", err) + } + + // Get the vault metadata to retrieve the derivation index + vaultDir := filepath.Join(stateDir, "vaults.d", "default") + + metadata, err := vault.LoadVaultMetadata(cli.fs, vaultDir) + if err != nil { + secret.Debug("Failed to load vault metadata", "error", err) + + return nil, nil, fmt.Errorf("failed to load vault metadata: %w", err) + } + + // Derive the long-term key using the same index that CreateVault used + ltIdentity, err := agehd.DeriveIdentity(mnemonicStr, metadata.DerivationIndex) + if err != nil { + secret.Debug("Failed to derive long-term key", "error", err) + + return nil, nil, fmt.Errorf( + "failed to derive long-term key from mnemonic: %w", err) + } + + return vlt, ltIdentity, nil +} + // Init initializes the secret manager func (cli *Instance) Init(cmd *cobra.Command) error { secret.Debug("Starting secret manager initialization") @@ -44,7 +111,8 @@ func (cli *Instance) Init(cmd *cobra.Command) error { stateDir := cli.GetStateDir() secret.DebugWith("Creating state directory", slog.String("path", stateDir)) - if err := cli.fs.MkdirAll(stateDir, secret.DirPerms); err != nil { + err := cli.fs.MkdirAll(stateDir, secret.DirPerms) + if err != nil { secret.Debug("Failed to create state directory", "error", err) return fmt.Errorf("failed to create state directory: %w", err) @@ -55,100 +123,55 @@ func (cli *Instance) Init(cmd *cobra.Command) error { } // Prompt for mnemonic - var mnemonicStr string - - if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" { - secret.Debug("Using mnemonic from environment variable") - mnemonicStr = envMnemonic - } else { - secret.Debug("Prompting user for mnemonic phrase") - // Read mnemonic securely without echo - mnemonicBuffer, err := secret.ReadPassphrase("Enter your BIP39 mnemonic phrase: ") - if err != nil { - secret.Debug("Failed to read mnemonic from stdin", "error", err) - - return fmt.Errorf("failed to read mnemonic: %w", err) - } - defer mnemonicBuffer.Destroy() - - mnemonicStr = mnemonicBuffer.String() - fmt.Fprintln(os.Stderr) // Add newline after hidden input + mnemonicStr, cleanupMnemonic, err := promptMnemonic() + if err != nil { + return err } + defer cleanupMnemonic() if mnemonicStr == "" { secret.Debug("Empty mnemonic provided") - return fmt.Errorf("mnemonic cannot be empty") + return errMnemonicEmpty } // Validate the mnemonic using BIP39 - secret.DebugWith("Validating BIP39 mnemonic", slog.Int("word_count", len(strings.Fields(mnemonicStr)))) + secret.DebugWith("Validating BIP39 mnemonic", + slog.Int("word_count", len(strings.Fields(mnemonicStr)))) + if !bip39.IsMnemonicValid(mnemonicStr) { secret.Debug("Invalid BIP39 mnemonic provided") - return fmt.Errorf("invalid BIP39 mnemonic phrase\nRun 'secret generate mnemonic' to create a valid mnemonic") + return fmt.Errorf( + "%w\nRun 'secret generate mnemonic' to create a valid mnemonic", + errInvalidMnemonicPhrase) } // Set mnemonic in environment for CreateVault to use - originalMnemonic := os.Getenv(secret.EnvMnemonic) - _ = os.Setenv(secret.EnvMnemonic, mnemonicStr) - defer func() { - if originalMnemonic != "" { - _ = os.Setenv(secret.EnvMnemonic, originalMnemonic) - } else { - _ = os.Unsetenv(secret.EnvMnemonic) - } - }() + restoreMnemonicEnv := setMnemonicEnv(mnemonicStr) + defer restoreMnemonicEnv() - // Create the default vault - it will handle key derivation internally - secret.Debug("Creating default vault") - vlt, err := vault.CreateVault(cli.fs, cli.stateDir, "default") + // Create the default vault and derive its long-term key + vlt, ltIdentity, err := cli.setupDefaultVault(stateDir, mnemonicStr) if err != nil { - secret.Debug("Failed to create default vault", "error", err) - - return fmt.Errorf("failed to create default vault: %w", err) + return err } - // Get the vault metadata to retrieve the derivation index - vaultDir := filepath.Join(stateDir, "vaults.d", "default") - metadata, err := vault.LoadVaultMetadata(cli.fs, vaultDir) - if err != nil { - secret.Debug("Failed to load vault metadata", "error", err) - - return fmt.Errorf("failed to load vault metadata: %w", err) - } - - // Derive the long-term key using the same index that CreateVault used - ltIdentity, err := agehd.DeriveIdentity(mnemonicStr, metadata.DerivationIndex) - if err != nil { - secret.Debug("Failed to derive long-term key", "error", err) - - return fmt.Errorf("failed to derive long-term key from mnemonic: %w", err) - } ltPubKey := ltIdentity.Recipient().String() // Unlock the vault with the derived long-term key vlt.Unlock(ltIdentity) // Prompt for passphrase for unlocker - var passphraseBuffer *memguard.LockedBuffer - if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" { - secret.Debug("Using unlock passphrase from environment variable") - passphraseBuffer = memguard.NewBufferFromBytes([]byte(envPassphrase)) - } else { - secret.Debug("Prompting user for unlock passphrase") - // Use secure passphrase input with confirmation - passphraseBuffer, err = readSecurePassphrase("Enter passphrase for unlocker: ") - if err != nil { - secret.Debug("Failed to read unlock passphrase", "error", err) - - return fmt.Errorf("failed to read passphrase: %w", err) - } + passphraseBuffer, err := resolvePassphrase() + if err != nil { + return err } defer passphraseBuffer.Destroy() // Create passphrase-protected unlocker secret.Debug("Creating passphrase-protected unlocker") + passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer) if err != nil { secret.Debug("Failed to create unlocker", "error", err) @@ -194,7 +217,7 @@ func readSecurePassphrase(prompt string) (*memguard.LockedBuffer, error) { passphraseBuffer1.Destroy() passphraseBuffer2.Destroy() - return nil, fmt.Errorf("passphrases do not match") + return nil, errPassphraseMismatch } // Clean up the second buffer, we'll return the first diff --git a/internal/cli/integration_test.go b/internal/cli/integration_test.go index 814ef25..adf4a6c 100644 --- a/internal/cli/integration_test.go +++ b/internal/cli/integration_test.go @@ -2,11 +2,14 @@ package cli_test import ( + "context" "encoding/json" + "errors" "fmt" "os" "os/exec" "path/filepath" + "slices" "strings" "testing" "time" @@ -20,9 +23,28 @@ import ( const ( // testMnemonic is a standard BIP39 mnemonic used for testing + //nolint:dupword // BIP39 test mnemonic intentionally repeats a word testMnemonic = "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about" ) +// errEmptyValue indicates a concurrent reader received an empty secret value. +var errEmptyValue = errors.New("got empty value") + +// runSecret runs the secret CLI in-process. +func runSecret(args ...string) (string, error) { + return cli.ExecuteCommandInProcess(args, "", nil) +} + +// runSecretWithEnv runs the secret CLI in-process with environment overrides. +func runSecretWithEnv(env map[string]string, args ...string) (string, error) { + return cli.ExecuteCommandInProcess(args, "", env) +} + +// runSecretWithStdin runs the secret CLI in-process with stdin and env. +func runSecretWithStdin(stdin string, env map[string]string, args ...string) (string, error) { + return cli.ExecuteCommandInProcess(args, stdin, env) +} + // TestMain runs before all tests and ensures the binary is built func TestMain(m *testing.M) { // Get the current working directory @@ -36,8 +58,9 @@ func TestMain(m *testing.M) { projectRoot := filepath.Join(wd, "..", "..") // Build the binary - cmd := exec.Command("go", "build", "-o", "secret", "./cmd/secret") + cmd := exec.CommandContext(context.Background(), "go", "build", "-o", "secret", "./cmd/secret") cmd.Dir = projectRoot + output, err := cmd.CombinedOutput() if err != nil { fmt.Fprintf(os.Stderr, "Failed to build secret binary: %v\nOutput: %s\n", err, output) @@ -73,30 +96,8 @@ func TestSecretManagerIntegration(t *testing.T) { // Set environment variables for the test t.Setenv("SB_SECRET_STATE_DIR", tempDir) - // Find the secret binary path (needed for tests that still use exec.Command) - wd, err := os.Getwd() - require.NoError(t, err, "should get working directory") - projectRoot := filepath.Join(wd, "..", "..") - secretPath := filepath.Join(projectRoot, "secret") - - // Helper function to run the secret command - runSecret := func(args ...string) (string, error) { - return cli.ExecuteCommandInProcess(args, "", nil) - } - - // Helper function to run secret with environment variables - runSecretWithEnv := func(env map[string]string, args ...string) (string, error) { - return cli.ExecuteCommandInProcess(args, "", env) - } - - // Helper function to run secret with stdin - runSecretWithStdin := func(stdin string, env map[string]string, args ...string) (string, error) { - return cli.ExecuteCommandInProcess(args, stdin, env) - } - - // Declare runSecret to avoid unused variable error - will be used in later tests - _ = runSecret - _ = runSecretWithStdin + // Find the secret binary path (needed for tests that still exec the binary) + secretPath := secretBinaryPath(t) // Test 1: Initialize secret manager // Command: secret init @@ -321,11 +322,23 @@ func TestSecretManagerIntegration(t *testing.T) { // Helper functions for each test section +// secretBinaryPath returns the path of the secret binary built by TestMain. +func secretBinaryPath(t *testing.T) string { + t.Helper() + + wd, err := os.Getwd() + require.NoError(t, err, "should get working directory") + + return filepath.Join(wd, "..", "..", "secret") +} + func test01Initialize(t *testing.T, tempDir, testMnemonic, testPassphrase string, runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Run init with environment variables to avoid prompts output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, - "SB_UNLOCK_PASSPHRASE": testPassphrase, + secret.EnvMnemonic: testMnemonic, + secret.EnvUnlockPassphrase: testPassphrase, }, "init") require.NoError(t, err, "init should succeed") @@ -340,8 +353,9 @@ func test01Initialize(t *testing.T, tempDir, testMnemonic, testPassphrase string // Check currentvault file contains the vault name currentVaultFile := filepath.Join(tempDir, "currentvault") - targetBytes, err := os.ReadFile(currentVaultFile) + targetBytes, err := os.ReadFile(filepath.Clean(currentVaultFile)) require.NoError(t, err, "should be able to read currentvault file") + target := string(targetBytes) assert.Equal(t, "default", target, "currentvault should contain vault name") @@ -383,14 +397,15 @@ func test01Initialize(t *testing.T, tempDir, testMnemonic, testPassphrase string metadataBytes := readFile(t, vaultMetadata) t.Logf("Vault metadata raw content: %s", string(metadataBytes)) - var metadata map[string]interface{} + var metadata map[string]any + err = json.Unmarshal(metadataBytes, &metadata) require.NoError(t, err, "vault metadata should be valid JSON") t.Logf("Parsed metadata: %+v", metadata) // Verify metadata fields - assert.Equal(t, float64(0), metadata["derivationIndex"], "first vault should have index 0") + assert.InDelta(t, float64(0), metadata["derivationIndex"], 0, "first vault should have index 0") assert.Contains(t, metadata, "publicKeyHash", "should contain public key hash") assert.Contains(t, metadata, "createdAt", "should contain creation timestamp") @@ -400,6 +415,8 @@ func test01Initialize(t *testing.T, tempDir, testMnemonic, testPassphrase string } func test02ListVaults(t *testing.T, runSecret func(...string) (string, error)) { + t.Helper() + // List vaults output, err := runSecret("vault", "list") require.NoError(t, err, "vault list should succeed") @@ -416,7 +433,8 @@ func test02ListVaults(t *testing.T, runSecret func(...string) (string, error)) { t.Logf("JSON output length: %d", len(jsonOutput)) // Parse JSON output - var response map[string]interface{} + var response map[string]any + err = json.Unmarshal([]byte(jsonOutput), &response) require.NoError(t, err, "JSON output should be valid") @@ -429,7 +447,7 @@ func test02ListVaults(t *testing.T, runSecret func(...string) (string, error)) { vaultsRaw, ok := response["vaults"] require.True(t, ok, "response should contain vaults key") - vaults, ok := vaultsRaw.([]interface{}) + vaults, ok := vaultsRaw.([]any) require.True(t, ok, "vaults should be an array") // Verify we have at least one vault @@ -437,23 +455,33 @@ func test02ListVaults(t *testing.T, runSecret func(...string) (string, error)) { // Find default vault in the list foundDefault := false + for _, v := range vaults { vaultName, ok := v.(string) require.True(t, ok, "vault should be a string") + if vaultName == "default" { foundDefault = true + break } } + require.True(t, foundDefault, "default vault should exist in vaults list") } func test03CreateVault(t *testing.T, tempDir string, runSecret func(...string) (string, error)) { - // Set environment variables for vault creation - _ = os.Setenv("SB_SECRET_MNEMONIC", testMnemonic) - _ = os.Setenv("SB_UNLOCK_PASSPHRASE", "test-passphrase") - defer func() { _ = os.Unsetenv("SB_SECRET_MNEMONIC") }() - defer func() { _ = os.Unsetenv("SB_UNLOCK_PASSPHRASE") }() + t.Helper() + + // Set environment variables for vault creation; unset them again before + // returning so later test steps run without ambient credentials + t.Setenv(secret.EnvMnemonic, testMnemonic) + t.Setenv(secret.EnvUnlockPassphrase, "test-passphrase") + + defer func() { + _ = os.Unsetenv(secret.EnvMnemonic) + _ = os.Unsetenv(secret.EnvUnlockPassphrase) + }() // Create work vault output, err := runSecret("vault", "create", "work") @@ -466,8 +494,9 @@ func test03CreateVault(t *testing.T, tempDir string, runSecret func(...string) ( // Check currentvault file was updated currentVaultFile := filepath.Join(tempDir, "currentvault") - targetBytes, err := os.ReadFile(currentVaultFile) + targetBytes, err := os.ReadFile(filepath.Clean(currentVaultFile)) require.NoError(t, err, "should be able to read currentvault file") + target := string(targetBytes) assert.Equal(t, "work", target, "currentvault should contain vault name") @@ -491,10 +520,12 @@ func test03CreateVault(t *testing.T, tempDir string, runSecret func(...string) ( //nolint:unused // TODO: re-enable when vault import is implemented func test04ImportMnemonic(t *testing.T, tempDir, testMnemonic, testPassphrase string, runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Import mnemonic into work vault output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, - "SB_UNLOCK_PASSPHRASE": testPassphrase, + secret.EnvMnemonic: testMnemonic, + secret.EnvUnlockPassphrase: testPassphrase, }, "vault", "import", "work") require.NoError(t, err, "vault import should succeed") @@ -527,7 +558,9 @@ func test04ImportMnemonic(t *testing.T, tempDir, testMnemonic, testPassphrase st verifyFileExists(t, vaultMetadata) metadataBytes := readFile(t, vaultMetadata) - var metadata map[string]interface{} + + var metadata map[string]any + err = json.Unmarshal(metadataBytes, &metadata) require.NoError(t, err, "vault metadata should be valid JSON") @@ -544,6 +577,8 @@ func test04ImportMnemonic(t *testing.T, tempDir, testMnemonic, testPassphrase st } func test05AddSecret(t *testing.T, tempDir, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { + t.Helper() + // Switch back to default vault which has derivation index 0 // matching our mnemonic environment variable _, err := runSecret("vault", "select", "default") @@ -552,7 +587,7 @@ func test05AddSecret(t *testing.T, tempDir, testMnemonic string, runSecret func( // Add a secret with environment variables set secretValue := "password123" output, err := runSecretWithStdin(secretValue, map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "database/password") require.NoError(t, err, "add secret should succeed: %s", output) @@ -601,30 +636,34 @@ func test05AddSecret(t *testing.T, tempDir, testMnemonic string, runSecret func( verifyFileExists(t, currentLink) // Verify current file contains the version name - targetBytes, err := os.ReadFile(currentLink) + targetBytes, err := os.ReadFile(filepath.Clean(currentLink)) require.NoError(t, err, "should read current file") + target := string(targetBytes) assert.Equal(t, versionName, target, "current file should contain version name") // Verify we can retrieve the secret getOutput, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "database/password") if err != nil { t.Logf("Get secret failed. Output: %s", getOutput) } + require.NoError(t, err, "get secret should succeed") assert.Equal(t, secretValue, strings.TrimSpace(getOutput), "retrieved value should match") } func test06GetSecret(t *testing.T, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") // Get the secret that was added in test 05 output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "database/password") t.Logf("Get secret output: %q (length=%d)", output, len(output)) @@ -635,11 +674,13 @@ func test06GetSecret(t *testing.T, testMnemonic string, runSecret func(...string // Test that without mnemonic, we get an error output, err = runSecret("get", "database/password") - assert.Error(t, err, "get should fail without unlock method") + require.Error(t, err, "get should fail without unlock method") assert.Contains(t, output, "failed to unlock vault", "should indicate unlock failure") } func test07AddSecretVersion(t *testing.T, tempDir, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") @@ -647,7 +688,7 @@ func test07AddSecretVersion(t *testing.T, tempDir, testMnemonic string, runSecre // Add new version of existing secret newSecretValue := "newpassword456" output, err := runSecretWithStdin(newSecretValue, map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "database/password", "--force") require.NoError(t, err, "add secret with --force should succeed: %s", output) @@ -664,6 +705,7 @@ func test07AddSecretVersion(t *testing.T, tempDir, testMnemonic string, runSecre // Find which version is newer var oldVersion, newVersion string + for _, entry := range entries { versionName := entry.Name() if strings.HasSuffix(versionName, ".001") { @@ -688,27 +730,30 @@ func test07AddSecretVersion(t *testing.T, tempDir, testMnemonic string, runSecre // Check current file points to new version currentLink := filepath.Join(secretDir, "current") - targetBytes, err := os.ReadFile(currentLink) + targetBytes, err := os.ReadFile(filepath.Clean(currentLink)) require.NoError(t, err, "should read current file") + target := string(targetBytes) assert.Equal(t, newVersion, target, "current file should contain version name") // Verify we get the new value when retrieving the secret getOutput, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "database/password") require.NoError(t, err, "get secret should succeed") assert.Equal(t, newSecretValue, strings.TrimSpace(getOutput), "should return new secret value") } func test08ListVersions(t *testing.T, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") // List versions output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "version", "list", "database/password") require.NoError(t, err, "version list should succeed") @@ -724,13 +769,16 @@ func test08ListVersions(t *testing.T, testMnemonic string, runSecret func(...str // The newer version should be marked as current lines := strings.Split(output, "\n") + var foundCurrent bool + var foundExpired bool for _, line := range lines { if strings.Contains(line, ".002") && strings.Contains(line, "current") { foundCurrent = true } + if strings.Contains(line, ".001") && strings.Contains(line, "expired") { foundExpired = true } @@ -741,6 +789,8 @@ func test08ListVersions(t *testing.T, testMnemonic string, runSecret func(...str } func test09GetSpecificVersion(t *testing.T, tempDir, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") @@ -753,17 +803,20 @@ func test09GetSpecificVersion(t *testing.T, tempDir, testMnemonic string, runSec require.NoError(t, err, "should read versions directory") var version001 string + for _, entry := range entries { if strings.HasSuffix(entry.Name(), ".001") { version001 = entry.Name() + break } } + require.NotEmpty(t, version001, "should find version .001") // Get the specific old version output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "--version", version001, "database/password") require.NoError(t, err, "get specific version should succeed") @@ -771,7 +824,7 @@ func test09GetSpecificVersion(t *testing.T, tempDir, testMnemonic string, runSec // Verify that getting without --version returns the new value output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "database/password") require.NoError(t, err, "get current version should succeed") @@ -779,6 +832,8 @@ func test09GetSpecificVersion(t *testing.T, tempDir, testMnemonic string, runSec } func test10PromoteVersion(t *testing.T, tempDir, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") @@ -791,6 +846,7 @@ func test10PromoteVersion(t *testing.T, tempDir, testMnemonic string, runSecret require.NoError(t, err, "should read versions directory") var version001, version002 string + for _, entry := range entries { if strings.HasSuffix(entry.Name(), ".001") { version001 = entry.Name() @@ -798,19 +854,21 @@ func test10PromoteVersion(t *testing.T, tempDir, testMnemonic string, runSecret version002 = entry.Name() } } + require.NotEmpty(t, version001, "should find version .001") require.NotEmpty(t, version002, "should find version .002") // Before promotion, current should point to .002 (from test 07) currentLink := filepath.Join(defaultVaultDir, "secrets.d", "database%password", "current") - targetBytes, err := os.ReadFile(currentLink) + targetBytes, err := os.ReadFile(filepath.Clean(currentLink)) require.NoError(t, err, "should read current file") + target := string(targetBytes) assert.Equal(t, version002, target, "current should initially point to .002") // Promote the old version output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "version", "promote", "database/password", version001) require.NoError(t, err, "version promote should succeed") @@ -818,14 +876,15 @@ func test10PromoteVersion(t *testing.T, tempDir, testMnemonic string, runSecret assert.Contains(t, output, version001, "should mention the promoted version") // Verify current file was updated - newTargetBytes, err := os.ReadFile(currentLink) + newTargetBytes, err := os.ReadFile(filepath.Clean(currentLink)) require.NoError(t, err, "should read current file after promotion") + newTarget := string(newTargetBytes) assert.Equal(t, version001, newTarget, "current file should now point to .001") // Verify we now get the old value when retrieving the secret getOutput, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "database/password") require.NoError(t, err, "get secret should succeed") assert.Equal(t, "password123", strings.TrimSpace(getOutput), "should return original secret value after promotion") @@ -842,14 +901,16 @@ func test10PromoteVersion(t *testing.T, tempDir, testMnemonic string, runSecret } func test11ListSecrets(t *testing.T, testMnemonic string, runSecret func(...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") // Add a couple more secrets to make the list more interesting for _, secretName := range []string{"api/key", "config/database.yaml"} { - _, err := runSecretWithStdin(fmt.Sprintf("test-value-%s", secretName), map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + _, err := runSecretWithStdin("test-value-"+secretName, map[string]string{ + secret.EnvMnemonic: testMnemonic, }, "add", secretName) require.NoError(t, err, "add %s should succeed", secretName) } @@ -914,6 +975,8 @@ func test11ListSecrets(t *testing.T, testMnemonic string, runSecret func(...stri } func test11bListSecretsQuiet(t *testing.T, testMnemonic string, runSecret func(...string) (string, error)) { + t.Helper() + // Test quiet output quietOutput, err := runSecret("list", "-q") require.NoError(t, err, "secret list -q should succeed") @@ -965,22 +1028,18 @@ func test11bListSecretsQuiet(t *testing.T, testMnemonic string, runSecret func(. require.NoError(t, err, "secret list --json -q should succeed") // Should be valid JSON, not quiet output - var jsonResponse map[string]interface{} + var jsonResponse map[string]any + err = json.Unmarshal([]byte(jsonQuietOutput), &jsonResponse) - assert.NoError(t, err, "output should be valid JSON when both flags are used") + require.NoError(t, err, "output should be valid JSON when both flags are used") // Test using quiet output in command substitution would work like: // secret get $(secret list -q | head -1) // We'll simulate this by getting the first secret name firstSecret := lines[0] - // Need to create a runSecretWithEnv to provide mnemonic for get operation - runSecretWithEnv := func(env map[string]string, args ...string) (string, error) { - return cli.ExecuteCommandInProcess(args, "", env) - } - getOutput, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", firstSecret) require.NoError(t, err, "get with secret name from quiet output should succeed") @@ -989,11 +1048,9 @@ func test11bListSecretsQuiet(t *testing.T, testMnemonic string, runSecret func(. } func test12SecretNameFormats(t *testing.T, tempDir, testMnemonic string, runSecretWithEnv func(map[string]string, ...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { - // Make sure we're in default vault - runSecret := func(args ...string) (string, error) { - return cli.ExecuteCommandInProcess(args, "", nil) - } + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") @@ -1020,7 +1077,7 @@ func test12SecretNameFormats(t *testing.T, tempDir, testMnemonic string, runSecr for _, tc := range testCases { t.Run(tc.secretName, func(t *testing.T) { output, err := runSecretWithStdin(tc.value, map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", tc.secretName) require.NoError(t, err, "add %s should succeed: %s", tc.secretName, output) @@ -1038,7 +1095,7 @@ func test12SecretNameFormats(t *testing.T, tempDir, testMnemonic string, runSecr // Verify we can retrieve the secret getOutput, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", tc.secretName) require.NoError(t, err, "get %s should succeed", tc.secretName) assert.Equal(t, tc.value, strings.TrimSpace(getOutput), "should return correct value for %s", tc.secretName) @@ -1046,6 +1103,13 @@ func test12SecretNameFormats(t *testing.T, tempDir, testMnemonic string, runSecr } // Test invalid secret names + testInvalidSecretNames(t, testMnemonic, runSecretWithStdin) +} + +// testInvalidSecretNames verifies how invalid secret names are handled. +func testInvalidSecretNames(t *testing.T, testMnemonic string, runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { + t.Helper() + invalidNames := []string{ "", // empty "with space", // spaces not allowed @@ -1062,28 +1126,24 @@ func test12SecretNameFormats(t *testing.T, tempDir, testMnemonic string, runSecr // Replace slashes in test name to avoid issues testName := strings.ReplaceAll(invalidName, "/", "_slash_") testName = strings.ReplaceAll(testName, " ", "_space_") + if testName == "" { testName = "empty" } t.Run("invalid_"+testName, func(t *testing.T) { output, err := runSecretWithStdin("test-value", map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", invalidName) // Some of these might not be invalid after all (e.g., leading/trailing slashes might be stripped, .hidden might be allowed) // For now, just check the ones we know should definitely fail definitelyInvalid := []string{"", "with space", "with@symbol", "with#hash", "with$dollar"} - shouldFail := false - for _, invalid := range definitelyInvalid { - if invalidName == invalid { - shouldFail = true - break - } - } + shouldFail := slices.Contains(definitelyInvalid, invalidName) if shouldFail { - assert.Error(t, err, "add '%s' should fail", invalidName) + require.Error(t, err, "add '%s' should fail", invalidName) + if err != nil { assert.Contains(t, output, "invalid secret name", "should indicate invalid name for '%s'", invalidName) } @@ -1097,9 +1157,11 @@ func test12SecretNameFormats(t *testing.T, tempDir, testMnemonic string, runSecr } func test12bMoveSecret(t *testing.T, testMnemonic string, runSecret func(...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { + t.Helper() + // First, create a secret to move _, err := runSecretWithStdin("original-value", map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "test/original") require.NoError(t, err, "add test/original should succeed") @@ -1108,27 +1170,22 @@ func test12bMoveSecret(t *testing.T, testMnemonic string, runSecret func(...stri require.NoError(t, err, "move should succeed") assert.Contains(t, output, "Moved secret 'test/original' to 'test/renamed'", "should show move confirmation") - // Need to create a runSecretWithEnv for get operations - runSecretWithEnv := func(env map[string]string, args ...string) (string, error) { - return cli.ExecuteCommandInProcess(args, "", env) - } - // Verify original doesn't exist _, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "test/original") - assert.Error(t, err, "get original should fail after move") + require.Error(t, err, "get original should fail after move") // Verify new location exists and has correct value getOutput, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "test/renamed") require.NoError(t, err, "get renamed should succeed") assert.Equal(t, "original-value", getOutput, "renamed secret should have original value") // Test mv alias _, err = runSecretWithStdin("another-value", map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "test/another") require.NoError(t, err, "add test/another should succeed") @@ -1138,7 +1195,7 @@ func test12bMoveSecret(t *testing.T, testMnemonic string, runSecret func(...stri // Test rename alias _, err = runSecretWithStdin("rename-test-value", map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "test/rename-me") require.NoError(t, err, "add test/rename-me should succeed") @@ -1149,30 +1206,32 @@ func test12bMoveSecret(t *testing.T, testMnemonic string, runSecret func(...stri // Test error cases // Try to move non-existent secret output, err = runSecret("move", "test/nonexistent", "test/destination") - assert.Error(t, err, "move non-existent should fail") + require.Error(t, err, "move non-existent should fail") assert.Contains(t, output, "not found", "should indicate source not found") // Try to move to existing destination _, err = runSecretWithStdin("dest-value", map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "test/existing-dest") require.NoError(t, err, "add test/existing-dest should succeed") output, err = runSecret("move", "test/renamed", "test/existing-dest") - assert.Error(t, err, "move to existing destination should fail") + require.Error(t, err, "move to existing destination should fail") assert.Contains(t, output, "already exists", "should indicate destination exists") // Verify the source wasn't removed since move failed getOutput, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "test/renamed") require.NoError(t, err, "get source should still work after failed move") assert.Equal(t, "original-value", getOutput, "source should still have original value") } func test12cCrossVaultMove(t *testing.T, testMnemonic string, runSecretWithEnv func(map[string]string, ...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { + t.Helper() + env := map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, } // Create a test secret in the work vault @@ -1206,7 +1265,7 @@ func test12cCrossVaultMove(t *testing.T, testMnemonic string, runSecretWithEnv f require.NoError(t, err, "select work vault should succeed") _, err = runSecretWithEnv(env, "get", "cross/move/test") - assert.Error(t, err, "get from work vault should fail after move") + require.Error(t, err, "get from work vault should fail after move") // Test cross-vault move with rename _, err = runSecretWithStdin("rename-test", env, "add", "rename/source") @@ -1236,7 +1295,7 @@ func test12cCrossVaultMove(t *testing.T, testMnemonic string, runSecretWithEnv f // Move without force should fail output, err = runSecretWithEnv(env, "move", "work:force/test", "default") - assert.Error(t, err, "move without force should fail when dest exists") + require.Error(t, err, "move without force should fail when dest exists") assert.Contains(t, output, "already exists", "should indicate destination exists") // Move with force should succeed @@ -1254,6 +1313,8 @@ func test12cCrossVaultMove(t *testing.T, testMnemonic string, runSecretWithEnv f } func test13UnlockerManagement(t *testing.T, tempDir, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") @@ -1268,12 +1329,13 @@ func test13UnlockerManagement(t *testing.T, tempDir, testMnemonic string, runSec // Create another passphrase unlocker output, err = runSecretWithEnv(map[string]string{ - "SB_UNLOCK_PASSPHRASE": "another-passphrase", - "SB_SECRET_MNEMONIC": testMnemonic, // Need mnemonic to get long-term key + secret.EnvUnlockPassphrase: "another-passphrase", + secret.EnvMnemonic: testMnemonic, // Need mnemonic to get long-term key }, "unlocker", "add", "passphrase") if err != nil { t.Logf("Error adding passphrase unlocker: %v, output: %s", err, output) } + require.NoError(t, err, "add passphrase unlocker should succeed") // List unlockers again - should have 2 now @@ -1283,6 +1345,7 @@ func test13UnlockerManagement(t *testing.T, tempDir, testMnemonic string, runSec // Count passphrase unlockers lines := strings.Split(output, "\n") passphraseCount := 0 + for _, line := range lines { if strings.Contains(line, "passphrase") { passphraseCount++ @@ -1297,11 +1360,12 @@ func test13UnlockerManagement(t *testing.T, tempDir, testMnemonic string, runSec jsonOutput, err := runSecret("unlocker", "list", "--json") require.NoError(t, err, "unlocker list --json should succeed") - var response map[string]interface{} + var response map[string]any + err = json.Unmarshal([]byte(jsonOutput), &response) require.NoError(t, err, "JSON output should be valid") - unlockers, ok := response["unlockers"].([]interface{}) + unlockers, ok := response["unlockers"].([]any) require.True(t, ok, "response should contain unlockers array") // Just verify we have at least 1 unlocker assert.GreaterOrEqual(t, len(unlockers), 1, "should have at least 1 unlocker") @@ -1317,14 +1381,17 @@ func test13UnlockerManagement(t *testing.T, tempDir, testMnemonic string, runSec } func test14SwitchVault(t *testing.T, tempDir string, runSecret func(...string) (string, error)) { + t.Helper() + // Start in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select default should succeed") // Verify current vault is default currentVaultFile := filepath.Join(tempDir, "currentvault") - targetBytes, err := os.ReadFile(currentVaultFile) + targetBytes, err := os.ReadFile(filepath.Clean(currentVaultFile)) require.NoError(t, err, "should read currentvault file") + target := string(targetBytes) assert.Equal(t, "default", target, "currentvault should contain vault name") @@ -1333,8 +1400,9 @@ func test14SwitchVault(t *testing.T, tempDir string, runSecret func(...string) ( require.NoError(t, err, "vault select work should succeed") // Verify current vault is now work - targetBytes, err = os.ReadFile(currentVaultFile) + targetBytes, err = os.ReadFile(filepath.Clean(currentVaultFile)) require.NoError(t, err, "should read currentvault file") + target = string(targetBytes) assert.Equal(t, "work", target, "currentvault should contain vault name") @@ -1344,18 +1412,20 @@ func test14SwitchVault(t *testing.T, tempDir string, runSecret func(...string) ( // Test selecting non-existent vault output, err := runSecret("vault", "select", "nonexistent") - assert.Error(t, err, "selecting non-existent vault should fail") + require.Error(t, err, "selecting non-existent vault should fail") assert.Contains(t, output, "does not exist", "should indicate vault doesn't exist") } func test15VaultIsolation(t *testing.T, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") // Add a unique secret to default vault _, err = runSecretWithStdin("default-vault-secret", map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "default-only/secret", "--force") require.NoError(t, err, "add secret to default vault should succeed") @@ -1365,14 +1435,14 @@ func test15VaultIsolation(t *testing.T, testMnemonic string, runSecret func(...s // Try to get the default-only secret (should fail) output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "default-only/secret") - assert.Error(t, err, "should not be able to get default vault secret from work vault") + require.Error(t, err, "should not be able to get default vault secret from work vault") assert.Contains(t, output, "not found", "should indicate secret not found") // Add a unique secret to work vault _, err = runSecretWithStdin("work-vault-secret", map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "work-only/secret", "--force") require.NoError(t, err, "add secret to work vault should succeed") @@ -1382,27 +1452,29 @@ func test15VaultIsolation(t *testing.T, testMnemonic string, runSecret func(...s // Try to get the work-only secret (should fail) output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "work-only/secret") - assert.Error(t, err, "should not be able to get work vault secret from default vault") + require.Error(t, err, "should not be able to get work vault secret from default vault") assert.Contains(t, output, "not found", "should indicate secret not found") // Verify we can still get the default-only secret output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "default-only/secret") require.NoError(t, err, "get default-only secret should succeed") assert.Equal(t, "default-vault-secret", strings.TrimSpace(output)) } func test16GenerateSecret(t *testing.T, tempDir, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") // Generate a base58 secret output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "generate", "secret", "generated/base58", "--length", "32", "--type", "base58") require.NoError(t, err, "generate secret should succeed") assert.Contains(t, output, "Generated and stored", "should confirm generation") @@ -1410,7 +1482,7 @@ func test16GenerateSecret(t *testing.T, tempDir, testMnemonic string, runSecret // Retrieve and verify the generated secret generatedValue, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "generated/base58") require.NoError(t, err, "get generated secret should succeed") @@ -1424,27 +1496,28 @@ func test16GenerateSecret(t *testing.T, tempDir, testMnemonic string, runSecret // Generate an alphanumeric secret _, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "generate", "secret", "generated/alnum", "--length", "16", "--type", "alnum") require.NoError(t, err, "generate alnum secret should succeed") // Retrieve and verify alnumValue, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "generated/alnum") require.NoError(t, err, "get alnum secret should succeed") + alnumValue = strings.TrimSpace(alnumValue) assert.Len(t, alnumValue, 16, "generated secret should be 16 characters") // Test overwrite protection _, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "generate", "secret", "generated/base58", "--length", "32", "--type", "base58") - assert.Error(t, err, "generate without --force should fail for existing secret") + require.Error(t, err, "generate without --force should fail for existing secret") // Test with --force _, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "generate", "secret", "generated/base58", "--length", "32", "--type", "base58", "--force") require.NoError(t, err, "generate with --force should succeed") @@ -1457,11 +1530,9 @@ func test16GenerateSecret(t *testing.T, tempDir, testMnemonic string, runSecret } func test17ImportFromFile(t *testing.T, tempDir, testMnemonic string, runSecretWithEnv func(map[string]string, ...string) (string, error)) { - // Make sure we're in default vault - runSecret := func(args ...string) (string, error) { - return cli.ExecuteCommandInProcess(args, "", nil) - } + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") @@ -1472,14 +1543,14 @@ func test17ImportFromFile(t *testing.T, tempDir, testMnemonic string, runSecretW // Import the file output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "import", "imported/file", "--source", testFile) require.NoError(t, err, "import should succeed") assert.Contains(t, output, "Successfully imported", "should confirm import") // Retrieve and verify the imported content importedValue, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "imported/file") require.NoError(t, err, "get imported secret should succeed") assert.Equal(t, testContent, strings.TrimSpace(importedValue), "imported content should match") @@ -1491,7 +1562,7 @@ func test17ImportFromFile(t *testing.T, tempDir, testMnemonic string, runSecretW // Import binary file _, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "import", "imported/binary", "--source", binaryFile) require.NoError(t, err, "import binary should succeed") @@ -1500,9 +1571,9 @@ func test17ImportFromFile(t *testing.T, tempDir, testMnemonic string, runSecretW // Test importing non-existent file output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "import", "imported/nonexistent", "--source", "/nonexistent/file") - assert.Error(t, err, "importing non-existent file should fail") + require.Error(t, err, "importing non-existent file should fail") assert.Contains(t, output, "failed", "should indicate failure") // Verify filesystem structure @@ -1512,13 +1583,16 @@ func test17ImportFromFile(t *testing.T, tempDir, testMnemonic string, runSecretW } func test18AgeKeyOperations(t *testing.T, tempDir, secretPath, testMnemonic string, runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault runSecret := func(args ...string) (string, error) { - cmd := exec.Command(secretPath, args...) + //nolint:gosec // G204: test executes the freshly built secret binary + cmd := exec.CommandContext(t.Context(), secretPath, args...) cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } output, err := cmd.CombinedOutput() @@ -1536,7 +1610,7 @@ func test18AgeKeyOperations(t *testing.T, tempDir, secretPath, testMnemonic stri // Encrypt the file using a stored age key encryptedFile := filepath.Join(tempDir, "test-encrypt.txt.age") _, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "encrypt", "encryption/key", "--input", testFile, "--output", encryptedFile) require.NoError(t, err, "encrypt should succeed") // Note: encrypt command doesn't output confirmation message @@ -1547,7 +1621,7 @@ func test18AgeKeyOperations(t *testing.T, tempDir, secretPath, testMnemonic stri // Decrypt the file decryptedFile := filepath.Join(tempDir, "test-decrypt.txt") _, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "decrypt", "encryption/key", "--input", encryptedFile, "--output", decryptedFile) require.NoError(t, err, "decrypt should succeed") // Note: decrypt command doesn't output confirmation message @@ -1558,7 +1632,7 @@ func test18AgeKeyOperations(t *testing.T, tempDir, secretPath, testMnemonic stri // Test encrypting to stdout output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "encrypt", "encryption/key", "--input", testFile) require.NoError(t, err, "encrypt to stdout should succeed") t.Logf("DEBUG: encrypt output: %q", output) @@ -1566,27 +1640,33 @@ func test18AgeKeyOperations(t *testing.T, tempDir, secretPath, testMnemonic stri // Test that the age key was stored as a secret keyValue, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "encryption/key") require.NoError(t, err, "get age key should succeed") assert.Contains(t, keyValue, "AGE-SECRET-KEY", "should be an age secret key") } func test19DisasterRecovery(t *testing.T, tempDir, secretPath, testMnemonic string, runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Skip if age CLI is not available - if _, err := exec.LookPath("age"); err != nil { + _, lookErr := exec.LookPath("age") + if lookErr != nil { t.Skip("age CLI not found in PATH, cannot test manual disaster recovery") + return } // Make sure we're in default vault runSecret := func(args ...string) (string, error) { - cmd := exec.Command(secretPath, args...) + //nolint:gosec // G204: test executes the freshly built secret binary + cmd := exec.CommandContext(t.Context(), secretPath, args...) cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } + output, err := cmd.CombinedOutput() return string(output), err @@ -1597,37 +1677,34 @@ func test19DisasterRecovery(t *testing.T, tempDir, secretPath, testMnemonic stri // Add a test secret testSecretValue := "disaster-recovery-test-secret-value-12345" - cmd := exec.Command(secretPath, "add", "test/disaster-recovery", "--force") + cmd := exec.CommandContext(t.Context(), secretPath, "add", "test/disaster-recovery", "--force") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader(testSecretValue) + output, err := cmd.CombinedOutput() require.NoError(t, err, "add test secret should succeed: %s", string(output)) // Get the vault metadata to know the derivation index defaultVaultDir := filepath.Join(tempDir, "vaults.d", "default") metadataPath := filepath.Join(defaultVaultDir, "vault-metadata.json") - metadataBytes, err := os.ReadFile(metadataPath) + metadataBytes, err := os.ReadFile(filepath.Clean(metadataPath)) require.NoError(t, err, "read vault metadata") var metadata struct { DerivationIndex uint32 `json:"derivationIndex"` } + err = json.Unmarshal(metadataBytes, &metadata) require.NoError(t, err, "parse vault metadata") - // Step 1: Derive the long-term private key from mnemonic using our code - ltIdentity, err := agehd.DeriveIdentity(testMnemonic, metadata.DerivationIndex) - require.NoError(t, err, "derive long-term identity from mnemonic") - - // Write the long-term private key to a file for age CLI - ltPrivKeyPath := filepath.Join(tempDir, "lt-private.key") - err = os.WriteFile(ltPrivKeyPath, []byte(ltIdentity.String()), 0o600) - require.NoError(t, err, "write long-term private key") + // Step 1: Derive the long-term private key from mnemonic and write it + // to a file for the age CLI + ltPrivKeyPath := writeLongTermKey(t, tempDir, testMnemonic, metadata.DerivationIndex) // Find the secret version directory secretDir := filepath.Join(defaultVaultDir, "secrets.d", "test%disaster-recovery") @@ -1642,46 +1719,67 @@ func test19DisasterRecovery(t *testing.T, tempDir, secretPath, testMnemonic stri // Step 2: Use age CLI to decrypt the version private key encryptedPrivKeyPath := filepath.Join(versionDir, "priv.age") versionPrivKeyPath := filepath.Join(tempDir, "version-private.key") - - ageDecryptCmd := exec.Command("age", "-d", "-i", ltPrivKeyPath, "-o", versionPrivKeyPath, encryptedPrivKeyPath) - output, err = ageDecryptCmd.CombinedOutput() - require.NoError(t, err, "age decrypt version private key: %s", string(output)) + ageDecryptFile(t, ltPrivKeyPath, versionPrivKeyPath, encryptedPrivKeyPath) // Step 3: Use age CLI to decrypt the secret value encryptedValuePath := filepath.Join(versionDir, "value.age") decryptedValuePath := filepath.Join(tempDir, "decrypted-value.txt") - - ageDecryptCmd = exec.Command("age", "-d", "-i", versionPrivKeyPath, "-o", decryptedValuePath, encryptedValuePath) - output, err = ageDecryptCmd.CombinedOutput() - require.NoError(t, err, "age decrypt secret value: %s", string(output)) + ageDecryptFile(t, versionPrivKeyPath, decryptedValuePath, encryptedValuePath) // Step 4: Verify the decrypted value matches the original - decryptedValue, err := os.ReadFile(decryptedValuePath) + decryptedValue, err := os.ReadFile(filepath.Clean(decryptedValuePath)) require.NoError(t, err, "read decrypted value") assert.Equal(t, testSecretValue, string(decryptedValue), "manually decrypted value should match original") // Also verify using our tool produces the same result toolOutput, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "test/disaster-recovery") require.NoError(t, err, "get secret using tool") assert.Equal(t, testSecretValue, strings.TrimSpace(toolOutput), "tool output should match original") + // The temporary key/value files live in tempDir and are removed with it. +} - // Clean up temporary files - _ = os.Remove(ltPrivKeyPath) - _ = os.Remove(versionPrivKeyPath) - _ = os.Remove(decryptedValuePath) +// writeLongTermKey derives the long-term identity from the mnemonic and +// writes it to a file for the age CLI, returning the file path. +func writeLongTermKey(t *testing.T, tempDir, mnemonic string, derivationIndex uint32) string { + t.Helper() + + ltIdentity, err := agehd.DeriveIdentity(mnemonic, derivationIndex) + require.NoError(t, err, "derive long-term identity from mnemonic") + + ltPrivKeyPath := filepath.Join(tempDir, "lt-private.key") + err = os.WriteFile(ltPrivKeyPath, []byte(ltIdentity.String()), 0o600) + require.NoError(t, err, "write long-term private key") + + return ltPrivKeyPath +} + +// ageDecryptFile decrypts inPath to outPath using the age CLI with the +// given identity file. +func ageDecryptFile(t *testing.T, identityPath, outPath, inPath string) { + t.Helper() + + //nolint:gosec // G204: age CLI invoked with test-controlled paths + cmd := exec.CommandContext(t.Context(), "age", "-d", "-i", identityPath, "-o", outPath, inPath) + + output, err := cmd.CombinedOutput() + require.NoError(t, err, "age decrypt should succeed: %s", string(output)) } func test20VersionTimestamps(t *testing.T, tempDir, secretPath, testMnemonic string, runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault runSecret := func(args ...string) (string, error) { - cmd := exec.Command(secretPath, args...) + //nolint:gosec // G204: test executes the freshly built secret binary + cmd := exec.CommandContext(t.Context(), secretPath, args...) cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } + output, err := cmd.CombinedOutput() return string(output), err @@ -1691,12 +1789,12 @@ func test20VersionTimestamps(t *testing.T, tempDir, secretPath, testMnemonic str require.NoError(t, err, "vault select should succeed") // Add a test secret - cmd := exec.Command(secretPath, "add", "timestamp/test", "--force") + cmd := exec.CommandContext(t.Context(), secretPath, "add", "timestamp/test", "--force") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader("version1") _, err = cmd.CombinedOutput() @@ -1706,12 +1804,12 @@ func test20VersionTimestamps(t *testing.T, tempDir, secretPath, testMnemonic str time.Sleep(100 * time.Millisecond) // Add second version - cmd = exec.Command(secretPath, "add", "timestamp/test", "--force") + cmd = exec.CommandContext(t.Context(), secretPath, "add", "timestamp/test", "--force") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader("version2") _, err = cmd.CombinedOutput() @@ -1719,7 +1817,7 @@ func test20VersionTimestamps(t *testing.T, tempDir, secretPath, testMnemonic str // List versions and check timestamps output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "version", "list", "timestamp/test") require.NoError(t, err, "version list should succeed") @@ -1734,13 +1832,16 @@ func test20VersionTimestamps(t *testing.T, tempDir, secretPath, testMnemonic str // The newer version should be marked as current lines := strings.Split(output, "\n") + var foundCurrent bool + var foundExpired bool for _, line := range lines { if strings.Contains(line, ".002") && strings.Contains(line, "current") { foundCurrent = true } + if strings.Contains(line, ".001") && strings.Contains(line, "expired") { foundExpired = true } @@ -1751,12 +1852,16 @@ func test20VersionTimestamps(t *testing.T, tempDir, secretPath, testMnemonic str } func test21MaxVersionsPerDay(t *testing.T) { + t.Helper() + // This test would create 999 versions which is too slow for regular testing // Just test that version numbers increment properly t.Log("Test for max versions per day limit - not implemented due to time constraints") } func test22JSONOutput(t *testing.T, runSecret func(...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") @@ -1765,7 +1870,8 @@ func test22JSONOutput(t *testing.T, runSecret func(...string) (string, error)) { output, err := runSecret("vault", "list", "--json") require.NoError(t, err, "vault list --json should succeed") - var vaultListResponse map[string]interface{} + var vaultListResponse map[string]any + err = json.Unmarshal([]byte(output), &vaultListResponse) require.NoError(t, err, "vault list JSON should be valid") assert.Contains(t, vaultListResponse, "vaults", "should have vaults key") @@ -1780,111 +1886,115 @@ func test22JSONOutput(t *testing.T, runSecret func(...string) (string, error)) { } func test23ErrorHandling(t *testing.T, tempDir, secretPath, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Get non-existent secret output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "nonexistent/secret") - assert.Error(t, err, "get non-existent secret should fail") + require.Error(t, err, "get non-existent secret should fail") assert.Contains(t, output, "not found", "should indicate secret not found") // Add secret without mnemonic or unlocker - unsetMnemonic := os.Getenv("SB_SECRET_MNEMONIC") - _ = os.Unsetenv("SB_SECRET_MNEMONIC") - cmd := exec.Command(secretPath, "add", "test/nomnemonic") + unsetMnemonic := os.Getenv(secret.EnvMnemonic) + _ = os.Unsetenv(secret.EnvMnemonic) + cmd := exec.CommandContext(t.Context(), secretPath, "add", "test/nomnemonic") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader("test-value") cmdOutput, err := cmd.CombinedOutput() require.NoError(t, err, "add without mnemonic should succeed - only needs public key: %s", string(cmdOutput)) // Verify we can't get it back without mnemonic - cmd = exec.Command(secretPath, "get", "test/nomnemonic") + cmd = exec.CommandContext(t.Context(), secretPath, "get", "test/nomnemonic") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmdOutput, err = cmd.CombinedOutput() - assert.Error(t, err, "get without mnemonic should fail") + require.Error(t, err, "get without mnemonic should fail") assert.Contains(t, string(cmdOutput), "failed to unlock", "should indicate unlock failure") - t.Setenv("SB_SECRET_MNEMONIC", unsetMnemonic) + t.Setenv(secret.EnvMnemonic, unsetMnemonic) // Invalid secret names (already tested in test 12) // Non-existent vault operations output, err = runSecret("vault", "select", "nonexistent") - assert.Error(t, err, "select non-existent vault should fail") + require.Error(t, err, "select non-existent vault should fail") assert.Contains(t, output, "does not exist", "should indicate vault doesn't exist") // Import to non-existent vault with test passphrase testPassphrase := "test-passphrase-123" // Define testPassphrase locally output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, - "SB_UNLOCK_PASSPHRASE": testPassphrase, + secret.EnvMnemonic: testMnemonic, + secret.EnvUnlockPassphrase: testPassphrase, }, "vault", "import", "nonexistent") - assert.Error(t, err, "import to non-existent vault should fail") + require.Error(t, err, "import to non-existent vault should fail") assert.Contains(t, output, "does not exist", "should indicate vault doesn't exist") // Get specific version that doesn't exist output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "--version", "99999999.999", "database/password") - assert.Error(t, err, "get non-existent version should fail") + require.Error(t, err, "get non-existent version should fail") assert.Contains(t, output, "not found", "should indicate version not found") // Promote non-existent version output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "version", "promote", "database/password", "99999999.999") - assert.Error(t, err, "promote non-existent version should fail") + require.Error(t, err, "promote non-existent version should fail") assert.Contains(t, output, "not found", "should indicate version not found") } func test24EnvironmentVariables(t *testing.T, tempDir, secretPath, testMnemonic, testPassphrase string) { + t.Helper() + // Create a new temporary directory for this test envTestDir := filepath.Join(tempDir, "env-test") err := os.MkdirAll(envTestDir, 0o700) require.NoError(t, err, "create env test dir should succeed") // Test init with both env vars set - _, err = exec.Command(secretPath, "init").Output() - assert.Error(t, err, "init without env vars should fail or prompt") + _, err = exec.CommandContext(t.Context(), secretPath, "init").Output() + require.Error(t, err, "init without env vars should fail or prompt") // Now with env vars - cmd := exec.Command(secretPath, "init") + cmd := exec.CommandContext(t.Context(), secretPath, "init") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", envTestDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("SB_UNLOCK_PASSPHRASE=%s", testPassphrase), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + envTestDir, + secret.EnvMnemonic + "=" + testMnemonic, + secret.EnvUnlockPassphrase + "=" + testPassphrase, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } output, err := cmd.CombinedOutput() require.NoError(t, err, "init with env vars should succeed: %s", string(output)) assert.Contains(t, string(output), "ready to use", "should confirm initialization") // Test that operations work with just mnemonic - cmd = exec.Command(secretPath, "add", "env/test") + cmd = exec.CommandContext(t.Context(), secretPath, "add", "env/test") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", envTestDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + envTestDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader("env-test-value") _, err = cmd.CombinedOutput() require.NoError(t, err, "add with mnemonic env var should succeed") // Verify we can get it back - cmd = exec.Command(secretPath, "get", "env/test") + cmd = exec.CommandContext(t.Context(), secretPath, "get", "env/test") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", envTestDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + envTestDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmdOutput2, err := cmd.CombinedOutput() require.NoError(t, err, "get with mnemonic env var should succeed") @@ -1892,32 +2002,37 @@ func test24EnvironmentVariables(t *testing.T, tempDir, secretPath, testMnemonic, } func test25ConcurrentOperations(t *testing.T, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") // Run multiple concurrent reads const numReaders = 5 - errors := make(chan error, numReaders) + + errCh := make(chan error, numReaders) for i := range numReaders { go func(id int) { output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "database/password") - if err != nil { - errors <- fmt.Errorf("reader %d failed: %v", id, err) - } else if strings.TrimSpace(output) == "" { - errors <- fmt.Errorf("reader %d got empty value", id) - } else { - errors <- nil + + switch { + case err != nil: + errCh <- fmt.Errorf("reader %d failed: %w", id, err) + case strings.TrimSpace(output) == "": + errCh <- fmt.Errorf("%w: reader %d", errEmptyValue, id) + default: + errCh <- nil } }(i) } // Wait for all readers for range numReaders { - err := <-errors + err := <-errCh assert.NoError(t, err, "concurrent read should succeed") } @@ -1925,7 +2040,9 @@ func test25ConcurrentOperations(t *testing.T, testMnemonic string, runSecret fun // to avoid conflicts, but reads should always work } -func test26LargeSecrets(t *testing.T, tempDir, secretPath, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { +func test26LargeSecrets(t *testing.T, _, _, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error), runSecretWithStdin func(string, map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") @@ -1939,13 +2056,13 @@ func test26LargeSecrets(t *testing.T, tempDir, secretPath, testMnemonic string, // Add large secret _, err = runSecretWithStdin(largeValue, map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "large/secret", "--force") require.NoError(t, err, "add large secret should succeed") // Retrieve and verify retrievedValue, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "large/secret") require.NoError(t, err, "get large secret should succeed") assert.Equal(t, largeValue, strings.TrimSpace(retrievedValue), "large secret should match") @@ -1958,31 +2075,34 @@ aWRnaXRzIFB0eSBMdGQwHhcNMTgwMjI4MTQwMzQ5WhcNMjgwMjI2MTQwMzQ5WjBF -----END CERTIFICATE-----` _, err = runSecretWithStdin(certValue, map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "add", "cert/test", "--force") require.NoError(t, err, "add certificate should succeed") // Retrieve and verify certificate retrievedCert, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "cert/test") require.NoError(t, err, "get certificate should succeed") assert.Equal(t, certValue, strings.TrimSpace(retrievedCert), "certificate should match") } func test27SpecialCharacters(t *testing.T, tempDir, secretPath, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // Make sure we're in default vault _, err := runSecret("vault", "select", "default") require.NoError(t, err, "vault select should succeed") // Test with unicode characters + //nolint:gosmopolitan // test intentionally exercises non-ASCII secret values unicodeValue := "Hello 世界! 🔐 Encryption test με UTF-8" - cmd := exec.Command(secretPath, "add", "special/unicode", "--force") + cmd := exec.CommandContext(t.Context(), secretPath, "add", "special/unicode", "--force") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader(unicodeValue) _, err = cmd.CombinedOutput() @@ -1990,19 +2110,19 @@ func test27SpecialCharacters(t *testing.T, tempDir, secretPath, testMnemonic str // Retrieve and verify retrievedUnicode, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "special/unicode") require.NoError(t, err, "get unicode secret should succeed") assert.Equal(t, unicodeValue, strings.TrimSpace(retrievedUnicode), "unicode should match") // Test with special shell characters shellValue := `$PATH; echo "test" && rm -rf / || true` - cmd = exec.Command(secretPath, "add", "special/shell", "--force") + cmd = exec.CommandContext(t.Context(), secretPath, "add", "special/shell", "--force") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader(shellValue) _, err = cmd.CombinedOutput() @@ -2010,19 +2130,19 @@ func test27SpecialCharacters(t *testing.T, tempDir, secretPath, testMnemonic str // Retrieve and verify retrievedShell, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "special/shell") require.NoError(t, err, "get shell chars secret should succeed") assert.Equal(t, shellValue, strings.TrimSpace(retrievedShell), "shell chars should match") // Test with newlines and tabs multilineValue := "Line 1\nLine 2\n\tIndented line 3\nLine 4" - cmd = exec.Command(secretPath, "add", "special/multiline", "--force") + cmd = exec.CommandContext(t.Context(), secretPath, "add", "special/multiline", "--force") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader(multilineValue) _, err = cmd.CombinedOutput() @@ -2030,24 +2150,28 @@ func test27SpecialCharacters(t *testing.T, tempDir, secretPath, testMnemonic str // Retrieve and verify retrievedMultiline, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "special/multiline") require.NoError(t, err, "get multiline secret should succeed") assert.Equal(t, multilineValue, strings.TrimSpace(retrievedMultiline), "multiline should match") } func test28VaultMetadata(t *testing.T, tempDir string) { + t.Helper() + // Check default vault metadata defaultMetadataPath := filepath.Join(tempDir, "vaults.d", "default", "vault-metadata.json") verifyFileExists(t, defaultMetadataPath) metadataBytes := readFile(t, defaultMetadataPath) - var defaultMetadata map[string]interface{} + + var defaultMetadata map[string]any + err := json.Unmarshal(metadataBytes, &defaultMetadata) require.NoError(t, err, "default vault metadata should be valid JSON") // Verify required fields - assert.Equal(t, float64(0), defaultMetadata["derivationIndex"]) + assert.InDelta(t, float64(0), defaultMetadata["derivationIndex"], 0) assert.Contains(t, defaultMetadata, "createdAt") assert.Contains(t, defaultMetadata, "publicKeyHash") assert.Contains(t, defaultMetadata, "mnemonicFamilyHash") @@ -2057,12 +2181,15 @@ func test28VaultMetadata(t *testing.T, tempDir string) { verifyFileExists(t, workMetadataPath) metadataBytes = readFile(t, workMetadataPath) - var workMetadata map[string]interface{} + + var workMetadata map[string]any + err = json.Unmarshal(metadataBytes, &workMetadata) require.NoError(t, err, "work vault metadata should be valid JSON") // Work vault should have different derivation index - workIndex := workMetadata["derivationIndex"].(float64) + workIndex, ok := workMetadata["derivationIndex"].(float64) + require.True(t, ok, "derivationIndex should be a number") assert.NotEqual(t, float64(0), workIndex, "work vault should have non-zero derivation index") // Both vaults created with same mnemonic should have same mnemonicFamilyHash @@ -2071,13 +2198,16 @@ func test28VaultMetadata(t *testing.T, tempDir string) { } func test29SymlinkHandling(t *testing.T, tempDir, secretPath, testMnemonic string) { + t.Helper() + // Test currentvault file currentVaultFile := filepath.Join(tempDir, "currentvault") verifyFileExists(t, currentVaultFile) // Read the file - should contain just the vault name - targetBytes, err := os.ReadFile(currentVaultFile) + targetBytes, err := os.ReadFile(filepath.Clean(currentVaultFile)) require.NoError(t, err, "should read currentvault file") + target := string(targetBytes) assert.NotContains(t, target, "/", "should be bare vault name without path") @@ -2087,55 +2217,74 @@ func test29SymlinkHandling(t *testing.T, tempDir, secretPath, testMnemonic strin currentLink := filepath.Join(secretDir, "current") verifyFileExists(t, currentLink) - targetBytes, err = os.ReadFile(currentLink) + targetBytes, err = os.ReadFile(filepath.Clean(currentLink)) require.NoError(t, err, "should read current version file") + target = string(targetBytes) assert.NotContains(t, target, "/", "should be bare version name without path") // Test that current file updates properly // Add new version - cmd := exec.Command(secretPath, "add", "database/password", "--force") + cmd := exec.CommandContext(t.Context(), secretPath, "add", "database/password", "--force") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader("new-symlink-test-value") _, err = cmd.CombinedOutput() require.NoError(t, err, "add new version should succeed") // Check that current file was updated - newTargetBytes, err := os.ReadFile(currentLink) + newTargetBytes, err := os.ReadFile(filepath.Clean(currentLink)) require.NoError(t, err, "should read updated current file") + newTarget := string(newTargetBytes) assert.NotEqual(t, target, newTarget, "current file should point to new version") assert.NotContains(t, newTarget, "/", "new current file should be bare version name") } -func test30BackupRestore(t *testing.T, tempDir, secretPath, testMnemonic string, runSecretWithEnv func(map[string]string, ...string) (string, error)) { - // Clean up any malformed secret directories from previous test runs - // (e.g., from test 12 when invalid names were accepted) - vaultsDir := filepath.Join(tempDir, "vaults.d") +// removeMalformedSecretDirs removes secret directories that lack a +// versions directory (e.g., left over from invalid-name test cases). +func removeMalformedSecretDirs(vaultsDir string) { vaultDirs, _ := os.ReadDir(vaultsDir) for _, vaultEntry := range vaultDirs { - if vaultEntry.IsDir() { - secretsDir := filepath.Join(vaultsDir, vaultEntry.Name(), "secrets.d") - if secretEntries, err := os.ReadDir(secretsDir); err == nil { - for _, secretEntry := range secretEntries { - if secretEntry.IsDir() { - secretPath := filepath.Join(secretsDir, secretEntry.Name()) - // Check if this is a malformed secret (no versions directory) - versionsPath := filepath.Join(secretPath, "versions") - if _, err := os.Stat(versionsPath); os.IsNotExist(err) { - // This is a malformed secret directory, remove it - _ = os.RemoveAll(secretPath) - } - } - } + if !vaultEntry.IsDir() { + continue + } + + secretsDir := filepath.Join(vaultsDir, vaultEntry.Name(), "secrets.d") + + secretEntries, err := os.ReadDir(secretsDir) + if err != nil { + continue + } + + for _, secretEntry := range secretEntries { + if !secretEntry.IsDir() { + continue + } + + // Check if this is a malformed secret (no versions directory) + secretDirPath := filepath.Join(secretsDir, secretEntry.Name()) + versionsPath := filepath.Join(secretDirPath, "versions") + + _, err := os.Stat(versionsPath) + if os.IsNotExist(err) { + // This is a malformed secret directory, remove it + _ = os.RemoveAll(secretDirPath) } } } +} + +func test30BackupRestore(t *testing.T, tempDir, secretPath, testMnemonic string, runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + + // Clean up any malformed secret directories from previous test runs + // (e.g., from test 12 when invalid names were accepted) + removeMalformedSecretDirs(filepath.Join(tempDir, "vaults.d")) // Create backup directory backupDir := filepath.Join(tempDir, "backup") @@ -2153,12 +2302,12 @@ func test30BackupRestore(t *testing.T, tempDir, secretPath, testMnemonic string, writeFile(t, currentVaultDst, data) // Add more secrets after backup - cmd := exec.Command(secretPath, "add", "post-backup/secret", "--force") + cmd := exec.CommandContext(t.Context(), secretPath, "add", "post-backup/secret", "--force") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader("post-backup-value") _, err = cmd.CombinedOutput() @@ -2166,7 +2315,7 @@ func test30BackupRestore(t *testing.T, tempDir, secretPath, testMnemonic string, // Verify the new secret exists output, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "post-backup/secret") require.NoError(t, err, "get post-backup secret should succeed") assert.Equal(t, "post-backup-value", strings.TrimSpace(output)) @@ -2185,25 +2334,28 @@ func test30BackupRestore(t *testing.T, tempDir, secretPath, testMnemonic string, // Verify original secrets are restored output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "database/password") if err != nil { t.Logf("Error getting restored secret: %v, output: %s", err, output) } + require.NoError(t, err, "get restored secret should succeed") assert.NotEmpty(t, output, "restored secret should have value") // Verify post-backup secret is gone output, err = runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "post-backup/secret") - assert.Error(t, err, "post-backup secret should not exist after restore") + require.Error(t, err, "post-backup secret should not exist after restore") assert.Contains(t, output, "not found", "should indicate secret not found") t.Log("Backup and restore completed successfully") } func test31EnvMnemonicUsesVaultDerivationIndex(t *testing.T, tempDir, secretPath, testMnemonic string, runSecret func(...string) (string, error), runSecretWithEnv func(map[string]string, ...string) (string, error)) { + t.Helper() + // This test demonstrates the bug where GetValue uses hardcoded index 0 // instead of the vault's actual derivation index when using environment mnemonic @@ -2214,17 +2366,21 @@ func test31EnvMnemonicUsesVaultDerivationIndex(t *testing.T, tempDir, secretPath // First, let's verify the derivation indices defaultMetadataPath := filepath.Join(tempDir, "vaults.d", "default", "vault-metadata.json") defaultMetadataBytes := readFile(t, defaultMetadataPath) - var defaultMetadata map[string]interface{} + + var defaultMetadata map[string]any + err := json.Unmarshal(defaultMetadataBytes, &defaultMetadata) require.NoError(t, err, "default vault metadata should be valid JSON") - assert.Equal(t, float64(0), defaultMetadata["derivationIndex"], "default vault should have index 0") + assert.InDelta(t, float64(0), defaultMetadata["derivationIndex"], 0, "default vault should have index 0") workMetadataPath := filepath.Join(tempDir, "vaults.d", "work", "vault-metadata.json") workMetadataBytes := readFile(t, workMetadataPath) - var workMetadata map[string]interface{} + + var workMetadata map[string]any + err = json.Unmarshal(workMetadataBytes, &workMetadata) require.NoError(t, err, "work vault metadata should be valid JSON") - assert.Equal(t, float64(1), workMetadata["derivationIndex"], "work vault should have index 1") + assert.InDelta(t, float64(1), workMetadata["derivationIndex"], 0, "work vault should have index 1") // Switch to work vault _, err = runSecret("vault", "select", "work") @@ -2232,12 +2388,12 @@ func test31EnvMnemonicUsesVaultDerivationIndex(t *testing.T, tempDir, secretPath // Add a secret to work vault using environment mnemonic secretValue := "work-vault-secret" //nolint:gosec // G101: This is test data, not a real credential - cmd := exec.Command(secretPath, "add", "test/derivation") + cmd := exec.CommandContext(t.Context(), secretPath, "add", "test/derivation") cmd.Env = []string{ - fmt.Sprintf("SB_SECRET_STATE_DIR=%s", tempDir), - fmt.Sprintf("SB_SECRET_MNEMONIC=%s", testMnemonic), - fmt.Sprintf("PATH=%s", os.Getenv("PATH")), - fmt.Sprintf("HOME=%s", os.Getenv("HOME")), + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + "PATH=" + os.Getenv("PATH"), + "HOME=" + os.Getenv("HOME"), } cmd.Stdin = strings.NewReader(secretValue) output, err := cmd.CombinedOutput() @@ -2247,7 +2403,7 @@ func test31EnvMnemonicUsesVaultDerivationIndex(t *testing.T, tempDir, secretPath // This is where the bug manifests: GetValue uses hardcoded index 0 // instead of reading the vault metadata to get index 1 getOutput, err := runSecretWithEnv(map[string]string{ - "SB_SECRET_MNEMONIC": testMnemonic, + secret.EnvMnemonic: testMnemonic, }, "get", "test/derivation") // With the bug, this will fail because it tries to decrypt with the wrong key @@ -2257,7 +2413,7 @@ func test31EnvMnemonicUsesVaultDerivationIndex(t *testing.T, tempDir, secretPath t.Logf("Output: %s", getOutput) // This is the expected behavior with the current bug - assert.Error(t, err, "get should fail due to wrong derivation index") + require.Error(t, err, "get should fail due to wrong derivation index") assert.Contains(t, getOutput, "derived public key does not match vault", "should indicate key derivation failure") // Document what should happen when the bug is fixed @@ -2280,6 +2436,7 @@ func test31EnvMnemonicUsesVaultDerivationIndex(t *testing.T, tempDir, secretPath // verifyFileExists checks if a file exists at the given path func verifyFileExists(t *testing.T, path string) { t.Helper() + _, err := os.Stat(path) require.NoError(t, err, "File should exist: %s", path) } @@ -2289,6 +2446,7 @@ func verifyFileExists(t *testing.T, path string) { //nolint:unused // kept for future use func verifyFileNotExists(t *testing.T, path string) { t.Helper() + _, err := os.Stat(path) require.True(t, os.IsNotExist(err), "File should not exist: %s", path) } @@ -2296,7 +2454,8 @@ func verifyFileNotExists(t *testing.T, path string) { // readFile reads and returns the contents of a file func readFile(t *testing.T, path string) []byte { t.Helper() - data, err := os.ReadFile(path) + + data, err := os.ReadFile(filepath.Clean(path)) require.NoError(t, err, "Should be able to read file: %s", path) return data @@ -2305,6 +2464,7 @@ func readFile(t *testing.T, path string) []byte { // writeFile writes data to a file func writeFile(t *testing.T, path string, data []byte) { t.Helper() + err := os.WriteFile(path, data, 0o600) require.NoError(t, err, "Should be able to write file: %s", path) } @@ -2316,7 +2476,7 @@ func copyDir(src, dst string) error { return err } - err = os.MkdirAll(dst, 0o755) + err = os.MkdirAll(dst, 0o750) if err != nil { return err } @@ -2348,17 +2508,19 @@ func copyFile(src, dst string) error { if err != nil { return err } + if srcInfo.IsDir() { // Skip directories, they should be handled by copyDir return nil } - srcData, err := os.ReadFile(src) + srcData, err := os.ReadFile(filepath.Clean(src)) if err != nil { return err } - err = os.WriteFile(dst, srcData, 0o600) + //nolint:gosec // G703: src and dst stay within the test's temp dir + err = os.WriteFile(filepath.Clean(dst), srcData, 0o600) if err != nil { return err } diff --git a/internal/cli/root.go b/internal/cli/root.go index a0aad85..94ddd13 100644 --- a/internal/cli/root.go +++ b/internal/cli/root.go @@ -10,17 +10,21 @@ import ( // Entry is the entry point for the secret CLI application func Entry() { cmd := newRootCmd() - if err := cmd.Execute(); err != nil { + + err := cmd.Execute() + if err != nil { os.Exit(1) } } func newRootCmd() *cobra.Command { secret.Debug("newRootCmd starting") + cmd := &cobra.Command{ Use: "secret", Short: "A simple secrets manager", - Long: `A simple secrets manager to store and retrieve sensitive information securely.`, + Long: `A simple secrets manager to store and retrieve sensitive ` + + `information securely.`, // Ensure usage is shown after errors SilenceUsage: false, SilenceErrors: false, diff --git a/internal/cli/secrets.go b/internal/cli/secrets.go index ee66aac..f62ecf0 100644 --- a/internal/cli/secrets.go +++ b/internal/cli/secrets.go @@ -2,10 +2,12 @@ package cli import ( "encoding/json" + "errors" "fmt" "io" "log" "path/filepath" + "slices" "strings" "git.eeqj.de/sneak/secret/internal/secret" @@ -20,12 +22,36 @@ const ( vaultSecretSeparator = ":" // vaultSecretParts is the number of parts when splitting vault:secret vaultSecretParts = 2 + + // initialBufferSize is the starting size for secret read buffers (4KB) + initialBufferSize = 4 * 1024 + // maxSecretSize is the maximum allowed size of a secret (100MB) + maxSecretSize = 100 * 1024 * 1024 ) +// Sentinel errors for secret operations +var ( + errSecretTooLarge = errors.New("secret too large: exceeds 100MB limit") + errSecretFileTooLarge = errors.New( + "secret file too large: exceeds 100MB limit") + errSecretNotFound = errors.New("not found") + errSecretExistsNoForce = errors.New( + "already exists (use --force to overwrite)") + errVaultDoesNotExist = errors.New("does not exist") + errCrossVaultSourceUnqualified = errors.New( + "source must specify vault (e.g., vault:secret) for cross-vault move") +) + +// bufferInfo tracks a protected buffer and the number of bytes used in it +type bufferInfo struct { + buffer *memguard.LockedBuffer + used int +} + // ParseVaultSecretRef parses a "vault:secret" or just "secret" reference // Returns (vaultName, secretName, isQualified) // If no vault is specified, returns empty vaultName and isQualified=false -func ParseVaultSecretRef(ref string) (vaultName, secretName string, isQualified bool) { +func ParseVaultSecretRef(ref string) (string, string, bool) { parts := strings.SplitN(ref, vaultSecretSeparator, vaultSecretParts) if len(parts) == vaultSecretParts { return parts[0], parts[1], true @@ -42,6 +68,7 @@ func newAddCmd() *cobra.Command { Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { secret.Debug("Add command RunE starting", "secret_name", args[0]) + force, _ := cmd.Flags().GetBool("force") secret.Debug("Got force flag", "force", force) @@ -49,7 +76,9 @@ func newAddCmd() *cobra.Command { if err != nil { return fmt.Errorf("failed to initialize CLI: %w", err) } + cli.cmd = cmd // Set the command for stdin access + secret.Debug("Created CLI instance, calling AddSecret") return cli.AddSecret(args[0], force) @@ -66,6 +95,7 @@ func newGetCmd() *cobra.Command { if err != nil { log.Fatalf("failed to initialize CLI: %v", err) } + cmd := &cobra.Command{ Use: "get ", Short: "Retrieve a secret from the vault", @@ -73,6 +103,7 @@ func newGetCmd() *cobra.Command { ValidArgsFunction: getSecretNamesCompletionFunc(cli.fs, cli.stateDir), RunE: func(cmd *cobra.Command, args []string) error { version, _ := cmd.Flags().GetString("version") + cli, err := NewCLIInstance() if err != nil { return fmt.Errorf("failed to initialize CLI: %w", err) @@ -92,8 +123,9 @@ func newListCmd() *cobra.Command { Use: "list [filter]", Aliases: []string{"ls"}, Short: "List all secrets in the current vault", - Long: `List all secrets in the current vault. Optionally filter by substring match in secret name.`, - Args: cobra.MaximumNArgs(1), + Long: `List all secrets in the current vault. Optionally filter ` + + `by substring match in secret name.`, + Args: cobra.MaximumNArgs(1), RunE: func(cmd *cobra.Command, args []string) error { jsonOutput, _ := cmd.Flags().GetBool("json") quietOutput, _ := cmd.Flags().GetBool("quiet") @@ -122,8 +154,9 @@ func newImportCmd() *cobra.Command { cmd := &cobra.Command{ Use: "import ", Short: "Import a secret from a file", - Long: `Import a secret from a file and store it in the current vault under the given name.`, - Args: cobra.ExactArgs(1), + Long: `Import a secret from a file and store it in the current ` + + `vault under the given name.`, + Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { sourceFile, _ := cmd.Flags().GetString("source") force, _ := cmd.Flags().GetBool("force") @@ -149,12 +182,13 @@ func newRemoveCmd() *cobra.Command { if err != nil { log.Fatalf("failed to initialize CLI: %v", err) } + cmd := &cobra.Command{ Use: "remove ", Aliases: []string{"rm"}, Short: "Remove a secret from the vault", - Long: `Remove a secret and all its versions from the current vault. This action is permanent and ` + - `cannot be undone.`, + Long: `Remove a secret and all its versions from the current ` + + `vault. This action is permanent and cannot be undone.`, Args: cobra.ExactArgs(1), ValidArgsFunction: getSecretNamesCompletionFunc(cli.fs, cli.stateDir), RunE: func(cmd *cobra.Command, args []string) error { @@ -175,6 +209,7 @@ func newMoveCmd() *cobra.Command { if err != nil { log.Fatalf("failed to initialize CLI: %v", err) } + cmd := &cobra.Command{ Use: "move ", Aliases: []string{"mv", "rename"}, @@ -190,13 +225,16 @@ For cross-vault moves: Cross-vault moves copy ALL versions of the secret, preserving history. The source secret is deleted after successful copy.`, - Args: cobra.ExactArgs(2), //nolint:mnd // Command requires exactly 2 arguments: source and destination - ValidArgsFunction: func(cmd *cobra.Command, args []string, toComplete string) ([]string, cobra.ShellCompDirective) { + Args: cobra.ExactArgs(2), //nolint:mnd // source and destination args + ValidArgsFunction: func( + cmd *cobra.Command, args []string, toComplete string, + ) ([]string, cobra.ShellCompDirective) { // Complete vault:secret format return getVaultSecretCompletionFunc(cli.fs, cli.stateDir)(cmd, args, toComplete) }, RunE: func(cmd *cobra.Command, args []string) error { force, _ := cmd.Flags().GetBool("force") + cli, err := NewCLIInstance() if err != nil { return fmt.Errorf("failed to initialize CLI: %w", err) @@ -206,16 +244,20 @@ The source secret is deleted after successful copy.`, }, } - cmd.Flags().BoolP("force", "f", false, "Overwrite if destination secret already exists") + cmd.Flags().BoolP("force", "f", false, + "Overwrite if destination secret already exists") return cmd } // updateBufferSize updates the buffer size based on usage pattern func updateBufferSize(currentSize int, sameSize *int) int { + const ( + doubleAfterBuffers = 2 + growthFactor = 2 + ) + *sameSize++ - const doubleAfterBuffers = 2 - const growthFactor = 2 if *sameSize >= doubleAfterBuffers { *sameSize = 0 @@ -225,40 +267,21 @@ func updateBufferSize(currentSize int, sameSize *int) int { return currentSize } -// AddSecret adds a secret to the current vault -func (cli *Instance) AddSecret(secretName string, force bool) error { - secret.Debug("CLI AddSecret starting", "secret_name", secretName, "force", force) - - // Get current vault - secret.Debug("Getting current vault") - vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) - if err != nil { - return err - } - - secret.Debug("Got current vault", "vault_name", vlt.GetName()) - - // Read secret value directly into protected buffers - secret.Debug("Reading secret value from stdin into protected buffers") - - const initialSize = 4 * 1024 // 4KB initial buffer - const maxSize = 100 * 1024 * 1024 // 100MB max - - type bufferInfo struct { - buffer *memguard.LockedBuffer - used int +// destroyBuffers destroys every buffer in the list +func destroyBuffers(buffers []bufferInfo) { + for _, b := range buffers { + b.buffer.Destroy() } +} +// readSecretFromReader reads all data from reader into protected buffers, +// enforcing the maximum secret size. On failure the accumulated buffers +// are destroyed; on success the caller must destroy them. +func readSecretFromReader(reader io.Reader) ([]bufferInfo, int, error) { var buffers []bufferInfo - defer func() { - for _, b := range buffers { - b.buffer.Destroy() - } - }() - reader := cli.cmd.InOrStdin() totalSize := 0 - currentBufferSize := initialSize + currentBufferSize := initialBufferSize sameSize := 0 for { @@ -273,8 +296,10 @@ func (cli *Instance) AddSecret(secretName string, force bool) error { buffers = append(buffers, bufferInfo{buffer: buffer, used: n}) totalSize += n - if totalSize > maxSize { - return fmt.Errorf("secret too large: exceeds 100MB limit") + if totalSize > maxSecretSize { + destroyBuffers(buffers) + + return nil, 0, errSecretTooLarge } // If we filled the buffer, consider growing for next iteration @@ -283,13 +308,59 @@ func (cli *Instance) AddSecret(secretName string, force bool) error { } } - if err == io.EOF || err == io.ErrUnexpectedEOF { + if err == io.EOF || errors.Is(err, io.ErrUnexpectedEOF) { break } else if err != nil { - return fmt.Errorf("failed to read secret value: %w", err) + destroyBuffers(buffers) + + return nil, 0, err } } + return buffers, totalSize, nil +} + +// combineBuffers copies the used portions of buffers into a single +// protected buffer of totalSize bytes +func combineBuffers(buffers []bufferInfo, totalSize int) *memguard.LockedBuffer { + valueBuffer := memguard.NewBuffer(totalSize) + + offset := 0 + for _, b := range buffers { + copy(valueBuffer.Bytes()[offset:], b.buffer.Bytes()[:b.used]) + offset += b.used + } + + return valueBuffer +} + +// AddSecret adds a secret to the current vault +func (cli *Instance) AddSecret(secretName string, force bool) error { + secret.Debug("CLI AddSecret starting", "secret_name", secretName, "force", force) + + // Get current vault + secret.Debug("Getting current vault") + + vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) + if err != nil { + return err + } + + secret.Debug("Got current vault", "vault_name", vlt.GetName()) + + // Read secret value directly into protected buffers + secret.Debug("Reading secret value from stdin into protected buffers") + + buffers, totalSize, err := readSecretFromReader(cli.cmd.InOrStdin()) + if err != nil { + if errors.Is(err, errSecretTooLarge) { + return err + } + + return fmt.Errorf("failed to read secret value: %w", err) + } + defer destroyBuffers(buffers) + // Check for trailing newline in the last buffer if len(buffers) > 0 && totalSize > 0 { lastBuffer := &buffers[len(buffers)-1] @@ -299,21 +370,19 @@ func (cli *Instance) AddSecret(secretName string, force bool) error { } } - secret.Debug("Read secret value from stdin", "value_length", totalSize, "buffers", len(buffers)) + secret.Debug("Read secret value from stdin", + "value_length", totalSize, "buffers", len(buffers)) // Combine all buffers into a single protected buffer - valueBuffer := memguard.NewBuffer(totalSize) + valueBuffer := combineBuffers(buffers, totalSize) defer valueBuffer.Destroy() - offset := 0 - for _, b := range buffers { - copy(valueBuffer.Bytes()[offset:], b.buffer.Bytes()[:b.used]) - offset += b.used - } - // Add the secret to the vault - secret.Debug("Calling vault.AddSecret", "secret_name", secretName, "value_length", valueBuffer.Size(), "force", force) - if err := vlt.AddSecret(secretName, valueBuffer, force); err != nil { + secret.Debug("Calling vault.AddSecret", "secret_name", secretName, + "value_length", valueBuffer.Size(), "force", force) + + err = vlt.AddSecret(secretName, valueBuffer, force) + if err != nil { secret.Debug("vault.AddSecret failed", "error", err) return err @@ -330,8 +399,11 @@ func (cli *Instance) GetSecret(cmd *cobra.Command, secretName string) error { } // GetSecretWithVersion retrieves and prints a specific version of a secret -func (cli *Instance) GetSecretWithVersion(cmd *cobra.Command, secretName string, version string) error { - secret.Debug("GetSecretWithVersion called", "secretName", secretName, "version", version) +func (cli *Instance) GetSecretWithVersion( + cmd *cobra.Command, secretName string, version string, +) error { + secret.Debug("GetSecretWithVersion called", + "secretName", secretName, "version", version) // Store the command for output cli.cmd = cmd @@ -351,6 +423,7 @@ func (cli *Instance) GetSecretWithVersion(cmd *cobra.Command, secretName string, } else { value, err = vlt.GetSecretVersion(secretName, version) } + if err != nil { secret.Debug("Failed to get secret", "error", err) @@ -361,6 +434,7 @@ func (cli *Instance) GetSecretWithVersion(cmd *cobra.Command, secretName string, // Print the secret value to stdout _, _ = cli.Print(string(value)) + secret.Debug("Printed value to stdout") // Debug: Log what we're actually printing @@ -375,7 +449,9 @@ func (cli *Instance) GetSecretWithVersion(cmd *cobra.Command, secretName string, } // ListSecrets lists all secrets in the current vault -func (cli *Instance) ListSecrets(cmd *cobra.Command, jsonOutput bool, quietOutput bool, filter string) error { +func (cli *Instance) ListSecrets( + cmd *cobra.Command, jsonOutput bool, quietOutput bool, filter string, +) error { // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -390,6 +466,7 @@ func (cli *Instance) ListSecrets(cmd *cobra.Command, jsonOutput bool, quietOutpu // Filter secrets if filter is provided var filteredSecrets []string + if filter != "" { for _, secretName := range secrets { if strings.Contains(secretName, filter) { @@ -400,100 +477,132 @@ func (cli *Instance) ListSecrets(cmd *cobra.Command, jsonOutput bool, quietOutpu filteredSecrets = secrets } - if jsonOutput { //nolint:nestif // Separate JSON and table output formatting logic - // For JSON output, get metadata for each secret - secretsWithMetadata := make([]map[string]interface{}, 0, len(filteredSecrets)) - - for _, secretName := range filteredSecrets { - secretInfo := map[string]interface{}{ - "name": secretName, - } - - // Try to get metadata using GetSecretObject - if secretObj, err := vlt.GetSecretObject(secretName); err == nil { - metadata := secretObj.GetMetadata() - secretInfo["created_at"] = metadata.CreatedAt - secretInfo["updated_at"] = metadata.UpdatedAt - } - - secretsWithMetadata = append(secretsWithMetadata, secretInfo) - } - - output := map[string]interface{}{ - "secrets": secretsWithMetadata, - } - if filter != "" { - output["filter"] = filter - } - - jsonBytes, err := json.MarshalIndent(output, "", " ") - if err != nil { - return fmt.Errorf("failed to marshal JSON: %w", err) - } - - _, _ = fmt.Fprintln(cmd.OutOrStdout(), string(jsonBytes)) - } else if quietOutput { + switch { + case jsonOutput: + return printSecretsJSON(cmd, vlt, filteredSecrets, filter) + case quietOutput: // Quiet output - just secret names for _, secretName := range filteredSecrets { _, _ = fmt.Fprintln(cmd.OutOrStdout(), secretName) } - } else { - // Pretty table output - out := cmd.OutOrStdout() - if len(filteredSecrets) == 0 { - if filter != "" { - _, _ = fmt.Fprintf(out, "No secrets found in vault '%s' matching filter '%s'.\n", vlt.GetName(), filter) - } else { - _, _ = fmt.Fprintln(out, "No secrets found in current vault.") - _, _ = fmt.Fprintln(out, "Run 'secret add ' to create one.") - } - return nil - } - - // Get current vault name for display - if filter != "" { - _, _ = fmt.Fprintf(out, "Secrets in vault '%s' matching '%s':\n\n", vlt.GetName(), filter) - } else { - _, _ = fmt.Fprintf(out, "Secrets in vault '%s':\n\n", vlt.GetName()) - } - - // Calculate the maximum name length for proper column alignment - maxNameLen := len("NAME") // Start with header length - for _, secretName := range filteredSecrets { - if len(secretName) > maxNameLen { - maxNameLen = len(secretName) - } - } - // Add some padding - maxNameLen += 2 - - // Print headers with dynamic width - nameFormat := fmt.Sprintf("%%-%ds", maxNameLen) - _, _ = fmt.Fprintf(out, nameFormat+" %-20s\n", "NAME", "LAST UPDATED") - _, _ = fmt.Fprintf(out, nameFormat+" %-20s\n", strings.Repeat("-", len("NAME")), "------------") - - for _, secretName := range filteredSecrets { - lastUpdated := "unknown" - if secretObj, err := vlt.GetSecretObject(secretName); err == nil { - metadata := secretObj.GetMetadata() - lastUpdated = metadata.UpdatedAt.Format("2006-01-02 15:04") - } - _, _ = fmt.Fprintf(out, nameFormat+" %-20s\n", secretName, lastUpdated) - } - - _, _ = fmt.Fprintf(out, "\nTotal: %d secret(s)", len(filteredSecrets)) - if filter != "" { - _, _ = fmt.Fprintf(out, " (filtered from %d)", len(secrets)) - } - _, _ = fmt.Fprintln(out) + return nil + default: + return printSecretsTable(cmd, vlt, filteredSecrets, filter, len(secrets)) } +} + +// printSecretsJSON prints the filtered secrets with metadata as JSON +func printSecretsJSON( + cmd *cobra.Command, vlt *vault.Vault, filteredSecrets []string, filter string, +) error { + // For JSON output, get metadata for each secret + secretsWithMetadata := make([]map[string]any, 0, len(filteredSecrets)) + + for _, secretName := range filteredSecrets { + secretInfo := map[string]any{ + "name": secretName, + } + + // Try to get metadata using GetSecretObject + secretObj, err := vlt.GetSecretObject(secretName) + if err == nil { + metadata := secretObj.GetMetadata() + secretInfo["created_at"] = metadata.CreatedAt + secretInfo["updated_at"] = metadata.UpdatedAt + } + + secretsWithMetadata = append(secretsWithMetadata, secretInfo) + } + + output := map[string]any{ + "secrets": secretsWithMetadata, + } + if filter != "" { + output["filter"] = filter + } + + jsonBytes, err := json.MarshalIndent(output, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal JSON: %w", err) + } + + _, _ = fmt.Fprintln(cmd.OutOrStdout(), string(jsonBytes)) + + return nil +} + +// printSecretsTable prints the filtered secrets as a formatted table +func printSecretsTable( + cmd *cobra.Command, vlt *vault.Vault, + filteredSecrets []string, filter string, totalCount int, +) error { + // Pretty table output + out := cmd.OutOrStdout() + + if len(filteredSecrets) == 0 { + if filter != "" { + _, _ = fmt.Fprintf(out, + "No secrets found in vault '%s' matching filter '%s'.\n", + vlt.GetName(), filter) + } else { + _, _ = fmt.Fprintln(out, "No secrets found in current vault.") + _, _ = fmt.Fprintln(out, "Run 'secret add ' to create one.") + } + + return nil + } + + // Get current vault name for display + if filter != "" { + _, _ = fmt.Fprintf(out, "Secrets in vault '%s' matching '%s':\n\n", + vlt.GetName(), filter) + } else { + _, _ = fmt.Fprintf(out, "Secrets in vault '%s':\n\n", vlt.GetName()) + } + + // Calculate the maximum name length for proper column alignment + maxNameLen := len("NAME") // Start with header length + for _, secretName := range filteredSecrets { + if len(secretName) > maxNameLen { + maxNameLen = len(secretName) + } + } + // Add some padding + maxNameLen += 2 + + // Print headers with dynamic width + nameFormat := fmt.Sprintf("%%-%ds", maxNameLen) + _, _ = fmt.Fprintf(out, nameFormat+" %-20s\n", "NAME", "LAST UPDATED") + _, _ = fmt.Fprintf(out, nameFormat+" %-20s\n", + strings.Repeat("-", len("NAME")), "------------") + + for _, secretName := range filteredSecrets { + lastUpdated := "unknown" + + secretObj, err := vlt.GetSecretObject(secretName) + if err == nil { + metadata := secretObj.GetMetadata() + lastUpdated = metadata.UpdatedAt.Format("2006-01-02 15:04") + } + + _, _ = fmt.Fprintf(out, nameFormat+" %-20s\n", secretName, lastUpdated) + } + + _, _ = fmt.Fprintf(out, "\nTotal: %d secret(s)", len(filteredSecrets)) + if filter != "" { + _, _ = fmt.Fprintf(out, " (filtered from %d)", totalCount) + } + + _, _ = fmt.Fprintln(out) return nil } // ImportSecret imports a secret from a file -func (cli *Instance) ImportSecret(cmd *cobra.Command, secretName, sourceFile string, force bool) error { +func (cli *Instance) ImportSecret( + cmd *cobra.Command, secretName, sourceFile string, force bool, +) error { // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -506,75 +615,34 @@ func (cli *Instance) ImportSecret(cmd *cobra.Command, secretName, sourceFile str return fmt.Errorf("failed to open file %s: %w", sourceFile, err) } defer func() { - if err := file.Close(); err != nil { - secret.Warn("Failed to close file", "error", err) + closeErr := file.Close() + if closeErr != nil { + secret.Warn("Failed to close file", "error", closeErr) } }() - const initialSize = 4 * 1024 // 4KB initial buffer - const maxSize = 100 * 1024 * 1024 // 100MB max + buffers, totalSize, err := readSecretFromReader(file) + if err != nil { + if errors.Is(err, errSecretTooLarge) { + return errSecretFileTooLarge + } - type bufferInfo struct { - buffer *memguard.LockedBuffer - used int - } - - var buffers []bufferInfo - defer func() { - for _, b := range buffers { - b.buffer.Destroy() - } - }() - - totalSize := 0 - currentBufferSize := initialSize - sameSize := 0 - - for { - // Create a new buffer - buffer := memguard.NewBuffer(currentBufferSize) - n, err := io.ReadFull(file, buffer.Bytes()) - - if n == 0 { - // No data read, destroy the unused buffer - buffer.Destroy() - } else { - buffers = append(buffers, bufferInfo{buffer: buffer, used: n}) - totalSize += n - - if totalSize > maxSize { - return fmt.Errorf("secret file too large: exceeds 100MB limit") - } - - // If we filled the buffer, consider growing for next iteration - if n == currentBufferSize { - currentBufferSize = updateBufferSize(currentBufferSize, &sameSize) - } - } - - if err == io.EOF || err == io.ErrUnexpectedEOF { - break - } else if err != nil { - return fmt.Errorf("failed to read secret from file %s: %w", sourceFile, err) - } + return fmt.Errorf("failed to read secret from file %s: %w", sourceFile, err) } + defer destroyBuffers(buffers) // Combine all buffers into a single protected buffer - valueBuffer := memguard.NewBuffer(totalSize) + valueBuffer := combineBuffers(buffers, totalSize) defer valueBuffer.Destroy() - offset := 0 - for _, b := range buffers { - copy(valueBuffer.Bytes()[offset:], b.buffer.Bytes()[:b.used]) - offset += b.used - } - // Store the secret in the vault - if err := vlt.AddSecret(secretName, valueBuffer, force); err != nil { + err = vlt.AddSecret(secretName, valueBuffer, force) + if err != nil { return err } - cmd.Printf("Successfully imported secret '%s' from file '%s'\n", secretName, sourceFile) + cmd.Printf("Successfully imported secret '%s' from file '%s'\n", + secretName, sourceFile) return nil } @@ -600,29 +668,36 @@ func (cli *Instance) RemoveSecret(cmd *cobra.Command, secretName string, _ bool) if err != nil { return fmt.Errorf("failed to check if secret exists: %w", err) } + if !exists { - return fmt.Errorf("secret '%s' not found", secretName) + return fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound) } // Count versions for information versionsDir := filepath.Join(secretDir, "versions") versionCount := 0 - if entries, err := afero.ReadDir(cli.fs, versionsDir); err == nil { + + entries, err := afero.ReadDir(cli.fs, versionsDir) + if err == nil { versionCount = len(entries) } // Remove the secret directory - if err := cli.fs.RemoveAll(secretDir); err != nil { + err = cli.fs.RemoveAll(secretDir) + if err != nil { return fmt.Errorf("failed to remove secret: %w", err) } - cmd.Printf("Removed secret '%s' (%d version(s) deleted)\n", secretName, versionCount) + cmd.Printf("Removed secret '%s' (%d version(s) deleted)\n", + secretName, versionCount) return nil } // MoveSecret moves or renames a secret (within or across vaults) -func (cli *Instance) MoveSecret(cmd *cobra.Command, source, dest string, force bool) error { +func (cli *Instance) MoveSecret( + cmd *cobra.Command, source, dest string, force bool, +) error { // Parse source and destination srcVaultName, srcSecretName, srcQualified := ParseVaultSecretRef(source) destVaultName, destSecretName, destQualified := ParseVaultSecretRef(dest) @@ -634,25 +709,20 @@ func (cli *Instance) MoveSecret(cmd *cobra.Command, source, dest string, force b // Cross-vault move requires source to be qualified if !srcQualified { - return fmt.Errorf("source must specify vault (e.g., vault:secret) for cross-vault move") + return errCrossVaultSourceUnqualified } // If destination is not qualified (no colon), check if it's a vault name // Format: "work:secret default" means move to vault "default" - // Format: "work:secret default:newname" means move to vault "default" with new name + // Format: "work:secret default:newname" means move to vault "default" + // with a new name if !destQualified { // Check if dest is actually a vault name vaults, err := vault.ListVaults(cli.fs, cli.stateDir) - if err == nil { - for _, v := range vaults { - if v == dest { - // dest is a vault name, use source secret name - destVaultName = dest - destSecretName = srcSecretName - - break - } - } + if err == nil && slices.Contains(vaults, dest) { + // dest is a vault name, use source secret name + destVaultName = dest + destSecretName = srcSecretName } // If destVaultName is still empty, dest is a secret name in source vault @@ -670,7 +740,8 @@ func (cli *Instance) MoveSecret(cmd *cobra.Command, source, dest string, force b // Same vault? Use simple rename if possible (optimization) if srcVaultName == destVaultName { // Select the vault and do a simple move - if err := vault.SelectVault(cli.fs, cli.stateDir, srcVaultName); err != nil { + err := vault.SelectVault(cli.fs, cli.stateDir, srcVaultName) + if err != nil { return fmt.Errorf("failed to select vault '%s': %w", srcVaultName, err) } @@ -678,11 +749,14 @@ func (cli *Instance) MoveSecret(cmd *cobra.Command, source, dest string, force b } // Cross-vault move - return cli.moveSecretCrossVault(cmd, srcVaultName, srcSecretName, destVaultName, destSecretName, force) + return cli.moveSecretCrossVault( + cmd, srcVaultName, srcSecretName, destVaultName, destSecretName, force) } // moveSecretWithinVault handles rename within the current vault -func (cli *Instance) moveSecretWithinVault(cmd *cobra.Command, source, dest string, force bool) error { +func (cli *Instance) moveSecretWithinVault( + cmd *cobra.Command, source, dest string, force bool, +) error { currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { return err @@ -702,7 +776,7 @@ func (cli *Instance) moveSecretWithinVault(cmd *cobra.Command, source, dest stri } if !exists { - return fmt.Errorf("secret '%s' not found", source) + return fmt.Errorf("secret '%s' %w", source, errSecretNotFound) } destEncoded := strings.ReplaceAll(dest, "/", "%") @@ -715,15 +789,17 @@ func (cli *Instance) moveSecretWithinVault(cmd *cobra.Command, source, dest stri if exists { if !force { - return fmt.Errorf("secret '%s' already exists (use --force to overwrite)", dest) + return fmt.Errorf("secret '%s' %w", dest, errSecretExistsNoForce) } - if err := cli.fs.RemoveAll(destDir); err != nil { + err = cli.fs.RemoveAll(destDir) + if err != nil { return fmt.Errorf("failed to remove existing destination: %w", err) } } - if err := cli.fs.Rename(sourceDir, destDir); err != nil { + err = cli.fs.Rename(sourceDir, destDir) + if err != nil { return fmt.Errorf("failed to move secret: %w", err) } @@ -741,8 +817,8 @@ func (cli *Instance) moveSecretCrossVault( ) error { // Get source vault srcVault := vault.NewVault(cli.fs, cli.stateDir, srcVaultName) - srcVaultDir, err := srcVault.GetDirectory() + srcVaultDir, err := srcVault.GetDirectory() if err != nil { return fmt.Errorf("failed to get source vault directory: %w", err) } @@ -750,7 +826,7 @@ func (cli *Instance) moveSecretCrossVault( // Verify source vault exists exists, err := afero.DirExists(cli.fs, srcVaultDir) if err != nil || !exists { - return fmt.Errorf("source vault '%s' does not exist", srcVaultName) + return fmt.Errorf("source vault '%s' %w", srcVaultName, errVaultDoesNotExist) } // Verify source secret exists @@ -759,13 +835,14 @@ func (cli *Instance) moveSecretCrossVault( exists, err = afero.DirExists(cli.fs, srcSecretDir) if err != nil || !exists { - return fmt.Errorf("secret '%s' not found in vault '%s'", srcSecretName, srcVaultName) + return fmt.Errorf("secret '%s' %w in vault '%s'", + srcSecretName, errSecretNotFound, srcVaultName) } // Get destination vault destVault := vault.NewVault(cli.fs, cli.stateDir, destVaultName) - destVaultDir, err := destVault.GetDirectory() + destVaultDir, err := destVault.GetDirectory() if err != nil { return fmt.Errorf("failed to get destination vault directory: %w", err) } @@ -773,7 +850,8 @@ func (cli *Instance) moveSecretCrossVault( // Verify destination vault exists exists, err = afero.DirExists(cli.fs, destVaultDir) if err != nil || !exists { - return fmt.Errorf("destination vault '%s' does not exist", destVaultName) + return fmt.Errorf("destination vault '%s' %w", + destVaultName, errVaultDoesNotExist) } // Unlock destination vault (will fail if neither mnemonic nor unlocker available) @@ -787,12 +865,15 @@ func (cli *Instance) moveSecretCrossVault( versionCount := len(versions) // Copy all versions - if err := destVault.CopySecretAllVersions(srcVault, srcSecretName, destSecretName, force); err != nil { + err = destVault.CopySecretAllVersions( + srcVault, srcSecretName, destSecretName, force) + if err != nil { return err } // Delete source secret - if err := cli.fs.RemoveAll(srcSecretDir); err != nil { + err = cli.fs.RemoveAll(srcSecretDir) + if err != nil { // Copy succeeded but delete failed - warn but don't fail cmd.Printf("Warning: copied secret but failed to remove source: %v\n", err) cmd.Printf("Moved secret '%s:%s' to '%s:%s' (%d version(s))\n", diff --git a/internal/cli/secrets_size_test.go b/internal/cli/secrets_size_test.go index dd882f8..670589d 100644 --- a/internal/cli/secrets_size_test.go +++ b/internal/cli/secrets_size_test.go @@ -1,3 +1,4 @@ +//nolint:testpackage // white-box test of unexported internals package cli import ( @@ -18,7 +19,144 @@ import ( "github.com/stretchr/testify/require" ) +// testVaultName is the vault name used by the size tests. +const testVaultName = "test-vault" + +// newSizeTestVault creates an in-memory vault unlocked with the test +// mnemonic and returns the filesystem and vault. +// +//nolint:ireturn // afero.Fs is the filesystem abstraction used throughout +func newSizeTestVault(t *testing.T) (afero.Fs, *vault.Vault) { + t.Helper() + + fs := afero.NewMemMapFs() + + // Set test mnemonic + t.Setenv(secret.EnvMnemonic, testMnemonic) + + // Create vault + _, err := vault.CreateVault(fs, testStateDir, testVaultName) + require.NoError(t, err) + + // Set current vault + currentVaultPath := filepath.Join(testStateDir, "currentvault") + vaultPath := filepath.Join(testStateDir, "vaults.d", testVaultName) + err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600) + require.NoError(t, err) + + // Get vault and set up long-term key + vlt, err := vault.GetCurrentVault(fs, testStateDir) + require.NoError(t, err) + + ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0) + require.NoError(t, err) + vlt.Unlock(ltIdentity) + + return fs, vlt +} + +// runAddSecretSizeCase adds a secret of the given size through stdin and +// verifies the outcome. +func runAddSecretSizeCase(t *testing.T, size int, wantErr bool, errMsg string) { + t.Helper() + + fs, vlt := newSizeTestVault(t) + + // Generate test data of specified size + testData := make([]byte, size) + _, err := rand.Read(testData) + require.NoError(t, err) + + // Add newline that will be stripped + testDataWithNewline := make([]byte, 0, len(testData)+1) + testDataWithNewline = append(testDataWithNewline, testData...) + testDataWithNewline = append(testDataWithNewline, '\n') + + // Create command with fake stdin + cmd := &cobra.Command{} + cmd.SetIn(bytes.NewReader(testDataWithNewline)) + + // Create CLI instance + cli, err := NewCLIInstance() + if err != nil { + t.Fatalf("failed to initialize CLI: %v", err) + } + + cli.fs = fs + cli.stateDir = testStateDir + cli.cmd = cmd + + // Test adding the secret + secretName := fmt.Sprintf("test-secret-%d", size) + err = cli.AddSecret(secretName, false) + + if wantErr { + require.Error(t, err) + assert.Contains(t, err.Error(), errMsg) + + return + } + + require.NoError(t, err) + + // Verify the secret was stored correctly + retrievedValue, err := vlt.GetSecret(secretName) + require.NoError(t, err) + assert.Equal(t, testData, retrievedValue, + "Retrieved secret should match original (without newline)") +} + +// runImportSecretSizeCase imports a secret file of the given size and +// verifies the outcome. +func runImportSecretSizeCase(t *testing.T, size int, wantErr bool, errMsg string) { + t.Helper() + + fs, vlt := newSizeTestVault(t) + + // Generate test data of specified size + testData := make([]byte, size) + _, err := rand.Read(testData) + require.NoError(t, err) + + // Write test data to file + testFile := fmt.Sprintf("/test/secret-%d.bin", size) + err = afero.WriteFile(fs, testFile, testData, 0o600) + require.NoError(t, err) + + // Create command + cmd := &cobra.Command{} + + // Create CLI instance + cli, err := NewCLIInstance() + if err != nil { + t.Fatalf("failed to initialize CLI: %v", err) + } + + cli.fs = fs + cli.stateDir = testStateDir + + // Test importing the secret + secretName := fmt.Sprintf("imported-secret-%d", size) + err = cli.ImportSecret(cmd, secretName, testFile, false) + + if wantErr { + require.Error(t, err) + assert.Contains(t, err.Error(), errMsg) + + return + } + + require.NoError(t, err) + + // Verify the secret was stored correctly + retrievedValue, err := vlt.GetSecret(secretName) + require.NoError(t, err) + assert.Equal(t, testData, retrievedValue, "Retrieved secret should match original") +} + // TestAddSecretVariousSizes tests adding secrets of various sizes through stdin +// +//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault func TestAddSecretVariousSizes(t *testing.T) { tests := []struct { name string @@ -71,76 +209,14 @@ func TestAddSecretVariousSizes(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - // Set up test environment - fs := afero.NewMemMapFs() - stateDir := "/test/state" - - // Set test mnemonic - t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about") - - // Create vault - vaultName := "test-vault" - _, err := vault.CreateVault(fs, stateDir, vaultName) - require.NoError(t, err) - - // Set current vault - currentVaultPath := filepath.Join(stateDir, "currentvault") - vaultPath := filepath.Join(stateDir, "vaults.d", vaultName) - err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600) - require.NoError(t, err) - - // Get vault and set up long-term key - vlt, err := vault.GetCurrentVault(fs, stateDir) - require.NoError(t, err) - - ltIdentity, err := agehd.DeriveIdentity("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about", 0) - require.NoError(t, err) - vlt.Unlock(ltIdentity) - - // Generate test data of specified size - testData := make([]byte, tt.size) - _, err = rand.Read(testData) - require.NoError(t, err) - - // Add newline that will be stripped - testDataWithNewline := append(testData, '\n') - - // Create fake stdin - stdin := bytes.NewReader(testDataWithNewline) - - // Create command with fake stdin - cmd := &cobra.Command{} - cmd.SetIn(stdin) - - // Create CLI instance - cli, err := NewCLIInstance() - if err != nil { - t.Fatalf("failed to initialize CLI: %v", err) - } - cli.fs = fs - cli.stateDir = stateDir - cli.cmd = cmd - - // Test adding the secret - secretName := fmt.Sprintf("test-secret-%d", tt.size) - err = cli.AddSecret(secretName, false) - - if tt.shouldError { - assert.Error(t, err) - assert.Contains(t, err.Error(), tt.errorMsg) - } else { - require.NoError(t, err) - - // Verify the secret was stored correctly - retrievedValue, err := vlt.GetSecret(secretName) - require.NoError(t, err) - assert.Equal(t, testData, retrievedValue, "Retrieved secret should match original (without newline)") - } + runAddSecretSizeCase(t, tt.size, tt.shouldError, tt.errorMsg) }) } } // TestImportSecretVariousSizes tests importing secrets of various sizes from files +// +//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault func TestImportSecretVariousSizes(t *testing.T) { tests := []struct { name string @@ -193,73 +269,14 @@ func TestImportSecretVariousSizes(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - // Set up test environment - fs := afero.NewMemMapFs() - stateDir := "/test/state" - - // Set test mnemonic - t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about") - - // Create vault - vaultName := "test-vault" - _, err := vault.CreateVault(fs, stateDir, vaultName) - require.NoError(t, err) - - // Set current vault - currentVaultPath := filepath.Join(stateDir, "currentvault") - vaultPath := filepath.Join(stateDir, "vaults.d", vaultName) - err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600) - require.NoError(t, err) - - // Get vault and set up long-term key - vlt, err := vault.GetCurrentVault(fs, stateDir) - require.NoError(t, err) - - ltIdentity, err := agehd.DeriveIdentity("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about", 0) - require.NoError(t, err) - vlt.Unlock(ltIdentity) - - // Generate test data of specified size - testData := make([]byte, tt.size) - _, err = rand.Read(testData) - require.NoError(t, err) - - // Write test data to file - testFile := fmt.Sprintf("/test/secret-%d.bin", tt.size) - err = afero.WriteFile(fs, testFile, testData, 0o600) - require.NoError(t, err) - - // Create command - cmd := &cobra.Command{} - - // Create CLI instance - cli, err := NewCLIInstance() - if err != nil { - t.Fatalf("failed to initialize CLI: %v", err) - } - cli.fs = fs - cli.stateDir = stateDir - - // Test importing the secret - secretName := fmt.Sprintf("imported-secret-%d", tt.size) - err = cli.ImportSecret(cmd, secretName, testFile, false) - - if tt.shouldError { - assert.Error(t, err) - assert.Contains(t, err.Error(), tt.errorMsg) - } else { - require.NoError(t, err) - - // Verify the secret was stored correctly - retrievedValue, err := vlt.GetSecret(secretName) - require.NoError(t, err) - assert.Equal(t, testData, retrievedValue, "Retrieved secret should match original") - } + runImportSecretSizeCase(t, tt.size, tt.shouldError, tt.errorMsg) }) } } // TestAddSecretBufferGrowth tests that our buffer growth strategy works correctly +// +//nolint:paralleltest // subtests use t.Setenv via newSizeTestVault func TestAddSecretBufferGrowth(t *testing.T) { // Test various sizes that should trigger buffer growth sizes := []int{ @@ -283,31 +300,7 @@ func TestAddSecretBufferGrowth(t *testing.T) { for _, size := range sizes { t.Run(fmt.Sprintf("size_%d", size), func(t *testing.T) { - // Set up test environment - fs := afero.NewMemMapFs() - stateDir := "/test/state" - - // Set test mnemonic - t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about") - - // Create vault - vaultName := "test-vault" - _, err := vault.CreateVault(fs, stateDir, vaultName) - require.NoError(t, err) - - // Set current vault - currentVaultPath := filepath.Join(stateDir, "currentvault") - vaultPath := filepath.Join(stateDir, "vaults.d", vaultName) - err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600) - require.NoError(t, err) - - // Get vault and set up long-term key - vlt, err := vault.GetCurrentVault(fs, stateDir) - require.NoError(t, err) - - ltIdentity, err := agehd.DeriveIdentity("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about", 0) - require.NoError(t, err) - vlt.Unlock(ltIdentity) + fs, vlt := newSizeTestVault(t) // Create test data of exactly the specified size // Use a pattern that's easy to verify @@ -316,20 +309,18 @@ func TestAddSecretBufferGrowth(t *testing.T) { testData[i] = byte(i % 256) } - // Create fake stdin without newline - stdin := bytes.NewReader(testData) - - // Create command with fake stdin + // Create command with fake stdin (no newline) cmd := &cobra.Command{} - cmd.SetIn(stdin) + cmd.SetIn(bytes.NewReader(testData)) // Create CLI instance cli, err := NewCLIInstance() if err != nil { t.Fatalf("failed to initialize CLI: %v", err) } + cli.fs = fs - cli.stateDir = stateDir + cli.stateDir = testStateDir cli.cmd = cmd // Test adding the secret @@ -340,58 +331,38 @@ func TestAddSecretBufferGrowth(t *testing.T) { // Verify the secret was stored correctly retrievedValue, err := vlt.GetSecret(secretName) require.NoError(t, err) - assert.Equal(t, testData, retrievedValue, "Retrieved secret should match original exactly") + assert.Equal(t, testData, retrievedValue, + "Retrieved secret should match original exactly") }) } } // TestAddSecretStreamingBehavior tests that we handle streaming input correctly +// +//nolint:paralleltest // uses t.Setenv via newSizeTestVault func TestAddSecretStreamingBehavior(t *testing.T) { - // Set up test environment - fs := afero.NewMemMapFs() - stateDir := "/test/state" - - // Set test mnemonic - t.Setenv(secret.EnvMnemonic, "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about") - - // Create vault - vaultName := "test-vault" - _, err := vault.CreateVault(fs, stateDir, vaultName) - require.NoError(t, err) - - // Set current vault - currentVaultPath := filepath.Join(stateDir, "currentvault") - vaultPath := filepath.Join(stateDir, "vaults.d", vaultName) - err = afero.WriteFile(fs, currentVaultPath, []byte(vaultPath), 0o600) - require.NoError(t, err) - - // Get vault and set up long-term key - vlt, err := vault.GetCurrentVault(fs, stateDir) - require.NoError(t, err) - - ltIdentity, err := agehd.DeriveIdentity("abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about", 0) - require.NoError(t, err) - vlt.Unlock(ltIdentity) + fs, vlt := newSizeTestVault(t) // Create a custom reader that simulates slow streaming input // This will help verify our buffer handling works correctly with partial reads testData := []byte(strings.Repeat("Hello, World! ", 1000)) // ~14KB - slowReader := &slowReader{ + streamingStdin := &slowReader{ data: testData, chunkSize: 1000, // Read 1KB at a time } // Create command with slow reader as stdin cmd := &cobra.Command{} - cmd.SetIn(slowReader) + cmd.SetIn(streamingStdin) // Create CLI instance cli, err := NewCLIInstance() if err != nil { t.Fatalf("failed to initialize CLI: %v", err) } + cli.fs = fs - cli.stateDir = stateDir + cli.stateDir = testStateDir cli.cmd = cmd // Test adding the secret @@ -411,27 +382,22 @@ type slowReader struct { chunkSize int } -func (r *slowReader) Read(p []byte) (n int, err error) { +func (r *slowReader) Read(p []byte) (int, error) { if r.offset >= len(r.data) { return 0, io.EOF } - // Read at most chunkSize bytes + // Read at most chunkSize bytes, bounded by the remaining data and + // the destination buffer remaining := len(r.data) - r.offset - toRead := r.chunkSize - if toRead > remaining { - toRead = remaining - } - if toRead > len(p) { - toRead = len(p) - } + toRead := min(r.chunkSize, remaining, len(p)) - n = copy(p, r.data[r.offset:r.offset+toRead]) + n := copy(p, r.data[r.offset:r.offset+toRead]) r.offset += n if r.offset >= len(r.data) { - err = io.EOF + return n, io.EOF } - return n, err + return n, nil } diff --git a/internal/cli/stdout_stderr_test.go b/internal/cli/stdout_stderr_test.go index 14042af..2176998 100644 --- a/internal/cli/stdout_stderr_test.go +++ b/internal/cli/stdout_stderr_test.go @@ -7,57 +7,64 @@ import ( "strings" "testing" + "git.eeqj.de/sneak/secret/internal/secret" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -// TestGetCommandOutputsToStdout tests that 'secret get' outputs the secret value to stdout, not stderr +// TestGetCommandOutputsToStdout tests that 'secret get' outputs the secret +// value to stdout, not stderr func TestGetCommandOutputsToStdout(t *testing.T) { // Create a temporary directory for our vault tempDir := t.TempDir() // Set environment variables for the test - t.Setenv("SB_SECRET_STATE_DIR", tempDir) + t.Setenv(secret.EnvStateDir, tempDir) // Find the secret binary path wd, err := filepath.Abs("../..") require.NoError(t, err, "should get working directory") - secretPath := filepath.Join(wd, "secret") - testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about" + secretPath := filepath.Join(wd, "secret") testPassphrase := "test-passphrase" // Initialize vault - cmd := exec.Command(secretPath, "init") + //nolint:gosec // G204: test executes the freshly built secret binary + cmd := exec.CommandContext(t.Context(), secretPath, "init") cmd.Env = []string{ - "SB_SECRET_STATE_DIR=" + tempDir, - "SB_SECRET_MNEMONIC=" + testMnemonic, - "SB_UNLOCK_PASSPHRASE=" + testPassphrase, + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, + secret.EnvUnlockPassphrase + "=" + testPassphrase, "PATH=" + "/usr/bin:/bin", } + output, err := cmd.CombinedOutput() require.NoError(t, err, "init should succeed: %s", string(output)) // Add a secret - cmd = exec.Command(secretPath, "add", "test/secret") + //nolint:gosec // G204: test executes the freshly built secret binary + cmd = exec.CommandContext(t.Context(), secretPath, "add", "test/secret") cmd.Env = []string{ - "SB_SECRET_STATE_DIR=" + tempDir, - "SB_SECRET_MNEMONIC=" + testMnemonic, + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, "PATH=" + "/usr/bin:/bin", } cmd.Stdin = strings.NewReader("test-secret-value") + output, err = cmd.CombinedOutput() require.NoError(t, err, "add should succeed: %s", string(output)) // Test that 'secret get' outputs to stdout, not stderr - cmd = exec.Command(secretPath, "get", "test/secret") + //nolint:gosec // G204: test executes the freshly built secret binary + cmd = exec.CommandContext(t.Context(), secretPath, "get", "test/secret") cmd.Env = []string{ - "SB_SECRET_STATE_DIR=" + tempDir, - "SB_SECRET_MNEMONIC=" + testMnemonic, + secret.EnvStateDir + "=" + tempDir, + secret.EnvMnemonic + "=" + testMnemonic, "PATH=" + "/usr/bin:/bin", } var stdout, stderr bytes.Buffer + cmd.Stdout = &stdout cmd.Stderr = &stderr @@ -65,7 +72,8 @@ func TestGetCommandOutputsToStdout(t *testing.T) { require.NoError(t, err, "get should succeed") // The secret value should be in stdout - assert.Equal(t, "test-secret-value", strings.TrimSpace(stdout.String()), "secret value should be in stdout") + assert.Equal(t, "test-secret-value", strings.TrimSpace(stdout.String()), + "secret value should be in stdout") // Nothing should be in stderr assert.Empty(t, stderr.String(), "stderr should be empty") diff --git a/internal/cli/test_helpers.go b/internal/cli/test_helpers.go index 06e4e66..c9e43e1 100644 --- a/internal/cli/test_helpers.go +++ b/internal/cli/test_helpers.go @@ -9,7 +9,9 @@ import ( ) // ExecuteCommandInProcess executes a CLI command in-process for testing -func ExecuteCommandInProcess(args []string, stdin string, env map[string]string) (string, error) { +func ExecuteCommandInProcess( + args []string, stdin string, env map[string]string, +) (string, error) { secret.Debug("ExecuteCommandInProcess called", "args", args) // Save current environment @@ -43,11 +45,13 @@ func ExecuteCommandInProcess(args []string, stdin string, env map[string]string) err := rootCmd.Execute() output := buf.String() - secret.Debug("Command execution completed", "error", err, "outputLength", len(output), "output", output) + secret.Debug("Command execution completed", + "error", err, "outputLength", len(output), "output", output) // Add debug info for troubleshooting if len(output) == 0 && err == nil { - secret.Debug("Warning: Command executed successfully but produced no output", "args", args) + secret.Debug("Warning: Command executed successfully but produced no output", + "args", args) } // Restore environment diff --git a/internal/cli/test_output_test.go b/internal/cli/test_output_test.go index 2365f6a..46e2f8e 100644 --- a/internal/cli/test_output_test.go +++ b/internal/cli/test_output_test.go @@ -1,21 +1,23 @@ -package cli +package cli_test import ( "testing" + "git.eeqj.de/sneak/secret/internal/cli" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) +//nolint:paralleltest // executes the CLI in-process against shared state func TestOutputCapture(t *testing.T) { // Test vault list command which we fixed - output, err := ExecuteCommandInProcess([]string{"vault", "list"}, "", nil) + output, err := cli.ExecuteCommandInProcess([]string{"vault", "list"}, "", nil) require.NoError(t, err) assert.Contains(t, output, "Available vaults", "should capture vault list output") t.Logf("vault list output: %q", output) // Test help command - output, err = ExecuteCommandInProcess([]string{"--help"}, "", nil) + output, err = cli.ExecuteCommandInProcess([]string{"--help"}, "", nil) require.NoError(t, err) assert.NotEmpty(t, output, "help output should not be empty") t.Logf("help output length: %d", len(output)) diff --git a/internal/cli/unlockers.go b/internal/cli/unlockers.go index 0c1d3b0..0ee849c 100644 --- a/internal/cli/unlockers.go +++ b/internal/cli/unlockers.go @@ -1,13 +1,16 @@ package cli import ( + "context" "encoding/json" + "errors" "fmt" "log" "os" "os/exec" "path/filepath" "runtime" + "slices" "strings" "time" @@ -18,6 +21,35 @@ import ( "github.com/spf13/cobra" ) +// Unlocker type names and platform identifiers shared across the CLI +const ( + unlockerTypePassphrase = "passphrase" + unlockerTypeKeychain = "keychain" + unlockerTypePGP = "pgp" + unlockerTypeSecureEnclave = "secure-enclave" + + platformDarwin = "darwin" + + cmdUseList = "list" +) + +// Sentinel errors for unlocker operations +var ( + errNoGPGSecretKeys = errors.New("no GPG secret keys found") + errInvalidUnlockerType = errors.New("invalid unlocker type") + errKeyIDOnlyForPGP = errors.New( + "--keyid flag is only valid for PGP unlockers") + errKeychainMacOSOnly = errors.New( + "keychain unlockers are only supported on macOS") + errSecureEnclaveMacOSOnly = errors.New( + "secure enclave unlockers are only supported on macOS") + errGPGKeyAlreadyUnlocker = errors.New( + "GPG key is already added as an unlocker") + errUnsupportedUnlockerType = errors.New("unsupported unlocker type") + errLastUnlocker = errors.New("refusing to remove last unlocker") + errUnlockerExists = errors.New("unlocker already exists") +) + // UnlockerInfo represents unlocker information for display type UnlockerInfo struct { ID string `json:"id"` @@ -37,12 +69,14 @@ const ( // getDefaultGPGKey returns the default GPG key ID if available func getDefaultGPGKey() (string, error) { + ctx := context.Background() + // First try to get the configured default key using gpgconf - cmd := exec.Command("gpgconf", "--list-options", "gpg") + cmd := exec.CommandContext(ctx, "gpgconf", "--list-options", "gpg") + output, err := cmd.Output() if err == nil { - lines := strings.Split(string(output), "\n") - for _, line := range lines { + for line := range strings.SplitSeq(string(output), "\n") { fields := strings.Split(line, ":") if len(fields) > 9 && fields[0] == "default-key" && fields[9] != "" { // The default key is in field 10 (index 9) @@ -52,15 +86,15 @@ func getDefaultGPGKey() (string, error) { } // If no default key is configured, get the first secret key - cmd = exec.Command("gpg", "--list-secret-keys", "--with-colons") + cmd = exec.CommandContext(ctx, "gpg", "--list-secret-keys", "--with-colons") + output, err = cmd.Output() if err != nil { return "", fmt.Errorf("failed to list GPG keys: %w", err) } // Parse output to find the first usable secret key - lines := strings.Split(string(output), "\n") - for _, line := range lines { + for line := range strings.SplitSeq(string(output), "\n") { // sec line indicates a secret key if strings.HasPrefix(line, "sec:") { fields := strings.Split(line, ":") @@ -71,7 +105,7 @@ func getDefaultGPGKey() (string, error) { } } - return "", fmt.Errorf("no GPG secret keys found") + return "", errNoGPGSecretKeys } func newUnlockerCmd() *cobra.Command { @@ -91,7 +125,7 @@ func newUnlockerCmd() *cobra.Command { func newUnlockerListCmd() *cobra.Command { cmd := &cobra.Command{ - Use: "list", + Use: cmdUseList, Aliases: []string{"ls"}, Short: "List unlockers in the current vault", RunE: func(cmd *cobra.Command, _ []string) error { @@ -101,6 +135,7 @@ func newUnlockerListCmd() *cobra.Command { if err != nil { return fmt.Errorf("failed to initialize CLI: %w", err) } + cli.cmd = cmd return cli.UnlockersList(jsonOutput) @@ -112,53 +147,80 @@ func newUnlockerListCmd() *cobra.Command { return cmd } -func newUnlockerAddCmd() *cobra.Command { +// unlockerAddHelp returns the supported unlocker types list and their +// descriptions for the current platform +func unlockerAddHelp() (string, string) { // Build the supported types list based on platform supportedTypes := "passphrase, pgp" - typeDescriptions := `Available unlocker types: + typeDescriptions := "Available unlocker types:\n" + + "\n" + + " passphrase - Traditional password-based encryption\n" + + " Prompts for a passphrase that will be used to " + + "encrypt/decrypt the vault's master key.\n" + + " The passphrase is never stored in plaintext.\n" + + "\n" + + " pgp - GNU Privacy Guard (GPG) key-based encryption \n" + + " Uses your existing GPG key to encrypt/decrypt " + + "the vault's master key.\n" + + " Requires gpg to be installed and configured " + + "with at least one secret key.\n" + + " Use --keyid to specify a particular key, " + + "otherwise uses your default GPG key." - passphrase - Traditional password-based encryption - Prompts for a passphrase that will be used to encrypt/decrypt the vault's master key. - The passphrase is never stored in plaintext. - - pgp - GNU Privacy Guard (GPG) key-based encryption - Uses your existing GPG key to encrypt/decrypt the vault's master key. - Requires gpg to be installed and configured with at least one secret key. - Use --keyid to specify a particular key, otherwise uses your default GPG key.` - - if runtime.GOOS == "darwin" { + if runtime.GOOS == platformDarwin { supportedTypes = "passphrase, keychain, pgp, secure-enclave" - typeDescriptions = `Available unlocker types: - - passphrase - Traditional password-based encryption - Prompts for a passphrase that will be used to encrypt/decrypt the vault's master key. - The passphrase is never stored in plaintext. - - keychain - macOS Keychain integration (macOS only) - Stores the vault's master key in the macOS Keychain, protected by your login password. - Automatically unlocks when your Keychain is unlocked (e.g., after login). - Provides seamless integration with macOS security features like Touch ID. - - pgp - GNU Privacy Guard (GPG) key-based encryption - Uses your existing GPG key to encrypt/decrypt the vault's master key. - Requires gpg to be installed and configured with at least one secret key. - Use --keyid to specify a particular key, otherwise uses your default GPG key. - - secure-enclave - Apple Secure Enclave hardware protection (macOS only) - Stores the vault's master key encrypted by a non-exportable P-256 key - held in the Secure Enclave. The key never leaves the hardware. - Uses ECIES encryption; decryption is performed inside the SE.` + typeDescriptions = "Available unlocker types:\n" + + "\n" + + " passphrase - Traditional password-based encryption\n" + + " Prompts for a passphrase that will be " + + "used to encrypt/decrypt the vault's master key.\n" + + " The passphrase is never stored in " + + "plaintext.\n" + + "\n" + + " keychain - macOS Keychain integration (macOS only)\n" + + " Stores the vault's master key in the " + + "macOS Keychain, protected by your login password.\n" + + " Automatically unlocks when your Keychain " + + "is unlocked (e.g., after login).\n" + + " Provides seamless integration with macOS " + + "security features like Touch ID.\n" + + "\n" + + " pgp - GNU Privacy Guard (GPG) key-based " + + "encryption\n" + + " Uses your existing GPG key to " + + "encrypt/decrypt the vault's master key.\n" + + " Requires gpg to be installed and " + + "configured with at least one secret key.\n" + + " Use --keyid to specify a particular key, " + + "otherwise uses your default GPG key.\n" + + "\n" + + " secure-enclave - Apple Secure Enclave hardware protection " + + "(macOS only)\n" + + " Stores the vault's master key encrypted " + + "by a non-exportable P-256 key\n" + + " held in the Secure Enclave. The key " + + "never leaves the hardware.\n" + + " Uses ECIES encryption; decryption is " + + "performed inside the SE." } + return supportedTypes, typeDescriptions +} + +func newUnlockerAddCmd() *cobra.Command { + supportedTypes, typeDescriptions := unlockerAddHelp() + cmd := &cobra.Command{ Use: "add ", Short: "Add a new unlocker", - Long: fmt.Sprintf(`Add a new unlocker to the current vault. - -%s - -Each vault can have multiple unlockers, allowing different authentication methods -to access the same vault. This provides flexibility and backup access options.`, typeDescriptions), + Long: "Add a new unlocker to the current vault.\n" + + "\n" + + typeDescriptions + "\n" + + "\n" + + "Each vault can have multiple unlockers, allowing different " + + "authentication methods\n" + + "to access the same vault. This provides flexibility and " + + "backup access options.", Args: cobra.ExactArgs(1), ValidArgs: strings.Split(supportedTypes, ", "), RunE: func(cmd *cobra.Command, args []string) error { @@ -166,33 +228,28 @@ to access the same vault. This provides flexibility and backup access options.`, if err != nil { return fmt.Errorf("failed to initialize CLI: %w", err) } + unlockerType := args[0] // Validate unlocker type validTypes := strings.Split(supportedTypes, ", ") - valid := false - for _, t := range validTypes { - if unlockerType == t { - valid = true - - break - } - } - if !valid { - return fmt.Errorf("invalid unlocker type '%s'\n\nSupported types: %s\n\n"+ - "Run 'secret unlocker add --help' for detailed descriptions", unlockerType, supportedTypes) + if !slices.Contains(validTypes, unlockerType) { + return fmt.Errorf("%w '%s'\n\nSupported types: %s\n\n"+ + "Run 'secret unlocker add --help' for detailed descriptions", + errInvalidUnlockerType, unlockerType, supportedTypes) } // Check if --keyid was used with non-PGP type - if unlockerType != "pgp" && cmd.Flags().Changed("keyid") { - return fmt.Errorf("--keyid flag is only valid for PGP unlockers") + if unlockerType != unlockerTypePGP && cmd.Flags().Changed("keyid") { + return errKeyIDOnlyForPGP } return cli.UnlockersAdd(unlockerType, cmd) }, } - cmd.Flags().String("keyid", "", "GPG key ID for PGP unlockers (optional, uses default key if not specified)") + cmd.Flags().String("keyid", "", + "GPG key ID for PGP unlockers (optional, uses default key if not specified)") return cmd } @@ -202,17 +259,20 @@ func newUnlockerRemoveCmd() *cobra.Command { if err != nil { log.Fatalf("failed to initialize CLI: %v", err) } + cmd := &cobra.Command{ Use: "remove ", Aliases: []string{"rm"}, Short: "Remove an unlocker", - Long: `Remove an unlocker from the current vault. Cannot remove the last unlocker if the vault has ` + - `secrets unless --force is used. Warning: Without unlockers and without your mnemonic, vault data ` + - `will be permanently inaccessible.`, + Long: `Remove an unlocker from the current vault. Cannot remove ` + + `the last unlocker if the vault has secrets unless --force is ` + + `used. Warning: Without unlockers and without your mnemonic, ` + + `vault data will be permanently inaccessible.`, Args: cobra.ExactArgs(1), ValidArgsFunction: getUnlockerIDsCompletionFunc(cli.fs, cli.stateDir), RunE: func(cmd *cobra.Command, args []string) error { force, _ := cmd.Flags().GetBool("force") + cli, err := NewCLIInstance() if err != nil { return fmt.Errorf("failed to initialize CLI: %w", err) @@ -222,7 +282,8 @@ func newUnlockerRemoveCmd() *cobra.Command { }, } - cmd.Flags().BoolP("force", "f", false, "Force removal of last unlocker even if vault has secrets") + cmd.Flags().BoolP("force", "f", false, + "Force removal of last unlocker even if vault has secrets") return cmd } @@ -249,6 +310,88 @@ func newUnlockerSelectCmd() *cobra.Command { } } +// unlockerIDFromDir constructs an unlocker of the given metadata type +// rooted at unlockerDir and returns its ID. Returns "" for unknown types +// and, when includeSecureEnclave is false, for secure enclave unlockers. +func unlockerIDFromDir( + fs afero.Fs, unlockerDir string, metadata secret.UnlockerMetadata, + includeSecureEnclave bool, +) string { + // Create the appropriate unlocker instance + var unlocker secret.Unlocker + + switch metadata.Type { + case unlockerTypePassphrase: + unlocker = secret.NewPassphraseUnlocker(fs, unlockerDir, metadata) + case unlockerTypeKeychain: + unlocker = secret.NewKeychainUnlocker(fs, unlockerDir, metadata) + case unlockerTypePGP: + unlocker = secret.NewPGPUnlocker(fs, unlockerDir, metadata) + case unlockerTypeSecureEnclave: + if includeSecureEnclave { + unlocker = secret.NewSecureEnclaveUnlocker(fs, unlockerDir, metadata) + } + } + + if unlocker == nil { + return "" + } + + return unlocker.GetID() +} + +// findUnlockerIDByMetadata scans unlockersDir for the directory whose +// stored metadata matches the given type and creation time and returns +// the matching unlocker's ID. Returns "" if no match is found. +func findUnlockerIDByMetadata( + fs afero.Fs, unlockersDir string, metadata secret.UnlockerMetadata, + includeSecureEnclave bool, +) string { + files, err := afero.ReadDir(fs, unlockersDir) + if err != nil { + secret.Warn("Could not read unlockers directory", "error", err) + + return "" + } + + for _, file := range files { + if !file.IsDir() { + continue + } + + unlockerDir := filepath.Join(unlockersDir, file.Name()) + metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") + + // Check if this is the right unlocker by comparing metadata + metadataBytes, err := afero.ReadFile(fs, metadataPath) + if err != nil { + secret.Warn("Could not read unlocker metadata file", + "path", metadataPath, "error", err) + + continue + } + + var diskMetadata secret.UnlockerMetadata + + err = json.Unmarshal(metadataBytes, &diskMetadata) + if err != nil { + secret.Warn("Could not parse unlocker metadata file", + "path", metadataPath, "error", err) + + continue + } + + // Match by type and creation time + if diskMetadata.Type == metadata.Type && + diskMetadata.CreatedAt.Equal(metadata.CreatedAt) { + return unlockerIDFromDir(fs, unlockerDir, diskMetadata, + includeSecureEnclave) + } + } + + return "" +} + // UnlockersList lists unlockers in the current vault func (cli *Instance) UnlockersList(jsonOutput bool) error { // Get current vault @@ -259,6 +402,7 @@ func (cli *Instance) UnlockersList(jsonOutput bool) error { // Get the current unlocker ID var currentUnlockerID string + currentUnlocker, err := vlt.GetCurrentUnlocker() if err == nil { currentUnlockerID = currentUnlocker.GetID() @@ -272,74 +416,31 @@ func (cli *Instance) UnlockersList(jsonOutput bool) error { // Load actual unlocker objects to get the proper IDs var unlockers []UnlockerInfo + for _, metadata := range unlockerMetadataList { // Create unlocker instance to get the proper ID vaultDir, err := vlt.GetDirectory() if err != nil { - secret.Warn("Could not get vault directory while listing unlockers", "error", err) + secret.Warn("Could not get vault directory while listing unlockers", + "error", err) continue } // Find the unlocker directory by type and created time unlockersDir := filepath.Join(vaultDir, "unlockers.d") - files, err := afero.ReadDir(cli.fs, unlockersDir) - if err != nil { - secret.Warn("Could not read unlockers directory", "error", err) - - continue - } - - var unlocker secret.Unlocker - for _, file := range files { - if !file.IsDir() { - continue - } - - unlockerDir := filepath.Join(unlockersDir, file.Name()) - metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") - - // Check if this is the right unlocker by comparing metadata - metadataBytes, err := afero.ReadFile(cli.fs, metadataPath) - if err != nil { - secret.Warn("Could not read unlocker metadata file", "path", metadataPath, "error", err) - - continue - } - - var diskMetadata secret.UnlockerMetadata - if err := json.Unmarshal(metadataBytes, &diskMetadata); err != nil { - secret.Warn("Could not parse unlocker metadata file", "path", metadataPath, "error", err) - - continue - } - - // Match by type and creation time - if diskMetadata.Type == metadata.Type && diskMetadata.CreatedAt.Equal(metadata.CreatedAt) { - // Create the appropriate unlocker instance - switch metadata.Type { - case "passphrase": - unlocker = secret.NewPassphraseUnlocker(cli.fs, unlockerDir, diskMetadata) - case "keychain": - unlocker = secret.NewKeychainUnlocker(cli.fs, unlockerDir, diskMetadata) - case "pgp": - unlocker = secret.NewPGPUnlocker(cli.fs, unlockerDir, diskMetadata) - case "secure-enclave": - unlocker = secret.NewSecureEnclaveUnlocker(cli.fs, unlockerDir, diskMetadata) - } - - break - } - } + unlockerID := findUnlockerIDByMetadata(cli.fs, unlockersDir, metadata, true) // Get the proper ID using the unlocker's ID() method var properID string - if unlocker != nil { - properID = unlocker.GetID() + if unlockerID != "" { + properID = unlockerID } else { // Generate ID as fallback - properID = fmt.Sprintf("%s-%s", metadata.CreatedAt.Format("2006-01-02.15.04"), metadata.Type) - secret.Warn("Could not create unlocker instance, using fallback ID", "fallback_id", properID, "type", metadata.Type) + properID = fmt.Sprintf("%s-%s", + metadata.CreatedAt.Format("2006-01-02.15.04"), metadata.Type) + secret.Warn("Could not create unlocker instance, using fallback ID", + "fallback_id", properID, "type", metadata.Type) } unlockerInfo := UnlockerInfo{ @@ -360,8 +461,10 @@ func (cli *Instance) UnlockersList(jsonOutput bool) error { } // printUnlockersJSON prints unlockers in JSON format -func (cli *Instance) printUnlockersJSON(unlockers []UnlockerInfo, currentUnlockerID string) error { - output := map[string]interface{}{ +func (cli *Instance) printUnlockersJSON( + unlockers []UnlockerInfo, currentUnlockerID string, +) error { + output := map[string]any{ "unlockers": unlockers, "currentUnlockerID": currentUnlockerID, } @@ -395,10 +498,12 @@ func (cli *Instance) printUnlockersTable(unlockers []UnlockerInfo) error { if len(unlocker.Flags) > 0 { flags = strings.Join(unlocker.Flags, ",") } + prefix := " " if unlocker.IsCurrent { prefix = "* " } + cli.cmd.Printf("%s%-40s %-12s %-20s %s\n", prefix, unlocker.ID, @@ -414,164 +519,186 @@ func (cli *Instance) printUnlockersTable(unlockers []UnlockerInfo) error { // UnlockersAdd adds a new unlocker func (cli *Instance) UnlockersAdd(unlockerType string, cmd *cobra.Command) error { - // Build the supported types list based on platform - supportedTypes := "passphrase, pgp" - if runtime.GOOS == "darwin" { - supportedTypes = "passphrase, keychain, pgp, secure-enclave" - } - switch unlockerType { - case "passphrase": - // Get current vault - vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) - if err != nil { - return fmt.Errorf("failed to get current vault: %w", err) - } - - // For passphrase unlockers, we don't need the vault to be unlocked - // The CreatePassphraseUnlocker method will handle getting the long-term key - - // Check if passphrase is set in environment variable - var passphraseBuffer *memguard.LockedBuffer - if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" { - passphraseBuffer = memguard.NewBufferFromBytes([]byte(envPassphrase)) - } else { - // Use secure passphrase input with confirmation - passphraseBuffer, err = readSecurePassphrase("Enter passphrase for unlocker: ") - if err != nil { - return fmt.Errorf("failed to read passphrase: %w", err) - } - } - defer passphraseBuffer.Destroy() - - passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer) - if err != nil { - return err - } - - cmd.Printf("Created passphrase unlocker: %s\n", passphraseUnlocker.GetID()) - - // Auto-select the newly created unlocker - if err := vlt.SelectUnlocker(passphraseUnlocker.GetID()); err != nil { - cmd.Printf("Warning: Failed to auto-select new unlocker: %v\n", err) - } else { - cmd.Printf("Automatically selected as current unlocker\n") - } - - return nil - - case "keychain": - if runtime.GOOS != "darwin" { - return fmt.Errorf("keychain unlockers are only supported on macOS") - } - - keychainUnlocker, err := secret.CreateKeychainUnlocker(cli.fs, cli.stateDir) - if err != nil { - return fmt.Errorf("failed to create macOS Keychain unlocker: %w", err) - } - - cmd.Printf("Created macOS Keychain unlocker: %s\n", keychainUnlocker.GetID()) - if keyName, err := keychainUnlocker.GetKeychainItemName(); err == nil { - cmd.Printf("Keychain Item Name: %s\n", keyName) - } - - // Auto-select the newly created unlocker - vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) - if err != nil { - return fmt.Errorf("failed to get current vault: %w", err) - } - if err := vlt.SelectUnlocker(keychainUnlocker.GetID()); err != nil { - cmd.Printf("Warning: Failed to auto-select new unlocker: %v\n", err) - } else { - cmd.Printf("Automatically selected as current unlocker\n") - } - - return nil - - case "secure-enclave": - if runtime.GOOS != "darwin" { - return fmt.Errorf("secure enclave unlockers are only supported on macOS") - } - - seUnlocker, err := secret.CreateSecureEnclaveUnlocker(cli.fs, cli.stateDir) - if err != nil { - return fmt.Errorf("failed to create Secure Enclave unlocker: %w", err) - } - - cmd.Printf("Created Secure Enclave unlocker: %s\n", seUnlocker.GetID()) - - vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) - if err != nil { - return fmt.Errorf("failed to get current vault: %w", err) - } - - if err := vlt.SelectUnlocker(seUnlocker.GetID()); err != nil { - cmd.Printf("Warning: Failed to auto-select new unlocker: %v\n", err) - } else { - cmd.Printf("Automatically selected as current unlocker\n") - } - - return nil - - case "pgp": - // Get GPG key ID from flag, environment, or default key - var gpgKeyID string - if flagKeyID, _ := cmd.Flags().GetString("keyid"); flagKeyID != "" { - gpgKeyID = flagKeyID - } else if envKeyID := os.Getenv(secret.EnvGPGKeyID); envKeyID != "" { - gpgKeyID = envKeyID - } else { - // Try to get the default GPG key - defaultKeyID, err := getDefaultGPGKey() - if err != nil { - return fmt.Errorf("no GPG key specified and no default key found: %w", err) - } - gpgKeyID = defaultKeyID - cmd.Printf("Using default GPG key: %s\n", gpgKeyID) - } - - // Check if this key is already added as an unlocker - vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) - if err != nil { - return fmt.Errorf("failed to get current vault: %w", err) - } - - // Resolve the GPG key ID to its fingerprint - fingerprint, err := secret.ResolveGPGKeyFingerprint(gpgKeyID) - if err != nil { - return fmt.Errorf("failed to resolve GPG key fingerprint: %w", err) - } - - // Check if this GPG key is already added - expectedID := fmt.Sprintf("pgp-%s", fingerprint) - if err := cli.checkUnlockerExists(vlt, expectedID); err != nil { - return fmt.Errorf("GPG key %s is already added as an unlocker", gpgKeyID) - } - - pgpUnlocker, err := secret.CreatePGPUnlocker(cli.fs, cli.stateDir, gpgKeyID) - if err != nil { - return err - } - - cmd.Printf("Created PGP unlocker: %s\n", pgpUnlocker.GetID()) - cmd.Printf("GPG Key ID: %s\n", gpgKeyID) - - // Auto-select the newly created unlocker - if err := vlt.SelectUnlocker(pgpUnlocker.GetID()); err != nil { - cmd.Printf("Warning: Failed to auto-select new unlocker: %v\n", err) - } else { - cmd.Printf("Automatically selected as current unlocker\n") - } - - return nil - + case unlockerTypePassphrase: + return cli.addPassphraseUnlocker(cmd) + case unlockerTypeKeychain: + return cli.addKeychainUnlocker(cmd) + case unlockerTypeSecureEnclave: + return cli.addSecureEnclaveUnlocker(cmd) + case unlockerTypePGP: + return cli.addPGPUnlocker(cmd) default: - return fmt.Errorf("unsupported unlocker type: %s (supported: %s)", unlockerType, supportedTypes) + // Build the supported types list based on platform + supportedTypes := "passphrase, pgp" + if runtime.GOOS == platformDarwin { + supportedTypes = "passphrase, keychain, pgp, secure-enclave" + } + + return fmt.Errorf("%w: %s (supported: %s)", + errUnsupportedUnlockerType, unlockerType, supportedTypes) } } +// autoSelectUnlocker selects the newly created unlocker as current, +// printing a warning if selection fails +func autoSelectUnlocker(cmd *cobra.Command, vlt *vault.Vault, unlockerID string) { + err := vlt.SelectUnlocker(unlockerID) + if err != nil { + cmd.Printf("Warning: Failed to auto-select new unlocker: %v\n", err) + } else { + cmd.Printf("Automatically selected as current unlocker\n") + } +} + +// addPassphraseUnlocker creates a passphrase unlocker in the current vault +func (cli *Instance) addPassphraseUnlocker(cmd *cobra.Command) error { + // Get current vault + vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) + if err != nil { + return fmt.Errorf("failed to get current vault: %w", err) + } + + // For passphrase unlockers, we don't need the vault to be unlocked + // The CreatePassphraseUnlocker method will handle getting the + // long-term key + + // Check if passphrase is set in environment variable + var passphraseBuffer *memguard.LockedBuffer + if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" { + passphraseBuffer = memguard.NewBufferFromBytes([]byte(envPassphrase)) + } else { + // Use secure passphrase input with confirmation + passphraseBuffer, err = readSecurePassphrase("Enter passphrase for unlocker: ") + if err != nil { + return fmt.Errorf("failed to read passphrase: %w", err) + } + } + defer passphraseBuffer.Destroy() + + passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer) + if err != nil { + return err + } + + cmd.Printf("Created passphrase unlocker: %s\n", passphraseUnlocker.GetID()) + + // Auto-select the newly created unlocker + autoSelectUnlocker(cmd, vlt, passphraseUnlocker.GetID()) + + return nil +} + +// addKeychainUnlocker creates a macOS Keychain unlocker in the current vault +func (cli *Instance) addKeychainUnlocker(cmd *cobra.Command) error { + if runtime.GOOS != platformDarwin { + return errKeychainMacOSOnly + } + + keychainUnlocker, err := secret.CreateKeychainUnlocker(cli.fs, cli.stateDir) + if err != nil { + return fmt.Errorf("failed to create macOS Keychain unlocker: %w", err) + } + + cmd.Printf("Created macOS Keychain unlocker: %s\n", keychainUnlocker.GetID()) + + keyName, err := keychainUnlocker.GetKeychainItemName() + if err == nil { + cmd.Printf("Keychain Item Name: %s\n", keyName) + } + + // Auto-select the newly created unlocker + vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) + if err != nil { + return fmt.Errorf("failed to get current vault: %w", err) + } + + autoSelectUnlocker(cmd, vlt, keychainUnlocker.GetID()) + + return nil +} + +// addSecureEnclaveUnlocker creates a Secure Enclave unlocker in the +// current vault +func (cli *Instance) addSecureEnclaveUnlocker(cmd *cobra.Command) error { + if runtime.GOOS != platformDarwin { + return errSecureEnclaveMacOSOnly + } + + seUnlocker, err := secret.CreateSecureEnclaveUnlocker(cli.fs, cli.stateDir) + if err != nil { + return fmt.Errorf("failed to create Secure Enclave unlocker: %w", err) + } + + cmd.Printf("Created Secure Enclave unlocker: %s\n", seUnlocker.GetID()) + + vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) + if err != nil { + return fmt.Errorf("failed to get current vault: %w", err) + } + + autoSelectUnlocker(cmd, vlt, seUnlocker.GetID()) + + return nil +} + +// addPGPUnlocker creates a PGP unlocker in the current vault +func (cli *Instance) addPGPUnlocker(cmd *cobra.Command) error { + // Get GPG key ID from flag, environment, or default key + var gpgKeyID string + if flagKeyID, _ := cmd.Flags().GetString("keyid"); flagKeyID != "" { + gpgKeyID = flagKeyID + } else if envKeyID := os.Getenv(secret.EnvGPGKeyID); envKeyID != "" { + gpgKeyID = envKeyID + } else { + // Try to get the default GPG key + defaultKeyID, err := getDefaultGPGKey() + if err != nil { + return fmt.Errorf("no GPG key specified and no default key found: %w", err) + } + + gpgKeyID = defaultKeyID + cmd.Printf("Using default GPG key: %s\n", gpgKeyID) + } + + // Check if this key is already added as an unlocker + vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) + if err != nil { + return fmt.Errorf("failed to get current vault: %w", err) + } + + // Resolve the GPG key ID to its fingerprint + fingerprint, err := secret.ResolveGPGKeyFingerprint(gpgKeyID) + if err != nil { + return fmt.Errorf("failed to resolve GPG key fingerprint: %w", err) + } + + // Check if this GPG key is already added + expectedID := "pgp-" + fingerprint + + err = cli.checkUnlockerExists(vlt, expectedID) + if err != nil { + return fmt.Errorf("%w: %s", errGPGKeyAlreadyUnlocker, gpgKeyID) + } + + pgpUnlocker, err := secret.CreatePGPUnlocker(cli.fs, cli.stateDir, gpgKeyID) + if err != nil { + return err + } + + cmd.Printf("Created PGP unlocker: %s\n", pgpUnlocker.GetID()) + cmd.Printf("GPG Key ID: %s\n", gpgKeyID) + + // Auto-select the newly created unlocker + autoSelectUnlocker(cmd, vlt, pgpUnlocker.GetID()) + + return nil +} + // UnlockersRemove removes an unlocker with safety checks -func (cli *Instance) UnlockersRemove(unlockerID string, force bool, cmd *cobra.Command) error { +func (cli *Instance) UnlockersRemove( + unlockerID string, force bool, cmd *cobra.Command, +) error { // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -593,20 +720,24 @@ func (cli *Instance) UnlockersRemove(unlockerID string, force bool, cmd *cobra.C } if numSecrets > 0 && !force { - cmd.Println("ERROR: Cannot remove the last unlocker when the vault contains secrets.") - cmd.Println("WARNING: Without unlockers, you MUST have your mnemonic phrase to decrypt the vault.") + cmd.Println("ERROR: Cannot remove the last unlocker when the " + + "vault contains secrets.") + cmd.Println("WARNING: Without unlockers, you MUST have your " + + "mnemonic phrase to decrypt the vault.") cmd.Println("If you want to proceed anyway, use --force") - return fmt.Errorf("refusing to remove last unlocker") + return errLastUnlocker } if numSecrets > 0 && force { - cmd.Println("WARNING: Removing the last unlocker. You MUST have your mnemonic phrase to access this vault again!") + cmd.Println("WARNING: Removing the last unlocker. You MUST " + + "have your mnemonic phrase to access this vault again!") } } // Remove the unlocker - if err := vlt.RemoveUnlocker(unlockerID); err != nil { + err = vlt.RemoveUnlocker(unlockerID) + if err != nil { return err } @@ -639,65 +770,20 @@ func (cli *Instance) checkUnlockerExists(vlt *vault.Vault, unlockerID string) er // Get vault directory to construct unlocker instances vaultDir, err := vlt.GetDirectory() if err != nil { - secret.Warn("Could not get vault directory during duplicate check", "error", err) + secret.Warn("Could not get vault directory during duplicate check", + "error", err) return nil } // Check each unlocker's ID + unlockersDir := filepath.Join(vaultDir, "unlockers.d") + for _, metadata := range unlockers { - // Construct the unlocker based on type to get its ID - unlockersDir := filepath.Join(vaultDir, "unlockers.d") - files, err := afero.ReadDir(cli.fs, unlockersDir) - if err != nil { - secret.Warn("Could not read unlockers directory during duplicate check", "error", err) - - continue - } - - for _, file := range files { - if !file.IsDir() { - continue - } - - unlockerDir := filepath.Join(unlockersDir, file.Name()) - metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") - - // Check if this matches our metadata - metadataBytes, err := afero.ReadFile(cli.fs, metadataPath) - if err != nil { - secret.Warn("Could not read unlocker metadata during duplicate check", "path", metadataPath, "error", err) - - continue - } - - var diskMetadata secret.UnlockerMetadata - if err := json.Unmarshal(metadataBytes, &diskMetadata); err != nil { - secret.Warn("Could not parse unlocker metadata during duplicate check", "path", metadataPath, "error", err) - - continue - } - - // Match by type and creation time - if diskMetadata.Type == metadata.Type && diskMetadata.CreatedAt.Equal(metadata.CreatedAt) { - var unlocker secret.Unlocker - switch metadata.Type { - case "passphrase": - unlocker = secret.NewPassphraseUnlocker(cli.fs, unlockerDir, diskMetadata) - case "keychain": - unlocker = secret.NewKeychainUnlocker(cli.fs, unlockerDir, diskMetadata) - case "pgp": - unlocker = secret.NewPGPUnlocker(cli.fs, unlockerDir, diskMetadata) - case "secure-enclave": - unlocker = secret.NewSecureEnclaveUnlocker(cli.fs, unlockerDir, diskMetadata) - } - - if unlocker != nil && unlocker.GetID() == unlockerID { - return fmt.Errorf("unlocker already exists") - } - - break - } + // Construct the unlocker matching this metadata to get its ID + id := findUnlockerIDByMetadata(cli.fs, unlockersDir, metadata, true) + if id != "" && id == unlockerID { + return errUnlockerExists } } diff --git a/internal/cli/vault.go b/internal/cli/vault.go index dcd54e0..63781d4 100644 --- a/internal/cli/vault.go +++ b/internal/cli/vault.go @@ -2,10 +2,12 @@ package cli import ( "encoding/json" + "errors" "fmt" "log" "os" "path/filepath" + "slices" "strings" "time" @@ -18,6 +20,22 @@ import ( "github.com/tyler-smith/go-bip39" ) +// Sentinel errors for vault operations +var ( + errMnemonicEmpty = errors.New("mnemonic cannot be empty") + errInvalidMnemonicPhrase = errors.New("invalid BIP39 mnemonic phrase") + errInvalidMnemonic = errors.New("invalid BIP39 mnemonic") + errVaultHasLongTermKey = errors.New( + "already has a long-term key configured") + errMnemonicEnvNotSet = errors.New( + "SB_SECRET_MNEMONIC environment variable not set") + errPassphraseEnvNotSet = errors.New( + "SB_UNLOCK_PASSPHRASE environment variable not set") + errCannotRemoveLastVault = errors.New("cannot remove the last vault") + errVaultContainsSecrets = errors.New( + "contains secrets; use --force to remove") +) + func newVaultCmd() *cobra.Command { cmd := &cobra.Command{ Use: "vault", @@ -36,7 +54,7 @@ func newVaultCmd() *cobra.Command { func newVaultListCmd() *cobra.Command { cmd := &cobra.Command{ - Use: "list", + Use: cmdUseList, Aliases: []string{"ls"}, Short: "List available vaults", RunE: func(cmd *cobra.Command, _ []string) error { @@ -101,9 +119,10 @@ func newVaultImportCmd() *cobra.Command { } return &cobra.Command{ - Use: "import ", - Short: "Import a mnemonic into a vault", - Long: `Import a BIP39 mnemonic phrase into the specified vault (default if not specified).`, + Use: "import ", + Short: "Import a mnemonic into a vault", + Long: `Import a BIP39 mnemonic phrase into the specified vault ` + + `(default if not specified).`, Args: cobra.MaximumNArgs(1), ValidArgsFunction: getVaultNamesCompletionFunc(cli.fs, cli.stateDir), RunE: func(cmd *cobra.Command, args []string) error { @@ -127,16 +146,19 @@ func newVaultRemoveCmd() *cobra.Command { if err != nil { log.Fatalf("failed to initialize CLI: %v", err) } + cmd := &cobra.Command{ Use: "remove ", Aliases: []string{"rm"}, Short: "Remove a vault", - Long: `Remove a vault. Requires --force if the vault contains secrets. Will automatically ` + - `switch to another vault if removing the currently selected one.`, + Long: `Remove a vault. Requires --force if the vault contains ` + + `secrets. Will automatically switch to another vault if ` + + `removing the currently selected one.`, Args: cobra.ExactArgs(1), ValidArgsFunction: getVaultNamesCompletionFunc(cli.fs, cli.stateDir), RunE: func(cmd *cobra.Command, args []string) error { force, _ := cmd.Flags().GetBool("force") + cli, err := NewCLIInstance() if err != nil { return fmt.Errorf("failed to initialize CLI: %w", err) @@ -161,11 +183,13 @@ func (cli *Instance) ListVaults(cmd *cobra.Command, jsonOutput bool) error { if jsonOutput { //nolint:nestif // Separate JSON and text output formatting logic // Get current vault name for context currentVault := "" - if currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir); err == nil { + + currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) + if err == nil { currentVault = currentVlt.GetName() } - result := map[string]interface{}{ + result := map[string]any{ "vaults": vaults, "currentVault": currentVault, } @@ -174,16 +198,20 @@ func (cli *Instance) ListVaults(cmd *cobra.Command, jsonOutput bool) error { if err != nil { return err } + cmd.Println(string(jsonBytes)) } else { // Text output cmd.Println("Available vaults:") + if len(vaults) == 0 { cmd.Println(" (none)") } else { // Try to get current vault for marking currentVault := "" - if currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir); err == nil { + + currentVlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) + if err == nil { currentVault = currentVlt.GetName() } @@ -200,19 +228,57 @@ func (cli *Instance) ListVaults(cmd *cobra.Command, jsonOutput bool) error { return nil } +// setMnemonicEnv sets the mnemonic environment variable and returns a +// function that restores the previous value +func setMnemonicEnv(mnemonicStr string) func() { + originalMnemonic := os.Getenv(secret.EnvMnemonic) + _ = os.Setenv(secret.EnvMnemonic, mnemonicStr) + + return func() { + if originalMnemonic != "" { + _ = os.Setenv(secret.EnvMnemonic, originalMnemonic) + } else { + _ = os.Unsetenv(secret.EnvMnemonic) + } + } +} + +// resolvePassphrase returns the unlock passphrase from the environment or +// prompts the user for it with confirmation +func resolvePassphrase() (*memguard.LockedBuffer, error) { + if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" { + secret.Debug("Using unlock passphrase from environment variable") + + return memguard.NewBufferFromBytes([]byte(envPassphrase)), nil + } + + secret.Debug("Prompting user for unlock passphrase") + + // Use secure passphrase input with confirmation + passphraseBuffer, err := readSecurePassphrase("Enter passphrase for unlocker: ") + if err != nil { + return nil, fmt.Errorf("failed to read passphrase: %w", err) + } + + return passphraseBuffer, nil +} + // CreateVault creates a new vault func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error { secret.Debug("Creating new vault", "name", name, "state_dir", cli.stateDir) // Get or prompt for mnemonic var mnemonicStr string + if envMnemonic := os.Getenv(secret.EnvMnemonic); envMnemonic != "" { secret.Debug("Using mnemonic from environment variable") + mnemonicStr = envMnemonic } else { secret.Debug("Prompting user for mnemonic phrase") // Read mnemonic securely without echo - mnemonicBuffer, err := secret.ReadPassphrase("Enter your BIP39 mnemonic phrase: ") + mnemonicBuffer, err := secret.ReadPassphrase( + "Enter your BIP39 mnemonic phrase: ") if err != nil { secret.Debug("Failed to read mnemonic from stdin", "error", err) @@ -221,30 +287,25 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error { defer mnemonicBuffer.Destroy() mnemonicStr = mnemonicBuffer.String() + fmt.Fprintln(os.Stderr) // Add newline after hidden input } if mnemonicStr == "" { - return fmt.Errorf("mnemonic cannot be empty") + return errMnemonicEmpty } // Validate the mnemonic mnemonicWords := strings.Fields(mnemonicStr) secret.Debug("Validating BIP39 mnemonic", "word_count", len(mnemonicWords)) + if !bip39.IsMnemonicValid(mnemonicStr) { - return fmt.Errorf("invalid BIP39 mnemonic phrase") + return errInvalidMnemonicPhrase } // Set mnemonic in environment for CreateVault to use - originalMnemonic := os.Getenv(secret.EnvMnemonic) - _ = os.Setenv(secret.EnvMnemonic, mnemonicStr) - defer func() { - if originalMnemonic != "" { - _ = os.Setenv(secret.EnvMnemonic, originalMnemonic) - } else { - _ = os.Unsetenv(secret.EnvMnemonic) - } - }() + restoreMnemonicEnv := setMnemonicEnv(mnemonicStr) + defer restoreMnemonicEnv() // Create the vault - it will handle key derivation internally vlt, err := vault.CreateVault(cli.fs, cli.stateDir, name) @@ -254,6 +315,7 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error { // Get the vault metadata to retrieve the derivation index vaultDir := filepath.Join(cli.stateDir, "vaults.d", name) + metadata, err := vault.LoadVaultMetadata(cli.fs, vaultDir) if err != nil { return fmt.Errorf("failed to load vault metadata: %w", err) @@ -269,22 +331,15 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error { vlt.Unlock(ltIdentity) // Get or prompt for passphrase - var passphraseBuffer *memguard.LockedBuffer - if envPassphrase := os.Getenv(secret.EnvUnlockPassphrase); envPassphrase != "" { - secret.Debug("Using unlock passphrase from environment variable") - passphraseBuffer = memguard.NewBufferFromBytes([]byte(envPassphrase)) - } else { - secret.Debug("Prompting user for unlock passphrase") - // Use secure passphrase input with confirmation - passphraseBuffer, err = readSecurePassphrase("Enter passphrase for unlocker: ") - if err != nil { - return fmt.Errorf("failed to read passphrase: %w", err) - } + passphraseBuffer, err := resolvePassphrase() + if err != nil { + return err } defer passphraseBuffer.Destroy() // Create passphrase-protected unlocker secret.Debug("Creating passphrase-protected unlocker") + passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer) if err != nil { return fmt.Errorf("failed to create unlocker: %w", err) @@ -299,7 +354,8 @@ func (cli *Instance) CreateVault(cmd *cobra.Command, name string) error { // SelectVault selects a vault as the current one func (cli *Instance) SelectVault(cmd *cobra.Command, name string) error { - if err := vault.SelectVault(cli.fs, cli.stateDir, name); err != nil { + err := vault.SelectVault(cli.fs, cli.stateDir, name) + if err != nil { return err } @@ -308,84 +364,60 @@ func (cli *Instance) SelectVault(cmd *cobra.Command, name string) error { return nil } -// VaultImport imports a mnemonic into a specific vault -func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error { - secret.Debug("Importing mnemonic into vault", "vault_name", vaultName, "state_dir", cli.stateDir) - - // Get the specific vault by name - vlt := vault.NewVault(cli.fs, cli.stateDir, vaultName) - +// vaultImportPreflight verifies the vault exists without a long-term key +// and returns the vault directory, public key path, and validated mnemonic +func (cli *Instance) vaultImportPreflight( + vlt *vault.Vault, vaultName string, +) (string, string, string, error) { // Check if vault exists vaultDir, err := vlt.GetDirectory() if err != nil { - return err + return "", "", "", err } exists, err := afero.DirExists(cli.fs, vaultDir) if err != nil { - return fmt.Errorf("failed to check if vault exists: %w", err) + return "", "", "", fmt.Errorf("failed to check if vault exists: %w", err) } + if !exists { - return fmt.Errorf("vault '%s' does not exist", vaultName) + return "", "", "", fmt.Errorf("vault '%s' %w", + vaultName, errVaultDoesNotExist) } // Check if vault already has a public key - pubKeyPath := fmt.Sprintf("%s/pub.age", vaultDir) - if _, err := cli.fs.Stat(pubKeyPath); err == nil { - return fmt.Errorf("vault '%s' already has a long-term key configured", vaultName) + pubKeyPath := vaultDir + "/pub.age" + + _, err = cli.fs.Stat(pubKeyPath) + if err == nil { + return "", "", "", fmt.Errorf("vault '%s' %w", + vaultName, errVaultHasLongTermKey) } // Get mnemonic from environment mnemonic := os.Getenv(secret.EnvMnemonic) if mnemonic == "" { - return fmt.Errorf("SB_SECRET_MNEMONIC environment variable not set") + return "", "", "", errMnemonicEnvNotSet } // Validate the mnemonic mnemonicWords := strings.Fields(mnemonic) secret.Debug("Validating BIP39 mnemonic", "word_count", len(mnemonicWords)) + if !bip39.IsMnemonicValid(mnemonic) { - return fmt.Errorf("invalid BIP39 mnemonic") + return "", "", "", errInvalidMnemonic } - // Get the next available derivation index for this mnemonic - derivationIndex, err := vault.GetNextDerivationIndex(cli.fs, cli.stateDir, mnemonic) - if err != nil { - secret.Debug("Failed to get next derivation index", "error", err) - - return fmt.Errorf("failed to get next derivation index: %w", err) - } - secret.Debug("Using derivation index", "index", derivationIndex) - - // Derive long-term key from mnemonic with the appropriate index - secret.Debug("Deriving long-term key from mnemonic", "index", derivationIndex) - ltIdentity, err := agehd.DeriveIdentity(mnemonic, derivationIndex) - if err != nil { - return fmt.Errorf("failed to derive long-term key: %w", err) - } - - // Store long-term public key in vault - ltPublicKey := ltIdentity.Recipient().String() - secret.Debug("Storing long-term public key", "pubkey", ltPublicKey, "vault_dir", vaultDir) - - if err := afero.WriteFile(cli.fs, pubKeyPath, []byte(ltPublicKey), secret.FilePerms); err != nil { - return fmt.Errorf("failed to store long-term public key: %w", err) - } - - // Calculate public key hash from the actual derivation index being used - // This is used to verify that the derived key matches what was stored - publicKeyHash := vault.ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String())) - - // Calculate 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 - identity0, err := agehd.DeriveIdentity(mnemonic, 0) - if err != nil { - return fmt.Errorf("failed to derive identity for index 0: %w", err) - } - familyHash := vault.ComputeDoubleSHA256([]byte(identity0.Recipient().String())) + return vaultDir, pubKeyPath, mnemonic, nil +} +// updateVaultImportMetadata stores the derivation info in vault metadata +func updateVaultImportMetadata( + fs afero.Fs, vaultDir string, derivationIndex uint32, + publicKeyHash, familyHash string, +) error { // Load existing metadata - existingMetadata, err := vault.LoadVaultMetadata(cli.fs, vaultDir) + existingMetadata, err := vault.LoadVaultMetadata(fs, vaultDir) if err != nil { // If metadata doesn't exist, create new existingMetadata = &vault.Metadata{ @@ -398,17 +430,83 @@ func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error { existingMetadata.PublicKeyHash = publicKeyHash existingMetadata.MnemonicFamilyHash = familyHash - if err := vault.SaveVaultMetadata(cli.fs, vaultDir, existingMetadata); err != nil { + err = vault.SaveVaultMetadata(fs, vaultDir, existingMetadata) + if err != nil { secret.Debug("Failed to save vault metadata", "error", err) return fmt.Errorf("failed to save vault metadata: %w", err) } + secret.Debug("Saved vault metadata with derivation index and public key hash") + return nil +} + +// VaultImport imports a mnemonic into a specific vault +func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error { + secret.Debug("Importing mnemonic into vault", + "vault_name", vaultName, "state_dir", cli.stateDir) + + // Get the specific vault by name + vlt := vault.NewVault(cli.fs, cli.stateDir, vaultName) + + vaultDir, pubKeyPath, mnemonic, err := cli.vaultImportPreflight(vlt, vaultName) + if err != nil { + return err + } + + // Get the next available derivation index for this mnemonic + derivationIndex, err := vault.GetNextDerivationIndex(cli.fs, cli.stateDir, mnemonic) + if err != nil { + secret.Debug("Failed to get next derivation index", "error", err) + + return fmt.Errorf("failed to get next derivation index: %w", err) + } + + secret.Debug("Using derivation index", "index", derivationIndex) + + // Derive long-term key from mnemonic with the appropriate index + secret.Debug("Deriving long-term key from mnemonic", "index", derivationIndex) + + ltIdentity, err := agehd.DeriveIdentity(mnemonic, derivationIndex) + if err != nil { + return fmt.Errorf("failed to derive long-term key: %w", err) + } + + // Store long-term public key in vault + ltPublicKey := ltIdentity.Recipient().String() + secret.Debug("Storing long-term public key", + "pubkey", ltPublicKey, "vault_dir", vaultDir) + + err = afero.WriteFile(cli.fs, pubKeyPath, []byte(ltPublicKey), secret.FilePerms) + if err != nil { + return fmt.Errorf("failed to store long-term public key: %w", err) + } + + // Calculate public key hash from the actual derivation index being used + // This is used to verify that the derived key matches what was stored + publicKeyHash := vault.ComputeDoubleSHA256([]byte(ltIdentity.Recipient().String())) + + // Calculate 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. + identity0, err := agehd.DeriveIdentity(mnemonic, 0) + if err != nil { + return fmt.Errorf("failed to derive identity for index 0: %w", err) + } + + familyHash := vault.ComputeDoubleSHA256([]byte(identity0.Recipient().String())) + + err = updateVaultImportMetadata( + cli.fs, vaultDir, derivationIndex, publicKeyHash, familyHash) + if err != nil { + return err + } + // Get passphrase from environment variable passphraseStr := os.Getenv(secret.EnvUnlockPassphrase) if passphraseStr == "" { - return fmt.Errorf("SB_UNLOCK_PASSPHRASE environment variable not set") + return errPassphraseEnvNotSet } secret.Debug("Using unlock passphrase from environment variable") @@ -422,6 +520,7 @@ func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error { // Create passphrase-protected unlocker secret.Debug("Creating passphrase-protected unlocker") + passphraseUnlocker, err := vlt.CreatePassphraseUnlocker(passphraseBuffer) if err != nil { secret.Debug("Failed to create unlocker", "error", err) @@ -436,6 +535,46 @@ func (cli *Instance) VaultImport(cmd *cobra.Command, vaultName string) error { return nil } +// vaultHasSecrets reports whether the vault directory contains any secrets +func (cli *Instance) vaultHasSecrets(vaultDir string) bool { + secretsDir := filepath.Join(vaultDir, "secrets.d") + + exists, _ := afero.DirExists(cli.fs, secretsDir) + if !exists { + return false + } + + entries, err := afero.ReadDir(cli.fs, secretsDir) + + return err == nil && len(entries) > 0 +} + +// switchAwayFromVault selects another vault as current before removal +func (cli *Instance) switchAwayFromVault( + cmd *cobra.Command, vaults []string, name string, +) error { + // Find another vault to switch to + var newVault string + + for _, v := range vaults { + if v != name { + newVault = v + + break + } + } + + // Switch to the new vault + err := vault.SelectVault(cli.fs, cli.stateDir, newVault) + if err != nil { + return fmt.Errorf("failed to switch to vault '%s': %w", newVault, err) + } + + cmd.Printf("Switched current vault to '%s'\n", newVault) + + return nil +} + // RemoveVault removes a vault with safety checks func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) error { // Get list of all vaults @@ -445,21 +584,13 @@ func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) er } // Check if vault exists - vaultExists := false - for _, v := range vaults { - if v == name { - vaultExists = true - - break - } - } - if !vaultExists { - return fmt.Errorf("vault '%s' does not exist", name) + if !slices.Contains(vaults, name) { + return fmt.Errorf("vault '%s' %w", name, errVaultDoesNotExist) } // Don't allow removing the last vault if len(vaults) == 1 { - return fmt.Errorf("cannot remove the last vault") + return errCannotRemoveLastVault } // Check if this is the current vault @@ -467,57 +598,44 @@ func (cli *Instance) RemoveVault(cmd *cobra.Command, name string, force bool) er if err != nil { return fmt.Errorf("failed to get current vault: %w", err) } + isCurrentVault := currentVault.GetName() == name // Load the vault to check for secrets vlt := vault.NewVault(cli.fs, cli.stateDir, name) + vaultDir, err := vlt.GetDirectory() if err != nil { return fmt.Errorf("failed to get vault directory: %w", err) } // Check if vault has secrets - secretsDir := filepath.Join(vaultDir, "secrets.d") - hasSecrets := false - if exists, _ := afero.DirExists(cli.fs, secretsDir); exists { - entries, err := afero.ReadDir(cli.fs, secretsDir) - if err == nil && len(entries) > 0 { - hasSecrets = true - } - } + hasSecrets := cli.vaultHasSecrets(vaultDir) // Require --force if vault has secrets if hasSecrets && !force { - return fmt.Errorf("vault '%s' contains secrets; use --force to remove", name) + return fmt.Errorf("vault '%s' %w", name, errVaultContainsSecrets) } // If removing current vault, switch to another vault first if isCurrentVault { - // Find another vault to switch to - var newVault string - for _, v := range vaults { - if v != name { - newVault = v - - break - } + err = cli.switchAwayFromVault(cmd, vaults, name) + if err != nil { + return err } - - // Switch to the new vault - if err := vault.SelectVault(cli.fs, cli.stateDir, newVault); err != nil { - return fmt.Errorf("failed to switch to vault '%s': %w", newVault, err) - } - cmd.Printf("Switched current vault to '%s'\n", newVault) } // Remove the vault directory - if err := cli.fs.RemoveAll(vaultDir); err != nil { + err = cli.fs.RemoveAll(vaultDir) + if err != nil { return fmt.Errorf("failed to remove vault directory: %w", err) } cmd.Printf("Removed vault '%s'\n", name) + if hasSecrets { - cmd.Printf("Warning: Vault contained secrets that have been permanently deleted\n") + cmd.Printf("Warning: Vault contained secrets that have been " + + "permanently deleted\n") } return nil diff --git a/internal/cli/version.go b/internal/cli/version.go index a9ded8b..17f113f 100644 --- a/internal/cli/version.go +++ b/internal/cli/version.go @@ -1,12 +1,16 @@ package cli import ( + "errors" "fmt" + "io" "log" "path/filepath" "strings" "text/tabwriter" + "time" + "filippo.io/age" "git.eeqj.de/sneak/secret/internal/secret" "git.eeqj.de/sneak/secret/internal/vault" "github.com/spf13/afero" @@ -17,6 +21,12 @@ const ( tabWriterPadding = 2 ) +// Sentinel errors for version operations +var ( + errVersionNotFound = errors.New("not found for secret") + errCannotRemoveCurrentVersion = errors.New("promote another version first") +) + // newVersionCmd returns the version management command func newVersionCmd() *cobra.Command { cli, err := NewCLIInstance() @@ -32,7 +42,8 @@ func VersionCommands(cli *Instance) *cobra.Command { versionCmd := &cobra.Command{ Use: "version", Short: "Manage secret versions", - Long: "Commands for managing secret versions including listing, promoting, and retrieving specific versions", + Long: "Commands for managing secret versions including listing, " + + "promoting, and retrieving specific versions", } // List versions command @@ -51,14 +62,17 @@ func VersionCommands(cli *Instance) *cobra.Command { promoteCmd := &cobra.Command{ Use: "promote ", Short: "Promote a specific version to current", - Long: "Updates the current symlink to point to the specified version without modifying timestamps", - Args: cobra.ExactArgs(2), //nolint:mnd // Command requires exactly 2 arguments: secret-name and version - ValidArgsFunction: func(cmd *cobra.Command, args []string, toComplete string) ([]string, cobra.ShellCompDirective) { + Long: "Updates the current symlink to point to the specified " + + "version without modifying timestamps", + Args: cobra.ExactArgs(2), //nolint:mnd // secret-name and version args + ValidArgsFunction: func( + cmd *cobra.Command, args []string, toComplete string, + ) ([]string, cobra.ShellCompDirective) { // Complete secret name for first arg if len(args) == 0 { return getSecretNamesCompletionFunc(cli.fs, cli.stateDir)(cmd, args, toComplete) } - // TODO: Complete version numbers for second arg + // Version number completion for the second arg is not implemented return nil, cobra.ShellCompDirectiveNoFileComp }, RunE: func(cmd *cobra.Command, args []string) error { @@ -71,14 +85,17 @@ func VersionCommands(cli *Instance) *cobra.Command { Use: "remove ", Aliases: []string{"rm"}, Short: "Remove a specific version of a secret", - Long: "Remove a specific version of a secret. Cannot remove the current version.", - Args: cobra.ExactArgs(2), //nolint:mnd // Command requires exactly 2 arguments: secret-name and version - ValidArgsFunction: func(cmd *cobra.Command, args []string, toComplete string) ([]string, cobra.ShellCompDirective) { + Long: "Remove a specific version of a secret. Cannot remove the " + + "current version.", + Args: cobra.ExactArgs(2), //nolint:mnd // secret-name and version args + ValidArgsFunction: func( + cmd *cobra.Command, args []string, toComplete string, + ) ([]string, cobra.ShellCompDirective) { // Complete secret name for first arg if len(args) == 0 { return getSecretNamesCompletionFunc(cli.fs, cli.stateDir)(cmd, args, toComplete) } - // TODO: Complete version numbers for second arg + // Version number completion for the second arg is not implemented return nil, cobra.ShellCompDirectiveNoFileComp }, RunE: func(cmd *cobra.Command, args []string) error { @@ -121,10 +138,11 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error { return fmt.Errorf("failed to check if secret exists: %w", err) } + if !exists { secret.Debug("Secret not found", "secret_name", secretName) - return fmt.Errorf("secret '%s' not found", secretName) + return fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound) } // List all versions @@ -145,6 +163,7 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error { currentVersion, err := secret.GetCurrentVersion(cli.fs, secretDir) if err != nil { secret.Debug("Failed to get current version", "error", err) + currentVersion = "" } @@ -160,44 +179,7 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error { // Load and display each version's metadata for _, version := range versions { - sv := secret.NewVersion(vlt, secretName, version) - - // Load metadata - if err := sv.LoadMetadata(ltIdentity); err != nil { - secret.Warn("Failed to load version metadata", "version", version, "error", err) - // Display version with error - status := "error" - if version == currentVersion { - status = "current (error)" - } - _, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n", version, "-", status, "-", "-") - - continue - } - - // Determine status - status := "expired" - if version == currentVersion { - status = "current" - } - - // Format timestamps - createdAt := "-" - if sv.Metadata.CreatedAt != nil { - createdAt = sv.Metadata.CreatedAt.Format("2006-01-02 15:04:05") - } - - notBefore := "-" - if sv.Metadata.NotBefore != nil { - notBefore = sv.Metadata.NotBefore.Format("2006-01-02 15:04:05") - } - - notAfter := "-" - if sv.Metadata.NotAfter != nil { - notAfter = sv.Metadata.NotAfter.Format("2006-01-02 15:04:05") - } - - _, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n", version, createdAt, status, notBefore, notAfter) + printVersionRow(w, vlt, secretName, version, currentVersion, ltIdentity) } _ = w.Flush() @@ -205,8 +187,58 @@ func (cli *Instance) ListVersions(cmd *cobra.Command, secretName string) error { return nil } +// printVersionRow loads one version's metadata and writes its table row +func printVersionRow( + w io.Writer, vlt *vault.Vault, + secretName, version, currentVersion string, + ltIdentity *age.X25519Identity, +) { + sv := secret.NewVersion(vlt, secretName, version) + + // Load metadata + err := sv.LoadMetadata(ltIdentity) + if err != nil { + secret.Warn("Failed to load version metadata", + "version", version, "error", err) + // Display version with error + status := "error" + if version == currentVersion { + status = "current (error)" + } + + _, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n", version, "-", status, "-", "-") + + return + } + + // Determine status + status := "expired" + if version == currentVersion { + status = "current" + } + + // Format timestamps + createdAt := formatVersionTime(sv.Metadata.CreatedAt) + notBefore := formatVersionTime(sv.Metadata.NotBefore) + notAfter := formatVersionTime(sv.Metadata.NotAfter) + + _, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n", + version, createdAt, status, notBefore, notAfter) +} + +// formatVersionTime formats an optional version timestamp, "-" when unset +func formatVersionTime(t *time.Time) string { + if t == nil { + return "-" + } + + return t.Format("2006-01-02 15:04:05") +} + // PromoteVersion promotes a specific version to current -func (cli *Instance) PromoteVersion(cmd *cobra.Command, secretName string, version string) error { +func (cli *Instance) PromoteVersion( + cmd *cobra.Command, secretName string, version string, +) error { // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -224,16 +256,20 @@ func (cli *Instance) PromoteVersion(cmd *cobra.Command, secretName string, versi // Check if version exists versionDir := filepath.Join(secretDir, "versions", version) + exists, err := afero.DirExists(cli.fs, versionDir) if err != nil { return fmt.Errorf("failed to check if version exists: %w", err) } + if !exists { - return fmt.Errorf("version '%s' not found for secret '%s'", version, secretName) + return fmt.Errorf("version '%s' %w '%s'", + version, errVersionNotFound, secretName) } // Update the current symlink using the proper function - if err := secret.SetCurrentVersion(cli.fs, secretDir, version); err != nil { + err = secret.SetCurrentVersion(cli.fs, secretDir, version) + if err != nil { return fmt.Errorf("failed to update current version: %w", err) } @@ -243,7 +279,9 @@ func (cli *Instance) PromoteVersion(cmd *cobra.Command, secretName string, versi } // RemoveVersion removes a specific version of a secret -func (cli *Instance) RemoveVersion(cmd *cobra.Command, secretName string, version string) error { +func (cli *Instance) RemoveVersion( + cmd *cobra.Command, secretName string, version string, +) error { // Get current vault vlt, err := vault.GetCurrentVault(cli.fs, cli.stateDir) if err != nil { @@ -264,18 +302,22 @@ func (cli *Instance) RemoveVersion(cmd *cobra.Command, secretName string, versio if err != nil { return fmt.Errorf("failed to check if secret exists: %w", err) } + if !exists { - return fmt.Errorf("secret '%s' not found", secretName) + return fmt.Errorf("secret '%s' %w", secretName, errSecretNotFound) } // Check if version exists versionDir := filepath.Join(secretDir, "versions", version) + exists, err = afero.DirExists(cli.fs, versionDir) if err != nil { return fmt.Errorf("failed to check if version exists: %w", err) } + if !exists { - return fmt.Errorf("version '%s' not found for secret '%s'", version, secretName) + return fmt.Errorf("version '%s' %w '%s'", + version, errVersionNotFound, secretName) } // Get current version @@ -286,11 +328,13 @@ func (cli *Instance) RemoveVersion(cmd *cobra.Command, secretName string, versio // Don't allow removing the current version if version == currentVersion { - return fmt.Errorf("cannot remove the current version '%s'; promote another version first", version) + return fmt.Errorf("cannot remove the current version '%s'; %w", + version, errCannotRemoveCurrentVersion) } // Remove the version directory - if err := cli.fs.RemoveAll(versionDir); err != nil { + err = cli.fs.RemoveAll(versionDir) + if err != nil { return fmt.Errorf("failed to remove version: %w", err) } diff --git a/internal/cli/version_test.go b/internal/cli/version_test.go index 2bebc6b..79b9e44 100644 --- a/internal/cli/version_test.go +++ b/internal/cli/version_test.go @@ -14,6 +14,7 @@ // - setupTestVault(): CLI test helper for vault initialization // - Uses consistent test mnemonic for reproducible testing +//nolint:testpackage // white-box test of unexported internals package cli import ( @@ -32,29 +33,41 @@ import ( "github.com/stretchr/testify/require" ) -// Helper function to add a secret to vault with proper buffer protection -func addTestSecret(t *testing.T, vlt *vault.Vault, name string, value []byte, force bool) { +const ( + // testMnemonic is the standard BIP39 mnemonic used for CLI tests. + //nolint:dupword // BIP39 test mnemonic intentionally repeats a word + testMnemonic = "abandon abandon abandon abandon abandon abandon " + + "abandon abandon abandon abandon abandon about" + + // testStateDir is the in-memory state directory used by CLI tests. + testStateDir = "/test/state" +) + +// Helper function to add a version of the "test/secret" secret to the +// vault with proper buffer protection +func addTestSecret(t *testing.T, vlt *vault.Vault, value []byte, force bool) { t.Helper() + buffer := memguard.NewBufferFromBytes(value) defer buffer.Destroy() - err := vlt.AddSecret(name, buffer, force) + + err := vlt.AddSecret("test/secret", buffer, force) require.NoError(t, err) } -// Helper function to set up a vault with long-term key -func setupTestVault(t *testing.T, fs afero.Fs, stateDir string) { +// Helper function to set up a vault with long-term key in testStateDir +func setupTestVault(t *testing.T, fs afero.Fs) { + t.Helper() + // Set mnemonic for testing - testMnemonic := "abandon abandon abandon abandon abandon abandon " + - "abandon abandon abandon abandon abandon about" t.Setenv(secret.EnvMnemonic, testMnemonic) // Create vault - vlt, err := vault.CreateVault(fs, stateDir, "default") + vlt, err := vault.CreateVault(fs, testStateDir, "default") 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 @@ -64,30 +77,32 @@ func setupTestVault(t *testing.T, fs afero.Fs, stateDir string) { require.NoError(t, err) // Select vault - err = vault.SelectVault(fs, stateDir, "default") + err = vault.SelectVault(fs, testStateDir, "default") require.NoError(t, err) } +//nolint:paralleltest // uses t.Setenv via setupTestVault func TestListVersionsCommand(t *testing.T) { fs := afero.NewMemMapFs() - stateDir := "/test/state" + stateDir := testStateDir cli := NewCLIInstanceWithStateDir(fs, stateDir) // Set up vault with long-term key - setupTestVault(t, fs, stateDir) + setupTestVault(t, fs) // Add a secret with multiple versions vlt, err := vault.GetCurrentVault(fs, stateDir) require.NoError(t, err) - addTestSecret(t, vlt, "test/secret", []byte("version-1"), false) + addTestSecret(t, vlt, []byte("version-1"), false) time.Sleep(10 * time.Millisecond) - addTestSecret(t, vlt, "test/secret", []byte("version-2"), true) + addTestSecret(t, vlt, []byte("version-2"), true) // Create a command for output capture cmd := newRootCmd() + var buf bytes.Buffer cmd.SetOut(&buf) cmd.SetErr(&buf) @@ -112,24 +127,28 @@ func TestListVersionsCommand(t *testing.T) { // Should have two version entries lines := strings.Split(outputStr, "\n") versionLines := 0 + for _, line := range lines { if strings.Contains(line, ".001") || strings.Contains(line, ".002") { versionLines++ } } + assert.Equal(t, 2, versionLines) } +//nolint:paralleltest // uses t.Setenv via setupTestVault func TestListVersionsNonExistentSecret(t *testing.T) { fs := afero.NewMemMapFs() - stateDir := "/test/state" + stateDir := testStateDir cli := NewCLIInstanceWithStateDir(fs, stateDir) // Set up vault with long-term key - setupTestVault(t, fs, stateDir) + setupTestVault(t, fs) // Create a command for output capture cmd := newRootCmd() + var buf bytes.Buffer cmd.SetOut(&buf) cmd.SetErr(&buf) @@ -140,23 +159,24 @@ func TestListVersionsNonExistentSecret(t *testing.T) { assert.Contains(t, err.Error(), "not found") } +//nolint:paralleltest // uses t.Setenv via setupTestVault func TestPromoteVersionCommand(t *testing.T) { fs := afero.NewMemMapFs() - stateDir := "/test/state" + stateDir := testStateDir cli := NewCLIInstanceWithStateDir(fs, stateDir) // Set up vault with long-term key - setupTestVault(t, fs, stateDir) + setupTestVault(t, fs) // Add a secret with multiple versions vlt, err := vault.GetCurrentVault(fs, stateDir) require.NoError(t, err) - addTestSecret(t, vlt, "test/secret", []byte("version-1"), false) + addTestSecret(t, vlt, []byte("version-1"), false) time.Sleep(10 * time.Millisecond) - addTestSecret(t, vlt, "test/secret", []byte("version-2"), true) + addTestSecret(t, vlt, []byte("version-2"), true) // Get versions vaultDir, _ := vlt.GetDirectory() @@ -175,6 +195,7 @@ func TestPromoteVersionCommand(t *testing.T) { // Create a command for output capture cmd := newRootCmd() + var buf bytes.Buffer cmd.SetOut(&buf) cmd.SetErr(&buf) @@ -195,22 +216,24 @@ func TestPromoteVersionCommand(t *testing.T) { assert.Equal(t, []byte("version-1"), value) } +//nolint:paralleltest // uses t.Setenv via setupTestVault func TestPromoteNonExistentVersion(t *testing.T) { fs := afero.NewMemMapFs() - stateDir := "/test/state" + stateDir := testStateDir cli := NewCLIInstanceWithStateDir(fs, stateDir) // Set up vault with long-term key - setupTestVault(t, fs, stateDir) + setupTestVault(t, fs) // Add a secret vlt, err := vault.GetCurrentVault(fs, stateDir) require.NoError(t, err) - addTestSecret(t, vlt, "test/secret", []byte("value"), false) + addTestSecret(t, vlt, []byte("value"), false) // Create a command for output capture cmd := newRootCmd() + var buf bytes.Buffer cmd.SetOut(&buf) cmd.SetErr(&buf) @@ -221,23 +244,24 @@ func TestPromoteNonExistentVersion(t *testing.T) { assert.Contains(t, err.Error(), "not found") } +//nolint:paralleltest // uses t.Setenv via setupTestVault func TestGetSecretWithVersion(t *testing.T) { fs := afero.NewMemMapFs() - stateDir := "/test/state" + stateDir := testStateDir cli := NewCLIInstanceWithStateDir(fs, stateDir) // Set up vault with long-term key - setupTestVault(t, fs, stateDir) + setupTestVault(t, fs) // Add a secret with multiple versions vlt, err := vault.GetCurrentVault(fs, stateDir) require.NoError(t, err) - addTestSecret(t, vlt, "test/secret", []byte("version-1"), false) + addTestSecret(t, vlt, []byte("version-1"), false) time.Sleep(10 * time.Millisecond) - addTestSecret(t, vlt, "test/secret", []byte("version-2"), true) + addTestSecret(t, vlt, []byte("version-2"), true) // Get versions vaultDir, _ := vlt.GetDirectory() @@ -248,6 +272,7 @@ func TestGetSecretWithVersion(t *testing.T) { // Create a command for output capture cmd := newRootCmd() + var buf bytes.Buffer cmd.SetOut(&buf) @@ -258,18 +283,21 @@ func TestGetSecretWithVersion(t *testing.T) { // Test getting specific version buf.Reset() + firstVersion := versions[1] // Older version err = cli.GetSecretWithVersion(cmd, "test/secret", firstVersion) require.NoError(t, err) assert.Equal(t, "version-1", buf.String()) } +//nolint:paralleltest // reads process environment to determine the state dir func TestVersionCommandStructure(t *testing.T) { // Test that version commands are properly structured cli, err := NewCLIInstance() if err != nil { t.Fatalf("failed to initialize CLI: %v", err) } + cmd := VersionCommands(cli) assert.Equal(t, "version", cmd.Use) @@ -285,13 +313,14 @@ func TestVersionCommandStructure(t *testing.T) { assert.Equal(t, "Promote a specific version to current", promoteCmd.Short) } +//nolint:paralleltest // uses t.Setenv via setupTestVault func TestListVersionsEmptyOutput(t *testing.T) { fs := afero.NewMemMapFs() - stateDir := "/test/state" + stateDir := testStateDir cli := NewCLIInstanceWithStateDir(fs, stateDir) // Set up vault with long-term key - setupTestVault(t, fs, stateDir) + setupTestVault(t, fs) // Create a secret directory without versions (edge case) vaultDir := stateDir + "/vaults.d/default" @@ -301,6 +330,7 @@ func TestListVersionsEmptyOutput(t *testing.T) { // Create a command for output capture cmd := newRootCmd() + var buf bytes.Buffer cmd.SetOut(&buf) cmd.SetErr(&buf) diff --git a/internal/macse/macse_stub.go b/internal/macse/macse_stub.go index 44fe611..90fb917 100644 --- a/internal/macse/macse_stub.go +++ b/internal/macse/macse_stub.go @@ -1,12 +1,11 @@ //go:build !darwin -// +build !darwin // Package macse provides Go bindings for macOS Secure Enclave operations. package macse -import "fmt" +import "errors" -var errNotSupported = fmt.Errorf("secure enclave is only supported on macOS") //nolint:gochecknoglobals +var errNotSupported = errors.New("secure enclave is only supported on macOS") // CreateKey is not supported on non-darwin platforms. func CreateKey(_ string) ([]byte, string, error) { diff --git a/internal/secret/constants.go b/internal/secret/constants.go index 7e82a80..8c55459 100644 --- a/internal/secret/constants.go +++ b/internal/secret/constants.go @@ -12,7 +12,8 @@ const ( // EnvMnemonic is the environment variable for providing the mnemonic phrase EnvMnemonic = "SB_SECRET_MNEMONIC" // EnvUnlockPassphrase is the environment variable for providing the unlock passphrase - EnvUnlockPassphrase = "SB_UNLOCK_PASSPHRASE" //nolint:gosec // G101: This is an env var name, not a credential + //nolint:gosec // G101: env var name, not a credential + EnvUnlockPassphrase = "SB_UNLOCK_PASSPHRASE" // EnvGPGKeyID is the environment variable for providing the GPG key ID EnvGPGKeyID = "SB_GPG_KEY_ID" ) diff --git a/internal/secret/crypto.go b/internal/secret/crypto.go index 1bf20a3..ad46ab5 100644 --- a/internal/secret/crypto.go +++ b/internal/secret/crypto.go @@ -2,6 +2,7 @@ package secret import ( "bytes" + "errors" "fmt" "io" "os" @@ -12,39 +13,61 @@ import ( "golang.org/x/term" ) +var ( + errNilPassphraseBuffer = errors.New("passphrase buffer is nil") + errStdinNotTerminal = errors.New( + "cannot read passphrase from non-terminal stdin " + + "(piped input or script). Please set the SB_UNLOCK_PASSPHRASE " + + "environment variable or run interactively") + errStderrNotTerminal = errors.New( + "cannot prompt for passphrase: stderr is not a terminal " + + "(running in non-interactive mode). Please set the " + + "SB_UNLOCK_PASSPHRASE environment variable") + errEmptyPassphrase = errors.New("passphrase cannot be empty") +) + // EncryptToRecipient encrypts data to a recipient using age // The data parameter should be a LockedBuffer for secure memory handling -func EncryptToRecipient(data *memguard.LockedBuffer, recipient age.Recipient) ([]byte, error) { +func EncryptToRecipient( + data *memguard.LockedBuffer, recipient age.Recipient, +) ([]byte, error) { if data == nil { - return nil, fmt.Errorf("data buffer is nil") + return nil, errNilDataBuffer } Debug("EncryptToRecipient starting", "data_length", data.Size()) var buf bytes.Buffer + Debug("Creating age encryptor") + w, err := age.Encrypt(&buf, recipient) if err != nil { Debug("Failed to create encryptor", "error", err) return nil, fmt.Errorf("failed to create encryptor: %w", err) } - Debug("Created age encryptor successfully") + Debug("Created age encryptor successfully") Debug("Writing data to encryptor") - if _, err := w.Write(data.Bytes()); err != nil { + + _, err = w.Write(data.Bytes()) + if err != nil { Debug("Failed to write data to encryptor", "error", err) return nil, fmt.Errorf("failed to write data: %w", err) } - Debug("Wrote data to encryptor successfully") + Debug("Wrote data to encryptor successfully") Debug("Closing encryptor") - if err := w.Close(); err != nil { + + err = w.Close() + if err != nil { Debug("Failed to close encryptor", "error", err) return nil, fmt.Errorf("failed to close encryptor: %w", err) } + Debug("Closed encryptor successfully") result := buf.Bytes() @@ -54,7 +77,9 @@ func EncryptToRecipient(data *memguard.LockedBuffer, recipient age.Recipient) ([ } // DecryptWithIdentity decrypts data with an identity using age -func DecryptWithIdentity(data []byte, identity age.Identity) (*memguard.LockedBuffer, error) { +func DecryptWithIdentity( + data []byte, identity age.Identity, +) (*memguard.LockedBuffer, error) { r, err := age.Decrypt(bytes.NewReader(data), identity) if err != nil { return nil, fmt.Errorf("failed to create decryptor: %w", err) @@ -68,7 +93,8 @@ func DecryptWithIdentity(data []byte, identity age.Identity) (*memguard.LockedBu // Create a secure buffer for the decrypted data resultBuffer := memguard.NewBufferFromBytes(result) - // Zero out the original slice to prevent plaintext from lingering in unprotected memory + // Zero out the original slice to prevent plaintext from lingering + // in unprotected memory for i := range result { result[i] = 0 } @@ -76,17 +102,22 @@ func DecryptWithIdentity(data []byte, identity age.Identity) (*memguard.LockedBu return resultBuffer, nil } -// EncryptWithPassphrase encrypts data using a passphrase with age's scrypt-based encryption -// Both data and passphrase parameters should be LockedBuffers for secure memory handling -func EncryptWithPassphrase(data *memguard.LockedBuffer, passphrase *memguard.LockedBuffer) ([]byte, error) { +// EncryptWithPassphrase encrypts data using a passphrase with age's +// scrypt-based encryption. Both data and passphrase parameters should +// be LockedBuffers for secure memory handling +func EncryptWithPassphrase( + data *memguard.LockedBuffer, passphrase *memguard.LockedBuffer, +) ([]byte, error) { if data == nil { - return nil, fmt.Errorf("data buffer is nil") - } - if passphrase == nil { - return nil, fmt.Errorf("passphrase buffer is nil") + return nil, errNilDataBuffer } - // Create recipient directly from passphrase - unavoidable string conversion due to age API + if passphrase == nil { + return nil, errNilPassphraseBuffer + } + + // Create recipient directly from passphrase - unavoidable string + // conversion due to age API recipient, err := age.NewScryptRecipient(passphrase.String()) if err != nil { return nil, fmt.Errorf("failed to create scrypt recipient: %w", err) @@ -95,14 +126,18 @@ func EncryptWithPassphrase(data *memguard.LockedBuffer, passphrase *memguard.Loc return EncryptToRecipient(data, recipient) } -// DecryptWithPassphrase decrypts data using a passphrase with age's scrypt-based decryption -// The passphrase parameter should be a LockedBuffer for secure memory handling -func DecryptWithPassphrase(encryptedData []byte, passphrase *memguard.LockedBuffer) (*memguard.LockedBuffer, error) { +// DecryptWithPassphrase decrypts data using a passphrase with age's +// scrypt-based decryption. The passphrase parameter should be a +// LockedBuffer for secure memory handling +func DecryptWithPassphrase( + encryptedData []byte, passphrase *memguard.LockedBuffer, +) (*memguard.LockedBuffer, error) { if passphrase == nil { - return nil, fmt.Errorf("passphrase buffer is nil") + return nil, errNilPassphraseBuffer } - // Create identity directly from passphrase - unavoidable string conversion due to age API + // Create identity directly from passphrase - unavoidable string + // conversion due to age API identity, err := age.NewScryptIdentity(passphrase.String()) if err != nil { return nil, fmt.Errorf("failed to create scrypt identity: %w", err) @@ -117,29 +152,30 @@ func DecryptWithPassphrase(encryptedData []byte, passphrase *memguard.LockedBuff func ReadPassphrase(prompt string) (*memguard.LockedBuffer, error) { // Check if stdin is a terminal if !term.IsTerminal(syscall.Stdin) { - // Not a terminal - never read passphrases from piped input for security reasons - return nil, fmt.Errorf("cannot read passphrase from non-terminal stdin " + - "(piped input or script). Please set the SB_UNLOCK_PASSPHRASE " + - "environment variable or run interactively") + // Not a terminal - never read passphrases from piped input + // for security reasons + return nil, errStdinNotTerminal } - // stdin is a terminal, check if stderr is also a terminal for interactive prompting + // stdin is a terminal, check if stderr is also a terminal for + // interactive prompting if !term.IsTerminal(syscall.Stderr) { - return nil, fmt.Errorf("cannot prompt for passphrase: stderr is not a terminal " + - "(running in non-interactive mode). Please set the SB_UNLOCK_PASSPHRASE " + - "environment variable") + return nil, errStderrNotTerminal } // Both stdin and stderr are terminals - use secure password reading fmt.Fprint(os.Stderr, prompt) // Write prompt to stderr, not stdout + passphrase, err := term.ReadPassword(syscall.Stdin) if err != nil { return nil, fmt.Errorf("failed to read passphrase: %w", err) } - fmt.Fprintln(os.Stderr) // Print newline to stderr since ReadPassword doesn't echo + + // Print newline to stderr since ReadPassword doesn't echo + fmt.Fprintln(os.Stderr) if len(passphrase) == 0 { - return nil, fmt.Errorf("passphrase cannot be empty") + return nil, errEmptyPassphrase } // Create a secure buffer and copy the passphrase diff --git a/internal/secret/debug.go b/internal/secret/debug.go index c38d4ad..9e3209f 100644 --- a/internal/secret/debug.go +++ b/internal/secret/debug.go @@ -13,28 +13,33 @@ import ( ) var ( - debugEnabled bool //nolint:gochecknoglobals // Package-wide debug state is necessary - debugLogger *slog.Logger //nolint:gochecknoglobals // Package-wide logger instance is necessary + debugEnabled bool //nolint:gochecknoglobals // package debug state + debugLogger *slog.Logger //nolint:gochecknoglobals // package debug logger ) +//nolint:gochecknoinits // debug logging must be ready before any package use func init() { InitDebugLogging() } -// InitDebugLogging initializes the debug logging system based on current GODEBUG environment variable +// InitDebugLogging initializes the debug logging system based on the +// current GODEBUG environment variable func InitDebugLogging() { godebug := os.Getenv("GODEBUG") debugEnabled = strings.Contains(godebug, "berlin.sneak.pkg.secret") if !debugEnabled { // Create a no-op logger that discards all output - debugLogger = slog.New(slog.NewTextHandler(io.Discard, nil)) + debugLogger = slog.New(slog.DiscardHandler) return } - // Disable stderr buffering for immediate debug output when debugging is enabled - _, _, _ = syscall.Syscall(syscall.SYS_FCNTL, os.Stderr.Fd(), syscall.F_SETFL, syscall.O_SYNC) + // Disable stderr buffering for immediate debug output when + // debugging is enabled + //nolint:dogsled // syscall.Syscall returns three values, none needed + _, _, _ = syscall.Syscall( + syscall.SYS_FCNTL, os.Stderr.Fd(), syscall.F_SETFL, syscall.O_SYNC) // Check if STDERR is a TTY isTTY := term.IsTerminal(syscall.Stderr) @@ -58,14 +63,19 @@ func IsDebugEnabled() bool { return debugEnabled } -// Warn logs a warning message to stderr unconditionally (visible without --verbose or debug flags) +// Warn logs a warning message to stderr unconditionally (visible +// without --verbose or debug flags) func Warn(msg string, args ...any) { - output := fmt.Sprintf("WARNING: %s", msg) + var output strings.Builder + + output.WriteString("WARNING: " + msg) + for i := 0; i+1 < len(args); i += 2 { - output += fmt.Sprintf(" %s=%v", args[i], args[i+1]) + fmt.Fprintf(&output, " %s=%v", args[i], args[i+1]) } - output += "\n" - fmt.Fprint(os.Stderr, output) + + output.WriteString("\n") + fmt.Fprint(os.Stderr, output.String()) } // Debug logs a debug message with optional attributes @@ -73,14 +83,16 @@ func Debug(msg string, args ...any) { if !debugEnabled { return } + debugLogger.Debug(msg, args...) } -// DebugF logs a formatted debug message with optional attributes -func DebugF(format string, args ...any) { +// Debugf logs a formatted debug message with optional attributes +func Debugf(format string, args ...any) { if !debugEnabled { return } + debugLogger.Debug(fmt.Sprintf(format, args...)) } @@ -89,6 +101,7 @@ func DebugWith(msg string, attrs ...slog.Attr) { if !debugEnabled { return } + debugLogger.LogAttrs(context.Background(), slog.LevelDebug, msg, attrs...) } @@ -118,15 +131,18 @@ func (h *colorizedHandler) Handle(_ context.Context, record slog.Record) error { if record.NumAttrs() > 0 { output += " \033[33m{" first := true + record.Attrs(func(attr slog.Attr) bool { if !first { output += ", " } + first = false output += fmt.Sprintf("%s=%#v", attr.Key, attr.Value.Any()) return true }) + output += "}\033[0m" } diff --git a/internal/secret/debug_test.go b/internal/secret/debug_test.go index f8a3601..3e11a9e 100644 --- a/internal/secret/debug_test.go +++ b/internal/secret/debug_test.go @@ -1,3 +1,4 @@ +//nolint:testpackage // white-box test of unexported debug internals package secret import ( @@ -90,9 +91,11 @@ func TestDebugLogging(t *testing.T) { } } +//nolint:paralleltest // exercises process-global debug logger state func TestDebugFunctions(t *testing.T) { // Enable debug for testing t.Setenv("GODEBUG", "berlin.sneak.pkg.secret") + defer InitDebugLogging() // Re-initialize after test InitDebugLogging() @@ -107,8 +110,8 @@ func TestDebugFunctions(t *testing.T) { Debug("test with args", "key", "value", "number", 42) }) - t.Run("DebugF", func(_ *testing.T) { - DebugF("formatted message: %s %d", "test", 123) + t.Run("Debugf", func(_ *testing.T) { + Debugf("formatted message: %s %d", "test", 123) }) t.Run("DebugWith", func(_ *testing.T) { diff --git a/internal/secret/helpers.go b/internal/secret/helpers.go index 5321f37..bd7713f 100644 --- a/internal/secret/helpers.go +++ b/internal/secret/helpers.go @@ -6,7 +6,8 @@ import ( "path/filepath" ) -// DetermineStateDir determines the state directory based on environment variables and OS. +// DetermineStateDir determines the state directory based on environment +// variables and OS. // It returns an error if no usable directory can be determined. func DetermineStateDir(customConfigDir string) (string, error) { // Check for environment variable first @@ -28,11 +29,14 @@ func DetermineStateDir(customConfigDir string) (string, error) { // Fallback to a reasonable default if we can't determine user config dir homeDir, homeErr := os.UserHomeDir() if homeErr != nil { - return "", fmt.Errorf("unable to determine state directory: config dir: %w, home dir: %w", err, homeErr) + return "", fmt.Errorf( + "unable to determine state directory: config dir: %w, home dir: %w", + err, homeErr) } fallbackDir := filepath.Join(homeDir, ".config", AppID) - Warn("Could not determine user config directory, falling back to default", "fallback", fallbackDir, "error", err) + Warn("Could not determine user config directory, falling back to default", + "fallback", fallbackDir, "error", err) return fallbackDir, nil } diff --git a/internal/secret/helpers_test.go b/internal/secret/helpers_test.go index b989b2d..51a46a2 100644 --- a/internal/secret/helpers_test.go +++ b/internal/secret/helpers_test.go @@ -1,7 +1,9 @@ -package secret +package secret_test import ( "testing" + + "git.eeqj.de/sneak/secret/internal/secret" ) func TestDetermineStateDir_ErrorsWhenHomeDirUnavailable(t *testing.T) { @@ -9,11 +11,11 @@ func TestDetermineStateDir_ErrorsWhenHomeDirUnavailable(t *testing.T) { // On Darwin, os.UserHomeDir may still succeed via the password // database, so we also test via an explicit empty-customConfigDir // path to exercise the fallback branch. - t.Setenv(EnvStateDir, "") + t.Setenv(secret.EnvStateDir, "") t.Setenv("HOME", "") t.Setenv("XDG_CONFIG_HOME", "") - result, err := DetermineStateDir("") + result, err := secret.DetermineStateDir("") // On systems where both lookups fail, we must get an error. // On systems where the OS provides a fallback (e.g. macOS pw db), // result should still be valid (non-empty, not root-relative). @@ -21,29 +23,36 @@ func TestDetermineStateDir_ErrorsWhenHomeDirUnavailable(t *testing.T) { // Good — the error case is handled. return } - if result == "/.config/"+AppID || result == "" { - t.Errorf("DetermineStateDir returned dangerous/empty path %q without error", result) + + if result == "/.config/"+secret.AppID || result == "" { + t.Errorf( + "DetermineStateDir returned dangerous/empty path %q without error", + result) } } func TestDetermineStateDir_UsesEnvVar(t *testing.T) { - t.Setenv(EnvStateDir, "/custom/state") - result, err := DetermineStateDir("") + t.Setenv(secret.EnvStateDir, "/custom/state") + + result, err := secret.DetermineStateDir("") if err != nil { t.Fatalf("unexpected error: %v", err) } + if result != "/custom/state" { t.Errorf("expected /custom/state, got %q", result) } } func TestDetermineStateDir_UsesCustomConfigDir(t *testing.T) { - t.Setenv(EnvStateDir, "") - result, err := DetermineStateDir("/my/config") + t.Setenv(secret.EnvStateDir, "") + + result, err := secret.DetermineStateDir("/my/config") if err != nil { t.Fatalf("unexpected error: %v", err) } - expected := "/my/config/" + AppID + + expected := "/my/config/" + secret.AppID if result != expected { t.Errorf("expected %q, got %q", expected, result) } diff --git a/internal/secret/keychainunlocker_stub.go b/internal/secret/keychainunlocker_stub.go index 6e79370..c6de515 100644 --- a/internal/secret/keychainunlocker_stub.go +++ b/internal/secret/keychainunlocker_stub.go @@ -1,10 +1,9 @@ //go:build !darwin -// +build !darwin package secret import ( - "fmt" + "errors" "filippo.io/age" "github.com/awnumar/memguard" @@ -14,6 +13,7 @@ import ( // KeychainUnlockerMetadata is a stub for non-Darwin platforms type KeychainUnlockerMetadata struct { UnlockerMetadata + KeychainItemName string `json:"keychainItemName"` } @@ -24,7 +24,21 @@ type KeychainUnlocker struct { fs afero.Fs } -var errKeychainNotSupported = fmt.Errorf("keychain unlockers are only supported on macOS") +var errKeychainNotSupported = errors.New( + "keychain unlockers are only supported on macOS") + +// NewKeychainUnlocker creates a stub KeychainUnlocker on non-Darwin +// platforms. The returned instance's methods that require macOS +// functionality will return errors. +func NewKeychainUnlocker( + fs afero.Fs, directory string, metadata UnlockerMetadata, +) *KeychainUnlocker { + return &KeychainUnlocker{ + Directory: directory, + Metadata: metadata, + fs: fs, + } +} // GetIdentity returns an error on non-Darwin platforms func (k *KeychainUnlocker) GetIdentity() (*age.X25519Identity, error) { @@ -48,7 +62,7 @@ func (k *KeychainUnlocker) GetDirectory() string { // GetID returns the unlocker ID func (k *KeychainUnlocker) GetID() string { - return fmt.Sprintf("%s-keychain", k.Metadata.CreatedAt.Format("2006-01-02.15.04")) + return k.Metadata.CreatedAt.Format("2006-01-02.15.04") + "-keychain" } // GetKeychainItemName returns an error on non-Darwin platforms @@ -61,22 +75,14 @@ func (k *KeychainUnlocker) Remove() error { return errKeychainNotSupported } -// NewKeychainUnlocker creates a stub KeychainUnlocker on non-Darwin platforms. -// The returned instance's methods that require macOS functionality will return errors. -func NewKeychainUnlocker(fs afero.Fs, directory string, metadata UnlockerMetadata) *KeychainUnlocker { - return &KeychainUnlocker{ - Directory: directory, - Metadata: metadata, - fs: fs, - } -} - // CreateKeychainUnlocker returns an error on non-Darwin platforms func CreateKeychainUnlocker(_ afero.Fs, _ string) (*KeychainUnlocker, error) { return nil, errKeychainNotSupported } // getLongTermPrivateKey returns an error on non-Darwin platforms -func getLongTermPrivateKey(_ afero.Fs, _ VaultInterface) (*memguard.LockedBuffer, error) { +func getLongTermPrivateKey( + _ afero.Fs, _ VaultInterface, +) (*memguard.LockedBuffer, error) { return nil, errKeychainNotSupported } diff --git a/internal/secret/passphrase_test.go b/internal/secret/passphrase_test.go index a0a2167..ebb2436 100644 --- a/internal/secret/passphrase_test.go +++ b/internal/secret/passphrase_test.go @@ -13,29 +13,134 @@ import ( "github.com/spf13/afero" ) -func TestPassphraseUnlockerWithRealFS(t *testing.T) { - // This test uses real filesystem - if os.Getenv("CI") == "true" { - t.Log("Running in CI environment with real filesystem") - } +// testMnemonic is the standard BIP39 test vector mnemonic. +// +//nolint:dupword // BIP39 test mnemonic repeats words by design +const testMnemonic = "abandon abandon abandon abandon abandon abandon " + + "abandon abandon abandon abandon abandon about" - // Create a temporary directory for our tests - tempDir, err := os.MkdirTemp("", "secret-passphrase-test-") +// writeTestPublicKey writes the unlocker public key and verifies it exists. +func writeTestPublicKey( + t *testing.T, fs afero.Fs, unlockerDir string, agePublicKey string, +) { + t.Helper() + + pubKeyPath := filepath.Join(unlockerDir, "pub.age") + + err := afero.WriteFile(fs, pubKeyPath, []byte(agePublicKey), secret.FilePerms) if err != nil { - t.Fatalf("Failed to create temp dir: %v", err) + t.Fatalf("Failed to write public key: %v", err) } - defer func() { _ = os.RemoveAll(tempDir) }() // Clean up after test - // Use the real filesystem - fs := afero.NewOsFs() + // Verify the file exists + exists, err := afero.Exists(fs, pubKeyPath) + if err != nil { + t.Fatalf("Failed to check if public key exists: %v", err) + } - // Test data - testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about" - testPassphrase := "test-passphrase-123" + if !exists { + t.Errorf("Public key file should exist at %s", pubKeyPath) + } +} - // Create the directory structure - unlockerDir := filepath.Join(tempDir, "unlocker") - if err := os.MkdirAll(unlockerDir, secret.DirPerms); err != nil { +// writeTestPrivateKey encrypts the private key with the passphrase, +// writes it, and verifies it exists. +func writeTestPrivateKey( + t *testing.T, + fs afero.Fs, + unlockerDir string, + agePrivateKey string, + testPassphrase string, +) { + t.Helper() + + privKeyBuffer := memguard.NewBufferFromBytes([]byte(agePrivateKey)) + defer privKeyBuffer.Destroy() + + passphraseBuffer := memguard.NewBufferFromBytes([]byte(testPassphrase)) + defer passphraseBuffer.Destroy() + + encryptedPrivKey, err := secret.EncryptWithPassphrase( + privKeyBuffer, passphraseBuffer) + if err != nil { + t.Fatalf("Failed to encrypt private key: %v", err) + } + + privKeyPath := filepath.Join(unlockerDir, "priv.age") + + err = afero.WriteFile(fs, privKeyPath, encryptedPrivKey, secret.FilePerms) + if err != nil { + t.Fatalf("Failed to write encrypted private key: %v", err) + } + + // Verify the file exists + exists, err := afero.Exists(fs, privKeyPath) + if err != nil { + t.Fatalf("Failed to check if private key exists: %v", err) + } + + if !exists { + t.Errorf("Encrypted private key file should exist at %s", privKeyPath) + } +} + +// writeTestLongTermKey encrypts the derived long-term key to the +// unlocker's recipient, writes it, and verifies it exists. +func writeTestLongTermKey( + t *testing.T, fs afero.Fs, unlockerDir string, agePublicKey string, +) { + t.Helper() + + // Derive a long-term identity from the test mnemonic + ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0) + if err != nil { + t.Fatalf("Failed to derive long-term identity: %v", err) + } + + // Encrypt long-term private key to the unlocker's recipient + recipient, err := age.ParseX25519Recipient(agePublicKey) + if err != nil { + t.Fatalf("Failed to parse recipient: %v", err) + } + + ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String())) + defer ltPrivKeyBuffer.Destroy() + + encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, recipient) + if err != nil { + t.Fatalf("Failed to encrypt long-term private key: %v", err) + } + + ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age") + + err = afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms) + if err != nil { + t.Fatalf("Failed to write encrypted long-term private key: %v", err) + } + + // Verify the file exists + exists, err := afero.Exists(fs, ltPrivKeyPath) + if err != nil { + t.Fatalf("Failed to check if long-term key exists: %v", err) + } + + if !exists { + t.Errorf("Encrypted long-term key file should exist at %s", ltPrivKeyPath) + } +} + +// newTestPassphraseUnlocker creates a temp unlocker directory and a +// passphrase unlocker with a fresh age identity for testing. +func newTestPassphraseUnlocker( + t *testing.T, fs afero.Fs, +) (*secret.PassphraseUnlocker, *age.X25519Identity, string) { + t.Helper() + + // Create the directory structure in a temp dir + unlockerDir := filepath.Join(t.TempDir(), "unlocker") + + err := os.MkdirAll(unlockerDir, secret.DirPerms) + if err != nil { t.Fatalf("Failed to create unlocker directory: %v", err) } @@ -54,86 +159,40 @@ func TestPassphraseUnlockerWithRealFS(t *testing.T) { if err != nil { t.Fatalf("Failed to generate age identity: %v", err) } + + return unlocker, ageIdentity, unlockerDir +} + +//nolint:paralleltest // subtests share real-FS state and t.Setenv, order matters +func TestPassphraseUnlockerWithRealFS(t *testing.T) { + // This test uses real filesystem + if os.Getenv("CI") == "true" { + t.Log("Running in CI environment with real filesystem") + } + + // Use the real filesystem + fs := afero.NewOsFs() + + // Test data + testPassphrase := "test-passphrase-123" + + unlocker, ageIdentity, unlockerDir := newTestPassphraseUnlocker(t, fs) agePrivateKey := ageIdentity.String() agePublicKey := ageIdentity.Recipient().String() // Test writing public key t.Run("WritePublicKey", func(t *testing.T) { - pubKeyPath := filepath.Join(unlockerDir, "pub.age") - if err := afero.WriteFile(fs, pubKeyPath, []byte(agePublicKey), secret.FilePerms); err != nil { - t.Fatalf("Failed to write public key: %v", err) - } - - // Verify the file exists - exists, err := afero.Exists(fs, pubKeyPath) - if err != nil { - t.Fatalf("Failed to check if public key exists: %v", err) - } - if !exists { - t.Errorf("Public key file should exist at %s", pubKeyPath) - } + writeTestPublicKey(t, fs, unlockerDir, agePublicKey) }) // Test encrypting private key with passphrase t.Run("EncryptPrivateKey", func(t *testing.T) { - privKeyBuffer := memguard.NewBufferFromBytes([]byte(agePrivateKey)) - defer privKeyBuffer.Destroy() - passphraseBuffer := memguard.NewBufferFromBytes([]byte(testPassphrase)) - defer passphraseBuffer.Destroy() - encryptedPrivKey, err := secret.EncryptWithPassphrase(privKeyBuffer, passphraseBuffer) - if err != nil { - t.Fatalf("Failed to encrypt private key: %v", err) - } - - privKeyPath := filepath.Join(unlockerDir, "priv.age") - if err := afero.WriteFile(fs, privKeyPath, encryptedPrivKey, secret.FilePerms); err != nil { - t.Fatalf("Failed to write encrypted private key: %v", err) - } - - // Verify the file exists - exists, err := afero.Exists(fs, privKeyPath) - if err != nil { - t.Fatalf("Failed to check if private key exists: %v", err) - } - if !exists { - t.Errorf("Encrypted private key file should exist at %s", privKeyPath) - } + writeTestPrivateKey(t, fs, unlockerDir, agePrivateKey, testPassphrase) }) // Test writing long-term key t.Run("WriteLongTermKey", func(t *testing.T) { - // Derive a long-term identity from the test mnemonic - ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0) - if err != nil { - t.Fatalf("Failed to derive long-term identity: %v", err) - } - - // Encrypt long-term private key to the unlocker's recipient - recipient, err := age.ParseX25519Recipient(agePublicKey) - if err != nil { - t.Fatalf("Failed to parse recipient: %v", err) - } - - ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String())) - defer ltPrivKeyBuffer.Destroy() - encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, recipient) - if err != nil { - t.Fatalf("Failed to encrypt long-term private key: %v", err) - } - - ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age") - if err := afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms); err != nil { - t.Fatalf("Failed to write encrypted long-term private key: %v", err) - } - - // Verify the file exists - exists, err := afero.Exists(fs, ltPrivKeyPath) - if err != nil { - t.Fatalf("Failed to check if long-term key exists: %v", err) - } - if !exists { - t.Errorf("Encrypted long-term key file should exist at %s", ltPrivKeyPath) - } + writeTestLongTermKey(t, fs, unlockerDir, agePublicKey) }) // Set test environment variable (cleaned up automatically) @@ -148,18 +207,21 @@ func TestPassphraseUnlockerWithRealFS(t *testing.T) { // Verify the identity matches what we expect expectedPubKey := ageIdentity.Recipient().String() + actualPubKey := identity.Recipient().String() if actualPubKey != expectedPubKey { - t.Errorf("Public key mismatch. Expected %s, got %s", expectedPubKey, actualPubKey) + t.Errorf("Public key mismatch. Expected %s, got %s", + expectedPubKey, actualPubKey) } }) // Unset the environment variable to test interactive prompt _ = os.Unsetenv(secret.EnvUnlockPassphrase) - // Test getting identity from prompt (this would require mocking the prompt) - // For real integration tests, we'd need to provide a way to mock the passphrase input - // Here we'll just verify the error is what we expect when no passphrase is available + // Test getting identity from prompt (this would require mocking the + // prompt). For real integration tests, we'd need a way to mock the + // passphrase input. Here we just verify the error is what we expect + // when no passphrase is available. t.Run("GetIdentityWithoutEnv", func(t *testing.T) { // This should fail since we're not in an interactive terminal _, err := unlocker.GetIdentity() @@ -180,6 +242,7 @@ func TestPassphraseUnlockerWithRealFS(t *testing.T) { if err != nil { t.Fatalf("Failed to check if unlocker directory exists: %v", err) } + if exists { t.Errorf("Unlocker directory should not exist after removal") } diff --git a/internal/secret/passphraseunlocker.go b/internal/secret/passphraseunlocker.go index ef7f22f..9711d9a 100644 --- a/internal/secret/passphraseunlocker.go +++ b/internal/secret/passphraseunlocker.go @@ -19,37 +19,15 @@ type PassphraseUnlocker struct { Passphrase *memguard.LockedBuffer // Secure buffer for passphrase } -// getPassphrase retrieves the passphrase from memory, environment, or user input -// Returns a LockedBuffer for secure memory handling -func (p *PassphraseUnlocker) getPassphrase() (*memguard.LockedBuffer, error) { - // First check if we already have the passphrase - if p.Passphrase != nil && p.Passphrase.IsAlive() { - Debug("Using in-memory passphrase", "unlocker_id", p.GetID()) - // Return a copy of the passphrase buffer - return memguard.NewBufferFromBytes(p.Passphrase.Bytes()), nil +// NewPassphraseUnlocker creates a new PassphraseUnlocker instance +func NewPassphraseUnlocker( + fs afero.Fs, directory string, metadata UnlockerMetadata, +) *PassphraseUnlocker { + return &PassphraseUnlocker{ + Directory: directory, + Metadata: metadata, + fs: fs, } - - Debug("No passphrase in memory, checking environment") - // Check environment variable for passphrase - passphraseStr := os.Getenv(EnvUnlockPassphrase) - if passphraseStr != "" { - Debug("Using passphrase from environment", "unlocker_id", p.GetID()) - // Convert to secure buffer - secureBuffer := memguard.NewBufferFromBytes([]byte(passphraseStr)) - - return secureBuffer, nil - } - - Debug("No passphrase in environment, prompting user") - // Prompt for passphrase - secureBuffer, err := ReadPassphrase("Enter unlock passphrase: ") - if err != nil { - Debug("Failed to read passphrase", "error", err, "unlocker_id", p.GetID()) - - return nil, fmt.Errorf("failed to read passphrase: %w", err) - } - - return secureBuffer, nil } // GetIdentity implements Unlocker interface for passphrase-based unlockers @@ -71,7 +49,8 @@ func (p *PassphraseUnlocker) GetIdentity() (*age.X25519Identity, error) { encryptedPrivKeyData, err := afero.ReadFile(p.fs, unlockerPrivPath) if err != nil { - Debug("Failed to read passphrase unlocker private key", "error", err, "path", unlockerPrivPath) + Debug("Failed to read passphrase unlocker private key", + "error", err, "path", unlockerPrivPath) return nil, fmt.Errorf("failed to read unlocker private key: %w", err) } @@ -86,7 +65,8 @@ func (p *PassphraseUnlocker) GetIdentity() (*age.X25519Identity, error) { // Decrypt the unlocker private key with passphrase privKeyBuffer, err := DecryptWithPassphrase(encryptedPrivKeyData, passphraseBuffer) if err != nil { - Debug("Failed to decrypt unlocker private key", "error", err, "unlocker_id", p.GetID()) + Debug("Failed to decrypt unlocker private key", + "error", err, "unlocker_id", p.GetID()) return nil, fmt.Errorf("failed to decrypt unlocker private key: %w", err) } @@ -135,7 +115,7 @@ func (p *PassphraseUnlocker) GetID() string { // Generate ID using creation timestamp: YYYY-MM-DD.HH.mm-passphrase createdAt := p.Metadata.CreatedAt - return fmt.Sprintf("%s-passphrase", createdAt.Format("2006-01-02.15.04")) + return createdAt.Format("2006-01-02.15.04") + "-passphrase" } // Remove implements Unlocker interface - removes the passphrase unlocker @@ -147,20 +127,45 @@ func (p *PassphraseUnlocker) Remove() error { // For passphrase unlockers, we just need to remove the directory // No external resources (like keychain items) to clean up - if err := p.fs.RemoveAll(p.Directory); err != nil { + err := p.fs.RemoveAll(p.Directory) + if err != nil { return fmt.Errorf("failed to remove passphrase unlocker directory: %w", err) } return nil } -// NewPassphraseUnlocker creates a new PassphraseUnlocker instance -func NewPassphraseUnlocker(fs afero.Fs, directory string, metadata UnlockerMetadata) *PassphraseUnlocker { - return &PassphraseUnlocker{ - Directory: directory, - Metadata: metadata, - fs: fs, +// getPassphrase retrieves the passphrase from memory, environment, or +// user input. Returns a LockedBuffer for secure memory handling +func (p *PassphraseUnlocker) getPassphrase() (*memguard.LockedBuffer, error) { + // First check if we already have the passphrase + if p.Passphrase != nil && p.Passphrase.IsAlive() { + Debug("Using in-memory passphrase", "unlocker_id", p.GetID()) + // Return a copy of the passphrase buffer + return memguard.NewBufferFromBytes(p.Passphrase.Bytes()), nil } + + Debug("No passphrase in memory, checking environment") + // Check environment variable for passphrase + passphraseStr := os.Getenv(EnvUnlockPassphrase) + if passphraseStr != "" { + Debug("Using passphrase from environment", "unlocker_id", p.GetID()) + // Convert to secure buffer + secureBuffer := memguard.NewBufferFromBytes([]byte(passphraseStr)) + + return secureBuffer, nil + } + + Debug("No passphrase in environment, prompting user") + // Prompt for passphrase + secureBuffer, err := ReadPassphrase("Enter unlock passphrase: ") + if err != nil { + Debug("Failed to read passphrase", "error", err, "unlocker_id", p.GetID()) + + return nil, fmt.Errorf("failed to read passphrase: %w", err) + } + + return secureBuffer, nil } // CreatePassphraseUnlocker creates a new passphrase-protected unlocker diff --git a/internal/secret/pgpunlocker.go b/internal/secret/pgpunlocker.go index bca1546..c6fcbc3 100644 --- a/internal/secret/pgpunlocker.go +++ b/internal/secret/pgpunlocker.go @@ -1,7 +1,9 @@ package secret import ( + "context" "encoding/json" + "errors" "fmt" "log/slog" "os" @@ -16,17 +18,28 @@ import ( "github.com/spf13/afero" ) +var ( + errGPGKeyIDEmpty = errors.New("GPG key ID cannot be empty") + errInvalidGPGKeyID = errors.New("invalid GPG key ID format") + errNoGPGFingerprint = errors.New("could not find fingerprint for GPG key") + errNilDataBuffer = errors.New("data buffer is nil") +) + // Variables to allow overriding in tests var ( // GPGEncryptFunc is the function used for GPG encryption // Can be overridden in tests to provide a non-interactive implementation //nolint:gochecknoglobals // Required for test mocking - GPGEncryptFunc func(data *memguard.LockedBuffer, keyID string) ([]byte, error) = gpgEncryptDefault + GPGEncryptFunc func( + data *memguard.LockedBuffer, keyID string, + ) ([]byte, error) = gpgEncryptDefault // GPGDecryptFunc is the function used for GPG decryption // Can be overridden in tests to provide a non-interactive implementation //nolint:gochecknoglobals // Required for test mocking - GPGDecryptFunc func(encryptedData []byte) (*memguard.LockedBuffer, error) = gpgDecryptDefault + GPGDecryptFunc func( + encryptedData []byte, + ) (*memguard.LockedBuffer, error) = gpgDecryptDefault // gpgKeyIDRegex validates GPG key IDs // Allows either: @@ -45,6 +58,7 @@ var ( // PGPUnlockerMetadata extends UnlockerMetadata with PGP-specific data type PGPUnlockerMetadata struct { UnlockerMetadata + // GPG key ID used for encryption GPGKeyID string `json:"gpgKeyId"` } @@ -56,6 +70,17 @@ type PGPUnlocker struct { fs afero.Fs } +// NewPGPUnlocker creates a new PGPUnlocker instance +func NewPGPUnlocker( + fs afero.Fs, directory string, metadata UnlockerMetadata, +) *PGPUnlocker { + return &PGPUnlocker{ + Directory: directory, + Metadata: metadata, + fs: fs, + } +} + // GetIdentity implements Unlocker interface for PGP-based unlockers func (p *PGPUnlocker) GetIdentity() (*age.X25519Identity, error) { DebugWith("Getting PGP unlocker identity", @@ -69,7 +94,8 @@ func (p *PGPUnlocker) GetIdentity() (*age.X25519Identity, error) { encryptedAgePrivKeyData, err := afero.ReadFile(p.fs, agePrivKeyPath) if err != nil { - Debug("Failed to read PGP-encrypted age private key", "error", err, "path", agePrivKeyPath) + Debug("Failed to read PGP-encrypted age private key", + "error", err, "path", agePrivKeyPath) return nil, fmt.Errorf("failed to read encrypted age private key: %w", err) } @@ -81,9 +107,11 @@ func (p *PGPUnlocker) GetIdentity() (*age.X25519Identity, error) { // Step 2: Decrypt the age private key using GPG Debug("Decrypting age private key with GPG", "unlocker_id", p.GetID()) + agePrivKeyBuffer, err := GPGDecryptFunc(encryptedAgePrivKeyData) if err != nil { - Debug("Failed to decrypt age private key with GPG", "error", err, "unlocker_id", p.GetID()) + Debug("Failed to decrypt age private key with GPG", + "error", err, "unlocker_id", p.GetID()) return nil, fmt.Errorf("failed to decrypt age private key with GPG: %w", err) } @@ -96,6 +124,7 @@ func (p *PGPUnlocker) GetIdentity() (*age.X25519Identity, error) { // Step 3: Parse the decrypted age private key Debug("Parsing decrypted age private key", "unlocker_id", p.GetID()) + ageIdentity, err := age.ParseX25519Identity(agePrivKeyBuffer.String()) if err != nil { Debug("Failed to parse age private key", "error", err, "unlocker_id", p.GetID()) @@ -136,47 +165,43 @@ func (p *PGPUnlocker) GetID() string { panic(fmt.Sprintf("PGP unlocker metadata is corrupt or missing GPG key ID: %v", err)) } - return fmt.Sprintf("pgp-%s", gpgKeyID) + return "pgp-" + gpgKeyID } // Remove implements Unlocker interface - removes the PGP unlocker func (p *PGPUnlocker) Remove() error { // For PGP unlockers, we just need to remove the directory // No external resources (like keychain items) to clean up - if err := p.fs.RemoveAll(p.Directory); err != nil { + err := p.fs.RemoveAll(p.Directory) + if err != nil { return fmt.Errorf("failed to remove PGP unlocker directory: %w", err) } return nil } -// NewPGPUnlocker creates a new PGPUnlocker instance -func NewPGPUnlocker(fs afero.Fs, directory string, metadata UnlockerMetadata) *PGPUnlocker { - return &PGPUnlocker{ - Directory: directory, - Metadata: metadata, - fs: fs, - } -} - // GetGPGKeyID returns the GPG key ID from metadata func (p *PGPUnlocker) GetGPGKeyID() (string, error) { // Load the metadata metadataPath := filepath.Join(p.Directory, "unlocker-metadata.json") + metadataData, err := afero.ReadFile(p.fs, metadataPath) if err != nil { return "", fmt.Errorf("failed to read PGP metadata: %w", err) } var pgpMetadata PGPUnlockerMetadata - if err := json.Unmarshal(metadataData, &pgpMetadata); err != nil { + + err = json.Unmarshal(metadataData, &pgpMetadata) + if err != nil { return "", fmt.Errorf("failed to parse PGP metadata: %w", err) } return pgpMetadata.GPGKeyID, nil } -// generatePGPUnlockerName generates a unique name for the PGP unlocker based on hostname and date +// generatePGPUnlockerName generates a unique name for the PGP unlocker +// based on hostname and date func generatePGPUnlockerName() (string, error) { hostname, err := os.Hostname() if err != nil { @@ -189,34 +214,55 @@ func generatePGPUnlockerName() (string, error) { return fmt.Sprintf("%s-pgp-%s", hostname, enrollmentDate), nil } -// CreatePGPUnlocker creates a new PGP unlocker and stores it in the vault -func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnlocker, error) { +// preparePGPUnlockerDir checks GPG availability and creates the +// unlocker directory in the current vault, returning the vault and the +// directory path. +// +//nolint:ireturn // the vault is only available behind VaultInterface +func preparePGPUnlockerDir( + fs afero.Fs, stateDir string, +) (VaultInterface, string, error) { // Check if GPG is available - if err := checkGPGAvailable(); err != nil { - return nil, err + err := checkGPGAvailable() + if err != nil { + return nil, "", err } // Get current vault vault, err := GetCurrentVault(fs, stateDir) if err != nil { - return nil, fmt.Errorf("failed to get current vault: %w", err) + return nil, "", fmt.Errorf("failed to get current vault: %w", err) } // Generate the unlocker name based on hostname and date unlockerName, err := generatePGPUnlockerName() if err != nil { - return nil, fmt.Errorf("failed to generate unlocker name: %w", err) + return nil, "", fmt.Errorf("failed to generate unlocker name: %w", err) } // Create unlocker directory using the generated name vaultDir, err := vault.GetDirectory() if err != nil { - return nil, fmt.Errorf("failed to get vault directory: %w", err) + return nil, "", fmt.Errorf("failed to get vault directory: %w", err) } unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerName) - if err := fs.MkdirAll(unlockerDir, DirPerms); err != nil { - return nil, fmt.Errorf("failed to create unlocker directory: %w", err) + + err = fs.MkdirAll(unlockerDir, DirPerms) + if err != nil { + return nil, "", fmt.Errorf("failed to create unlocker directory: %w", err) + } + + return vault, unlockerDir, nil +} + +// CreatePGPUnlocker creates a new PGP unlocker and stores it in the vault +func CreatePGPUnlocker( + fs afero.Fs, stateDir string, gpgKeyID string, +) (*PGPUnlocker, error) { + vault, unlockerDir, err := preparePGPUnlockerDir(fs, stateDir) + if err != nil { + return nil, err } // Step 1: Generate a new age keypair for the PGP unlocker @@ -228,7 +274,9 @@ func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnloc // Step 2: Store age recipient as plaintext ageRecipient := ageIdentity.Recipient().String() recipientPath := filepath.Join(unlockerDir, "pub.txt") - if err := afero.WriteFile(fs, recipientPath, []byte(ageRecipient), FilePerms); err != nil { + + err = afero.WriteFile(fs, recipientPath, []byte(ageRecipient), FilePerms) + if err != nil { return nil, fmt.Errorf("failed to write age recipient: %w", err) } @@ -240,14 +288,18 @@ func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnloc defer ltPrivKeyData.Destroy() // Step 7: Encrypt long-term private key to the new age unlocker - encryptedLtPrivKeyToAge, err := EncryptToRecipient(ltPrivKeyData, ageIdentity.Recipient()) + encryptedLtPrivKeyToAge, err := EncryptToRecipient( + ltPrivKeyData, ageIdentity.Recipient()) if err != nil { - return nil, fmt.Errorf("failed to encrypt long-term private key to age unlocker: %w", err) + return nil, fmt.Errorf( + "failed to encrypt long-term private key to age unlocker: %w", err) } // Write encrypted long-term private key ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age") - if err := afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKeyToAge, FilePerms); err != nil { + + err = afero.WriteFile(fs, ltPrivKeyPath, encryptedLtPrivKeyToAge, FilePerms) + if err != nil { return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err) } @@ -262,17 +314,35 @@ func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnloc } agePrivKeyPath := filepath.Join(unlockerDir, "priv.age.gpg") - if err := afero.WriteFile(fs, agePrivKeyPath, encryptedAgePrivKey, FilePerms); err != nil { + + err = afero.WriteFile(fs, agePrivKeyPath, encryptedAgePrivKey, FilePerms) + if err != nil { return nil, fmt.Errorf("failed to write encrypted age private key: %w", err) } - // Step 9: Resolve the GPG key ID to its full fingerprint + // Steps 9-10: Resolve the fingerprint and write enhanced metadata + pgpMetadata, err := writePGPUnlockerMetadata(fs, unlockerDir, gpgKeyID) + if err != nil { + return nil, err + } + + return &PGPUnlocker{ + Directory: unlockerDir, + Metadata: pgpMetadata.UnlockerMetadata, + fs: fs, + }, nil +} + +// writePGPUnlockerMetadata resolves the GPG key fingerprint and writes +// the unlocker metadata file, returning the metadata written. +func writePGPUnlockerMetadata( + fs afero.Fs, unlockerDir string, gpgKeyID string, +) (*PGPUnlockerMetadata, error) { fingerprint, err := ResolveGPGKeyFingerprint(gpgKeyID) if err != nil { return nil, fmt.Errorf("failed to resolve GPG key fingerprint: %w", err) } - // Step 10: Create and write enhanced metadata with full fingerprint pgpMetadata := PGPUnlockerMetadata{ UnlockerMetadata: UnlockerMetadata{ Type: "pgp", @@ -287,27 +357,24 @@ func CreatePGPUnlocker(fs afero.Fs, stateDir string, gpgKeyID string) (*PGPUnloc return nil, fmt.Errorf("failed to marshal unlocker metadata: %w", err) } - if err := afero.WriteFile(fs, + err = afero.WriteFile(fs, filepath.Join(unlockerDir, "unlocker-metadata.json"), - metadataBytes, FilePerms); err != nil { + metadataBytes, FilePerms) + if err != nil { return nil, fmt.Errorf("failed to write unlocker metadata: %w", err) } - return &PGPUnlocker{ - Directory: unlockerDir, - Metadata: pgpMetadata.UnlockerMetadata, - fs: fs, - }, nil + return &pgpMetadata, nil } // validateGPGKeyID validates that a GPG key ID is safe for command execution func validateGPGKeyID(keyID string) error { if keyID == "" { - return fmt.Errorf("GPG key ID cannot be empty") + return errGPGKeyIDEmpty } if !gpgKeyIDRegex.MatchString(keyID) { - return fmt.Errorf("invalid GPG key ID format: %s", keyID) + return fmt.Errorf("%w: %s", errInvalidGPGKeyID, keyID) } return nil @@ -315,22 +382,24 @@ func validateGPGKeyID(keyID string) error { // ResolveGPGKeyFingerprint resolves any GPG key identifier to its full fingerprint func ResolveGPGKeyFingerprint(keyID string) (string, error) { - if err := validateGPGKeyID(keyID); err != nil { + err := validateGPGKeyID(keyID) + if err != nil { return "", fmt.Errorf("invalid GPG key ID: %w", err) } // Use GPG to get the full fingerprint for the key - cmd := exec.Command( // #nosec G204 -- keyID validated + cmd := exec.CommandContext( //nolint:gosec // G204: keyID validated above + context.Background(), "gpg", "--list-keys", "--with-colons", "--fingerprint", keyID, ) + output, err := cmd.Output() if err != nil { return "", fmt.Errorf("failed to resolve GPG key fingerprint: %w", err) } // Parse the output to extract the fingerprint - lines := strings.Split(string(output), "\n") - for _, line := range lines { + for line := range strings.SplitSeq(string(output), "\n") { if strings.HasPrefix(line, "fpr:") { fields := strings.Split(line, ":") if len(fields) >= 10 && fields[9] != "" { @@ -339,14 +408,18 @@ func ResolveGPGKeyFingerprint(keyID string) (string, error) { } } - return "", fmt.Errorf("could not find fingerprint for GPG key: %s", keyID) + return "", fmt.Errorf("%w: %s", errNoGPGFingerprint, keyID) } // checkGPGAvailable verifies that GPG is available func checkGPGAvailable() error { - cmd := exec.Command("gpg", "--version") - if err := cmd.Run(); err != nil { - return fmt.Errorf("GPG not available: %w (make sure 'gpg' command is installed and in PATH)", err) + cmd := exec.CommandContext(context.Background(), "gpg", "--version") + + err := cmd.Run() + if err != nil { + return fmt.Errorf( + "GPG not available: %w (make sure 'gpg' command is installed and in PATH)", + err) } return nil @@ -355,13 +428,16 @@ func checkGPGAvailable() error { // gpgEncryptDefault is the default implementation of GPG encryption func gpgEncryptDefault(data *memguard.LockedBuffer, keyID string) ([]byte, error) { if data == nil { - return nil, fmt.Errorf("data buffer is nil") + return nil, errNilDataBuffer } - if err := validateGPGKeyID(keyID); err != nil { + + err := validateGPGKeyID(keyID) + if err != nil { return nil, fmt.Errorf("invalid GPG key ID: %w", err) } - cmd := exec.Command( // #nosec G204 -- keyID validated + cmd := exec.CommandContext( //nolint:gosec // G204: keyID validated above + context.Background(), "gpg", "--trust-model", "always", "--armor", "--encrypt", "-r", keyID, ) cmd.Stdin = strings.NewReader(data.String()) @@ -376,7 +452,7 @@ func gpgEncryptDefault(data *memguard.LockedBuffer, keyID string) ([]byte, error // gpgDecryptDefault is the default implementation of GPG decryption func gpgDecryptDefault(encryptedData []byte) (*memguard.LockedBuffer, error) { - cmd := exec.Command("gpg", "--quiet", "--decrypt") + cmd := exec.CommandContext(context.Background(), "gpg", "--quiet", "--decrypt") cmd.Stdin = strings.NewReader(string(encryptedData)) output, err := cmd.Output() diff --git a/internal/secret/secret.go b/internal/secret/secret.go index a37ba7f..b8215ff 100644 --- a/internal/secret/secret.go +++ b/internal/secret/secret.go @@ -2,6 +2,7 @@ package secret import ( "encoding/json" + "errors" "fmt" "log/slog" "os" @@ -15,6 +16,15 @@ import ( "github.com/spf13/afero" ) +var ( + errSecretNotFound = errors.New("secret not found") + errUnlockerRequired = errors.New("unlocker required to decrypt secret") + errGetEncryptedDataDeprecated = errors.New( + "GetEncryptedData is deprecated - use version-specific methods") + errGetCurrentVaultNotRegistered = errors.New( + "GetCurrentVault function not registered") +) + // VaultInterface defines the interface that vault implementations must satisfy type VaultInterface interface { GetDirectory() (string, error) @@ -22,7 +32,8 @@ type VaultInterface interface { GetName() string GetFilesystem() afero.Fs GetCurrentUnlocker() (Unlocker, error) - CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*PassphraseUnlocker, error) + CreatePassphraseUnlocker( + passphrase *memguard.LockedBuffer) (*PassphraseUnlocker, error) } // Secret represents a secret in a vault @@ -62,7 +73,8 @@ func NewSecret(vault VaultInterface, name string) *Secret { } } -// GetValue retrieves and decrypts the current version's value using the provided unlocker +// GetValue retrieves and decrypts the current version's value using the +// provided unlocker func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) { DebugWith("Getting secret value", slog.String("secret_name", s.Name), @@ -72,14 +84,17 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) { // Check if secret exists exists, err := s.Exists() if err != nil { - Debug("Failed to check if secret exists during GetValue", "error", err, "secret_name", s.Name) + Debug("Failed to check if secret exists during GetValue", + "error", err, "secret_name", s.Name) return nil, fmt.Errorf("failed to check if secret exists: %w", err) } - if !exists { - Debug("Secret not found during GetValue", "secret_name", s.Name, "vault_name", s.vault.GetName()) - return nil, fmt.Errorf("secret %s not found", s.Name) + if !exists { + Debug("Secret not found during GetValue", + "secret_name", s.Name, "vault_name", s.vault.GetName()) + + return nil, fmt.Errorf("%w: %s", errSecretNotFound, s.Name) } Debug("Secret exists, getting current version", "secret_name", s.Name) @@ -95,52 +110,9 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) { // Create version object version := NewVersion(s.vault, s.Name, currentVersion) - // Check if we have SB_SECRET_MNEMONIC environment variable for direct decryption + // Check for SB_SECRET_MNEMONIC environment variable for direct decryption if envMnemonic := os.Getenv(EnvMnemonic); envMnemonic != "" { - Debug("Using mnemonic from environment for direct long-term key derivation", "secret_name", s.Name) - - // Get vault directory to read metadata - vaultDir, err := s.vault.GetDirectory() - if err != nil { - Debug("Failed to get vault directory", "error", err, "secret_name", s.Name) - - return nil, fmt.Errorf("failed to get vault directory: %w", err) - } - - // Load vault metadata to get the correct derivation index - metadataPath := filepath.Join(vaultDir, "vault-metadata.json") - metadataBytes, err := afero.ReadFile(s.vault.GetFilesystem(), metadataPath) - if err != nil { - Debug("Failed to read vault metadata", "error", err, "path", metadataPath) - - return nil, fmt.Errorf("failed to read vault metadata: %w", err) - } - - var metadata VaultMetadata - if err := json.Unmarshal(metadataBytes, &metadata); err != nil { - Debug("Failed to parse vault metadata", "error", err, "secret_name", s.Name) - - return nil, fmt.Errorf("failed to parse vault metadata: %w", err) - } - - DebugWith("Using vault derivation index from metadata", - slog.String("secret_name", s.Name), - slog.String("vault_name", s.vault.GetName()), - slog.Uint64("derivation_index", uint64(metadata.DerivationIndex)), - ) - - // Use mnemonic with the vault's derivation index from metadata - ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex) - if err != nil { - Debug("Failed to derive long-term key from mnemonic for secret", "error", err, "secret_name", s.Name) - - return nil, fmt.Errorf("failed to derive long-term key from mnemonic: %w", err) - } - - Debug("Successfully derived long-term key from mnemonic", "secret_name", s.Name) - - // Use the long-term key to decrypt the version - return version.GetValue(ltIdentity) + return s.getValueViaMnemonic(version, envMnemonic) } Debug("Using unlocker for vault access", "secret_name", s.Name) @@ -149,51 +121,12 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) { if unlocker == nil { Debug("No unlocker provided for secret decryption", "secret_name", s.Name) - return nil, fmt.Errorf("unlocker required to decrypt secret") + return nil, errUnlockerRequired } - DebugWith("Getting vault's long-term key using unlocker", - slog.String("secret_name", s.Name), - slog.String("unlocker_type", unlocker.GetType()), - slog.String("unlocker_id", unlocker.GetID()), - ) - - // Step 1: Use the unlocker to get the vault's long-term private key - unlockIdentity, err := unlocker.GetIdentity() + ltIdentity, err := s.getLongTermIdentityFromUnlocker(unlocker) if err != nil { - Debug("Failed to get unlocker identity", "error", err, "secret_name", s.Name, "unlocker_type", unlocker.GetType()) - - return nil, fmt.Errorf("failed to get unlocker identity: %w", err) - } - - // Read the encrypted long-term private key from the unlocker directory - encryptedLtPrivKeyPath := filepath.Join(unlocker.GetDirectory(), "longterm.age") - Debug("Reading encrypted long-term private key", "path", encryptedLtPrivKeyPath) - - encryptedLtPrivKey, err := afero.ReadFile(s.vault.GetFilesystem(), encryptedLtPrivKeyPath) - if err != nil { - Debug("Failed to read encrypted long-term private key", "error", err, "path", encryptedLtPrivKeyPath) - - return nil, fmt.Errorf("failed to read encrypted long-term private key: %w", err) - } - - // Decrypt the encrypted long-term private key using the unlocker - Debug("Decrypting long-term private key using unlocker", "secret_name", s.Name) - ltPrivKeyBuffer, err := DecryptWithIdentity(encryptedLtPrivKey, unlockIdentity) - if err != nil { - Debug("Failed to decrypt long-term private key", "error", err, "secret_name", s.Name) - - return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err) - } - defer ltPrivKeyBuffer.Destroy() - - // Parse the long-term private key - Debug("Parsing long-term private key", "secret_name", s.Name) - ltIdentity, err := age.ParseX25519Identity(ltPrivKeyBuffer.String()) - if err != nil { - Debug("Failed to parse long-term private key", "error", err, "secret_name", s.Name) - - return nil, fmt.Errorf("failed to parse long-term private key: %w", err) + return nil, err } DebugWith("Successfully obtained vault's long-term key", @@ -207,7 +140,8 @@ func (s *Secret) GetValue(unlocker Unlocker) (*memguard.LockedBuffer, error) { // LoadMetadata is deprecated - metadata is now per-version and encrypted func (s *Secret) LoadMetadata() error { - Debug("LoadMetadata called but is deprecated in versioned model", "secret_name", s.Name) + Debug("LoadMetadata called but is deprecated in versioned model", + "secret_name", s.Name) // For backward compatibility, we'll populate with basic info now := time.Now() s.Metadata = Metadata{ @@ -227,9 +161,10 @@ func (s *Secret) GetMetadata() Metadata { // GetEncryptedData is deprecated - data is now stored in versions func (s *Secret) GetEncryptedData() ([]byte, error) { - Debug("GetEncryptedData called but is deprecated in versioned model", "secret_name", s.Name) + Debug("GetEncryptedData called but is deprecated in versioned model", + "secret_name", s.Name) - return nil, fmt.Errorf("GetEncryptedData is deprecated - use version-specific methods") + return nil, errGetEncryptedDataDeprecated } // Exists checks if the secret exists on disk @@ -242,7 +177,8 @@ func (s *Secret) Exists() (bool, error) { // Check if the secret directory exists and has a current symlink exists, err := afero.DirExists(s.vault.GetFilesystem(), s.Directory) if err != nil { - Debug("Failed to check secret directory existence", "error", err, "secret_dir", s.Directory) + Debug("Failed to check secret directory existence", + "error", err, "secret_dir", s.Directory) return false, err } @@ -269,14 +205,134 @@ func (s *Secret) Exists() (bool, error) { return true, nil } +// getValueViaMnemonic derives the vault's long-term key from the +// mnemonic in the environment and decrypts the version value with it. +func (s *Secret) getValueViaMnemonic( + version *Version, envMnemonic string, +) (*memguard.LockedBuffer, error) { + Debug("Using mnemonic from environment for direct long-term key derivation", + "secret_name", s.Name) + + // Get vault directory to read metadata + vaultDir, err := s.vault.GetDirectory() + if err != nil { + Debug("Failed to get vault directory", "error", err, "secret_name", s.Name) + + return nil, fmt.Errorf("failed to get vault directory: %w", err) + } + + // Load vault metadata to get the correct derivation index + metadataPath := filepath.Join(vaultDir, "vault-metadata.json") + + metadataBytes, err := afero.ReadFile(s.vault.GetFilesystem(), metadataPath) + if err != nil { + Debug("Failed to read vault metadata", "error", err, "path", metadataPath) + + return nil, fmt.Errorf("failed to read vault metadata: %w", err) + } + + var metadata VaultMetadata + + err = json.Unmarshal(metadataBytes, &metadata) + if err != nil { + Debug("Failed to parse vault metadata", "error", err, "secret_name", s.Name) + + return nil, fmt.Errorf("failed to parse vault metadata: %w", err) + } + + DebugWith("Using vault derivation index from metadata", + slog.String("secret_name", s.Name), + slog.String("vault_name", s.vault.GetName()), + slog.Uint64("derivation_index", uint64(metadata.DerivationIndex)), + ) + + // Use mnemonic with the vault's derivation index from metadata + ltIdentity, err := agehd.DeriveIdentity(envMnemonic, metadata.DerivationIndex) + if err != nil { + Debug("Failed to derive long-term key from mnemonic for secret", + "error", err, "secret_name", s.Name) + + return nil, fmt.Errorf( + "failed to derive long-term key from mnemonic: %w", err) + } + + Debug("Successfully derived long-term key from mnemonic", "secret_name", s.Name) + + // Use the long-term key to decrypt the version + return version.GetValue(ltIdentity) +} + +// getLongTermIdentityFromUnlocker uses the unlocker to obtain and parse +// the vault's long-term private key. +func (s *Secret) getLongTermIdentityFromUnlocker( + unlocker Unlocker, +) (*age.X25519Identity, error) { + DebugWith("Getting vault's long-term key using unlocker", + slog.String("secret_name", s.Name), + slog.String("unlocker_type", unlocker.GetType()), + slog.String("unlocker_id", unlocker.GetID()), + ) + + // Step 1: Use the unlocker to get the vault's long-term private key + unlockIdentity, err := unlocker.GetIdentity() + if err != nil { + Debug("Failed to get unlocker identity", + "error", err, "secret_name", s.Name, + "unlocker_type", unlocker.GetType()) + + return nil, fmt.Errorf("failed to get unlocker identity: %w", err) + } + + // Read the encrypted long-term private key from the unlocker directory + encryptedLtPrivKeyPath := filepath.Join(unlocker.GetDirectory(), "longterm.age") + Debug("Reading encrypted long-term private key", "path", encryptedLtPrivKeyPath) + + encryptedLtPrivKey, err := afero.ReadFile( + s.vault.GetFilesystem(), encryptedLtPrivKeyPath) + if err != nil { + Debug("Failed to read encrypted long-term private key", + "error", err, "path", encryptedLtPrivKeyPath) + + return nil, fmt.Errorf( + "failed to read encrypted long-term private key: %w", err) + } + + // Decrypt the encrypted long-term private key using the unlocker + Debug("Decrypting long-term private key using unlocker", "secret_name", s.Name) + + ltPrivKeyBuffer, err := DecryptWithIdentity(encryptedLtPrivKey, unlockIdentity) + if err != nil { + Debug("Failed to decrypt long-term private key", + "error", err, "secret_name", s.Name) + + return nil, fmt.Errorf("failed to decrypt long-term private key: %w", err) + } + defer ltPrivKeyBuffer.Destroy() + + // Parse the long-term private key + Debug("Parsing long-term private key", "secret_name", s.Name) + + ltIdentity, err := age.ParseX25519Identity(ltPrivKeyBuffer.String()) + if err != nil { + Debug("Failed to parse long-term private key", + "error", err, "secret_name", s.Name) + + return nil, fmt.Errorf("failed to parse long-term private key: %w", err) + } + + return ltIdentity, nil +} + // GetCurrentVault gets the current vault from the file system // This function is a wrapper around the actual implementation in the vault package // and exists to break the import cycle. +// +//nolint:ireturn // must return the interface to break the import cycle func GetCurrentVault(fs afero.Fs, stateDir string) (VaultInterface, error) { // This is a forward declaration. The actual implementation is provided // by the vault package when it calls RegisterGetCurrentVaultFunc. if getCurrentVaultFunc == nil { - return nil, fmt.Errorf("GetCurrentVault function not registered") + return nil, errGetCurrentVaultNotRegistered } return getCurrentVaultFunc(fs, stateDir) @@ -288,8 +344,10 @@ func GetCurrentVault(fs afero.Fs, stateDir string) (VaultInterface, error) { //nolint:gochecknoglobals // Required to break import cycle var getCurrentVaultFunc func(fs afero.Fs, stateDir string) (VaultInterface, error) -// RegisterGetCurrentVaultFunc allows the vault package to register its implementation -// of GetCurrentVault to break the import cycle -func RegisterGetCurrentVaultFunc(fn func(fs afero.Fs, stateDir string) (VaultInterface, error)) { +// RegisterGetCurrentVaultFunc allows the vault package to register its +// implementation of GetCurrentVault to break the import cycle +func RegisterGetCurrentVaultFunc( + fn func(fs afero.Fs, stateDir string) (VaultInterface, error), +) { getCurrentVaultFunc = fn } diff --git a/internal/secret/secret_test.go b/internal/secret/secret_test.go index a8560a1..81b6fd2 100644 --- a/internal/secret/secret_test.go +++ b/internal/secret/secret_test.go @@ -1,7 +1,8 @@ +//nolint:testpackage // white-box test of unexported internals package secret import ( - "fmt" + "errors" "os" "path/filepath" "strings" @@ -14,6 +15,17 @@ import ( "github.com/stretchr/testify/require" ) +// testMnemonicValue is the standard BIP39 test vector mnemonic. +// +//nolint:dupword // BIP39 test mnemonic repeats words by design +const testMnemonicValue = "abandon abandon abandon abandon abandon abandon " + + "abandon abandon abandon abandon abandon about" + +var ( + errMnemonicNotSet = errors.New("SB_SECRET_MNEMONIC not set") + errNotImplementedInMock = errors.New("not implemented in mock") +) + // MockVault is a test implementation of the VaultInterface type MockVault struct { name string @@ -30,14 +42,18 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool) // Create secret directory with proper storage name conversion storageName := strings.ReplaceAll(name, "/", "%") secretDir := filepath.Join(m.directory, "secrets.d", storageName) - if err := m.fs.MkdirAll(secretDir, 0o700); err != nil { + + err := m.fs.MkdirAll(secretDir, 0o700) + if err != nil { return err } // Create version directory with proper path versionName := "20240101.001" // Use a fixed version name for testing versionDir := filepath.Join(secretDir, "versions", versionName) - if err := m.fs.MkdirAll(versionDir, 0o700); err != nil { + + err = m.fs.MkdirAll(versionDir, 0o700) + if err != nil { return err } @@ -47,7 +63,7 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool) // Derive long-term key using the vault's derivation index mnemonic := os.Getenv(EnvMnemonic) if mnemonic == "" { - return fmt.Errorf("SB_SECRET_MNEMONIC not set") + return errMnemonicNotSet } ltIdentity, err := agehd.DeriveIdentity(mnemonic, m.derivationIndex) @@ -56,13 +72,54 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool) } // Write long-term public key if it doesn't exist - if _, err := m.fs.Stat(ltPubKeyPath); os.IsNotExist(err) { + _, err = m.fs.Stat(ltPubKeyPath) + if os.IsNotExist(err) { pubKey := ltIdentity.Recipient().String() - if err := afero.WriteFile(m.fs, ltPubKeyPath, []byte(pubKey), 0o600); err != nil { + + err = afero.WriteFile(m.fs, ltPubKeyPath, []byte(pubKey), 0o600) + if err != nil { return err } } + err = m.writeVersionFiles(versionDir, value, ltIdentity) + if err != nil { + return err + } + + // Create current file pointing to the version (just the version name) + currentLink := filepath.Join(secretDir, "current") + + return afero.WriteFile(m.fs, currentLink, []byte(versionName), 0o600) +} + +func (m *MockVault) GetName() string { + return m.name +} + +//nolint:ireturn // implements VaultInterface +func (m *MockVault) GetFilesystem() afero.Fs { + return m.fs +} + +//nolint:ireturn // implements VaultInterface +func (m *MockVault) GetCurrentUnlocker() (Unlocker, error) { + return nil, errNotImplementedInMock +} + +func (m *MockVault) CreatePassphraseUnlocker( + _ *memguard.LockedBuffer, +) (*PassphraseUnlocker, error) { + return nil, errNotImplementedInMock +} + +// writeVersionFiles generates a version keypair and writes the version +// key and value files for the mock vault. +func (m *MockVault) writeVersionFiles( + versionDir string, + value *memguard.LockedBuffer, + ltIdentity *age.X25519Identity, +) error { // Generate version-specific keypair versionIdentity, err := age.GenerateX25519Identity() if err != nil { @@ -71,7 +128,10 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool) // Write version public key pubKeyPath := filepath.Join(versionDir, "pub.age") - if err := afero.WriteFile(m.fs, pubKeyPath, []byte(versionIdentity.Recipient().String()), 0o600); err != nil { + + err = afero.WriteFile( + m.fs, pubKeyPath, []byte(versionIdentity.Recipient().String()), 0o600) + if err != nil { return err } @@ -83,60 +143,32 @@ func (m *MockVault) AddSecret(name string, value *memguard.LockedBuffer, _ bool) // Write encrypted value valuePath := filepath.Join(versionDir, "value.age") - if err := afero.WriteFile(m.fs, valuePath, encryptedValue, 0o600); err != nil { + + err = afero.WriteFile(m.fs, valuePath, encryptedValue, 0o600) + if err != nil { return err } // Encrypt version private key to long-term public key versionPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(versionIdentity.String())) defer versionPrivKeyBuffer.Destroy() - encryptedPrivKey, err := EncryptToRecipient(versionPrivKeyBuffer, ltIdentity.Recipient()) + + encryptedPrivKey, err := EncryptToRecipient( + versionPrivKeyBuffer, ltIdentity.Recipient()) if err != nil { return err } // Write encrypted version private key privKeyPath := filepath.Join(versionDir, "priv.age") - if err := afero.WriteFile(m.fs, privKeyPath, encryptedPrivKey, 0o600); err != nil { - return err - } - // Create current file pointing to the version (just the version name) - currentLink := filepath.Join(secretDir, "current") - if err := afero.WriteFile(m.fs, currentLink, []byte(versionName), 0o600); err != nil { - return err - } - - return nil + return afero.WriteFile(m.fs, privKeyPath, encryptedPrivKey, 0o600) } -func (m *MockVault) GetName() string { - return m.name -} - -func (m *MockVault) GetFilesystem() afero.Fs { - return m.fs -} - -func (m *MockVault) GetCurrentUnlocker() (Unlocker, error) { - return nil, nil -} - -func (m *MockVault) CreatePassphraseUnlocker(_ *memguard.LockedBuffer) (*PassphraseUnlocker, error) { - return nil, nil -} - -func TestPerSecretKeyFunctionality(t *testing.T) { - // Create an in-memory filesystem for testing - fs := afero.NewMemMapFs() - - // Set test mnemonic for direct encryption/decryption - testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about" - t.Setenv(EnvMnemonic, testMnemonic) - - // Set up a test vault structure - baseDir := "/test-config/berlin.sneak.pkg.secret" - vaultDir := filepath.Join(baseDir, "vaults.d", "test-vault") +// setupMockVaultDirs creates the vault directory structure, long-term +// public key, and current vault pointer for tests. +func setupMockVaultDirs(t *testing.T, fs afero.Fs, baseDir, vaultDir string) { + t.Helper() // Create vault directory structure err := fs.MkdirAll(filepath.Join(vaultDir, "secrets.d"), DirPerms) @@ -145,13 +177,14 @@ func TestPerSecretKeyFunctionality(t *testing.T) { } // Generate a long-term keypair for the vault using the test mnemonic - ltIdentity, err := agehd.DeriveIdentity(testMnemonic, 0) + ltIdentity, err := agehd.DeriveIdentity(testMnemonicValue, 0) if err != nil { t.Fatalf("Failed to generate long-term identity: %v", err) } // Write long-term public key ltPubKeyPath := filepath.Join(vaultDir, "pub.age") + err = afero.WriteFile( fs, ltPubKeyPath, @@ -164,10 +197,56 @@ func TestPerSecretKeyFunctionality(t *testing.T) { // Set current vault currentVaultPath := filepath.Join(baseDir, "currentvault") + err = afero.WriteFile(fs, currentVaultPath, []byte(vaultDir), FilePerms) if err != nil { t.Fatalf("Failed to set current vault: %v", err) } +} + +// verifySecretFiles checks that AddSecret created the expected version +// files for the secret. +func verifySecretFiles(t *testing.T, fs afero.Fs, vaultDir, secretName string) { + t.Helper() + + secretDir := filepath.Join(vaultDir, "secrets.d", secretName) + + // Check versions directory exists + versionsDir := filepath.Join(secretDir, "versions") + + versionsDirExists, err := afero.DirExists(fs, versionsDir) + if err != nil || !versionsDirExists { + t.Fatalf("versions directory was not created") + } + + // Check current file exists and points at a version + currentVersion, err := GetCurrentVersion(fs, secretDir) + if err != nil { + t.Fatalf("Failed to get current version: %v", err) + } + + // Check value.age exists in the version directory + versionDir := filepath.Join(versionsDir, currentVersion) + + valueExists, err := afero.Exists(fs, filepath.Join(versionDir, "value.age")) + if err != nil || !valueExists { + t.Fatalf("value.age file was not created in version directory") + } +} + +//nolint:paralleltest // uses t.Setenv (process-global environment) +func TestPerSecretKeyFunctionality(t *testing.T) { + // Create an in-memory filesystem for testing + fs := afero.NewMemMapFs() + + // Set test mnemonic for direct encryption/decryption + t.Setenv(EnvMnemonic, testMnemonicValue) + + // Set up a test vault structure + baseDir := "/test-config/berlin.sneak.pkg.secret" + vaultDir := filepath.Join(baseDir, "vaults.d", "test-vault") + + setupMockVaultDirs(t, fs, baseDir, vaultDir) // Create vault instance using the mock vault vault := &MockVault{ @@ -193,30 +272,7 @@ func TestPerSecretKeyFunctionality(t *testing.T) { } // Verify that all expected files were created - secretDir := filepath.Join(vaultDir, "secrets.d", secretName) - - // Check versions directory exists - versionsDir := filepath.Join(secretDir, "versions") - versionsDirExists, err := afero.DirExists(fs, versionsDir) - if err != nil || !versionsDirExists { - t.Fatalf("versions directory was not created") - } - - // Check current symlink exists - currentVersion, err := GetCurrentVersion(fs, secretDir) - if err != nil { - t.Fatalf("Failed to get current version: %v", err) - } - - // Check value.age exists in the version directory - versionDir := filepath.Join(versionsDir, currentVersion) - valueExists, err := afero.Exists( - fs, - filepath.Join(versionDir, "value.age"), - ) - if err != nil || !valueExists { - t.Fatalf("value.age file was not created in version directory") - } + verifySecretFiles(t, fs, vaultDir, secretName) t.Logf("All expected files created successfully with versioning") }) @@ -245,9 +301,11 @@ func TestPerSecretKeyFunctionality(t *testing.T) { if err != nil { t.Fatalf("Error checking if secret exists: %v", err) } + if !exists { t.Fatalf("Secret should exist but Exists() returned false") } + t.Logf("Secret.Exists() works correctly") }) } @@ -274,6 +332,8 @@ func isValidSecretName(name string) bool { } func TestSecretNameValidation(t *testing.T) { + t.Parallel() + tests := []struct { name string valid bool @@ -293,6 +353,8 @@ func TestSecretNameValidation(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { + t.Parallel() + result := isValidSecretName(test.name) if result != test.valid { t.Errorf( @@ -311,13 +373,13 @@ func TestSecretGetValueWithEnvMnemonicUsesVaultDerivationIndex(t *testing.T) { // instead of the vault's actual derivation index when using environment mnemonic // Set up test mnemonic - testMnemonic := "abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon abandon about" - t.Setenv(EnvMnemonic, testMnemonic) + t.Setenv(EnvMnemonic, testMnemonicValue) // Create temporary directory for vaults fs := afero.NewOsFs() tempDir, err := afero.TempDir(fs, "", "secret-test-") require.NoError(t, err) + defer func() { _ = fs.RemoveAll(tempDir) }() diff --git a/internal/secret/seunlocker_stub.go b/internal/secret/seunlocker_stub.go index e1f819a..fc5bc04 100644 --- a/internal/secret/seunlocker_stub.go +++ b/internal/secret/seunlocker_stub.go @@ -1,22 +1,25 @@ //go:build !darwin -// +build !darwin package secret import ( - "fmt" + "errors" "filippo.io/age" "github.com/spf13/afero" ) -var errSENotSupported = fmt.Errorf( +// seUnlockerType is the type string for Secure Enclave unlockers. +const seUnlockerType = "secure-enclave" + +var errSENotSupported = errors.New( "secure enclave unlockers are only supported on macOS", ) // SecureEnclaveUnlockerMetadata is a stub for non-Darwin platforms. type SecureEnclaveUnlockerMetadata struct { UnlockerMetadata + SEKeyLabel string `json:"seKeyLabel"` SEKeyHash string `json:"seKeyHash"` } @@ -28,6 +31,21 @@ type SecureEnclaveUnlocker struct { fs afero.Fs } +// NewSecureEnclaveUnlocker creates a stub SecureEnclaveUnlocker on +// non-Darwin platforms. The returned instance's methods that require +// macOS functionality will return errors. +func NewSecureEnclaveUnlocker( + fs afero.Fs, + directory string, + metadata UnlockerMetadata, +) *SecureEnclaveUnlocker { + return &SecureEnclaveUnlocker{ + Directory: directory, + Metadata: metadata, + fs: fs, + } +} + // GetIdentity returns an error on non-Darwin platforms. func (s *SecureEnclaveUnlocker) GetIdentity() (*age.X25519Identity, error) { return nil, errSENotSupported @@ -35,7 +53,7 @@ func (s *SecureEnclaveUnlocker) GetIdentity() (*age.X25519Identity, error) { // GetType returns the unlocker type. func (s *SecureEnclaveUnlocker) GetType() string { - return "secure-enclave" + return seUnlockerType } // GetMetadata returns the unlocker metadata. @@ -50,10 +68,7 @@ func (s *SecureEnclaveUnlocker) GetDirectory() string { // GetID returns the unlocker ID. func (s *SecureEnclaveUnlocker) GetID() string { - return fmt.Sprintf( - "%s-secure-enclave", - s.Metadata.CreatedAt.Format("2006-01-02.15.04"), - ) + return s.Metadata.CreatedAt.Format("2006-01-02.15.04") + "-" + seUnlockerType } // Remove returns an error on non-Darwin platforms. @@ -61,20 +76,6 @@ func (s *SecureEnclaveUnlocker) Remove() error { return errSENotSupported } -// NewSecureEnclaveUnlocker creates a stub SecureEnclaveUnlocker on non-Darwin platforms. -// The returned instance's methods that require macOS functionality will return errors. -func NewSecureEnclaveUnlocker( - fs afero.Fs, - directory string, - metadata UnlockerMetadata, -) *SecureEnclaveUnlocker { - return &SecureEnclaveUnlocker{ - Directory: directory, - Metadata: metadata, - fs: fs, - } -} - // CreateSecureEnclaveUnlocker returns an error on non-Darwin platforms. func CreateSecureEnclaveUnlocker( _ afero.Fs, diff --git a/internal/secret/seunlocker_stub_test.go b/internal/secret/seunlocker_stub_test.go index bac86dc..1b04a69 100644 --- a/internal/secret/seunlocker_stub_test.go +++ b/internal/secret/seunlocker_stub_test.go @@ -1,6 +1,6 @@ //go:build !darwin -// +build !darwin +//nolint:testpackage // white-box test asserting unexported sentinel errors package secret import ( @@ -13,19 +13,21 @@ import ( ) func TestNewSecureEnclaveUnlocker(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() dir := "/tmp/test-se-unlocker" metadata := UnlockerMetadata{ - Type: "secure-enclave", + Type: seUnlockerType, CreatedAt: time.Date(2026, 1, 15, 10, 30, 0, 0, time.UTC), - Flags: []string{"secure-enclave", "macos"}, + Flags: []string{seUnlockerType, "macos"}, } unlocker := NewSecureEnclaveUnlocker(fs, dir, metadata) require.NotNil(t, unlocker, "NewSecureEnclaveUnlocker should return a valid instance") // Test GetType returns correct type - assert.Equal(t, "secure-enclave", unlocker.GetType()) + assert.Equal(t, seUnlockerType, unlocker.GetType()) // Test GetMetadata returns the metadata we passed in assert.Equal(t, metadata, unlocker.GetMetadata()) @@ -39,9 +41,11 @@ func TestNewSecureEnclaveUnlocker(t *testing.T) { } func TestSecureEnclaveUnlockerGetIdentityReturnsError(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() metadata := UnlockerMetadata{ - Type: "secure-enclave", + Type: seUnlockerType, CreatedAt: time.Now().UTC(), } @@ -49,37 +53,43 @@ func TestSecureEnclaveUnlockerGetIdentityReturnsError(t *testing.T) { identity, err := unlocker.GetIdentity() assert.Nil(t, identity) - assert.Error(t, err) - assert.ErrorIs(t, err, errSENotSupported) + require.Error(t, err) + require.ErrorIs(t, err, errSENotSupported) } func TestSecureEnclaveUnlockerRemoveReturnsError(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() metadata := UnlockerMetadata{ - Type: "secure-enclave", + Type: seUnlockerType, CreatedAt: time.Now().UTC(), } unlocker := NewSecureEnclaveUnlocker(fs, "/tmp/test", metadata) err := unlocker.Remove() - assert.Error(t, err) - assert.ErrorIs(t, err, errSENotSupported) + require.Error(t, err) + require.ErrorIs(t, err, errSENotSupported) } func TestCreateSecureEnclaveUnlockerReturnsError(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() unlocker, err := CreateSecureEnclaveUnlocker(fs, "/tmp/test") assert.Nil(t, unlocker) - assert.Error(t, err) - assert.ErrorIs(t, err, errSENotSupported) + require.Error(t, err) + require.ErrorIs(t, err, errSENotSupported) } func TestSecureEnclaveUnlockerImplementsInterface(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() metadata := UnlockerMetadata{ - Type: "secure-enclave", + Type: seUnlockerType, CreatedAt: time.Now().UTC(), } diff --git a/internal/secret/validation_test.go b/internal/secret/validation_test.go index 9ab1909..4423547 100644 --- a/internal/secret/validation_test.go +++ b/internal/secret/validation_test.go @@ -1,3 +1,4 @@ +//nolint:testpackage // white-box test of unexported internals package secret import ( @@ -5,148 +6,60 @@ import ( ) func TestValidateGPGKeyID(t *testing.T) { + t.Parallel() + tests := []struct { name string keyID string wantErr bool }{ // Valid cases + {"valid email address", "test@example.com", false}, + {"valid email with dots and hyphens", "test.user-name@example-domain.co.uk", false}, + {"valid email with plus", "test+tag@example.com", false}, + {"valid short key ID (8 hex chars)", "ABCDEF12", false}, + {"valid long key ID (16 hex chars)", "ABCDEF1234567890", false}, { - name: "valid email address", - keyID: "test@example.com", - wantErr: false, + "valid fingerprint (40 hex chars)", + "ABCDEF1234567890ABCDEF1234567890ABCDEF12", false, }, { - name: "valid email with dots and hyphens", - keyID: "test.user-name@example-domain.co.uk", - wantErr: false, - }, - { - name: "valid email with plus", - keyID: "test+tag@example.com", - wantErr: false, - }, - { - name: "valid short key ID (8 hex chars)", - keyID: "ABCDEF12", - wantErr: false, - }, - { - name: "valid long key ID (16 hex chars)", - keyID: "ABCDEF1234567890", - wantErr: false, - }, - { - name: "valid fingerprint (40 hex chars)", - keyID: "ABCDEF1234567890ABCDEF1234567890ABCDEF12", - wantErr: false, - }, - { - name: "valid lowercase hex fingerprint", - keyID: "abcdef1234567890abcdef1234567890abcdef12", - wantErr: false, - }, - { - name: "valid mixed case hex", - keyID: "AbCdEf1234567890", - wantErr: false, + "valid lowercase hex fingerprint", + "abcdef1234567890abcdef1234567890abcdef12", false, }, + {"valid mixed case hex", "AbCdEf1234567890", false}, // Invalid cases + {"empty key ID", "", true}, + {"key ID with spaces", "test user@example.com", true}, + {"key ID with semicolon (command injection)", "test@example.com; rm -rf /", true}, { - name: "empty key ID", - keyID: "", - wantErr: true, + "key ID with pipe (command injection)", + "test@example.com | cat /etc/passwd", true, }, + {"key ID with backticks (command injection)", "test@example.com`whoami`", true}, { - name: "key ID with spaces", - keyID: "test user@example.com", - wantErr: true, - }, - { - name: "key ID with semicolon (command injection)", - keyID: "test@example.com; rm -rf /", - wantErr: true, - }, - { - name: "key ID with pipe (command injection)", - keyID: "test@example.com | cat /etc/passwd", - wantErr: true, - }, - { - name: "key ID with backticks (command injection)", - keyID: "test@example.com`whoami`", - wantErr: true, - }, - { - name: "key ID with dollar sign (command injection)", - keyID: "test@example.com$(whoami)", - wantErr: true, - }, - { - name: "key ID with quotes", - keyID: "test\"@example.com", - wantErr: true, - }, - { - name: "key ID with single quotes", - keyID: "test'@example.com", - wantErr: true, - }, - { - name: "key ID with backslash", - keyID: "test\\@example.com", - wantErr: true, - }, - { - name: "key ID with newline", - keyID: "test@example.com\nrm -rf /", - wantErr: true, - }, - { - name: "key ID with carriage return", - keyID: "test@example.com\rrm -rf /", - wantErr: true, - }, - { - name: "hex with invalid length (7 chars)", - keyID: "ABCDEF1", - wantErr: true, - }, - { - name: "hex with invalid length (9 chars)", - keyID: "ABCDEF123", - wantErr: true, - }, - { - name: "hex with non-hex characters", - keyID: "ABCDEFGH", - wantErr: true, - }, - { - name: "mixed format (email with hex)", - keyID: "test@ABCDEF12", - wantErr: true, - }, - { - name: "key ID with ampersand", - keyID: "test@example.com & echo test", - wantErr: true, - }, - { - name: "key ID with redirect", - keyID: "test@example.com > /tmp/test", - wantErr: true, - }, - { - name: "key ID with null byte", - keyID: "test@example.com\x00", - wantErr: true, + "key ID with dollar sign (command injection)", + "test@example.com$(whoami)", true, }, + {"key ID with quotes", "test\"@example.com", true}, + {"key ID with single quotes", "test'@example.com", true}, + {"key ID with backslash", "test\\@example.com", true}, + {"key ID with newline", "test@example.com\nrm -rf /", true}, + {"key ID with carriage return", "test@example.com\rrm -rf /", true}, + {"hex with invalid length (7 chars)", "ABCDEF1", true}, + {"hex with invalid length (9 chars)", "ABCDEF123", true}, + {"hex with non-hex characters", "ABCDEFGH", true}, + {"mixed format (email with hex)", "test@ABCDEF12", true}, + {"key ID with ampersand", "test@example.com & echo test", true}, + {"key ID with redirect", "test@example.com > /tmp/test", true}, + {"key ID with null byte", "test@example.com\x00", true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + t.Parallel() + err := validateGPGKeyID(tt.keyID) if (err != nil) != tt.wantErr { t.Errorf("validateGPGKeyID() error = %v, wantErr %v", err, tt.wantErr) diff --git a/internal/secret/version.go b/internal/secret/version.go index 0525021..39efed2 100644 --- a/internal/secret/version.go +++ b/internal/secret/version.go @@ -2,6 +2,7 @@ package secret import ( "encoding/json" + "errors" "fmt" "log/slog" "path/filepath" @@ -20,12 +21,17 @@ const ( maxVersionsPerDay = 999 ) +var ( + errMaxVersionsPerDay = errors.New("exceeded maximum versions per day (999)") + errNilValueBuffer = errors.New("value buffer is nil") +) + // VersionMetadata contains information about a secret version type VersionMetadata struct { ID string `json:"id"` // ULID CreatedAt *time.Time `json:"createdAt,omitempty"` // When version was created NotBefore *time.Time `json:"notBefore,omitempty"` // When this version becomes active - NotAfter *time.Time `json:"notAfter,omitempty"` // When this version expires (nil = current) + NotAfter *time.Time `json:"notAfter,omitempty"` // Expiry (nil = current) } // Version represents a version of a secret @@ -75,7 +81,8 @@ func GenerateVersionName(fs afero.Fs, secretDir string) (string, error) { versionsDir := filepath.Join(secretDir, "versions") // Ensure versions directory exists - if err := fs.MkdirAll(versionsDir, DirPerms); err != nil { + err := fs.MkdirAll(versionsDir, DirPerms) + if err != nil { return "", fmt.Errorf("failed to create versions directory: %w", err) } @@ -101,8 +108,11 @@ func GenerateVersionName(fs afero.Fs, secretDir string) (string, error) { } var serial int - if _, err := fmt.Sscanf(parts[1], "%03d", &serial); err != nil { - Warn("Skipping malformed version directory name", "name", entry.Name(), "error", err) + + _, err := fmt.Sscanf(parts[1], "%03d", &serial) + if err != nil { + Warn("Skipping malformed version directory name", + "name", entry.Name(), "error", err) continue } @@ -115,7 +125,7 @@ func GenerateVersionName(fs afero.Fs, secretDir string) (string, error) { // Generate new version name newSerial := maxSerial + 1 if newSerial > maxVersionsPerDay { - return "", fmt.Errorf("exceeded maximum versions per day (999)") + return "", errMaxVersionsPerDay } return fmt.Sprintf("%s.%03d", today, newSerial), nil @@ -124,7 +134,7 @@ func GenerateVersionName(fs afero.Fs, secretDir string) (string, error) { // Save saves the version metadata and value func (sv *Version) Save(value *memguard.LockedBuffer) error { if value == nil { - return fmt.Errorf("value buffer is nil") + return errNilValueBuffer } DebugWith("Saving secret version", @@ -136,14 +146,16 @@ func (sv *Version) Save(value *memguard.LockedBuffer) error { fs := sv.vault.GetFilesystem() // Create version directory - if err := fs.MkdirAll(sv.Directory, DirPerms); err != nil { + err := fs.MkdirAll(sv.Directory, DirPerms) + if err != nil { Debug("Failed to create version directory", "error", err, "dir", sv.Directory) return fmt.Errorf("failed to create version directory: %w", err) } - // Step 1: Generate a new keypair for this version + // Generate a new keypair for this version Debug("Generating version-specific keypair", "version", sv.Version) + versionIdentity, err := age.GenerateX25519Identity() if err != nil { Debug("Failed to generate version keypair", "error", err, "version", sv.Version) @@ -151,110 +163,33 @@ func (sv *Version) Save(value *memguard.LockedBuffer) error { return fmt.Errorf("failed to generate version keypair: %w", err) } - versionPublicKey := versionIdentity.Recipient().String() // Store private key in memguard buffer immediately - versionPrivateKeyBuffer := memguard.NewBufferFromBytes([]byte(versionIdentity.String())) + versionPrivateKeyBuffer := memguard.NewBufferFromBytes( + []byte(versionIdentity.String())) defer versionPrivateKeyBuffer.Destroy() DebugWith("Generated version keypair", slog.String("version", sv.Version), - slog.String("public_key", versionPublicKey), + slog.String("public_key", versionIdentity.Recipient().String()), ) - // Step 2: Store the version's public key - pubKeyPath := filepath.Join(sv.Directory, "pub.age") - Debug("Writing version public key", "path", pubKeyPath) - if err := afero.WriteFile(fs, pubKeyPath, []byte(versionPublicKey), FilePerms); err != nil { - Debug("Failed to write version public key", "error", err, "path", pubKeyPath) - - return fmt.Errorf("failed to write version public key: %w", err) - } - - // Step 3: Encrypt the value to the version's public key - Debug("Encrypting value to version's public key", "version", sv.Version) - encryptedValue, err := EncryptToRecipient(value, versionIdentity.Recipient()) + err = sv.writePublicKeyAndValue(fs, versionIdentity, value) if err != nil { - Debug("Failed to encrypt version value", "error", err, "version", sv.Version) - - return fmt.Errorf("failed to encrypt version value: %w", err) + return err } - // Step 4: Store the encrypted value - valuePath := filepath.Join(sv.Directory, "value.age") - Debug("Writing encrypted version value", "path", valuePath) - if err := afero.WriteFile(fs, valuePath, encryptedValue, FilePerms); err != nil { - Debug("Failed to write encrypted version value", "error", err, "path", valuePath) - - return fmt.Errorf("failed to write encrypted version value: %w", err) - } - - // Step 5: Get vault's long-term public key for encrypting the version's private key - vaultDir, _ := sv.vault.GetDirectory() - ltPubKeyPath := filepath.Join(vaultDir, "pub.age") - Debug("Reading long-term public key", "path", ltPubKeyPath) - - ltPubKeyData, err := afero.ReadFile(fs, ltPubKeyPath) + err = sv.writeEncryptedPrivateKey(fs, versionPrivateKeyBuffer) if err != nil { - Debug("Failed to read long-term public key", "error", err, "path", ltPubKeyPath) - - return fmt.Errorf("failed to read long-term public key: %w", err) + return err } - Debug("Parsing long-term public key") - ltRecipient, err := age.ParseX25519Recipient(string(ltPubKeyData)) + err = sv.writeEncryptedMetadata(fs, versionIdentity) if err != nil { - Debug("Failed to parse long-term public key", "error", err) - - return fmt.Errorf("failed to parse long-term public key: %w", err) + return err } - // Step 6: Encrypt the version's private key to the long-term public key - Debug("Encrypting version private key to long-term public key", "version", sv.Version) - encryptedPrivKey, err := EncryptToRecipient(versionPrivateKeyBuffer, ltRecipient) - if err != nil { - Debug("Failed to encrypt version private key", "error", err, "version", sv.Version) - - return fmt.Errorf("failed to encrypt version private key: %w", err) - } - - // Step 7: Store the encrypted private key - privKeyPath := filepath.Join(sv.Directory, "priv.age") - Debug("Writing encrypted version private key", "path", privKeyPath) - if err := afero.WriteFile(fs, privKeyPath, encryptedPrivKey, FilePerms); err != nil { - Debug("Failed to write encrypted version private key", "error", err, "path", privKeyPath) - - return fmt.Errorf("failed to write encrypted version private key: %w", err) - } - - // Step 8: Encrypt and store metadata - Debug("Encrypting version metadata", "version", sv.Version) - metadataBytes, err := json.MarshalIndent(sv.Metadata, "", " ") - if err != nil { - Debug("Failed to marshal version metadata", "error", err) - - return fmt.Errorf("failed to marshal version metadata: %w", err) - } - - // Encrypt metadata to the version's public key - metadataBuffer := memguard.NewBufferFromBytes(metadataBytes) - defer metadataBuffer.Destroy() - - encryptedMetadata, err := EncryptToRecipient(metadataBuffer, versionIdentity.Recipient()) - if err != nil { - Debug("Failed to encrypt version metadata", "error", err, "version", sv.Version) - - return fmt.Errorf("failed to encrypt version metadata: %w", err) - } - - metadataPath := filepath.Join(sv.Directory, "metadata.age") - Debug("Writing encrypted version metadata", "path", metadataPath) - if err := afero.WriteFile(fs, metadataPath, encryptedMetadata, FilePerms); err != nil { - Debug("Failed to write encrypted version metadata", "error", err, "path", metadataPath) - - return fmt.Errorf("failed to write encrypted version metadata: %w", err) - } - - Debug("Successfully saved secret version", "version", sv.Version, "secret_name", sv.SecretName) + Debug("Successfully saved secret version", + "version", sv.Version, "secret_name", sv.SecretName) return nil } @@ -270,9 +205,11 @@ func (sv *Version) LoadMetadata(ltIdentity *age.X25519Identity) error { // Step 1: Read encrypted version private key encryptedPrivKeyPath := filepath.Join(sv.Directory, "priv.age") + encryptedPrivKey, err := afero.ReadFile(fs, encryptedPrivKeyPath) if err != nil { - Debug("Failed to read encrypted version private key", "error", err, "path", encryptedPrivKeyPath) + Debug("Failed to read encrypted version private key", + "error", err, "path", encryptedPrivKeyPath) return fmt.Errorf("failed to read encrypted version private key: %w", err) } @@ -296,9 +233,11 @@ func (sv *Version) LoadMetadata(ltIdentity *age.X25519Identity) error { // Step 4: Read encrypted metadata encryptedMetadataPath := filepath.Join(sv.Directory, "metadata.age") + encryptedMetadata, err := afero.ReadFile(fs, encryptedMetadataPath) if err != nil { - Debug("Failed to read encrypted version metadata", "error", err, "path", encryptedMetadataPath) + Debug("Failed to read encrypted version metadata", + "error", err, "path", encryptedMetadataPath) return fmt.Errorf("failed to read encrypted version metadata: %w", err) } @@ -314,20 +253,25 @@ func (sv *Version) LoadMetadata(ltIdentity *age.X25519Identity) error { // Step 6: Unmarshal metadata var metadata VersionMetadata - if err := json.Unmarshal(metadataBuffer.Bytes(), &metadata); err != nil { + + err = json.Unmarshal(metadataBuffer.Bytes(), &metadata) + if err != nil { Debug("Failed to unmarshal version metadata", "error", err, "version", sv.Version) return fmt.Errorf("failed to unmarshal version metadata: %w", err) } sv.Metadata = metadata + Debug("Successfully loaded version metadata", "version", sv.Version) return nil } // GetValue retrieves and decrypts the version value -func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuffer, error) { +func (sv *Version) GetValue( + ltIdentity *age.X25519Identity, +) (*memguard.LockedBuffer, error) { DebugWith("Getting version value", slog.String("secret_name", sv.SecretName), slog.String("version", sv.Version), @@ -345,16 +289,22 @@ func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuf // Step 1: Read encrypted version private key encryptedPrivKeyPath := filepath.Join(sv.Directory, "priv.age") Debug("Reading encrypted version private key", "path", encryptedPrivKeyPath) + encryptedPrivKey, err := afero.ReadFile(fs, encryptedPrivKeyPath) if err != nil { - Debug("Failed to read encrypted version private key", "error", err, "path", encryptedPrivKeyPath) + Debug("Failed to read encrypted version private key", + "error", err, "path", encryptedPrivKeyPath) - return nil, fmt.Errorf("failed to read encrypted version private key: %w", err) + return nil, fmt.Errorf( + "failed to read encrypted version private key: %w", err) } - Debug("Successfully read encrypted version private key", "path", encryptedPrivKeyPath, "size", len(encryptedPrivKey)) + + Debug("Successfully read encrypted version private key", + "path", encryptedPrivKeyPath, "size", len(encryptedPrivKey)) // Step 2: Decrypt version private key using long-term key Debug("Decrypting version private key with long-term identity", "version", sv.Version) + versionPrivKeyBuffer, err := DecryptWithIdentity(encryptedPrivKey, ltIdentity) if err != nil { Debug("Failed to decrypt version private key", "error", err, "version", sv.Version) @@ -362,7 +312,9 @@ func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuf return nil, fmt.Errorf("failed to decrypt version private key: %w", err) } defer versionPrivKeyBuffer.Destroy() - Debug("Successfully decrypted version private key", "version", sv.Version, "size", versionPrivKeyBuffer.Size()) + + Debug("Successfully decrypted version private key", + "version", sv.Version, "size", versionPrivKeyBuffer.Size()) // Step 3: Parse version private key versionIdentity, err := age.ParseX25519Identity(versionPrivKeyBuffer.String()) @@ -375,16 +327,21 @@ func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuf // Step 4: Read encrypted value encryptedValuePath := filepath.Join(sv.Directory, "value.age") Debug("Reading encrypted value", "path", encryptedValuePath) + encryptedValue, err := afero.ReadFile(fs, encryptedValuePath) if err != nil { - Debug("Failed to read encrypted version value", "error", err, "path", encryptedValuePath) + Debug("Failed to read encrypted version value", + "error", err, "path", encryptedValuePath) return nil, fmt.Errorf("failed to read encrypted version value: %w", err) } - Debug("Successfully read encrypted value", "path", encryptedValuePath, "size", len(encryptedValue)) + + Debug("Successfully read encrypted value", + "path", encryptedValuePath, "size", len(encryptedValue)) // Step 5: Decrypt value using version key Debug("Decrypting value with version identity", "version", sv.Version) + valueBuffer, err := DecryptWithIdentity(encryptedValue, versionIdentity) if err != nil { Debug("Failed to decrypt version value", "error", err, "version", sv.Version) @@ -400,6 +357,139 @@ func (sv *Version) GetValue(ltIdentity *age.X25519Identity) (*memguard.LockedBuf return valueBuffer, nil } +// writePublicKeyAndValue stores the version's public key and the value +// encrypted to it. +func (sv *Version) writePublicKeyAndValue( + fs afero.Fs, + versionIdentity *age.X25519Identity, + value *memguard.LockedBuffer, +) error { + versionPublicKey := versionIdentity.Recipient().String() + pubKeyPath := filepath.Join(sv.Directory, "pub.age") + Debug("Writing version public key", "path", pubKeyPath) + + err := afero.WriteFile(fs, pubKeyPath, []byte(versionPublicKey), FilePerms) + if err != nil { + Debug("Failed to write version public key", "error", err, "path", pubKeyPath) + + return fmt.Errorf("failed to write version public key: %w", err) + } + + // Encrypt the value to the version's public key + Debug("Encrypting value to version's public key", "version", sv.Version) + + encryptedValue, err := EncryptToRecipient(value, versionIdentity.Recipient()) + if err != nil { + Debug("Failed to encrypt version value", "error", err, "version", sv.Version) + + return fmt.Errorf("failed to encrypt version value: %w", err) + } + + valuePath := filepath.Join(sv.Directory, "value.age") + Debug("Writing encrypted version value", "path", valuePath) + + err = afero.WriteFile(fs, valuePath, encryptedValue, FilePerms) + if err != nil { + Debug("Failed to write encrypted version value", "error", err, "path", valuePath) + + return fmt.Errorf("failed to write encrypted version value: %w", err) + } + + return nil +} + +// writeEncryptedPrivateKey encrypts the version's private key to the +// vault's long-term public key and stores it. +func (sv *Version) writeEncryptedPrivateKey( + fs afero.Fs, + versionPrivateKeyBuffer *memguard.LockedBuffer, +) error { + vaultDir, _ := sv.vault.GetDirectory() + ltPubKeyPath := filepath.Join(vaultDir, "pub.age") + Debug("Reading long-term public key", "path", ltPubKeyPath) + + ltPubKeyData, err := afero.ReadFile(fs, ltPubKeyPath) + if err != nil { + Debug("Failed to read long-term public key", "error", err, "path", ltPubKeyPath) + + return fmt.Errorf("failed to read long-term public key: %w", err) + } + + Debug("Parsing long-term public key") + + ltRecipient, err := age.ParseX25519Recipient(string(ltPubKeyData)) + if err != nil { + Debug("Failed to parse long-term public key", "error", err) + + return fmt.Errorf("failed to parse long-term public key: %w", err) + } + + Debug("Encrypting version private key to long-term public key", + "version", sv.Version) + + encryptedPrivKey, err := EncryptToRecipient(versionPrivateKeyBuffer, ltRecipient) + if err != nil { + Debug("Failed to encrypt version private key", + "error", err, "version", sv.Version) + + return fmt.Errorf("failed to encrypt version private key: %w", err) + } + + privKeyPath := filepath.Join(sv.Directory, "priv.age") + Debug("Writing encrypted version private key", "path", privKeyPath) + + err = afero.WriteFile(fs, privKeyPath, encryptedPrivKey, FilePerms) + if err != nil { + Debug("Failed to write encrypted version private key", + "error", err, "path", privKeyPath) + + return fmt.Errorf("failed to write encrypted version private key: %w", err) + } + + return nil +} + +// writeEncryptedMetadata encrypts the version metadata to the version's +// public key and stores it. +func (sv *Version) writeEncryptedMetadata( + fs afero.Fs, + versionIdentity *age.X25519Identity, +) error { + Debug("Encrypting version metadata", "version", sv.Version) + + metadataBytes, err := json.MarshalIndent(sv.Metadata, "", " ") + if err != nil { + Debug("Failed to marshal version metadata", "error", err) + + return fmt.Errorf("failed to marshal version metadata: %w", err) + } + + // Encrypt metadata to the version's public key + metadataBuffer := memguard.NewBufferFromBytes(metadataBytes) + defer metadataBuffer.Destroy() + + encryptedMetadata, err := EncryptToRecipient( + metadataBuffer, versionIdentity.Recipient()) + if err != nil { + Debug("Failed to encrypt version metadata", "error", err, "version", sv.Version) + + return fmt.Errorf("failed to encrypt version metadata: %w", err) + } + + metadataPath := filepath.Join(sv.Directory, "metadata.age") + Debug("Writing encrypted version metadata", "path", metadataPath) + + err = afero.WriteFile(fs, metadataPath, encryptedMetadata, FilePerms) + if err != nil { + Debug("Failed to write encrypted version metadata", + "error", err, "path", metadataPath) + + return fmt.Errorf("failed to write encrypted version metadata: %w", err) + } + + return nil +} + // ListVersions lists all versions of a secret func ListVersions(fs afero.Fs, secretDir string) ([]string, error) { versionsDir := filepath.Join(secretDir, "versions") @@ -409,6 +499,7 @@ func ListVersions(fs afero.Fs, secretDir string) ([]string, error) { if err != nil { return nil, fmt.Errorf("failed to check versions directory: %w", err) } + if !exists { return []string{}, nil } @@ -420,6 +511,7 @@ func ListVersions(fs afero.Fs, secretDir string) ([]string, error) { } var versions []string + for _, entry := range entries { if entry.IsDir() { versions = append(versions, entry.Name()) @@ -456,7 +548,8 @@ func SetCurrentVersion(fs afero.Fs, secretDir string, version string) error { _ = fs.Remove(currentPath) // Write just the version name to the file - if err := afero.WriteFile(fs, currentPath, []byte(version), FilePerms); err != nil { + err := afero.WriteFile(fs, currentPath, []byte(version), FilePerms) + if err != nil { return fmt.Errorf("failed to create current version file: %w", err) } diff --git a/internal/secret/version_test.go b/internal/secret/version_test.go index 17ee4d0..a1f2cda 100644 --- a/internal/secret/version_test.go +++ b/internal/secret/version_test.go @@ -32,22 +32,32 @@ // - Long-term key required for all operations // - Concurrent reads handled safely -package secret +package secret_test import ( + "errors" "fmt" "path/filepath" "testing" "time" "filippo.io/age" + "git.eeqj.de/sneak/secret/internal/secret" "github.com/awnumar/memguard" "github.com/spf13/afero" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -// MockVault implements VaultInterface for testing +const ( + testSecretDir = "/test/secret" + testVaultName = "test" + testVaultStateDir = "/test" +) + +var errNotImplementedInMock = errors.New("not implemented in mock") + +// MockVersionVault implements VaultInterface for testing type MockVersionVault struct { Name string fs afero.Fs @@ -60,31 +70,37 @@ func (m *MockVersionVault) GetDirectory() (string, error) { } func (m *MockVersionVault) AddSecret(_ string, _ *memguard.LockedBuffer, _ bool) error { - return fmt.Errorf("not implemented in mock") + return errNotImplementedInMock } func (m *MockVersionVault) GetName() string { return m.Name } +//nolint:ireturn // implements VaultInterface func (m *MockVersionVault) GetFilesystem() afero.Fs { return m.fs } -func (m *MockVersionVault) GetCurrentUnlocker() (Unlocker, error) { - return nil, fmt.Errorf("not implemented in mock") +//nolint:ireturn // implements VaultInterface +func (m *MockVersionVault) GetCurrentUnlocker() (secret.Unlocker, error) { + return nil, errNotImplementedInMock } -func (m *MockVersionVault) CreatePassphraseUnlocker(_ *memguard.LockedBuffer) (*PassphraseUnlocker, error) { - return nil, fmt.Errorf("not implemented in mock") +func (m *MockVersionVault) CreatePassphraseUnlocker( + _ *memguard.LockedBuffer, +) (*secret.PassphraseUnlocker, error) { + return nil, errNotImplementedInMock } func TestGenerateVersionName(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() - secretDir := "/test/secret" + secretDir := testSecretDir // Test first version generation - version1, err := GenerateVersionName(fs, secretDir) + version1, err := secret.GenerateVersionName(fs, secretDir) require.NoError(t, err) assert.Regexp(t, `^\d{8}\.001$`, version1) @@ -94,7 +110,7 @@ func TestGenerateVersionName(t *testing.T) { require.NoError(t, err) // Test second version generation on same day - version2, err := GenerateVersionName(fs, secretDir) + version2, err := secret.GenerateVersionName(fs, secretDir) require.NoError(t, err) assert.Regexp(t, `^\d{8}\.002$`, version2) @@ -104,8 +120,10 @@ func TestGenerateVersionName(t *testing.T) { } func TestGenerateVersionNameMaxSerial(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() - secretDir := "/test/secret" + secretDir := testSecretDir versionsDir := filepath.Join(secretDir, "versions") // Create 999 versions @@ -117,20 +135,22 @@ func TestGenerateVersionNameMaxSerial(t *testing.T) { } // Try to create one more - should fail - _, err := GenerateVersionName(fs, secretDir) - assert.Error(t, err) + _, err := secret.GenerateVersionName(fs, secretDir) + require.Error(t, err) assert.Contains(t, err.Error(), "exceeded maximum versions per day") } func TestNewVersion(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() vault := &MockVersionVault{ - Name: "test", + Name: testVaultName, fs: fs, - stateDir: "/test", + stateDir: testVaultStateDir, } - sv := NewVersion(vault, "test/secret", "20231215.001") + sv := secret.NewVersion(vault, "test/secret", "20231215.001") assert.Equal(t, "test/secret", sv.SecretName) assert.Equal(t, "20231215.001", sv.Version) @@ -140,11 +160,13 @@ func TestNewVersion(t *testing.T) { } func TestSecretVersionSave(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() vault := &MockVersionVault{ - Name: "test", + Name: testVaultName, fs: fs, - stateDir: "/test", + stateDir: testVaultStateDir, } // Create vault directory structure and long-term key @@ -155,18 +177,21 @@ func TestSecretVersionSave(t *testing.T) { // Generate and store long-term public key ltIdentity, err := age.GenerateX25519Identity() require.NoError(t, err) + vault.longTermKey = ltIdentity 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) // Create and save a version - sv := NewVersion(vault, "test/secret", "20231215.001") + sv := secret.NewVersion(vault, "test/secret", "20231215.001") testValue := []byte("test-secret-value") testBuffer := memguard.NewBufferFromBytes(testValue) defer testBuffer.Destroy() + err = sv.Save(testBuffer) require.NoError(t, err) @@ -178,11 +203,13 @@ func TestSecretVersionSave(t *testing.T) { } func TestSecretVersionLoadMetadata(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() vault := &MockVersionVault{ - Name: "test", + Name: testVaultName, fs: fs, - stateDir: "/test", + stateDir: testVaultStateDir, } // Setup vault with long-term key @@ -192,14 +219,16 @@ func TestSecretVersionLoadMetadata(t *testing.T) { ltIdentity, err := age.GenerateX25519Identity() require.NoError(t, err) + vault.longTermKey = ltIdentity 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) // Create and save a version with custom metadata - sv := NewVersion(vault, "test/secret", "20231215.001") + sv := secret.NewVersion(vault, "test/secret", "20231215.001") now := time.Now() epochPlusOne := time.Unix(1, 0) sv.Metadata.NotBefore = &epochPlusOne @@ -207,11 +236,12 @@ func TestSecretVersionLoadMetadata(t *testing.T) { testBuffer := memguard.NewBufferFromBytes([]byte("test-value")) defer testBuffer.Destroy() + err = sv.Save(testBuffer) require.NoError(t, err) // Create new version object and load metadata - sv2 := NewVersion(vault, "test/secret", "20231215.001") + sv2 := secret.NewVersion(vault, "test/secret", "20231215.001") err = sv2.LoadMetadata(ltIdentity) require.NoError(t, err) @@ -223,11 +253,13 @@ func TestSecretVersionLoadMetadata(t *testing.T) { } func TestSecretVersionGetValue(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() vault := &MockVersionVault{ - Name: "test", + Name: testVaultName, fs: fs, - stateDir: "/test", + stateDir: testVaultStateDir, } // Setup vault with long-term key @@ -237,64 +269,77 @@ func TestSecretVersionGetValue(t *testing.T) { ltIdentity, err := age.GenerateX25519Identity() require.NoError(t, err) + vault.longTermKey = ltIdentity 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) // Create and save a version - sv := NewVersion(vault, "test/secret", "20231215.001") + sv := secret.NewVersion(vault, "test/secret", "20231215.001") originalValue := []byte("test-secret-value-12345") expectedValue := make([]byte, len(originalValue)) copy(expectedValue, originalValue) originalBuffer := memguard.NewBufferFromBytes(originalValue) defer originalBuffer.Destroy() + err = sv.Save(originalBuffer) require.NoError(t, err) // Retrieve the value retrievedBuffer, err := sv.GetValue(ltIdentity) require.NoError(t, err) + defer retrievedBuffer.Destroy() assert.Equal(t, expectedValue, retrievedBuffer.Bytes()) } func TestListVersions(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() - secretDir := "/test/secret" + secretDir := testSecretDir versionsDir := filepath.Join(secretDir, "versions") // No versions directory - versions, err := ListVersions(fs, secretDir) + versions, err := secret.ListVersions(fs, secretDir) require.NoError(t, err) assert.Empty(t, versions) // Create some versions - testVersions := []string{"20231215.001", "20231215.002", "20231216.001", "20231214.001"} + testVersions := []string{ + "20231215.001", "20231215.002", "20231216.001", "20231214.001", + } for _, v := range testVersions { err := fs.MkdirAll(filepath.Join(versionsDir, v), 0o755) require.NoError(t, err) } // Create a file (not directory) that should be ignored - err = afero.WriteFile(fs, filepath.Join(versionsDir, "ignore.txt"), []byte("test"), 0o600) + err = afero.WriteFile( + fs, filepath.Join(versionsDir, "ignore.txt"), []byte("test"), 0o600) require.NoError(t, err) // List versions - versions, err = ListVersions(fs, secretDir) + versions, err = secret.ListVersions(fs, secretDir) require.NoError(t, err) // Should be sorted in reverse chronological order - expected := []string{"20231216.001", "20231215.002", "20231215.001", "20231214.001"} + expected := []string{ + "20231216.001", "20231215.002", "20231215.001", "20231214.001", + } assert.Equal(t, expected, versions) } func TestGetCurrentVersion(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() - secretDir := "/test/secret" + secretDir := testSecretDir // The current file contains just the version name currentPath := filepath.Join(secretDir, "current") @@ -304,39 +349,43 @@ func TestGetCurrentVersion(t *testing.T) { err = afero.WriteFile(fs, currentPath, []byte("20231216.001"), 0o600) require.NoError(t, err) - version, err := GetCurrentVersion(fs, secretDir) + version, err := secret.GetCurrentVersion(fs, secretDir) require.NoError(t, err) assert.Equal(t, "20231216.001", version) } func TestSetCurrentVersion(t *testing.T) { + t.Parallel() + fs := afero.NewMemMapFs() - secretDir := "/test/secret" + secretDir := testSecretDir err := fs.MkdirAll(secretDir, 0o755) require.NoError(t, err) // Set current version - err = SetCurrentVersion(fs, secretDir, "20231216.002") + err = secret.SetCurrentVersion(fs, secretDir, "20231216.002") require.NoError(t, err) // Verify it was set - version, err := GetCurrentVersion(fs, secretDir) + version, err := secret.GetCurrentVersion(fs, secretDir) require.NoError(t, err) assert.Equal(t, "20231216.002", version) // Update to different version - err = SetCurrentVersion(fs, secretDir, "20231217.001") + err = secret.SetCurrentVersion(fs, secretDir, "20231217.001") require.NoError(t, err) - version, err = GetCurrentVersion(fs, secretDir) + version, err = secret.GetCurrentVersion(fs, secretDir) require.NoError(t, err) assert.Equal(t, "20231217.001", version) } func TestVersionMetadataTimestamps(t *testing.T) { + t.Parallel() + // Test that all timestamp fields behave consistently as pointers - vm := VersionMetadata{ + vm := secret.VersionMetadata{ ID: "test-id", } @@ -368,5 +417,6 @@ func TestVersionMetadataTimestamps(t *testing.T) { // Helper function func fileExists(fs afero.Fs, path string) bool { exists, _ := afero.Exists(fs, path) + return exists } diff --git a/internal/vault/errors.go b/internal/vault/errors.go new file mode 100644 index 0000000..c9d8cdc --- /dev/null +++ b/internal/vault/errors.go @@ -0,0 +1,53 @@ +package vault + +import "errors" + +// Sentinel errors returned by vault operations. +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.\-_]+. + ErrInvalidVaultName = errors.New( + "invalid vault name: must match pattern [a-z0-9.\\-_]+", + ) + + // ErrVaultNotFound indicates the named vault does not exist. + ErrVaultNotFound = errors.New("vault 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.\-_/]+. + ErrInvalidSecretName = errors.New( + "invalid secret name: must match pattern [a-z0-9.\\-_/]+", + ) + + // ErrSecretExists indicates the secret already exists and --force + // was not supplied. + ErrSecretExists = errors.New( + "secret already exists (use --force to overwrite)", + ) + + // ErrSecretNotFound indicates the named secret does not exist. + ErrSecretNotFound = errors.New("secret not found") + + // ErrVersionNotFound indicates the requested secret version does not + // exist. + ErrVersionNotFound = errors.New("version not found") + + // ErrNoVersions indicates the source secret has no versions. + ErrNoVersions = errors.New("source secret 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. + ErrUnlockerNotFound = errors.New("unlocker not found") +) diff --git a/internal/vault/integration_test.go b/internal/vault/integration_test.go index 01fc497..e173254 100644 --- a/internal/vault/integration_test.go +++ b/internal/vault/integration_test.go @@ -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") + } +} diff --git a/internal/vault/integration_version_test.go b/internal/vault/integration_version_test.go index ec83858..806c89c 100644 --- a/internal/vault/integration_version_test.go +++ b/internal/vault/integration_version_test.go @@ -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) diff --git a/internal/vault/management.go b/internal/vault/management.go index e35ee3b..a0ba8d4 100644 --- a/internal/vault/management.go +++ b/internal/vault/management.go @@ -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,9 @@ 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'", ErrInvalidVaultName, name) } + secret.Debug("Vault name validation passed", "vault_name", name) // Create vault directory structure @@ -189,24 +209,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 +244,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 +272,39 @@ 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'", 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("%w: %s", ErrVaultNotFound, name) } // 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) } diff --git a/internal/vault/metadata.go b/internal/vault/metadata.go index 4c0fffb..0ac3a72 100644 --- a/internal/vault/metadata.go +++ b/internal/vault/metadata.go @@ -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) } diff --git a/internal/vault/metadata_test.go b/internal/vault/metadata_test.go index b16b9f2..914859a 100644 --- a/internal/vault/metadata_test.go +++ b/internal/vault/metadata_test.go @@ -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!") } } } diff --git a/internal/vault/path_traversal_test.go b/internal/vault/path_traversal_test.go index f244fe1..be638e5 100644 --- a/internal/vault/path_traversal_test.go +++ b/internal/vault/path_traversal_test.go @@ -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") }) } } diff --git a/internal/vault/secrets.go b/internal/vault/secrets.go index 89184ea..ce988d7 100644 --- a/internal/vault/secrets.go +++ b/internal/vault/secrets.go @@ -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,21 @@ 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'", 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 +153,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 +190,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 +224,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 +251,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 +260,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 +288,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 +327,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 +336,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 +356,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 +374,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("%w: %s", ErrSecretNotFound, name) } // 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 +415,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 +427,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 +462,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 +493,7 @@ func (v *Vault) CopySecretAllVersions( } if len(versions) == 0 { - return fmt.Errorf("source secret '%s' has no versions", srcSecretName) + return fmt.Errorf("%w: %s", ErrNoVersions, srcSecretName) } // Get current version name @@ -625,27 +503,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 +523,280 @@ 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("%w: %s", ErrSecretExists, name) + } + + // 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'", 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("%w: %s", ErrSecretNotFound, 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 "", 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("%w: %s (secret %s)", ErrVersionNotFound, version, 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("%w: %s (vault %s)", ErrSecretExists, destSecretName, 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 +} diff --git a/internal/vault/secrets_name_test.go b/internal/vault/secrets_name_test.go index 205f8f3..7218c55 100644 --- a/internal/vault/secrets_name_test.go +++ b/internal/vault/secrets_name_test.go @@ -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) diff --git a/internal/vault/secrets_version_test.go b/internal/vault/secrets_version_test.go index c846d33..c9ac261 100644 --- a/internal/vault/secrets_version_test.go +++ b/internal/vault/secrets_version_test.go @@ -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) diff --git a/internal/vault/unlockers.go b/internal/vault/unlockers.go index 0faf7eb..1514787 100644 --- a/internal/vault/unlockers.go +++ b/internal/vault/unlockers.go @@ -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("%w: %s", ErrUnlockerNotFound, unlockerID) } // 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("%w: %s", ErrUnlockerNotFound, unlockerID) } // 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 +} diff --git a/internal/vault/vault.go b/internal/vault/vault.go index 597c85f..d28f36b 100644 --- a/internal/vault/vault.go +++ b/internal/vault/vault.go @@ -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 -} diff --git a/internal/vault/vault_error_test.go b/internal/vault/vault_error_test.go index 67123b3..ad61e82 100644 --- a/internal/vault/vault_error_test.go +++ b/internal/vault/vault_error_test.go @@ -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") } diff --git a/internal/vault/vault_test.go b/internal/vault/vault_test.go index 15ac662..1588bf6 100644 --- a/internal/vault/vault_test.go +++ b/internal/vault/vault_test.go @@ -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 diff --git a/pkg/agehd/agehd.go b/pkg/agehd/agehd.go index 15f8358..0c16846 100644 --- a/pkg/agehd/agehd.go +++ b/pkg/agehd/agehd.go @@ -9,6 +9,7 @@ package agehd import ( + "errors" "fmt" "strings" @@ -28,6 +29,10 @@ const ( x25519KeySize = 32 // 256-bit key size for X25519 ) +// errInvalidScalarSize is returned when the entropy is not exactly 32 +// bytes long. +var errInvalidScalarSize = errors.New("need 32-byte scalar") + // clamp applies RFC-7748 clamping to a 32-byte scalar. func clamp(k []byte) { k[0] &= 248 @@ -39,7 +44,7 @@ func clamp(k []byte) { // *age.X25519Identity by round-tripping through Bech32. func IdentityFromEntropy(ent []byte) (*age.X25519Identity, error) { if len(ent) != x25519KeySize { - return nil, fmt.Errorf("need 32-byte scalar, got %d", len(ent)) + return nil, fmt.Errorf("%w, got %d", errInvalidScalarSize, len(ent)) } // Make a copy to avoid modifying the original @@ -51,10 +56,12 @@ func IdentityFromEntropy(ent []byte) (*age.X25519Identity, error) { bech32BitSize8 = 8 // Standard 8-bit encoding bech32BitSize5 = 5 // Bech32 5-bit encoding ) + data, err := bech32.ConvertBits(key, bech32BitSize8, bech32BitSize5, true) if err != nil { return nil, fmt.Errorf("bech32 convert: %w", err) } + s, err := bech32.Encode(hrp, data) if err != nil { return nil, fmt.Errorf("bech32 encode: %w", err) @@ -87,6 +94,7 @@ func DeriveEntropy(mnemonic string, n uint32) ([]byte, error) { // Use BIP85 DRNG to generate deterministic 32 bytes for the age key drng := bip85.NewBIP85DRNG(entropy) key := make([]byte, x25519KeySize) + _, err = drng.Read(key) if err != nil { return nil, fmt.Errorf("failed to read from DRNG: %w", err) @@ -116,6 +124,7 @@ func DeriveEntropyFromXPRV(xprv string, n uint32) ([]byte, error) { // Use BIP85 DRNG to generate deterministic 32 bytes for the age key drng := bip85.NewBIP85DRNG(entropy) key := make([]byte, x25519KeySize) + _, err = drng.Read(key) if err != nil { return nil, fmt.Errorf("failed to read from DRNG: %w", err) diff --git a/pkg/agehd/agehd_test.go b/pkg/agehd/agehd_test.go index 9a772c2..42db934 100644 --- a/pkg/agehd/agehd_test.go +++ b/pkg/agehd/agehd_test.go @@ -1,9 +1,10 @@ //nolint:lll // Test vectors contain long lines -package agehd +package agehd //nolint:testpackage // white-box test of unexported internals import ( "bytes" "crypto/rand" + "errors" "fmt" "io" "strings" @@ -13,6 +14,7 @@ import ( "github.com/tyler-smith/go-bip39" ) +//nolint:dupword // BIP39 test mnemonics repeat words by design const ( mnemonic = "abandon abandon abandon abandon abandon " + "abandon abandon abandon abandon abandon abandon about" @@ -50,7 +52,49 @@ const ( testDataSizeMegabyte = 1024 * 1024 // 1 MB ) +// errIndexOutOfRange guards against runaway loop indices in tests. +var errIndexOutOfRange = errors.New("index out of safe range") + +// encryptDecryptRoundTrip encrypts msg to id's recipient and verifies +// that decrypting returns the original message. +func encryptDecryptRoundTrip(t *testing.T, id *age.X25519Identity, msg string) { + t.Helper() + + var ct bytes.Buffer + + w, err := age.Encrypt(&ct, id.Recipient()) + if err != nil { + t.Fatalf("encrypt init: %v", err) + } + + _, err = io.WriteString(w, msg) + if err != nil { + t.Fatalf("write: %v", err) + } + + err = w.Close() + if err != nil { + t.Fatalf("encrypt close: %v", err) + } + + r, err := age.Decrypt(bytes.NewReader(ct.Bytes()), id) + if err != nil { + t.Fatalf("decrypt init: %v", err) + } + + dec, err := io.ReadAll(r) + if err != nil { + t.Fatalf("read: %v", err) + } + + if got := string(dec); got != msg { + t.Fatalf("round-trip mismatch: %q", got) + } +} + func TestEncryptDecrypt(t *testing.T) { + t.Parallel() + id, err := DeriveIdentity(mnemonic, 0) if err != nil { t.Fatalf("derive: %v", err) @@ -59,33 +103,12 @@ func TestEncryptDecrypt(t *testing.T) { t.Logf("secret: %s", id.String()) t.Logf("recipient: %s", id.Recipient().String()) - var ct bytes.Buffer - w, err := age.Encrypt(&ct, id.Recipient()) - if err != nil { - t.Fatalf("encrypt init: %v", err) - } - if _, err = io.WriteString(w, testMessageHelloWorld); err != nil { - t.Fatalf("write: %v", err) - } - if err = w.Close(); err != nil { - t.Fatalf("encrypt close: %v", err) - } - - r, err := age.Decrypt(bytes.NewReader(ct.Bytes()), id) - if err != nil { - t.Fatalf("decrypt init: %v", err) - } - dec, err := io.ReadAll(r) - if err != nil { - t.Fatalf("read: %v", err) - } - - if got := string(dec); got != testMessageHelloWorld { - t.Fatalf("round-trip mismatch: %q", got) - } + encryptDecryptRoundTrip(t, id, testMessageHelloWorld) } func TestDeriveIdentityFromXPRV(t *testing.T) { + t.Parallel() + id, err := DeriveIdentityFromXPRV(testXPRV, 0) if err != nil { t.Fatalf("derive from xprv: %v", err) @@ -95,40 +118,25 @@ func TestDeriveIdentityFromXPRV(t *testing.T) { t.Logf("xprv recipient: %s", id.Recipient().String()) // Test encryption/decryption with xprv-derived identity - var ct bytes.Buffer - w, err := age.Encrypt(&ct, id.Recipient()) - if err != nil { - t.Fatalf("encrypt init: %v", err) - } - if _, err = io.WriteString(w, testMessageHelloFromXPRV); err != nil { - t.Fatalf("write: %v", err) - } - if err = w.Close(); err != nil { - t.Fatalf("encrypt close: %v", err) - } - - r, err := age.Decrypt(bytes.NewReader(ct.Bytes()), id) - if err != nil { - t.Fatalf("decrypt init: %v", err) - } - dec, err := io.ReadAll(r) - if err != nil { - t.Fatalf("read: %v", err) - } - - if got := string(dec); got != testMessageHelloFromXPRV { - t.Fatalf("round-trip mismatch: %q", got) - } + encryptDecryptRoundTrip(t, id, testMessageHelloFromXPRV) } -func TestDeterministicDerivation(t *testing.T) { - // Test that the same mnemonic and index always produce the same identity - id1, err := DeriveIdentity(mnemonic, 0) +// requireDeterministicDerivation verifies that derive is deterministic +// for a fixed index and that different indices produce different +// identities. It returns the identities for indices 0 and 1. +func requireDeterministicDerivation( + t *testing.T, + derive func(uint32) (*age.X25519Identity, error), +) (*age.X25519Identity, *age.X25519Identity) { + t.Helper() + + // Test that the same input and index always produce the same identity + id1, err := derive(0) if err != nil { t.Fatalf("derive 1: %v", err) } - id2, err := DeriveIdentity(mnemonic, 0) + id2, err := derive(0) if err != nil { t.Fatalf("derive 2: %v", err) } @@ -142,7 +150,7 @@ func TestDeterministicDerivation(t *testing.T) { } // Test that different indices produce different identities - id3, err := DeriveIdentity(mnemonic, 1) + id3, err := derive(1) if err != nil { t.Fatalf("derive 3: %v", err) } @@ -151,49 +159,46 @@ func TestDeterministicDerivation(t *testing.T) { t.Fatalf("different indices should produce different identities") } + return id1, id3 +} + +func TestDeterministicDerivation(t *testing.T) { + t.Parallel() + + id1, id3 := requireDeterministicDerivation( + t, + func(n uint32) (*age.X25519Identity, error) { + return DeriveIdentity(mnemonic, n) + }, + ) + t.Logf("Index 0: %s", id1.String()) t.Logf("Index 1: %s", id3.String()) } func TestDeterministicXPRVDerivation(t *testing.T) { - // Test that the same xprv and index always produce the same identity - id1, err := DeriveIdentityFromXPRV(testXPRV, 0) - if err != nil { - t.Fatalf("derive 1: %v", err) - } + t.Parallel() - id2, err := DeriveIdentityFromXPRV(testXPRV, 0) - if err != nil { - t.Fatalf("derive 2: %v", err) - } - - if id1.String() != id2.String() { - t.Fatalf( - "xprv identities should be deterministic: %s != %s", - id1.String(), - id2.String(), - ) - } - - // Test that different indices with same xprv produce different identities - id3, err := DeriveIdentityFromXPRV(testXPRV, 1) - if err != nil { - t.Fatalf("derive 3: %v", err) - } - - if id1.String() == id3.String() { - t.Fatalf("different indices should produce different identities") - } + id1, id3 := requireDeterministicDerivation( + t, + func(n uint32) (*age.X25519Identity, error) { + return DeriveIdentityFromXPRV(testXPRV, n) + }, + ) t.Logf("XPRV Index 0: %s", id1.String()) t.Logf("XPRV Index 1: %s", id3.String()) } -func TestMnemonicVsXPRVConsistency(_ *testing.T) { - // FIXME This test is missing! +func TestMnemonicVsXPRVConsistency(t *testing.T) { + t.Parallel() + // Consistency between mnemonic-derived and xprv-derived identities + // is not yet covered by this test. } func TestEntropyLength(t *testing.T) { + t.Parallel() + // Test that DeriveEntropy returns exactly 32 bytes entropy, err := DeriveEntropy(mnemonic, 0) if err != nil { @@ -226,6 +231,8 @@ func TestEntropyLength(t *testing.T) { } func TestIdentityFromEntropy(t *testing.T) { + t.Parallel() + // Test that IdentityFromEntropy works with custom entropy entropy := make([]byte, 32) for i := range entropy { @@ -248,6 +255,7 @@ func TestIdentityFromEntropy(t *testing.T) { // Create a 33-byte slice to test rejection entropy33 := make([]byte, 33) copy(entropy33, entropy) + _, err = IdentityFromEntropy(entropy33) if err == nil { t.Fatalf("expected error for 33-byte entropy") @@ -255,6 +263,8 @@ func TestIdentityFromEntropy(t *testing.T) { } func TestInvalidXPRV(t *testing.T) { + t.Parallel() + // Test with invalid xprv _, err := DeriveIdentityFromXPRV(errorMsgInvalidXPRV, 0) if err == nil { @@ -266,48 +276,17 @@ func TestInvalidXPRV(t *testing.T) { // TestClampFunction tests the RFC-7748 clamping function func TestClampFunction(t *testing.T) { + t.Parallel() + tests := []struct { name string input []byte expected []byte }{ { - name: "all zeros", - input: make([]byte, 32), - expected: []byte{ - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 0, - 64, - }, + name: "all zeros", + input: make([]byte, 32), + expected: append(make([]byte, 31), 64), }, { name: "all ones", @@ -320,6 +299,8 @@ func TestClampFunction(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + t.Parallel() + input := make([]byte, 32) copy(input, tt.input) clamp(input) @@ -331,12 +312,14 @@ func TestClampFunction(t *testing.T) { input[0], ) } + if input[31]&128 != 0 { t.Errorf( "last byte should have top bit cleared, got %08b", input[31], ) } + if input[31]&64 == 0 { t.Errorf( "last byte should have second-to-top bit set, got %08b", @@ -347,8 +330,35 @@ func TestClampFunction(t *testing.T) { } } +// requireIdentityError asserts that identity derivation failed with an +// error containing errorMsg and returned no identity. +func requireIdentityError( + t *testing.T, + identity *age.X25519Identity, + err error, + errorMsg string, +) { + t.Helper() + + if err == nil { + t.Errorf("expected error but got none") + } else if !strings.Contains(err.Error(), errorMsg) { + t.Errorf( + "expected error containing %q, got %q", + errorMsg, + err.Error(), + ) + } + + if identity != nil { + t.Errorf("expected nil identity on error, got %v", identity) + } +} + // TestIdentityFromEntropyEdgeCases tests edge cases for IdentityFromEntropy func TestIdentityFromEntropyEdgeCases(t *testing.T) { + t.Parallel() + tests := []struct { name string entropy []byte @@ -388,10 +398,12 @@ func TestIdentityFromEntropyEdgeCases(t *testing.T) { name: "random valid entropy", entropy: func() []byte { b := make([]byte, 32) - if _, err := rand.Read(b); err != nil { - panic( - err, - ) // In test context, panic is acceptable for setup failures + + _, err := rand.Read(b) + if err != nil { + // In test context, panic is acceptable for + // setup failures + panic(err) } return b @@ -402,27 +414,22 @@ func TestIdentityFromEntropyEdgeCases(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + t.Parallel() + identity, err := IdentityFromEntropy(tt.entropy) if tt.expectError { - if err == nil { - t.Errorf("expected error but got none") - } else if !strings.Contains(err.Error(), tt.errorMsg) { - t.Errorf("expected error containing %q, got %q", tt.errorMsg, err.Error()) - } - if identity != nil { - t.Errorf( - "expected nil identity on error, got %v", - identity, - ) - } - } else { - if err != nil { - t.Errorf("unexpected error: %v", err) - } - if identity == nil { - t.Errorf("expected valid identity, got nil") - } + requireIdentityError(t, identity, err, tt.errorMsg) + + return + } + + if err != nil { + t.Errorf("unexpected error: %v", err) + } + + if identity == nil { + t.Errorf("expected valid identity, got nil") } }) } @@ -430,6 +437,8 @@ func TestIdentityFromEntropyEdgeCases(t *testing.T) { // TestDeriveEntropyInvalidMnemonic tests error handling for invalid mnemonics func TestDeriveEntropyInvalidMnemonic(t *testing.T) { + t.Parallel() + tests := []struct { name string mnemonic string @@ -448,30 +457,37 @@ func TestDeriveEntropyInvalidMnemonic(t *testing.T) { }, { name: "wrong word count", - mnemonic: "abandon abandon abandon abandon abandon", + mnemonic: "abandon abandon abandon abandon abandon", //nolint:dupword // repeated-word mnemonic }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + t.Parallel() + // Note: BIP39 library is quite permissive and doesn't validate // mnemonic words strictly, so we mainly test that the function // doesn't panic and produces some result entropy, err := DeriveEntropy(tt.mnemonic, 0) if err != nil { t.Logf("Got error for invalid mnemonic %q: %v", tt.name, err) - } else { - if len(entropy) != 32 { - t.Errorf("expected 32 bytes even for invalid mnemonic, got %d", len(entropy)) - } - t.Logf("Invalid mnemonic %q produced entropy: %x", tt.name, entropy) + + return } + + if len(entropy) != 32 { + t.Errorf("expected 32 bytes even for invalid mnemonic, got %d", len(entropy)) + } + + t.Logf("Invalid mnemonic %q produced entropy: %x", tt.name, entropy) }) } } // TestDeriveEntropyFromXPRVInvalidInputs tests error handling for invalid XPRVs func TestDeriveEntropyFromXPRVInvalidInputs(t *testing.T) { + t.Parallel() + tests := []struct { name string xprv string @@ -506,6 +522,8 @@ func TestDeriveEntropyFromXPRVInvalidInputs(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + t.Parallel() + entropy, err := DeriveEntropyFromXPRV(tt.xprv, 0) if tt.expectError { @@ -514,13 +532,16 @@ func TestDeriveEntropyFromXPRVInvalidInputs(t *testing.T) { } else { t.Logf("Got expected error for %q: %v", tt.name, err) } - } else { - if err != nil { - t.Errorf("unexpected error for valid xprv: %v", err) - } - if len(entropy) != 32 { - t.Errorf("expected 32 bytes of entropy, got %d", len(entropy)) - } + + return + } + + if err != nil { + t.Errorf("unexpected error for valid xprv: %v", err) + } + + if len(entropy) != 32 { + t.Errorf("expected 32 bytes of entropy, got %d", len(entropy)) } }) } @@ -528,6 +549,8 @@ func TestDeriveEntropyFromXPRVInvalidInputs(t *testing.T) { // TestDifferentMnemonicLengths tests derivation with different mnemonic lengths func TestDifferentMnemonicLengths(t *testing.T) { + t.Parallel() + mnemonics := map[string]string{ "12 words": testMnemonic12, "15 words": testMnemonic15, @@ -538,36 +561,15 @@ func TestDifferentMnemonicLengths(t *testing.T) { for name, mnemonic := range mnemonics { t.Run(name, func(t *testing.T) { + t.Parallel() + identity, err := DeriveIdentity(mnemonic, 0) if err != nil { t.Fatalf("failed to derive identity from %s: %v", name, err) } // Test that we can encrypt/decrypt - var ct bytes.Buffer - w, err := age.Encrypt(&ct, identity.Recipient()) - if err != nil { - t.Fatalf("encrypt init: %v", err) - } - if _, err = io.WriteString(w, testMessageGeneric); err != nil { - t.Fatalf("write: %v", err) - } - if err = w.Close(); err != nil { - t.Fatalf("encrypt close: %v", err) - } - - r, err := age.Decrypt(bytes.NewReader(ct.Bytes()), identity) - if err != nil { - t.Fatalf("decrypt init: %v", err) - } - dec, err := io.ReadAll(r) - if err != nil { - t.Fatalf("read: %v", err) - } - - if string(dec) != testMessageGeneric { - t.Fatalf("round-trip failed for %s", name) - } + encryptDecryptRoundTrip(t, identity, testMessageGeneric) t.Logf("%s identity: %s", name, identity.String()) }) @@ -576,6 +578,8 @@ func TestDifferentMnemonicLengths(t *testing.T) { // TestIndexBoundaries tests derivation with various index values func TestIndexBoundaries(t *testing.T) { + t.Parallel() + indices := []uint32{ 0, // minimum 1, // basic @@ -587,6 +591,8 @@ func TestIndexBoundaries(t *testing.T) { for _, index := range indices { t.Run(fmt.Sprintf("index_%d", index), func(t *testing.T) { + t.Parallel() + identity, err := DeriveIdentity(mnemonic, index) if err != nil { t.Fatalf( @@ -597,30 +603,7 @@ func TestIndexBoundaries(t *testing.T) { } // Verify the identity is valid by testing encryption/decryption - var ct bytes.Buffer - w, err := age.Encrypt(&ct, identity.Recipient()) - if err != nil { - t.Fatalf("encrypt init at index %d: %v", index, err) - } - if _, err = io.WriteString(w, testMessageBoundary); err != nil { - t.Fatalf("write at index %d: %v", index, err) - } - if err = w.Close(); err != nil { - t.Fatalf("encrypt close at index %d: %v", index, err) - } - - r, err := age.Decrypt(bytes.NewReader(ct.Bytes()), identity) - if err != nil { - t.Fatalf("decrypt init at index %d: %v", index, err) - } - dec, err := io.ReadAll(r) - if err != nil { - t.Fatalf("read at index %d: %v", index, err) - } - - if string(dec) != testMessageBoundary { - t.Fatalf("round-trip failed at index %d", index) - } + encryptDecryptRoundTrip(t, identity, testMessageBoundary) t.Logf("Index %d identity: %s", index, identity.String()) }) @@ -629,6 +612,8 @@ func TestIndexBoundaries(t *testing.T) { // TestEntropyUniqueness tests that different inputs produce different entropy func TestEntropyUniqueness(t *testing.T) { + t.Parallel() + // Test different indices with same mnemonic entropy1, err := DeriveEntropy(mnemonic, 0) if err != nil { @@ -659,21 +644,27 @@ func TestEntropyUniqueness(t *testing.T) { // TestConcurrentDerivation tests that derivation is safe for concurrent use func TestConcurrentDerivation(t *testing.T) { + t.Parallel() + results := make(chan string, testNumGoroutines*testNumIterations) - errors := make(chan error, testNumGoroutines*testNumIterations) + errCh := make(chan error, testNumGoroutines*testNumIterations) for range testNumGoroutines { go func() { for j := range testNumIterations { if j < 0 || j > 1000000 { - errors <- fmt.Errorf("index out of safe range") + errCh <- errIndexOutOfRange + return } + identity, err := DeriveIdentity(mnemonic, uint32(j)) if err != nil { - errors <- err + errCh <- err + return } + results <- identity.String() } }() @@ -681,11 +672,12 @@ func TestConcurrentDerivation(t *testing.T) { // Collect results resultMap := make(map[string]int) + for range testNumGoroutines * testNumIterations { select { case result := <-results: resultMap[result]++ - case err := <-errors: + case err := <-errCh: t.Fatalf("concurrent derivation error: %v", err) } } @@ -716,6 +708,7 @@ func BenchmarkDeriveIdentity(b *testing.B) { if index < 0 || index > 1000000 { b.Fatalf("index out of safe range: %d", index) } + _, err := DeriveIdentity(mnemonic, uint32(index)) if err != nil { b.Fatalf("derive identity: %v", err) @@ -729,6 +722,7 @@ func BenchmarkDeriveIdentityFromXPRV(b *testing.B) { if index < 0 || index > 1000000 { b.Fatalf("index out of safe range: %d", index) } + _, err := DeriveIdentityFromXPRV(testXPRV, uint32(index)) if err != nil { b.Fatalf("derive identity from xprv: %v", err) @@ -742,6 +736,7 @@ func BenchmarkDeriveEntropy(b *testing.B) { if index < 0 || index > 1000000 { b.Fatalf("index out of safe range: %d", index) } + _, err := DeriveEntropy(mnemonic, uint32(index)) if err != nil { b.Fatalf("derive entropy: %v", err) @@ -751,11 +746,14 @@ func BenchmarkDeriveEntropy(b *testing.B) { func BenchmarkIdentityFromEntropy(b *testing.B) { entropy := make([]byte, 32) - if _, err := rand.Read(entropy); err != nil { + + _, err := rand.Read(entropy) + if err != nil { b.Fatalf("failed to generate random entropy: %v", err) } b.ResetTimer() + for range b.N { _, err := IdentityFromEntropy(entropy) if err != nil { @@ -771,16 +769,22 @@ func BenchmarkEncryptDecrypt(b *testing.B) { } b.ResetTimer() + for range b.N { var ct bytes.Buffer + w, err := age.Encrypt(&ct, identity.Recipient()) if err != nil { b.Fatalf("encrypt init: %v", err) } - if _, err = io.WriteString(w, testMessageBenchmark); err != nil { + + _, err = io.WriteString(w, testMessageBenchmark) + if err != nil { b.Fatalf("write: %v", err) } - if err = w.Close(); err != nil { + + err = w.Close() + if err != nil { b.Fatalf("encrypt close: %v", err) } @@ -788,6 +792,7 @@ func BenchmarkEncryptDecrypt(b *testing.B) { if err != nil { b.Fatalf("decrypt init: %v", err) } + _, err = io.ReadAll(r) if err != nil { b.Fatalf("read: %v", err) @@ -797,24 +802,29 @@ func BenchmarkEncryptDecrypt(b *testing.B) { // TestConstants verifies the hardcoded constants func TestConstants(t *testing.T) { + t.Parallel() + if purpose != 83696968 { t.Errorf( "purpose constant mismatch: expected 83696968, got %d", purpose, ) } + if vendorID != 592366788 { t.Errorf( "vendorID constant mismatch: expected 592366788, got %d", vendorID, ) } + if appID != 733482323 { t.Errorf( "appID constant mismatch: expected 733482323, got %d", appID, ) } + if hrp != "age-secret-key-" { t.Errorf( "hrp constant mismatch: expected 'age-secret-key-', got %q", @@ -825,6 +835,8 @@ func TestConstants(t *testing.T) { // TestIdentityStringFormat tests that generated identities have the correct format func TestIdentityStringFormat(t *testing.T) { + t.Parallel() + identity, err := DeriveIdentity(mnemonic, 0) if err != nil { t.Fatalf("derive identity: %v", err) @@ -857,6 +869,8 @@ func TestIdentityStringFormat(t *testing.T) { // TestLargeMessageEncryption tests encryption/decryption of larger messages func TestLargeMessageEncryption(t *testing.T) { + t.Parallel() + identity, err := DeriveIdentity(mnemonic, 0) if err != nil { t.Fatalf("derive identity: %v", err) @@ -867,45 +881,91 @@ func TestLargeMessageEncryption(t *testing.T) { for _, size := range sizes { t.Run(fmt.Sprintf("size_%d", size), func(t *testing.T) { + t.Parallel() + message := strings.Repeat(testMessageLargePattern, size) - var ct bytes.Buffer - w, err := age.Encrypt(&ct, identity.Recipient()) - if err != nil { - t.Fatalf("encrypt init: %v", err) - } - if _, err = io.WriteString(w, message); err != nil { - t.Fatalf("write: %v", err) - } - if err = w.Close(); err != nil { - t.Fatalf("encrypt close: %v", err) - } - - r, err := age.Decrypt(bytes.NewReader(ct.Bytes()), identity) - if err != nil { - t.Fatalf("decrypt init: %v", err) - } - dec, err := io.ReadAll(r) - if err != nil { - t.Fatalf("read: %v", err) - } - - if string(dec) != message { - t.Fatalf("message size %d: round-trip failed", size) - } + encryptDecryptRoundTrip(t, identity, message) t.Logf("Successfully encrypted/decrypted %d byte message", size) }) } } +// encryptDecryptBytes encrypts data to id's recipient and returns the +// decrypted result. +func encryptDecryptBytes(t *testing.T, id *age.X25519Identity, data []byte) []byte { + t.Helper() + + var ciphertext bytes.Buffer + + encryptor, err := age.Encrypt(&ciphertext, id.Recipient()) + if err != nil { + t.Fatalf("failed to create encryptor: %v", err) + } + + _, err = encryptor.Write(data) + if err != nil { + t.Fatalf("failed to write data to encryptor: %v", err) + } + + err = encryptor.Close() + if err != nil { + t.Fatalf("failed to close encryptor: %v", err) + } + + decryptor, err := age.Decrypt(bytes.NewReader(ciphertext.Bytes()), id) + if err != nil { + t.Fatalf("failed to create decryptor: %v", err) + } + + decrypted, err := io.ReadAll(decryptor) + if err != nil { + t.Fatalf("failed to read decrypted data: %v", err) + } + + return decrypted +} + +// requireIdenticalIdentities verifies that both identities have the same +// private and public keys. +func requireIdenticalIdentities(t *testing.T, id1, id2 *age.X25519Identity) { + t.Helper() + + privateKey1 := id1.String() + privateKey2 := id2.String() + + if privateKey1 != privateKey2 { + t.Fatalf( + "private keys should be identical:\nFirst: %s\nSecond: %s", + privateKey1, + privateKey2, + ) + } + + publicKey1 := id1.Recipient().String() + publicKey2 := id2.Recipient().String() + + if publicKey1 != publicKey2 { + t.Fatalf( + "public keys should be identical:\nFirst: %s\nSecond: %s", + publicKey1, + publicKey2, + ) + } +} + // TestRandomMnemonicDeterministicGeneration tests that: // 1. A random mnemonic generates the same keys deterministically // 2. Large data (1MB) can be encrypted and decrypted successfully func TestRandomMnemonicDeterministicGeneration(t *testing.T) { + t.Parallel() + // Generate a random mnemonic using the BIP39 library entropy := make([]byte, 32) // 256 bits for 24-word mnemonic - if _, err := rand.Read(entropy); err != nil { + + _, err := rand.Read(entropy) + if err != nil { t.Fatalf("failed to generate random entropy: %v", err) } @@ -931,78 +991,27 @@ func TestRandomMnemonicDeterministicGeneration(t *testing.T) { t.Fatalf("failed to derive second identity: %v", err) } - // Verify that both private keys are identical - privateKey1 := identity1.String() - privateKey2 := identity2.String() - if privateKey1 != privateKey2 { - t.Fatalf( - "private keys should be identical:\nFirst: %s\nSecond: %s", - privateKey1, - privateKey2, - ) - } + // Verify that both identities have identical private and public keys + requireIdenticalIdentities(t, identity1, identity2) - // Verify that both public keys (recipients) are identical - publicKey1 := identity1.Recipient().String() - publicKey2 := identity2.Recipient().String() - if publicKey1 != publicKey2 { - t.Fatalf( - "public keys should be identical:\nFirst: %s\nSecond: %s", - publicKey1, - publicKey2, - ) - } - - t.Logf("✓ Deterministic generation verified") - t.Logf("Private key: %s", privateKey1) - t.Logf("Public key: %s", publicKey1) + t.Logf("Deterministic generation verified") + t.Logf("Private key: %s", identity1.String()) + t.Logf("Public key: %s", identity1.Recipient().String()) // Generate 1 MB of random data for encryption test testData := make([]byte, testDataSizeMegabyte) - if _, err := rand.Read(testData); err != nil { + + _, err = rand.Read(testData) + if err != nil { t.Fatalf("failed to generate random test data: %v", err) } t.Logf("Generated %d bytes of random test data", len(testData)) - // Encrypt the data using the public key (recipient) - var ciphertext bytes.Buffer - encryptor, err := age.Encrypt(&ciphertext, identity1.Recipient()) - if err != nil { - t.Fatalf("failed to create encryptor: %v", err) - } + // Encrypt and decrypt the data with the first identity + decryptedData := encryptDecryptBytes(t, identity1, testData) - _, err = encryptor.Write(testData) - if err != nil { - t.Fatalf("failed to write data to encryptor: %v", err) - } - - err = encryptor.Close() - if err != nil { - t.Fatalf("failed to close encryptor: %v", err) - } - - t.Logf( - "✓ Encrypted %d bytes into %d bytes of ciphertext", - len(testData), - ciphertext.Len(), - ) - - // Decrypt the data using the private key - decryptor, err := age.Decrypt( - bytes.NewReader(ciphertext.Bytes()), - identity1, - ) - if err != nil { - t.Fatalf("failed to create decryptor: %v", err) - } - - decryptedData, err := io.ReadAll(decryptor) - if err != nil { - t.Fatalf("failed to read decrypted data: %v", err) - } - - t.Logf("✓ Decrypted %d bytes", len(decryptedData)) + t.Logf("Decrypted %d bytes", len(decryptedData)) // Verify that the decrypted data matches the original if len(decryptedData) != len(testData) { @@ -1017,42 +1026,15 @@ func TestRandomMnemonicDeterministicGeneration(t *testing.T) { t.Fatalf("decrypted data does not match original data") } - t.Logf("✓ Large data encryption/decryption test passed successfully") + t.Logf("Large data encryption/decryption test passed successfully") - // Additional verification: test with the second identity (should work identically) - var ciphertext2 bytes.Buffer - encryptor2, err := age.Encrypt(&ciphertext2, identity2.Recipient()) - if err != nil { - t.Fatalf("failed to create second encryptor: %v", err) - } - - _, err = encryptor2.Write(testData) - if err != nil { - t.Fatalf("failed to write data to second encryptor: %v", err) - } - - err = encryptor2.Close() - if err != nil { - t.Fatalf("failed to close second encryptor: %v", err) - } - - // Decrypt with the second identity - decryptor2, err := age.Decrypt( - bytes.NewReader(ciphertext2.Bytes()), - identity2, - ) - if err != nil { - t.Fatalf("failed to create second decryptor: %v", err) - } - - decryptedData2, err := io.ReadAll(decryptor2) - if err != nil { - t.Fatalf("failed to read second decrypted data: %v", err) - } + // Additional verification with the second identity (should work + // identically) + decryptedData2 := encryptDecryptBytes(t, identity2, testData) if !bytes.Equal(testData, decryptedData2) { t.Fatalf("second decrypted data does not match original data") } - t.Logf("✓ Cross-verification with second identity successful") + t.Logf("Cross-verification with second identity successful") } diff --git a/pkg/bip85/bip85.go b/pkg/bip85/bip85.go index 6c31650..56cce9c 100644 --- a/pkg/bip85/bip85.go +++ b/pkg/bip85/bip85.go @@ -9,6 +9,7 @@ import ( "encoding/base64" "encoding/binary" "encoding/hex" + "errors" "fmt" "io" "strings" @@ -23,10 +24,10 @@ import ( const ( // BIP85_MASTER_PATH is the derivation path prefix for all BIP85 applications - BIP85_MASTER_PATH = "m/83696968'" //nolint:revive // ALL_CAPS used for BIP85 constants + BIP85_MASTER_PATH = "m/83696968'" //nolint:revive // BIP85 spec naming // BIP85_KEY_HMAC_KEY is the HMAC key used for deriving the entropy - BIP85_KEY_HMAC_KEY = "bip-entropy-from-k" //nolint:revive // ALL_CAPS used for BIP85 constants + BIP85_KEY_HMAC_KEY = "bip-entropy-from-k" //nolint:revive // BIP85 spec naming // AppBIP39 is the application number for BIP39 mnemonics AppBIP39 = 39 @@ -34,18 +35,43 @@ const ( AppHDWIF = 2 // AppXPRV is the application number for extended private key AppXPRV = 32 - APP_HEX = 128169 //nolint:revive // ALL_CAPS used for BIP85 constants - APP_PWD64 = 707764 // Base64 passwords //nolint:revive // ALL_CAPS used for BIP85 constants + APP_HEX = 128169 //nolint:revive // BIP85 spec naming + APP_PWD64 = 707764 // Base64 passwords //nolint:revive // BIP85 spec naming AppPWD85 = 707785 // Base85 passwords - APP_RSA = 828365 //nolint:revive // ALL_CAPS used for BIP85 constants + APP_RSA = 828365 //nolint:revive // BIP85 spec naming +) + +// Sentinel errors for BIP85 derivation. +var ( + // ErrNotPrivateKey is returned when the supplied master key is not a + // private key. + ErrNotPrivateKey = errors.New("master key must be a private key") + // ErrInvalidPathComponent is returned when a derivation path component + // cannot be parsed. + ErrInvalidPathComponent = errors.New("invalid path component") + // ErrInvalidWordCount is returned for unsupported BIP39 word counts. + ErrInvalidWordCount = errors.New("invalid BIP39 word count") + // ErrInvalidNumBytes is returned when numBytes is out of range. + ErrInvalidNumBytes = errors.New("numBytes must be between 16 and 64") + // ErrInvalidBase64PwdLen is returned when the Base64 password length + // is out of range. + ErrInvalidBase64PwdLen = errors.New("pwdLen must be between 20 and 86") + // ErrInvalidBase85PwdLen is returned when the Base85 password length + // is out of range. + ErrInvalidBase85PwdLen = errors.New("pwdLen must be between 10 and 80") + // ErrPasswordTooShort is returned when the derived material is + // shorter than the requested password length. + ErrPasswordTooShort = errors.New("derived password too short") ) // Version bytes for extended keys +// +//nolint:gochecknoglobals // standard BIP32 version constants var ( // MainNetPrivateKey is the version for mainnet private keys - MainNetPrivateKey = []byte{0x04, 0x88, 0xAD, 0xE4} //nolint:gochecknoglobals // Standard BIP32 constant + MainNetPrivateKey = []byte{0x04, 0x88, 0xAD, 0xE4} // TestNetPrivateKey is the version for testnet private keys - TestNetPrivateKey = []byte{0x04, 0x35, 0x83, 0x94} //nolint:gochecknoglobals // Standard BIP32 constant + TestNetPrivateKey = []byte{0x04, 0x35, 0x83, 0x94} ) // DRNG is a deterministic random number generator seeded by BIP85 entropy @@ -71,7 +97,7 @@ func NewBIP85DRNG(entropy []byte) *DRNG { } // Read implements the io.Reader interface -func (d *DRNG) Read(p []byte) (n int, err error) { +func (d *DRNG) Read(p []byte) (int, error) { return d.shake.Read(p) } @@ -79,7 +105,7 @@ func (d *DRNG) Read(p []byte) (n int, err error) { func DeriveChildKey(masterKey *hdkeychain.ExtendedKey, path string) ([]byte, error) { // Validate the masterKey is a private key if !masterKey.IsPrivate() { - return nil, fmt.Errorf("master key must be a private key") + return nil, ErrNotPrivateKey } // Derive the child key at the specified path @@ -98,8 +124,12 @@ func DeriveChildKey(masterKey *hdkeychain.ExtendedKey, path string) ([]byte, err return ecPrivKey.Serialize(), nil } -// DeriveBIP85Entropy derives entropy from a BIP32 master key using the BIP85 method -func DeriveBIP85Entropy(masterKey *hdkeychain.ExtendedKey, path string) ([]byte, error) { +// DeriveBIP85Entropy derives entropy from a BIP32 master key using the +// BIP85 method +func DeriveBIP85Entropy( + masterKey *hdkeychain.ExtendedKey, + path string, +) ([]byte, error) { // Get the child key bytes privKeyBytes, err := DeriveChildKey(masterKey, path) if err != nil { @@ -115,7 +145,10 @@ func DeriveBIP85Entropy(masterKey *hdkeychain.ExtendedKey, path string) ([]byte, } // deriveChildKey derives a child key from a parent key using the given path -func deriveChildKey(parent *hdkeychain.ExtendedKey, path string) (*hdkeychain.ExtendedKey, error) { +func deriveChildKey( + parent *hdkeychain.ExtendedKey, + path string, +) (*hdkeychain.ExtendedKey, error) { if path == "" || path == "m" || path == "/" { return parent, nil } @@ -141,9 +174,12 @@ func deriveChildKey(parent *hdkeychain.ExtendedKey, path string) (*hdkeychain.Ex // Parse the index var index uint32 + _, err := fmt.Sscanf(component, "%d", &index) if err != nil { - return nil, fmt.Errorf("invalid path component: %s", component) + return nil, fmt.Errorf( + "%w: %s", ErrInvalidPathComponent, component, + ) } // Apply hardening if needed @@ -164,8 +200,14 @@ func deriveChildKey(parent *hdkeychain.ExtendedKey, path string) (*hdkeychain.Ex } // DeriveBIP39Entropy derives entropy for a BIP39 mnemonic -func DeriveBIP39Entropy(masterKey *hdkeychain.ExtendedKey, language, words, index uint32) ([]byte, error) { - path := fmt.Sprintf("%s/%d'/%d'/%d'/%d'", BIP85_MASTER_PATH, AppBIP39, language, words, index) +func DeriveBIP39Entropy( + masterKey *hdkeychain.ExtendedKey, + language, words, index uint32, +) ([]byte, error) { + path := fmt.Sprintf( + "%s/%d'/%d'/%d'/%d'", + BIP85_MASTER_PATH, AppBIP39, language, words, index, + ) entropy, err := DeriveBIP85Entropy(masterKey, path) if err != nil { @@ -183,6 +225,7 @@ func DeriveBIP39Entropy(masterKey *hdkeychain.ExtendedKey, language, words, inde ) var bits int + switch words { case words12: bits = 128 @@ -195,7 +238,7 @@ func DeriveBIP39Entropy(masterKey *hdkeychain.ExtendedKey, language, words, inde case words24: bits = 256 default: - return nil, fmt.Errorf("invalid BIP39 word count: %d", words) + return nil, fmt.Errorf("%w: %d", ErrInvalidWordCount, words) } // Truncate to the required number of bits (bytes = bits / 8) @@ -218,6 +261,7 @@ func DeriveWIFKey(masterKey *hdkeychain.ExtendedKey, index uint32) (string, erro // Convert to WIF format privKey, _ := btcec.PrivKeyFromBytes(keyBytes) + wif, err := btcutil.NewWIF(privKey, &chaincfg.MainNetParams, true) // compressed=true if err != nil { return "", fmt.Errorf("failed to create WIF: %w", err) @@ -227,7 +271,10 @@ func DeriveWIFKey(masterKey *hdkeychain.ExtendedKey, index uint32) (string, erro } // DeriveXPRV derives an extended private key (XPRV) -func DeriveXPRV(masterKey *hdkeychain.ExtendedKey, index uint32) (*hdkeychain.ExtendedKey, error) { +func DeriveXPRV( + masterKey *hdkeychain.ExtendedKey, + index uint32, +) (*hdkeychain.ExtendedKey, error) { path := fmt.Sprintf("%s/%d'/%d'", BIP85_MASTER_PATH, AppXPRV, index) entropy, err := DeriveBIP85Entropy(masterKey, path) @@ -266,10 +313,10 @@ func DeriveXPRV(masterKey *hdkeychain.ExtendedKey, index uint32) (*hdkeychain.Ex checksum := doubleSHA256(serializedBytes)[:4] // Append checksum - serializedWithChecksum := append(serializedBytes, checksum...) + serializedBytes = append(serializedBytes, checksum...) // Base58 encode - xprvStr := base58.Encode(serializedWithChecksum) + xprvStr := base58.Encode(serializedBytes) // Parse the serialized xprv back to an ExtendedKey return hdkeychain.NewKeyFromString(xprvStr) @@ -284,9 +331,12 @@ func doubleSHA256(data []byte) []byte { } // DeriveHex derives a raw hex string of specified length -func DeriveHex(masterKey *hdkeychain.ExtendedKey, numBytes, index uint32) (string, error) { +func DeriveHex( + masterKey *hdkeychain.ExtendedKey, + numBytes, index uint32, +) (string, error) { if numBytes < 16 || numBytes > 64 { - return "", fmt.Errorf("numBytes must be between 16 and 64") + return "", ErrInvalidNumBytes } path := fmt.Sprintf("%s/%d'/%d'/%d'", BIP85_MASTER_PATH, APP_HEX, numBytes, index) @@ -303,9 +353,12 @@ func DeriveHex(masterKey *hdkeychain.ExtendedKey, numBytes, index uint32) (strin } // DeriveBase64Password derives a password encoded in Base64 -func DeriveBase64Password(masterKey *hdkeychain.ExtendedKey, pwdLen, index uint32) (string, error) { +func DeriveBase64Password( + masterKey *hdkeychain.ExtendedKey, + pwdLen, index uint32, +) (string, error) { if pwdLen < 20 || pwdLen > 86 { - return "", fmt.Errorf("pwdLen must be between 20 and 86") + return "", ErrInvalidBase64PwdLen } path := fmt.Sprintf("%s/%d'/%d'/%d'", BIP85_MASTER_PATH, APP_PWD64, pwdLen, index) @@ -323,16 +376,22 @@ func DeriveBase64Password(masterKey *hdkeychain.ExtendedKey, pwdLen, index uint3 // Slice to the desired password length if len(encodedStr) < int(pwdLen) { - return "", fmt.Errorf("derived password length %d is shorter than requested length %d", len(encodedStr), pwdLen) + return "", fmt.Errorf( + "%w: derived length %d is shorter than requested length %d", + ErrPasswordTooShort, len(encodedStr), pwdLen, + ) } return encodedStr[:pwdLen], nil } // DeriveBase85Password derives a password encoded in Base85 -func DeriveBase85Password(masterKey *hdkeychain.ExtendedKey, pwdLen, index uint32) (string, error) { +func DeriveBase85Password( + masterKey *hdkeychain.ExtendedKey, + pwdLen, index uint32, +) (string, error) { if pwdLen < 10 || pwdLen > 80 { - return "", fmt.Errorf("pwdLen must be between 10 and 80") + return "", ErrInvalidBase85PwdLen } path := fmt.Sprintf("%s/%d'/%d'/%d'", BIP85_MASTER_PATH, AppPWD85, pwdLen, index) @@ -347,16 +406,21 @@ func DeriveBase85Password(masterKey *hdkeychain.ExtendedKey, pwdLen, index uint3 // Slice to the desired password length if len(encoded) < int(pwdLen) { - return "", fmt.Errorf("encoded length %d is less than requested length %d", len(encoded), pwdLen) + return "", fmt.Errorf( + "%w: encoded length %d is less than requested length %d", + ErrPasswordTooShort, len(encoded), pwdLen, + ) } return encoded[:pwdLen], nil } -// encodeBase85WithRFC1924Charset encodes data using Base85 with the RFC1924 character set +// encodeBase85WithRFC1924Charset encodes data using Base85 with the +// RFC1924 character set func encodeBase85WithRFC1924Charset(data []byte) string { // RFC1924 character set - charset := "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz!#$%&()*+-;<=>?@^_`{|}~" + charset := "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ" + + "abcdefghijklmnopqrstuvwxyz!#$%&()*+-;<=>?@^_`{|}~" const ( base85ChunkSize = 4 // Process 4 bytes at a time @@ -369,7 +433,9 @@ func encodeBase85WithRFC1924Charset(data []byte) string { copy(padded, data) var buf strings.Builder - buf.Grow(len(padded) * base85DigitCount / base85ChunkSize) // Each 4 bytes becomes 5 Base85 characters + + // Each 4 bytes becomes 5 Base85 characters + buf.Grow(len(padded) * base85DigitCount / base85ChunkSize) // Process in 4-byte chunks for i := 0; i < len(padded); i += base85ChunkSize { diff --git a/pkg/bip85/bip85_test.go b/pkg/bip85/bip85_test.go index ca6a79f..f4107cf 100644 --- a/pkg/bip85/bip85_test.go +++ b/pkg/bip85/bip85_test.go @@ -1,8 +1,5 @@ //nolint:gosec // G101: Test file contains BIP85 test vectors, not real credentials -//nolint:lll // Test vectors contain long lines -package bip85 - -//nolint:revive,unparam // Test file with BIP85 test vectors +package bip85_test import ( "bytes" @@ -11,336 +8,352 @@ import ( "strings" "testing" + "git.eeqj.de/sneak/secret/pkg/bip85" + "github.com/btcsuite/btcd/btcutil/hdkeychain" "github.com/tyler-smith/go-bip39" ) const ( // Test master BIP32 root key from the BIP85 specification - testMasterKey = "xprv9s21ZrQH143K2LBWUUQRFXhucrQqBpKdRRxNVq2zBqsx8HVqFk2uYo8kmbaLLHRdqtQpUm98uKfu3vca1LqdGhUtyoFnCNkfmXRyPXLjbKb" + testMasterKey = "xprv9s21ZrQH143K2LBWUUQRFXhucrQqBpKdRRxNVq2zBqsx8HVqFk2uYo8" + + "kmbaLLHRdqtQpUm98uKfu3vca1LqdGhUtyoFnCNkfmXRyPXLjbKb" // Test Case 1 - Basic entropy derivation with path m/83696968'/0'/0' testCase1Path = "m/83696968'/0'/0'" - testCase1ExpectedDerivedKey = "cca20ccb0e9a90feb0912870c3323b24874b0ca3d8018c4b96d0b97c0e82ded0" - testCase1ExpectedEntropy = "efecfbccffea313214232d29e71563d941229afb4338c21f9517c41aaa0d16f00b83d2a09ef747e7a64e8e2bd5a14869e693da66ce94ac2da570ab7ee48618f7" + testCase1ExpectedDerivedKey = "cca20ccb0e9a90feb0912870c3323b24" + + "874b0ca3d8018c4b96d0b97c0e82ded0" + testCase1ExpectedEntropy = "efecfbccffea313214232d29e71563d941229afb4338c21f" + + "9517c41aaa0d16f00b83d2a09ef747e7a64e8e2bd5a14869" + + "e693da66ce94ac2da570ab7ee48618f7" // Test Case 2 - Basic entropy derivation with path m/83696968'/0'/1' testCase2Path = "m/83696968'/0'/1'" - testCase2ExpectedDerivedKey = "503776919131758bb7de7beb6c0ae24894f4ec042c26032890c29359216e21ba" - testCase2ExpectedEntropy = "70c6e3e8ebee8dc4c0dbba66076819bb8c09672527c4277ca8729532ad711872218f826919f6b67218adde99018a6df9095ab2b58d803b5b93ec9802085a690e" + testCase2ExpectedDerivedKey = "503776919131758bb7de7beb6c0ae248" + + "94f4ec042c26032890c29359216e21ba" + testCase2ExpectedEntropy = "70c6e3e8ebee8dc4c0dbba66076819bb8c09672527c4277c" + + "a8729532ad711872218f826919f6b67218adde99018a6df9" + + "095ab2b58d803b5b93ec9802085a690e" // BIP85-DRNG-SHAKE256 test vector drngTestPath = "m/83696968'/0'/0'" - drngExpected80Bytes = "b78b1ee6b345eae6836c2d53d33c64cdaf9a696487be81b03e822dc84b3f1cd883d7559e53d175f243e4c349e822a957bbff9224bc5dde9492ef54e8a439f6bc8c7355b87a925a37ee405a7502991111" + drngExpected80Bytes = "b78b1ee6b345eae6836c2d53d33c64cdaf9a6964" + + "87be81b03e822dc84b3f1cd883d7559e53d175f2" + + "43e4c349e822a957bbff9224bc5dde9492ef54e8" + + "a439f6bc8c7355b87a925a37ee405a7502991111" // Python DRNG test vectors - pythonDRNG50BytesExpected = "b78b1ee6b345eae6836c2d53d33c64cdaf9a696487be81b03e822dc84b3f1cd883d7559e53d175f243e4c349e822a957bbff" - pythonDRNG100BytesExpected = "9224bc5dde9492ef54e8a439f6bc8c7355b87a925a37ee405a7502991111cd2dddaf1883f4e962abf4fb4b31cd28d5cf6b14f6ddcc9c19fd56d7f960a4b27f1d423a55dda4865aa6ddd6b4c26f18d400bb0a593e6c785d6d7e28c9c64608624318eddc01" - pythonDRNG150BytesExpected = "23750caa2a271f35faa6a3ca292b4be357404eca6842c69a3717dc3e41f7b38c67be492395b32221470aa08a2c489018c635a175f731245330e1f47091dbfb26f2923d10bd2e09280bffd1d94eb2a88f964aeb1774da04aad3bb1fdde0f77cd5ca79617ae317375417a51339523057bebef434c4400303890332e458425242f56a4293dad4f632b82713467b18ed6e1dab633220523d" - pythonDRNG20BytesExpected = "b78b1ee6b345eae6836c2d53d33c64cdaf9a6964" - pythonDRNG25BytesExpected = "87be81b03e822dc84b3f1cd883d7559e53d175f243e4c349e8" + pythonDRNG50BytesExpected = "b78b1ee6b345eae6836c2d53d33c64cdaf9a6964" + + "87be81b03e822dc84b3f1cd883d7559e53d175f2" + + "43e4c349e822a957bbff" + pythonDRNG100BytesExpected = "9224bc5dde9492ef54e8a439f6bc8c7355b87a92" + + "5a37ee405a7502991111cd2dddaf1883f4e962ab" + + "f4fb4b31cd28d5cf6b14f6ddcc9c19fd56d7f960" + + "a4b27f1d423a55dda4865aa6ddd6b4c26f18d400" + + "bb0a593e6c785d6d7e28c9c64608624318eddc01" + pythonDRNG150BytesExpected = "23750caa2a271f35faa6a3ca292b4be357404eca" + + "6842c69a3717dc3e41f7b38c67be492395b32221" + + "470aa08a2c489018c635a175f731245330e1f470" + + "91dbfb26f2923d10bd2e09280bffd1d94eb2a88f" + + "964aeb1774da04aad3bb1fdde0f77cd5ca79617a" + + "e317375417a51339523057bebef434c440030389" + + "0332e458425242f56a4293dad4f632b82713467b" + + "18ed6e1dab633220523d" + pythonDRNG20BytesExpected = "b78b1ee6b345eae6836c2d53d33c64cdaf9a6964" + pythonDRNG25BytesExpected = "87be81b03e822dc84b3f1cd883d7559e53d175f243e4c349e8" // BIP39 12 English words test vector bip39_12WordsPath = "m/83696968'/39'/0'/12'/0'" bip39_12WordsExpectedEntropy = "6250b68daf746d12a24d58b4787a714b" - bip39_12WordsExpectedMnemonic = "girl mad pet galaxy egg matter matrix prison refuse sense ordinary nose" + bip39_12WordsExpectedMnemonic = "girl mad pet galaxy egg matter matrix prison " + + "refuse sense ordinary nose" // BIP39 18 English words test vector bip39_18WordsPath = "m/83696968'/39'/0'/18'/0'" bip39_18WordsExpectedEntropy = "938033ed8b12698449d4bbca3c853c66b293ea1b1ce9d9dc" - bip39_18WordsExpectedMnemonic = "near account window bike charge season chef number sketch tomorrow excuse sniff circle vital hockey outdoor supply token" + bip39_18WordsExpectedMnemonic = "near account window bike charge season chef " + + "number sketch tomorrow excuse sniff circle vital hockey " + + "outdoor supply token" // BIP39 24 English words test vector - bip39_24WordsPath = "m/83696968'/39'/0'/24'/0'" - bip39_24WordsExpectedEntropy = "ae131e2312cdc61331542efe0d1077bac5ea803adf24b313a4f0e48e9c51f37f" - bip39_24WordsExpectedMnemonic = "puppy ocean match cereal symbol another shed magic wrap hammer bulb intact gadget divorce twin tonight reason outdoor destroy simple truth cigar social volcano" + bip39_24WordsPath = "m/83696968'/39'/0'/24'/0'" + bip39_24WordsExpectedEntropy = "ae131e2312cdc61331542efe0d1077ba" + + "c5ea803adf24b313a4f0e48e9c51f37f" + bip39_24WordsExpectedMnemonic = "puppy ocean match cereal symbol another " + + "shed magic wrap hammer bulb intact gadget divorce twin tonight " + + "reason outdoor destroy simple truth cigar social volcano" // HD-Seed WIF test vector hdWifPath = "m/83696968'/2'/0'" - hdWifExpectedEntropy = "7040bb53104f27367f317558e78a994ada7296c6fde36a364e5baf206e502bb1" - hdWifExpectedWIF = "Kzyv4uF39d4Jrw2W7UryTHwZr1zQVNk4dAFyqE6BuMrMh1Za7uhp" + hdWifExpectedEntropy = "7040bb53104f27367f317558e78a994a" + + "da7296c6fde36a364e5baf206e502bb1" + hdWifExpectedWIF = "Kzyv4uF39d4Jrw2W7UryTHwZr1zQVNk4dAFyqE6BuMrMh1Za7uhp" // XPRV test vector xprvPath = "m/83696968'/32'/0'" - xprvExpectedKey = "xprv9s21ZrQH143K2srSbCSg4m4kLvPMzcWydgmKEnMmoZUurYuBuYG46c6P71UGXMzmriLzCCBvKQWBUv3vPB3m1SATMhp3uEjXHJ42jFg7myX" + xprvExpectedKey = "xprv9s21ZrQH143K2srSbCSg4m4kLvPMzcWydgmKEnMmoZUurYuBuYG46c6" + + "P71UGXMzmriLzCCBvKQWBUv3vPB3m1SATMhp3uEjXHJ42jFg7myX" // HEX test vector hexPath = "m/83696968'/128169'/64'/0'" - hexExpectedEntropy = "492db4698cf3b73a5a24998aa3e9d7fa96275d85724a91e71aa2d645442f878555d078fd1f1f67e368976f04137b1f7a0d19232136ca50c44614af72b5582a5c" + hexExpectedEntropy = "492db4698cf3b73a5a24998aa3e9d7fa96275d85724a91e7" + + "1aa2d645442f878555d078fd1f1f67e368976f04137b1f7a" + + "0d19232136ca50c44614af72b5582a5c" // PWD Base64 test vector - pwdBase64Path = "m/83696968'/707764'/21'/0'" - pwdBase64ExpectedEntropy = "74a2e87a9ba0cdd549bdd2f9ea880d554c6c355b08ed25088cfa88f3f1c4f74632b652fd4a8f5fda43074c6f6964a3753b08bb5210c8f5e75c07a4c2a20bf6e9" + pwdBase64Path = "m/83696968'/707764'/21'/0'" + pwdBase64ExpectedEntropy = "74a2e87a9ba0cdd549bdd2f9ea880d554c6c355b08ed2508" + + "8cfa88f3f1c4f74632b652fd4a8f5fda43074c6f6964a375" + + "3b08bb5210c8f5e75c07a4c2a20bf6e9" pwdBase64ExpectedPassword = "dKLoepugzdVJvdL56ogNV" // PWD Base85 test vector - pwdBase85Path = "m/83696968'/707785'/12'/0'" - pwdBase85ExpectedEntropy = "f7cfe56f63dca2490f65fcbf9ee63dcd85d18f751b6b5e1c1b8733af6459c904a75e82b4a22efff9b9e69de2144b293aa8714319a054b6cb55826a8e51425209" + pwdBase85Path = "m/83696968'/707785'/12'/0'" + pwdBase85ExpectedEntropy = "f7cfe56f63dca2490f65fcbf9ee63dcd85d18f751b6b5e1c" + + "1b8733af6459c904a75e82b4a22efff9b9e69de2144b293a" + + "a8714319a054b6cb55826a8e51425209" pwdBase85ExpectedPassword = "_s`{TW89)i4`" // Test keys for parsing tests - testInvalidMasterKey = "xprv9s21ZrQH143K2LBWUUQRFXhucrQqBpKdRRxNVq2zBqsx8HVqFk2uYo8kmbaLLHRdqtQpUm98uKfu3vca1LqdGhUtyoFnCNkfmXRyPXLjbXX" - testTestnetMasterKey = "tprv8ZgxMBicQKsPeWHBt7a68nPnvgTnuDhUgDWC8wZCgA8GahrQ3f3uWpq7wE7Uc1dLBnCe1hhCZ886K6ND37memRDWqsA9HgSKDXtwh2Qxo6J" + testInvalidMasterKey = "xprv9s21ZrQH143K2LBWUUQRFXhucrQqBpKdRRxNVq2zBqsx8HVqFk2uYo8" + + "kmbaLLHRdqtQpUm98uKfu3vca1LqdGhUtyoFnCNkfmXRyPXLjbXX" + testTestnetMasterKey = "tprv8ZgxMBicQKsPeWHBt7a68nPnvgTnuDhUgDWC8wZCgA8GahrQ3f3uWpq7" + + "wE7Uc1dLBnCe1hhCZ886K6ND37memRDWqsA9HgSKDXtwh2Qxo6J" ) // logTestVector logs test information in a cleaner, more concise format func logTestVector(t *testing.T, title string) { + t.Helper() t.Logf("=== TEST: %s ===", title) } -// TestDerivedKey is a helper function to test the derived key directly -func TestDerivedKey(t *testing.T) { - logTestVector(t, "Derived Child Keys") +// mustParseTestMasterKey parses the shared test master key. +func mustParseTestMasterKey(t *testing.T) *hdkeychain.ExtendedKey { + t.Helper() - masterKey, err := ParseMasterKey(testMasterKey) + masterKey, err := bip85.ParseMasterKey(testMasterKey) if err != nil { t.Fatalf("Failed to parse master key: %v", err) } - // Test case 1 - t.Logf("Deriving key for path: %s", testCase1Path) - derivedKeyBytes, err := DeriveChildKey(masterKey, testCase1Path) + return masterKey +} + +// checkDerivedChildKey derives the child key at path and compares it +// against the expected hex value. +func checkDerivedChildKey(t *testing.T, path, expected string) { + t.Helper() + + masterKey := mustParseTestMasterKey(t) + + t.Logf("Deriving key for path: %s", path) + + derivedKeyBytes, err := bip85.DeriveChildKey(masterKey, path) if err != nil { t.Fatalf("Failed to derive child key: %v", err) } derivedKeyHex := hex.EncodeToString(derivedKeyBytes) - t.Logf("EXPECTED: %s", testCase1ExpectedDerivedKey) + + t.Logf("EXPECTED: %s", expected) t.Logf("ACTUAL: %s", derivedKeyHex) - if derivedKeyHex != testCase1ExpectedDerivedKey { - t.Errorf("Expected derived key bytes %s, got %s", testCase1ExpectedDerivedKey, derivedKeyHex) + if derivedKeyHex != expected { + t.Errorf( + "Expected derived key bytes %s, got %s", + expected, + derivedKeyHex, + ) } else { - t.Logf("RESULT: PASS ✓") + t.Logf("RESULT: PASS") } +} + +// TestDerivedKey tests the derived key directly +func TestDerivedKey(t *testing.T) { + t.Parallel() + + logTestVector(t, "Derived Child Keys") + + // Test case 1 + checkDerivedChildKey(t, testCase1Path, testCase1ExpectedDerivedKey) // Test case 2 - t.Logf("Deriving key for path: %s", testCase2Path) - derivedKeyBytes, err = DeriveChildKey(masterKey, testCase2Path) + checkDerivedChildKey(t, testCase2Path, testCase2ExpectedDerivedKey) +} + +// runEntropyVectorTest checks a BIP85 entropy derivation test vector. +func runEntropyVectorTest(t *testing.T, title, path, expectedEntropy string) { + t.Helper() + + logTestVector(t, title) + + masterKey := mustParseTestMasterKey(t) + + t.Logf("Test path: %s", path) + + entropy, err := bip85.DeriveBIP85Entropy(masterKey, path) if err != nil { - t.Fatalf("Failed to derive child key: %v", err) + t.Fatalf("Failed to derive entropy: %v", err) } - derivedKeyHex = hex.EncodeToString(derivedKeyBytes) - t.Logf("EXPECTED: %s", testCase2ExpectedDerivedKey) - t.Logf("ACTUAL: %s", derivedKeyHex) + derivedEntropyHex := hex.EncodeToString(entropy) - if derivedKeyHex != testCase2ExpectedDerivedKey { - t.Errorf("Expected derived key bytes %s, got %s", testCase2ExpectedDerivedKey, derivedKeyHex) + t.Logf("EXPECTED: %s", expectedEntropy) + t.Logf("ACTUAL: %s", derivedEntropyHex) + + if derivedEntropyHex != expectedEntropy { + t.Errorf( + "Expected derived entropy %s, got %s", + expectedEntropy, + derivedEntropyHex, + ) } else { - t.Logf("RESULT: PASS ✓") + t.Logf("RESULT: PASS") } } // TestCase1 tests the first test vector from the BIP85 specification func TestCase1(t *testing.T) { - logTestVector(t, "Test Case 1") + t.Parallel() - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } - - t.Logf("Test path: %s", testCase1Path) - entropy, err := DeriveBIP85Entropy(masterKey, testCase1Path) - if err != nil { - t.Fatalf("Failed to derive entropy: %v", err) - } - - derivedEntropyHex := hex.EncodeToString(entropy) - t.Logf("EXPECTED: %s", testCase1ExpectedEntropy) - t.Logf("ACTUAL: %s", derivedEntropyHex) - - if derivedEntropyHex != testCase1ExpectedEntropy { - t.Errorf("Expected derived entropy %s, got %s", testCase1ExpectedEntropy, derivedEntropyHex) - } else { - t.Logf("RESULT: PASS ✓") - } + runEntropyVectorTest(t, "Test Case 1", testCase1Path, testCase1ExpectedEntropy) } // TestCase2 tests the second test vector from the BIP85 specification func TestCase2(t *testing.T) { - logTestVector(t, "Test Case 2") + t.Parallel() - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + runEntropyVectorTest(t, "Test Case 2", testCase2Path, testCase2ExpectedEntropy) +} - t.Logf("Test path: %s", testCase2Path) - entropy, err := DeriveBIP85Entropy(masterKey, testCase2Path) +// runBIP39VectorTest checks a BIP39 mnemonic derivation test vector. +func runBIP39VectorTest( + t *testing.T, + title, path string, + words uint32, + expectedEntropy, expectedMnemonic string, +) { + t.Helper() + + logTestVector(t, title) + + masterKey := mustParseTestMasterKey(t) + + t.Logf("Path: %s", path) + t.Logf("Parameters: Language=English(0), Words=%d, Index=0", words) + + // Derive the BIP39 mnemonic entropy + entropy, err := bip85.DeriveBIP39Entropy(masterKey, 0, words, 0) if err != nil { - t.Fatalf("Failed to derive entropy: %v", err) + t.Fatalf("Failed to derive BIP39 entropy: %v", err) } derivedEntropyHex := hex.EncodeToString(entropy) - t.Logf("EXPECTED: %s", testCase2ExpectedEntropy) - t.Logf("ACTUAL: %s", derivedEntropyHex) - if derivedEntropyHex != testCase2ExpectedEntropy { - t.Errorf("Expected derived entropy %s, got %s", testCase2ExpectedEntropy, derivedEntropyHex) + t.Logf("EXPECTED ENTROPY: %s", expectedEntropy) + t.Logf("ACTUAL ENTROPY: %s", derivedEntropyHex) + + if derivedEntropyHex != expectedEntropy { + t.Errorf( + "Expected derived entropy %s, got %s", + expectedEntropy, + derivedEntropyHex, + ) } else { - t.Logf("RESULT: PASS ✓") + t.Logf("ENTROPY MATCH: PASS") + } + + // Convert entropy to mnemonic + mnemonic, err := bip39.NewMnemonic(entropy) + if err != nil { + t.Fatalf("Failed to create mnemonic: %v", err) + } + + t.Logf("EXPECTED MNEMONIC: %s", expectedMnemonic) + t.Logf("ACTUAL MNEMONIC: %s", mnemonic) + + if mnemonic != expectedMnemonic { + t.Errorf( + "Expected mnemonic '%s', got '%s'", + expectedMnemonic, + mnemonic, + ) + } else { + t.Logf("MNEMONIC MATCH: PASS") } } // TestBIP39_12EnglishWords tests the BIP39 12 English words test vector func TestBIP39_12EnglishWords(t *testing.T) { - logTestVector(t, "BIP39 12 English Words") + t.Parallel() - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } - - t.Logf("Path: %s", bip39_12WordsPath) - t.Logf("Parameters: Language=English(0), Words=12, Index=0") - - // BIP39 English 12 word mnemonic - entropy, err := DeriveBIP39Entropy(masterKey, 0, 12, 0) - if err != nil { - t.Fatalf("Failed to derive BIP39 entropy: %v", err) - } - - derivedEntropyHex := hex.EncodeToString(entropy) - t.Logf("EXPECTED ENTROPY: %s", bip39_12WordsExpectedEntropy) - t.Logf("ACTUAL ENTROPY: %s", derivedEntropyHex) - - if derivedEntropyHex != bip39_12WordsExpectedEntropy { - t.Errorf("Expected derived entropy %s, got %s", bip39_12WordsExpectedEntropy, derivedEntropyHex) - } else { - t.Logf("ENTROPY MATCH: PASS ✓") - } - - // Convert entropy to mnemonic - mnemonic, err := bip39.NewMnemonic(entropy) - if err != nil { - t.Fatalf("Failed to create mnemonic: %v", err) - } - - t.Logf("EXPECTED MNEMONIC: %s", bip39_12WordsExpectedMnemonic) - t.Logf("ACTUAL MNEMONIC: %s", mnemonic) - - if mnemonic != bip39_12WordsExpectedMnemonic { - t.Errorf("Expected mnemonic '%s', got '%s'", bip39_12WordsExpectedMnemonic, mnemonic) - } else { - t.Logf("MNEMONIC MATCH: PASS ✓") - } + runBIP39VectorTest( + t, + "BIP39 12 English Words", + bip39_12WordsPath, + 12, + bip39_12WordsExpectedEntropy, + bip39_12WordsExpectedMnemonic, + ) } // TestBIP39_18EnglishWords tests the BIP39 18 English words test vector func TestBIP39_18EnglishWords(t *testing.T) { - logTestVector(t, "BIP39 18 English Words") + t.Parallel() - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } - - t.Logf("Path: %s", bip39_18WordsPath) - t.Logf("Parameters: Language=English(0), Words=18, Index=0") - - // BIP39 English 18 word mnemonic - entropy, err := DeriveBIP39Entropy(masterKey, 0, 18, 0) - if err != nil { - t.Fatalf("Failed to derive BIP39 entropy: %v", err) - } - - derivedEntropyHex := hex.EncodeToString(entropy) - t.Logf("EXPECTED ENTROPY: %s", bip39_18WordsExpectedEntropy) - t.Logf("ACTUAL ENTROPY: %s", derivedEntropyHex) - - if derivedEntropyHex != bip39_18WordsExpectedEntropy { - t.Errorf("Expected derived entropy %s, got %s", bip39_18WordsExpectedEntropy, derivedEntropyHex) - } else { - t.Logf("ENTROPY MATCH: PASS ✓") - } - - // Convert entropy to mnemonic - mnemonic, err := bip39.NewMnemonic(entropy) - if err != nil { - t.Fatalf("Failed to create mnemonic: %v", err) - } - - t.Logf("EXPECTED MNEMONIC: %s", bip39_18WordsExpectedMnemonic) - t.Logf("ACTUAL MNEMONIC: %s", mnemonic) - - if mnemonic != bip39_18WordsExpectedMnemonic { - t.Errorf("Expected mnemonic '%s', got '%s'", bip39_18WordsExpectedMnemonic, mnemonic) - } else { - t.Logf("MNEMONIC MATCH: PASS ✓") - } + runBIP39VectorTest( + t, + "BIP39 18 English Words", + bip39_18WordsPath, + 18, + bip39_18WordsExpectedEntropy, + bip39_18WordsExpectedMnemonic, + ) } // TestBIP39_24EnglishWords tests the BIP39 24 English words test vector func TestBIP39_24EnglishWords(t *testing.T) { - logTestVector(t, "BIP39 24 English Words") + t.Parallel() - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } - - t.Logf("Path: %s", bip39_24WordsPath) - t.Logf("Parameters: Language=English(0), Words=24, Index=0") - - // BIP39 English 24 word mnemonic - entropy, err := DeriveBIP39Entropy(masterKey, 0, 24, 0) - if err != nil { - t.Fatalf("Failed to derive BIP39 entropy: %v", err) - } - - derivedEntropyHex := hex.EncodeToString(entropy) - t.Logf("EXPECTED ENTROPY: %s", bip39_24WordsExpectedEntropy) - t.Logf("ACTUAL ENTROPY: %s", derivedEntropyHex) - - if derivedEntropyHex != bip39_24WordsExpectedEntropy { - t.Errorf("Expected derived entropy %s, got %s", bip39_24WordsExpectedEntropy, derivedEntropyHex) - } else { - t.Logf("ENTROPY MATCH: PASS ✓") - } - - // Convert entropy to mnemonic - mnemonic, err := bip39.NewMnemonic(entropy) - if err != nil { - t.Fatalf("Failed to create mnemonic: %v", err) - } - - t.Logf("EXPECTED MNEMONIC: %s", bip39_24WordsExpectedMnemonic) - t.Logf("ACTUAL MNEMONIC: %s", mnemonic) - - if mnemonic != bip39_24WordsExpectedMnemonic { - t.Errorf("Expected mnemonic '%s', got '%s'", bip39_24WordsExpectedMnemonic, mnemonic) - } else { - t.Logf("MNEMONIC MATCH: PASS ✓") - } + runBIP39VectorTest( + t, + "BIP39 24 English Words", + bip39_24WordsPath, + 24, + bip39_24WordsExpectedEntropy, + bip39_24WordsExpectedMnemonic, + ) } // TestHD_WIF tests the WIF test vector func TestHD_WIF(t *testing.T) { + t.Parallel() + logTestVector(t, "HD-Seed WIF") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // First verify the entropy derivation t.Logf("Path: %s", hdWifPath) - entropy, err := DeriveBIP85Entropy(masterKey, hdWifPath) + entropy, err := bip85.DeriveBIP85Entropy(masterKey, hdWifPath) if err != nil { t.Fatalf("Failed to derive entropy: %v", err) } - // Expected entropy from BIP85 spec - derivedEntropyHex := hex.EncodeToString(entropy[:32]) // WIF uses first 32 bytes + // Expected entropy from BIP85 spec; WIF uses first 32 bytes + derivedEntropyHex := hex.EncodeToString(entropy[:32]) if derivedEntropyHex != hdWifExpectedEntropy { - t.Errorf("Entropy mismatch!\nExpected: %s\nGot: %s", hdWifExpectedEntropy, derivedEntropyHex) + t.Errorf( + "Entropy mismatch!\nExpected: %s\nGot: %s", + hdWifExpectedEntropy, + derivedEntropyHex, + ) } // Now test the WIF derivation - wif, err := DeriveWIFKey(masterKey, 0) + wif, err := bip85.DeriveWIFKey(masterKey, 0) if err != nil { t.Fatalf("Failed to derive WIF key: %v", err) } @@ -351,60 +364,62 @@ func TestHD_WIF(t *testing.T) { if wif != hdWifExpectedWIF { t.Errorf("Expected WIF %s, got %s", hdWifExpectedWIF, wif) } else { - t.Logf("RESULT: PASS ✓") + t.Logf("RESULT: PASS") } } // TestXPRV tests the XPRV test vector func TestXPRV(t *testing.T) { + t.Parallel() + logTestVector(t, "XPRV") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) t.Logf("Path: %s", xprvPath) - derivedKey, err := DeriveXPRV(masterKey, 0) + + derivedKey, err := bip85.DeriveXPRV(masterKey, 0) if err != nil { t.Fatalf("Failed to derive XPRV: %v", err) } derivedXPRV := derivedKey.String() + t.Logf("EXPECTED XPRV: %s", xprvExpectedKey) t.Logf("ACTUAL XPRV: %s", derivedXPRV) if derivedXPRV != xprvExpectedKey { t.Errorf("Expected XPRV %s, got %s", xprvExpectedKey, derivedXPRV) } else { - t.Logf("RESULT: PASS ✓") + t.Logf("RESULT: PASS") } } // TestDRNG_SHAKE256 tests the BIP85-DRNG-SHAKE256 test vector func TestDRNG_SHAKE256(t *testing.T) { + t.Parallel() + logTestVector(t, "DRNG-SHAKE256") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // Derive entropy for the DRNG - entropy, err := DeriveBIP85Entropy(masterKey, drngTestPath) + entropy, err := bip85.DeriveBIP85Entropy(masterKey, drngTestPath) if err != nil { t.Fatalf("Failed to derive entropy: %v", err) } // Create DRNG - drng := NewBIP85DRNG(entropy) + drng := bip85.NewBIP85DRNG(entropy) // Read 80 bytes buffer := make([]byte, 80) + n, err := drng.Read(buffer) if err != nil { t.Fatalf("Failed to read from DRNG: %v", err) } + if n != 80 { t.Errorf("Expected to read 80 bytes, got %d", n) } @@ -412,164 +427,154 @@ func TestDRNG_SHAKE256(t *testing.T) { hexOutput := hex.EncodeToString(buffer) if !strings.EqualFold(hexOutput, drngExpected80Bytes) { - t.Errorf("Expected DRNG output:\n%s\n\nGot:\n%s", drngExpected80Bytes, hexOutput) + t.Errorf( + "Expected DRNG output:\n%s\n\nGot:\n%s", + drngExpected80Bytes, + hexOutput, + ) } } +// readDRNGHex reads size bytes from drng and returns the hex encoding. +func readDRNGHex(t *testing.T, drng *bip85.DRNG, size int, label string) string { + t.Helper() + + buf := make([]byte, size) + + _, err := drng.Read(buf) + if err != nil { + t.Fatalf("Failed to read %s from DRNG: %v", label, err) + } + + return hex.EncodeToString(buf) +} + // TestPythonDRNGVectors tests the DRNG vectors from the Python implementation func TestPythonDRNGVectors(t *testing.T) { + t.Parallel() + logTestVector(t, "Python DRNG Vectors") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // Derive entropy for the DRNG - entropy, err := DeriveBIP85Entropy(masterKey, drngTestPath) + entropy, err := bip85.DeriveBIP85Entropy(masterKey, drngTestPath) if err != nil { t.Fatalf("Failed to derive entropy: %v", err) } // Create DRNG - drng := NewBIP85DRNG(entropy) + drng := bip85.NewBIP85DRNG(entropy) // Test vector 1: Read 50 bytes - buffer1 := make([]byte, 50) - _, err = drng.Read(buffer1) - if err != nil { - t.Fatalf("Failed to read 50 bytes from DRNG: %v", err) - } - actual1 := hex.EncodeToString(buffer1) + actual1 := readDRNGHex(t, drng, 50, "50 bytes") if actual1 != pythonDRNG50BytesExpected { - t.Errorf("Test vector 1 failed. Expected:\n%s\n\nGot:\n%s", pythonDRNG50BytesExpected, actual1) + t.Errorf( + "Test vector 1 failed. Expected:\n%s\n\nGot:\n%s", + pythonDRNG50BytesExpected, + actual1, + ) } // Test vector 2: Read 100 bytes - buffer2 := make([]byte, 100) - _, err = drng.Read(buffer2) - if err != nil { - t.Fatalf("Failed to read 100 bytes from DRNG: %v", err) - } - actual2 := hex.EncodeToString(buffer2) + actual2 := readDRNGHex(t, drng, 100, "100 bytes") if actual2 != pythonDRNG100BytesExpected { - t.Errorf("Test vector 2 failed. Expected:\n%s\n\nGot:\n%s", pythonDRNG100BytesExpected, actual2) + t.Errorf( + "Test vector 2 failed. Expected:\n%s\n\nGot:\n%s", + pythonDRNG100BytesExpected, + actual2, + ) } // Test vector 3: Read 150 bytes - buffer3 := make([]byte, 150) - _, err = drng.Read(buffer3) - if err != nil { - t.Fatalf("Failed to read 150 bytes from DRNG: %v", err) - } - actual3 := hex.EncodeToString(buffer3) + actual3 := readDRNGHex(t, drng, 150, "150 bytes") if actual3 != pythonDRNG150BytesExpected { - t.Errorf("Test vector 3 failed. Expected:\n%s\n\nGot:\n%s", pythonDRNG150BytesExpected, actual3) + t.Errorf( + "Test vector 3 failed. Expected:\n%s\n\nGot:\n%s", + pythonDRNG150BytesExpected, + actual3, + ) } // Test with fresh DRNG - drng2 := NewBIP85DRNG(entropy) - buffer4 := make([]byte, 20) - _, err = drng2.Read(buffer4) - if err != nil { - t.Fatalf("Failed to read 20 bytes from DRNG: %v", err) - } - actual4 := hex.EncodeToString(buffer4) + drng2 := bip85.NewBIP85DRNG(entropy) + + actual4 := readDRNGHex(t, drng2, 20, "20 bytes") if actual4 != pythonDRNG20BytesExpected { - t.Errorf("Test vector 4 failed. Expected:\n%s\n\nGot:\n%s", pythonDRNG20BytesExpected, actual4) + t.Errorf( + "Test vector 4 failed. Expected:\n%s\n\nGot:\n%s", + pythonDRNG20BytesExpected, + actual4, + ) } // Read another 25 bytes - buffer5 := make([]byte, 25) - _, err = drng2.Read(buffer5) - if err != nil { - t.Fatalf("Failed to read 25 bytes from DRNG: %v", err) - } - actual5 := hex.EncodeToString(buffer5) + actual5 := readDRNGHex(t, drng2, 25, "25 bytes") if actual5 != pythonDRNG25BytesExpected { - t.Errorf("Test vector 5 failed. Expected:\n%s\n\nGot:\n%s", pythonDRNG25BytesExpected, actual5) + t.Errorf( + "Test vector 5 failed. Expected:\n%s\n\nGot:\n%s", + pythonDRNG25BytesExpected, + actual5, + ) } } +// drngReadChunks performs sequential reads of the given sizes from drng +// and returns the concatenated output. +func drngReadChunks(t *testing.T, drng *bip85.DRNG, name string, sizes ...int) []byte { + t.Helper() + + total := 0 + + for _, size := range sizes { + total += size + } + + out := make([]byte, 0, total) + + for _, size := range sizes { + buf := make([]byte, size) + + _, err := drng.Read(buf) + if err != nil { + t.Fatalf("Failed to read from %s: %v", name, err) + } + + out = append(out, buf...) + } + + return out +} + // TestDRNGDeterminism tests the deterministic behavior of the DRNG func TestDRNGDeterminism(t *testing.T) { + t.Parallel() + logTestVector(t, "DRNG Determinism") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // Derive entropy for the DRNG - entropy, err := DeriveBIP85Entropy(masterKey, drngTestPath) + entropy, err := bip85.DeriveBIP85Entropy(masterKey, drngTestPath) if err != nil { t.Fatalf("Failed to derive entropy: %v", err) } // Create 3 DRNGs with the same seed - drng1 := NewBIP85DRNG(entropy) - drng2 := NewBIP85DRNG(entropy) - drng3 := NewBIP85DRNG(entropy) + drng1 := bip85.NewBIP85DRNG(entropy) + drng2 := bip85.NewBIP85DRNG(entropy) + drng3 := bip85.NewBIP85DRNG(entropy) - // Read from drng1 with multiple calls - buf1a := make([]byte, 10) - buf1b := make([]byte, 20) - buf1c := make([]byte, 30) - buf1d := make([]byte, 40) - _, err = drng1.Read(buf1a) - if err != nil { - t.Fatalf("Failed to read from drng1: %v", err) - } - _, err = drng1.Read(buf1b) - if err != nil { - t.Fatalf("Failed to read from drng1: %v", err) - } - _, err = drng1.Read(buf1c) - if err != nil { - t.Fatalf("Failed to read from drng1: %v", err) - } - _, err = drng1.Read(buf1d) - if err != nil { - t.Fatalf("Failed to read from drng1: %v", err) - } - - // Read from drng2 with multiple calls in different order - buf2a := make([]byte, 40) - buf2b := make([]byte, 30) - buf2c := make([]byte, 20) - buf2d := make([]byte, 10) - _, err = drng2.Read(buf2a) - if err != nil { - t.Fatalf("Failed to read from drng2: %v", err) - } - _, err = drng2.Read(buf2b) - if err != nil { - t.Fatalf("Failed to read from drng2: %v", err) - } - _, err = drng2.Read(buf2c) - if err != nil { - t.Fatalf("Failed to read from drng2: %v", err) - } - _, err = drng2.Read(buf2d) - if err != nil { - t.Fatalf("Failed to read from drng2: %v", err) - } - - // Read from drng3 with a single call - buf3 := make([]byte, 100) - _, err = drng3.Read(buf3) - if err != nil { - t.Fatalf("Failed to read from drng3: %v", err) - } - - // Combine the results from the multiple reads - result1 := append(append(append(buf1a, buf1b...), buf1c...), buf1d...) - result2 := append(append(append(buf2a, buf2b...), buf2c...), buf2d...) + // Read the same amount of data with differently sized read calls + result1 := drngReadChunks(t, drng1, "drng1", 10, 20, 30, 40) + result2 := drngReadChunks(t, drng2, "drng2", 40, 30, 20, 10) + buf3 := drngReadChunks(t, drng3, "drng3", 100) // All results should be identical if !bytes.Equal(result1, result2) { t.Errorf("Expected drng1 and drng2 to produce identical outputs") } + if !bytes.Equal(result2, buf3) { t.Errorf("Expected drng2 and drng3 to produce identical outputs") } @@ -577,31 +582,33 @@ func TestDRNGDeterminism(t *testing.T) { // TestDRNGLengths tests the DRNG with different lengths func TestDRNGLengths(t *testing.T) { + t.Parallel() + logTestVector(t, "DRNG Lengths") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // Derive entropy for the DRNG - entropy, err := DeriveBIP85Entropy(masterKey, drngTestPath) + entropy, err := bip85.DeriveBIP85Entropy(masterKey, drngTestPath) if err != nil { t.Fatalf("Failed to derive entropy: %v", err) } // Create DRNG - drng := NewBIP85DRNG(entropy) + drng := bip85.NewBIP85DRNG(entropy) // Test various lengths lengths := []int{1, 10, 100, 1000, 10000} for _, length := range lengths { buffer := make([]byte, length) + n, err := drng.Read(buffer) if err != nil { t.Errorf("Failed to read %d bytes: %v", length, err) + continue } + if n != length { t.Errorf("Expected to read %d bytes, got %d", length, n) } @@ -610,6 +617,8 @@ func TestDRNGLengths(t *testing.T) { // TestDRNGExceptions tests error handling in the DRNG func TestDRNGExceptions(t *testing.T) { + t.Parallel() + logTestVector(t, "DRNG Exceptions") // Test with entropy of the wrong size @@ -617,6 +626,8 @@ func TestDRNGExceptions(t *testing.T) { for _, size := range testCases { t.Run(fmt.Sprintf("EntropySize_%d", size), func(t *testing.T) { + t.Parallel() + entropy := make([]byte, size) // Use a function to capture the panic @@ -629,10 +640,13 @@ func TestDRNGExceptions(t *testing.T) { }() // This should panic for any size != 64 - _ = NewBIP85DRNG(entropy) + _ = bip85.NewBIP85DRNG(entropy) // If we get here without panic, it's an error - t.Errorf("Expected panic for entropy length %d, but it didn't happen", size) + t.Errorf( + "Expected panic for entropy length %d, but it didn't happen", + size, + ) } testPanic() @@ -642,36 +656,38 @@ func TestDRNGExceptions(t *testing.T) { // TestDRNGDifferentSizes tests the DRNG with different buffer sizes func TestDRNGDifferentSizes(t *testing.T) { + t.Parallel() + logTestVector(t, "DRNG Different Sizes") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) - entropy, err := DeriveBIP85Entropy(masterKey, drngTestPath) + entropy, err := bip85.DeriveBIP85Entropy(masterKey, drngTestPath) if err != nil { t.Fatalf("Failed to derive entropy: %v", err) } // Create DRNG - drng := NewBIP85DRNG(entropy) + drng := bip85.NewBIP85DRNG(entropy) // Test reading different sizes for _, size := range []int{32, 64, 128, 256} { buffer := make([]byte, size) + n, err := drng.Read(buffer) if err != nil { t.Fatalf("Failed to read %d bytes from DRNG: %v", size, err) } + if n != size { t.Errorf("Expected to read %d bytes, got %d", size, n) } } - // Test deterministic behavior - two DRNGs with the same seed should produce the same output - drng1 := NewBIP85DRNG(entropy) - drng2 := NewBIP85DRNG(entropy) + // Test deterministic behavior - two DRNGs with the same seed should + // produce the same output + drng1 := bip85.NewBIP85DRNG(entropy) + drng2 := bip85.NewBIP85DRNG(entropy) buffer1 := make([]byte, 32) buffer2 := make([]byte, 32) @@ -690,8 +706,10 @@ func TestDRNGDifferentSizes(t *testing.T) { t.Errorf("Expected identical outputs from DRNGs with same seed") } - // Reading another 32 bytes should produce different output from the first read + // Reading another 32 bytes should produce different output from the + // first read buffer3 := make([]byte, 32) + _, err = drng1.Read(buffer3) if err != nil { t.Fatalf("Failed to read second buffer from DRNG: %v", err) @@ -704,59 +722,69 @@ func TestDRNGDifferentSizes(t *testing.T) { // TestMasterKeyParsing tests parsing of different master key formats func TestMasterKeyParsing(t *testing.T) { + t.Parallel() + logTestVector(t, "Master Key Parsing") // Test valid master key t.Logf("Testing valid master key") - _, err := ParseMasterKey(testMasterKey) + + _, err := bip85.ParseMasterKey(testMasterKey) if err != nil { t.Errorf("Failed to parse valid master key: %v", err) } else { - t.Logf("Valid master key parsed successfully: PASS ✓") + t.Logf("Valid master key parsed successfully: PASS") } // Test invalid master key (wrong checksum) t.Logf("Testing invalid master key (corrupted)") - _, err = ParseMasterKey(testInvalidMasterKey) + + _, err = bip85.ParseMasterKey(testInvalidMasterKey) if err == nil { t.Errorf("Expected error for invalid master key, but got nil") } else { t.Logf("Got expected error for invalid master key: %v", err) - t.Logf("RESULT: PASS ✓") + t.Logf("RESULT: PASS") } // Test testnet master key (tprv) t.Logf("Testing testnet master key format") - testnetMasterKey, err := ParseMasterKey(testTestnetMasterKey) + + testnetMasterKey, err := bip85.ParseMasterKey(testTestnetMasterKey) if err != nil { t.Errorf("Failed to parse testnet master key: %v", err) + + return + } + + t.Logf("Testnet master key parsed successfully: PASS") + + // Test that XPRV derivation using a testnet master key produces a + // testnet XPRV + derivedKey, err := bip85.DeriveXPRV(testnetMasterKey, 0) + if err != nil { + t.Fatalf("Failed to derive XPRV from testnet key: %v", err) + } + + derivedKeyStr := derivedKey.String() + if !strings.HasPrefix(derivedKeyStr, "tprv") { + t.Errorf( + "Expected derived key to be testnet (tprv prefix), got: %s", + derivedKeyStr, + ) } else { - t.Logf("Testnet master key parsed successfully: PASS ✓") - - // Test that XPRV derivation using a testnet master key produces a testnet XPRV - derivedKey, err := DeriveXPRV(testnetMasterKey, 0) - if err != nil { - t.Fatalf("Failed to derive XPRV from testnet key: %v", err) - } - - derivedKeyStr := derivedKey.String() - if !strings.HasPrefix(derivedKeyStr, "tprv") { - t.Errorf("Expected derived key to be testnet (tprv prefix), got: %s", derivedKeyStr) - } else { - t.Logf("Testnet XPRV derived successfully: %s", derivedKeyStr) - t.Logf("RESULT: PASS ✓") - } + t.Logf("Testnet XPRV derived successfully: %s", derivedKeyStr) + t.Logf("RESULT: PASS") } } // TestDifferentPathFormats tests different path format expressions func TestDifferentPathFormats(t *testing.T) { + t.Parallel() + logTestVector(t, "Path Formats") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // Define equivalent paths in different formats paths := []string{ @@ -766,50 +794,57 @@ func TestDifferentPathFormats(t *testing.T) { "83696968'/0'/0'", } - var results [][]byte + results := make([][]byte, 0, len(paths)) // Derive entropy using each path for i, path := range paths { t.Logf("Testing path format %d: %s", i+1, path) - entropy, err := DeriveBIP85Entropy(masterKey, path) + + entropy, err := bip85.DeriveBIP85Entropy(masterKey, path) if err != nil { t.Errorf("Failed to derive entropy with path %s: %v", path, err) + continue } results = append(results, entropy) - t.Logf("Derivation succeeded: PASS ✓") + + t.Logf("Derivation succeeded: PASS") } // Verify all results are the same for i := 1; i < len(results); i++ { if !bytes.Equal(results[0], results[i]) { - t.Errorf("Path %s produced different entropy than path %s", paths[0], paths[i]) + t.Errorf( + "Path %s produced different entropy than path %s", + paths[0], + paths[i], + ) } } if len(results) > 1 { - t.Logf("All equivalent path formats produced the same entropy: PASS ✓") + t.Logf("All equivalent path formats produced the same entropy: PASS") } } // TestDirectBase85Encoding tests direct Base85 encoding with the test vector entropy func TestDirectBase85Encoding(t *testing.T) { + t.Parallel() + logTestVector(t, "Direct Base85 Encoding") // Parse the master key - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // First, derive the entropy and verify it matches the test vector - derivedEntropy, err := DeriveBIP85Entropy(masterKey, pwdBase85Path) + derivedEntropy, err := bip85.DeriveBIP85Entropy(masterKey, pwdBase85Path) if err != nil { t.Fatalf("Failed to derive entropy: %v", err) } - // This is the expected entropy from the BIP85 spec for the Base85 test vector + // This is the expected entropy from the BIP85 spec for the Base85 + // test vector expectedEntropy, err := hex.DecodeString(pwdBase85ExpectedEntropy) if err != nil { t.Fatalf("Failed to decode expected entropy hex: %v", err) @@ -818,7 +853,11 @@ func TestDirectBase85Encoding(t *testing.T) { // Verify the derived entropy matches the expected entropy derivedEntropyHex := hex.EncodeToString(derivedEntropy) if derivedEntropyHex != pwdBase85ExpectedEntropy { - t.Errorf("Entropy mismatch!\nExpected: %s\nGot: %s", pwdBase85ExpectedEntropy, derivedEntropyHex) + t.Errorf( + "Entropy mismatch!\nExpected: %s\nGot: %s", + pwdBase85ExpectedEntropy, + derivedEntropyHex, + ) } // Verify the entropy bytes match @@ -827,152 +866,156 @@ func TestDirectBase85Encoding(t *testing.T) { } // Now test the password generation - pwd, err := DeriveBase85Password(masterKey, 12, 0) + pwd, err := bip85.DeriveBase85Password(masterKey, 12, 0) if err != nil { t.Fatalf("Failed to derive Base85 password: %v", err) } // Expected password from the test vector if pwd != pwdBase85ExpectedPassword { - t.Errorf("Password mismatch!\nExpected: '%s'\nGot: '%s'", pwdBase85ExpectedPassword, pwd) + t.Errorf( + "Password mismatch!\nExpected: '%s'\nGot: '%s'", + pwdBase85ExpectedPassword, + pwd, + ) } } +// runPasswordVectorTest checks a BIP85 password derivation test vector. +func runPasswordVectorTest( + t *testing.T, + title, path string, + pwdLen uint32, + derive func(*hdkeychain.ExtendedKey, uint32, uint32) (string, error), + expectedEntropy, expectedPassword string, +) { + t.Helper() + + logTestVector(t, title) + + masterKey := mustParseTestMasterKey(t) + + // Testing with the example from the BIP85 spec + t.Logf("Path: %s", path) + t.Logf("Parameters: Length=%d, Index=0", pwdLen) + + // First verify the entropy derivation + entropy, err := bip85.DeriveBIP85Entropy(masterKey, path) + if err != nil { + t.Fatalf("Failed to derive entropy: %v", err) + } + + // Expected entropy from BIP85 spec + derivedEntropyHex := hex.EncodeToString(entropy) + + if derivedEntropyHex != expectedEntropy { + t.Errorf( + "Entropy mismatch!\nExpected: %s\nGot: %s", + expectedEntropy, + derivedEntropyHex, + ) + } + + // Now test the password generation + pwd, err := derive(masterKey, pwdLen, 0) + if err != nil { + t.Fatalf("Failed to derive password: %v", err) + } + + // The test vector from the BIP85 specification + t.Logf("EXPECTED PASSWORD: %s", expectedPassword) + t.Logf("ACTUAL PASSWORD: %s", pwd) + + if pwd != expectedPassword { + t.Errorf("Expected password '%s', got '%s'", expectedPassword, pwd) + } else { + t.Logf("RESULT: PASS") + } + + t.Logf("Password length: %d characters", len(pwd)) +} + // TestPWDBase64 tests the Base64 password test vector func TestPWDBase64(t *testing.T) { - logTestVector(t, "PWD Base64") + t.Parallel() - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } - - // Testing with the example from the BIP85 spec - 21 characters - t.Logf("Path: %s", pwdBase64Path) - t.Logf("Parameters: Length=21, Index=0") - - // First verify the entropy derivation - entropy, err := DeriveBIP85Entropy(masterKey, pwdBase64Path) - if err != nil { - t.Fatalf("Failed to derive entropy: %v", err) - } - - // Expected entropy from BIP85 spec - derivedEntropyHex := hex.EncodeToString(entropy) - - if derivedEntropyHex != pwdBase64ExpectedEntropy { - t.Errorf("Entropy mismatch!\nExpected: %s\nGot: %s", pwdBase64ExpectedEntropy, derivedEntropyHex) - } - - // Now test the password generation - pwd, err := DeriveBase64Password(masterKey, 21, 0) - if err != nil { - t.Fatalf("Failed to derive Base64 password: %v", err) - } - - // The test vector from the BIP85 specification - t.Logf("EXPECTED PASSWORD: %s", pwdBase64ExpectedPassword) - t.Logf("ACTUAL PASSWORD: %s", pwd) - - if pwd != pwdBase64ExpectedPassword { - t.Errorf("Expected password '%s', got '%s'", pwdBase64ExpectedPassword, pwd) - } else { - t.Logf("RESULT: PASS ✓") - } - - t.Logf("Password length: %d characters", len(pwd)) + runPasswordVectorTest( + t, + "PWD Base64", + pwdBase64Path, + 21, + bip85.DeriveBase64Password, + pwdBase64ExpectedEntropy, + pwdBase64ExpectedPassword, + ) } // TestPWDBase85 tests the Base85 password test vector func TestPWDBase85(t *testing.T) { - logTestVector(t, "PWD Base85") + t.Parallel() - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } - - // Testing with the example from the BIP85 spec - 12 characters - t.Logf("Path: %s", pwdBase85Path) - t.Logf("Parameters: Length=12, Index=0") - - // First verify the entropy derivation - entropy, err := DeriveBIP85Entropy(masterKey, pwdBase85Path) - if err != nil { - t.Fatalf("Failed to derive entropy: %v", err) - } - - // Expected entropy from BIP85 spec - derivedEntropyHex := hex.EncodeToString(entropy) - - if derivedEntropyHex != pwdBase85ExpectedEntropy { - t.Errorf("Entropy mismatch!\nExpected: %s\nGot: %s", pwdBase85ExpectedEntropy, derivedEntropyHex) - } - - // Now test the password generation - pwd, err := DeriveBase85Password(masterKey, 12, 0) - if err != nil { - t.Fatalf("Failed to derive Base85 password: %v", err) - } - - // The test vector from the BIP85 specification - t.Logf("EXPECTED PASSWORD: %s", pwdBase85ExpectedPassword) - t.Logf("ACTUAL PASSWORD: %s", pwd) - - if pwd != pwdBase85ExpectedPassword { - t.Errorf("Expected password '%s', got '%s'", pwdBase85ExpectedPassword, pwd) - } else { - t.Logf("RESULT: PASS ✓") - } - - t.Logf("Password length: %d characters", len(pwd)) + runPasswordVectorTest( + t, + "PWD Base85", + pwdBase85Path, + 12, + bip85.DeriveBase85Password, + pwdBase85ExpectedEntropy, + pwdBase85ExpectedPassword, + ) } // TestHexDerivation tests the HEX derivation test vector func TestHexDerivation(t *testing.T) { + t.Parallel() + logTestVector(t, "HEX Derivation") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // Test vector from BIP85 spec t.Logf("Path: %s", hexPath) t.Logf("Parameters: NumBytes=64, Index=0") // First verify the entropy derivation - entropy, err := DeriveBIP85Entropy(masterKey, hexPath) + entropy, err := bip85.DeriveBIP85Entropy(masterKey, hexPath) if err != nil { t.Fatalf("Failed to derive entropy: %v", err) } - // Expected entropy from BIP85 spec - derivedEntropyHex := hex.EncodeToString(entropy[:64]) // HEX uses first 64 bytes + // Expected entropy from BIP85 spec; HEX uses first 64 bytes + derivedEntropyHex := hex.EncodeToString(entropy[:64]) if derivedEntropyHex != hexExpectedEntropy { - t.Errorf("Entropy mismatch!\nExpected: %s\nGot: %s", hexExpectedEntropy, derivedEntropyHex) + t.Errorf( + "Entropy mismatch!\nExpected: %s\nGot: %s", + hexExpectedEntropy, + derivedEntropyHex, + ) } // Now test the hex derivation - hexData, err := DeriveHex(masterKey, 64, 0) + hexData, err := bip85.DeriveHex(masterKey, 64, 0) if err != nil { t.Fatalf("Failed to derive hex data: %v", err) } if hexData != hexExpectedEntropy { - t.Errorf("Hex data mismatch!\nExpected: %s\nGot: %s", hexExpectedEntropy, hexData) + t.Errorf( + "Hex data mismatch!\nExpected: %s\nGot: %s", + hexExpectedEntropy, + hexData, + ) } } // TestInvalidParameters tests error conditions for parameter validation func TestInvalidParameters(t *testing.T) { + t.Parallel() + logTestVector(t, "Invalid Parameters") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // Test cases for parameter validation testCases := []struct { @@ -982,49 +1025,63 @@ func TestInvalidParameters(t *testing.T) { { name: "BIP39 invalid word count", testFunc: func() error { - _, err := DeriveBIP39Entropy(masterKey, 0, 13, 0) // 13 is not valid (must be 12, 15, 18, 21, 24) + // 13 is not valid (must be 12, 15, 18, 21, 24) + _, err := bip85.DeriveBIP39Entropy(masterKey, 0, 13, 0) + return err }, }, { name: "Base64 password too short", testFunc: func() error { - _, err := DeriveBase64Password(masterKey, 19, 0) // Min is 20 + // Min is 20 + _, err := bip85.DeriveBase64Password(masterKey, 19, 0) + return err }, }, { name: "Base64 password too long", testFunc: func() error { - _, err := DeriveBase64Password(masterKey, 87, 0) // Max is 86 + // Max is 86 + _, err := bip85.DeriveBase64Password(masterKey, 87, 0) + return err }, }, { name: "Base85 password too short", testFunc: func() error { - _, err := DeriveBase85Password(masterKey, 9, 0) // Min is 10 + // Min is 10 + _, err := bip85.DeriveBase85Password(masterKey, 9, 0) + return err }, }, { name: "Base85 password too long", testFunc: func() error { - _, err := DeriveBase85Password(masterKey, 81, 0) // Max is 80 + // Max is 80 + _, err := bip85.DeriveBase85Password(masterKey, 81, 0) + return err }, }, { name: "Hex data too small", testFunc: func() error { - _, err := DeriveHex(masterKey, 15, 0) // Min is 16 + // Min is 16 + _, err := bip85.DeriveHex(masterKey, 15, 0) + return err }, }, { name: "Hex data too large", testFunc: func() error { - _, err := DeriveHex(masterKey, 65, 0) // Max is 64 + // Max is 64 + _, err := bip85.DeriveHex(masterKey, 65, 0) + return err }, }, @@ -1033,56 +1090,66 @@ func TestInvalidParameters(t *testing.T) { // Run all validation test cases for _, tc := range testCases { t.Logf("Testing: %s", tc.name) + err := tc.testFunc() if err == nil { t.Errorf("Expected error for %s, but got nil", tc.name) } else { t.Logf("Got expected error: %v", err) - t.Logf("RESULT: PASS ✓") + t.Logf("RESULT: PASS") } } } // TestAdditionalDeriveHex tests additional hex derivation scenarios func TestAdditionalDeriveHex(t *testing.T) { + t.Parallel() + logTestVector(t, "Additional Hex Derivation") - masterKey, err := ParseMasterKey(testMasterKey) - if err != nil { - t.Fatalf("Failed to parse master key: %v", err) - } + masterKey := mustParseTestMasterKey(t) // Test min size (16 bytes) - hexMinBytes, err := DeriveHex(masterKey, 16, 0) + hexMinBytes, err := bip85.DeriveHex(masterKey, 16, 0) if err != nil { t.Fatalf("Failed to derive 16-byte hex: %v", err) } + t.Logf("16-byte hex: %s", hexMinBytes) + if len(hexMinBytes) != 32 { // 16 bytes = 32 hex chars - t.Errorf("Expected 32 hex chars (16 bytes), got %d chars", len(hexMinBytes)) + t.Errorf( + "Expected 32 hex chars (16 bytes), got %d chars", + len(hexMinBytes), + ) } else { - t.Logf("RESULT: PASS ✓") + t.Logf("RESULT: PASS") } // Test max size (64 bytes) - hexMaxBytes, err := DeriveHex(masterKey, 64, 0) + hexMaxBytes, err := bip85.DeriveHex(masterKey, 64, 0) if err != nil { t.Fatalf("Failed to derive 64-byte hex: %v", err) } + t.Logf("64-byte hex: %s", hexMaxBytes) + if len(hexMaxBytes) != 128 { // 64 bytes = 128 hex chars - t.Errorf("Expected 128 hex chars (64 bytes), got %d chars", len(hexMaxBytes)) + t.Errorf( + "Expected 128 hex chars (64 bytes), got %d chars", + len(hexMaxBytes), + ) } else { - t.Logf("RESULT: PASS ✓") + t.Logf("RESULT: PASS") } // Test different index values - hex1, err := DeriveHex(masterKey, 32, 0) + hex1, err := bip85.DeriveHex(masterKey, 32, 0) if err != nil { t.Fatalf("Failed to derive hex with index 0: %v", err) } - hex2, err := DeriveHex(masterKey, 32, 1) + hex2, err := bip85.DeriveHex(masterKey, 32, 1) if err != nil { t.Fatalf("Failed to derive hex with index 1: %v", err) } @@ -1093,6 +1160,6 @@ func TestAdditionalDeriveHex(t *testing.T) { if hex1 == hex2 { t.Errorf("Expected different hex values for different indexes") } else { - t.Logf("Different indexes produced different outputs: PASS ✓") + t.Logf("Different indexes produced different outputs: PASS") } }