This repository has no description
0

Configure Feed

Select the types of activity you want to include in your feed.

core / spindle / engines / microvm / dns_proxy.go
7.8 kB 382 lines
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}