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 }