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
1 change: 1 addition & 0 deletions config/rewriter.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ type RewriterConfig struct {
// Domain rewrite rules applied before resolution; keys are rewritten to their values.
Rewrite map[string]string `yaml:"rewrite"`
// If true, the original query is sent upstream when the mapped resolver returns an empty answer.
// Only has an effect together with Rewrite.
FallbackUpstream bool `default:"false" yaml:"fallbackUpstream"`
}

Expand Down
4 changes: 2 additions & 2 deletions docs/config.schema.json
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@
},
"fallbackUpstream": {
"type": "boolean",
"description": "If true, the original query is sent upstream when the mapped resolver returns an empty answer.",
"description": "If true, the original query is sent upstream when the mapped resolver returns an empty answer.\nOnly has an effect together with Rewrite.",
"default": false
},
"customTTL": {
Expand Down Expand Up @@ -193,7 +193,7 @@
},
"fallbackUpstream": {
"type": "boolean",
"description": "If true, the original query is sent upstream when the mapped resolver returns an empty answer.",
"description": "If true, the original query is sent upstream when the mapped resolver returns an empty answer.\nOnly has an effect together with Rewrite.",
"default": false
},
"mapping": {
Expand Down
10 changes: 6 additions & 4 deletions docs/configuration.md
Original file line number Diff line number Diff line change
Expand Up @@ -438,10 +438,11 @@ them to the allowlist. Entries match the domain itself and all of its subdomains
done on the name of the inspected query, so a `CNAME` inside an upstream answer pointing at an
allowlisted name does not bypass the protection (for `customDNS` `CNAME` entries the inspected
query is the lookup of the CNAME target, so allowlist the **target** name, not the entry's
name). If `customDNS.rewrite` rules apply to the query, matching uses the rewritten name —
allowlist the rewritten form (`conditional.rewrite` rules do not affect matching). Entries must
be plain domain names: wildcards (`*.example.com`), regexes and whitespace are rejected at
startup, and internationalized domains must be given in punycode (`xn--…`) form.
name). Rewrite rules do not affect matching: a `customDNS.rewrite` or `conditional.rewrite`
target is used only for that resolver's own lookup and is never handed down the chain, so
allowlist the name the client asks for. Entries must be plain domain names: wildcards
(`*.example.com`), regexes and whitespace are rejected at startup, and internationalized domains
must be given in punycode (`xn--…`) form.

!!! example

Expand Down Expand Up @@ -650,6 +651,7 @@ hostname belongs to which IP address, all DNS queries for the local network shou
The optional parameter `rewrite` behaves the same as with custom DNS.

The optional parameter `fallbackUpstream`, if false (default), return empty result if after rewrite, the mapped resolver returned an empty answer. If true, the original query will be sent to the upstream resolver.
It only has an effect together with `rewrite`; without any rewrite rules it is ignored.
Comment on lines 653 to +654

**Usage:** One usecase when having split DNS for internal and external (internet facing) users, but not all subdomains are listed in the internal domain

Expand Down
19 changes: 10 additions & 9 deletions resolver/conditional_upstream_resolver.go
Original file line number Diff line number Diff line change
Expand Up @@ -96,20 +96,21 @@ func (r *ConditionalUpstreamResolver) Resolve(ctx context.Context, request *mode
resolved, response, err = r.processRequest(ctx, request)
}

if !resolved && err == nil {
// Revert the request before it leaves this resolver: the rewritten name is
// meant for the mapped upstream, not for the rest of the chain.
request.Req = original

// processRequest only reports `resolved` for a query it sent to a mapped
// upstream, so anything else continues down the chain under its original name.
if !resolved {
logger.WithField("next_resolver", Name(r.next)).Trace("go to next resolver")
response, err = r.next.Resolve(ctx, request)
if err != nil {
return nil, err
}
}

// Revert the request
request.Req = original
return r.next.Resolve(ctx, request)
}

// The mapped resolver failed or had nothing: ask the rest of the chain,
// using the original name (`fallbackUpstream`).
if shouldFallbackUpstream(&r.cfg.RewriterConfig, resolved, response, err) {
if shouldFallbackUpstream(&r.cfg.RewriterConfig, response, err) {
logger.WithField("next_resolver", Name(r.next)).Trace("fallback to next resolver")

return r.next.Resolve(ctx, request)
Expand Down
87 changes: 53 additions & 34 deletions resolver/conditional_upstream_resolver_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package resolver

import (
"context"
"errors"

"github.com/0xERR0R/blocky/config"
. "github.com/0xERR0R/blocky/helpertest"
Expand Down Expand Up @@ -272,17 +273,9 @@ var _ = Describe("ConditionalUpstreamResolver", Label("conditionalResolver"), fu
})

It("should ask the next resolver with the original name", func() {
var seen string
var seen *string

m = &mockResolver{}
m.On("Resolve", mock.Anything).Return(&Response{Res: new(dns.Msg)}, nil)
m.ResolveFn = func(_ context.Context, req *Request) (*Response, error) {
seen = req.Req.Question[0].Name
resp, err := util.NewMsgWithAnswer(seen, 250, A, "192.192.192.192")
Expect(err).Should(Succeed())

return &Response{Res: resp, RType: ResponseTypeRESOLVED, Reason: "RESOLVED"}, nil
}
m, seen = newRecordingResolver(A, "192.192.192.192")
sut.Next(m)

Expect(sut.Resolve(ctx, newRequest("www.source.test.", A))).
Expand All @@ -293,50 +286,34 @@ var _ = Describe("ConditionalUpstreamResolver", Label("conditionalResolver"), fu
HaveReturnCode(dns.RcodeSuccess),
))

Expect(seen).Should(Equal("www.source.test."))
Expect(*seen).Should(Equal("www.source.test."))
})
})

