main
c8b91bc · 4 months ago 7 commits
 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}