package traverse import ( "context" "fmt" "net" "strings" "github.com/hits/ExploreDNS/internal/dns" miekgdns "github.com/miekg/dns" "golang.org/x/net/idna" ) type ResolutionState int const ( StateUnresolved ResolutionState = iota StateResolving StateResolved ) func (s ResolutionState) String() string { switch s { case StateUnresolved: return "unresolved" case StateResolving: return "resolving" case StateResolved: return "resolved" default: return "unknown" } } type Referral struct { Name string Qtype uint16 Qclass uint16 Bailiwick string Addresses []net.IP State ResolutionState NSName string Parent *Referral Depth int Prob float64 } // idnaLookup is the IDN lookup profile used to convert internationalised domain // names (unicode labels) to their ACE/punycode equivalents before querying. var idnaLookup = idna.New( idna.MapForLookup(), idna.BidiRule(), idna.StrictDomainName(false), ) // toASCII converts a domain name that may contain unicode labels to its // punycode (ACE) representation. Pure-ASCII names are returned unchanged. // On conversion errors the original name is returned so the caller can still // attempt a query (the server will reject it if truly invalid). func toASCII(name string) string { if name == "" || name == "." { return name } ascii, err := idnaLookup.ToASCII(name) if err != nil { return name } return ascii } func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral { return &Referral{ Name: miekgdns.Fqdn(strings.ToLower(toASCII(name))), Qtype: qtype, Qclass: miekgdns.ClassINET, Bailiwick: miekgdns.Fqdn(strings.ToLower(toASCII(bailiwick))), Depth: depth, Prob: prob, Parent: parent, State: StateUnresolved, } } func (r *Referral) InBailiwick(name string) bool { if r.Bailiwick == "" || r.Bailiwick == "." { return true } fqdn := miekgdns.Fqdn(strings.ToLower(name)) return miekgdns.IsSubDomain(r.Bailiwick, fqdn) } func (r *Referral) HasAddresses() bool { return len(r.Addresses) > 0 } func (r *Referral) SetAddresses(addrs []net.IP) { r.Addresses = addrs if len(addrs) > 0 { r.State = StateResolved } else { 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 } // 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 }