tools(skills): add encrypted-dns-skill (ednsdiag CLI + agent skill)
This commit is contained in:
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user