From 3a194f1705b2e56c04fbe868d9edf2c226e77975 Mon Sep 17 00:00:00 2001 From: radhitya Date: Mon, 6 Jul 2026 19:16:57 +0700 Subject: recursive mode --- internal/resolver/recursive.go | 227 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 227 insertions(+) create mode 100644 internal/resolver/recursive.go (limited to 'internal/resolver/recursive.go') 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 +} -- cgit v1.2.3