[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/locator_getters.hpp"
#include "common/logging.hpp"
#include "common/string.hpp"
#include "net/srp_server.hpp"
#include "net/udp6.hpp"
@@ -48,8 +49,8 @@ namespace ot {
namespace Dns {
namespace ServiceDiscovery {
const char Server::kDnssdProtocolUdp[4] = {'_', 'u', 'd', 'p'};
const char Server::kDnssdProtocolTcp[4] = {'_', 't', 'c', 'p'};
const char Server::kDnssdProtocolUdp[] = "_udp";
const char Server::kDnssdProtocolTcp[] = "_tcp";
const char Server::kDnssdSubTypeLabel[] = "._sub.";
const char Server::kDefaultDomainName[] = "default.service.arpa.";
@@ -210,11 +211,11 @@ void Server::SendResponse(Header aHeader,
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
{
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;
// Check whether `aName` is a sub-type service name.
serviceName = StringFind(aName, kDnssdSubTypeLabel);
serviceName = StringFind(aName, kDnssdSubTypeLabel, kStringCaseInsensitiveMatch);
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));
if (labelEnd == labelBegin + kProtocolLabelLength &&
(memcmp(&aName[labelBegin], kDnssdProtocolUdp, kProtocolLabelLength) == 0 ||
memcmp(&aName[labelBegin], kDnssdProtocolTcp, kProtocolLabelLength) == 0))
(StringStartsWith(&aName[labelBegin], kDnssdProtocolUdp, kStringCaseInsensitiveMatch) ||
StringStartsWith(&aName[labelBegin], kDnssdProtocolTcp, kStringCaseInsensitiveMatch)))
{
// <Protocol> label found
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
// label only so we want to compare it with "_sub".
if ((labelEnd == labelBegin + kSubTypeLabelLength) &&
(memcmp(&aName[labelBegin], kDnssdSubTypeLabel + 1, kSubTypeLabelLength) == 0))
StringStartsWith(&aName[labelBegin], kDnssdSubTypeLabel + 1, kStringCaseInsensitiveMatch))
{
SuccessOrExit(error = FindPreviousLabel(aName, labelBegin, labelEnd));
VerifyOrExit(labelBegin == 0, error = kErrorInvalidArgs);
@@ -864,10 +865,10 @@ bool Server::CanAnswerQuery(const QueryTransaction & aQuery,
switch (sdType)
{
case kDnsQueryBrowse:
canAnswer = (strcmp(name, aServiceFullName) == 0);
canAnswer = StringMatch(name, aServiceFullName, kStringCaseInsensitiveMatch);
break;
case kDnsQueryResolve:
canAnswer = (strcmp(name, aInstanceInfo.mFullName) == 0);
canAnswer = StringMatch(name, aInstanceInfo.mFullName, kStringCaseInsensitiveMatch);
break;
default:
break;
@@ -882,7 +883,7 @@ bool Server::CanAnswerQuery(const Server::QueryTransaction &aQuery, const char *
DnsQueryType sdType;
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,
+2 -2
View File
@@ -392,8 +392,8 @@ private:
void HandleTimer(void);
void ResetTimer(void);
static const char kDnssdProtocolUdp[4];
static const char kDnssdProtocolTcp[4];
static const char kDnssdProtocolUdp[];
static const char kDnssdProtocolTcp[];
static const char kDnssdSubTypeLabel[];
static const char kDefaultDomainName[];
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()]
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(client3, server, 'ins3', 'host3', 33333, 3, 3, client3_addrs, ",_s1")
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")
#---------------------------------------------------------------
# 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),
set(map(ipaddress.IPv6Address, client1_addrs)))
@@ -169,36 +169,36 @@ class TestDnssd(thread_cert.TestCase):
}
# 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._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['ins3'], instance3_verify_info)
# 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._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info)
# 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._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info)
#---------------------------------------------------------------
# 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)
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)
#---------------------------------------------------------------
# Add another service with TXT entries to the existing host and
# 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)
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)
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)
verify_addresses = info['address']