diff --git a/core/include/join/protocol.hpp b/core/include/join/protocol.hpp index 5cd4e5d1..ee2d53f3 100644 --- a/core/include/join/protocol.hpp +++ b/core/include/join/protocol.hpp @@ -60,6 +60,8 @@ namespace join template class BasicDatagramNameServer; + template + class BasicDatagramPeer; template class BasicHttpClient; @@ -779,6 +781,115 @@ namespace join return !(a == b); } + /** + * @brief Multicast DNS protocol class + */ + class Mdns + { + public: + using Endpoint = BasicInternetEndpoint; + using Socket = BasicDatagramSocket; + using Peer = BasicDatagramPeer; + + /** + * @brief construct the mDNS protocol instance. + * @param family IP address family. + */ + constexpr Mdns (int family = AF_INET) noexcept + : _family (family) + { + } + + /** + * @brief get protocol suitable for IPv4 address family. + * @return an IPv4 address family suitable protocol. + */ + static inline Mdns& v4 () noexcept + { + static Mdns mdnsv4 (AF_INET); + return mdnsv4; + } + + /** + * @brief get protocol suitable for IPv6 address family. + * @return an IPv6 address family suitable protocol. + */ + static inline Mdns& v6 () noexcept + { + static Mdns mdnsv6 (AF_INET6); + return mdnsv6; + } + + /** + * @brief get the protocol IP address family. + * @return the protocol IP address family. + */ + constexpr int family () const noexcept + { + return _family; + } + + /** + * @brief get the protocol communication semantic. + * @return the protocol communication semantic. + */ + constexpr int type () const noexcept + { + return SOCK_DGRAM; + } + + /** + * @brief get the protocol type. + * @return the protocol type. + */ + constexpr int protocol () const noexcept + { + return IPPROTO_UDP; + } + + /** + * @brief get multicast address for the given address family. + * @param family IP address family. + * @return multicast IP address. + */ + static IpAddress multicastAddress (int family) noexcept + { + return (family == AF_INET6) ? "ff02::fb" : "224.0.0.251"; + } + + /// default DNS port. + static constexpr uint16_t defaultPort = 5353; + + /// maximum DNS message size. + static constexpr size_t maxMsgSize = 8192; + + private: + /// IP address family. + int _family; + }; + + /** + * @brief check if equals. + * @param a protocol to check. + * @param b protocol to check. + * @return true if equals. + */ + constexpr bool operator== (const Mdns& a, const Mdns& b) noexcept + { + return a.family () == b.family (); + } + + /** + * @brief check if not equals. + * @param a protocol to check. + * @param b protocol to check. + * @return true if not equals. + */ + constexpr bool operator!= (const Mdns& a, const Mdns& b) noexcept + { + return !(a == b); + } + /** * @brief DNS over TLS protocol class. */ diff --git a/core/tests/protocol_test.cpp b/core/tests/protocol_test.cpp index 550fc006..42e13f25 100644 --- a/core/tests/protocol_test.cpp +++ b/core/tests/protocol_test.cpp @@ -36,6 +36,7 @@ using join::Icmp; using join::Tcp; using join::Tls; using join::Dns; +using join::Mdns; using join::Dot; using join::Http; using join::Https; @@ -66,6 +67,9 @@ TEST (Protocol, family) ASSERT_EQ (Dns ().family (), AF_INET); ASSERT_EQ (Dns::v6 ().family (), AF_INET6); ASSERT_EQ (Dns::v4 ().family (), AF_INET); + ASSERT_EQ (Mdns ().family (), AF_INET); + ASSERT_EQ (Mdns::v6 ().family (), AF_INET6); + ASSERT_EQ (Mdns::v4 ().family (), AF_INET); ASSERT_EQ (Dot ().family (), AF_INET); ASSERT_EQ (Dot::v6 ().family (), AF_INET6); ASSERT_EQ (Dot::v4 ().family (), AF_INET); @@ -97,6 +101,7 @@ TEST (Protocol, type) ASSERT_EQ (Tcp ().type (), SOCK_STREAM); ASSERT_EQ (Tls ().type (), SOCK_STREAM); ASSERT_EQ (Dns ().type (), SOCK_DGRAM); + ASSERT_EQ (Mdns ().type (), SOCK_DGRAM); ASSERT_EQ (Dot ().type (), SOCK_STREAM); ASSERT_EQ (Http ().type (), SOCK_STREAM); ASSERT_EQ (Https ().type (), SOCK_STREAM); @@ -119,6 +124,7 @@ TEST (Protocol, protocol) ASSERT_EQ (Tcp ().protocol (), IPPROTO_TCP); ASSERT_EQ (Tls ().protocol (), IPPROTO_TCP); ASSERT_EQ (Dns ().protocol (), IPPROTO_UDP); + ASSERT_EQ (Mdns ().protocol (), IPPROTO_UDP); ASSERT_EQ (Dot ().protocol (), IPPROTO_TCP); ASSERT_EQ (Http ().protocol (), IPPROTO_TCP); ASSERT_EQ (Https ().protocol (), IPPROTO_TCP); @@ -159,6 +165,11 @@ TEST (Protocol, equal) ASSERT_EQ (Dns::v6 (), Dns::v6 ()); ASSERT_NE (Dns::v6 (), Dns::v4 ()); + ASSERT_EQ (Mdns::v4 (), Mdns::v4 ()); + ASSERT_NE (Mdns::v4 (), Mdns::v6 ()); + ASSERT_EQ (Mdns::v6 (), Mdns::v6 ()); + ASSERT_NE (Mdns::v6 (), Mdns::v4 ()); + ASSERT_EQ (Dot::v4 (), Dot::v4 ()); ASSERT_NE (Dot::v4 (), Dot::v6 ()); ASSERT_EQ (Dot::v6 (), Dot::v6 ()); diff --git a/fabric/include/join/arp.hpp b/fabric/include/join/arp.hpp index ec019b78..6eef60c1 100644 --- a/fabric/include/join/arp.hpp +++ b/fabric/include/join/arp.hpp @@ -204,17 +204,45 @@ namespace join ScopedLock lock (_syncMutex); + _reactor->addHandler (handle (), this); + + uint32_t tip; + ::memcpy (&tip, &out.arp.ar_tip, sizeof (tip)); + auto inserted = _pending.emplace (tip, std::make_unique ()); + if (!inserted.second) + { + // LCOV_EXCL_START + _reactor->delHandler (handle ()); + close (); + lastError = make_error_code (Errc::OperationFailed); + return {}; + // LCOV_EXCL_STOP + } + if (write (reinterpret_cast (&out), sizeof (Packet)) == -1) { + // LCOV_EXCL_START + _pending.erase (inserted.first); + _reactor->delHandler (handle ()); + close (); + return {}; + // LCOV_EXCL_STOP + } + + if (!inserted.first->second->cond.timedWait (lock, timeout)) + { + _pending.erase (inserted.first); + _reactor->delHandler (handle ()); close (); + lastError = std::make_error_code (std::errc::no_such_device_or_address); return {}; } - _reactor->addHandler (handle (), this); - MacAddress mac = waitResponse (lock, out.arp.ar_tip, timeout); + MacAddress mac = inserted.first->second->mac; + _pending.erase (inserted.first); _reactor->delHandler (handle ()); - close (); + return mac; } @@ -326,38 +354,6 @@ namespace join ArpPacket arp; }; - /** - * @brief wait for ARP response. - * @param lock mutex previously locked by the calling thread. - * @param tip target IP. - * @param timeout wait timeout. - * @return the MAC address. - */ - template - MacAddress waitResponse (ScopedLock& lock, uint32_t tip, std::chrono::duration timeout) - { - auto inserted = _pending.emplace (tip, std::make_unique ()); - if (!inserted.second) - { - // LCOV_EXCL_START - lastError = make_error_code (Errc::OperationFailed); - return {}; - // LCOV_EXCL_STOP - } - - if (!inserted.first->second->cond.timedWait (lock, timeout)) - { - _pending.erase (inserted.first); - lastError = std::make_error_code (std::errc::no_such_device_or_address); - return {}; - } - - MacAddress mac = inserted.first->second->mac; - _pending.erase (inserted.first); - - return mac; - } - /** * @brief method called when data are ready to be read. * @param fd file descriptor. diff --git a/fabric/include/join/nameserver.hpp b/fabric/include/join/nameserver.hpp index 1cb1f5e8..cd2017a8 100644 --- a/fabric/include/join/nameserver.hpp +++ b/fabric/include/join/nameserver.hpp @@ -27,7 +27,8 @@ // libjoin. #include -#include +#include +#include namespace join { @@ -46,7 +47,8 @@ namespace join * @param reactor event loop reactor. */ explicit BasicDatagramNameServer (Reactor* reactor = nullptr) - : _reactor (reactor ? reactor : ReactorThread::reactor ()) + : Socket () + , _reactor (reactor ? reactor : ReactorThread::reactor ()) , _buffer (std::make_unique (Protocol::maxMsgSize)) { } @@ -121,16 +123,62 @@ namespace join const std::vector& authorities = {}, const std::vector& additionals = {}, uint16_t rcode = 0) { - DnsPacket response{}; - response.id = query.id; - response.flags = (uint16_t (1) << 15) | (query.flags & 0x7800) | (uint16_t (1) << 10) | (rcode & 0x000F); - response.questions = query.questions; - response.answers = answers; - response.authorities = authorities; - response.additionals = additionals; + DnsPacket packet{}; + packet.id = query.id; + packet.flags = (uint16_t (1) << 15) | (query.flags & 0x7800) | (uint16_t (1) << 10) | (rcode & 0x000F); + packet.dest = query.src; + packet.port = query.port; + packet.questions = query.questions; + packet.answers = answers; + packet.authorities = authorities; + packet.additionals = additionals; + + return send (packet); + } + + /** + * @brief method called when a DNS query is received. + * @param packet parsed DNS query received. + */ + virtual void onQuery (const DnsPacket& packet) = 0; + + protected: + /** + * @brief method called when data are ready to be read on handle. + * @param fd file descriptor. + */ + virtual void onReceive ([[maybe_unused]] int fd) override + { + Endpoint from; + int size = this->readFrom (_buffer.get (), Protocol::maxMsgSize, &from); + if (size >= int (_headerSize)) + { + std::stringstream data; + data.rdbuf ()->pubsetbuf (_buffer.get (), size); + + DnsPacket packet; + _message.deserialize (packet, data); + packet.src = from.ip (); + packet.dest = this->localEndpoint ().ip (); + packet.port = from.port (); + + if ((packet.flags & 0x8000) == 0) + { + this->onQuery (packet); + } + } + } + + /** + * @brief serialize and send a DNS packet. + * @param packet DNS packet to send. + * @return 0 on success, -1 on error. + */ + int send (DnsPacket& packet) + { std::stringstream data; - if (_message.serialize (response, data) == -1) + if (_message.serialize (packet, data) == -1) { // LCOV_EXCL_START lastError = make_error_code (Errc::InvalidParam); @@ -147,59 +195,713 @@ namespace join // LCOV_EXCL_STOP } - Endpoint to (query.src, query.port); - if (this->writeTo (buffer.data (), buffer.size (), to) == -1) + if (this->writeTo (buffer.data (), buffer.size (), {packet.dest, packet.port}) == -1) { return -1; // LCOV_EXCL_LINE } return 0; + }; + + /// DNS message header size. + static constexpr size_t _headerSize = 12; + + /// DNS message codec. + DnsMessage _message; + + /// event loop reactor. + Reactor* _reactor; + + /// reception buffer. + std::unique_ptr _buffer; + }; + + /** + * @brief mDNS peer. + */ + template + class BasicDatagramPeer : public BasicDatagramNameServer + { + public: + using Socket = typename BasicDatagramNameServer::Socket; + using Endpoint = typename BasicDatagramNameServer::Endpoint; + + /// DNS notification callback type. + using DnsNotify = std::function; + + /// callback called when a lookup sequence succeed. + DnsNotify onSuccess; + + /// callback called when a lookup sequence failed. + DnsNotify onFailure; + + /** + * @brief construct the mDNS peer instance. + * @param ifindex interface index. + * @param reactor event loop reactor. + */ + explicit BasicDatagramPeer (unsigned int ifindex, Reactor* reactor = nullptr) + : BasicDatagramNameServer (reactor) +#ifdef DEBUG + , onSuccess (defaultOnSuccess) + , onFailure (defaultOnFailure) +#else + , onSuccess (nullptr) + , onFailure (nullptr) +#endif + , _ifindex (ifindex) + { + } + + /** + * @brief construct the mDNS peer instance. + * @param interface interface name. + * @param reactor event loop reactor. + */ + explicit BasicDatagramPeer (const std::string& interface, Reactor* reactor = nullptr) + : BasicDatagramPeer (if_nametoindex (interface.c_str ()), reactor) + { + } + + /** + * @brief copy constructor. + * @param other other object to copy. + */ + BasicDatagramPeer (const BasicDatagramPeer& other) = delete; + + /** + * @brief copy assignment operator. + * @param other other object to copy. + * @return a reference to the current object. + */ + BasicDatagramPeer& operator= (const BasicDatagramPeer& other) = delete; + + /** + * @brief move constructor. + * @param other other object to move. + */ + BasicDatagramPeer (BasicDatagramPeer&& other) = delete; + + /** + * @brief move assignment operator. + * @param other other object to move. + * @return a reference to the current object. + */ + BasicDatagramPeer& operator= (BasicDatagramPeer&& other) = delete; + + /** + * @brief destroy instance. + */ + virtual ~BasicDatagramPeer () noexcept = default; + + using BasicDatagramNameServer::bind; + + /** + * @brief bind the socket to specified address family. + * @param family address family. + * @return 0 on success, -1 on failure. + */ + int bind (int family) noexcept + { + IpAddress maddress = Protocol::multicastAddress (family); + Endpoint endpoint{IpAddress (family), Protocol::defaultPort}; + + if ((this->_state == Socket::State::Closed) && (this->open (endpoint.protocol ()) == -1)) + { + return -1; // LCOV_EXCL_LINE + } + + if (this->setOption (Socket::ReusePort, 1) == -1) + { + // LCOV_EXCL_START + this->close (); + return -1; + // LCOV_EXCL_STOP + } + + if (Socket::bind (endpoint) == -1) + { + // LCOV_EXCL_START + this->close (); + return -1; + // LCOV_EXCL_STOP + } + + if (endpoint.protocol ().family () == AF_INET6) + { + ipv6_mreq mreq{}; + ::memcpy (&mreq.ipv6mr_multiaddr, maddress.addr (), maddress.length ()); + mreq.ipv6mr_interface = _ifindex; + if (::setsockopt (this->handle (), IPPROTO_IPV6, IPV6_ADD_MEMBERSHIP, &mreq, sizeof (mreq)) == -1) + { + // LCOV_EXCL_START + lastError = std::error_code (errno, std::generic_category ()); + this->close (); + return -1; + // LCOV_EXCL_STOP + } + if (::setsockopt (this->handle (), IPPROTO_IPV6, IPV6_MULTICAST_IF, &_ifindex, sizeof (_ifindex)) == -1) + { + // LCOV_EXCL_START + lastError = std::error_code (errno, std::generic_category ()); + this->close (); + return -1; + // LCOV_EXCL_STOP + } + } + else + { + // LCOV_EXCL_START: IPv4 multicast not supported by github action containers. + ip_mreqn mreq{}; + ::memcpy (&mreq.imr_multiaddr, maddress.addr (), maddress.length ()); + mreq.imr_ifindex = static_cast (_ifindex); + if (::setsockopt (this->handle (), IPPROTO_IP, IP_ADD_MEMBERSHIP, &mreq, sizeof (mreq)) == -1) + { + lastError = std::error_code (errno, std::generic_category ()); + this->close (); + return -1; + } + if (::setsockopt (this->handle (), IPPROTO_IP, IP_MULTICAST_IF, &mreq, sizeof (mreq)) == -1) + { + lastError = std::error_code (errno, std::generic_category ()); + this->close (); + return -1; + } + // LCOV_EXCL_STOP + } + +#ifndef DEBUG + if (this->setOption (Socket::MulticastLoop, 0) == -1) + { + // LCOV_EXCL_START + this->close (); + return -1; + // LCOV_EXCL_STOP + } +#endif + + this->_reactor->addHandler (this->handle (), this); + + return 0; + } + + /** + * @brief probe the local network for the presence of a service. + * @param records resource records to query for. + * @return 0 on success, -1 on error. + */ + int probe (const std::vector& records) + { + if (records.empty ()) + { + lastError = make_error_code (Errc::InvalidParam); + return -1; + } + + DnsPacket packet{}; + packet.id = 0; + packet.flags = 0; + IpAddress mcast = Protocol::multicastAddress (this->family ()); + packet.dest = IpAddress (mcast.addr (), mcast.length (), _ifindex); + packet.port = Protocol::defaultPort; + + for (auto const& record : records) + { + QuestionRecord question; + question.host = record.host; + question.type = DnsMessage::RecordType::ANY; + question.dnsclass = DnsMessage::RecordClass::IN | 0x8000; + packet.questions.push_back (question); + + packet.authorities.push_back (record); + } + + return this->send (packet); + } + + /** + * @brief announce the presence of a service on the local network. + * @param records resource records to announce. + * @return 0 on success, -1 on error. + */ + int announce (const std::vector& records) + { + if (records.empty ()) + { + lastError = make_error_code (Errc::InvalidParam); + return -1; + } + + DnsPacket packet{}; + packet.id = 0; + packet.flags = (uint16_t (1) << 15) | (uint16_t (1) << 10); + IpAddress mcast = Protocol::multicastAddress (this->family ()); + packet.dest = IpAddress (mcast.addr (), mcast.length (), _ifindex); + packet.port = Protocol::defaultPort; + + for (auto const& record : records) + { + packet.answers.push_back (record); + } + + return this->send (packet); + } + + /** + * @brief send a goodbye message. + * @param records resource records to send in goodbye message. + * @return 0 on success, -1 on error. + */ + int goodbye (const std::vector& records) + { + if (records.empty ()) + { + lastError = make_error_code (Errc::InvalidParam); + return -1; + } + + DnsPacket packet{}; + packet.id = 0; + packet.flags = (uint16_t (1) << 15) | (uint16_t (1) << 10); + IpAddress mcast = Protocol::multicastAddress (this->family ()); + packet.dest = IpAddress (mcast.addr (), mcast.length (), _ifindex); + packet.port = Protocol::defaultPort; + + for (auto const& record : records) + { + ResourceRecord goodbye = record; + goodbye.ttl = 0; + packet.answers.push_back (goodbye); + } + + return this->send (packet); + } + + /** + * @brief browse for services on the local network. + * @param serviceType service type to browse for (e.g. "_http._tcp.local"). + * @return 0 on success, -1 on error. + */ + int browse (const std::string& serviceType) + { + if (serviceType.empty ()) + { + lastError = make_error_code (Errc::InvalidParam); + return -1; + } + + DnsPacket packet{}; + packet.id = 0; + packet.flags = 0; + IpAddress mcast = Protocol::multicastAddress (this->family ()); + packet.dest = IpAddress (mcast.addr (), mcast.length (), _ifindex); + packet.port = Protocol::defaultPort; + + QuestionRecord question; + question.host = serviceType; + question.type = DnsMessage::RecordType::PTR; + question.dnsclass = DnsMessage::RecordClass::IN; + packet.questions.push_back (question); + + return this->send (packet); + } + + /** + * @brief resolve host name and return all IP addresses found. + * @param host host name to resolve. + * @param family address family. + * @param timeout timeout in milliseconds (default: 5000). + * @return the resolved IP address list. + */ + IpAddressList resolveAllAddress (const std::string& host, int family, + std::chrono::milliseconds timeout = std::chrono::seconds (5)) + { + if (host.empty ()) + { + return {}; + } + + DnsPacket packet{}; + packet.id = join::randomize (); + packet.flags = 1 << 8; + + QuestionRecord question; + question.host = host; + question.type = (family == AF_INET6) ? DnsMessage::RecordType::AAAA : DnsMessage::RecordType::A; + question.dnsclass = DnsMessage::RecordClass::IN; + packet.questions.push_back (question); + + if (query (packet, timeout) == -1) + { + return {}; + } + + IpAddressList addresses; + + for (auto const& answer : packet.answers) + { + if (!answer.addr.isWildcard () && (answer.type == question.type)) + { + addresses.push_back (answer.addr); + } + } + + return addresses; + } + + /** + * @brief resolve host name and return all IP addresses found. + * @param host host name to resolve. + * @param timeout timeout in milliseconds (default: 5000). + * @return the resolved IP address list. + */ + IpAddressList resolveAllAddress (const std::string& host, + std::chrono::milliseconds timeout = std::chrono::seconds (5)) + { + IpAddressList addresses; + + for (auto const& family : {AF_INET, AF_INET6}) + { + IpAddressList tmp = resolveAllAddress (host, family, timeout); + addresses.insert (addresses.end (), tmp.begin (), tmp.end ()); + } + + return addresses; + } + + /** + * @brief resolve host name using address family. + * @param host host name to resolve. + * @param family address family. + * @param timeout timeout in milliseconds (default: 5000). + * @return the first resolved IP address found matching address family. + */ + IpAddress resolveAddress (const std::string& host, int family, + std::chrono::milliseconds timeout = std::chrono::seconds (5)) + { + for (auto const& address : resolveAllAddress (host, family, timeout)) + { + return address; + } + + return IpAddress (family); + } + + /** + * @brief resolve host name. + * @param host host name to resolve. + * @param timeout timeout in milliseconds (default: 5000). + * @return the first resolved IP address found. + */ + IpAddress resolveAddress (const std::string& host, std::chrono::milliseconds timeout = std::chrono::seconds (5)) + { + for (auto const& address : resolveAllAddress (host, timeout)) + { + return address; + } + + return {}; + } + + /** + * @brief resolve all host address. + * @param address host address to resolve. + * @param timeout timeout in milliseconds (default: 5000). + * @return the resolved alias list. + */ + AliasList resolveAllName (const IpAddress& address, + std::chrono::milliseconds timeout = std::chrono::seconds (5)) + { + if (address.isWildcard ()) + { + return {}; + } + + DnsPacket packet{}; + packet.id = join::randomize (); + packet.flags = 1 << 8; + + QuestionRecord question; + question.host = address.toArpa (); + question.type = DnsMessage::RecordType::PTR; + question.dnsclass = DnsMessage::RecordClass::IN; + packet.questions.push_back (question); + + if (query (packet, timeout) == -1) + { + return {}; + } + + AliasList aliases; + + for (auto const& answer : packet.answers) + { + if (!answer.name.empty () && (answer.type == DnsMessage::RecordType::PTR)) + { + aliases.insert (answer.name); + } + } + + return aliases; + } + + /** + * @brief resolve host address. + * @param address host address to resolve. + * @param timeout timeout in milliseconds (default: 5000). + * @return the first resolved alias. + */ + std::string resolveName (const IpAddress& address, std::chrono::milliseconds timeout = std::chrono::seconds (5)) + { + for (auto const& alias : resolveAllName (address, timeout)) + { + return alias; + } + + return {}; } /** * @brief method called when a DNS query is received. * @param packet parsed DNS query received. */ - virtual void onQuery (const DnsPacket& packet) = 0; + virtual void onAnnouncement (const DnsPacket& packet) = 0; protected: /** * @brief method called when data are ready to be read on handle. * @param fd file descriptor. */ - virtual void onReceive ([[maybe_unused]] int fd) override + void onReceive ([[maybe_unused]] int fd) override final { Endpoint from; - int size = this->readFrom (_buffer.get (), Protocol::maxMsgSize, &from); - if (size >= int (_headerSize)) + int size = this->readFrom (this->_buffer.get (), Protocol::maxMsgSize, &from); + if (size >= int (this->_headerSize)) { std::stringstream data; - data.rdbuf ()->pubsetbuf (_buffer.get (), size); + data.rdbuf ()->pubsetbuf (this->_buffer.get (), size); DnsPacket packet; - _message.deserialize (packet, data); + this->_message.deserialize (packet, data); + IpAddress mcast = Protocol::multicastAddress (this->family ()); packet.src = from.ip (); + packet.dest = IpAddress (mcast.addr (), mcast.length (), _ifindex); packet.port = from.port (); - packet.dest = this->localEndpoint ().ip (); if ((packet.flags & 0x8000) == 0) { + bool unicast = false; + for (auto const& q : packet.questions) + { + if (q.dnsclass & 0x8000) + { + unicast = true; + break; + } + } + + if (!unicast) + { + packet.src = IpAddress (mcast.addr (), mcast.length (), _ifindex); + } + this->onQuery (packet); + return; + } + + { + ScopedLock lock (_syncMutex); + + auto it = _pending.find (packet.id); + if (it != _pending.end ()) + { + it->second->packet = packet; + it->second->ec = DnsMessage::decodeError (packet.flags & 0x000F); + it->second->cond.signal (); + return; + } } + + onAnnouncement (packet); } } - /// DNS message header size. - static constexpr size_t _headerSize = 12; +#ifdef DEBUG + /* + * @brief default callback called when a lookup sequence succeed. + * @param packet DNS packet. + */ + static void defaultOnSuccess (const DnsPacket& packet) + { + std::cout << std::endl; + std::cout << "PEER: " << packet.dest << "#" << packet.port << std::endl; - /// DNS message codec. - DnsMessage _message; + std::cout << std::endl; + std::cout << ";; QUESTION SECTION: " << std::endl; + for (auto const& question : packet.questions) + { + std::cout << question.host; + std::cout << " " << DnsMessage::typeName (question.type); + std::cout << " " << DnsMessage::className (question.dnsclass); + std::cout << std::endl; + } - /// event loop reactor. - Reactor* _reactor; + std::cout << std::endl; + std::cout << ";; ANSWER SECTION: " << std::endl; + for (auto const& answer : packet.answers) + { + std::cout << answer.host; + std::cout << " " << DnsMessage::typeName (answer.type); + std::cout << " " << DnsMessage::className (answer.dnsclass); + std::cout << " " << answer.ttl; + if (answer.type == DnsMessage::RecordType::A) + { + std::cout << " " << answer.addr; + } + else if (answer.type == DnsMessage::RecordType::PTR) + { + std::cout << " " << answer.name; + } + else if (answer.type == DnsMessage::RecordType::AAAA) + { + std::cout << " " << answer.addr; + } + std::cout << std::endl; + } + } - /// reception buffer. - std::unique_ptr _buffer; + /* + * @brief default callback called when a lookup sequence failed. + * @param packet DNS packet. + */ + static void defaultOnFailure (const DnsPacket& packet) + { + std::cout << std::endl; + std::cout << "PEER: " << packet.dest << "#" << packet.port << std::endl; + + std::cout << std::endl; + std::cout << ";; QUESTION SECTION: " << std::endl; + for (auto const& question : packet.questions) + { + std::cout << question.host; + std::cout << " " << DnsMessage::typeName (question.type); + std::cout << " " << DnsMessage::className (question.dnsclass); + std::cout << std::endl; + } + + std::cout << std::endl; + std::cout << lastError.message () << std::endl; + } +#endif + + /** + * @brief safe way to notify DNS events. + * @param func function to call. + * @param packet DNS packet. + */ + void notify (const DnsNotify& func, const DnsPacket& packet) const noexcept + { + if (func) + { + func (packet); + } + } + + /** + * @brief serialize and send a DNS query, waiting for a response. + * @param packet DNS packet to send, filled with the response on success + * @param timeout query timeout. + * @return 0 on success, -1 on error. + */ + int query (DnsPacket& packet, std::chrono::milliseconds timeout) + { + IpAddress mcast = Protocol::multicastAddress (this->family ()); + packet.dest = IpAddress (mcast.addr (), mcast.length (), _ifindex); + packet.port = Protocol::defaultPort; + + std::stringstream data; + if (this->_message.serialize (packet, data) == -1) + { + // LCOV_EXCL_START + lastError = make_error_code (Errc::InvalidParam); + return -1; + // LCOV_EXCL_STOP + } + + std::string buffer = data.str (); + if (buffer.size () > Protocol::maxMsgSize) + { + // LCOV_EXCL_START + lastError = make_error_code (Errc::MessageTooLong); + return -1; + // LCOV_EXCL_STOP + } + + ScopedLock lock (_syncMutex); + + auto inserted = _pending.emplace (packet.id, std::make_unique ()); + if (!inserted.second) + { + // LCOV_EXCL_START + lastError = make_error_code (Errc::OperationFailed); + notify (onFailure, packet); + return -1; + // LCOV_EXCL_STOP + } + + if (this->writeTo (buffer.data (), buffer.size (), {packet.dest, packet.port}) == -1) + { + // LCOV_EXCL_START + _pending.erase (inserted.first); + notify (onFailure, packet); + return -1; + // LCOV_EXCL_STOP + } + + if (!inserted.first->second->cond.timedWait (lock, timeout)) + { + // LCOV_EXCL_START + _pending.erase (inserted.first); + lastError = make_error_code (Errc::TimedOut); + notify (onFailure, packet); + return -1; + // LCOV_EXCL_STOP + } + + auto pendingReq = std::move (inserted.first->second); + _pending.erase (inserted.first); + + if (pendingReq->ec) + { + // LCOV_EXCL_START + lastError = pendingReq->ec; + notify (onFailure, packet); + return -1; + // LCOV_EXCL_STOP + } + + packet = std::move (pendingReq->packet); + notify (onSuccess, packet); + + return 0; + } + + /// interface index. + unsigned int _ifindex; + + /// pending synchronous request. + struct PendingRequest + { + Condition cond; /**< condition variable to signal response reception. */ + DnsPacket packet; /**< received response packet. */ + std::error_code ec; /**< error code from the response. */ + }; + + /// synchronous requests indexed by sequence number. + std::unordered_map> _pending; + + /// protection mutex. + Mutex _syncMutex; }; } diff --git a/fabric/include/join/netlinkmanager.hpp b/fabric/include/join/netlinkmanager.hpp index 6b4f7805..ce36670a 100644 --- a/fabric/include/join/netlinkmanager.hpp +++ b/fabric/include/join/netlinkmanager.hpp @@ -108,42 +108,6 @@ namespace join */ int sendRequest (struct nlmsghdr* nlh, bool sync, std::chrono::milliseconds timeout = std::chrono::seconds (5)); - /** - * @brief wait for specific netlink response. - * @param lock mutex previously locked by the calling thread. - * @param seq sequence number to wait for. - * @param timeout maximum wait duration. - * @return 0 on success, -1 on failure. - */ - template - int waitResponse (ScopedLock& lock, uint32_t seq, std::chrono::duration timeout) - { - auto inserted = _pending.emplace (seq, std::make_unique ()); - if (!inserted.second) - { - lastError = make_error_code (Errc::OperationFailed); - return -1; - } - - if (!inserted.first->second->cond.timedWait (lock, timeout)) - { - _pending.erase (inserted.first); - lastError = make_error_code (Errc::TimedOut); - return -1; - } - - if (inserted.first->second->error != 0) - { - int err = inserted.first->second->error; - _pending.erase (inserted.first); - lastError = std::error_code (err, std::generic_category ()); - return -1; - } - - _pending.erase (inserted.first); - return 0; - } - /** * @brief push a job to be executed on the reactor thread. * @param func function to execute on the reactor thread. diff --git a/fabric/include/join/resolver.hpp b/fabric/include/join/resolver.hpp index cedfd3d1..441d7413 100644 --- a/fabric/include/join/resolver.hpp +++ b/fabric/include/join/resolver.hpp @@ -54,6 +54,15 @@ namespace join using Endpoint = typename Protocol::Endpoint; using State = typename Socket::State; + /// notification callback definition. + using DnsNotify = std::function; + + /// callback called when a lookup sequence succeed. + DnsNotify onSuccess; + + /// callback called when a lookup sequence failed. + DnsNotify onFailure; + /** * @brief construct the resolver instance. * @param server remote DNS server hostname or IP address. @@ -64,11 +73,11 @@ namespace join Reactor* reactor = nullptr) : Socket () #ifdef DEBUG - , _onSuccess (defaultOnSuccess) - , _onFailure (defaultOnFailure) + , onSuccess (defaultOnSuccess) + , onFailure (defaultOnFailure) #else - , _onSuccess (nullptr) - , _onFailure (nullptr) + , onSuccess (nullptr) + , onFailure (nullptr) #endif , _server (server) , _port (port) @@ -702,15 +711,6 @@ namespace join return 0; } - /// notification callback definition. - using DnsNotify = std::function; - - /// callback called when a lookup sequence succeed. - DnsNotify _onSuccess; - - /// callback called when a lookup sequence failed. - DnsNotify _onFailure; - protected: /** * @brief check if client must reconnect. @@ -731,8 +731,10 @@ namespace join { if (this->disconnect () == -1) { + // LCOV_EXCL_START this->close (); return -1; + // LCOV_EXCL_STOP } if (this->connect (endpoint) == -1) @@ -760,7 +762,7 @@ namespace join if (ip.isWildcard ()) { lastError = make_error_code (Errc::InvalidParam); - notify (_onFailure, packet); + notify (onFailure, packet); return -1; } @@ -778,7 +780,7 @@ namespace join if (this->reconnect (endpoint, timeout) == -1) { - notify (_onFailure, packet); + notify (onFailure, packet); return -1; } } @@ -789,7 +791,7 @@ namespace join if (_message.serialize (packet, data) == -1) { lastError = make_error_code (Errc::InvalidParam); - notify (_onFailure, packet); + notify (onFailure, packet); return -1; } @@ -797,7 +799,7 @@ namespace join if (buffer.size () > Protocol::maxMsgSize) { lastError = make_error_code (Errc::MessageTooLong); - notify (_onFailure, packet); + notify (onFailure, packet); return -1; } @@ -806,23 +808,27 @@ namespace join auto inserted = _pending.emplace (packet.id, std::make_unique ()); if (!inserted.second) { + // LCOV_EXCL_START lastError = make_error_code (Errc::OperationFailed); - notify (_onFailure, packet); + notify (onFailure, packet); return -1; + // LCOV_EXCL_STOP } if (this->write (buffer.data (), buffer.size ()) == -1) { + // LCOV_EXCL_START _pending.erase (inserted.first); - notify (_onFailure, packet); + notify (onFailure, packet); return -1; + // LCOV_EXCL_STOP } if (!inserted.first->second->cond.timedWait (lock, timeout)) { _pending.erase (inserted.first); lastError = make_error_code (Errc::TimedOut); - notify (_onFailure, packet); + notify (onFailure, packet); return -1; } @@ -832,12 +838,12 @@ namespace join if (pendingReq->ec) { lastError = pendingReq->ec; - notify (_onFailure, packet); + notify (onFailure, packet); return -1; } packet = std::move (pendingReq->packet); - notify (_onSuccess, packet); + notify (onSuccess, packet); return 0; } @@ -1375,8 +1381,10 @@ namespace join { if (lastError != Errc::TemporaryError) { + // LCOV_EXCL_START this->close (); return -1; + // LCOV_EXCL_STOP } if (!this->waitDisconnected (timeout.count ())) diff --git a/fabric/src/netlinkmanager.cpp b/fabric/src/netlinkmanager.cpp index cc956567..00c00428 100644 --- a/fabric/src/netlinkmanager.cpp +++ b/fabric/src/netlinkmanager.cpp @@ -98,7 +98,7 @@ int NetlinkManager::sendRequest (struct nlmsghdr* nlh, bool sync, std::chrono::m { if (write (reinterpret_cast (nlh), nlh->nlmsg_len) == -1) { - return -1; + return -1; // LCOV_EXCL_LINE } return 0; @@ -106,12 +106,41 @@ int NetlinkManager::sendRequest (struct nlmsghdr* nlh, bool sync, std::chrono::m ScopedLock lock (_syncMutex); + auto inserted = _pending.emplace (nlh->nlmsg_seq, std::make_unique ()); + if (!inserted.second) + { + // LCOV_EXCL_START + lastError = make_error_code (Errc::OperationFailed); + return -1; + // LCOV_EXCL_STOP + } + if (write (reinterpret_cast (nlh), nlh->nlmsg_len) == -1) { + // LCOV_EXCL_START + _pending.erase (inserted.first); return -1; + // LCOV_EXCL_STOP } - return waitResponse (lock, nlh->nlmsg_seq, timeout); + if (!inserted.first->second->cond.timedWait (lock, timeout)) + { + _pending.erase (inserted.first); + lastError = make_error_code (Errc::TimedOut); + return -1; + } + + if (inserted.first->second->error != 0) + { + int err = inserted.first->second->error; + _pending.erase (inserted.first); + lastError = std::error_code (err, std::generic_category ()); + return -1; + } + + _pending.erase (inserted.first); + + return 0; } // ========================================================================= diff --git a/fabric/tests/CMakeLists.txt b/fabric/tests/CMakeLists.txt index 4b4cddf9..272b5235 100644 --- a/fabric/tests/CMakeLists.txt +++ b/fabric/tests/CMakeLists.txt @@ -12,6 +12,11 @@ target_link_libraries(dns.gtest ${JOIN_FABRIC} GTest::gtest_main) add_test(NAME dns.gtest COMMAND dns.gtest) install(TARGETS dns.gtest RUNTIME DESTINATION ${CMAKE_INSTALL_DATADIR}/${PROJECT_NAME}/test) +add_executable(mdns.gtest mdns_test.cpp) +target_link_libraries(mdns.gtest ${JOIN_FABRIC} GTest::gtest_main) +add_test(NAME mdns.gtest COMMAND mdns.gtest) +install(TARGETS mdns.gtest RUNTIME DESTINATION ${CMAKE_INSTALL_DATADIR}/${PROJECT_NAME}/test) + add_executable(dot.gtest dot_test.cpp) target_link_libraries(dot.gtest ${JOIN_FABRIC} GTest::gtest_main) add_test(NAME dot.gtest COMMAND dot.gtest) diff --git a/fabric/tests/dns_test.cpp b/fabric/tests/dns_test.cpp index 79923ff8..4b0734d6 100644 --- a/fabric/tests/dns_test.cpp +++ b/fabric/tests/dns_test.cpp @@ -347,6 +347,10 @@ TEST_F (DnsTest, resolveAllName) aliases = Dns::Resolver ("8.8.8.8", 53) .resolveAllName (Dns::Resolver::lookupAddress ("joinframework.net", AF_INET6), 1ms); EXPECT_EQ (aliases.size (), 0); + + aliases = + Dns::Resolver ("8.8.8.8", 53).resolveAllName (Dns::Resolver::lookupAddress ("www.joinframework.net", AF_INET)); + EXPECT_GT (aliases.size (), 0); } /** @@ -394,6 +398,9 @@ TEST_F (DnsTest, resolveName) alias = Dns::Resolver ("8.8.8.8", 53).resolveName (Dns::Resolver::lookupAddress ("joinframework.net", AF_INET6), 1ms); EXPECT_TRUE (alias.empty ()); + + alias = Dns::Resolver ("8.8.8.8", 53).resolveName (Dns::Resolver::lookupAddress ("www.joinframework.net", AF_INET)); + EXPECT_FALSE (alias.empty ()); } /** diff --git a/fabric/tests/mdns_test.cpp b/fabric/tests/mdns_test.cpp new file mode 100644 index 00000000..c958f146 --- /dev/null +++ b/fabric/tests/mdns_test.cpp @@ -0,0 +1,787 @@ +/** + * MIT License + * + * Copyright (c) 2026 Mathieu Rabine + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to deal + * in the Software without restriction, including without limitation the rights + * to use, copy, modify, merge, publish, distribute, sublicense, and/or sell + * copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, + * OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE + * SOFTWARE. + */ + +// libjoin. +#include + +// Libraries. +#include + +// C++. +#include + +// #define SUPPORT_MULTICAST_IPV4 + +using join::lastError; +using join::Errc; +using join::Mdns; +using join::IpAddress; +using join::IpAddressList; +using join::AliasList; +using join::DnsPacket; +using join::ResourceRecord; +using join::DnsMessage; +using join::Mutex; +using join::Condition; +using join::ScopedLock; + +using namespace std::chrono_literals; + +/** + * @brief mDNS peer acting as an announcer. + */ +class MdnsAnnouncer : public Mdns::Peer +{ +public: + /** + * @brief construct the announcer instance. + */ + MdnsAnnouncer () + : Mdns::Peer (_iface) + { + } + + /** + * @brief handle a mDNS query by replying with matching local records. + * @param query the mDNS query received. + */ + void onQuery (const DnsPacket& query) override + { + std::vector answers; + + for (auto const& question : query.questions) + { + if (question.host != _host && question.host != _service && question.host != _serviceType && + question.host != _hostIp4.toArpa () && question.host != _hostIp6.toArpa ()) + { + continue; + } + + for (auto const& kv : _records) + { + const ResourceRecord& record = kv.second; + + if (record.host == question.host && + (question.type == DnsMessage::RecordType::ANY || record.type == question.type)) + { + answers.push_back (record); + } + } + } + + if (!answers.empty ()) + { + this->reply (query, answers); + } + } + + /** + * @brief handle a mDNS announcement. + * @param packet the mDNS announcement received. + */ + void onAnnouncement ([[maybe_unused]] const DnsPacket& packet) override + { + // ignored by the announcer. + } + + /// local resource records indexed by "host/type". + std::map _records; + + /// network interface to use. + static const std::string _iface; + + /// local hostname. + static const std::string _host; + + /// local IPv4 address. + static const IpAddress _hostIp4; + + /// local IPv6 address. + static const IpAddress _hostIp6; + + /// service instance name. + static const std::string _service; + + /// service type. + static const std::string _serviceType; +}; + +const std::string MdnsAnnouncer::_iface = "veth0"; +const std::string MdnsAnnouncer::_host = "mytest.local"; +const IpAddress MdnsAnnouncer::_hostIp4 = "192.168.10.1"; +const IpAddress MdnsAnnouncer::_hostIp6 = "fd00:10::1"; +const std::string MdnsAnnouncer::_service = "MyTest._foobar._tcp.local"; +const std::string MdnsAnnouncer::_serviceType = "_foobar._tcp.local"; + +/** + * @brief mDNS peer acting as a resolver (collects unsolicited announcements). + */ +class MdnsResolver : public Mdns::Peer +{ +public: + /** + * @brief construct the resolver instance. + */ + MdnsResolver () + : Mdns::Peer (_iface) + { + } + + /** + * @brief handle a mDNS query (ignored by the resolver). + * @param query the mDNS query received. + */ + void onQuery ([[maybe_unused]] const DnsPacket& query) override + { + // ignored by the resolver. + } + + /** + * @brief handle a mDNS announcement by storing received records. + * @param packet the mDNS announcement received. + */ + void onAnnouncement (const DnsPacket& packet) override + { + ScopedLock lock (_mutex); + for (auto const& answer : packet.answers) + { + if (answer.host == MdnsAnnouncer::_host || answer.host == MdnsAnnouncer::_service || + (answer.host == MdnsAnnouncer::_serviceType && answer.name == MdnsAnnouncer::_service)) + { + _received.push_back (answer); + } + } + _cond.signal (); + } + + /** + * @brief wait until a record of the given type is received or timeout expires. + * @param type record type to wait for. + * @param timeout maximum time to wait. + * @return true if the record was received, false on timeout. + */ + bool waitForRecord (uint16_t type, std::chrono::milliseconds timeout = 1000ms) + { + ScopedLock lock (_mutex); + return _cond.timedWait (lock, timeout, [&] { + for (auto const& r : _received) + { + if (r.type == type) + { + return true; + } + } + return false; + }); + } + + /// network interface to use. + static const std::string _iface; + + /// records received via onAnnouncement. + std::vector _received; + + /// protection mutex. + Mutex _mutex; + + /// condition variable to signal record reception. + Condition _cond; +}; + +const std::string MdnsResolver::_iface = "veth1"; + +/** + * @brief mDNS test class. + */ +class MdnsTest : public ::testing::Test +{ +protected: + /** + * @brief set up the test suite. + */ + static void SetUpTestSuite () + { + [[maybe_unused]] int result; + + result = ::system ("ip link del veth0 2>/dev/null"); + result = ::system ( + "ip link add veth0 address 06:d8:b9:37:c2:76 type veth peer name veth1 address 52:37:c5:9d:1d:b6"); + + result = ::system ("sysctl -w net.ipv6.conf.veth0.accept_dad=0"); + result = ::system ("sysctl -w net.ipv6.conf.veth1.accept_dad=0"); + + result = ::system ("ip addr add 192.168.10.1/24 dev veth0"); + result = ::system ("ip addr add 192.168.10.2/24 dev veth1"); + + result = ::system ("ip link set veth0 multicast on"); + result = ::system ("ip link set veth1 multicast on"); + + result = ::system ("ip link set veth0 up"); + result = ::system ("ip link set veth1 up"); + + result = ::system ("ip route add 224.0.0.0/4 dev veth0"); + result = ::system ("ip route add 224.0.0.0/4 dev veth1"); + + result = ::system ("ip neigh add 192.168.10.2 lladdr 52:37:c5:9d:1d:b6 dev veth0"); + result = ::system ("ip neigh add 192.168.10.1 lladdr 06:d8:b9:37:c2:76 dev veth1"); + + result = ::system ("ip neigh add fe80::5037:c5ff:fe9d:1db6 lladdr 52:37:c5:9d:1d:b6 dev veth0"); + result = ::system ("ip neigh add fe80::4d8:b9ff:fe37:c276 lladdr 06:d8:b9:37:c2:76 dev veth1"); + + result = ::system ("sysctl -w net.ipv6.conf.veth1.accept_ra=0"); + result = ::system ("sysctl -w net.ipv4.conf.veth0.rp_filter=0"); + result = ::system ("sysctl -w net.ipv4.conf.veth1.rp_filter=0"); + } + + /** + * @brief tear down the test suite. + */ + static void TearDownTestSuite () + { + [[maybe_unused]] int result; + + result = ::system ("ip link del veth0 2>/dev/null"); + } + + /** + * @brief set up the test fixture. + */ + void SetUp () override + { + // IPv6 records + _resolver6._received.clear (); + + ResourceRecord a6; + a6.host = MdnsAnnouncer::_host; + a6.type = DnsMessage::RecordType::A; + a6.dnsclass = DnsMessage::RecordClass::IN; + a6.ttl = 120; + a6.addr = MdnsAnnouncer::_hostIp4; + _announcer6._records[a6.host + "/A"] = a6; + + ResourceRecord aaaa6; + aaaa6.host = MdnsAnnouncer::_host; + aaaa6.type = DnsMessage::RecordType::AAAA; + aaaa6.dnsclass = DnsMessage::RecordClass::IN; + aaaa6.ttl = 120; + aaaa6.addr = MdnsAnnouncer::_hostIp6; + _announcer6._records[aaaa6.host + "/AAAA"] = aaaa6; + + ResourceRecord ptr6_arpa6; + ptr6_arpa6.host = MdnsAnnouncer::_hostIp6.toArpa (); + ptr6_arpa6.type = DnsMessage::RecordType::PTR; + ptr6_arpa6.dnsclass = DnsMessage::RecordClass::IN; + ptr6_arpa6.ttl = 120; + ptr6_arpa6.name = MdnsAnnouncer::_host; + _announcer6._records[ptr6_arpa6.host + "/PTR"] = ptr6_arpa6; + + ResourceRecord ptr6_arpa4; + ptr6_arpa4.host = MdnsAnnouncer::_hostIp4.toArpa (); + ptr6_arpa4.type = DnsMessage::RecordType::PTR; + ptr6_arpa4.dnsclass = DnsMessage::RecordClass::IN; + ptr6_arpa4.ttl = 120; + ptr6_arpa4.name = MdnsAnnouncer::_host; + _announcer6._records[ptr6_arpa4.host + "/PTR"] = ptr6_arpa4; + + ResourceRecord ptr6; + ptr6.host = MdnsAnnouncer::_serviceType; + ptr6.type = DnsMessage::RecordType::PTR; + ptr6.dnsclass = DnsMessage::RecordClass::IN; + ptr6.ttl = 120; + ptr6.name = MdnsAnnouncer::_service; + _announcer6._records[ptr6.host + "/PTR"] = ptr6; + + ResourceRecord srv6; + srv6.host = MdnsAnnouncer::_service; + srv6.type = DnsMessage::RecordType::SRV; + srv6.dnsclass = DnsMessage::RecordClass::IN; + srv6.ttl = 120; + srv6.priority = 0; + srv6.weight = 0; + srv6.port = 80; + srv6.name = MdnsAnnouncer::_host; + _announcer6._records[srv6.host + "/SRV"] = srv6; + + ResourceRecord txt6; + txt6.host = MdnsAnnouncer::_service; + txt6.type = DnsMessage::RecordType::TXT; + txt6.dnsclass = DnsMessage::RecordClass::IN; + txt6.ttl = 120; + txt6.txts = {"path=/", "version=1.0"}; + _announcer6._records[txt6.host + "/TXT"] = txt6; + + ASSERT_EQ (_announcer6.bind (AF_INET6), 0) << lastError.message (); + ASSERT_EQ (_resolver6.bind (AF_INET6), 0) << lastError.message (); + +#ifdef SUPPORT_MULTICAST_IPV4 + // IPv4 records + _resolver4._received.clear (); + + ResourceRecord a4; + a4.host = MdnsAnnouncer::_host; + a4.type = DnsMessage::RecordType::A; + a4.dnsclass = DnsMessage::RecordClass::IN; + a4.ttl = 120; + a4.addr = MdnsAnnouncer::_hostIp4; + _announcer4._records[a4.host + "/A"] = a4; + + ResourceRecord aaaa4; + aaaa4.host = MdnsAnnouncer::_host; + aaaa4.type = DnsMessage::RecordType::AAAA; + aaaa4.dnsclass = DnsMessage::RecordClass::IN; + aaaa4.ttl = 120; + aaaa4.addr = MdnsAnnouncer::_hostIp6; + _announcer4._records[aaaa4.host + "/AAAA"] = aaaa4; + + ResourceRecord ptr4_arpa6; + ptr4_arpa6.host = MdnsAnnouncer::_hostIp6.toArpa (); + ptr4_arpa6.type = DnsMessage::RecordType::PTR; + ptr4_arpa6.dnsclass = DnsMessage::RecordClass::IN; + ptr4_arpa6.ttl = 120; + ptr4_arpa6.name = MdnsAnnouncer::_host; + _announcer4._records[ptr4_arpa6.host + "/PTR"] = ptr4_arpa6; + + ResourceRecord ptr4_arpa4; + ptr4_arpa4.host = MdnsAnnouncer::_hostIp4.toArpa (); + ptr4_arpa4.type = DnsMessage::RecordType::PTR; + ptr4_arpa4.dnsclass = DnsMessage::RecordClass::IN; + ptr4_arpa4.ttl = 120; + ptr4_arpa4.name = MdnsAnnouncer::_host; + _announcer4._records[ptr4_arpa4.host + "/PTR"] = ptr4_arpa4; + + ResourceRecord ptr4; + ptr4.host = MdnsAnnouncer::_serviceType; + ptr4.type = DnsMessage::RecordType::PTR; + ptr4.dnsclass = DnsMessage::RecordClass::IN; + ptr4.ttl = 120; + ptr4.name = MdnsAnnouncer::_service; + _announcer4._records[ptr4.host + "/PTR"] = ptr4; + + ResourceRecord srv4; + srv4.host = MdnsAnnouncer::_service; + srv4.type = DnsMessage::RecordType::SRV; + srv4.dnsclass = DnsMessage::RecordClass::IN; + srv4.ttl = 120; + srv4.priority = 0; + srv4.weight = 0; + srv4.port = 80; + srv4.name = MdnsAnnouncer::_host; + _announcer4._records[srv4.host + "/SRV"] = srv4; + + ResourceRecord txt4; + txt4.host = MdnsAnnouncer::_service; + txt4.type = DnsMessage::RecordType::TXT; + txt4.dnsclass = DnsMessage::RecordClass::IN; + txt4.ttl = 120; + txt4.txts = {"path=/", "version=1.0"}; + _announcer4._records[txt4.host + "/TXT"] = txt4; + + ASSERT_EQ (_announcer4.bind (AF_INET), 0) << lastError.message (); + ASSERT_EQ (_resolver4.bind (AF_INET), 0) << lastError.message (); +#endif + } + + /** + * @brief tear down the test fixture. + */ + void TearDown () override + { + _announcer6.close (); + _resolver6.close (); + +#ifdef SUPPORT_MULTICAST_IPV4 + _announcer4.close (); + _resolver4.close (); +#endif + } + + /// mDNS IPv6 announcer instance. + MdnsAnnouncer _announcer6; + + /// mDNS IPv6 resolver instance. + MdnsResolver _resolver6; + +#ifdef SUPPORT_MULTICAST_IPV4 + /// mDNS IPv4 announcer instance. + MdnsAnnouncer _announcer4; + + /// mDNS IPv4 resolver instance. + MdnsResolver _resolver4; +#endif +}; + +/** + * @brief test the probe method. + */ +TEST_F (MdnsTest, probe) +{ + std::vector records6; + for (auto const& kv : _announcer6._records) + { + records6.push_back (kv.second); + } + EXPECT_EQ (_announcer6.probe (records6), 0) << lastError.message (); + + EXPECT_EQ (_announcer6.probe ({}), -1); + EXPECT_EQ (lastError, make_error_code (Errc::InvalidParam)); + +#ifdef SUPPORT_MULTICAST_IPV4 + std::vector records4; + for (auto const& kv : _announcer4._records) + { + records4.push_back (kv.second); + } + EXPECT_EQ (_announcer4.probe (records4), 0) << lastError.message (); + + EXPECT_EQ (_announcer4.probe ({}), -1); + EXPECT_EQ (lastError, make_error_code (Errc::InvalidParam)); +#endif +} + +/** + * @brief test the announce method. + */ +TEST_F (MdnsTest, announce) +{ + std::vector records6; + for (auto const& kv : _announcer6._records) + { + records6.push_back (kv.second); + } + EXPECT_EQ (_announcer6.announce (records6), 0) << lastError.message (); + EXPECT_TRUE (_resolver6.waitForRecord (DnsMessage::RecordType::AAAA)); + + EXPECT_EQ (_announcer6.announce ({}), -1); + EXPECT_EQ (lastError, make_error_code (Errc::InvalidParam)); + +#ifdef SUPPORT_MULTICAST_IPV4 + std::vector records4; + for (auto const& kv : _announcer4._records) + { + records4.push_back (kv.second); + } + EXPECT_EQ (_announcer4.announce (records4), 0) << lastError.message (); + EXPECT_TRUE (_resolver4.waitForRecord (DnsMessage::RecordType::A)); + + EXPECT_EQ (_announcer4.announce ({}), -1); + EXPECT_EQ (lastError, make_error_code (Errc::InvalidParam)); +#endif +} + +/** + * @brief test the goodbye method. + */ +TEST_F (MdnsTest, goodbye) +{ + std::vector records6; + for (auto const& kv : _announcer6._records) + { + records6.push_back (kv.second); + } + EXPECT_EQ (_announcer6.goodbye (records6), 0) << lastError.message (); + EXPECT_TRUE (_resolver6.waitForRecord (DnsMessage::RecordType::AAAA)); + for (auto const& r : _resolver6._received) + { + if (r.type == DnsMessage::RecordType::AAAA) + { + EXPECT_EQ (r.ttl, 0); + } + } + + EXPECT_EQ (_announcer6.goodbye ({}), -1); + EXPECT_EQ (lastError, make_error_code (Errc::InvalidParam)); + +#ifdef SUPPORT_MULTICAST_IPV4 + std::vector records4; + for (auto const& kv : _announcer4._records) + { + records4.push_back (kv.second); + } + EXPECT_EQ (_announcer4.goodbye (records4), 0) << lastError.message (); + EXPECT_TRUE (_resolver4.waitForRecord (DnsMessage::RecordType::A)); + for (auto const& r : _resolver4._received) + { + if (r.type == DnsMessage::RecordType::A) + { + EXPECT_EQ (r.ttl, 0); + } + } + + EXPECT_EQ (_announcer4.goodbye ({}), -1); + EXPECT_EQ (lastError, make_error_code (Errc::InvalidParam)); +#endif +} + +/** + * @brief test the browse method. + */ +TEST_F (MdnsTest, browse) +{ + EXPECT_EQ (_resolver6.browse (MdnsAnnouncer::_serviceType), 0) << lastError.message (); + EXPECT_TRUE (_resolver6.waitForRecord (DnsMessage::RecordType::PTR)); + bool found6 = false; + for (auto const& r : _resolver6._received) + { + if (r.type == DnsMessage::RecordType::PTR && r.name == MdnsAnnouncer::_service) + { + found6 = true; + break; + } + } + EXPECT_TRUE (found6); + + EXPECT_EQ (_resolver6.browse (""), -1); + EXPECT_EQ (lastError, make_error_code (Errc::InvalidParam)); + +#ifdef SUPPORT_MULTICAST_IPV4 + EXPECT_EQ (_resolver4.browse (MdnsAnnouncer::_serviceType), 0) << lastError.message (); + EXPECT_TRUE (_resolver4.waitForRecord (DnsMessage::RecordType::PTR)); + bool found4 = false; + for (auto const& r : _resolver4._received) + { + if (r.type == DnsMessage::RecordType::PTR && r.name == MdnsAnnouncer::_service) + { + found4 = true; + break; + } + } + EXPECT_TRUE (found4); + + EXPECT_EQ (_resolver4.browse (""), -1); + EXPECT_EQ (lastError, make_error_code (Errc::InvalidParam)); +#endif +} + +/** + * @brief test the resolveAddress method. + */ +TEST_F (MdnsTest, resolveAddress) +{ + IpAddress addr = _resolver6.resolveAddress ("", AF_INET6, 500ms); + EXPECT_TRUE (addr.isWildcard ()); + + addr = _resolver6.resolveAddress (MdnsAnnouncer::_host, AF_INET6, 500ms); + EXPECT_FALSE (addr.isWildcard ()); + EXPECT_EQ (addr, MdnsAnnouncer::_hostIp6); + + addr = _resolver6.resolveAddress (MdnsAnnouncer::_host, AF_INET, 500ms); + EXPECT_FALSE (addr.isWildcard ()); + EXPECT_EQ (addr, MdnsAnnouncer::_hostIp4); + + addr = _resolver6.resolveAddress ("unknown.local", AF_INET6, 500ms); + EXPECT_TRUE (addr.isWildcard ()); + EXPECT_EQ (lastError, make_error_code (Errc::TimedOut)); + + addr = _resolver6.resolveAddress ("", 500ms); + EXPECT_TRUE (addr.isWildcard ()); + + addr = _resolver6.resolveAddress (MdnsAnnouncer::_host, 500ms); + EXPECT_FALSE (addr.isWildcard ()); + + addr = _resolver6.resolveAddress (MdnsAnnouncer::_host, 500ms); + EXPECT_FALSE (addr.isWildcard ()); + + addr = _resolver6.resolveAddress ("unknown.local", 500ms); + EXPECT_TRUE (addr.isWildcard ()); + EXPECT_EQ (lastError, make_error_code (Errc::TimedOut)); + +#ifdef SUPPORT_MULTICAST_IPV4 + addr = _resolver4.resolveAddress ("", AF_INET, 500ms); + EXPECT_TRUE (addr.isWildcard ()); + + addr = _resolver4.resolveAddress (MdnsAnnouncer::_host, AF_INET6, 500ms); + EXPECT_FALSE (addr.isWildcard ()); + EXPECT_EQ (addr, MdnsAnnouncer::_hostIp6); + + addr = _resolver4.resolveAddress (MdnsAnnouncer::_host, AF_INET, 500ms); + EXPECT_FALSE (addr.isWildcard ()); + EXPECT_EQ (addr, MdnsAnnouncer::_hostIp4); + + addr = _resolver4.resolveAddress ("unknown.local", AF_INET, 500ms); + EXPECT_TRUE (addr.isWildcard ()); + EXPECT_EQ (lastError, make_error_code (Errc::TimedOut)); + + addr = _resolver4.resolveAddress ("", 500ms); + EXPECT_TRUE (addr.isWildcard ()); + + addr = _resolver4.resolveAddress (MdnsAnnouncer::_host, 500ms); + EXPECT_FALSE (addr.isWildcard ()); + + addr = _resolver4.resolveAddress (MdnsAnnouncer::_host, 500ms); + EXPECT_FALSE (addr.isWildcard ()); + + addr = _resolver4.resolveAddress ("unknown.local", 500ms); + EXPECT_TRUE (addr.isWildcard ()); + EXPECT_EQ (lastError, make_error_code (Errc::TimedOut)); +#endif +} + +/** + * @brief test the resolveAllAddress method. + */ +TEST_F (MdnsTest, resolveAllAddress) +{ + IpAddressList addrs = _resolver6.resolveAllAddress ("", AF_INET6, 500ms); + EXPECT_EQ (addrs.size (), 0); + + addrs = _resolver6.resolveAllAddress (MdnsAnnouncer::_host, AF_INET6, 500ms); + ASSERT_GT (addrs.size (), 0); + EXPECT_FALSE (addrs.front ().isWildcard ()); + + addrs = _resolver6.resolveAllAddress (MdnsAnnouncer::_host, AF_INET, 500ms); + ASSERT_GT (addrs.size (), 0); + EXPECT_FALSE (addrs.front ().isWildcard ()); + + addrs = _resolver6.resolveAllAddress ("unknown.local", AF_INET6, 500ms); + EXPECT_EQ (addrs.size (), 0); + + addrs = _resolver6.resolveAllAddress ("", 500ms); + EXPECT_EQ (addrs.size (), 0); + + addrs = _resolver6.resolveAllAddress (MdnsAnnouncer::_host, 500ms); + ASSERT_GT (addrs.size (), 0); + EXPECT_FALSE (addrs.front ().isWildcard ()); + + addrs = _resolver6.resolveAllAddress (MdnsAnnouncer::_host, 500ms); + ASSERT_GT (addrs.size (), 0); + EXPECT_FALSE (addrs.front ().isWildcard ()); + + addrs = _resolver6.resolveAllAddress ("unknown.local", 500ms); + EXPECT_EQ (addrs.size (), 0); + +#ifdef SUPPORT_MULTICAST_IPV4 + addrs = _resolver4.resolveAllAddress ("", AF_INET, 500ms); + EXPECT_EQ (addrs.size (), 0); + + addrs = _resolver4.resolveAllAddress (MdnsAnnouncer::_host, AF_INET6, 500ms); + ASSERT_GT (addrs.size (), 0); + EXPECT_FALSE (addrs.front ().isWildcard ()); + + addrs = _resolver4.resolveAllAddress (MdnsAnnouncer::_host, AF_INET, 500ms); + ASSERT_GT (addrs.size (), 0); + EXPECT_FALSE (addrs.front ().isWildcard ()); + + addrs = _resolver4.resolveAllAddress ("unknown.local", AF_INET, 500ms); + EXPECT_EQ (addrs.size (), 0); + + addrs = _resolver4.resolveAllAddress ("", 500ms); + EXPECT_EQ (addrs.size (), 0); + + addrs = _resolver4.resolveAllAddress (MdnsAnnouncer::_host, 500ms); + ASSERT_GT (addrs.size (), 0); + EXPECT_FALSE (addrs.front ().isWildcard ()); + + addrs = _resolver4.resolveAllAddress (MdnsAnnouncer::_host, 500ms); + ASSERT_GT (addrs.size (), 0); + EXPECT_FALSE (addrs.front ().isWildcard ()); + + addrs = _resolver4.resolveAllAddress ("unknown.local", 500ms); + EXPECT_EQ (addrs.size (), 0); +#endif +} + +/** + * @brief test the resolveName method. + */ +TEST_F (MdnsTest, resolveName) +{ + std::string name = _resolver6.resolveName (IpAddress ("::"), 500ms); + EXPECT_TRUE (name.empty ()); + + name = _resolver6.resolveName (MdnsAnnouncer::_hostIp6, 500ms); + EXPECT_FALSE (name.empty ()); + EXPECT_EQ (name, MdnsAnnouncer::_host); + + name = _resolver6.resolveName (MdnsAnnouncer::_hostIp4, 500ms); + EXPECT_FALSE (name.empty ()); + EXPECT_EQ (name, MdnsAnnouncer::_host); + + name = _resolver6.resolveName (IpAddress ("fd00::1"), 500ms); + EXPECT_TRUE (name.empty ()); + EXPECT_EQ (lastError, make_error_code (Errc::TimedOut)); + +#ifdef SUPPORT_MULTICAST_IPV4 + name = _resolver4.resolveName (IpAddress ("0.0.0.0"), 500ms); + EXPECT_TRUE (name.empty ()); + + name = _resolver4.resolveName (MdnsAnnouncer::_hostIp6, 500ms); + EXPECT_FALSE (name.empty ()); + EXPECT_EQ (name, MdnsAnnouncer::_host); + + name = _resolver4.resolveName (MdnsAnnouncer::_hostIp4, 500ms); + EXPECT_FALSE (name.empty ()); + EXPECT_EQ (name, MdnsAnnouncer::_host); + + name = _resolver4.resolveName (IpAddress ("192.168.1.99"), 500ms); + EXPECT_TRUE (name.empty ()); + EXPECT_EQ (lastError, make_error_code (Errc::TimedOut)); +#endif +} + +/** + * @brief test the resolveAllName method. + */ +TEST_F (MdnsTest, resolveAllName) +{ + AliasList aliases = _resolver6.resolveAllName (IpAddress ("::"), 500ms); + EXPECT_EQ (aliases.size (), 0); + + aliases = _resolver6.resolveAllName (MdnsAnnouncer::_hostIp6, 500ms); + ASSERT_GT (aliases.size (), 0); + EXPECT_NE (aliases.find (MdnsAnnouncer::_host), aliases.end ()); + + aliases = _resolver6.resolveAllName (MdnsAnnouncer::_hostIp4, 500ms); + ASSERT_GT (aliases.size (), 0); + EXPECT_NE (aliases.find (MdnsAnnouncer::_host), aliases.end ()); + + aliases = _resolver6.resolveAllName (IpAddress ("fd00::1"), 500ms); + EXPECT_EQ (aliases.size (), 0); + EXPECT_EQ (lastError, make_error_code (Errc::TimedOut)); + +#ifdef SUPPORT_MULTICAST_IPV4 + aliases = _resolver4.resolveAllName (IpAddress ("0.0.0.0"), 500ms); + EXPECT_EQ (aliases.size (), 0); + + aliases = _resolver4.resolveAllName (MdnsAnnouncer::_hostIp6, 500ms); + ASSERT_GT (aliases.size (), 0); + EXPECT_NE (aliases.find (MdnsAnnouncer::_host), aliases.end ()); + + aliases = _resolver4.resolveAllName (MdnsAnnouncer::_hostIp4, 500ms); + ASSERT_GT (aliases.size (), 0); + EXPECT_NE (aliases.find (MdnsAnnouncer::_host), aliases.end ()); + + aliases = _resolver4.resolveAllName (IpAddress ("192.168.1.99"), 500ms); + EXPECT_EQ (aliases.size (), 0); + EXPECT_EQ (lastError, make_error_code (Errc::TimedOut)); +#endif +} + +/** + * @brief main function. + */ +int main (int argc, char** argv) +{ + testing::InitGoogleTest (&argc, argv); + return RUN_ALL_TESTS (); +}