[dnssd-server] use case-insensitive DNS name match (#7189)

This commit updates `Dns::ServiceDiscovery::Server` to use
case-insensitive string match when comparing DNS names. It also
updates `test_dnssd.py` to use mixed case names when browsing for or
resolving services. This validates the case-insensitive treatment
of names by SRP client and server and DNS-SD server (resolver) and
DNS client.
This commit is contained in:
Abtin Keshavarzian
2021-11-23 16:31:18 -08:00
committed by Jonathan Hui
parent 9d81f99b94
commit 4ac6b504a4
3 changed files with 25 additions and 23 deletions
+12 -11
View File
@@ -41,6 +41,7 @@
#include "common/instance.hpp" #include "common/instance.hpp"
#include "common/locator_getters.hpp" #include "common/locator_getters.hpp"
#include "common/logging.hpp" #include "common/logging.hpp"
#include "common/string.hpp"
#include "net/srp_server.hpp" #include "net/srp_server.hpp"
#include "net/udp6.hpp" #include "net/udp6.hpp"
@@ -48,8 +49,8 @@ namespace ot {
namespace Dns { namespace Dns {
namespace ServiceDiscovery { namespace ServiceDiscovery {
const char Server::kDnssdProtocolUdp[4] = {'_', 'u', 'd', 'p'}; const char Server::kDnssdProtocolUdp[] = "_udp";
const char Server::kDnssdProtocolTcp[4] = {'_', 't', 'c', 'p'}; const char Server::kDnssdProtocolTcp[] = "_tcp";
const char Server::kDnssdSubTypeLabel[] = "._sub."; const char Server::kDnssdSubTypeLabel[] = "._sub.";
const char Server::kDefaultDomainName[] = "default.service.arpa."; const char Server::kDefaultDomainName[] = "default.service.arpa.";
@@ -210,11 +211,11 @@ void Server::SendResponse(Header aHeader,
if (error != kErrorNone) if (error != kErrorNone)
{ {
otLogWarnDns("[server] failed to send DNS-SD reply: %s", otThreadErrorToString(error)); otLogWarnDns("[server] failed to send DNS-SD reply: %s", ErrorToString(error));
} }
else else
{ {
otLogInfoDns("[server] send DNS-SD reply: %s, RCODE=%d", otThreadErrorToString(error), aResponseCode); otLogInfoDns("[server] send DNS-SD reply: %s, RCODE=%d", ErrorToString(error), aResponseCode);
} }
} }
@@ -395,7 +396,7 @@ Error Server::AppendServiceName(Message &aMessage, const char *aName, NameCompre
const char *serviceName; const char *serviceName;
// Check whether `aName` is a sub-type service name. // Check whether `aName` is a sub-type service name.
serviceName = StringFind(aName, kDnssdSubTypeLabel); serviceName = StringFind(aName, kDnssdSubTypeLabel, kStringCaseInsensitiveMatch);
if (serviceName != nullptr) if (serviceName != nullptr)
{ {
@@ -577,8 +578,8 @@ Error Server::FindNameComponents(const char *aName, const char *aDomain, NameCom
VerifyOrExit(error == kErrorNone, error = (error == kErrorNotFound ? kErrorNone : error)); VerifyOrExit(error == kErrorNone, error = (error == kErrorNotFound ? kErrorNone : error));
if (labelEnd == labelBegin + kProtocolLabelLength && if (labelEnd == labelBegin + kProtocolLabelLength &&
(memcmp(&aName[labelBegin], kDnssdProtocolUdp, kProtocolLabelLength) == 0 || (StringStartsWith(&aName[labelBegin], kDnssdProtocolUdp, kStringCaseInsensitiveMatch) ||
memcmp(&aName[labelBegin], kDnssdProtocolTcp, kProtocolLabelLength) == 0)) StringStartsWith(&aName[labelBegin], kDnssdProtocolTcp, kStringCaseInsensitiveMatch)))
{ {
// <Protocol> label found // <Protocol> label found
aInfo.mProtocolOffset = labelBegin; aInfo.mProtocolOffset = labelBegin;
@@ -599,7 +600,7 @@ Error Server::FindNameComponents(const char *aName, const char *aDomain, NameCom
// Note that `kDnssdSubTypeLabel` is "._sub.". Here we get the // Note that `kDnssdSubTypeLabel` is "._sub.". Here we get the
// label only so we want to compare it with "_sub". // label only so we want to compare it with "_sub".
if ((labelEnd == labelBegin + kSubTypeLabelLength) && if ((labelEnd == labelBegin + kSubTypeLabelLength) &&
(memcmp(&aName[labelBegin], kDnssdSubTypeLabel + 1, kSubTypeLabelLength) == 0)) StringStartsWith(&aName[labelBegin], kDnssdSubTypeLabel + 1, kStringCaseInsensitiveMatch))
{ {
SuccessOrExit(error = FindPreviousLabel(aName, labelBegin, labelEnd)); SuccessOrExit(error = FindPreviousLabel(aName, labelBegin, labelEnd));
VerifyOrExit(labelBegin == 0, error = kErrorInvalidArgs); VerifyOrExit(labelBegin == 0, error = kErrorInvalidArgs);
@@ -864,10 +865,10 @@ bool Server::CanAnswerQuery(const QueryTransaction & aQuery,
switch (sdType) switch (sdType)
{ {
case kDnsQueryBrowse: case kDnsQueryBrowse:
canAnswer = (strcmp(name, aServiceFullName) == 0); canAnswer = StringMatch(name, aServiceFullName, kStringCaseInsensitiveMatch);
break; break;
case kDnsQueryResolve: case kDnsQueryResolve:
canAnswer = (strcmp(name, aInstanceInfo.mFullName) == 0); canAnswer = StringMatch(name, aInstanceInfo.mFullName, kStringCaseInsensitiveMatch);
break; break;
default: default:
break; break;
@@ -882,7 +883,7 @@ bool Server::CanAnswerQuery(const Server::QueryTransaction &aQuery, const char *
DnsQueryType sdType; DnsQueryType sdType;
sdType = GetQueryTypeAndName(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), name); sdType = GetQueryTypeAndName(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), name);
return (sdType == kDnsQueryResolveHost) && (strcmp(name, aHostFullName) == 0); return (sdType == kDnsQueryResolveHost) && StringMatch(name, aHostFullName, kStringCaseInsensitiveMatch);
} }
void Server::AnswerQuery(QueryTransaction & aQuery, void Server::AnswerQuery(QueryTransaction & aQuery,
+2 -2
View File
@@ -392,8 +392,8 @@ private:
void HandleTimer(void); void HandleTimer(void);
void ResetTimer(void); void ResetTimer(void);
static const char kDnssdProtocolUdp[4]; static const char kDnssdProtocolUdp[];
static const char kDnssdProtocolTcp[4]; static const char kDnssdProtocolTcp[];
static const char kDnssdSubTypeLabel[]; static const char kDnssdSubTypeLabel[];
static const char kDefaultDomainName[]; static const char kDefaultDomainName[];
Ip6::Udp::Socket mSocket; Ip6::Udp::Socket mSocket;
+11 -10
View File
@@ -103,13 +103,13 @@ class TestDnssd(thread_cert.TestCase):
client3_addrs = [client3.get_mleid(), client2.get_rloc()] client3_addrs = [client3.get_mleid(), client2.get_rloc()]
self._config_srp_client_services(client1, server, 'ins1', 'host1', 11111, 1, 1, client1_addrs, ",_s1,_s2") self._config_srp_client_services(client1, server, 'ins1', 'host1', 11111, 1, 1, client1_addrs, ",_s1,_s2")
self._config_srp_client_services(client2, server, 'ins2', 'host2', 22222, 2, 2, client2_addrs) self._config_srp_client_services(client2, server, 'ins2', 'HOST2', 22222, 2, 2, client2_addrs)
self._config_srp_client_services(client3, server, 'ins3', 'host3', 33333, 3, 3, client3_addrs, ",_s1") self._config_srp_client_services(client3, server, 'ins3', 'host3', 33333, 3, 3, client3_addrs, ",_S1")
#--------------------------------------------------------------- #---------------------------------------------------------------
# Resolve address (AAAA records) # Resolve address (AAAA records)
answers = client1.dns_resolve(f"host1.{DOMAIN}", server.get_mleid(), 53) answers = client1.dns_resolve(f"host1.{DOMAIN}".upper(), server.get_mleid(), 53)
self.assertEqual(set(ipaddress.IPv6Address(ip) for ip, _ in answers), self.assertEqual(set(ipaddress.IPv6Address(ip) for ip, _ in answers),
set(map(ipaddress.IPv6Address, client1_addrs))) set(map(ipaddress.IPv6Address, client1_addrs)))
@@ -169,36 +169,36 @@ class TestDnssd(thread_cert.TestCase):
} }
# Browse for main service # Browse for main service
service_instances = client1.dns_browse(f'{SERVICE}.{DOMAIN}', server.get_mleid(), 53) service_instances = client1.dns_browse(f'{SERVICE}.{DOMAIN}'.upper(), server.get_mleid(), 53)
self.assertEqual({'ins1', 'ins2', 'ins3'}, set(service_instances.keys())) self.assertEqual({'ins1', 'ins2', 'ins3'}, set(service_instances.keys()))
self._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info) self._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info)
self._assert_service_instance_equal(service_instances['ins2'], instance2_verify_info) self._assert_service_instance_equal(service_instances['ins2'], instance2_verify_info)
self._assert_service_instance_equal(service_instances['ins3'], instance3_verify_info) self._assert_service_instance_equal(service_instances['ins3'], instance3_verify_info)
# Browse for service sub-type _s1. # Browse for service sub-type _s1.
service_instances = client1.dns_browse(f'_s1._sub.{SERVICE}.{DOMAIN}', server.get_mleid(), 53) service_instances = client1.dns_browse(f'_s1._sub.{SERVICE}.{DOMAIN}'.upper(), server.get_mleid(), 53)
self.assertEqual({'ins1', 'ins3'}, set(service_instances.keys())) self.assertEqual({'ins1', 'ins3'}, set(service_instances.keys()))
self._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info) self._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info)
# Browse for service sub-type _s2. # Browse for service sub-type _s2.
service_instances = client1.dns_browse(f'_s2._sub.{SERVICE}.{DOMAIN}', server.get_mleid(), 53) service_instances = client1.dns_browse(f'_s2._sub.{SERVICE}.{DOMAIN}'.upper(), server.get_mleid(), 53)
self.assertEqual({'ins1'}, set(service_instances.keys())) self.assertEqual({'ins1'}, set(service_instances.keys()))
self._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info) self._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info)
#--------------------------------------------------------------- #---------------------------------------------------------------
# Resolve service # Resolve service
service_instance = client1.dns_resolve_service('ins1', f'{SERVICE}.{DOMAIN}', server.get_mleid(), 53) service_instance = client1.dns_resolve_service('ins1', f'{SERVICE}.{DOMAIN}'.upper(), server.get_mleid(), 53)
self._assert_service_instance_equal(service_instance, instance1_verify_info) self._assert_service_instance_equal(service_instance, instance1_verify_info)
service_instance = client1.dns_resolve_service('ins2', f'{SERVICE}.{DOMAIN}', server.get_mleid(), 53) service_instance = client1.dns_resolve_service('ins2', f'{SERVICE}.{DOMAIN}'.upper(), server.get_mleid(), 53)
self._assert_service_instance_equal(service_instance, instance2_verify_info) self._assert_service_instance_equal(service_instance, instance2_verify_info)
#--------------------------------------------------------------- #---------------------------------------------------------------
# Add another service with TXT entries to the existing host and # Add another service with TXT entries to the existing host and
# verify that it is properly merged. # verify that it is properly merged.
client3.srp_client_add_service('ins4', SERVICE + ",_s1", 44444, 4, 4, txt_entries=['KEY=ABC']) client3.srp_client_add_service('ins4', (SERVICE + ",_s1").upper(), 44444, 4, 4, txt_entries=['KEY=ABC'])
self.simulator.go(5) self.simulator.go(5)
service_instances = client1.dns_browse(f'{SERVICE}.{DOMAIN}', server.get_mleid(), 53) service_instances = client1.dns_browse(f'{SERVICE}.{DOMAIN}', server.get_mleid(), 53)
@@ -209,7 +209,8 @@ class TestDnssd(thread_cert.TestCase):
self._assert_service_instance_equal(service_instances['ins4'], instance4_verify_info) self._assert_service_instance_equal(service_instances['ins4'], instance4_verify_info)
def _assert_service_instance_equal(self, instance, info): def _assert_service_instance_equal(self, instance, info):
for f in ('port', 'priority', 'weight', 'host', 'txt_data'): self.assertEqual(instance['host'].lower(), info['host'].lower(), instance)
for f in ('port', 'priority', 'weight', 'txt_data'):
self.assertEqual(instance[f], info[f], instance) self.assertEqual(instance[f], info[f], instance)
verify_addresses = info['address'] verify_addresses = info['address']