Home | History | Annotate | Line # | Download | only in isctest
      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