summaryrefslogtreecommitdiff
path: root/internal/server/server.go
diff options
context:
space:
mode:
authorradhitya <alif@radhitya.org>2026-07-05 14:05:24 +0700
committerradhitya <alif@radhitya.org>2026-07-05 14:05:24 +0700
commit1799ef47cedcfed8c979e145783fff7c9d1944cc (patch)
tree4634846c9035339d03f783543ac83056686a9f5b /internal/server/server.go
parent52a0d3845ead07f840e9e99cb4a8c3507c84ab30 (diff)
foreward upstream, in memory cache, udp listener, handler, cli flags
Diffstat (limited to 'internal/server/server.go')
-rw-r--r--internal/server/server.go161
1 files changed, 161 insertions, 0 deletions
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)
+}