tools(skills): add encrypted-dns-skill (ednsdiag CLI + agent skill)

This commit is contained in:
windyboy
2026-08-13 12:52:20 +08:00
parent 95ec2350af
commit c0cf82d4af
24 changed files with 2086 additions and 0 deletions
@@ -0,0 +1,229 @@
package edns
import (
"crypto/rand"
"encoding/base64"
"encoding/binary"
"fmt"
"net"
"strings"
"golang.org/x/net/dns/dnsmessage"
"golang.org/x/net/idna"
)
var recordTypes = map[string]dnsmessage.Type{
"A": dnsmessage.TypeA,
"AAAA": dnsmessage.TypeAAAA,
"CNAME": dnsmessage.TypeCNAME,
"MX": dnsmessage.TypeMX,
"TXT": dnsmessage.TypeTXT,
"NS": dnsmessage.TypeNS,
"SOA": dnsmessage.TypeSOA,
"CAA": dnsmessage.Type(257),
"SRV": dnsmessage.TypeSRV,
"SVCB": dnsmessage.TypeSVCB,
"HTTPS": dnsmessage.TypeHTTPS,
}
func BuildQuery(name, recordType string) ([]byte, QueryInfo, uint16, error) {
canonical, err := canonicalName(name)
if err != nil {
return nil, QueryInfo{}, 0, err
}
typeName := strings.ToUpper(recordType)
qtype, ok := recordTypes[typeName]
if !ok {
return nil, QueryInfo{}, 0, fmt.Errorf("unsupported record type %q", recordType)
}
dnsName, err := dnsmessage.NewName(canonical + ".")
if err != nil {
return nil, QueryInfo{}, 0, fmt.Errorf("encode domain name: %w", err)
}
var randomID [2]byte
if _, err := rand.Read(randomID[:]); err != nil {
return nil, QueryInfo{}, 0, fmt.Errorf("generate DNS transaction ID: %w", err)
}
id := binary.BigEndian.Uint16(randomID[:])
message := dnsmessage.Message{
Header: dnsmessage.Header{ID: id, RecursionDesired: true},
Questions: []dnsmessage.Question{{
Name: dnsName,
Type: qtype,
Class: dnsmessage.ClassINET,
}},
}
wire, err := message.Pack()
if err != nil {
return nil, QueryInfo{}, 0, fmt.Errorf("pack DNS query: %w", err)
}
return wire, QueryInfo{Name: canonical, Type: typeName}, id, nil
}
func ParseResponse(wire []byte, expectedID uint16, query QueryInfo) (DNSInfo, error) {
var message dnsmessage.Message
if err := message.Unpack(wire); err != nil {
return DNSInfo{}, fmt.Errorf("unpack DNS response: %w", err)
}
if !message.Header.Response {
return DNSInfo{}, fmt.Errorf("received a DNS query instead of a response")
}
if message.Header.ID != expectedID {
return DNSInfo{}, fmt.Errorf("DNS transaction ID mismatch")
}
if len(message.Questions) != 1 {
return DNSInfo{}, fmt.Errorf("DNS response contains %d questions, want 1", len(message.Questions))
}
wantType := recordTypes[query.Type]
question := message.Questions[0]
if trimRoot(question.Name.String()) != query.Name || question.Type != wantType {
return DNSInfo{}, fmt.Errorf("DNS response question does not match request")
}
answers := make([]AnswerRecord, 0, len(message.Answers))
for _, resource := range message.Answers {
answers = append(answers, normalizeAnswer(resource))
}
return DNSInfo{
RCode: rcodeName(message.Header.RCode),
RCodeValue: int(message.Header.RCode),
ResolverReportsDNSSECAuthenticated: message.Header.AuthenticData,
ClientValidatedDNSSEC: false,
Answers: answers,
}, nil
}
func canonicalName(input string) (string, error) {
name := strings.TrimSuffix(strings.TrimSpace(input), ".")
if name == "" {
return "", fmt.Errorf("domain name is empty")
}
if net.ParseIP(name) != nil {
return "", fmt.Errorf("IP literals are not accepted as domain names")
}
ascii, err := idna.Lookup.ToASCII(name)
if err != nil {
return "", fmt.Errorf("convert domain name to IDNA ASCII: %w", err)
}
ascii = strings.ToLower(ascii)
if len(ascii) > 253 {
return "", fmt.Errorf("domain name exceeds 253 bytes")
}
for _, label := range strings.Split(ascii, ".") {
if label == "" || len(label) > 63 {
return "", fmt.Errorf("domain name contains an invalid label")
}
}
blocked := []string{"localhost", ".local", ".internal", ".lan", ".arpa"}
for _, suffix := range blocked {
if ascii == strings.TrimPrefix(suffix, ".") || strings.HasSuffix(ascii, suffix) {
return "", fmt.Errorf("domain name is blocked by the local-name policy")
}
}
return ascii, nil
}
func normalizeAnswer(resource dnsmessage.Resource) AnswerRecord {
record := AnswerRecord{
"name": trimRoot(resource.Header.Name.String()),
"type": typeName(resource.Header.Type),
"ttl": resource.Header.TTL,
}
switch body := resource.Body.(type) {
case *dnsmessage.AResource:
record["address"] = net.IP(body.A[:]).String()
case *dnsmessage.AAAAResource:
record["address"] = net.IP(body.AAAA[:]).String()
case *dnsmessage.CNAMEResource:
record["target"] = trimRoot(body.CNAME.String())
case *dnsmessage.MXResource:
record["priority"] = body.Pref
record["exchange"] = trimRoot(body.MX.String())
case *dnsmessage.TXTResource:
record["strings"] = body.TXT
case *dnsmessage.NSResource:
record["host"] = trimRoot(body.NS.String())
case *dnsmessage.PTRResource:
record["target"] = trimRoot(body.PTR.String())
case *dnsmessage.SOAResource:
record["primary_ns"] = trimRoot(body.NS.String())
record["responsible_mailbox"] = trimRoot(body.MBox.String())
record["serial"] = body.Serial
record["refresh"] = body.Refresh
record["retry"] = body.Retry
record["expire"] = body.Expire
record["minimum_ttl"] = body.MinTTL
case *dnsmessage.SRVResource:
record["priority"] = body.Priority
record["weight"] = body.Weight
record["port"] = body.Port
record["target"] = trimRoot(body.Target.String())
case *dnsmessage.SVCBResource:
addSVCBFields(record, body.Priority, body.Target, body.Params)
case *dnsmessage.HTTPSResource:
addSVCBFields(record, body.Priority, body.Target, body.Params)
case *dnsmessage.UnknownResource:
if resource.Header.Type == dnsmessage.Type(257) && len(body.Data) >= 2 {
record["flags"] = body.Data[0]
tagLength := int(body.Data[1])
if 2+tagLength <= len(body.Data) {
record["tag"] = string(body.Data[2 : 2+tagLength])
record["value"] = string(body.Data[2+tagLength:])
} else {
record["rdata_base64"] = base64.StdEncoding.EncodeToString(body.Data)
}
} else {
record["rdata_base64"] = base64.StdEncoding.EncodeToString(body.Data)
}
}
return record
}
func addSVCBFields(record AnswerRecord, priority uint16, target dnsmessage.Name, params []dnsmessage.SVCParam) {
record["priority"] = priority
record["target"] = trimRoot(target.String())
values := make([]map[string]any, 0, len(params))
for _, param := range params {
values = append(values, map[string]any{
"key": param.Key.String(),
"key_value": uint16(param.Key),
"value_base64": base64.StdEncoding.EncodeToString(param.Value),
})
}
record["params"] = values
}
func trimRoot(name string) string {
return strings.TrimSuffix(strings.ToLower(name), ".")
}
func typeName(recordType dnsmessage.Type) string {
for name, value := range recordTypes {
if value == recordType {
return name
}
}
return fmt.Sprintf("TYPE%d", recordType)
}
func rcodeName(rcode dnsmessage.RCode) string {
names := map[dnsmessage.RCode]string{
dnsmessage.RCodeSuccess: "NOERROR",
dnsmessage.RCodeFormatError: "FORMERR",
dnsmessage.RCodeServerFailure: "SERVFAIL",
dnsmessage.RCodeNameError: "NXDOMAIN",
dnsmessage.RCodeNotImplemented: "NOTIMP",
dnsmessage.RCodeRefused: "REFUSED",
}
if name, ok := names[rcode]; ok {
return name
}
return fmt.Sprintf("RCODE%d", rcode)
}
@@ -0,0 +1,100 @@
package edns
import (
"encoding/binary"
"testing"
"golang.org/x/net/dns/dnsmessage"
)
func TestBuildAndParseResponse(t *testing.T) {
queryWire, query, transactionID, err := BuildQuery("Example.COM.", "A")
if err != nil {
t.Fatalf("build query: %v", err)
}
if query.Name != "example.com" || query.Type != "A" {
t.Fatalf("canonical query = %#v", query)
}
var request dnsmessage.Message
if err := request.Unpack(queryWire); err != nil {
t.Fatalf("unpack query: %v", err)
}
response := dnsmessage.Message{
Header: dnsmessage.Header{
ID: transactionID,
Response: true,
RecursionDesired: true,
RecursionAvailable: true,
AuthenticData: true,
},
Questions: request.Questions,
Answers: []dnsmessage.Resource{{
Header: dnsmessage.ResourceHeader{Name: request.Questions[0].Name, Class: dnsmessage.ClassINET, TTL: 60},
Body: &dnsmessage.AResource{A: [4]byte{192, 0, 2, 1}},
}},
}
responseWire, err := response.Pack()
if err != nil {
t.Fatalf("pack response: %v", err)
}
dnsResult, err := ParseResponse(responseWire, transactionID, query)
if err != nil {
t.Fatalf("parse response: %v", err)
}
if dnsResult.RCode != "NOERROR" || !dnsResult.ResolverReportsDNSSECAuthenticated {
t.Fatalf("unexpected DNS result: %#v", dnsResult)
}
if got := dnsResult.Answers[0]["address"]; got != "192.0.2.1" {
t.Fatalf("address = %v, want 192.0.2.1", got)
}
}
func TestBuildQueryIDNAAndBlockedNames(t *testing.T) {
_, query, _, err := BuildQuery("bücher.example", "AAAA")
if err != nil {
t.Fatalf("build IDNA query: %v", err)
}
if query.Name != "xn--bcher-kva.example" {
t.Fatalf("IDNA name = %q", query.Name)
}
blocked := []string{"localhost", "router.local", "service.internal", "host.lan", "1.0.0.127.in-addr.arpa", "127.0.0.1"}
for _, name := range blocked {
if _, _, _, err := BuildQuery(name, "A"); err == nil {
t.Errorf("BuildQuery(%q) succeeded, want policy error", name)
}
}
}
func TestNormalizeCAA(t *testing.T) {
name := dnsmessage.MustNewName("example.com.")
data := append([]byte{0, 5}, []byte("issueletsencrypt.org")...)
record := normalizeAnswer(dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{Name: name, Type: dnsmessage.Type(257), Class: dnsmessage.ClassINET, TTL: 300},
Body: &dnsmessage.UnknownResource{Type: dnsmessage.Type(257), Data: data},
})
if record["tag"] != "issue" || record["value"] != "letsencrypt.org" {
t.Fatalf("unexpected CAA normalization: %#v", record)
}
}
func TestParseResponseRejectsTransactionMismatch(t *testing.T) {
name := dnsmessage.MustNewName("example.com.")
message := dnsmessage.Message{
Header: dnsmessage.Header{ID: 2, Response: true},
Questions: []dnsmessage.Question{{Name: name, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET}},
}
wire, err := message.Pack()
if err != nil {
t.Fatalf("pack response: %v", err)
}
if _, err := ParseResponse(wire, 1, QueryInfo{Name: "example.com", Type: "A"}); err == nil {
t.Fatal("transaction mismatch was accepted")
}
if binary.BigEndian.Uint16(wire[:2]) != 2 {
t.Fatal("test response ID was not encoded")
}
}
@@ -0,0 +1,131 @@
package edns
import (
"bytes"
"context"
"crypto/tls"
"encoding/base64"
"fmt"
"io"
"mime"
"net"
"net/http"
"net/url"
"strings"
"time"
)
const maxDNSMessageSize = 65535
func exchangeDoH(ctx context.Context, provider Provider, wire []byte, method string) ([]byte, TransportInfo, error) {
client := newDoHClient(provider.DoHURL)
return exchangeDoHWithClient(ctx, client, provider.DoHURL, wire, method)
}
func newDoHClient(endpoint string) *http.Client {
origin, _ := url.Parse(endpoint)
transport := &http.Transport{
ForceAttemptHTTP2: true,
DialContext: (&net.Dialer{Timeout: 5 * time.Second, KeepAlive: 30 * time.Second}).DialContext,
TLSClientConfig: &tls.Config{
MinVersion: tls.VersionTLS12,
},
TLSHandshakeTimeout: 5 * time.Second,
}
return &http.Client{
Transport: transport,
CheckRedirect: func(request *http.Request, via []*http.Request) error {
if len(via) >= 3 {
return fmt.Errorf("too many DoH redirects")
}
if request.URL.Scheme != "https" {
return fmt.Errorf("DoH redirect changed to a non-HTTPS scheme")
}
if !strings.EqualFold(request.URL.Hostname(), origin.Hostname()) {
return fmt.Errorf("DoH redirect changed authentication domain")
}
return nil
},
}
}
func exchangeDoHWithClient(ctx context.Context, client *http.Client, endpoint string, wire []byte, method string) ([]byte, TransportInfo, error) {
started := time.Now()
info := TransportInfo{
Protocol: "doh",
Encrypted: true,
Bootstrap: "system_resolver",
}
requestURL := endpoint
var body io.Reader
switch strings.ToLower(method) {
case "get":
parsed, err := url.Parse(endpoint)
if err != nil {
return nil, info, fmt.Errorf("parse DoH endpoint: %w", err)
}
query := parsed.Query()
query.Set("dns", base64.RawURLEncoding.EncodeToString(wire))
parsed.RawQuery = query.Encode()
requestURL = parsed.String()
case "post", "":
method = "post"
body = bytes.NewReader(wire)
default:
return nil, info, fmt.Errorf("unsupported DoH method %q", method)
}
request, err := http.NewRequestWithContext(ctx, strings.ToUpper(method), requestURL, body)
if err != nil {
return nil, info, fmt.Errorf("create DoH request: %w", err)
}
request.Header.Set("Accept", "application/dns-message")
if strings.EqualFold(method, "post") {
request.Header.Set("Content-Type", "application/dns-message")
}
request.Header.Set("User-Agent", "ednsdiag/0.1.0-dev")
response, err := client.Do(request)
info.ElapsedMS = time.Since(started).Milliseconds()
if err != nil {
return nil, info, fmt.Errorf("perform DoH exchange: %w", err)
}
defer response.Body.Close()
info.HTTPVersion = response.Proto
if response.TLS == nil || len(response.TLS.VerifiedChains) == 0 {
return nil, info, fmt.Errorf("DoH server TLS identity was not verified")
}
info.ServerAuthenticated = true
info.TLSVersion = tlsVersionName(response.TLS.Version)
info.ALPN = response.TLS.NegotiatedProtocol
if response.StatusCode < 200 || response.StatusCode > 299 {
return nil, info, fmt.Errorf("DoH server returned HTTP status %d", response.StatusCode)
}
mediaType, _, err := mime.ParseMediaType(response.Header.Get("Content-Type"))
if err != nil || !strings.EqualFold(mediaType, "application/dns-message") {
return nil, info, fmt.Errorf("DoH server returned unsupported content type %q", response.Header.Get("Content-Type"))
}
payload, err := io.ReadAll(io.LimitReader(response.Body, maxDNSMessageSize+1))
if err != nil {
return nil, info, fmt.Errorf("read DoH response: %w", err)
}
if len(payload) > maxDNSMessageSize {
return nil, info, fmt.Errorf("DoH response exceeds %d bytes", maxDNSMessageSize)
}
return payload, info, nil
}
func tlsVersionName(version uint16) string {
switch version {
case tls.VersionTLS13:
return "TLS1.3"
case tls.VersionTLS12:
return "TLS1.2"
default:
return fmt.Sprintf("0x%04x", version)
}
}
@@ -0,0 +1,75 @@
package edns
import (
"io"
"net/http"
"net/http/httptest"
"testing"
"golang.org/x/net/dns/dnsmessage"
)
func TestExchangeDoHGETAndPOST(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
var payload []byte
var err error
if request.Method == http.MethodGet {
payload, err = decodeGETQuery(request.URL.Query().Get("dns"))
} else {
payload, err = io.ReadAll(request.Body)
}
if err != nil {
http.Error(writer, err.Error(), http.StatusBadRequest)
return
}
if request.Header.Get("Accept") != "application/dns-message" {
http.Error(writer, "missing accept", http.StatusNotAcceptable)
return
}
var query dnsmessage.Message
if err := query.Unpack(payload); err != nil {
http.Error(writer, err.Error(), http.StatusBadRequest)
return
}
response := dnsmessage.Message{
Header: dnsmessage.Header{ID: query.Header.ID, Response: true, RecursionAvailable: true},
Questions: query.Questions,
}
responseWire, err := response.Pack()
if err != nil {
http.Error(writer, err.Error(), http.StatusInternalServerError)
return
}
writer.Header().Set("Content-Type", "application/dns-message")
_, _ = writer.Write(responseWire)
}))
defer server.Close()
wire, _, _, err := BuildQuery("example.com", "A")
if err != nil {
t.Fatalf("build query: %v", err)
}
for _, method := range []string{"get", "post"} {
t.Run(method, func(t *testing.T) {
response, info, err := exchangeDoHWithClient(t.Context(), server.Client(), server.URL, wire, method)
if err != nil {
t.Fatalf("exchange DoH: %v", err)
}
if len(response) == 0 || !info.Encrypted || !info.ServerAuthenticated {
t.Fatalf("unexpected result: response=%d info=%#v", len(response), info)
}
})
}
}
func TestExchangeDoHRejectsHTTPError(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
http.Error(writer, "unavailable", http.StatusServiceUnavailable)
}))
defer server.Close()
if _, _, err := exchangeDoHWithClient(t.Context(), server.Client(), server.URL, []byte{1}, "post"); err == nil {
t.Fatal("HTTP error was accepted")
}
}
@@ -0,0 +1,106 @@
package edns
import (
"context"
"crypto/tls"
"encoding/binary"
"fmt"
"io"
"net"
"time"
)
func exchangeDoT(ctx context.Context, provider Provider, wire []byte) ([]byte, TransportInfo, error) {
return exchangeDoTWithTLSConfig(ctx, provider, wire, &tls.Config{
ServerName: provider.DoTName,
MinVersion: tls.VersionTLS12,
NextProtos: []string{"dot"},
})
}
func exchangeDoTWithTLSConfig(ctx context.Context, provider Provider, wire []byte, tlsConfig *tls.Config) ([]byte, TransportInfo, error) {
started := time.Now()
info := TransportInfo{
Protocol: "dot",
Encrypted: true,
Bootstrap: "system_resolver",
}
rawConnection, err := (&net.Dialer{}).DialContext(ctx, "tcp", provider.DoTAddr)
if err != nil {
info.ElapsedMS = time.Since(started).Milliseconds()
return nil, info, fmt.Errorf("connect to DoT server: %w", err)
}
defer rawConnection.Close()
if deadline, ok := ctx.Deadline(); ok {
if err := rawConnection.SetDeadline(deadline); err != nil {
return nil, info, fmt.Errorf("set DoT deadline: %w", err)
}
}
tlsConfig = tlsConfig.Clone()
tlsConfig.ServerName = provider.DoTName
tlsConnection := tls.Client(rawConnection, tlsConfig)
if err := tlsConnection.HandshakeContext(ctx); err != nil {
info.ElapsedMS = time.Since(started).Milliseconds()
return nil, info, fmt.Errorf("authenticate DoT server: %w", err)
}
state := tlsConnection.ConnectionState()
if len(state.VerifiedChains) == 0 {
return nil, info, fmt.Errorf("DoT server TLS identity was not verified")
}
info.ServerAuthenticated = true
info.TLSVersion = tlsVersionName(state.Version)
info.ALPN = state.NegotiatedProtocol
if info.ALPN != "" && info.ALPN != "dot" {
info.ElapsedMS = time.Since(started).Milliseconds()
return nil, info, fmt.Errorf("DoT server negotiated unexpected ALPN protocol %q", info.ALPN)
}
response, err := exchangeTCPFrame(tlsConnection, wire)
info.ElapsedMS = time.Since(started).Milliseconds()
if err != nil {
return nil, info, fmt.Errorf("perform DoT exchange: %w", err)
}
return response, info, nil
}
func exchangeTCPFrame(connection io.ReadWriter, wire []byte) ([]byte, error) {
if len(wire) == 0 || len(wire) > maxDNSMessageSize {
return nil, fmt.Errorf("invalid DNS message length %d", len(wire))
}
frame := make([]byte, 2+len(wire))
binary.BigEndian.PutUint16(frame[:2], uint16(len(wire)))
copy(frame[2:], wire)
if err := writeAll(connection, frame); err != nil {
return nil, fmt.Errorf("write framed DNS query: %w", err)
}
var lengthBytes [2]byte
if _, err := io.ReadFull(connection, lengthBytes[:]); err != nil {
return nil, fmt.Errorf("read DNS response length: %w", err)
}
length := int(binary.BigEndian.Uint16(lengthBytes[:]))
if length == 0 {
return nil, fmt.Errorf("DoT server returned an empty DNS message")
}
response := make([]byte, length)
if _, err := io.ReadFull(connection, response); err != nil {
return nil, fmt.Errorf("read DNS response: %w", err)
}
return response, nil
}
func writeAll(writer io.Writer, payload []byte) error {
for len(payload) > 0 {
written, err := writer.Write(payload)
if err != nil {
return err
}
if written == 0 {
return io.ErrShortWrite
}
payload = payload[written:]
}
return nil
}
@@ -0,0 +1,254 @@
package edns
import (
"bytes"
"context"
"crypto/ed25519"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"encoding/binary"
"io"
"math/big"
"net"
"strings"
"testing"
"time"
"golang.org/x/net/dns/dnsmessage"
)
type scriptedReadWriter struct {
read *bytes.Reader
written bytes.Buffer
}
func (stream *scriptedReadWriter) Read(payload []byte) (int, error) {
return stream.read.Read(payload)
}
func (stream *scriptedReadWriter) Write(payload []byte) (int, error) {
return stream.written.Write(payload)
}
func TestExchangeTCPFrame(t *testing.T) {
responsePayload := []byte{9, 8, 7}
framedResponse := make([]byte, 2+len(responsePayload))
binary.BigEndian.PutUint16(framedResponse[:2], uint16(len(responsePayload)))
copy(framedResponse[2:], responsePayload)
stream := &scriptedReadWriter{read: bytes.NewReader(framedResponse)}
query := []byte{1, 2, 3, 4}
response, err := exchangeTCPFrame(stream, query)
if err != nil {
t.Fatalf("exchange TCP frame: %v", err)
}
if !bytes.Equal(response, responsePayload) {
t.Fatalf("response = %v, want %v", response, responsePayload)
}
written := stream.written.Bytes()
if int(binary.BigEndian.Uint16(written[:2])) != len(query) || !bytes.Equal(written[2:], query) {
t.Fatalf("invalid query frame: %v", written)
}
}
func TestExchangeDoTAuthenticatesServer(t *testing.T) {
certificate, roots := newTestCertificate(t, "resolver.test")
listener, err := tls.Listen("tcp", "127.0.0.1:0", &tls.Config{
Certificates: []tls.Certificate{certificate},
MinVersion: tls.VersionTLS12,
NextProtos: []string{"dot"},
})
if err != nil {
t.Fatalf("listen for DoT: %v", err)
}
defer listener.Close()
serverError := make(chan error, 1)
go func() {
connection, err := listener.Accept()
if err != nil {
serverError <- err
return
}
defer connection.Close()
response, err := serveOneDoTQuery(connection)
if err == nil {
err = writeAll(connection, response)
}
serverError <- err
}()
queryWire, query, transactionID, err := BuildQuery("example.com", "A")
if err != nil {
t.Fatalf("build query: %v", err)
}
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
response, info, err := exchangeDoTWithTLSConfig(ctx, Provider{DoTAddr: listener.Addr().String(), DoTName: "resolver.test"}, queryWire, &tls.Config{
RootCAs: roots,
MinVersion: tls.VersionTLS12,
NextProtos: []string{"dot"},
})
if err != nil {
t.Fatalf("exchange DoT: %v", err)
}
if err := <-serverError; err != nil {
t.Fatalf("serve DoT: %v", err)
}
if !info.ServerAuthenticated || info.ALPN != "dot" {
t.Fatalf("unexpected transport info: %#v", info)
}
if _, err := ParseResponse(response, transactionID, query); err != nil {
t.Fatalf("parse response: %v", err)
}
}
func TestExchangeDoTAllowsMissingALPN(t *testing.T) {
certificate, roots := newTestCertificate(t, "resolver.test")
listener, err := tls.Listen("tcp", "127.0.0.1:0", &tls.Config{
Certificates: []tls.Certificate{certificate},
MinVersion: tls.VersionTLS12,
})
if err != nil {
t.Fatalf("listen for DoT: %v", err)
}
defer listener.Close()
serverError := make(chan error, 1)
go func() {
connection, err := listener.Accept()
if err != nil {
serverError <- err
return
}
defer connection.Close()
response, err := serveOneDoTQuery(connection)
if err == nil {
err = writeAll(connection, response)
}
serverError <- err
}()
queryWire, _, _, err := BuildQuery("example.com", "A")
if err != nil {
t.Fatalf("build query: %v", err)
}
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
_, info, err := exchangeDoTWithTLSConfig(ctx, Provider{DoTAddr: listener.Addr().String(), DoTName: "resolver.test"}, queryWire, &tls.Config{
RootCAs: roots,
MinVersion: tls.VersionTLS12,
NextProtos: []string{"dot"},
})
if err != nil {
t.Fatalf("exchange DoT without server ALPN: %v", err)
}
if err := <-serverError; err != nil {
t.Fatalf("serve DoT: %v", err)
}
if !info.ServerAuthenticated || info.ALPN != "" {
t.Fatalf("unexpected transport info: %#v", info)
}
}
func TestExchangeDoTRejectsUnexpectedALPN(t *testing.T) {
certificate, roots := newTestCertificate(t, "resolver.test")
listener, err := tls.Listen("tcp", "127.0.0.1:0", &tls.Config{
Certificates: []tls.Certificate{certificate},
MinVersion: tls.VersionTLS12,
NextProtos: []string{"http/1.1"},
})
if err != nil {
t.Fatalf("listen for TLS: %v", err)
}
defer listener.Close()
serverError := make(chan error, 1)
go func() {
connection, err := listener.Accept()
if err != nil {
serverError <- err
return
}
defer connection.Close()
serverError <- connection.(*tls.Conn).Handshake()
}()
queryWire, _, _, err := BuildQuery("example.com", "A")
if err != nil {
t.Fatalf("build query: %v", err)
}
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
_, info, err := exchangeDoTWithTLSConfig(ctx, Provider{DoTAddr: listener.Addr().String(), DoTName: "resolver.test"}, queryWire, &tls.Config{
RootCAs: roots,
MinVersion: tls.VersionTLS12,
NextProtos: []string{"http/1.1"},
})
if err == nil || !strings.Contains(err.Error(), "unexpected ALPN protocol") {
t.Fatalf("exchange DoT error = %v, want unexpected ALPN error", err)
}
if err := <-serverError; err != nil {
t.Fatalf("complete TLS handshake: %v", err)
}
if !info.ServerAuthenticated || info.ALPN != "http/1.1" {
t.Fatalf("unexpected transport info: %#v", info)
}
}
func newTestCertificate(t *testing.T, name string) (tls.Certificate, *x509.CertPool) {
t.Helper()
publicKey, privateKey, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
t.Fatalf("generate key: %v", err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(1),
DNSNames: []string{name},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
IsCA: true,
BasicConstraintsValid: true,
}
der, err := x509.CreateCertificate(rand.Reader, template, template, publicKey, privateKey)
if err != nil {
t.Fatalf("create certificate: %v", err)
}
parsed, err := x509.ParseCertificate(der)
if err != nil {
t.Fatalf("parse certificate: %v", err)
}
roots := x509.NewCertPool()
roots.AddCert(parsed)
return tls.Certificate{Certificate: [][]byte{der}, PrivateKey: privateKey}, roots
}
func serveOneDoTQuery(connection net.Conn) ([]byte, error) {
var lengthBytes [2]byte
if _, err := io.ReadFull(connection, lengthBytes[:]); err != nil {
return nil, err
}
wire := make([]byte, int(binary.BigEndian.Uint16(lengthBytes[:])))
if _, err := io.ReadFull(connection, wire); err != nil {
return nil, err
}
var query dnsmessage.Message
if err := query.Unpack(wire); err != nil {
return nil, err
}
response := dnsmessage.Message{
Header: dnsmessage.Header{ID: query.Header.ID, Response: true, RecursionAvailable: true},
Questions: query.Questions,
}
responseWire, err := response.Pack()
if err != nil {
return nil, err
}
framed := make([]byte, 2+len(responseWire))
binary.BigEndian.PutUint16(framed[:2], uint16(len(responseWire)))
copy(framed[2:], responseWire)
return framed, nil
}
@@ -0,0 +1,58 @@
package edns
type QueryOptions struct {
Name string
RecordType string
Protocol string
Provider string
Method string
}
type Result struct {
SchemaVersion int `json:"schema_version"`
Operation string `json:"operation"`
Completed bool `json:"completed"`
Query QueryInfo `json:"query"`
Resolver ResolverInfo `json:"resolver"`
Transport TransportInfo `json:"transport"`
DNS DNSInfo `json:"dns"`
Warnings []string `json:"warnings,omitempty"`
Error *ErrorInfo `json:"error,omitempty"`
}
type QueryInfo struct {
Name string `json:"name"`
Type string `json:"type"`
}
type ResolverInfo struct {
Provider string `json:"provider"`
Endpoint string `json:"endpoint"`
Profile string `json:"profile"`
}
type TransportInfo struct {
Protocol string `json:"protocol"`
Encrypted bool `json:"encrypted"`
ServerAuthenticated bool `json:"server_authenticated"`
ElapsedMS int64 `json:"elapsed_ms"`
Bootstrap string `json:"bootstrap"`
TLSVersion string `json:"tls_version,omitempty"`
ALPN string `json:"alpn,omitempty"`
HTTPVersion string `json:"http_version,omitempty"`
}
type DNSInfo struct {
RCode string `json:"rcode"`
RCodeValue int `json:"rcode_value"`
ResolverReportsDNSSECAuthenticated bool `json:"resolver_reports_dnssec_authenticated"`
ClientValidatedDNSSEC bool `json:"client_validated_dnssec"`
Answers []AnswerRecord `json:"answers"`
}
type AnswerRecord map[string]any
type ErrorInfo struct {
Class string `json:"class"`
Message string `json:"message"`
}
@@ -0,0 +1,53 @@
package edns
import (
"fmt"
"strings"
)
type Provider struct {
ID string
Profile string
DoHURL string
DoTAddr string
DoTName string
}
var providers = map[string]Provider{
"cloudflare": {
ID: "cloudflare",
Profile: "unfiltered",
DoHURL: "https://cloudflare-dns.com/dns-query",
DoTAddr: "one.one.one.one:853",
DoTName: "one.one.one.one",
},
"google": {
ID: "google",
Profile: "unfiltered",
DoHURL: "https://dns.google/dns-query",
DoTAddr: "dns.google:853",
DoTName: "dns.google",
},
"quad9": {
ID: "quad9",
Profile: "security-filtered",
DoHURL: "https://dns.quad9.net/dns-query",
DoTAddr: "dns.quad9.net:853",
DoTName: "dns.quad9.net",
},
"adguard": {
ID: "adguard",
Profile: "ad-and-security-filtered",
DoHURL: "https://dns.adguard-dns.com/dns-query",
DoTAddr: "dns.adguard-dns.com:853",
DoTName: "dns.adguard-dns.com",
},
}
func FindProvider(name string) (Provider, error) {
provider, ok := providers[strings.ToLower(name)]
if !ok {
return Provider{}, fmt.Errorf("unknown provider %q", name)
}
return provider, nil
}
@@ -0,0 +1,18 @@
package edns
import "testing"
func TestBuiltInProvidersHaveStrictEndpoints(t *testing.T) {
for _, name := range []string{"cloudflare", "google", "quad9", "adguard"} {
provider, err := FindProvider(name)
if err != nil {
t.Fatalf("find provider %s: %v", name, err)
}
if provider.DoHURL == "" || provider.DoTAddr == "" || provider.DoTName == "" {
t.Fatalf("provider %s is incomplete: %#v", name, provider)
}
}
if _, err := FindProvider("custom"); err == nil {
t.Fatal("unapproved custom provider was accepted")
}
}
@@ -0,0 +1,57 @@
package edns
import (
"context"
"fmt"
)
func Query(ctx context.Context, options QueryOptions) Result {
wire, query, transactionID, err := BuildQuery(options.Name, options.RecordType)
result := Result{
SchemaVersion: 1,
Operation: "query",
Query: query,
Transport: TransportInfo{
Protocol: options.Protocol,
Encrypted: true,
Bootstrap: "system_resolver",
},
DNS: DNSInfo{Answers: []AnswerRecord{}},
}
if err != nil {
result.Query = QueryInfo{Name: options.Name, Type: options.RecordType}
result.Error = &ErrorInfo{Class: "input", Message: err.Error()}
return result
}
provider, err := FindProvider(options.Provider)
if err != nil {
result.Error = &ErrorInfo{Class: "input", Message: err.Error()}
return result
}
result.Resolver = ResolverInfo{Provider: provider.ID, Profile: provider.Profile}
var response []byte
switch options.Protocol {
case "doh":
result.Resolver.Endpoint = provider.DoHURL
response, result.Transport, err = exchangeDoH(ctx, provider, wire, options.Method)
case "dot":
result.Resolver.Endpoint = provider.DoTAddr
response, result.Transport, err = exchangeDoT(ctx, provider, wire)
default:
err = fmt.Errorf("protocol %q is not available; run ednsdiag capabilities", options.Protocol)
}
if err != nil {
result.Error = &ErrorInfo{Class: "transport", Message: err.Error()}
return result
}
result.DNS, err = ParseResponse(response, transactionID, query)
if err != nil {
result.Error = &ErrorInfo{Class: "protocol", Message: err.Error()}
return result
}
result.Completed = true
return result
}
@@ -0,0 +1,7 @@
package edns
import "encoding/base64"
func decodeGETQuery(value string) ([]byte, error) {
return base64.RawURLEncoding.DecodeString(value)
}