1 # Copyright (C) Internet Systems Consortium, Inc. ("ISC") 2 # 3 # SPDX-License-Identifier: MPL-2.0 4 # 5 # This Source Code Form is subject to the terms of the Mozilla Public 6 # License, v. 2.0. If a copy of the MPL was not distributed with this 7 # file, you can obtain one at https://mozilla.org/MPL/2.0/. 8 # 9 # See the COPYRIGHT file distributed with this work for additional 10 # information regarding copyright ownership. 11 12 from typing import cast 13 14 import difflib 15 import shutil 16 17 from dns.edns import EDECode, EDEOption 18 19 import dns.edns 20 import dns.flags 21 import dns.message 22 import dns.rcode 23 import dns.rrset 24 import dns.zone 25 26 import isctest.log 27 28 29 def rcode(message: dns.message.Message, expected_rcode) -> None: 30 assert message.rcode() == expected_rcode, str(message) 31 32 33 def noerror(message: dns.message.Message) -> None: 34 rcode(message, dns.rcode.NOERROR) 35 36 37 def notimp(message: dns.message.Message) -> None: 38 rcode(message, dns.rcode.NOTIMP) 39 40 41 def refused(message: dns.message.Message) -> None: 42 rcode(message, dns.rcode.REFUSED) 43 44 45 def servfail(message: dns.message.Message) -> None: 46 rcode(message, dns.rcode.SERVFAIL) 47 48 49 def formerr(message: dns.message.Message) -> None: 50 rcode(message, dns.rcode.FORMERR) 51 52 53 def adflag(message: dns.message.Message) -> None: 54 assert (message.flags & dns.flags.AD) != 0, str(message) 55 56 57 def noadflag(message: dns.message.Message) -> None: 58 assert (message.flags & dns.flags.AD) == 0, str(message) 59 60 61 def rdflag(message: dns.message.Message) -> None: 62 assert (message.flags & dns.flags.RD) != 0, str(message) 63 64 65 def nordflag(message: dns.message.Message) -> None: 66 assert (message.flags & dns.flags.RD) == 0, str(message) 67 68 69 def raflag(message: dns.message.Message) -> None: 70 assert (message.flags & dns.flags.RA) != 0, str(message) 71 72 73 def noraflag(message: dns.message.Message) -> None: 74 assert (message.flags & dns.flags.RA) == 0, str(message) 75 76 77 def _extract_ede_options( 78 message: dns.message.Message, 79 ) -> list[EDEOption]: 80 """Extract EDE options from the DNS message.""" 81 return cast( 82 list[EDEOption], 83 [ 84 option 85 for option in message.options 86 if option.otype == dns.edns.OptionType.EDE 87 ], 88 ) 89 90 91 def noede(message: dns.message.Message) -> None: 92 """Check that message contains no EDE option.""" 93 ede_options = _extract_ede_options(message) 94 assert not ede_options, f"unexpected EDE options {ede_options} in {message}" 95 96 97 def ede(message: dns.message.Message, code: EDECode, text: str | None = None) -> None: 98 """Check if message contains expected EDE code (and its text).""" 99 msg_opts = _extract_ede_options(message) 100 matching_opts = [opt for opt in msg_opts if opt.code == code] 101 102 assert matching_opts, f"missing EDE code {code} in {message}" 103 104 if text is None: 105 return 106 107 # check at least one matching EDE option has the required text 108 for opt in matching_opts: 109 if opt.text == text: 110 return 111 opt_str = ", ".join([opt.to_text() for opt in matching_opts]) 112 assert False, f'EDE text "{text}" not found in [{opt_str}]' 113 114 115 def section_equal(first_section: list, second_section: list) -> None: 116 for rrset in first_section: 117 assert ( 118 rrset in second_section 119 ), f"No corresponding RRset found in second section: {rrset}" 120 for rrset in second_section: 121 assert ( 122 rrset in first_section 123 ), f"No corresponding RRset found in first section: {rrset}" 124 125 126 def same_data(res1: dns.message.Message, res2: dns.message.Message): 127 section_equal(res1.question, res2.question) 128 section_equal(res1.answer, res2.answer) 129 section_equal(res1.authority, res2.authority) 130 section_equal(res1.additional, res2.additional) 131 assert res1.rcode() == res2.rcode() 132 133 134 def same_answer(res1: dns.message.Message, res2: dns.message.Message): 135 section_equal(res1.question, res2.question) 136 section_equal(res1.answer, res2.answer) 137 assert res1.rcode() == res2.rcode() 138 139 140 def rrsets_equal( 141 first_rrset: dns.rrset.RRset, 142 second_rrset: dns.rrset.RRset, 143 compare_ttl: bool | None = False, 144 ) -> None: 145 """Compare two RRset (optionally including TTL)""" 146 147 def compare_rrs(rr1, rrset): 148 rr2 = next((other_rr for other_rr in rrset if rr1 == other_rr), None) 149 assert rr2 is not None, f"No corresponding RR found for: {rr1}" 150 if compare_ttl: 151 assert rr1.ttl == rr2.ttl 152 153 isctest.log.debug( 154 "%s() first RRset:\n%s", 155 rrsets_equal.__name__, 156 "\n".join([str(rr) for rr in first_rrset]), 157 ) 158 isctest.log.debug( 159 "%s() second RRset:\n%s", 160 rrsets_equal.__name__, 161 "\n".join([str(rr) for rr in second_rrset]), 162 ) 163 for rr in first_rrset: 164 compare_rrs(rr, second_rrset) 165 for rr in second_rrset: 166 compare_rrs(rr, first_rrset) 167 168 169 def zones_equal( 170 first_zone: dns.zone.Zone, 171 second_zone: dns.zone.Zone, 172 compare_ttl: bool | None = False, 173 ) -> None: 174 """Compare two zones (optionally including TTL)""" 175 176 isctest.log.debug( 177 "%s() first zone:\n%s", 178 zones_equal.__name__, 179 first_zone.to_text(relativize=False), 180 ) 181 isctest.log.debug( 182 "%s() second zone:\n%s", 183 zones_equal.__name__, 184 second_zone.to_text(relativize=False), 185 ) 186 assert first_zone == second_zone 187 if compare_ttl: 188 for name, node in first_zone.nodes.items(): 189 for rdataset in node: 190 found_rdataset = second_zone.find_rdataset( 191 name=name, rdtype=rdataset.rdtype 192 ) 193 assert found_rdataset 194 assert found_rdataset.ttl == rdataset.ttl 195 196 197 def is_executable(cmd: str, errmsg: str) -> None: 198 executable = shutil.which(cmd) 199 assert executable is not None, errmsg 200 201 202 def named_alive(named_proc, resolver_ip): 203 assert named_proc.poll() is None, "named isn't running" 204 msg = isctest.query.create("version.bind", "TXT", "CH") 205 isctest.query.tcp(msg, resolver_ip, expected_rcode=dns.rcode.NOERROR) 206 207 208 def notauth(message: dns.message.Message) -> None: 209 rcode(message, dns.rcode.NOTAUTH) 210 211 212 def nxdomain(message: dns.message.Message) -> None: 213 rcode(message, dns.rcode.NXDOMAIN) 214 215 216 def single_question(message: dns.message.Message) -> None: 217 assert len(message.question) == 1, str(message) 218 219 220 def empty_answer(message: dns.message.Message) -> None: 221 assert not message.answer, str(message) 222 223 224 def rr_count_eq(section: list, expected: int): 225 # NOTE: OPT and TSIG records aren't included in the count for ADDITIONAL section 226 count = sum(len(rrset) for rrset in section) 227 assert count == expected, str(section) 228 229 230 def is_response_to(response: dns.message.Message, query: dns.message.Message) -> None: 231 single_question(response) 232 single_question(query) 233 assert query.is_response(response), str(response) 234 235 236 def file_contents_equal(file1, file2): 237 def normalize_line(line): 238 # remove trailing&leading whitespace and replace multiple whitespaces 239 return " ".join(line.split()) 240 241 def read_lines(file_path): 242 with open(file_path, "r", encoding="utf-8") as file: 243 return [normalize_line(line) for line in file.readlines()] 244 245 lines1 = read_lines(file1) 246 lines2 = read_lines(file2) 247 248 differ = difflib.Differ() 249 diff = differ.compare(lines1, lines2) 250 251 for line in diff: 252 assert not line.startswith("+ ") and not line.startswith( 253 "- " 254 ), f'file contents of "{file1}" and "{file2}" differ' 255