166 lines
4.5 KiB
Go
166 lines
4.5 KiB
Go
//go:build darwin
|
|
|
|
package local
|
|
|
|
import (
|
|
"cmp"
|
|
"context"
|
|
"net"
|
|
"os"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
mDNS "github.com/miekg/dns"
|
|
)
|
|
|
|
// "localhost" is answered by the mDNSResponder daemon itself, so these tests need
|
|
// no external network.
|
|
|
|
func requireMDNSResponder(t *testing.T) {
|
|
t.Helper()
|
|
socketPath := cmp.Or(os.Getenv(mdnsResponderSocketEnv), mdnsResponderSocketPath)
|
|
conn, err := net.DialTimeout("unix", socketPath, time.Second)
|
|
if err != nil {
|
|
t.Skipf("mDNSResponder not reachable at %s: %v", socketPath, err)
|
|
}
|
|
conn.Close()
|
|
}
|
|
|
|
func systemExchangeForTest(ctx context.Context, transport *Transport, message *mDNS.Msg) (*mDNS.Msg, error) {
|
|
done := make(chan struct{})
|
|
var (
|
|
response *mDNS.Msg
|
|
err error
|
|
)
|
|
transport.systemExchangeAsync(ctx, message, func(callbackResponse *mDNS.Msg, callbackErr error) {
|
|
response = callbackResponse
|
|
err = callbackErr
|
|
close(done)
|
|
})
|
|
<-done
|
|
return response, err
|
|
}
|
|
|
|
func TestSystemExchangeLoopback(t *testing.T) {
|
|
requireMDNSResponder(t)
|
|
transport := &Transport{}
|
|
defer transport.system.close()
|
|
for _, testCase := range []struct {
|
|
qtype uint16
|
|
expected net.IP
|
|
}{
|
|
{mDNS.TypeA, net.IPv4(127, 0, 0, 1)},
|
|
{mDNS.TypeAAAA, net.IPv6loopback},
|
|
} {
|
|
message := new(mDNS.Msg)
|
|
message.SetQuestion("localhost.", testCase.qtype)
|
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
response, err := systemExchangeForTest(ctx, transport, message)
|
|
cancel()
|
|
if err != nil {
|
|
t.Fatalf("%s localhost: %v", mDNS.TypeToString[testCase.qtype], err)
|
|
}
|
|
if response.Id != message.Id {
|
|
t.Fatalf("%s response id %d != request id %d", mDNS.TypeToString[testCase.qtype], response.Id, message.Id)
|
|
}
|
|
if !response.Response {
|
|
t.Fatalf("%s response flag not set", mDNS.TypeToString[testCase.qtype])
|
|
}
|
|
var found bool
|
|
for _, answer := range response.Answer {
|
|
switch record := answer.(type) {
|
|
case *mDNS.A:
|
|
found = found || record.A.Equal(testCase.expected)
|
|
case *mDNS.AAAA:
|
|
found = found || record.AAAA.Equal(testCase.expected)
|
|
}
|
|
}
|
|
if !found {
|
|
t.Fatalf("%s localhost: expected %s in answer, got %v", mDNS.TypeToString[testCase.qtype], testCase.expected, response.Answer)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestSystemExchangeNoData(t *testing.T) {
|
|
requireMDNSResponder(t)
|
|
transport := &Transport{}
|
|
defer transport.system.close()
|
|
message := new(mDNS.Msg)
|
|
// localhost has no MX record, so the daemon reports NoSuchRecord, which must
|
|
// surface as an empty NOERROR response rather than an error.
|
|
message.SetQuestion("localhost.", mDNS.TypeMX)
|
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
defer cancel()
|
|
response, err := systemExchangeForTest(ctx, transport, message)
|
|
if err != nil {
|
|
t.Fatalf("MX localhost: %v", err)
|
|
}
|
|
if response.Rcode != mDNS.RcodeSuccess {
|
|
t.Fatalf("MX localhost: rcode %s, want NOERROR", mDNS.RcodeToString[response.Rcode])
|
|
}
|
|
if len(response.Answer) != 0 {
|
|
t.Fatalf("MX localhost: expected no answers, got %v", response.Answer)
|
|
}
|
|
}
|
|
|
|
func TestSystemExchangeCancel(t *testing.T) {
|
|
requireMDNSResponder(t)
|
|
transport := &Transport{}
|
|
defer transport.system.close()
|
|
message := new(mDNS.Msg)
|
|
message.SetQuestion("localhost.", mDNS.TypeA)
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
cancel()
|
|
start := time.Now()
|
|
_, err := systemExchangeForTest(ctx, transport, message)
|
|
elapsed := time.Since(start)
|
|
if err == nil {
|
|
t.Fatal("expected error for cancelled context")
|
|
}
|
|
if elapsed > time.Second {
|
|
t.Fatalf("cancellation too slow: %s", elapsed)
|
|
}
|
|
}
|
|
|
|
func TestSystemExchangeConcurrent(t *testing.T) {
|
|
requireMDNSResponder(t)
|
|
transport := &Transport{}
|
|
defer transport.system.close()
|
|
var waitGroup sync.WaitGroup
|
|
errors := make(chan error, 16)
|
|
for i := range 16 {
|
|
qtype := mDNS.TypeA
|
|
if i%2 == 1 {
|
|
qtype = mDNS.TypeAAAA
|
|
}
|
|
waitGroup.Add(1)
|
|
go func() {
|
|
defer waitGroup.Done()
|
|
message := new(mDNS.Msg)
|
|
message.SetQuestion("localhost.", qtype)
|
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
defer cancel()
|
|
response, exchangeErr := systemExchangeForTest(ctx, transport, message)
|
|
if exchangeErr != nil {
|
|
errors <- exchangeErr
|
|
return
|
|
}
|
|
if len(response.Answer) == 0 {
|
|
errors <- context.DeadlineExceeded
|
|
}
|
|
}()
|
|
}
|
|
waitGroup.Wait()
|
|
close(errors)
|
|
for exchangeErr := range errors {
|
|
t.Fatal("concurrent query failed: ", exchangeErr)
|
|
}
|
|
transport.system.queryAccess.Lock()
|
|
pendingCount := len(transport.system.queries)
|
|
transport.system.queryAccess.Unlock()
|
|
if pendingCount != 0 {
|
|
t.Fatalf("expected no pending queries after completion, got %d", pendingCount)
|
|
}
|
|
}
|