summaryrefslogtreecommitdiff
path: root/internal/resolver/recursive.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/resolver/recursive.go')
-rw-r--r--internal/resolver/recursive.go227
1 files changed, 227 insertions, 0 deletions
diff --git a/internal/resolver/recursive.go b/internal/resolver/recursive.go
new file mode 100644
index 0000000..e9ab3b4
--- /dev/null
+++ b/internal/resolver/recursive.go
@@ -0,0 +1,227 @@
+package resolver
+
+import (
+ "fmt"
+ "math/rand"
+ "net"
+ "time"
+
+ "linum/internal/cache"
+ "linum/internal/dns"
+)
+
+const (
+ maxDelegations = 30
+ maxCNAME = 10
+ maxDepth = 20
+ queryTimeout = 2 * time.Second
+)
+
+type Recursive struct {
+ hints []NS
+ cache *cache.Cache
+}
+
+func NewRecursive(hintsPath string) (*Recursive, error) {
+ hints, err := LoadHints(hintsPath)
+ if err != nil {
+ return nil, fmt.Errorf("load hints: %w", err)
+ }
+ return &Recursive{
+ hints: hints,
+ cache: cache.New(),
+ }, nil
+}
+
+func (r *Recursive) Resolve(query *dns.Msg) (*dns.Msg, error) {
+ if len(query.Question) == 0 {
+ return nil, fmt.Errorf("no question")
+ }
+ q := query.Question[0]
+ resp, err := r.resolve(q.Name, q.Type, q.Class, query.Header.ID, 0)
+ if err != nil {
+ return nil, err
+ }
+ resp.Header.SetQR(true)
+ resp.Header.SetRA(true)
+ return resp, nil
+}
+
+func (r *Recursive) resolve(qname dns.Name, qtype, qclass uint16, id uint16, depth int) (*dns.Msg, error) {
+ if depth > maxDepth {
+ return nil, fmt.Errorf("max depth exceeded")
+ }
+
+ if cached, ok := r.cache.Get(qname, qtype, qclass); ok {
+ resp := *cached
+ resp.Header.ID = id
+ return &resp, nil
+ }
+
+ resp, err := r.followDelegations(qname, qtype, qclass, id, depth)
+ if err != nil {
+ return nil, err
+ }
+
+ if resp.Header.Rcode() == dns.RcodeNoError || resp.Header.Rcode() == dns.RcodeNxDomain {
+ ttl := minTTL(resp)
+ r.cache.Set(qname, qtype, qclass, resp, ttl)
+ }
+
+ if resp.Header.Rcode() == dns.RcodeNoError && len(resp.Answer) > 0 {
+ for _, rr := range resp.Answer {
+ if rr.Type == dns.TypeCNAME {
+ if cname, ok := rr.Data.(*dns.CNAME); ok {
+ targetResp, err := r.resolve(cname.Cname, qtype, qclass, id, depth+1)
+ if err != nil {
+ return resp, nil
+ }
+ merged := *resp
+ merged.Answer = append(merged.Answer, targetResp.Answer...)
+ return &merged, nil
+ }
+ }
+ }
+ }
+ return resp, nil
+}
+
+func (r *Recursive) followDelegations(qname dns.Name, qtype, qclass, id uint16, depth int) (*dns.Msg, error) {
+ nsList := make([]NS, len(r.hints))
+ copy(nsList, r.hints)
+ rand.Shuffle(len(nsList), func(i, j int) {
+ nsList[i], nsList[j] = nsList[j], nsList[i]
+ })
+
+ for delegations := 0; delegations < maxDelegations; delegations++ {
+ resp, err := r.queryNSList(nsList, qname, qtype, qclass, id, depth)
+ if err != nil {
+ return nil, err
+ }
+
+ if len(resp.Answer) > 0 {
+ return resp, nil
+ }
+ if resp.Header.Rcode() == dns.RcodeNxDomain {
+ return resp, nil
+ }
+ if len(resp.Ns) > 0 {
+ newNS := extractDelegation(resp)
+ if len(newNS) > 0 {
+ nsList = newNS
+ rand.Shuffle(len(nsList), func(i, j int) {
+ nsList[i], nsList[j] = nsList[j], nsList[i]
+ })
+ continue
+ }
+ }
+ return resp, nil
+ }
+ return nil, fmt.Errorf("max delegations exceeded")
+}
+
+func (r *Recursive) queryNSList(nsList []NS, qname dns.Name, qtype, qclass, id uint16, depth int) (*dns.Msg, error) {
+ var lastErr error
+ for _, ns := range nsList {
+ addr := ns.Addr
+ if addr == "" {
+ nsName, err := dns.NewName(ns.Name)
+ if err != nil {
+ lastErr = err
+ continue
+ }
+ ip, err := r.resolveNS(nsName, id, depth)
+ if err != nil {
+ lastErr = err
+ continue
+ }
+ addr = ip
+ }
+ resp, err := r.query(addr, qname, qtype, qclass, id)
+ if err != nil {
+ lastErr = err
+ continue
+ }
+ return resp, nil
+ }
+ return nil, fmt.Errorf("all NS failed: %w", lastErr)
+}
+func (r *Recursive) query(server string, qname dns.Name, qtype, qclass, id uint16) (*dns.Msg, error) {
+ q := &dns.Msg{
+ Header: dns.Header{ID: id},
+ Question: []dns.Question{{Name: qname, Type: qtype, Class: qclass}},
+ }
+ wire, err := dns.Pack(q)
+ if err != nil {
+ return nil, err
+ }
+
+ conn, err := net.DialTimeout("udp", net.JoinHostPort(server, "53"), queryTimeout)
+ if err != nil {
+ return nil, err
+ }
+ defer conn.Close()
+ conn.SetDeadline(time.Now().Add(queryTimeout))
+
+ if _, err := conn.Write(wire); err != nil {
+ return nil, err
+ }
+
+ buf := make([]byte, 4096)
+ n, err := conn.Read(buf)
+ if err != nil {
+ return nil, err
+ }
+ return dns.Unpack(buf[:n])
+}
+
+func (r *Recursive) resolveNS(name dns.Name, id uint16, depth int) (string, error) {
+ resp, err := r.resolve(name, dns.TypeA, dns.ClassIN, id, depth+1)
+ if err != nil {
+ return "", err
+ }
+ for _, rr := range resp.Answer {
+ if a, ok := rr.Data.(*dns.A); ok {
+ return a.Addr.String(), nil
+ }
+ }
+ return "", fmt.Errorf("no A record for %s", name.String())
+}
+
+func extractDelegation(resp *dns.Msg) []NS {
+ var nsList []NS
+ glue := make(map[string]string)
+ for _, rr := range resp.Extra {
+ if rr.Type == dns.TypeA {
+ if a, ok := rr.Data.(*dns.A); ok {
+ glue[rr.Name.String()] = a.Addr.String()
+ }
+ }
+ }
+ for _, rr := range resp.Ns {
+ if rr.Type == dns.TypeNS {
+ if ns, ok := rr.Data.(*dns.NS); ok {
+ nsList = append(nsList, NS{
+ Name: ns.Ns.String(),
+ Addr: glue[ns.Ns.String()],
+ })
+ }
+ }
+ }
+ return nsList
+}
+
+func minTTL(resp *dns.Msg) uint32 {
+ min := uint32(300)
+ for _, rr := range resp.Answer {
+ if rr.TTL < min {
+ min = rr.TTL
+ }
+ }
+ for _, rr := range resp.Ns {
+ if rr.TTL < min {
+ min = rr.TTL
+ }
+ }
+ return min
+}