Phase 4.1: Error handling, edge cases, and robustness (#11)
CI / test (push) Failing after 2m36s
CI / test (push) Failing after 2m36s
This commit was merged in pull request #11.
This commit is contained in:
@@ -164,6 +164,7 @@ func main() {
|
|||||||
QueryType: queryTypeValue,
|
QueryType: queryTypeValue,
|
||||||
RootConfig: rootConfig,
|
RootConfig: rootConfig,
|
||||||
QueryConfig: queryConfig,
|
QueryConfig: queryConfig,
|
||||||
|
Fast: cfg.Fast,
|
||||||
}
|
}
|
||||||
|
|
||||||
traverser := traverse.NewTraverser(traverserConfig)
|
traverser := traverse.NewTraverser(traverserConfig)
|
||||||
|
|||||||
@@ -2,12 +2,15 @@ module github.com/hits/ExploreDNS
|
|||||||
|
|
||||||
go 1.25.6
|
go 1.25.6
|
||||||
|
|
||||||
require github.com/miekg/dns v1.1.72
|
require (
|
||||||
|
github.com/miekg/dns v1.1.72
|
||||||
|
golang.org/x/net v0.48.0
|
||||||
|
)
|
||||||
|
|
||||||
require (
|
require (
|
||||||
golang.org/x/mod v0.31.0 // indirect
|
golang.org/x/mod v0.31.0 // indirect
|
||||||
golang.org/x/net v0.48.0 // indirect
|
|
||||||
golang.org/x/sync v0.19.0 // indirect
|
golang.org/x/sync v0.19.0 // indirect
|
||||||
golang.org/x/sys v0.39.0 // indirect
|
golang.org/x/sys v0.39.0 // indirect
|
||||||
|
golang.org/x/text v0.32.0 // indirect
|
||||||
golang.org/x/tools v0.40.0 // indirect
|
golang.org/x/tools v0.40.0 // indirect
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -10,5 +10,7 @@ golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
|
|||||||
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||||
golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk=
|
golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk=
|
||||||
golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
||||||
|
golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU=
|
||||||
|
golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY=
|
||||||
golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
|
golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
|
||||||
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
|
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package dns
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
)
|
)
|
||||||
@@ -14,6 +15,8 @@ const (
|
|||||||
ResponseNODATA
|
ResponseNODATA
|
||||||
ResponseNXDOMAIN
|
ResponseNXDOMAIN
|
||||||
ResponseSERVFAIL
|
ResponseSERVFAIL
|
||||||
|
ResponseREFUSED
|
||||||
|
ResponseNOTIMPL
|
||||||
ResponseOther
|
ResponseOther
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -29,6 +32,10 @@ func (rc ResponseClassification) String() string {
|
|||||||
return "nxdomain"
|
return "nxdomain"
|
||||||
case ResponseSERVFAIL:
|
case ResponseSERVFAIL:
|
||||||
return "servfail"
|
return "servfail"
|
||||||
|
case ResponseREFUSED:
|
||||||
|
return "refused"
|
||||||
|
case ResponseNOTIMPL:
|
||||||
|
return "notimp"
|
||||||
default:
|
default:
|
||||||
return "other"
|
return "other"
|
||||||
}
|
}
|
||||||
@@ -45,6 +52,13 @@ type DecodedResponse struct {
|
|||||||
Authority []dns.RR
|
Authority []dns.RR
|
||||||
Additional []dns.RR
|
Additional []dns.RR
|
||||||
CNAMEChain []string
|
CNAMEChain []string
|
||||||
|
DNAMEMappings []DNAMEMapping
|
||||||
|
}
|
||||||
|
|
||||||
|
// DNAMEMapping holds a DNAME record's owner and target for redirect synthesis.
|
||||||
|
type DNAMEMapping struct {
|
||||||
|
Owner string // e.g., "example.com."
|
||||||
|
Target string // e.g., "example.net."
|
||||||
}
|
}
|
||||||
|
|
||||||
func DecodeResponse(msg *dns.Msg) *DecodedResponse {
|
func DecodeResponse(msg *dns.Msg) *DecodedResponse {
|
||||||
@@ -62,6 +76,7 @@ func DecodeResponse(msg *dns.Msg) *DecodedResponse {
|
|||||||
Authority: msg.Ns,
|
Authority: msg.Ns,
|
||||||
Additional: msg.Extra,
|
Additional: msg.Extra,
|
||||||
CNAMEChain: extractCNAMEChain(msg),
|
CNAMEChain: extractCNAMEChain(msg),
|
||||||
|
DNAMEMappings: extractDNAMEMappings(msg),
|
||||||
}
|
}
|
||||||
|
|
||||||
d.Classification = classify(msg)
|
d.Classification = classify(msg)
|
||||||
@@ -75,6 +90,10 @@ func classify(msg *dns.Msg) ResponseClassification {
|
|||||||
return ResponseNXDOMAIN
|
return ResponseNXDOMAIN
|
||||||
case dns.RcodeServerFailure:
|
case dns.RcodeServerFailure:
|
||||||
return ResponseSERVFAIL
|
return ResponseSERVFAIL
|
||||||
|
case dns.RcodeRefused:
|
||||||
|
return ResponseREFUSED
|
||||||
|
case dns.RcodeNotImplemented:
|
||||||
|
return ResponseNOTIMPL
|
||||||
case dns.RcodeSuccess:
|
case dns.RcodeSuccess:
|
||||||
return classifySuccess(msg)
|
return classifySuccess(msg)
|
||||||
default:
|
default:
|
||||||
@@ -125,6 +144,33 @@ func extractCNAMEChain(msg *dns.Msg) []string {
|
|||||||
return chain
|
return chain
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func extractDNAMEMappings(msg *dns.Msg) []DNAMEMapping {
|
||||||
|
var mappings []DNAMEMapping
|
||||||
|
for _, rr := range msg.Answer {
|
||||||
|
if dname, ok := rr.(*dns.DNAME); ok {
|
||||||
|
mappings = append(mappings, DNAMEMapping{
|
||||||
|
Owner: dns.Fqdn(dname.Hdr.Name),
|
||||||
|
Target: dns.Fqdn(dname.Target),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return mappings
|
||||||
|
}
|
||||||
|
|
||||||
|
// SynthesizeCNAMEFromDNAME computes the CNAME target for queryName given a DNAME mapping.
|
||||||
|
// Returns empty string if queryName is not a strict subdomain of dnameOwner.
|
||||||
|
func SynthesizeCNAMEFromDNAME(queryName, dnameOwner, dnameTarget string) string {
|
||||||
|
q := strings.ToLower(dns.Fqdn(queryName))
|
||||||
|
owner := strings.ToLower(dns.Fqdn(dnameOwner))
|
||||||
|
target := strings.ToLower(dns.Fqdn(dnameTarget))
|
||||||
|
|
||||||
|
if !dns.IsSubDomain(owner, q) || q == owner {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
prefix := strings.TrimSuffix(q, owner)
|
||||||
|
return prefix + target
|
||||||
|
}
|
||||||
|
|
||||||
func IsTruncated(msg *dns.Msg) bool {
|
func IsTruncated(msg *dns.Msg) bool {
|
||||||
return msg != nil && msg.Truncated
|
return msg != nil && msg.Truncated
|
||||||
}
|
}
|
||||||
|
|||||||
+16
-2
@@ -93,7 +93,7 @@ func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype ui
|
|||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err())
|
return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err())
|
||||||
case <-time.After(100 * time.Millisecond):
|
case <-time.After(backoffDelay(attempt)):
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -156,7 +156,7 @@ func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string,
|
|||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err())
|
return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err())
|
||||||
case <-time.After(100 * time.Millisecond):
|
case <-time.After(backoffDelay(attempt)):
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -195,6 +195,20 @@ func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string,
|
|||||||
return nil, fmt.Errorf("iterative query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr)
|
return nil, fmt.Errorf("iterative query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// backoffDelay computes the wait duration before the given retry attempt (1-indexed).
|
||||||
|
// Delays: attempt=1 → 100ms, attempt=2 → 200ms, attempt=3 → 400ms, capped at 2s.
|
||||||
|
func backoffDelay(attempt int) time.Duration {
|
||||||
|
if attempt <= 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
delay := time.Duration(uint(1)<<uint(attempt-1)) * 100 * time.Millisecond
|
||||||
|
const maxDelay = 2 * time.Second
|
||||||
|
if delay > maxDelay {
|
||||||
|
return maxDelay
|
||||||
|
}
|
||||||
|
return delay
|
||||||
|
}
|
||||||
|
|
||||||
func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg {
|
func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg {
|
||||||
m := new(dns.Msg)
|
m := new(dns.Msg)
|
||||||
m.SetQuestion(dns.Fqdn(name), qtype)
|
m.SetQuestion(dns.Fqdn(name), qtype)
|
||||||
|
|||||||
@@ -0,0 +1,139 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDecodeResponseREFUSED(t *testing.T) {
|
||||||
|
msg := newTestMsg(dns.RcodeRefused)
|
||||||
|
d := DecodeResponse(msg)
|
||||||
|
if d.Classification != ResponseREFUSED {
|
||||||
|
t.Errorf("classification = %v, want ResponseREFUSED", d.Classification)
|
||||||
|
}
|
||||||
|
if d.RcodeName != "REFUSED" {
|
||||||
|
t.Errorf("RcodeName = %q, want REFUSED", d.RcodeName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeResponseNOTIMPL(t *testing.T) {
|
||||||
|
msg := newTestMsg(dns.RcodeNotImplemented)
|
||||||
|
d := DecodeResponse(msg)
|
||||||
|
if d.Classification != ResponseNOTIMPL {
|
||||||
|
t.Errorf("classification = %v, want ResponseNOTIMPL", d.Classification)
|
||||||
|
}
|
||||||
|
if d.RcodeName != "NOTIMP" {
|
||||||
|
t.Errorf("RcodeName = %q, want NOTIMP", d.RcodeName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeResponseREFUSEDString(t *testing.T) {
|
||||||
|
if got := ResponseREFUSED.String(); got != "refused" {
|
||||||
|
t.Errorf("ResponseREFUSED.String() = %q, want \"refused\"", got)
|
||||||
|
}
|
||||||
|
if got := ResponseNOTIMPL.String(); got != "notimp" {
|
||||||
|
t.Errorf("ResponseNOTIMPL.String() = %q, want \"notimp\"", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractDNAMEMappings(t *testing.T) {
|
||||||
|
t.Run("no DNAME", func(t *testing.T) {
|
||||||
|
msg := new(dns.Msg)
|
||||||
|
msg.Answer = append(msg.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA},
|
||||||
|
A: MustParseIP("1.2.3.4"),
|
||||||
|
})
|
||||||
|
d := DecodeResponse(msg)
|
||||||
|
if len(d.DNAMEMappings) != 0 {
|
||||||
|
t.Errorf("expected 0 DNAME mappings, got %d", len(d.DNAMEMappings))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("DNAME in answer", func(t *testing.T) {
|
||||||
|
msg := new(dns.Msg)
|
||||||
|
msg.SetReply(new(dns.Msg))
|
||||||
|
msg.Answer = append(msg.Answer, &dns.DNAME{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeDNAME, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
Target: "example.net.",
|
||||||
|
})
|
||||||
|
d := DecodeResponse(msg)
|
||||||
|
if len(d.DNAMEMappings) != 1 {
|
||||||
|
t.Fatalf("expected 1 DNAME mapping, got %d", len(d.DNAMEMappings))
|
||||||
|
}
|
||||||
|
if d.DNAMEMappings[0].Owner != "example.com." {
|
||||||
|
t.Errorf("Owner = %q, want %q", d.DNAMEMappings[0].Owner, "example.com.")
|
||||||
|
}
|
||||||
|
if d.DNAMEMappings[0].Target != "example.net." {
|
||||||
|
t.Errorf("Target = %q, want %q", d.DNAMEMappings[0].Target, "example.net.")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSynthesizeCNAMEFromDNAME(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
queryName string
|
||||||
|
dnameOwner string
|
||||||
|
dnameTarget string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
queryName: "foo.example.com.",
|
||||||
|
dnameOwner: "example.com.",
|
||||||
|
dnameTarget: "example.net.",
|
||||||
|
want: "foo.example.net.",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
queryName: "bar.foo.example.com.",
|
||||||
|
dnameOwner: "example.com.",
|
||||||
|
dnameTarget: "example.net.",
|
||||||
|
want: "bar.foo.example.net.",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Owner itself is not redirected
|
||||||
|
queryName: "example.com.",
|
||||||
|
dnameOwner: "example.com.",
|
||||||
|
dnameTarget: "example.net.",
|
||||||
|
want: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Not a subdomain
|
||||||
|
queryName: "other.com.",
|
||||||
|
dnameOwner: "example.com.",
|
||||||
|
dnameTarget: "example.net.",
|
||||||
|
want: "",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
got := SynthesizeCNAMEFromDNAME(tt.queryName, tt.dnameOwner, tt.dnameTarget)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("SynthesizeCNAMEFromDNAME(%q, %q, %q) = %q, want %q",
|
||||||
|
tt.queryName, tt.dnameOwner, tt.dnameTarget, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBackoffDelay(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
attempt int
|
||||||
|
want time.Duration
|
||||||
|
}{
|
||||||
|
{0, 0},
|
||||||
|
{1, 100 * time.Millisecond},
|
||||||
|
{2, 200 * time.Millisecond},
|
||||||
|
{3, 400 * time.Millisecond},
|
||||||
|
{4, 800 * time.Millisecond},
|
||||||
|
{5, 1600 * time.Millisecond},
|
||||||
|
{6, 2000 * time.Millisecond}, // capped at 2s
|
||||||
|
{10, 2000 * time.Millisecond}, // still capped
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
got := backoffDelay(tt.attempt)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("backoffDelay(%d) = %v, want %v", tt.attempt, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -138,6 +138,12 @@ func summaryTypeLabel(respType string) string {
|
|||||||
return "name does not exist"
|
return "name does not exist"
|
||||||
case "servfail":
|
case "servfail":
|
||||||
return "resulted in SERVFAIL"
|
return "resulted in SERVFAIL"
|
||||||
|
case "refused":
|
||||||
|
return "query refused by server"
|
||||||
|
case "notimp":
|
||||||
|
return "query type not implemented by server"
|
||||||
|
case "cname_loop":
|
||||||
|
return "resulted in a CNAME loop"
|
||||||
case "error":
|
case "error":
|
||||||
return "resulted in an error"
|
return "resulted in an error"
|
||||||
case "referral":
|
case "referral":
|
||||||
|
|||||||
+15
-1
@@ -210,8 +210,22 @@ func (f *textFormatter) formatResultLine(result traverse.TraversalResult) string
|
|||||||
return f.colorize(fmt.Sprintf("%s name does not exist", prob), colorYellow)
|
return f.colorize(fmt.Sprintf("%s name does not exist", prob), colorYellow)
|
||||||
case traverse.RespSERVFAIL:
|
case traverse.RespSERVFAIL:
|
||||||
return f.colorize(fmt.Sprintf("%s resulted in SERVFAIL", prob), colorRed)
|
return f.colorize(fmt.Sprintf("%s resulted in SERVFAIL", prob), colorRed)
|
||||||
|
case traverse.RespREFUSED:
|
||||||
|
return f.colorize(fmt.Sprintf("%s query refused by server", prob), colorRed)
|
||||||
|
case traverse.RespNOTIMPL:
|
||||||
|
return f.colorize(fmt.Sprintf("%s query type not implemented by server", prob), colorRed)
|
||||||
|
case traverse.RespCNAMELoop:
|
||||||
|
msg := "CNAME loop detected"
|
||||||
|
if result.Response.ErrorMessage != "" {
|
||||||
|
msg = result.Response.ErrorMessage
|
||||||
|
}
|
||||||
|
return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorRed)
|
||||||
case traverse.RespError:
|
case traverse.RespError:
|
||||||
return f.colorize(fmt.Sprintf("%s resulted in an error", prob), colorRed)
|
msg := "resulted in an error"
|
||||||
|
if result.Response.ErrorMessage != "" {
|
||||||
|
msg = result.Response.ErrorMessage
|
||||||
|
}
|
||||||
|
return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorRed)
|
||||||
default:
|
default:
|
||||||
return fmt.Sprintf("%s %s", prob, result.Response.Type)
|
return fmt.Sprintf("%s %s", prob, result.Response.Type)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
|
|
||||||
"github.com/hits/ExploreDNS/internal/dns"
|
"github.com/hits/ExploreDNS/internal/dns"
|
||||||
miekgdns "github.com/miekg/dns"
|
miekgdns "github.com/miekg/dns"
|
||||||
|
"golang.org/x/net/idna"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ResolutionState int
|
type ResolutionState int
|
||||||
@@ -46,12 +47,35 @@ type Referral struct {
|
|||||||
Prob float64
|
Prob float64
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// idnaLookup is the IDN lookup profile used to convert internationalised domain
|
||||||
|
// names (unicode labels) to their ACE/punycode equivalents before querying.
|
||||||
|
var idnaLookup = idna.New(
|
||||||
|
idna.MapForLookup(),
|
||||||
|
idna.BidiRule(),
|
||||||
|
idna.StrictDomainName(false),
|
||||||
|
)
|
||||||
|
|
||||||
|
// toASCII converts a domain name that may contain unicode labels to its
|
||||||
|
// punycode (ACE) representation. Pure-ASCII names are returned unchanged.
|
||||||
|
// On conversion errors the original name is returned so the caller can still
|
||||||
|
// attempt a query (the server will reject it if truly invalid).
|
||||||
|
func toASCII(name string) string {
|
||||||
|
if name == "" || name == "." {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
ascii, err := idnaLookup.ToASCII(name)
|
||||||
|
if err != nil {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
return ascii
|
||||||
|
}
|
||||||
|
|
||||||
func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral {
|
func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral {
|
||||||
return &Referral{
|
return &Referral{
|
||||||
Name: miekgdns.Fqdn(strings.ToLower(name)),
|
Name: miekgdns.Fqdn(strings.ToLower(toASCII(name))),
|
||||||
Qtype: qtype,
|
Qtype: qtype,
|
||||||
Qclass: miekgdns.ClassINET,
|
Qclass: miekgdns.ClassINET,
|
||||||
Bailiwick: miekgdns.Fqdn(strings.ToLower(bailiwick)),
|
Bailiwick: miekgdns.Fqdn(strings.ToLower(toASCII(bailiwick))),
|
||||||
Depth: depth,
|
Depth: depth,
|
||||||
Prob: prob,
|
Prob: prob,
|
||||||
Parent: parent,
|
Parent: parent,
|
||||||
@@ -236,3 +260,17 @@ func getVisitedNames(visited map[string]bool) []string {
|
|||||||
}
|
}
|
||||||
return names
|
return names
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// IsNameInChain reports whether name appears anywhere in this referral's ancestor
|
||||||
|
// chain, including this referral itself. Used for CNAME loop detection.
|
||||||
|
func (r *Referral) IsNameInChain(name string) bool {
|
||||||
|
n := miekgdns.Fqdn(strings.ToLower(name))
|
||||||
|
curr := r
|
||||||
|
for curr != nil {
|
||||||
|
if curr.Name == n {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
curr = curr.Parent
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|||||||
@@ -16,6 +16,9 @@ const (
|
|||||||
RespNODATA
|
RespNODATA
|
||||||
RespNXDOMAIN
|
RespNXDOMAIN
|
||||||
RespSERVFAIL
|
RespSERVFAIL
|
||||||
|
RespREFUSED
|
||||||
|
RespNOTIMPL
|
||||||
|
RespCNAMELoop
|
||||||
RespError
|
RespError
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -33,6 +36,12 @@ func (rt ResponseType) String() string {
|
|||||||
return "nxdomain"
|
return "nxdomain"
|
||||||
case RespSERVFAIL:
|
case RespSERVFAIL:
|
||||||
return "servfail"
|
return "servfail"
|
||||||
|
case RespREFUSED:
|
||||||
|
return "refused"
|
||||||
|
case RespNOTIMPL:
|
||||||
|
return "notimp"
|
||||||
|
case RespCNAMELoop:
|
||||||
|
return "cname_loop"
|
||||||
case RespError:
|
case RespError:
|
||||||
return "error"
|
return "error"
|
||||||
default:
|
default:
|
||||||
@@ -41,11 +50,12 @@ func (rt ResponseType) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type Response struct {
|
type Response struct {
|
||||||
Referral *Referral
|
Referral *Referral
|
||||||
Server net.IP
|
Server net.IP
|
||||||
Cache *InfoCache
|
Cache *InfoCache
|
||||||
Decoded *dns.DecodedResponse
|
Decoded *dns.DecodedResponse
|
||||||
Type ResponseType
|
Type ResponseType
|
||||||
|
ErrorMessage string
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response {
|
func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response {
|
||||||
@@ -59,15 +69,28 @@ func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response {
|
|||||||
func (r *Response) Process(msg *miekgdns.Msg) *Response {
|
func (r *Response) Process(msg *miekgdns.Msg) *Response {
|
||||||
if msg == nil {
|
if msg == nil {
|
||||||
r.Type = RespError
|
r.Type = RespError
|
||||||
|
r.ErrorMessage = "nil DNS response"
|
||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
r.Decoded = dns.DecodeResponse(msg)
|
r.Decoded = dns.DecodeResponse(msg)
|
||||||
if r.Decoded == nil {
|
if r.Decoded == nil {
|
||||||
r.Type = RespError
|
r.Type = RespError
|
||||||
|
r.ErrorMessage = "failed to decode DNS response"
|
||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Synthesize CNAME from DNAME when the server didn't include a synthesized CNAME record.
|
||||||
|
if len(r.Decoded.CNAMEChain) == 0 && r.Referral != nil && len(r.Decoded.DNAMEMappings) > 0 {
|
||||||
|
for _, dm := range r.Decoded.DNAMEMappings {
|
||||||
|
synthesized := dns.SynthesizeCNAMEFromDNAME(r.Referral.Name, dm.Owner, dm.Target)
|
||||||
|
if synthesized != "" {
|
||||||
|
r.Decoded.CNAMEChain = append(r.Decoded.CNAMEChain, synthesized)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
r.Type = r.classify()
|
r.Type = r.classify()
|
||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
@@ -78,6 +101,10 @@ func (r *Response) classify() ResponseType {
|
|||||||
return RespNXDOMAIN
|
return RespNXDOMAIN
|
||||||
case dns.ResponseSERVFAIL:
|
case dns.ResponseSERVFAIL:
|
||||||
return RespSERVFAIL
|
return RespSERVFAIL
|
||||||
|
case dns.ResponseREFUSED:
|
||||||
|
return RespREFUSED
|
||||||
|
case dns.ResponseNOTIMPL:
|
||||||
|
return RespNOTIMPL
|
||||||
case dns.ResponseAnswer:
|
case dns.ResponseAnswer:
|
||||||
if len(r.Decoded.CNAMEChain) > 0 && !r.hasFinalAnswer() {
|
if len(r.Decoded.CNAMEChain) > 0 && !r.hasFinalAnswer() {
|
||||||
return RespCNAMEFollow
|
return RespCNAMEFollow
|
||||||
@@ -94,7 +121,11 @@ func (r *Response) classify() ResponseType {
|
|||||||
|
|
||||||
func (r *Response) hasFinalAnswer() bool {
|
func (r *Response) hasFinalAnswer() bool {
|
||||||
for _, rr := range r.Decoded.Answers {
|
for _, rr := range r.Decoded.Answers {
|
||||||
if _, ok := rr.(*miekgdns.CNAME); ok {
|
switch rr.(type) {
|
||||||
|
case *miekgdns.CNAME, *miekgdns.DNAME, *miekgdns.RRSIG:
|
||||||
|
// CNAME and DNAME are redirect records, not final answers.
|
||||||
|
// RRSIG is a DNSSEC signature record — it covers the CNAME/DNAME
|
||||||
|
// but is not itself the answer to the original question type.
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
return true
|
return true
|
||||||
@@ -200,7 +231,7 @@ func (r *Response) resolveGlue(child *Referral) {
|
|||||||
|
|
||||||
func (r *Response) IsTerminal() bool {
|
func (r *Response) IsTerminal() bool {
|
||||||
switch r.Type {
|
switch r.Type {
|
||||||
case RespAnswer, RespNODATA, RespNXDOMAIN, RespSERVFAIL, RespError:
|
case RespAnswer, RespNODATA, RespNXDOMAIN, RespSERVFAIL, RespREFUSED, RespNOTIMPL, RespCNAMELoop, RespError:
|
||||||
return true
|
return true
|
||||||
default:
|
default:
|
||||||
return false
|
return false
|
||||||
|
|||||||
@@ -0,0 +1,768 @@
|
|||||||
|
package traverse
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"net"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestCNAMELoopDetected verifies that a two-step CNAME loop (A → B → A) is
|
||||||
|
// detected without infinite recursion and produces a RespCNAMELoop result.
|
||||||
|
func TestCNAMELoopDetected(t *testing.T) {
|
||||||
|
// www.example.com → CNAME → alias.example.com → CNAME → www.example.com (loop)
|
||||||
|
cnameToAlias := new(dns.Msg)
|
||||||
|
cnameToAlias.SetReply(new(dns.Msg))
|
||||||
|
cnameToAlias.Answer = append(cnameToAlias.Answer, &dns.CNAME{
|
||||||
|
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
|
||||||
|
Target: "alias.example.com.",
|
||||||
|
})
|
||||||
|
|
||||||
|
cnameBack := new(dns.Msg)
|
||||||
|
cnameBack.SetReply(new(dns.Msg))
|
||||||
|
cnameBack.Answer = append(cnameBack.Answer, &dns.CNAME{
|
||||||
|
Hdr: dns.RR_Header{Name: "alias.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
|
||||||
|
Target: "www.example.com.",
|
||||||
|
})
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 10,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
q := msg.Question[0]
|
||||||
|
switch q.Name {
|
||||||
|
case "www.example.com.":
|
||||||
|
return cnameToAlias.Copy(), nil
|
||||||
|
case "alias.example.com.":
|
||||||
|
return cnameBack.Copy(), nil
|
||||||
|
}
|
||||||
|
return nil, nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "www.example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
foundLoop := false
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response != nil && r.Response.Type == RespCNAMELoop {
|
||||||
|
foundLoop = true
|
||||||
|
if r.Response.ErrorMessage == "" {
|
||||||
|
t.Error("expected non-empty ErrorMessage on CNAME loop result")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !foundLoop {
|
||||||
|
t.Error("expected RespCNAMELoop result for CNAME loop A → B → A")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCNAMEDirectLoop verifies that a direct self-loop (A → A) is handled.
|
||||||
|
func TestCNAMEDirectLoop(t *testing.T) {
|
||||||
|
selfLoop := new(dns.Msg)
|
||||||
|
selfLoop.SetReply(new(dns.Msg))
|
||||||
|
selfLoop.Answer = append(selfLoop.Answer, &dns.CNAME{
|
||||||
|
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET},
|
||||||
|
Target: "www.example.com.",
|
||||||
|
})
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 10,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return selfLoop.Copy(), nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "www.example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
foundLoop := false
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response != nil && r.Response.Type == RespCNAMELoop {
|
||||||
|
foundLoop = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !foundLoop {
|
||||||
|
t.Error("expected RespCNAMELoop for direct self-referencing CNAME")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestREFUSEDResponse verifies that a REFUSED rcode is classified as RespREFUSED.
|
||||||
|
func TestREFUSEDResponse(t *testing.T) {
|
||||||
|
refusedResp := new(dns.Msg)
|
||||||
|
refusedResp.Rcode = dns.RcodeRefused
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return refusedResp.Copy(), nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(results) == 0 {
|
||||||
|
t.Fatal("expected at least 1 result")
|
||||||
|
}
|
||||||
|
if results[0].Response.Type != RespREFUSED {
|
||||||
|
t.Errorf("Type = %s, want refused", results[0].Response.Type)
|
||||||
|
}
|
||||||
|
if !results[0].Response.IsTerminal() {
|
||||||
|
t.Error("REFUSED should be a terminal response")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestNOTIMPLResponse verifies that a NOTIMP rcode is classified as RespNOTIMPL.
|
||||||
|
func TestNOTIMPLResponse(t *testing.T) {
|
||||||
|
notImplResp := new(dns.Msg)
|
||||||
|
notImplResp.Rcode = dns.RcodeNotImplemented
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return notImplResp.Copy(), nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(results) == 0 {
|
||||||
|
t.Fatal("expected at least 1 result")
|
||||||
|
}
|
||||||
|
if results[0].Response.Type != RespNOTIMPL {
|
||||||
|
t.Errorf("Type = %s, want notimp", results[0].Response.Type)
|
||||||
|
}
|
||||||
|
if !results[0].Response.IsTerminal() {
|
||||||
|
t.Error("NOTIMP should be a terminal response")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGracefulDegradationUnreachableServer verifies that when some servers are
|
||||||
|
// unreachable, traversal continues with the remaining servers and does not panic.
|
||||||
|
func TestGracefulDegradationUnreachableServer(t *testing.T) {
|
||||||
|
answerResp := new(dns.Msg)
|
||||||
|
answerResp.SetReply(new(dns.Msg))
|
||||||
|
answerResp.Answer = append(answerResp.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
A: net.ParseIP("93.184.216.34"),
|
||||||
|
})
|
||||||
|
|
||||||
|
// Referral with two nameservers; first always fails, second provides the answer.
|
||||||
|
referralMsg := new(dns.Msg)
|
||||||
|
referralMsg.Rcode = dns.RcodeSuccess
|
||||||
|
referralMsg.Authoritative = false
|
||||||
|
referralMsg.Ns = append(referralMsg.Ns,
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."},
|
||||||
|
)
|
||||||
|
referralMsg.Extra = append(referralMsg.Extra,
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")},
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")},
|
||||||
|
)
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
switch server {
|
||||||
|
case "198.41.0.4":
|
||||||
|
return referralMsg.Copy(), nil
|
||||||
|
case "10.0.0.1":
|
||||||
|
return nil, errors.New("connection refused")
|
||||||
|
case "10.0.0.2":
|
||||||
|
return answerResp.Copy(), nil
|
||||||
|
}
|
||||||
|
return nil, errors.New("unexpected server")
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
foundAnswer := false
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response != nil && r.Response.Type == RespAnswer {
|
||||||
|
foundAnswer = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !foundAnswer {
|
||||||
|
t.Error("expected an answer result from the reachable server")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestGracefulDegradationAllUnreachable verifies that when ALL servers fail,
|
||||||
|
// the traversal returns a SERVFAIL result without panicking.
|
||||||
|
func TestGracefulDegradationAllUnreachable(t *testing.T) {
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return nil, errors.New("network unreachable")
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("traversal must not return a top-level error: %v", err)
|
||||||
|
}
|
||||||
|
if len(results) == 0 {
|
||||||
|
t.Fatal("expected at least one result even on total failure")
|
||||||
|
}
|
||||||
|
last := results[len(results)-1]
|
||||||
|
if last.Response == nil {
|
||||||
|
t.Fatal("last result must have a response")
|
||||||
|
}
|
||||||
|
if last.Response.Type != RespSERVFAIL && last.Response.Type != RespError {
|
||||||
|
t.Errorf("expected SERVFAIL or error when all servers unreachable, got %s", last.Response.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDNAMEFollowNoSynthesizedCNAME verifies that a DNAME record in the answer
|
||||||
|
// section synthesizes a CNAME follow when the server doesn't include one.
|
||||||
|
func TestDNAMEFollowNoSynthesizedCNAME(t *testing.T) {
|
||||||
|
// Server returns DNAME only (no synthesized CNAME).
|
||||||
|
dnameResp := new(dns.Msg)
|
||||||
|
dnameResp.SetReply(new(dns.Msg))
|
||||||
|
dnameResp.Answer = append(dnameResp.Answer, &dns.DNAME{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeDNAME, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
Target: "example.net.",
|
||||||
|
})
|
||||||
|
|
||||||
|
answerResp := new(dns.Msg)
|
||||||
|
answerResp.SetReply(new(dns.Msg))
|
||||||
|
answerResp.Answer = append(answerResp.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "www.example.net.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
A: net.ParseIP("203.0.113.1"),
|
||||||
|
})
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
q := msg.Question[0]
|
||||||
|
if q.Name == "www.example.com." {
|
||||||
|
return dnameResp.Copy(), nil
|
||||||
|
}
|
||||||
|
if q.Name == "www.example.net." {
|
||||||
|
return answerResp.Copy(), nil
|
||||||
|
}
|
||||||
|
return nil, nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "www.example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
foundCNAMEFollow := false
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response != nil && r.Response.Type == RespCNAMEFollow {
|
||||||
|
foundCNAMEFollow = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !foundCNAMEFollow {
|
||||||
|
t.Error("expected RespCNAMEFollow synthesized from DNAME record")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestIsNameInChain verifies the ancestor chain lookup.
|
||||||
|
func TestIsNameInChain(t *testing.T) {
|
||||||
|
root := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil)
|
||||||
|
child := NewReferral("www.example.com", dnsTypeA, "example.com.", 1, 1.0, root)
|
||||||
|
grandchild := NewReferral("sub.www.example.com", dnsTypeA, "www.example.com.", 2, 1.0, child)
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
ref *Referral
|
||||||
|
name string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{grandchild, "sub.www.example.com", true}, // self
|
||||||
|
{grandchild, "www.example.com", true}, // parent
|
||||||
|
{grandchild, "example.com", true}, // grandparent
|
||||||
|
{grandchild, "other.example.com", false}, // not in chain
|
||||||
|
{root, "example.com", true}, // root matches itself
|
||||||
|
{root, "www.example.com", false}, // child not in chain from root
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
got := tt.ref.IsNameInChain(tt.name)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("IsNameInChain(%q) from %q = %v, want %v", tt.name, tt.ref.Name, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestResponseTypeStrings verifies String() for new response types.
|
||||||
|
func TestResponseTypeStrings(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
rt ResponseType
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{RespReferral, "referral"},
|
||||||
|
{RespAnswer, "answer"},
|
||||||
|
{RespCNAMEFollow, "cname_follow"},
|
||||||
|
{RespNODATA, "nodata"},
|
||||||
|
{RespNXDOMAIN, "nxdomain"},
|
||||||
|
{RespSERVFAIL, "servfail"},
|
||||||
|
{RespREFUSED, "refused"},
|
||||||
|
{RespNOTIMPL, "notimp"},
|
||||||
|
{RespCNAMELoop, "cname_loop"},
|
||||||
|
{RespError, "error"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := tt.rt.String(); got != tt.want {
|
||||||
|
t.Errorf("ResponseType(%d).String() = %q, want %q", tt.rt, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMalformedResponseNoPanic verifies that a nil response from the exchange
|
||||||
|
// function does not cause a panic, and produces an error result.
|
||||||
|
func TestMalformedResponseNoPanic(t *testing.T) {
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return nil, nil // nil response, no error
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(results) == 0 {
|
||||||
|
t.Fatal("expected at least one result")
|
||||||
|
}
|
||||||
|
// Should produce error/servfail, not panic
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response == nil {
|
||||||
|
t.Error("result has nil response")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDNSSECRRSIGDoesNotBlockCNAMEFollow verifies that a DNSSEC RRSIG record
|
||||||
|
// accompanying a CNAME in the answer section is treated as metadata and does
|
||||||
|
// NOT prevent the traversal from following the CNAME.
|
||||||
|
func TestDNSSECRRSIGDoesNotBlockCNAMEFollow(t *testing.T) {
|
||||||
|
// Server returns CNAME + RRSIG (DNSSEC-signed zone response).
|
||||||
|
cnameWithRRSIG := new(dns.Msg)
|
||||||
|
cnameWithRRSIG.SetReply(new(dns.Msg))
|
||||||
|
cnameWithRRSIG.Answer = append(cnameWithRRSIG.Answer,
|
||||||
|
&dns.CNAME{
|
||||||
|
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
Target: "example.com.",
|
||||||
|
},
|
||||||
|
&dns.RRSIG{
|
||||||
|
Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeRRSIG, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
TypeCovered: dns.TypeCNAME,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
finalAnswer := new(dns.Msg)
|
||||||
|
finalAnswer.SetReply(new(dns.Msg))
|
||||||
|
finalAnswer.Answer = append(finalAnswer.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
A: net.ParseIP("93.184.216.34"),
|
||||||
|
})
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 10,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
Fast: true,
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
q := msg.Question[0]
|
||||||
|
if q.Name == "www.example.com." {
|
||||||
|
return cnameWithRRSIG.Copy(), nil
|
||||||
|
}
|
||||||
|
return finalAnswer.Copy(), nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "www.example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
foundCNAMEFollow := false
|
||||||
|
foundAnswer := false
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response != nil {
|
||||||
|
switch r.Response.Type {
|
||||||
|
case RespCNAMEFollow:
|
||||||
|
foundCNAMEFollow = true
|
||||||
|
case RespAnswer:
|
||||||
|
foundAnswer = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !foundCNAMEFollow {
|
||||||
|
t.Error("expected RespCNAMEFollow: RRSIG should not block CNAME following")
|
||||||
|
}
|
||||||
|
if !foundAnswer {
|
||||||
|
t.Error("expected final RespAnswer after CNAME follow")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFastModeOn verifies that Fast=true uses the shared root cache (default
|
||||||
|
// behaviour): a child branch can see glue stored by the root referral.
|
||||||
|
func TestFastModeOn(t *testing.T) {
|
||||||
|
// Root referral returns two nameservers with glue. Each NS branch returns
|
||||||
|
// an answer. We verify both branches are queried.
|
||||||
|
referralMsg := new(dns.Msg)
|
||||||
|
referralMsg.Rcode = dns.RcodeSuccess
|
||||||
|
referralMsg.Authoritative = false
|
||||||
|
referralMsg.Ns = append(referralMsg.Ns,
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."},
|
||||||
|
)
|
||||||
|
referralMsg.Extra = append(referralMsg.Extra,
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")},
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")},
|
||||||
|
)
|
||||||
|
|
||||||
|
answerMsg := new(dns.Msg)
|
||||||
|
answerMsg.SetReply(new(dns.Msg))
|
||||||
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
A: net.ParseIP("93.184.216.34"),
|
||||||
|
})
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
Fast: true,
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
if server == "198.41.0.4" {
|
||||||
|
return referralMsg.Copy(), nil
|
||||||
|
}
|
||||||
|
return answerMsg.Copy(), nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
answers := 0
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response != nil && r.Response.Type == RespAnswer {
|
||||||
|
answers++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if answers == 0 {
|
||||||
|
t.Error("expected at least one answer with Fast=true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFastModeOff verifies that Fast=false gives each referral its own
|
||||||
|
// independent cache — no cross-branch glue contamination.
|
||||||
|
func TestFastModeOff(t *testing.T) {
|
||||||
|
referralMsg := new(dns.Msg)
|
||||||
|
referralMsg.Rcode = dns.RcodeSuccess
|
||||||
|
referralMsg.Authoritative = false
|
||||||
|
referralMsg.Ns = append(referralMsg.Ns,
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."},
|
||||||
|
)
|
||||||
|
referralMsg.Extra = append(referralMsg.Extra,
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")},
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")},
|
||||||
|
)
|
||||||
|
|
||||||
|
answerMsg := new(dns.Msg)
|
||||||
|
answerMsg.SetReply(new(dns.Msg))
|
||||||
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
A: net.ParseIP("93.184.216.34"),
|
||||||
|
})
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
Fast: false,
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
if server == "198.41.0.4" {
|
||||||
|
return referralMsg.Copy(), nil
|
||||||
|
}
|
||||||
|
return answerMsg.Copy(), nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
// Traversal must complete without panic and produce results.
|
||||||
|
if len(results) == 0 {
|
||||||
|
t.Fatal("expected at least one result with Fast=false")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFastModeDefaultIsTrue verifies that DefaultTraverserConfig has Fast=true.
|
||||||
|
func TestFastModeDefaultIsTrue(t *testing.T) {
|
||||||
|
cfg := DefaultTraverserConfig()
|
||||||
|
if !cfg.Fast {
|
||||||
|
t.Error("DefaultTraverserConfig().Fast should be true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestManyNSRecords verifies that a referral with more than 10 nameservers is
|
||||||
|
// handled gracefully — no panics, results are produced.
|
||||||
|
func TestManyNSRecords(t *testing.T) {
|
||||||
|
referralMsg := new(dns.Msg)
|
||||||
|
referralMsg.Rcode = dns.RcodeSuccess
|
||||||
|
referralMsg.Authoritative = false
|
||||||
|
for i := 1; i <= 12; i++ {
|
||||||
|
ns := &dns.NS{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS},
|
||||||
|
Ns: net.ParseIP(string(rune('a'+i-1))).String() + ".ns.example.com.",
|
||||||
|
}
|
||||||
|
// Use a distinct IP for each NS so glue is resolved.
|
||||||
|
ip := net.IP{10, 0, 0, byte(i)}
|
||||||
|
referralMsg.Ns = append(referralMsg.Ns, &dns.NS{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS},
|
||||||
|
Ns: ns.Ns,
|
||||||
|
})
|
||||||
|
referralMsg.Extra = append(referralMsg.Extra, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: ns.Ns, Rrtype: dnsTypeA},
|
||||||
|
A: ip,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
answerMsg := new(dns.Msg)
|
||||||
|
answerMsg.SetReply(new(dns.Msg))
|
||||||
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
A: net.ParseIP("93.184.216.34"),
|
||||||
|
})
|
||||||
|
|
||||||
|
var queries int64
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
Fast: true,
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
atomic.AddInt64(&queries, 1)
|
||||||
|
if server == "198.41.0.4" {
|
||||||
|
return referralMsg.Copy(), nil
|
||||||
|
}
|
||||||
|
return answerMsg.Copy(), nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error with 12 NS records: %v", err)
|
||||||
|
}
|
||||||
|
if len(results) == 0 {
|
||||||
|
t.Fatal("expected results with many NS records")
|
||||||
|
}
|
||||||
|
foundAnswer := false
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response != nil && r.Response.Type == RespAnswer {
|
||||||
|
foundAnswer = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !foundAnswer {
|
||||||
|
t.Error("expected at least one answer from the 12-NS referral")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestIDNPunycodeConversion verifies that a unicode (IDN) domain name is
|
||||||
|
// converted to its punycode/ACE form before querying.
|
||||||
|
func TestIDNPunycodeConversion(t *testing.T) {
|
||||||
|
// "münchen.de" → "xn--mnchen-3ya.de" (after punycode encoding)
|
||||||
|
ref := NewReferral("münchen.de", dnsTypeA, ".", 0, 1.0, nil)
|
||||||
|
if ref.Name == "münchen.de." {
|
||||||
|
t.Errorf("IDN name was not converted to punycode: got %q", ref.Name)
|
||||||
|
}
|
||||||
|
// Verify it starts with the expected punycode label.
|
||||||
|
if ref.Name != "xn--mnchen-3ya.de." {
|
||||||
|
t.Errorf("unexpected punycode result: got %q, want %q", ref.Name, "xn--mnchen-3ya.de.")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestASCIIDomainUnchanged verifies that a plain ASCII domain is not mangled
|
||||||
|
// by the IDN conversion path.
|
||||||
|
func TestASCIIDomainUnchanged(t *testing.T) {
|
||||||
|
ref := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil)
|
||||||
|
if ref.Name != "example.com." {
|
||||||
|
t.Errorf("ASCII domain was mangled: got %q, want %q", ref.Name, "example.com.")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWildcardResponse verifies that a wildcard answer (e.g. *.example.com
|
||||||
|
// returning an A record for sub.example.com) is handled as a regular answer.
|
||||||
|
func TestWildcardResponse(t *testing.T) {
|
||||||
|
wildcardAnswer := new(dns.Msg)
|
||||||
|
wildcardAnswer.SetReply(new(dns.Msg))
|
||||||
|
wildcardAnswer.Authoritative = true
|
||||||
|
wildcardAnswer.Answer = append(wildcardAnswer.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "sub.example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
A: net.ParseIP("1.2.3.4"),
|
||||||
|
})
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
return wildcardAnswer.Copy(), nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "sub.example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(results) == 0 {
|
||||||
|
t.Fatal("expected results for wildcard response")
|
||||||
|
}
|
||||||
|
if results[0].Response.Type != RespAnswer {
|
||||||
|
t.Errorf("Type = %s, want answer", results[0].Response.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLongCNAMEChainDepthLimit verifies that a very long CNAME chain is
|
||||||
|
// terminated by the MaxDepth limit without infinite recursion or a panic.
|
||||||
|
func TestLongCNAMEChainDepthLimit(t *testing.T) {
|
||||||
|
// Every query returns a CNAME to the next label. The MaxDepth setting
|
||||||
|
// must stop the chain.
|
||||||
|
counter := 0
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
counter++
|
||||||
|
q := msg.Question[0]
|
||||||
|
resp := new(dns.Msg)
|
||||||
|
resp.SetReply(msg)
|
||||||
|
next := "next" + q.Name
|
||||||
|
resp.Answer = append(resp.Answer, &dns.CNAME{
|
||||||
|
Hdr: dns.RR_Header{Name: q.Name, Rrtype: dnsTypeCNAME, Class: dns.ClassINET},
|
||||||
|
Target: next,
|
||||||
|
})
|
||||||
|
return resp, nil
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "start.example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(results) == 0 {
|
||||||
|
t.Fatal("expected results")
|
||||||
|
}
|
||||||
|
// Traversal must have stopped — counter should not be unbounded.
|
||||||
|
if counter > 50 {
|
||||||
|
t.Errorf("too many exchange calls (%d): chain depth limit not enforced", counter)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestPartialBranchFailureReturnsResults verifies the graceful degradation
|
||||||
|
// requirement: when some NS branches fail completely, the partial results from
|
||||||
|
// successful branches are still returned.
|
||||||
|
func TestPartialBranchFailureReturnsResults(t *testing.T) {
|
||||||
|
// Three nameservers: first two error, third succeeds.
|
||||||
|
referralMsg := new(dns.Msg)
|
||||||
|
referralMsg.Rcode = dns.RcodeSuccess
|
||||||
|
referralMsg.Authoritative = false
|
||||||
|
referralMsg.Ns = append(referralMsg.Ns,
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."},
|
||||||
|
&dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns3.example.com."},
|
||||||
|
)
|
||||||
|
referralMsg.Extra = append(referralMsg.Extra,
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")},
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")},
|
||||||
|
&dns.A{Hdr: dns.RR_Header{Name: "ns3.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.3")},
|
||||||
|
)
|
||||||
|
|
||||||
|
answerMsg := new(dns.Msg)
|
||||||
|
answerMsg.SetReply(new(dns.Msg))
|
||||||
|
answerMsg.Answer = append(answerMsg.Answer, &dns.A{
|
||||||
|
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300},
|
||||||
|
A: net.ParseIP("93.184.216.34"),
|
||||||
|
})
|
||||||
|
|
||||||
|
tr := NewTraverser(&TraverserConfig{
|
||||||
|
MaxDepth: 5,
|
||||||
|
QueryType: dnsTypeA,
|
||||||
|
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
|
||||||
|
})
|
||||||
|
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||||
|
switch server {
|
||||||
|
case "198.41.0.4":
|
||||||
|
return referralMsg.Copy(), nil
|
||||||
|
case "10.0.0.1", "10.0.0.2":
|
||||||
|
return nil, errors.New("server unreachable")
|
||||||
|
case "10.0.0.3":
|
||||||
|
return answerMsg.Copy(), nil
|
||||||
|
}
|
||||||
|
return nil, errors.New("unexpected server")
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
results, err := tr.Traverse(ctx, "example.com")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("traversal must not return a top-level error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
foundAnswer := false
|
||||||
|
for _, r := range results {
|
||||||
|
if r.Response != nil && r.Response.Type == RespAnswer {
|
||||||
|
foundAnswer = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !foundAnswer {
|
||||||
|
t.Error("expected an answer from the third (reachable) nameserver despite others failing")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/hits/ExploreDNS/internal/dns"
|
"github.com/hits/ExploreDNS/internal/dns"
|
||||||
miekgdns "github.com/miekg/dns"
|
miekgdns "github.com/miekg/dns"
|
||||||
@@ -17,6 +18,12 @@ type TraverserConfig struct {
|
|||||||
QueryConfig *dns.QueryConfig
|
QueryConfig *dns.QueryConfig
|
||||||
RootAddrs []net.IP
|
RootAddrs []net.IP
|
||||||
Hooks *TraverserHooks
|
Hooks *TraverserHooks
|
||||||
|
// Fast controls cache sharing across branches. When true (default), child
|
||||||
|
// branches inherit glue discovered by earlier branches via the shared root
|
||||||
|
// cache, trading accuracy for speed. When false, each branch gets a
|
||||||
|
// completely independent cache — slower but results are not contaminated by
|
||||||
|
// sibling branch observations.
|
||||||
|
Fast bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func DefaultTraverserConfig() *TraverserConfig {
|
func DefaultTraverserConfig() *TraverserConfig {
|
||||||
@@ -26,6 +33,7 @@ func DefaultTraverserConfig() *TraverserConfig {
|
|||||||
RootConfig: nil,
|
RootConfig: nil,
|
||||||
QueryConfig: nil,
|
QueryConfig: nil,
|
||||||
RootAddrs: nil,
|
RootAddrs: nil,
|
||||||
|
Fast: true,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -97,9 +105,19 @@ func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResul
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
cache := rootCache
|
var cache *InfoCache
|
||||||
if ref.Parent != nil {
|
if t.config.Fast {
|
||||||
cache = rootCache.Child()
|
// Fast mode: inherit glue from the shared root cache so earlier
|
||||||
|
// branch discoveries are visible to later branches.
|
||||||
|
cache = rootCache
|
||||||
|
if ref.Parent != nil {
|
||||||
|
cache = rootCache.Child()
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Non-fast mode: every referral gets its own independent cache so
|
||||||
|
// no cross-branch glue is reused, ensuring each path is resolved
|
||||||
|
// from scratch.
|
||||||
|
cache = NewInfoCache(nil)
|
||||||
}
|
}
|
||||||
|
|
||||||
if t.config.Hooks != nil {
|
if t.config.Hooks != nil {
|
||||||
@@ -141,7 +159,19 @@ func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResul
|
|||||||
if resp.Type == RespCNAMEFollow {
|
if resp.Type == RespCNAMEFollow {
|
||||||
follow := resp.CNAMEFollowReferral()
|
follow := resp.CNAMEFollowReferral()
|
||||||
if follow != nil {
|
if follow != nil {
|
||||||
if !stack.Push(follow) {
|
// Detect CNAME loop: target name already appears in the ancestor chain.
|
||||||
|
if follow.Parent != nil && follow.Parent.IsNameInChain(follow.Name) {
|
||||||
|
mu.Lock()
|
||||||
|
results = append(results, TraversalResult{
|
||||||
|
Referral: follow,
|
||||||
|
Response: &Response{
|
||||||
|
Referral: follow,
|
||||||
|
Type: RespCNAMELoop,
|
||||||
|
ErrorMessage: fmt.Sprintf("CNAME loop detected: %s already in traversal chain", follow.Name),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
mu.Unlock()
|
||||||
|
} else if !stack.Push(follow) {
|
||||||
mu.Lock()
|
mu.Lock()
|
||||||
results = append(results, TraversalResult{
|
results = append(results, TraversalResult{
|
||||||
Referral: follow,
|
Referral: follow,
|
||||||
@@ -412,12 +442,16 @@ func (t *Traverser) resolveGlueViaSystem(ctx context.Context, name string, cache
|
|||||||
|
|
||||||
c := &miekgdns.Client{
|
c := &miekgdns.Client{
|
||||||
Net: "udp",
|
Net: "udp",
|
||||||
ReadTimeout: 5,
|
ReadTimeout: 5 * time.Second,
|
||||||
WriteTimeout: 5,
|
WriteTimeout: 5 * time.Second,
|
||||||
}
|
}
|
||||||
if deadline, ok := ctx.Deadline(); ok {
|
if deadline, ok := ctx.Deadline(); ok {
|
||||||
c.ReadTimeout = deadline.Sub(deadline)
|
remaining := time.Until(deadline)
|
||||||
c.WriteTimeout = deadline.Sub(deadline)
|
if remaining <= 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
c.ReadTimeout = remaining
|
||||||
|
c.WriteTimeout = remaining
|
||||||
}
|
}
|
||||||
|
|
||||||
fqdn := miekgdns.Fqdn(name)
|
fqdn := miekgdns.Fqdn(name)
|
||||||
|
|||||||
Reference in New Issue
Block a user