Skip to content
Merged
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
2 changes: 1 addition & 1 deletion blocklistloader-http.go
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,7 @@ func (l *HTTPLoader) Load(reset func(), fn func(rule string) error) error {
return listFailed(log, l.opt.AllowFailure, l.loaded, reset, err)
}
l.loaded = true
log.With("load-time", time.Since(start)).Debug("completed loading blocklist")
log.Debug("completed loading blocklist", "load-time", time.Since(start))
return nil
}

Expand Down
2 changes: 1 addition & 1 deletion cache.go
Original file line number Diff line number Diff line change
Expand Up @@ -264,7 +264,7 @@ func (r *Cache) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
}
r.metrics.miss.Add(1)

log.With("resolver", r.resolver.String()).Debug("cache-miss, forwarding")
log.Debug("cache-miss, forwarding", "resolver", r.resolver.String())

// Get a response from upstream. The query is keyed on again below to store
// the answer, which is safe because a resolver may not modify what it is
Expand Down
11 changes: 7 additions & 4 deletions cidr-db.go
Original file line number Diff line number Diff line change
Expand Up @@ -98,12 +98,15 @@ func (m *CidrDB) Match(ip net.IP) (*BlocklistMatch, bool) {
if len(ip) == 0 {
return nil, false
}
trie := m.ip4
if addr := ip.To4(); addr == nil {
rule, ok := m.ip6.hasIP(ip)
return &BlocklistMatch{List: m.name, Rule: rule}, ok
trie = m.ip6
}
rule, ok := m.ip4.hasIP(ip)
return &BlocklistMatch{List: m.name, Rule: rule}, ok
rule, ok := trie.hasIP(ip)
if !ok {
return nil, false
}
return &BlocklistMatch{List: m.name, Rule: rule}, true
}

