This repository has no description
1//go:build linux
2
3package microvm
4
5import (
6 "context"
7 "errors"
8 "fmt"
9 "log/slog"
10 "net"
11 "sync"
12 "time"
13
14 "github.com/miekg/dns"
15)
16
17const (
18 dnsProxyIOTimeout = 10 * time.Second
19 dnsProxyIdleTimeout = 30 * time.Second
20 dnsProxyShutdownTimeout = 10 * time.Second
21 dnsProxyMaxConnections = 64
22 dnsProxyMaxTCPQueries = 128
23 dnsProxyResolvConfPath = "/etc/resolv.conf"
24)
25
26type DNSProxy struct {
27 port uint32
28 srv *dns.Server
29
30 closeOnce sync.Once
31 closeErr error
32}
33
34func StartDNSProxy(ctx context.Context, cid uint32, logger *slog.Logger) (*DNSProxy, error) {
35 if ctx == nil {
36 ctx = context.Background()
37 }
38
39 if logger == nil {
40 logger = slog.Default()
41 }
42 logger = logger.With("where", "dns_proxy", "cid", cid)
43
44 ln, port, err := listenRandomVsockPort(ctx)
45 if err != nil {
46 return nil, fmt.Errorf("listen for dns proxy: %w", err)
47 }
48
49 resolver, err := newHostDNSResolver(dnsProxyResolvConfPath, logger)
50 if err != nil {
51 _ = ln.Close()
52 return nil, err
53 }
54
55 listener := newLimitedListener(
56 &cidFilteredVsockListener{
57 Listener: ln,
58 cid: cid,
59 logger: logger,
60 },
61 dnsProxyMaxConnections,
62 logger,
63 )
64
65 proxy := &DNSProxy{
66 port: port,
67 srv: &dns.Server{
68 Net: "tcp",
69 Listener: listener,
70 Handler: dns.HandlerFunc(resolver.ServeDNS),
71 ReadTimeout: dnsProxyIOTimeout,
72 WriteTimeout: dnsProxyIOTimeout,
73 IdleTimeout: func() time.Duration { return dnsProxyIdleTimeout },
74 MaxTCPQueries: dnsProxyMaxTCPQueries,
75 MsgInvalidFunc: func(_ []byte, err error) {
76 logger.Warn("dns proxy invalid message", "error", err)
77 },
78 },
79 }
80
81 go func() {
82 <-ctx.Done()
83 _ = proxy.Close()
84 }()
85
86 go func() {
87 if err := proxy.srv.ActivateAndServe(); err != nil && !errors.Is(err, net.ErrClosed) {
88 logger.Warn("dns proxy stopped", "error", err)
89 }
90 }()
91
92 logger.Info("started dns proxy", "port", port)
93 return proxy, nil
94}
95
96func (p *DNSProxy) Port() uint32 {
97 if p == nil {
98 return 0
99 }
100 return p.port
101}
102
103func (p *DNSProxy) Close() error {
104 if p == nil || p.srv == nil {
105 return nil
106 }
107
108 p.closeOnce.Do(func() {
109 shutdownCtx, cancel := context.WithTimeout(context.Background(), dnsProxyShutdownTimeout)
110 defer cancel()
111
112 p.closeErr = p.srv.ShutdownContext(shutdownCtx)
113 })
114 return p.closeErr
115}
116
117type limitedListener struct {
118 net.Listener
119 slots chan struct{}
120 logger *slog.Logger
121}
122
123func newLimitedListener(listener net.Listener, limit int, logger *slog.Logger) net.Listener {
124 if limit <= 0 {
125 return listener
126 }
127 return &limitedListener{
128 Listener: listener,
129 slots: make(chan struct{}, limit),
130 logger: logger,
131 }
132}
133
134func (l *limitedListener) Accept() (net.Conn, error) {
135 for {
136 conn, err := l.Listener.Accept()
137 if err != nil {
138 return nil, err
139 }
140
141 select {
142 case l.slots <- struct{}{}:
143 return &limitedConn{
144 Conn: conn,
145 release: func() {
146 <-l.slots
147 },
148 }, nil
149 default:
150 l.logger.Warn("dns proxy dropped connection because workers are busy")
151 _ = conn.Close()
152 }
153 }
154}
155
156type limitedConn struct {
157 net.Conn
158 once sync.Once
159 release func()
160}
161
162func (c *limitedConn) Close() error {
163 err := c.Conn.Close()
164 c.once.Do(c.release)
165 return err
166}
167
168type hostDNSResolver struct {
169 upstreams []string
170 attempts int
171 timeout time.Duration
172 logger *slog.Logger
173}
174
175func newHostDNSResolver(path string, logger *slog.Logger) (*hostDNSResolver, error) {
176 config, err := dns.ClientConfigFromFile(path)
177 if err != nil {
178 return nil, fmt.Errorf("read host resolv.conf: %w", err)
179 }
180 if len(config.Servers) == 0 {
181 return nil, fmt.Errorf("host resolv.conf has no nameservers")
182 }
183
184 port := config.Port
185 if port == "" {
186 port = "53"
187 }
188
189 upstreams := make([]string, 0, len(config.Servers))
190 for _, server := range config.Servers {
191 upstreams = append(upstreams, net.JoinHostPort(server, port))
192 }
193
194 timeout := time.Duration(config.Timeout) * time.Second
195 if timeout <= 0 {
196 timeout = dnsProxyIOTimeout
197 }
198
199 return &hostDNSResolver{
200 upstreams: upstreams,
201 attempts: max(config.Attempts, 1),
202 timeout: timeout,
203 logger: logger,
204 }, nil
205}
206
207func (r *hostDNSResolver) ServeDNS(w dns.ResponseWriter, req *dns.Msg) {
208 resp, err := r.exchange(req)
209 if err != nil {
210 r.logger.Warn(
211 "dns upstream exchange failed",
212 "question", dnsQuestionLogValue(req),
213 "error", err,
214 )
215 if err := w.WriteMsg(rcodeResponse(req, dns.RcodeServerFailure)); err != nil {
216 r.logger.Warn("dns proxy response write failed", "error", err)
217 }
218 return
219 }
220
221 filterDNSResponse(resp)
222
223 if err := w.WriteMsg(resp); err != nil {
224 r.logger.Warn("dns proxy response write failed", "error", err)
225 }
226}
227
228func (r *hostDNSResolver) exchange(req *dns.Msg) (*dns.Msg, error) {
229 var errs []error
230
231 for range r.attempts {
232 for _, upstream := range r.upstreams {
233 resp, err := exchangeDNSAt(req, upstream, r.timeout)
234 if err == nil {
235 return resp, nil
236 }
237 errs = append(errs, fmt.Errorf("%s: %w", upstream, err))
238 }
239 }
240
241 return nil, errors.Join(errs...)
242}
243
244func exchangeDNSAt(req *dns.Msg, addr string, timeout time.Duration) (*dns.Msg, error) {
245 resp, _, err := (&dns.Client{Net: "udp", Timeout: timeout}).Exchange(req, addr)
246 if err != nil {
247 return nil, err
248 }
249 if resp == nil {
250 return nil, fmt.Errorf("empty udp response")
251 }
252 if !resp.Truncated {
253 return resp, nil
254 }
255
256 resp, _, err = (&dns.Client{Net: "tcp", Timeout: timeout}).Exchange(req, addr)
257 if err != nil {
258 return nil, err
259 }
260 if resp == nil {
261 return nil, fmt.Errorf("empty tcp response")
262 }
263 return resp, nil
264}
265
266func filterDNSResponse(msg *dns.Msg) {
267 if msg == nil {
268 return
269 }
270 msg.Answer = filterDNSRRs(msg.Answer)
271 msg.Ns = filterDNSRRs(msg.Ns)
272 msg.Extra = filterDNSRRs(msg.Extra)
273}
274
275func filterDNSRRs(rrs []dns.RR) []dns.RR {
276 filtered := rrs[:0]
277 for _, rr := range rrs {
278 if rr := filterDNSRR(rr); rr != nil {
279 filtered = append(filtered, rr)
280 }
281 }
282 return filtered
283}
284
285func filterDNSRR(rr dns.RR) dns.RR {
286 switch rr := rr.(type) {
287 case *dns.A:
288 if isBlockedNamespaceIP(rr.A) {
289 return nil
290 }
291 case *dns.AAAA:
292 if isBlockedNamespaceIP(rr.AAAA) {
293 return nil
294 }
295 case *dns.SVCB:
296 filterSVCBValues(&rr.Value)
297 case *dns.HTTPS:
298 filterSVCBValues(&rr.Value)
299 }
300 return rr
301}
302
303// this removes any blocked namespaces in ipv4/v6 hints
304func filterSVCBValues(values *[]dns.SVCBKeyValue) {
305 filtered := (*values)[:0]
306 for _, value := range *values {
307 switch value := value.(type) {
308 case *dns.SVCBIPv4Hint:
309 value.Hint = filterDNSIPs(value.Hint)
310 if len(value.Hint) == 0 {
311 continue
312 }
313 case *dns.SVCBIPv6Hint:
314 value.Hint = filterDNSIPs(value.Hint)
315 if len(value.Hint) == 0 {
316 continue
317 }
318 }
319 filtered = append(filtered, value)
320 }
321 *values = filtered
322}
323
324func filterDNSIPs(ips []net.IP) []net.IP {
325 filtered := ips[:0]
326 for _, ip := range ips {
327 if !isBlockedNamespaceIP(ip) {
328 filtered = append(filtered, ip)
329 }
330 }
331 return filtered
332}
333
334func isBlockedNamespaceIP(ip net.IP) bool {
335 if ip == nil {
336 return true
337 }
338 if ip4 := ip.To4(); ip4 != nil {
339 return isBlockedByNamespaceNets(ip4, 32)
340 }
341 return isBlockedByNamespaceNets(ip, 128)
342}
343
344func isBlockedByNamespaceNets(ip net.IP, bits int) bool {
345 for _, blockedNet := range blockedNamespaceNets {
346 if blockedNet == nil {
347 continue
348 }
349
350 _, blockedBits := blockedNet.Mask.Size()
351 if blockedBits != bits {
352 continue
353 }
354 if blockedNet.Contains(ip) {
355 return true
356 }
357 }
358 return false
359}
360
361func rcodeResponse(req *dns.Msg, rcode int) *dns.Msg {
362 resp := new(dns.Msg)
363 if req == nil {
364 resp.Rcode = rcode
365 return resp
366 }
367 resp.SetRcode(req, rcode)
368 return resp
369}
370
371func dnsQuestionLogValue(msg *dns.Msg) string {
372 if msg == nil || len(msg.Question) == 0 {
373 return ""
374 }
375
376 q := msg.Question[0]
377 qtype := dns.TypeToString[q.Qtype]
378 if qtype == "" {
379 qtype = fmt.Sprintf("TYPE%d", q.Qtype)
380 }
381 return fmt.Sprintf("%s/%s", q.Name, qtype)
382}