From c2a8d29ae673f71683c3bc0a515e95a4bced7d26 Mon Sep 17 00:00:00 2001 From: sydneyli Date: Thu, 28 Mar 2019 10:23:49 -0700 Subject: [PATCH 1/3] Regular validator goes through MTA-STS-dependent policies For each MTA-STS-dependent policy, checks MTA-STS policy against internal DB hostname store. --- checker/domain.go | 2 +- checker/mta_sts.go | 4 ++- db/sqldb.go | 27 ++++++++++------- db/sqldb_test.go | 42 +++++++++++++------------- main.go | 9 +++++- models/domain.go | 12 ++++++++ policy/policy.go | 19 +++++++++--- policy/policy_test.go | 6 ++-- util/util.go | 27 +++++++++++++++++ util/util_test.go | 29 ++++++++++++++++++ validator/validator.go | 59 +++++++++++++++++++++++++++---------- validator/validator_test.go | 25 ++++++++-------- 12 files changed, 192 insertions(+), 69 deletions(-) create mode 100644 util/util.go create mode 100644 util/util_test.go diff --git a/checker/domain.go b/checker/domain.go index ed63a46a..7964d89e 100644 --- a/checker/domain.go +++ b/checker/domain.go @@ -127,7 +127,7 @@ func (c *Checker) CheckDomain(domain string, expectedHostnames []string) DomainR } } result.PreferredHostnames = checkedHostnames - result.MTASTSResult = c.checkMTASTS(domain, result.HostnameResults) + result.MTASTSResult = c.CheckMTASTS(domain, result.HostnameResults) // Derive Domain code from Hostname results. if len(checkedHostnames) == 0 { diff --git a/checker/mta_sts.go b/checker/mta_sts.go index 170382ae..f4469bf3 100644 --- a/checker/mta_sts.go +++ b/checker/mta_sts.go @@ -176,7 +176,9 @@ func validateMTASTSMXs(policyFileMXs []string, dnsMXs map[string]HostnameResult, } } -func (c Checker) checkMTASTS(domain string, hostnameResults map[string]HostnameResult) *MTASTSResult { +// CheckMTASTS performs all associated checks for a particular domain's +// MTA-STS support. +func (c Checker) CheckMTASTS(domain string, hostnameResults map[string]HostnameResult) *MTASTSResult { if c.checkMTASTSOverride != nil { // Allow the Checker to mock this function. return c.checkMTASTSOverride(domain, hostnameResults) diff --git a/db/sqldb.go b/db/sqldb.go index 2fa5aa87..f8141c18 100644 --- a/db/sqldb.go +++ b/db/sqldb.go @@ -216,6 +216,14 @@ func (db *SQLDatabase) PutDomain(domain models.Domain) error { return err } +// UpdateDomainPolicy allows us to update the internal data about a particular domain. +func (db *SQLDatabase) UpdateDomainPolicy(domain models.Domain) error { + _, err := db.conn.Exec("UPDATE domains SET data=$2, status=$3 WHERE domain=$1 AND mta_sts=TRUE", + domain.Name, strings.Join(domain.MXs[:], ","), domain.State) + return err + +} + // GetDomain retrieves the status and information associated with a particular // mailserver domain. func (db SQLDatabase) GetDomain(domain string) (models.Domain, error) { @@ -234,7 +242,7 @@ func (db SQLDatabase) GetDomain(domain string) (models.Domain, error) { // GetDomains retrieves all the domains which match a particular state, // that are not in MTA_STS mode func (db SQLDatabase) GetDomains(state models.DomainState) ([]models.Domain, error) { - return db.getDomainsWhere("status=$1", state) + return db.getDomainsWhere("status=$1 AND mta_sts=FALSE", state) } // GetMTASTSDomains retrieves domains which wish their policy to be queued with their MTASTS. @@ -312,20 +320,17 @@ func (db SQLDatabase) DomainsToValidate() ([]string, error) { if err != nil { return domains, err } + dataMTASTS, err := db.GetMTASTSDomains() + if err != nil { + return domains, err + } for _, domainInfo := range data { domains = append(domains, domainInfo.Name) } - return domains, nil -} - -// HostnamesForDomain [interface Validator] retrieves the hostname policy for -// a particular domain. -func (db SQLDatabase) HostnamesForDomain(domain string) ([]string, error) { - data, err := db.GetDomain(domain) - if err != nil { - return []string{}, err + for _, domainInfo := range dataMTASTS { + domains = append(domains, domainInfo.Name) } - return data.MXs, nil + return domains, nil } // GetHostnameScan retrives most recent scan from database. diff --git a/db/sqldb_test.go b/db/sqldb_test.go index 9a9e8ac8..b334eef4 100644 --- a/db/sqldb_test.go +++ b/db/sqldb_test.go @@ -265,26 +265,6 @@ func TestDomainsToValidate(t *testing.T) { } } -func TestHostnamesForDomain(t *testing.T) { - database.ClearTables() - database.PutDomain(models.Domain{Name: "x", MXs: []string{"x.com", "y.org"}}) - database.PutDomain(models.Domain{Name: "y"}) - result, err := database.HostnamesForDomain("x") - if err != nil { - t.Fatalf("HostnamesForDomain failed: %v\n", err) - } - if len(result) != 2 || result[0] != "x.com" || result[1] != "y.org" { - t.Errorf("Expected two hostnames, x.com and y.org\n") - } - result, err = database.HostnamesForDomain("y") - if err != nil { - t.Fatalf("HostnamesForDomain failed: %v\n", err) - } - if len(result) > 0 { - t.Errorf("Expected no hostnames to be returned, got %s\n", result[0]) - } -} - func TestPutAndIsBlacklistedEmail(t *testing.T) { defer database.ClearTables() @@ -447,3 +427,25 @@ func TestGetMTASTSDomains(t *testing.T) { } } } + +func TestUpdateDomainPolicy(t *testing.T) { + database.ClearTables() + database.PutDomain(models.Domain{Name: "no-mtasts"}) + database.PutDomain(models.Domain{Name: "mtasts", MTASTSMode: "on", Email: "real-email"}) + database.UpdateDomainPolicy(models.Domain{Name: "no-mtasts", State: models.StateEnforce}) + database.UpdateDomainPolicy(models.Domain{Name: "mtasts", State: models.StateEnforce, MXs: []string{"hostname"}, Email: "fake-email"}) + domain, _ := database.GetDomain("no-mtasts") + if domain.State == models.StateEnforce { + t.Errorf("Expected State to not update since unicorns isn't MTASTS") + } + domain, _ = database.GetDomain("mtasts") + if domain.State != models.StateEnforce { + t.Errorf("Expected State to update after UpdateDomainPolicy") + } + if len(domain.MXs) != 1 || domain.MXs[0] != "hostname" { + t.Errorf("Expected MXs to update after UpdateDomainPolicy") + } + if domain.Email != "real-email" { + t.Errorf("Did not expect Email to update after UpdateDomainPolicy") + } +} diff --git a/main.go b/main.go index bfab917c..00e95fcb 100644 --- a/main.go +++ b/main.go @@ -119,7 +119,14 @@ func main() { } if os.Getenv("VALIDATE_QUEUED") == "1" { log.Println("[Starting queued validator]") - go validator.ValidateRegularly("Testing domains", db, 24*time.Hour) + v := validator.Validator{ + Name: "Testing and enforced domains", + Store: db, + Interval: 24 * time.Hour, + CheckPerformer: validator.GetDBCheck(db.UpdateDomainPolicy), + } + go v.Run() + // go validator.ValidateRegularly("Testing domains", db, 24*time.Hour) } ServePublicEndpoints(&api, &cfg) } diff --git a/models/domain.go b/models/domain.go index 2c7b9212..56a32ea6 100644 --- a/models/domain.go +++ b/models/domain.go @@ -6,6 +6,7 @@ import ( "time" "github.com/EFForg/starttls-backend/checker" + "github.com/EFForg/starttls-backend/util" ) // Domain stores the preload state of a single domain. @@ -128,3 +129,14 @@ func (d Domain) AsyncPolicyListCheck(store domainStore, list policyList) <-chan go func() { result <- *d.PolicyListCheck(store, list) }() return result } + +// SamePolicy checks whether the underlying policy represented by Domain +// and the one picked up by the MTA-STS check represent the same policy. +func (d *Domain) SamePolicy(result *checker.MTASTSResult) bool { + if (result.Mode == "enforce" && d.State != StateEnforce) || + (result.Mode == "testing" && d.State != StateTesting) || + result.Mode == "none" { + return false + } + return util.ListsEqual(d.MXs, result.MXs) +} diff --git a/policy/policy.go b/policy/policy.go index 8eab88ac..4b5ab46c 100644 --- a/policy/policy.go +++ b/policy/policy.go @@ -8,6 +8,8 @@ import ( "net/http" "sync" "time" + + "github.com/EFForg/starttls-backend/models" ) // policyURL is the default URL from which to fetch the policy JSON. @@ -80,14 +82,23 @@ func (l *UpdatedList) DomainsToValidate() ([]string, error) { return domains, nil } -// HostnamesForDomain [interface Validator] retrieves the hostname policy for +// GetDomain [interface Validator] retrieves the domain object for // a particular domain. -func (l *UpdatedList) HostnamesForDomain(domain string) ([]string, error) { +func (l *UpdatedList) GetDomain(domain string) (models.Domain, error) { policy, err := l.Get(domain) if err != nil { - return []string{}, err + return models.Domain{}, err + } + domainObj := models.Domain{ + Name: domain, + MXs: policy.MXs, + } + if policy.Mode == "enforce" { + domainObj.State = models.StateEnforce + } else if policy.Mode == "testing" { + domainObj.State = models.StateTesting } - return policy.MXs, nil + return domainObj, nil } // Get safely reads from the underlying policy list and returns a TLSPolicy for a domain diff --git a/policy/policy_test.go b/policy/policy_test.go index f79e65bb..cc3a1107 100644 --- a/policy/policy_test.go +++ b/policy/policy_test.go @@ -102,12 +102,12 @@ func TestHostnamesForDomain(t *testing.T) { var updatedList = List{Policies: map[string]TLSPolicy{ "eff.org": TLSPolicy{MXs: hostnames}}} list := makeUpdatedList(func() (List, error) { return updatedList, nil }, time.Second) - returned, err := list.HostnamesForDomain("eff.org") + returned, err := list.GetDomain("eff.org") if err != nil { t.Fatalf("Encountered %v", err) } - if !reflect.DeepEqual(returned, hostnames) { - t.Errorf("Expected %s, got %s", hostnames, returned) + if !reflect.DeepEqual(returned.MXs, hostnames) { + t.Errorf("Expected %s, got %s", hostnames, returned.MXs) } } diff --git a/util/util.go b/util/util.go new file mode 100644 index 00000000..dc057342 --- /dev/null +++ b/util/util.go @@ -0,0 +1,27 @@ +package util + +import ( + "reflect" +) + +// ListsEqual checks that two lists have the same elements, +// regardless of order. +func ListsEqual(x []string, y []string) bool { + // Transform each list into a histogram + xMap := make(map[string]uint) + yMap := make(map[string]uint) + for _, element := range x { + if _, ok := xMap[element]; !ok { + xMap[element] = 0 + } + xMap[element]++ + } + for _, element := range y { + if _, ok := yMap[element]; !ok { + yMap[element] = 0 + } + yMap[element]++ + } + // Compare the histogram maps + return reflect.DeepEqual(xMap, yMap) +} diff --git a/util/util_test.go b/util/util_test.go new file mode 100644 index 00000000..ad2ca66e --- /dev/null +++ b/util/util_test.go @@ -0,0 +1,29 @@ +package util + +import ( + "testing" +) + +func TestListsEqual(t *testing.T) { + testCases := []struct { + x []string + y []string + expected bool + }{ + {[]string{}, []string{}, true}, + {[]string{"a"}, []string{}, false}, + {[]string{"a"}, []string{"a"}, true}, + {[]string{"a", "a"}, []string{"a"}, false}, + {[]string{"a", "b", "c"}, []string{"a", "b", "c"}, true}, + {[]string{"b", "a", "c"}, []string{"a", "b", "c"}, true}, + {[]string{"b", "a", "b", "c"}, []string{"a", "b", "c"}, false}, + {[]string{"b", "a", "b", "c"}, []string{"a", "b", "a", "c"}, false}, + {[]string{"a", "a", "b", "c"}, []string{"a", "b", "a", "c"}, true}, + } + for _, testCase := range testCases { + got := ListsEqual(testCase.x, testCase.y) + if got != testCase.expected { + t.Errorf("Compared %v and %v, expected %v, got %v", testCase.x, testCase.y, testCase.expected, got) + } + } +} diff --git a/validator/validator.go b/validator/validator.go index 87d0b606..197035fd 100644 --- a/validator/validator.go +++ b/validator/validator.go @@ -6,6 +6,7 @@ import ( "time" "github.com/EFForg/starttls-backend/checker" + "github.com/EFForg/starttls-backend/models" "github.com/getsentry/raven-go" ) @@ -14,7 +15,7 @@ import ( // expected hostnames). type DomainPolicyStore interface { DomainsToValidate() ([]string, error) - HostnamesForDomain(string) ([]string, error) + GetDomain(string) (models.Domain, error) } // Called with failure by defaault. @@ -28,8 +29,10 @@ func reportToSentry(name string, domain string, result checker.DomainResult) { result) } -type checkPerformer func(string, []string) checker.DomainResult -type resultCallback func(string, string, checker.DomainResult) +type resultCallback func(string, models.Domain, checker.DomainResult) + +// CheckPerformer defines a function that performs a security check on a domain. +type CheckPerformer func(models.Domain) checker.DomainResult // Validator runs checks regularly against domain policies. This structure // defines the configurations. @@ -47,18 +50,42 @@ type Validator struct { OnFailure resultCallback // OnSuccess: optional. Called when a particular policy validation succeeds. OnSuccess resultCallback - // checkPerformer: performs the check. - checkPerformer checkPerformer + // CheckPerformer: performs the check. + CheckPerformer CheckPerformer +} + +// UpdatePolicy is a callback we can provide to GetDBCheck in order to perform a policy +// update if we notice a discrepancy between our view and the MTA-STS policy. +type UpdatePolicy func(models.Domain) error + +// GetDBCheck returns a CheckPerformer that performs an MTASTS check and update if +// the policy is updated, or performs a regular security check if MTASTS is not supported. +func GetDBCheck(update UpdatePolicy) CheckPerformer { + c := checker.Checker{Cache: checker.MakeSimpleCache(time.Hour)} + return func(domain models.Domain) checker.DomainResult { + if domain.MTASTSMode == "on" { + result := c.CheckDomain(domain.Name, []string{}) + if !domain.SamePolicy(result.MTASTSResult) { + if update(domain) != nil { + reportToSentry("Couldn't update policy in DB", domain.Name, result) + } + } + return result + } + return c.CheckDomain(domain.Name, domain.MXs) + } } -func (v *Validator) checkPolicy(domain string, hostnames []string) checker.DomainResult { - if v.checkPerformer == nil { +func (v *Validator) checkPolicy(domain models.Domain) checker.DomainResult { + if v.CheckPerformer == nil { c := checker.Checker{ Cache: checker.MakeSimpleCache(time.Hour), } - v.checkPerformer = c.CheckDomain + v.CheckPerformer = func(domain models.Domain) checker.DomainResult { + return c.CheckDomain(domain.Name, domain.MXs) + } } - return v.checkPerformer(domain, hostnames) + return v.CheckPerformer(domain) } func (v *Validator) interval() time.Duration { @@ -68,14 +95,14 @@ func (v *Validator) interval() time.Duration { return time.Hour * 24 } -func (v *Validator) policyFailed(name string, domain string, result checker.DomainResult) { +func (v *Validator) policyFailed(name string, domain models.Domain, result checker.DomainResult) { if v.OnFailure != nil { v.OnFailure(name, domain, result) } - reportToSentry(name, domain, result) + reportToSentry(name, domain.Name, result) } -func (v *Validator) policyPassed(name string, domain string, result checker.DomainResult) { +func (v *Validator) policyPassed(name string, domain models.Domain, result checker.DomainResult) { if v.OnSuccess != nil { v.OnSuccess(name, domain, result) } @@ -93,17 +120,17 @@ func (v *Validator) Run() { continue } for _, domain := range domains { - hostnames, err := v.Store.HostnamesForDomain(domain) + domainData, err := v.Store.GetDomain(domain) if err != nil { log.Printf("[%s validator] Could not retrieve policy for domain %s: %v", v.Name, domain, err) continue } - result := v.checkPolicy(domain, hostnames) + result := v.checkPolicy(domainData) if result.Status != 0 { log.Printf("[%s validator] %s failed; sending report", v.Name, domain) - v.policyFailed(v.Name, domain, result) + v.policyFailed(v.Name, domainData, result) } else { - v.policyPassed(v.Name, domain, result) + v.policyPassed(v.Name, domainData, result) } } } diff --git a/validator/validator_test.go b/validator/validator_test.go index eec55a9f..ef16ba2c 100644 --- a/validator/validator_test.go +++ b/validator/validator_test.go @@ -5,6 +5,7 @@ import ( "time" "github.com/EFForg/starttls-backend/checker" + "github.com/EFForg/starttls-backend/models" ) type mockDomainPolicyStore struct { @@ -19,21 +20,21 @@ func (m mockDomainPolicyStore) DomainsToValidate() ([]string, error) { return domains, nil } -func (m mockDomainPolicyStore) HostnamesForDomain(domain string) ([]string, error) { - return m.hostnames[domain], nil +func (m mockDomainPolicyStore) GetDomain(domain string) (models.Domain, error) { + return models.Domain{Name: domain, MXs: m.hostnames[domain]}, nil } -func noop(_ string, _ string, _ checker.DomainResult) {} +func noop(_ string, _ models.Domain, _ checker.DomainResult) {} func TestRegularValidationValidates(t *testing.T) { called := make(chan bool) - fakeChecker := func(domain string, hostnames []string) checker.DomainResult { + fakeChecker := func(_ models.Domain) checker.DomainResult { called <- true return checker.DomainResult{} } mock := mockDomainPolicyStore{ hostnames: map[string][]string{"a": []string{"hostname"}}} - v := Validator{Store: mock, Interval: 100 * time.Millisecond, checkPerformer: fakeChecker, OnFailure: noop} + v := Validator{Store: mock, Interval: 100 * time.Millisecond, CheckPerformer: fakeChecker, OnFailure: noop} go v.Run() select { @@ -46,25 +47,25 @@ func TestRegularValidationValidates(t *testing.T) { func TestRegularValidationReportsErrors(t *testing.T) { reports := make(chan string) - fakeChecker := func(domain string, hostnames []string) checker.DomainResult { - if domain == "fail" || domain == "error" { + fakeChecker := func(domain models.Domain) checker.DomainResult { + if domain.Name == "fail" || domain.Name == "error" { return checker.DomainResult{Status: 5} } return checker.DomainResult{Status: 0} } - fakeReporter := func(name string, domain string, result checker.DomainResult) { - reports <- domain + fakeReporter := func(name string, domain models.Domain, result checker.DomainResult) { + reports <- domain.Name } successReports := make(chan string) - fakeSuccessReporter := func(name string, domain string, result checker.DomainResult) { - successReports <- domain + fakeSuccessReporter := func(name string, domain models.Domain, result checker.DomainResult) { + successReports <- domain.Name } mock := mockDomainPolicyStore{ hostnames: map[string][]string{ "fail": []string{"hostname"}, "error": []string{"hostname"}, "normal": []string{"hostname"}}} - v := Validator{Store: mock, Interval: 100 * time.Millisecond, checkPerformer: fakeChecker, + v := Validator{Store: mock, Interval: 100 * time.Millisecond, CheckPerformer: fakeChecker, OnFailure: fakeReporter, OnSuccess: fakeSuccessReporter, } go v.Run() From 321f55e6bc7bcb1b11d4b5b477fc29fafc35575e Mon Sep 17 00:00:00 2001 From: sydneyli Date: Mon, 6 May 2019 14:12:20 -0700 Subject: [PATCH 2/3] fix golint path --- .travis.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.travis.yml b/.travis.yml index 3e3d41cf..c568cd01 100644 --- a/.travis.yml +++ b/.travis.yml @@ -10,7 +10,7 @@ env: - TEST_DB_NAME=starttls_test install: - - go get -u github.com/golang/lint/golint + - go get -u golang.org/x/lint/golint - go get github.com/mattn/goveralls - go get -t ./... From fccb65455133196157a769440bef247a07bcdb71 Mon Sep 17 00:00:00 2001 From: sydneyli Date: Mon, 6 May 2019 15:40:53 -0700 Subject: [PATCH 3/3] Revert name changes from GetDomain => GetDomainInState --- db/db.go | 2 +- db/sqldb.go | 14 ++++++++------ db/sqldb_test.go | 16 ++++++++-------- models/domain.go | 12 ++++++------ models/domain_test.go | 2 +- models/token.go | 2 +- policy/policy.go | 6 +++--- policy/policy_test.go | 2 +- validator/validator.go | 4 ++-- validator/validator_test.go | 2 +- 10 files changed, 32 insertions(+), 30 deletions(-) diff --git a/db/db.go b/db/db.go index dc6307db..7f302d55 100644 --- a/db/db.go +++ b/db/db.go @@ -36,7 +36,7 @@ type Database interface { // Upserts domain state. PutDomain(models.Domain) error // Retrieves state of a domain - GetDomainInState(string, models.DomainState) (models.Domain, error) + GetDomain(string, models.DomainState) (models.Domain, error) // Retrieves all domains in a particular state. GetDomains(models.DomainState) ([]models.Domain, error) SetStatus(string, models.DomainState) error diff --git a/db/sqldb.go b/db/sqldb.go index e9f85131..8f6a80bc 100644 --- a/db/sqldb.go +++ b/db/sqldb.go @@ -224,9 +224,9 @@ func (db *SQLDatabase) UpdateDomainPolicy(domain models.Domain) error { } -// GetDomainInState retrieves the status and information associated with a particular +// GetDomain retrieves the status and information associated with a particular // mailserver domain. -func (db SQLDatabase) GetDomainInState(domain string, state models.DomainState) (models.Domain, error) { +func (db SQLDatabase) GetDomain(domain string, state models.DomainState) (models.Domain, error) { return db.queryDomain("SELECT %s FROM domains WHERE domain=$1 AND status=$2", domain, state) } @@ -332,7 +332,7 @@ func (db SQLDatabase) queryDomainsWhere(condition string, args ...interface{}) ( return domains, nil } -// DomainsToValidate [interface Validator] retrieves domains from the +// DomainsToValidate [interface DomainPolicyStore] retrieves domains from the // DB whose policies should be validated. func (db SQLDatabase) DomainsToValidate() ([]string, error) { domains := []string{} @@ -353,10 +353,12 @@ func (db SQLDatabase) DomainsToValidate() ([]string, error) { return domains, nil } -func (db SQLDatabase) GetDomain(domain string) (models.Domain, error) { - data, err := db.GetDomainInState(domain, models.StateEnforce) +// GetDomainPolicy [interface DomainPolicyStore] retrieves the domain object for +// a particular domain. +func (db SQLDatabase) GetDomainPolicy(domain string) (models.Domain, error) { + data, err := db.GetDomain(domain, models.StateEnforce) if err != nil { - data, err = db.GetDomainInState(domain, models.StateTesting) + data, err = db.GetDomain(domain, models.StateTesting) } return data, err } diff --git a/db/sqldb_test.go b/db/sqldb_test.go index e6eff2db..c8db7ead 100644 --- a/db/sqldb_test.go +++ b/db/sqldb_test.go @@ -147,7 +147,7 @@ func TestPutGetDomain(t *testing.T) { if err != nil { t.Errorf("PutDomain failed: %v\n", err) } - retrievedData, err := database.GetDomainInState(data.Name, models.StateUnconfirmed) + retrievedData, err := database.GetDomain(data.Name, models.StateUnconfirmed) if err != nil { t.Errorf("GetDomain(%s) failed: %v\n", data.Name, err) } @@ -171,7 +171,7 @@ func TestUpsertDomain(t *testing.T) { if err != nil { t.Errorf("PutDomain(%s) failed: %v\n", data.Name, err) } - retrievedData, err := database.GetDomainInState(data.Name, models.StateUnconfirmed) + retrievedData, err := database.GetDomain(data.Name, models.StateUnconfirmed) if retrievedData.MXs[0] != "hello_darkness_my_old_friend" || retrievedData.Email != "actual_admin@testing.com" { t.Errorf("Email and MXs should have been rewritten: %v\n", retrievedData) } @@ -220,11 +220,11 @@ func TestLastUpdatedFieldUpdates(t *testing.T) { State: models.StateUnconfirmed, } database.PutDomain(data) - retrievedData, _ := database.GetDomainInState(data.Name, models.StateUnconfirmed) + retrievedData, _ := database.GetDomain(data.Name, models.StateUnconfirmed) lastUpdated := retrievedData.LastUpdated data.State = models.StateTesting database.PutDomain(models.Domain{Name: data.Name, Email: "new fone who dis"}) - retrievedData, _ = database.GetDomainInState(data.Name, models.StateUnconfirmed) + retrievedData, _ = database.GetDomain(data.Name, models.StateUnconfirmed) if lastUpdated.Equal(retrievedData.LastUpdated) { t.Errorf("Expected last_updated to be updated on change: %v", lastUpdated) } @@ -238,10 +238,10 @@ func TestLastUpdatedFieldDoesntUpdate(t *testing.T) { State: models.StateUnconfirmed, } database.PutDomain(data) - retrievedData, _ := database.GetDomainInState(data.Name, models.StateUnconfirmed) + retrievedData, _ := database.GetDomain(data.Name, models.StateUnconfirmed) lastUpdated := retrievedData.LastUpdated database.PutDomain(data) - retrievedData, _ = database.GetDomainInState(data.Name, models.StateUnconfirmed) + retrievedData, _ = database.GetDomain(data.Name, models.StateUnconfirmed) if !lastUpdated.Equal(retrievedData.LastUpdated) { t.Errorf("Expected last_updated to stay the same if no changes were made") } @@ -439,11 +439,11 @@ func TestUpdateDomainPolicy(t *testing.T) { database.PutDomain(models.Domain{Name: "mtasts", MTASTS: true, Email: "real-email"}) database.UpdateDomainPolicy(models.Domain{Name: "no-mtasts", State: models.StateEnforce}) database.UpdateDomainPolicy(models.Domain{Name: "mtasts", State: models.StateEnforce, MXs: []string{"hostname"}, Email: "fake-email"}) - domain, _ := database.GetDomain("no-mtasts") + domain, _ := database.GetDomainPolicy("no-mtasts") if domain.State == models.StateEnforce { t.Errorf("Expected State to not update since unicorns isn't MTASTS") } - domain, _ = database.GetDomain("mtasts") + domain, _ = database.GetDomainPolicy("mtasts") if domain.State != models.StateEnforce { t.Errorf("Expected State to update after UpdateDomainPolicy") } diff --git a/models/domain.go b/models/domain.go index eb53fc14..68158a33 100644 --- a/models/domain.go +++ b/models/domain.go @@ -30,7 +30,7 @@ type Domain struct { // domainStore is a simple interface for fetching and adding domain objects. type domainStore interface { PutDomain(Domain) error - GetDomainInState(string, DomainState) (Domain, error) + GetDomain(string, DomainState) (Domain, error) GetDomains(DomainState) ([]Domain, error) SetStatus(string, DomainState) error RemoveDomain(string, DomainState) (Domain, error) @@ -70,7 +70,7 @@ func (d *Domain) IsQueueable(domains domainStore, scans scanStore, list policyLi if list.HasDomain(d.Name) { return false, "Domain is already on the policy list!", scan } - if _, err := domains.GetDomainInState(d.Name, StateEnforce); err == nil { + if _, err := domains.GetDomain(d.Name, StateEnforce); err == nil { return false, "Domain is already on the policy list!", scan } // Domains without submitted MTA-STS support must match provided mx patterns. @@ -159,17 +159,17 @@ func (d *Domain) SamePolicy(result *checker.MTASTSResult) bool { // or StateTesting. If that domain exists in the store, return that one. // Otherwise, look for a Domain policy in the unconfirmed state. func GetDomain(store domainStore, name string) (Domain, error) { - domain, err := store.GetDomainInState(name, StateEnforce) + domain, err := store.GetDomain(name, StateEnforce) if err == nil { return domain, nil } - domain, err = store.GetDomainInState(name, StateTesting) + domain, err = store.GetDomain(name, StateTesting) if err == nil { return domain, nil } - domain, err = store.GetDomainInState(name, StateUnconfirmed) + domain, err = store.GetDomain(name, StateUnconfirmed) if err == nil { return domain, nil } - return store.GetDomainInState(name, StateFailed) + return store.GetDomain(name, StateFailed) } diff --git a/models/domain_test.go b/models/domain_test.go index 02f13ed6..0dea7490 100644 --- a/models/domain_test.go +++ b/models/domain_test.go @@ -24,7 +24,7 @@ func (m *mockDomainStore) SetStatus(d string, status DomainState) error { return m.err } -func (m *mockDomainStore) GetDomainInState(d string, state DomainState) (Domain, error) { +func (m *mockDomainStore) GetDomain(d string, state DomainState) (Domain, error) { domain := m.domain if state != domain.State { return m.domain, errors.New("") diff --git a/models/token.go b/models/token.go index e51e9e33..a4caf7a5 100644 --- a/models/token.go +++ b/models/token.go @@ -23,7 +23,7 @@ func (t *Token) Redeem(store domainStore, tokens tokenStore) (ret string, userEr if err != nil { return domain, err, nil } - domainData, err := store.GetDomainInState(domain, StateUnconfirmed) + domainData, err := store.GetDomain(domain, StateUnconfirmed) if err != nil { return domain, nil, err } diff --git a/policy/policy.go b/policy/policy.go index 4b5ab46c..c9494748 100644 --- a/policy/policy.go +++ b/policy/policy.go @@ -70,7 +70,7 @@ type UpdatedList struct { *List } -// DomainsToValidate [interface Validator] retrieves domains from the +// DomainsToValidate [interface DomainPolicyStore] retrieves domains from the // DB whose policies should be validated. func (l *UpdatedList) DomainsToValidate() ([]string, error) { l.mu.RLock() @@ -82,9 +82,9 @@ func (l *UpdatedList) DomainsToValidate() ([]string, error) { return domains, nil } -// GetDomain [interface Validator] retrieves the domain object for +// GetDomainPolicy [interface DomainPolicyStore] retrieves the domain object for // a particular domain. -func (l *UpdatedList) GetDomain(domain string) (models.Domain, error) { +func (l *UpdatedList) GetDomainPolicy(domain string) (models.Domain, error) { policy, err := l.Get(domain) if err != nil { return models.Domain{}, err diff --git a/policy/policy_test.go b/policy/policy_test.go index cc3a1107..5cff375c 100644 --- a/policy/policy_test.go +++ b/policy/policy_test.go @@ -102,7 +102,7 @@ func TestHostnamesForDomain(t *testing.T) { var updatedList = List{Policies: map[string]TLSPolicy{ "eff.org": TLSPolicy{MXs: hostnames}}} list := makeUpdatedList(func() (List, error) { return updatedList, nil }, time.Second) - returned, err := list.GetDomain("eff.org") + returned, err := list.GetDomainPolicy("eff.org") if err != nil { t.Fatalf("Encountered %v", err) } diff --git a/validator/validator.go b/validator/validator.go index b829923e..fb8e27e9 100644 --- a/validator/validator.go +++ b/validator/validator.go @@ -15,7 +15,7 @@ import ( // expected hostnames). type DomainPolicyStore interface { DomainsToValidate() ([]string, error) - GetDomain(string) (models.Domain, error) + GetDomainPolicy(string) (models.Domain, error) } // Called with failure by defaault. @@ -120,7 +120,7 @@ func (v *Validator) Run() { continue } for _, domain := range domains { - domainData, err := v.Store.GetDomain(domain) + domainData, err := v.Store.GetDomainPolicy(domain) if err != nil { log.Printf("[%s validator] Could not retrieve policy for domain %s: %v", v.Name, domain, err) continue diff --git a/validator/validator_test.go b/validator/validator_test.go index ef16ba2c..3457faed 100644 --- a/validator/validator_test.go +++ b/validator/validator_test.go @@ -20,7 +20,7 @@ func (m mockDomainPolicyStore) DomainsToValidate() ([]string, error) { return domains, nil } -func (m mockDomainPolicyStore) GetDomain(domain string) (models.Domain, error) { +func (m mockDomainPolicyStore) GetDomainPolicy(domain string) (models.Domain, error) { return models.Domain{Name: domain, MXs: m.hostnames[domain]}, nil }