From 397011a5922fa3dde1a5d424adb7e61a130f2b8a Mon Sep 17 00:00:00 2001 From: sneak Date: Fri, 7 Aug 2026 17:27:23 +0000 Subject: [PATCH] Update golangci-lint to v2.12.2 with canonical config (closes #30) - Replace .golangci.yml with the canonical strict config (all linters enabled except the standard disable list; lll 88, funlen 80/50, cyclop 15, dupl 100; test files now linted) - Pin the Dockerfile lint stage to golangci/golangci-lint:v2.12.2 by tag and digest (Debian-based) - Fix all ~1550 findings surfaced by the new config: line wrapping, wsl_v5/nlreturn blank lines, noinlineerr splits, err113 sentinel errors, perfsprint/modernize rewrites, goconst constants, thelper, testifylint, noctx CommandContext, testpackage conversions, t.Parallel() where safe, and complexity/dupl helper extraction - Record the change and follow-up items in TODO.md User-visible strings -------------------- No user-visible string changes remain. Every error message this branch composes is byte-identical to the one main composes. The err113 sentinels are shaped so that fmt.Errorf reassembles the original text around them: a sentinel carries the fixed words of the message and the caller supplies the interpolated value in the position it has always occupied. Where the value sits in the middle of the sentence the sentinel therefore holds only a fragment (for example vault.ErrVaultNotFound is "does not exist", composed by its caller as "vault does not exist"); each such sentinel documents the message it participates in. Verified mechanically rather than by inspection: every fmt.Errorf and errors.New call site in both trees was parsed, the Error() text of any sentinel passed to %w substituted in, and the resulting sets of composed message templates compared. All 350 templates main produces are still produced, character for character; the set of messages lost or altered is empty. unlocker list ------------- findUnlockerIDByMetadata now returns (string, error) instead of signalling failure with an empty ID. An unreadable unlockers.d is no longer indistinguishable from "no matching entry", so UnlockersList skips the entry with a warning naming the directory, as it did before the scan was extracted into a helper, rather than emitting a row under a synthesized fallback ID that no unlocker remove or unlocker select can match and that suppresses the current-unlocker marker. The duplicate-check and shell-completion callers skip on the same condition, matching their pre-extraction behavior. Covered by tests in internal/cli/unlockers_list_test.go. --- .golangci.yml | 144 +-- Dockerfile | 4 +- TODO.md | 25 + internal/cli/cli.go | 9 +- internal/cli/cli_test.go | 24 +- internal/cli/completion.go | 6 +- internal/cli/completions.go | 197 ++-- internal/cli/crypto.go | 197 ++-- internal/cli/generate.go | 50 +- internal/cli/info.go | 28 +- internal/cli/info_helper.go | 158 +-- internal/cli/init.go | 161 +-- internal/cli/integration_test.go | 786 +++++++++------ internal/cli/root.go | 8 +- internal/cli/secrets.go | 537 +++++----- internal/cli/secrets_size_test.go | 364 +++---- internal/cli/stdout_stderr_test.go | 38 +- internal/cli/test_helpers.go | 10 +- internal/cli/test_output_test.go | 8 +- internal/cli/unlockers.go | 750 ++++++++------ internal/cli/unlockers_list_test.go | 229 +++++ internal/cli/vault.go | 362 ++++--- internal/cli/version.go | 156 ++- internal/cli/version_test.go | 90 +- internal/macse/macse_stub.go | 5 +- internal/secret/constants.go | 3 +- internal/secret/crypto.go | 98 +- internal/secret/debug.go | 42 +- internal/secret/debug_test.go | 7 +- internal/secret/helpers.go | 10 +- internal/secret/helpers_test.go | 29 +- internal/secret/keychainunlocker_stub.go | 36 +- internal/secret/passphrase_test.go | 239 +++-- internal/secret/passphraseunlocker.go | 85 +- internal/secret/pgpunlocker.go | 184 +++- internal/secret/secret.go | 263 +++-- internal/secret/secret_test.go | 208 ++-- internal/secret/seunlocker_stub.go | 45 +- internal/secret/seunlocker_stub_test.go | 36 +- internal/secret/validation_test.go | 157 +-- internal/secret/version.go | 307 ++++-- internal/secret/version_test.go | 136 ++- internal/vault/errors.go | 65 ++ internal/vault/integration_test.go | 779 +++++++------- internal/vault/integration_version_test.go | 451 +++++---- internal/vault/management.go | 93 +- internal/vault/metadata.go | 18 +- internal/vault/metadata_test.go | 470 +++++---- internal/vault/path_traversal_test.go | 43 +- internal/vault/secrets.go | 579 +++++++---- internal/vault/secrets_name_test.go | 5 + internal/vault/secrets_version_test.go | 160 +-- internal/vault/unlockers.go | 250 +++-- internal/vault/vault.go | 241 +++-- internal/vault/vault_error_test.go | 44 +- internal/vault/vault_test.go | 466 +++++---- pkg/agehd/agehd.go | 11 +- pkg/agehd/agehd_test.go | 638 ++++++------ pkg/bip85/bip85.go | 133 ++- pkg/bip85/bip85_test.go | 1065 +++++++++++--------- 60 files changed, 6867 insertions(+), 4875 deletions(-) create mode 100644 internal/cli/unlockers_list_test.go create mode 100644 internal/vault/errors.go 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..a0eae99 100644 --- a/TODO.md +++ b/TODO.md @@ -25,6 +25,20 @@ 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`. Reworked after review: + the `err113` sentinels in `internal/vault`, `internal/secret`, + `internal/cli` and `pkg/bip85` were reshaped so every composed + error message is byte-identical to `main`, and + `findUnlockerIDByMetadata` now returns an error so `unlocker list` + skips an unreadable `unlockers.d` entry with a warning instead of + emitting a fabricated fallback ID. - 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 +64,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..576fa2f 100644 --- a/internal/cli/completions.go +++ b/internal/cli/completions.go @@ -1,7 +1,6 @@ package cli import ( - "encoding/json" "path/filepath" "strings" @@ -11,11 +10,14 @@ import ( "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 +32,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 +43,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 +72,24 @@ 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) + id, err := findUnlockerIDByMetadata( + fs, unlockersDir, metadata, false, + ) if err != nil { - secret.Warn("Could not read unlockers directory during completion", "error", err) + secret.Warn( + "Could not read unlockers directory during completion, "+ + "skipping unlocker", + "unlockers_dir", unlockersDir, "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 - } + if id != "" && strings.HasPrefix(id, toComplete) { + completions = append(completions, id) } } @@ -128,17 +97,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 +122,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..593f462 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,37 @@ 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 carries only the message tail; the caller + // composes "GPG key is already added as an unlocker". + errGPGKeyAlreadyUnlocker = errors.New( + "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 +71,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 +88,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 +107,7 @@ func getDefaultGPGKey() (string, error) { } } - return "", fmt.Errorf("no GPG secret keys found") + return "", errNoGPGSecretKeys } func newUnlockerCmd() *cobra.Command { @@ -91,7 +127,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 +137,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 +149,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 +230,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 +261,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 +284,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 +312,92 @@ 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. It returns ("", nil) when the directory is +// readable but holds no match, and a non-nil error when the directory +// itself cannot be read. Callers must distinguish the two: an unreadable +// directory means the unlocker's real ID is unknowable, so the entry has +// to be skipped rather than reported under a synthesized ID. +func findUnlockerIDByMetadata( + fs afero.Fs, unlockersDir string, metadata secret.UnlockerMetadata, + includeSecureEnclave bool, +) (string, error) { + files, err := afero.ReadDir(fs, unlockersDir) + if err != nil { + return "", fmt.Errorf( + "failed to read unlockers directory %s: %w", unlockersDir, err, + ) + } + + 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), nil + } + } + + return "", nil +} + // UnlockersList lists unlockers in the current vault func (cli *Instance) UnlockersList(jsonOutput bool) error { // Get current vault @@ -259,6 +408,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 +422,40 @@ 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) + + unlockerID, err := findUnlockerIDByMetadata( + cli.fs, unlockersDir, metadata, true, + ) if err != nil { - secret.Warn("Could not read unlockers directory", "error", err) + secret.Warn("Could not read unlockers directory, skipping unlocker", + "unlockers_dir", unlockersDir, "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 - } - } - // 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 +476,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 +513,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 +534,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("GPG key %s %w", gpgKeyID, errGPGKeyAlreadyUnlocker) + } + + 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 +735,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 +785,29 @@ 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) + // Construct the unlocker matching this metadata to get its ID + id, err := findUnlockerIDByMetadata(cli.fs, unlockersDir, metadata, true) if err != nil { - secret.Warn("Could not read unlockers directory during duplicate check", "error", err) + secret.Warn( + "Could not read unlockers directory during duplicate check, "+ + "skipping unlocker", + "unlockers_dir", unlockersDir, "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 - } + if id != "" && id == unlockerID { + return errUnlockerExists } } diff --git a/internal/cli/unlockers_list_test.go b/internal/cli/unlockers_list_test.go new file mode 100644 index 0000000..bf2c564 --- /dev/null +++ b/internal/cli/unlockers_list_test.go @@ -0,0 +1,229 @@ +// Unlocker List Tests +// +// Tests for `secret unlocker list` behavior when the unlockers.d directory +// cannot be read while the listing is being rendered: +// +// - TestUnlockersListSkipsUnreadableUnlockersDir: an unreadable +// unlockers.d yields no rows rather than rows bearing synthesized IDs. +// - TestUnlockersListSkipsOnlyUnreadableEntries: a readable entry is +// still listed, with its real ID and its current-unlocker marker, +// when a later entry's scan fails. +// +// The listing resolves each unlocker's real ID by rescanning unlockers.d +// after the vault has already enumerated it. If that rescan fails the ID +// is unknowable, so the entry must be skipped: a synthesized ID matches +// no `unlocker remove` or `unlocker select` argument and would also +// suppress the current-unlocker marker. + +//nolint:testpackage // white-box test of unexported internals +package cli + +import ( + "bytes" + "encoding/json" + "errors" + "path/filepath" + "testing" + "time" + + "git.eeqj.de/sneak/secret/internal/secret" + "github.com/spf13/afero" + "github.com/spf13/cobra" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const ( + // listTestStateDir is the state directory of the synthetic vault used + // by the unlocker listing tests. + listTestStateDir = "/state" + + // listTestVaultName is the name of that synthetic vault. + listTestVaultName = "default" + + // listTestGPGKeyID is the GPG key ID recorded in the readable PGP + // unlocker's metadata. The unlocker's real ID is derived from it, and + // differs from the timestamp-derived fallback ID. + listTestGPGKeyID = "DEADBEEFDEADBEEF" + + // listTestUnlockerDirOne and listTestUnlockerDirTwo are the unlocker + // directory names under unlockers.d. + listTestUnlockerDirOne = "host-pgp-2026-08-09" + listTestUnlockerDirTwo = "host-pgp-2026-08-10" + + // listTestUnlockersDirName is the directory the listing rescans to + // resolve unlocker IDs. + listTestUnlockersDirName = "unlockers.d" + + // listTestMetadataFileName is the per-unlocker metadata file name. + listTestMetadataFileName = "unlocker-metadata.json" + + // listTestDirPerm and listTestFilePerm are the fixture permissions. + listTestDirPerm = 0o700 + listTestFilePerm = 0o600 +) + +// errUnlockersDirUnreadable is returned by the test filesystem in place of +// a successful open of unlockers.d. +var errUnlockersDirUnreadable = errors.New("permission denied") + +// unlockersDirFailFs makes unlockers.d unreadable once it has been opened +// successfully openBudget times. This reproduces the directory becoming +// unreadable (permission change, partially restored backup, EIO) between +// the vault's own enumeration and the per-entry rescan that resolves +// unlocker IDs. +type unlockersDirFailFs struct { + afero.Fs + + openBudget int + opens int +} + +//nolint:ireturn // afero.File is the interface required by afero.Fs +func (f *unlockersDirFailFs) Open(name string) (afero.File, error) { + if filepath.Base(name) == listTestUnlockersDirName { + f.opens++ + if f.opens > f.openBudget { + return nil, errUnlockersDirUnreadable + } + } + + //nolint:wrapcheck // test double must return the wrapped Fs error as-is + return f.Fs.Open(name) +} + +// writePGPUnlocker writes a PGP unlocker directory with metadata that +// yields the real ID "pgp-". +func writePGPUnlocker( + t *testing.T, fs afero.Fs, unlockersDir, dirName string, + createdAt time.Time, keyID string, +) { + t.Helper() + + metadata := secret.PGPUnlockerMetadata{ + UnlockerMetadata: secret.UnlockerMetadata{ + Type: unlockerTypePGP, + CreatedAt: createdAt, + }, + GPGKeyID: keyID, + } + + encoded, err := json.Marshal(metadata) + require.NoError(t, err) + + dir := filepath.Join(unlockersDir, dirName) + require.NoError(t, fs.MkdirAll(dir, listTestDirPerm)) + require.NoError(t, afero.WriteFile( + fs, filepath.Join(dir, listTestMetadataFileName), encoded, + listTestFilePerm, + )) +} + +// newListTestVault builds a synthetic vault on a MemMapFs containing the +// given number of PGP unlockers, with the first one selected as current. +func newListTestVault(t *testing.T, unlockerCount int) *afero.MemMapFs { + t.Helper() + + base := &afero.MemMapFs{} + vaultDir := filepath.Join(listTestStateDir, "vaults.d", listTestVaultName) + unlockersDir := filepath.Join(vaultDir, listTestUnlockersDirName) + + require.NoError(t, afero.WriteFile( + base, filepath.Join(listTestStateDir, "currentvault"), + []byte(listTestVaultName), listTestFilePerm, + )) + + names := []string{listTestUnlockerDirOne, listTestUnlockerDirTwo} + names = names[:unlockerCount] + + for i, name := range names { + writePGPUnlocker(t, base, unlockersDir, name, + time.Date(2026, time.August, 9+i, 12, 30, 0, 0, time.UTC), + listTestGPGKeyID+string(rune('A'+i)), + ) + } + + require.NoError(t, afero.WriteFile( + base, filepath.Join(vaultDir, "current-unlocker"), + []byte(names[0]), listTestFilePerm, + )) + + return base +} + +// listUnlockersJSON runs UnlockersList in JSON mode against the given +// filesystem and decodes the emitted unlocker rows. +func listUnlockersJSON(t *testing.T, fs afero.Fs) []UnlockerInfo { + t.Helper() + + var buf bytes.Buffer + + cmd := &cobra.Command{} + cmd.SetOut(&buf) + cmd.SetErr(&buf) + + instance := &Instance{fs: fs, stateDir: listTestStateDir, cmd: cmd} + require.NoError(t, instance.UnlockersList(true)) + + var decoded struct { + Unlockers []UnlockerInfo `json:"unlockers"` + } + + require.NoError(t, json.Unmarshal(buf.Bytes(), &decoded)) + + return decoded.Unlockers +} + +// TestUnlockersListSkipsUnreadableUnlockersDir asserts that an unlockers.d +// which becomes unreadable after the vault enumerated it produces no rows, +// rather than rows carrying fabricated fallback IDs. +func TestUnlockersListSkipsUnreadableUnlockersDir(t *testing.T) { + t.Parallel() + + base := newListTestVault(t, 1) + // Budget of one: the vault's own ListUnlockers scan succeeds, the + // per-entry rescan that resolves the ID fails. + fs := &unlockersDirFailFs{Fs: base, openBudget: 1} + + unlockers := listUnlockersJSON(t, fs) + + assert.Empty(t, unlockers, + "an unreadable unlockers.d must yield no rows, not fabricated IDs") +} + +// TestUnlockersListSkipsOnlyUnreadableEntries asserts that a readable +// entry survives with its real ID and current-unlocker marker when a later +// entry's rescan fails. +func TestUnlockersListSkipsOnlyUnreadableEntries(t *testing.T) { + t.Parallel() + + base := newListTestVault(t, 2) + // Budget of two: ListUnlockers plus the first entry's rescan succeed, + // the second entry's rescan fails. + fs := &unlockersDirFailFs{Fs: base, openBudget: 2} + + unlockers := listUnlockersJSON(t, fs) + + require.Len(t, unlockers, 1, + "only the entry whose directory was readable may be listed") + assert.Equal(t, "pgp-"+listTestGPGKeyID+"A", unlockers[0].ID, + "the surviving row must carry the real unlocker ID") + assert.True(t, unlockers[0].IsCurrent, + "the current-unlocker marker must survive the skip") +} + +// TestUnlockersListReadableEntriesAreListed is the control case: with a +// fully readable unlockers.d every entry is listed with its real ID. +func TestUnlockersListReadableEntriesAreListed(t *testing.T) { + t.Parallel() + + base := newListTestVault(t, 2) + + unlockers := listUnlockersJSON(t, base) + + require.Len(t, unlockers, 2) + assert.Equal(t, "pgp-"+listTestGPGKeyID+"A", unlockers[0].ID) + assert.Equal(t, "pgp-"+listTestGPGKeyID+"B", unlockers[1].ID) + assert.True(t, unlockers[0].IsCurrent) + assert.False(t, unlockers[1].IsCurrent) +} 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..b6aa789 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,18 @@ import ( "github.com/spf13/afero" ) +var ( + // errSecretNotFound carries only the message tail; callers compose + // "secret not found" around it so the emitted text is + // unchanged. + errSecretNotFound = errors.New("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 +35,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 +76,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 +87,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("secret %s %w", s.Name, errSecretNotFound) } Debug("Secret exists, getting current version", "secret_name", s.Name) @@ -95,52 +113,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 +124,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 +143,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 +164,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 +180,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 +208,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 +347,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..7d51e2e --- /dev/null +++ b/internal/vault/errors.go @@ -0,0 +1,65 @@ +package vault + +import "errors" + +// Sentinel errors returned by vault operations. +// +// Several of these carry deliberately partial text: the message a caller +// composes with fmt.Errorf places the interpolated value where it has +// always appeared, and the sentinel supplies only the surrounding fixed +// words. This keeps every composed message byte-identical to the dynamic +// errors these sentinels replaced. Each such sentinel notes the message it +// participates in. +var ( + // ErrMnemonicMismatch indicates the mnemonic-derived public key does + // not match the vault's stored public key hash. + ErrMnemonicMismatch = errors.New( + "derived public key does not match vault: mnemonic may be incorrect", + ) + + // ErrInvalidVaultName indicates a vault name that does not match the + // allowed pattern [a-z0-9.\-_]+. Composed as + // "invalid vault name '': must match pattern [a-z0-9.\-_]+". + ErrInvalidVaultName = errors.New("invalid vault name") + + // ErrVaultNotFound indicates the named vault does not exist. Composed + // as "vault does not exist". + ErrVaultNotFound = errors.New("does not exist") + + // ErrNilValueBuffer indicates a nil value buffer was supplied. + ErrNilValueBuffer = errors.New("value buffer is nil") + + // ErrInvalidSecretName indicates a secret name that does not match + // the allowed pattern [a-z0-9.\-_/]+. Composed as + // "invalid secret name '': must match pattern [a-z0-9.\-_/]+", + // or as "invalid secret name: " by GetSecretObject. + ErrInvalidSecretName = errors.New("invalid secret name") + + // ErrSecretExists indicates the secret already exists and --force + // was not supplied. Composed as + // "secret already exists (use --force to overwrite)", or as + // "secret '' already exists in vault '' (use --force to + // overwrite)" when copying between vaults. + ErrSecretExists = errors.New("already exists") + + // ErrSecretNotFound indicates the named secret does not exist. + // Composed as "secret not found". + ErrSecretNotFound = errors.New("not found") + + // ErrVersionNotFound indicates the requested secret version does not + // exist. Composed as + // "version not found for secret ". + ErrVersionNotFound = errors.New("not found for secret") + + // ErrNoVersions indicates the source secret has no versions. Composed + // as "source secret '' has no versions". + ErrNoVersions = errors.New("has no versions") + + // ErrUnsupportedUnlockerType indicates an unlocker metadata type + // that this build does not support. + ErrUnsupportedUnlockerType = errors.New("unsupported unlocker type") + + // ErrUnlockerNotFound indicates no unlocker with the given ID exists. + // Composed as "unlocker with ID not found". + ErrUnlockerNotFound = errors.New("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..f112da8 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,12 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) { if !isValidVaultName(name) { secret.Debug("Invalid vault name provided", "vault_name", name) - return nil, fmt.Errorf("invalid vault name '%s': must match pattern [a-z0-9.\\-_]+", name) + return nil, fmt.Errorf( + "%w '%s': must match pattern [a-z0-9.\\-_]+", + ErrInvalidVaultName, name, + ) } + secret.Debug("Vault name validation passed", "vault_name", name) // Create vault directory structure @@ -189,24 +212,30 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) { secret.Debug("Creating vault directory structure", "vault_dir", vaultDir) // Create main vault directory - if err := fs.MkdirAll(vaultDir, secret.DirPerms); err != nil { + err := fs.MkdirAll(vaultDir, secret.DirPerms) + if err != nil { return nil, fmt.Errorf("failed to create vault directory: %w", err) } // Create secrets directory secretsDir := filepath.Join(vaultDir, "secrets.d") - if err := fs.MkdirAll(secretsDir, secret.DirPerms); err != nil { + + err = fs.MkdirAll(secretsDir, secret.DirPerms) + if err != nil { return nil, fmt.Errorf("failed to create secrets directory: %w", err) } // Create unlockers directory unlockersDir := filepath.Join(vaultDir, "unlockers.d") - if err := fs.MkdirAll(unlockersDir, secret.DirPerms); err != nil { + + err = fs.MkdirAll(unlockersDir, secret.DirPerms) + if err != nil { return nil, fmt.Errorf("failed to create unlockers directory: %w", err) } // Process mnemonic if available - derivationIndex, publicKeyHash, familyHash, err := processMnemonicForVault(fs, stateDir, vaultDir, name) + derivationIndex, publicKeyHash, familyHash, err := processMnemonicForVault( + fs, stateDir, vaultDir, name) if err != nil { return nil, err } @@ -218,13 +247,17 @@ func CreateVault(fs afero.Fs, stateDir string, name string) (*Vault, error) { PublicKeyHash: publicKeyHash, MnemonicFamilyHash: familyHash, } - if err := SaveVaultMetadata(fs, vaultDir, metadata); err != nil { + + err = SaveVaultMetadata(fs, vaultDir, metadata) + if err != nil { return nil, fmt.Errorf("failed to save vault metadata: %w", err) } // Select the newly created vault as current secret.Debug("Selecting newly created vault as current", "name", name) - if err := SelectVault(fs, stateDir, name); err != nil { + + err = SelectVault(fs, stateDir, name) + if err != nil { return nil, fmt.Errorf("failed to select vault: %w", err) } @@ -242,32 +275,42 @@ func SelectVault(fs afero.Fs, stateDir string, name string) error { if !isValidVaultName(name) { secret.Debug("Invalid vault name provided", "vault_name", name) - return fmt.Errorf("invalid vault name '%s': must match pattern [a-z0-9.\\-_]+", name) + return fmt.Errorf( + "%w '%s': must match pattern [a-z0-9.\\-_]+", + ErrInvalidVaultName, name, + ) } + secret.Debug("Vault name validation passed", "vault_name", name) // Check if vault exists vaultDir := filepath.Join(stateDir, "vaults.d", name) + exists, err := afero.DirExists(fs, vaultDir) if err != nil { return fmt.Errorf("failed to check if vault exists: %w", err) } + if !exists { - return fmt.Errorf("vault %s does not exist", name) + return fmt.Errorf("vault %s %w", name, ErrVaultNotFound) } // Create or update the currentvault file with just the vault name currentVaultPath := filepath.Join(stateDir, "currentvault") // Remove existing file if it exists - if _, err := fs.Stat(currentVaultPath); err == nil { + _, err = fs.Stat(currentVaultPath) + if err == nil { secret.Debug("Removing existing currentvault file", "path", currentVaultPath) + _ = fs.Remove(currentVaultPath) } // Write just the vault name to the file secret.Debug("Writing currentvault file", "vault_name", name) - if err := afero.WriteFile(fs, currentVaultPath, []byte(name), secret.FilePerms); err != nil { + + err = afero.WriteFile(fs, currentVaultPath, []byte(name), secret.FilePerms) + if err != nil { return fmt.Errorf("failed to select vault: %w", err) } 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..b8ab559 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,24 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool) if !isValidSecretName(name) { secret.Debug("Invalid secret name provided", "secret_name", name) - return fmt.Errorf("invalid secret name '%s': must match pattern [a-z0-9.\\-_/]+", name) + return fmt.Errorf( + "%w '%s': must match pattern [a-z0-9.\\-_/]+", + ErrInvalidSecretName, name, + ) } + secret.Debug("Secret name validation passed", "secret_name", name) secret.Debug("Getting vault directory") + vaultDir, err := v.GetDirectory() if err != nil { - secret.Debug("Failed to get vault directory for secret addition", "error", err, "vault_name", v.Name) + secret.Debug("Failed to get vault directory for secret addition", + "error", err, "vault_name", v.Name) return err } + secret.Debug("Got vault directory", "vault_dir", vaultDir) // Convert slashes to percent signs for storage @@ -144,112 +156,30 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool) slog.String("secret_dir", secretDir), ) - // Check if secret already exists - secret.Debug("Checking if secret already exists", "secret_dir", secretDir) - exists, err := afero.DirExists(v.fs, secretDir) + // Check for an existing secret and prepare its directory + exists, previousVersion, err := v.prepareSecretDir(name, secretDir, force) if err != nil { - secret.Debug("Failed to check if secret exists", "error", err, "secret_dir", secretDir) - - return fmt.Errorf("failed to check if secret exists: %w", err) + return err } - secret.Debug("Secret existence check complete", "exists", exists) - // Handle existing secret case now := time.Now() - var previousVersion *secret.Version - if exists { - if !force { - secret.Debug("Secret already exists and force not specified", "secret_name", name, "secret_dir", secretDir) - - return fmt.Errorf("secret %s already exists (use --force to overwrite)", name) - } - - // Get the current version to update its notAfter timestamp - currentVersionName, err := secret.GetCurrentVersion(v.fs, secretDir) - if err == nil && currentVersionName != "" { - previousVersion = secret.NewVersion(v, name, currentVersionName) - // We'll need to load and update its metadata after we unlock the vault - } - } else { - // Create secret directory for new secret - secret.Debug("Creating secret directory", "secret_dir", secretDir) - if err := v.fs.MkdirAll(secretDir, secret.DirPerms); err != nil { - secret.Debug("Failed to create secret directory", "error", err, "secret_dir", secretDir) - - return fmt.Errorf("failed to create secret directory: %w", err) - } - secret.Debug("Created secret directory successfully") - } - - // Generate new version name - versionName, err := secret.GenerateVersionName(v.fs, secretDir) + // Create the new version and save the encrypted value + versionName, err := v.createAndSaveVersion( + name, secretDir, value, previousVersion, &now, exists) if err != nil { - secret.Debug("Failed to generate version name", "error", err, "secret_name", name) - - return fmt.Errorf("failed to generate version name: %w", err) + return err } - secret.Debug("Generated new version name", "version", versionName, "secret_name", name) - - // Create new version - newVersion := secret.NewVersion(v, name, versionName) - - // Set version timestamps - if previousVersion == nil { - // First version: notBefore = epoch + 1 second - epochPlusOne := time.Unix(1, 0) - newVersion.Metadata.NotBefore = &epochPlusOne - } else { - // New version: notBefore = now - newVersion.Metadata.NotBefore = &now - - // We'll update the previous version's notAfter after we save the new version - } - - // Save the new version - pass the LockedBuffer directly - if err := newVersion.Save(value); err != nil { - secret.Debug("Failed to save new version", "error", err, "version", versionName) - - // Clean up the secret directory if this was a new secret - if !exists { - secret.Debug("Cleaning up secret directory due to save failure", "secret_dir", secretDir) - _ = v.fs.RemoveAll(secretDir) - } - - return fmt.Errorf("failed to save version: %w", err) - } - - // Update previous version if it exists - if previousVersion != nil { - // Get long-term key to decrypt/encrypt metadata - ltIdentity, err := v.GetOrDeriveLongTermKey() - if err != nil { - secret.Debug("Failed to get long-term key for metadata update", "error", err) - - return fmt.Errorf("failed to get long-term key: %w", err) - } - - // Load previous version metadata - if err := previousVersion.LoadMetadata(ltIdentity); err != nil { - secret.Debug("Failed to load previous version metadata", "error", err) - - return fmt.Errorf("failed to load previous version metadata: %w", err) - } - - // Update notAfter timestamp - previousVersion.Metadata.NotAfter = &now - - // Re-save the metadata (we need to implement an update method) - if err := updateVersionMetadata(v.fs, previousVersion, ltIdentity); err != nil { - secret.Debug("Failed to update previous version metadata", "error", err) - - return fmt.Errorf("failed to update previous version metadata: %w", err) - } + // Update previous version's notAfter timestamp if it exists + err = v.updatePreviousVersion(previousVersion, &now) + if err != nil { + return err } // Set current symlink to new version - if err := secret.SetCurrentVersion(v.fs, secretDir, versionName); err != nil { + err = secret.SetCurrentVersion(v.fs, secretDir, versionName) + if err != nil { secret.Debug("Failed to set current version", "error", err, "version", versionName) return fmt.Errorf("failed to set current version: %w", err) @@ -263,9 +193,12 @@ func (v *Vault) AddSecret(name string, value *memguard.LockedBuffer, force bool) } // updateVersionMetadata updates the metadata of an existing version -func updateVersionMetadata(fs afero.Fs, version *secret.Version, ltIdentity *age.X25519Identity) error { +func updateVersionMetadata( + fs afero.Fs, version *secret.Version, ltIdentity *age.X25519Identity, +) error { // Read the version's encrypted private key encryptedPrivKeyPath := filepath.Join(version.Directory, "priv.age") + encryptedPrivKey, err := afero.ReadFile(fs, encryptedPrivKeyPath) if err != nil { return fmt.Errorf("failed to read encrypted version private key: %w", err) @@ -294,14 +227,17 @@ func updateVersionMetadata(fs afero.Fs, version *secret.Version, ltIdentity *age metadataBuffer := memguard.NewBufferFromBytes(metadataBytes) defer metadataBuffer.Destroy() - encryptedMetadata, err := secret.EncryptToRecipient(metadataBuffer, versionIdentity.Recipient()) + encryptedMetadata, err := secret.EncryptToRecipient(metadataBuffer, + versionIdentity.Recipient()) if err != nil { return fmt.Errorf("failed to encrypt version metadata: %w", err) } // Write encrypted metadata metadataPath := filepath.Join(version.Directory, "metadata.age") - if err := afero.WriteFile(fs, metadataPath, encryptedMetadata, secret.FilePerms); err != nil { + + err = afero.WriteFile(fs, metadataPath, encryptedMetadata, secret.FilePerms) + if err != nil { return fmt.Errorf("failed to write encrypted version metadata: %w", err) } @@ -318,7 +254,8 @@ func (v *Vault) GetSecret(name string) ([]byte, error) { return v.GetSecretVersion(name, "") } -// GetSecretVersion retrieves a specific version of a secret (empty version means current) +// GetSecretVersion retrieves a specific version of a secret (empty version +// means current) func (v *Vault) GetSecretVersion(name string, version string) ([]byte, error) { secret.DebugWith("Getting secret version from vault", slog.String("vault_name", v.Name), @@ -326,69 +263,17 @@ func (v *Vault) GetSecretVersion(name string, version string) ([]byte, error) { slog.String("version", version), ) - // Validate secret name to prevent path traversal - if !isValidSecretName(name) { - secret.Debug("Invalid secret name provided", "secret_name", name) - - return nil, fmt.Errorf("invalid secret name '%s': must match pattern [a-z0-9.\\-_/]+", name) - } - - // Get vault directory - vaultDir, err := v.GetDirectory() + // Validate the name and resolve the version to fetch + version, err := v.resolveSecretVersion(name, version) if err != nil { - secret.Debug("Failed to get vault directory", "error", err, "vault_name", v.Name) - return nil, err } - // Convert slashes to percent signs for storage - storageName := strings.ReplaceAll(name, "/", "%") - secretDir := filepath.Join(vaultDir, "secrets.d", storageName) - - // Check if secret exists - exists, err := afero.DirExists(v.fs, secretDir) - if err != nil { - secret.Debug("Failed to check if secret exists", "error", err, "secret_name", name) - - return nil, fmt.Errorf("failed to check if secret exists: %w", err) - } - if !exists { - secret.Debug("Secret not found in vault", "secret_name", name, "vault_name", v.Name) - - return nil, fmt.Errorf("secret %s not found", name) - } - - // Determine which version to get - if version == "" { - // Get current version - currentVersion, err := secret.GetCurrentVersion(v.fs, secretDir) - if err != nil { - secret.Debug("Failed to get current version", "error", err, "secret_name", name) - - return nil, fmt.Errorf("failed to get current version: %w", err) - } - version = currentVersion - secret.Debug("Using current version", "version", version, "secret_name", name) - } - // Create version object secretVersion := secret.NewVersion(v, name, version) - // Check if version exists - versionPath := filepath.Join(secretDir, "versions", version) - exists, err = afero.DirExists(v.fs, versionPath) - if err != nil { - secret.Debug("Failed to check if version exists", "error", err, "version", version) - - return nil, fmt.Errorf("failed to check if version exists: %w", err) - } - if !exists { - secret.Debug("Version not found", "version", version, "secret_name", name) - - return nil, fmt.Errorf("version %s not found for secret %s", version, name) - } - - secret.Debug("Version exists, proceeding with vault unlock and decryption", "version", version, "secret_name", name) + secret.Debug("Version exists, proceeding with vault unlock and decryption", + "version", version, "secret_name", name) // Unlock the vault (get long-term key in memory) longTermIdentity, err := v.UnlockVault() @@ -406,10 +291,13 @@ func (v *Vault) GetSecretVersion(name string, version string) ([]byte, error) { ) // Get the version's value - secret.Debug("About to call secretVersion.GetValue", "version", version, "secret_name", name) + secret.Debug("About to call secretVersion.GetValue", + "version", version, "secret_name", name) + decryptedValue, err := secretVersion.GetValue(longTermIdentity) if err != nil { - secret.Debug("Failed to decrypt version value", "error", err, "version", version, "secret_name", name) + secret.Debug("Failed to decrypt version value", + "error", err, "version", version, "secret_name", name) return nil, fmt.Errorf("failed to decrypt version: %w", err) } @@ -442,7 +330,8 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) { // If vault is already unlocked, return the cached key if !v.Locked() { - secret.Debug("Vault already unlocked, returning cached long-term key", "vault_name", v.Name) + secret.Debug("Vault already unlocked, returning cached long-term key", + "vault_name", v.Name) return v.longTermKey, nil } @@ -450,7 +339,8 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) { // Get or derive the long-term key (but don't store it yet) longTermIdentity, err := v.GetOrDeriveLongTermKey() if err != nil { - secret.Debug("Failed to get or derive long-term key", "error", err, "vault_name", v.Name) + secret.Debug("Failed to get or derive long-term key", + "error", err, "vault_name", v.Name) return nil, fmt.Errorf("failed to get long-term key: %w", err) } @@ -469,7 +359,7 @@ func (v *Vault) UnlockVault() (*age.X25519Identity, error) { // GetSecretObject retrieves a Secret object with metadata loaded from this vault func (v *Vault) GetSecretObject(name string) (*secret.Secret, error) { if !isValidSecretName(name) { - return nil, fmt.Errorf("invalid secret name: %s", name) + return nil, fmt.Errorf("%w: %s", ErrInvalidSecretName, name) } // First check if the secret exists by checking for the metadata file @@ -487,15 +377,17 @@ func (v *Vault) GetSecretObject(name string) (*secret.Secret, error) { if err != nil { return nil, fmt.Errorf("failed to check if secret exists: %w", err) } + if !exists { - return nil, fmt.Errorf("secret %s not found", name) + return nil, fmt.Errorf("secret %s %w", name, ErrSecretNotFound) } // Create a Secret object secretObj := secret.NewSecret(v, name) // Load the metadata from disk - if err := secretObj.LoadMetadata(); err != nil { + err = secretObj.LoadMetadata() + if err != nil { return nil, err } @@ -526,7 +418,8 @@ func (v *Vault) CopySecretVersion( defer valueBuffer.Destroy() // Load source metadata - if err := srcVersion.LoadMetadata(srcIdentity); err != nil { + err = srcVersion.LoadMetadata(srcIdentity) + if err != nil { return fmt.Errorf("failed to load source metadata: %w", err) } @@ -537,7 +430,8 @@ func (v *Vault) CopySecretVersion( destVersion.Metadata = srcVersion.Metadata // Save the version (encrypts to this vault's LT key) - if err := destVersion.Save(valueBuffer); err != nil { + err = destVersion.Save(valueBuffer) + if err != nil { return fmt.Errorf("failed to save destination version: %w", err) } @@ -571,26 +465,13 @@ func (v *Vault) CopySecretAllVersions( return fmt.Errorf("failed to get destination vault directory: %w", err) } - // Check if destination secret already exists + // Check if destination secret already exists and clear it if forced destStorageName := strings.ReplaceAll(destSecretName, "/", "%") destSecretDir := filepath.Join(destVaultDir, "secrets.d", destStorageName) - exists, err := afero.DirExists(v.fs, destSecretDir) + err = v.prepareCopyDestination(destSecretDir, destSecretName, force) if err != nil { - return fmt.Errorf("failed to check destination: %w", err) - } - - if exists && !force { - return fmt.Errorf("secret '%s' already exists in vault '%s' (use --force to overwrite)", - destSecretName, v.Name) - } - - if exists && force { - // Remove existing secret - secret.Debug("Removing existing destination secret", "path", destSecretDir) - if err := v.fs.RemoveAll(destSecretDir); err != nil { - return fmt.Errorf("failed to remove existing destination secret: %w", err) - } + return err } // Get source vault's long-term key @@ -615,7 +496,7 @@ func (v *Vault) CopySecretAllVersions( } if len(versions) == 0 { - return fmt.Errorf("source secret '%s' has no versions", srcSecretName) + return fmt.Errorf("source secret '%s' %w", srcSecretName, ErrNoVersions) } // Get current version name @@ -625,27 +506,16 @@ func (v *Vault) CopySecretAllVersions( } // Create destination secret directory - if err := v.fs.MkdirAll(destSecretDir, secret.DirPerms); err != nil { + err = v.fs.MkdirAll(destSecretDir, secret.DirPerms) + if err != nil { return fmt.Errorf("failed to create destination secret directory: %w", err) } - // Copy each version - for _, versionName := range versions { - srcVersion := secret.NewVersion(srcVault, srcSecretName, versionName) - if err := v.CopySecretVersion(srcVersion, srcIdentity, destSecretName, versionName); err != nil { - // Rollback: remove partial copy - secret.Debug("Rolling back partial copy due to error", "error", err) - _ = v.fs.RemoveAll(destSecretDir) - - return fmt.Errorf("failed to copy version %s: %w", versionName, err) - } - } - - // Set current version - if err := secret.SetCurrentVersion(v.fs, destSecretDir, currentVersion); err != nil { - _ = v.fs.RemoveAll(destSecretDir) - - return fmt.Errorf("failed to set current version: %w", err) + // Copy each version and set the current pointer, rolling back on error + err = v.copyVersionsWithRollback(srcVault, srcIdentity, + srcSecretName, destSecretName, destSecretDir, versions, currentVersion) + if err != nil { + return err } secret.DebugWith("Successfully copied all secret versions", @@ -656,3 +526,292 @@ func (v *Vault) CopySecretAllVersions( return nil } + +// prepareSecretDir checks for an existing secret directory and prepares it +// for a new version. It returns whether the secret already existed and the +// current version to be superseded, if any. +func (v *Vault) prepareSecretDir( + name, secretDir string, force bool, +) (bool, *secret.Version, error) { + // Check if secret already exists + secret.Debug("Checking if secret already exists", "secret_dir", secretDir) + + exists, err := afero.DirExists(v.fs, secretDir) + if err != nil { + secret.Debug("Failed to check if secret exists", + "error", err, "secret_dir", secretDir) + + return false, nil, fmt.Errorf("failed to check if secret exists: %w", err) + } + + secret.Debug("Secret existence check complete", "exists", exists) + + if !exists { + // Create secret directory for new secret + secret.Debug("Creating secret directory", "secret_dir", secretDir) + + err = v.fs.MkdirAll(secretDir, secret.DirPerms) + if err != nil { + secret.Debug("Failed to create secret directory", + "error", err, "secret_dir", secretDir) + + return false, nil, fmt.Errorf("failed to create secret directory: %w", err) + } + + secret.Debug("Created secret directory successfully") + + return false, nil, nil + } + + if !force { + secret.Debug("Secret already exists and force not specified", + "secret_name", name, "secret_dir", secretDir) + + return true, nil, fmt.Errorf( + "secret %s %w (use --force to overwrite)", + name, ErrSecretExists, + ) + } + + // Get the current version to update its notAfter timestamp + var previousVersion *secret.Version + + currentVersionName, err := secret.GetCurrentVersion(v.fs, secretDir) + if err == nil && currentVersionName != "" { + previousVersion = secret.NewVersion(v, name, currentVersionName) + // We'll need to load and update its metadata after we unlock the vault + } + + return true, previousVersion, nil +} + +// updatePreviousVersion sets the notAfter timestamp on the version being +// superseded. It is a no-op when previousVersion is nil. +func (v *Vault) updatePreviousVersion( + previousVersion *secret.Version, now *time.Time, +) error { + if previousVersion == nil { + return nil + } + + // Get long-term key to decrypt/encrypt metadata + ltIdentity, err := v.GetOrDeriveLongTermKey() + if err != nil { + secret.Debug("Failed to get long-term key for metadata update", "error", err) + + return fmt.Errorf("failed to get long-term key: %w", err) + } + + // Load previous version metadata + err = previousVersion.LoadMetadata(ltIdentity) + if err != nil { + secret.Debug("Failed to load previous version metadata", "error", err) + + return fmt.Errorf("failed to load previous version metadata: %w", err) + } + + // Update notAfter timestamp + previousVersion.Metadata.NotAfter = now + + // Re-save the metadata (we need to implement an update method) + err = updateVersionMetadata(v.fs, previousVersion, ltIdentity) + if err != nil { + secret.Debug("Failed to update previous version metadata", "error", err) + + return fmt.Errorf("failed to update previous version metadata: %w", err) + } + + return nil +} + +// resolveSecretVersion validates the secret name, verifies the secret and +// version exist, and resolves an empty version to the current one. +func (v *Vault) resolveSecretVersion(name, version string) (string, error) { + // Validate secret name to prevent path traversal + if !isValidSecretName(name) { + secret.Debug("Invalid secret name provided", "secret_name", name) + + return "", fmt.Errorf( + "%w '%s': must match pattern [a-z0-9.\\-_/]+", + ErrInvalidSecretName, name, + ) + } + + // Get vault directory + vaultDir, err := v.GetDirectory() + if err != nil { + secret.Debug("Failed to get vault directory", "error", err, "vault_name", v.Name) + + return "", err + } + + // Convert slashes to percent signs for storage + storageName := strings.ReplaceAll(name, "/", "%") + secretDir := filepath.Join(vaultDir, "secrets.d", storageName) + + // Check if secret exists + exists, err := afero.DirExists(v.fs, secretDir) + if err != nil { + secret.Debug("Failed to check if secret exists", "error", err, "secret_name", name) + + return "", fmt.Errorf("failed to check if secret exists: %w", err) + } + + if !exists { + secret.Debug("Secret not found in vault", "secret_name", name, "vault_name", v.Name) + + return "", fmt.Errorf("secret %s %w", name, ErrSecretNotFound) + } + + // Determine which version to get + if version == "" { + // Get current version + currentVersion, err := secret.GetCurrentVersion(v.fs, secretDir) + if err != nil { + secret.Debug("Failed to get current version", "error", err, "secret_name", name) + + return "", fmt.Errorf("failed to get current version: %w", err) + } + + version = currentVersion + + secret.Debug("Using current version", "version", version, "secret_name", name) + } + + // Check if version exists + versionPath := filepath.Join(secretDir, "versions", version) + + exists, err = afero.DirExists(v.fs, versionPath) + if err != nil { + secret.Debug("Failed to check if version exists", "error", err, "version", version) + + return "", fmt.Errorf("failed to check if version exists: %w", err) + } + + if !exists { + secret.Debug("Version not found", "version", version, "secret_name", name) + + return "", fmt.Errorf( + "version %s %w %s", + version, ErrVersionNotFound, name, + ) + } + + return version, nil +} + +// createAndSaveVersion generates a new version name, sets the version +// timestamps, and saves the encrypted value. When saving fails for a newly +// created secret, the secret directory is removed again. +func (v *Vault) createAndSaveVersion( + name, secretDir string, value *memguard.LockedBuffer, + previousVersion *secret.Version, now *time.Time, exists bool, +) (string, error) { + // Generate new version name + versionName, err := secret.GenerateVersionName(v.fs, secretDir) + if err != nil { + secret.Debug("Failed to generate version name", "error", err, "secret_name", name) + + return "", fmt.Errorf("failed to generate version name: %w", err) + } + + secret.Debug("Generated new version name", "version", versionName, "secret_name", name) + + // Create new version + newVersion := secret.NewVersion(v, name, versionName) + + // Set version timestamps + if previousVersion == nil { + // First version: notBefore = epoch + 1 second + epochPlusOne := time.Unix(1, 0) + newVersion.Metadata.NotBefore = &epochPlusOne + } else { + // New version: notBefore = now + newVersion.Metadata.NotBefore = now + + // We'll update the previous version's notAfter after we save the + // new version + } + + // Save the new version - pass the LockedBuffer directly + err = newVersion.Save(value) + if err != nil { + secret.Debug("Failed to save new version", "error", err, "version", versionName) + + // Clean up the secret directory if this was a new secret + if !exists { + secret.Debug("Cleaning up secret directory due to save failure", + "secret_dir", secretDir) + + _ = v.fs.RemoveAll(secretDir) + } + + return "", fmt.Errorf("failed to save version: %w", err) + } + + return versionName, nil +} + +// copyVersionsWithRollback copies each version of the source secret into the +// destination directory and sets the current version pointer, removing the +// partial copy when any step fails. +func (v *Vault) copyVersionsWithRollback( + srcVault *Vault, srcIdentity *age.X25519Identity, + srcSecretName, destSecretName, destSecretDir string, + versions []string, currentVersion string, +) error { + // Copy each version + for _, versionName := range versions { + srcVersion := secret.NewVersion(srcVault, srcSecretName, versionName) + + err := v.CopySecretVersion(srcVersion, srcIdentity, destSecretName, versionName) + if err != nil { + // Rollback: remove partial copy + secret.Debug("Rolling back partial copy due to error", "error", err) + + _ = v.fs.RemoveAll(destSecretDir) + + return fmt.Errorf("failed to copy version %s: %w", versionName, err) + } + } + + // Set current version + err := secret.SetCurrentVersion(v.fs, destSecretDir, currentVersion) + if err != nil { + _ = v.fs.RemoveAll(destSecretDir) + + return fmt.Errorf("failed to set current version: %w", err) + } + + return nil +} + +// prepareCopyDestination ensures the destination secret directory can be +// created, removing an existing secret when force is set. +func (v *Vault) prepareCopyDestination( + destSecretDir, destSecretName string, force bool, +) error { + exists, err := afero.DirExists(v.fs, destSecretDir) + if err != nil { + return fmt.Errorf("failed to check destination: %w", err) + } + + if exists && !force { + return fmt.Errorf( + "secret '%s' %w in vault '%s' (use --force to overwrite)", + destSecretName, ErrSecretExists, v.Name, + ) + } + + if exists && force { + // Remove existing secret + secret.Debug("Removing existing destination secret", "path", destSecretDir) + + err = v.fs.RemoveAll(destSecretDir) + if err != nil { + return fmt.Errorf("failed to remove existing destination secret: %w", err) + } + } + + return nil +} 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..e2c562b 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("unlocker with ID %s %w", unlockerID, ErrUnlockerNotFound) } // Use the unlocker's Remove method @@ -281,17 +307,21 @@ func (v *Vault) SelectUnlocker(unlockerID string) error { } if targetUnlockerDir == "" { - return fmt.Errorf("unlocker with ID %s not found", unlockerID) + return fmt.Errorf("unlocker with ID %s %w", unlockerID, ErrUnlockerNotFound) } // Create/update current-unlocker file with just the unlocker name currentUnlockerPath := filepath.Join(vaultDir, "current-unlocker") // Remove existing file if it exists - if exists, err := afero.Exists(v.fs, currentUnlockerPath); err != nil { + exists, err := afero.Exists(v.fs, currentUnlockerPath) + if err != nil { return fmt.Errorf("failed to check if current-unlocker file exists: %w", err) - } else if exists { - if err := v.fs.Remove(currentUnlockerPath); err != nil { + } + + if exists { + err = v.fs.Remove(currentUnlockerPath) + if err != nil { return fmt.Errorf("failed to remove existing current-unlocker file: %w", err) } } @@ -301,7 +331,10 @@ func (v *Vault) SelectUnlocker(unlockerID string) error { // Write just the unlocker name to the file secret.Debug("Writing current-unlocker file", "unlocker_name", unlockerName) - if err := afero.WriteFile(v.fs, currentUnlockerPath, []byte(unlockerName), secret.FilePerms); err != nil { + + err = afero.WriteFile(v.fs, currentUnlockerPath, []byte(unlockerName), + secret.FilePerms) + if err != nil { return fmt.Errorf("failed to create current-unlocker file: %w", err) } @@ -310,15 +343,19 @@ func (v *Vault) SelectUnlocker(unlockerID string) error { // CreatePassphraseUnlocker creates a new passphrase-protected unlocker // The passphrase must be provided as a LockedBuffer for security -func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*secret.PassphraseUnlocker, error) { +func (v *Vault) CreatePassphraseUnlocker( + passphrase *memguard.LockedBuffer, +) (*secret.PassphraseUnlocker, error) { vaultDir, err := v.GetDirectory() if err != nil { return nil, fmt.Errorf("failed to get vault directory: %w", err) } // Create unlocker directory - unlockerDir := filepath.Join(vaultDir, "unlockers.d", "passphrase") - if err := v.fs.MkdirAll(unlockerDir, secret.DirPerms); err != nil { + unlockerDir := filepath.Join(vaultDir, "unlockers.d", unlockerTypePassphrase) + + err = v.fs.MkdirAll(unlockerDir, secret.DirPerms) + if err != nil { return nil, fmt.Errorf("failed to create unlocker directory: %w", err) } @@ -328,32 +365,15 @@ func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*se return nil, fmt.Errorf("failed to generate unlocker: %w", err) } - // Write public key - pubKeyPath := filepath.Join(unlockerDir, "pub.age") - if err := afero.WriteFile(v.fs, pubKeyPath, - []byte(unlockerIdentity.Recipient().String()), - secret.FilePerms); err != nil { - return nil, fmt.Errorf("failed to write unlocker public key: %w", err) - } - - // Encrypt private key with passphrase - privKeyStr := unlockerIdentity.String() - privKeyBuffer := memguard.NewBufferFromBytes([]byte(privKeyStr)) - defer privKeyBuffer.Destroy() - encryptedPrivKey, err := secret.EncryptWithPassphrase(privKeyBuffer, passphrase) + // Write the unlocker keypair (public and passphrase-encrypted private) + err = v.writeUnlockerKeypair(unlockerDir, unlockerIdentity, passphrase) if err != nil { - return nil, fmt.Errorf("failed to encrypt unlocker private key: %w", err) - } - - // Write encrypted private key - privKeyPath := filepath.Join(unlockerDir, "priv.age") - if err := afero.WriteFile(v.fs, privKeyPath, encryptedPrivKey, secret.FilePerms); err != nil { - return nil, fmt.Errorf("failed to write encrypted unlocker private key: %w", err) + return nil, err } // Create metadata metadata := UnlockerMetadata{ - Type: "passphrase", + Type: unlockerTypePassphrase, CreatedAt: time.Now(), Flags: []string{}, } @@ -365,7 +385,9 @@ func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*se } metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") - if err := afero.WriteFile(v.fs, metadataPath, metadataBytes, secret.FilePerms); err != nil { + + err = afero.WriteFile(v.fs, metadataPath, metadataBytes, secret.FilePerms) + if err != nil { return nil, fmt.Errorf("failed to write unlocker metadata: %w", err) } @@ -379,13 +401,16 @@ func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*se ltPrivKeyBuffer := memguard.NewBufferFromBytes([]byte(ltIdentity.String())) defer ltPrivKeyBuffer.Destroy() - encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, unlockerIdentity.Recipient()) + encryptedLtPrivKey, err := secret.EncryptToRecipient(ltPrivKeyBuffer, + unlockerIdentity.Recipient()) if err != nil { return nil, fmt.Errorf("failed to encrypt long-term private key: %w", err) } ltPrivKeyPath := filepath.Join(unlockerDir, "longterm.age") - if err := afero.WriteFile(v.fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms); err != nil { + + err = afero.WriteFile(v.fs, ltPrivKeyPath, encryptedLtPrivKey, secret.FilePerms) + if err != nil { return nil, fmt.Errorf("failed to write encrypted long-term private key: %w", err) } @@ -393,9 +418,80 @@ func (v *Vault) CreatePassphraseUnlocker(passphrase *memguard.LockedBuffer) (*se unlocker := secret.NewPassphraseUnlocker(v.fs, unlockerDir, metadata) // Select this unlocker as current - if err := v.SelectUnlocker(unlocker.GetID()); err != nil { + err = v.SelectUnlocker(unlocker.GetID()) + if err != nil { return nil, fmt.Errorf("failed to select new unlocker: %w", err) } return unlocker, nil } + +// readUnlockerMetadata reads and parses the unlocker-metadata.json file in +// the given unlocker directory. +func (v *Vault) readUnlockerMetadata(unlockerDir string) (UnlockerMetadata, error) { + metadataPath := filepath.Join(unlockerDir, "unlocker-metadata.json") + secret.Debug("Reading unlocker metadata", "path", metadataPath) + + var metadata UnlockerMetadata + + metadataBytes, err := afero.ReadFile(v.fs, metadataPath) + if err != nil { + secret.Debug("Failed to read unlocker metadata", "error", err, "path", metadataPath) + + return metadata, fmt.Errorf("failed to read unlocker metadata: %w", err) + } + + err = json.Unmarshal(metadataBytes, &metadata) + if err != nil { + secret.Debug("Failed to parse unlocker metadata", "error", err, "path", metadataPath) + + return metadata, fmt.Errorf("failed to parse unlocker metadata: %w", err) + } + + secret.DebugWith("Parsed unlocker metadata", + slog.String("unlocker_type", metadata.Type), + slog.Time("created_at", metadata.CreatedAt), + slog.Any("flags", metadata.Flags), + ) + + return metadata, nil +} + +// writeUnlockerKeypair writes the unlocker's public key and its +// passphrase-encrypted private key into the unlocker directory. +func (v *Vault) writeUnlockerKeypair( + unlockerDir string, + unlockerIdentity *age.X25519Identity, + passphrase *memguard.LockedBuffer, +) error { + // Write public key + pubKeyPath := filepath.Join(unlockerDir, "pub.age") + + err := afero.WriteFile(v.fs, pubKeyPath, + []byte(unlockerIdentity.Recipient().String()), + secret.FilePerms) + if err != nil { + return fmt.Errorf("failed to write unlocker public key: %w", err) + } + + // Encrypt private key with passphrase + privKeyStr := unlockerIdentity.String() + + privKeyBuffer := memguard.NewBufferFromBytes([]byte(privKeyStr)) + defer privKeyBuffer.Destroy() + + encryptedPrivKey, err := secret.EncryptWithPassphrase(privKeyBuffer, passphrase) + if err != nil { + return fmt.Errorf("failed to encrypt unlocker private key: %w", err) + } + + // Write encrypted private key + privKeyPath := filepath.Join(unlockerDir, "priv.age") + + err = afero.WriteFile(v.fs, privKeyPath, encryptedPrivKey, secret.FilePerms) + if err != nil { + return fmt.Errorf("failed to write encrypted unlocker private key: %w", err) + } + + return nil +} 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..c5ff617 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,50 @@ 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. It carries only the + // middle of the message, which the caller composes as + // "derived password length is shorter than requested length ", + // so the emitted text is unchanged. + ErrPasswordTooShort = errors.New("is shorter than requested length") + // ErrEncodedTooShort is returned when the encoded material is shorter + // than the requested password length. Composed as + // "encoded length is less than requested length ". + ErrEncodedTooShort = errors.New("is less than requested length") ) // 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 +104,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 +112,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 +131,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 +152,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 +181,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 +207,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 +232,7 @@ func DeriveBIP39Entropy(masterKey *hdkeychain.ExtendedKey, language, words, inde ) var bits int + switch words { case words12: bits = 128 @@ -195,7 +245,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 +268,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 +278,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 +320,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 +338,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 +360,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 +383,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( + "derived password length %d %w %d", + len(encodedStr), ErrPasswordTooShort, 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 +413,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( + "encoded length %d %w %d", + len(encoded), ErrEncodedTooShort, 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 +440,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") } } -- 2.49.1