feat(everything): did a massive dnssec refactor to improve robustness
This commit is contained in:
@@ -61,25 +61,21 @@ func (c *Client) Close() {
|
||||
// The dnscrypt library doesn't require explicit cleanup
|
||||
}
|
||||
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) {
|
||||
if len(msg.Question) > 0 {
|
||||
question := msg.Question[0]
|
||||
logger.Debug("DNSCrypt query: %s %s", question.Name, dns.TypeToString[question.Qtype])
|
||||
}
|
||||
|
||||
if c.config.DNSSEC {
|
||||
msg.SetEdns0(4096, true)
|
||||
}
|
||||
|
||||
response, err := c.resolver.Exchange(msg, c.ri)
|
||||
if err != nil {
|
||||
logger.Error("DNSCrypt query failed: %v", err)
|
||||
return nil, fmt.Errorf("dnscrypt: query failed: %w", err)
|
||||
return msg, nil, fmt.Errorf("dnscrypt: query failed: %w", err)
|
||||
}
|
||||
|
||||
if len(response.Answer) > 0 {
|
||||
logger.Debug("DNSCrypt response: %d answers", len(response.Answer))
|
||||
}
|
||||
|
||||
return response, nil
|
||||
return msg, response, nil
|
||||
}
|
||||
|
||||
+10
-26
@@ -61,7 +61,7 @@ func New(config Config) (*Client, error) {
|
||||
|
||||
tlsConfig := &tls.Config{
|
||||
ServerName: config.Host,
|
||||
MinVersion: tls.VersionTLS13,
|
||||
MinVersion: tls.VersionTLS13,
|
||||
ClientSessionCache: tls.NewLRUClientSessionCache(100),
|
||||
}
|
||||
|
||||
@@ -114,13 +114,6 @@ func New(config Config) (*Client, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
// newSingleUseClient builds a throwaway *http.Client whose transport is torn
|
||||
// down by the returned closeFn. Used for non-persistent (non-keep-alive) mode
|
||||
// so that each query establishes and tears down its own connection, and thus
|
||||
// pays its own TCP+TLS (or QUIC) handshake. Without this, the pooled
|
||||
// HTTP/2 transport silently reuses a single connection across all queries
|
||||
// (HTTP/2 multiplexes and ignores the Connection: close header), which would
|
||||
// make the non-persistent measurement indistinguishable from persistent.
|
||||
func (c *Client) newSingleUseClient() (*http.Client, func()) {
|
||||
tlsConfig := &tls.Config{
|
||||
ServerName: c.config.Host,
|
||||
@@ -162,27 +155,18 @@ func (c *Client) Close() {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) {
|
||||
if len(msg.Question) > 0 {
|
||||
question := msg.Question[0]
|
||||
logger.Debug("DoH query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.upstreamURL.Host)
|
||||
}
|
||||
|
||||
if c.config.DNSSEC {
|
||||
msg.SetEdns0(4096, true)
|
||||
}
|
||||
|
||||
packedMsg, err := msg.Pack()
|
||||
if err != nil {
|
||||
logger.Error("DoH failed to pack DNS message: %v", err)
|
||||
return nil, fmt.Errorf("doh: failed to pack DNS message: %w", err)
|
||||
return msg, nil, fmt.Errorf("doh: failed to pack DNS message: %w", err)
|
||||
}
|
||||
|
||||
// Select the HTTP client. In persistent mode, reuse the pooled client built
|
||||
// in New(). In non-persistent mode, build a fresh client (and transport)
|
||||
// per query and tear it down afterwards, so every query pays its own
|
||||
// handshake — otherwise HTTP/2 would reuse one connection and the
|
||||
// non-persistent measurement would be wrong.
|
||||
httpClient := c.httpClient
|
||||
if !c.config.KeepAlive {
|
||||
freshClient, closeFn := c.newSingleUseClient()
|
||||
@@ -193,7 +177,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
httpReq, err := http.NewRequest(http.MethodPost, c.upstreamURL.String(), bytes.NewReader(packedMsg))
|
||||
if err != nil {
|
||||
logger.Error("DoH failed to create HTTP request: %v", err)
|
||||
return nil, fmt.Errorf("doh: failed to create HTTP request object: %w", err)
|
||||
return msg, nil, fmt.Errorf("doh: failed to create HTTP request object: %w", err)
|
||||
}
|
||||
|
||||
httpReq.Header.Set("User-Agent", "sdns-proxy")
|
||||
@@ -203,36 +187,36 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
httpResp, err := httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
logger.Error("DoH request failed to %s: %v", c.upstreamURL.Host, err)
|
||||
return nil, fmt.Errorf("doh: failed executing HTTP request to %s: %w", c.upstreamURL.Host, err)
|
||||
return msg, nil, fmt.Errorf("doh: failed executing HTTP request to %s: %w", c.upstreamURL.Host, err)
|
||||
}
|
||||
defer httpResp.Body.Close()
|
||||
|
||||
if httpResp.StatusCode != http.StatusOK {
|
||||
logger.Error("DoH received non-200 status from %s: %s", c.upstreamURL.Host, httpResp.Status)
|
||||
return nil, fmt.Errorf("doh: received non-200 HTTP status from %s: %s", c.upstreamURL.Host, httpResp.Status)
|
||||
return msg, nil, fmt.Errorf("doh: received non-200 HTTP status from %s: %s", c.upstreamURL.Host, httpResp.Status)
|
||||
}
|
||||
|
||||
if ct := httpResp.Header.Get("Content-Type"); ct != dnsMessageContentType {
|
||||
logger.Error("DoH unexpected Content-Type from %s: %s", c.upstreamURL.Host, ct)
|
||||
return nil, fmt.Errorf("doh: unexpected Content-Type from %s: got %q, want %q", c.upstreamURL.Host, ct, dnsMessageContentType)
|
||||
return msg, nil, fmt.Errorf("doh: unexpected Content-Type from %s: got %q, want %q", c.upstreamURL.Host, ct, dnsMessageContentType)
|
||||
}
|
||||
|
||||
responseBody, err := io.ReadAll(httpResp.Body)
|
||||
if err != nil {
|
||||
logger.Error("DoH failed reading response from %s: %v", c.upstreamURL.Host, err)
|
||||
return nil, fmt.Errorf("doh: failed reading response body from %s: %w", c.upstreamURL.Host, err)
|
||||
return msg, nil, fmt.Errorf("doh: failed reading response body from %s: %w", c.upstreamURL.Host, err)
|
||||
}
|
||||
|
||||
recvMsg := new(dns.Msg)
|
||||
err = recvMsg.Unpack(responseBody)
|
||||
if err != nil {
|
||||
logger.Error("DoH failed to unpack response from %s: %v", c.upstreamURL.Host, err)
|
||||
return nil, fmt.Errorf("doh: failed to unpack DNS response from %s: %w", c.upstreamURL.Host, err)
|
||||
return msg, nil, fmt.Errorf("doh: failed to unpack DNS response from %s: %w", c.upstreamURL.Host, err)
|
||||
}
|
||||
|
||||
if len(recvMsg.Answer) > 0 {
|
||||
logger.Debug("DoH response from %s: %d answers", c.upstreamURL.Host, len(recvMsg.Answer))
|
||||
}
|
||||
|
||||
return recvMsg, nil
|
||||
return msg, recvMsg, nil
|
||||
}
|
||||
|
||||
+14
-16
@@ -95,7 +95,7 @@ func (c *Client) OpenConnection() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) {
|
||||
if len(msg.Question) > 0 {
|
||||
question := msg.Question[0]
|
||||
logger.Debug("DoQ query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.targetAddr)
|
||||
@@ -104,19 +104,17 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
if c.quicConn == nil {
|
||||
err := c.OpenConnection()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return msg, nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// Prepare DNS message
|
||||
// DoQ requires Id to be 0
|
||||
msg.Id = 0
|
||||
if c.config.DNSSEC {
|
||||
msg.SetEdns0(4096, true)
|
||||
}
|
||||
packed, err := msg.Pack()
|
||||
if err != nil {
|
||||
logger.Error("DoQ failed to pack message: %v", err)
|
||||
return nil, fmt.Errorf("doq: failed to pack message: %w", err)
|
||||
return msg, nil, fmt.Errorf("doq: failed to pack message: %w", err)
|
||||
}
|
||||
|
||||
var quicStream *quic.Stream
|
||||
@@ -125,12 +123,12 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
logger.Debug("DoQ stream failed, reconnecting: %v", err)
|
||||
err = c.OpenConnection()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return msg, nil, err
|
||||
}
|
||||
quicStream, err = c.quicConn.OpenStream()
|
||||
if err != nil {
|
||||
logger.Error("DoQ failed to open stream after reconnect: %v", err)
|
||||
return nil, err
|
||||
return msg, nil, err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -138,18 +136,18 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
err = binary.Write(&lengthPrefixedMessage, binary.BigEndian, uint16(len(packed)))
|
||||
if err != nil {
|
||||
logger.Error("DoQ failed to write message length: %v", err)
|
||||
return nil, fmt.Errorf("failed to write message length: %w", err)
|
||||
return msg, nil, fmt.Errorf("failed to write message length: %w", err)
|
||||
}
|
||||
_, err = lengthPrefixedMessage.Write(packed)
|
||||
if err != nil {
|
||||
logger.Error("DoQ failed to write DNS message: %v", err)
|
||||
return nil, fmt.Errorf("failed to write DNS message: %w", err)
|
||||
return msg, nil, fmt.Errorf("failed to write DNS message: %w", err)
|
||||
}
|
||||
|
||||
_, err = quicStream.Write(lengthPrefixedMessage.Bytes())
|
||||
if err != nil {
|
||||
logger.Error("DoQ failed to write to stream: %v", err)
|
||||
return nil, fmt.Errorf("failed writing to QUIC stream: %w", err)
|
||||
return msg, nil, fmt.Errorf("failed writing to QUIC stream: %w", err)
|
||||
}
|
||||
quicStream.Close()
|
||||
|
||||
@@ -157,32 +155,32 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
_, err = io.ReadFull(quicStream, lengthBuf)
|
||||
if err != nil {
|
||||
logger.Error("DoQ failed to read response length: %v", err)
|
||||
return nil, fmt.Errorf("failed reading response length: %w", err)
|
||||
return msg, nil, fmt.Errorf("failed reading response length: %w", err)
|
||||
}
|
||||
|
||||
messageLength := binary.BigEndian.Uint16(lengthBuf)
|
||||
if messageLength == 0 {
|
||||
logger.Error("DoQ received zero-length message")
|
||||
return nil, fmt.Errorf("received zero-length message")
|
||||
return msg, nil, fmt.Errorf("received zero-length message")
|
||||
}
|
||||
|
||||
responseBuf := make([]byte, messageLength)
|
||||
_, err = io.ReadFull(quicStream, responseBuf)
|
||||
if err != nil {
|
||||
logger.Error("DoQ failed to read response data: %v", err)
|
||||
return nil, fmt.Errorf("failed reading response data: %w", err)
|
||||
return msg, nil, fmt.Errorf("failed reading response data: %w", err)
|
||||
}
|
||||
|
||||
recvMsg := new(dns.Msg)
|
||||
err = recvMsg.Unpack(responseBuf)
|
||||
if err != nil {
|
||||
logger.Error("DoQ failed to parse response: %v", err)
|
||||
return nil, fmt.Errorf("failed to parse DNS response: %w", err)
|
||||
return msg, nil, fmt.Errorf("failed to parse DNS response: %w", err)
|
||||
}
|
||||
|
||||
if len(recvMsg.Answer) > 0 {
|
||||
logger.Debug("DoQ response from %s: %d answers", c.targetAddr, len(recvMsg.Answer))
|
||||
}
|
||||
|
||||
return recvMsg, nil
|
||||
return msg, recvMsg, nil
|
||||
}
|
||||
|
||||
+15
-19
@@ -124,7 +124,7 @@ func (c *Client) ensureConnection() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) {
|
||||
if len(msg.Question) > 0 {
|
||||
question := msg.Question[0]
|
||||
logger.Debug("DoT query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort)
|
||||
@@ -133,7 +133,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
// Ensure we have a connection (either persistent or new)
|
||||
if c.config.KeepAlive {
|
||||
if err := c.ensureConnection(); err != nil {
|
||||
return nil, fmt.Errorf("dot: failed to ensure connection: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to ensure connection: %w", err)
|
||||
}
|
||||
} else {
|
||||
// For non-keepalive mode, create a fresh connection for each query
|
||||
@@ -145,18 +145,14 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
c.connMutex.Unlock()
|
||||
|
||||
if err := c.ensureConnection(); err != nil {
|
||||
return nil, fmt.Errorf("dot: failed to create connection: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to create connection: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Prepare DNS message
|
||||
if c.config.DNSSEC {
|
||||
msg.SetEdns0(4096, true)
|
||||
}
|
||||
packed, err := msg.Pack()
|
||||
if err != nil {
|
||||
logger.Error("DoT failed to pack message: %v", err)
|
||||
return nil, fmt.Errorf("dot: failed to pack message: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to pack message: %w", err)
|
||||
}
|
||||
|
||||
// Prepend message length (DNS over TCP format)
|
||||
@@ -171,7 +167,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
// Write query
|
||||
if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
|
||||
logger.Error("DoT failed to set write deadline: %v", err)
|
||||
return nil, fmt.Errorf("dot: failed to set write deadline: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to set write deadline: %w", err)
|
||||
}
|
||||
|
||||
if _, err := conn.Write(data); err != nil {
|
||||
@@ -181,7 +177,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
if c.config.KeepAlive {
|
||||
logger.Debug("DoT write failed with keep-alive, attempting reconnect")
|
||||
if reconnectErr := c.ensureConnection(); reconnectErr != nil {
|
||||
return nil, fmt.Errorf("dot: failed to reconnect: %w", reconnectErr)
|
||||
return msg, nil, fmt.Errorf("dot: failed to reconnect: %w", reconnectErr)
|
||||
}
|
||||
|
||||
c.connMutex.Lock()
|
||||
@@ -189,48 +185,48 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
c.connMutex.Unlock()
|
||||
|
||||
if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
|
||||
return nil, fmt.Errorf("dot: failed to set write deadline after reconnect: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to set write deadline after reconnect: %w", err)
|
||||
}
|
||||
|
||||
if _, err := conn.Write(data); err != nil {
|
||||
return nil, fmt.Errorf("dot: failed to write message after reconnect: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to write message after reconnect: %w", err)
|
||||
}
|
||||
} else {
|
||||
return nil, fmt.Errorf("dot: failed to write message: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to write message: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Read response
|
||||
if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil {
|
||||
logger.Error("DoT failed to set read deadline: %v", err)
|
||||
return nil, fmt.Errorf("dot: failed to set read deadline: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to set read deadline: %w", err)
|
||||
}
|
||||
|
||||
// Read message length
|
||||
lengthBuf := make([]byte, 2)
|
||||
if _, err := io.ReadFull(conn, lengthBuf); err != nil {
|
||||
logger.Error("DoT failed to read response length from %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("dot: failed to read response length: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to read response length: %w", err)
|
||||
}
|
||||
|
||||
msgLen := binary.BigEndian.Uint16(lengthBuf)
|
||||
if msgLen > dns.MaxMsgSize {
|
||||
logger.Error("DoT response too large from %s: %d bytes", c.hostAndPort, msgLen)
|
||||
return nil, fmt.Errorf("dot: response message too large: %d", msgLen)
|
||||
return msg, nil, fmt.Errorf("dot: response message too large: %d", msgLen)
|
||||
}
|
||||
|
||||
// Read message body
|
||||
buffer := make([]byte, msgLen)
|
||||
if _, err := io.ReadFull(conn, buffer); err != nil {
|
||||
logger.Error("DoT failed to read response from %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("dot: failed to read response: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to read response: %w", err)
|
||||
}
|
||||
|
||||
// Parse response
|
||||
response := new(dns.Msg)
|
||||
if err := response.Unpack(buffer); err != nil {
|
||||
logger.Error("DoT failed to unpack response from %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("dot: failed to unpack response: %w", err)
|
||||
return msg, nil, fmt.Errorf("dot: failed to unpack response: %w", err)
|
||||
}
|
||||
|
||||
if len(response.Answer) > 0 {
|
||||
@@ -247,5 +243,5 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
c.connMutex.Unlock()
|
||||
}
|
||||
|
||||
return response, nil
|
||||
return msg, response, nil
|
||||
}
|
||||
|
||||
@@ -104,7 +104,7 @@ func (c *Client) ensureConnection() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) {
|
||||
if len(msg.Question) > 0 {
|
||||
question := msg.Question[0]
|
||||
logger.Debug("DoTCP query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort)
|
||||
@@ -112,7 +112,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
|
||||
if c.config.KeepAlive {
|
||||
if err := c.ensureConnection(); err != nil {
|
||||
return nil, fmt.Errorf("dotcp: failed to ensure connection: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to ensure connection: %w", err)
|
||||
}
|
||||
} else {
|
||||
c.connMutex.Lock()
|
||||
@@ -123,18 +123,14 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
c.connMutex.Unlock()
|
||||
|
||||
if err := c.ensureConnection(); err != nil {
|
||||
return nil, fmt.Errorf("dotcp: failed to create connection: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to create connection: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if c.config.DNSSEC {
|
||||
msg.SetEdns0(4096, true)
|
||||
}
|
||||
|
||||
packed, err := msg.Pack()
|
||||
if err != nil {
|
||||
logger.Error("DoTCP failed to pack message: %v", err)
|
||||
return nil, fmt.Errorf("dotcp: failed to pack message: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to pack message: %w", err)
|
||||
}
|
||||
|
||||
// DNS over TCP uses 2-byte length prefix
|
||||
@@ -148,7 +144,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
|
||||
if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
|
||||
logger.Error("DoTCP failed to set write deadline: %v", err)
|
||||
return nil, fmt.Errorf("dotcp: failed to set write deadline: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to set write deadline: %w", err)
|
||||
}
|
||||
|
||||
if _, err := conn.Write(data); err != nil {
|
||||
@@ -157,7 +153,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
if c.config.KeepAlive {
|
||||
logger.Debug("DoTCP write failed with keep-alive, attempting reconnect")
|
||||
if reconnectErr := c.ensureConnection(); reconnectErr != nil {
|
||||
return nil, fmt.Errorf("dotcp: failed to reconnect: %w", reconnectErr)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to reconnect: %w", reconnectErr)
|
||||
}
|
||||
|
||||
c.connMutex.Lock()
|
||||
@@ -165,44 +161,44 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
c.connMutex.Unlock()
|
||||
|
||||
if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
|
||||
return nil, fmt.Errorf("dotcp: failed to set write deadline after reconnect: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to set write deadline after reconnect: %w", err)
|
||||
}
|
||||
|
||||
if _, err := conn.Write(data); err != nil {
|
||||
return nil, fmt.Errorf("dotcp: failed to write message after reconnect: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to write message after reconnect: %w", err)
|
||||
}
|
||||
} else {
|
||||
return nil, fmt.Errorf("dotcp: failed to write message: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to write message: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil {
|
||||
logger.Error("DoTCP failed to set read deadline: %v", err)
|
||||
return nil, fmt.Errorf("dotcp: failed to set read deadline: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to set read deadline: %w", err)
|
||||
}
|
||||
|
||||
lengthBuf := make([]byte, 2)
|
||||
if _, err := io.ReadFull(conn, lengthBuf); err != nil {
|
||||
logger.Error("DoTCP failed to read response length from %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("dotcp: failed to read response length: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to read response length: %w", err)
|
||||
}
|
||||
|
||||
msgLen := binary.BigEndian.Uint16(lengthBuf)
|
||||
if msgLen > dns.MaxMsgSize {
|
||||
logger.Error("DoTCP response too large from %s: %d bytes", c.hostAndPort, msgLen)
|
||||
return nil, fmt.Errorf("dotcp: response message too large: %d", msgLen)
|
||||
return msg, nil, fmt.Errorf("dotcp: response message too large: %d", msgLen)
|
||||
}
|
||||
|
||||
buffer := make([]byte, msgLen)
|
||||
if _, err := io.ReadFull(conn, buffer); err != nil {
|
||||
logger.Error("DoTCP failed to read response from %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("dotcp: failed to read response: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to read response: %w", err)
|
||||
}
|
||||
|
||||
response := new(dns.Msg)
|
||||
if err := response.Unpack(buffer); err != nil {
|
||||
logger.Error("DoTCP failed to unpack response from %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("dotcp: failed to unpack response: %w", err)
|
||||
return msg, nil, fmt.Errorf("dotcp: failed to unpack response: %w", err)
|
||||
}
|
||||
|
||||
if len(response.Answer) > 0 {
|
||||
@@ -218,5 +214,5 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
c.connMutex.Unlock()
|
||||
}
|
||||
|
||||
return response, nil
|
||||
return msg, response, nil
|
||||
}
|
||||
|
||||
@@ -64,7 +64,7 @@ func (c *Client) createConnection() (*net.UDPConn, error) {
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) {
|
||||
if len(msg.Question) > 0 {
|
||||
question := msg.Question[0]
|
||||
logger.Debug("DoUDP query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort)
|
||||
@@ -72,56 +72,52 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
|
||||
|
||||
conn, err := c.createConnection()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("doudp: failed to create connection: %w", err)
|
||||
return msg, nil, fmt.Errorf("doudp: failed to create connection: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
if c.config.DNSSEC {
|
||||
msg.SetEdns0(4096, true)
|
||||
}
|
||||
|
||||
packedMsg, err := msg.Pack()
|
||||
if err != nil {
|
||||
logger.Error("DoUDP failed to pack message: %v", err)
|
||||
return nil, fmt.Errorf("doudp: failed to pack DNS message: %w", err)
|
||||
return msg, nil, fmt.Errorf("doudp: failed to pack DNS message: %w", err)
|
||||
}
|
||||
|
||||
if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
|
||||
logger.Error("DoUDP failed to set write deadline: %v", err)
|
||||
return nil, fmt.Errorf("doudp: failed to set write deadline: %w", err)
|
||||
return msg, nil, fmt.Errorf("doudp: failed to set write deadline: %w", err)
|
||||
}
|
||||
|
||||
if _, err := conn.Write(packedMsg); err != nil {
|
||||
logger.Error("DoUDP failed to send query to %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("doudp: failed to send DNS query: %w", err)
|
||||
return msg, nil, fmt.Errorf("doudp: failed to send DNS query: %w", err)
|
||||
}
|
||||
|
||||
if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil {
|
||||
logger.Error("DoUDP failed to set read deadline: %v", err)
|
||||
return nil, fmt.Errorf("doudp: failed to set read deadline: %w", err)
|
||||
return msg, nil, fmt.Errorf("doudp: failed to set read deadline: %w", err)
|
||||
}
|
||||
|
||||
var buffer []byte
|
||||
if c.config.DNSSEC {
|
||||
buffer = make([]byte, 4096)
|
||||
}else {
|
||||
buffer = make([]byte, 512)
|
||||
bufSize := 512
|
||||
if opt := msg.IsEdns0(); opt != nil {
|
||||
bufSize = max(int(opt.UDPSize()), 512)
|
||||
}
|
||||
buffer := make([]byte, bufSize)
|
||||
|
||||
n, err := conn.Read(buffer)
|
||||
if err != nil {
|
||||
logger.Error("DoUDP failed to read response from %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("doudp: failed to read DNS response: %w", err)
|
||||
return msg, nil, fmt.Errorf("doudp: failed to read DNS response: %w", err)
|
||||
}
|
||||
|
||||
response := new(dns.Msg)
|
||||
if err := response.Unpack(buffer[:n]); err != nil {
|
||||
logger.Error("DoUDP failed to unpack response from %s: %v", c.hostAndPort, err)
|
||||
return nil, fmt.Errorf("doudp: failed to unpack DNS response: %w", err)
|
||||
return msg, nil, fmt.Errorf("doudp: failed to unpack DNS response: %w", err)
|
||||
}
|
||||
|
||||
if len(response.Answer) > 0 {
|
||||
logger.Debug("DoUDP response from %s: %d answers", c.hostAndPort, len(response.Answer))
|
||||
}
|
||||
|
||||
return response, nil
|
||||
return msg, response, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user