When("the mapped upstream fails and fallbackUpstream is set", func() {
BeforeEach(func() {
deadUpstream := NewMockUDPUpstreamServer().WithAnswerRR("dead.box. 123 IN A 1.2.3.4")
upstream := deadUpstream.Start()
deadUpstream.Close()

sutConfig.FallbackUpstream = true
sutConfig.Mapping.Upstreams["dead.box"] = []config.Upstream{upstream}
sutConfig.Rewrite = map[string]string{"source.test": "dead.box"}
sutConfig.Mapping.Upstreams["broken.box"] = []config.Upstream{NewBrokenUDPUpstreamServer()}
sutConfig.Rewrite = map[string]string{"source.test": "broken.box"}
})

It("should ask the next resolver with the original name", func() {
var seen string
var seen *string

m = &mockResolver{}
m.On("Resolve", mock.Anything).Return(&Response{Res: new(dns.Msg)}, nil)
m.ResolveFn = func(_ context.Context, req *Request) (*Response, error) {
seen = req.Req.Question[0].Name
resp, err := util.NewMsgWithAnswer(seen, 250, A, "192.192.192.192")
Expect(err).Should(Succeed())

return &Response{Res: resp, RType: ResponseTypeRESOLVED, Reason: "RESOLVED"}, nil
}
m, seen = newRecordingResolver(A, "192.192.192.192")
sut.Next(m)

Expect(sut.Resolve(ctx, newRequest("www.source.test.", A))).
Should(HaveResponseType(ResponseTypeRESOLVED))

Expect(seen).Should(Equal("www.source.test."))
Expect(*seen).Should(Equal("www.source.test."))
})
})

