feat: implement referral resolution for nameserver names without glue records (HAN-381) #7

Merged
multica-agent merged 1 commits from agent/go-expert-developer/43a58bd0 into main 2026-06-07 06:09:50 +00:00
3 changed files with 360 additions and 3 deletions
Showing only changes of commit 4e806ce5e3 - Show all commits
+160
View File
@@ -1,9 +1,12 @@
package traverse
import (
"context"
"fmt"
"net"
"strings"
"github.com/hits/ExploreDNS/internal/dns"
miekgdns "github.com/miekg/dns"
)
@@ -76,3 +79,160 @@ func (r *Referral) SetAddresses(addrs []net.IP) {
r.State = StateUnresolved
}
}
type CircularReferralError struct {
Name string
Chain []string
}
func (e *CircularReferralError) Error() string {
return fmt.Sprintf("circular referral detected for %s: %v", e.Name, e.Chain)
}
type UnresolvableNameserverError struct {
Name string
Reason string
}
func (e *UnresolvableNameserverError) Error() string {
return fmt.Sprintf("unresolvable nameserver %s: %s", e.Name, e.Reason)
}
func (r *Referral) Resolve(ctx context.Context, traverser *Traverser, cache *InfoCache, visited map[string]bool, depth int) error {
if r.HasAddresses() {
r.State = StateResolved
return nil
}
if cache != nil {
if addrs := cache.LookupGlue(r.Name); len(addrs) > 0 {
r.Addresses = addrs
r.State = StateResolved
return nil
}
}
if visited != nil {
if visited[r.Name] {
return &CircularReferralError{
Name: r.Name,
Chain: getVisitedNames(visited),
}
}
visited[r.Name] = true
}
if depth > DefaultMaxDepth {
return &UnresolvableNameserverError{
Name: r.Name,
Reason: "max depth exceeded",
}
}
roots, err := traverser.discoverRoots(ctx)
if err != nil {
return fmt.Errorf("root discovery: %w", err)
}
initial := NewReferral(r.Name, dns.TypeA, ".", 0, 1.0, nil)
initial.Addresses = roots
initial.State = StateResolved
stack := NewStack(DefaultMaxDepth)
stack.Push(initial)
traversalCache := NewInfoCache(nil)
if visited != nil {
for name := range visited {
traversalCache.StoreGlue(name, []net.IP{})
}
}
var lastErr error
for {
select {
case <-ctx.Done():
return fmt.Errorf("resolution cancelled: %w", ctx.Err())
default:
}
ref := stack.Pop()
if ref == nil {
break
}
cacheForStep := traversalCache
if ref.Parent != nil {
cacheForStep = traversalCache.Child()
}
resp := traverser.processReferral(ctx, ref, cacheForStep)
if resp.Type == RespAnswer {
if len(resp.Decoded.Answers) > 0 {
var addrs []net.IP
for _, rr := range resp.Decoded.Answers {
if a, ok := rr.(*miekgdns.A); ok {
addrs = append(addrs, a.A)
}
if aaaa, ok := rr.(*miekgdns.AAAA); ok {
addrs = append(addrs, aaaa.AAAA)
}
}
if len(addrs) > 0 {
r.Addresses = addrs
r.State = StateResolved
if cache != nil {
cache.StoreGlue(r.Name, addrs)
}
return nil
}
}
}
if resp.Type == RespNXDOMAIN {
lastErr = &UnresolvableNameserverError{
Name: r.Name,
Reason: "NXDOMAIN",
}
break
}
if resp.Type == RespSERVFAIL || resp.Type == RespError {
lastErr = fmt.Errorf("server error resolving %s: %s", r.Name, resp.Type)
continue
}
if resp.Type == RespReferral {
children := resp.ChildReferrals()
for _, child := range children {
if visited != nil && visited[child.Name] {
continue
}
if !stack.Push(child) {
lastErr = &UnresolvableNameserverError{
Name: r.Name,
Reason: "max depth exceeded during resolution",
}
}
}
}
}
if lastErr != nil {
return lastErr
}
return &UnresolvableNameserverError{
Name: r.Name,
Reason: "resolution exhausted without answer",
}
}
func getVisitedNames(visited map[string]bool) []string {
var names []string
for name := range visited {
names = append(names, name)
}
return names
}
+48
View File
@@ -109,3 +109,51 @@ func TestResolutionStateString(t *testing.T) {
})
}
}
func TestCircularReferralError(t *testing.T) {
err := &CircularReferralError{
Name: "ns.example.com.",
Chain: []string{"ns1.example.com.", "ns2.example.com."},
}
expected := "circular referral detected for ns.example.com.: [ns1.example.com. ns2.example.com.]"
if err.Error() != expected {
t.Errorf("Error() = %q, want %q", err.Error(), expected)
}
}
func TestUnresolvableNameserverError(t *testing.T) {
err := &UnresolvableNameserverError{
Name: "ns.example.com.",
Reason: "NXDOMAIN",
}
expected := "unresolvable nameserver ns.example.com.: NXDOMAIN"
if err.Error() != expected {
t.Errorf("Error() = %q, want %q", err.Error(), expected)
}
}
func TestGetVisitedNames(t *testing.T) {
visited := map[string]bool{
"ns1.example.com.": true,
"ns2.example.com.": true,
"ns3.example.com.": true,
}
names := getVisitedNames(visited)
if len(names) != 3 {
t.Errorf("got %d names, want 3", len(names))
}
seen := make(map[string]bool)
for _, name := range names {
if seen[name] {
t.Errorf("duplicate name: %s", name)
}
seen[name] = true
if !visited[name] {
t.Errorf("unexpected name: %s", name)
}
}
}
+152 -3
View File
@@ -36,6 +36,9 @@ type TraversalResult struct {
type Traverser struct {
config *TraverserConfig
exchange dns.ExchangeFunc
visited map[string]bool
depth int
mu sync.Mutex
}
func NewTraverser(cfg *TraverserConfig) *Traverser {
@@ -45,6 +48,8 @@ func NewTraverser(cfg *TraverserConfig) *Traverser {
return &Traverser{
config: cfg,
exchange: nil,
visited: make(map[string]bool),
depth: 0,
}
}
@@ -155,13 +160,32 @@ func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) {
func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response {
if !ref.HasAddresses() {
t.mu.Lock()
visitedCopy := make(map[string]bool)
for k, v := range t.visited {
visitedCopy[k] = v
}
t.mu.Unlock()
ref.Addresses = t.resolveGlueViaSystem(ctx, ref.Name, cache)
if len(ref.Addresses) > 0 {
ref.State = StateResolved
} else {
return &Response{
Referral: ref,
Type: RespError,
addrs, err := t.ResolveNS(ctx, ref.Name, cache, visitedCopy, t.depth)
if err != nil {
return &Response{
Referral: ref,
Type: RespError,
}
}
if len(addrs) > 0 {
ref.Addresses = addrs
ref.State = StateResolved
} else {
return &Response{
Referral: ref,
Type: RespError,
}
}
}
}
@@ -179,6 +203,131 @@ func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *I
}
}
func (t *Traverser) ResolveNS(ctx context.Context, nsName string, cache *InfoCache, visited map[string]bool, depth int) ([]net.IP, error) {
if cache != nil {
if addrs := cache.LookupGlue(nsName); len(addrs) > 0 {
return addrs, nil
}
}
if visited != nil {
if visited[nsName] {
return nil, &CircularReferralError{
Name: nsName,
Chain: getVisitedNames(visited),
}
}
visited[nsName] = true
}
if depth > DefaultMaxDepth {
return nil, &UnresolvableNameserverError{
Name: nsName,
Reason: "max depth exceeded",
}
}
roots, err := t.discoverRoots(ctx)
if err != nil {
return nil, fmt.Errorf("root discovery: %w", err)
}
ref := NewReferral(nsName, dns.TypeA, ".", 0, 1.0, nil)
ref.Addresses = roots
ref.State = StateResolved
traversalCache := NewInfoCache(nil)
if visited != nil {
for name := range visited {
traversalCache.StoreGlue(name, []net.IP{})
}
}
var addrs []net.IP
var lastErr error
stack := NewStack(DefaultMaxDepth)
stack.Push(ref)
for {
select {
case <-ctx.Done():
return nil, fmt.Errorf("resolution cancelled: %w", ctx.Err())
default:
}
current := stack.Pop()
if current == nil {
break
}
cacheForStep := traversalCache
if current.Parent != nil {
cacheForStep = traversalCache.Child()
}
resp := t.processReferral(ctx, current, cacheForStep)
if resp.Type == RespAnswer && len(resp.Decoded.Answers) > 0 {
for _, rr := range resp.Decoded.Answers {
if a, ok := rr.(*miekgdns.A); ok {
addrs = append(addrs, a.A)
}
if aaaa, ok := rr.(*miekgdns.AAAA); ok {
addrs = append(addrs, aaaa.AAAA)
}
}
if len(addrs) > 0 {
if cache != nil {
cache.StoreGlue(nsName, addrs)
}
return addrs, nil
}
}
if resp.Type == RespNXDOMAIN {
lastErr = &UnresolvableNameserverError{
Name: nsName,
Reason: "NXDOMAIN",
}
break
}
if resp.Type == RespSERVFAIL || resp.Type == RespError {
lastErr = fmt.Errorf("server error resolving %s: %s", nsName, resp.Type)
continue
}
if resp.Type == RespReferral {
children := resp.ChildReferrals()
for _, child := range children {
if visited != nil && visited[child.Name] {
continue
}
if !stack.Push(child) {
lastErr = &UnresolvableNameserverError{
Name: nsName,
Reason: "max depth exceeded during resolution",
}
}
}
}
}
if len(addrs) > 0 {
return addrs, nil
}
if lastErr != nil {
return nil, lastErr
}
return nil, &UnresolvableNameserverError{
Name: nsName,
Reason: "resolution exhausted without answer",
}
}
func (t *Traverser) queryServer(ctx context.Context, ref *Referral, server net.IP, cache *InfoCache) *Response {
var msg *miekgdns.Msg
var err error