From 4ef9ea5a4a3c79e744b7c46001a57c6033c0344d Mon Sep 17 00:00:00 2001 From: afonsofrancof Date: Fri, 10 Jul 2026 18:59:49 +0100 Subject: [PATCH] feat(everything): did a massive dnssec refactor to improve robustness --- Makefile | 50 ++- client/client.go | 113 +---- client/validating.go | 100 +++++ cmd/sdns-proxy/sdns-proxy.go | 49 +- common/dnssec/anchor.go | 95 ++++ common/dnssec/authchain.go | 226 ---------- common/dnssec/authoritative.go | 212 --------- common/dnssec/validator.go | 136 +++--- common/dnssec/walker_iterative.go | 248 ++++++++++ common/dnssec/walker_trust.go | 136 ++++++ common/protocols/dnscrypt/dnscrypt.go | 10 +- common/protocols/doh/doh.go | 36 +- common/protocols/doq/doq.go | 30 +- common/protocols/dot/dot.go | 34 +- common/protocols/dotcp/dotcp.go | 34 +- common/protocols/doudp/doudp.go | 32 +- go.mod | 2 +- internal/qol/measurement.go | 37 +- internal/qol/results/writer.go | 8 +- internal/qol/stats/runtime.go | 14 +- scripts/post_processing/merge_files.py | 4 + scripts/post_processing/merge_mem.py | 3 +- server/server.go | 598 ------------------------- 23 files changed, 809 insertions(+), 1398 deletions(-) create mode 100644 client/validating.go create mode 100644 common/dnssec/anchor.go delete mode 100644 common/dnssec/authchain.go delete mode 100644 common/dnssec/authoritative.go create mode 100644 common/dnssec/walker_iterative.go create mode 100644 common/dnssec/walker_trust.go delete mode 100644 server/server.go diff --git a/Makefile b/Makefile index 5792227..001c5b3 100644 --- a/Makefile +++ b/Makefile @@ -1,10 +1,44 @@ -.PHONY: all clean +PY := python3 +PP := scripts/post_processing -all: - python3 scripts/post_processing/merge_files.py - python3 scripts/post_processing/merge_cpu.py - python3 scripts/post_processing/merge_mem.py - python3 scripts/post_processing/merge_pcap.py +# Input trees produced by run.sh +GEN_IN := results +DNS_IN := results-dnssec -clean: - rm -f ./dns_results.csv ./dns_results_cpu.csv ./dns_results_mem.csv ./dns_results_pcap.csv +# Output layout +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: general dnssec + +# ----- 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 diff --git a/client/client.go b/client/client.go index 9d91108..985c6cd 100644 --- a/client/client.go +++ b/client/client.go @@ -6,7 +6,6 @@ import ( "net/url" "strings" - "github.com/afonsofrancof/sdns-proxy/common/dnssec" "github.com/afonsofrancof/sdns-proxy/common/logger" "github.com/afonsofrancof/sdns-proxy/common/protocols/dnscrypt" "github.com/afonsofrancof/sdns-proxy/common/protocols/doh" @@ -18,16 +17,10 @@ import ( ) type DNSClient interface { - Query(msg *dns.Msg) (*dns.Msg, error) + Query(msg *dns.Msg) (sent *dns.Msg, resp *dns.Msg, err error) Close() } -type ValidatingDNSClient struct { - client DNSClient - validator *dnssec.Validator - options Options -} - type Options struct { DNSSEC bool AuthoritativeDNSSEC bool @@ -39,7 +32,6 @@ type Options struct { func New(upstream string, opts Options) (DNSClient, error) { logger.Debug("Creating DNS client for upstream: %s with options: %+v", upstream, opts) - // Try to parse as URL parsedURL, err := url.Parse(upstream) if err != nil { logger.Error("Invalid upstream format: %v", err) @@ -48,12 +40,10 @@ func New(upstream string, opts Options) (DNSClient, error) { var baseClient DNSClient - // If it has a scheme, treat it as a full URL if parsedURL.Scheme != "" { logger.Debug("Parsing %s as URL with scheme %s", upstream, parsedURL.Scheme) baseClient, err = createClientFromURL(parsedURL, opts) } else { - // No scheme - treat as plain DNS address (defaults to UDP) logger.Debug("Parsing %s as plain DNS address", upstream) baseClient, err = createClientFromPlainAddress(upstream, opts) } @@ -63,105 +53,14 @@ func New(upstream string, opts Options) (DNSClient, error) { return nil, err } - // If DNSSEC is not enabled, return the base client + // Without DNSSEC, the base protocol client is returned directly. if !opts.DNSSEC { logger.Debug("DNSSEC disabled, returning base client") return baseClient, nil } - logger.Debug("DNSSEC enabled, wrapping with validator (AuthoritativeDNSSEC: %v)", opts.AuthoritativeDNSSEC) - - 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() - } + // With DNSSEC, wrap the base client with a validating client. + return NewValidating(baseClient, opts), nil } 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) - // Default to UDP for plain addresses return createClient("udp", host, port, "", opts) } @@ -290,9 +188,6 @@ func createClient(scheme, host, port, path string, opts Options) (DNSClient, err case "sdns": 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), DNSSEC: opts.DNSSEC, } diff --git a/client/validating.go b/client/validating.go new file mode 100644 index 0000000..413f8c7 --- /dev/null +++ b/client/validating.go @@ -0,0 +1,100 @@ +package client + +import ( + "fmt" + + "github.com/afonsofrancof/sdns-proxy/common/dnssec" + "github.com/afonsofrancof/sdns-proxy/common/logger" + "github.com/miekg/dns" +) + +type ValidatingDNSClient struct { + client DNSClient + validator *dnssec.Validator + options Options +} + +func NewValidating(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() + } else { + validator = dnssec.NewValidator(func(m *dns.Msg) (*dns.Msg, error) { _, r, e := base.Query(m); return r, e }) + } + + return &ValidatingDNSClient{ + client: base, + validator: validator, + options: opts, + } +} + +func (v *ValidatingDNSClient) LastValidation() dnssec.ValidationStats { + if v.validator == nil { + return dnssec.ValidationStats{} + } + return v.validator.TakeStats() +} + +func (v *ValidatingDNSClient) Query(msg *dns.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) + } + + // DNSSEC policy lives here + msg.SetEdns0(4096, true) + // 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 *ValidatingDNSClient) Close() { + logger.Debug("Closing ValidatingDNSClient") + if v.client != nil { + v.client.Close() + } +} diff --git a/cmd/sdns-proxy/sdns-proxy.go b/cmd/sdns-proxy/sdns-proxy.go index 408bfae..3a1b00b 100644 --- a/cmd/sdns-proxy/sdns-proxy.go +++ b/cmd/sdns-proxy/sdns-proxy.go @@ -7,16 +7,14 @@ import ( "github.com/afonsofrancof/sdns-proxy/client" "github.com/afonsofrancof/sdns-proxy/common/logger" - "github.com/afonsofrancof/sdns-proxy/server" "github.com/alecthomas/kong" "github.com/miekg/dns" ) var cli struct { - Debug bool `help:"Enable debug logging globally." short:"D" env:"DEBUG"` - Query QueryCmd `cmd:"" help:"Perform a DNS query (client mode)."` - Listen ListenCmd `cmd:"" help:"Run as a DNS listener/resolver (server mode)."` + Debug bool `help:"Enable debug logging globally." short:"D" env:"DEBUG"` + Query QueryCmd `cmd:"" help:"Perform a DNS query (client mode)."` } 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"` } -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 { 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) @@ -77,7 +63,7 @@ func (q *QueryCmd) Run() error { msg.RecursionDesired = true 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 { logger.Error("DNS query failed: %v", err) return err @@ -89,35 +75,6 @@ func (q *QueryCmd) Run() error { 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) { fmt.Println(";; QUESTION SECTION:") diff --git a/common/dnssec/anchor.go b/common/dnssec/anchor.go new file mode 100644 index 0000000..75bda24 --- /dev/null +++ b/common/dnssec/anchor.go @@ -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 +} diff --git a/common/dnssec/authchain.go b/common/dnssec/authchain.go deleted file mode 100644 index 19b59ab..0000000 --- a/common/dnssec/authchain.go +++ /dev/null @@ -1,226 +0,0 @@ -package dnssec - -// CODE ADAPTED FROM THIS - -// ISC License -// -// Copyright (c) 2012-2016 Peter Banik -// -// 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 -} diff --git a/common/dnssec/authoritative.go b/common/dnssec/authoritative.go deleted file mode 100644 index 18c132f..0000000 --- a/common/dnssec/authoritative.go +++ /dev/null @@ -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 -} diff --git a/common/dnssec/validator.go b/common/dnssec/validator.go index caeb00c..2a75185 100644 --- a/common/dnssec/validator.go +++ b/common/dnssec/validator.go @@ -1,100 +1,84 @@ package dnssec -// CODE ADAPTED FROM THIS - -// ISC License -// -// Copyright (c) 2012-2016 Peter Banik -// -// 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 ( "github.com/afonsofrancof/sdns-proxy/common/logger" "github.com/miekg/dns" ) -type Validator struct { - queryFunc func(string, uint16) (*dns.Msg, error) +type ValidationStats struct { + Queries int + BytesSent int + BytesReceived int + Validated bool } -func NewValidator(queryFunc func(string, uint16) (*dns.Msg, error)) *Validator { - return &Validator{ - queryFunc: queryFunc, - } +type SendFunc func(msg *dns.Msg) (*dns.Msg, error) + +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() *Validator { + return &Validator{walker: newIterativeWalker()} } 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 { - logger.Debug("No result for %s %s", qname, dns.TypeToString[qtype]) return ErrNoResult } - // Extract RRSet from response - rrset := NewRRSet() - for _, rr := range msg.Answer { + answer := extractRRSet(msg, qname, qtype) + if answer.IsEmpty() { + return ErrNoResult + } + if !answer.IsSigned() { + return ErrResourceNotSigned + } + if err := answer.CheckHeaderIntegrity(dns.Fqdn(qname)); err != nil { + return err + } + + logger.Debug("Validating %s %s (signer: %s)", qname, dns.TypeToString[qtype], answer.SignerName()) + return v.walker.validate(answer, qname, qtype) +} + +func (v *Validator) TakeStats() ValidationStats { + s := v.walker.stats() + v.walker.resetStats() + return s +} + +func extractRRSet(msg *dns.Msg, name string, qtype uint16) *RRSet { + if msg == nil { + return NewRRSet() + } + return extractRRSetFrom(msg.Answer, name, qtype) +} + +func extractRRSetFrom(rrs []dns.RR, name string, qtype uint16) *RRSet { + set := NewRRSet() + fq := dns.Fqdn(name) + for _, rr := range rrs { 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) + if t.TypeCovered == qtype && dns.Fqdn(t.Header().Name) == fq { + set.RRSig = t } 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 rr.Header().Rrtype == qtype && dns.Fqdn(rr.Header().Name) == fq { + set.RRs = append(set.RRs, rr) } } } - - if rrset.IsEmpty() { - logger.Debug("Empty RRSet for %s %s", qname, dns.TypeToString[qtype]) - return ErrNoResult - } - - if !rrset.IsSigned() { - logger.Debug("RRSet for %s %s is not signed", qname, dns.TypeToString[qtype]) - return ErrResourceNotSigned - } - - // 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 - } - - // Build and verify authentication chain - signerName := rrset.SignerName() - 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 { - logger.Debug("DNSSEC validation failed for %s %s: %v", qname, dns.TypeToString[qtype], err) - return err - } - - logger.Debug("DNSSEC validation successful for %s %s", qname, dns.TypeToString[qtype]) - return nil -} - -func NewValidatorWithAuthoritativeQueries() *Validator { - querier := NewAuthoritativeQuerier() - return NewValidator(querier.QueryAuthoritative) + return set } diff --git a/common/dnssec/walker_iterative.go b/common/dnssec/walker_iterative.go new file mode 100644 index 0000000..6fe579d --- /dev/null +++ b/common/dnssec/walker_iterative.go @@ -0,0 +1,248 @@ +package dnssec + +import ( + "fmt" + "net" + "strings" + "time" + + "github.com/afonsofrancof/sdns-proxy/common/logger" + "github.com/miekg/dns" +) + +// iterativeWalker validates top-down. +type iterativeWalker struct { + client *dns.Client + st ValidationStats + ipCache map[string]string +} + +func newIterativeWalker() *iterativeWalker { + return &iterativeWalker{ + client: &dns.Client{Timeout: 5 * time.Second}, + 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(4096, true) + m.RecursionDesired = false + + if b, err := m.Pack(); err == nil { + w.st.BytesSent += len(b) + } + w.st.Queries++ + + resp, _, err := w.client.Exchange(m, server) + 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 +} diff --git a/common/dnssec/walker_trust.go b/common/dnssec/walker_trust.go new file mode 100644 index 0000000..4a5d742 --- /dev/null +++ b/common/dnssec/walker_trust.go @@ -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(4096, 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) +} diff --git a/common/protocols/dnscrypt/dnscrypt.go b/common/protocols/dnscrypt/dnscrypt.go index 61b59e9..a34761c 100644 --- a/common/protocols/dnscrypt/dnscrypt.go +++ b/common/protocols/dnscrypt/dnscrypt.go @@ -61,25 +61,21 @@ func (c *Client) Close() { // The dnscrypt library doesn't require explicit cleanup } -func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { +func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) { if len(msg.Question) > 0 { question := msg.Question[0] logger.Debug("DNSCrypt query: %s %s", question.Name, dns.TypeToString[question.Qtype]) } - if c.config.DNSSEC { - msg.SetEdns0(4096, true) - } - response, err := c.resolver.Exchange(msg, c.ri) if err != nil { logger.Error("DNSCrypt query failed: %v", err) - return nil, fmt.Errorf("dnscrypt: query failed: %w", err) + return msg, nil, fmt.Errorf("dnscrypt: query failed: %w", err) } if len(response.Answer) > 0 { logger.Debug("DNSCrypt response: %d answers", len(response.Answer)) } - return response, nil + return msg, response, nil } diff --git a/common/protocols/doh/doh.go b/common/protocols/doh/doh.go index 0966003..ea19e99 100644 --- a/common/protocols/doh/doh.go +++ b/common/protocols/doh/doh.go @@ -61,7 +61,7 @@ func New(config Config) (*Client, error) { tlsConfig := &tls.Config{ ServerName: config.Host, - MinVersion: tls.VersionTLS13, + MinVersion: tls.VersionTLS13, ClientSessionCache: tls.NewLRUClientSessionCache(100), } @@ -114,13 +114,6 @@ func New(config Config) (*Client, error) { }, nil } -// newSingleUseClient builds a throwaway *http.Client whose transport is torn -// down by the returned closeFn. Used for non-persistent (non-keep-alive) mode -// so that each query establishes and tears down its own connection, and thus -// pays its own TCP+TLS (or QUIC) handshake. Without this, the pooled -// HTTP/2 transport silently reuses a single connection across all queries -// (HTTP/2 multiplexes and ignores the Connection: close header), which would -// make the non-persistent measurement indistinguishable from persistent. func (c *Client) newSingleUseClient() (*http.Client, func()) { tlsConfig := &tls.Config{ ServerName: c.config.Host, @@ -162,27 +155,18 @@ func (c *Client) Close() { } } -func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { +func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) { if len(msg.Question) > 0 { question := msg.Question[0] logger.Debug("DoH query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.upstreamURL.Host) } - if c.config.DNSSEC { - msg.SetEdns0(4096, true) - } - packedMsg, err := msg.Pack() if err != nil { logger.Error("DoH failed to pack DNS message: %v", err) - return nil, fmt.Errorf("doh: failed to pack DNS message: %w", err) + return msg, nil, fmt.Errorf("doh: failed to pack DNS message: %w", err) } - // Select the HTTP client. In persistent mode, reuse the pooled client built - // in New(). In non-persistent mode, build a fresh client (and transport) - // per query and tear it down afterwards, so every query pays its own - // handshake — otherwise HTTP/2 would reuse one connection and the - // non-persistent measurement would be wrong. httpClient := c.httpClient if !c.config.KeepAlive { freshClient, closeFn := c.newSingleUseClient() @@ -193,7 +177,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { httpReq, err := http.NewRequest(http.MethodPost, c.upstreamURL.String(), bytes.NewReader(packedMsg)) if err != nil { logger.Error("DoH failed to create HTTP request: %v", err) - return nil, fmt.Errorf("doh: failed to create HTTP request object: %w", err) + return msg, nil, fmt.Errorf("doh: failed to create HTTP request object: %w", err) } httpReq.Header.Set("User-Agent", "sdns-proxy") @@ -203,36 +187,36 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { httpResp, err := httpClient.Do(httpReq) if err != nil { logger.Error("DoH request failed to %s: %v", c.upstreamURL.Host, err) - return nil, fmt.Errorf("doh: failed executing HTTP request to %s: %w", c.upstreamURL.Host, err) + return msg, nil, fmt.Errorf("doh: failed executing HTTP request to %s: %w", c.upstreamURL.Host, err) } defer httpResp.Body.Close() if httpResp.StatusCode != http.StatusOK { logger.Error("DoH received non-200 status from %s: %s", c.upstreamURL.Host, httpResp.Status) - return nil, fmt.Errorf("doh: received non-200 HTTP status from %s: %s", c.upstreamURL.Host, httpResp.Status) + return msg, nil, fmt.Errorf("doh: received non-200 HTTP status from %s: %s", c.upstreamURL.Host, httpResp.Status) } if ct := httpResp.Header.Get("Content-Type"); ct != dnsMessageContentType { logger.Error("DoH unexpected Content-Type from %s: %s", c.upstreamURL.Host, ct) - return nil, fmt.Errorf("doh: unexpected Content-Type from %s: got %q, want %q", c.upstreamURL.Host, ct, dnsMessageContentType) + return msg, nil, fmt.Errorf("doh: unexpected Content-Type from %s: got %q, want %q", c.upstreamURL.Host, ct, dnsMessageContentType) } responseBody, err := io.ReadAll(httpResp.Body) if err != nil { logger.Error("DoH failed reading response from %s: %v", c.upstreamURL.Host, err) - return nil, fmt.Errorf("doh: failed reading response body from %s: %w", c.upstreamURL.Host, err) + return msg, nil, fmt.Errorf("doh: failed reading response body from %s: %w", c.upstreamURL.Host, err) } recvMsg := new(dns.Msg) err = recvMsg.Unpack(responseBody) if err != nil { logger.Error("DoH failed to unpack response from %s: %v", c.upstreamURL.Host, err) - return nil, fmt.Errorf("doh: failed to unpack DNS response from %s: %w", c.upstreamURL.Host, err) + return msg, nil, fmt.Errorf("doh: failed to unpack DNS response from %s: %w", c.upstreamURL.Host, err) } if len(recvMsg.Answer) > 0 { logger.Debug("DoH response from %s: %d answers", c.upstreamURL.Host, len(recvMsg.Answer)) } - return recvMsg, nil + return msg, recvMsg, nil } diff --git a/common/protocols/doq/doq.go b/common/protocols/doq/doq.go index 83d9063..d08b942 100644 --- a/common/protocols/doq/doq.go +++ b/common/protocols/doq/doq.go @@ -95,7 +95,7 @@ func (c *Client) OpenConnection() error { return nil } -func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { +func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) { if len(msg.Question) > 0 { question := msg.Question[0] logger.Debug("DoQ query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.targetAddr) @@ -104,19 +104,17 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { if c.quicConn == nil { err := c.OpenConnection() if err != nil { - return nil, err + return msg, nil, err } } // Prepare DNS message + // DoQ requires Id to be 0 msg.Id = 0 - if c.config.DNSSEC { - msg.SetEdns0(4096, true) - } packed, err := msg.Pack() if err != nil { logger.Error("DoQ failed to pack message: %v", err) - return nil, fmt.Errorf("doq: failed to pack message: %w", err) + return msg, nil, fmt.Errorf("doq: failed to pack message: %w", err) } var quicStream *quic.Stream @@ -125,12 +123,12 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { logger.Debug("DoQ stream failed, reconnecting: %v", err) err = c.OpenConnection() if err != nil { - return nil, err + return msg, nil, err } quicStream, err = c.quicConn.OpenStream() if err != nil { logger.Error("DoQ failed to open stream after reconnect: %v", err) - return nil, err + return msg, nil, err } } @@ -138,18 +136,18 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { err = binary.Write(&lengthPrefixedMessage, binary.BigEndian, uint16(len(packed))) if err != nil { logger.Error("DoQ failed to write message length: %v", err) - return nil, fmt.Errorf("failed to write message length: %w", err) + return msg, nil, fmt.Errorf("failed to write message length: %w", err) } _, err = lengthPrefixedMessage.Write(packed) if err != nil { logger.Error("DoQ failed to write DNS message: %v", err) - return nil, fmt.Errorf("failed to write DNS message: %w", err) + return msg, nil, fmt.Errorf("failed to write DNS message: %w", err) } _, err = quicStream.Write(lengthPrefixedMessage.Bytes()) if err != nil { logger.Error("DoQ failed to write to stream: %v", err) - return nil, fmt.Errorf("failed writing to QUIC stream: %w", err) + return msg, nil, fmt.Errorf("failed writing to QUIC stream: %w", err) } quicStream.Close() @@ -157,32 +155,32 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { _, err = io.ReadFull(quicStream, lengthBuf) if err != nil { logger.Error("DoQ failed to read response length: %v", err) - return nil, fmt.Errorf("failed reading response length: %w", err) + return msg, nil, fmt.Errorf("failed reading response length: %w", err) } messageLength := binary.BigEndian.Uint16(lengthBuf) if messageLength == 0 { logger.Error("DoQ received zero-length message") - return nil, fmt.Errorf("received zero-length message") + return msg, nil, fmt.Errorf("received zero-length message") } responseBuf := make([]byte, messageLength) _, err = io.ReadFull(quicStream, responseBuf) if err != nil { logger.Error("DoQ failed to read response data: %v", err) - return nil, fmt.Errorf("failed reading response data: %w", err) + return msg, nil, fmt.Errorf("failed reading response data: %w", err) } recvMsg := new(dns.Msg) err = recvMsg.Unpack(responseBuf) if err != nil { logger.Error("DoQ failed to parse response: %v", err) - return nil, fmt.Errorf("failed to parse DNS response: %w", err) + return msg, nil, fmt.Errorf("failed to parse DNS response: %w", err) } if len(recvMsg.Answer) > 0 { logger.Debug("DoQ response from %s: %d answers", c.targetAddr, len(recvMsg.Answer)) } - return recvMsg, nil + return msg, recvMsg, nil } diff --git a/common/protocols/dot/dot.go b/common/protocols/dot/dot.go index 1c2c309..9338768 100644 --- a/common/protocols/dot/dot.go +++ b/common/protocols/dot/dot.go @@ -124,7 +124,7 @@ func (c *Client) ensureConnection() error { return nil } -func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { +func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) { if len(msg.Question) > 0 { question := msg.Question[0] logger.Debug("DoT query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort) @@ -133,7 +133,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { // Ensure we have a connection (either persistent or new) if c.config.KeepAlive { if err := c.ensureConnection(); err != nil { - return nil, fmt.Errorf("dot: failed to ensure connection: %w", err) + return msg, nil, fmt.Errorf("dot: failed to ensure connection: %w", err) } } else { // For non-keepalive mode, create a fresh connection for each query @@ -145,18 +145,14 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { c.connMutex.Unlock() if err := c.ensureConnection(); err != nil { - return nil, fmt.Errorf("dot: failed to create connection: %w", err) + return msg, nil, fmt.Errorf("dot: failed to create connection: %w", err) } } - // Prepare DNS message - if c.config.DNSSEC { - msg.SetEdns0(4096, true) - } packed, err := msg.Pack() if err != nil { logger.Error("DoT failed to pack message: %v", err) - return nil, fmt.Errorf("dot: failed to pack message: %w", err) + return msg, nil, fmt.Errorf("dot: failed to pack message: %w", err) } // Prepend message length (DNS over TCP format) @@ -171,7 +167,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { // Write query if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil { logger.Error("DoT failed to set write deadline: %v", err) - return nil, fmt.Errorf("dot: failed to set write deadline: %w", err) + return msg, nil, fmt.Errorf("dot: failed to set write deadline: %w", err) } if _, err := conn.Write(data); err != nil { @@ -181,7 +177,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { if c.config.KeepAlive { logger.Debug("DoT write failed with keep-alive, attempting reconnect") if reconnectErr := c.ensureConnection(); reconnectErr != nil { - return nil, fmt.Errorf("dot: failed to reconnect: %w", reconnectErr) + return msg, nil, fmt.Errorf("dot: failed to reconnect: %w", reconnectErr) } c.connMutex.Lock() @@ -189,48 +185,48 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { c.connMutex.Unlock() if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil { - return nil, fmt.Errorf("dot: failed to set write deadline after reconnect: %w", err) + return msg, nil, fmt.Errorf("dot: failed to set write deadline after reconnect: %w", err) } if _, err := conn.Write(data); err != nil { - return nil, fmt.Errorf("dot: failed to write message after reconnect: %w", err) + return msg, nil, fmt.Errorf("dot: failed to write message after reconnect: %w", err) } } else { - return nil, fmt.Errorf("dot: failed to write message: %w", err) + return msg, nil, fmt.Errorf("dot: failed to write message: %w", err) } } // Read response if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil { logger.Error("DoT failed to set read deadline: %v", err) - return nil, fmt.Errorf("dot: failed to set read deadline: %w", err) + return msg, nil, fmt.Errorf("dot: failed to set read deadline: %w", err) } // Read message length lengthBuf := make([]byte, 2) if _, err := io.ReadFull(conn, lengthBuf); err != nil { logger.Error("DoT failed to read response length from %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("dot: failed to read response length: %w", err) + return msg, nil, fmt.Errorf("dot: failed to read response length: %w", err) } msgLen := binary.BigEndian.Uint16(lengthBuf) if msgLen > dns.MaxMsgSize { logger.Error("DoT response too large from %s: %d bytes", c.hostAndPort, msgLen) - return nil, fmt.Errorf("dot: response message too large: %d", msgLen) + return msg, nil, fmt.Errorf("dot: response message too large: %d", msgLen) } // Read message body buffer := make([]byte, msgLen) if _, err := io.ReadFull(conn, buffer); err != nil { logger.Error("DoT failed to read response from %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("dot: failed to read response: %w", err) + return msg, nil, fmt.Errorf("dot: failed to read response: %w", err) } // Parse response response := new(dns.Msg) if err := response.Unpack(buffer); err != nil { logger.Error("DoT failed to unpack response from %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("dot: failed to unpack response: %w", err) + return msg, nil, fmt.Errorf("dot: failed to unpack response: %w", err) } if len(response.Answer) > 0 { @@ -247,5 +243,5 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { c.connMutex.Unlock() } - return response, nil + return msg, response, nil } diff --git a/common/protocols/dotcp/dotcp.go b/common/protocols/dotcp/dotcp.go index 0e130a7..f5e7888 100644 --- a/common/protocols/dotcp/dotcp.go +++ b/common/protocols/dotcp/dotcp.go @@ -104,7 +104,7 @@ func (c *Client) ensureConnection() error { return nil } -func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { +func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) { if len(msg.Question) > 0 { question := msg.Question[0] logger.Debug("DoTCP query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort) @@ -112,7 +112,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { if c.config.KeepAlive { if err := c.ensureConnection(); err != nil { - return nil, fmt.Errorf("dotcp: failed to ensure connection: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to ensure connection: %w", err) } } else { c.connMutex.Lock() @@ -123,18 +123,14 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { c.connMutex.Unlock() if err := c.ensureConnection(); err != nil { - return nil, fmt.Errorf("dotcp: failed to create connection: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to create connection: %w", err) } } - if c.config.DNSSEC { - msg.SetEdns0(4096, true) - } - packed, err := msg.Pack() if err != nil { logger.Error("DoTCP failed to pack message: %v", err) - return nil, fmt.Errorf("dotcp: failed to pack message: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to pack message: %w", err) } // DNS over TCP uses 2-byte length prefix @@ -148,7 +144,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil { logger.Error("DoTCP failed to set write deadline: %v", err) - return nil, fmt.Errorf("dotcp: failed to set write deadline: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to set write deadline: %w", err) } if _, err := conn.Write(data); err != nil { @@ -157,7 +153,7 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { if c.config.KeepAlive { logger.Debug("DoTCP write failed with keep-alive, attempting reconnect") if reconnectErr := c.ensureConnection(); reconnectErr != nil { - return nil, fmt.Errorf("dotcp: failed to reconnect: %w", reconnectErr) + return msg, nil, fmt.Errorf("dotcp: failed to reconnect: %w", reconnectErr) } c.connMutex.Lock() @@ -165,44 +161,44 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { c.connMutex.Unlock() if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil { - return nil, fmt.Errorf("dotcp: failed to set write deadline after reconnect: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to set write deadline after reconnect: %w", err) } if _, err := conn.Write(data); err != nil { - return nil, fmt.Errorf("dotcp: failed to write message after reconnect: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to write message after reconnect: %w", err) } } else { - return nil, fmt.Errorf("dotcp: failed to write message: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to write message: %w", err) } } if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil { logger.Error("DoTCP failed to set read deadline: %v", err) - return nil, fmt.Errorf("dotcp: failed to set read deadline: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to set read deadline: %w", err) } lengthBuf := make([]byte, 2) if _, err := io.ReadFull(conn, lengthBuf); err != nil { logger.Error("DoTCP failed to read response length from %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("dotcp: failed to read response length: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to read response length: %w", err) } msgLen := binary.BigEndian.Uint16(lengthBuf) if msgLen > dns.MaxMsgSize { logger.Error("DoTCP response too large from %s: %d bytes", c.hostAndPort, msgLen) - return nil, fmt.Errorf("dotcp: response message too large: %d", msgLen) + return msg, nil, fmt.Errorf("dotcp: response message too large: %d", msgLen) } buffer := make([]byte, msgLen) if _, err := io.ReadFull(conn, buffer); err != nil { logger.Error("DoTCP failed to read response from %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("dotcp: failed to read response: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to read response: %w", err) } response := new(dns.Msg) if err := response.Unpack(buffer); err != nil { logger.Error("DoTCP failed to unpack response from %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("dotcp: failed to unpack response: %w", err) + return msg, nil, fmt.Errorf("dotcp: failed to unpack response: %w", err) } if len(response.Answer) > 0 { @@ -218,5 +214,5 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { c.connMutex.Unlock() } - return response, nil + return msg, response, nil } diff --git a/common/protocols/doudp/doudp.go b/common/protocols/doudp/doudp.go index bded89d..f38e626 100644 --- a/common/protocols/doudp/doudp.go +++ b/common/protocols/doudp/doudp.go @@ -64,7 +64,7 @@ func (c *Client) createConnection() (*net.UDPConn, error) { return conn, nil } -func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { +func (c *Client) Query(msg *dns.Msg) (*dns.Msg, *dns.Msg, error) { if len(msg.Question) > 0 { question := msg.Question[0] logger.Debug("DoUDP query: %s %s to %s", question.Name, dns.TypeToString[question.Qtype], c.hostAndPort) @@ -72,56 +72,52 @@ func (c *Client) Query(msg *dns.Msg) (*dns.Msg, error) { conn, err := c.createConnection() if err != nil { - return nil, fmt.Errorf("doudp: failed to create connection: %w", err) + return msg, nil, fmt.Errorf("doudp: failed to create connection: %w", err) } defer conn.Close() - if c.config.DNSSEC { - msg.SetEdns0(4096, true) - } - packedMsg, err := msg.Pack() if err != nil { logger.Error("DoUDP failed to pack message: %v", err) - return nil, fmt.Errorf("doudp: failed to pack DNS message: %w", err) + return msg, nil, fmt.Errorf("doudp: failed to pack DNS message: %w", err) } if err := conn.SetWriteDeadline(time.Now().Add(c.config.WriteTimeout)); err != nil { logger.Error("DoUDP failed to set write deadline: %v", err) - return nil, fmt.Errorf("doudp: failed to set write deadline: %w", err) + return msg, nil, fmt.Errorf("doudp: failed to set write deadline: %w", err) } if _, err := conn.Write(packedMsg); err != nil { logger.Error("DoUDP failed to send query to %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("doudp: failed to send DNS query: %w", err) + return msg, nil, fmt.Errorf("doudp: failed to send DNS query: %w", err) } if err := conn.SetReadDeadline(time.Now().Add(c.config.ReadTimeout)); err != nil { logger.Error("DoUDP failed to set read deadline: %v", err) - return nil, fmt.Errorf("doudp: failed to set read deadline: %w", err) + return msg, nil, fmt.Errorf("doudp: failed to set read deadline: %w", err) } - var buffer []byte - if c.config.DNSSEC { - buffer = make([]byte, 4096) - }else { - buffer = make([]byte, 512) + bufSize := 512 + if opt := msg.IsEdns0(); opt != nil { + bufSize = max(int(opt.UDPSize()), 512) } + buffer := make([]byte, bufSize) + n, err := conn.Read(buffer) if err != nil { logger.Error("DoUDP failed to read response from %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("doudp: failed to read DNS response: %w", err) + return msg, nil, fmt.Errorf("doudp: failed to read DNS response: %w", err) } response := new(dns.Msg) if err := response.Unpack(buffer[:n]); err != nil { logger.Error("DoUDP failed to unpack response from %s: %v", c.hostAndPort, err) - return nil, fmt.Errorf("doudp: failed to unpack DNS response: %w", err) + return msg, nil, fmt.Errorf("doudp: failed to unpack DNS response: %w", err) } if len(response.Answer) > 0 { logger.Debug("DoUDP response from %s: %d answers", c.hostAndPort, len(response.Answer)) } - return response, nil + return msg, response, nil } diff --git a/go.mod b/go.mod index 3b6d09d..92c11af 100644 --- a/go.mod +++ b/go.mod @@ -9,6 +9,7 @@ require ( github.com/miekg/dns v1.1.72 github.com/quic-go/quic-go v0.60.0 golang.org/x/net v0.56.0 + golang.org/x/sys v0.46.0 ) require ( @@ -19,7 +20,6 @@ require ( golang.org/x/exp v0.0.0-20260611194520-c48552f49976 // indirect golang.org/x/mod v0.37.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/tools v0.47.0 // indirect ) diff --git a/internal/qol/measurement.go b/internal/qol/measurement.go index 225161b..175ceb8 100644 --- a/internal/qol/measurement.go +++ b/internal/qol/measurement.go @@ -10,6 +10,7 @@ import ( "time" "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/results" "github.com/afonsofrancof/sdns-proxy/internal/qol/stats" @@ -202,34 +203,44 @@ func (r *MeasurementRunner) performQuery(dnsClient client.DNSClient, domain, ups msg.RecursionDesired = true msg.SetQuestion(dns.Fqdn(domain), qType) - 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() metric.Timestamp = start - resp, err := dnsClient.Query(msg) + sent, resp, err := dnsClient.Query(msg) metric.Duration = time.Since(start).Nanoseconds() 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 { metric.ResponseCode = "ERROR" metric.Error = err.Error() return metric } - respBytes, err := resp.Pack() - if err != nil { + if b, perr := resp.Pack(); perr == nil { + metric.ResponseSize += len(b) + } else { metric.ResponseCode = "ERROR" - metric.Error = fmt.Sprintf("pack response: %v", err) + metric.Error = fmt.Sprintf("pack response: %v", perr) return metric } - metric.ResponseSize = len(respBytes) metric.ResponseCode = dns.RcodeToString[resp.Rcode] return metric } diff --git a/internal/qol/results/writer.go b/internal/qol/results/writer.go index 2e7ed81..3db6423 100644 --- a/internal/qol/results/writer.go +++ b/internal/qol/results/writer.go @@ -13,6 +13,8 @@ type DNSMetric struct { QueryType string `json:"query_type"` Protocol string `json:"protocol"` DNSSEC bool `json:"dnssec"` + DNSSECValidated bool `json:"dnssec_validated"` + DNSSECQueries int `json:"dnssec_queries"` AuthoritativeDNSSEC bool `json:"auth_dnssec"` KeepAlive bool `json:"keep_alive"` DNSServer string `json:"dns_server"` @@ -48,8 +50,8 @@ func NewMetricsWriter(path string) (*MetricsWriter, error) { // Only write header if file is new if !fileExists { header := []string{ - "domain", "query_type", "protocol", "dnssec", "auth_dnssec", "keep_alive", - "dns_server", "timestamp", "duration_ns", "duration_ms", + "domain", "query_type", "protocol", "dnssec", "dnssec_validated", "dnssec_queries", + "auth_dnssec", "keep_alive", "dns_server", "timestamp", "duration_ns", "duration_ms", "request_size_bytes", "response_size_bytes", "response_code", "error", } @@ -72,6 +74,8 @@ func (mw *MetricsWriter) WriteMetric(metric DNSMetric) error { metric.QueryType, metric.Protocol, strconv.FormatBool(metric.DNSSEC), + strconv.FormatBool(metric.DNSSECValidated), + strconv.Itoa(metric.DNSSECQueries), strconv.FormatBool(metric.AuthoritativeDNSSEC), strconv.FormatBool(metric.KeepAlive), metric.DNSServer, diff --git a/internal/qol/stats/runtime.go b/internal/qol/stats/runtime.go index cf0b78d..3ef93c7 100644 --- a/internal/qol/stats/runtime.go +++ b/internal/qol/stats/runtime.go @@ -3,6 +3,7 @@ package stats import ( "encoding/csv" "fmt" + "golang.org/x/sys/unix" "os" "runtime" "time" @@ -15,6 +16,7 @@ type RuntimeStats struct { AllocDelta uint64 MallocsDelta uint64 GCDelta uint32 + PeakRSSKB int64 } type RuntimeCollector struct { @@ -43,6 +45,7 @@ func (rc *RuntimeCollector) Collect() RuntimeStats { AllocDelta: current.TotalAlloc - rc.startStats.TotalAlloc, MallocsDelta: current.Mallocs - rc.startStats.Mallocs, GCDelta: current.NumGC - rc.startStats.NumGC, + PeakRSSKB: peakRSSKB(), } } @@ -69,7 +72,7 @@ func (rc *RuntimeCollector) WriteStats() error { if !fileExists { header := []string{ "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 { 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.MallocsDelta), fmt.Sprintf("%d", stats.GCDelta), + fmt.Sprintf("%d", stats.PeakRSSKB), } if err := writer.Write(row); err != nil { return fmt.Errorf("failed to write mem.csv row: %w", err) @@ -93,3 +97,11 @@ func (rc *RuntimeCollector) WriteStats() error { writer.Flush() 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) +} diff --git a/scripts/post_processing/merge_files.py b/scripts/post_processing/merge_files.py index a6f4294..e497a99 100644 --- a/scripts/post_processing/merge_files.py +++ b/scripts/post_processing/merge_files.py @@ -90,6 +90,8 @@ def merge_all_csvs(input_dir: Path, output_path: Path): 'provider', 'protocol', 'dnssec_mode', + 'dnssec_validated', + 'dnssec_queries', 'domain', 'query_type', 'keep_alive', @@ -135,6 +137,8 @@ def merge_all_csvs(input_dir: Path, output_path: Path): 'provider': provider, 'protocol': config['protocol'], 'dnssec_mode': config['dnssec_mode'], + 'dnssec_validated': row.get('dnssec_validated', ''), + 'dnssec_queries': row.get('dnssec_queries', ''), 'keep_alive': config['keep_alive'], 'domain': row.get('domain', ''), 'query_type': row.get('query_type', ''), diff --git a/scripts/post_processing/merge_mem.py b/scripts/post_processing/merge_mem.py index 9c6c195..71fd59c 100644 --- a/scripts/post_processing/merge_mem.py +++ b/scripts/post_processing/merge_mem.py @@ -52,7 +52,7 @@ def merge_mem_files(input_dir: Path, output_path: Path): output_columns = [ 'id','provider', 'protocol', 'dnssec_mode', 'keep_alive', '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 @@ -85,6 +85,7 @@ def merge_mem_files(input_dir: Path, output_path: Path): 'alloc_delta': row.get('alloc_delta', ''), 'mallocs_delta': row.get('mallocs_delta', ''), 'gc_delta': row.get('gc_delta', ''), + 'peak_rss_kb': row.get('peak_rss_kb', ''), } writer.writerow(out_row) diff --git a/server/server.go b/server/server.go deleted file mode 100644 index 5a1b6a3..0000000 --- a/server/server.go +++ /dev/null @@ -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") -}