Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 17 additions & 1 deletion pkg/spflib/parse.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
Expand Down Expand Up @@ -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)
}
Expand All @@ -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
}
81 changes: 81 additions & 0 deletions pkg/spflib/parse_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
Loading