diff options
| -rw-r--r-- | Makefile | 8 | ||||
| -rw-r--r-- | internal/blocklist/blocklist.go | 2 | ||||
| -rw-r--r-- | internal/blocklist/blocklist_test.go | 6 | ||||
| -rw-r--r-- | internal/cache/cache.go | 6 | ||||
| -rw-r--r-- | internal/dns/header.go | 2 | ||||
| -rw-r--r-- | internal/resolver/forward.go | 2 | ||||
| -rw-r--r-- | internal/resolver/hints.go | 49 | ||||
| -rw-r--r-- | internal/resolver/recursive.go | 227 | ||||
| -rw-r--r-- | internal/resolver/resolver.go | 7 | ||||
| -rw-r--r-- | internal/server/server.go | 30 | ||||
| -rw-r--r-- | main.go | 4 |
11 files changed, 325 insertions, 18 deletions
diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..f9f8242 --- /dev/null +++ b/Makefile @@ -0,0 +1,8 @@ +BINARY = linum +GO = go +MAIN = . +OUTPUT = build/$(BINARY) +LDFLAGS = -ldflags="-s -w -X main.version=$(VERSION)" + +default: + $(GO) build $(LDFLAGS) -o $(OUTPUT) $(MAIN) diff --git a/internal/blocklist/blocklist.go b/internal/blocklist/blocklist.go index 84d758c..7154b33 100644 --- a/internal/blocklist/blocklist.go +++ b/internal/blocklist/blocklist.go @@ -29,7 +29,7 @@ func Load(path string) (*Blocklist, error) { } domain := parseLine(line) if domain != "" && !strings.HasPrefix(domain, "*.") { - bl.domains[strings.ToLower(domain)] = struct {}{} + bl.domains[strings.ToLower(domain)] = struct{}{} } } return bl, sc.Err() diff --git a/internal/blocklist/blocklist_test.go b/internal/blocklist/blocklist_test.go index 0104bc9..578c42b 100644 --- a/internal/blocklist/blocklist_test.go +++ b/internal/blocklist/blocklist_test.go @@ -6,13 +6,13 @@ import ( ) func TestParseLine(t *testing.T) { - tests := []struct{in,out string}{ + tests := []struct{ in, out string }{ {"0.0.0.0 domain.com", "domain.com"}, - {"0.0.0.0 domain.com # foo", "domain.com"}, + {"0.0.0.0 domain.com # foo", "domain.com"}, {"0.0.0.0 *.domain.com", "*.domain.com"}, {"||domain.com^", "domain.com"}, {"||domain.com^$script", "domain.com"}, - {"",""}, + {"", ""}, {"# comment", ""}, {"||*.domain.com^", "*.domain.com"}, } diff --git a/internal/cache/cache.go b/internal/cache/cache.go index 4d6b45c..b650ba1 100644 --- a/internal/cache/cache.go +++ b/internal/cache/cache.go @@ -8,7 +8,7 @@ import ( ) type Entry struct { - Msg *dns.Msg + Msg *dns.Msg Expires time.Time } @@ -17,7 +17,7 @@ func (e *Entry) expired() bool { } type Cache struct { - mu sync.RWMutex + mu sync.RWMutex data map[string]*Entry } @@ -67,7 +67,7 @@ func (c *Cache) Set(qname dns.Name, qtype, qclass uint16, msg *dns.Msg, ttl uint key := makeKey(qname, qtype, qclass) c.mu.Lock() c.data[key] = &Entry{ - Msg: msg, + Msg: msg, Expires: time.Now().Add(time.Duration(ttl) * time.Second), } c.mu.Unlock() diff --git a/internal/dns/header.go b/internal/dns/header.go index 0f462a8..eafe89f 100644 --- a/internal/dns/header.go +++ b/internal/dns/header.go @@ -47,7 +47,7 @@ func (h *Header) SetZ(v uint8) { h.Flags = (h.Flags &^ (0x7 << 4)) | ((uint1 func (h *Header) SetAD(v bool) { h.set(0x0020, v) } func (h *Header) SetCD(v bool) { h.set(0x0010, v) } func (h *Header) SetRCode(v uint8) { h.Flags = (h.Flags &^ 0xF) | (uint16(v) & 0xF) } -func (h *Header) Rcode() uint8 { return uint8(h.Flags & 0xF) } +func (h *Header) Rcode() uint8 { return uint8(h.Flags & 0xF) } func (h *Header) set(mask uint16, v bool) { if v { h.Flags |= mask diff --git a/internal/resolver/forward.go b/internal/resolver/forward.go index 68edc5d..05daad6 100644 --- a/internal/resolver/forward.go +++ b/internal/resolver/forward.go @@ -41,7 +41,7 @@ func (f *Forwarder) Resolve(query *dns.Msg) (*dns.Msg, error) { resp, err := dns.Unpack(buf[:n]) if err != nil { - return nil, fmt.Errorf("unpack response: $w", err) + return nil, fmt.Errorf("unpack response: %w", err) } return resp, nil } diff --git a/internal/resolver/hints.go b/internal/resolver/hints.go new file mode 100644 index 0000000..c7b23b4 --- /dev/null +++ b/internal/resolver/hints.go @@ -0,0 +1,49 @@ +package resolver + +import ( + "bufio" + "fmt" + "os" + "strings" +) + +type NS struct { + Name string + Addr string +} + +func LoadHints(path string) ([]NS, error) { + f, err := os.Open(path) + if err != nil { + return nil, fmt.Errorf("hints file %s not found — download from https://www.internic.net/domain/named.root: %w", path, err) + } + defer f.Close() + + var nsNames []string + glue := make(map[string]string) + sc := bufio.NewScanner(f) + for sc.Scan() { + fields := strings.Fields(sc.Text()) + if len(fields) < 4 || fields[0] == ";" { + continue + } + if fields[0] == "." && fields[2] == "NS" { + nsNames = append(nsNames, fields[3]) + } + if fields[2] == "A" { + glue[fields[0]] = fields[3] + } + } + if err := sc.Err(); err != nil { + return nil, err + } + + var hints []NS + for _, name := range nsNames { + hints = append(hints, NS{Name: name, Addr: glue[name]}) + } + if len(hints) == 0 { + return nil, fmt.Errorf("no root hints parsed from %s", path) + } + return hints, nil +} 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 +} diff --git a/internal/resolver/resolver.go b/internal/resolver/resolver.go new file mode 100644 index 0000000..43d9a0f --- /dev/null +++ b/internal/resolver/resolver.go @@ -0,0 +1,7 @@ +package resolver + +import "linum/internal/dns" + +type Resolver interface { + Resolve(query *dns.Msg) (*dns.Msg, error) +} diff --git a/internal/server/server.go b/internal/server/server.go index 5e580e6..fdb4364 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -13,15 +13,17 @@ import ( ) type Config struct { - Addr string - Upstream string + Mode string + Addr string + Upstream string BlockFile string + HintsFile string } type Server struct { - cfg Config - resolver *resolver.Forwarder - cache *cache.Cache + cfg Config + resolver resolver.Resolver + cache *cache.Cache blocklist *blocklist.Blocklist } @@ -30,10 +32,20 @@ func New(cfg Config) (*Server, error) { if err != nil { return nil, fmt.Errorf("blocklist: %w", err) } + var res resolver.Resolver + if cfg.Mode == "recursive" { + rec, err := resolver.NewRecursive(cfg.HintsFile) + if err != nil { + return nil, fmt.Errorf("recursive: %w", err) + } + res = rec + } else { + res = resolver.NewForwarder(cfg.Upstream) + } return &Server{ - cfg: cfg, - resolver: resolver.NewForwarder(cfg.Upstream), - cache: cache.New(), + cfg: cfg, + resolver: res, + cache: cache.New(), blocklist: bl, }, nil } @@ -49,7 +61,7 @@ func (s *Server) ListenAndServe() error { return fmt.Errorf("listen: %w", err) } defer conn.Close() - log.Printf("linum listening on %s, upstream %s", s.cfg.Addr, s.cfg.Upstream) + log.Printf("linum listening on %s, mode=%s, upstream %s", s.cfg.Addr, s.cfg.Mode, s.cfg.Upstream) buf := make([]byte, 4096) for { @@ -12,14 +12,18 @@ import ( func main() { addr := flag.String("addr", ":5353", "listen address") + mode := flag.String("mode", "forward", "resolver mode: forward or recursive") upstream := flag.String("upstream", "1.1.1.1:53", "upstream dns server") blockfile := flag.String("blocklist", "blocklist.txt", "path to blocklist file") + hints := flag.String("hints", "named.cache", "path to root hints file (download from internic.net)") flag.Parse() srv, err := server.New(server.Config{ + Mode: *mode, Addr: *addr, Upstream: *upstream, BlockFile: *blockfile, + HintsFile: *hints, }) if err != nil { log.Fatalf("create server: %v", err) |
