Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 20 additions & 5 deletions client.go
Original file line number Diff line number Diff line change
Expand Up @@ -209,6 +209,19 @@ func newClient(v4 bool, v6 bool, logger *log.Logger) (*client, error) {
return nil, fmt.Errorf("at least one of IPv4 and IPv6 must be enabled for querying")
}

// sendQuery writes from the multicast sockets. Loopback is required so a
// server on the same host still receives those queries.
if mconn4 != nil {
if err := ipv4.NewPacketConn(mconn4).SetMulticastLoopback(true); err != nil {
logger.Printf("[ERR] mdns: Failed to set IPv4 multicast loopback: %v", err)
}
}
if mconn6 != nil {
if err := ipv6.NewPacketConn(mconn6).SetMulticastLoopback(true); err != nil {
logger.Printf("[ERR] mdns: Failed to set IPv6 multicast loopback: %v", err)
}
}

c := &client{
use_ipv4: v4,
use_ipv6: v6,
Expand Down Expand Up @@ -397,20 +410,22 @@ func (c *client) query(params *QueryParam) error {
}
}

// sendQuery is used to multicast a query out
// sendQuery is used to multicast a query out.
// Queries are written from the multicast sockets (UDP source port 5353) so
// devices that ignore queries from ephemeral unicast ports will still answer.
func (c *client) sendQuery(q *dns.Msg) error {
buf, err := q.Pack()
if err != nil {
return err
}
if c.ipv4UnicastConn != nil {
_, err = c.ipv4UnicastConn.WriteToUDP(buf, ipv4Addr)
if c.ipv4MulticastConn != nil {
_, err = c.ipv4MulticastConn.WriteToUDP(buf, ipv4Addr)
if err != nil {
return err
}
}
if c.ipv6UnicastConn != nil {
_, err = c.ipv6UnicastConn.WriteToUDP(buf, ipv6Addr)
if c.ipv6MulticastConn != nil {
_, err = c.ipv6MulticastConn.WriteToUDP(buf, ipv6Addr)
if err != nil {
return err
}
Expand Down
88 changes: 88 additions & 0 deletions client_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
// Copyright IBM Corp. 2014, 2026
// SPDX-License-Identifier: MIT

package mdns

import (
"log"
"net"
"testing"
"time"

"github.com/miekg/dns"
)

// TestSendQuery_UsesMulticastSocket is a regression for hashicorp/mdns#144:
// sendQuery used to WriteToUDP only on the unicast sockets. Some devices
// ignore queries that do not originate on a multicast socket bound to 5353.
func TestSendQuery_UsesMulticastSocket(t *testing.T) {
ln, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
if err != nil {
t.Fatalf("listen dest: %v", err)
}
defer func() {
if err := ln.Close(); err != nil {
t.Errorf("close dest: %v", err)
}
}()

orig := ipv4Addr
ipv4Addr = ln.LocalAddr().(*net.UDPAddr)
t.Cleanup(func() { ipv4Addr = orig })

mconn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
if err != nil {
t.Fatalf("listen multicast stand-in: %v", err)
}
defer func() {
if err := mconn.Close(); err != nil {
t.Errorf("close multicast stand-in: %v", err)
}
}()

uconn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
if err != nil {
t.Fatalf("listen unicast stand-in: %v", err)
}
// Close the unicast socket so writes on it fail. sendQuery must still
// deliver the query from the multicast socket.
if err := uconn.Close(); err != nil {
t.Fatalf("close unicast: %v", err)
}

c := &client{
ipv4MulticastConn: mconn,
ipv4UnicastConn: uconn,
log: log.Default(),
}

got := make(chan *net.UDPAddr, 1)
errCh := make(chan error, 1)
go func() {
buf := make([]byte, 65536)
_ = ln.SetReadDeadline(time.Now().Add(2 * time.Second))
_, addr, err := ln.ReadFromUDP(buf)
if err != nil {
errCh <- err
return
}
got <- addr
}()

q := new(dns.Msg)
q.SetQuestion("_foobar._tcp.local.", dns.TypePTR)
q.RecursionDesired = false
if err := c.sendQuery(q); err != nil {
t.Fatalf("sendQuery: %v", err)
}

select {
case addr := <-got:
mport := mconn.LocalAddr().(*net.UDPAddr).Port
if addr.Port != mport {
t.Fatalf("query source port = %d, want multicast socket port %d", addr.Port, mport)
}
case err := <-errCh:
t.Fatalf("did not receive query on destination: %v", err)
}
}