Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 11 additions & 8 deletions proxy/cache.go
Original file line number Diff line number Diff line change
Expand Up @@ -51,9 +51,10 @@ type cache struct {
// cacheItem is a single cache entry. It's a helper type to aggregate the
// item-specific logic.
type cacheItem struct {
m *dns.Msg
u string
ttl uint32
m *dns.Msg
u string
ttl uint32
responseAD bool
}

// respToItem converts the pair of the response and upstream resolved the one
Expand All @@ -70,9 +71,10 @@ func (c *cache) respToItem(m *dns.Msg, u upstream.Upstream, l *slog.Logger) (ite
}

return &cacheItem{
m: m,
u: upsAddr,
ttl: ttl,
m: m,
u: upsAddr,
ttl: ttl,
responseAD: m.AuthenticatedData,
}
}

Expand Down Expand Up @@ -158,8 +160,9 @@ func (c *cache) unpackItem(data []byte, req *dns.Msg) (ci *cacheItem, expired bo
filterMsg(res, m, req.AuthenticatedData, doBit, ttl)

return &cacheItem{
m: res,
u: string(b.Next(b.Len())),
m: res,
u: string(b.Next(b.Len())),
responseAD: m.AuthenticatedData,
}, expired
}

Expand Down
10 changes: 10 additions & 0 deletions proxy/dnscontext.go
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,10 @@ type DNSContext struct {
// instance.
RequestID uint64

// responseAD is the authenticated-data flag from the response before
// client-specific filtering.
responseAD bool

// udpSize is the UDP buffer size from request's EDNS0 RR if presented,
// or default otherwise.
udpSize uint16
Expand Down Expand Up @@ -142,6 +146,12 @@ func (dctx *DNSContext) QueryStatistics() (s *QueryStatistics) {
return dctx.queryStatistics
}

// ResponseAD reports whether the response from an upstream or the cache had
// the authenticated-data flag set before client-specific filtering.
func (dctx *DNSContext) ResponseAD() (ok bool) {
return dctx.responseAD
}

// calcFlagsAndSize lazily calculates some values required for Resolve method.
func (dctx *DNSContext) calcFlagsAndSize() {
if dctx.udpSize != 0 || dctx.Req == nil {
Expand Down
2 changes: 2 additions & 0 deletions proxy/pending.go
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,7 @@ func (pr *defaultPendingRequests) queue(
// each request.
dctx.queryStatistics = origDNSCtx.queryStatistics
dctx.Upstream = origDNSCtx.Upstream
dctx.responseAD = origDNSCtx.responseAD
if origDNSCtx.Res != nil {
// TODO(e.burkov): Add cloner for DNS messages.
dctx.Res = origDNSCtx.Res.Copy().SetReply(dctx.Req)
Expand Down Expand Up @@ -117,6 +118,7 @@ func (pr *defaultPendingRequests) done(ctx context.Context, dctx *DNSContext, er
cloneCtx := &DNSContext{
Upstream: dctx.Upstream,
queryStatistics: dctx.queryStatistics,
responseAD: dctx.responseAD,
}

if dctx.Res != nil {
Expand Down
37 changes: 37 additions & 0 deletions proxy/pending_internal_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
package proxy

import (
"testing"

"github.com/AdguardTeam/golibs/testutil"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestDefaultPendingRequests_ResponseAD(t *testing.T) {
ctx := testutil.ContextWithTimeout(t, defaultTimeout)
pr := newDefaultPendingRequests()
first := &DNSContext{
Req: newHostTestMessage("example.com"),
responseAD: true,
}

loaded, err := pr.queue(ctx, first)
require.NoError(t, err)
assert.False(t, loaded)

key := string(msgToKey(first.Req))
pending, ok := pr.storage.Load(key)
require.True(t, ok)

pr.done(ctx, first, nil)
require.NotNil(t, pending.cloneDNSCtx)
assert.True(t, pending.cloneDNSCtx.responseAD)

pr.storage.Store(key, pending)
second := &DNSContext{Req: first.Req.Copy()}
loaded, err = pr.queue(ctx, second)
require.NoError(t, err)
assert.True(t, loaded)
assert.True(t, second.ResponseAD())
}
3 changes: 3 additions & 0 deletions proxy/proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -659,6 +659,7 @@ func (p *Proxy) handleExchangeResult(
}

d.Upstream = u
d.responseAD = resp.AuthenticatedData
d.Res = resp
d.Res.Authoritative = false

Expand Down Expand Up @@ -690,6 +691,8 @@ const defaultUDPBufSize = 2048
// Resolve is the default resolving method used by the DNS proxy to query
// upstream servers. It expects dctx is filled with the client's request.
func (p *Proxy) Resolve(ctx context.Context, dctx *DNSContext) (err error) {
dctx.responseAD = false

if p.EnableEDNSClientSubnet {
dctx.processECS(p.EDNSAddr, p.logger)
}
Expand Down
65 changes: 65 additions & 0 deletions proxy/proxy_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -587,6 +587,71 @@ func TestProxy_Resolve_dnssecCache(t *testing.T) {
}
}

func TestDNSContext_ResponseAD(t *testing.T) {
const host = "example.com"

var exchangeCount int
u := &dnsproxytest.Upstream{
OnExchange: func(req *dns.Msg) (resp *dns.Msg, err error) {
exchangeCount++

resp = (&dns.Msg{
MsgHdr: dns.MsgHdr{
AuthenticatedData: true,
},
Answer: []dns.RR{&dns.A{
Hdr: dns.RR_Header{
Name: dns.Fqdn(host),
Rrtype: dns.TypeA,
Class: dns.ClassINET,
Ttl: defaultTestTTL,
},
A: net.IP{1, 2, 3, 4},
}},
}).SetReply(req)

return resp, nil
},
OnAddress: func() (addr string) { return "stub" },
OnClose: func() (err error) { return nil },
}

p := mustNew(t, &Config{
Logger: testLogger,
UDPListenAddr: []*net.UDPAddr{net.UDPAddrFromAddrPort(localhostAnyPort)},
TCPListenAddr: []*net.TCPAddr{net.TCPAddrFromAddrPort(localhostAnyPort)},
UpstreamConfig: &UpstreamConfig{Upstreams: []upstream.Upstream{u}},
TrustedProxies: defaultTrustedProxies,
CacheEnabled: true,
DNSSECEnabled: true,
CacheSizeBytes: defaultCacheSize,
})
dctx := &DNSContext{}

for _, name := range []string{"upstream", "cache"} {
t.Run(name, func(t *testing.T) {
dctx.Req = newHostTestMessage(host)
err := p.Resolve(testutil.ContextWithTimeout(t, defaultTimeout), dctx)
require.NoError(t, err)

assert.False(t, dctx.Res.AuthenticatedData)
assert.True(t, dctx.ResponseAD())
})
}

assert.Equal(t, 1, exchangeCount)

u.OnExchange = func(_ *dns.Msg) (resp *dns.Msg, err error) {
return nil, assert.AnError
}
dctx.Req = newHostTestMessage("failure.example")
err := p.Resolve(testutil.ContextWithTimeout(t, defaultTimeout), dctx)
require.ErrorIs(t, err, assert.AnError)

assert.Equal(t, dns.RcodeServerFailure, dctx.Res.Rcode)
assert.False(t, dctx.ResponseAD())
}

func TestExchangeWithReservedDomains(t *testing.T) {
t.Parallel()

Expand Down
1 change: 1 addition & 0 deletions proxy/proxycache.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ func (p *Proxy) replyFromCache(d *DNSContext) (hit bool) {
}

d.Res = ci.m
d.responseAD = ci.responseAD
d.queryStatistics = cachedQueryStatistics(ci.u)

p.logger.Debug(
Expand Down