1package dns
2
3import (
4 "fmt"
5 "log"
6 "net"
7)
8
9type Handler func(req *Message) *Message
10
11type Server struct {
12 addr string
13 conn *net.UDPConn
14 handler Handler
15}
16
17func NewServer(addr string, handler Handler) *Server {
18 return &Server{
19 addr: addr,
20 handler: handler,
21 }
22}
23
24func (s *Server) Start() error {
25 udpAddr, err := net.ResolveUDPAddr("udp", s.addr)
26 if err != nil {
27 return fmt.Errorf("resolve addr: %w", err)
28 }
29
30 // Bind the UDP socket
31 conn, err := net.ListenUDP("udp", udpAddr)
32 if err != nil {
33 return fmt.Errorf("listen udp: %w", err)
34 }
35 s.conn = conn
36
37 log.Printf("DNS server listening on %s", s.addr)
38
39 // Read loop — runs forever
40 buf := make([]byte, 512) // RFC 1035: max UDP DNS message is 512 bytes
41 for {
42 n, clientAddr, err := conn.ReadFromUDP(buf)
43 if err != nil {
44 log.Printf("read error: %v", err)
45 continue
46 }
47
48 // Handle each request in a goroutine
49 // so slow requests don't block other clients
50 go s.handlePacket(buf[:n], clientAddr)
51 }
52}
53
54func (s *Server) handlePacket(buf []byte, clientAddr *net.UDPAddr) {
55 // Parse the incoming message
56 req, err := UnpackMessage(buf)
57 if err != nil {
58 log.Printf("failed to parse DNS message from %s: %v", clientAddr, err)
59 return
60 }
61
62 log.Printf("query from %s: %s type=%d", clientAddr, req.Questions[0].Name, req.Questions[0].Type)
63
64 // Call our handler to build a response
65 resp := s.handler(req)
66 if resp == nil {
67 return
68 }
69
70 // Serialize response to wire format
71 respBytes, err := resp.Pack()
72 if err != nil {
73 log.Printf("failed to pack response: %v", err)
74 return
75 }
76
77 // Send it back to the client
78 _, err = s.conn.WriteToUDP(respBytes, clientAddr)
79 if err != nil {
80 log.Printf("failed to send response to %s: %v", clientAddr, err)
81 }
82}