mirror of
https://github.com/coredns/coredns.git
synced 2026-08-20 23:08:28 -04:00
fix(rewrite): preserve original request during rewrites (#8235)
This commit is contained in:
@@ -44,6 +44,8 @@ type edns0SetResponseRule struct {
|
|||||||
code uint16
|
code uint16
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *edns0SetResponseRule) revertRequestExtra() {}
|
||||||
|
|
||||||
func (r *edns0SetResponseRule) RewriteResponse(res *dns.Msg, _ dns.RR) {
|
func (r *edns0SetResponseRule) RewriteResponse(res *dns.Msg, _ dns.RR) {
|
||||||
ednsOpt := res.IsEdns0()
|
ednsOpt := res.IsEdns0()
|
||||||
if ednsOpt == nil {
|
if ednsOpt == nil {
|
||||||
@@ -62,6 +64,8 @@ type edns0ReplaceResponseRule[T dns.EDNS0] struct {
|
|||||||
source T
|
source T
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *edns0ReplaceResponseRule[T]) revertRequestExtra() {}
|
||||||
|
|
||||||
func (r *edns0ReplaceResponseRule[T]) RewriteResponse(res *dns.Msg, _ dns.RR) {
|
func (r *edns0ReplaceResponseRule[T]) RewriteResponse(res *dns.Msg, _ dns.RR) {
|
||||||
ednsOpt := res.IsEdns0()
|
ednsOpt := res.IsEdns0()
|
||||||
if ednsOpt == nil {
|
if ednsOpt == nil {
|
||||||
|
|||||||
@@ -44,6 +44,11 @@ type ResponseRule interface {
|
|||||||
RewriteResponse(res *dns.Msg, rr dns.RR)
|
RewriteResponse(res *dns.Msg, rr dns.RR)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type requestExtraRevertRule interface {
|
||||||
|
ResponseRule
|
||||||
|
revertRequestExtra()
|
||||||
|
}
|
||||||
|
|
||||||
// ResponseRules describes an ordered list of response rules to apply
|
// ResponseRules describes an ordered list of response rules to apply
|
||||||
// after a name rewrite
|
// after a name rewrite
|
||||||
type ResponseRules = []ResponseRule
|
type ResponseRules = []ResponseRule
|
||||||
@@ -54,6 +59,7 @@ type ResponseRules = []ResponseRule
|
|||||||
type ResponseReverter struct {
|
type ResponseReverter struct {
|
||||||
dns.ResponseWriter
|
dns.ResponseWriter
|
||||||
originalQuestion dns.Question
|
originalQuestion dns.Question
|
||||||
|
request *dns.Msg
|
||||||
ResponseRules ResponseRules
|
ResponseRules ResponseRules
|
||||||
revertPolicy RevertPolicy
|
revertPolicy RevertPolicy
|
||||||
}
|
}
|
||||||
@@ -63,6 +69,7 @@ func NewResponseReverter(w dns.ResponseWriter, r *dns.Msg, policy RevertPolicy)
|
|||||||
return &ResponseReverter{
|
return &ResponseReverter{
|
||||||
ResponseWriter: w,
|
ResponseWriter: w,
|
||||||
originalQuestion: r.Question[0],
|
originalQuestion: r.Question[0],
|
||||||
|
request: r,
|
||||||
revertPolicy: policy,
|
revertPolicy: policy,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -90,9 +97,48 @@ func (r *ResponseReverter) WriteMsg(res1 *dns.Msg) error {
|
|||||||
r.rewriteResourceRecord(res, rr)
|
r.rewriteResourceRecord(res, rr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
return r.writeWithRevertedRequestExtra(res)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *ResponseReverter) writeWithRevertedRequestExtra(res *dns.Msg) error {
|
||||||
|
if r.request == nil || !r.hasRequestExtraRevertRule() {
|
||||||
|
return r.ResponseWriter.WriteMsg(res)
|
||||||
|
}
|
||||||
|
|
||||||
|
currentExtra := r.request.Extra
|
||||||
|
req := new(dns.Msg)
|
||||||
|
req.Extra = copyRRs(currentExtra)
|
||||||
|
for _, rr := range req.Extra {
|
||||||
|
r.rewriteRequestExtra(req, rr)
|
||||||
|
}
|
||||||
|
r.request.Extra = req.Extra
|
||||||
|
defer func() {
|
||||||
|
r.request.Extra = currentExtra
|
||||||
|
}()
|
||||||
|
|
||||||
return r.ResponseWriter.WriteMsg(res)
|
return r.ResponseWriter.WriteMsg(res)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *ResponseReverter) hasRequestExtraRevertRule() bool {
|
||||||
|
for _, rule := range r.ResponseRules {
|
||||||
|
if _, ok := rule.(requestExtraRevertRule); ok {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func copyRRs(rrs []dns.RR) []dns.RR {
|
||||||
|
if len(rrs) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
copied := make([]dns.RR, len(rrs))
|
||||||
|
for i, rr := range rrs {
|
||||||
|
copied[i] = dns.Copy(rr)
|
||||||
|
}
|
||||||
|
return copied
|
||||||
|
}
|
||||||
|
|
||||||
func (r *ResponseReverter) rewriteResourceRecord(res *dns.Msg, rr dns.RR) {
|
func (r *ResponseReverter) rewriteResourceRecord(res *dns.Msg, rr dns.RR) {
|
||||||
// The reverting rules need to be done in reversed order.
|
// The reverting rules need to be done in reversed order.
|
||||||
for i := len(r.ResponseRules) - 1; i >= 0; i-- {
|
for i := len(r.ResponseRules) - 1; i >= 0; i-- {
|
||||||
@@ -100,6 +146,17 @@ func (r *ResponseReverter) rewriteResourceRecord(res *dns.Msg, rr dns.RR) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *ResponseReverter) rewriteRequestExtra(req *dns.Msg, rr dns.RR) {
|
||||||
|
// The reverting rules need to be done in reversed order.
|
||||||
|
for i := len(r.ResponseRules) - 1; i >= 0; i-- {
|
||||||
|
rule, ok := r.ResponseRules[i].(requestExtraRevertRule)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
rule.RewriteResponse(req, rr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Write is a wrapper that records the size of the message that gets written.
|
// Write is a wrapper that records the size of the message that gets written.
|
||||||
func (r *ResponseReverter) Write(buf []byte) (int, error) {
|
func (r *ResponseReverter) Write(buf []byte) (int, error) {
|
||||||
n, err := r.ResponseWriter.Write(buf)
|
n, err := r.ResponseWriter.Write(buf)
|
||||||
|
|||||||
@@ -1153,8 +1153,56 @@ func TestRewriteEDNS0Unset(t *testing.T) {
|
|||||||
rec := dnstest.NewRecorder(&test.ResponseWriter{})
|
rec := dnstest.NewRecorder(&test.ResponseWriter{})
|
||||||
rw.ServeDNS(ctx, rec, m)
|
rw.ServeDNS(ctx, rec, m)
|
||||||
|
|
||||||
if !optsEqual(o.Option, tc.toOpts) {
|
respOpt := rec.Msg.IsEdns0()
|
||||||
t.Errorf("Test %d: Expected %v but got %v", i, tc.toOpts, o)
|
if respOpt == nil {
|
||||||
|
t.Errorf("Test %d: EDNS0 options not set", i)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if !optsEqual(respOpt.Option, tc.toOpts) {
|
||||||
|
t.Errorf("Test %d: Expected %v but got %v", i, tc.toOpts, respOpt)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRewriteEDNS0RevertDoesNotLeakThroughScrubWriter(t *testing.T) {
|
||||||
|
rw := Rewrite{
|
||||||
|
Next: plugin.HandlerFunc(func(_ctx context.Context, w dns.ResponseWriter, r *dns.Msg) (int, error) {
|
||||||
|
resp := new(dns.Msg)
|
||||||
|
resp.SetReply(r)
|
||||||
|
return 0, w.WriteMsg(resp)
|
||||||
|
}),
|
||||||
|
RevertPolicy: NewRevertPolicy(false, false),
|
||||||
|
}
|
||||||
|
|
||||||
|
r, err := newEdns0Rule("stop", "local", "set", "0xffee", "0xabcdef", "revert")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Error creating test rule: %s", err)
|
||||||
|
}
|
||||||
|
rw.Rules = []Rule{r}
|
||||||
|
|
||||||
|
m := new(dns.Msg)
|
||||||
|
m.SetQuestion("example.com.", dns.TypeA)
|
||||||
|
m.SetEdns0(4096, false)
|
||||||
|
m.IsEdns0().Option = append(m.IsEdns0().Option, &dns.EDNS0_COOKIE{Code: dns.EDNS0COOKIE, Cookie: "abcdef0123456789"})
|
||||||
|
|
||||||
|
rec := dnstest.NewRecorder(&test.ResponseWriter{})
|
||||||
|
scrub := request.NewScrubWriter(m, rec)
|
||||||
|
rw.ServeDNS(context.TODO(), scrub, m)
|
||||||
|
|
||||||
|
o := rec.Msg.IsEdns0()
|
||||||
|
if o == nil {
|
||||||
|
t.Fatal("expected EDNS0 option record in response")
|
||||||
|
}
|
||||||
|
var foundCookie bool
|
||||||
|
for _, opt := range o.Option {
|
||||||
|
if opt.Option() == 0xffee {
|
||||||
|
t.Fatalf("expected rewritten EDNS0 option to be reverted, got %v", o.Option)
|
||||||
|
}
|
||||||
|
if opt.Option() == dns.EDNS0COOKIE {
|
||||||
|
foundCookie = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !foundCookie {
|
||||||
|
t.Fatalf("expected original EDNS0 cookie option to be preserved, got %v", o.Option)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user