diff --git a/pkg/spflib/parse.go b/pkg/spflib/parse.go index 12415fd762..2f324b5d23 100644 --- a/pkg/spflib/parse.go +++ b/pkg/spflib/parse.go @@ -43,6 +43,10 @@ var qualifiers = map[byte]bool{ // Parse parses a raw SPF record. func Parse(text string, dnsres Resolver) (*SPFRecord, error) { + return parse(text, dnsres, nil) +} + +func parse(text string, dnsres Resolver, chain []string) (*SPFRecord, error) { if !strings.HasPrefix(text, "v=spf1 ") { return nil, errors.New("not an SPF record") } @@ -84,11 +88,14 @@ func Parse(text string, dnsres Resolver) (*SPFRecord, error) { } p.IsLookup = true if dnsres != nil { + if inChain(chain, p.IncludeDomain) { + return nil, fmt.Errorf("SPF include loop: %s", strings.Join(append(chain, p.IncludeDomain), " -> ")) + } subRecord, err := dnsres.GetSPF(p.IncludeDomain) if err != nil { return nil, err } - p.IncludeRecord, err = Parse(subRecord, dnsres) + p.IncludeRecord, err = parse(subRecord, dnsres, append(chain, p.IncludeDomain)) if err != nil { return nil, fmt.Errorf("in included SPF: %w", err) } @@ -101,3 +108,12 @@ func Parse(text string, dnsres Resolver) (*SPFRecord, error) { } return rec, nil } + +func inChain(chain []string, domain string) bool { + for _, d := range chain { + if strings.EqualFold(d, domain) { + return true + } + } + return false +} diff --git a/pkg/spflib/parse_test.go b/pkg/spflib/parse_test.go index fe5e76d569..a47f6f627d 100644 --- a/pkg/spflib/parse_test.go +++ b/pkg/spflib/parse_test.go @@ -151,6 +151,87 @@ func TestParseQualifiedMechanisms(t *testing.T) { } } +func TestParseIncludeLoop(t *testing.T) { + tests := []struct { + description string + dnsres fakeResolver + input string + wantErr string + }{ + { + description: "a domain that includes itself", + dnsres: fakeResolver{ + "a.example.com": "v=spf1 include:a.example.com ~all", + }, + input: "v=spf1 include:a.example.com ~all", + wantErr: "in included SPF: SPF include loop: a.example.com -> a.example.com", + }, + { + description: "two domains that include each other", + dnsres: fakeResolver{ + "a.example.com": "v=spf1 include:b.example.com ~all", + "b.example.com": "v=spf1 include:a.example.com ~all", + }, + input: "v=spf1 include:a.example.com ~all", + wantErr: "in included SPF: in included SPF: SPF include loop: a.example.com -> b.example.com -> a.example.com", + }, + { + description: "a domain that redirects to itself", + dnsres: fakeResolver{ + "a.example.com": "v=spf1 redirect=a.example.com", + }, + input: "v=spf1 redirect=a.example.com", + wantErr: "in included SPF: SPF include loop: a.example.com -> a.example.com", + }, + { + description: "a qualified include that loops", + dnsres: fakeResolver{ + "a.example.com": "v=spf1 +include:a.example.com ~all", + }, + input: "v=spf1 ?include:a.example.com ~all", + wantErr: "in included SPF: SPF include loop: a.example.com -> a.example.com", + }, + { + description: "a loop that changes the case of the domain", + dnsres: fakeResolver{ + "a.example.com": "v=spf1 include:A.EXAMPLE.COM ~all", + "A.EXAMPLE.COM": "v=spf1 include:a.example.com ~all", + }, + input: "v=spf1 include:a.example.com ~all", + wantErr: "in included SPF: SPF include loop: a.example.com -> A.EXAMPLE.COM", + }, + } + for _, tt := range tests { + t.Run(tt.description, func(t *testing.T) { + _, err := Parse(tt.input, tt.dnsres) + if err == nil { + t.Fatalf("Parse(%q) error = nil, want %q", tt.input, tt.wantErr) + } + if err.Error() != tt.wantErr { + t.Errorf("Parse(%q) error = %q, want %q", tt.input, err, tt.wantErr) + } + }) + } +} + +func TestParseSharedIncludeIsNotALoop(t *testing.T) { + dnsres := fakeResolver{ + "a.example.com": "v=spf1 include:shared.example.com ~all", + "b.example.com": "v=spf1 include:shared.example.com ~all", + "shared.example.com": "v=spf1 ip4:192.0.2.0/24 ~all", + } + rec, err := Parse("v=spf1 include:a.example.com include:b.example.com ~all", dnsres) + if err != nil { + t.Fatal(err) + } + if got, want := len(rec.Parts), 3; got != want { + t.Errorf("len(Parts) = %d, want %d", got, want) + } + if got, want := rec.Lookups(), 4; got != want { + t.Errorf("Lookups() = %d, want %d", got, want) + } +} + func TestParseRedirectLast(t *testing.T) { dnsres, err := NewCache("testdata-dns1.json") if err != nil {