fix(dnssec): fix DNSSEC by adding TCP fallback to DoUDP
This commit is contained in:
@@ -5,6 +5,12 @@ import (
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
type Exchanger interface {
|
||||
Query(msg *dns.Msg) (sent, resp *dns.Msg, err error)
|
||||
Close()
|
||||
}
|
||||
type ExchangeFactory func(server string) (Exchanger, error)
|
||||
|
||||
type ValidationStats struct {
|
||||
Queries int
|
||||
BytesSent int
|
||||
@@ -28,8 +34,8 @@ func NewValidator(send SendFunc) *Validator {
|
||||
return &Validator{walker: newTrustWalker(send)}
|
||||
}
|
||||
|
||||
func NewAuthoritativeValidator() *Validator {
|
||||
return &Validator{walker: newIterativeWalker()}
|
||||
func NewAuthoritativeValidator(f ExchangeFactory) *Validator {
|
||||
return &Validator{walker: newIterativeWalker(f)}
|
||||
}
|
||||
|
||||
func (v *Validator) ValidateResponse(msg *dns.Msg, qname string, qtype uint16) error {
|
||||
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/afonsofrancof/sdns-proxy/common/logger"
|
||||
"github.com/miekg/dns"
|
||||
@@ -12,15 +11,15 @@ import (
|
||||
|
||||
// iterativeWalker validates top-down.
|
||||
type iterativeWalker struct {
|
||||
client *dns.Client
|
||||
st ValidationStats
|
||||
ipCache map[string]string
|
||||
newClient ExchangeFactory
|
||||
st ValidationStats
|
||||
ipCache map[string]string
|
||||
}
|
||||
|
||||
func newIterativeWalker() *iterativeWalker {
|
||||
func newIterativeWalker(f ExchangeFactory) *iterativeWalker {
|
||||
return &iterativeWalker{
|
||||
client: &dns.Client{Timeout: 5 * time.Second},
|
||||
ipCache: make(map[string]string),
|
||||
newClient: f,
|
||||
ipCache: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -38,7 +37,13 @@ func (w *iterativeWalker) exchange(server, name string, qtype uint16) (*dns.Msg,
|
||||
}
|
||||
w.st.Queries++
|
||||
|
||||
resp, _, err := w.client.Exchange(m, server)
|
||||
cl, err := w.newClient(server)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer cl.Close()
|
||||
|
||||
_, resp, err := cl.Query(m)
|
||||
if err == nil && resp != nil {
|
||||
if b, perr := resp.Pack(); perr == nil {
|
||||
w.st.BytesReceived += len(b)
|
||||
|
||||
@@ -115,9 +115,43 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) {
|
||||
return msg, nil, fmt.Errorf("doudp: failed to unpack DNS response: %w", err)
|
||||
}
|
||||
|
||||
// RFC 1123 / 7766: a truncated UDP answer must be retried over TCP.
|
||||
if response.Truncated {
|
||||
logger.Debug("DoUDP response from %s truncated (TC set), retrying over TCP", c.hostAndPort)
|
||||
tcpResp, terr := c.queryTCP(msg)
|
||||
if terr != nil {
|
||||
return msg, nil, fmt.Errorf("doudp: TCP fallback failed: %w", terr)
|
||||
}
|
||||
response = tcpResp
|
||||
}
|
||||
|
||||
if len(response.Answer) > 0 {
|
||||
logger.Debug("DoUDP response from %s: %d answers", c.hostAndPort, len(response.Answer))
|
||||
}
|
||||
|
||||
return msg, response, nil
|
||||
}
|
||||
|
||||
func (c *Client) queryTCP(msg *dns.Msg) (*dns.Msg, error) {
|
||||
conn, err := net.DialTimeout("tcp", c.hostAndPort, c.config.WriteTimeout)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("dial tcp: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
co := &dns.Conn{Conn: conn}
|
||||
if err := co.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
|
||||
return nil, fmt.Errorf("set write deadline: %w", err)
|
||||
}
|
||||
if err := co.WriteMsg(msg); err != nil {
|
||||
return nil, fmt.Errorf("write: %w", err)
|
||||
}
|
||||
if err := co.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil {
|
||||
return nil, fmt.Errorf("set read deadline: %w", err)
|
||||
}
|
||||
resp, err := co.ReadMsg()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read: %w", err)
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user