feat: implement referral resolution for nameserver names without glue records (HAN-381) (#7)
CI / test (push) Failing after 2m58s
CI / test (push) Failing after 2m58s
Co-authored-by: Hansen IT Solutions <gary@hansenits.com> Co-committed-by: Hansen IT Solutions <gary@hansenits.com>
This commit was merged in pull request #7.
This commit is contained in:
@@ -1,9 +1,12 @@
|
|||||||
package traverse
|
package traverse
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/hits/ExploreDNS/internal/dns"
|
||||||
miekgdns "github.com/miekg/dns"
|
miekgdns "github.com/miekg/dns"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -76,3 +79,160 @@ func (r *Referral) SetAddresses(addrs []net.IP) {
|
|||||||
r.State = StateUnresolved
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -36,6 +36,9 @@ type TraversalResult struct {
|
|||||||
type Traverser struct {
|
type Traverser struct {
|
||||||
config *TraverserConfig
|
config *TraverserConfig
|
||||||
exchange dns.ExchangeFunc
|
exchange dns.ExchangeFunc
|
||||||
|
visited map[string]bool
|
||||||
|
depth int
|
||||||
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewTraverser(cfg *TraverserConfig) *Traverser {
|
func NewTraverser(cfg *TraverserConfig) *Traverser {
|
||||||
@@ -45,6 +48,8 @@ func NewTraverser(cfg *TraverserConfig) *Traverser {
|
|||||||
return &Traverser{
|
return &Traverser{
|
||||||
config: cfg,
|
config: cfg,
|
||||||
exchange: nil,
|
exchange: nil,
|
||||||
|
visited: make(map[string]bool),
|
||||||
|
depth: 0,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -155,9 +160,27 @@ func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) {
|
|||||||
|
|
||||||
func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response {
|
func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response {
|
||||||
if !ref.HasAddresses() {
|
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)
|
ref.Addresses = t.resolveGlueViaSystem(ctx, ref.Name, cache)
|
||||||
if len(ref.Addresses) > 0 {
|
if len(ref.Addresses) > 0 {
|
||||||
ref.State = StateResolved
|
ref.State = StateResolved
|
||||||
|
} else {
|
||||||
|
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 {
|
} else {
|
||||||
return &Response{
|
return &Response{
|
||||||
Referral: ref,
|
Referral: ref,
|
||||||
@@ -165,6 +188,7 @@ func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *I
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
for _, addr := range ref.Addresses {
|
for _, addr := range ref.Addresses {
|
||||||
resp := t.queryServer(ctx, ref, addr, cache)
|
resp := t.queryServer(ctx, ref, addr, cache)
|
||||||
@@ -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 {
|
func (t *Traverser) queryServer(ctx context.Context, ref *Referral, server net.IP, cache *InfoCache) *Response {
|
||||||
var msg *miekgdns.Msg
|
var msg *miekgdns.Msg
|
||||||
var err error
|
var err error
|
||||||
|
|||||||
Reference in New Issue
Block a user