diff --git a/extras/outbounds/dns_standard.go b/extras/outbounds/dns_standard.go index 9dec606..a97b624 100644 --- a/extras/outbounds/dns_standard.go +++ b/extras/outbounds/dns_standard.go @@ -2,7 +2,9 @@ package outbounds import ( "crypto/tls" + "errors" "net" + "strings" "time" "github.com/miekg/dns" @@ -11,8 +13,11 @@ import ( const ( resolverDefaultTimeout = 2 * time.Second standardResolverRetryTimes = 2 + maxCNAMEDepth = 16 ) +var errCNAMEChainTooLong = errors.New("CNAME chain too long") + // standardResolver is a PluggableOutbound DNS resolver that resolves hostnames // using the user-provided DNS server. // Based on "github.com/miekg/dns", it supports UDP, TCP & DNS-over-TLS (TCP). @@ -109,6 +114,18 @@ func (r *standardResolver) skipCNAMEChain(answers []dns.RR) string { // lookup4 resolves a hostname to an IPv4 address. // If there's no IPv4 address, it returns (nil, nil), no error. func (r *standardResolver) lookup4(host string) (net.IP, error) { + return r.lookup4WithCNAMEDepth(host, 0, make(map[string]struct{})) +} + +func (r *standardResolver) lookup4WithCNAMEDepth(host string, depth int, seen map[string]struct{}) (net.IP, error) { + if depth > maxCNAMEDepth { + return nil, errCNAMEChainTooLong + } + key := strings.ToLower(dns.Fqdn(host)) + if _, ok := seen[key]; ok { + return nil, errCNAMEChainTooLong + } + seen[key] = struct{}{} m := new(dns.Msg) m.SetQuestion(dns.Fqdn(host), dns.TypeA) m.RecursionDesired = true @@ -129,7 +146,7 @@ func (r *standardResolver) lookup4(host string) (net.IP, error) { } } if hasCNAME { - return r.lookup4(r.skipCNAMEChain(resp.Answer)) + return r.lookup4WithCNAMEDepth(r.skipCNAMEChain(resp.Answer), depth+1, seen) } else { // Should not happen return nil, nil @@ -139,6 +156,18 @@ func (r *standardResolver) lookup4(host string) (net.IP, error) { // lookup6 resolves a hostname to an IPv6 address. // If there's no IPv6 address, it returns (nil, nil), no error. func (r *standardResolver) lookup6(host string) (net.IP, error) { + return r.lookup6WithCNAMEDepth(host, 0, make(map[string]struct{})) +} + +func (r *standardResolver) lookup6WithCNAMEDepth(host string, depth int, seen map[string]struct{}) (net.IP, error) { + if depth > maxCNAMEDepth { + return nil, errCNAMEChainTooLong + } + key := strings.ToLower(dns.Fqdn(host)) + if _, ok := seen[key]; ok { + return nil, errCNAMEChainTooLong + } + seen[key] = struct{}{} m := new(dns.Msg) m.SetQuestion(dns.Fqdn(host), dns.TypeAAAA) m.RecursionDesired = true @@ -159,7 +188,7 @@ func (r *standardResolver) lookup6(host string) (net.IP, error) { } } if hasCNAME { - return r.lookup6(r.skipCNAMEChain(resp.Answer)) + return r.lookup6WithCNAMEDepth(r.skipCNAMEChain(resp.Answer), depth+1, seen) } else { // Should not happen return nil, nil diff --git a/extras/outbounds/dns_standard_test.go b/extras/outbounds/dns_standard_test.go new file mode 100644 index 0000000..4b3d983 --- /dev/null +++ b/extras/outbounds/dns_standard_test.go @@ -0,0 +1,44 @@ +package outbounds + +import ( + "errors" + "net" + "sync/atomic" + "testing" + "time" + + "github.com/miekg/dns" +) + +func TestStandardResolverRejectsCNAMECycle(t *testing.T) { + var queries atomic.Int32 + mux := dns.NewServeMux() + mux.HandleFunc(".", func(w dns.ResponseWriter, req *dns.Msg) { + queries.Add(1) + q := req.Question[0] + resp := new(dns.Msg) + resp.SetReply(req) + resp.Answer = append(resp.Answer, &dns.CNAME{ + Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 1}, + Target: q.Name, + }) + _ = w.WriteMsg(resp) + }) + + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + server := &dns.Server{PacketConn: pc, Handler: mux} + go func() { _ = server.ActivateAndServe() }() + defer server.Shutdown() + + r := &standardResolver{Addr: pc.LocalAddr().String(), Client: &dns.Client{Timeout: time.Second}} + _, err = r.lookup4("loop.example") + if !errors.Is(err, errCNAMEChainTooLong) { + t.Fatalf("lookup4 error = %v, want %v", err, errCNAMEChainTooLong) + } + if got := queries.Load(); got > 2 { + t.Fatalf("lookup4 followed CNAME cycle for %d queries", got) + } +}