diff --git a/docs/config.md b/docs/config.md index 1684825..e923da6 100644 --- a/docs/config.md +++ b/docs/config.md @@ -76,8 +76,10 @@ The prompt-facing location timezone is derived from the effective ### `secrets` `secrets.directory` defaults to empty, which disables secret loading. When it -is set, every regular file directly in that directory is loaded after the file -and command-line overrides. A file basename must match +is set, every regular file directly in that directory is staged after the file +and command-line overrides, then applied only after the complete configuration +has validated successfully. A rejected load leaves the existing environment +unchanged. A file basename must match `[A-Za-z_][A-Za-z0-9_]*`; it becomes an environment variable name, and the file contents replace any existing value. One trailing LF or CRLF is removed. diff --git a/internal/config/config_test.go b/internal/config/config_test.go index a4acfb6..b3de296 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -1706,9 +1706,13 @@ func TestLoadFileLoadsSecretsBeforeReturningNotifyConfig(t *testing.T) { func TestLoadSecretsDisabledLeavesEnvironmentUnchanged(t *testing.T) { t.Setenv("WEATHERREPORTER_DISABLED_SECRET", "original") - if err := loadSecrets(SecretsConfig{}); err != nil { + secrets, err := loadSecrets(SecretsConfig{}) + if err != nil { t.Fatalf("loadSecrets() error = %v", err) } + if len(secrets) != 0 { + t.Fatalf("staged secrets = %#v, want none", secrets) + } if got := os.Getenv("WEATHERREPORTER_DISABLED_SECRET"); got != "original" { t.Fatalf("environment value = %q, want original", got) } @@ -1744,9 +1748,13 @@ func TestLoadSecretsOverwritesExistingEnvironment(t *testing.T) { } t.Setenv("WEATHERREPORTER_SECRET", "existing") - if err := loadSecrets(SecretsConfig{Directory: dir}); err != nil { + secrets, err := loadSecrets(SecretsConfig{Directory: dir}) + if err != nil { t.Fatalf("loadSecrets() error = %v", err) } + if err := applySecrets(secrets); err != nil { + t.Fatalf("applySecrets() error = %v", err) + } if got := os.Getenv("WEATHERREPORTER_SECRET"); got != "from-file" { t.Fatalf("environment value = %q, want from-file", got) } @@ -1773,9 +1781,13 @@ func TestLoadSecretsTrimsOneTrailingLineEnding(t *testing.T) { } t.Setenv("WEATHERREPORTER_SECRET", "") - if err := loadSecrets(SecretsConfig{Directory: dir}); err != nil { + secrets, err := loadSecrets(SecretsConfig{Directory: dir}) + if err != nil { t.Fatalf("loadSecrets() error = %v", err) } + if err := applySecrets(secrets); err != nil { + t.Fatalf("applySecrets() error = %v", err) + } if got := os.Getenv("WEATHERREPORTER_SECRET"); got != tt.want { t.Fatalf("environment value = %q, want %q", got, tt.want) } @@ -1846,7 +1858,7 @@ func TestLoadSecretsRejectsInvalidDirectoryEntries(t *testing.T) { dir := t.TempDir() tt.setup(t, dir) - err := loadSecrets(SecretsConfig{Directory: dir}) + _, err := loadSecrets(SecretsConfig{Directory: dir}) if err == nil { t.Fatal("loadSecrets() error = nil, want error") } @@ -1861,7 +1873,7 @@ func TestLoadSecretsRejectsInvalidDirectoryEntries(t *testing.T) { } func TestLoadSecretsRejectsMissingDirectory(t *testing.T) { - err := loadSecrets(SecretsConfig{Directory: filepath.Join(t.TempDir(), "missing")}) + _, err := loadSecrets(SecretsConfig{Directory: filepath.Join(t.TempDir(), "missing")}) if err == nil { t.Fatal("loadSecrets() error = nil, want missing directory error") } @@ -1869,3 +1881,82 @@ func TestLoadSecretsRejectsMissingDirectory(t *testing.T) { t.Fatalf("error = %q, want read secrets directory context", err.Error()) } } + +func TestLoadFileSecretDirectoryFailureLeavesEnvironmentUnchanged(t *testing.T) { + dir := t.TempDir() + secretsDir := filepath.Join(dir, "secrets") + if err := os.Mkdir(secretsDir, 0o700); err != nil { + t.Fatalf("create secrets directory: %v", err) + } + if err := os.WriteFile(filepath.Join(secretsDir, "A_SECRET"), []byte("new-value"), 0o600); err != nil { + t.Fatalf("write secret: %v", err) + } + if err := os.WriteFile(filepath.Join(secretsDir, "Z-INVALID"), []byte("unused"), 0o600); err != nil { + t.Fatalf("write invalid secret: %v", err) + } + path := writeConfig(t, "secrets:\n directory: "+secretsDir+"\n") + t.Setenv("A_SECRET", "original-value") + + _, err := LoadFile(path) + if err == nil || !strings.Contains(err.Error(), "invalid environment variable name") { + t.Fatalf("LoadFile() error = %v, want invalid secret filename", err) + } + if got := os.Getenv("A_SECRET"); got != "original-value" { + t.Fatalf("environment value = %q, want original-value after rejected load", got) + } +} + +func TestLoadFileValidationFailureLeavesSecretEnvironmentUnset(t *testing.T) { + dir := t.TempDir() + secretsDir := filepath.Join(dir, "secrets") + if err := os.Mkdir(secretsDir, 0o700); err != nil { + t.Fatalf("create secrets directory: %v", err) + } + if err := os.WriteFile(filepath.Join(secretsDir, "WEATHERREPORTER_SECRET"), []byte("new-value"), 0o600); err != nil { + t.Fatalf("write secret: %v", err) + } + path := writeConfig(t, "secrets:\n directory: "+secretsDir+"\nmissing_source:\n default: invalid\n") + unsetEnvironment(t, "WEATHERREPORTER_SECRET") + + _, err := LoadFile(path) + if err == nil || !strings.Contains(err.Error(), "missing_source.default") { + t.Fatalf("LoadFile() error = %v, want configuration validation error", err) + } + if _, set := os.LookupEnv("WEATHERREPORTER_SECRET"); set { + t.Fatal("WEATHERREPORTER_SECRET was set by a rejected configuration") + } +} + +func TestApplySecretsRollsBackOnEnvironmentFailure(t *testing.T) { + t.Setenv("A_SECRET", "original-value") + unsetEnvironment(t, "Z_SECRET") + + err := applySecrets([]secretValue{ + {name: "A_SECRET", value: "new-value"}, + {name: "Z_SECRET", value: "invalid\x00value"}, + }) + if err == nil || !strings.Contains(err.Error(), `secret file "Z_SECRET"`) { + t.Fatalf("applySecrets() error = %v, want Z_SECRET context", err) + } + if got := os.Getenv("A_SECRET"); got != "original-value" { + t.Fatalf("A_SECRET = %q, want original-value after rollback", got) + } + if _, set := os.LookupEnv("Z_SECRET"); set { + t.Fatal("Z_SECRET was set after failed environment application") + } +} + +func unsetEnvironment(t *testing.T, name string) { + t.Helper() + value, set := os.LookupEnv(name) + if err := os.Unsetenv(name); err != nil { + t.Fatalf("unset environment variable %q: %v", name, err) + } + t.Cleanup(func() { + if set { + _ = os.Setenv(name, value) + return + } + _ = os.Unsetenv(name) + }) +} diff --git a/internal/config/load.go b/internal/config/load.go index 0986bd5..9d1266a 100644 --- a/internal/config/load.go +++ b/internal/config/load.go @@ -40,13 +40,17 @@ func Load(opts LoadOptions) (Config, error) { return Config{}, err } - if err := loadSecrets(cfg.Secrets); err != nil { + secrets, err := loadSecrets(cfg.Secrets) + if err != nil { return Config{}, err } if err := Validate(cfg); err != nil { return Config{}, err } + if err := applySecrets(secrets); err != nil { + return Config{}, err + } return cfg, nil } diff --git a/internal/config/secrets.go b/internal/config/secrets.go index 9509b67..ce18f4e 100644 --- a/internal/config/secrets.go +++ b/internal/config/secrets.go @@ -10,43 +10,54 @@ import ( var secretNamePattern = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) -func loadSecrets(cfg SecretsConfig) error { +type secretValue struct { + name string + value string +} + +type environmentValue struct { + value string + set bool +} + +func loadSecrets(cfg SecretsConfig) ([]secretValue, error) { if cfg.Directory == "" { - return nil + return nil, nil } entries, err := os.ReadDir(cfg.Directory) if err != nil { - return fmt.Errorf("read secrets directory %q: %w", cfg.Directory, err) + return nil, fmt.Errorf("read secrets directory %q: %w", cfg.Directory, err) } + secrets := make([]secretValue, 0, len(entries)) for _, entry := range entries { name := entry.Name() if name == "" { - return fmt.Errorf("secrets directory %q contains an empty filename", cfg.Directory) + return nil, fmt.Errorf("secrets directory %q contains an empty filename", cfg.Directory) } if !secretNamePattern.MatchString(name) { - return fmt.Errorf("secret file %q has invalid environment variable name", name) + return nil, fmt.Errorf("secret file %q has invalid environment variable name", name) } if entry.Type()&os.ModeSymlink != 0 { - return fmt.Errorf("secret file %q must be a regular file, not a symlink", name) + return nil, fmt.Errorf("secret file %q must be a regular file, not a symlink", name) } if entry.IsDir() { - return fmt.Errorf("secret file %q must be a regular file, not a directory", name) + return nil, fmt.Errorf("secret file %q must be a regular file, not a directory", name) } info, err := entry.Info() if err != nil { - return fmt.Errorf("inspect secret file %q: %w", name, err) + return nil, fmt.Errorf("inspect secret file %q: %w", name, err) } if !info.Mode().IsRegular() { - return fmt.Errorf("secret file %q must be a regular file", name) + return nil, fmt.Errorf("secret file %q must be a regular file", name) } path := filepath.Join(cfg.Directory, name) data, err := os.ReadFile(path) if err != nil { - return fmt.Errorf("read secret file %q: %w", name, err) + return nil, fmt.Errorf("read secret file %q: %w", name, err) } value := string(data) if strings.HasSuffix(value, "\r\n") { @@ -54,10 +65,44 @@ func loadSecrets(cfg SecretsConfig) error { } else { value = strings.TrimSuffix(value, "\n") } - if err := os.Setenv(name, value); err != nil { - return fmt.Errorf("set environment variable from secret file %q: %w", name, err) - } + secrets = append(secrets, secretValue{name: name, value: value}) } + return secrets, nil +} + +func applySecrets(secrets []secretValue) error { + previous := make(map[string]environmentValue, len(secrets)) + applied := make([]string, 0, len(secrets)) + for _, secret := range secrets { + if _, ok := previous[secret.name]; !ok { + value, set := os.LookupEnv(secret.name) + previous[secret.name] = environmentValue{value: value, set: set} + } + if err := os.Setenv(secret.name, secret.value); err != nil { + if rollbackErr := restoreEnvironment(previous, applied); rollbackErr != nil { + return fmt.Errorf("set environment variable from secret file %q: %w; restore environment: %v", secret.name, err, rollbackErr) + } + return fmt.Errorf("set environment variable from secret file %q: %w", secret.name, err) + } + applied = append(applied, secret.name) + } + return nil +} + +func restoreEnvironment(previous map[string]environmentValue, names []string) error { + for i := len(names) - 1; i >= 0; i-- { + name := names[i] + value := previous[name] + var err error + if value.set { + err = os.Setenv(name, value.value) + } else { + err = os.Unsetenv(name) + } + if err != nil { + return fmt.Errorf("restore environment variable %q: %w", name, err) + } + } return nil }