When("the mapped upstream fails and fallbackUpstream is not set", func() {
BeforeEach(func() {
deadUpstream := NewMockUDPUpstreamServer().WithAnswerRR("dead.box. 123 IN A 1.2.3.4")
upstream := deadUpstream.Start()
deadUpstream.Close()

sutConfig.Mapping.Upstreams["dead.box"] = []config.Upstream{upstream}
sutConfig.Rewrite = map[string]string{"source.test": "dead.box"}
sutConfig.Mapping.Upstreams["broken.box"] = []config.Upstream{NewBrokenUDPUpstreamServer()}
sutConfig.Rewrite = map[string]string{"source.test": "broken.box"}
})

It("should return the error", func() {
Expand Down Expand Up @@ -372,6 +349,48 @@ var _ = Describe("ConditionalUpstreamResolver", Label("conditionalResolver"), fu
})
})

When("the rewritten name matches no conditional mapping", func() {
BeforeEach(func() {
sutConfig.Rewrite = map[string]string{"source.test": "nomatch.example"}
})

It("should ask the next resolver with the original name", func() {
var seen *string

m, seen = newRecordingResolver(A, "192.192.192.192")
sut.Next(m)

Expect(sut.Resolve(ctx, newRequest("www.source.test.", A))).
Should(
SatisfyAll(
BeDNSRecord("www.source.test.", A, "192.192.192.192"),
HaveResponseType(ResponseTypeRESOLVED),
HaveReturnCode(dns.RcodeSuccess),
))

Expect(*seen).Should(Equal("www.source.test."))
Expect(m.Calls).Should(HaveLen(1))
})
})

When("the next resolver fails for a rewritten query", func() {
BeforeEach(func() {
sutConfig.Rewrite = map[string]string{"source.test": "nomatch.example"}
})

It("should restore the original request before returning", func() {
m = &mockResolver{}
m.On("Resolve", mock.Anything).Return(nil, errors.New("next resolver failed"))
sut.Next(m)

request := newRequest("www.source.test.", A)

_, err := sut.Resolve(ctx, request)
Expect(err).Should(HaveOccurred())
Expect(request.Req.Question[0].Name).Should(Equal("www.source.test."))
})
})

When("request does not match rewrite rule but matches conditional mapping", func() {
It("should not rewrite and resolve via conditional upstream", func() {
// Direct request to fritz.box (no rewrite)
Expand Down
52 changes: 34 additions & 18 deletions resolver/custom_dns_resolver.go
Original file line number Diff line number Diff line change
Expand Up @@ -127,19 +127,23 @@ func (r *CustomDNSResolver) handleReverseDNS(request *model.Request) *model.Resp
return nil
}

// processRequest resolves a query from the configured mapping. It reports
// whether the mapping handled the query at all; when it did not, the caller
// continues down the chain itself, so that it can restore a rewritten request
// first.
func (r *CustomDNSResolver) processRequest(
ctx context.Context,
logger *logrus.Entry,
request *model.Request,
resolvedCnames []string,
) (*model.Response, error) {
) (handled bool, response *model.Response, err error) {
question := request.Req.Question[0]
domain := util.ExtractDomain(question)
var answers []dns.RR

for len(domain) > 0 {
if err := ctx.Err(); err != nil {
return nil, fmt.Errorf("context cancelled during custom DNS resolution: %w", err)
return true, nil, fmt.Errorf("context cancelled during custom DNS resolution: %w", err)
}

entries, found := r.mapping[domain]
Expand All @@ -148,7 +152,7 @@ func (r *CustomDNSResolver) processRequest(
for _, entry := range entries {
result, err := r.processDNSEntry(ctx, logger, request, resolvedCnames, question, entry)
if err != nil {
return nil, err
return true, nil, err
}

answers = append(answers, result...)
Expand All @@ -160,7 +164,7 @@ func (r *CustomDNSResolver) processRequest(
logFieldDomain: util.Obfuscate(domain),
}).Debugf("returning custom dns entry")

return model.NewResponseWithAnswers(request, answers, model.ResponseTypeCUSTOMDNS, "CUSTOM DNS"), nil
return true, model.NewResponseWithAnswers(request, answers, model.ResponseTypeCUSTOMDNS, "CUSTOM DNS"), nil
}

// Mapping exists for this domain, but for another type
Expand All @@ -171,11 +175,11 @@ func (r *CustomDNSResolver) processRequest(

// return NOERROR/NODATA with an SOA record in the authority section (RFC 2308),
// so caching resolvers know how long to cache the negative result
response := model.NewResponseWithReason(request, model.ResponseTypeCUSTOMDNS, "CUSTOM DNS")
nodata := model.NewResponseWithReason(request, model.ResponseTypeCUSTOMDNS, "CUSTOM DNS")
soa := util.CreateSOAForNegativeResponse(question, r.cfg.CustomTTL.SecondsU32())
response.Res.Ns = []dns.RR{soa}
nodata.Res.Ns = []dns.RR{soa}

return response, nil
return true, nodata, nil
}

if i := strings.IndexRune(domain, '.'); i >= 0 {
Expand All @@ -185,9 +189,7 @@ func (r *CustomDNSResolver) processRequest(
}
}

logger.WithField("next_resolver", Name(r.next)).Trace("go to next resolver")

return r.next.Resolve(ctx, request)
return false, nil, nil
}

func (r *CustomDNSResolver) processDNSEntry(
Expand Down Expand Up @@ -230,16 +232,21 @@ func (r *CustomDNSResolver) Resolve(ctx context.Context, request *model.Request)
request.Req = rewritten
}

response, err := r.processRequest(ctx, logger, request, make([]string, 0, len(r.cfg.Mapping)))
handled, response, err := r.processRequest(ctx, logger, request, make([]string, 0, len(r.cfg.Mapping)))

// Revert the request
// Revert the request before it leaves this resolver: the rewritten name is
// meant for the mapping, not for the rest of the chain.
request.Req = original

// A response we produced ourselves without an answer means the mapping had
// nothing for this query: ask the rest of the chain, using the original
// name (`fallbackUpstream`).
answered := err == nil && response != nil && response.RType == model.ResponseTypeCUSTOMDNS
if shouldFallbackUpstream(&r.cfg.RewriterConfig, answered, response, nil) {
if !handled {
logger.WithField("next_resolver", Name(r.next)).Trace("go to next resolver")

return r.next.Resolve(ctx, request)
}

// The mapping failed or had nothing for this query: ask the rest of the
// chain, using the original name (`fallbackUpstream`).
if shouldFallbackUpstream(&r.cfg.RewriterConfig, response, err) {
logger.WithField("next_resolver", Name(r.next)).Trace("fallback to next resolver")

return r.next.Resolve(ctx, request)
Expand Down Expand Up @@ -333,11 +340,20 @@ func (r *CustomDNSResolver) processCNAME(
targetRequest := newRequestWithClientID(targetWithoutDot, dns.Type(question.Qtype), clientIP, clientID)

// resolve the target recursively
targetResp, err := r.processRequest(ctx, logger, targetRequest, cnames)
handled, targetResp, err := r.processRequest(ctx, logger, targetRequest, cnames)
if err != nil {
return nil, err
}

if !handled {
// the target is outside the mapping: resolve it via the rest of the chain
logger.WithField("next_resolver", Name(r.next)).Trace("go to next resolver")

if targetResp, err = r.next.Resolve(ctx, targetRequest); err != nil {
return nil, err
}
}

// If target resolution returns NoResponse, just return the CNAME record itself
if targetResp == NoResponse {
return result, nil
Expand Down
Loading