mirror of
https://github.com/soypat/lneto.git
synced 2026-09-11 09:09:30 +00:00
Add mdns package and Client implementation (#55)
* add mdns * define mdns.Client * fix some mdns stuff * refine dns package for use with mdns * remove Querier and Responder and replace with Client * add record setting methods on dns.record to reuse record buffer * protect against unbounded client answer growth * round off sharp mdns edges; reduce allocs * add mutlicast to stacks; add xnet tests for mdns * fix documentation on acceptmulticast field * mdns working
This commit is contained in:
+1
-1
@@ -34,7 +34,7 @@ func (sudp *Client) ConnectionID() *uint64 { return &sudp.connID }
|
|||||||
func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error {
|
func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error {
|
||||||
nd := len(cfg.Questions)
|
nd := len(cfg.Questions)
|
||||||
if nd > math.MaxUint16 {
|
if nd > math.MaxUint16 {
|
||||||
return lneto.ErrBufferFull
|
return lneto.ErrInvalidConfig
|
||||||
}
|
}
|
||||||
c.reset(localPort, txid, dnsSendQuery, cfg.EnableRecursion)
|
c.reset(localPort, txid, dnsSendQuery, cfg.EnableRecursion)
|
||||||
c.msg.LimitResourceDecoding(uint16(nd), uint16(nd), 0, 0)
|
c.msg.LimitResourceDecoding(uint16(nd), uint16(nd), 0, 0)
|
||||||
|
|||||||
+115
-37
@@ -57,6 +57,13 @@ type Name struct {
|
|||||||
data []byte
|
data []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// NamesEqual reports whether two DNS names are equal by comparing
|
||||||
|
// their wire-format representations directly. This is case-sensitive;
|
||||||
|
// for case-insensitive comparison use [NamesEqualFold].
|
||||||
|
func NamesEqual(a, b Name) bool {
|
||||||
|
return internal.BytesEqual(a.data, b.data)
|
||||||
|
}
|
||||||
|
|
||||||
type ZFlags uint16
|
type ZFlags uint16
|
||||||
|
|
||||||
func NewResource(name Name, typ Type, class Class, ttl uint32, data []byte) Resource {
|
func NewResource(name Name, typ Type, class Class, ttl uint32, data []byte) Resource {
|
||||||
@@ -86,83 +93,91 @@ func (r *Resource) SetEDNS0(UDPlength uint16, rcode RCode, zflags ZFlags, data [
|
|||||||
r.data = append(r.data[:0], data...)
|
r.data = append(r.data[:0], data...)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Decode decodes the DNS message in b into m. It returns the number of bytes
|
// DecodeMessage decodes the DNS message into question, answer, authority and additional resources.
|
||||||
|
// It returns the number of bytes
|
||||||
// consumed from b (0 if no bytes were consumed) and any error encountered.
|
// consumed from b (0 if no bytes were consumed) and any error encountered.
|
||||||
// If the message was not completely parsed due to LimitResourceDecoding,
|
// If the message was not completely parsed due to LimitResourceDecoding,
|
||||||
// incompleteButOK is true and an error is returned, though the message is still usable.
|
// incompleteButOK is true and an error is returned, though the message is still usable.
|
||||||
func (m *Message) Decode(msg []byte) (_ uint16, incompleteButOK bool, err error) {
|
//
|
||||||
if len(msg) < SizeHeader {
|
// The slice memory is overwritten and capacity used as the limit of encoding.
|
||||||
return 0, false, errBaseLen
|
// If the argument slice is nil it is skipped for decoding but does not prevent further decoding
|
||||||
} else if len(msg) > math.MaxUint16 {
|
// of other answers, authorities or additionals from being decoded.
|
||||||
return 0, false, errResTooLong
|
func DecodeMessage(q *[]Question, answers, authorities, additionals *[]Resource, msg []byte) (_ uint16, incompleteButOK bool, err error) {
|
||||||
}
|
|
||||||
m.Reset()
|
|
||||||
hdr, err := NewFrame(msg)
|
hdr, err := NewFrame(msg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, false, err
|
return 0, false, err
|
||||||
}
|
}
|
||||||
nq := int(hdr.QDCount())
|
qd := hdr.QDCount()
|
||||||
|
nq := int(qd)
|
||||||
off := uint16(SizeHeader)
|
off := uint16(SizeHeader)
|
||||||
// Return tooManyErr if found to flag to the caller that the message was
|
// Return tooManyErr if found to flag to the caller that the message was
|
||||||
// decoded but contained too many resources to decode completely.
|
// decoded but contained too many resources to decode completely.
|
||||||
|
|
||||||
var tooManyErr error
|
var tooManyErr error
|
||||||
switch {
|
switch {
|
||||||
case nq > cap(m.Questions):
|
case nq > caporzero(q):
|
||||||
tooManyErr = errTooManyQuestions
|
tooManyErr = errTooManyQuestions
|
||||||
case hdr.ANCount() > uint16(cap(m.Answers)):
|
case int(hdr.ANCount()) > caporzero(answers):
|
||||||
tooManyErr = errTooManyAnswers
|
tooManyErr = errTooManyAnswers
|
||||||
case hdr.NSCount() > uint16(cap(m.Authorities)):
|
case int(hdr.NSCount()) > caporzero(authorities):
|
||||||
tooManyErr = errTooManyAuthorities
|
tooManyErr = errTooManyAuthorities
|
||||||
case hdr.ARCount() > uint16(cap(m.Additionals)):
|
case int(hdr.ARCount()) > caporzero(additionals):
|
||||||
tooManyErr = errTooManyAdditionals
|
tooManyErr = errTooManyAdditionals
|
||||||
}
|
}
|
||||||
if nq > cap(m.Questions) {
|
if q != nil {
|
||||||
nq = cap(m.Questions)
|
if nq > cap(*q) {
|
||||||
}
|
nq = cap(*q)
|
||||||
m.Questions = m.Questions[:nq]
|
|
||||||
for i := 0; i < nq; i++ {
|
|
||||||
off, err = m.Questions[i].Decode(msg, off)
|
|
||||||
if err != nil {
|
|
||||||
m.Questions = m.Questions[:i] // Trim non-decoded/failed questions.
|
|
||||||
return off, false, err
|
|
||||||
}
|
}
|
||||||
|
*q = (*q)[:nq]
|
||||||
|
for i := 0; i < nq; i++ {
|
||||||
|
off, err = (*q)[i].Decode(msg, off)
|
||||||
|
if err != nil {
|
||||||
|
*q = (*q)[:i] // Trim non-decoded/failed questions.
|
||||||
|
return off, false, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
nq = 0 // No question slice provided, skip all questions below.
|
||||||
}
|
}
|
||||||
// Skip undecoded questions.
|
// Skip undecoded questions.
|
||||||
qd := hdr.QDCount()
|
|
||||||
for i := 0; i < int(qd)-nq; i++ {
|
for i := 0; i < int(qd)-nq; i++ {
|
||||||
off, err = skipQuestion(msg, off)
|
off, err = skipQuestion(msg, off)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return off, false, err
|
return off, false, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
off, err = decodeToCapResources(answers, msg, hdr.ANCount(), off)
|
||||||
off, err = decodeToCapResources(&m.Answers, msg, hdr.ANCount(), off)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return off, false, err
|
return off, false, err
|
||||||
}
|
}
|
||||||
off, err = decodeToCapResources(&m.Authorities, msg, hdr.NSCount(), off)
|
off, err = decodeToCapResources(authorities, msg, hdr.NSCount(), off)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return off, false, err
|
return off, false, err
|
||||||
}
|
}
|
||||||
off, err = decodeToCapResources(&m.Additionals, msg, hdr.ARCount(), off)
|
off, err = decodeToCapResources(additionals, msg, hdr.ARCount(), off)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return off, false, err
|
return off, false, err
|
||||||
}
|
}
|
||||||
return off, tooManyErr != nil, tooManyErr
|
return off, tooManyErr != nil, tooManyErr
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Decode decodes the DNS message in b into m. It is a convenience wrapper for [DecodeMessage].
|
||||||
|
func (m *Message) Decode(msg []byte) (_ uint16, incompleteButOK bool, err error) {
|
||||||
|
return DecodeMessage(&m.Questions, &m.Answers, &m.Authorities, &m.Additionals, msg)
|
||||||
|
}
|
||||||
|
|
||||||
func decodeToCapResources(dst *[]Resource, msg []byte, nrec, off uint16) (_ uint16, err error) {
|
func decodeToCapResources(dst *[]Resource, msg []byte, nrec, off uint16) (_ uint16, err error) {
|
||||||
originalRec := nrec
|
originalRec := nrec
|
||||||
if nrec > uint16(cap(*dst)) {
|
if dst != nil {
|
||||||
nrec = uint16(cap(*dst)) // Decode up to cap. Caller will return an error flag.
|
if nrec > uint16(cap(*dst)) {
|
||||||
}
|
nrec = uint16(cap(*dst)) // Decode up to cap. Caller will return an error flag.
|
||||||
*dst = (*dst)[:nrec]
|
}
|
||||||
for i := uint16(0); i < nrec; i++ {
|
*dst = (*dst)[:nrec]
|
||||||
off, err = (*dst)[i].Decode(msg, off)
|
for i := uint16(0); i < nrec; i++ {
|
||||||
if err != nil {
|
off, err = (*dst)[i].Decode(msg, off)
|
||||||
*dst = (*dst)[:i] // Trim non-decoded/failed resources.
|
if err != nil {
|
||||||
return off, err
|
*dst = (*dst)[:i] // Trim non-decoded/failed resources.
|
||||||
|
return off, err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// Parse undecoded resources, effectively skipping them.
|
// Parse undecoded resources, effectively skipping them.
|
||||||
@@ -468,6 +483,20 @@ func NewName(domain string) (Name, error) {
|
|||||||
return name, nil
|
return name, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TrimLabels returns a Name sharing the same backing data with the first n labels removed.
|
||||||
|
// For example, trimming 1 label from "My Web._http._tcp.local" yields "_http._tcp.local".
|
||||||
|
// Returns an empty Name if n exceeds the number of labels.
|
||||||
|
func (n Name) TrimLabels(skip int) Name {
|
||||||
|
off := 0
|
||||||
|
for i := 0; i < skip; i++ {
|
||||||
|
if off >= len(n.data) {
|
||||||
|
return Name{}
|
||||||
|
}
|
||||||
|
off += 1 + int(n.data[off])
|
||||||
|
}
|
||||||
|
return Name{data: n.data[off:]}
|
||||||
|
}
|
||||||
|
|
||||||
// Len returns the length over-the-wire of the encoded Name.
|
// Len returns the length over-the-wire of the encoded Name.
|
||||||
func (n *Name) Len() uint16 {
|
func (n *Name) Len() uint16 {
|
||||||
if len(n.data) > math.MaxUint16 {
|
if len(n.data) > math.MaxUint16 {
|
||||||
@@ -666,6 +695,48 @@ func (dst *Resource) CopyFrom(r Resource) {
|
|||||||
dst.data = append(dst.data[:0], r.data...)
|
dst.data = append(dst.data[:0], r.data...)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetA sets an A (IPv4 address) resource record, reusing internal buffers.
|
||||||
|
func (r *Resource) SetA(name Name, class Class, ttl uint32, addr []byte) {
|
||||||
|
r.setHeader(name, TypeA, class, ttl)
|
||||||
|
r.data = append(r.data[:0], addr...)
|
||||||
|
r.header.Length = uint16(len(r.data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetPTR sets a PTR (pointer) resource record, reusing internal buffers.
|
||||||
|
func (r *Resource) SetPTR(name Name, class Class, ttl uint32, target Name) {
|
||||||
|
r.setHeader(name, TypePTR, class, ttl)
|
||||||
|
r.data, _ = target.AppendTo(r.data[:0])
|
||||||
|
r.header.Length = uint16(len(r.data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetSRV sets a SRV (service locator) resource record, reusing internal buffers.
|
||||||
|
func (r *Resource) SetSRV(name Name, class Class, ttl uint32, priority, weight, port uint16, target Name) {
|
||||||
|
r.setHeader(name, TypeSRV, class, ttl)
|
||||||
|
r.data = binary.BigEndian.AppendUint16(r.data[:0], priority)
|
||||||
|
r.data = binary.BigEndian.AppendUint16(r.data, weight)
|
||||||
|
r.data = binary.BigEndian.AppendUint16(r.data, port)
|
||||||
|
r.data, _ = target.AppendTo(r.data)
|
||||||
|
r.header.Length = uint16(len(r.data))
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetTXT sets a TXT resource record, reusing internal buffers.
|
||||||
|
func (r *Resource) SetTXT(name Name, class Class, ttl uint32, txt []byte) {
|
||||||
|
r.setHeader(name, TypeTXT, class, ttl)
|
||||||
|
if len(txt) == 0 {
|
||||||
|
r.data = append(r.data[:0], 0)
|
||||||
|
} else {
|
||||||
|
r.data = append(r.data[:0], txt...)
|
||||||
|
}
|
||||||
|
r.header.Length = uint16(len(r.data))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Resource) setHeader(name Name, typ Type, class Class, ttl uint32) {
|
||||||
|
r.header.Name.CopyFrom(name)
|
||||||
|
r.header.Type = typ
|
||||||
|
r.header.Class = class
|
||||||
|
r.header.TTL = ttl
|
||||||
|
}
|
||||||
|
|
||||||
func (dst *ResourceHeader) CopyFrom(rh ResourceHeader) {
|
func (dst *ResourceHeader) CopyFrom(rh ResourceHeader) {
|
||||||
dst.Name.CopyFrom(rh.Name)
|
dst.Name.CopyFrom(rh.Name)
|
||||||
dst.Type = rh.Type
|
dst.Type = rh.Type
|
||||||
@@ -673,3 +744,10 @@ func (dst *ResourceHeader) CopyFrom(rh ResourceHeader) {
|
|||||||
dst.TTL = rh.TTL
|
dst.TTL = rh.TTL
|
||||||
dst.Length = rh.Length
|
dst.Length = rh.Length
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func caporzero[T any](v *[]T) int {
|
||||||
|
if v == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return cap(*v)
|
||||||
|
}
|
||||||
|
|||||||
@@ -117,6 +117,7 @@ func run() (err error) {
|
|||||||
}
|
}
|
||||||
fmt.Println("NIC hardware address:", net.HardwareAddr(nicHW[:]).String(), "bridgeHW:", net.HardwareAddr(brHW[:]).String(), "mtu:", mtu, "addr:", nicAddr.String())
|
fmt.Println("NIC hardware address:", net.HardwareAddr(nicHW[:]).String(), "bridgeHW:", net.HardwareAddr(brHW[:]).String(), "mtu:", mtu, "addr:", nicAddr.String())
|
||||||
var stack xnet.StackAsync
|
var stack xnet.StackAsync
|
||||||
|
|
||||||
err = stack.Reset(xnet.StackConfig{
|
err = stack.Reset(xnet.StackConfig{
|
||||||
Hostname: "xnet-test",
|
Hostname: "xnet-test",
|
||||||
RandSeed: softRand,
|
RandSeed: softRand,
|
||||||
@@ -221,6 +222,7 @@ func run() (err error) {
|
|||||||
return fmt.Errorf("DHCP failed: %w", err)
|
return fmt.Errorf("DHCP failed: %w", err)
|
||||||
}
|
}
|
||||||
timeDHCP()
|
timeDHCP()
|
||||||
|
|
||||||
err = stack.AssimilateDHCPResults(results)
|
err = stack.AssimilateDHCPResults(results)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("assimilating DHCP results: %w", err)
|
return fmt.Errorf("assimilating DHCP results: %w", err)
|
||||||
|
|||||||
@@ -295,12 +295,14 @@ func (h *Header) reuseOrAppend(tok headerSlice, value string) headerSlice {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (h *Header) appendSlice(value string) headerSlice {
|
func (h *Header) appendSlice(value string) headerSlice {
|
||||||
|
debuglog("http:appendslice:start")
|
||||||
free := h.hbuf.free()
|
free := h.hbuf.free()
|
||||||
if len(value) > free {
|
if len(value) > free {
|
||||||
if h.flags.hasAny(flagNoBufferGrow) {
|
if h.flags.hasAny(flagNoBufferGrow) {
|
||||||
h.flags |= flagOOMReached
|
h.flags |= flagOOMReached
|
||||||
return headerSlice{}
|
return headerSlice{}
|
||||||
}
|
}
|
||||||
|
debuglog("http:appendslice:grow-buf")
|
||||||
h.hbuf.buf = slices.Grow(h.hbuf.buf, len(value)+1) // Grow 1 beyond due to slice validity.
|
h.hbuf.buf = slices.Grow(h.hbuf.buf, len(value)+1) // Grow 1 beyond due to slice validity.
|
||||||
}
|
}
|
||||||
h.flags |= flagMangledBuffer
|
h.flags |= flagMangledBuffer
|
||||||
|
|||||||
@@ -52,3 +52,41 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) {
|
|||||||
copy(dstaddr, dst)
|
copy(dstaddr, dst)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetMulticast sets the IP destination to multicastAddr and derives the
|
||||||
|
// Ethernet destination MAC from it. It supports IPv4 (RFC 1112 §6.4) and
|
||||||
|
// IPv6 (RFC 2464 §7) multicast MAC mapping.
|
||||||
|
func SetMulticast(ethernetCarrier []byte, ipOff int, multicastAddr []byte) (err error) {
|
||||||
|
ip := ethernetCarrier[ipOff:]
|
||||||
|
mac := ethernetCarrier[0:6]
|
||||||
|
version := ip[0] >> 4
|
||||||
|
switch version {
|
||||||
|
case 4:
|
||||||
|
if len(multicastAddr) != 4 {
|
||||||
|
return lneto.ErrMismatchLen
|
||||||
|
}
|
||||||
|
copy(ip[16:20], multicastAddr)
|
||||||
|
// IPv4 multicast MAC: 01:00:5e + low 23 bits of IP destination.
|
||||||
|
mac[0] = 0x01
|
||||||
|
mac[1] = 0x00
|
||||||
|
mac[2] = 0x5e
|
||||||
|
mac[3] = multicastAddr[1] & 0x7f
|
||||||
|
mac[4] = multicastAddr[2]
|
||||||
|
mac[5] = multicastAddr[3]
|
||||||
|
case 6:
|
||||||
|
if len(multicastAddr) != 16 {
|
||||||
|
return lneto.ErrMismatchLen
|
||||||
|
}
|
||||||
|
copy(ip[24:40], multicastAddr)
|
||||||
|
// IPv6 multicast MAC: 33:33 + last 4 bytes of IP destination.
|
||||||
|
mac[0] = 0x33
|
||||||
|
mac[1] = 0x33
|
||||||
|
mac[2] = multicastAddr[12]
|
||||||
|
mac[3] = multicastAddr[13]
|
||||||
|
mac[4] = multicastAddr[14]
|
||||||
|
mac[5] = multicastAddr[15]
|
||||||
|
default:
|
||||||
|
return lneto.ErrUnsupported
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ import (
|
|||||||
"github.com/soypat/lneto/ipv4"
|
"github.com/soypat/lneto/ipv4"
|
||||||
"github.com/soypat/lneto/ipv4/icmpv4"
|
"github.com/soypat/lneto/ipv4/icmpv4"
|
||||||
"github.com/soypat/lneto/ipv6"
|
"github.com/soypat/lneto/ipv6"
|
||||||
|
"github.com/soypat/lneto/mdns"
|
||||||
"github.com/soypat/lneto/ntp"
|
"github.com/soypat/lneto/ntp"
|
||||||
"github.com/soypat/lneto/tcp"
|
"github.com/soypat/lneto/tcp"
|
||||||
"github.com/soypat/lneto/udp"
|
"github.com/soypat/lneto/udp"
|
||||||
@@ -359,7 +360,7 @@ func (pc *PacketBreakdown) CaptureUDP(dst []Frame, pkt []byte, bitOffset int) ([
|
|||||||
srcport := ufrm.SourcePort()
|
srcport := ufrm.SourcePort()
|
||||||
if dhcpv4.PayloadIsDHCPv4(payload) {
|
if dhcpv4.PayloadIsDHCPv4(payload) {
|
||||||
dst, err = pc.CaptureDHCPv4(dst, pkt, end)
|
dst, err = pc.CaptureDHCPv4(dst, pkt, end)
|
||||||
} else if dstport == dns.ServerPort || srcport == dns.ServerPort {
|
} else if dstport == dns.ServerPort || srcport == dns.ServerPort || dstport == mdns.Port || srcport == mdns.Port {
|
||||||
dst, err = pc.CaptureDNS(dst, pkt, end)
|
dst, err = pc.CaptureDNS(dst, pkt, end)
|
||||||
} else if dstport == ntp.ServerPort || srcport == ntp.ServerPort {
|
} else if dstport == ntp.ServerPort || srcport == ntp.ServerPort {
|
||||||
dst, err = pc.CaptureNTP(dst, pkt, end)
|
dst, err = pc.CaptureNTP(dst, pkt, end)
|
||||||
@@ -437,7 +438,7 @@ func (pc *PacketBreakdown) CaptureDNS(dst []Frame, pkt []byte, bitOffset int) ([
|
|||||||
return dst, errNotByteAligned
|
return dst, errNotByteAligned
|
||||||
}
|
}
|
||||||
dnsData := pkt[bitOffset/8:]
|
dnsData := pkt[bitOffset/8:]
|
||||||
pc.dmsg.LimitResourceDecoding(20, 20, 20, 20)
|
pc.dmsg.LimitResourceDecoding(4, 4, 4, 4)
|
||||||
off, incomplete, err := pc.dmsg.Decode(dnsData)
|
off, incomplete, err := pc.dmsg.Decode(dnsData)
|
||||||
if err != nil && !incomplete {
|
if err != nil && !incomplete {
|
||||||
return dst, err
|
return dst, err
|
||||||
@@ -1265,6 +1266,43 @@ var baseDHCPv4Fields = [...]FrameField{
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var baseDNSFields = [...]FrameField{
|
||||||
|
{
|
||||||
|
Class: FieldClassID,
|
||||||
|
FrameBitOffset: 0,
|
||||||
|
BitLength: 2 * octet,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Class: FieldClassFlags,
|
||||||
|
FrameBitOffset: 2 * octet,
|
||||||
|
BitLength: 2 * octet,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "Questions",
|
||||||
|
Class: FieldClassSize,
|
||||||
|
FrameBitOffset: 4 * octet,
|
||||||
|
BitLength: 2 * octet,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "Answers",
|
||||||
|
Class: FieldClassSize,
|
||||||
|
FrameBitOffset: 6 * octet,
|
||||||
|
BitLength: 2 * octet,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "Authorities",
|
||||||
|
Class: FieldClassSize,
|
||||||
|
FrameBitOffset: 8 * octet,
|
||||||
|
BitLength: 2 * octet,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "Additionals",
|
||||||
|
Class: FieldClassSize,
|
||||||
|
FrameBitOffset: 10 * octet,
|
||||||
|
BitLength: 2 * octet,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
var baseNTPFields = [...]FrameField{
|
var baseNTPFields = [...]FrameField{
|
||||||
{
|
{
|
||||||
Name: "Mode",
|
Name: "Mode",
|
||||||
|
|||||||
+19
-11
@@ -33,11 +33,12 @@ type StackEthernetConfig struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type StackEthernet struct {
|
type StackEthernet struct {
|
||||||
connID uint64
|
connID uint64
|
||||||
handlers handlers
|
handlers handlers
|
||||||
mac [6]byte
|
mac [6]byte
|
||||||
gwmac [6]byte
|
gwmac [6]byte
|
||||||
mtu uint16
|
mtu uint16
|
||||||
|
acceptMulticast bool
|
||||||
// crcupdate set when crc32 has been configured to be appended.
|
// crcupdate set when crc32 has been configured to be appended.
|
||||||
crcupdate func(crc uint32, p []byte) uint32
|
crcupdate func(crc uint32, p []byte) uint32
|
||||||
}
|
}
|
||||||
@@ -50,6 +51,10 @@ func (ls *StackEthernet) Gateway6() (gw [6]byte) {
|
|||||||
return ls.gwmac
|
return ls.gwmac
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (ls *StackEthernet) SetAcceptMulticast(accept bool) {
|
||||||
|
ls.acceptMulticast = accept
|
||||||
|
}
|
||||||
|
|
||||||
func (ls *StackEthernet) SetHardwareAddr6(mac [6]byte) {
|
func (ls *StackEthernet) SetHardwareAddr6(mac [6]byte) {
|
||||||
ls.mac = mac
|
ls.mac = mac
|
||||||
}
|
}
|
||||||
@@ -83,11 +88,12 @@ func (ls *StackEthernet) Configure(cfg StackEthernetConfig) error {
|
|||||||
}
|
}
|
||||||
ls.handlers.reset("StackEthernet", cfg.MaxNodes)
|
ls.handlers.reset("StackEthernet", cfg.MaxNodes)
|
||||||
*ls = StackEthernet{
|
*ls = StackEthernet{
|
||||||
connID: ls.connID + 1,
|
connID: ls.connID + 1,
|
||||||
handlers: ls.handlers,
|
handlers: ls.handlers,
|
||||||
mac: cfg.MAC,
|
mac: cfg.MAC,
|
||||||
gwmac: cfg.Gateway,
|
gwmac: cfg.Gateway,
|
||||||
mtu: uint16(cfg.MTU),
|
mtu: uint16(cfg.MTU),
|
||||||
|
acceptMulticast: ls.acceptMulticast,
|
||||||
}
|
}
|
||||||
if cfg.AppendCRC32 {
|
if cfg.AppendCRC32 {
|
||||||
ls.crcupdate = cfg.CRC32Update
|
ls.crcupdate = cfg.CRC32Update
|
||||||
@@ -121,7 +127,9 @@ func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error)
|
|||||||
dstaddr := efrm.DestinationHardwareAddr()
|
dstaddr := efrm.DestinationHardwareAddr()
|
||||||
var vld lneto.Validator
|
var vld lneto.Validator
|
||||||
if !efrm.IsBroadcast() && ls.mac != *dstaddr {
|
if !efrm.IsBroadcast() && ls.mac != *dstaddr {
|
||||||
goto DROP
|
if !ls.acceptMulticast || dstaddr[0]&1 == 0 {
|
||||||
|
goto DROP
|
||||||
|
}
|
||||||
}
|
}
|
||||||
efrm.ValidateSize(&vld)
|
efrm.ValidateSize(&vld)
|
||||||
if vld.HasError() {
|
if vld.HasError() {
|
||||||
|
|||||||
+19
-11
@@ -16,11 +16,12 @@ import (
|
|||||||
var _ StackNode = (*StackIP)(nil)
|
var _ StackNode = (*StackIP)(nil)
|
||||||
|
|
||||||
type StackIP struct {
|
type StackIP struct {
|
||||||
connID uint64
|
connID uint64
|
||||||
ipID uint16
|
ipID uint16
|
||||||
ip [4]byte
|
ip [4]byte
|
||||||
validator lneto.Validator
|
acceptMulticast bool
|
||||||
handlers handlers
|
validator lneto.Validator
|
||||||
|
handlers handlers
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error {
|
func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error {
|
||||||
@@ -33,10 +34,11 @@ func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error {
|
|||||||
}
|
}
|
||||||
sb.handlers.reset("StackIP", maxNodes)
|
sb.handlers.reset("StackIP", maxNodes)
|
||||||
*sb = StackIP{
|
*sb = StackIP{
|
||||||
connID: sb.connID + 1,
|
connID: sb.connID + 1,
|
||||||
validator: sb.validator,
|
validator: sb.validator,
|
||||||
handlers: sb.handlers,
|
handlers: sb.handlers,
|
||||||
ip: sb.ip,
|
ip: sb.ip,
|
||||||
|
acceptMulticast: sb.acceptMulticast,
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -65,6 +67,10 @@ func (sb *StackIP) Addr() netip.Addr {
|
|||||||
return netip.AddrFrom4(sb.ip)
|
return netip.AddrFrom4(sb.ip)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (sb *StackIP) SetAcceptMulticast(accept bool) {
|
||||||
|
sb.acceptMulticast = accept
|
||||||
|
}
|
||||||
|
|
||||||
func (sb *StackIP) SetLogger(logger *slog.Logger) {
|
func (sb *StackIP) SetLogger(logger *slog.Logger) {
|
||||||
sb.handlers.log = logger
|
sb.handlers.log = logger
|
||||||
}
|
}
|
||||||
@@ -79,8 +85,10 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
|
|||||||
}
|
}
|
||||||
dst := ifrm.DestinationAddr()
|
dst := ifrm.DestinationAddr()
|
||||||
if sb.ip != ([4]byte{}) && *dst != sb.ip {
|
if sb.ip != ([4]byte{}) && *dst != sb.ip {
|
||||||
sb.handlers.debug("ip:not-for-us")
|
if !sb.acceptMulticast || dst[0]&0xF0 != 0xE0 {
|
||||||
return lneto.ErrPacketDrop // Not meant for us.
|
sb.handlers.debug("ip:not-for-us")
|
||||||
|
return lneto.ErrPacketDrop // Not meant for us.
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
sb.validator.ResetErr()
|
sb.validator.ResetErr()
|
||||||
|
|||||||
@@ -6,6 +6,14 @@ const (
|
|||||||
sizeHeader = 20
|
sizeHeader = 20
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// IsMulticast reports whether addr is an IPv4 multicast address (224.0.0.0/4),
|
||||||
|
// i.e. the most significant nibble is 0xE (1110 in binary) as defined in [RFC1112].
|
||||||
|
//
|
||||||
|
// [RFC1112]: https://datatracker.ietf.org/doc/html/rfc1112
|
||||||
|
func IsMulticast(addr [4]byte) bool {
|
||||||
|
return addr[0]&0xf0 == 0xe0
|
||||||
|
}
|
||||||
|
|
||||||
// ToS represents the Traffic Class (a.k.a Type of Service). It is 8 bits long. 6 MSB are Differentiated Services; 2 LSB are Explicit Congenstion Notification.
|
// ToS represents the Traffic Class (a.k.a Type of Service). It is 8 bits long. 6 MSB are Differentiated Services; 2 LSB are Explicit Congenstion Notification.
|
||||||
type ToS uint8
|
type ToS uint8
|
||||||
|
|
||||||
|
|||||||
+263
@@ -0,0 +1,263 @@
|
|||||||
|
package mdns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"github.com/soypat/lneto"
|
||||||
|
"github.com/soypat/lneto/dns"
|
||||||
|
"github.com/soypat/lneto/internal"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// Port is the mDNS UDP port (RFC 6762 §1).
|
||||||
|
Port = 5353
|
||||||
|
|
||||||
|
// Default TTL for mDNS records (RFC 6762 §11).
|
||||||
|
DefaultTTL uint32 = 120
|
||||||
|
// classCacheFlush is bit 15 of the Class field, indicating the record
|
||||||
|
// is from a unique source and should replace cached entries (RFC 6762 §10.2).
|
||||||
|
classCacheFlush uint16 = 1 << 15
|
||||||
|
mdnsTxID = 0
|
||||||
|
mdnsFlags = 0
|
||||||
|
)
|
||||||
|
|
||||||
|
type querierState uint8
|
||||||
|
|
||||||
|
const (
|
||||||
|
querierIdle querierState = iota
|
||||||
|
querierSendQuery // Query ready to be sent.
|
||||||
|
querierAwaitResponse // Waiting for answers.
|
||||||
|
querierFailed // failed query
|
||||||
|
querierDone // Answers collected.
|
||||||
|
)
|
||||||
|
|
||||||
|
// Client provides both querying and service multicast DNS functionality
|
||||||
|
// once configured and attached to MDNS port 5353.
|
||||||
|
//
|
||||||
|
// Clients are attached to MDNS ports and function until manual detachment
|
||||||
|
// due to their dual design: they double as a querier and service discovery.
|
||||||
|
type Client struct {
|
||||||
|
connID uint64
|
||||||
|
closed bool
|
||||||
|
lport uint16
|
||||||
|
ip []byte
|
||||||
|
// Query State:
|
||||||
|
qstate querierState
|
||||||
|
qcode dns.RCode
|
||||||
|
qerr error
|
||||||
|
// qmsg is used for queries to marshal/unmarshal our
|
||||||
|
// outgoing queries and responses to our queries.
|
||||||
|
qans []dns.Resource
|
||||||
|
qqst []dns.Question
|
||||||
|
// Response state:
|
||||||
|
services []Service // Stores services we'd broadcast.
|
||||||
|
rans []dns.Resource
|
||||||
|
rqst []dns.Question
|
||||||
|
}
|
||||||
|
|
||||||
|
type ClientConfig struct {
|
||||||
|
LocalPort uint16
|
||||||
|
Services []Service
|
||||||
|
MulticastAddr []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) Configure(cfg ClientConfig) error {
|
||||||
|
if cfg.LocalPort == 0 {
|
||||||
|
return lneto.ErrZeroSource
|
||||||
|
}
|
||||||
|
c.reset(cfg.LocalPort)
|
||||||
|
c.services = append(c.services[:0], cfg.Services...)
|
||||||
|
c.ip = append(c.ip[:0], cfg.MulticastAddr...)
|
||||||
|
internal.SliceReuse(&c.rqst, len(cfg.Services))
|
||||||
|
// Each service can produce up to 4 answer records (PTR+SRV+TXT+A).
|
||||||
|
nrans := 2 * len(cfg.Services)
|
||||||
|
if nrans > 0 {
|
||||||
|
nrans = max(4, nrans)
|
||||||
|
}
|
||||||
|
internal.SliceReuse(&c.rans, nrans)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) }
|
||||||
|
|
||||||
|
func (c *Client) LocalPort() uint16 { return c.lport }
|
||||||
|
|
||||||
|
func (c *Client) ConnectionID() *uint64 { return &c.connID }
|
||||||
|
|
||||||
|
type ResolveConfig struct {
|
||||||
|
Questions []dns.Question
|
||||||
|
MaxResponseAnswers uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) StartResolve(cfg ResolveConfig) error {
|
||||||
|
nq := len(cfg.Questions)
|
||||||
|
if nq > math.MaxUint16 || nq == 0 || cfg.MaxResponseAnswers == 0 {
|
||||||
|
return lneto.ErrInvalidConfig
|
||||||
|
}
|
||||||
|
c.qreset(querierSendQuery)
|
||||||
|
internal.SliceReuse(&c.qans, int(cfg.MaxResponseAnswers))
|
||||||
|
internal.SliceReuse(&c.qqst, nq)
|
||||||
|
c.qqst = c.qqst[:nq]
|
||||||
|
for i := range c.qqst {
|
||||||
|
c.qqst[i].CopyFrom(cfg.Questions[i])
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) reset(localport uint16) {
|
||||||
|
*c = Client{
|
||||||
|
connID: c.connID + 1,
|
||||||
|
lport: localport,
|
||||||
|
// Ensure memory reused:
|
||||||
|
qqst: c.qqst[:0],
|
||||||
|
qans: c.qans[:0],
|
||||||
|
services: c.services[:0],
|
||||||
|
rans: c.rans[:0],
|
||||||
|
ip: c.ip[:0],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// qreset resets the current query state. It is only a partial reset of a Client.
|
||||||
|
func (c *Client) qreset(state querierState) {
|
||||||
|
c.qstate = state
|
||||||
|
}
|
||||||
|
|
||||||
|
// Encapsulate writes a pending mDNS packet into carrierData[offsetToFrame:].
|
||||||
|
// Pending responses take priority over outgoing queries.
|
||||||
|
func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) {
|
||||||
|
if c.isClosed() {
|
||||||
|
return 0, net.ErrClosed
|
||||||
|
}
|
||||||
|
if len(c.rans) > 0 {
|
||||||
|
// Pending response to an incoming query.
|
||||||
|
n, err = c.encapsResponse(carrierData[offsetToFrame:])
|
||||||
|
} else if c.qstate == querierSendQuery {
|
||||||
|
n, err = c.encapsQuery(carrierData[offsetToFrame:])
|
||||||
|
}
|
||||||
|
if n > 0 && offsetToIP >= 0 {
|
||||||
|
// Set Multicast IP destination and Ethernet MAC.
|
||||||
|
internal.SetMulticast(carrierData, offsetToIP, c.ip)
|
||||||
|
}
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Demux processes an incoming mDNS response packet. Answers are accumulated
|
||||||
|
// into the internal message. Once sufficient answers are collected or a
|
||||||
|
// timeout occurs the querier transitions to querierDone.
|
||||||
|
func (c *Client) Demux(carrierData []byte, frameOffset int) error {
|
||||||
|
if c.isClosed() {
|
||||||
|
return net.ErrClosed
|
||||||
|
}
|
||||||
|
frame := carrierData[frameOffset:]
|
||||||
|
f, err := dns.NewFrame(frame)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
} else if f.TxID() != 0 {
|
||||||
|
return lneto.ErrPacketDrop
|
||||||
|
}
|
||||||
|
flags := f.Flags()
|
||||||
|
isresponse := flags.IsResponse()
|
||||||
|
if isresponse && c.qstate == querierAwaitResponse {
|
||||||
|
c.qcode = flags.ResponseCode()
|
||||||
|
// Decode response into our message, collecting answers.
|
||||||
|
_, _, c.qerr = dns.DecodeMessage(nil, &c.qans, nil, nil, frame)
|
||||||
|
if c.qerr != nil {
|
||||||
|
c.qstate = querierFailed
|
||||||
|
return c.qerr
|
||||||
|
}
|
||||||
|
c.qstate = querierDone
|
||||||
|
return nil // success.
|
||||||
|
}
|
||||||
|
freeAns := cap(c.rans) - len(c.rans)
|
||||||
|
if !isresponse && len(c.services) > 0 && freeAns > 0 {
|
||||||
|
// Incoming query — match against our services.
|
||||||
|
var query dns.Message
|
||||||
|
query.LimitResourceDecoding(f.QDCount(), 0, 0, 0)
|
||||||
|
_, _, err = query.Decode(frame)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for i := range query.Questions {
|
||||||
|
q := &query.Questions[i]
|
||||||
|
for j := range c.services {
|
||||||
|
if matchQuestion(q, &c.services[j]) {
|
||||||
|
addServiceAnswers(&c.rans, q, &c.services[j])
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) encapsQuery(frame []byte) (int, error) {
|
||||||
|
msg := dns.Message{
|
||||||
|
Questions: c.qqst,
|
||||||
|
}
|
||||||
|
msglen := msg.Len()
|
||||||
|
if int(msglen) > len(frame) {
|
||||||
|
c.qerr = lneto.ErrShortBuffer
|
||||||
|
c.qstate = querierFailed
|
||||||
|
return 0, c.qerr
|
||||||
|
}
|
||||||
|
// mDNS queries use txid=0 and no flags (RFC 6762 §18.1).
|
||||||
|
data, err := msg.AppendTo(frame[:0], mdnsTxID, mdnsFlags)
|
||||||
|
if err != nil {
|
||||||
|
c.qerr = err
|
||||||
|
c.qstate = querierFailed
|
||||||
|
return 0, err
|
||||||
|
} else if len(data) != int(msglen) {
|
||||||
|
panic("bad dns length calculation") // panic since this is a big bug in lneto.
|
||||||
|
}
|
||||||
|
c.qstate = querierAwaitResponse
|
||||||
|
return len(data), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) encapsResponse(frame []byte) (int, error) {
|
||||||
|
var msg dns.Message
|
||||||
|
msg.Answers = c.rans
|
||||||
|
msglen := msg.Len()
|
||||||
|
if int(msglen) > len(frame) {
|
||||||
|
return 0, lneto.ErrShortBuffer
|
||||||
|
}
|
||||||
|
// mDNS responses: txid=0, QR=1, AA=1 (RFC 6762 §18.4, §6).
|
||||||
|
flags := dns.HeaderFlags(1<<15 | 1<<10)
|
||||||
|
data, err := msg.AppendTo(frame[:0], 0, flags)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return len(data), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Abort closes the client, causing all subsequent Encapsulate/Demux calls to return [net.ErrClosed].
|
||||||
|
func (c *Client) Abort() {
|
||||||
|
c.closed = true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) isClosed() bool {
|
||||||
|
return c.closed
|
||||||
|
}
|
||||||
|
|
||||||
|
// AnswersCopyTo checks if [Client.StartResolve] ended succesfully before
|
||||||
|
// doing a deep copy of answers received to the argument buffer using [dns.Resource.CopyFrom].
|
||||||
|
func (c *Client) AnswersCopyTo(dst []dns.Resource) (n int, done bool, err error) {
|
||||||
|
if len(dst) == 0 {
|
||||||
|
return 0, false, lneto.ErrShortBuffer
|
||||||
|
} else if c.qstate == querierIdle {
|
||||||
|
return 0, false, net.ErrClosed
|
||||||
|
} else if c.qstate == querierFailed {
|
||||||
|
return 0, false, c.qerr
|
||||||
|
} else if c.qstate != querierDone {
|
||||||
|
return 0, false, nil
|
||||||
|
}
|
||||||
|
for i := range min(len(dst), len(c.qans)) {
|
||||||
|
dst[i].CopyFrom(c.qans[i])
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
rcode := c.qcode
|
||||||
|
if rcode != 0 {
|
||||||
|
return n, true, rcode
|
||||||
|
}
|
||||||
|
return n, true, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,122 @@
|
|||||||
|
package mdns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/soypat/lneto/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Service describes a service to advertise via mDNS.
|
||||||
|
// A single Service produces PTR, SRV, TXT, and A resource records.
|
||||||
|
//
|
||||||
|
// i.e: To generate a hostname styled A record like the one
|
||||||
|
// linux machines provide to reach them at hostname.local:
|
||||||
|
//
|
||||||
|
// s := Service{
|
||||||
|
// Host: dns.NewName("yourhostname.local"),
|
||||||
|
// Addr: ipAddressSlice,
|
||||||
|
// }
|
||||||
|
type Service struct {
|
||||||
|
// Name is the fully-qualified service instance name in wire format,
|
||||||
|
// e.g. "My Web Server._http._tcp.local".
|
||||||
|
Name dns.Name
|
||||||
|
// Host is the hostname in wire format, e.g. "mydevice.local".
|
||||||
|
Host dns.Name
|
||||||
|
// TXTData is raw TXT record data (length-prefixed strings).
|
||||||
|
TXTData []byte
|
||||||
|
// Addr is the IP address for the A record.
|
||||||
|
Addr []byte
|
||||||
|
// TTL is the record TTL in seconds. Zero uses DefaultTTL.
|
||||||
|
TTL uint32
|
||||||
|
// Port is the TCP/UDP port for the SRV record.
|
||||||
|
Port uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) ttl() uint32 {
|
||||||
|
if s.TTL == 0 {
|
||||||
|
return DefaultTTL
|
||||||
|
}
|
||||||
|
return s.TTL
|
||||||
|
}
|
||||||
|
|
||||||
|
// serviceType extracts the service type portion of the instance name.
|
||||||
|
// For "_http._tcp.local" it returns the same; for "My Web._http._tcp.local"
|
||||||
|
// it returns "_http._tcp.local" by trimming the first label.
|
||||||
|
// Returns a view into the original Name data — zero allocation.
|
||||||
|
func (s *Service) serviceType() dns.Name {
|
||||||
|
var totalLabels int
|
||||||
|
s.Name.VisitLabels(func(label []byte) {
|
||||||
|
totalLabels++
|
||||||
|
})
|
||||||
|
if totalLabels <= 3 {
|
||||||
|
return s.Name
|
||||||
|
}
|
||||||
|
return s.Name.TrimLabels(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// matchQuestion reports whether the question matches the given service.
|
||||||
|
func matchQuestion(q *dns.Question, svc *Service) bool {
|
||||||
|
switch q.Type {
|
||||||
|
case dns.TypePTR:
|
||||||
|
svcType := svc.serviceType()
|
||||||
|
return dns.NamesEqual(q.Name, svcType)
|
||||||
|
case dns.TypeSRV, dns.TypeTXT:
|
||||||
|
return dns.NamesEqual(q.Name, svc.Name)
|
||||||
|
case dns.TypeA:
|
||||||
|
return dns.NamesEqual(q.Name, svc.Host)
|
||||||
|
case dns.TypeALL:
|
||||||
|
svcType := svc.serviceType()
|
||||||
|
return dns.NamesEqual(q.Name, svc.Name) || dns.NamesEqual(q.Name, svc.Host) || dns.NamesEqual(q.Name, svcType)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// addServiceAnswers adds the appropriate answer records for a matched question.
|
||||||
|
// It grows ans in-place, reusing existing Resource buffers when available.
|
||||||
|
func addServiceAnswers(dst *[]dns.Resource, q *dns.Question, svc *Service) {
|
||||||
|
cacheFlush := dns.Class(uint16(dns.ClassINET) | classCacheFlush)
|
||||||
|
ttl := svc.ttl()
|
||||||
|
txtData := svc.TXTData
|
||||||
|
avail := cap(*dst) - len(*dst)
|
||||||
|
switch q.Type {
|
||||||
|
case dns.TypePTR:
|
||||||
|
if avail < 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
setPTR(growSlice(dst), svc)
|
||||||
|
case dns.TypeSRV:
|
||||||
|
if avail < 2 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
growSlice(dst).SetSRV(svc.Name, cacheFlush, ttl, 0, 0, svc.Port, svc.Host)
|
||||||
|
growSlice(dst).SetA(svc.Host, cacheFlush, ttl, svc.Addr)
|
||||||
|
case dns.TypeTXT:
|
||||||
|
if avail < 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
growSlice(dst).SetTXT(svc.Name, cacheFlush, ttl, txtData)
|
||||||
|
case dns.TypeA:
|
||||||
|
if avail < 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
growSlice(dst).SetA(svc.Host, cacheFlush, ttl, svc.Addr)
|
||||||
|
case dns.TypeALL:
|
||||||
|
if avail < 4 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
setPTR(growSlice(dst), svc)
|
||||||
|
growSlice(dst).SetSRV(svc.Name, cacheFlush, ttl, 0, 0, svc.Port, svc.Host)
|
||||||
|
growSlice(dst).SetTXT(svc.Name, cacheFlush, ttl, txtData)
|
||||||
|
growSlice(dst).SetA(svc.Host, cacheFlush, ttl, svc.Addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// growSlice grows the slice by one element and returns a pointer to the new last element.
|
||||||
|
// Panics if at capacity — callers must check available space before calling.
|
||||||
|
func growSlice[T any](s *[]T) *T {
|
||||||
|
*s = (*s)[:len(*s)+1]
|
||||||
|
return &(*s)[len(*s)-1]
|
||||||
|
}
|
||||||
|
|
||||||
|
func setPTR(ans *dns.Resource, svc *Service) {
|
||||||
|
svcType := svc.serviceType()
|
||||||
|
ans.SetPTR(svcType, dns.ClassINET, svc.ttl(), svc.Name)
|
||||||
|
}
|
||||||
@@ -0,0 +1,409 @@
|
|||||||
|
package mdns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"net"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/soypat/lneto/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
func mustNewName(s string) dns.Name {
|
||||||
|
n, err := dns.NewName(s)
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func testService() Service {
|
||||||
|
return Service{
|
||||||
|
Name: mustNewName("My Web._http._tcp.local"),
|
||||||
|
Host: mustNewName("mydevice.local"),
|
||||||
|
Addr: []byte{192, 168, 1, 50},
|
||||||
|
Port: 80,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newQuerier(t *testing.T, questions []dns.Question, maxAnswers uint16) *Client {
|
||||||
|
t.Helper()
|
||||||
|
var c Client
|
||||||
|
err := c.Configure(ClientConfig{LocalPort: Port})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
err = c.StartResolve(ResolveConfig{
|
||||||
|
Questions: questions,
|
||||||
|
MaxResponseAnswers: maxAnswers,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return &c
|
||||||
|
}
|
||||||
|
|
||||||
|
func newResponder(t *testing.T, services []Service) *Client {
|
||||||
|
t.Helper()
|
||||||
|
var c Client
|
||||||
|
err := c.Configure(ClientConfig{
|
||||||
|
LocalPort: Port,
|
||||||
|
Services: services,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return &c
|
||||||
|
}
|
||||||
|
|
||||||
|
// queryRespond performs the full query->demux->encapsulate->demux cycle
|
||||||
|
// between a querier and responder, returning the response length.
|
||||||
|
func queryRespond(t *testing.T, querier, responder *Client, buf []byte) int {
|
||||||
|
t.Helper()
|
||||||
|
n, err := querier.Encapsulate(buf, -1, 0)
|
||||||
|
if err != nil || n == 0 {
|
||||||
|
t.Fatal("querier encapsulate:", err, n)
|
||||||
|
}
|
||||||
|
if err = responder.Demux(buf[:n], 0); err != nil {
|
||||||
|
t.Fatal("responder demux:", err)
|
||||||
|
}
|
||||||
|
n, err = responder.Encapsulate(buf, -1, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("responder encapsulate:", err)
|
||||||
|
}
|
||||||
|
if n > 0 {
|
||||||
|
if err = querier.Demux(buf[:n], 0); err != nil {
|
||||||
|
t.Fatal("querier demux:", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientQueryPTR(t *testing.T) {
|
||||||
|
svc := testService()
|
||||||
|
responder := newResponder(t, []Service{svc})
|
||||||
|
querier := newQuerier(t, []dns.Question{{
|
||||||
|
Name: mustNewName("_http._tcp.local"),
|
||||||
|
Type: dns.TypePTR,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}}, 8)
|
||||||
|
|
||||||
|
var buf [1024]byte
|
||||||
|
|
||||||
|
// Querier encapsulates query.
|
||||||
|
n, err := querier.Encapsulate(buf[:], -1, 0)
|
||||||
|
if err != nil || n == 0 {
|
||||||
|
t.Fatal("querier encapsulate:", err, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify query wire format: txid=0, flags=0.
|
||||||
|
f, _ := dns.NewFrame(buf[:n])
|
||||||
|
if f.TxID() != 0 {
|
||||||
|
t.Errorf("mDNS query txid=%d, want 0", f.TxID())
|
||||||
|
}
|
||||||
|
if f.Flags() != 0 {
|
||||||
|
t.Errorf("mDNS query flags=%d, want 0", f.Flags())
|
||||||
|
}
|
||||||
|
if f.QDCount() != 1 {
|
||||||
|
t.Errorf("mDNS query QDCount=%d, want 1", f.QDCount())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Responder demuxes and encapsulates response.
|
||||||
|
if err = responder.Demux(buf[:n], 0); err != nil {
|
||||||
|
t.Fatal("responder demux:", err)
|
||||||
|
}
|
||||||
|
n, err = responder.Encapsulate(buf[:], -1, 0)
|
||||||
|
if err != nil || n == 0 {
|
||||||
|
t.Fatal("responder encapsulate:", err, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify response wire format: QR=1, AA=1.
|
||||||
|
f, _ = dns.NewFrame(buf[:n])
|
||||||
|
flags := f.Flags()
|
||||||
|
if !flags.IsResponse() {
|
||||||
|
t.Error("mDNS response missing QR bit")
|
||||||
|
}
|
||||||
|
if !flags.IsAuthorativeAnswer() {
|
||||||
|
t.Error("mDNS response missing AA bit")
|
||||||
|
}
|
||||||
|
if f.ANCount() == 0 {
|
||||||
|
t.Fatal("mDNS response has 0 answers")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Querier demuxes response and reads answers.
|
||||||
|
if err = querier.Demux(buf[:n], 0); err != nil {
|
||||||
|
t.Fatal("querier demux:", err)
|
||||||
|
}
|
||||||
|
var answers [8]dns.Resource
|
||||||
|
nans, done, err := querier.AnswersCopyTo(answers[:])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
} else if !done {
|
||||||
|
t.Fatal("expected done")
|
||||||
|
} else if nans == 0 {
|
||||||
|
t.Fatal("querier got 0 answers")
|
||||||
|
}
|
||||||
|
|
||||||
|
// The PTR answer data should contain the instance name in wire format.
|
||||||
|
ptrData := answers[0].RawData()
|
||||||
|
if len(ptrData) == 0 {
|
||||||
|
t.Fatal("PTR answer has no data")
|
||||||
|
}
|
||||||
|
var ptrName dns.Name
|
||||||
|
_, err = ptrName.Decode(ptrData, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("decode PTR name:", err)
|
||||||
|
}
|
||||||
|
if ptrName.String() != svc.Name.String() {
|
||||||
|
t.Errorf("PTR target=%q, want %q", ptrName.String(), svc.Name.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientQuerySRV(t *testing.T) {
|
||||||
|
svc := testService()
|
||||||
|
responder := newResponder(t, []Service{svc})
|
||||||
|
querier := newQuerier(t, []dns.Question{{
|
||||||
|
Name: mustNewName("My Web._http._tcp.local"),
|
||||||
|
Type: dns.TypeSRV,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}}, 8)
|
||||||
|
|
||||||
|
var buf [1024]byte
|
||||||
|
queryRespond(t, querier, responder, buf[:])
|
||||||
|
|
||||||
|
var answers [8]dns.Resource
|
||||||
|
nans, done, err := querier.AnswersCopyTo(answers[:])
|
||||||
|
if err != nil || !done {
|
||||||
|
t.Fatal("expected done without error:", err)
|
||||||
|
}
|
||||||
|
// SRV query should return SRV + A record.
|
||||||
|
if nans < 2 {
|
||||||
|
t.Fatalf("expected at least 2 answers (SRV+A), got %d", nans)
|
||||||
|
}
|
||||||
|
|
||||||
|
// First answer should be SRV. Parse priority(2)+weight(2)+port(2)+target.
|
||||||
|
srvData := answers[0].RawData()
|
||||||
|
if len(srvData) < 6 {
|
||||||
|
t.Fatalf("SRV data too short: %d bytes", len(srvData))
|
||||||
|
}
|
||||||
|
gotPort := binary.BigEndian.Uint16(srvData[4:6])
|
||||||
|
if gotPort != svc.Port {
|
||||||
|
t.Errorf("SRV port=%d, want %d", gotPort, svc.Port)
|
||||||
|
}
|
||||||
|
var srvTarget dns.Name
|
||||||
|
_, err = srvTarget.Decode(srvData, 6)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("decode SRV target:", err)
|
||||||
|
}
|
||||||
|
if srvTarget.String() != svc.Host.String() {
|
||||||
|
t.Errorf("SRV target=%q, want %q", srvTarget.String(), svc.Host.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Second answer should be A record with 4-byte IP.
|
||||||
|
aData := answers[1].RawData()
|
||||||
|
if len(aData) != 4 {
|
||||||
|
t.Fatalf("A record data length=%d, want 4", len(aData))
|
||||||
|
}
|
||||||
|
if [4]byte(aData) != [4]byte(svc.Addr) {
|
||||||
|
t.Errorf("A record addr=%v, want %v", aData, svc.Addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientQueryARecord(t *testing.T) {
|
||||||
|
svc := testService()
|
||||||
|
responder := newResponder(t, []Service{svc})
|
||||||
|
querier := newQuerier(t, []dns.Question{{
|
||||||
|
Name: mustNewName("mydevice.local"),
|
||||||
|
Type: dns.TypeA,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}}, 4)
|
||||||
|
|
||||||
|
var buf [1024]byte
|
||||||
|
queryRespond(t, querier, responder, buf[:])
|
||||||
|
|
||||||
|
var answers [4]dns.Resource
|
||||||
|
nans, done, err := querier.AnswersCopyTo(answers[:])
|
||||||
|
if err != nil || !done {
|
||||||
|
t.Fatal("expected done without error:", err)
|
||||||
|
}
|
||||||
|
if nans != 1 {
|
||||||
|
t.Fatalf("expected 1 answer, got %d", nans)
|
||||||
|
}
|
||||||
|
aData := answers[0].RawData()
|
||||||
|
if [4]byte(aData) != [4]byte(svc.Addr) {
|
||||||
|
t.Errorf("A record addr=%v, want %v", aData, svc.Addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: TestClientAnnounce — test unsolicited announcement of all registered services.
|
||||||
|
// func TestClientAnnounce(t *testing.T) { ... }
|
||||||
|
|
||||||
|
// TODO: TestClientProbeFinish — test probing sequence for name uniqueness (RFC 6762 §8.1).
|
||||||
|
// func TestClientProbeFinish(t *testing.T) { ... }
|
||||||
|
|
||||||
|
func TestClientIgnoresQueriesWithoutServices(t *testing.T) {
|
||||||
|
// A querier-only client should ignore incoming queries.
|
||||||
|
querier := newQuerier(t, []dns.Question{{
|
||||||
|
Name: mustNewName("_http._tcp.local"),
|
||||||
|
Type: dns.TypePTR,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}}, 4)
|
||||||
|
|
||||||
|
// Encapsulate query, then feed it back — should be ignored (not a response).
|
||||||
|
var buf [512]byte
|
||||||
|
n, _ := querier.Encapsulate(buf[:], -1, 0)
|
||||||
|
err := querier.Demux(buf[:n], 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("demux query:", err)
|
||||||
|
}
|
||||||
|
// No response should be pending.
|
||||||
|
n, _ = querier.Encapsulate(buf[:], -1, 0)
|
||||||
|
if n != 0 {
|
||||||
|
t.Error("querier without services should not respond to queries")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientResponderIgnoresResponses(t *testing.T) {
|
||||||
|
svc := testService()
|
||||||
|
responder := newResponder(t, []Service{svc})
|
||||||
|
|
||||||
|
// Build a response packet (QR=1).
|
||||||
|
var msg dns.Message
|
||||||
|
var buf [512]byte
|
||||||
|
const responseFlags = dns.HeaderFlags(1 << 15)
|
||||||
|
data, err := msg.AppendTo(buf[:0], 0, responseFlags)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Feed a response to the responder — should be ignored.
|
||||||
|
err = responder.Demux(data, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Nothing should be pending.
|
||||||
|
n, _ := responder.Encapsulate(buf[:], -1, 0)
|
||||||
|
if n != 0 {
|
||||||
|
t.Error("responder should not respond to responses")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientUnmatchedQuery(t *testing.T) {
|
||||||
|
svc := testService()
|
||||||
|
responder := newResponder(t, []Service{svc})
|
||||||
|
|
||||||
|
// Query for a name the responder doesn't know.
|
||||||
|
querier := newQuerier(t, []dns.Question{{
|
||||||
|
Name: mustNewName("_ftp._tcp.local"),
|
||||||
|
Type: dns.TypePTR,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}}, 4)
|
||||||
|
|
||||||
|
var buf [1024]byte
|
||||||
|
n, _ := querier.Encapsulate(buf[:], -1, 0)
|
||||||
|
responder.Demux(buf[:n], 0)
|
||||||
|
|
||||||
|
// Responder should have nothing pending.
|
||||||
|
n, _ = responder.Encapsulate(buf[:], -1, 0)
|
||||||
|
if n != 0 {
|
||||||
|
t.Error("responder should not respond to unmatched query")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientMultipleResponders(t *testing.T) {
|
||||||
|
// Two responder clients advertising different instances of the same service type.
|
||||||
|
svc1 := Service{
|
||||||
|
Name: mustNewName("Device A._http._tcp.local"),
|
||||||
|
Host: mustNewName("device-a.local"),
|
||||||
|
Addr: []byte{192, 168, 1, 10},
|
||||||
|
Port: 80,
|
||||||
|
}
|
||||||
|
svc2 := Service{
|
||||||
|
Name: mustNewName("Device B._http._tcp.local"),
|
||||||
|
Host: mustNewName("device-b.local"),
|
||||||
|
Addr: []byte{192, 168, 1, 11},
|
||||||
|
Port: 8080,
|
||||||
|
}
|
||||||
|
|
||||||
|
responder1 := newResponder(t, []Service{svc1})
|
||||||
|
responder2 := newResponder(t, []Service{svc2})
|
||||||
|
querier := newQuerier(t, []dns.Question{{
|
||||||
|
Name: mustNewName("_http._tcp.local"),
|
||||||
|
Type: dns.TypePTR,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}}, 8)
|
||||||
|
|
||||||
|
var buf [1024]byte
|
||||||
|
|
||||||
|
// Querier sends query.
|
||||||
|
n, _ := querier.Encapsulate(buf[:], -1, 0)
|
||||||
|
query := make([]byte, n)
|
||||||
|
copy(query, buf[:n])
|
||||||
|
|
||||||
|
// Both responders process the query.
|
||||||
|
responder1.Demux(query, 0)
|
||||||
|
responder2.Demux(query, 0)
|
||||||
|
|
||||||
|
// Querier receives response from responder 1.
|
||||||
|
n, _ = responder1.Encapsulate(buf[:], -1, 0)
|
||||||
|
querier.Demux(buf[:n], 0)
|
||||||
|
|
||||||
|
var answers [8]dns.Resource
|
||||||
|
nans, _, err := querier.AnswersCopyTo(answers[:])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if nans < 1 {
|
||||||
|
t.Fatalf("expected at least 1 answer from first responder, got %d", nans)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Start a new resolve to receive from responder 2.
|
||||||
|
querier.StartResolve(ResolveConfig{
|
||||||
|
Questions: []dns.Question{{
|
||||||
|
Name: mustNewName("_http._tcp.local"),
|
||||||
|
Type: dns.TypePTR,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}},
|
||||||
|
MaxResponseAnswers: 8,
|
||||||
|
})
|
||||||
|
// Re-send query for second responder.
|
||||||
|
n, _ = querier.Encapsulate(buf[:], -1, 0)
|
||||||
|
|
||||||
|
n, _ = responder2.Encapsulate(buf[:], -1, 0)
|
||||||
|
querier.Demux(buf[:n], 0)
|
||||||
|
|
||||||
|
nans, _, err = querier.AnswersCopyTo(answers[:])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if nans < 1 {
|
||||||
|
t.Fatalf("expected at least 1 answer from second responder, got %d", nans)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientAbort(t *testing.T) {
|
||||||
|
svc := testService()
|
||||||
|
responder := newResponder(t, []Service{svc})
|
||||||
|
|
||||||
|
responder.Abort()
|
||||||
|
|
||||||
|
// Demux on aborted client should return ErrClosed.
|
||||||
|
var buf [512]byte
|
||||||
|
var msg dns.Message
|
||||||
|
msg.AddQuestions([]dns.Question{{
|
||||||
|
Name: mustNewName("_http._tcp.local"),
|
||||||
|
Type: dns.TypePTR,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}})
|
||||||
|
data, _ := msg.AppendTo(buf[:0], 0, 0)
|
||||||
|
err := responder.Demux(data, 0)
|
||||||
|
if err != net.ErrClosed {
|
||||||
|
t.Errorf("expected net.ErrClosed, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Encapsulate on aborted client should return ErrClosed.
|
||||||
|
_, err = responder.Encapsulate(buf[:], -1, 0)
|
||||||
|
if err != net.ErrClosed {
|
||||||
|
t.Errorf("expected net.ErrClosed, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+32
-5
@@ -46,6 +46,8 @@ type StackAsync struct {
|
|||||||
ntpUDP internet.StackUDPPort
|
ntpUDP internet.StackUDPPort
|
||||||
ntp ntp.Client
|
ntp ntp.Client
|
||||||
|
|
||||||
|
userUDPs []internet.StackUDPPort
|
||||||
|
|
||||||
sysprec int8 // NTP system precision.
|
sysprec int8 // NTP system precision.
|
||||||
|
|
||||||
prng uint32
|
prng uint32
|
||||||
@@ -60,12 +62,16 @@ type StackConfig struct {
|
|||||||
StaticAddress netip.Addr
|
StaticAddress netip.Addr
|
||||||
DNSServer netip.Addr
|
DNSServer netip.Addr
|
||||||
NTPServer netip.Addr
|
NTPServer netip.Addr
|
||||||
|
RandSeed int64
|
||||||
Hostname string
|
Hostname string
|
||||||
MaxTCPConns int
|
MaxTCPConns int
|
||||||
RandSeed int64
|
MaxUDPConns int
|
||||||
HardwareAddress [6]byte
|
|
||||||
MTU uint16
|
|
||||||
EthernetTxCRC32Update func(crc uint32, b []byte) uint32
|
EthernetTxCRC32Update func(crc uint32, b []byte) uint32
|
||||||
|
|
||||||
|
HardwareAddress [6]byte
|
||||||
|
MTU uint16
|
||||||
|
// Accept multicast ethernet and IP packets. Needed for MDNS.
|
||||||
|
AcceptMulticast bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *StackAsync) Hostname() string {
|
func (s *StackAsync) Hostname() string {
|
||||||
@@ -120,18 +126,24 @@ func (s *StackAsync) Reset(cfg StackConfig) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
s.link.SetAcceptMulticast(cfg.AcceptMulticast)
|
||||||
const ipNodes = 2 // UDP, TCP ports.
|
const ipNodes = 2 // UDP, TCP ports.
|
||||||
err = s.ip.Reset(addr, ipNodes)
|
err = s.ip.Reset(addr, ipNodes)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
s.ip.SetAcceptMulticast(cfg.AcceptMulticast)
|
||||||
//
|
//
|
||||||
err = s.resetARP()
|
err = s.resetARP()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
const udpMaintenanceConns = 3 // DHCP, DNS, NTP.
|
udpConns := 3 + cfg.MaxUDPConns // DHCP, DNS, NTP + user-registered.
|
||||||
err = s.udps.ResetUDP(udpMaintenanceConns)
|
err = s.udps.ResetUDP(udpConns)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
internal.SliceReuse(&s.userUDPs, cfg.MaxUDPConns)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -322,6 +334,21 @@ func (s *StackAsync) RegisterListener(listener *tcp.Listener) (err error) {
|
|||||||
return s.tcps.Register(listener, nil)
|
return s.tcps.Register(listener, nil)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RegisterUDP registers a StackNode on a UDP port with the given remote address and port.
|
||||||
|
// The StackUDPPort wrapping is handled internally. The number of user-registered UDP ports
|
||||||
|
// is limited by [StackConfig.MaxUDPConns].
|
||||||
|
func (s *StackAsync) RegisterUDP(node internet.StackNode, remoteAddr []byte, remotePort uint16) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
idx := len(s.userUDPs)
|
||||||
|
if idx >= cap(s.userUDPs) {
|
||||||
|
return lneto.ErrBufferFull
|
||||||
|
}
|
||||||
|
s.userUDPs = s.userUDPs[:idx+1]
|
||||||
|
s.userUDPs[idx].SetStackNode(node, remoteAddr, remotePort)
|
||||||
|
return s.udps.Register(&s.userUDPs[idx])
|
||||||
|
}
|
||||||
|
|
||||||
var errNoDNSServer = errors.New("no DNS server- did DHCP complete? You can set a predetermined DNS server in Stack configuration")
|
var errNoDNSServer = errors.New("no DNS server- did DHCP complete? You can set a predetermined DNS server in Stack configuration")
|
||||||
|
|
||||||
func (s *StackAsync) StartLookupIP(host string) error {
|
func (s *StackAsync) StartLookupIP(host string) error {
|
||||||
|
|||||||
@@ -0,0 +1,332 @@
|
|||||||
|
package xnet
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/soypat/lneto/dns"
|
||||||
|
"github.com/soypat/lneto/ethernet"
|
||||||
|
"github.com/soypat/lneto/mdns"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMDNS_QueryResponse(t *testing.T) {
|
||||||
|
const MTU = 1500
|
||||||
|
svcName, err := dns.NewName("My Web._http._tcp.local")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
hostName, err := dns.NewName("mydevice.local")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
svcType, err := dns.NewName("_http._tcp.local")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
svc := mdns.Service{
|
||||||
|
Name: svcName,
|
||||||
|
Host: hostName,
|
||||||
|
Addr: []byte{192, 168, 1, 50},
|
||||||
|
Port: 80,
|
||||||
|
}
|
||||||
|
|
||||||
|
responderAddr := netip.AddrFrom4([4]byte{192, 168, 1, 50})
|
||||||
|
responderMAC := [6]byte{0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0x01}
|
||||||
|
querierAddr := netip.AddrFrom4([4]byte{192, 168, 1, 100})
|
||||||
|
querierMAC := [6]byte{0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0x02}
|
||||||
|
mcastAddr := []byte{224, 0, 0, 251}
|
||||||
|
|
||||||
|
// Setup responder stack with mDNS service.
|
||||||
|
responderStack := new(StackAsync)
|
||||||
|
err = responderStack.Reset(StackConfig{
|
||||||
|
Hostname: "responder",
|
||||||
|
RandSeed: 1234,
|
||||||
|
StaticAddress: responderAddr,
|
||||||
|
HardwareAddress: responderMAC,
|
||||||
|
MTU: MTU,
|
||||||
|
MaxUDPConns: 1,
|
||||||
|
AcceptMulticast: true,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("responder reset:", err)
|
||||||
|
}
|
||||||
|
responderStack.SetGateway6(querierMAC)
|
||||||
|
|
||||||
|
var responderClient mdns.Client
|
||||||
|
err = responderClient.Configure(mdns.ClientConfig{
|
||||||
|
LocalPort: mdns.Port,
|
||||||
|
Services: []mdns.Service{svc},
|
||||||
|
MulticastAddr: mcastAddr,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("responder configure:", err)
|
||||||
|
}
|
||||||
|
err = responderStack.RegisterUDP(&responderClient, mcastAddr, mdns.Port)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("responder register:", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Setup querier stack.
|
||||||
|
querierStack := new(StackAsync)
|
||||||
|
err = querierStack.Reset(StackConfig{
|
||||||
|
Hostname: "querier",
|
||||||
|
RandSeed: 5678,
|
||||||
|
StaticAddress: querierAddr,
|
||||||
|
HardwareAddress: querierMAC,
|
||||||
|
MTU: MTU,
|
||||||
|
MaxUDPConns: 1,
|
||||||
|
AcceptMulticast: true,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("querier reset:", err)
|
||||||
|
}
|
||||||
|
querierStack.SetGateway6(responderMAC)
|
||||||
|
|
||||||
|
var querierClient mdns.Client
|
||||||
|
err = querierClient.Configure(mdns.ClientConfig{
|
||||||
|
LocalPort: mdns.Port,
|
||||||
|
MulticastAddr: mcastAddr,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("querier configure:", err)
|
||||||
|
}
|
||||||
|
err = querierClient.StartResolve(mdns.ResolveConfig{
|
||||||
|
Questions: []dns.Question{{
|
||||||
|
Name: svcType,
|
||||||
|
Type: dns.TypePTR,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}},
|
||||||
|
MaxResponseAnswers: 4,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("start resolve:", err)
|
||||||
|
}
|
||||||
|
err = querierStack.RegisterUDP(&querierClient, mcastAddr, mdns.Port)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("querier register:", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
const carrierDataSize = MTU + ethernet.MaxOverheadSize
|
||||||
|
var buf [carrierDataSize]byte
|
||||||
|
|
||||||
|
// Querier encapsulates query through full stack (Ethernet+IP+UDP+mDNS).
|
||||||
|
n, err := querierStack.Encapsulate(buf[:], -1, 0)
|
||||||
|
if err != nil || n == 0 {
|
||||||
|
t.Fatal("querier encapsulate:", err, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify mDNS query wire format at DNS layer.
|
||||||
|
const ethHdrLen = 14
|
||||||
|
ipIHL := int(buf[ethHdrLen]&0x0f) * 4
|
||||||
|
dnsStart := ethHdrLen + ipIHL + 8
|
||||||
|
dnsFrame, err := dns.NewFrame(buf[dnsStart:n])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("parse query dns frame:", err)
|
||||||
|
}
|
||||||
|
if dnsFrame.TxID() != 0 {
|
||||||
|
t.Errorf("mDNS query txid=%d, want 0", dnsFrame.TxID())
|
||||||
|
}
|
||||||
|
if dnsFrame.Flags() != 0 {
|
||||||
|
t.Errorf("mDNS query flags=%d, want 0", dnsFrame.Flags())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Responder demuxes the query (multicast MAC+IP accepted via AcceptMulticast).
|
||||||
|
err = responderStack.Demux(buf[:n], 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("responder demux:", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Responder encapsulates response.
|
||||||
|
n, err = responderStack.Encapsulate(buf[:], -1, 0)
|
||||||
|
if err != nil || n == 0 {
|
||||||
|
t.Fatal("responder encapsulate:", err, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify response DNS flags.
|
||||||
|
ipIHL = int(buf[ethHdrLen]&0x0f) * 4
|
||||||
|
dnsStart = ethHdrLen + ipIHL + 8
|
||||||
|
dnsFrame, err = dns.NewFrame(buf[dnsStart:n])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("parse response dns frame:", err)
|
||||||
|
}
|
||||||
|
flags := dnsFrame.Flags()
|
||||||
|
if !flags.IsResponse() {
|
||||||
|
t.Error("mDNS response missing QR bit")
|
||||||
|
}
|
||||||
|
if !flags.IsAuthorativeAnswer() {
|
||||||
|
t.Error("mDNS response missing AA bit")
|
||||||
|
}
|
||||||
|
if dnsFrame.ANCount() == 0 {
|
||||||
|
t.Fatal("mDNS response has 0 answers")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Querier demuxes response.
|
||||||
|
err = querierStack.Demux(buf[:n], 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("querier demux:", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read answers.
|
||||||
|
var answers [4]dns.Resource
|
||||||
|
nans, done, err := querierClient.AnswersCopyTo(answers[:])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("answers:", err)
|
||||||
|
}
|
||||||
|
if !done {
|
||||||
|
t.Fatal("expected done")
|
||||||
|
}
|
||||||
|
if nans == 0 {
|
||||||
|
t.Fatal("got 0 answers")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify PTR answer points to our service instance name.
|
||||||
|
ptrData := answers[0].RawData()
|
||||||
|
var ptrTarget dns.Name
|
||||||
|
_, err = ptrTarget.Decode(ptrData, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("decode PTR target:", err)
|
||||||
|
}
|
||||||
|
if !dns.NamesEqual(ptrTarget, svcName) {
|
||||||
|
t.Errorf("PTR target=%q, want %q", ptrTarget.String(), svcName.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMDNS_SRVThroughStack(t *testing.T) {
|
||||||
|
const MTU = 1500
|
||||||
|
svcName, err := dns.NewName("My Web._http._tcp.local")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
hostName, err := dns.NewName("mydevice.local")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
svc := mdns.Service{
|
||||||
|
Name: svcName,
|
||||||
|
Host: hostName,
|
||||||
|
Addr: []byte{192, 168, 1, 50},
|
||||||
|
Port: 80,
|
||||||
|
}
|
||||||
|
mcastAddr := []byte{224, 0, 0, 251}
|
||||||
|
|
||||||
|
responderMAC := [6]byte{0x01, 0x02, 0x03, 0x04, 0x05, 0x01}
|
||||||
|
querierMAC := [6]byte{0x01, 0x02, 0x03, 0x04, 0x05, 0x02}
|
||||||
|
|
||||||
|
// Create responder.
|
||||||
|
responderStack, _ := newMDNSStack(t, "responder", 1111,
|
||||||
|
netip.AddrFrom4([4]byte{192, 168, 1, 50}), responderMAC, querierMAC,
|
||||||
|
mdns.ClientConfig{LocalPort: mdns.Port, Services: []mdns.Service{svc}, MulticastAddr: mcastAddr},
|
||||||
|
)
|
||||||
|
|
||||||
|
// Create querier.
|
||||||
|
querierStack, querierClient := newMDNSStack(t, "querier", 2222,
|
||||||
|
netip.AddrFrom4([4]byte{192, 168, 1, 100}), querierMAC, responderMAC,
|
||||||
|
mdns.ClientConfig{LocalPort: mdns.Port, MulticastAddr: mcastAddr},
|
||||||
|
)
|
||||||
|
err = querierClient.StartResolve(mdns.ResolveConfig{
|
||||||
|
Questions: []dns.Question{{
|
||||||
|
Name: svcName,
|
||||||
|
Type: dns.TypeSRV,
|
||||||
|
Class: dns.ClassINET,
|
||||||
|
}},
|
||||||
|
MaxResponseAnswers: 4,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("start resolve:", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Full round-trip through both stacks.
|
||||||
|
var buf [MTU + ethernet.MaxOverheadSize]byte
|
||||||
|
mdnsQueryRespond(t, querierStack, responderStack, buf[:])
|
||||||
|
|
||||||
|
var answers [4]dns.Resource
|
||||||
|
nans, done, err := querierClient.AnswersCopyTo(answers[:])
|
||||||
|
if err != nil || !done {
|
||||||
|
t.Fatal("expected done:", err)
|
||||||
|
}
|
||||||
|
if nans < 2 {
|
||||||
|
t.Fatalf("expected at least 2 answers (SRV+A), got %d", nans)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify SRV port.
|
||||||
|
srvData := answers[0].RawData()
|
||||||
|
if len(srvData) < 6 {
|
||||||
|
t.Fatalf("SRV data too short: %d", len(srvData))
|
||||||
|
}
|
||||||
|
gotPort := binary.BigEndian.Uint16(srvData[4:6])
|
||||||
|
if gotPort != svc.Port {
|
||||||
|
t.Errorf("SRV port=%d, want %d", gotPort, svc.Port)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify A record.
|
||||||
|
aData := answers[1].RawData()
|
||||||
|
if [4]byte(aData) != [4]byte(svc.Addr) {
|
||||||
|
t.Errorf("A record addr=%v, want %v", aData, svc.Addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// newMDNSStack creates a StackAsync with an mDNS client registered on its UDP ports.
|
||||||
|
func newMDNSStack(t *testing.T, hostname string, seed int64,
|
||||||
|
addr netip.Addr, mac, gatewayMAC [6]byte,
|
||||||
|
mdnsCfg mdns.ClientConfig,
|
||||||
|
) (*StackAsync, *mdns.Client) {
|
||||||
|
t.Helper()
|
||||||
|
const MTU = 1500
|
||||||
|
stack := new(StackAsync)
|
||||||
|
err := stack.Reset(StackConfig{
|
||||||
|
Hostname: hostname,
|
||||||
|
RandSeed: seed,
|
||||||
|
StaticAddress: addr,
|
||||||
|
HardwareAddress: mac,
|
||||||
|
MTU: MTU,
|
||||||
|
MaxUDPConns: 1,
|
||||||
|
AcceptMulticast: true,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(hostname, "reset:", err)
|
||||||
|
}
|
||||||
|
stack.SetGateway6(gatewayMAC)
|
||||||
|
|
||||||
|
var client mdns.Client
|
||||||
|
err = client.Configure(mdnsCfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(hostname, "mdns configure:", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = stack.RegisterUDP(&client, mdnsCfg.MulticastAddr, mdns.Port)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(hostname, "register udp:", err)
|
||||||
|
}
|
||||||
|
return stack, &client
|
||||||
|
}
|
||||||
|
|
||||||
|
// mdnsQueryRespond performs a full Ethernet+IP+UDP+mDNS query→response cycle
|
||||||
|
// between two stacks with AcceptMulticast enabled.
|
||||||
|
func mdnsQueryRespond(t *testing.T, querier, responder *StackAsync, buf []byte) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
// Querier encapsulates query.
|
||||||
|
n, err := querier.Encapsulate(buf, -1, 0)
|
||||||
|
if err != nil || n == 0 {
|
||||||
|
t.Fatal("querier encapsulate:", err, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Responder demuxes multicast query directly.
|
||||||
|
err = responder.Demux(buf[:n], 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("responder demux:", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Responder encapsulates response.
|
||||||
|
n, err = responder.Encapsulate(buf, -1, 0)
|
||||||
|
if err != nil || n == 0 {
|
||||||
|
t.Fatal("responder encapsulate:", err, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Querier demuxes multicast response.
|
||||||
|
err = querier.Demux(buf[:n], 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal("querier demux:", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user