diff --git a/feature/conn25/conn25.go b/feature/conn25/conn25.go index 9c02dd350..90c7300db 100644 --- a/feature/conn25/conn25.go +++ b/feature/conn25/conn25.go @@ -1069,9 +1069,9 @@ func (e *extension) sendAddressAssignment(ctx context.Context, as addrs) (tailcf } type dnsResponseRewrite struct { - domain dnsname.FQDN - dst netip.Addr - ttl time.Duration + domain dnsname.FQDN + dst netip.Addr + ttlSeconds uint32 } func makeServFail(logf logger.Logf, h dnsmessage.Header, q dnsmessage.Question) []byte { @@ -1263,7 +1263,7 @@ func (c *Conn25) mapDNSResponse(buf []byte) []byte { } dstAddr = netip.AddrFrom16(r.AAAA) } - answers = append(answers, dnsResponseRewrite{domain: queriedDomain, dst: dstAddr, ttl: time.Second * time.Duration(h.TTL)}) + answers = append(answers, dnsResponseRewrite{domain: queriedDomain, dst: dstAddr, ttlSeconds: h.TTL}) default: // we already checked the question was for a supported type, this answer is unexpected if err := p.SkipAnswer(); err != nil { @@ -1297,7 +1297,7 @@ func (c *client) rewriteDNSResponse(appName string, hdr dnsmessage.Header, quest // make an answer for each rewrite for _, rw := range answers { - as, err := c.reserveAddresses(appName, rw.domain, rw.dst, rw.ttl) + as, err := c.reserveAddresses(appName, rw.domain, rw.dst, time.Duration(rw.ttlSeconds)*time.Second) if err != nil { return nil, err } @@ -1309,12 +1309,12 @@ func (c *client) rewriteDNSResponse(appName string, hdr dnsmessage.Header, quest return nil, err } if rw.dst.Is4() { - rhdr := dnsmessage.ResourceHeader{Name: name, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET, TTL: 0} + rhdr := dnsmessage.ResourceHeader{Name: name, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET, TTL: rw.ttlSeconds} if err := b.AResource(rhdr, dnsmessage.AResource{A: as.magic.As4()}); err != nil { return nil, err } } else if rw.dst.Is6() { - rhdr := dnsmessage.ResourceHeader{Name: name, Type: dnsmessage.TypeAAAA, Class: dnsmessage.ClassINET, TTL: 0} + rhdr := dnsmessage.ResourceHeader{Name: name, Type: dnsmessage.TypeAAAA, Class: dnsmessage.ClassINET, TTL: rw.ttlSeconds} if err := b.AAAAResource(rhdr, dnsmessage.AAAAResource{AAAA: as.magic.As16()}); err != nil { return nil, err } diff --git a/feature/conn25/conn25_test.go b/feature/conn25/conn25_test.go index 0c404c76b..05ba31e00 100644 --- a/feature/conn25/conn25_test.go +++ b/feature/conn25/conn25_test.go @@ -1272,6 +1272,101 @@ func TestMapDNSResponseSetsExpiryBasedOnTTL(t *testing.T) { } +func TestMapDNSResponsePreservesTTL(t *testing.T) { + configuredDomain := "example.com" + domainName := configuredDomain + "." + dnsMessageName := dnsmessage.MustNewName(domainName) + sn := makeSelfNode(t, []appctype.Conn25Attr{{ + Name: "app1", + Connectors: []string{"tag:connector"}, + Domains: []string{configuredDomain}, + }}, appctype.Conn25PoolsAttr{ + V4MagicIPPool: []netipx.IPRange{v4RangeFrom("0", "10")}, + V4TransitIPPool: []netipx.IPRange{v4RangeFrom("40", "50")}, + V6MagicIPPool: []netipx.IPRange{netipx.IPRangeFrom(netip.MustParseAddr("2606:4700::6812:100"), netip.MustParseAddr("2606:4700::6812:1ff"))}, + V6TransitIPPool: []netipx.IPRange{netipx.IPRangeFrom(netip.MustParseAddr("2606:4700::6813:100"), netip.MustParseAddr("2606:4700::6813:1ff"))}, + }, nil) + cfg := mustConfig(t, sn) + + const wantTTL uint32 = 300 + + for _, tt := range []struct { + name string + toMap []byte + }{ + { + name: "typeA", + toMap: makeDNSResponseForSections(t, + []dnsmessage.Question{{Name: dnsMessageName, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET}}, + []dnsmessage.Resource{{ + Header: dnsmessage.ResourceHeader{Name: dnsMessageName, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET, TTL: wantTTL}, + Body: &dnsmessage.AResource{A: netip.MustParseAddr("1.2.3.4").As4()}, + }}, + nil, + ), + }, + { + name: "typeAAAA", + toMap: makeDNSResponseForSections(t, + []dnsmessage.Question{{Name: dnsMessageName, Type: dnsmessage.TypeAAAA, Class: dnsmessage.ClassINET}}, + []dnsmessage.Resource{{ + Header: dnsmessage.ResourceHeader{Name: dnsMessageName, Type: dnsmessage.TypeAAAA, Class: dnsmessage.ClassINET, TTL: wantTTL}, + Body: &dnsmessage.AAAAResource{AAAA: netip.MustParseAddr("2606:4700::6812:1a78").As16()}, + }}, + nil, + ), + }, + { + // Use the TTL in the A record, not the CNAME. + name: "typeA-cname-chain", + toMap: makeDNSResponseForSections(t, + []dnsmessage.Question{{Name: dnsMessageName, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET}}, + []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{Name: dnsMessageName, Type: dnsmessage.TypeCNAME, Class: dnsmessage.ClassINET, TTL: wantTTL + 999}, + Body: &dnsmessage.CNAMEResource{CNAME: dnsmessage.MustNewName("cdn.example.net.")}, + }, + { + Header: dnsmessage.ResourceHeader{Name: dnsmessage.MustNewName("cdn.example.net."), Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET, TTL: wantTTL}, + Body: &dnsmessage.AResource{A: netip.MustParseAddr("1.2.3.4").As4()}, + }, + }, + nil, + ), + }, + { + // Use the TTL in the AAAA record, not the CNAME. + name: "typeAAAA-cname-chain", + toMap: makeDNSResponseForSections(t, + []dnsmessage.Question{{Name: dnsMessageName, Type: dnsmessage.TypeAAAA, Class: dnsmessage.ClassINET}}, + []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{Name: dnsMessageName, Type: dnsmessage.TypeCNAME, Class: dnsmessage.ClassINET, TTL: wantTTL + 999}, + Body: &dnsmessage.CNAMEResource{CNAME: dnsmessage.MustNewName("cdn.example.net.")}, + }, + { + Header: dnsmessage.ResourceHeader{Name: dnsmessage.MustNewName("cdn.example.net."), Type: dnsmessage.TypeAAAA, Class: dnsmessage.ClassINET, TTL: wantTTL}, + Body: &dnsmessage.AAAAResource{AAAA: netip.MustParseAddr("2606:4700::6812:1a78").As16()}, + }, + }, + nil, + ), + }, + } { + t.Run(tt.name, func(t *testing.T) { + c := newConn25(logger.Discard) + c.reconfig(cfg) + answers, _ := parseResponse(t, c.mapDNSResponse(tt.toMap)) + if len(answers) != 1 { + t.Fatalf("got %d answers, want 1", len(answers)) + } + if got := answers[0].Header.TTL; got != wantTTL { + t.Fatalf("rewritten answer TTL = %d, want %d", got, wantTTL) + } + }) + } +} + func TestNormalizedDNSNames(t *testing.T) { tests := []struct { name string