summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--internal/blocklist/blocklist.go53
-rw-r--r--internal/cache/cache.go80
-rw-r--r--internal/dns/.pack.go.swpbin0 -> 12288 bytes
-rw-r--r--internal/dns/header.go2
-rw-r--r--internal/resolver/forward.go47
-rw-r--r--internal/server/server.go161
-rw-r--r--main.go37
7 files changed, 379 insertions, 1 deletions
diff --git a/internal/blocklist/blocklist.go b/internal/blocklist/blocklist.go
new file mode 100644
index 0000000..34a456d
--- /dev/null
+++ b/internal/blocklist/blocklist.go
@@ -0,0 +1,53 @@
+package blocklist
+
+import (
+ "bufio"
+ "os"
+ "strings"
+)
+
+type Blocklist struct {
+ domains map[string]struct{}
+}
+
+func Load(path string) (*Blocklist, error) {
+ bl := &Blocklist{domains: make(map[string]struct{})}
+ f, err := os.Open(path)
+ if err != nil {
+ if os.IsNotExist(err) {
+ return bl, nil
+ }
+ return nil, err
+ }
+ defer f.Close()
+
+ sc := bufio.NewScanner(f)
+ for sc.Scan() {
+ line := strings.TrimSpace(sc.Text())
+ if line == "" || line[0] == '#' {
+ continue
+ }
+ line = strings.ToLower(line)
+ line = strings.TrimSuffix(line, ".")
+ bl.domains[line] = struct{}{}
+ }
+ return bl, sc.Err()
+}
+
+func (bl *Blocklist) Blocked(qname string) bool {
+ qname = strings.ToLower(qname)
+ qname = strings.TrimSuffix(qname, ".")
+
+ if _, ok := bl.domains[qname]; ok {
+ return true
+ }
+
+ for i := 0; i < len(qname); i++ {
+ if qname[i] == '.' {
+ if _, ok := bl.domains[qname[i+1:]]; ok {
+ return true
+ }
+ }
+ }
+ return false
+}
diff --git a/internal/cache/cache.go b/internal/cache/cache.go
new file mode 100644
index 0000000..4d6b45c
--- /dev/null
+++ b/internal/cache/cache.go
@@ -0,0 +1,80 @@
+package cache
+
+import (
+ "sync"
+ "time"
+
+ "linum/internal/dns"
+)
+
+type Entry struct {
+ Msg *dns.Msg
+ Expires time.Time
+}
+
+func (e *Entry) expired() bool {
+ return time.Now().After(e.Expires)
+}
+
+type Cache struct {
+ mu sync.RWMutex
+ data map[string]*Entry
+}
+
+func New() *Cache {
+ return &Cache{
+ data: make(map[string]*Entry),
+ }
+}
+
+func makeKey(qname dns.Name, qtype, qclass uint16) string {
+ return qname.String() + ":" + itoa(qtype) + ":" + itoa(qclass)
+}
+
+func itoa(v uint16) string {
+ if v == 0 {
+ return "0"
+ }
+
+ var buf [5]byte
+ i := len(buf)
+ for v > 0 {
+ i--
+ buf[i] = byte(v%10) + '0'
+ v /= 10
+ }
+ return string(buf[i:])
+}
+
+func (c *Cache) Get(qname dns.Name, qtype, qclass uint16) (*dns.Msg, bool) {
+ key := makeKey(qname, qtype, qclass)
+ c.mu.RLock()
+ e, ok := c.data[key]
+ c.mu.RUnlock()
+ if !ok || e.expired() {
+ return nil, false
+ }
+ return e.Msg, true
+}
+
+func (c *Cache) Set(qname dns.Name, qtype, qclass uint16, msg *dns.Msg, ttl uint32) {
+ if ttl == 0 {
+ return
+ }
+ if ttl > 3600 {
+ ttl = 3600
+ }
+ key := makeKey(qname, qtype, qclass)
+ c.mu.Lock()
+ c.data[key] = &Entry{
+ Msg: msg,
+ Expires: time.Now().Add(time.Duration(ttl) * time.Second),
+ }
+ c.mu.Unlock()
+}
+
+func (c *Cache) Len() int {
+ c.mu.RLock()
+ defer c.mu.RUnlock()
+ return len(c.data)
+}
diff --git a/internal/dns/.pack.go.swp b/internal/dns/.pack.go.swp
new file mode 100644
index 0000000..dd841fb
--- /dev/null
+++ b/internal/dns/.pack.go.swp
Binary files differ
diff --git a/internal/dns/header.go b/internal/dns/header.go
index 3b51f67..0f462a8 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) set(mask uint16, v bool) {
if v {
h.Flags |= mask
diff --git a/internal/resolver/forward.go b/internal/resolver/forward.go
new file mode 100644
index 0000000..68edc5d
--- /dev/null
+++ b/internal/resolver/forward.go
@@ -0,0 +1,47 @@
+package resolver
+
+import (
+ "fmt"
+ "net"
+
+ "linum/internal/dns"
+)
+
+type Forwarder struct {
+ upstream string
+}
+
+func NewForwarder(upstream string) *Forwarder {
+ return &Forwarder{upstream: upstream}
+}
+
+func (f *Forwarder) Resolve(query *dns.Msg) (*dns.Msg, error) {
+ query.Header.SetRD(true)
+
+ wire, err := dns.Pack(query)
+ if err != nil {
+ return nil, fmt.Errorf("pack: %w", err)
+ }
+
+ conn, err := net.Dial("udp", f.upstream)
+ if err != nil {
+ return nil, fmt.Errorf("dial %s: %w", f.upstream, err)
+ }
+ defer conn.Close()
+
+ if _, err := conn.Write(wire); err != nil {
+ return nil, fmt.Errorf("write: %w", err)
+ }
+
+ buf := make([]byte, 4096)
+ n, err := conn.Read(buf)
+ if err != nil {
+ return nil, fmt.Errorf("read: %w", err)
+ }
+
+ resp, err := dns.Unpack(buf[:n])
+ if err != nil {
+ return nil, fmt.Errorf("unpack response: $w", err)
+ }
+ return resp, nil
+}
diff --git a/internal/server/server.go b/internal/server/server.go
new file mode 100644
index 0000000..5e580e6
--- /dev/null
+++ b/internal/server/server.go
@@ -0,0 +1,161 @@
+package server
+
+import (
+ "fmt"
+ "log"
+ "net"
+ "time"
+
+ "linum/internal/blocklist"
+ "linum/internal/cache"
+ "linum/internal/dns"
+ "linum/internal/resolver"
+)
+
+type Config struct {
+ Addr string
+ Upstream string
+ BlockFile string
+}
+
+type Server struct {
+ cfg Config
+ resolver *resolver.Forwarder
+ cache *cache.Cache
+ blocklist *blocklist.Blocklist
+}
+
+func New(cfg Config) (*Server, error) {
+ bl, err := blocklist.Load(cfg.BlockFile)
+ if err != nil {
+ return nil, fmt.Errorf("blocklist: %w", err)
+ }
+ return &Server{
+ cfg: cfg,
+ resolver: resolver.NewForwarder(cfg.Upstream),
+ cache: cache.New(),
+ blocklist: bl,
+ }, nil
+}
+
+func (s *Server) ListenAndServe() error {
+ addr, err := net.ResolveUDPAddr("udp", s.cfg.Addr)
+ if err != nil {
+ return fmt.Errorf("resolve addr: %w", err)
+ }
+
+ conn, err := net.ListenUDP("udp", addr)
+ if err != nil {
+ return fmt.Errorf("listen: %w", err)
+ }
+ defer conn.Close()
+ log.Printf("linum listening on %s, upstream %s", s.cfg.Addr, s.cfg.Upstream)
+
+ buf := make([]byte, 4096)
+ for {
+ n, remote, err := conn.ReadFromUDP(buf)
+ if err != nil {
+ log.Printf("read error: %v", err)
+ continue
+ }
+ go s.handle(conn, remote, buf[:n])
+ }
+}
+
+func (s *Server) handle(conn *net.UDPConn, remote *net.UDPAddr, data []byte) {
+ start := time.Now()
+
+ query, err := dns.Unpack(data)
+ if err != nil {
+ log.Printf("unpack from %s: %v", remote, err)
+ return
+ }
+ if len(query.Question) == 0 {
+ return
+ }
+
+ q := query.Question[0]
+ log.Printf("query %s %s from %s", qtypeName(q.Type), q.Name.String(), remote)
+
+ if s.blocklist != nil && s.blocklist.Blocked(q.Name.String()) {
+ s.reply(conn, remote, s.blockedResponse(query))
+ log.Printf("blocked %s (%v)", q.Name.String(), time.Since(start))
+ return
+ }
+
+ if cached, ok := s.cache.Get(q.Name, q.Type, q.Class); ok {
+ resp := *cached
+ resp.Header.ID = query.Header.ID
+ s.reply(conn, remote, &resp)
+ log.Printf("cache hit %s (%v)", q.Name.String(), time.Since(start))
+ return
+ }
+
+ resp, err := s.resolver.Resolve(query)
+ if err != nil {
+ log.Printf("resolve %s: %v", q.Name.String(), err)
+ s.reply(conn, remote, s.servfailResponse(query))
+ return
+ }
+
+ if resp.Header.Rcode() == dns.RcodeNoError || resp.Header.Rcode() == dns.RcodeNxDomain {
+ ttl := s.minTTL(resp)
+ s.cache.Set(q.Name, q.Type, q.Class, resp, ttl)
+ }
+
+ s.reply(conn, remote, resp)
+ log.Printf("resolved %s (%v)", q.Name.String(), time.Since(start))
+}
+
+func (s *Server) blockedResponse(query *dns.Msg) *dns.Msg {
+ resp := &dns.Msg{}
+ resp.Header.ID = query.Header.ID
+ resp.Header.SetQR(true)
+ resp.Header.SetRD(true)
+ resp.Header.SetRA(true)
+ resp.Header.SetRCode(dns.RcodeNxDomain)
+ resp.Question = query.Question
+ return resp
+}
+
+func (s *Server) servfailResponse(query *dns.Msg) *dns.Msg {
+ resp := &dns.Msg{}
+ resp.Header.ID = query.Header.ID
+ resp.Header.SetQR(true)
+ resp.Header.SetRCode(dns.RcodeServFail)
+ resp.Question = query.Question
+ return resp
+}
+
+func (s *Server) minTTL(resp *dns.Msg) uint32 {
+ min := uint32(300)
+ for _, rr := range resp.Answer {
+ if rr.TTL < min {
+ min = rr.TTL
+ }
+ }
+
+ return min
+}
+func (s *Server) reply(conn *net.UDPConn, remote *net.UDPAddr, msg *dns.Msg) {
+ wire, err := dns.Pack(msg)
+ if err != nil {
+ log.Printf("pack error: %v", err)
+ return
+ }
+ if _, err := conn.WriteToUDP(wire, remote); err != nil {
+ log.Printf("write error to %s: %v", remote, err)
+ }
+}
+
+func qtypeName(t uint16) string {
+ names := map[uint16]string{
+ dns.TypeA: "A", dns.TypeNS: "NS", dns.TypeCNAME: "CNAME",
+ dns.TypeSOA: "SOA", dns.TypePTR: "PTR", dns.TypeMX: "MX",
+ dns.TypeTXT: "TXT", dns.TypeAAAA: "AAAA",
+ }
+ if name, ok := names[t]; ok {
+ return name
+ }
+ return fmt.Sprintf("TYPE%d", t)
+}
diff --git a/main.go b/main.go
new file mode 100644
index 0000000..904f436
--- /dev/null
+++ b/main.go
@@ -0,0 +1,37 @@
+package main
+
+import (
+ "flag"
+ "log"
+ "os"
+ "os/signal"
+ "syscall"
+
+ "linum/internal/server"
+)
+
+func main() {
+ addr := flag.String("addr", ":5353", "listen address")
+ upstream := flag.String("upstream", "1.1.1.1:53", "upstream dns server")
+ blockfile := flag.String("blocklist", "blocklist.txt", "path to blocklist file")
+ flag.Parse()
+
+ srv, err := server.New(server.Config{
+ Addr: *addr,
+ Upstream: *upstream,
+ BlockFile: *blockfile,
+ })
+ if err != nil {
+ log.Fatalf("create server: %v", err)
+ }
+
+ go func() {
+ sig := make(chan os.Signal, 1)
+ signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM)
+ <-sig
+ log.Println("shutting down")
+ os.Exit(0)
+ }()
+
+ log.Fatal(srv.ListenAndServe())
+}