diff --git a/dns/client.go b/dns/client.go index e0151d5..7613df5 100644 --- a/dns/client.go +++ b/dns/client.go @@ -34,7 +34,7 @@ func (sudp *Client) ConnectionID() *uint64 { return &sudp.connID } func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error { nd := len(cfg.Questions) if nd > math.MaxUint16 { - return lneto.ErrBufferFull + return lneto.ErrInvalidConfig } c.reset(localPort, txid, dnsSendQuery, cfg.EnableRecursion) c.msg.LimitResourceDecoding(uint16(nd), uint16(nd), 0, 0) diff --git a/dns/dns.go b/dns/dns.go index 025ff00..487331e 100644 --- a/dns/dns.go +++ b/dns/dns.go @@ -57,6 +57,13 @@ type Name struct { 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 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...) } -// 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. // If the message was not completely parsed due to LimitResourceDecoding, // 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 { - return 0, false, errBaseLen - } else if len(msg) > math.MaxUint16 { - return 0, false, errResTooLong - } - m.Reset() +// +// The slice memory is overwritten and capacity used as the limit of encoding. +// If the argument slice is nil it is skipped for decoding but does not prevent further decoding +// of other answers, authorities or additionals from being decoded. +func DecodeMessage(q *[]Question, answers, authorities, additionals *[]Resource, msg []byte) (_ uint16, incompleteButOK bool, err error) { hdr, err := NewFrame(msg) if err != nil { return 0, false, err } - nq := int(hdr.QDCount()) + qd := hdr.QDCount() + nq := int(qd) off := uint16(SizeHeader) // Return tooManyErr if found to flag to the caller that the message was // decoded but contained too many resources to decode completely. - var tooManyErr error switch { - case nq > cap(m.Questions): + case nq > caporzero(q): tooManyErr = errTooManyQuestions - case hdr.ANCount() > uint16(cap(m.Answers)): + case int(hdr.ANCount()) > caporzero(answers): tooManyErr = errTooManyAnswers - case hdr.NSCount() > uint16(cap(m.Authorities)): + case int(hdr.NSCount()) > caporzero(authorities): tooManyErr = errTooManyAuthorities - case hdr.ARCount() > uint16(cap(m.Additionals)): + case int(hdr.ARCount()) > caporzero(additionals): tooManyErr = errTooManyAdditionals } - if nq > cap(m.Questions) { - nq = cap(m.Questions) - } - 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 + if q != nil { + if nq > cap(*q) { + nq = cap(*q) } + *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. - qd := hdr.QDCount() for i := 0; i < int(qd)-nq; i++ { off, err = skipQuestion(msg, off) if err != nil { return off, false, err } } - - off, err = decodeToCapResources(&m.Answers, msg, hdr.ANCount(), off) + off, err = decodeToCapResources(answers, msg, hdr.ANCount(), off) if err != nil { return off, false, err } - off, err = decodeToCapResources(&m.Authorities, msg, hdr.NSCount(), off) + off, err = decodeToCapResources(authorities, msg, hdr.NSCount(), off) if err != nil { return off, false, err } - off, err = decodeToCapResources(&m.Additionals, msg, hdr.ARCount(), off) + off, err = decodeToCapResources(additionals, msg, hdr.ARCount(), off) if err != nil { return off, false, err } 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) { originalRec := nrec - 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++ { - off, err = (*dst)[i].Decode(msg, off) - if err != nil { - *dst = (*dst)[:i] // Trim non-decoded/failed resources. - return off, err + if dst != nil { + 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++ { + off, err = (*dst)[i].Decode(msg, off) + if err != nil { + *dst = (*dst)[:i] // Trim non-decoded/failed resources. + return off, err + } } } // Parse undecoded resources, effectively skipping them. @@ -468,6 +483,20 @@ func NewName(domain string) (Name, error) { 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. func (n *Name) Len() uint16 { if len(n.data) > math.MaxUint16 { @@ -666,6 +695,48 @@ func (dst *Resource) CopyFrom(r Resource) { 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) { dst.Name.CopyFrom(rh.Name) dst.Type = rh.Type @@ -673,3 +744,10 @@ func (dst *ResourceHeader) CopyFrom(rh ResourceHeader) { dst.TTL = rh.TTL dst.Length = rh.Length } + +func caporzero[T any](v *[]T) int { + if v == nil { + return 0 + } + return cap(*v) +} diff --git a/examples/xcurl/main.go b/examples/xcurl/main.go index de34c4c..80192db 100644 --- a/examples/xcurl/main.go +++ b/examples/xcurl/main.go @@ -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()) var stack xnet.StackAsync + err = stack.Reset(xnet.StackConfig{ Hostname: "xnet-test", RandSeed: softRand, @@ -221,6 +222,7 @@ func run() (err error) { return fmt.Errorf("DHCP failed: %w", err) } timeDHCP() + err = stack.AssimilateDHCPResults(results) if err != nil { return fmt.Errorf("assimilating DHCP results: %w", err) diff --git a/http/httpraw/parse.go b/http/httpraw/parse.go index 92b46e0..5243764 100644 --- a/http/httpraw/parse.go +++ b/http/httpraw/parse.go @@ -295,12 +295,14 @@ func (h *Header) reuseOrAppend(tok headerSlice, value string) headerSlice { } func (h *Header) appendSlice(value string) headerSlice { + debuglog("http:appendslice:start") free := h.hbuf.free() if len(value) > free { if h.flags.hasAny(flagNoBufferGrow) { h.flags |= flagOOMReached 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.flags |= flagMangledBuffer diff --git a/internal/ip.go b/internal/ip.go index e1bdc35..42e354c 100644 --- a/internal/ip.go +++ b/internal/ip.go @@ -52,3 +52,41 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) { copy(dstaddr, dst) 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 +} diff --git a/internet/pcap/capture.go b/internet/pcap/capture.go index 0c2259f..7826c7c 100644 --- a/internet/pcap/capture.go +++ b/internet/pcap/capture.go @@ -19,6 +19,7 @@ import ( "github.com/soypat/lneto/ipv4" "github.com/soypat/lneto/ipv4/icmpv4" "github.com/soypat/lneto/ipv6" + "github.com/soypat/lneto/mdns" "github.com/soypat/lneto/ntp" "github.com/soypat/lneto/tcp" "github.com/soypat/lneto/udp" @@ -359,7 +360,7 @@ func (pc *PacketBreakdown) CaptureUDP(dst []Frame, pkt []byte, bitOffset int) ([ srcport := ufrm.SourcePort() if dhcpv4.PayloadIsDHCPv4(payload) { 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) } else if dstport == ntp.ServerPort || srcport == ntp.ServerPort { dst, err = pc.CaptureNTP(dst, pkt, end) @@ -437,7 +438,7 @@ func (pc *PacketBreakdown) CaptureDNS(dst []Frame, pkt []byte, bitOffset int) ([ return dst, errNotByteAligned } 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) if err != nil && !incomplete { 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{ { Name: "Mode", diff --git a/internet/stack-ethernet.go b/internet/stack-ethernet.go index 2dd0e45..b45b2fd 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -33,11 +33,12 @@ type StackEthernetConfig struct { } type StackEthernet struct { - connID uint64 - handlers handlers - mac [6]byte - gwmac [6]byte - mtu uint16 + connID uint64 + handlers handlers + mac [6]byte + gwmac [6]byte + mtu uint16 + acceptMulticast bool // crcupdate set when crc32 has been configured to be appended. crcupdate func(crc uint32, p []byte) uint32 } @@ -50,6 +51,10 @@ func (ls *StackEthernet) Gateway6() (gw [6]byte) { return ls.gwmac } +func (ls *StackEthernet) SetAcceptMulticast(accept bool) { + ls.acceptMulticast = accept +} + func (ls *StackEthernet) SetHardwareAddr6(mac [6]byte) { ls.mac = mac } @@ -83,11 +88,12 @@ func (ls *StackEthernet) Configure(cfg StackEthernetConfig) error { } ls.handlers.reset("StackEthernet", cfg.MaxNodes) *ls = StackEthernet{ - connID: ls.connID + 1, - handlers: ls.handlers, - mac: cfg.MAC, - gwmac: cfg.Gateway, - mtu: uint16(cfg.MTU), + connID: ls.connID + 1, + handlers: ls.handlers, + mac: cfg.MAC, + gwmac: cfg.Gateway, + mtu: uint16(cfg.MTU), + acceptMulticast: ls.acceptMulticast, } if cfg.AppendCRC32 { ls.crcupdate = cfg.CRC32Update @@ -121,7 +127,9 @@ func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error) dstaddr := efrm.DestinationHardwareAddr() var vld lneto.Validator if !efrm.IsBroadcast() && ls.mac != *dstaddr { - goto DROP + if !ls.acceptMulticast || dstaddr[0]&1 == 0 { + goto DROP + } } efrm.ValidateSize(&vld) if vld.HasError() { diff --git a/internet/stack-ip.go b/internet/stack-ip.go index 3029d17..c67966f 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -16,11 +16,12 @@ import ( var _ StackNode = (*StackIP)(nil) type StackIP struct { - connID uint64 - ipID uint16 - ip [4]byte - validator lneto.Validator - handlers handlers + connID uint64 + ipID uint16 + ip [4]byte + acceptMulticast bool + validator lneto.Validator + handlers handlers } 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 = StackIP{ - connID: sb.connID + 1, - validator: sb.validator, - handlers: sb.handlers, - ip: sb.ip, + connID: sb.connID + 1, + validator: sb.validator, + handlers: sb.handlers, + ip: sb.ip, + acceptMulticast: sb.acceptMulticast, } return nil } @@ -65,6 +67,10 @@ func (sb *StackIP) Addr() netip.Addr { return netip.AddrFrom4(sb.ip) } +func (sb *StackIP) SetAcceptMulticast(accept bool) { + sb.acceptMulticast = accept +} + func (sb *StackIP) SetLogger(logger *slog.Logger) { sb.handlers.log = logger } @@ -79,8 +85,10 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { } dst := ifrm.DestinationAddr() if sb.ip != ([4]byte{}) && *dst != sb.ip { - sb.handlers.debug("ip:not-for-us") - return lneto.ErrPacketDrop // Not meant for us. + if !sb.acceptMulticast || dst[0]&0xF0 != 0xE0 { + sb.handlers.debug("ip:not-for-us") + return lneto.ErrPacketDrop // Not meant for us. + } } sb.validator.ResetErr() diff --git a/ipv4/definitions.go b/ipv4/definitions.go index 266fef2..087c4e2 100644 --- a/ipv4/definitions.go +++ b/ipv4/definitions.go @@ -6,6 +6,14 @@ const ( 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. type ToS uint8 diff --git a/mdns/client.go b/mdns/client.go new file mode 100644 index 0000000..65bcba4 --- /dev/null +++ b/mdns/client.go @@ -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 +} diff --git a/mdns/definitions.go b/mdns/definitions.go new file mode 100644 index 0000000..d6ca4ae --- /dev/null +++ b/mdns/definitions.go @@ -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) +} diff --git a/mdns/mdns_test.go b/mdns/mdns_test.go new file mode 100644 index 0000000..3c9fcd2 --- /dev/null +++ b/mdns/mdns_test.go @@ -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) + } +} diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 499f2e3..2255348 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -46,6 +46,8 @@ type StackAsync struct { ntpUDP internet.StackUDPPort ntp ntp.Client + userUDPs []internet.StackUDPPort + sysprec int8 // NTP system precision. prng uint32 @@ -60,12 +62,16 @@ type StackConfig struct { StaticAddress netip.Addr DNSServer netip.Addr NTPServer netip.Addr + RandSeed int64 Hostname string MaxTCPConns int - RandSeed int64 - HardwareAddress [6]byte - MTU uint16 + MaxUDPConns int 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 { @@ -120,18 +126,24 @@ func (s *StackAsync) Reset(cfg StackConfig) error { if err != nil { return err } + s.link.SetAcceptMulticast(cfg.AcceptMulticast) const ipNodes = 2 // UDP, TCP ports. err = s.ip.Reset(addr, ipNodes) if err != nil { return err } + s.ip.SetAcceptMulticast(cfg.AcceptMulticast) // err = s.resetARP() if err != nil { return err } - const udpMaintenanceConns = 3 // DHCP, DNS, NTP. - err = s.udps.ResetUDP(udpMaintenanceConns) + udpConns := 3 + cfg.MaxUDPConns // DHCP, DNS, NTP + user-registered. + err = s.udps.ResetUDP(udpConns) + if err != nil { + return err + } + internal.SliceReuse(&s.userUDPs, cfg.MaxUDPConns) if err != nil { return err } @@ -322,6 +334,21 @@ func (s *StackAsync) RegisterListener(listener *tcp.Listener) (err error) { 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") func (s *StackAsync) StartLookupIP(host string) error { diff --git a/x/xnet/xnet_mdns_test.go b/x/xnet/xnet_mdns_test.go new file mode 100644 index 0000000..e9d68dd --- /dev/null +++ b/x/xnet/xnet_mdns_test.go @@ -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) + } +}