add support for config dir Env & unexport functions

This commit is contained in:
Sahilb315
2025-12-17 23:19:11 +05:30
parent 187d5909b3
commit 1655bf1921
7 changed files with 125 additions and 109 deletions
+81 -73
View File
@@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"os"
"strings"
"sync"
@@ -31,7 +32,12 @@ type Config struct {
InsecureInstallation bool `mapstructure:"insecure_installation"`
// TrustedPackages allows for trusting an suspicious package and ignoring the suspicious behaviour for the package in future installations
TrustedPackages map[string][]string `mapstructure:"trusted_packages"`
TrustedPackages TrustedPackage `mapstructure:"trusted_packages"`
}
type TrustedPackage struct {
// Purl of the trusted package. Eg. pkg:npm/express@5.2.1
Purl []string `mapstructure:"purls"`
}
var (
@@ -51,85 +57,17 @@ func DefaultConfig() Config {
Paranoid: false,
DryRun: false,
InsecureInstallation: false,
TrustedPackages: map[string][]string{},
TrustedPackages: TrustedPackage{Purl: []string{}},
}
}
func configAsMap(cfg Config) map[string]any {
return map[string]any{
"transitive": cfg.Transitive,
"transitive_depth": cfg.TransitiveDepth,
"include_dev_dependencies": cfg.IncludeDevDependencies,
"dry_run": cfg.DryRun,
"paranoid": cfg.Paranoid,
"insecure_installation": cfg.InsecureInstallation,
"trusted_packages": cfg.TrustedPackages,
}
}
func SetupViper() (string, error) {
if err := ensureViperConfigured(); err != nil {
return "", err
}
cfgPath, err := ConfigFilePath()
if err != nil {
return "", err
}
return cfgPath, nil
}
func ensureViperConfigured() error {
setupOnce.Do(func() {
dir, err := PmgConfigDir()
if err != nil {
setupErr = err
return
}
v := viper.GetViper()
v.SetConfigName(pmgConfigName)
v.SetConfigType(pmgConfigType)
v.AddConfigPath(dir)
v.SetEnvPrefix("PMG")
v.SetEnvKeyReplacer(strings.NewReplacer("-", "_"))
v.AutomaticEnv()
for key, value := range configAsMap(DefaultConfig()) {
v.SetDefault(key, value)
}
})
return setupErr
}
func BindFlags(fs *pflag.FlagSet) {
if fs == nil {
return
}
// Helper binds a flag if it exists
bind := func(key, flag string) {
if f := fs.Lookup(flag); f != nil {
_ = viper.BindPFlag(key, f)
}
}
bind("transitive", "transitive")
bind("transitive_depth", "transitive-depth")
bind("include_dev_dependencies", "include-dev-dependencies")
bind("dry_run", "dry-run")
bind("paranoid", "paranoid")
}
func Load(fs *pflag.FlagSet) (Config, error) {
if _, err := SetupViper(); err != nil {
if err := ensureViperConfigured(); err != nil {
return Config{}, err
}
// Bind CLI flags so they override config/env
BindFlags(fs)
bindFlags(fs)
// Read the config file if it exists
if err := viper.ReadInConfig(); err != nil {
@@ -148,7 +86,7 @@ func Load(fs *pflag.FlagSet) (Config, error) {
// CreateConfig writes the PMG config file and returns its absolute path.
func CreateConfig() (string, error) {
if _, err := CreatePmgConfigDir(); err != nil {
if _, err := createConfigDir(); err != nil {
return "", err
}
@@ -184,6 +122,19 @@ func CreateConfig() (string, error) {
return cfgFile, nil
}
// RemoveConfig removes the PMG configuration directory and its contents.
func RemoveConfig() error {
dir, err := ConfigDir()
if err != nil {
return err
}
if err := os.RemoveAll(dir); err != nil {
return fmt.Errorf("failed to remove config directory %s: %w", dir, err)
}
return nil
}
// Inject config into context while protecting against context poisoning
func (c Config) Inject(ctx context.Context) context.Context {
return context.WithValue(ctx, configKey{}, &contextValue{Config: c})
@@ -198,3 +149,60 @@ func FromContext(ctx context.Context) (Config, error) {
return c.Config, nil
}
func ensureViperConfigured() error {
setupOnce.Do(func() {
dir, err := ConfigDir()
if err != nil {
setupErr = err
return
}
v := viper.GetViper()
v.SetConfigName(pmgConfigName)
v.SetConfigType(pmgConfigType)
v.AddConfigPath(dir)
v.SetEnvPrefix("PMG")
v.SetEnvKeyReplacer(strings.NewReplacer("-", "_"))
v.AutomaticEnv()
for key, value := range configAsMap(DefaultConfig()) {
v.SetDefault(key, value)
}
})
return setupErr
}
func bindFlags(fs *pflag.FlagSet) {
if fs == nil {
return
}
// Helper binds a flag if it exists
bind := func(key, flag string) {
if f := fs.Lookup(flag); f != nil {
_ = viper.BindPFlag(key, f)
}
}
bind("transitive", "transitive")
bind("transitive_depth", "transitive-depth")
bind("include_dev_dependencies", "include-dev-dependencies")
bind("dry_run", "dry-run")
bind("paranoid", "paranoid")
}
// Helper function to map the provided config for setting key/values in viper
func configAsMap(cfg Config) map[string]any {
return map[string]any{
"transitive": cfg.Transitive,
"transitive_depth": cfg.TransitiveDepth,
"include_dev_dependencies": cfg.IncludeDevDependencies,
"dry_run": cfg.DryRun,
"paranoid": cfg.Paranoid,
"insecure_installation": cfg.InsecureInstallation,
"trusted_packages": cfg.TrustedPackages,
}
}
+1 -1
View File
@@ -59,7 +59,7 @@ func TestLoad_FlagsOverrideDefaults(t *testing.T) {
func TestLoad_ConfigFileOverridesDefaults(t *testing.T) {
t.Setenv("XDG_CONFIG_HOME", t.TempDir())
dir, err := config.PmgConfigDir()
dir, err := config.ConfigDir()
assert.NoError(t, err)
assert.NoError(t, os.MkdirAll(dir, 0o755))
cfgFile, _ := config.ConfigFilePath()
+19 -22
View File
@@ -14,17 +14,27 @@ const (
pmgConfigName = "config"
pmgConfigType = "yml"
pmgConfigPath = "safedep/pmg"
PMG_CONFIG_DIR_ENV = "PMG_CONFIG_DIR"
)
// defaultRcFileName is the default name for the shell RC file that contains PMG aliases.
const defaultRcFileName = ".pmg.rc"
const (
defaultRcFileName = "pmg.rc"
)
// PmgConfigDir returns the base application config directory.
// By default, this is:
// ConfigDir returns the base application config directory.
// If the PMG_CONFIG_DIR environment variable is set, its value is used as the base before appending safedep/pmg.
// Otherwise, the defaults are:
// - macOS: ~/Library/Application Support/safedep/pmg
// - Linux: ~/.config/safedep/pmg
// - Windows: %AppData%\safedep\pmg
func PmgConfigDir() (string, error) {
func ConfigDir() (string, error) {
dir := os.Getenv(PMG_CONFIG_DIR_ENV)
if dir != "" {
return filepath.Join(dir, pmgConfigPath), nil
}
userConfigDir, err := os.UserConfigDir()
if err != nil {
return "", fmt.Errorf("failed to retrieve user config directory: %w", err)
@@ -33,9 +43,9 @@ func PmgConfigDir() (string, error) {
return filepath.Join(userConfigDir, pmgConfigPath), nil
}
// CreatePmgConfigDir ensures the application config directory exists and returns its path.
func CreatePmgConfigDir() (string, error) {
dir, err := PmgConfigDir()
// createConfigDir ensures the application config directory exists and returns its path.
func createConfigDir() (string, error) {
dir, err := ConfigDir()
if err != nil {
return "", err
}
@@ -46,23 +56,10 @@ func CreatePmgConfigDir() (string, error) {
return dir, nil
}
// RemovePmgConfigDir removes the PMG configuration directory and its contents.
func RemovePmgConfigDir() error {
dir, err := PmgConfigDir()
if err != nil {
return err
}
if err := os.RemoveAll(dir); err != nil {
return fmt.Errorf("failed to remove config directory %s: %w", dir, err)
}
return nil
}
// ConfigFilePath returns the absolute path to the main PMG config file (e.g., config.yml),
// without creating any directories.
func ConfigFilePath() (string, error) {
dir, err := PmgConfigDir()
dir, err := ConfigDir()
if err != nil {
return "", err
}
@@ -77,7 +74,7 @@ func RcFileName() string {
// RcFilePath returns the absolute path to the PMG RC file under the app config directory,
// without creating any directories.
func RcFilePath() (string, error) {
dir, err := PmgConfigDir()
dir, err := ConfigDir()
if err != nil {
return "", err
}