From 0ef49e9516db7e453de1234d902266bde20703af Mon Sep 17 00:00:00 2001 From: Dustin Spicuzza Date: Wed, 11 Jul 2018 00:37:12 -0400 Subject: [PATCH] Add support for retrieving all IdentityFile directives via DefaultAll --- config.go | 22 +++++++++++++++++----- config_test.go | 3 +-- validators.go | 20 ++++++++++++++++++++ 3 files changed, 38 insertions(+), 7 deletions(-) diff --git a/config.go b/config.go index 7c2c8b6..8fb0bb6 100644 --- a/config.go +++ b/config.go @@ -50,6 +50,13 @@ var _ = version type configFinder func() string +type config interface { + getinternal(alias, key string) string +} + +var _ config = &UserSettings{} +var _ config = &Config{} + // UserSettings checks ~/.ssh and /etc/ssh for configuration files. The config // files are parsed and cached the first time Get() or GetStrict() is called. type UserSettings struct { @@ -180,6 +187,10 @@ func (u *UserSettings) Get(alias, key string) string { return val } +func (u *UserSettings) getinternal(alias, key string) string { + return u.Get(alias, key) +} + // GetAll retrieves zero or more directives for key for the given alias. GetAll // returns nil if no value was found, or if IgnoreErrors is false and we could // not parse the configuration file. Use GetStrict to disambiguate the latter @@ -237,11 +248,7 @@ func (u *UserSettings) GetAllStrict(alias, key string) ([]string, error) { if err2 != nil || val2 != nil { return val2, err2 } - // TODO: IdentityFile has multiple default values that we should return. - if def := Default(key); def != "" { - return []string{def}, nil - } - return []string{}, nil + return DefaultAll(key, alias, u), nil } func (u *UserSettings) doLoadConfigs() { @@ -365,6 +372,11 @@ func (c *Config) Get(alias, key string) (string, error) { return "", nil } +func (c *Config) getinternal(alias, key string) string { + v, _ := c.Get(alias, key) + return v +} + // GetAll returns all values in the configuration that match the alias and // contains key, or nil if none are present. func (c *Config) GetAll(alias, key string) ([]string, error) { diff --git a/config_test.go b/config_test.go index f086600..4ea0eaf 100644 --- a/config_test.go +++ b/config_test.go @@ -107,8 +107,7 @@ func TestGetIdentities(t *testing.T) { t.Errorf("expected nil err, got %v", err) } if len(val) != len(defaultProtocol2Identities) { - // TODO: return the right values here. - log.Printf("expected defaults, got %v", val) + t.Errorf("expected defaults, got %v", val) } else { for i, v := range defaultProtocol2Identities { if val[i] != v { diff --git a/validators.go b/validators.go index 5977f90..57b897f 100644 --- a/validators.go +++ b/validators.go @@ -15,6 +15,26 @@ func Default(keyword string) string { return defaults[strings.ToLower(keyword)] } +// DefaultAll returns the default value for the given keyword, but as a slice. If +// there is no default for the keyword, nil is returned. +// +// Some multi-valued settings have different defaults based on other settings, so +// you must provide the host alias and a config to retrieve a setting from +func DefaultAll(keyword string, alias string, cfg config) []string { + if strings.ToLower(keyword) == "identityfile" && cfg.getinternal(alias, "Protocol") == "2" { + def := make([]string, len(defaultProtocol2Identities)) + copy(def, defaultProtocol2Identities) + return def + } + + def := Default(keyword) + if def != "" { + return []string{def} + } + + return nil +} + // Arguments where the value must be "yes" or "no" and *only* yes or no. var yesnos = map[string]bool{ strings.ToLower("BatchMode"): true,