CI / test (push) Failing after 2m58s
Co-authored-by: Hansen IT Solutions <gary@hansenits.com> Co-committed-by: Hansen IT Solutions <gary@hansenits.com>
424 lines
8.7 KiB
Go
424 lines
8.7 KiB
Go
package traverse
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"net"
|
|
"sync"
|
|
|
|
"github.com/hits/ExploreDNS/internal/dns"
|
|
miekgdns "github.com/miekg/dns"
|
|
)
|
|
|
|
type TraverserConfig struct {
|
|
MaxDepth int
|
|
QueryType uint16
|
|
RootConfig *dns.RootDiscoveryConfig
|
|
QueryConfig *dns.QueryConfig
|
|
RootAddrs []net.IP
|
|
}
|
|
|
|
func DefaultTraverserConfig() *TraverserConfig {
|
|
return &TraverserConfig{
|
|
MaxDepth: DefaultMaxDepth,
|
|
QueryType: dns.TypeA,
|
|
RootConfig: nil,
|
|
QueryConfig: nil,
|
|
RootAddrs: nil,
|
|
}
|
|
}
|
|
|
|
type TraversalResult struct {
|
|
Referral *Referral
|
|
Response *Response
|
|
}
|
|
|
|
type Traverser struct {
|
|
config *TraverserConfig
|
|
exchange dns.ExchangeFunc
|
|
visited map[string]bool
|
|
depth int
|
|
mu sync.Mutex
|
|
}
|
|
|
|
func NewTraverser(cfg *TraverserConfig) *Traverser {
|
|
if cfg == nil {
|
|
cfg = DefaultTraverserConfig()
|
|
}
|
|
return &Traverser{
|
|
config: cfg,
|
|
exchange: nil,
|
|
visited: make(map[string]bool),
|
|
depth: 0,
|
|
}
|
|
}
|
|
|
|
func (t *Traverser) SetExchange(fn dns.ExchangeFunc) {
|
|
t.exchange = fn
|
|
}
|
|
|
|
func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResult, error) {
|
|
name = miekgdns.Fqdn(name)
|
|
|
|
roots, err := t.discoverRoots(ctx)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("root discovery: %w", err)
|
|
}
|
|
|
|
initial := NewReferral(name, t.config.QueryType, ".", 0, 1.0, nil)
|
|
initial.Addresses = roots
|
|
|
|
stack := NewStack(t.config.MaxDepth)
|
|
stack.Push(initial)
|
|
|
|
rootCache := NewInfoCache(nil)
|
|
var (
|
|
mu sync.Mutex
|
|
results []TraversalResult
|
|
)
|
|
|
|
for {
|
|
select {
|
|
case <-ctx.Done():
|
|
return results, fmt.Errorf("traversal cancelled: %w", ctx.Err())
|
|
default:
|
|
}
|
|
|
|
ref := stack.Pop()
|
|
if ref == nil {
|
|
break
|
|
}
|
|
|
|
cache := rootCache
|
|
if ref.Parent != nil {
|
|
cache = rootCache.Child()
|
|
}
|
|
|
|
resp := t.processReferral(ctx, ref, cache)
|
|
mu.Lock()
|
|
results = append(results, TraversalResult{Referral: ref, Response: resp})
|
|
mu.Unlock()
|
|
|
|
if resp.IsTerminal() {
|
|
continue
|
|
}
|
|
|
|
if resp.Type == RespReferral {
|
|
children := resp.ChildReferrals()
|
|
for _, child := range children {
|
|
if !stack.Push(child) {
|
|
mu.Lock()
|
|
results = append(results, TraversalResult{
|
|
Referral: child,
|
|
Response: &Response{
|
|
Referral: child,
|
|
Type: RespError,
|
|
},
|
|
})
|
|
mu.Unlock()
|
|
}
|
|
}
|
|
}
|
|
|
|
if resp.Type == RespCNAMEFollow {
|
|
follow := resp.CNAMEFollowReferral()
|
|
if follow != nil {
|
|
if !stack.Push(follow) {
|
|
mu.Lock()
|
|
results = append(results, TraversalResult{
|
|
Referral: follow,
|
|
Response: &Response{
|
|
Referral: follow,
|
|
Type: RespError,
|
|
},
|
|
})
|
|
mu.Unlock()
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
return results, nil
|
|
}
|
|
|
|
func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) {
|
|
if len(t.config.RootAddrs) > 0 {
|
|
return t.config.RootAddrs, nil
|
|
}
|
|
|
|
servers, err := dns.DiscoverRoots(ctx, t.config.RootConfig)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
var addrs []net.IP
|
|
for _, srv := range servers {
|
|
addrs = append(addrs, srv.AllIPs(false)...)
|
|
}
|
|
return addrs, nil
|
|
}
|
|
|
|
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 {
|
|
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,
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
for _, addr := range ref.Addresses {
|
|
resp := t.queryServer(ctx, ref, addr, cache)
|
|
if resp != nil && resp.Type != RespSERVFAIL {
|
|
return resp
|
|
}
|
|
}
|
|
|
|
return &Response{
|
|
Referral: ref,
|
|
Type: RespSERVFAIL,
|
|
}
|
|
}
|
|
|
|
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
|
|
|
|
if t.exchange != nil {
|
|
msg, err = t.iterativeQueryWithExchange(ctx, server, ref.Name, ref.Qtype)
|
|
} else {
|
|
msg, err = dns.Query(ctx, server, ref.Name, ref.Qtype, t.config.QueryConfig)
|
|
if err == nil {
|
|
msg = t.ensureRDFalse(msg, server, ref.Name, ref.Qtype, t.config.QueryConfig)
|
|
}
|
|
}
|
|
|
|
if err != nil {
|
|
return &Response{
|
|
Referral: ref,
|
|
Server: server,
|
|
Type: RespError,
|
|
}
|
|
}
|
|
|
|
resp := NewResponse(ref, server, cache)
|
|
resp.Process(msg)
|
|
return resp
|
|
}
|
|
|
|
func (t *Traverser) iterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16) (*miekgdns.Msg, error) {
|
|
if t.config.QueryConfig == nil {
|
|
return dns.IterativeQueryWithExchange(ctx, server, name, qtype, nil, t.exchange)
|
|
}
|
|
return dns.IterativeQueryWithExchange(ctx, server, name, qtype, t.config.QueryConfig, t.exchange)
|
|
}
|
|
|
|
func (t *Traverser) ensureRDFalse(msg *miekgdns.Msg, server net.IP, name string, qtype uint16, cfg *dns.QueryConfig) *miekgdns.Msg {
|
|
if msg != nil && msg.RecursionDesired {
|
|
if t.exchange != nil {
|
|
ctx := context.Background()
|
|
var err error
|
|
msg, err = t.iterativeQueryWithExchange(ctx, server, name, qtype)
|
|
if err != nil {
|
|
return nil
|
|
}
|
|
return msg
|
|
}
|
|
msg.RecursionDesired = false
|
|
}
|
|
return msg
|
|
}
|
|
|
|
func (t *Traverser) resolveGlueViaSystem(ctx context.Context, name string, cache *InfoCache) []net.IP {
|
|
if cache != nil {
|
|
if addrs := cache.LookupGlue(name); len(addrs) > 0 {
|
|
return addrs
|
|
}
|
|
}
|
|
|
|
c := &miekgdns.Client{
|
|
Net: "udp",
|
|
ReadTimeout: 5,
|
|
WriteTimeout: 5,
|
|
}
|
|
if deadline, ok := ctx.Deadline(); ok {
|
|
c.ReadTimeout = deadline.Sub(deadline)
|
|
c.WriteTimeout = deadline.Sub(deadline)
|
|
}
|
|
|
|
fqdn := miekgdns.Fqdn(name)
|
|
|
|
aMsg, _, err := c.ExchangeContext(ctx, newAQuery(fqdn), "127.0.0.1:53")
|
|
if err == nil {
|
|
var addrs []net.IP
|
|
for _, rr := range aMsg.Answer {
|
|
if a, ok := rr.(*miekgdns.A); ok {
|
|
addrs = append(addrs, a.A)
|
|
}
|
|
}
|
|
if len(addrs) > 0 {
|
|
if cache != nil {
|
|
cache.StoreGlue(name, addrs)
|
|
}
|
|
return addrs
|
|
}
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func newAQuery(name string) *miekgdns.Msg {
|
|
m := new(miekgdns.Msg)
|
|
m.SetQuestion(name, miekgdns.TypeA)
|
|
m.RecursionDesired = true
|
|
return m
|
|
}
|