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 #6.
This commit is contained in:
@@ -0,0 +1,274 @@
|
||||
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
|
||||
}
|
||||
|
||||
func NewTraverser(cfg *TraverserConfig) *Traverser {
|
||||
if cfg == nil {
|
||||
cfg = DefaultTraverserConfig()
|
||||
}
|
||||
return &Traverser{
|
||||
config: cfg,
|
||||
exchange: nil,
|
||||
}
|
||||
}
|
||||
|
||||
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() {
|
||||
ref.Addresses = t.resolveGlueViaSystem(ctx, ref.Name, cache)
|
||||
if len(ref.Addresses) > 0 {
|
||||
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) 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
|
||||
}
|
||||
Reference in New Issue
Block a user