func (m *CidrDB) Close() error {
Expand Down
17 changes: 17 additions & 0 deletions cidr-db_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -99,3 +99,20 @@ func TestCidrDBV4MappedBoundary(t *testing.T) {
_, ok = def.Match(net.ParseIP("203.0.113.9"))
require.False(t, ok, "::/0 is not a v4 rule")
}

// A query that matches nothing carries no match, as the name databases already
// did. It is checked against every address in a response, so a miss is the
// common case and must not allocate one to throw away.
func TestCidrDBNoMatchIsNil(t *testing.T) {
db, err := NewCidrDB("testlist", NewStaticLoader([]string{"10.0.0.0/8", "2001:db8::/32"}))
require.NoError(t, err)

for _, ip := range []string{"1.2.3.4", "2001:db9::1"} {
match, ok := db.Match(net.ParseIP(ip))
require.False(t, ok, "ip: %s", ip)
require.Nil(t, match, "ip: %s", ip)
}
match, ok := db.Match(net.ParseIP("10.1.2.3"))
require.True(t, ok)
require.Equal(t, "10.0.0.0/8", match.Rule)
}
2 changes: 1 addition & 1 deletion dnslistener.go
Original file line number Diff line number Diff line change
Expand Up @@ -117,7 +117,7 @@ func listenHandler(id, protocol, addr string, r Resolver, allowedNet []*net.IPNe

a := new(dns.Msg)
if isAllowed(allowedNet, ci.SourceIP) {
log.With("resolver", r.String()).Debug("forwarding query to resolver")
log.Debug("forwarding query to resolver", "resolver", r.String())
a, err = r.Resolve(req, ci)
if err != nil {
metrics.err.Add("resolve", 1)
Expand Down
24 changes: 13 additions & 11 deletions dohclient.go
Original file line number Diff line number Diff line change
Expand Up @@ -96,9 +96,12 @@ type DoHClient struct {
id string
endpoint string
template *uritemplates.UriTemplate
client *http.Client
opt DoHClientOptions
metrics *listenerMetrics
// The endpoint a POST goes to. POST carries the query in the body, so the
// template takes no values and expands to the same URL every time.
postURL string
client *http.Client
opt DoHClientOptions
metrics *listenerMetrics
}

var _ Resolver = &DoHClient{}
Expand Down Expand Up @@ -166,10 +169,16 @@ func NewDoHClient(id, endpoint string, opt DoHClientOptions) (*DoHClient, error)
opt.QueryTimeout = defaultQueryTimeout
}

postURL, err := template.Expand(map[string]any{})
if err != nil {
return nil, err
}

return &DoHClient{
id: id,
endpoint: endpoint,
template: template,
postURL: postURL,
client: client,
opt: opt,
metrics: newListenerMetrics("client", id),
Expand Down Expand Up @@ -239,14 +248,7 @@ func (d *DoHClient) do(req *http.Request) (*http.Response, error) {
}

func (d *DoHClient) buildPostRequest(ctx context.Context, msg []byte) (*http.Request, error) {
// The URL could be a template. Process it without values since POST doesn't use variables in the URL.
u, err := d.template.Expand(map[string]any{})
if err != nil {
d.metrics.err.Add("template", 1)
return nil, err
}

req, err := http.NewRequestWithContext(ctx, "POST", u, bytes.NewReader(msg))
req, err := http.NewRequestWithContext(ctx, "POST", d.postURL, bytes.NewReader(msg))
if err != nil {
d.metrics.err.Add("http", 1)
return nil, err
Expand Down
2 changes: 1 addition & 1 deletion dohlistener.go
Original file line number Diff line number Diff line change
Expand Up @@ -362,7 +362,7 @@ func (s *DoHListener) parseAndRespond(b []byte, w http.ResponseWriter, r *http.R
var err error
a := new(dns.Msg)
if isAllowed(s.opt.AllowedNet, ci.SourceIP) {
log.With("resolver", s.r.String()).Debug("forwarding query to resolver")
log.Debug("forwarding query to resolver", "resolver", s.r.String())
a, err = s.r.Resolve(q, ci)
if err != nil {
log.Warn("failed to resolve", "error", err)
Expand Down
6 changes: 3 additions & 3 deletions failback.go
Original file line number Diff line number Diff line change
Expand Up @@ -81,14 +81,14 @@ func (r *FailBack) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
)
for i := 0; i < len(r.resolvers); i++ {
resolver, active := r.current(i)
log.With("resolver", resolver.String()).Debug("forwarding query to resolver")
log.Debug("forwarding query to resolver", "resolver", resolver.String())
r.metrics.route.Add(resolver.String(), 1)
a, err = resolver.Resolve(q, ci)
if err == nil && r.isSuccessResponse(a) { // Return immediately if successful
return a, err
}
log.With("resolver", resolver.String()).Debug("resolver returned failure",
"error", err)
log.Debug("resolver returned failure",
"resolver", resolver.String(), "error", err)
r.metrics.failure.Add(resolver.String(), 1)

r.errorFrom(active)
Expand Down
6 changes: 3 additions & 3 deletions failrotate.go
Original file line number Diff line number Diff line change
Expand Up @@ -53,14 +53,14 @@ func (r *FailRotate) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
)
for i := 0; i < len(r.resolvers); i++ {
resolver, active := r.current()
log.With("resolver", resolver.String()).Debug("forwarding query to resolver")
log.Debug("forwarding query to resolver", "resolver", resolver.String())
r.metrics.route.Add(resolver.String(), 1)
a, err = resolver.Resolve(q, ci)
if err == nil && r.isSuccessResponse(a) { // Return immediately if successful
return a, err
}
log.With("resolver", resolver.String()).Debug("resolver returned failure",
"error", err)
log.Debug("resolver returned failure",
"resolver", resolver.String(), "error", err)
r.metrics.failure.Add(resolver.String(), 1)

r.errorFrom(active)
Expand Down
2 changes: 1 addition & 1 deletion fastest-tcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -214,7 +214,7 @@ func (r *FastestTCP) probe(ctx context.Context, log queryLogger, rrs []dns.RR) <
resultCh <- tcpProbeResult{err: err}
return
}
log.With("ip", ip).With("response-time", time.Since(start)).Debug("tcp probe finished")
log.Debug("tcp probe finished", "ip", ip, "response-time", time.Since(start))
defer c.Close()
resultCh <- tcpProbeResult{rr: rr}
}(rr)
Expand Down
6 changes: 3 additions & 3 deletions fastest.go
Original file line number Diff line number Diff line change
Expand Up @@ -51,11 +51,11 @@ func (r *Fastest) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
for resolverResponse := range responseCh {
resolver, a, err := resolverResponse.r, resolverResponse.a, resolverResponse.err
if err == nil && (a == nil || a.Rcode != dns.RcodeServerFailure) { // Return immediately if successful
log.With("resolver", resolver.String()).Debug("using response from resolver")
log.Debug("using response from resolver", "resolver", resolver.String())
return a, err
}
log.With("resolver", resolver.String()).Debug("resolver returned failure, waiting for next response",
"error", err)
log.Debug("resolver returned failure, waiting for next response",
"resolver", resolver.String(), "error", err)

// If all responses were bad, return the last one
if i++; i >= len(r.resolvers) {
Expand Down
4 changes: 2 additions & 2 deletions geoip-db.go
Original file line number Diff line number Diff line change
Expand Up @@ -100,8 +100,8 @@ func (m *GeoIPDB) Match(ip net.IP) (*BlocklistMatch, bool) {
}

if err := m.geoDB.Lookup(ip, &record); err != nil {
Log.With("ip", ip).Error("failed to lookup ip in geo location database",
"error", err)
Log.Error("failed to lookup ip in geo location database",
"ip", ip, "error", err)
return nil, false
}

Expand Down
6 changes: 3 additions & 3 deletions loadbalance.go
Original file line number Diff line number Diff line change
Expand Up @@ -123,7 +123,7 @@ func (r *LoadBalance) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
resolver := r.resolvers[idx]

r.metrics.route.Add(resolver.String(), 1)
log.With("resolver", resolver.String()).Debug("forwarding query to resolver")
log.Debug("forwarding query to resolver", "resolver", resolver.String())

start := time.Now()
a, err = resolver.Resolve(q.Copy(), ci)
Expand All @@ -133,8 +133,8 @@ func (r *LoadBalance) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
return a, nil
}

log.With("resolver", resolver.String()).Debug("resolver returned failure",
"error", err)
log.Debug("resolver returned failure",
"resolver", resolver.String(), "error", err)
r.metrics.failure.Add(resolver.String(), 1)
// Count every failure as a failover to stay consistent with FailRotate
// (which increments on every failure, including the last/only resolver).
Expand Down
8 changes: 4 additions & 4 deletions pipeline.go
Original file line number Diff line number Diff line change
Expand Up @@ -124,7 +124,7 @@ func (c *Pipeline) start() {
c.metrics.err.Add("inflight_full", 1)
continue
}
log.With("qname", qName(query)).Debug("sending query")
log.Debug("sending query", "qname", qName(query))
c.metrics.query.Add(1)
if err := conn.WriteMsg(query); err != nil {
// Take the request back out of the in-flight queue before
Expand All @@ -138,8 +138,8 @@ func (c *Pipeline) start() {
conn.Close() // throw away this connection, should wake up the reader as well
wg.Done()
c.metrics.err.Add("send_query", 1)
log.With("qname", qName(query)).Debug("failed sending query",
"error", err)
log.Debug("failed sending query",
"qname", qName(query), "error", err)
return
}
case <-done: // the reader ran into an error and we want to stop using this connection
Expand Down Expand Up @@ -193,7 +193,7 @@ func (c *Pipeline) start() {
req := c.inFlight.get(a) // match the answer to an in-flight query
if req == nil {
c.metrics.err.Add("unexpected_a", 1)
log.With("qname", qName(a)).Warn("unexpected answer received, ignoring")
log.Warn("unexpected answer received, ignoring", "qname", qName(a))
continue
}
c.metrics.response.Add(rCode(a), 1)
Expand Down
6 changes: 3 additions & 3 deletions random.go
Original file line number Diff line number Diff line change
Expand Up @@ -57,13 +57,13 @@ func (r *Random) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
}

r.metrics.route.Add(resolver.String(), 1)
log.With("resolver", resolver.String()).Debug("forwarding query to resolver")
log.Debug("forwarding query to resolver", "resolver", resolver.String())
a, err := resolver.Resolve(q, ci)
if err == nil && r.isSuccessResponse(a) { // Return immediately if successful
return a, err
}
log.With("resolver", resolver.String()).Debug("resolver returned failure",
"error", err)
log.Debug("resolver returned failure",
"resolver", resolver.String(), "error", err)
r.metrics.failure.Add(resolver.String(), 1)
r.deactivate(resolver)
}
Expand Down
52 changes: 41 additions & 11 deletions rate-limiter.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,12 +16,24 @@ type RateLimiter struct {
resolver Resolver
RateLimiterOptions

// The masks the prefix lengths describe, built once rather than per query.
mask4, mask6 net.IPMask

mu sync.RWMutex
currWinID int64
counters map[string]*uint
counters map[clientNetwork]*uint
metrics *rateLimiterMetrics
}

// clientNetwork identifies the network a query came from, the client address
// with the configured prefix applied. An array rather than a string so that
// using it as a map key costs no allocation; a v4 address sits in the first
// four bytes with the length saying so.
type clientNetwork struct {
addr [net.IPv6len]byte
len uint8
}

var _ Resolver = &RateLimiter{}

type RateLimiterOptions struct {
Expand Down Expand Up @@ -56,6 +68,8 @@ func NewRateLimiter(id string, resolver Resolver, opt RateLimiterOptions) *RateL
id: id,
resolver: resolver,
RateLimiterOptions: opt,
mask4: net.CIDRMask(int(opt.Prefix4), 32),
mask6: net.CIDRMask(int(opt.Prefix6), 128),
metrics: &rateLimiterMetrics{
query: getVarInt("router", id, "query"),
exceed: getVarInt("router", id, "exceed"),
Expand All @@ -70,13 +84,7 @@ func (r *RateLimiter) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
r.metrics.query.Add(1)

// Apply the desired mask to the client IP to build a key it identify the client (network)
source := ci.SourceIP
if ip4 := source.To4(); len(ip4) == net.IPv4len {
source = source.Mask(net.CIDRMask(int(r.Prefix4), 32))
} else {
source = source.Mask(net.CIDRMask(int(r.Prefix6), 128))
}
key := source.String()
key := r.clientKey(ci.SourceIP)

// Calculate the current (fixed) window
windowID := time.Now().Unix() / int64(r.Window)
Expand All @@ -87,7 +95,7 @@ func (r *RateLimiter) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
// If we have moved on to the next window, re-initialize the counters
if windowID != r.currWinID {
r.currWinID = windowID
r.counters = make(map[string]*uint)
r.counters = make(map[clientNetwork]*uint)
}

// Load the current counter for this client or make a new one
Expand All @@ -107,17 +115,39 @@ func (r *RateLimiter) Resolve(q *dns.Msg, ci ClientInfo) (*dns.Msg, error) {
if reject {
r.metrics.exceed.Add(1)
if r.LimitResolver != nil {
log.With("resolver", r.LimitResolver).Debug("rate-limit exceeded, forwarding to limit-resolver")
log.Debug("rate-limit exceeded, forwarding to limit-resolver", "resolver", r.LimitResolver)
return r.LimitResolver.Resolve(q, ci)
}
r.metrics.drop.Add(1)
log.Debug("rate-limit reached, dropping")
return nil, nil
}
log.With("resolver", r.resolver).Debug("forwarding query to resolver")
log.Debug("forwarding query to resolver", "resolver", r.resolver)
return r.resolver.Resolve(q, ci)
}

// clientKey masks a client address down to the network the limit counts, in a
// form that can be a map key without allocating.
func (r *RateLimiter) clientKey(ip net.IP) clientNetwork {
var k clientNetwork
if ip4 := ip.To4(); len(ip4) == net.IPv4len {
k.len = net.IPv4len
for i := range ip4 {
k.addr[i] = ip4[i] & r.mask4[i]
}
return k
}
ip16 := ip.To16()
if ip16 == nil { // not an address at all, all such queries share a counter
return k
}
k.len = net.IPv6len
for i := range ip16 {
k.addr[i] = ip16[i] & r.mask6[i]
}
return k
}

func (r *RateLimiter) String() string {
return r.id
}
Loading
Loading