Files
weatherreporter/internal/config/secrets.go

109 lines
2.9 KiB
Go

package config
import (
"fmt"
"os"
"path/filepath"
"regexp"
"strings"
)
var secretNamePattern = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
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, nil
}
entries, err := os.ReadDir(cfg.Directory)
if err != nil {
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 nil, fmt.Errorf("secrets directory %q contains an empty filename", cfg.Directory)
}
if !secretNamePattern.MatchString(name) {
return nil, fmt.Errorf("secret file %q has invalid environment variable name", name)
}
if entry.Type()&os.ModeSymlink != 0 {
return nil, fmt.Errorf("secret file %q must be a regular file, not a symlink", name)
}
if entry.IsDir() {
return nil, fmt.Errorf("secret file %q must be a regular file, not a directory", name)
}
info, err := entry.Info()
if err != nil {
return nil, fmt.Errorf("inspect secret file %q: %w", name, err)
}
if !info.Mode().IsRegular() {
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 nil, fmt.Errorf("read secret file %q: %w", name, err)
}
value := string(data)
if strings.HasSuffix(value, "\r\n") {
value = strings.TrimSuffix(value, "\r\n")
} else {
value = strings.TrimSuffix(value, "\n")
}
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
}