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) }