Compare commits

...

10 Commits

28 changed files with 1078 additions and 1702 deletions
+7 -2
View File
@@ -17,5 +17,10 @@
**/tls-key-log.txt **/tls-key-log.txt
/results /results
/results.bak /results.bak
/results_merged /out
dns_results* /results-dnssec
/figures
/tables
/qol
/dns_analysis.ipynb
/sdns-proxy
+46 -8
View File
@@ -1,10 +1,48 @@
.PHONY: all clean PY := python3
PP := scripts/post_processing
all: # Input trees produced by run.sh
python3 scripts/post_processing/merge_files.py GEN_IN := results
python3 scripts/post_processing/merge_cpu.py DNS_IN := results-dnssec
python3 scripts/post_processing/merge_mem.py
python3 scripts/post_processing/merge_pcap.py
clean: # Output layout
rm -f ./dns_results.csv ./dns_results_cpu.csv ./dns_results_mem.csv ./dns_results_pcap.csv OUT := out
GEN_OUT := $(OUT)/general
DNS_OUT := $(OUT)/dnssec
# netns veth IP for pcap direction detection (must match setup-netns.sh NS_IP).
LOCAL_IP := 192.168.100.2
.PHONY: all general dnssec clean clean-general clean-dnssec
all: get_files general dnssec
get_files:
rsync -a --progress afonso@afonso-pi:~/sdns-proxy/results/ $(GEN_IN)
rsync -a --progress afonso@afonso-pi:~/sdns-proxy/results-dnssec/ $(DNS_IN)
# ----- general workload: results/ -> out/general/ -----
general:
mkdir -p $(GEN_OUT)
$(PY) $(PP)/merge_files.py $(GEN_IN) -o $(GEN_OUT)/dns_results.csv
$(PY) $(PP)/merge_cpu.py $(GEN_IN) -o $(GEN_OUT)/dns_results_cpu.csv
$(PY) $(PP)/merge_mem.py $(GEN_IN) -o $(GEN_OUT)/dns_results_mem.csv
$(PY) $(PP)/merge_pcap.py $(GEN_IN) -o $(GEN_OUT)/dns_results_pcap.csv --local-ip $(LOCAL_IP)
# ----- DNSSEC workload: results-dnssec/ -> out/dnssec/ -----
dnssec:
mkdir -p $(DNS_OUT)
$(PY) $(PP)/merge_files.py $(DNS_IN) -o $(DNS_OUT)/dns_results.csv
$(PY) $(PP)/merge_cpu.py $(DNS_IN) -o $(DNS_OUT)/dns_results_cpu.csv
$(PY) $(PP)/merge_mem.py $(DNS_IN) -o $(DNS_OUT)/dns_results_mem.csv
$(PY) $(PP)/merge_pcap.py $(DNS_IN) -o $(DNS_OUT)/dns_results_pcap.csv --local-ip $(LOCAL_IP)
clean: clean-general clean-dnssec
clean-general:
rm -f $(GEN_OUT)/dns_results.csv $(GEN_OUT)/dns_results_cpu.csv \
$(GEN_OUT)/dns_results_mem.csv $(GEN_OUT)/dns_results_pcap.csv
clean-dnssec:
rm -f $(DNS_OUT)/dns_results.csv $(DNS_OUT)/dns_results_cpu.csv \
$(DNS_OUT)/dns_results_mem.csv $(DNS_OUT)/dns_results_pcap.csv
+4 -109
View File
@@ -6,7 +6,6 @@ import (
"net/url" "net/url"
"strings" "strings"
"github.com/afonsofrancof/sdns-proxy/common/dnssec"
"github.com/afonsofrancof/sdns-proxy/common/logger" "github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/afonsofrancof/sdns-proxy/common/protocols/dnscrypt" "github.com/afonsofrancof/sdns-proxy/common/protocols/dnscrypt"
"github.com/afonsofrancof/sdns-proxy/common/protocols/doh" "github.com/afonsofrancof/sdns-proxy/common/protocols/doh"
@@ -18,16 +17,10 @@ import (
) )
type DNSClient interface { type DNSClient interface {
Query(msg *dns.Msg) (*dns.Msg, error) Query(msg *dns.Msg) (sent *dns.Msg, resp *dns.Msg, err error)
Close() Close()
} }
type ValidatingDNSClient struct {
client DNSClient
validator *dnssec.Validator
options Options
}
type Options struct { type Options struct {
DNSSEC bool DNSSEC bool
AuthoritativeDNSSEC bool AuthoritativeDNSSEC bool
@@ -39,7 +32,6 @@ type Options struct {
func New(upstream string, opts Options) (DNSClient, error) { func New(upstream string, opts Options) (DNSClient, error) {
logger.Debug("Creating DNS client for upstream: %s with options: %+v", upstream, opts) logger.Debug("Creating DNS client for upstream: %s with options: %+v", upstream, opts)
// Try to parse as URL
parsedURL, err := url.Parse(upstream) parsedURL, err := url.Parse(upstream)
if err != nil { if err != nil {
logger.Error("Invalid upstream format: %v", err) logger.Error("Invalid upstream format: %v", err)
@@ -48,12 +40,10 @@ func New(upstream string, opts Options) (DNSClient, error) {
var baseClient DNSClient var baseClient DNSClient
// If it has a scheme, treat it as a full URL
if parsedURL.Scheme != "" { if parsedURL.Scheme != "" {
logger.Debug("Parsing %s as URL with scheme %s", upstream, parsedURL.Scheme) logger.Debug("Parsing %s as URL with scheme %s", upstream, parsedURL.Scheme)
baseClient, err = createClientFromURL(parsedURL, opts) baseClient, err = createClientFromURL(parsedURL, opts)
} else { } else {
// No scheme - treat as plain DNS address (defaults to UDP)
logger.Debug("Parsing %s as plain DNS address", upstream) logger.Debug("Parsing %s as plain DNS address", upstream)
baseClient, err = createClientFromPlainAddress(upstream, opts) baseClient, err = createClientFromPlainAddress(upstream, opts)
} }
@@ -63,105 +53,14 @@ func New(upstream string, opts Options) (DNSClient, error) {
return nil, err return nil, err
} }
// If DNSSEC is not enabled, return the base client // Without DNSSEC, the base protocol client is returned directly.
if !opts.DNSSEC { if !opts.DNSSEC {
logger.Debug("DNSSEC disabled, returning base client") logger.Debug("DNSSEC disabled, returning base client")
return baseClient, nil return baseClient, nil
} }
logger.Debug("DNSSEC enabled, wrapping with validator (AuthoritativeDNSSEC: %v)", opts.AuthoritativeDNSSEC) // With DNSSEC, wrap the base client with a validating client.
return NewDNSSECClient(baseClient, opts), nil
var validator *dnssec.Validator
if opts.AuthoritativeDNSSEC {
validator = dnssec.NewValidatorWithAuthoritativeQueries()
} else {
validator = dnssec.NewValidator(func(qname string, qtype uint16) (*dns.Msg, error) {
msg := new(dns.Msg)
msg.SetQuestion(dns.Fqdn(qname), qtype)
msg.Id = dns.Id()
msg.RecursionDesired = true
msg.SetEdns0(4096, true)
return baseClient.Query(msg)
})
}
return &ValidatingDNSClient{
client: baseClient,
validator: validator,
options: opts,
}, nil
}
func (v *ValidatingDNSClient) Query(msg *dns.Msg) (*dns.Msg, error) {
if len(msg.Question) > 0 {
question := msg.Question[0]
logger.Debug("ValidatingDNSClient query: %s %s (DNSSEC: %v, AuthoritativeDNSSEC: %v, ValidateOnly: %v, StrictValidation: %v)",
question.Name, dns.TypeToString[question.Qtype], v.options.DNSSEC, v.options.AuthoritativeDNSSEC, v.options.ValidateOnly, v.options.StrictValidation)
}
// Always query the upstream first
response, err := v.client.Query(msg)
if err != nil {
logger.Debug("Base client query failed: %v", err)
return nil, err
}
// If DNSSEC validation is disabled, return response as-is
if !v.options.DNSSEC {
return response, nil
}
// Extract question details for validation
if len(msg.Question) == 0 {
logger.Debug("No questions in message, skipping DNSSEC validation")
return response, nil
}
question := msg.Question[0]
qname := question.Name
qtype := question.Qtype
logger.Debug("Starting DNSSEC validation for %s %s", qname, dns.TypeToString[qtype])
// Validate the response
validationErr := v.validator.ValidateResponse(response, qname, qtype)
// Handle validation results based on options
if validationErr != nil {
// Check if it's a "not signed" error
if validationErr == dnssec.ErrResourceNotSigned {
logger.Debug("Domain %s is not DNSSEC signed", qname)
if v.options.ValidateOnly {
logger.Error("Domain %s is not DNSSEC signed (ValidateOnly mode)", qname)
return nil, fmt.Errorf("domain %s is not DNSSEC signed", qname)
}
// Return unsigned response if not in validate-only mode
logger.Debug("Returning unsigned response for %s", qname)
return response, nil
}
// For other validation errors
logger.Debug("DNSSEC validation failed for %s: %v", qname, validationErr)
if v.options.StrictValidation {
logger.Error("DNSSEC validation failed for %s (strict mode): %v", qname, validationErr)
return nil, fmt.Errorf("DNSSEC validation failed for %s: %w", qname, validationErr)
}
// In non-strict mode, log the error but return the response
logger.Debug("DNSSEC validation failed for %s (non-strict mode), returning response anyway: %v", qname, validationErr)
return response, nil
}
// Validation successful
logger.Debug("DNSSEC validation successful for %s %s", qname, dns.TypeToString[qtype])
return response, nil
}
func (v *ValidatingDNSClient) Close() {
logger.Debug("Closing ValidatingDNSClient")
if v.client != nil {
v.client.Close()
}
} }
func createClientFromURL(parsedURL *url.URL, opts Options) (DNSClient, error) { func createClientFromURL(parsedURL *url.URL, opts Options) (DNSClient, error) {
@@ -202,7 +101,6 @@ func createClientFromPlainAddress(address string, opts Options) (DNSClient, erro
} }
logger.Debug("Creating client from plain address: host=%s, port=%s", host, port) logger.Debug("Creating client from plain address: host=%s, port=%s", host, port)
// Default to UDP for plain addresses
return createClient("udp", host, port, "", opts) return createClient("udp", host, port, "", opts)
} }
@@ -290,9 +188,6 @@ func createClient(scheme, host, port, path string, opts Options) (DNSClient, err
case "sdns": case "sdns":
config := dnscrypt.Config{ config := dnscrypt.Config{
// Janky solution but whatever
// Here we rejoin them as the client wants them together
// The host is not really a host but whatever
ServerStamp: fmt.Sprintf("%v://%v", scheme, host), ServerStamp: fmt.Sprintf("%v://%v", scheme, host),
DNSSEC: opts.DNSSEC, DNSSEC: opts.DNSSEC,
} }
+103
View File
@@ -0,0 +1,103 @@
package client
import (
"fmt"
"net"
"github.com/afonsofrancof/sdns-proxy/common/dnssec"
"github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/afonsofrancof/sdns-proxy/common/protocols/doudp"
"github.com/miekg/dns"
)
type DNSSECClient struct {
client DNSClient
options Options
validator *dnssec.Validator
}
func NewDNSSECClient(base DNSClient, opts Options) DNSClient {
logger.Debug("Wrapping base client with DNSSEC validator (AuthoritativeDNSSEC: %v)", opts.AuthoritativeDNSSEC)
var validator *dnssec.Validator
if opts.AuthoritativeDNSSEC {
validator = dnssec.NewAuthoritativeValidator(func(server string) (dnssec.Exchanger, error) {
host, port, err := net.SplitHostPort(server)
if err != nil {
host, port = server, "53"
}
return doudp.New(doudp.Config{
HostAndPort: net.JoinHostPort(host, port),
DNSSEC: true,
})
})
} else {
validator = dnssec.NewValidator(func(m *dns.Msg) (*dns.Msg, error) { _, r, e := base.Query(m); return r, e })
}
return &DNSSECClient{client: base, validator: validator, options: opts}
}
func (v *DNSSECClient) LastValidation() dnssec.ValidationStats {
if v.validator == nil {
return dnssec.ValidationStats{}
}
return v.validator.TakeStats()
}
func (v *DNSSECClient) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) {
if len(msg.Question) > 0 {
question := msg.Question[0]
logger.Debug("DNSSECClient query: %s %s (DNSSEC: %v, AuthoritativeDNSSEC: %v, ValidateOnly: %v, StrictValidation: %v)",
question.Name, dns.TypeToString[question.Qtype], v.options.DNSSEC, v.options.AuthoritativeDNSSEC, v.options.ValidateOnly, v.options.StrictValidation)
}
// Query the base client
sent, response, err := v.client.Query(msg)
if err != nil {
logger.Debug("Base client query failed: %v", err)
return sent, nil, err
}
if len(msg.Question) == 0 {
logger.Debug("No questions in message, skipping DNSSEC validation")
return sent, response, nil
}
question := sent.Question[0]
qname := question.Name
qtype := question.Qtype
logger.Debug("Starting DNSSEC validation for %s %s", qname, dns.TypeToString[qtype])
validationErr := v.validator.ValidateResponse(response, qname, qtype)
if validationErr != nil {
// Unsigned domain: return the response unless validate-only is set.
if validationErr == dnssec.ErrResourceNotSigned {
logger.Debug("Domain %s is not DNSSEC signed", qname)
if v.options.ValidateOnly {
logger.Error("Domain %s is not DNSSEC signed (ValidateOnly mode)", qname)
return sent, nil, fmt.Errorf("domain %s is not DNSSEC signed", qname)
}
return sent, response, nil
}
// Any other validation error.
logger.Debug("DNSSEC validation failed for %s: %v", qname, validationErr)
if v.options.StrictValidation {
logger.Error("DNSSEC validation failed for %s (strict mode): %v", qname, validationErr)
return sent, nil, fmt.Errorf("DNSSEC validation failed for %s: %w", qname, validationErr)
}
logger.Debug("DNSSEC validation failed for %s (non-strict mode), returning response anyway: %v", qname, validationErr)
return sent, response, nil
}
logger.Debug("DNSSEC validation successful for %s %s", qname, dns.TypeToString[qtype])
return sent, response, nil
}
func (v *DNSSECClient) Close() {
logger.Debug("Closing DNSSECClient")
if v.client != nil {
v.client.Close()
}
}
+2 -44
View File
@@ -7,7 +7,6 @@ import (
"github.com/afonsofrancof/sdns-proxy/client" "github.com/afonsofrancof/sdns-proxy/client"
"github.com/afonsofrancof/sdns-proxy/common/logger" "github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/afonsofrancof/sdns-proxy/server"
"github.com/alecthomas/kong" "github.com/alecthomas/kong"
"github.com/miekg/dns" "github.com/miekg/dns"
@@ -16,7 +15,6 @@ import (
var cli struct { var cli struct {
Debug bool `help:"Enable debug logging globally." short:"D" env:"DEBUG"` Debug bool `help:"Enable debug logging globally." short:"D" env:"DEBUG"`
Query QueryCmd `cmd:"" help:"Perform a DNS query (client mode)."` Query QueryCmd `cmd:"" help:"Perform a DNS query (client mode)."`
Listen ListenCmd `cmd:"" help:"Run as a DNS listener/resolver (server mode)."`
} }
type QueryCmd struct { type QueryCmd struct {
@@ -32,18 +30,6 @@ type QueryCmd struct {
KeyLogFile string `help:"Path to TLS key log file (for DoT/DoH/DoQ)." env:"SSLKEYLOGFILE"` KeyLogFile string `help:"Path to TLS key log file (for DoT/DoH/DoQ)." env:"SSLKEYLOGFILE"`
} }
type ListenCmd struct {
Address string `help:"Address to listen on (e.g., :53, :8053)." default:":53"`
Upstream string `help:"Upstream DNS server (e.g., https://1.1.1.1/dns-query, tls://8.8.8.8)." short:"u" required:""`
Fallback string `help:"Fallback DNS server (e.g., https://1.1.1.1/dns-query, tls://8.8.8.8)." short:"f"`
Bootstrap string `help:"Bootstrap DNS server (must be an IP address, e.g., 8.8.8.8, 1.1.1.1)." short:"b"`
DNSSEC bool `help:"Enable DNSSEC for upstream queries." short:"d"`
AuthoritativeDNSSEC bool `help:"Use authoritative DNSSEC validation instead of trusting resolver." short:"a"`
KeepAlive bool `help:"Use persistent connections to upstream servers." short:"k"`
Timeout time.Duration `help:"Timeout for upstream queries." default:"5s"`
Verbose bool `help:"Enable verbose logging." short:"v"`
}
func (q *QueryCmd) Run() error { func (q *QueryCmd) Run() error {
logger.Info("Querying %s for %s type %s (DNSSEC: %v, AuthoritativeDNSSEC: %v, ValidateOnly: %v, StrictValidation: %v, KeepAlive: %v, Timeout: %v)", logger.Info("Querying %s for %s type %s (DNSSEC: %v, AuthoritativeDNSSEC: %v, ValidateOnly: %v, StrictValidation: %v, KeepAlive: %v, Timeout: %v)",
q.Server, q.DomainName, q.QueryType, q.DNSSEC, q.AuthoritativeDNSSEC, q.ValidateOnly, q.StrictValidation, q.KeepAlive, q.Timeout) q.Server, q.DomainName, q.QueryType, q.DNSSEC, q.AuthoritativeDNSSEC, q.ValidateOnly, q.StrictValidation, q.KeepAlive, q.Timeout)
@@ -75,9 +61,10 @@ func (q *QueryCmd) Run() error {
msg.SetQuestion(dns.Fqdn(q.DomainName), qTypeUint) msg.SetQuestion(dns.Fqdn(q.DomainName), qTypeUint)
msg.Id = dns.Id() msg.Id = dns.Id()
msg.RecursionDesired = true msg.RecursionDesired = true
msg.SetEdns0(1232, q.DNSSEC)
logger.Debug("Sending DNS query: ID=%d, Question=%s %s", msg.Id, q.DomainName, q.QueryType) logger.Debug("Sending DNS query: ID=%d, Question=%s %s", msg.Id, q.DomainName, q.QueryType)
recvMsg, err := dnsClient.Query(msg) _, recvMsg, err := dnsClient.Query(msg)
if err != nil { if err != nil {
logger.Error("DNS query failed: %v", err) logger.Error("DNS query failed: %v", err)
return err return err
@@ -89,35 +76,6 @@ func (q *QueryCmd) Run() error {
return nil return nil
} }
func (l *ListenCmd) Run() error {
config := server.Config{
Address: l.Address,
Upstream: l.Upstream,
Fallback: l.Fallback,
Bootstrap: l.Bootstrap,
DNSSEC: l.DNSSEC,
AuthoritativeDNSSEC: l.AuthoritativeDNSSEC,
KeepAlive: l.KeepAlive,
Timeout: l.Timeout,
Verbose: l.Verbose,
}
logger.Debug("Server config: %+v", config)
srv, err := server.New(config)
if err != nil {
logger.Error("Failed to create server: %v", err)
return fmt.Errorf("failed to create server: %w", err)
}
logger.Info("Starting DNS proxy server on %s", l.Address)
logger.Info("Upstream server: %v", l.Upstream)
logger.Info("Fallback server: %v", l.Fallback)
logger.Info("Bootstrap server: %v", l.Bootstrap)
logger.Info("KeepAlive: %v", l.KeepAlive)
return srv.Start()
}
func printResponse(domain, qtype string, msg *dns.Msg) { func printResponse(domain, qtype string, msg *dns.Msg) {
fmt.Println(";; QUESTION SECTION:") fmt.Println(";; QUESTION SECTION:")
+95
View File
@@ -0,0 +1,95 @@
package dnssec
import (
"strings"
"github.com/miekg/dns"
)
// rootAnchor is one IANA root zone trust anchor: a KSK key tag and the SHA-256
// digest of that key. A fetched root KSK is trusted if it matches any anchor.
type rootAnchor struct {
keyTag uint16
digest string // SHA-256, uppercase hex (matches DNSKEY.ToDS output)
}
// rootTrustAnchors holds the currently-valid root KSK trust anchors, as
// published by IANA (root-anchors.xml). Both the KSK-2017 and KSK-2024 keys
// are listed so validation keeps working across the KSK rollover, during which
// the root publishes and may sign with either key. The expired 2010 anchor
// (key tag 19036) is intentionally omitted.
var rootTrustAnchors = []rootAnchor{
{keyTag: 20326, digest: "E06D44B80B8F1D39A95C0B0D7C65D08458E880409BBC683457104237C7F8EC8D"}, // KSK-2017
{keyTag: 38696, digest: "683D2D0ACB8C9B712A1948B27F741219298D0A450D612C483AF444A4C0FB2B16"}, // KSK-2024
}
const rootDigestType uint8 = dns.SHA256
// rootHints holds the IPv4 addresses of the DNS root servers, used to
// bootstrap iterative resolution in chain-of-trust (authoritative) mode.
var rootHints = []string{
// Root servers. {a..m}.root-servers.net
"198.41.0.4:53",
"170.247.170.2:53",
"192.33.4.12:53",
"199.7.91.13:53",
"192.203.230.10:53",
"192.5.5.241:53",
"192.112.36.4:53",
"198.97.190.53:53",
"192.36.148.17:53",
"192.58.128.30:53",
"193.0.14.129:53",
"199.7.83.42:53",
"202.12.27.33:53",
}
// matchesRootAnchor reports whether a root DNSKEY matches any trusted anchor,
// by key tag and SHA-256 digest.
func matchesRootAnchor(k *dns.DNSKEY) bool {
tag := k.KeyTag()
for _, a := range rootTrustAnchors {
if tag != a.keyTag {
continue
}
ds := k.ToDS(rootDigestType)
if ds != nil && strings.EqualFold(ds.Digest, a.digest) {
return true
}
}
return false
}
// verifyRootAnchor confirms that a fetched root DNSKEY RRset contains a KSK
// matching one of the hardcoded trust anchors, and that the RRset is signed by
// a key in the set. On success the returned zone can validate child DS records.
func verifyRootAnchor(dnskeys *RRSet) (*SignedZone, error) {
if dnskeys == nil || dnskeys.IsEmpty() || !dnskeys.IsSigned() {
return nil, ErrDnskeyNotAvailable
}
zone := NewSignedZone(".")
zone.DNSKey = dnskeys
matched := false
for _, rr := range dnskeys.RRs {
k, ok := rr.(*dns.DNSKEY)
if !ok {
continue
}
zone.AddPubKey(k)
if matchesRootAnchor(k) {
matched = true
}
}
if !matched {
return nil, ErrDsInvalid
}
// The DNSKEY RRset must be signed by one of its own keys (the KSK).
if err := zone.VerifyRRSIG(dnskeys); err != nil {
return nil, err
}
return zone, nil
}
-226
View File
@@ -1,226 +0,0 @@
package dnssec
// CODE ADAPTED FROM THIS
// ISC License
//
// Copyright (c) 2012-2016 Peter Banik <peter@froggle.org>
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES
// WITH REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF
// MERCHANTABILITY AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR
// ANY SPECIAL, DIRECT, INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES
// WHATSOEVER RESULTING FROM LOSS OF USE, DATA OR PROFITS, WHETHER IN AN
// ACTION OF CONTRACT, NEGLIGENCE OR OTHER TORTIOUS ACTION, ARISING OUT OF
// OR IN CONNECTION WITH THE USE OR PERFORMANCE OF THIS SOFTWARE.
import (
"fmt"
"strings"
"github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/miekg/dns"
)
type AuthenticationChain struct {
DelegationChain []SignedZone
}
func NewAuthenticationChain() *AuthenticationChain {
return &AuthenticationChain{}
}
func (ac *AuthenticationChain) Populate(domainName string, queryFunc func(string, uint16) (*dns.Msg, error)) error {
// Clean domain name and split into components
domainName = strings.TrimSuffix(domainName, ".")
qnameComponents := strings.Split(domainName, ".")
// Remove empty components
var cleanComponents []string
for _, comp := range qnameComponents {
if comp != "" {
cleanComponents = append(cleanComponents, comp)
}
}
// Build zones from root down to target
// For example.com: [".","com.","example.com."]
zones := []string{"."} // Start with root
// Add each level from TLD down to target
for i := len(cleanComponents) - 1; i >= 0; i-- {
zone := dns.Fqdn(strings.Join(cleanComponents[i:], "."))
zones = append(zones, zone)
}
logger.Debug("Building DNSSEC chain for zones: %v", zones)
ac.DelegationChain = make([]SignedZone, 0, len(zones))
// Query each zone from root down
for i, zoneName := range zones {
logger.Debug("Querying zone: %s", zoneName)
delegation, err := ac.queryDelegation(zoneName, queryFunc)
if err != nil {
return fmt.Errorf("failed to query zone %s: %w", zoneName, err)
}
// Set parent relationship (previous zone in chain is parent)
if i > 0 {
delegation.ParentZone = &ac.DelegationChain[i-1]
}
ac.DelegationChain = append(ac.DelegationChain, *delegation)
}
return nil
}
func (ac *AuthenticationChain) queryDelegation(domainName string, queryFunc func(string, uint16) (*dns.Msg, error)) (*SignedZone, error) {
signedZone := NewSignedZone(domainName)
// Query DNSKEY records
dnskeyRRset, err := ac.queryRRset(domainName, dns.TypeDNSKEY, queryFunc)
if err != nil {
return nil, err
}
signedZone.DNSKey = dnskeyRRset
logger.Debug("Found %d DNSKEY records for %s", len(dnskeyRRset.RRs), domainName)
// Populate public key lookup
for _, rr := range signedZone.DNSKey.RRs {
if dnskey, ok := rr.(*dns.DNSKEY); ok {
signedZone.AddPubKey(dnskey)
logger.Debug("Added DNSKEY for %s: keytag=%d, flags=%d, algorithm=%d", domainName, dnskey.KeyTag(), dnskey.Flags, dnskey.Algorithm)
}
}
// Only query DS records for non-root zones
if domainName != "." {
dsRRset, _ := ac.queryRRset(domainName, dns.TypeDS, queryFunc)
signedZone.DS = dsRRset
if dsRRset != nil && len(dsRRset.RRs) > 0 {
logger.Debug("Found %d DS records for %s", len(dsRRset.RRs), domainName)
for _, rr := range dsRRset.RRs {
if ds, ok := rr.(*dns.DS); ok {
logger.Debug("DS record for %s: keytag=%d", domainName, ds.KeyTag)
}
}
}
} else {
// Root zone has no DS records - trusted by default
signedZone.DS = NewRRSet()
logger.Debug("Root zone - no DS records, trusted by default")
}
return signedZone, nil
}
func (ac *AuthenticationChain) queryRRset(qname string, qtype uint16, queryFunc func(string, uint16) (*dns.Msg, error)) (*RRSet, error) {
r, err := queryFunc(qname, qtype)
if err != nil {
logger.Debug("cannot lookup %v", err)
return NewRRSet(), nil // Return empty RRSet instead of nil
}
if r.Rcode == dns.RcodeNameError {
logger.Debug("no such domain %s", qname)
return NewRRSet(), nil // Return empty RRSet instead of nil
}
result := NewRRSet()
if r.Answer == nil {
return result, nil
}
result.RRs = make([]dns.RR, 0, len(r.Answer))
for _, rr := range r.Answer {
switch t := rr.(type) {
case *dns.RRSIG:
if result.RRSig == nil || t.TypeCovered == qtype {
result.RRSig = t
}
default:
if rr != nil && rr.Header().Rrtype == qtype {
result.RRs = append(result.RRs, rr)
}
}
}
return result, nil
}
func (ac *AuthenticationChain) Verify(answerRRset *RRSet) error {
if len(ac.DelegationChain) == 0 {
return ErrDnskeyNotAvailable
}
// Find the target zone (last in chain)
targetZone := &ac.DelegationChain[len(ac.DelegationChain)-1]
// Verify the answer RRset against target zone's keys
err := targetZone.VerifyRRSIG(answerRRset)
if err != nil {
logger.Debug("Answer RRSIG verification failed: %v", err)
return ErrInvalidRRsig
}
// Validate the chain from root down
for _, zone := range ac.DelegationChain {
logger.Debug("Validating zone: %s", zone.Zone)
// Verify DNSKEY RRset signature
if !zone.HasDNSKeys() {
logger.Debug("No DNSKEYs for zone %s", zone.Zone)
return ErrDnskeyNotAvailable
}
err := zone.VerifyRRSIG(zone.DNSKey)
if err != nil {
logger.Debug("DNSKEY validation failed for %s: %v", zone.Zone, err)
return ErrRrsigValidationError
}
// Skip ALL validation for root - just trust it
if zone.Zone == "." {
logger.Debug("Root zone - trusted by default, no validation performed")
continue
}
// For non-root zones, validate DS records against parent zone
if zone.ParentZone == nil {
logger.Debug("Non-root zone %s has no parent", zone.Zone)
return fmt.Errorf("non-root zone %s has no parent", zone.Zone)
}
if zone.DS == nil || zone.DS.IsEmpty() {
logger.Debug("No DS records for zone %s", zone.Zone)
return ErrDsNotAvailable
}
// Verify DS signature using parent's key
err = zone.ParentZone.VerifyRRSIG(zone.DS)
if err != nil {
logger.Debug("DS signature validation failed for %s: %v", zone.Zone, err)
return ErrRrsigValidationError
}
// Verify DS matches this zone's DNSKEY
err = zone.VerifyDS(zone.DS.RRs)
if err != nil {
logger.Debug("DS-DNSKEY validation failed for %s: %v", zone.Zone, err)
return ErrDsInvalid
}
logger.Debug("Zone %s validated successfully", zone.Zone)
}
logger.Debug("DNSSEC validation successful for entire chain!")
return nil
}
-212
View File
@@ -1,212 +0,0 @@
package dnssec
import (
"fmt"
"net"
"strings"
"time"
"github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/miekg/dns"
)
type AuthoritativeQuerier struct {
client *dns.Client
// Cache of NS records to avoid repeated lookups
nsCache map[string][]string
ipCache map[string]string
}
func NewAuthoritativeQuerier() *AuthoritativeQuerier {
return &AuthoritativeQuerier{
client: &dns.Client{
Timeout: 10 * time.Second,
},
nsCache: make(map[string][]string),
ipCache: make(map[string]string),
}
}
func (aq *AuthoritativeQuerier) QueryAuthoritative(qname string, qtype uint16) (*dns.Msg, error) {
logger.Debug("Querying authoritative servers for %s type %d", qname, qtype)
var zone string
if qtype == dns.TypeDS {
zone = aq.getParentZone(qname)
if zone == "" {
logger.Debug("No parent zone for %s - returning NXDOMAIN for DS query", qname)
msg := &dns.Msg{}
msg.SetRcode(&dns.Msg{}, dns.RcodeNameError)
return msg, nil
}
} else {
zone = aq.findZone(qname)
}
logger.Debug("Determined zone: %s for query %s type %d", zone, qname, qtype)
// Get NS names (not IPs yet)
nsNames, err := aq.findAuthoritativeNSNames(zone)
if err != nil {
return nil, fmt.Errorf("failed to find authoritative servers: %w", err)
}
// Try servers one by one, resolving IPs lazily
var lastErr error
for _, nsName := range nsNames {
server := aq.resolveNSToIP(nsName)
if server == "" {
continue
}
logger.Debug("Trying server: %s (%s)", server, nsName)
msg, err := aq.queryServer(server, qname, qtype)
if err != nil {
logger.Debug("Server %s failed: %v", server, err)
lastErr = err
continue
}
logger.Debug("Server %s responded, authoritative: %v, rcode: %d, answers: %d", server, msg.Authoritative, msg.Rcode, len(msg.Answer))
if (msg.Rcode == dns.RcodeSuccess && len(msg.Answer) > 0) || msg.Rcode == dns.RcodeNameError {
return msg, nil
}
}
if lastErr != nil {
return nil, fmt.Errorf("all servers failed, last error: %w", lastErr)
}
return nil, fmt.Errorf("no authoritative response received")
}
func (aq *AuthoritativeQuerier) findAuthoritativeNSNames(zone string) ([]string, error) {
if nsNames, exists := aq.nsCache[zone]; exists {
logger.Debug("Using cached NS names for %s: %v", zone, nsNames)
return nsNames, nil
}
logger.Debug("Looking for NS records for zone: %s", zone)
// Use a public resolver to find the NS records
resolver := &dns.Client{Timeout: 5 * time.Second}
msg, _, err := resolver.Exchange(&dns.Msg{
MsgHdr: dns.MsgHdr{
Id: dns.Id(),
RecursionDesired: true,
},
Question: []dns.Question{{Name: dns.Fqdn(zone), Qtype: dns.TypeNS, Qclass: dns.ClassINET}},
}, "8.8.8.8:53")
if err != nil {
return nil, fmt.Errorf("failed to query NS records: %w", err)
}
var nsNames []string
// Collect NS records from answer section
for _, rr := range msg.Answer {
if ns, ok := rr.(*dns.NS); ok {
nsNames = append(nsNames, ns.Ns)
}
}
// Also check authority section if answer is empty
if len(nsNames) == 0 {
for _, rr := range msg.Ns {
if ns, ok := rr.(*dns.NS); ok {
nsNames = append(nsNames, ns.Ns)
}
}
}
if len(nsNames) == 0 {
return nil, fmt.Errorf("no NS servers found for %s", zone)
}
logger.Debug("Found NS names for %s: %v", zone, nsNames)
aq.nsCache[zone] = nsNames
logger.Debug("Cached NS names for %s: %v", zone, nsNames)
return nsNames, nil
}
func (aq *AuthoritativeQuerier) getParentZone(qname string) string {
logger.Debug("Getting parent zone for: %s", qname)
// Clean the qname
qname = strings.TrimSuffix(qname, ".")
// Root zone has no parent
if qname == "" || qname == "." {
logger.Debug("Root zone has no parent")
return ""
}
labels := dns.SplitDomainName(qname)
logger.Debug("Labels for %s: %v", qname, labels)
if len(labels) <= 1 {
logger.Debug("Parent of TLD %s is root", qname)
return "." // Parent of TLD is root
}
parentLabels := labels[1:]
parent := dns.Fqdn(strings.Join(parentLabels, "."))
logger.Debug("Parent zone of %s is %s", qname, parent)
return parent
}
func (aq *AuthoritativeQuerier) findZone(qname string) string {
// For now, assume the zone is the domain itself
// In a more sophisticated implementation, you'd walk up the hierarchy
labels := dns.SplitDomainName(qname)
if len(labels) >= 2 {
return labels[len(labels)-2] + "." + labels[len(labels)-1] + "."
}
return qname
}
func (aq *AuthoritativeQuerier) resolveNSToIP(nsName string) string {
if ip, exists := aq.ipCache[nsName]; exists {
logger.Debug("Using cached IP for %s: %s", nsName, ip)
return ip
}
nsName = strings.TrimSuffix(nsName, ".")
logger.Debug("Resolving NS %s to IP", nsName)
ips, err := net.LookupIP(nsName)
if err != nil {
logger.Debug("Failed to resolve %s: %v", nsName, err)
return ""
}
for _, ip := range ips {
if ip.To4() != nil { // Prefer IPv4
result := ip.String() + ":53"
logger.Debug("Resolved %s to %s", nsName, result)
// Cache the result before returning
aq.ipCache[nsName] = result
return result
}
}
return ""
}
func (aq *AuthoritativeQuerier) queryServer(server, qname string, qtype uint16) (*dns.Msg, error) {
m := new(dns.Msg)
m.SetQuestion(dns.Fqdn(qname), qtype)
m.SetEdns0(4096, true) // Enable DNSSEC
logger.Debug("Querying %s for %s type %d", server, qname, qtype)
msg, _, err := aq.client.Exchange(m, server)
if err != nil {
return nil, err
}
logger.Debug("Response from %s: rcode=%d, answers=%d", server, msg.Rcode, len(msg.Answer))
return msg, err
}
+59 -69
View File
@@ -1,100 +1,90 @@
package dnssec package dnssec
// CODE ADAPTED FROM THIS
// ISC License
//
// Copyright (c) 2012-2016 Peter Banik <peter@froggle.org>
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES
// WITH REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF
// MERCHANTABILITY AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR
// ANY SPECIAL, DIRECT, INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES
// WHATSOEVER RESULTING FROM LOSS OF USE, DATA OR PROFITS, WHETHER IN AN
// ACTION OF CONTRACT, NEGLIGENCE OR OTHER TORTIOUS ACTION, ARISING OUT OF
// OR IN CONNECTION WITH THE USE OR PERFORMANCE OF THIS SOFTWARE.
// ./common/dnssec/validator.go
import ( import (
"github.com/afonsofrancof/sdns-proxy/common/logger" "github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/miekg/dns" "github.com/miekg/dns"
) )
type Validator struct { type Exchanger interface {
queryFunc func(string, uint16) (*dns.Msg, error) 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
BytesReceived int
Validated bool
} }
func NewValidator(queryFunc func(string, uint16) (*dns.Msg, error)) *Validator { type SendFunc func(msg *dns.Msg) (*dns.Msg, error)
return &Validator{
queryFunc: queryFunc, type walker interface {
validate(answer *RRSet, qname string, qtype uint16) error
stats() ValidationStats
resetStats()
} }
type Validator struct {
walker walker
}
func NewValidator(send SendFunc) *Validator {
return &Validator{walker: newTrustWalker(send)}
}
func NewAuthoritativeValidator(f ExchangeFactory) *Validator {
return &Validator{walker: newIterativeWalker(f)}
} }
func (v *Validator) ValidateResponse(msg *dns.Msg, qname string, qtype uint16) error { func (v *Validator) ValidateResponse(msg *dns.Msg, qname string, qtype uint16) error {
logger.Debug("Starting DNSSEC validation for %s %s", qname, dns.TypeToString[qtype])
if msg == nil || len(msg.Answer) == 0 { if msg == nil || len(msg.Answer) == 0 {
logger.Debug("No result for %s %s", qname, dns.TypeToString[qtype])
return ErrNoResult return ErrNoResult
} }
// Extract RRSet from response answer := extractRRSet(msg, qname, qtype)
rrset := NewRRSet() if answer.IsEmpty() {
for _, rr := range msg.Answer {
switch t := rr.(type) {
case *dns.RRSIG:
if t.TypeCovered == qtype {
rrset.RRSig = t
logger.Debug("Found RRSIG for %s %s (keytag: %d)", qname, dns.TypeToString[qtype], t.KeyTag)
}
default:
if rr.Header().Rrtype == qtype {
rrset.RRs = append(rrset.RRs, rr)
logger.Debug("Found RR for %s %s: %s", qname, dns.TypeToString[qtype], rr.String())
}
}
}
if rrset.IsEmpty() {
logger.Debug("Empty RRSet for %s %s", qname, dns.TypeToString[qtype])
return ErrNoResult return ErrNoResult
} }
if !answer.IsSigned() {
if !rrset.IsSigned() {
logger.Debug("RRSet for %s %s is not signed", qname, dns.TypeToString[qtype])
return ErrResourceNotSigned return ErrResourceNotSigned
} }
if err := answer.CheckHeaderIntegrity(dns.Fqdn(qname)); err != nil {
// Check header integrity
if err := rrset.CheckHeaderIntegrity(qname); err != nil {
logger.Debug("Header integrity check failed for %s %s: %v", qname, dns.TypeToString[qtype], err)
return err return err
} }
// Build and verify authentication chain logger.Debug("Validating %s %s (signer: %s)", qname, dns.TypeToString[qtype], answer.SignerName())
signerName := rrset.SignerName() return v.walker.validate(answer, qname, qtype)
logger.Debug("Building authentication chain for signer: %s", signerName)
authChain := NewAuthenticationChain()
if err := authChain.Populate(signerName, v.queryFunc); err != nil {
logger.Debug("Cannot populate authentication chain for %s: %v", signerName, err)
return err
} }
if err := authChain.Verify(rrset); err != nil { func (v *Validator) TakeStats() ValidationStats {
logger.Debug("DNSSEC validation failed for %s %s: %v", qname, dns.TypeToString[qtype], err) s := v.walker.stats()
return err v.walker.resetStats()
return s
} }
logger.Debug("DNSSEC validation successful for %s %s", qname, dns.TypeToString[qtype]) func extractRRSet(msg *dns.Msg, name string, qtype uint16) *RRSet {
return nil if msg == nil {
return NewRRSet()
}
return extractRRSetFrom(msg.Answer, name, qtype)
} }
func NewValidatorWithAuthoritativeQueries() *Validator { func extractRRSetFrom(rrs []dns.RR, name string, qtype uint16) *RRSet {
querier := NewAuthoritativeQuerier() set := NewRRSet()
return NewValidator(querier.QueryAuthoritative) fq := dns.Fqdn(name)
for _, rr := range rrs {
switch t := rr.(type) {
case *dns.RRSIG:
if t.TypeCovered == qtype && dns.Fqdn(t.Header().Name) == fq {
set.RRSig = t
}
default:
if rr.Header().Rrtype == qtype && dns.Fqdn(rr.Header().Name) == fq {
set.RRs = append(set.RRs, rr)
}
}
}
return set
} }
+253
View File
@@ -0,0 +1,253 @@
package dnssec
import (
"fmt"
"net"
"strings"
"github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/miekg/dns"
)
// iterativeWalker validates top-down.
type iterativeWalker struct {
newClient ExchangeFactory
st ValidationStats
ipCache map[string]string
}
func newIterativeWalker(f ExchangeFactory) *iterativeWalker {
return &iterativeWalker{
newClient: f,
ipCache: make(map[string]string),
}
}
func (w *iterativeWalker) stats() ValidationStats { return w.st }
func (w *iterativeWalker) resetStats() { w.st = ValidationStats{} }
func (w *iterativeWalker) exchange(server, name string, qtype uint16) (*dns.Msg, error) {
m := new(dns.Msg)
m.SetQuestion(dns.Fqdn(name), qtype)
m.SetEdns0(1232, true)
m.RecursionDesired = false
if b, err := m.Pack(); err == nil {
w.st.BytesSent += len(b)
}
w.st.Queries++
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)
}
}
return resp, err
}
func (w *iterativeWalker) queryAny(servers []string, name string, qtype uint16) (*dns.Msg, error) {
var lastErr error
for _, s := range servers {
resp, err := w.exchange(s, name, qtype)
if err == nil && resp != nil {
return resp, nil
}
lastErr = err
}
if lastErr == nil {
lastErr = fmt.Errorf("no response from servers")
}
return nil, lastErr
}
func (w *iterativeWalker) validate(answer *RRSet, qname string, qtype uint16) error {
current, servers, err := w.trustRoot()
if err != nil {
return err
}
currentServers := servers
// Follow delegations from the root toward qname.
for {
resp, err := w.queryAny(currentServers, qname, qtype)
if err != nil {
return err
}
// Authoritative answer for the record: verify it and finish.
if hasAnswer(resp, qname, qtype) {
ans := extractRRSet(resp, qname, qtype)
if ans.IsEmpty() || !ans.IsSigned() {
return ErrResourceNotSigned
}
if err := current.VerifyRRSIG(ans); err != nil {
logger.Debug("answer RRSIG verification failed for %s: %v", qname, err)
return ErrInvalidRRsig
}
w.st.Validated = true
return nil
}
// Otherwise expect a referral to a child zone.
childZone, childServers, dsSet, err := w.parseReferral(resp)
if err != nil {
return err
}
if dsSet.IsEmpty() || !dsSet.IsSigned() {
// No signed DS at the cut means the child is not securely
// delegated; the chain cannot continue.
return ErrDsNotAvailable
}
// The DS must be signed by the current (parent) zone.
if err := current.VerifyRRSIG(dsSet); err != nil {
logger.Debug("DS RRSIG verification failed for %s: %v", childZone, err)
return ErrInvalidRRsig
}
// Fetch the child DNSKEY and verify it matches the DS.
child, err := w.fetchZoneKeys(childServers, childZone)
if err != nil {
return err
}
if err := child.VerifyDS(dsSet.RRs); err != nil {
return err
}
current = child
currentServers = childServers
}
}
func (w *iterativeWalker) trustRoot() (*SignedZone, []string, error) {
var lastErr error
for _, server := range rootHints {
resp, err := w.exchange(server, ".", dns.TypeDNSKEY)
if err != nil || resp == nil {
lastErr = err
continue
}
keyset := extractRRSet(resp, ".", dns.TypeDNSKEY)
zone, verr := verifyRootAnchor(keyset)
if verr != nil {
lastErr = verr
continue
}
return zone, rootHints, nil
}
if lastErr == nil {
lastErr = ErrDnskeyNotAvailable
}
return nil, nil, fmt.Errorf("root trust anchor validation failed: %w", lastErr)
}
func (w *iterativeWalker) fetchZoneKeys(servers []string, zoneName string) (*SignedZone, error) {
resp, err := w.queryAny(servers, zoneName, dns.TypeDNSKEY)
if err != nil {
return nil, err
}
keyset := extractRRSet(resp, zoneName, dns.TypeDNSKEY)
if keyset.IsEmpty() || !keyset.IsSigned() {
return nil, ErrDnskeyNotAvailable
}
zone := NewSignedZone(zoneName)
zone.DNSKey = keyset
for _, rr := range keyset.RRs {
if k, ok := rr.(*dns.DNSKEY); ok {
zone.AddPubKey(k)
}
}
if err := zone.VerifyRRSIG(keyset); err != nil {
return nil, err
}
return zone, nil
}
func (w *iterativeWalker) parseReferral(msg *dns.Msg) (childZone string, childServers []string, ds *RRSet, err error) {
if msg == nil {
return "", nil, nil, fmt.Errorf("nil referral")
}
// The delegated zone is the owner name of the NS records in the authority
// section. Collect NS names per zone.
nsNamesByZone := map[string][]string{}
for _, rr := range msg.Ns {
if ns, ok := rr.(*dns.NS); ok {
z := dns.Fqdn(ns.Header().Name)
nsNamesByZone[z] = append(nsNamesByZone[z], ns.Ns)
}
}
if len(nsNamesByZone) == 0 {
return "", nil, nil, fmt.Errorf("no referral NS records present")
}
// There should be exactly one delegated zone in a referral.
var nsNames []string
for z, names := range nsNamesByZone {
childZone = z
nsNames = names
break
}
ds = extractRRSetFrom(msg.Ns, childZone, dns.TypeDS)
childServers = w.resolveNS(nsNames, msg.Extra)
if len(childServers) == 0 {
return "", nil, nil, fmt.Errorf("could not resolve nameservers for %s", childZone)
}
return childZone, childServers, ds, nil
}
func (w *iterativeWalker) resolveNS(nsNames []string, extra []dns.RR) []string {
glue := map[string]string{}
for _, rr := range extra {
if a, ok := rr.(*dns.A); ok {
glue[dns.Fqdn(a.Header().Name)] = a.A.String()
}
}
var servers []string
for _, ns := range nsNames {
fq := dns.Fqdn(ns)
if ip, ok := glue[fq]; ok {
servers = append(servers, net.JoinHostPort(ip, "53"))
continue
}
if ip, ok := w.ipCache[fq]; ok {
servers = append(servers, net.JoinHostPort(ip, "53"))
continue
}
ips, e := net.LookupIP(strings.TrimSuffix(fq, "."))
if e != nil {
continue
}
for _, ip := range ips {
if v4 := ip.To4(); v4 != nil {
addr := v4.String()
w.ipCache[fq] = addr
servers = append(servers, net.JoinHostPort(addr, "53"))
break
}
}
}
return servers
}
func hasAnswer(msg *dns.Msg, qname string, qtype uint16) bool {
if msg == nil {
return false
}
fq := dns.Fqdn(qname)
for _, rr := range msg.Answer {
if rr.Header().Rrtype == qtype && dns.Fqdn(rr.Header().Name) == fq {
return true
}
}
return false
}
+136
View File
@@ -0,0 +1,136 @@
package dnssec
import (
"github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/miekg/dns"
)
// trustWalker validates bottom-up.
// The parent of each zone is discovered from the DS record's RRSIG signer name.
type trustWalker struct {
send SendFunc
st ValidationStats
}
func newTrustWalker(send SendFunc) *trustWalker {
return &trustWalker{send: send}
}
func (w *trustWalker) stats() ValidationStats { return w.st }
func (w *trustWalker) resetStats() { w.st = ValidationStats{} }
func (w *trustWalker) ask(name string, qtype uint16) (*dns.Msg, error) {
m := new(dns.Msg)
m.SetQuestion(dns.Fqdn(name), qtype)
m.Id = dns.Id()
m.RecursionDesired = true
m.SetEdns0(1232, true)
if b, err := m.Pack(); err == nil {
w.st.BytesSent += len(b)
}
w.st.Queries++
resp, err := w.send(m)
logger.Debug("trust walker got %s %s: %d answers", name, dns.TypeToString[qtype], len(resp.Answer))
if err == nil && resp != nil {
if b, perr := resp.Pack(); perr == nil {
w.st.BytesReceived += len(b)
}
}
return resp, err
}
func (w *trustWalker) validate(answer *RRSet, qname string, qtype uint16) error {
signer := dns.Fqdn(answer.SignerName())
if signer == "" {
return ErrResourceNotSigned
}
// Fetch the signing zone's keys and verify the answer against them.
signingZone, err := w.fetchSelfSignedKeys(signer)
if err != nil {
return err
}
if err := signingZone.VerifyRRSIG(answer); err != nil {
logger.Debug("answer RRSIG verification failed for %s: %v", qname, err)
return ErrInvalidRRsig
}
// Walk up: signing zone -> ... -> root, verifying DS linkage each step.
current := signingZone
name := signer
for name != "." {
dsResp, err := w.ask(name, dns.TypeDS)
if err != nil {
return err
}
dsSet := extractRRSet(dsResp, name, dns.TypeDS)
if dsSet.IsEmpty() || !dsSet.IsSigned() {
return ErrDsNotAvailable
}
// The DS must match the current zone's key.
if err := current.VerifyDS(dsSet.RRs); err != nil {
return err
}
// The parent zone is whoever signed the DS RRset.
parentName := dns.Fqdn(dsSet.SignerName())
var parent *SignedZone
if parentName == "." {
parent, err = w.fetchRoot()
} else {
parent, err = w.fetchSelfSignedKeys(parentName)
}
if err != nil {
return err
}
// The parent must have signed the DS RRset.
if err := parent.VerifyRRSIG(dsSet); err != nil {
logger.Debug("DS RRSIG verification failed for %s: %v", name, err)
return ErrInvalidRRsig
}
current = parent
name = parentName
}
w.st.Validated = true
return nil
}
func (w *trustWalker) fetchSelfSignedKeys(zoneName string) (*SignedZone, error) {
resp, err := w.ask(zoneName, dns.TypeDNSKEY)
if err != nil {
return nil, err
}
keyset := extractRRSet(resp, zoneName, dns.TypeDNSKEY)
if keyset.IsEmpty() || !keyset.IsSigned() {
return nil, ErrDnskeyNotAvailable
}
zone := NewSignedZone(zoneName)
zone.DNSKey = keyset
for _, rr := range keyset.RRs {
if k, ok := rr.(*dns.DNSKEY); ok {
zone.AddPubKey(k)
}
}
if err := zone.VerifyRRSIG(keyset); err != nil {
return nil, err
}
return zone, nil
}
// fetchRoot fetches the root DNSKEY set and verifies it against the anchor.
func (w *trustWalker) fetchRoot() (*SignedZone, error) {
resp, err := w.ask(".", dns.TypeDNSKEY)
if err != nil {
return nil, err
}
keyset := extractRRSet(resp, ".", dns.TypeDNSKEY)
return verifyRootAnchor(keyset)
}
+3 -7
View File
@@ -61,25 +61,21 @@ func (c *Client) Close() {
// The dnscrypt library doesn't require explicit cleanup // 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 { if len(msg.Question) > 0 {
question := msg.Question[0] question := msg.Question[0]
logger.Debug("DNSCrypt query: %s %s", question.Name, dns.TypeToString[question.Qtype]) 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) response, err := c.resolver.Exchange(msg, c.ri)
if err != nil { if err != nil {
logger.Error("DNSCrypt query failed: %v", err) 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 { if len(response.Answer) > 0 {
logger.Debug("DNSCrypt response: %d answers", len(response.Answer)) logger.Debug("DNSCrypt response: %d answers", len(response.Answer))
} }
return response, nil return msg, response, nil
} }
+9 -25
View File
@@ -114,13 +114,6 @@ func New(config Config) (*Client, error) {
}, nil }, 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()) { func (c *Client) newSingleUseClient() (*http.Client, func()) {
tlsConfig := &tls.Config{ tlsConfig := &tls.Config{
ServerName: c.config.Host, 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 { if len(msg.Question) > 0 {
question := msg.Question[0] question := msg.Question[0]
logger.Debug("DoH query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.upstreamURL.Host) 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() packedMsg, err := msg.Pack()
if err != nil { if err != nil {
logger.Error("DoH failed to pack DNS message: %v", err) 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 httpClient := c.httpClient
if !c.config.KeepAlive { if !c.config.KeepAlive {
freshClient, closeFn := c.newSingleUseClient() 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)) httpReq, err := http.NewRequest(http.MethodPost, c.upstreamURL.String(), bytes.NewReader(packedMsg))
if err != nil { if err != nil {
logger.Error("DoH failed to create HTTP request: %v", err) 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") 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) httpResp, err := httpClient.Do(httpReq)
if err != nil { if err != nil {
logger.Error("DoH request failed to %s: %v", c.upstreamURL.Host, err) 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() defer httpResp.Body.Close()
if httpResp.StatusCode != http.StatusOK { if httpResp.StatusCode != http.StatusOK {
logger.Error("DoH received non-200 status from %s: %s", c.upstreamURL.Host, httpResp.Status) 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 { if ct := httpResp.Header.Get("Content-Type"); ct != dnsMessageContentType {
logger.Error("DoH unexpected Content-Type from %s: %s", c.upstreamURL.Host, ct) 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) responseBody, err := io.ReadAll(httpResp.Body)
if err != nil { if err != nil {
logger.Error("DoH failed reading response from %s: %v", c.upstreamURL.Host, err) 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) recvMsg := new(dns.Msg)
err = recvMsg.Unpack(responseBody) err = recvMsg.Unpack(responseBody)
if err != nil { if err != nil {
logger.Error("DoH failed to unpack response from %s: %v", c.upstreamURL.Host, err) 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 { if len(recvMsg.Answer) > 0 {
logger.Debug("DoH response from %s: %d answers", c.upstreamURL.Host, len(recvMsg.Answer)) logger.Debug("DoH response from %s: %d answers", c.upstreamURL.Host, len(recvMsg.Answer))
} }
return recvMsg, nil return msg, recvMsg, nil
} }
+14 -16
View File
@@ -95,7 +95,7 @@ func (c *Client) OpenConnection() error {
return nil 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 { if len(msg.Question) > 0 {
question := msg.Question[0] question := msg.Question[0]
logger.Debug("DoQ query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.targetAddr) 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 { if c.quicConn == nil {
err := c.OpenConnection() err := c.OpenConnection()
if err != nil { if err != nil {
return nil, err return msg, nil, err
} }
} }
// Prepare DNS message // Prepare DNS message
// DoQ requires Id to be 0
msg.Id = 0 msg.Id = 0
if c.config.DNSSEC {
msg.SetEdns0(4096, true)
}
packed, err := msg.Pack() packed, err := msg.Pack()
if err != nil { if err != nil {
logger.Error("DoQ failed to pack message: %v", err) 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 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) logger.Debug("DoQ stream failed, reconnecting: %v", err)
err = c.OpenConnection() err = c.OpenConnection()
if err != nil { if err != nil {
return nil, err return msg, nil, err
} }
quicStream, err = c.quicConn.OpenStream() quicStream, err = c.quicConn.OpenStream()
if err != nil { if err != nil {
logger.Error("DoQ failed to open stream after reconnect: %v", err) 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))) err = binary.Write(&lengthPrefixedMessage, binary.BigEndian, uint16(len(packed)))
if err != nil { if err != nil {
logger.Error("DoQ failed to write message length: %v", err) 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) _, err = lengthPrefixedMessage.Write(packed)
if err != nil { if err != nil {
logger.Error("DoQ failed to write DNS message: %v", err) 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()) _, err = quicStream.Write(lengthPrefixedMessage.Bytes())
if err != nil { if err != nil {
logger.Error("DoQ failed to write to stream: %v", err) 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() quicStream.Close()
@@ -157,32 +155,32 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
_, err = io.ReadFull(quicStream, lengthBuf) _, err = io.ReadFull(quicStream, lengthBuf)
if err != nil { if err != nil {
logger.Error("DoQ failed to read response length: %v", err) 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) messageLength := binary.BigEndian.Uint16(lengthBuf)
if messageLength == 0 { if messageLength == 0 {
logger.Error("DoQ received zero-length message") 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) responseBuf := make([]byte, messageLength)
_, err = io.ReadFull(quicStream, responseBuf) _, err = io.ReadFull(quicStream, responseBuf)
if err != nil { if err != nil {
logger.Error("DoQ failed to read response data: %v", err) 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) recvMsg := new(dns.Msg)
err = recvMsg.Unpack(responseBuf) err = recvMsg.Unpack(responseBuf)
if err != nil { if err != nil {
logger.Error("DoQ failed to parse response: %v", err) 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 { if len(recvMsg.Answer) > 0 {
logger.Debug("DoQ response from %s: %d answers", c.targetAddr, len(recvMsg.Answer)) logger.Debug("DoQ response from %s: %d answers", c.targetAddr, len(recvMsg.Answer))
} }
return recvMsg, nil return msg, recvMsg, nil
} }
+15 -19
View File
@@ -124,7 +124,7 @@ func (c *Client) ensureConnection() error {
return nil 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 { if len(msg.Question) > 0 {
question := msg.Question[0] question := msg.Question[0]
logger.Debug("DoT query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort) 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) // Ensure we have a connection (either persistent or new)
if c.config.KeepAlive { if c.config.KeepAlive {
if err := c.ensureConnection(); err != nil { 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 { } else {
// For non-keepalive mode, create a fresh connection for each query // 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() c.connMutex.Unlock()
if err := c.ensureConnection(); err != nil { 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() packed, err := msg.Pack()
if err != nil { if err != nil {
logger.Error("DoT failed to pack message: %v", err) 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) // Prepend message length (DNS over TCP format)
@@ -171,7 +167,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
// Write query // Write query
if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil { if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
logger.Error("DoT failed to set write deadline: %v", err) 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 { 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 { if c.config.KeepAlive {
logger.Debug("DoT write failed with keep-alive, attempting reconnect") logger.Debug("DoT write failed with keep-alive, attempting reconnect")
if reconnectErr := c.ensureConnection(); reconnectErr != nil { 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() c.connMutex.Lock()
@@ -189,48 +185,48 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
c.connMutex.Unlock() c.connMutex.Unlock()
if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil { 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 { 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 { } 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 // Read response
if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil { if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil {
logger.Error("DoT failed to set read deadline: %v", err) 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 // Read message length
lengthBuf := make([]byte, 2) lengthBuf := make([]byte, 2)
if _, err := io.ReadFull(conn, lengthBuf); err != nil { if _, err := io.ReadFull(conn, lengthBuf); err != nil {
logger.Error("DoT failed to read response length from %s: %v", c.hostAndPort, err) 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) msgLen := binary.BigEndian.Uint16(lengthBuf)
if msgLen > dns.MaxMsgSize { if msgLen > dns.MaxMsgSize {
logger.Error("DoT response too large from %s: %d bytes", c.hostAndPort, msgLen) 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 // Read message body
buffer := make([]byte, msgLen) buffer := make([]byte, msgLen)
if _, err := io.ReadFull(conn, buffer); err != nil { if _, err := io.ReadFull(conn, buffer); err != nil {
logger.Error("DoT failed to read response from %s: %v", c.hostAndPort, err) 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 // Parse response
response := new(dns.Msg) response := new(dns.Msg)
if err := response.Unpack(buffer); err != nil { if err := response.Unpack(buffer); err != nil {
logger.Error("DoT failed to unpack response from %s: %v", c.hostAndPort, err) 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 { if len(response.Answer) > 0 {
@@ -247,5 +243,5 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
c.connMutex.Unlock() c.connMutex.Unlock()
} }
return response, nil return msg, response, nil
} }
+15 -19
View File
@@ -104,7 +104,7 @@ func (c *Client) ensureConnection() error {
return nil 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 { if len(msg.Question) > 0 {
question := msg.Question[0] question := msg.Question[0]
logger.Debug("DoTCP query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort) 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 c.config.KeepAlive {
if err := c.ensureConnection(); err != nil { 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 { } else {
c.connMutex.Lock() c.connMutex.Lock()
@@ -123,18 +123,14 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
c.connMutex.Unlock() c.connMutex.Unlock()
if err := c.ensureConnection(); err != nil { 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() packed, err := msg.Pack()
if err != nil { if err != nil {
logger.Error("DoTCP failed to pack message: %v", err) 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 // 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 { if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
logger.Error("DoTCP failed to set write deadline: %v", err) 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 { 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 { if c.config.KeepAlive {
logger.Debug("DoTCP write failed with keep-alive, attempting reconnect") logger.Debug("DoTCP write failed with keep-alive, attempting reconnect")
if reconnectErr := c.ensureConnection(); reconnectErr != nil { 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() c.connMutex.Lock()
@@ -165,44 +161,44 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
c.connMutex.Unlock() c.connMutex.Unlock()
if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil { 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 { 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 { } 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 { if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil {
logger.Error("DoTCP failed to set read deadline: %v", err) 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) lengthBuf := make([]byte, 2)
if _, err := io.ReadFull(conn, lengthBuf); err != nil { if _, err := io.ReadFull(conn, lengthBuf); err != nil {
logger.Error("DoTCP failed to read response length from %s: %v", c.hostAndPort, err) 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) msgLen := binary.BigEndian.Uint16(lengthBuf)
if msgLen > dns.MaxMsgSize { if msgLen > dns.MaxMsgSize {
logger.Error("DoTCP response too large from %s: %d bytes", c.hostAndPort, msgLen) 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) buffer := make([]byte, msgLen)
if _, err := io.ReadFull(conn, buffer); err != nil { if _, err := io.ReadFull(conn, buffer); err != nil {
logger.Error("DoTCP failed to read response from %s: %v", c.hostAndPort, err) 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) response := new(dns.Msg)
if err := response.Unpack(buffer); err != nil { if err := response.Unpack(buffer); err != nil {
logger.Error("DoTCP failed to unpack response from %s: %v", c.hostAndPort, err) 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 { if len(response.Answer) > 0 {
@@ -218,5 +214,5 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
c.connMutex.Unlock() c.connMutex.Unlock()
} }
return response, nil return msg, response, nil
} }
+49 -14
View File
@@ -64,7 +64,7 @@ func (c *Client) createConnection() (*net.UDPConn, error) {
return conn, nil 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 { if len(msg.Question) > 0 {
question := msg.Question[0] question := msg.Question[0]
logger.Debug("DoUDP query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort) logger.Debug("DoUDP query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort)
@@ -72,51 +72,86 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) {
conn, err := c.createConnection() conn, err := c.createConnection()
if err != nil { 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() defer conn.Close()
if c.config.DNSSEC {
msg.SetEdns0(4096, true)
}
packedMsg, err := msg.Pack() packedMsg, err := msg.Pack()
if err != nil { if err != nil {
logger.Error("DoUDP failed to pack message: %v", err) 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 { if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil {
logger.Error("DoUDP failed to set write deadline: %v", err) 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 { if _, err := conn.Write(packedMsg); err != nil {
logger.Error("DoUDP failed to send query to %s: %v", c.hostAndPort, err) 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 { if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil {
logger.Error("DoUDP failed to set read deadline: %v", err) 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)
} }
buffer := make([]byte, dns.MaxMsgSize) bufSize := 512
if opt := msg.IsEdns0(); opt != nil {
bufSize = max(int(opt.UDPSize()), 512)
}
buffer := make([]byte, bufSize)
n, err := conn.Read(buffer) n, err := conn.Read(buffer)
if err != nil { if err != nil {
logger.Error("DoUDP failed to read response from %s: %v", c.hostAndPort, err) 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) response := new(dns.Msg)
if err := response.Unpack(buffer[:n]); err != nil { if err := response.Unpack(buffer[:n]); err != nil {
logger.Error("DoUDP failed to unpack response from %s: %v", c.hostAndPort, err) 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)
}
// 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 { if len(response.Answer) > 0 {
logger.Debug("DoUDP response from %s: %d answers", c.hostAndPort, len(response.Answer)) logger.Debug("DoUDP response from %s: %d answers", c.hostAndPort, len(response.Answer))
} }
return response, nil 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
} }
+1 -1
View File
@@ -9,6 +9,7 @@ require (
github.com/miekg/dns v1.1.72 github.com/miekg/dns v1.1.72
github.com/quic-go/quic-go v0.60.0 github.com/quic-go/quic-go v0.60.0
golang.org/x/net v0.56.0 golang.org/x/net v0.56.0
golang.org/x/sys v0.46.0
) )
require ( require (
@@ -19,7 +20,6 @@ require (
golang.org/x/exp v0.0.0-20260611194520-c48552f49976 // indirect golang.org/x/exp v0.0.0-20260611194520-c48552f49976 // indirect
golang.org/x/mod v0.37.0 // indirect golang.org/x/mod v0.37.0 // indirect
golang.org/x/sync v0.21.0 // indirect golang.org/x/sync v0.21.0 // indirect
golang.org/x/sys v0.46.0 // indirect
golang.org/x/text v0.39.0 // indirect golang.org/x/text v0.39.0 // indirect
golang.org/x/tools v0.47.0 // indirect golang.org/x/tools v0.47.0 // indirect
) )
+89
View File
@@ -0,0 +1,89 @@
office.com
live.com
ecs.office.com
officeapps.live.com
config.edge.skype.com
login.live.com
teams.microsoft.com
cloudflare.com
pki.goog
office365.com
c.pki.goog
outlook.office365.com
config.teams.microsoft.com
nexusrules.officeapps.live.com
outlook.office.com
substrate.office.com
nist.gov
config.office.net
clients.config.office.net
officeclient.microsoft.com
storage.live.com
time-a.nist.gov
time-b.nist.gov
onedrive.live.com
g.live.com
exo.nel.measure.office.net
odc.officeapps.live.com
ocsp.entrust.net
scorecardresearch.com
outlook.com
mrodevicemgr.officeapps.live.com
cloud.microsoft
roaming.officeapps.live.com
sb.scorecardresearch.com
r4.res.office365.com
a.nel.cloudflare.com
apple-relay.cloudflare.com
cdnjs.cloudflare.com
teams.cloud.microsoft
augloop.office.com
cdn.jsdelivr.net
statics.teams.cdn.office.net
bam.nr-data.net
activity.windows.com
taboola.com
dns.google
res.cdn.office.net
res.public.onecdn.static.microsoft
tenable.com
newrelic.com
augloop.svc.cloud.microsoft
id5-sync.com
ads.linkedin.com
m365.cloud.microsoft
px.ads.linkedin.com
outlook.cloud.microsoft
chatgpt.com
autodiscover-s.outlook.com
copilot.cloud.microsoft
substrate.svc.cloud.microsoft
connectivity-test.usercontent.microsoft
connectivity-test.static.microsoft
connectivity-test.cloud.microsoft
verisign.com
ipv4probe.office.com
tr-ssc-mira.office.com
onetrust.com
acdc-direct.office.com
delve.office.com
loki.delve.office.com
config.edge.skype.com.trafficmanager.net
enterprise.activity.windows.com
nleditor.osi.office.net
atm-fp-direct.office.com
res-1.cdn.office.net
crl.verisign.com
geolocation.onetrust.com
pendo.io
www.linkedin.com
webshell.suite.office.com
qualys.com
cdn.id5-sync.com
trc.taboola.com
nexus.officeapps.live.com
ocws.officeapps.live.com
ow1.res.office365.com
www.temu.com
waconatm.officeapps.live.com
waconafd.officeapps.live.com
+25 -13
View File
@@ -10,6 +10,7 @@ import (
"time" "time"
"github.com/afonsofrancof/sdns-proxy/client" "github.com/afonsofrancof/sdns-proxy/client"
"github.com/afonsofrancof/sdns-proxy/common/dnssec"
"github.com/afonsofrancof/sdns-proxy/internal/qol/capture" "github.com/afonsofrancof/sdns-proxy/internal/qol/capture"
"github.com/afonsofrancof/sdns-proxy/internal/qol/results" "github.com/afonsofrancof/sdns-proxy/internal/qol/results"
"github.com/afonsofrancof/sdns-proxy/internal/qol/stats" "github.com/afonsofrancof/sdns-proxy/internal/qol/stats"
@@ -201,35 +202,46 @@ func (r *MeasurementRunner) performQuery(dnsClient client.DNSClient, domain, ups
msg.Id = dns.Id() msg.Id = dns.Id()
msg.RecursionDesired = true msg.RecursionDesired = true
msg.SetQuestion(dns.Fqdn(domain), qType) msg.SetQuestion(dns.Fqdn(domain), qType)
msg.SetEdns0(1232, r.config.DNSSEC)
packed, err := msg.Pack()
if err != nil {
metric.ResponseCode = "ERROR"
metric.Error = fmt.Sprintf("pack request: %v", err)
return metric
}
metric.RequestSize = len(packed)
start := time.Now() start := time.Now()
metric.Timestamp = start metric.Timestamp = start
resp, err := dnsClient.Query(msg) sent, resp, err := dnsClient.Query(msg)
metric.Duration = time.Since(start).Nanoseconds() metric.Duration = time.Since(start).Nanoseconds()
metric.DurationMs = float64(metric.Duration) / 1e6 metric.DurationMs = float64(metric.Duration) / 1e6
if sent != nil {
if b, perr := sent.Pack(); perr == nil {
metric.RequestSize = len(b)
}
} else if b, perr := msg.Pack(); perr == nil {
metric.RequestSize = len(b)
}
if reporter, ok := dnsClient.(interface {
LastValidation() dnssec.ValidationStats
}); ok {
st := reporter.LastValidation()
metric.DNSSECValidated = st.Validated
metric.DNSSECQueries = st.Queries
metric.RequestSize += st.BytesSent
metric.ResponseSize += st.BytesReceived
}
if err != nil { if err != nil {
metric.ResponseCode = "ERROR" metric.ResponseCode = "ERROR"
metric.Error = err.Error() metric.Error = err.Error()
return metric return metric
} }
respBytes, err := resp.Pack() if b, perr := resp.Pack(); perr == nil {
if err != nil { metric.ResponseSize += len(b)
} else {
metric.ResponseCode = "ERROR" metric.ResponseCode = "ERROR"
metric.Error = fmt.Sprintf("pack response: %v", err) metric.Error = fmt.Sprintf("pack response: %v", perr)
return metric return metric
} }
metric.ResponseSize = len(respBytes)
metric.ResponseCode = dns.RcodeToString[resp.Rcode] metric.ResponseCode = dns.RcodeToString[resp.Rcode]
return metric return metric
} }
+6 -2
View File
@@ -13,6 +13,8 @@ type DNSMetric struct {
QueryType string `json:"query_type"` QueryType string `json:"query_type"`
Protocol string `json:"protocol"` Protocol string `json:"protocol"`
DNSSEC bool `json:"dnssec"` DNSSEC bool `json:"dnssec"`
DNSSECValidated bool `json:"dnssec_validated"`
DNSSECQueries int `json:"dnssec_queries"`
AuthoritativeDNSSEC bool `json:"auth_dnssec"` AuthoritativeDNSSEC bool `json:"auth_dnssec"`
KeepAlive bool `json:"keep_alive"` KeepAlive bool `json:"keep_alive"`
DNSServer string `json:"dns_server"` DNSServer string `json:"dns_server"`
@@ -48,8 +50,8 @@ func NewMetricsWriter(path string) (*MetricsWriter, error) {
// Only write header if file is new // Only write header if file is new
if !fileExists { if !fileExists {
header := []string{ header := []string{
"domain", "query_type", "protocol", "dnssec", "auth_dnssec", "keep_alive", "domain", "query_type", "protocol", "dnssec", "dnssec_validated", "dnssec_queries",
"dns_server", "timestamp", "duration_ns", "duration_ms", "auth_dnssec", "keep_alive", "dns_server", "timestamp", "duration_ns", "duration_ms",
"request_size_bytes", "response_size_bytes", "response_code", "error", "request_size_bytes", "response_size_bytes", "response_code", "error",
} }
@@ -72,6 +74,8 @@ func (mw *MetricsWriter) WriteMetric(metric DNSMetric) error {
metric.QueryType, metric.QueryType,
metric.Protocol, metric.Protocol,
strconv.FormatBool(metric.DNSSEC), strconv.FormatBool(metric.DNSSEC),
strconv.FormatBool(metric.DNSSECValidated),
strconv.Itoa(metric.DNSSECQueries),
strconv.FormatBool(metric.AuthoritativeDNSSEC), strconv.FormatBool(metric.AuthoritativeDNSSEC),
strconv.FormatBool(metric.KeepAlive), strconv.FormatBool(metric.KeepAlive),
metric.DNSServer, metric.DNSServer,
+13 -1
View File
@@ -3,6 +3,7 @@ package stats
import ( import (
"encoding/csv" "encoding/csv"
"fmt" "fmt"
"golang.org/x/sys/unix"
"os" "os"
"runtime" "runtime"
"time" "time"
@@ -15,6 +16,7 @@ type RuntimeStats struct {
AllocDelta uint64 AllocDelta uint64
MallocsDelta uint64 MallocsDelta uint64
GCDelta uint32 GCDelta uint32
PeakRSSKB int64
} }
type RuntimeCollector struct { type RuntimeCollector struct {
@@ -43,6 +45,7 @@ func (rc *RuntimeCollector) Collect() RuntimeStats {
AllocDelta: current.TotalAlloc - rc.startStats.TotalAlloc, AllocDelta: current.TotalAlloc - rc.startStats.TotalAlloc,
MallocsDelta: current.Mallocs - rc.startStats.Mallocs, MallocsDelta: current.Mallocs - rc.startStats.Mallocs,
GCDelta: current.NumGC - rc.startStats.NumGC, GCDelta: current.NumGC - rc.startStats.NumGC,
PeakRSSKB: peakRSSKB(),
} }
} }
@@ -69,7 +72,7 @@ func (rc *RuntimeCollector) WriteStats() error {
if !fileExists { if !fileExists {
header := []string{ header := []string{
"timestamp", "total_alloc_bytes", "mallocs", "gc_cycles", "timestamp", "total_alloc_bytes", "mallocs", "gc_cycles",
"alloc_delta", "mallocs_delta", "gc_delta", "alloc_delta", "mallocs_delta", "gc_delta", "peak_rss_kb",
} }
if err := writer.Write(header); err != nil { if err := writer.Write(header); err != nil {
return fmt.Errorf("failed to write mem.csv header: %w", err) return fmt.Errorf("failed to write mem.csv header: %w", err)
@@ -85,6 +88,7 @@ func (rc *RuntimeCollector) WriteStats() error {
fmt.Sprintf("%d", stats.AllocDelta), fmt.Sprintf("%d", stats.AllocDelta),
fmt.Sprintf("%d", stats.MallocsDelta), fmt.Sprintf("%d", stats.MallocsDelta),
fmt.Sprintf("%d", stats.GCDelta), fmt.Sprintf("%d", stats.GCDelta),
fmt.Sprintf("%d", stats.PeakRSSKB),
} }
if err := writer.Write(row); err != nil { if err := writer.Write(row); err != nil {
return fmt.Errorf("failed to write mem.csv row: %w", err) return fmt.Errorf("failed to write mem.csv row: %w", err)
@@ -93,3 +97,11 @@ func (rc *RuntimeCollector) WriteStats() error {
writer.Flush() writer.Flush()
return writer.Error() return writer.Error()
} }
func peakRSSKB() int64 {
var ru unix.Rusage
if err := unix.Getrusage(unix.RUSAGE_SELF, &ru); err != nil {
return 0
}
return int64(ru.Maxrss)
}
+3
View File
@@ -61,6 +61,9 @@ func cleanServerName(server string) string {
"dns.adguard.com": "adguard", "dns.adguard.com": "adguard",
"dns.adguard-dns.com": "adguard", "dns.adguard-dns.com": "adguard",
"AQMAAAAAAAAAETk0LjE0MC4xNC4xNDo1NDQzINErR_JS3PLCu_iZEIbq95zkSV2LFsigxDIuUso_OQhzIjIuZG5zY3J5cHQuZGVmYXVsdC5uczEuYWRndWFyZC5jb20": "adguard", "AQMAAAAAAAAAETk0LjE0MC4xNC4xNDo1NDQzINErR_JS3PLCu_iZEIbq95zkSV2LFsigxDIuUso_OQhzIjIuZG5zY3J5cHQuZGVmYXVsdC5uczEuYWRndWFyZC5jb20": "adguard",
"94.140.14.140": "adguard",
"9.9.9.10": "quad9",
"unfiltered.adguard-dns.com": "adguard",
} }
serverName := "" serverName := ""
@@ -1,258 +0,0 @@
#!/usr/bin/env python3
"""
Fast PCAP Preprocessor for DNS QoS Analysis.
Loads PCAP into memory, then uses binary search to match packets to query
windows. Direction is determined by a configured local IP (the netns veth IP).
"""
import argparse
import bisect
import csv
import shutil
import socket
import time
from pathlib import Path
from typing import Dict, List, NamedTuple
import dpkt
from dateutil import parser as date_parser
BANDWIDTH_COLUMNS = [
"bytes_sent",
"bytes_received",
"packets_sent",
"packets_received",
"total_bytes",
]
DEFAULT_LOCAL_IP = "192.168.100.2" # netns veth1 address
DEFAULT_PROVIDERS = ["adguard", "cloudflare", "google", "quad9"]
class Packet(NamedTuple):
timestamp: float
size: int
is_outbound: bool
def needs_processing(csv_path: Path) -> bool:
"""True if file lacks bandwidth columns OR all values are empty/zero."""
try:
with open(csv_path, "r", encoding="utf-8") as f:
reader = csv.DictReader(f)
if not reader.fieldnames:
return True
if not all(c in reader.fieldnames for c in BANDWIDTH_COLUMNS):
return True
for row in reader:
for col in BANDWIDTH_COLUMNS:
val = (row.get(col) or "").strip()
if val and val != "0":
return False # has real data
return True # columns exist but all empty/zero
except Exception:
return True
def parse_csv_timestamp(ts_str: str) -> float:
return date_parser.isoparse(ts_str).timestamp()
def load_pcap(pcap_path: Path, local_ip_bytes: bytes) -> List[Packet]:
"""Load PCAP into a list of Packets sorted by timestamp."""
print(" Loading PCAP...")
t0 = time.time()
packets: List[Packet] = []
with open(pcap_path, "rb") as f:
try:
reader = dpkt.pcap.Reader(f)
except ValueError:
f.seek(0)
reader = dpkt.pcapng.Reader(f)
for ts, buf in reader:
try:
eth = dpkt.ethernet.Ethernet(buf)
ip = eth.data
if not isinstance(ip, dpkt.ip.IP):
continue
if ip.src == local_ip_bytes:
is_outbound = True
elif ip.dst == local_ip_bytes:
is_outbound = False
else:
continue # not our traffic
packets.append(Packet(float(ts), len(buf), is_outbound))
except (dpkt.dpkt.NeedData, dpkt.dpkt.UnpackError, AttributeError):
continue
packets.sort(key=lambda p: p.timestamp)
print(f" Loaded {len(packets):,} packets in {time.time() - t0:.2f}s")
return packets
def load_csv_queries(csv_path: Path) -> List[Dict]:
queries = []
with open(csv_path, "r", encoding="utf-8") as f:
for row in csv.DictReader(f):
try:
start = parse_csv_timestamp(row["timestamp"])
duration = float(row["duration_ns"]) / 1e9
queries.append(
{"data": row, "start_time": start, "end_time": start + duration}
)
except Exception as e:
print(f" Warning: skipping row - {e}")
return queries
def match_packets(packets: List[Packet], queries: List[Dict]) -> int:
"""Assign bandwidth metrics to each query. Returns total matched packets."""
if not packets or not queries:
for q in queries:
q.update({c: 0 for c in BANDWIDTH_COLUMNS})
return 0
print(" Matching packets to queries...")
t0 = time.time()
timestamps = [p.timestamp for p in packets]
matched = 0
for q in queries:
lo = bisect.bisect_left(timestamps, q["start_time"])
hi = bisect.bisect_right(timestamps, q["end_time"])
bs = br = ps = pr = 0
for pkt in packets[lo:hi]:
if pkt.is_outbound:
bs += pkt.size
ps += 1
else:
br += pkt.size
pr += 1
q["bytes_sent"] = bs
q["bytes_received"] = br
q["packets_sent"] = ps
q["packets_received"] = pr
q["total_bytes"] = bs + br
matched += ps + pr
print(f" Matched {matched:,} packets in {time.time() - t0:.2f}s")
total_sent = sum(q["bytes_sent"] for q in queries)
total_recv = sum(q["bytes_received"] for q in queries)
with_data = sum(1 for q in queries if q["total_bytes"] > 0)
print(f" Total: {total_sent:,} B sent, {total_recv:,} B received")
print(f" Queries with data: {with_data}/{len(queries)}")
return matched
def write_csv(csv_path: Path, queries: List[Dict], backup: bool = True):
if backup and csv_path.exists():
bak = csv_path.with_suffix(".csv.bak")
if not bak.exists():
shutil.copy2(csv_path, bak)
print(f" Backup: {bak.name}")
original_fields = [
f for f in queries[0]["data"].keys() if f not in BANDWIDTH_COLUMNS
]
fieldnames = original_fields + BANDWIDTH_COLUMNS
with open(csv_path, "w", encoding="utf-8", newline="") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
for q in queries:
row = {k: q["data"][k] for k in original_fields}
for c in BANDWIDTH_COLUMNS:
row[c] = q[c]
writer.writerow(row)
print(f" Written: {csv_path.name}")
def process_provider(provider_path: Path, local_ip_bytes: bytes):
print(f"\n{'=' * 60}\nProcessing: {provider_path.name.upper()}\n{'=' * 60}")
processed = skipped = 0
total_time = 0.0
for csv_path in sorted(provider_path.glob("*.csv")):
name = csv_path.name.lower()
if ".bak" in name or name.endswith((".cpu.csv", ".mem.csv")):
continue
pcap_path = csv_path.with_suffix(".pcap")
if not pcap_path.exists():
print(f"\n{csv_path.name}: no matching PCAP")
continue
if not needs_processing(csv_path):
print(f"\n{csv_path.name}: already processed")
skipped += 1
continue
print(f"\n 📁 {csv_path.name}")
t0 = time.time()
packets = load_pcap(pcap_path, local_ip_bytes)
if not packets:
print(" ⚠ No usable packets in PCAP")
continue
queries = load_csv_queries(csv_path)
if not queries:
print(" ⚠ No valid queries in CSV")
continue
print(f" Loaded {len(queries):,} queries")
match_packets(packets, queries)
write_csv(csv_path, queries)
dt = time.time() - t0
total_time += dt
processed += 1
print(f" ✓ Completed in {dt:.2f}s")
print(
f"\n {provider_path.name}: {processed} processed, "
f"{skipped} skipped, {total_time:.2f}s"
)
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument(
"--local-ip",
default=DEFAULT_LOCAL_IP,
help=f"Local (netns veth) IP used to determine direction "
f"(default: {DEFAULT_LOCAL_IP})",
)
ap.add_argument("--results-dir", default="results", type=Path)
ap.add_argument("--providers", nargs="+", default=DEFAULT_PROVIDERS)
args = ap.parse_args()
local_ip_bytes = socket.inet_aton(args.local_ip)
print(f"\n{'=' * 60}\nDNS PCAP PREPROCESSOR\n{'=' * 60}")
print(f"Local IP: {args.local_ip}")
print(f"Results: {args.results_dir}")
if not args.results_dir.exists():
print(f"\n❌ Directory not found: {args.results_dir}")
return
t0 = time.time()
for provider in args.providers:
path = args.results_dir / provider
if path.exists():
process_provider(path, local_ip_bytes)
else:
print(f"\n⚠ Missing provider directory: {provider}")
total = time.time() - t0
print(f"\n{'=' * 60}")
print(f"✓ DONE in {total:.2f}s ({total / 60:.1f} min)")
print(f"{'=' * 60}\n")
if __name__ == "__main__":
main()
+4
View File
@@ -90,6 +90,8 @@ def merge_all_csvs(input_dir: Path, output_path: Path):
'provider', 'provider',
'protocol', 'protocol',
'dnssec_mode', 'dnssec_mode',
'dnssec_validated',
'dnssec_queries',
'domain', 'domain',
'query_type', 'query_type',
'keep_alive', 'keep_alive',
@@ -135,6 +137,8 @@ def merge_all_csvs(input_dir: Path, output_path: Path):
'provider': provider, 'provider': provider,
'protocol': config['protocol'], 'protocol': config['protocol'],
'dnssec_mode': config['dnssec_mode'], 'dnssec_mode': config['dnssec_mode'],
'dnssec_validated': row.get('dnssec_validated', ''),
'dnssec_queries': row.get('dnssec_queries', ''),
'keep_alive': config['keep_alive'], 'keep_alive': config['keep_alive'],
'domain': row.get('domain', ''), 'domain': row.get('domain', ''),
'query_type': row.get('query_type', ''), 'query_type': row.get('query_type', ''),
+2 -1
View File
@@ -52,7 +52,7 @@ def merge_mem_files(input_dir: Path, output_path: Path):
output_columns = [ output_columns = [
'id','provider', 'protocol', 'dnssec_mode', 'keep_alive', 'id','provider', 'protocol', 'dnssec_mode', 'keep_alive',
'timestamp', 'total_alloc_bytes', 'mallocs', 'gc_cycles', 'timestamp', 'total_alloc_bytes', 'mallocs', 'gc_cycles',
'alloc_delta', 'mallocs_delta', 'gc_delta' 'alloc_delta', 'mallocs_delta', 'gc_delta', 'peak_rss_kb'
] ]
total_rows = 0 total_rows = 0
@@ -85,6 +85,7 @@ def merge_mem_files(input_dir: Path, output_path: Path):
'alloc_delta': row.get('alloc_delta', ''), 'alloc_delta': row.get('alloc_delta', ''),
'mallocs_delta': row.get('mallocs_delta', ''), 'mallocs_delta': row.get('mallocs_delta', ''),
'gc_delta': row.get('gc_delta', ''), 'gc_delta': row.get('gc_delta', ''),
'peak_rss_kb': row.get('peak_rss_kb', ''),
} }
writer.writerow(out_row) writer.writerow(out_row)
+109 -42
View File
@@ -12,6 +12,9 @@ TIMEOUT="5s"
SLEEP_TIME="1" SLEEP_TIME="1"
DRY_RUN="false" DRY_RUN="false"
DNSSEC_DOMAINS_FILE="./domains_signed.txt"
DNSSEC_OUTPUT_DIR="./results-dnssec"
# Parse arguments # Parse arguments
while [[ $# -gt 0 ]]; do while [[ $# -gt 0 ]]; do
case $1 in case $1 in
@@ -43,17 +46,27 @@ while [[ $# -gt 0 ]]; do
DRY_RUN="true" DRY_RUN="true"
shift shift
;; ;;
--dnssec-domains)
DNSSEC_DOMAINS_FILE="$2"
shift 2
;;
--dnssec-output)
DNSSEC_OUTPUT_DIR="$2"
shift 2
;;
--help) --help)
echo "Usage: $0 [OPTIONS]" echo "Usage: $0 [OPTIONS]"
echo "" echo ""
echo "Options:" echo "Options:"
echo " -t, --tool-path PATH Path to qol tool (default: ./qol)" echo " -t, --tool-path PATH Path to qol tool (default: ./qol)"
echo " -d, --domains-file PATH Path to domains file (default: ./domains.txt)" echo " -d, --domains-file PATH Domains file for the general runs (default: ./domains.txt)"
echo " -o, --output-dir PATH Output directory (default: ./results)" echo " -o, --output-dir PATH Output dir for the general runs (default: ./results)"
echo " -I, --interface NAME Network interface (default: veth1)" echo " -I, --interface NAME Network interface (default: veth1)"
echo " -T, --timeout DURATION Timeout duration (default: 5s)" echo " -T, --timeout DURATION Timeout duration (default: 5s)"
echo " -s, --sleep SECONDS Sleep between runs (default: 1)" echo " -s, --sleep SECONDS Sleep between runs (default: 1)"
echo " -n, --dry-run Print the scenarios that would run, then exit" echo " -n, --dry-run Print the scenarios that would run, then exit"
echo " --dnssec-domains PATH Signed-domain list for DNSSEC runs (default: ./domains_signed.txt)"
echo " --dnssec-output PATH Output dir for DNSSEC runs (default: ./results-dnssec)"
echo " --help Show this help" echo " --help Show this help"
exit 0 exit 0
;; ;;
@@ -69,6 +82,8 @@ echo "Configuration:"
echo " Tool path: $TOOL_PATH" echo " Tool path: $TOOL_PATH"
echo " Domains file: $DOMAINS_FILE" echo " Domains file: $DOMAINS_FILE"
echo " Output dir: $OUTPUT_DIR" echo " Output dir: $OUTPUT_DIR"
echo " DNSSEC domains file: $DNSSEC_DOMAINS_FILE"
echo " DNSSEC output dir: $DNSSEC_OUTPUT_DIR"
echo " Interface: $INTERFACE" echo " Interface: $INTERFACE"
echo " Timeout: $TIMEOUT" echo " Timeout: $TIMEOUT"
echo " Sleep time: ${SLEEP_TIME}s" echo " Sleep time: ${SLEEP_TIME}s"
@@ -80,13 +95,17 @@ SC_URL=()
SC_DNSSEC=() SC_DNSSEC=()
SC_AUTH=() SC_AUTH=()
SC_KEEP=() SC_KEEP=()
SC_DOMAINS=() # per-scenario domains file
SC_OUTDIR=() # per-scenario output dir
add() { # add <name> <url> <dnssec> <auth> <keepalive> add() { # add <name> <url> <dnssec> <auth> <keepalive> (general workload)
SC_NAME+=("$1") SC_NAME+=("$1"); SC_URL+=("$2"); SC_DNSSEC+=("$3"); SC_AUTH+=("$4"); SC_KEEP+=("$5")
SC_URL+=("$2") SC_DOMAINS+=("$DOMAINS_FILE"); SC_OUTDIR+=("$OUTPUT_DIR")
SC_DNSSEC+=("$3") }
SC_AUTH+=("$4")
SC_KEEP+=("$5") add_dnssec() { # same signature, but pins the signed list + DNSSEC output dir
SC_NAME+=("$1"); SC_URL+=("$2"); SC_DNSSEC+=("$3"); SC_AUTH+=("$4"); SC_KEEP+=("$5")
SC_DOMAINS+=("$DNSSEC_DOMAINS_FILE"); SC_OUTDIR+=("$DNSSEC_OUTPUT_DIR")
} }
# Protocol comparison - DNSSEC off, non-persistent. # Protocol comparison - DNSSEC off, non-persistent.
@@ -120,8 +139,7 @@ add adguard-dnscrypt "sdns://AQMAAAAAAAAAETk0LjE0MC4xNC4xNDo1NDQzINErR_JS3PLCu_
add quad9-dnscrypt "sdns://AQMAAAAAAAAAFDE0OS4xMTIuMTEyLjExMjo4NDQzIGfIR7jIdYzRICRVQ751Z0bfNN8dhMALjEcDaN-CHYY-GTIuZG5zY3J5cHQtY2VydC5xdWFkOS5uZXQ" false false false add quad9-dnscrypt "sdns://AQMAAAAAAAAAFDE0OS4xMTIuMTEyLjExMjo4NDQzIGfIR7jIdYzRICRVQ751Z0bfNN8dhMALjEcDaN-CHYY-GTIuZG5zY3J5cHQtY2VydC5xdWFkOS5uZXQ" false false false
# Persistence contrast - DNSSEC off, persistent. # Persistence contrast - DNSSEC off, persistent.
# Teste for TCP based protocols. # Tested for TCP-based protocols. QUIC makes no sense because of 0-RTT.
# QUIC makes no sense cause of 0-RTT.
add google-dotcp "dotcp://8.8.8.8:53" false false true add google-dotcp "dotcp://8.8.8.8:53" false false true
add cloudflare-dotcp "dotcp://1.1.1.1:53" false false true add cloudflare-dotcp "dotcp://1.1.1.1:53" false false true
add quad9-dotcp "dotcp://9.9.9.10:53" false false true add quad9-dotcp "dotcp://9.9.9.10:53" false false true
@@ -137,37 +155,72 @@ add cloudflare-doh "https://cloudflare-dns.com/dns-query" false false true
add quad9-doh "https://dns10.quad9.net/dns-query" false false true add quad9-doh "https://dns10.quad9.net/dns-query" false false true
add adguard-doh "https://unfiltered.adguard-dns.com/dns-query" false false true add adguard-doh "https://unfiltered.adguard-dns.com/dns-query" false false true
# DNSSEC trust - non-persistent. # ==========================================================================
add google-dotcp "dotcp://8.8.8.8:53" true false false # DNSSEC scenarios: signed domain list + separate output dir (results-dnssec).
add cloudflare-dotcp "dotcp://1.1.1.1:53" true false false # off / trust / auth all share the signed list, so the domain set is held
add quad9-dotcp "dotcp://9.9.9.10:53" true false false # constant across the three modes. Non-persistent throughout.
add adguard-dotcp "dotcp://94.140.14.140:53" true false false # ==========================================================================
add google-dot "tls://8.8.8.8:853" true false false # DNSSEC off - baseline ON THE SIGNED LIST (separate from the protocol-comparison
add cloudflare-dot "tls://1.1.1.1:853" true false false # off runs above, which use the general list).
add quad9-dot "tls://9.9.9.10:853" true false false add_dnssec google-dotcp "dotcp://8.8.8.8:53" false false false
add adguard-dot "tls://94.140.14.140:853" true false false add_dnssec cloudflare-dotcp "dotcp://1.1.1.1:53" false false false
add_dnssec quad9-dotcp "dotcp://9.9.9.10:53" false false false
add_dnssec adguard-dotcp "dotcp://94.140.14.140:53" false false false
add google-doh "https://dns.google/dns-query" true false false add_dnssec google-dot "tls://8.8.8.8:853" false false false
add cloudflare-doh "https://cloudflare-dns.com/dns-query" true false false add_dnssec cloudflare-dot "tls://1.1.1.1:853" false false false
add quad9-doh "https://dns10.quad9.net/dns-query" true false false add_dnssec quad9-dot "tls://9.9.9.10:853" false false false
add adguard-doh "https://unfiltered.adguard-dns.com/dns-query" true false false add_dnssec adguard-dot "tls://94.140.14.140:853" false false false
add google-doh3 "doh3://dns.google/dns-query" true false false add_dnssec google-doh "https://dns.google/dns-query" false false false
add cloudflare-doh3 "doh3://cloudflare-dns.com/dns-query" true false false add_dnssec cloudflare-doh "https://cloudflare-dns.com/dns-query" false false false
add adguard-doh3 "doh3://unfiltered.adguard-dns.com/dns-query" true false false add_dnssec quad9-doh "https://dns10.quad9.net/dns-query" false false false
add adguard-doq "doq://unfiltered.adguard-dns.com:853" true false false add_dnssec adguard-doh "https://unfiltered.adguard-dns.com/dns-query" false false false
add google-doudp "udp://8.8.8.8:53" true false false add_dnssec google-doh3 "doh3://dns.google/dns-query" false false false
add cloudflare-doudp "udp://1.1.1.1:53" true false false add_dnssec cloudflare-doh3 "doh3://cloudflare-dns.com/dns-query" false false false
add quad9-doudp "udp://9.9.9.10:53" true false false add_dnssec adguard-doh3 "doh3://unfiltered.adguard-dns.com/dns-query" false false false
add adguard-doudp "udp://94.140.14.140:53" true false false add_dnssec adguard-doq "doq://unfiltered.adguard-dns.com:853" false false false
add adguard-dnscrypt "sdns://AQMAAAAAAAAAETk0LjE0MC4xNC4xNDo1NDQzINErR_JS3PLCu_iZEIbq95zkSV2LFsigxDIuUso_OQhzIjIuZG5zY3J5cHQuZGVmYXVsdC5uczEuYWRndWFyZC5jb20" true false false
add quad9-dnscrypt "sdns://AQMAAAAAAAAAFDE0OS4xMTIuMTEyLjExMjo4NDQzIGfIR7jIdYzRICRVQ751Z0bfNN8dhMALjEcDaN-CHYY-GTIuZG5zY3J5cHQtY2VydC5xdWFkOS5uZXQ" true false false add_dnssec google-doudp "udp://8.8.8.8:53" false false false
add_dnssec cloudflare-doudp "udp://1.1.1.1:53" false false false
add_dnssec quad9-doudp "udp://9.9.9.10:53" false false false
add_dnssec adguard-doudp "udp://94.140.14.140:53" false false false
add_dnssec adguard-dnscrypt "sdns://AQMAAAAAAAAAETk0LjE0MC4xNC4xNDo1NDQzINErR_JS3PLCu_iZEIbq95zkSV2LFsigxDIuUso_OQhzIjIuZG5zY3J5cHQuZGVmYXVsdC5uczEuYWRndWFyZC5jb20" false false false
add_dnssec quad9-dnscrypt "sdns://AQMAAAAAAAAAFDE0OS4xMTIuMTEyLjExMjo4NDQzIGfIR7jIdYzRICRVQ751Z0bfNN8dhMALjEcDaN-CHYY-GTIuZG5zY3J5cHQtY2VydC5xdWFkOS5uZXQ" false false false
# DNSSEC trust - non-persistent, signed list.
add_dnssec google-dotcp "dotcp://8.8.8.8:53" true false false
add_dnssec cloudflare-dotcp "dotcp://1.1.1.1:53" true false false
add_dnssec quad9-dotcp "dotcp://9.9.9.10:53" true false false
add_dnssec adguard-dotcp "dotcp://94.140.14.140:53" true false false
add_dnssec google-dot "tls://8.8.8.8:853" true false false
add_dnssec cloudflare-dot "tls://1.1.1.1:853" true false false
add_dnssec quad9-dot "tls://9.9.9.10:853" true false false
add_dnssec adguard-dot "tls://94.140.14.140:853" true false false
add_dnssec google-doh "https://dns.google/dns-query" true false false
add_dnssec cloudflare-doh "https://cloudflare-dns.com/dns-query" true false false
add_dnssec quad9-doh "https://dns10.quad9.net/dns-query" true false false
add_dnssec adguard-doh "https://unfiltered.adguard-dns.com/dns-query" true false false
add_dnssec google-doh3 "doh3://dns.google/dns-query" true false false
add_dnssec cloudflare-doh3 "doh3://cloudflare-dns.com/dns-query" true false false
add_dnssec adguard-doh3 "doh3://unfiltered.adguard-dns.com/dns-query" true false false
add_dnssec adguard-doq "doq://unfiltered.adguard-dns.com:853" true false false
add_dnssec google-doudp "udp://8.8.8.8:53" true false false
add_dnssec cloudflare-doudp "udp://1.1.1.1:53" true false false
add_dnssec quad9-doudp "udp://9.9.9.10:53" true false false
add_dnssec adguard-doudp "udp://94.140.14.140:53" true false false
add_dnssec adguard-dnscrypt "sdns://AQMAAAAAAAAAETk0LjE0MC4xNC4xNDo1NDQzINErR_JS3PLCu_iZEIbq95zkSV2LFsigxDIuUso_OQhzIjIuZG5zY3J5cHQuZGVmYXVsdC5uczEuYWRndWFyZC5jb20" true false false
add_dnssec quad9-dnscrypt "sdns://AQMAAAAAAAAAFDE0OS4xMTIuMTEyLjExMjo4NDQzIGfIR7jIdYzRICRVQ751Z0bfNN8dhMALjEcDaN-CHYY-GTIuZG5zY3J5cHQtY2VydC5xdWFkOS5uZXQ" true false false
# DNSSEC auth - DoUDP only. # DNSSEC auth - DoUDP only.
# auth walks the chain of trust over plain DoUDP regardless of the protocol. # auth walks the chain of trust over plain DoUDP regardless of the protocol.
add adguard-doudp "udp://unfiltered.adguard-dns.com:53" true true false add_dnssec adguard-doudp "udp://94.140.14.140:53" true true false
# Function to get flags suffix for filename # Function to get flags suffix for filename
get_flags_suffix() { get_flags_suffix() {
@@ -200,6 +253,8 @@ run_with_perf() {
local dnssec="$3" local dnssec="$3"
local auth="$4" local auth="$4"
local keepalive="$5" local keepalive="$5"
local domains_file="$6"
local output_dir="$7"
local suffix=$(get_flags_suffix "$dnssec" "$auth" "$keepalive") local suffix=$(get_flags_suffix "$dnssec" "$auth" "$keepalive")
local provider="${name%%-*}" local provider="${name%%-*}"
@@ -208,7 +263,7 @@ run_with_perf() {
if [[ -n "$suffix" ]]; then if [[ -n "$suffix" ]]; then
base_name="${protocol}-${suffix}" base_name="${protocol}-${suffix}"
fi fi
local cpu_csv_file="$OUTPUT_DIR/${provider}/${base_name}.cpu.csv" local cpu_csv_file="$output_dir/${provider}/${base_name}.cpu.csv"
# Create directory if needed # Create directory if needed
mkdir -p "$(dirname "$cpu_csv_file")" mkdir -p "$(dirname "$cpu_csv_file")"
@@ -220,8 +275,8 @@ run_with_perf() {
# Build command arguments # Build command arguments
local cmd_args=( local cmd_args=(
"$DOMAINS_FILE" "$domains_file"
--output-dir "$OUTPUT_DIR" --output-dir "$output_dir"
--interface "$INTERFACE" --interface "$INTERFACE"
--timeout "$TIMEOUT" --timeout "$TIMEOUT"
-s "$url" -s "$url"
@@ -259,18 +314,29 @@ run_with_perf() {
# Cleanup # Cleanup
rm -f "$perf_tmp" rm -f "$perf_tmp"
echo " -> CPU metrics saved to ${provider}/${base_name}.cpu.csv" echo " -> CPU metrics saved to ${output_dir}/${provider}/${base_name}.cpu.csv"
} }
echo "Total scenarios: ${#SC_NAME[@]}" echo "Total scenarios: ${#SC_NAME[@]}"
echo "" echo ""
# Guard: all scenario arrays must be the same length (they desync only if a row
# was added without going through add()/add_dnssec()).
_n=${#SC_NAME[@]}
if (( ${#SC_URL[@]} != _n || ${#SC_DNSSEC[@]} != _n || ${#SC_AUTH[@]} != _n \
|| ${#SC_KEEP[@]} != _n || ${#SC_DOMAINS[@]} != _n || ${#SC_OUTDIR[@]} != _n )); then
echo "ERROR: scenario arrays are misaligned (lengths differ). Check add()/add_dnssec() calls." >&2
exit 1
fi
for i in "${!SC_NAME[@]}"; do for i in "${!SC_NAME[@]}"; do
name="${SC_NAME[$i]}" name="${SC_NAME[$i]}"
url="${SC_URL[$i]}" url="${SC_URL[$i]}"
dnssec="${SC_DNSSEC[$i]}" dnssec="${SC_DNSSEC[$i]}"
auth="${SC_AUTH[$i]}" auth="${SC_AUTH[$i]}"
keepalive="${SC_KEEP[$i]}" keepalive="${SC_KEEP[$i]}"
domains_file="${SC_DOMAINS[$i]}"
output_dir="${SC_OUTDIR[$i]}"
suffix=$(get_flags_suffix "$dnssec" "$auth" "$keepalive") suffix=$(get_flags_suffix "$dnssec" "$auth" "$keepalive")
provider="${name%%-*}" provider="${name%%-*}"
@@ -279,13 +345,14 @@ for i in "${!SC_NAME[@]}"; do
[[ -n "$suffix" ]] && label="${protocol}-${suffix}" [[ -n "$suffix" ]] && label="${protocol}-${suffix}"
if [[ "$DRY_RUN" == "true" ]]; then if [[ "$DRY_RUN" == "true" ]]; then
printf '[dry] %-28s dnssec=%-5s auth=%-5s keepalive=%-5s -> %s/%s\n' \ printf '[dry] %-28s dnssec=%-5s auth=%-5s keepalive=%-5s [%s] -> %s/%s/%s\n' \
"$name" "$dnssec" "$auth" "$keepalive" "$provider" "$label" "$name" "$dnssec" "$auth" "$keepalive" \
"$(basename "$domains_file")" "$output_dir" "$provider" "$label"
continue continue
fi fi
echo "Processing: $name ($url) [dnssec=$dnssec auth=$auth keepalive=$keepalive]" echo "Processing: $name ($url) [dnssec=$dnssec auth=$auth keepalive=$keepalive, domains=$(basename "$domains_file"), out=$output_dir]"
run_with_perf "$name" "$url" "$dnssec" "$auth" "$keepalive" run_with_perf "$name" "$url" "$dnssec" "$auth" "$keepalive" "$domains_file" "$output_dir"
sleep "$SLEEP_TIME" sleep "$SLEEP_TIME"
done done
-598
View File
@@ -1,598 +0,0 @@
package server
import (
"context"
"fmt"
"net"
"net/url"
"os"
"os/signal"
"strings"
"sync"
"syscall"
"time"
"github.com/afonsofrancof/sdns-proxy/client"
"github.com/afonsofrancof/sdns-proxy/common/logger"
"github.com/miekg/dns"
)
type Config struct {
Address string
Upstream string
Fallback string
Bootstrap string
DNSSEC bool
AuthoritativeDNSSEC bool
KeepAlive bool
Timeout time.Duration
Verbose bool
}
type cacheKey struct {
domain string
qtype uint16
}
type cacheEntry struct {
records []dns.RR
expiresAt time.Time
}
type Server struct {
config Config
upstreamClient client.DNSClient
fallbackClient client.DNSClient
bootstrapClient client.DNSClient
resolvedHosts map[string]string
queryCache map[cacheKey]*cacheEntry
hostsMutex sync.RWMutex
cacheMutex sync.RWMutex
dnsServer *dns.Server
}
func New(config Config) (*Server, error) {
logger.Debug("Creating new server with config: %+v", config)
if config.Upstream == "" {
logger.Error("Upstream server is required")
return nil, fmt.Errorf("upstream server is required")
}
// Check if we need bootstrap server
needsBootstrap := containsHostname(config.Upstream)
if config.Fallback != "" {
needsBootstrap = needsBootstrap || containsHostname(config.Fallback)
}
logger.Debug("Bootstrap needed: %v (upstream has hostname: %v, fallback has hostname: %v)",
needsBootstrap, containsHostname(config.Upstream),
config.Fallback != "" && containsHostname(config.Fallback))
if needsBootstrap && config.Bootstrap == "" {
logger.Error("Bootstrap server is required when upstream or fallback contains hostnames")
return nil, fmt.Errorf("bootstrap server is required when upstream or fallback contains hostnames")
}
if config.Bootstrap != "" && containsHostname(config.Bootstrap) {
logger.Error("Bootstrap server cannot contain hostnames: %s", config.Bootstrap)
return nil, fmt.Errorf("bootstrap server cannot contain hostnames: %s", config.Bootstrap)
}
s := &Server{
config: config,
resolvedHosts: make(map[string]string),
queryCache: make(map[cacheKey]*cacheEntry),
}
// Create bootstrap client if needed
if config.Bootstrap != "" {
logger.Debug("Creating bootstrap client for %s", config.Bootstrap)
bootstrapClient, err := client.New(config.Bootstrap, client.Options{
DNSSEC: false,
KeepAlive: config.KeepAlive, // Pass KeepAlive to bootstrap client
})
if err != nil {
logger.Error("Failed to create bootstrap client: %v", err)
return nil, fmt.Errorf("failed to create bootstrap client: %w", err)
}
s.bootstrapClient = bootstrapClient
logger.Debug("Bootstrap client created successfully")
}
// Initialize upstream and fallback clients
if err := s.initClients(); err != nil {
logger.Error("Failed to initialize clients: %v", err)
return nil, fmt.Errorf("failed to initialize clients: %w", err)
}
// Setup DNS server
mux := dns.NewServeMux()
mux.HandleFunc(".", s.handleDNSRequest)
s.dnsServer = &dns.Server{
Addr: config.Address,
Net: "udp",
Handler: mux,
}
logger.Debug("Server created successfully, listening on %s", config.Address)
return s, nil
}
func containsHostname(serverAddr string) bool {
logger.Debug("Checking if %s contains hostname", serverAddr)
// Use the same parsing logic as the client package
parsedURL, err := url.Parse(serverAddr)
if err != nil {
logger.Debug("URL parsing failed for %s, treating as plain address", serverAddr)
// If URL parsing fails, assume it's a plain address
host, _, err := net.SplitHostPort(serverAddr)
if err != nil {
// Assume it's just a host
isHostname := net.ParseIP(serverAddr) == nil
logger.Debug("Address %s is hostname: %v", serverAddr, isHostname)
return isHostname
}
isHostname := net.ParseIP(host) == nil
logger.Debug("Host %s from %s is hostname: %v", host, serverAddr, isHostname)
return isHostname
}
host := parsedURL.Hostname()
if host == "" {
logger.Debug("No hostname found in URL %s", serverAddr)
return false
}
isHostname := net.ParseIP(host) == nil
logger.Debug("Host %s from URL %s is hostname: %v", host, serverAddr, isHostname)
return isHostname
}
func (s *Server) initClients() error {
logger.Debug("Initializing DNS clients")
// Initialize upstream client
resolvedUpstream, err := s.resolveServerAddress(s.config.Upstream)
if err != nil {
logger.Error("Failed to resolve upstream %s: %v", s.config.Upstream, err)
return fmt.Errorf("failed to resolve upstream %s: %w", s.config.Upstream, err)
}
logger.Debug("Creating upstream client for %s (resolved: %s)", s.config.Upstream, resolvedUpstream)
upstreamClient, err := client.New(resolvedUpstream, client.Options{
DNSSEC: s.config.DNSSEC,
AuthoritativeDNSSEC: s.config.AuthoritativeDNSSEC,
KeepAlive: s.config.KeepAlive,
})
if err != nil {
logger.Error("Failed to create upstream client: %v", err)
return fmt.Errorf("failed to create upstream client: %w", err)
}
s.upstreamClient = upstreamClient
if s.config.Verbose {
logger.Info("Initialized upstream client: %s -> %s (KeepAlive: %v)", s.config.Upstream, resolvedUpstream, s.config.KeepAlive)
}
// Initialize fallback client if specified
if s.config.Fallback != "" {
resolvedFallback, err := s.resolveServerAddress(s.config.Fallback)
if err != nil {
logger.Error("Failed to resolve fallback %s: %v", s.config.Fallback, err)
return fmt.Errorf("failed to resolve fallback %s: %w", s.config.Fallback, err)
}
logger.Debug("Creating fallback client for %s (resolved: %s)", s.config.Fallback, resolvedFallback)
fallbackClient, err := client.New(resolvedFallback, client.Options{
DNSSEC: s.config.DNSSEC,
KeepAlive: s.config.KeepAlive, // Pass KeepAlive to fallback client
})
if err != nil {
logger.Error("Failed to create fallback client: %v", err)
return fmt.Errorf("failed to create fallback client: %w", err)
}
s.fallbackClient = fallbackClient
if s.config.Verbose {
logger.Info("Initialized fallback client: %s -> %s (KeepAlive: %v)", s.config.Fallback, resolvedFallback, s.config.KeepAlive)
}
}
logger.Debug("All DNS clients initialized successfully")
return nil
}
func (s *Server) resolveServerAddress(serverAddr string) (string, error) {
logger.Debug("Resolving server address: %s", serverAddr)
// If it doesn't contain hostnames, return as-is
if !containsHostname(serverAddr) {
logger.Debug("Address %s contains no hostnames, returning as-is", serverAddr)
return serverAddr, nil
}
// If no bootstrap client, we can't resolve hostnames
if s.bootstrapClient == nil {
logger.Error("Cannot resolve hostname in %s: no bootstrap server configured", serverAddr)
return "", fmt.Errorf("cannot resolve hostname in %s: no bootstrap server configured", serverAddr)
}
// Use the same parsing logic as the client package
parsedURL, err := url.Parse(serverAddr)
if err != nil {
logger.Debug("Parsing %s as plain host:port format", serverAddr)
// Handle plain host:port format
host, port, err := net.SplitHostPort(serverAddr)
if err != nil {
// Assume it's just a hostname
resolvedIP, err := s.resolveHostname(serverAddr)
if err != nil {
return "", err
}
logger.Debug("Resolved %s to %s", serverAddr, resolvedIP)
return resolvedIP, nil
}
resolvedIP, err := s.resolveHostname(host)
if err != nil {
return "", err
}
resolved := net.JoinHostPort(resolvedIP, port)
logger.Debug("Resolved %s to %s", serverAddr, resolved)
return resolved, nil
}
// Handle URL format
hostname := parsedURL.Hostname()
if hostname == "" {
logger.Error("No hostname in URL: %s", serverAddr)
return "", fmt.Errorf("no hostname in URL: %s", serverAddr)
}
resolvedIP, err := s.resolveHostname(hostname)
if err != nil {
return "", err
}
// Replace hostname with IP in the URL
port := parsedURL.Port()
if port == "" {
parsedURL.Host = resolvedIP
} else {
parsedURL.Host = net.JoinHostPort(resolvedIP, port)
}
resolved := parsedURL.String()
logger.Debug("Resolved URL %s to %s", serverAddr, resolved)
return resolved, nil
}
func (s *Server) resolveHostname(hostname string) (string, error) {
logger.Debug("Resolving hostname: %s", hostname)
// Check cache first
s.hostsMutex.RLock()
if ip, exists := s.resolvedHosts[hostname]; exists {
s.hostsMutex.RUnlock()
logger.Debug("Found cached resolution for %s: %s", hostname, ip)
return ip, nil
}
s.hostsMutex.RUnlock()
// Resolve using bootstrap
if s.config.Verbose {
logger.Info("Resolving hostname %s using bootstrap server", hostname)
}
msg := new(dns.Msg)
msg.SetQuestion(dns.Fqdn(hostname), dns.TypeA)
msg.Id = dns.Id()
msg.RecursionDesired = true
logger.Debug("Sending bootstrap query for %s (ID: %d)", hostname, msg.Id)
msg, err := s.bootstrapClient.Query(msg)
if err != nil {
logger.Error("Bootstrap query failed for %s: %v", hostname, err)
return "", fmt.Errorf("failed to resolve %s via bootstrap: %w", hostname, err)
}
logger.Debug("Bootstrap response for %s: %d answers", hostname, len(msg.Answer))
if len(msg.Answer) == 0 {
logger.Error("No A records found for %s", hostname)
return "", fmt.Errorf("no A records found for %s", hostname)
}
// Find first A record
for _, rr := range msg.Answer {
if a, ok := rr.(*dns.A); ok {
ip := a.A.String()
// Cache the result
s.hostsMutex.Lock()
s.resolvedHosts[hostname] = ip
s.hostsMutex.Unlock()
if s.config.Verbose {
logger.Info("Resolved %s to %s", hostname, ip)
}
logger.Debug("Cached resolution: %s -> %s", hostname, ip)
return ip, nil
}
}
logger.Error("No valid A record found for %s", hostname)
return "", fmt.Errorf("no valid A record found for %s", hostname)
}
func (s *Server) handleDNSRequest(w dns.ResponseWriter, r *dns.Msg) {
if len(r.Question) == 0 {
logger.Debug("Received request with no questions from %s", w.RemoteAddr())
dns.HandleFailed(w, r)
return
}
question := r.Question[0]
domain := strings.ToLower(question.Name)
qtype := question.Qtype
logger.Debug("Handling DNS request: %s %s from %s (ID: %d)",
question.Name, dns.TypeToString[qtype], w.RemoteAddr(), r.Id)
if s.config.Verbose {
logger.Info("Query: %s %s from %s",
question.Name,
dns.TypeToString[qtype],
w.RemoteAddr())
}
// Check cache first
if cachedRecords := s.getCachedRecords(domain, qtype); cachedRecords != nil {
response := s.buildResponse(r, cachedRecords)
if s.config.Verbose {
logger.Info("Cache hit: %s %s -> %d records",
question.Name,
dns.TypeToString[qtype],
len(cachedRecords))
}
logger.Debug("Serving cached response for %s %s (%d records)",
question.Name, dns.TypeToString[qtype], len(cachedRecords))
w.WriteMsg(response)
return
}
logger.Debug("Cache miss for %s %s, querying upstream", question.Name, dns.TypeToString[qtype])
// Try upstream first
response, err := s.queryUpstream(s.upstreamClient, question.Name, qtype)
if err != nil {
if s.config.Verbose {
logger.Info("Upstream query failed: %v", err)
}
logger.Debug("Upstream query failed for %s %s: %v", question.Name, dns.TypeToString[qtype], err)
// Try fallback if available
if s.fallbackClient != nil {
if s.config.Verbose {
logger.Info("Trying fallback server")
}
logger.Debug("Attempting fallback query for %s %s", question.Name, dns.TypeToString[qtype])
response, err = s.queryUpstream(s.fallbackClient, question.Name, qtype)
if err != nil {
logger.Error("Both upstream and fallback failed for %s %s: %v",
question.Name,
dns.TypeToString[qtype],
err)
} else {
logger.Debug("Fallback query succeeded for %s %s", question.Name, dns.TypeToString[qtype])
}
}
// If still failed, return SERVFAIL
if err != nil {
logger.Error("All servers failed for %s %s: %v",
question.Name,
dns.TypeToString[qtype],
err)
m := new(dns.Msg)
m.SetReply(r)
m.Rcode = dns.RcodeServerFailure
w.WriteMsg(m)
return
}
} else {
logger.Debug("Upstream query succeeded for %s %s", question.Name, dns.TypeToString[qtype])
}
// Cache successful response
s.cacheResponse(domain, qtype, response)
// Copy request ID to response
response.Id = r.Id
if s.config.Verbose {
logger.Info("Response: %s %s -> %d answers",
question.Name,
dns.TypeToString[qtype],
len(response.Answer))
}
logger.Debug("Sending response for %s %s: %d answers, rcode: %s",
question.Name, dns.TypeToString[qtype], len(response.Answer), dns.RcodeToString[response.Rcode])
w.WriteMsg(response)
}
func (s *Server) getCachedRecords(domain string, qtype uint16) []dns.RR {
key := cacheKey{domain: domain, qtype: qtype}
s.cacheMutex.RLock()
entry, exists := s.queryCache[key]
s.cacheMutex.RUnlock()
if !exists {
logger.Debug("No cache entry for %s %s", domain, dns.TypeToString[qtype])
return nil
}
// Check if expired and clean up on the spot
if time.Now().After(entry.expiresAt) {
logger.Debug("Cache entry expired for %s %s", domain, dns.TypeToString[qtype])
s.cacheMutex.Lock()
delete(s.queryCache, key)
s.cacheMutex.Unlock()
return nil
}
logger.Debug("Cache hit for %s %s (%d records, expires in %v)",
domain, dns.TypeToString[qtype], len(entry.records), time.Until(entry.expiresAt))
// Return a copy of the cached records
records := make([]dns.RR, len(entry.records))
for i, rr := range entry.records {
records[i] = dns.Copy(rr)
}
return records
}
func (s *Server) buildResponse(request *dns.Msg, records []dns.RR) *dns.Msg {
response := new(dns.Msg)
response.SetReply(request)
response.Answer = records
logger.Debug("Built response with %d records", len(records))
return response
}
func (s *Server) cacheResponse(domain string, qtype uint16, msg *dns.Msg) {
if msg == nil || len(msg.Answer) == 0 {
logger.Debug("Not caching empty response for %s %s", domain, dns.TypeToString[qtype])
return
}
var validRecords []dns.RR
minTTL := uint32(3600)
// Find minimum TTL from answer records
for _, rr := range msg.Answer {
// Only cache records that match our query type or are CNAMEs
if rr.Header().Rrtype == qtype || rr.Header().Rrtype == dns.TypeCNAME {
validRecords = append(validRecords, dns.Copy(rr))
if rr.Header().Ttl < minTTL {
minTTL = rr.Header().Ttl
}
}
}
if len(validRecords) == 0 {
logger.Debug("No valid records to cache for %s %s", domain, dns.TypeToString[qtype])
return
}
// Don't cache responses with very low TTL
if minTTL < 10 {
logger.Debug("TTL too low (%ds) for caching %s %s", minTTL, domain, dns.TypeToString[qtype])
return
}
key := cacheKey{domain: domain, qtype: qtype}
entry := &cacheEntry{
records: validRecords,
expiresAt: time.Now().Add(time.Duration(minTTL) * time.Second),
}
s.cacheMutex.Lock()
s.queryCache[key] = entry
s.cacheMutex.Unlock()
if s.config.Verbose {
logger.Info("Cached %d records for %s %s (TTL: %ds)",
len(validRecords), domain, dns.TypeToString[qtype], minTTL)
}
logger.Debug("Cached %d records for %s %s (TTL: %ds, expires: %v)",
len(validRecords), domain, dns.TypeToString[qtype], minTTL, entry.expiresAt)
}
func (s *Server) queryUpstream(upstreamClient client.DNSClient, domain string, qtype uint16) (*dns.Msg, error) {
logger.Debug("Querying upstream for %s %s", domain, dns.TypeToString[qtype])
// Create context with timeout
ctx, cancel := context.WithTimeout(context.Background(), s.config.Timeout)
defer cancel()
// Channel to receive result
type result struct {
msg *dns.Msg
err error
}
resultChan := make(chan result, 1)
// Query in goroutine to respect context timeout
go func() {
msg := new(dns.Msg)
msg.SetQuestion(dns.Fqdn(domain), qtype)
msg.Id = dns.Id()
msg.RecursionDesired = true
logger.Debug("Sending upstream query: %s %s (ID: %d)", domain, dns.TypeToString[qtype], msg.Id)
recvMsg, err := upstreamClient.Query(msg)
if err != nil {
logger.Debug("Upstream query error for %s %s: %v", domain, dns.TypeToString[qtype], err)
} else {
logger.Debug("Upstream query response for %s %s: %d answers, rcode: %s",
domain, dns.TypeToString[qtype], len(recvMsg.Answer), dns.RcodeToString[recvMsg.Rcode])
}
resultChan <- result{msg: recvMsg, err: err}
}()
select {
case res := <-resultChan:
return res.msg, res.err
case <-ctx.Done():
logger.Debug("Upstream query timeout for %s %s after %v", domain, dns.TypeToString[qtype], s.config.Timeout)
return nil, fmt.Errorf("upstream query timeout")
}
}
func (s *Server) Start() error {
go func() {
sigChan := make(chan os.Signal, 1)
signal.Notify(sigChan, syscall.SIGINT, syscall.SIGTERM)
sig := <-sigChan
logger.Info("Received signal %v, shutting down DNS server...", sig)
s.Shutdown()
}()
logger.Info("DNS proxy server listening on %s", s.config.Address)
logger.Debug("Server starting with timeout: %v, DNSSEC: %v, KeepAlive: %v", s.config.Timeout, s.config.DNSSEC, s.config.KeepAlive)
return s.dnsServer.ListenAndServe()
}
func (s *Server) Shutdown() {
logger.Debug("Shutting down server components")
if s.dnsServer != nil {
logger.Debug("Shutting down DNS server")
s.dnsServer.Shutdown()
}
if s.upstreamClient != nil {
logger.Debug("Closing upstream client")
s.upstreamClient.Close()
}
if s.fallbackClient != nil {
logger.Debug("Closing fallback client")
s.fallbackClient.Close()
}
if s.bootstrapClient != nil {
logger.Debug("Closing bootstrap client")
s.bootstrapClient.Close()
}
logger.Info("Server shutdown complete")
}