From c56385602e93b2866aaee456405ff3ef4b14e6eb Mon Sep 17 00:00:00 2001 From: Pavel Sviderski Date: Tue, 29 Apr 2025 20:20:55 +1000 Subject: [PATCH] chore(dns-server): authoritative responses for internal DNS queries and forwarding to ustream servers --- go.mod | 16 +- go.sum | 32 ++-- internal/machine/dns/server.go | 334 ++++++++++++++++++++++++++++++++- 3 files changed, 353 insertions(+), 29 deletions(-) diff --git a/go.mod b/go.mod index 477ad102..540afa4c 100644 --- a/go.mod +++ b/go.mod @@ -36,6 +36,7 @@ require ( github.com/ipfs/go-log/v2 v2.5.1 github.com/jmoiron/sqlx v1.4.0 github.com/lmittmann/tint v1.0.5 + github.com/miekg/dns v1.1.65 github.com/opencontainers/go-digest v1.0.0 github.com/opencontainers/image-spec v1.1.0 github.com/siderolabs/discovery-api v0.1.4 @@ -46,9 +47,9 @@ require ( github.com/vishvananda/netlink v1.3.0 go.uber.org/zap v1.27.0 go4.org/netipx v0.0.0-20231129151722-fdeea329fbba - golang.org/x/crypto v0.32.0 - golang.org/x/net v0.34.0 - golang.org/x/sync v0.10.0 + golang.org/x/crypto v0.33.0 + golang.org/x/net v0.35.0 + golang.org/x/sync v0.11.0 golang.org/x/sys v0.31.0 golang.zx2c4.com/wireguard v0.0.0-20231211153847-12269c276173 golang.zx2c4.com/wireguard/wgctrl v0.0.0-20230429144221-925a1e7659e6 @@ -197,7 +198,6 @@ require ( github.com/mdlayher/socket v0.5.1 // indirect github.com/mgutz/ansi v0.0.0-20200706080929-d51e80ef957d // indirect github.com/mholt/acmez/v2 v2.0.3 // indirect - github.com/miekg/dns v1.1.62 // indirect github.com/miekg/pkcs11 v1.1.1 // indirect github.com/minio/sha256-simd v1.0.1 // indirect github.com/mitchellh/cli v1.1.5 // indirect @@ -303,11 +303,11 @@ require ( go.uber.org/zap/exp v0.2.0 // indirect golang.org/x/crypto/x509roots/fallback v0.0.0-20240507223354-67b13616a595 // indirect golang.org/x/exp v0.0.0-20241215155358-4a5509556b9e // indirect - golang.org/x/mod v0.22.0 // indirect - golang.org/x/term v0.28.0 // indirect - golang.org/x/text v0.21.0 // indirect + golang.org/x/mod v0.23.0 // indirect + golang.org/x/term v0.29.0 // indirect + golang.org/x/text v0.22.0 // indirect golang.org/x/time v0.8.0 // indirect - golang.org/x/tools v0.29.0 // indirect + golang.org/x/tools v0.30.0 // indirect golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect google.golang.org/genproto v0.0.0-20240401170217-c3f982113cda // indirect google.golang.org/genproto/googleapis/api v0.0.0-20241209162323-e6fa225c2576 // indirect diff --git a/go.sum b/go.sum index a0f882cc..2c743bc6 100644 --- a/go.sum +++ b/go.sum @@ -721,8 +721,8 @@ github.com/mholt/acmez/v2 v2.0.3 h1:CgDBlEwg3QBp6s45tPQmFIBrkRIkBT4rW4orMM6p4sw= github.com/mholt/acmez/v2 v2.0.3/go.mod h1:pQ1ysaDeGrIMvJ9dfJMk5kJNkn7L2sb3UhyrX6Q91cw= github.com/miekg/dns v1.1.26/go.mod h1:bPDLeHnStXmXAq1m/Ch/hvfNHr14JKNPMBo3VZKjuso= github.com/miekg/dns v1.1.41/go.mod h1:p6aan82bvRIyn+zDIv9xYNUpwa73JcSh9BKwknJysuI= -github.com/miekg/dns v1.1.62 h1:cN8OuEF1/x5Rq6Np+h1epln8OiyPWV+lROx9LxcGgIQ= -github.com/miekg/dns v1.1.62/go.mod h1:mvDlcItzm+br7MToIKqkglaGhlFMHJ9DTNNWONWXbNQ= +github.com/miekg/dns v1.1.65 h1:0+tIPHzUW0GCge7IiK3guGP57VAw7hoPDfApjkMD1Fc= +github.com/miekg/dns v1.1.65/go.mod h1:Dzw9769uoKVaLuODMDZz9M6ynFU6Em65csPuoi8G0ck= github.com/miekg/pkcs11 v1.0.2/go.mod h1:XsNlhZGX73bx86s2hdc/FuaLm2CPZJemRLMA+WTFxgs= github.com/miekg/pkcs11 v1.1.1 h1:Ugu9pdy6vAYku5DEpVWVFPYnzV+bxB+iRdbuFSu7TvU= github.com/miekg/pkcs11 v1.1.1/go.mod h1:XsNlhZGX73bx86s2hdc/FuaLm2CPZJemRLMA+WTFxgs= @@ -1196,8 +1196,8 @@ golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97/go.mod h1:GvvjBRRGRdwPK5y golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= golang.org/x/crypto v0.3.0/go.mod h1:hebNnKkNXi2UzZN1eVRvBB7co0a+JxK6XbPiWVs/3J4= golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU= -golang.org/x/crypto v0.32.0 h1:euUpcYgM8WcP71gNpTqQCn6rC2t6ULUPiOzfWaXVVfc= -golang.org/x/crypto v0.32.0/go.mod h1:ZnnJkOaASj8g0AjIduWNlq2NRxL0PlBrbKVyZ6V/Ugc= +golang.org/x/crypto v0.33.0 h1:IOBPskki6Lysi0lo9qQvbxiQ+FvsCC/YWOecCHAixus= +golang.org/x/crypto v0.33.0/go.mod h1:bVdXmD7IV/4GdElGPozy6U7lWdRXA4qyRVGJV57uQ5M= golang.org/x/crypto/x509roots/fallback v0.0.0-20240507223354-67b13616a595 h1:TgSqweA595vD0Zt86JzLv3Pb/syKg8gd5KMGGbJPYFw= golang.org/x/crypto/x509roots/fallback v0.0.0-20240507223354-67b13616a595/go.mod h1:kNa9WdvYnzFwC79zRpLRMJbdEFlhyM5RPFBBZp/wWH8= golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= @@ -1214,8 +1214,8 @@ golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.4.2/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= -golang.org/x/mod v0.22.0 h1:D4nJWe9zXqHOmWqj4VMOJhvzj7bEZg4wEYa759z1pH4= -golang.org/x/mod v0.22.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY= +golang.org/x/mod v0.23.0 h1:Zb7khfcRGKk+kqfxFaP5tZqCnDZMjC5VtUBs87Hr6QM= +golang.org/x/mod v0.23.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180906233101-161cd47e91fd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= @@ -1237,8 +1237,8 @@ golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug golang.org/x/net v0.2.0/go.mod h1:KqCZLdyyvdV855qA2rE3GC2aiw5xGR5TEjj8smXukLY= golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg= -golang.org/x/net v0.34.0 h1:Mb7Mrk043xzHgnRM88suvJFwzVrRfHEHJEl5/71CKw0= -golang.org/x/net v0.34.0/go.mod h1:di0qlW3YNM5oh6GqDGQr92MyTozJPmybPK4Ev/Gm31k= +golang.org/x/net v0.35.0 h1:T5GQRQb2y08kTAByq9L4/bz8cipCdA8FbRTXewonqY8= +golang.org/x/net v0.35.0/go.mod h1:EglIi67kWsHKlRzzVMUD93VMSWGFOMSZgxFjparz1Qk= golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= golang.org/x/oauth2 v0.24.0 h1:KTBBxWqUa0ykRPLtV69rRto9TLXcqYkeswu48x/gvNE= golang.org/x/oauth2 v0.24.0/go.mod h1:XYTD2NtWslqkgxebSiOHnXEap4TF09sJSc7H1sXbhtI= @@ -1252,8 +1252,8 @@ golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJ golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.10.0 h1:3NQrjDixjgGwUOCaF8w2+VYHv0Ve/vGYSbdkTa98gmQ= -golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= +golang.org/x/sync v0.11.0 h1:GGz8+XQP4FvTTrjZPzNKTMFtSXH80RAzG+5ghFPgK9w= +golang.org/x/sync v0.11.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk= golang.org/x/sys v0.0.0-20180823144017-11551d06cbcc/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= @@ -1310,8 +1310,8 @@ golang.org/x/term v0.2.0/go.mod h1:TVmDHMZPmdnySmBfhjOoOdhjzdE1h4u1VwSiw2l1Nuc= golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k= golang.org/x/term v0.8.0/go.mod h1:xPskH00ivmX89bAKVGSKKtLOWNx2+17Eiy94tnKShWo= golang.org/x/term v0.17.0/go.mod h1:lLRBjIVuehSbZlaOtGMbcMncT+aqLLLmKrsjNrUguwk= -golang.org/x/term v0.28.0 h1:/Ts8HFuMR2E6IP/jlo7QVLZHggjKQbhu/7H0LJFr3Gg= -golang.org/x/term v0.28.0/go.mod h1:Sw/lC2IAUZ92udQNf3WodGtn4k/XoLyZoh8v/8uiwek= +golang.org/x/term v0.29.0 h1:L6pJp37ocefwRRtYPKSWOWzOtWSxVajvz2ldH/xi3iU= +golang.org/x/term v0.29.0/go.mod h1:6bl4lRlvVuDgSf3179VpIxBF0o10JUpXWOnI7nErv7s= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= @@ -1322,8 +1322,8 @@ golang.org/x/text v0.4.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8= golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= -golang.org/x/text v0.21.0 h1:zyQAAkrwaneQ066sspRyJaG9VNi/YJ1NfzcGB3hZ/qo= -golang.org/x/text v0.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ= +golang.org/x/text v0.22.0 h1:bofq7m3/HAFvbF51jz3Q9wLg3jkvSPuiZu/pD1XwgtM= +golang.org/x/text v0.22.0/go.mod h1:YRoo4H8PVmsu+E3Ou7cqLVH8oXWIHVoX0jqUWALQhfY= golang.org/x/time v0.8.0 h1:9i3RxcPv3PZnitoVGMPDKZSq1xW1gK1Xy3ArNOGZfEg= golang.org/x/time v0.8.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= @@ -1345,8 +1345,8 @@ golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4f golang.org/x/tools v0.1.5/go.mod h1:o0xws9oXOQQZyjljx8fwUC0k7L1pTE6eaCbjGeHmOkk= golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc= golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU= -golang.org/x/tools v0.29.0 h1:Xx0h3TtM9rzQpQuR4dKLrdglAmCEN5Oi+P74JdhdzXE= -golang.org/x/tools v0.29.0/go.mod h1:KMQVMRsVxU6nHCFXrBPhDB8XncLNLM0lIy/F14RP588= +golang.org/x/tools v0.30.0 h1:BgcpHewrV5AUp2G9MebG4XPFI1E2W41zU1SaqVA9vJY= +golang.org/x/tools v0.30.0/go.mod h1:c347cR/OJfw5TI+GfX7RUPNMdDRRbjvYTS0jPyvsVtY= golang.org/x/xerrors v0.0.0-20190410155217-1f06c39b4373/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20190513163551-3ee3066db522/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= diff --git a/internal/machine/dns/server.go b/internal/machine/dns/server.go index 91396cd1..90e23e49 100644 --- a/internal/machine/dns/server.go +++ b/internal/machine/dns/server.go @@ -2,14 +2,338 @@ package dns import ( "context" + "errors" + "fmt" + "log/slog" + "math/rand/v2" + "net" "net/netip" + "strconv" + "strings" + "sync" + "time" + + "github.com/miekg/dns" ) -// Server provides a DNS server for service discovery and proxying to upstream DNS servers. -type Server struct { - store RecordsStore +const ( + // InternalDomain is the cluster internal domain for service discovery. All DNS queries ending with this suffix + // will be resolved using the internal DNS server. + InternalDomain = "internal." + // dnsPort is the standard DNS port. + dnsPort = 53 + // maxConcurrentForwards is the maximum number of concurrent forwarded queries to upstream DNS servers. + // 1024 is the default used by the Docker internal DNS server. + maxConcurrentForwards = 1024 + // forwardingTimeout is the timeout for forwarding a DNS query to an upstream server. + forwardingTimeout = 3 * time.Second +) + +// Resolver is an interface for resolving service names to IP addresses. +type Resolver interface { + // Resolve returns a list of IP addresses of service containers. An empty list is returned if no service is found. + Resolve(serviceName string) []netip.Addr } -type RecordsStore interface { - Resolve(ctx context.Context, name string) ([]netip.Addr, error) +// Server is an embedded internal DNS server for service discovery and forwarding external queries +// to upstream DNS servers. +type Server struct { + listenAddr netip.Addr + resolver Resolver + upstreamServers []netip.AddrPort + + udpServer *dns.Server + tcpServer *dns.Server + inProgressReqs sync.WaitGroup + forwardSemaphore chan struct{} + log *slog.Logger +} + +// NewServer creates a new DNS server with the given configuration. +// If upstreams is nil, nameservers from /etc/resolv.conf will be used. An empty upstreams list means to only resolve +// internal DNS queries and not forward any external queries. +func NewServer(listenAddr netip.Addr, resolver Resolver, upstreams []netip.AddrPort) (*Server, error) { + if listenAddr.IsValid() { + return nil, fmt.Errorf("invalid listen address: %s", listenAddr) + } + if resolver == nil { + return nil, fmt.Errorf("resolver must be provided") + } + + // Load upstreams from /etc/resolv.conf and set fallback servers only if upstreams is not provided. + if upstreams == nil { + if nameservers, err := parseNameserversFromResolvConf(); err != nil { + slog.Warn("Failed to parse nameservers from /etc/resolv.conf.", "err", err) + } else { + for _, ns := range nameservers { + if ns.Compare(listenAddr) == 0 { + // Skip this internal DNS server address if it has been configured in /etc/resolv.conf + // to not create a forwarding loop. + continue + } + upstreams = append(upstreams, netip.AddrPortFrom(ns, dnsPort)) + } + } + + // Fallback to common public DNS servers if no nameservers were found in /etc/resolv.conf. + if upstreams == nil { + upstreams = []netip.AddrPort{ + netip.AddrPortFrom(netip.MustParseAddr("1.1.1.1"), dnsPort), // Cloudflare DNS + netip.AddrPortFrom(netip.MustParseAddr("8.8.8.8"), dnsPort), // Google DNS + } + } + } + + return &Server{ + listenAddr: listenAddr, + resolver: resolver, + upstreamServers: upstreams, + forwardSemaphore: make(chan struct{}, maxConcurrentForwards), + log: slog.With("component", "dns-server"), + }, nil +} + +// Run starts the DNS server listening on both UDP and TCP ports. It will stop the server when the context is canceled. +func (s *Server) Run(ctx context.Context) error { + addr := net.JoinHostPort(s.listenAddr.String(), strconv.Itoa(dnsPort)) + s.udpServer = &dns.Server{ + Addr: addr, + Net: "udp", + Handler: dns.HandlerFunc(s.handleRequest), + } + s.tcpServer = &dns.Server{ + Addr: addr, + Net: "tcp", + Handler: dns.HandlerFunc(s.handleRequest), + } + + errCh := make(chan error, 2) // Buffer size 2 for UDP and TCP errors. + + go func() { + s.log.Info("Starting DNS server (UDP).", "addr", addr, "proto", "udp") + if err := s.udpServer.ListenAndServe(); err != nil { + errCh <- fmt.Errorf("listen and serve on %s/udp: %w", addr, err) + } + }() + + go func() { + s.log.Info("Starting DNS server (TCP).", "addr", addr, "proto", "tcp") + if err := s.tcpServer.ListenAndServe(); err != nil { + errCh <- fmt.Errorf("listen and serve on %s/udp: %w", addr, err) + } + }() + + select { + case err := <-errCh: + // Stop the servers if one of them fails. + s.stop() + return err + case <-ctx.Done(): + s.log.Info("Stopping DNS server.") + return s.stop() + } +} + +// stop gracefully shuts down the DNS server. +func (s *Server) stop() error { + var udpErr, tcpErr error + + if s.udpServer != nil { + udpErr = s.udpServer.Shutdown() + } + if s.tcpServer != nil { + tcpErr = s.tcpServer.Shutdown() + } + + // Wait for all in-progress requests to finish. + s.inProgressReqs.Wait() + + if udpErr != nil { + return fmt.Errorf("shutdown DNS server (UDP): %w", udpErr) + } + + if tcpErr != nil { + return fmt.Errorf("shutdown DNS server (TCP): %w", tcpErr) + } + + return nil +} + +// handleRequest processes a DNS query and returns an appropriate response. +func (s *Server) handleRequest(w dns.ResponseWriter, req *dns.Msg) { + s.inProgressReqs.Add(1) + defer s.inProgressReqs.Done() + + if len(req.Question) == 0 { + resp := new(dns.Msg).SetRcode(req, dns.RcodeFormatError) + s.reply(w, req, resp) + return + } + + // While the original DNS RFCs allow multiple questions, in practice it never works. So handle only the first one. + q := req.Question[0] + log := s.log.With("name", q.Name, "type", dns.TypeToString[q.Qtype]) + log.Debug("Received DNS query.") + + if !dns.IsSubDomain(InternalDomain, dns.CanonicalName(q.Name)) { + log.Debug("Forwarding non-internal DNS query to upstream DNS servers.") + + // Use the same transport for the forwarded request as the original request. + resp, err := s.forwardRequest(req, w.LocalAddr().Network()) + if err != nil { + log.Error("Failed to forward DNS query.", "err", err) + resp = new(dns.Msg).SetRcode(req, dns.RcodeServerFailure) + } + + s.reply(w, req, resp) + return + } + + // Handle the query for the internal domain. + resp := new(dns.Msg).SetReply(req) + resp.Authoritative = true + resp.RecursionAvailable = true + + switch q.Qtype { + case dns.TypeA: + records := s.handleAQuery(q.Name) + if len(records) > 0 { + log.Debug("Found A records for internal DNS query.", "count", len(records)) + resp.Answer = append(resp.Answer, records...) + } else { + log.Debug("No records found for internal DNS query.") + resp.SetRcode(req, dns.RcodeNameError) + } + // TODO: Handle other query types (SRV, TXT, etc.) as needed. + } + + // Truncate the response if it exceeds the maximum size for the transport protocol. + maxSize := dns.MinMsgSize + if w.LocalAddr().Network() == "tcp" { + maxSize = dns.MaxMsgSize + } else { + // Retrieve the UDP buffer size from the EDNS0 record if present. + if opt := req.IsEdns0(); opt != nil { + if udpSize := int(opt.UDPSize()); udpSize > maxSize { + maxSize = udpSize + } + } + } + resp.Truncate(maxSize) + + s.reply(w, req, resp) +} + +func (s *Server) reply(w dns.ResponseWriter, req *dns.Msg, resp *dns.Msg) error { + if err := w.WriteMsg(resp); err != nil { + s.log.Error("Failed to write DNS response.", "err", err, "msg", resp) + // It may fail due to a malformed response message, e.g. exceeding the maximum size. In that case, + // attempt to send a server failure response instead so the client doesn't have to wait for a timeout. + if resp.Rcode != dns.RcodeServerFailure { + resp = new(dns.Msg).SetRcode(req, dns.RcodeServerFailure) + if err2 := w.WriteMsg(resp); err2 != nil { + s.log.Error("Failed to write DNS error response.", "err", err2, "msg", resp) + } + } + return err + } + return nil +} + +// forwardRequest forwards a DNS query to system DNS servers +func (s *Server) forwardRequest(req *dns.Msg, proto string) (*dns.Msg, error) { + if len(s.upstreamServers) == 0 { + return nil, errors.New("no upstream DNS servers configured") + } + + // Apply concurrency control for forwarded queries. + select { + case s.forwardSemaphore <- struct{}{}: + defer func() { + <-s.forwardSemaphore + }() + default: + return nil, fmt.Errorf("too many concurrent forwarded queries (max: %d)", maxConcurrentForwards) + } + + // Create DNS client with timeout and same transport as original request + client := &dns.Client{ + Net: proto, + Timeout: forwardingTimeout, + } + + var lastErr error + for _, server := range s.upstreamServers { + resp, _, err := client.Exchange(req, server.String()) + if err == nil { + return resp, nil + } + lastErr = err + s.log.Debug("Failed to forward DNS query to upstream server.", "server", server, "err", err) + } + + return nil, lastErr +} + +// handleAQuery processes an A query for the internal domain and returns A records for the requested name. +// The internal domain suffix is already stripped from the name. An empty list is returned if no records are found. +func (s *Server) handleAQuery(name string) []dns.RR { + serviceName := trimInternalDomain(name) + ips := s.resolver.Resolve(serviceName) + if len(ips) == 0 { + s.log.Debug("Failed to resolve service name.", "service", serviceName) + return nil + } + // TODO: verify the formatting of the IP addresses. + s.log.Debug("Resolved service name.", "service", serviceName, "ips", ips) + + if len(ips) > 1 { + // TODO: sort by proximity to the requesting container/machine. For now, just shuffle the IPs. + rand.Shuffle(len(ips), func(i, j int) { + ips[i], ips[j] = ips[j], ips[i] + }) + } + + // Create A records for each IP. + records := make([]dns.RR, 0, len(ips)) + for _, ip := range ips { + records = append(records, &dns.A{ + Hdr: dns.RR_Header{ + Name: name, + Rrtype: dns.TypeA, + Class: dns.ClassINET, + // TODO: should we increate the TTL to some reasonably small value like 5-30 seconds to allow + // at least some caching? + Ttl: 0, + }, + A: net.ParseIP(ip.String()), + }) + } + return records +} + +// parseNameserversFromResolvConf parses the nameservers from /etc/resolv.conf. +func parseNameserversFromResolvConf() ([]netip.Addr, error) { + config, err := dns.ClientConfigFromFile("/etc/resolv.conf") + if err != nil { + return nil, err + } + + var servers []netip.Addr + for _, server := range config.Servers { + if addr, err := netip.ParseAddr(server); err == nil { + servers = append(servers, addr) + } + } + + return servers, nil +} + +func trimInternalDomain(name string) string { + name = dns.CanonicalName(name) + if !dns.IsSubDomain(InternalDomain, name) { + return name + } + + return strings.TrimSuffix(name, "."+InternalDomain) }