package traverse import ( "fmt" "net" "sync" "testing" "github.com/miekg/dns" ) func nsRR(zone, target string) dns.RR { return &dns.NS{ Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: dns.Fqdn(target), } } func aRR(name, ip string) dns.RR { return &dns.A{ Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET}, A: net.ParseIP(ip).To4(), } } func aaaaRR(name, ip string) dns.RR { return &dns.AAAA{ Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeAAAA, Class: dns.ClassINET}, AAAA: net.ParseIP(ip), } } func TestCanonicalName(t *testing.T) { tests := []struct{ in, want string }{ {"example.com", "example.com"}, {"Example.COM.", "example.com"}, {".", ""}, {"", ""}, {"WWW.Example.Com", "www.example.com"}, } for _, tt := range tests { if got := canonicalName(tt.in); got != tt.want { t.Errorf("canonicalName(%q) = %q, want %q", tt.in, got, tt.want) } } } func TestInfoCacheAddAndGet(t *testing.T) { c := NewInfoCache(nil) c.Add([]dns.RR{nsRR("com", "a.gtld-servers.net"), nsRR("com", "b.gtld-servers.net")}) rrs := c.Get("com", dns.ClassINET, dns.TypeNS) if len(rrs) != 2 { t.Fatalf("expected 2 NS records, got %d", len(rrs)) } } func TestInfoCacheAddReplacesSameKey(t *testing.T) { c := NewInfoCache(nil) c.Add([]dns.RR{nsRR("com", "old1.example.net"), nsRR("com", "old2.example.net")}) c.Add([]dns.RR{nsRR("com", "new.example.net")}) rrs := c.Get("com", dns.ClassINET, dns.TypeNS) if len(rrs) != 1 { t.Fatalf("add should replace same name:class:type entry, got %d records", len(rrs)) } if rrs[0].(*dns.NS).Ns != "new.example.net." { t.Errorf("NS = %q, want new.example.net.", rrs[0].(*dns.NS).Ns) } } func TestInfoCacheAddKeepsDistinctKeys(t *testing.T) { c := NewInfoCache(nil) c.Add([]dns.RR{nsRR("com", "a.gtld-servers.net"), aRR("a.gtld-servers.net", "192.5.6.30")}) c.Add([]dns.RR{nsRR("org", "a0.org-servers.net")}) if got := c.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { t.Errorf("com NS lost after unrelated add: %v", got) } if got := c.Get("a.gtld-servers.net", dns.ClassINET, dns.TypeA); len(got) != 1 { t.Errorf("glue lost after unrelated add: %v", got) } } func TestInfoCacheGetCaseInsensitive(t *testing.T) { c := NewInfoCache(nil) c.Add([]dns.RR{nsRR("COM.", "A.GTLD-SERVERS.NET.")}) if got := c.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { t.Fatalf("expected case-insensitive hit, got %v", got) } } func TestInfoCacheGetMiss(t *testing.T) { c := NewInfoCache(nil) if got := c.Get("org", dns.ClassINET, dns.TypeNS); got != nil { t.Errorf("expected nil for miss, got %v", got) } } func TestInfoCacheGetRecursesToParent(t *testing.T) { parent := NewInfoCache(nil) parent.Add([]dns.RR{nsRR("com", "a.gtld-servers.net")}) child := parent.Child() if child.parent != parent { t.Fatal("Child() should link to parent") } if got := child.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { t.Fatalf("expected parent hit through child, got %v", got) } } func TestInfoCacheChildShadowsParent(t *testing.T) { parent := NewInfoCache(nil) parent.Add([]dns.RR{nsRR("com", "parent.example.net")}) child := parent.Child() child.Add([]dns.RR{nsRR("com", "child.example.net")}) got := child.Get("com", dns.ClassINET, dns.TypeNS) if len(got) != 1 || got[0].(*dns.NS).Ns != "child.example.net." { t.Errorf("child entry should shadow parent, got %v", got) } // The parent must be untouched. got = parent.Get("com", dns.ClassINET, dns.TypeNS) if len(got) != 1 || got[0].(*dns.NS).Ns != "parent.example.net." { t.Errorf("parent entry modified, got %v", got) } } func TestGetStartServersWalksLabels(t *testing.T) { c := NewInfoCache(nil) c.Add([]dns.RR{ nsRR("com", "a.gtld-servers.net"), nsRR("com", "b.gtld-servers.net"), aRR("a.gtld-servers.net", "192.5.6.30"), }) starters, bw, err := c.GetStartServers("www.deep.example.com") if err != nil { t.Fatalf("GetStartServers: %v", err) } if bw != "com" { t.Errorf("newbailiwick = %q, want com", bw) } if len(starters) != 2 { t.Fatalf("expected 2 starters, got %d", len(starters)) } if starters[0].Name != "a.gtld-servers.net" { t.Errorf("starter[0] = %q", starters[0].Name) } if len(starters[0].IPs) != 1 || starters[0].IPs[0] != "192.5.6.30" { t.Errorf("starter[0] IPs = %v, want [192.5.6.30]", starters[0].IPs) } if starters[1].IPs != nil { t.Errorf("glueless starter should have nil IPs, got %v", starters[1].IPs) } } func TestGetStartServersPrefersDeepestZone(t *testing.T) { c := NewInfoCache(nil) c.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}}) c.Add([]dns.RR{nsRR("example.com", "ns1.example.com"), aRR("ns1.example.com", "1.2.3.4")}) starters, bw, err := c.GetStartServers("www.example.com") if err != nil { t.Fatalf("GetStartServers: %v", err) } if bw != "example.com" { t.Errorf("newbailiwick = %q, want example.com", bw) } if len(starters) != 1 || starters[0].Name != "ns1.example.com" { t.Errorf("starters = %v", starters) } } func TestGetStartServersRootHints(t *testing.T) { c := NewInfoCache(nil) c.AddHints("", []StartServer{ {Name: "a.root-servers.net", IPs: []string{"198.41.0.4", "2001:503:ba3e::2:30"}}, {Name: "b.root-servers.net", IPs: []string{"170.247.170.2"}}, }) starters, bw, err := c.GetStartServers("anything.example.org") if err != nil { t.Fatalf("GetStartServers: %v", err) } if bw != "" { t.Errorf("root bailiwick should be \"\", got %q", bw) } if len(starters) != 2 { t.Fatalf("expected 2 root starters, got %d", len(starters)) } // Only the IPv4 address surfaces (IPv4-only transport); the AAAA is // cached but not returned as a start address. if len(starters[0].IPs) != 1 || starters[0].IPs[0] != "198.41.0.4" { t.Errorf("starter[0].IPs = %v, want [198.41.0.4]", starters[0].IPs) } if got := c.Get("a.root-servers.net", dns.ClassINET, dns.TypeAAAA); len(got) != 1 { t.Errorf("AAAA hint should be cached, got %v", got) } } func TestGetStartServersNoRootHints(t *testing.T) { c := NewInfoCache(nil) if _, _, err := c.GetStartServers("example.com"); err == nil { t.Fatal("expected error with no NS cached anywhere") } } func TestGetStartServersExactDomainMatch(t *testing.T) { c := NewInfoCache(nil) c.Add([]dns.RR{nsRR("example.com", "ns1.example.net")}) _, bw, err := c.GetStartServers("example.com") if err != nil { t.Fatalf("GetStartServers: %v", err) } if bw != "example.com" { t.Errorf("newbailiwick = %q, want example.com", bw) } } func TestGetStartServersUsesBranchCache(t *testing.T) { parent := NewInfoCache(nil) parent.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}}) child := parent.Child() child.Add([]dns.RR{nsRR("com", "a.gtld-servers.net")}) _, bw, err := child.GetStartServers("www.example.com") if err != nil { t.Fatalf("GetStartServers: %v", err) } if bw != "com" { t.Errorf("newbailiwick = %q, want com (child cache hit)", bw) } // A sibling branch must not see the child's records. sibling := parent.Child() _, bw, err = sibling.GetStartServers("www.example.com") if err != nil { t.Fatalf("GetStartServers: %v", err) } if bw != "" { t.Errorf("sibling newbailiwick = %q, want \"\" (root only)", bw) } } func TestInfoCacheConcurrentAccess(t *testing.T) { c := NewInfoCache(nil) var wg sync.WaitGroup for i := 0; i < 100; i++ { wg.Add(1) go func(i int) { defer wg.Done() name := fmt.Sprintf("ns%d.example.com", i) c.Add([]dns.RR{nsRR("example.com", name), aRR(name, fmt.Sprintf("1.2.3.%d", i%256))}) _, _, _ = c.GetStartServers("www.example.com") _ = c.Get(name, dns.ClassINET, dns.TypeA) }(i) } wg.Wait() }