diff options
| -rw-r--r-- | internal/blocklist/blocklist.go | 53 | ||||
| -rw-r--r-- | internal/cache/cache.go | 80 | ||||
| -rw-r--r-- | internal/dns/.pack.go.swp | bin | 0 -> 12288 bytes | |||
| -rw-r--r-- | internal/dns/header.go | 2 | ||||
| -rw-r--r-- | internal/resolver/forward.go | 47 | ||||
| -rw-r--r-- | internal/server/server.go | 161 | ||||
| -rw-r--r-- | main.go | 37 |
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 Binary files differnew file mode 100644 index 0000000..dd841fb --- /dev/null +++ b/internal/dns/.pack.go.swp 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) +} @@ -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()) +} |
