Phase 4.1: Error handling, edge cases, and robustness #11

Merged
multica-agent merged 2 commits from feature/phase-4.1-error-handling into main 2026-06-07 17:24:02 +00:00
9 changed files with 670 additions and 14 deletions
Showing only changes of commit 368f200d23 - Show all commits
+46
View File
@@ -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
View File
@@ -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)
+139
View File
@@ -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)
}
}
}
+6
View File
@@ -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
View File
@@ -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)
} }
+14
View File
@@ -236,3 +236,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
}
+36 -6
View File
@@ -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
@@ -97,6 +124,9 @@ func (r *Response) hasFinalAnswer() bool {
if _, ok := rr.(*miekgdns.CNAME); ok { if _, ok := rr.(*miekgdns.CNAME); ok {
continue continue
} }
if _, ok := rr.(*miekgdns.DNAME); ok {
continue
}
return true return true
} }
return false return false
@@ -200,7 +230,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
+376
View File
@@ -0,0 +1,376 @@
package traverse
import (
"context"
"errors"
"net"
"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")
}
}
}
+22 -5
View File
@@ -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"
@@ -141,7 +142,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 +425,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)