summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--Makefile8
-rw-r--r--internal/blocklist/blocklist.go2
-rw-r--r--internal/blocklist/blocklist_test.go6
-rw-r--r--internal/cache/cache.go6
-rw-r--r--internal/dns/header.go2
-rw-r--r--internal/resolver/forward.go2
-rw-r--r--internal/resolver/hints.go49
-rw-r--r--internal/resolver/recursive.go227
-rw-r--r--internal/resolver/resolver.go7
-rw-r--r--internal/server/server.go30
-rw-r--r--main.go4
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 {
diff --git a/main.go b/main.go
index 904f436..1633b2a 100644
--- a/main.go
+++ b/main.go
@@ -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)