diff options
| author | radhitya <alif@radhitya.org> | 2026-07-05 14:05:24 +0700 |
|---|---|---|
| committer | radhitya <alif@radhitya.org> | 2026-07-05 14:05:24 +0700 |
| commit | 1799ef47cedcfed8c979e145783fff7c9d1944cc (patch) | |
| tree | 4634846c9035339d03f783543ac83056686a9f5b /internal/server | |
| parent | 52a0d3845ead07f840e9e99cb4a8c3507c84ab30 (diff) | |
foreward upstream, in memory cache, udp listener, handler, cli flags
Diffstat (limited to 'internal/server')
| -rw-r--r-- | internal/server/server.go | 161 |
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) +} |
